Skip to content

Fix QuantMultiheadAttention ONNX export crash from stale tracing-state guard - #1617

Merged
Giuseppe5 merged 4 commits into
Xilinx:masterfrom
Sebbruunski:fix/quant-mha-export-rename-guard
Sep 23, 2026
Merged

Giuseppe5 merged 4 commits into
Xilinx:masterfrom
Sebbruunski:fix/quant-mha-export-rename-guard

Conversation

@Sebbruunski

Copy link
Copy Markdown
Contributor

Problem

QuantMultiheadAttention's multi_head_attention() fails to export via
brevitas.export.export_qonnx with:

TypeError: _create_graph_by_tracing(): incompatible function arguments ...
tensor(..., names=('L', 'N', 'E'))

This is a different bug from #1574 / #1565 (PyTorch 2.13's removal of named
tensors) — it reproduces on current PyTorch releases that still support named
tensors (e.g. torch 2.7.1), with brevitas at HEAD.

Root cause

multi_head_attention() guards its named-tensor bookkeeping
(query.rename_('L', 'N', 'E') etc., used so PTQ/GPTQ-style calibration code
can locate the batch dimension by name) with:

if not torch._C._get_tracing_state():
    ...
    query.rename_('L', 'N', 'E')

The guard is meant to skip this during ONNX export tracing. But
export_qonnx's own ONNXBaseManager.export_onnx first runs an eager
shape-caching pass (_cache_inp_out) on the real input tensors, before any
tracing flag is set. At that point torch._C._get_tracing_state() is None,
so the guard's condition is true and the in-place rename executes anyway —
permanently renaming the caller's real query/key/value tensor objects.
A few lines later, export_onnx calls torch.onnx.export(module, args, ...)
with those same (now-named) tensor objects, and the legacy tracer can't bind
named tensors, producing the crash above.

The existing cleanup (for t in [q, k, v]: t.rename_(None)) only clears
names off the newly-created q, k, v projections — never off the original
query/key/value objects, so they stay named for the rest of their
lifetime outside this function.

Fix

Use the non-mutating Tensor.rename(...) instead of the in-place
.rename_(...), rebinding the local variable. This function's own downstream
logic (self.in_proj(query) / self.q_proj(query) etc.) still sees a named
tensor, but no tensor object held outside this function is ever mutated —
so export_qonnx's eager pre-pass no longer has any observable side effect
on the caller's tensors.

Also gives query/key/value in the separate-projection
(use_separate_proj_weight=True) branch their own individual isinstance
checks, matching the pattern already used in the packed/self-attention
branch (previously a shared loop over [query, key, value]).

Testing

Added tools/verify_quant_mha_export_fix.py: builds a small
QuantMultiheadAttention (embed_dim=8, num_heads=2), exercises both the
packed self-attention path and the separate-projection
(packed_in_proj=False) path, exports each via export_qonnx, executes the
result, and checks it against an eager forward pass.

self_attention: export succeeded, max_abs_error=1.490e-08, output_shape=(3, 2, 8)
cross_attention: export succeeded, max_abs_error=5.960e-08, output_shape=(3, 2, 8)
QuantMultiheadAttention export regression: PASS

Also checked graph/gpxq.py and graph/equalize.py, which read
inp.names.index('N') for batch-dimension detection — they're unaffected,
since they depend on the existing (unchanged) name-clearing behavior on
q/k/v, not on query/key/value's own names.

@nickfraser
nickfraser self-requested a review September 21, 2026 10:51
@nickfraser

Copy link
Copy Markdown
Collaborator

Your fix is applied to an old version - can you confirm whether the fix is still required on master? The .rename logic changes in recent versions, since this API is no longer available on PyTorch>=2.13.

@nickfraser nickfraser added the bug Something isn't working label Sep 21, 2026
Sebbruunski and others added 3 commits September 21, 2026 15:56
…e guard

multi_head_attention()'s `if not torch._C._get_tracing_state():` guard was
meant to skip named-tensor bookkeeping during export, but Brevitas's own
export_qonnx runs an eager shape-caching pre-pass before any tracing flag is
set, so the guard fired anyway and the in-place `.rename_('L','N','E')` calls
permanently renamed the caller's real query/key/value tensor objects. A
subsequent torch.onnx.export on those same objects then crashed with
TypeError: _create_graph_by_tracing(): incompatible function arguments
(named tensors aren't supported by the legacy tracer).

Fix: use the non-mutating .rename(...) and rebind the local variable instead
of .rename_(...), so this function's own downstream logic still sees named
tensors but no tensor object held outside the function is ever mutated. Also
gives each of query/key/value in the cross-attention/separate-projection
branch its own isinstance check (previously shared one loop), matching the
self-attention branch's pattern.

Verified via new tools/verify_quant_mha_export_fix.py: both the
self-attention and cross-attention/separate-projection code paths now export
via export_qonnx without raising, with output matching an eager forward pass
to ~1e-8 max abs error. Confirmed gpxq.py/equalize.py's `inp.names.index('N')`
batch-dimension lookups are unaffected (they only depend on the existing,
unchanged q/k/v name-cleanup, not on query/key/value's own names).
@Giuseppe5

Copy link
Copy Markdown
Collaborator

Hey, I am rebasing/pushing some commits to this PR since we're planning a release soon and we'd like to fix this bug, although not exactly the way you implemented this.

Many thanks.

@Giuseppe5
Giuseppe5 force-pushed the fix/quant-mha-export-rename-guard branch from 5526cb1 to 2384eed Compare September 21, 2026 15:09
del m.cache_inference_quant_act_backup


def _cache_inp_out(module, *args, **kwargs):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

This is a side effect so we can use the function without having to use ONNXBaseManager._cache_inp_out, with the idea that maybe this lives in some utils once we decouple onnx from the rest of the codebase.

@Giuseppe5
Giuseppe5 self-requested a review September 22, 2026 08:24
@Giuseppe5
Giuseppe5 requested review from Giuseppe5 and removed request for Giuseppe5 September 22, 2026 09:48
@Giuseppe5
Giuseppe5 merged commit 3c3e27e into Xilinx:master Sep 23, 2026
478 of 483 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

bug Something isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants