InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
scaled_softmax.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SCALED_SOFTMAX_H_
2#define INFINI_OPS_BASE_SCALED_SOFTMAX_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.")]]
18ScaledSoftmax : public Operator<ScaledSoftmax> {
19 public:
20 ScaledSoftmax(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 "`ScaledSoftmax` currently supports 2D `[batch, vocab]` input");
29 assert(input.shape() == out.shape() &&
30 "`ScaledSoftmax` requires `input` and `out` to have the same shape");
31 assert(input.dtype() == out.dtype() &&
32 "`ScaledSoftmax` requires `input` and `out` to have the same dtype");
33 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16 ||
34 dtype_ == DataType::kFloat32 || dtype_ == DataType::kFloat64) &&
35 "`ScaledSoftmax` requires a floating point dtype");
36 assert(std::isfinite(scale_) &&
37 "`ScaledSoftmax` requires a finite `scale`");
38 }
39
40 virtual void operator()(const Tensor input, double scale,
41 Tensor out) const = 0;
42
43 protected:
44 double scale_{1.0};
45
46 Tensor::Size batch_size_{0};
47
48 Tensor::Size vocab_size_{0};
49
50 DataType dtype_;
51
52 Tensor::Strides input_strides_;
53
54 Tensor::Strides out_strides_;
55};
56
57} // namespace infini::ops
58
59#endif // INFINI_OPS_BASE_SCALED_SOFTMAX_H_
Definition generated/include/operator.h:282
Definition scaled_softmax.h:18
DataType dtype_
Definition scaled_softmax.h:50
virtual void operator()(const Tensor input, double scale, Tensor out) const =0
ScaledSoftmax(const Tensor input, double scale, Tensor out)
Definition scaled_softmax.h:20
Tensor::Strides input_strides_
Definition scaled_softmax.h:52
Tensor::Strides out_strides_
Definition scaled_softmax.h:54
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8