[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",
|
"score_mod",
|
||||||
"aux_tensors",
|
"aux_tensors",
|
||||||
"rel_bias",
|
"rel_bias",
|
||||||
|
"return_lse",
|
||||||
"q_descale",
|
"q_descale",
|
||||||
"k_descale",
|
"k_descale",
|
||||||
"v_descale",
|
"v_descale",
|
||||||
@@ -240,13 +241,16 @@ class RadixAttention(nn.Module):
|
|||||||
# schema; route this backend's extend attention through the plain
|
# schema; route this backend's extend attention through the plain
|
||||||
# eager path.
|
# eager path.
|
||||||
if is_in_breakable_cuda_graph():
|
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
|
q, k, v, output, save_kv_cache, self.layer_id, kwargs
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
attention_with_output_extra_kwargs(
|
lse = attention_with_output_extra_kwargs(
|
||||||
q, k, v, output, save_kv_cache, self.layer_id, 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
|
return output
|
||||||
# Chunked-prefix MHA needs LSE to merge independently normalized
|
# Chunked-prefix MHA needs LSE to merge independently normalized
|
||||||
# suffix and cached-prefix attention states.
|
# suffix and cached-prefix attention states.
|
||||||
@@ -586,7 +590,7 @@ def attention_with_output_extra_kwargs(
|
|||||||
save_kv_cache: bool,
|
save_kv_cache: bool,
|
||||||
layer_id: int,
|
layer_id: int,
|
||||||
extra_kwargs: dict,
|
extra_kwargs: dict,
|
||||||
) -> None:
|
) -> Optional[torch.Tensor]:
|
||||||
"""Breakable/tc_piecewise attention for backends whose forward needs kwargs
|
"""Breakable/tc_piecewise attention for backends whose forward needs kwargs
|
||||||
that cannot cross the ``unified_attention_with_output`` custom-op schema --
|
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
|
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
|
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():
|
if ret.data_ptr() != output.data_ptr():
|
||||||
output[:real_num_tokens].view(ret.shape).copy_(ret)
|
output[:real_num_tokens].view(ret.shape).copy_(ret)
|
||||||
|
|
||||||
if _is_hip:
|
if _is_hip:
|
||||||
_zero_padded_pcg_tail(output, context)
|
_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)(
|
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),
|
out_cache_loc=torch.arange(num_tokens, dtype=torch.int64),
|
||||||
positions=torch.arange(num_tokens, dtype=torch.int64),
|
positions=torch.arange(num_tokens, dtype=torch.int64),
|
||||||
_attn_output=None,
|
_attn_output=None,
|
||||||
|
mha_return_lse=False,
|
||||||
)
|
)
|
||||||
return SimpleNamespace(
|
return SimpleNamespace(
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
@@ -221,6 +222,34 @@ class TestRadixAttentionGraphInterface(CustomTestCase):
|
|||||||
self.assertTrue(torch.all(lse[2:] == 0))
|
self.assertTrue(torch.all(lse[2:] == 0))
|
||||||
self.assertIs(forward_batch.out_cache_loc, original_out_cache_loc)
|
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):
|
def test_impl_uses_independent_query_and_key_value_extents(self):
|
||||||
attention_layer = SimpleNamespace()
|
attention_layer = SimpleNamespace()
|
||||||
context = self._new_impl_context([attention_layer])
|
context = self._new_impl_context([attention_layer])
|
||||||
|
|||||||
Reference in New Issue
Block a user