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`");