InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
scaled_softmax_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SCALED_SOFTMAX_INFINILM_H_
2#define INFINI_OPS_BASE_SCALED_SOFTMAX_INFINILM_H_
3
4#include <cassert>
5#include <cmath>
6#include <cstddef>
7
8#include "data_type.h"
9#include "operator.h"
10#include "tensor.h"
11
12namespace infini::ops {
13
16class [[deprecated(
17 "Migrate to an open-source-aligned operator when available.")]]
18ScaledSoftmaxInfinilm : public Operator<ScaledSoftmaxInfinilm> {
19 public:
20 ScaledSoftmaxInfinilm(const Tensor input, double scale, Tensor out)
21 : scale_{scale},
22 batch_size_{input.size(0)},
23 vocab_size_{input.size(1)},
24 dtype_{input.dtype()},
25 input_strides_{input.strides()},
26 out_strides_{out.strides()} {
27 assert(input.ndim() == 2 &&
28 "`ScaledSoftmaxInfinilm` currently supports 2D `[batch, vocab]` "
29 "input");
30 assert(input.shape() == out.shape() &&
31 "`ScaledSoftmaxInfinilm` requires `input` and `out` to have the "
32 "same shape");
33 assert(input.dtype() == out.dtype() &&
34 "`ScaledSoftmaxInfinilm` requires `input` and `out` to have the "
35 "same dtype");
36 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16 ||
37 dtype_ == DataType::kFloat32 || dtype_ == DataType::kFloat64) &&
38 "`ScaledSoftmaxInfinilm` requires a floating point dtype");
39 assert(std::isfinite(scale_) &&
40 "`ScaledSoftmaxInfinilm` requires a finite `scale`");
41 }
42
43 virtual void operator()(const Tensor input, double scale,
44 Tensor out) const = 0;
45
46 protected:
47 double scale_{1.0};
48
49 Tensor::Size batch_size_{0};
50
51 Tensor::Size vocab_size_{0};
52
53 DataType dtype_;
54
55 Tensor::Strides input_strides_;
56
57 Tensor::Strides out_strides_;
58};
59
60} // namespace infini::ops
61
62#endif // INFINI_OPS_BASE_SCALED_SOFTMAX_INFINILM_H_
Definition generated/include/operator.h:282
Definition scaled_softmax_infinilm.h:18
DataType dtype_
Definition scaled_softmax_infinilm.h:53
virtual void operator()(const Tensor input, double scale, Tensor out) const =0
Tensor::Strides input_strides_
Definition scaled_softmax_infinilm.h:55
ScaledSoftmaxInfinilm(const Tensor input, double scale, Tensor out)
Definition scaled_softmax_infinilm.h:20
Tensor::Strides out_strides_
Definition scaled_softmax_infinilm.h:57
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8