diff --git a/python/sglang/srt/layers/radix_linear_attention.py b/python/sglang/srt/layers/radix_linear_attention.py index ad07aebe9..019b981b1 100644 --- a/python/sglang/srt/layers/radix_linear_attention.py +++ b/python/sglang/srt/layers/radix_linear_attention.py @@ -22,6 +22,12 @@ from torch import nn from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.compilation.piecewise_context_manager import get_forward_context +from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import ( + eager_on_graph, +) +from sglang.srt.model_executor.breakable_cuda_graph.context import ( + is_in_breakable_cuda_graph, +) from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.utils.custom_op import register_custom_op @@ -84,13 +90,22 @@ class RadixLinearAttention(nn.Module): dtype=mixed_qkv.dtype, device=mixed_qkv.device, ) - unified_linear_attention_with_output( - mixed_qkv, - a, - b, - output, - self.layer_id, - ) + if is_in_breakable_cuda_graph(): + bcg_unified_linear_attention_with_output( + mixed_qkv, + a, + b, + output, + self.layer_id, + ) + else: + unified_linear_attention_with_output( + mixed_qkv, + a, + b, + output, + self.layer_id, + ) return output else: return get_attn_backend().forward( @@ -136,3 +151,8 @@ def unified_linear_attention_with_output( output[:, :real_num_tokens].copy_(ret) return + + +bcg_unified_linear_attention_with_output = eager_on_graph(True)( + unified_linear_attention_with_output +)