InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
softmax_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SOFTMAX_INFINILM_H_
2#define INFINI_OPS_BASE_SOFTMAX_INFINILM_H_
3
4#include <cassert>
5#include <optional>
6
7#include "operator.h"
8
9namespace infini::ops {
10
13class [[deprecated("Use `Softmax` instead.")]] SoftmaxInfinilm
14 : public Operator<SoftmaxInfinilm> {
15 public:
16 SoftmaxInfinilm(const Tensor input, const int64_t dim,
17 const std::optional<DataType> dtype, Tensor out)
18 : input_shape_{input.shape()},
19 input_strides_{input.strides()},
20 input_type_{input.dtype()},
21 out_shape_{out.shape()},
22 out_strides_{out.strides()},
23 out_type_{out.dtype()},
24 dim_{dim < 0 ? dim + static_cast<int64_t>(input.ndim()) : dim},
25 dtype_{dtype},
26 ndim_{out.ndim()},
27 dim_size_{out.size(dim_)},
28 row_count_{out.numel() / dim_size_},
29 device_index_{out.device().index()} {
30 assert(input_shape_ == out_shape_ &&
31 "`SoftmaxInfinilm` input and output shapes must match");
32 assert(dim_ >= 0 && dim_ < static_cast<int64_t>(ndim_) &&
33 "`SoftmaxInfinilm` dim out of range");
34 assert(!dtype_.has_value() || dtype_.value() == out_type_);
35 assert(input_type_ == out_type_ &&
36 "`SoftmaxInfinilm` input and output dtypes must match");
37 assert(!out.HasBroadcastDim() &&
38 "`SoftmaxInfinilm` output must not have broadcasted dimensions");
39 }
40
41 virtual void operator()(const Tensor input, const int64_t dim,
42 const std::optional<DataType> dtype,
43 Tensor out) const = 0;
44
45 protected:
46 Tensor::Shape input_shape_;
47
48 Tensor::Strides input_strides_;
49
50 DataType input_type_;
51
52 Tensor::Shape out_shape_;
53
54 Tensor::Strides out_strides_;
55
56 DataType out_type_;
57
58 int64_t dim_{};
59
60 std::optional<DataType> dtype_{};
61
62 Tensor::Size ndim_{0};
63
64 Tensor::Size dim_size_{0};
65
66 Tensor::Size row_count_{0};
67
68 int device_index_{0};
69};
70
71} // namespace infini::ops
72
73#endif
Definition generated/include/operator.h:282
Definition softmax_infinilm.h:14
Tensor::Shape out_shape_
Definition softmax_infinilm.h:52
Tensor::Shape input_shape_
Definition softmax_infinilm.h:46
DataType out_type_
Definition softmax_infinilm.h:56
SoftmaxInfinilm(const Tensor input, const int64_t dim, const std::optional< DataType > dtype, Tensor out)
Definition softmax_infinilm.h:16
DataType input_type_
Definition softmax_infinilm.h:50
Tensor::Strides out_strides_
Definition softmax_infinilm.h:54
virtual void operator()(const Tensor input, const int64_t dim, const std::optional< DataType > dtype, Tensor out) const =0
Tensor::Strides input_strides_
Definition softmax_infinilm.h:48
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8