Skip to content

Commit 55f151c

Browse files
ycui1984facebook-github-bot
authored andcommitted
int64 permute support in permute_1D_sparse_data (#6029)
Summary: X-link: facebookresearch/FBGEMM#2933 # Context permute_1D_sparse_data_cuda hardcoded int32_t for the permute argument, so when upstream (TorchRec _get_recat) produces an int64 permute for large variable-batch KJT all-to-all, it failed with "expected scalar type Int but found Long". Add a permute_t template parameter (defaulting to int32_t) to permute_1D_lengths_kernel, permute_1D_data_kernel, and permute_1D_data_kernel_vec, and wrap the kernel launch sites with AT_DISPATCH_INDEX_TYPES(permute.scalar_type(), ...), reading the permute via data_ptr<permute_t>(). Fully backward compatible: existing int32 callers use the default template parameter (bit-identical). Companion to the TorchRec _get_recat int64 fix. Differential Revision: D112628954
1 parent b10c8f4 commit 55f151c

1 file changed

Lines changed: 84 additions & 71 deletions

File tree

fbgemm_gpu/src/sparse_ops/sparse_permute_1d.cu

Lines changed: 84 additions & 71 deletions
Original file line numberDiff line numberDiff line change
@@ -13,11 +13,11 @@ using Tensor = at::Tensor;
1313
namespace 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

Comments
 (0)