Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions fbgemm_gpu/src/sparse_ops/sparse_ops_cpu.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -952,7 +952,7 @@ std::tuple<Tensor, Tensor, std::optional<Tensor>> permute_1D_sparse_data_cpu(
// repetitions
Tensor permuted_lengths;
Tensor permuted_indices;
Tensor permuted_weights;
std::optional<Tensor> permuted_weights;

const auto permuted_lengths_size = permute.numel();
permuted_lengths = at::empty({permuted_lengths_size}, lengths.options());
Expand Down Expand Up @@ -1013,7 +1013,7 @@ std::tuple<Tensor, Tensor, std::optional<Tensor>> permute_1D_sparse_data_cpu(
permuted_lengths.mutable_data_ptr<offsets_t>(),
output_offsets.mutable_data_ptr<offsets_t>(),
permuted_indices.mutable_data_ptr<indices_t>(),
permuted_weights.mutable_data_ptr<weights_t>());
permuted_weights->mutable_data_ptr<weights_t>());
} else {
_permute_1D_indices_weights_kernel_cpu<
false,
Expand Down Expand Up @@ -3022,7 +3022,7 @@ std::tuple<Tensor, Tensor, std::optional<Tensor>> permute_sparse_features_cpu(

Tensor permuted_lengths;
Tensor permuted_indices;
Tensor permuted_weights;
std::optional<Tensor> permuted_weights;

permuted_lengths = at::empty({num_output_features, B}, lengths.options());

Expand Down Expand Up @@ -3073,7 +3073,7 @@ std::tuple<Tensor, Tensor, std::optional<Tensor>> permute_sparse_features_cpu(
input_offsets.const_data_ptr<index_t>(),
output_offsets_per_thread_cumsum.data(),
permuted_indices.mutable_data_ptr<index_t>(),
permuted_weights.mutable_data_ptr<scalar_t>(),
permuted_weights->mutable_data_ptr<scalar_t>(),
permuted_lengths.const_data_ptr<index_t>());
} else {
_permute_data_kernel_cpu<false, index_t, scalar_t>(
Expand Down
4 changes: 2 additions & 2 deletions fbgemm_gpu/src/sparse_ops/sparse_permute_1d.cu
Original file line number Diff line number Diff line change
Expand Up @@ -412,7 +412,7 @@ permute_1D_sparse_data_cuda(

Tensor permuted_lengths;
Tensor permuted_indices;
Tensor permuted_weights;
std::optional<Tensor> permuted_weights;
TORCH_CHECK(
permuted_lengths_size >= 0 &&
permuted_lengths_size <= std::numeric_limits<int32_t>::max(),
Expand Down Expand Up @@ -607,7 +607,7 @@ permute_1D_sparse_data_cuda(
input_offsets.data_ptr<offsets_t>(),
output_offsets.data_ptr<offsets_t>(),
permuted_indices.data_ptr<indices_t>(),
permuted_weights.data_ptr<weights_t>(),
permuted_weights->data_ptr<weights_t>(),
weights_columns);
}); // for each weights_t
} else {
Expand Down
10 changes: 5 additions & 5 deletions fbgemm_gpu/src/sparse_ops/sparse_permute_2d.cu
Original file line number Diff line number Diff line change
Expand Up @@ -489,7 +489,7 @@ permute_2D_sparse_preallocated_out_cuda(

Tensor permuted_lengths;
Tensor permuted_indices;
Tensor permuted_weights;
std::optional<Tensor> permuted_weights;

permuted_lengths = permuted_lengths_out.has_value()
? permuted_lengths_out.value()
Expand Down Expand Up @@ -681,7 +681,7 @@ permute_2D_sparse_preallocated_out_cuda(
input_offsets.data_ptr<offsets_t>(),
output_offsets.data_ptr<offsets_t>(),
permuted_indices.data_ptr<indices_t>(),
permuted_weights.data_ptr<weights_t>());
permuted_weights->data_ptr<weights_t>());
} else {
FBGEMM_LAUNCH_KERNEL(
(permute_2D_data_kernel<
Expand All @@ -703,7 +703,7 @@ permute_2D_sparse_preallocated_out_cuda(
input_offsets.data_ptr<offsets_t>(),
output_offsets.data_ptr<offsets_t>(),
permuted_indices.data_ptr<indices_t>(),
permuted_weights.data_ptr<weights_t>());
permuted_weights->data_ptr<weights_t>());
}
}); // for each weights_t
} else {
Expand Down Expand Up @@ -831,7 +831,7 @@ permute_sparse_features_cuda(

Tensor permuted_lengths;
Tensor permuted_indices;
Tensor permuted_weights;
std::optional<Tensor> permuted_weights;

permuted_lengths = at::empty({num_output_features, B}, lengths.options());

Expand Down Expand Up @@ -913,7 +913,7 @@ permute_sparse_features_cuda(
input_offsets.data_ptr<index_t>(),
output_offsets.data_ptr<index_t>(),
permuted_indices.data_ptr<index_t>(),
permuted_weights.data_ptr<scalar_t>());
permuted_weights->data_ptr<scalar_t>());
});
});
} else {
Expand Down
145 changes: 145 additions & 0 deletions fbgemm_gpu/test/sparse/permute_optional_weights_test.py
Original file line number Diff line number Diff line change
@@ -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()
Loading