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
29 changes: 26 additions & 3 deletions backends/cuda/cuda_weight_collector.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why would offgraph kv be in the weights fqn?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

That check isn't new in this change. It came with the off-graph KV lowering (#23292). That pass puts its cache buffers into the program as constants named __et_offgraph_kv_*, so they reach the weight collector together with the real weights. The collector skips them because their data is managed elsewhere. This change only moves that existing line up, so the new check can use it too.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It is for removing kvcache buffer from weight: once we detected is_offgraph_kv we will remove the buffer from ptd to make runtime control it.

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 (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

so basically iiuc if you have a fused weight aot like qkv aoti wants to split them into q k and v and there was some metadata getting trashed in this process that you fix? @shoumikhin

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Close, but it's about views rather than splitting. Take a buffer that is a view of a bigger tensor, like big[2:5]. AOTInductor copies it into a compact clone: 48 bytes, starting at offset 0. But the TensorProperties it reports still describe the original: 128 bytes at offset 8. The collector compared the two and rejected the export, even though the clone holds every byte the view needs. A fused QKV weight split into views is one way to get such a buffer, so yes, that's a real-world case.

The fix uses the clone's own size and offset when it's a compact clone with the same shape and strides. The bounds check still runs, so a truncated clone is still rejected.

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)"
Expand All @@ -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()
)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

can we add a check here that the required_bytes is the same as the tensor.nbytes()?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added, slightly adjusted: a compact clone must now span exactly its storage. Checking against tensor.nbytes() would reject strided views. For example, a column slice like big[:, 1:3] of an 8x8 tensor has 64 bytes of elements, but its clone spans 232 bytes, because the clone keeps the gaps between rows. A new test covers a clone with extra bytes past its view, which is now rejected, and a strided view, which is still accepted.

Expand All @@ -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(
Expand Down
100 changes: 98 additions & 2 deletions backends/cuda/tests/test_cuda_partitioner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)),
}
)
Expand All @@ -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)
Expand Down
Loading