Hacktoberfest 2026:メンテナが10月に向けて印を付けた、オープンで初心者向けの issue。 Hacktoberfest の issue を見る

FA4 returns an all-zero output for causal attention on SM100 after flash-attention #2490

オープン
#3,528 コメント 3 件 リアクション 0 件 担当者 0 名 GitHub で見る

メンテナーはふだん 2 日以内に返信

@nvegesna-netizen がすでに取り組んでいます。

2026年9月17日 から。

  • #3532 @nvegesna-netizen による — オープン

評価

難易度
4/5
見積もり時間
3〜5日
初心者へのやさしさ
20/100
issue の種類
バグ
明瞭さ
明確に書かれている
活発さ
停滞
技術スタック
python, pytorch

調査の方向性

dot_product_attention/utils.py と issue で指定されている FA4 のエントリーポイントから始め、続いて run_attention_with_cp.py と QA テスト設定を確認します。SM100 上で causal attention を使用し、CP と non-CP の両方のパスについて動作を検証し、報告された fix と回帰カバレッジについて #3532 を確認します。

索引モデルが issue の本文から書いたものです。

説明

Summary

With a flash-attn-4 build that includes #2490,
FlashAttention 4 returns an all-zero output for causal attention on SM100 when called the way
TransformerEngine calls it. No context parallelism is needed to reproduce it. It is silent: no
exception, no warning, finite values, every element zero.

Context parallelism is only where it was first noticed, because the p2p ring merges the
accompanying -inf log-sum-exp and turns a silent wrong answer into a visible NaN.

Edited 2026-09-17. This issue originally described a CP-specific NaN and, in a later edit,
argued the window encoding was not at fault. Both were wrong. Direct measurement against a float64
reference shows the defect is general to FA4 causal attention on SM100. History of both earlier
readings is kept at the bottom. Fixed by #3532, which also repairs head_dim 512 under context
parallelism -- see below; the earlier suggestion that 512 might be a second defect is retracted.

Reproduction, without context parallelism

import torch, transformer_engine.pytorch as te
b, s, h, d = 2, 1024, 8, 128
q, k, v = (torch.randn(b, s, h, d, device="cuda", dtype=torch.bfloat16) for _ in range(3))
dpa = te.DotProductAttention(h, d, qkv_format="bshd", attn_mask_type="causal",
                             attention_dropout=0.0).to(dtype=torch.bfloat16, device="cuda")
print(dpa(q, k, v).abs().max())   # tensor(0., device='cuda:0')

Run with NVTE_FUSED_ATTN=0 NVTE_UNFUSED_ATTN=0 so the selection cannot fall through to cuDNN.
Measured on B200:

Selected backend = FlashAttention (4.0.0b31.dev9+g346aa9e)
context_parallel: False, cp_size: 1
  flash_attn_func_v4  causal=True  window_size=(-1, 0)
  out.abs().max()=0.000000e+00  zeros=YES  relerr vs float64 reference = 1.000e+00

With NVTE_FLASH_ATTN_V4=0 the same script selects FlashAttention 2 and returns
out.abs().max()=3.453125e+00, relerr=2.046e-03. The reference is float64, not float32: torch
uses TF32 for fp32 matmuls on Ampere and newer, and TF32 carries an 11-bit significand, the same as
fp16, so an fp32 reference cannot judge a bf16 kernel.

The kernel, with TransformerEngine out of the picture

from flash_attn.cute.interface import flash_attn_func
flash_attn_func(q, k, v, causal=True, window_size=(-1, 0))     # out.abs().max() = 0.0
flash_attn_func(q, k, v, causal=True, window_size=(None, None))# out.abs().max() = 3.453125, relerr 2.0e-03

Why

TransformerEngine normalises causal masking to window_size = (-1, 0)
(dot_product_attention/utils.py, check_set_window_size) and passes it to FA4, whose sentinel for
"unbounded" is None, not -1. Dao-AILab/flash-attention#2490 ("cute: don't widen intentionally-empty offset windows to full
attention", merged 2026-09-09) narrowed the shim that used to absorb the mismatch:

# before #2490 -- causal applied first, then widen if the SUM is negative
if causal: window_size_right = 0
if wsl is not None and wsr is not None and wsl + wsr < 0: wsl = wsr = None

# after #2490 -- widen only if BOTH are negative, then apply causal
if wsl is not None and wsr is not None and wsl < 0 and wsr < 0: wsl = wsr = None
if causal: window_size_right = 0

Under the old rule (-1, 0) summed to -1 and always collapsed to (None, None), so TE's window
encoding was discarded and causal=True alone set the mask. Under the new rule (-1, 0) survives:
_resolve_causal_local_window(True, -1, 0) returns (False, True, -1, 0), so causal is dropped and a
local band is applied with left bound -1. mask.py guards on window_size_left is not None and
uses the value arithmetically, giving a band of [row+1, row] -- empty. Measured directly:

lse   n=131072  -inf=131072  +inf=0  nan=0     every element
out   n=16777216  -inf=0  +inf=0  nan=0        every element zero

