-
Notifications
You must be signed in to change notification settings - Fork 1.2k
Accept AOTI's compact clones of view buffers in the CUDA weight collector #23355
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 ( | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 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)" | ||
|
|
@@ -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() | ||
| ) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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()?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 |
||
|
|
@@ -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( | ||
|
|
||
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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.There was a problem hiding this comment.
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_kvwe will remove the buffer from ptd to make runtime control it.