Plumb FP8+THD - #2994
Conversation
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Take upstream versions for fused_attn.cpp and fused_attn_fp8.cu APIs. Keep branch's test_attention.py THD parametrization. FP8+THD ragged-offset plumbing is re-applied in the following commit. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Mirrors the F16 arbitrary_seqlen ragged-offset pattern in the FP8 path: - Backend selector: enable FP8+THD for cuDNN >= 9.23 on sm >= 100 - fwd/bwd _impl: ragged detection, batch/seqlen bucketing, set_ragged_offset() on Q/K/V/O/dO/dQ/dK/dV/Stats, workspace allocation for ragged offsets, cu_seqlens_padded_to_offsets kernel - fwd/bwd dispatchers: accept num_tokens_q/kv, cu_seqlens_padded, compute max_batch/max_tokens, THD Stats shape Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
for more information, see https://pre-commit.ci
Greptile SummaryThis PR enables FP8 attention with THD (token × head × dim, ragged/variable-length) format for both standalone and context-parallel (CP) attention paths by wiring ragged-offset tensors into the cuDNN FP8 SDPA graph, adding quantized batch/token sizing helpers, and lifting the blanket
Confidence Score: 4/5
Important Files Changed
Sequence DiagramsequenceDiagram
participant Caller
participant fused_attn_fp8_fwd as fused_attn_fp8_fwd (C++)
participant impl as fused_attn_fp8_fwd_impl
participant q_kernel as cu_seqlens_to_actual_seqlens
participant off_kernel as cu_seqlens_padded_to_offsets
participant cuDNN as cuDNN SDPA Graph
Caller->>fused_attn_fp8_fwd: batch, max_seqlen, num_tokens, cu_seqlens, cu_seqlens_padded
fused_attn_fp8_fwd->>fused_attn_fp8_fwd: detect THD format → get_max_batch_size / get_max_tokens
fused_attn_fp8_fwd->>impl: max_b, max_t_q, max_t_kv, devPtrSeqOffsetsQ/KV
impl->>impl: "b = max_b (if !use_cu_seqlens_directly)"
impl->>impl: build cuDNN graph with ragged_offset tensors (offset_q/k/v/o/stats)
impl->>q_kernel: actual_b, max_b → fill actual-seqlen array
impl->>off_kernel: actual_b, max_b, cu_seqlens_padded → fill int64 ragged offsets
impl->>cuDNN: variant_pack with ragged offsets + seq_q/kv + FP8 scales
cuDNN-->>impl: FP8 output O, amax, stats
impl-->>Caller: output tensors
Reviews (14): Last reviewed commit: "Merge branch 'main' into fp8_thd_attenti..." | Re-trigger Greptile |
| checkpoint_core_attention=False, | ||
| core_attention_bias_type=config.attn_bias_type, | ||
| fp8_output=fp8_dpa, | ||
| fast_zero_fill=False, |
There was a problem hiding this comment.
cuDNN doesn't touch the pad tokens (between seqs or at the end of the batch) so we had to zero out the entire output for F16 THD (see here). I wonder if we need to do the same for FP8?
cyanguwa
left a comment
There was a problem hiding this comment.
Overall looks good, but please address the few comments and pass the CI. Thanks for the PR!
| std::shared_ptr<fe::graph::Tensor_attributes>, // offset_o | ||
| std::shared_ptr<fe::graph::Tensor_attributes>, // offset_k | ||
| std::shared_ptr<fe::graph::Tensor_attributes>, // offset_v | ||
| std::shared_ptr<fe::graph::Tensor_attributes>, // offset_stats |
There was a problem hiding this comment.
Nit: do we want to follow the same order of listing as in F16, i.e. offset_q/k/v/o/stats?
|
|
||
| auto plan_workspace_size = mha_graph->get_workspace_size(); | ||
| attn_scale, O, amax_s, amax_o, Stats, bias, softmax_offset, seq_q, seq_kv, offset_q, | ||
| offset_o, offset_k, offset_v, offset_stats, dropout_seed, dropout_offset] = |
There was a problem hiding this comment.
Same comment about the ordering for q/k/v/o/stats.
Use cuDNN 9.23 as the FP8 THD ragged-offset gate and prefer int64 offsets, matching cuDNN guidance for the new path. Restrict FP8 THD backend selection to padding masks, align ragged offset tuple order with the F16 convention, and enable zero-fill for FP8 THD comparison tests. Suppress the forward FP8 graph-builder fn_size lint using the same local pattern already used by the backward builder, because refactoring the full graph construction is outside this review cleanup. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
df0f69f to
46153a4
Compare
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
|
/te-ci pytorch L1 |
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
| auto offset_q_tuple = is_ragged_q ? std::make_tuple(offset_q) : std::make_tuple(nullptr); | ||
| auto offset_kv_tuple = | ||
| is_ragged_kv ? std::make_tuple(offset_k, offset_v) : std::make_tuple(nullptr, nullptr); | ||
| auto offset_o_tuple = is_ragged_q ? std::make_tuple(offset_o) : std::make_tuple(nullptr); |
There was a problem hiding this comment.
Ok, maybe I was making a bigger deal out of this than it is. It looks like we are following the q/k/v/o/stats order in the shared_ptr declaration but not in the tuple creation, in the F16 file. So feel free to revert this change, or just change the shared_ptr order if you like. Thanks.
|
@sudhakarsingh27 Are there any plans for cudnn thd sm90 fp8 support? |
Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
|
/te-ci pytorch L1 |
The shared conversion kernel now accepts RaggedOffsetMultipliers, so construct it in FP8 forward and backward instead of passing the removed scalar argument list. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
8e4cbfc to
01144a0
Compare
|
/te-ci pytorch L1 |
Keep the actual batch when cuDNN consumes user cu-seqlens because a bucketed batch would read past the buffers. Reuse the aligned fallback workspace and keep SM120 stats dense so allocation matches the graph layout. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
|
/te-ci pytorch L1 |
The optimized mha_fill path reads CUDA cu_seqlens from host C++ and segfaults for THD. A controlled A/B passed with False while the enabled path exited 139. Keep the comparison test on the safe path until a graph-safe zero-fill implementation lands. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
0d77900 to
7d785e8
Compare
|
/te-ci pytorch L1 |
|
/te-ci pytorch L1 |
FP8 Linear flattens THD inputs to [t, h*d], so adjust generated sequence lengths before building cu_seqlens to make the total token count divisible by eight for both forward and backward. Keep fast zero fill disabled in the MHA helper to avoid the known host dereference path. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
|
/te-ci pytorch L1 |
Remove obsolete FP8+THD+CP gates now that the fused-attention path supports this combination. Replace the unsafe host-side suffix calculation with stream-ordered zeroing, and preserve FP8 metadata and token-major THD layout in the all-gather path. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
for more information, see https://pre-commit.ci
P2P FP8 backward densely combines rank-local partial gradients, which can repopulate inter-sequence dK/dV padding after native zero initialization. Reconstruct every local padding interval from actual and padded cu-seqlens and clear dQ/dK/dV after reduction. Gate the cleanup on FP8 backward because EOS H100 and Prenyx B200 controls showed BF16 and high-precision backward padding already remains zero. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
for more information, see https://pre-commit.ci
|
/te-ci pytorch L1 |
|
@bbuschkaemper, are you looking for both fwd and bwd in fp8+thd? SM90 in cudnn only supports fp8+thd in fwd right now |
cuDNN 9.25 provides a working Hopper FP8 THD forward kernel, while 9.23 selects a plan that traps with an illegal instruction. Keep Hopper backward gated and preserve the existing cuDNN 9.23 requirement on Blackwell. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Reuse four existing Hopper cases for forward-only validation instead of adding a new Cartesian test axis. Blackwell retains forward-and-backward coverage, while Hopper runs the cuDNN 9.25-supported inference path. Signed-off-by: Sudhakar Singh <sudhakars@nvidia.com>
Description
Please include a brief summary of the changes, relevant motivation and context.
Fixes # (issue)
Type of change
Changes
Please list the changes introduced in this PR:
Checklist: