Skip to content

[PyTorch][torch.compile] Support for UnfusedDotProductAttention#3201

Open
pggPL wants to merge 15 commits into
NVIDIA:mainfrom
pggPL:unfused_dpa_torch_compile
Open

[PyTorch][torch.compile] Support for UnfusedDotProductAttention#3201
pggPL wants to merge 15 commits into
NVIDIA:mainfrom
pggPL:unfused_dpa_torch_compile

Conversation

@pggPL

@pggPL pggPL commented Jul 10, 2026

Copy link
Copy Markdown
Collaborator

Description

Make the UnfusedDotProductAttention backend traceable by torch.compile(fullgraph=True, mode="reduce-overhead"), so the forward and backward can be captured into CUDA graphs without graph breaks.

Scope:

  1. bf16/fp16 path (compile-supported): register the TE softmax kernels and THD<->BSHD conversion helpers as torch.library.custom_ops with fake impls and autograd bindings; remove an unbacked-SymInt .item() from the hot path of ConvertBSHDtoTHD.
  2. FP8 is explicitly NOT supported under torch.compile: with fp8=True (emulation) and/or fp8_output=True (Float8Tensor output, a tensor subclass that cannot cross a graph boundary) the backend runs as an eager island — the forward dispatches to a torch._dynamo.disable'd wrapper, the same mechanism DotProductAttention and FusedAttention use module-wide. FP8 attention always involves delayed scaling regardless of the recipe: S and dP are produced inside the kernel, so their amax cannot be known before quantization and they use delayed-scaling quantizers even under Float8CurrentScaling (see DPA.init_fp8_metadata) — and delayed scaling (Float8Quantizer, tensor scale/amax state) is not supported under torch.compile.

Type of change

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

  • softmax.py / softmax.cpp: scaled_*_softmax_{forward,backward} as custom ops; C++ backward kernels allocate a fresh output buffer instead of writing in-place into output_grad (custom ops and cudagraph trees forbid input aliasing).
  • utils.py: ConvertTHDtoBSHD / ConvertBSHDtoTHD as custom ops; num_tokens passed by the caller instead of cu_seqlens[-1].item().
  • backends.py: UnfusedDotProductAttention.forward dispatches FP8 calls to an eager (dynamo-disabled) wrapper; the non-FP8 path is traced with no graph breaks.
  • tests/pytorch/test_torch_compile.py: test_unfused_dpa_torch_compile (5 qkv layouts, fullgraph + reduce-overhead, fwd+bwd captured into CUDA graphs and replayed).

Checklist:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes

🤖 Generated with Claude Code

pggPL added 11 commits June 25, 2026 14:39
…le + CUDA graphs

Refactor TE custom kernels used by the unfused attention path so that
`torch.compile(fullgraph=True, mode="reduce-overhead")` can trace the
forward and backward and capture them into CUDA graphs without graph
breaks.

- softmax.py / softmax.cpp: register all `scaled_*_softmax_{forward,backward}`
  kernels as `torch.library.custom_op`s with fake impls and an autograd
  binding that mirrors the previous `torch.autograd.Function`s. The C++
  backward kernels now allocate a fresh output buffer instead of writing
  in-place into `output_grad`, so the ops no longer alias their inputs
  (required by `torch.library.custom_op` and inductor cudagraph trees).
- utils.py: convert `ConvertTHDtoBSHD` / `ConvertBSHDtoTHD` to
  `torch.library.custom_op`s, with thin wrapper classes that keep the
  existing `.apply(...)` callsite syntax. Drop the
  `int(cu_seqlens[-1].item())` from the hot path of `ConvertBSHDtoTHD.apply`
  -- under `torch.compile` it created an unbacked SymInt, which made the
  Inductor partitioner emit `None` placeholders for output buffers and
  caused `cudagraph_trees` to assert. `num_tokens` is now passed in by
  the caller as a regular (Sym)Int.
