[Fix] Support LSE on the RadixAttention extra-kwargs graph path (#35453)
This commit is contained in:
@@ -230,6 +230,7 @@ class RadixAttention(nn.Module):
|
||||
"score_mod",
|
||||
"aux_tensors",
|
||||
"rel_bias",
|
||||
"return_lse",
|
||||
"q_descale",
|
||||
"k_descale",
|
||||
"v_descale",
|
||||
@@ -240,13 +241,16 @@ class RadixAttention(nn.Module):
|
||||
# schema; route this backend's extend attention through the plain
|
||||
# eager path.
|
||||
if is_in_breakable_cuda_graph():
|
||||
breakable_attention_with_output_extra_kwargs(
|
||||
lse = breakable_attention_with_output_extra_kwargs(
|
||||
q, k, v, output, save_kv_cache, self.layer_id, kwargs
|
||||
)
|
||||
else:
|
||||
attention_with_output_extra_kwargs(
|
||||
lse = attention_with_output_extra_kwargs(
|
||||
q, k, v, output, save_kv_cache, self.layer_id, kwargs
|
||||
)
|
||||
if kwargs.get("return_lse") or forward_batch.mha_return_lse:
|
||||
assert lse is not None
|
||||
return output.view(-1, self.tp_q_head_num, self.v_head_dim), lse
|
||||
return output
|
||||
# Chunked-prefix MHA needs LSE to merge independently normalized
|
||||
# suffix and cached-prefix attention states.
|
||||
@@ -586,7 +590,7 @@ def attention_with_output_extra_kwargs(
|
||||
save_kv_cache: bool,
|
||||
layer_id: int,
|
||||
extra_kwargs: dict,
|
||||
) -> None:
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""Breakable/tc_piecewise attention for backends whose forward needs kwargs
|
||||
that cannot cross the ``unified_attention_with_output`` custom-op schema --
|
||||
a ``score_mod`` callable and/or ``aux_tensors`` (e.g. Inkling's relative-bias
|
||||
@@ -628,12 +632,24 @@ def attention_with_output_extra_kwargs(
|
||||
)
|
||||
forward_batch.out_cache_loc = original_out_cache_loc
|
||||
|
||||
return_lse = bool(kwargs.get("return_lse") or forward_batch.mha_return_lse)
|
||||
if return_lse:
|
||||
assert isinstance(ret, tuple)
|
||||
ret, lse, *_ = ret
|
||||
else:
|
||||
assert isinstance(ret, torch.Tensor)
|
||||
lse = None
|
||||
|
||||
if ret.data_ptr() != output.data_ptr():
|
||||
output[:real_num_tokens].view(ret.shape).copy_(ret)
|
||||
|
||||
if _is_hip:
|
||||
_zero_padded_pcg_tail(output, context)
|
||||
return
|
||||
if lse is not None and lse.shape[0] != output.shape[0]:
|
||||
padded_lse = lse.new_zeros((output.shape[0], *lse.shape[1:]))
|
||||
padded_lse[:real_num_tokens].copy_(lse)
|
||||
lse = padded_lse
|
||||
return lse
|
||||
|
||||
|
||||
breakable_attention_with_output_extra_kwargs = eager_on_graph(True)(
|
||||
|
||||
@@ -73,6 +73,7 @@ class TestRadixAttentionGraphInterface(CustomTestCase):
|
||||
out_cache_loc=torch.arange(num_tokens, dtype=torch.int64),
|
||||
positions=torch.arange(num_tokens, dtype=torch.int64),
|
||||
_attn_output=None,
|
||||
mha_return_lse=False,
|
||||
)
|
||||
return SimpleNamespace(
|
||||
forward_batch=forward_batch,
|
||||
@@ -221,6 +222,34 @@ class TestRadixAttentionGraphInterface(CustomTestCase):
|
||||
self.assertTrue(torch.all(lse[2:] == 0))
|
||||
self.assertIs(forward_batch.out_cache_loc, original_out_cache_loc)
|
||||
|
||||
def test_extra_kwargs_path_returns_bucket_shaped_lse(self):
|
||||
attention_layer = SimpleNamespace()
|
||||
context = self._new_impl_context([attention_layer])
|
||||
backend = _RecordingAttentionBackend()
|
||||
query = torch.zeros((4, 2, 3))
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
radix_attention_module,
|
||||
"get_tc_piecewise_forward_context",
|
||||
return_value=context,
|
||||
),
|
||||
patch.object(
|
||||
radix_attention_module, "get_attn_backend", return_value=backend
|
||||
),
|
||||
):
|
||||
lse = radix_attention_module.attention_with_output_extra_kwargs(
|
||||
query,
|
||||
query,
|
||||
query,
|
||||
torch.empty_like(query),
|
||||
False,
|
||||
0,
|
||||
{"return_lse": True},
|
||||
)
|
||||
|
||||
self.assertEqual(lse.tolist(), [[7, 7], [7, 7], [0, 0], [0, 0]])
|
||||
|
||||
def test_impl_uses_independent_query_and_key_value_extents(self):
|
||||
attention_layer = SimpleNamespace()
|
||||
context = self._new_impl_context([attention_layer])
|
||||
|
||||
Reference in New Issue
Block a user