InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
swiglu.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SWIGLU_H_
2#define INFINI_OPS_BASE_SWIGLU_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
12class [[deprecated("Use `SiluAndMul` instead.")]] Swiglu
13 : public Operator<Swiglu> {
14 public:
15 Swiglu(const Tensor input, const Tensor gate, Tensor out)
16 : ndim_{out.ndim()},
17 output_size_{out.numel()},
18 input_type_{input.dtype()},
19 gate_type_{gate.dtype()},
20 out_type_{out.dtype()},
21 input_shape_{input.shape()},
22 gate_shape_{gate.shape()},
23 out_shape_{out.shape()},
24 input_strides_{input.strides()},
25 gate_strides_{gate.strides()},
26 out_strides_{out.strides()},
27 is_input_contiguous_{input.IsContiguous()},
28 is_gate_contiguous_{gate.IsContiguous()},
29 is_out_contiguous_{out.IsContiguous()} {
30 assert(
31 input_type_ == gate_type_ && gate_type_ == out_type_ &&
32 "operator `Swiglu` requires all input and output tensors to have the "
33 "same dtype");
34 }
35
36 virtual void operator()(const Tensor input, const Tensor gate,
37 Tensor out) const = 0;
38
39 protected:
40 Tensor::Size ndim_{0};
41
42 Tensor::Size output_size_{0};
43
44 const DataType input_type_;
45
46 const DataType gate_type_;
47
48 const DataType out_type_;
49
50 Tensor::Shape input_shape_;
51
52 Tensor::Shape gate_shape_;
53
54 Tensor::Shape out_shape_;
55
56 Tensor::Strides input_strides_;
57
58 Tensor::Strides gate_strides_;
59
60 Tensor::Strides out_strides_;
61
62 bool is_input_contiguous_{false};
63
64 bool is_gate_contiguous_{false};
65
66 bool is_out_contiguous_{false};
67};
68
69} // namespace infini::ops
70
71#endif
Definition generated/include/operator.h:282
Definition swiglu.h:13
Tensor::Strides out_strides_
Definition swiglu.h:60
Swiglu(const Tensor input, const Tensor gate, Tensor out)
Definition swiglu.h:15
const DataType out_type_
Definition swiglu.h:48
Tensor::Shape gate_shape_
Definition swiglu.h:52
const DataType gate_type_
Definition swiglu.h:46
virtual void operator()(const Tensor input, const Tensor gate, Tensor out) const =0
const DataType input_type_
Definition swiglu.h:44
Tensor::Strides input_strides_
Definition swiglu.h:56
Tensor::Strides gate_strides_
Definition swiglu.h:58
Tensor::Shape out_shape_
Definition swiglu.h:54
Tensor::Shape input_shape_
Definition swiglu.h:50
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8