From 844a8f168151ae28c95244c5f12fca2d4fafb1b8 Mon Sep 17 00:00:00 2001 From: Naren Dasan Date: Fri, 28 Aug 2026 01:57:44 +0000 Subject: [PATCH 1/2] fix(dynamo): handle dynamic scatter and index updates --- py/torch_tensorrt/dynamo/_compiler.py | 3 + py/torch_tensorrt/dynamo/_refit.py | 1 + py/torch_tensorrt/dynamo/backend/backends.py | 3 + .../dynamo/conversion/impl/arange.py | 23 ++++-- .../dynamo/conversion/impl/select.py | 30 ++++++++ .../dynamo/lowering/_decompositions.py | 36 ++++++++- .../passes/complex_decomposition_adapter.py | 7 +- .../dynamo/partitioning/common.py | 12 ++- .../dynamo/conversion/test_index_put_aten.py | 74 +++++++++++++++++++ .../py/dynamo/lowering/test_decompositions.py | 53 +++++++++++++ .../models/test_dynamic_shape_user_bounds.py | 27 +++++++ 11 files changed, 257 insertions(+), 12 deletions(-) diff --git a/py/torch_tensorrt/dynamo/_compiler.py b/py/torch_tensorrt/dynamo/_compiler.py index 0e5e959350..2e0a6bbf17 100644 --- a/py/torch_tensorrt/dynamo/_compiler.py +++ b/py/torch_tensorrt/dynamo/_compiler.py @@ -400,6 +400,7 @@ def cross_compile_for_windows( decompose_attention, use_distributed_mode_trace, use_fp32_acc=use_fp32_acc, + graph_module=exported_program.graph_module, ) ) @@ -806,6 +807,7 @@ def compile( enable_experimental_decompositions, decompose_attention, use_distributed_mode_trace, + graph_module=exported_program.graph_module, use_fp32_acc=use_fp32_acc, ) ) @@ -2113,6 +2115,7 @@ def convert_exported_program_to_serialized_trt_engine( decompose_attention, use_distributed_mode_trace, use_fp32_acc=use_fp32_acc, + graph_module=exported_program.graph_module, ) ) diff --git a/py/torch_tensorrt/dynamo/_refit.py b/py/torch_tensorrt/dynamo/_refit.py index 2b3ddce96e..b74275de72 100644 --- a/py/torch_tensorrt/dynamo/_refit.py +++ b/py/torch_tensorrt/dynamo/_refit.py @@ -262,6 +262,7 @@ def refit_module_weights( settings.decompose_attention, settings.use_distributed_mode_trace, use_fp32_acc=settings.use_fp32_acc, + graph_module=new_weight_module.graph_module, ) ) new_gm = new_weight_module.module() diff --git a/py/torch_tensorrt/dynamo/backend/backends.py b/py/torch_tensorrt/dynamo/backend/backends.py index 38282f13c3..1efb7a3c6b 100644 --- a/py/torch_tensorrt/dynamo/backend/backends.py +++ b/py/torch_tensorrt/dynamo/backend/backends.py @@ -86,6 +86,7 @@ def aot_torch_tensorrt_aten_backend( settings.decompose_attention, settings.use_distributed_mode_trace, use_fp32_acc=settings.use_fp32_acc, + graph_module=gm, ) # This is added since detach lowering leads to alias nodes # Error - View operation returned a tensor that is the same as the input base tensor @@ -135,6 +136,7 @@ def aot_torch_tensorrt_aten_backend( aot_decomps = get_decompositions( settings.enable_experimental_decompositions, settings.decompose_attention, + graph_module=gm, use_fp32_acc=settings.use_fp32_acc, ) # Remove detach decompositions to avoid alias node errors. @@ -338,6 +340,7 @@ def _pretraced_backend( settings.decompose_attention, settings.use_distributed_mode_trace, use_fp32_acc=settings.use_fp32_acc, + graph_module=gm, ), ) diff --git a/py/torch_tensorrt/dynamo/conversion/impl/arange.py b/py/torch_tensorrt/dynamo/conversion/impl/arange.py index 81fbbae590..3e6352296c 100644 --- a/py/torch_tensorrt/dynamo/conversion/impl/arange.py +++ b/py/torch_tensorrt/dynamo/conversion/impl/arange.py @@ -1,9 +1,9 @@ from typing import Optional, Union +import numpy as np import tensorrt as trt import torch from tensorrt import ITensor as TRTTensor -from torch._subclasses.fake_tensor import unset_fake_temporarily from torch.fx.node import Target from torch_tensorrt import _enums from torch_tensorrt.dynamo.conversion import impl @@ -112,10 +112,19 @@ def arange( else: # All arguments are static, so evaluate the sequence eagerly and freeze it - # into the engine as a constant. Letting torch pick the dtype preserves - # PyTorch's promotion rules (float result if any argument is a float). - with unset_fake_temporarily(): - values = torch.arange(start, end, step, dtype=dtype) - if values.dtype == torch.int64: - values = values.to(torch.int32) + # into the engine as a constant. NumPy avoids creating a FakeTensor when + # conversion runs inside torch.compile's active FakeTensorMode. + resolved_dtype = dtype + if resolved_dtype is None and any( + isinstance(value, float) for value in (start, end, step) + ): + resolved_dtype = torch.get_default_dtype() + np_dtype = ( + _enums.dtype._from(resolved_dtype).to(np.dtype) + if resolved_dtype is not None + else None + ) + values = np.arange(start, end, step, dtype=np_dtype) + if values.dtype == np.int64: + values = values.astype(np.int32) return get_trt_tensor(ctx, values, f"{name}_arange_const") diff --git a/py/torch_tensorrt/dynamo/conversion/impl/select.py b/py/torch_tensorrt/dynamo/conversion/impl/select.py index 945a5f19d2..7360bb03d6 100644 --- a/py/torch_tensorrt/dynamo/conversion/impl/select.py +++ b/py/torch_tensorrt/dynamo/conversion/impl/select.py @@ -1005,6 +1005,29 @@ def index_put_converter( values_permuted, expected_shape, ) + elif K == 1 and len(values.shape) == 1 and F and I[0] > max(F): + # For data[..., idx] = values, a 1-D values tensor carries the + # index extent and broadcasts across the preceding free dims. + # The converter's internal layout is (N, *F), so make that + # index axis explicit before expanding. This is essential when N + # is dynamic: using -1 for both axes would keep the original N + # extent instead of broadcasting the free dimensions. + values_reshaped = impl.shuffle.reshape( + ctx, + target, + source_ir, + f"{name}_reshape_index_values", + values, + (N,) + (1,) * len(F), + ) + values_expanded = impl.slice.expand( + ctx, + target, + source_ir, + f"{name}_expand_values", + values_reshaped, + expected_shape, + ) elif ( K > 0 and N in values_shape @@ -1089,6 +1112,13 @@ def index_put_converter( values_expanded, (-1,), ) + if flattened_values.dtype != input_tensor.dtype: + flattened_values = cast_trt_tensor( + ctx, + flattened_values, + input_tensor.dtype, + f"{name}_values_cast", + ) indices_cat = cast_trt_tensor(ctx, indices_cat, trt.int32, f"{name}_idx_int32") scatter_layer = ctx.net.add_scatter( input_tensor, diff --git a/py/torch_tensorrt/dynamo/lowering/_decompositions.py b/py/torch_tensorrt/dynamo/lowering/_decompositions.py index a32aef1b28..1ede81abac 100644 --- a/py/torch_tensorrt/dynamo/lowering/_decompositions.py +++ b/py/torch_tensorrt/dynamo/lowering/_decompositions.py @@ -717,11 +717,35 @@ def fp32_accumulation_decomposition(*args: Any, **kwargs: Any) -> Any: } +def _has_symbolic_scatter_add_extent( + graph_module: Optional[torch.fx.GraphModule], +) -> bool: + """Return whether scatter_add would require unrolling a symbolic extent.""" + if graph_module is None: + return False + + for node in graph_module.graph.nodes: + if node.target != torch.ops.aten.scatter_add.default or len(node.args) < 4: + continue + dim = node.args[1] + src_node = node.args[3] + if not isinstance(dim, int) or not isinstance(src_node, torch.fx.Node): + continue + src_val = src_node.meta.get("val", src_node.meta.get("example_value")) + if not isinstance(src_val, torch.Tensor) or not src_val.ndim: + continue + if isinstance(src_val.shape[get_positive_dim(dim, src_val.ndim)], torch.SymInt): + return True + + return False + + def get_decompositions( enable_experimental_decompositions: bool = False, decompose_attention: bool = False, use_distributed_mode_trace: bool = False, use_fp32_acc: bool = False, + graph_module: Optional[torch.fx.GraphModule] = None, ) -> Dict[OpOverload, Callable[[Any], Any]]: trt_decomps = ( dict(TORCH_TRT_DECOMPOSITIONS) @@ -749,7 +773,7 @@ def get_decompositions( for decomp in _core_aten_decompositions if decomp not in discard_decompositions } - return {**CORE_ATEN_DECOMPOSITIONS_FILTERED, **trt_decomps} + decompositions = {**CORE_ATEN_DECOMPOSITIONS_FILTERED, **trt_decomps} else: # changes made here due to torch2.6 changes https://github.com/pytorch/pytorch/pull/135080 decomp_table = {} @@ -763,8 +787,16 @@ def get_decompositions( and decomp not in ATTENTION_DECOMPOSITION_OPS } - return { + decompositions = { **ENABLED_TORCH_DECOMPOSITIONS, **DECOMP_TABLE_FILTERED, **trt_decomps, } + + if _has_symbolic_scatter_add_extent(graph_module): + # The custom decomposition uses a Python range over this extent. + # Keeping the op lets partitioning fall back to Torch without trying + # to specialize an unbacked or otherwise dynamic SymInt. + decompositions.pop(torch.ops.aten.scatter_add.default, None) + + return decompositions diff --git a/py/torch_tensorrt/dynamo/lowering/passes/complex_decomposition_adapter.py b/py/torch_tensorrt/dynamo/lowering/passes/complex_decomposition_adapter.py index 05a6f4f2ff..3eec4b62b4 100644 --- a/py/torch_tensorrt/dynamo/lowering/passes/complex_decomposition_adapter.py +++ b/py/torch_tensorrt/dynamo/lowering/passes/complex_decomposition_adapter.py @@ -100,7 +100,9 @@ def complex_decomposition_adapter( # anything, since it builds a new GraphModule rather than editing # the existing one in place). decomposed_gm = decompose_complex_in_graph( - gm, flat_args, decompositions=_trt_decomposition_table(settings) + gm, + flat_args, + decompositions=_trt_decomposition_table(settings, gm), ) except Exception as e: # decompose_complex_in_graph is upstream, experimental PyTorch code @@ -141,7 +143,7 @@ def complex_decomposition_adapter( def _trt_decomposition_table( - settings: CompilationSettings, + settings: CompilationSettings, graph_module: GraphModule ) -> dict[Any, Any]: """The op set the rest of the TRT flow expects to see. @@ -160,6 +162,7 @@ def _trt_decomposition_table( settings.decompose_attention, settings.use_distributed_mode_trace, use_fp32_acc=settings.use_fp32_acc, + graph_module=graph_module, ) diff --git a/py/torch_tensorrt/dynamo/partitioning/common.py b/py/torch_tensorrt/dynamo/partitioning/common.py index a0942a5d0b..07468d3037 100644 --- a/py/torch_tensorrt/dynamo/partitioning/common.py +++ b/py/torch_tensorrt/dynamo/partitioning/common.py @@ -118,7 +118,17 @@ def construct_dynamic_input( ) unwrapped_min_max_opt["min"] = 1 else: - unwrapped_min_max_opt["min"] = min_max_opt["min"] + min_bound = int(min_max_opt["min"]) + if min_bound < 1: + logger.warning( + "Dynamic input %s (shape: %s) has lower bound %d for dim %d. " + "TensorRT profiles require dimensions >= 1; clamping it to 1.", + name, + input_shape, + min_bound, + d, + ) + unwrapped_min_max_opt["min"] = max(1, min_bound) if "max" not in min_max_opt or min_max_opt["max"] is None: logger.warning( diff --git a/tests/py/dynamo/conversion/test_index_put_aten.py b/tests/py/dynamo/conversion/test_index_put_aten.py index 5536c63a3c..51ebb7726e 100644 --- a/tests/py/dynamo/conversion/test_index_put_aten.py +++ b/tests/py/dynamo/conversion/test_index_put_aten.py @@ -509,6 +509,80 @@ def forward(self, source_tensor, indices_tensor, value_tensor): torch.allclose(result, torch_output, atol=1e-4, rtol=1e-4) + def test_updates_are_cast_to_int64_destination_dtype(self): + class IndexPutArgmax(torch.nn.Module): + def forward(self, data, idx, scores): + values = torch.ops.aten.argmax.default(scores, 1) + return torch.ops.aten.index_put.default(data, [idx], values) + + data = torch.zeros(8, dtype=torch.int64, device="cuda") + idx = torch.tensor([1, 3], dtype=torch.int64, device="cuda") + scores = torch.randn((2, 5), device="cuda") + model = IndexPutArgmax().eval().cuda() + expected = model(data, idx, scores) + + exported = torch.export.export(model, (data, idx, scores)) + compiled = torchtrt.dynamo.compile( + exported, + arg_inputs=[data, idx, scores], + min_block_size=1, + pass_through_build_failures=True, + ) + + self.assertEqual(compiled(data, idx, scores), expected) + + def test_dynamic_index_broadcasts_1d_values_across_free_dim(self): + class IndexPutFreeDim(torch.nn.Module): + def forward(self, data, idx): + values = idx.to(torch.float32) + return torch.ops.aten.index_put.default(data, [None, idx], values) + + data = torch.zeros((32, 64), device="cuda") + idx = torch.tensor([1, 3, 5, 7], dtype=torch.int64, device="cuda") + index_length = torch.export.Dim("index_length", min=1, max=16) + model = IndexPutFreeDim().eval().cuda() + + exported = torch.export.export( + model, + (data, idx), + dynamic_shapes={"data": {}, "idx": {0: index_length}}, + ) + compiled = torchtrt.dynamo.compile( + exported, + arg_inputs=[ + torchtrt.Input(shape=(32, 64), dtype=torch.float32), + torchtrt.Input( + min_shape=(1,), + opt_shape=(4,), + max_shape=(16,), + dtype=torch.int64, + ), + ], + min_block_size=1, + pass_through_build_failures=True, + ) + + for runtime_idx in ( + idx, + torch.tensor([2, 6], dtype=torch.int64, device="cuda"), + ): + self.assertEqual( + compiled(data, runtime_idx), + model(data, runtime_idx), + ) + torch._dynamo.reset() + compile_idx = idx.clone() + torch._dynamo.mark_dynamic(compile_idx, 0) + compiled_backend = torch.compile( + model, + backend="tensorrt", + options={"pass_through_build_failures": True, "min_block_size": 1}, + ) + self.assertEqual( + compiled_backend(data, compile_idx), + model(data, compile_idx), + ) + def test_index_put_dynamic_index_length(self): """index_put where the index tensor itself has a dynamic length (N dynamic). diff --git a/tests/py/dynamo/lowering/test_decompositions.py b/tests/py/dynamo/lowering/test_decompositions.py index 89aacfa262..02c467f335 100644 --- a/tests/py/dynamo/lowering/test_decompositions.py +++ b/tests/py/dynamo/lowering/test_decompositions.py @@ -1236,6 +1236,59 @@ def forward(self, input): f"The optimized model results shape and torch model results shape should be equal in empty_stride", ) + def test_scatter_add_symbolic_extent_skips_unrolled_decomposition(self): + class ScatterAdd(torch.nn.Module): + def forward(self, base, index, src): + return torch.ops.aten.scatter_add.default(base, 0, index, src) + + model = ScatterAdd().eval() + base = torch.zeros((16, 8)) + index = torch.zeros((4, 8), dtype=torch.int64) + src = torch.ones((4, 8)) + src_rows = torch.export.Dim("src_rows", min=1, max=16) + + dynamic_export = torch.export.export( + model, + (base, index, src), + dynamic_shapes={ + "base": {}, + "index": {0: src_rows}, + "src": {0: src_rows}, + }, + ) + dynamic_decompositions = get_decompositions( + graph_module=dynamic_export.graph_module + ) + self.assertNotIn(torch.ops.aten.scatter_add.default, dynamic_decompositions) + + dynamic_lowered = dynamic_export.run_decompositions(dynamic_decompositions) + self.assertTrue( + any( + node.target == torch.ops.aten.scatter_add.default + for node in dynamic_lowered.graph_module.graph.nodes + ) + ) + + # Dynamo records fake tensor metadata under example_value, rather than val. + dynamic_scatter = next( + node + for node in dynamic_export.graph_module.graph.nodes + if node.target == torch.ops.aten.scatter_add.default + ) + dynamic_src = dynamic_scatter.args[3] + self.assertIsInstance(dynamic_src, torch.fx.Node) + dynamic_src.meta["example_value"] = dynamic_src.meta.pop("val") + dynamo_decompositions = get_decompositions( + graph_module=dynamic_export.graph_module + ) + self.assertNotIn(torch.ops.aten.scatter_add.default, dynamo_decompositions) + + static_export = torch.export.export(model, (base, index, src)) + static_decompositions = get_decompositions( + graph_module=static_export.graph_module + ) + self.assertIn(torch.ops.aten.scatter_add.default, static_decompositions) + @parameterized.expand( [ ( diff --git a/tests/py/dynamo/models/test_dynamic_shape_user_bounds.py b/tests/py/dynamo/models/test_dynamic_shape_user_bounds.py index 584e8b3227..5eda524d8f 100644 --- a/tests/py/dynamo/models/test_dynamic_shape_user_bounds.py +++ b/tests/py/dynamo/models/test_dynamic_shape_user_bounds.py @@ -11,6 +11,7 @@ import torch_tensorrt as torchtrt from torch_tensorrt._Input import Input from torch_tensorrt.dynamo._compiler import _build_user_symbol_bounds +from torch_tensorrt.dynamo.partitioning.common import construct_dynamic_input from torch_tensorrt.dynamo.utils import extract_var_range_info assertions = unittest.TestCase() @@ -560,5 +561,31 @@ def test_dim_dynamic_save_preserves_range_constraints(tmpdir): trt_module(too_big) +@pytest.mark.unit +def test_construct_dynamic_input_clamps_zero_minimum(): + """Data-dependent extents can include zero, but TRT profiles cannot.""" + + class Nonzero(torch.nn.Module): + def forward(self, x): + return torch.nonzero(x) + + exported = torch.export.export( + Nonzero(), (torch.tensor([True, False, True, False]),) + ) + nonzero = next( + node + for node in exported.graph.nodes + if node.target == torch.ops.aten.nonzero.default + ) + fake_value = nonzero.meta["val"] + assert isinstance(fake_value.shape[0], torch.SymInt) + + input_spec = construct_dynamic_input( + fake_value.shape, fake_value.dtype, name="nonzero_output" + ) + assert input_spec.shape["min_shape"] == (1, 1) + assert input_spec.shape["max_shape"] == (4, 1) + + if __name__ == "__main__": pytest.main([__file__, "-v"]) From 15dcad7599a421a12488faaeac046b8cf7f2bc27 Mon Sep 17 00:00:00 2001 From: apbose Date: Tue, 1 Sep 2026 23:49:14 -0700 Subject: [PATCH 2/2] fix(dynamo): preserve BF16 static arange constants Generate static arange values through a NumPy-supported staging dtype when the requested TensorRT dtype has no NumPy representation. Cast during constant creation and add BF16 converter coverage. --- .../dynamo/conversion/impl/arange.py | 19 +++++++++++++------ .../py/dynamo/conversion/test_arange_aten.py | 13 +++++++++++++ 2 files changed, 26 insertions(+), 6 deletions(-) diff --git a/py/torch_tensorrt/dynamo/conversion/impl/arange.py b/py/torch_tensorrt/dynamo/conversion/impl/arange.py index 3e6352296c..a263167484 100644 --- a/py/torch_tensorrt/dynamo/conversion/impl/arange.py +++ b/py/torch_tensorrt/dynamo/conversion/impl/arange.py @@ -119,12 +119,19 @@ def arange( isinstance(value, float) for value in (start, end, step) ): resolved_dtype = torch.get_default_dtype() - np_dtype = ( - _enums.dtype._from(resolved_dtype).to(np.dtype) - if resolved_dtype is not None - else None - ) + constant_dtype = None + if resolved_dtype is not None: + try: + np_dtype = _enums.dtype._from(resolved_dtype).to(np.dtype) + except TypeError: + # Some TensorRT dtypes, such as BF16, have no NumPy + # representation. Build the sequence in NumPy's inferred dtype + # and let constant creation cast it to the requested dtype. + np_dtype = None + constant_dtype = resolved_dtype + else: + np_dtype = None values = np.arange(start, end, step, dtype=np_dtype) if values.dtype == np.int64: values = values.astype(np.int32) - return get_trt_tensor(ctx, values, f"{name}_arange_const") + return get_trt_tensor(ctx, values, f"{name}_arange_const", dtype=constant_dtype) diff --git a/tests/py/dynamo/conversion/test_arange_aten.py b/tests/py/dynamo/conversion/test_arange_aten.py index e48f14b93e..e4ea09cd81 100644 --- a/tests/py/dynamo/conversion/test_arange_aten.py +++ b/tests/py/dynamo/conversion/test_arange_aten.py @@ -43,6 +43,19 @@ def forward(self, x): use_dynamo_tracer=True, ) + def test_arange_static_non_numpy_type(self): + class Arange(nn.Module): + def forward(self, x): + return torch.ops.aten.arange.start_step( + 0, 5, 1, dtype=torch.bfloat16, device=x.device + ) + + self.run_test( + Arange(), + [torch.randn(1, 1)], + use_dynamo_tracer=True, + ) + def test_arange_dynamic_int32(self): class Arange(nn.Module): def forward(self, end_tensor):