Fix QuantMultiheadAttention ONNX export crash from stale tracing-state guard - #1617
Merged
Giuseppe5 merged 4 commits intoSep 23, 2026
Merged
Conversation
nickfraser
self-requested a review
September 21, 2026 10:51
Collaborator
|
Your fix is applied to an old version - can you confirm whether the fix is still required on |
…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).
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
force-pushed
the
fix/quant-mha-export-rename-guard
branch
from
September 21, 2026 15:09
5526cb1 to
2384eed
Compare
Giuseppe5
reviewed
Sep 21, 2026
| del m.cache_inference_quant_act_backup | ||
|
|
||
|
|
||
| def _cache_inp_out(module, *args, **kwargs): |
Collaborator
There was a problem hiding this comment.
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
self-requested a review
September 22, 2026 08:24
Giuseppe5
requested review from
Giuseppe5
and removed request for
Giuseppe5
September 22, 2026 09:48
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
QuantMultiheadAttention'smulti_head_attention()fails to export viabrevitas.export.export_qonnxwith: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 codecan locate the batch dimension by name) with:
The guard is meant to skip this during ONNX export tracing. But
export_qonnx's ownONNXBaseManager.export_onnxfirst runs an eagershape-caching pass (
_cache_inp_out) on the real input tensors, before anytracing flag is set. At that point
torch._C._get_tracing_state()isNone,so the guard's condition is true and the in-place rename executes anyway —
permanently renaming the caller's real
query/key/valuetensor objects.A few lines later,
export_onnxcallstorch.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 clearsnames off the newly-created
q, k, vprojections — never off the originalquery/key/valueobjects, so they stay named for the rest of theirlifetime outside this function.
Fix
Use the non-mutating
Tensor.rename(...)instead of the in-place.rename_(...), rebinding the local variable. This function's own downstreamlogic (
self.in_proj(query)/self.q_proj(query)etc.) still sees a namedtensor, but no tensor object held outside this function is ever mutated —
so
export_qonnx's eager pre-pass no longer has any observable side effecton the caller's tensors.
Also gives
query/key/valuein the separate-projection(
use_separate_proj_weight=True) branch their own individualisinstancechecks, 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 smallQuantMultiheadAttention(embed_dim=8, num_heads=2), exercises both thepacked self-attention path and the separate-projection
(
packed_in_proj=False) path, exports each viaexport_qonnx, executes theresult, and checks it against an eager forward pass.
Also checked
graph/gpxq.pyandgraph/equalize.py, which readinp.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 onquery/key/value's own names.