- backends.py: in the THD branch of unfused DPA, capture
  `total_tokens_q = query_layer.shape[0]` before overwriting
  `query_layer` with the BSHD form, and thread it back into
  `ConvertBSHDtoTHD.apply` at the end of the forward.
- test_torch_compile.py: add `test_unfused_dpa_torch_compile`,
  parametrized over qkv layouts (`bshd_bshd_bshd`, `sbhd_sbhd_sbhd`,
  `thd_thd_thd`, `bs3hd`, `sbh3d`), that compiles
  `UnfusedDotProductAttention.forward` directly with `fullgraph=True,
  mode="reduce-overhead"` and runs forward+backward several times so the
  CUDA graphs are recorded and replayed.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Made-with: Cursor
…to unfused_attention_torch_compile

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>

# Conflicts:
#	tests/pytorch/test_torch_compile.py
…DotProductAttention

Make the FP8-emulation path (NVTE_UnfusedDPA_Emulate_FP8=1) of
UnfusedDotProductAttention traceable by torch.compile(fullgraph=True).

- backends.py: register the quantize+dequantize roundtrips used by
  FP8EmulationFunc as torch.library custom ops
  (te_fp8_emu::roundtrip_<QuantizerClass> and
  te_fp8_emu::roundtrip_qkv_<QuantizerClass>) taking the quantizer as a
  value-opaque argument, with fake impls for tracing. Ops are registered
  only for the value-opaque quantizer classes
  (Float8CurrentScalingQuantizer, MXFP8Quantizer); Float8Quantizer
  (delayed scaling) carries scale/amax tensor state, is not
  value-opaque, and deliberately keeps the plain eager path -- FP8
  emulation with delayed scaling is not supported under torch.compile.
- backends.py: dispatch helpers `_fp8_emu_roundtrip{,_qkv}` key on
  `type(quantizer).__qualname__` so they stay traceable for opaque
  quantizer arguments; FP8EmulationFunc forward/backward now call them
  (onnx_forward unchanged).
- backends.py: the joint q/k/v roundtrip clones any output whose
  storage is shared with an input or another output, checking storage
  identity directly -- the dequantized q/k/v can be views into one
  combined buffer, and view metadata (`_base`) is not populated under
  the torch-dispatch mode AOTAutograd runs custom ops with, so a
  `_base`-guarded clone triggered the custom-op aliasing deprecation
  warning under torch.compile.
- UnfusedDotProductAttention.forward: only query
  FP8GlobalStateManager.get_fp8_recipe() when
  fp8_meta["local_recipes"] is absent.
- test_torch_compile.py: add test_unfused_dpa_fp8_emulation_torch_compile
  (current scaling + mxfp8, sbhd/bshd layouts; compiled fullgraph
  forward+backward must match eager) and
  test_unfused_dpa_fp8_emulation_delayed_scaling_eager guarding the
  eager delayed-scaling path after the refactor.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…output=True

With fp8_output=True the backend returns a Float8Tensor -- a tensor
subclass that cannot cross a torch.compile graph boundary -- so the
forward dispatches to a torch._dynamo.disable'd wrapper, the same
mechanism DotProductAttention and FusedAttention use module-wide.
With fp8_output=False the dispatcher is resolved at trace time and
adds no graph break.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…cudagraphs)

Parametrize test_unfused_dpa_fp8_emulation_torch_compile over compile
mode (default, reduce-overhead), run 3 iterations so the CUDA graphs
are recorded and replayed. The te_fp8_emu roundtrip ops for current
scaling are pure (no mutated args), so inductor cudagraphs capture
them; verified no cudagraph skips with TORCH_LOGS=cudagraphs.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…n; run FP8 as an eager island

FP8 in the unfused backend (emulation and Float8Tensor output) is not
supported under torch.compile: the forward dispatcher routes fp8=True
and/or fp8_output=True to a torch._dynamo.disable'd wrapper, same as
DotProductAttention does module-wide. Remove the FP8-emulation compile
tests. The te_fp8_emu::* custom ops taking value-opaque quantizers stay
as the eager implementation of FP8EmulationFunc.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…Func

