| 62 | using ArgsGenerator = typename TestArgs::ArgsGenerator; |
| 63 | |
| 64 | void PrepareData(TestTensorList<InputType, Dims>& test_data) { |
| 65 | std::vector<int> sample_dims(Dims, static_cast<int>(DimSize)); |
| 66 | sample_dims[0] = DimSize0; |
| 67 | if (Dims > 1) { |
| 68 | sample_dims[1] = DimSize1; |
| 69 | } |
| 70 | if (Dims > 2) { |
| 71 | sample_dims[2] = DimSize2; |
| 72 | } |
| 73 | TensorListShape<Dims> shape = uniform_list_shape<Dims>(NumSamples, sample_dims); |
| 74 | test_data.reshape(shape); |
| 75 | |
| 76 | InputType num = 0; |
| 77 | auto seq_gen = [&num]() { return num++; }; |
| 78 | Fill(test_data.cpu(), seq_gen); |
| 79 | } |
| 80 | |
| 81 | void PrepareExpectedOutput(TestTensorList<InputType, Dims>& input_data, |
| 82 | std::vector<SliceArgs<OutputType, Dims>>& slice_args, |