From 3ae4df4e1d4143280f0f5bcb3baff16567df1115 Mon Sep 17 00:00:00 2001 From: gecheng Date: Sun, 20 Sep 2026 11:55:44 +0800 Subject: [PATCH 1/2] [bugfix] preserve absent KJT weights in exported graphs FBGEMM CUDA can return an undefined tensor for absent weights, causing native TorchScript inference to fail on a subsequent permutation. Insert FX guards after TorchRec MC and quantized permutations so unweighted inputs retain None in exported models. --- tzrec/utils/fx_util.py | 42 +++++++++- tzrec/utils/fx_util_test.py | 158 +++++++++++++++++++++++++++++++++++- 2 files changed, 197 insertions(+), 3 deletions(-) diff --git a/tzrec/utils/fx_util.py b/tzrec/utils/fx_util.py index 4e863866..7787b111 100644 --- a/tzrec/utils/fx_util.py +++ b/tzrec/utils/fx_util.py @@ -13,8 +13,10 @@ 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 +from torchrec.modules.mc_modules import _mcc_lazy_init_inplace +from torchrec.quant.embedding_modules import _permute_kjt # Modules whose forward FX cannot record -- they branch on tensor values or # turn them into Python ints -- so tracing keeps them opaque and TorchScript @@ -22,6 +24,21 @@ 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. + + 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. + """ + if source.weights_or_none() is None: + permuted._weights = None + return permuted + + def symbolic_trace( # pyre-ignore[24] root: Union[torch.nn.Module, Callable], @@ -48,7 +65,28 @@ def symbolic_trace( _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) + 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"] + 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) + ) + restored.meta["is_wrapped"] = True + node.replace_all_uses_with(restored) + restored.args = (source, node) + gm.graph.lint() + gm.recompile() + return gm @torch.fx.wrap diff --git a/tzrec/utils/fx_util_test.py b/tzrec/utils/fx_util_test.py index 288cd71a..22fd4f87 100644 --- a/tzrec/utils/fx_util_test.py +++ b/tzrec/utils/fx_util_test.py @@ -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 @@ -18,14 +18,170 @@ ``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, +) + +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 +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] +model = torch.jit.load(sys.argv[1], map_location=device) +cases = torch.load(sys.argv[2], weights_only=True) +with torch.no_grad(): + for inputs, expected in cases: + 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) + completed = subprocess.run( + [ + sys.executable, + "-c", + _NATIVE_SCRIPT_RUNNER, + model_path, + cases_path, + device, + ], + capture_output=True, + text=True, + timeout=120, + ) + 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]) class FxAvgBatchSizeTest(unittest.TestCase): From 701fbc54a6c76c904c7d3958d4707a96e0756c4f Mon Sep 17 00:00:00 2001 From: gecheng Date: Sun, 20 Sep 2026 15:47:11 +0800 Subject: [PATCH 2/2] [refactor] refine KJT export guards and diagnostics Keep private TorchRec imports within tracing and skip graph recompilation when no guard is inserted. Document guard mutation and retracing, preserve native-test timeout output, and cover torch.export with and without weights. --- tzrec/utils/fx_util.py | 22 +++++++++++--- tzrec/utils/fx_util_test.py | 58 ++++++++++++++++++++++++++++--------- 2 files changed, 62 insertions(+), 18 deletions(-) diff --git a/tzrec/utils/fx_util.py b/tzrec/utils/fx_util.py index 7787b111..7b6fdd63 100644 --- a/tzrec/utils/fx_util.py +++ b/tzrec/utils/fx_util.py @@ -15,8 +15,6 @@ import torch.distributed as dist from torchrec import JaggedTensor, KeyedJaggedTensor, KeyedTensor from torchrec.fx import symbolic_trace as _symbolic_trace -from torchrec.modules.mc_modules import _mcc_lazy_init_inplace -from torchrec.quant.embedding_modules import _permute_kjt # Modules whose forward FX cannot record -- they branch on tensor values or # turn them into Python ints -- so tracing keeps them opaque and TorchScript @@ -33,6 +31,9 @@ def _restore_unweighted_kjt( 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 @@ -53,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. @@ -62,10 +66,15 @@ 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) gm = _symbolic_trace(root, concrete_args, _leaf_modules) + inserted = False for node in list(gm.graph.nodes): if node.op != "call_function" or node.target not in ( _mcc_lazy_init_inplace, @@ -73,6 +82,7 @@ def symbolic_trace( ): 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): @@ -81,11 +91,15 @@ def symbolic_trace( 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) - gm.graph.lint() - gm.recompile() + inserted = True + if inserted: + gm.graph.lint() + gm.recompile() return gm diff --git a/tzrec/utils/fx_util_test.py b/tzrec/utils/fx_util_test.py index 22fd4f87..ee462cc2 100644 --- a/tzrec/utils/fx_util_test.py +++ b/tzrec/utils/fx_util_test.py @@ -42,6 +42,7 @@ 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) @@ -92,14 +93,17 @@ def forward( 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 inputs, expected in cases: + 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) @@ -152,19 +156,30 @@ def test_native_scripted_permutation(self, device, cached_order) -> None: ) cases_path = os.path.join(self.test_dir, "cases.pt") torch.save(cases, cases_path) - completed = subprocess.run( - [ - sys.executable, - "-c", - _NATIVE_SCRIPT_RUNNER, - model_path, - cases_path, - device, - ], - capture_output=True, - text=True, - timeout=120, - ) + 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: @@ -183,6 +198,21 @@ def test_retrace_preserves_weights_guards(self) -> None: 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): """The local/global rescaling factor of ``enable_global_average_loss``."""