The ops existed solely to make the FP8-emulation path traceable by
torch.compile; since FP8 in the unfused backend now always runs as an
eager island, they are dead machinery (plus import-time registration
and output clones the plain eager path never needed).

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
The tex softmax kernels take 'float scale_factor' directly. The 0-D
tensor wrapping was a leftover of the old autograd.Function idiom,
where the float had to be a tensor only to fit save_for_backward;
the custom ops keep the scale on ctx as a plain attribute.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…the callsite)

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…cate dict, silence W0613

- run black over the four changed files (earlier commits skipped pre-commit)
- drop unused 'import os' in test_torch_compile.py
- drop duplicated module-level _default_causal_mask dict in softmax.py
- del unused 'output' arg in the conversion setup_context helpers

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
@pggPL
pggPL requested review from cyanguwa and ksivaman as code owners July 10, 2026 13:21
@pggPL pggPL changed the title [PyTorch] torch.compile support for UnfusedDotProductAttention [PyTorch][torch.compile] Support for UnfusedDotProductAttention Jul 10, 2026
@greptile-apps

greptile-apps Bot commented Jul 10, 2026

Copy link
Copy Markdown
Contributor

Greptile Summary

This PR makes UnfusedDotProductAttention traceable by torch.compile(fullgraph=True, mode="reduce-overhead") for bf16/fp16 training paths. FP8 paths are explicitly excluded via a @no_torch_dynamo()-wrapped eager island. Four softmax kernel families and the THD↔BSHD layout-conversion helpers are refactored from torch.autograd.Function subclasses into torch.library.custom_ops with fake implementations and register_autograd bindings. C++ backward kernels now allocate fresh output buffers to satisfy the no-aliasing contract, and the .item() call that produced an unbacked SymInt from cu_seqlens[-1] is replaced by a shape read captured before any layout transformation.

  • softmax.py / softmax.cpp: Each of the four fused softmax kernels is wrapped as a pair of custom_ops (forward + backward), with autograd linked via register_autograd; C++ backward functions allocate fresh output tensors instead of writing in-place into the gradient input.
  • utils.py / backends.py: ConvertTHDtoBSHD and ConvertBSHDtoTHD become thin .apply() wrappers around new te_attention::convert_* custom ops; total_tokens_q is hoisted from the cu_seqlens[-1].item() call to a shape read before any layout conversion, fixing a NameError on inference THD→BSHD paths and avoiding an unbacked-SymInt GPU sync.
  • test_torch_compile.py: A new test_unfused_dpa_torch_compile parametrized test covers 5 QKV layouts with fullgraph=True + reduce-overhead, validating forward and backward through CUDA-graph replay.

Confidence Score: 5/5

Safe to merge; the non-FP8 training paths are well-covered by the new fullgraph compile test, and FP8 is correctly fenced behind a dynamo-disabled eager island.

The custom-op refactor is mechanically sound: fake impls correctly reflect output shapes, autograd hooks use register_autograd with setup_context (the compile-compatible pattern), and the C++ backward allocations remove input aliasing. The total_tokens_q hoisting is logically correct — it reads query_layer.shape[0] before any layout transformation, so the value is right for both the thd training path and the thd_2bshd inference path. No correctness or runtime-error issues were found in the new code.

No files require special attention; the only note is a stale docstring in forward_fused_softmax that still names the removed autograd.Function classes.

Important Files Changed

