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 adds FP8 fused-attention support for THD ragged inputs.
Confidence Score: 5/5The PR appears safe to merge. No blocking failure remains; the prior mask-selection mismatch, debug output, and forward sequence-length out-of-bounds issue are fixed at the current head. Important Files Changed
Flowchart%%{init: {'theme': 'neutral'}}%%
flowchart LR
A[PyTorch FP8 attention with THD tensors] --> B[FP8 backend selection]
B -->|cuDNN 9.23+, supported SM, padding mask| C[FP8 fused-attention dispatcher]
C --> D[Convert sequence lengths]
C --> E[Generate 64-bit ragged offsets]
D --> F[cuDNN FP8 SDPA graph]
E --> F
F --> G[THD output and training context]
G --> H[FP8 backward graph]
Reviews (9): Last reviewed commit: "Merge branch 'main' into fp8_thd_attenti..." | Re-trigger Greptile |
| print(f"qkv_format: {qkv_format}") | ||
| print(f"inp shape: {inp[0].shape}, {inp[1].shape}, {inp[2].shape}") |
There was a problem hiding this comment.
Debug prints committed to test code
Two print statements were left in _run_dpa_fp8_vs_f16, which will produce noise on every test run and in CI logs. These look like leftover debug instrumentation from development.
| print(f"qkv_format: {qkv_format}") | |
| print(f"inp shape: {inp[0].shape}, {inp[1].shape}, {inp[2].shape}") | |
| with autocast(enabled=fp8_dpa, recipe=fp8_recipe): |
There was a problem hiding this comment.
We can probably remove these debug prints once everything is verified to be fine?
| 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!
| bool is_ragged_kv = (kv_format == NVTE_QKV_Format::NVTE_THD); | ||
| const int device_id = cuda::current_device(); | ||
| const int sm_arch_ = cuda::sm_arch(device_id); | ||
| bool use_ragged_stats = is_ragged_q && cudnn_runtime_version >= 90600 && sm_arch_ != 120; |
There was a problem hiding this comment.
I wonder if all the 90600 occurrences should be 92300, given that we only start to support FP8 THD from 9.23?
| bool use_ragged_stats = is_ragged_q && cudnn_runtime_version >= 90600 && sm_arch_ != 120; | ||
|
|
||
| NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout); | ||
| const DType ragged_offset_type = cudnn_runtime_version >= 90500 ? DType::kInt64 : DType::kInt32; |
There was a problem hiding this comment.
Could you check with cuDNN see if FP8 THD is supported with both types, or just one, i.e. kInt64? Also, should the runtime version here really be 92300 (if both types were supported)?
| const DType ragged_offset_type = cudnn_runtime_version >= 90500 ? DType::kInt64 : DType::kInt32; | ||
|
|
||
| int64_t actual_b = b; | ||
| if ((is_ragged_q || is_ragged_kv) && cudnn_runtime_version >= 90600) { |
| 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 |
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: