Skip to content

Commit 98a82a1

Browse files
committed
header for autograd engine
1 parent 65beb13 commit 98a82a1

3 files changed

Lines changed: 73 additions & 0 deletions

File tree

.gitignore

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,7 @@ install_manifest.txt
7777
# Test artifacts
7878
Testing/
7979
googletest-*/
80+
test_tensor
8081

8182
# Auto-generated files
8283
/include/config.h

include/core/autograd/function.h

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,72 @@
1+
#pragma once
2+
#ifndef AUTOGRAD_FUNCTION_H
3+
#define AUTOGRAD_FUNCTION_H
4+
5+
#include <memory>
6+
#include <string>
7+
#include <vector>
8+
9+
#include "core/tensor/tensor.h"
10+
11+
namespace torchscratch {
12+
namespace core {
13+
namespace autograd {
14+
class Variable;
15+
16+
class Function {
17+
public:
18+
virtual ~Function() = default;
19+
Function() = default;
20+
Function(const Function&) = default;
21+
Function& operator=(const Function&) = default;
22+
Function(Function&&) noexcept = default;
23+
Function& operator=(Function&&) noexcept = default;
24+
25+
virtual std::vector<tensor::Tensor> forward(const std::vector<tensor::Tensor>& inputs) = 0;
26+
27+
virtual std::vector<tensor::Tensor> backward(const std::vector<tensor::Tensor>& grad_outputs) = 0;
28+
29+
virtual std::string name() const = 0;
30+
31+
void save_for_backward(const std::vector<tensor::Tensor>& inputs);
32+
33+
const std::vector<Variable*>& get_saved_variables() const;
34+
35+
private:
36+
std::vector<Variable*> saved_variables_;
37+
};
38+
class AddFunction:public Function{
39+
public:
40+
std::vector<tensor::Tensor> forward(const std::vector<tensor::Tensor>& inputs) override;
41+
std::vector<tensor::Tensor> backward(const std::vector<tensor::Tensor>& grad_outputs) override;
42+
std::string name() const override {
43+
return "AddFunction";
44+
}
45+
};
46+
class MulFunction:public Function{
47+
public:
48+
std::vector<tensor::Tensor> forward(const std::vector<tensor::Tensor>& inputs) override;
49+
std::vector<tensor::Tensor> backward(const std::vector<tensor::Tensor>& grad_outputs) override;
50+
std::string name() const override {
51+
return "MulFunction";
52+
}
53+
private:
54+
tensor::Tensor input1_;
55+
tensor::Tensor input2_;
56+
};
57+
class MatMulFunction:public Function{
58+
public:
59+
std::vector<tensor::Tensor> forward(const std::vector<tensor::Tensor>& inputs) override;
60+
std::vector<tensor::Tensor> backward(const std::vector<tensor::Tensor>& grad_outputs) override;
61+
std::string name() const override {
62+
return "MatMulFunction";
63+
}
64+
private:
65+
tensor::Tensor input1_;
66+
tensor::Tensor input2_;
67+
};
68+
69+
} // namespace autograd
70+
} // namespace core
71+
} // namespace torchscratch
72+
#endif

test_tensor

0 Bytes
Binary file not shown.

0 commit comments

Comments
 (0)