diff --git a/fbgemm_gpu/src/sparse_ops/sparse_ops_cpu.cpp b/fbgemm_gpu/src/sparse_ops/sparse_ops_cpu.cpp index 630f497f14..dee77ad82b 100644 --- a/fbgemm_gpu/src/sparse_ops/sparse_ops_cpu.cpp +++ b/fbgemm_gpu/src/sparse_ops/sparse_ops_cpu.cpp @@ -952,7 +952,7 @@ std::tuple> permute_1D_sparse_data_cpu( // repetitions Tensor permuted_lengths; Tensor permuted_indices; - Tensor permuted_weights; + std::optional permuted_weights; const auto permuted_lengths_size = permute.numel(); permuted_lengths = at::empty({permuted_lengths_size}, lengths.options()); @@ -1013,7 +1013,7 @@ std::tuple> permute_1D_sparse_data_cpu( permuted_lengths.mutable_data_ptr(), output_offsets.mutable_data_ptr(), permuted_indices.mutable_data_ptr(), - permuted_weights.mutable_data_ptr()); + permuted_weights->mutable_data_ptr()); } else { _permute_1D_indices_weights_kernel_cpu< false, @@ -3022,7 +3022,7 @@ std::tuple> permute_sparse_features_cpu( Tensor permuted_lengths; Tensor permuted_indices; - Tensor permuted_weights; + std::optional permuted_weights; permuted_lengths = at::empty({num_output_features, B}, lengths.options()); @@ -3073,7 +3073,7 @@ std::tuple> permute_sparse_features_cpu( input_offsets.const_data_ptr(), output_offsets_per_thread_cumsum.data(), permuted_indices.mutable_data_ptr(), - permuted_weights.mutable_data_ptr(), + permuted_weights->mutable_data_ptr(), permuted_lengths.const_data_ptr()); } else { _permute_data_kernel_cpu( diff --git a/fbgemm_gpu/src/sparse_ops/sparse_permute_1d.cu b/fbgemm_gpu/src/sparse_ops/sparse_permute_1d.cu index 0f6c19875e..21ceff408f 100644 --- a/fbgemm_gpu/src/sparse_ops/sparse_permute_1d.cu +++ b/fbgemm_gpu/src/sparse_ops/sparse_permute_1d.cu @@ -412,7 +412,7 @@ permute_1D_sparse_data_cuda( Tensor permuted_lengths; Tensor permuted_indices; - Tensor permuted_weights; + std::optional permuted_weights; TORCH_CHECK( permuted_lengths_size >= 0 && permuted_lengths_size <= std::numeric_limits::max(), @@ -607,7 +607,7 @@ permute_1D_sparse_data_cuda( input_offsets.data_ptr(), output_offsets.data_ptr(), permuted_indices.data_ptr(), - permuted_weights.data_ptr(), + permuted_weights->data_ptr(), weights_columns); }); // for each weights_t } else { diff --git a/fbgemm_gpu/src/sparse_ops/sparse_permute_2d.cu b/fbgemm_gpu/src/sparse_ops/sparse_permute_2d.cu index 1ec93912db..6a12869bea 100644 --- a/fbgemm_gpu/src/sparse_ops/sparse_permute_2d.cu +++ b/fbgemm_gpu/src/sparse_ops/sparse_permute_2d.cu @@ -489,7 +489,7 @@ permute_2D_sparse_preallocated_out_cuda( Tensor permuted_lengths; Tensor permuted_indices; - Tensor permuted_weights; + std::optional permuted_weights; permuted_lengths = permuted_lengths_out.has_value() ? permuted_lengths_out.value() @@ -681,7 +681,7 @@ permute_2D_sparse_preallocated_out_cuda( input_offsets.data_ptr(), output_offsets.data_ptr(), permuted_indices.data_ptr(), - permuted_weights.data_ptr()); + permuted_weights->data_ptr()); } else { FBGEMM_LAUNCH_KERNEL( (permute_2D_data_kernel< @@ -703,7 +703,7 @@ permute_2D_sparse_preallocated_out_cuda( input_offsets.data_ptr(), output_offsets.data_ptr(), permuted_indices.data_ptr(), - permuted_weights.data_ptr()); + permuted_weights->data_ptr()); } }); // for each weights_t } else { @@ -831,7 +831,7 @@ permute_sparse_features_cuda( Tensor permuted_lengths; Tensor permuted_indices; - Tensor permuted_weights; + std::optional permuted_weights; permuted_lengths = at::empty({num_output_features, B}, lengths.options()); @@ -913,7 +913,7 @@ permute_sparse_features_cuda( input_offsets.data_ptr(), output_offsets.data_ptr(), permuted_indices.data_ptr(), - permuted_weights.data_ptr()); + permuted_weights->data_ptr()); }); }); } else { diff --git a/fbgemm_gpu/test/sparse/permute_optional_weights_test.py b/fbgemm_gpu/test/sparse/permute_optional_weights_test.py new file mode 100644 index 0000000000..fedee51d8d --- /dev/null +++ b/fbgemm_gpu/test/sparse/permute_optional_weights_test.py @@ -0,0 +1,145 @@ +#!/usr/bin/env python3 +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +# pyre-strict + +import json +import os +import subprocess +import sys +import unittest + +import torch + + +def _check_native_permutations(device: str) -> None: + # Python autograd registrations can normalize an undefined Tensor to None. + # Keep this process free of those registrations, as in native inference. + assert "fbgemm_gpu" not in sys.modules + calls = { + "permute_2D_sparse_data": "permute, lengths, values, weights", + "permute_1D_sparse_data": "permute, lengths.view(-1), values, weights", + "permute_sparse_data": "permute, lengths, values, weights", + "permute_2D_sparse_data_input1D": ( + "permute, lengths.view(-1), values, 1, weights" + ), + "permute_sparse_features": "permute, lengths, values, weights", + "permute_2D_sparse_preallocated_out": ( + "permute, lengths, values, weights, values.numel(), " + "torch.empty_like(lengths), torch.empty_like(values), weights_out" + ), + } + passed = 0 + for op, arguments in calls.items(): + # Check None inside TorchScript before crossing the Python boundary, + # and pass the result directly to another native permutation. + functions = torch.jit.CompilationUnit( + f""" +def permute_twice(permute: Tensor, lengths: Tensor, values: Tensor, + weights: Optional[Tensor]): + weights_out: Optional[Tensor] = None + if weights is not None: + weights_out = torch.empty_like(weights) + out_lengths, out_values, out_weights = torch.ops.fbgemm.{op}({arguments}) + out_weights_none = out_weights is None + twice_lengths, twice_values, twice_weights = torch.ops.fbgemm.permute_2D_sparse_data( + permute, out_lengths.view(2, 1), out_values, out_weights) + return (out_lengths.view(2, 1), out_values, out_weights, out_weights_none, + twice_lengths, twice_values, twice_weights, twice_weights is None) +""" + ) + for dtype in (torch.int32, torch.int64): + for sizes in ([2, 1], [0, 2], [0, 0]): + for weighted in (False, True): + case = ( + f"{device}, {op}, {dtype}, lengths={sizes}, weighted={weighted}" + ) + print(case, flush=True) + permute = torch.tensor([1, 0], dtype=torch.int32, device=device) + lengths = torch.tensor(sizes, dtype=dtype, device=device).view(2, 1) + values = torch.arange(sum(sizes), dtype=dtype, device=device) + weights = values.float() * 0.25 if weighted else None + result = functions.permute_twice(permute, lengths, values, weights) + expected_values = torch.cat( + (values[sizes[0] :], values[: sizes[0]]) + ) + torch.testing.assert_close(result[0], lengths.flip(0)) + torch.testing.assert_close(result[1], expected_values) + torch.testing.assert_close(result[4], lengths) + torch.testing.assert_close(result[5], values) + assert result[3] == (not weighted), case + assert result[7] == (not weighted), case + if weighted: + torch.testing.assert_close( + result[2], expected_values.float() * 0.25 + ) + torch.testing.assert_close(result[6], weights) + else: + assert result[2] is None, case + assert result[6] is None, case + passed += 1 + print(f"{passed} native TorchScript cases passed on {device}", flush=True) + + +class PermuteOptionalWeightsTest(unittest.TestCase): + def _run_native_test(self, device: str, flat: bool) -> None: + # Resolve the shared libraries in the usual test environment, but load + # only their native registrations in a fresh process below. + import fbgemm_gpu + + if not getattr(fbgemm_gpu, "open_source", False): + torch.ops.load_library("//deeplearning/fbgemm/fbgemm_gpu:sparse_ops") + + # These flags are cached by C++; a fresh process tests both kernel paths + # independently of any other permutation tests in the parent process. + env = os.environ.copy() + env["FBGEMM_FLAT_PERMUTE_1D"] = str(int(flat)) + env["FBGEMM_FLAT_PERMUTE_2D"] = str(int(flat)) + script = """ +import json +import runpy +import sys +import torch + +for library in json.loads(sys.argv[1]): + torch.ops.load_library(library) +runpy.run_path(sys.argv[2])["_check_native_permutations"](sys.argv[3]) +""" + try: + result = subprocess.run( + [ + sys.executable, + "-c", + script, + json.dumps(sorted(torch.ops.loaded_libraries)), + os.path.abspath(__file__), + device, + ], + env=env, + capture_output=True, + text=True, + timeout=180, + ) + except subprocess.TimeoutExpired as error: + self.fail( + f"Native permutation tests timed out: {error.stdout}\n{error.stderr}" + ) + self.assertEqual(result.returncode, 0, f"{result.stdout}\n{result.stderr}") + self.assertIn("72 native TorchScript cases passed", result.stdout) + + def test_cpu(self) -> None: + self._run_native_test("cpu", flat=False) + + @unittest.skipIf(not torch.cuda.is_available(), "CUDA is not available") + def test_cuda(self) -> None: + for flat in (False, True): + with self.subTest(flat=flat): + self._run_native_test("cuda", flat=flat) + + +if __name__ == "__main__": + unittest.main()