FA4 is internally inconsistent about this: the d512 dispatch at interface.py:783 tests
window_size_left in (None, -1), treating -1 as "no window", while the resolver treats it as a
bound. But TE should not be sending -1 regardless.

Why CI and the CP tests do not catch it

Three independent reasons, all worth fixing on their own:

  1. NVTE_FLASH_ATTN_V4=0 is exported in qa/L1_pytorch_distributed_unittest/test.sh and the L0
    suites, so FA4 is never exercised under context parallelism.
  2. CI pins flash-attn-4==4.0.0b11, which predates Dao-AILab/flash-attention#2490.
  3. run_attention_with_cp.py compares a CP run against a non-CP run of the same backend
    (tensors_no_cp vs tensors_cp). When both sides return zeros they agree, and the test passes
    while measuring nothing. This is why all_gather and a2a report PASS on a backend that is
    returning zeros -- a systematic error appears identically on both sides and cancels. Only p2p
    fails, because the ring merge computes exp(lse_step - lse) over -inf and produces NaN, which
    then differs from the zeros on the other side.

Point 3 means the CP suite cannot detect any error that affects CP and non-CP equally. That is a
gap independent of this bug.

Scope

Any configuration on SM100 where TE selects FA4 for a causal mask, with a flash-attn-4 that includes
Dao-AILab/flash-attention#2490. cuDNN FusedAttention usually wins backend selection for common shapes, which is likely why
this has gone unnoticed -- but it does not win where cuDNN has no engine, which includes large head
dims, and it does not win when fused attention is disabled.

Suggested fix

TE should express "unbounded" to FA4 as None rather than -1, at every FA4 entry point --
both the public flash_attn_func_v4 used on the non-CP path and _flash_attn_fwd_v4 /
_flash_attn_bwd_v4 used by context parallelism. Genuine sliding windows already pass non-negative
bounds and are unaffected.

A partial fix is worse than none here, and this is not hypothetical: I first tried translating the
sentinel only at _flash_attn_fwd_v4 / _flash_attn_bwd_v4, which corrected the CP side while
leaving the non-CP reference side returning zeros. The self-referential comparison above then
reported the fix as a regression in all_gather and a2a.

head_dim 512 (previously flagged as possibly a second defect -- it is not)

At symmetric head_dim=512, with flash-attention Dao-AILab/flash-attention#2877 supplying the
kernels and TE's d512 gate lifted, all three cp_comm_type values failed with dq has nan values.
I originally suggested that might be an independent defect, reasoning that all_gather and a2a
perform no cross-step LSE merge so the ring correction could not be the source.

That reasoning was wrong. The backward consumes the log-sum-exp regardless of comm type, so an
all -inf LSE poisons exp(scores - lse) in any configuration. Measured as an A/B on one build,
one set of kernels, with the normalizer switched off and on at runtime:

cp_comm_type sentinel normalization off on
p2p FAIL dq has nan values PASS
all_gather FAIL dq has nan values PASS
a2a FAIL dq has nan values PASS

and at the kernel level, with no TransformerEngine involved:

head_dim 512, causal, window_size=(-1, 0)      out.abs().max() = 0.0
head_dim 512, causal, window_size=(None, None) out.abs().max() = 3.875

Same empty band as head_dim 128. One root cause, one fix. At head_dim 128 the degenerate input
yields zeros rather than NaN, which is why those arms passed against an equally-zero reference.

Ruled out

  • TE passing uninitialised dq. context_parallel.py uses torch.empty_like for the gradient
    buffers; patched to zeros_like and the NaN persisted. FA4's backward substitutes its own work
    buffer and copies back.
  • An LSE convention mismatch. FA4's forward LSE measured at (2, 4, 512) against an fp32
    reference: natural log, max error 1.3e-5.

Earlier readings, kept for the record

The first version of this issue attributed the NaN to the ring correction consuming an -inf LSE.
That mechanism is real and is why p2p fails loudly, but it is a consequence, not the root cause,
and it described the problem as CP-specific when it is not.

A later edit argued the window encoding was not at fault, reasoning that all_gather and a2a
send the same (-1, 0) and pass. That inference was wrong: they pass because the test compares them
against an equally broken reference, not because their output is correct.

Environment

  • B200 (SM100), nvcr.io/nvidia/pytorch:26.08-py3, cuDNN 9.25.0
  • TransformerEngine 2.18.0
  • flash-attn-4 4.0.0b31.dev9+g346aa9e; any build including Dao-AILab/flash-attention#2490 should reproduce the non-CP case
主要言語
Python
スター
3.6k
フォーク
851
平均マージ
4日 14時間
マージ済み PR(30日)
54

環境構築

はじめの一歩

  1. issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
  2. 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
  3. リポジトリをフォークし、ブランチを切って変更します。
  4. issue 番号を参照したプルリクエストを送ります。

NVIDIA/TransformerEngine のほかの issue

NVIDIA/TransformerEngine の issue をすべて見る

似ている issue

Python の issue をもっと見る

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。