InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
set.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SET_H_
2#define INFINI_OPS_BASE_SET_H_
3
4#include <vector>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class Set : public Operator<Set> {
11 public:
12 Set(const Tensor input, const Tensor source, const int64_t storage_offset,
13 const std::vector<int64_t> size, const std::vector<int64_t> stride,
14 Tensor out)
15 : input_shape_{input.shape()},
16 input_strides_{input.strides()},
17 input_type_{input.dtype()},
18 source_shape_{source.shape()},
19 source_strides_{source.strides()},
20 source_type_{source.dtype()},
21 storage_offset_{storage_offset},
22 size_{size},
23 stride_{stride},
24 out_shape_{out.shape()},
25 out_strides_{out.strides()},
26 out_type_{out.dtype()},
27 device_index_{out.device().index()} {}
28
29 Set(const Tensor input, const Tensor source, Tensor out)
30 : input_shape_{input.shape()},
31 input_strides_{input.strides()},
32 input_type_{input.dtype()},
33 source_shape_{source.shape()},
34 source_strides_{source.strides()},
35 source_type_{source.dtype()},
36 out_shape_{out.shape()},
37 out_strides_{out.strides()},
38 out_type_{out.dtype()},
39 device_index_{out.device().index()} {}
40
41 virtual void operator()(const Tensor input, const Tensor source,
42 const int64_t storage_offset,
43 const std::vector<int64_t> size,
44 const std::vector<int64_t> stride,
45 Tensor out) const = 0;
46
47 virtual void operator()(const Tensor input, const Tensor source,
48 Tensor out) const = 0;
49
50 protected:
51 Tensor::Shape input_shape_;
52
53 Tensor::Strides input_strides_;
54
55 DataType input_type_;
56
57 Tensor::Shape source_shape_;
58
59 Tensor::Strides source_strides_;
60
61 DataType source_type_;
62
63 int64_t storage_offset_{};
64
65 std::vector<int64_t> size_{};
66
67 std::vector<int64_t> stride_{};
68
69 Tensor::Shape out_shape_;
70
71 Tensor::Strides out_strides_;
72
73 DataType out_type_;
74
76};
77
78} // namespace infini::ops
79
80#endif
Definition generated/include/operator.h:282
Definition set.h:10
Tensor::Strides out_strides_
Definition set.h:71
Tensor::Shape source_shape_
Definition set.h:57
Tensor::Shape out_shape_
Definition set.h:69
Set(const Tensor input, const Tensor source, Tensor out)
Definition set.h:29
Tensor::Strides input_strides_
Definition set.h:53
Tensor::Strides source_strides_
Definition set.h:59
Set(const Tensor input, const Tensor source, const int64_t storage_offset, const std::vector< int64_t > size, const std::vector< int64_t > stride, Tensor out)
Definition set.h:12
int64_t storage_offset_
Definition set.h:63
int device_index_
Definition set.h:75
DataType source_type_
Definition set.h:61
Tensor::Shape input_shape_
Definition set.h:51
DataType input_type_
Definition set.h:55
std::vector< int64_t > size_
Definition set.h:65
virtual void operator()(const Tensor input, const Tensor source, Tensor out) const =0
virtual void operator()(const Tensor input, const Tensor source, const int64_t storage_offset, const std::vector< int64_t > size, const std::vector< int64_t > stride, Tensor out) const =0
DataType out_type_
Definition set.h:73
std::vector< int64_t > stride_
Definition set.h:67
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8