remove redundant code of has_block

This commit is contained in:
Kaitao Yang 2026-01-03 22:38:37 -08:00
commit 1ea6585b0c

View file

@ -219,16 +219,10 @@ def run_attention(
)
if config.n_groups != 1 and not requires_grad:
if has_block:
out = out.view(bsz, q_len, config.n_kv_heads, config.n_groups, head_dim)
else:
out = out.view(bsz, q_len, config.n_kv_heads, config.n_groups, head_dim)
out = out.view(bsz, q_len, config.n_kv_heads, config.n_groups, head_dim)
out = out.reshape(bsz, q_len, n_heads, head_dim)
else:
if has_block:
out = out.view(bsz, q_len, n_heads, head_dim)
else:
out = out.view(bsz, q_len, n_heads, head_dim)
out = out.view(bsz, q_len, n_heads, head_dim)
return out
else:
local_mask = context.attention_mask