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()},
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");
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");
51 !values.HasBroadcastDim() && !indices.HasBroadcastDim() &&
52 "`TopksoftmaxInfinilm` outputs must not have broadcasted dimensions");