InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
replication_pad1d_backward.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_REPLICATION_PAD1D_BACKWARD_H_
2#define INFINI_OPS_BASE_REPLICATION_PAD1D_BACKWARD_H_
3
4#include <vector>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class ReplicationPad1dBackward : public Operator<ReplicationPad1dBackward> {
11 public:
12 ReplicationPad1dBackward(const Tensor grad_output, const Tensor input,
13 const std::vector<int64_t> padding,
14 Tensor grad_input)
15 : grad_output_shape_{grad_output.shape()},
16 grad_output_strides_{grad_output.strides()},
17 grad_output_type_{grad_output.dtype()},
18 input_shape_{input.shape()},
19 input_strides_{input.strides()},
20 input_type_{input.dtype()},
21 grad_input_shape_{grad_input.shape()},
22 grad_input_strides_{grad_input.strides()},
23 grad_input_type_{grad_input.dtype()},
24 padding_{padding},
25 device_index_{grad_input.device().index()} {}
26
27 virtual void operator()(const Tensor grad_output, const Tensor input,
28 const std::vector<int64_t> padding,
29 Tensor grad_input) const = 0;
30
31 protected:
32 Tensor::Shape grad_output_shape_;
33
34 Tensor::Strides grad_output_strides_;
35
37
38 Tensor::Shape input_shape_;
39
40 Tensor::Strides input_strides_;
41
42 DataType input_type_;
43
44 Tensor::Shape grad_input_shape_;
45
46 Tensor::Strides grad_input_strides_;
47
49
50 std::vector<int64_t> padding_{};
51
53};
54
55} // namespace infini::ops
56
57#endif
Definition generated/include/operator.h:282
Definition replication_pad1d_backward.h:10
DataType grad_output_type_
Definition replication_pad1d_backward.h:36
virtual void operator()(const Tensor grad_output, const Tensor input, const std::vector< int64_t > padding, Tensor grad_input) const =0
Tensor::Strides grad_output_strides_
Definition replication_pad1d_backward.h:34
Tensor::Strides input_strides_
Definition replication_pad1d_backward.h:40
DataType grad_input_type_
Definition replication_pad1d_backward.h:48
std::vector< int64_t > padding_
Definition replication_pad1d_backward.h:50
Tensor::Shape grad_output_shape_
Definition replication_pad1d_backward.h:32
DataType input_type_
Definition replication_pad1d_backward.h:42
ReplicationPad1dBackward(const Tensor grad_output, const Tensor input, const std::vector< int64_t > padding, Tensor grad_input)
Definition replication_pad1d_backward.h:12
Tensor::Strides grad_input_strides_
Definition replication_pad1d_backward.h:46
Tensor::Shape grad_input_shape_
Definition replication_pad1d_backward.h:44
Tensor::Shape input_shape_
Definition replication_pad1d_backward.h:38
int device_index_
Definition replication_pad1d_backward.h:52
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8