InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
silu_and_mul_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SILU_AND_MUL_INFINILM_H_
2#define INFINI_OPS_BASE_SILU_AND_MUL_INFINILM_H_
3
4#include <cassert>
5
6#include "operator.h"
7
8namespace infini::ops {
9
12class [[deprecated("Use `SiluAndMul` instead.")]] SiluAndMulInfinilm
13 : public Operator<SiluAndMulInfinilm> {
14 public:
16 : input_shape_{input.shape()},
17 input_strides_{input.strides()},
18 input_type_{input.dtype()},
19 out_shape_{out.shape()},
20 out_strides_{out.strides()},
21 out_type_{out.dtype()},
22 output_size_{out.numel()},
23 ndim_{out.ndim()},
24 hidden_size_{out.size(out.ndim() - 1)},
25 row_count_{out.numel() / hidden_size_},
26 device_index_{out.device().index()} {
27 assert(input.ndim() == out.ndim() &&
28 "`SiluAndMulInfinilm` input and output ranks must match");
29 assert(input_type_ == out_type_ &&
30 "`SiluAndMulInfinilm` input and output dtypes must match");
31 assert(
32 input.size(input.ndim() - 1) == 2 * hidden_size_ &&
33 "`SiluAndMulInfinilm` input last dimension must be twice output last "
34 "dimension");
35 for (Tensor::Size i = 0; i + 1 < ndim_; ++i) {
36 assert(input.size(i) == out.size(i) &&
37 "`SiluAndMulInfinilm` leading dimensions must match");
38 }
39 assert(input.IsContiguous() && out.IsContiguous() &&
40 "`SiluAndMulInfinilm` only supports contiguous tensors");
41 assert(!out.HasBroadcastDim() &&
42 "`SiluAndMulInfinilm` output must not have broadcasted dimensions");
43 }
44
45 virtual void operator()(const Tensor input, Tensor out) const = 0;
46
47 protected:
48 Tensor::Shape input_shape_;
49
50 Tensor::Strides input_strides_;
51
52 DataType input_type_;
53
54 Tensor::Shape out_shape_;
55
56 Tensor::Strides out_strides_;
57
58 DataType out_type_;
59
60 Tensor::Size output_size_{0};
61
62 Tensor::Size ndim_{0};
63
64 Tensor::Size hidden_size_{0};
65
66 Tensor::Size row_count_{0};
67
68 int device_index_{0};
69};
70
71} // namespace infini::ops
72
73#endif
Definition generated/include/operator.h:282
Definition silu_and_mul_infinilm.h:13
Tensor::Shape out_shape_
Definition silu_and_mul_infinilm.h:54
virtual void operator()(const Tensor input, Tensor out) const =0
SiluAndMulInfinilm(const Tensor input, Tensor out)
Definition silu_and_mul_infinilm.h:15
Tensor::Strides input_strides_
Definition silu_and_mul_infinilm.h:50
DataType input_type_
Definition silu_and_mul_infinilm.h:52
Tensor::Strides out_strides_
Definition silu_and_mul_infinilm.h:56
DataType out_type_
Definition silu_and_mul_infinilm.h:58
Tensor::Shape input_shape_
Definition silu_and_mul_infinilm.h:48
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8