InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
silu_and_mul.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SILU_AND_MUL_H_
2#define INFINI_OPS_BASE_SILU_AND_MUL_H_
3
4#include "data_type.h"
5#include "operator.h"
6
7namespace infini::ops {
8
9class SiluAndMul : public Operator<SiluAndMul> {
10 public:
11 SiluAndMul(const Tensor input, Tensor out)
12 : ndim_{out.ndim()},
13 output_size_{out.numel()},
14 input_type_{input.dtype()},
15 out_type_{out.dtype()},
16 input_shape_{input.shape()},
17 out_shape_{out.shape()},
18 input_strides_{input.strides()},
19 out_strides_{out.strides()},
20 hidden_size_{out.size(-1)},
21 is_input_contiguous_{input.IsContiguous()},
22 is_out_contiguous_{out.IsContiguous()} {
23 assert(input.ndim() == out.ndim() &&
24 "`SiluAndMul` requires input and output ranks to match");
25 assert(input_type_ == out_type_ &&
26 "`SiluAndMul` requires input and output dtypes to match");
27 assert(input.size(-1) == 2 * hidden_size_ &&
28 "`SiluAndMul` requires input last dimension to be twice output "
29 "last dimension");
30 for (Tensor::Size i = 0; i + 1 < ndim_; ++i) {
31 assert(input.size(i) == out.size(i) &&
32 "`SiluAndMul` requires matching leading dimensions");
33 }
34 }
35
36 virtual void operator()(const Tensor input, Tensor out) const = 0;
37
38 template <typename TensorLike>
39 static auto MakeReturnValue(const TensorLike& input) {
40 typename TensorLike::Shape out_shape{input.shape()};
41 assert(!out_shape.empty() && out_shape.back() % 2 == 0 &&
42 "`SiluAndMul` requires an even input last dimension");
43 out_shape.back() /= 2;
44
45 return TensorLike::Empty(out_shape, input.dtype(), input.device());
46 }
47
48 protected:
49 Tensor::Size ndim_{0};
50
51 Tensor::Size output_size_{0};
52
53 DataType input_type_;
54
55 DataType out_type_;
56
57 Tensor::Shape input_shape_;
58
59 Tensor::Shape out_shape_;
60
61 Tensor::Strides input_strides_;
62
63 Tensor::Strides out_strides_;
64
65 Tensor::Size hidden_size_{0};
66
68
69 bool is_out_contiguous_{false};
70};
71
72} // namespace infini::ops
73
74#endif
Definition generated/include/operator.h:282
Definition silu_and_mul.h:9
Tensor::Size ndim_
Definition silu_and_mul.h:49
DataType out_type_
Definition silu_and_mul.h:55
static auto MakeReturnValue(const TensorLike &input)
Definition silu_and_mul.h:39
bool is_input_contiguous_
Definition silu_and_mul.h:67
SiluAndMul(const Tensor input, Tensor out)
Definition silu_and_mul.h:11
DataType input_type_
Definition silu_and_mul.h:53
Tensor::Shape out_shape_
Definition silu_and_mul.h:59
Tensor::Strides input_strides_
Definition silu_and_mul.h:61
Tensor::Strides out_strides_
Definition silu_and_mul.h:63
bool is_out_contiguous_
Definition silu_and_mul.h:69
Tensor::Size hidden_size_
Definition silu_and_mul.h:65
Tensor::Size output_size_
Definition silu_and_mul.h:51
virtual void operator()(const Tensor input, Tensor out) const =0
Tensor::Shape input_shape_
Definition silu_and_mul.h:57
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8