The calculation of scores_max in the flash_attn example #948
Replies: 10 comments
|
You're right that the code you quoted never combines the new block max with the previous one. That has since been added: #1269 (merged 2025-11-17) put for i in T.Parallel(block_M):
scores_max[i] = T.max(scores_max[i], scores_max_prev[i])right after the Without that line,
I ran
On your second question: yes. The |
|
|
|
|
|
|
|
|
|
Uh oh!
There was an error while loading. Please reload this page.
In the FA2 algorithm, there is a step that updates$max_{new}$ , which requires taking the element-wise maximum between the current block's max and the previous $max_{new}$ .
However, in the code at
tilelang/examples/flash_attention/example_gqa_fwd_bshd.py
Line 116 in 7fb0677
It seems that the operation
scores_max = max(scores_max_prev, scores_max)is not actually performed due toT.fill. Could you clarify where my understanding might be incorrect?Additionally, would
T.fill(scores_max, -T.infinity(accum_dtype))be equivalent to settingclear=TrueinT.reduce_max?All reactions