Skip to content
Merged
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
56 changes: 54 additions & 2 deletions tzrec/utils/fx_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@

import torch
import torch.distributed as dist
from torchrec import JaggedTensor, KeyedTensor
from torchrec import JaggedTensor, KeyedJaggedTensor, KeyedTensor
from torchrec.fx import symbolic_trace as _symbolic_trace

# Modules whose forward FX cannot record -- they branch on tensor values or
Expand All @@ -22,6 +22,24 @@
UNTRACEABLE_MODULES = ["ComputeJTDictToKJT", "PromptAssembler", "HoleKeyBuilder"]


@torch.fx.wrap
def _restore_unweighted_kjt(
source: KeyedJaggedTensor, permuted: KeyedJaggedTensor
) -> KeyedJaggedTensor:
"""Preserve absent weights after a TorchRec feature permutation.
Comment thread
eric-gecheng marked this conversation as resolved.

FBGEMM CUDA can return an undefined Tensor instead of None for absent
weights. Python converts it to None, but native TorchScript retains it and
fails when another permutation consumes it.

Mutates and returns ``permuted`` without copying it. ``source`` and
``permuted`` may be the same object on identity-permutation paths.
"""
if source.weights_or_none() is None:
permuted._weights = None
return permuted


def symbolic_trace(
# pyre-ignore[24]
root: Union[torch.nn.Module, Callable],
Expand All @@ -36,6 +54,9 @@ def symbolic_trace(
`concrete_args` allows you to partially specialize your function, whether it's to
remove control flow or data structures.

Inserts absent-weight guards after FX-wrapped TorchRec KJT permutations.
This post-processing is idempotent when tracing an already guarded graph.

Args:
root (Union[torch.nn.Module, Callable]): Module or function to be traced and
converted into a Graph representation.
Expand All @@ -45,10 +66,41 @@ def symbolic_trace(
Returns:
GraphModule: a Module created from the recorded operations from ``root``.
"""
# Resolve private TorchRec helpers only when tracing.
from torchrec.modules.mc_modules import _mcc_lazy_init_inplace
from torchrec.quant.embedding_modules import _permute_kjt

_leaf_modules = list(UNTRACEABLE_MODULES)
if leaf_modules:
_leaf_modules.extend(leaf_modules)
return _symbolic_trace(root, concrete_args, _leaf_modules)
gm = _symbolic_trace(root, concrete_args, _leaf_modules)
Comment thread
eric-gecheng marked this conversation as resolved.
inserted = False
for node in list(gm.graph.nodes):
if node.op != "call_function" or node.target not in (
_mcc_lazy_init_inplace,
_permute_kjt,
):
continue
source = node.args[0] if node.args else node.kwargs["features"]
# Split exporters can trace an already guarded graph again.
if len(node.users) == 1:
user = next(iter(node.users))
if user.target == _restore_unweighted_kjt and user.args == (source, node):
continue
with gm.graph.inserting_after(node):
restored = gm.graph.call_function(
_restore_unweighted_kjt, args=(source, node)
)
# Keep the guard opaque when FX traces the generated module again.
restored.meta["is_wrapped"] = True
node.replace_all_uses_with(restored)
# Replace-all also rewrites the guard's input; restore it to avoid a cycle.
restored.args = (source, node)
Comment thread
eric-gecheng marked this conversation as resolved.
inserted = True
if inserted:
gm.graph.lint()
gm.recompile()
return gm


@torch.fx.wrap
Expand Down
188 changes: 187 additions & 1 deletion tzrec/utils/fx_util_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.

"""Unit tests for ``tzrec.utils.fx_util`` (CPU only).
"""Unit tests for ``tzrec.utils.fx_util``.

``fx_avg_batch_size`` was promoted here from a ``DlrmHSTU`` private helper
when ``listwise_rank_loss`` was generalized to ``RankModel``, so every rank
Expand All @@ -18,14 +18,200 @@
``mock.patch.object`` pattern as ``tzrec.utils.predict_util_test``.
"""

import os
import shutil
import subprocess
import sys
import unittest
from typing import Optional
from unittest import mock

import torch
import torch.distributed as dist
from parameterized import parameterized
from torchrec import KeyedJaggedTensor
from torchrec.modules.mc_modules import _mcc_lazy_init_inplace
from torchrec.quant.embedding_modules import _permute_kjt

from tzrec.utils import fx_util
from tzrec.utils.fx_util import fx_avg_batch_size
from tzrec.utils.test_util import (
gpu_unavailable,
make_test_dir,
mark_ci_scope,
parameterized_name_func,
)

# FX wrap registration is local to the caller's globals.
torch.fx.wrap(_mcc_lazy_init_inplace)
torch.fx.wrap(_permute_kjt)


class _PermuteKJT(torch.nn.Module):
def __init__(self, cached_order: bool = False):
super().__init__()
self.register_buffer(
"order",
torch.tensor([0, 4, 5, 1, 2, 3], dtype=torch.int32)
if cached_order
else None,
)

def forward(
self,
values: torch.Tensor,
lengths: torch.Tensor,
weights: Optional[torch.Tensor] = None,
):
features = KeyedJaggedTensor(
keys=["a", "b", "c", "d", "e", "f"],
values=values,
lengths=lengths,
weights=weights,
stride=1,
)
features = _mcc_lazy_init_inplace(
features=features,
feature_names=["a", "b", "c", "f", "d", "e"],
features_order=[0, 1, 2, 5, 3, 4],
created_feature_order=[True],
)
features = _permute_kjt(features, [0, 4, 5, 1, 2, 3], self.order)
return (
features.keys(),
features.values(),
features.lengths(),
features.weights_or_none(),
)


_NATIVE_SCRIPT_RUNNER = """
import importlib.util
from pathlib import Path
import sys
import torch

torch.set_num_threads(1)
root = Path(importlib.util.find_spec("fbgemm_gpu").origin).parent
print("Loading FBGEMM native operators", flush=True)
torch.ops.load_library(str(root / "fbgemm_gpu_py.so"))
assert "fbgemm_gpu" not in sys.modules
assert "tzrec" not in sys.modules
device = sys.argv[3]
print(f"Loading TorchScript on {device}", flush=True)
model = torch.jit.load(sys.argv[1], map_location=device)
cases = torch.load(sys.argv[2], weights_only=True)
with torch.no_grad():
for index, (inputs, expected) in enumerate(cases):
print(f"Running native case {index + 1}/{len(cases)}", flush=True)
inputs = [x.to(device) if isinstance(x, torch.Tensor) else x for x in inputs]
actual = model(*inputs)
assert len(actual) == len(expected)
for result, reference in zip(actual, expected):
if isinstance(reference, torch.Tensor):
torch.testing.assert_close(result.cpu(), reference)
else:
assert result == reference, (result, reference)
print(f"{len(cases)} native TorchScript cases passed on {device}")
"""


@mark_ci_scope("gpu", "h20")
class KJTPermutationTest(unittest.TestCase):
def setUp(self) -> None:
self.test_dir = make_test_dir("fx_kjt_")
self.addCleanup(shutil.rmtree, self.test_dir)

@parameterized.expand(
[("cpu", False), ("cpu", True), ("cuda:0", False), ("cuda:0", True)],
name_func=parameterized_name_func,
)
def test_native_scripted_permutation(self, device, cached_order) -> None:
if device.startswith("cuda") and gpu_unavailable[0]:
self.skipTest(gpu_unavailable[1])
gm = fx_util.symbolic_trace(_PermuteKJT(cached_order))
model_path = os.path.join(self.test_dir, "model.pt")
torch.jit.script(gm).save(model_path)

cases = []
order = [0, 3, 4, 1, 2, 5]
for sizes in ([1] * 6, [2, 0, 1, 3, 0, 1], [0] * 6):
for length_dtype in (torch.int32, torch.int64):
lengths = torch.tensor(sizes, dtype=length_dtype)
values = torch.arange(sum(sizes), dtype=torch.int64)
segments = torch.split(values, sizes)
expected_values = torch.cat([segments[i] for i in order])
for weighted in (False, True):
weights = values.float() + 0.25 if weighted else None
cases.append(
(
[values, lengths, weights],
(
["a", "d", "e", "b", "c", "f"],
expected_values,
lengths[order],
expected_values.float() + 0.25 if weighted else None,
),
)
)
cases_path = os.path.join(self.test_dir, "cases.pt")
torch.save(cases, cases_path)
try:
completed = subprocess.run(
[
sys.executable,
"-c",
_NATIVE_SCRIPT_RUNNER,
model_path,
cases_path,
device,
],
capture_output=True,
text=True,
timeout=120,
)
except subprocess.TimeoutExpired as error:
diagnostics = []
for name, output in (("stdout", error.stdout), ("stderr", error.stderr)):
if isinstance(output, bytes):
output = output.decode(errors="replace")
diagnostics.append(f"{name}:\n{output or ''}")
self.fail(
f"Native TorchScript timed out after {error.timeout}s\n"
+ "\n".join(diagnostics)
)
self.assertEqual(completed.returncode, 0, completed.stdout + completed.stderr)

def test_retrace_preserves_weights_guards(self) -> None:
gm = fx_util.symbolic_trace(_PermuteKJT())
retraced = fx_util.symbolic_trace(gm)
for graph in (gm.graph, retraced.graph):
guards = [
node
for node in graph.nodes
if node.target == fx_util._restore_unweighted_kjt
]
self.assertEqual(len(guards), 2)
graph.lint()
model = torch.jit.script(retraced)
actual = model(torch.arange(6), torch.ones(6, dtype=torch.int64))
torch.testing.assert_close(actual[1], torch.tensor([0, 3, 4, 1, 2, 5]))
self.assertIsNone(actual[3])

@parameterized.expand([(False,), (True,)], name_func=parameterized_name_func)
def test_torch_export_preserves_optional_weights(self, weighted) -> None:
values = torch.arange(6)
lengths = torch.ones(6, dtype=torch.int64)
weights = values.float() + 0.25 if weighted else None
args = (values, lengths, weights)
model = _PermuteKJT(cached_order=True)
expected = model(*args)
gm = fx_util.symbolic_trace(model)
exported = torch.export.export(gm, args)
actual = exported.module()(*args)
self.assertEqual(actual[0], expected[0])
for result, reference in zip(actual[1:], expected[1:]):
torch.testing.assert_close(result, reference)


class FxAvgBatchSizeTest(unittest.TestCase):
Expand Down
Loading