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},
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");