1#ifndef INFINI_OPS_BASE_CAUSAL_SOFTMAX_H_
2#define INFINI_OPS_BASE_CAUSAL_SOFTMAX_H_
15 "Migrate to an open-source-aligned operator when available.")]]
19 : dtype_{input.dtype()},
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 "`CausalSoftmax` requires `input` and `out` same shape");
28 assert(input.dtype() == out.dtype() &&
29 "`CausalSoftmax` requires `input` and `out` same dtype");
30 assert((ndim_ == 2 || ndim_ == 3) &&
31 "`CausalSoftmax` requires 2D or 3D tensor");
32 assert(seq_len_ <= total_seq_len_ &&
33 "`CausalSoftmax` requires shape[-2] <= shape[-1]");
41 Tensor::Size ndim_{0};
43 Tensor::Size batch_size_{0};
45 Tensor::Size seq_len_{0};
47 Tensor::Size total_seq_len_{0};
Definition causal_softmax.h:16
Tensor::Strides input_strides_
Definition causal_softmax.h:49
CausalSoftmax(const Tensor input, Tensor out)
Definition causal_softmax.h:18
Tensor::Strides out_strides_
Definition causal_softmax.h:51
const DataType dtype_
Definition causal_softmax.h:39
virtual void operator()(const Tensor input, Tensor out) const =0
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8