InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
causal_softmax_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_CAUSAL_SOFTMAX_INFINILM_H_
2#define INFINI_OPS_BASE_CAUSAL_SOFTMAX_INFINILM_H_
3
4#include <cassert>
5#include <cstddef>
6
7#include "operator.h"
8#include "tensor.h"
9
10namespace infini::ops {
11
14class [[deprecated(
15 "Migrate to an open-source-aligned operator when available.")]]
16CausalSoftmaxInfinilm : public Operator<CausalSoftmaxInfinilm> {
17 public:
19 : dtype_{input.dtype()},
20 ndim_{out.ndim()},
21 batch_size_{ndim_ == 2 ? 1 : out.size(-3)},
22 seq_len_{out.size(-2)},
23 total_seq_len_{out.size(-1)},
24 input_strides_{input.strides()},
25 out_strides_{out.strides()} {
26 assert(input.shape() == out.shape() &&
27 "`CausalSoftmaxInfinilm` requires `input` and `out` same shape");
28 assert(input.dtype() == out.dtype() &&
29 "`CausalSoftmaxInfinilm` requires `input` and `out` same dtype");
30 assert((ndim_ == 2 || ndim_ == 3) &&
31 "`CausalSoftmaxInfinilm` requires 2D or 3D tensor");
32 assert(seq_len_ <= total_seq_len_ &&
33 "`CausalSoftmaxInfinilm` requires shape[-2] <= shape[-1]");
34 }
35
36 virtual void operator()(const Tensor input, Tensor out) const = 0;
37
38 protected:
39 const DataType dtype_;
40
41 Tensor::Size ndim_{0};
42
43 Tensor::Size batch_size_{0};
44
45 Tensor::Size seq_len_{0};
46
47 Tensor::Size total_seq_len_{0};
48
49 Tensor::Strides input_strides_;
50
51 Tensor::Strides out_strides_;
52};
53
54} // namespace infini::ops
55
56#endif
Definition causal_softmax_infinilm.h:16
CausalSoftmaxInfinilm(const Tensor input, Tensor out)
Definition causal_softmax_infinilm.h:18
Tensor::Strides input_strides_
Definition causal_softmax_infinilm.h:49
Tensor::Strides out_strides_
Definition causal_softmax_infinilm.h:51
virtual void operator()(const Tensor input, Tensor out) const =0
const DataType dtype_
Definition causal_softmax_infinilm.h:39
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8