@@ -13,11 +13,11 @@ using Tensor = at::Tensor;
1313namespace fbgemm_gpu {
1414
1515// Kernel for permuting 1D lengths. Used for permutation of sparse features.
16- template <typename index_t >
16+ template <typename index_t , typename permute_t = int32_t >
1717__global__ __launch_bounds__ (kMaxThreads ) void permute_1D_lengths_kernel(
1818 const index_t * __restrict__ lengths,
1919 int32_t permuted_lengths_size,
20- const int32_t * __restrict__ permute,
20+ const permute_t * __restrict__ permute,
2121 index_t * __restrict__ permuted_lengths) {
2222 CUDA_KERNEL_LOOP (i, permuted_lengths_size) {
2323 permuted_lengths[i] = lengths[permute[i]];
@@ -30,13 +30,14 @@ template <
3030 bool has_weight,
3131 typename offsets_t ,
3232 typename indices_t ,
33- typename weights_t >
33+ typename weights_t ,
34+ typename permute_t = int32_t >
3435__global__ __launch_bounds__ (kMaxThreads ) void permute_1D_data_kernel(
3536 int32_t permuted_indices_size,
3637 int32_t permuted_lengths_size,
3738 const indices_t * __restrict__ indices,
3839 const weights_t * __restrict__ weights,
39- const int32_t * __restrict__ permute,
40+ const permute_t * __restrict__ permute,
4041 const offsets_t * __restrict__ input_offsets,
4142 const offsets_t * __restrict__ output_offsets,
4243 indices_t * __restrict__ permuted_indices,
@@ -71,13 +72,14 @@ template <
7172 bool has_weight,
7273 typename offsets_t ,
7374 typename indices_t ,
74- typename weights_t >
75+ typename weights_t ,
76+ typename permute_t = int32_t >
7577__global__ __launch_bounds__ (kMaxThreads ) void permute_1D_data_kernel_vec(
7678 int32_t permuted_indices_size,
7779 int32_t permuted_lengths_size,
7880 const indices_t * __restrict__ indices,
7981 const weights_t * __restrict__ weights,
80- const int32_t * __restrict__ permute,
82+ const permute_t * __restrict__ permute,
8183 const offsets_t * __restrict__ input_offsets,
8284 const offsets_t * __restrict__ output_offsets,
8385 indices_t * __restrict__ permuted_indices,
@@ -244,17 +246,21 @@ permute_1D_sparse_data_cuda(
244246 at::cuda::getCurrentCUDAStream (),
245247 utils::cuda::BlockCapPolicy::OverflowOnly);
246248 AT_DISPATCH_INDEX_TYPES (
247- lengths.scalar_type (), " permute_1D_lengths_kernel" , [&] {
248- FBGEMM_LAUNCH_KERNEL (
249- (permute_1D_lengths_kernel<index_t >),
250- blocks_1,
251- threads_1,
252- 0 ,
253- at::cuda::getCurrentCUDAStream (),
254- lengths_contig.data_ptr <index_t >(),
255- permuted_lengths_size,
256- permute_contig.data_ptr <int32_t >(),
257- permuted_lengths.data_ptr <index_t >());
249+ permute.scalar_type (), " permute_1D_lengths_permute_type" , [&] {
250+ using permute_t = index_t ;
251+ AT_DISPATCH_INDEX_TYPES (
252+ lengths.scalar_type (), " permute_1D_lengths_kernel" , [&] {
253+ FBGEMM_LAUNCH_KERNEL (
254+ (permute_1D_lengths_kernel<index_t , permute_t >),
255+ blocks_1,
256+ threads_1,
257+ 0 ,
258+ at::cuda::getCurrentCUDAStream (),
259+ lengths_contig.data_ptr <index_t >(),
260+ permuted_lengths_size,
261+ permute_contig.data_ptr <permute_t >(),
262+ permuted_lengths.data_ptr <index_t >());
263+ });
258264 });
259265
260266 // convert lengths to offsets
@@ -289,74 +295,81 @@ permute_1D_sparse_data_cuda(
289295 permuted_indices = at::empty (permuted_indices_size, indices.options ());
290296
291297 AT_DISPATCH_INDEX_TYPES (
292- input_offsets.scalar_type (), " permute_1D_data_kernel_vec_1" , [&] {
293- using offsets_t = index_t ;
294- FBGEMM_DISPATCH_ALL_TYPES (
295- indices.scalar_type (), " permute_1D_data_kernel_vec_2" , [&] {
296- using indices_t = scalar_t ;
297- if (weights.has_value ()) {
298- const Tensor weights_value = weights.value ();
299- const auto weights_value_contig = weights_value.contiguous ();
300- int32_t weights_columns = 1 ;
301- if (weights_value.dense_dim () > 1 ) {
302- weights_columns = weights_value.size (1 );
303- permuted_weights = at::empty (
304- {permuted_indices_size, weights_columns},
305- weights_value.options ());
306- } else {
307- permuted_weights =
308- at::empty (permuted_indices_size, weights_value.options ());
309- }
310- FBGEMM_DISPATCH_ALL_TYPES_AND_DOUBLE (
311- weights_value.scalar_type (),
312- " permute_1D_data_kernel_vec_3" ,
313- [&] {
314- using weights_t = scalar_t ;
298+ permute.scalar_type (), " permute_1D_data_permute_type" , [&] {
299+ using permute_t = index_t ;
300+ AT_DISPATCH_INDEX_TYPES (
301+ input_offsets.scalar_type (), " permute_1D_data_kernel_vec_1" , [&] {
302+ using offsets_t = index_t ;
303+ FBGEMM_DISPATCH_ALL_TYPES (
304+ indices.scalar_type (), " permute_1D_data_kernel_vec_2" , [&] {
305+ using indices_t = scalar_t ;
306+ if (weights.has_value ()) {
307+ const Tensor weights_value = weights.value ();
308+ const auto weights_value_contig =
309+ weights_value.contiguous ();
310+ int32_t weights_columns = 1 ;
311+ if (weights_value.dense_dim () > 1 ) {
312+ weights_columns = weights_value.size (1 );
313+ permuted_weights = at::empty (
314+ {permuted_indices_size, weights_columns},
315+ weights_value.options ());
316+ } else {
317+ permuted_weights = at::empty (
318+ permuted_indices_size, weights_value.options ());
319+ }
320+ FBGEMM_DISPATCH_ALL_TYPES_AND_DOUBLE (
321+ weights_value.scalar_type (),
322+ " permute_1D_data_kernel_vec_3" ,
323+ [&] {
324+ using weights_t = scalar_t ;
325+ FBGEMM_LAUNCH_KERNEL (
326+ (permute_1D_data_kernel_vec<
327+ true ,
328+ offsets_t ,
329+ indices_t ,
330+ weights_t ,
331+ permute_t >),
332+ blocks_2,
333+ threads_2,
334+ 0 ,
335+ at::cuda::getCurrentCUDAStream (),
336+ permuted_indices_size,
337+ permuted_lengths_size,
338+ indices_contig.data_ptr <indices_t >(),
339+ weights_value_contig.data_ptr <weights_t >(),
340+ permute_contig.data_ptr <permute_t >(),
341+ input_offsets.data_ptr <offsets_t >(),
342+ output_offsets.data_ptr <offsets_t >(),
343+ permuted_indices.data_ptr <indices_t >(),
344+ permuted_weights.data_ptr <weights_t >(),
345+ weights_columns);
346+ }); // for each weights_t
347+ } else {
315348 FBGEMM_LAUNCH_KERNEL (
316349 (permute_1D_data_kernel_vec<
317- true ,
350+ false ,
318351 offsets_t ,
319352 indices_t ,
320- weights_t >),
353+ std::nullptr_t ,
354+ permute_t >),
321355 blocks_2,
322356 threads_2,
323357 0 ,
324358 at::cuda::getCurrentCUDAStream (),
325359 permuted_indices_size,
326360 permuted_lengths_size,
327361 indices_contig.data_ptr <indices_t >(),
328- weights_value_contig. data_ptr < weights_t >() ,
329- permute_contig.data_ptr <int32_t >(),
362+ nullptr ,
363+ permute_contig.data_ptr <permute_t >(),
330364 input_offsets.data_ptr <offsets_t >(),
331365 output_offsets.data_ptr <offsets_t >(),
332366 permuted_indices.data_ptr <indices_t >(),
333- permuted_weights.data_ptr <weights_t >(),
334- weights_columns);
335- }); // for each weights_t
336- } else {
337- FBGEMM_LAUNCH_KERNEL (
338- (permute_1D_data_kernel_vec<
339- false ,
340- offsets_t ,
341- indices_t ,
342- std::nullptr_t >),
343- blocks_2,
344- threads_2,
345- 0 ,
346- at::cuda::getCurrentCUDAStream (),
347- permuted_indices_size,
348- permuted_lengths_size,
349- indices_contig.data_ptr <indices_t >(),
350- nullptr ,
351- permute_contig.data_ptr <int32_t >(),
352- input_offsets.data_ptr <offsets_t >(),
353- output_offsets.data_ptr <offsets_t >(),
354- permuted_indices.data_ptr <indices_t >(),
355- nullptr ,
356- 1 );
357- }
358- }); // for each indices_t
359- }); // for each offsets_t
367+ nullptr ,
368+ 1 );
369+ }
370+ }); // for each indices_t
371+ }); // for each offsets_t
372+ }); // for each permute_t
360373
361374 return {permuted_lengths, permuted_indices, permuted_weights};
362375}
0 commit comments