InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
internal_addmm_activation.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_INTERNAL_ADDMM_ACTIVATION_H_
2#define INFINI_OPS_BASE_INTERNAL_ADDMM_ACTIVATION_H_
3
4#include "operator.h"
5
6namespace infini::ops::internal {
7
8class AddmmActivation : public Operator<AddmmActivation> {
9 public:
10 AddmmActivation(const Tensor input, const Tensor mat1, const Tensor mat2,
11 const double beta, const double alpha, const bool use_gelu,
12 Tensor out)
13 : input_shape_{input.shape()},
14 input_strides_{input.strides()},
15 input_type_{input.dtype()},
16 mat1_shape_{mat1.shape()},
17 mat1_strides_{mat1.strides()},
18 mat1_type_{mat1.dtype()},
19 mat2_shape_{mat2.shape()},
20 mat2_strides_{mat2.strides()},
21 mat2_type_{mat2.dtype()},
22 out_shape_{out.shape()},
23 out_strides_{out.strides()},
24 out_type_{out.dtype()},
25 beta_{beta},
26 alpha_{alpha},
27 use_gelu_{use_gelu},
28 device_index_{out.device().index()} {}
29
30 virtual void operator()(const Tensor input, const Tensor mat1,
31 const Tensor mat2, const double beta,
32 const double alpha, const bool use_gelu,
33 Tensor out) const = 0;
34
35 protected:
36 Tensor::Shape input_shape_;
37
38 Tensor::Strides input_strides_;
39
40 DataType input_type_;
41
42 Tensor::Shape mat1_shape_;
43
44 Tensor::Strides mat1_strides_;
45
46 DataType mat1_type_;
47
48 Tensor::Shape mat2_shape_;
49
50 Tensor::Strides mat2_strides_;
51
52 DataType mat2_type_;
53
54 Tensor::Shape out_shape_;
55
56 Tensor::Strides out_strides_;
57
58 DataType out_type_;
59
60 double beta_{};
61
62 double alpha_{};
63
64 bool use_gelu_{};
65
67};
68
69} // namespace infini::ops::internal
70
71#endif
Definition generated/include/operator.h:282
Definition internal_addmm_activation.h:8
AddmmActivation(const Tensor input, const Tensor mat1, const Tensor mat2, const double beta, const double alpha, const bool use_gelu, Tensor out)
Definition internal_addmm_activation.h:10
DataType mat1_type_
Definition internal_addmm_activation.h:46
Tensor::Shape out_shape_
Definition internal_addmm_activation.h:54
bool use_gelu_
Definition internal_addmm_activation.h:64
DataType out_type_
Definition internal_addmm_activation.h:58
Tensor::Shape mat1_shape_
Definition internal_addmm_activation.h:42
virtual void operator()(const Tensor input, const Tensor mat1, const Tensor mat2, const double beta, const double alpha, const bool use_gelu, Tensor out) const =0
DataType input_type_
Definition internal_addmm_activation.h:40
double beta_
Definition internal_addmm_activation.h:60
double alpha_
Definition internal_addmm_activation.h:62
int device_index_
Definition internal_addmm_activation.h:66
DataType mat2_type_
Definition internal_addmm_activation.h:52
Tensor::Strides mat1_strides_
Definition internal_addmm_activation.h:44
Tensor::Strides input_strides_
Definition internal_addmm_activation.h:38
Tensor::Strides out_strides_
Definition internal_addmm_activation.h:56
Tensor::Strides mat2_strides_
Definition internal_addmm_activation.h:50
Tensor::Shape mat2_shape_
Definition internal_addmm_activation.h:48
Tensor::Shape input_shape_
Definition internal_addmm_activation.h:36
Definition internal_add_relu.h:6
infini::rt::TensorView Tensor
Definition tensor.h:8