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