Filename Overview
transformer_engine/pytorch/attention/dot_product_attention/softmax.py Replaces four autograd.Function softmax classes with torch.library custom_op pairs (fwd + bwd), each with a fake impl and register_autograd binding; adds float() coercion on the scale parameter; docstring in forward_fused_softmax still names the removed classes.
transformer_engine/pytorch/csrc/extensions/softmax.cpp Each of the four backward kernel wrappers now allocates a fresh output buffer (torch::empty) instead of writing in-place into output_grads, satisfying the custom_op no-aliasing contract.
transformer_engine/pytorch/attention/dot_product_attention/utils.py ConvertTHDtoBSHD and ConvertBSHDtoTHD are refactored from autograd.Function subclasses to thin .apply() wrappers around te_attention custom ops; the custom ops carry their own fake impls and register_autograd hooks.
transformer_engine/pytorch/attention/dot_product_attention/backends.py Splits forward into a public dispatch shim and _forward; FP8 paths routed to @no_torch_dynamo _forward_eager; total_tokens_q is captured from query_layer.shape[0] before any layout transformation, avoiding the unbacked-SymInt .item() call and fixing NameError on thd_2bshd inference paths.
tests/pytorch/test_torch_compile.py Adds test_unfused_dpa_torch_compile parametrized over 5 QKV layouts with fullgraph=True + reduce-overhead; tests fwd+bwd CUDA-graph capture and replay; dtype is hardcoded to bfloat16 (fp16 not exercised).

Sequence Diagram

sequenceDiagram
    participant Caller
    participant forward as UnfusedDPA.forward
    participant _fwd as UnfusedDPA._forward
    participant _fwd_eager as UnfusedDPA._forward_eager (@no_torch_dynamo)
    participant thd2bshd as te_attention::convert_thd_to_bshd (custom_op)
    participant softmax as te_softmax::scaled_*_softmax_fwd (custom_op)
    participant bshd2thd as te_attention::convert_bshd_to_thd (custom_op)

    Caller->>forward: call(fp8, fp8_output, args...)
    alt "fp8=True or fp8_output=True"
        forward->>_fwd_eager: eager island (graph break)
        _fwd_eager->>_fwd: delegate
    else non-FP8 compile-compatible path
        forward->>_fwd: direct call (fully traced)
    end
    alt "qkv_format == thd (training)"
        _fwd->>thd2bshd: ConvertTHDtoBSHD.apply(q/k/v, cu_seqlens, max_seqlen)
    end
    _fwd->>softmax: FusedScaleMaskSoftmax
    alt "q_format == thd"
        _fwd->>bshd2thd: ConvertBSHDtoTHD.apply(ctx, cu_seqlens_q, total_tokens_q)
        Note over _fwd,bshd2thd: total_tokens_q from query_layer.shape[0] before layout conversion
    end
    _fwd-->>Caller: context_layer
Loading

Reviews (3): Last reviewed commit: "[PyTorch] Fix UnboundLocalError on thd i..." | Re-trigger Greptile

Comment thread transformer_engine/pytorch/attention/dot_product_attention/utils.py

@cyanguwa cyanguwa left a comment

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.

LGTM but please take a look at Greptile's few comments. Thanks.

pggPL added 2 commits July 23, 2026 13:42
… docstrings

Greptile P2: the ConvertTHDtoBSHD/ConvertBSHDtoTHD class docstrings said
callsites keep the .apply(...) syntax without reflecting the actual
argument list. Spell out the apply() signatures so the required args
(incl. num_tokens / max_seqlen) are explicit.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
@pggPL

pggPL commented Jul 23, 2026

Copy link
Copy Markdown
Collaborator Author

/te-ci pytorch

…Attention

The thd output conversion (q_format=='thd') passes total_tokens_q to the
new ConvertBSHDtoTHD custom op, but total_tokens_q was only assigned on the
training 'thd' input branch, not the inference 'thd_2bshd' branch, so thd
KV-cache inference raised UnboundLocalError.

Capture total_tokens_q once right after q_format is known, before any
layout conversion: for both 'thd' and 'thd_2bshd' the query enters in thd
layout so query_layer.shape[0] is the total query token count (a backed
SymInt, unlike cu_seqlens_q[-1].item() which would sync the GPU and break
torch.compile + cudagraphs).

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants