From e1704e2897206259ee89aa5e5292759cf49c21d4 Mon Sep 17 00:00:00 2001 From: Anthony Shoumikhin Date: Fri, 2 Oct 2026 00:12:26 -0700 Subject: [PATCH] Accept AOTI's compact clones of view buffers in the CUDA weight collector AOTInductor (AOTI) clones every buffer into compact storage that starts at the view, but builds the buffer's TensorProperties from the original. When a buffer is a view into a larger tensor, the CUDA weight collector judged the clone by the original storage and rejected the export: a contiguous view failed the storage size check ("cloned storage is smaller than its TensorProperties"), and a strided view, whose TensorProperties has no storage size, failed the bounds check ("requires 236 bytes from a 232-byte cloned storage"). Yet the clone holds every byte the view needs. When the value is such a clone (a different storage, with the same shape and strides as its TensorProperties), skip the storage size check and take the offset from the clone that is written. The clone must also span its storage exactly. Any other value takes the existing path. The check that the view fits inside the written bytes still runs, so a truncated clone is still rejected. --- backends/cuda/cuda_weight_collector.py | 29 +++++- backends/cuda/tests/test_cuda_partitioner.py | 100 ++++++++++++++++++- 2 files changed, 124 insertions(+), 5 deletions(-) diff --git a/backends/cuda/cuda_weight_collector.py b/backends/cuda/cuda_weight_collector.py index 81109d743f5..ae361c3aef7 100644 --- a/backends/cuda/cuda_weight_collector.py +++ b/backends/cuda/cuda_weight_collector.py @@ -399,15 +399,29 @@ def materialize( storages: Dict[str, FileBackedData] = {} for fqn, (tensor, properties) in weights.items(): + is_offgraph_kv = _is_offgraph_kv_fqn(fqn) storage = tensor.untyped_storage() storage_nbytes = storage.nbytes() + # AOTI clones buffers into compact storage that starts at the view, + # with the same shape and strides, so the original storage size and + # offset do not describe it. + is_compact_clone = ( + not is_offgraph_kv + and getattr(properties, "storage_ptr", None) + not in (None, storage.data_ptr()) + and tuple(tensor.shape) == tuple(getattr(properties, "shape", ())) + and tuple(tensor.stride()) == tuple(getattr(properties, "stride", ())) + ) del storage device_type = device_type_for_weight(tensor) - is_offgraph_kv = _is_offgraph_kv_fqn(fqn) expected_storage_nbytes = int( getattr(properties, "storage_size", None) or 0 ) - if not is_offgraph_kv and storage_nbytes < expected_storage_nbytes: + if ( + not is_offgraph_kv + and not is_compact_clone + and storage_nbytes < expected_storage_nbytes + ): raise RuntimeError( "AOTI cloned storage is smaller than its TensorProperties " f"({storage_nbytes} < {expected_storage_nbytes} bytes)" @@ -427,7 +441,11 @@ def materialize( strides = tuple( int(stride) for stride in getattr(properties, "stride", tensor.stride()) ) - storage_offset = int(getattr(properties, "offset", tensor.storage_offset())) + storage_offset = ( + tensor.storage_offset() + if is_compact_clone + else int(getattr(properties, "offset", tensor.storage_offset())) + ) required_nbytes = _required_view_nbytes( fqn, sizes, strides, storage_offset, tensor.element_size() ) @@ -436,6 +454,11 @@ def materialize( f"AOTI view {fqn!r} requires {required_nbytes} bytes from a " f"{storage_nbytes}-byte cloned storage" ) + if is_compact_clone and required_nbytes != storage_nbytes: + raise RuntimeError( + f"AOTI compact clone of {fqn!r} should span exactly its " + f"{storage_nbytes}-byte storage, but its view needs {required_nbytes}" + ) if is_offgraph_kv: # Preserve the AOTI view contract; the runtime supplies storage. storage_nbytes = max( diff --git a/backends/cuda/tests/test_cuda_partitioner.py b/backends/cuda/tests/test_cuda_partitioner.py index 2f3800743d8..af49ed25e77 100644 --- a/backends/cuda/tests/test_cuda_partitioner.py +++ b/backends/cuda/tests/test_cuda_partitioner.py @@ -39,6 +39,7 @@ from executorch.exir.backend.partitioner import PartitionResult from executorch.exir.delegate import executorch_call_delegate from torch._export.utils import is_buffer, is_lifted_tensor_constant, is_param +from torch._inductor.compile_fx import clone_preserve_strides from torch.export import export from torch.export.pt2_archive._package_weights import TensorProperties, Weights from torch.fx.passes.utils.fuser_utils import validate_partition @@ -232,8 +233,8 @@ def test_different_fqn_views_have_distinct_logical_storage(self) -> None: weights = Weights( { "base": (base, TensorProperties(base)), - # AOTI may return a cloned value tensor; TensorProperties is - # the source of truth for reconstructing the original view. + # The value shares the view's storage, so the offset from + # TensorProperties applies. "view": (base, TensorProperties(view)), } ) @@ -251,6 +252,101 @@ def test_different_fqn_views_have_distinct_logical_storage(self) -> None: for storage in artifact.storages.values(): storage.close() + def test_compact_clone_of_a_view_uses_the_written_storage(self) -> None: + # AOTI clones a buffer that views a larger tensor into compact storage + # but reports TensorProperties of the original view. + view = torch.arange(32, dtype=torch.float32).reshape(8, 4)[2:5] + clone = view.clone() + weights = Weights({"w": (clone, TensorProperties(view))}) + + with tempfile.TemporaryDirectory() as directory: + artifact = self._materialize(weights, directory) + entry = artifact.entries[0] + self.assertEqual(0, entry.storage_offset) + self.assertEqual(48, entry.storage_nbytes) + self.assertEqual((3, 4), entry.sizes) + data = artifact.storages[entry.storage_key].to_bytes() + self.assertEqual(bytes(clone.untyped_storage()), data) + for storage in artifact.storages.values(): + storage.close() + + def test_compact_clone_with_too_little_storage_is_rejected(self) -> None: + view = torch.arange(32, dtype=torch.float32).reshape(8, 4)[2:5] + truncated = view.clone() + truncated.untyped_storage().resize_(8 * truncated.element_size()) + with tempfile.TemporaryDirectory() as directory: + with self.assertRaisesRegex(RuntimeError, "requires 48 bytes"): + self._materialize( + Weights({"w": (truncated, TensorProperties(view))}), directory + ) + + def test_compact_clone_with_bytes_past_its_view_is_rejected(self) -> None: + view = torch.arange(32, dtype=torch.float32).reshape(8, 4)[2:5] + padded = view.clone() + padded.untyped_storage().resize_(16 * padded.element_size()) + with tempfile.TemporaryDirectory() as directory: + with self.assertRaisesRegex(RuntimeError, "span exactly its 64-byte"): + self._materialize( + Weights({"w": (padded, TensorProperties(view))}), directory + ) + + def test_value_with_the_same_shape_but_other_strides_is_not_a_clone(self) -> None: + view = torch.arange(32, dtype=torch.float32).reshape(8, 4)[2:5] + transposed = torch.arange(12, dtype=torch.float32).reshape(4, 3).t() + with tempfile.TemporaryDirectory() as directory: + with self.assertRaisesRegex( + RuntimeError, "smaller than its TensorProperties" + ): + self._materialize( + Weights({"w": (transposed, TensorProperties(view))}), directory + ) + + def test_compact_clone_of_a_strided_view_spans_its_storage(self) -> None: + # A strided view's clone keeps the gaps between its rows, so it spans + # more bytes than its elements take. + view = torch.arange(64, dtype=torch.float32).reshape(8, 8)[:, 1:3] + clone = clone_preserve_strides(view) + weights = Weights({"w": (clone, TensorProperties(view))}) + + with tempfile.TemporaryDirectory() as directory: + artifact = self._materialize(weights, directory) + entry = artifact.entries[0] + self.assertEqual(0, entry.storage_offset) + self.assertEqual(232, entry.storage_nbytes) + self.assertEqual((8, 1), entry.strides) + data = artifact.storages[entry.storage_key].to_bytes() + self.assertEqual(bytes(clone.untyped_storage()), data) + for storage in artifact.storages.values(): + storage.close() + + def test_view_sharing_its_storage_keeps_the_view_offset(self) -> None: + # AOTI does not clone parameters, so a parameter that slices a fused + # tensor arrives in the fused tensor's storage. + fused = torch.arange(36, dtype=torch.float32).reshape(9, 4) + view = fused[3:6] + weights = Weights({"w": (view, TensorProperties(view))}) + + with tempfile.TemporaryDirectory() as directory: + artifact = self._materialize(weights, directory) + entry = artifact.entries[0] + self.assertEqual(12, entry.storage_offset) + self.assertEqual(144, entry.storage_nbytes) + for storage in artifact.storages.values(): + storage.close() + + def test_value_with_other_layout_keeps_the_view_offset(self) -> None: + # A value in a different storage that is not a compact clone of the view + # (other shape) keeps the offset its TensorProperties records. + base = torch.arange(32, dtype=torch.float32).reshape(8, 4) + view = base[:, 1:] + weights = Weights({"w": (base.clone(), TensorProperties(view))}) + + with tempfile.TemporaryDirectory() as directory: + artifact = self._materialize(weights, directory) + self.assertEqual(1, artifact.entries[0].storage_offset) + for storage in artifact.storages.values(): + storage.close() + def test_identical_values_keep_distinct_fqn_keys(self) -> None: first = torch.zeros(4) second = torch.zeros(4)