Skip to content

[Test] Relax varlen tolerance for bf16 head_dim=256 - #2864

Open
machero wants to merge 1 commit into
Dao-AILab:mainfrom
machero:fix/varlen-bf16-hd256-tolerance
Open

machero wants to merge 1 commit into
Dao-AILab:mainfrom
machero:fix/varlen-bf16-hd256-tolerance

Conversation

@machero

@machero machero commented Sep 7, 2026 •

Copy link
Copy Markdown

For bf16 backward, increasing head_dim from 128 to 256 (#2412) increases the reduction length of the dot products involved in computing dP (and consequently dS), which can lead to larger floating-point discrepancies between the FA4 and PyTorch SDPA implementations. These discrepancies can propagate to dK/dV, and are further exposed in MQA/GQA where gradients for shared KV heads require additional accumulation.

test_varlen compares flash_attn_varlen_func against F.scaled_dot_product_attention
with a fixed atol=rtol=3e-2 that was calibrated for head_dim <= 128. Since
head_dim=256 support landed (Dao-AILab#2412), the bf16 backward dK/dV gradients
(dK = dS^T @ Q, dV = P^T @ dO, summed over 256 dims) accumulate ~2x the rounding
error of D=128 and reach ~6e-2, failing three MQA/GQA varlen cases. fp16 is
unaffected (11-bit mantissa).

Loosen the tolerance to 1e-1 for bf16 + head_dim=256 only; all other configs keep 3e-2.

Co-Authored-By: Claude Code <noreply@anthropic.com>

This branch has not been deployed

No deployments
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.

1 participant