| 212 | |
| 213 | template <typename T, int sp0, int sp1, int sp2, bool conjugate = false> |
| 214 | __global__ void ShuffleInTensor3SimpleVector(int nthreads, |
| 215 | const T* __restrict__ input, |
| 216 | Dimension<3> input_dims, |
| 217 | T* __restrict__ output) { |
| 218 | Dimension<3> output_dims; |
| 219 | output_dims[sp0] = input_dims[0]; |
| 220 | output_dims[sp1] = input_dims[1]; |
| 221 | output_dims[sp2] = input_dims[2]; |
| 222 | |
| 223 | const int stride = blockDim.x * gridDim.x * kUnroll; |
| 224 | const int tid = blockIdx.x * blockDim.x + threadIdx.x; |
| 225 | T buf[kUnroll]; |
| 226 | |
| 227 | int output_index; |
| 228 | for (output_index = tid * kUnroll; output_index + kUnroll - 1 < nthreads; |
| 229 | output_index += stride) { |
| 230 | #pragma unroll |
| 231 | for (int i = 0; i < kUnroll; i++) { |
| 232 | int output_index_i = output_index + i; |
| 233 | Index<3> output_tensor_index = FlatToTensorIndex(output_index_i, |
| 234 | output_dims); |
| 235 | Index<3> input_tensor_index; |
| 236 | input_tensor_index[0] = output_tensor_index[sp0]; |
| 237 | input_tensor_index[1] = output_tensor_index[sp1]; |
| 238 | input_tensor_index[2] = output_tensor_index[sp2]; |
| 239 | |
| 240 | int input_index_i = TensorIndexToFlat(input_tensor_index, input_dims); |
| 241 | buf[i] = maybe_conj<T, conjugate>::run(ldg(input + input_index_i)); |
| 242 | } |
| 243 | float2 *out = reinterpret_cast<float2*>(output + output_index); |
| 244 | *out = *reinterpret_cast<float2*>(buf); |
| 245 | } |
| 246 | |
| 247 | for(; output_index < nthreads; output_index++) { |
| 248 | Index<3> output_tensor_index = FlatToTensorIndex(output_index, output_dims); |
| 249 | |
| 250 | Index<3> input_tensor_index; |
| 251 | input_tensor_index[0] = output_tensor_index[sp0]; |
| 252 | input_tensor_index[1] = output_tensor_index[sp1]; |
| 253 | input_tensor_index[2] = output_tensor_index[sp2]; |
| 254 | |
| 255 | int input_index = TensorIndexToFlat(input_tensor_index, input_dims); |
| 256 | |
| 257 | output[output_index] = |
| 258 | maybe_conj<T, conjugate>::run(ldg(input + input_index)); |
| 259 | } |
| 260 | } |
| 261 | |
| 262 | // Use shared memory tiles to swap dimension-1 and dimension-2 of a 3D tensor, |
| 263 | // where dimensions are zero-based: output[i][j][k] = input[i][k][j]. |
nothing calls this directly
no test coverage detected