16 : input_shape_{input.shape()},
17 input_strides_{input.strides()},
18 input_type_{input.dtype()},
19 out_shape_{out.shape()},
20 out_strides_{out.strides()},
21 out_type_{out.dtype()},
22 output_size_{out.numel()},
24 hidden_size_{out.size(out.ndim() - 1)},
25 row_count_{out.numel() / hidden_size_},
26 device_index_{out.device().index()} {
27 assert(input.ndim() == out.ndim() &&
28 "`SiluAndMulInfinilm` input and output ranks must match");
29 assert(input_type_ == out_type_ &&
30 "`SiluAndMulInfinilm` input and output dtypes must match");
32 input.size(input.ndim() - 1) == 2 * hidden_size_ &&
33 "`SiluAndMulInfinilm` input last dimension must be twice output last "
35 for (Tensor::Size i = 0; i + 1 < ndim_; ++i) {
36 assert(input.size(i) == out.size(i) &&
37 "`SiluAndMulInfinilm` leading dimensions must match");
39 assert(input.IsContiguous() && out.IsContiguous() &&
40 "`SiluAndMulInfinilm` only supports contiguous tensors");
41 assert(!out.HasBroadcastDim() &&
42 "`SiluAndMulInfinilm` output must not have broadcasted dimensions");