InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
silu.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SILU_H_
2#define INFINI_OPS_BASE_SILU_H_
3
4#include <cstddef>
5
6#include "data_type.h"
7#include "operator.h"
8
9namespace infini::ops {
10
11// Aligned with InfiniCore and `torch.nn.functional.silu`.
12class Silu : public Operator<Silu> {
13 public:
14 Silu(const Tensor input, Tensor out)
15 : ndim_{out.ndim()},
16 output_size_{out.numel()},
17 input_type_{input.dtype()},
18 out_type_{out.dtype()},
19 input_shape_{input.shape()},
20 out_shape_{out.shape()},
21 input_strides_{input.strides()},
22 out_strides_{out.strides()},
23 is_input_contiguous_{input.IsContiguous()},
24 is_out_contiguous_{out.IsContiguous()} {
25 assert(input.shape() == out.shape() &&
26 "`Silu` requires `input` and `out` to have the same shape");
27 assert(input_type_ == out_type_ &&
28 "`Silu` requires `input` and `out` to have the same dtype");
29 assert((input_type_ == DataType::kFloat16 ||
30 input_type_ == DataType::kBFloat16 ||
31 input_type_ == DataType::kFloat32 ||
32 input_type_ == DataType::kFloat64) &&
33 "`Silu` supports float16, bfloat16, float32, and float64 only");
34 }
35
36 virtual void operator()(const Tensor input, Tensor out) const = 0;
37
38 protected:
39 Tensor::Size ndim_{0};
40
41 Tensor::Size output_size_{0};
42
43 DataType input_type_;
44
45 DataType out_type_;
46
47 Tensor::Shape input_shape_;
48
49 Tensor::Shape out_shape_;
50
51 Tensor::Strides input_strides_;
52
53 Tensor::Strides out_strides_;
54
56
57 bool is_out_contiguous_{false};
58};
59
60} // namespace infini::ops
61
62#endif
Definition generated/include/operator.h:282
Definition silu.h:12
bool is_input_contiguous_
Definition silu.h:55
Tensor::Strides input_strides_
Definition silu.h:51
Tensor::Strides out_strides_
Definition silu.h:53
bool is_out_contiguous_
Definition silu.h:57
virtual void operator()(const Tensor input, Tensor out) const =0
Tensor::Shape out_shape_
Definition silu.h:49
Tensor::Shape input_shape_
Definition silu.h:47
Silu(const Tensor input, Tensor out)
Definition silu.h:14
Tensor::Size output_size_
Definition silu.h:41
DataType out_type_
Definition silu.h:45
Tensor::Size ndim_
Definition silu.h:39
DataType input_type_
Definition silu.h:43
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8