InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
topksoftmax_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_TOPKSOFTMAX_INFINILM_H_
2#define INFINI_OPS_BASE_TOPKSOFTMAX_INFINILM_H_
3
4#include <cassert>
5
6#include "operator.h"
7
8namespace infini::ops {
9
12class [[deprecated("Use `TopkSoftmax` instead.")]] TopksoftmaxInfinilm
13 : public Operator<TopksoftmaxInfinilm> {
14 public:
15 TopksoftmaxInfinilm(const Tensor input, const int64_t topk, const bool norm,
16 Tensor values, Tensor indices)
17 : input_shape_{input.shape()},
18 input_strides_{input.strides()},
19 input_type_{input.dtype()},
20 values_shape_{values.shape()},
21 values_strides_{values.strides()},
22 values_type_{values.dtype()},
23 indices_shape_{indices.shape()},
24 indices_strides_{indices.strides()},
25 indices_type_{indices.dtype()},
26 topk_{topk},
27 norm_{norm},
28 row_count_{input.size(0)},
29 width_{input.size(1)},
30 device_index_{values.device().index()} {
31 assert(input.ndim() == 2 &&
32 "`TopksoftmaxInfinilm` input must be a 2D tensor");
33 assert(topk_ > 0 && topk_ <= static_cast<int64_t>(width_) &&
34 "`TopksoftmaxInfinilm` topk must be in (0, input.size(1)]");
35 assert(values_shape_ == indices_shape_ &&
36 "`TopksoftmaxInfinilm` values and indices shapes must match");
37 assert(
38 values_shape_.size() == 2 && values_shape_[0] == row_count_ &&
39 values_shape_[1] == static_cast<Tensor::Size>(topk_) &&
40 "`TopksoftmaxInfinilm` outputs must have shape (input.size(0), topk)");
41 assert(values_type_ == DataType::kFloat32 &&
42 "`TopksoftmaxInfinilm` values output must be float32");
43 assert(indices_type_ == DataType::kInt32 &&
44 "`TopksoftmaxInfinilm` indices output must be int32");
45 assert((input_type_ == DataType::kFloat16 ||
46 input_type_ == DataType::kBFloat16 ||
47 input_type_ == DataType::kFloat32 ||
48 input_type_ == DataType::kFloat64) &&
49 "`TopksoftmaxInfinilm` input must be a floating point tensor");
50 assert(
51 !values.HasBroadcastDim() && !indices.HasBroadcastDim() &&
52 "`TopksoftmaxInfinilm` outputs must not have broadcasted dimensions");
53 }
54
55 virtual void operator()(const Tensor input, const int64_t topk,
56 const bool norm, Tensor values,
57 Tensor indices) const = 0;
58
59 protected:
60 Tensor::Shape input_shape_;
61
62 Tensor::Strides input_strides_;
63
64 DataType input_type_;
65
66 Tensor::Shape values_shape_;
67
68 Tensor::Strides values_strides_;
69
70 DataType values_type_;
71
72 Tensor::Shape indices_shape_;
73
74 Tensor::Strides indices_strides_;
75
76 DataType indices_type_;
77
78 int64_t topk_{0};
79
80 bool norm_{false};
81
82 Tensor::Size row_count_{0};
83
84 Tensor::Size width_{0};
85
86 int device_index_{0};
87};
88
89} // namespace infini::ops
90
91#endif
Definition generated/include/operator.h:282
Definition topksoftmax_infinilm.h:13
TopksoftmaxInfinilm(const Tensor input, const int64_t topk, const bool norm, Tensor values, Tensor indices)
Definition topksoftmax_infinilm.h:15
Tensor::Shape values_shape_
Definition topksoftmax_infinilm.h:66
Tensor::Strides indices_strides_
Definition topksoftmax_infinilm.h:74
Tensor::Shape indices_shape_
Definition topksoftmax_infinilm.h:72
DataType indices_type_
Definition topksoftmax_infinilm.h:76
DataType values_type_
Definition topksoftmax_infinilm.h:70
Tensor::Strides input_strides_
Definition topksoftmax_infinilm.h:62
DataType input_type_
Definition topksoftmax_infinilm.h:64
virtual void operator()(const Tensor input, const int64_t topk, const bool norm, Tensor values, Tensor indices) const =0
Tensor::Shape input_shape_
Definition topksoftmax_infinilm.h:60
Tensor::Strides values_strides_
Definition topksoftmax_infinilm.h:68
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8