From cd0d9e1a91c92fcab3989da4a4b1856a46ca4b3f Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 30 Jul 2026 16:16:56 +0000 Subject: [PATCH] [OPT] Fix online softmax bug in Triton V2 Phase 1 Bug: if l_i > 0 branch in Triton is invalid (compiled as constexpr). Also: p = exp(scores - m_i_new) computed after m_i_new update was using the wrong reference max (should subtract m_ij first, then rescale). Fix: Adapted exactly from prefix_prefill.py's proven-correct pattern: p = exp(scores - m_ij) # probs relative to chunk max l_ij = sum(p) # chunk sum m_i_new = max(m_i, m_ij) # new running max alpha = exp(m_i - m_i_new) # old accumulator rescale beta = exp(m_ij - m_i_new) # new chunk rescale l_i_new = alpha*l_i + beta*l_ij acc = acc*(alpha*l_i/l_i_new) + (p*beta/l_i_new) @ V This is the Flash Attention online softmax tiling algorithm. Same math as CCCL's parallel_reduce with compound accumulators. --- paged_attention_v2_triton.py | 32 ++++++++++++++------------------ 1 file changed, 14 insertions(+), 18 deletions(-) diff --git a/paged_attention_v2_triton.py b/paged_attention_v2_triton.py index 36e7fc3c..b751674f 100644 --- a/paged_attention_v2_triton.py +++ b/paged_attention_v2_triton.py @@ -146,34 +146,30 @@ def _paged_attn_v2_partition_kernel( scores = tl.sum(q[:, None] * k, axis=0) * scale # [BLOCK_N] scores = tl.where(valid_mask, scores, float('-inf')) - # Online softmax update + # Online softmax (adapted from prefix_prefill.py — proven correct) m_ij = tl.max(scores, axis=0) # scalar: max of this chunk + p = tl.exp(scores - m_ij) # [BLOCK_N] — unnormalized probs + l_ij = tl.sum(p, axis=0) # scalar: sum of exp for this chunk + m_i_new = tl.maximum(m_i, m_ij) - - alpha = tl.exp(m_i - m_i_new) - beta = tl.exp(m_ij - m_i_new) - - p = tl.exp(scores - m_i_new) # [BLOCK_N] - l_ij = tl.sum(p, axis=0) - - l_i_new = alpha * l_i + beta * l_ij if l_i > 0 else l_ij + alpha = tl.exp(m_i - m_i_new) # rescale factor for old accumulator + beta = tl.exp(m_ij - m_i_new) # rescale factor for new chunk + l_i_new = alpha * l_i + beta * l_ij # === Paged V gather === - # V offsets: value_cache[bn, kv_head, d, within_block] - # Layout: [num_blocks, num_kv_heads, head_size, block_size] off_v = (bn[:, None] * stride_vc_b + kv_head_idx * stride_vc_h + offs_d[None, :] * stride_vc_d + within_block[:, None] * stride_vc_bs) - v = tl.load(value_cache_ptr + off_v, mask=valid_mask[:, None], other=0.0) # [N, D] - # Update accumulator: acc = (acc * alpha * l_i / l_i_new) + (p @ V * beta / l_i_new) - if l_i > 0: - acc_scale = l_i / l_i_new * alpha - acc = acc * acc_scale - - p_scaled = p / l_i_new * beta # [BLOCK_N] + # Update accumulator (Flash Attention online softmax pattern): + # acc = acc * (alpha * l_i / l_i_new) + (p * beta / l_i_new) @ V + # Safe division: if l_i_new == 0, this is the first chunk + acc_scale = alpha * l_i / tl.maximum(l_i_new, 1e-6) + acc = acc * acc_scale + p_scale = beta / tl.maximum(l_i_new, 1e-6) + p_scaled = p * p_scale # [BLOCK_N] acc += tl.sum(p_scaled[:, None] * v, axis=0) # [HEAD_DIM] l_i = l_i_new