Uh oh!
There was an error while loading. Please reload this page.
[PyTorch][torch.compile] DotProductAttention: declarative packed QKV/KV inputs - #3200
Conversation
Fused QKV projections naturally produce one packed buffer, but DotProductAttention forces callers to slice it into q/k/v views that TE then reverse-engineers with pointer-based layout detection (get_qkv_layout inspects data_ptr/storage_offset on every forward, which graph-breaks under torch.compile and adds CPU overhead). Let callers declare the packing instead (JAX-style): * DotProductAttention.forward gains optional qkv_layer (fully packed QKV: [b,s,3,h,d]/[s,b,3,h,d]/[b,s,h,3,d]/[s,b,h,3,d] dense, [t,3,h,d]/ [t,h,3,d] thd), kv_layer (packed KV used with query_layer), and qkv_interleave_dim (-3 or -2; explicit knob rather than shape inference since h==3 or hg==2 would be ambiguous). * Q/K/V are derived as zero-copy select() views and the exact layout enum (bs3hd, bsh3d, sb3hd, bshd_bs2hd, t3hd, ...) is constructed declaratively -- it is truthful by construction, so get_qkv_layout is never called on this path, including for thd and FP8 DPA. * combine_and_quantize no longer re-combines what is already combined: a new optional combined= argument carries the caller's original packed buffer, which is quantized directly instead of rebuilding the packed buffer from q/k/v views via combine_tensors (a raw set_ with a silent adjacency/interleave assumption). The packed original is threaded from DPA.forward through FusedAttention to FusedAttnFunc.forward; all legacy call sites are untouched (combined=None preserves exact behavior), and backward combine calls are unchanged (gradients have no pre-packed original). Tests: dense fwd+grad bit-exactness vs separate contiguous q/k/v for bs3hd/bsh3d/sb3hd/kv-packed/GQA (fused + flash), validation errors, torch.compile (no data_ptr/UntypedStorage graph breaks), FP8 combined-vs-views bit equivalence, and detection-free declared t3hd. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…claratively Adopt the new DotProductAttention packed API inside MultiheadAttention: * self-attention (np == ng): the fused QKV projection output, already viewed as [.., h, 3, d] (qkv_weight_interleaved) or [.., 3, h, d], is handed to DPA directly as qkv_layer with the matching qkv_interleave_dim (-2 / -3) -- no SplitAlongDim slicing in MHA and no pointer-based layout detection in DPA. * cross-attention: the packed KV projection output is exposed as [.., hg, 2, d] / [.., 2, hg, d] and passed as kv_layer. * The pass-through only engages when no per-tensor operation needs the individual q/k/v slices: it is skipped for RoPE, QK normalization, KV caching (inference_params), CPU offloading, GQA (np != ng, not a uniform 3-interleave) and quantized (FP8) projection outputs; those keep the legacy sliced-views path unchanged. Tests: MHA self (interleaved + non-interleaved) and cross (both interleaves) are bit-exact vs the same MHA with packed inputs converted back to separate contiguous q/k/v (output, input grad, weight grads); spy asserts the packed argument and interleave dim actually reach DPA; GQA and RoPE fall back to the views path. TransformerLayer regression suite unchanged; test_kv_cache failures on this device are pre-existing on origin/main (verified). Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…ack_packed_qkv Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
… arguments Rename the single layout-dependent packed_qkv plumbing argument (which held the full QKV buffer for *3* layouts but the KV buffer for *_2* layouts) into explicit packed_qkv/packed_kv, mirroring the public qkv_layer/kv_layer API. combine_and_quantize's combined= is split into combined_qkv=/combined_kv= accordingly, and _unpack_packed_qkv no longer returns the packed tensor since the callers already hold qkv_layer/kv_layer. Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Fold the tests from the new test_dpa_packed_inputs.py file into the existing attention test suite as a dedicated section, reusing its imports. No new test file; test logic unchanged apart from renaming the module-level constants (_B/_S/_H/_D/_DTYPE -> _PACKED_*) to avoid collisions. Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
get_qkv_layout now emits a DeprecationWarning when it recognizes q/k/v as views of a packed buffer purely from data pointers/strides/offsets (detected *3*/*_2* layouts), pointing callers at the declarative qkv_layer/kv_layer API. Separate q/k/v tensors (hd_hd_hd layouts) never warn since there is nothing to declare. Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Drop the API-validation, torch.compile graph-break, combine_and_quantize equivalence and thd no-detection spy tests; keep the dense/flash bit-exact equivalence tests and the MHA packed pass-through/fallback tests. Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…arative param
Parametrize test_dpa_qkv_layout and test_dpa_qkv_layout_thd with
declarative={views,declarative}: the declarative mode passes the packed buffer
to DotProductAttention via qkv_layer/kv_layer (declared layout, gradients read
off the packed buffer) instead of slicing it into q/k/v views for pointer-based
detection. This reuses the whole existing config matrix (masks, bias, SWA,
cross-attention, thd, all backends) for the declarative API, replacing the
dedicated packed-input test section.
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>…matrix Revert the declarative parametrization of test_dpa_qkv_layout(_thd) (which doubled their whole config x layout product) and instead add test_dpa_qkv_layout(_thd)_declarative covering all packed layouts on a trimmed config dimension: one self-attention and one cross-attention config (kv_layer path) for dense, one config for thd. Past the input handling the backend code is identical to the views mode, so the full config matrix added no coverage. Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Replace the .grad-holder objects substituted for q/k/v in declarative packed mode with q_grad/k_grad/v_grad variables computed right after backward, used uniformly by all return paths. Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…to dpa_packed_qkv_api Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com> # Conflicts: # tests/pytorch/attention/test_attention.py
for more information, see https://pre-commit.ci
Greptile SummaryThis PR adds declarative packed QKV/KV inputs to
Confidence Score: 5/5Safe to merge. All new code paths are additive: existing callers are unchanged, the packed path is guarded by explicit non-None checks, and the FP8 optimization in The layout-string construction ( The two new Important Files Changed
Flowchart%%{init: {'theme': 'neutral'}}%%
flowchart TD
A["DotProductAttention.forward\n(query_layer, key_layer, value_layer,\nqkv_layer?, kv_layer?, qkv_interleave_dim)"]
A --> B["_unpack_packed_qkv()"]
B --> C{"qkv_layer or kv_layer?"}
C -- "neither (legacy)" --> D["return q, k, v unchanged\ndeclared_layout = None"]
C -- "qkv_layer provided" --> E["Validate: rank, shape[interleave_dim]==3,\nstride(-1)==1, q/k/v must be None"]
C -- "kv_layer provided" --> F["Validate: rank, shape[interleave_dim]==2,\nquery_layer required, k/v must be None"]
E --> G["q=select(dim,0), k=select(dim,1), v=select(dim,2)\nlayout = e.g. 'bs3hd'"]
F --> H["k=select(dim,0), v=select(dim,1)\nlayout = e.g. 'bshd_bs2hd'"]
D --> I{"Layout branch"}
G --> I
H --> I
I -- "declared_layout not None" --> J["Skip get_qkv_layout\nqkv_layout = declared_layout"]
I -- "FP8 tensors" --> K["get_qkv_layout on _data"]
I -- "normal" --> L["get_qkv_layout on q/k/v\n(may emit DeprecationWarning)"]
J --> M["Select backend"]
K --> M
L --> M
M -- "FusedAttention" --> N["FusedAttnFunc.forward\n(packed_qkv=qkv_layer, packed_kv=kv_layer)"]
M -- "FlashAttention" --> O["flash_attention.forward"]
M -- "Unfused" --> P["unfused_attention.forward"]
N --> Q{"FP8 quantization?"}
Q -- "Yes + packed buffer" --> R["combine_and_quantize(combined_qkv/combined_kv)\n→ quantize buffer directly"]
Q -- "Yes, rebuild" --> S["combine_and_quantize\n→ combine_tensors then quantize"]
Q -- "No" --> T["standard forward"]
Reviews (3): Last reviewed commit: "[PyTorch] Address review nits: clarify D..." | Re-trigger Greptile |
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
…plicit k_norm - get_qkv_layout deprecation warning: add stacklevel=2 and skip it while CPU offloading is enabled (offloading forces MultiheadAttention onto the sliced-views fallback, so the caller has no migration option there). - MultiheadAttention: gate the packed pass-through on k_norm explicitly instead of relying on q_norm/k_norm being created together. Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
- DotProductAttention.forward: mark query/key/value_layer as Optional and note they are required only when no packed input (qkv_layer/kv_layer) is given (Charlene). - combine_and_quantize: describe combined_qkv/combined_kv in terms of qkv_group=1 / qkv_group=2 layouts instead of '3'/'2' layouts (Charlene). Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
pggPL
commented
Jul 23, 2026
/te-ci pytorch |
Uh oh!
There was an error while loading. Please reload this page.
Description
A fused QKV projection produces one packed buffer, but callers must slice it into
query_layer/key_layer/value_layerviews, whichDotProductAttentionthen reverse-engineers via pointer-based layout detection (get_qkv_layoutinspecting data pointers, strides and offsets on every forward). That detection graph-breaks undertorch.compileand re-derives information the caller already knows.This PR lets callers declare the packing instead. The declared layout is truthful by construction, so no detection runs on this path — including for thd and FP8 DPA.
DotProductAttention.forwardgains three optional arguments (existing call sites unchanged):qkv_layer— fully packed QKV, e.g.[b,s,3,h,d]or[s,b,h,3,d](thd:[t,3,h,d]).kv_layer— packed KV (e.g.[b,s,2,hg,d]), used withquery_layer; supports GQA.qkv_interleave_dim(default-3) — where the 3/2 interleave sits (-3or-2); explicit since shapes can be ambiguous.Handling lives in a single helper (
_unpack_packed_qkv): q/k/v come out as zero-copyselect()views, the exact layout string (bs3hd,sbhd_sb2hd,th3d, ...) is built from the declaration, and inputs are strictly validated (mutual exclusion, rank/size at the interleave dim,stride(-1) == 1so the declared layout cannot lie about memory, noinference_params).Type of change
Changes
DotProductAttention.forward: new optionalqkv_layer/kv_layer/qkv_interleave_dimarguments; packed inputs are resolved by_unpack_packed_qkvinto zero-copy q/k/v views plus a declared layout, skippingget_qkv_layoutentirely.get_qkv_layoutemits aDeprecationWarningwhen it detects a packed layout purely from pointers, pointing callers at the new arguments. Separate q/k/v never warn.packed_qkv/packed_kvandcombine_and_quantize(combined_qkv=..., combined_kv=...)quantizes it directly instead of rebuilding it from views viacombine_tensors(rawset_with a silent adjacency assumption). Backward combine paths untouched.MultiheadAttentionpasses its fused projection output straight to DPA (self-attention:qkv_layer; cross-attention:kv_layer). The legacy sliced-views path is kept for RoPE, QK norm, KV caching, CPU offloading, GQA and FP8 projection outputs._run_dot_product_attentiongains adeclarative_packedmode (packed buffer passed viaqkv_layer/kv_layer, input grads read off the packed buffer);test_dpa_qkv_layout_declarativecovers all 8 packed dense layouts on self- and cross-attention configs,test_dpa_qkv_layout_thd_declarativethe 4 packed thd layouts (Hopper+). Verified on sm89 (bf16, fused+flash+unfused) with the existingtest_dpa_qkv_layoutmatrix unchanged and green; lint clean.Checklist: