diff --git a/py/torch_tensorrt/dynamo/_compiler.py b/py/torch_tensorrt/dynamo/_compiler.py index 0e5e959350..233a935942 100644 --- a/py/torch_tensorrt/dynamo/_compiler.py +++ b/py/torch_tensorrt/dynamo/_compiler.py @@ -40,6 +40,7 @@ from torch_tensorrt.dynamo.debug._DebuggerConfig import DebuggerConfig from torch_tensorrt.dynamo.debug._supports_debugger import fn_supports_debugger from torch_tensorrt.dynamo.lowering import ( + filter_decomposition_table, get_decompositions, post_lowering, pre_export_lowering, @@ -395,11 +396,14 @@ def cross_compile_for_windows( logger.info("Compilation Settings: %s\n", settings) exported_program = pre_export_lowering(exported_program, settings) exported_program = exported_program.run_decompositions( - get_decompositions( - enable_experimental_decompositions, - decompose_attention, - use_distributed_mode_trace, - use_fp32_acc=use_fp32_acc, + filter_decomposition_table( + get_decompositions( + enable_experimental_decompositions, + decompose_attention, + use_distributed_mode_trace, + use_fp32_acc=use_fp32_acc, + ), + exported_program.graph_module, ) ) @@ -802,11 +806,14 @@ def compile( logger.info("Compilation Settings: %s\n", settings) exported_program = pre_export_lowering(exported_program, settings) exported_program = exported_program.run_decompositions( - get_decompositions( - enable_experimental_decompositions, - decompose_attention, - use_distributed_mode_trace, - use_fp32_acc=use_fp32_acc, + filter_decomposition_table( + get_decompositions( + enable_experimental_decompositions, + decompose_attention, + use_distributed_mode_trace, + use_fp32_acc=use_fp32_acc, + ), + exported_program.graph_module, ) ) @@ -2108,11 +2115,14 @@ def convert_exported_program_to_serialized_trt_engine( logger.info("Compilation Settings: %s\n", settings) exported_program = pre_export_lowering(exported_program, settings) exported_program = exported_program.run_decompositions( - get_decompositions( - enable_experimental_decompositions, - decompose_attention, - use_distributed_mode_trace, - use_fp32_acc=use_fp32_acc, + filter_decomposition_table( + get_decompositions( + enable_experimental_decompositions, + decompose_attention, + use_distributed_mode_trace, + use_fp32_acc=use_fp32_acc, + ), + exported_program.graph_module, ) ) diff --git a/py/torch_tensorrt/dynamo/_refit.py b/py/torch_tensorrt/dynamo/_refit.py index 2b3ddce96e..689d12a795 100644 --- a/py/torch_tensorrt/dynamo/_refit.py +++ b/py/torch_tensorrt/dynamo/_refit.py @@ -25,6 +25,7 @@ from torch_tensorrt.dynamo.conversion.truncate_double import repair_double_inputs from torch_tensorrt.dynamo.lowering import ( clean_up_graph_after_modifications, + filter_decomposition_table, get_decompositions, post_lowering, pre_export_lowering, @@ -257,11 +258,14 @@ def refit_module_weights( ) new_weight_module = pre_export_lowering(new_weight_module, settings) new_weight_module = new_weight_module.run_decompositions( - get_decompositions( - settings.enable_experimental_decompositions, - settings.decompose_attention, - settings.use_distributed_mode_trace, - use_fp32_acc=settings.use_fp32_acc, + filter_decomposition_table( + get_decompositions( + settings.enable_experimental_decompositions, + settings.decompose_attention, + settings.use_distributed_mode_trace, + use_fp32_acc=settings.use_fp32_acc, + ), + 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..4b9d6833f4 100644 --- a/py/torch_tensorrt/dynamo/backend/backends.py +++ b/py/torch_tensorrt/dynamo/backend/backends.py @@ -15,6 +15,7 @@ from torch_tensorrt.dynamo import CompilationSettings from torch_tensorrt.dynamo._compiler import compile_module from torch_tensorrt.dynamo.lowering import ( + filter_decomposition_table, get_decompositions, post_lowering, remove_detach, @@ -81,11 +82,14 @@ def aot_torch_tensorrt_aten_backend( _pretraced_backend, settings=settings, engine_cache=engine_cache ) settings_aot_autograd = {} - settings_aot_autograd["decompositions"] = get_decompositions( - settings.enable_experimental_decompositions, - settings.decompose_attention, - settings.use_distributed_mode_trace, - use_fp32_acc=settings.use_fp32_acc, + settings_aot_autograd["decompositions"] = filter_decomposition_table( + get_decompositions( + settings.enable_experimental_decompositions, + settings.decompose_attention, + settings.use_distributed_mode_trace, + use_fp32_acc=settings.use_fp32_acc, + ), + 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 @@ -132,10 +136,13 @@ def aot_torch_tensorrt_aten_backend( _pretraced_backend_autograd = functools.partial( _pretraced_backend, settings=aot_settings, engine_cache=engine_cache ) - aot_decomps = get_decompositions( - settings.enable_experimental_decompositions, - settings.decompose_attention, - use_fp32_acc=settings.use_fp32_acc, + aot_decomps = filter_decomposition_table( + get_decompositions( + settings.enable_experimental_decompositions, + settings.decompose_attention, + use_fp32_acc=settings.use_fp32_acc, + ), + gm, ) # Remove detach decompositions to avoid alias node errors. to_delete = {k for k in aot_decomps if "detach" in k._name} @@ -333,11 +340,14 @@ def _pretraced_backend( gm, sample_inputs, trace_joint=False, - decompositions=get_decompositions( - settings.enable_experimental_decompositions, - settings.decompose_attention, - settings.use_distributed_mode_trace, - use_fp32_acc=settings.use_fp32_acc, + decompositions=filter_decomposition_table( + get_decompositions( + settings.enable_experimental_decompositions, + settings.decompose_attention, + settings.use_distributed_mode_trace, + use_fp32_acc=settings.use_fp32_acc, + ), + gm, ), ) diff --git a/py/torch_tensorrt/dynamo/conversion/impl/arange.py b/py/torch_tensorrt/dynamo/conversion/impl/arange.py index 81fbbae590..a263167484 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,26 @@ 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) - return get_trt_tensor(ctx, values, f"{name}_arange_const") + # 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() + 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", dtype=constant_dtype) 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/__init__.py b/py/torch_tensorrt/dynamo/lowering/__init__.py index 3c73e42f86..e13f65ac79 100644 --- a/py/torch_tensorrt/dynamo/lowering/__init__.py +++ b/py/torch_tensorrt/dynamo/lowering/__init__.py @@ -1,3 +1,6 @@ +from torch_tensorrt.dynamo.lowering._SubgraphBuilder import SubgraphBuilder + +from ._decomposition_filter import filter_decomposition_table from ._decomposition_groups import ( TORCH_TRT_DECOMPOSITIONS, torch_disabled_decompositions, @@ -5,4 +8,3 @@ ) from ._decompositions import get_decompositions # noqa: F401 from .passes import * -from torch_tensorrt.dynamo.lowering._SubgraphBuilder import SubgraphBuilder diff --git a/py/torch_tensorrt/dynamo/lowering/_decomposition_filter.py b/py/torch_tensorrt/dynamo/lowering/_decomposition_filter.py new file mode 100644 index 0000000000..fff8d74cfd --- /dev/null +++ b/py/torch_tensorrt/dynamo/lowering/_decomposition_filter.py @@ -0,0 +1,45 @@ +from __future__ import annotations + +from typing import Any, Callable, Dict + +import torch +from torch._ops import OpOverload +from torch_tensorrt.dynamo.conversion.converter_utils import get_positive_dim + + +def _has_symbolic_scatter_add_extent(graph_module: torch.fx.GraphModule) -> bool: + """Return whether scatter_add would require unrolling a symbolic extent.""" + 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 filter_decomposition_table( + decompositions: Dict[OpOverload, Callable[[Any], Any]], + graph_module: torch.fx.GraphModule, +) -> Dict[OpOverload, Callable[[Any], Any]]: + """Filter graph-incompatible entries from a decomposition table. + + Decomposition tables are keyed by operator rather than individual FX node, + so encountering one dynamic ``scatter_add`` conservatively keeps every + ``scatter_add`` in this graph intact for partitioning to handle. + """ + filtered_decompositions = dict(decompositions) + 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. + filtered_decompositions.pop(torch.ops.aten.scatter_add.default, None) + + return filtered_decompositions diff --git a/py/torch_tensorrt/dynamo/lowering/_decompositions.py b/py/torch_tensorrt/dynamo/lowering/_decompositions.py index a32aef1b28..eee96c84b8 100644 --- a/py/torch_tensorrt/dynamo/lowering/_decompositions.py +++ b/py/torch_tensorrt/dynamo/lowering/_decompositions.py @@ -749,7 +749,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 +763,10 @@ def get_decompositions( and decomp not in ATTENTION_DECOMPOSITION_OPS } - return { + decompositions = { **ENABLED_TORCH_DECOMPOSITIONS, **DECOMP_TABLE_FILTERED, **trt_decomps, } + + 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..52124b5974 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. @@ -153,13 +155,19 @@ def _trt_decomposition_table( would have handled them, and the partitioner then cuts the graph around them. Handing the retrace the same table keeps the two in step. """ - from torch_tensorrt.dynamo.lowering import get_decompositions + from torch_tensorrt.dynamo.lowering import ( + filter_decomposition_table, + get_decompositions, + ) - return get_decompositions( - settings.enable_experimental_decompositions, - settings.decompose_attention, - settings.use_distributed_mode_trace, - use_fp32_acc=settings.use_fp32_acc, + return filter_decomposition_table( + get_decompositions( + settings.enable_experimental_decompositions, + settings.decompose_attention, + settings.use_distributed_mode_trace, + use_fp32_acc=settings.use_fp32_acc, + ), + 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_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): 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..5669f728a2 100644 --- a/tests/py/dynamo/lowering/test_decompositions.py +++ b/tests/py/dynamo/lowering/test_decompositions.py @@ -9,7 +9,10 @@ ) from torch.testing._internal.common_utils import TestCase, run_tests from torch_tensorrt.dynamo._settings import CompilationSettings -from torch_tensorrt.dynamo.lowering import get_decompositions +from torch_tensorrt.dynamo.lowering import ( + filter_decomposition_table, + get_decompositions, +) from torch_tensorrt.dynamo.lowering.passes._aten_lowering_pass import post_lowering from torch_tensorrt.dynamo.utils import ATOL, RTOL @@ -1236,6 +1239,61 @@ 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_filters_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}, + }, + ) + base_decompositions = get_decompositions() + dynamic_decompositions = filter_decomposition_table( + base_decompositions, dynamic_export.graph_module + ) + self.assertNotIn(torch.ops.aten.scatter_add.default, dynamic_decompositions) + self.assertIn(torch.ops.aten.scatter_add.default, base_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 = filter_decomposition_table( + get_decompositions(), 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 = filter_decomposition_table( + get_decompositions(), 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"])