[All] Refactor fused attention APIs with cuDNN-frontend support checks and opaque config/params handles#2964
[All] Refactor fused attention APIs with cuDNN-frontend support checks and opaque config/params handles#2964cyanguwa wants to merge 70 commits into
Conversation
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
for more information, see https://pre-commit.ci
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Greptile SummaryThe PR replaces hand-maintained fused-attention selection with cuDNN-frontend graph support checks and introduces opaque v2 configuration and parameter handles.
Confidence Score: 3/5The PR does not yet appear safe to merge because several previously reported compatibility and probe-versus-execution mismatches remain unverified at the current head. The device-aware cache fix and null-pointer initialization are established, but the available evidence does not verify that the deprecated query supplies batch, training, output-format, and gradient-layout fields consistently or that JAX FP8 probing reflects the active recipe; those unresolved paths can still reject supported configurations or accept configurations that fail during execution. Files Needing Attention: transformer_engine/common/fused_attn/fused_attn.cpp, transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.cu, transformer_engine/common/include/transformer_engine/fused_attn.h, transformer_engine/jax/cpp_extensions/attention.py Important Files Changed
Sequence DiagramsequenceDiagram
participant Frontend as PyTorch/JAX
participant Config as Opaque config/params
participant Probe as Backend support probe
participant Cache as Process-wide graph cache
participant cuDNN as cuDNN frontend
Frontend->>Config: Populate attention attributes
Config->>Probe: Derive normalized configuration
Probe->>Cache: Lookup device-keyed graph
alt cache miss
Cache->>cuDNN: Build and validate graph
cuDNN-->>Cache: Cache supported graph
end
Probe-->>Frontend: Backend and diagnostic
Frontend->>Cache: Execute through v2 API
Reviews (30): Last reviewed commit: "fix merge with torch.compile PRs" | Re-trigger Greptile |
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
for more information, see https://pre-commit.ci
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
for more information, see https://pre-commit.ci
|
/te-ci L1 |
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
for more information, see https://pre-commit.ci
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
for more information, see https://pre-commit.ci
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
…, document diagnostic semantics, add reject-path test for pre-scale bias Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
for more information, see https://pre-commit.ci
|
/te-ci L1 |
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
…e defaults Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
| None, # no window | ||
| ): | ||
| pytest.skip("No FusedAttn backend found") | ||
|
|
There was a problem hiding this comment.
Note: The explicit is_fused_attn_kernel_available() pre-check is no longer needed here:
FusedAttnRunner._check_configs() now runs the cuDNN-frontend support probe during
setup and skips with a diagnostic message when no fused-attn backend is available.
|
/te-ci L1 |
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
for more information, see https://pre-commit.ci
Signed-off-by: Charlene Yang <8636796+cyanguwa@users.noreply.github.com>
Description
TE currently hand-maintains the fused-attention backend-selection logic in
nvte_get_fused_attn_backend, duplicating cuDNN's support rules. This list drifts out of sync as cuDNN evolves, and the support check can disagree with what actually runs.This PR replaces that logic with cuDNN-frontend's production-grade support checks. The new
nvte_get_fused_attn_backend_v2builds the same graph cuDNN executes at runtime, so the probe and execution can no longer diverge. It caches the graph on success and returns a diagnostic message on failure, giving users actionable guidance (e.g. adjust the config, GPU architecture, or cuDNN version).This PR also reworks
nvte_fused_attn_fwd/nvte_fused_attn_bwdintonvte_fused_attn_fwd_v2/nvte_fused_attn_bwd_v2, which take opaque, attribute-based config/params handles instead of long flat argument lists — improving TE's API and ABI stability.Legacy APIs are retained as deprecated shims that route through the v2 APIs, so existing callers keep working.
Type of change
Changes
API rework (opaque config/params + v2 entry points)
common/fused_attn/config_and_params.{h,cpp},common/include/transformer_engine/fused_attn.h): newNVTEFusedAttnConfig/NVTEFusedAttnFwdParams/NVTEFusedAttnBwdParamswithcreate/destroy/get/setattribute accessors, for better API/ABI stability. The cache key, probe, and execution now all originate from one place viamake_config/derive/make_cache_key.common/fused_attn/fused_attn*.{cpp,cu}):nvte_get_fused_attn_backend_v2,nvte_fused_attn_fwd_v2, andnvte_fused_attn_bwd_v2. The F16 and FP8is_supported_*probes copy the config, set direction,derive(), and attempt a null-pointer graph build viacheck_support— i.e. the same graph cuDNN builds at runtime, so probe and execution can't diverge.nvte_get_fused_attn_backend/nvte_fused_attn_fwd/nvte_fused_attn_bwdare retained, routed through the v2 APIs.csrc/extensions/attention.cpp) and JAX (jax/csrc/extensions/attention.cpp).Correctness & backend selection
cp_per_step_configsprobes each context-parallel step instead of only the global, non-CP config.log2(0)guard: avoids UB when casting-inftosize_tinget_max_batch_size/get_max_tokens.Diagnostics
NVTE_DEBUG/NVTE_DEBUG_LEVELfor JAX (parity with PyTorch): level 1 reports the selected backend; level 2 adds a diagnostic message explaining why fused attention was rejected.Cleanup / removals
NVTE_FUSED_ATTN_BACKEND— the two remaining backends (F16, FP8) are mutually exclusive now that max512 is gone.Q_ID/.../MASK_VAL_IDmacros (used only by the max512 backend).cudnn_frontend::xxxutility functions (used only byfp8_impl_v0and max512).fused_attn/headers.Tests
tests/pytorch/test_attention.py).Checklist: