[Kimi-Linear] Refactor Kimi-Linear to support RadixLinearAttention (#17506)
Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
from typing import Optional, Union
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import triton
|
||||
@@ -624,32 +624,24 @@ class KimiLinearAttnBackend(MambaAttnBackendBase):
|
||||
|
||||
def forward_decode(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
layer: RadixAttention,
|
||||
forward_batch: ForwardBatch,
|
||||
save_kv_cache: bool = True,
|
||||
layer: RadixLinearAttention,
|
||||
mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]],
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
**kwargs,
|
||||
):
|
||||
q_proj_states = kwargs["q_proj_states"]
|
||||
k_proj_states = kwargs["k_proj_states"]
|
||||
v_proj_states = kwargs["v_proj_states"]
|
||||
q_conv_weights = kwargs["q_conv_weights"]
|
||||
k_conv_weights = kwargs["k_conv_weights"]
|
||||
v_conv_weights = kwargs["v_conv_weights"]
|
||||
assert isinstance(mixed_qkv, Tuple)
|
||||
(q_proj_states, k_proj_states, v_proj_states) = mixed_qkv
|
||||
(q_conv_weights, k_conv_weights, v_conv_weights) = layer.conv_weights
|
||||
(q_conv_bias, k_conv_bias, v_conv_bias) = layer.bias
|
||||
|
||||
q_conv_bias = kwargs["q_conv_bias"]
|
||||
k_conv_bias = kwargs["k_conv_bias"]
|
||||
v_conv_bias = kwargs["v_conv_bias"]
|
||||
head_dim = layer.head_qk_dim
|
||||
layer_id = layer.layer_id
|
||||
beta = b
|
||||
g = a
|
||||
|
||||
head_dim = kwargs["head_dim"]
|
||||
layer_id = kwargs["layer_id"]
|
||||
beta = kwargs["beta"]
|
||||
g = kwargs["gate"]
|
||||
|
||||
A_log = kwargs["A_log"]
|
||||
dt_bias = kwargs["dt_bias"]
|
||||
A_log = layer.A_log
|
||||
dt_bias = layer.dt_bias
|
||||
|
||||
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
||||
q_conv_state, k_conv_state, v_conv_state = layer_cache.conv
|
||||
@@ -711,33 +703,29 @@ class KimiLinearAttnBackend(MambaAttnBackendBase):
|
||||
|
||||
def forward_extend(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
layer: RadixAttention,
|
||||
layer: RadixLinearAttention,
|
||||
forward_batch: ForwardBatch,
|
||||
save_kv_cache: bool = True,
|
||||
**kwargs,
|
||||
mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]],
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
**kwargs, # Unused, for compatibility with HybridLinearAttnBackend
|
||||
):
|
||||
from sglang.srt.layers.attention.mamba.causal_conv1d_triton import (
|
||||
causal_conv1d_fn,
|
||||
)
|
||||
|
||||
q_proj_states = kwargs["q_proj_states"]
|
||||
k_proj_states = kwargs["k_proj_states"]
|
||||
v_proj_states = kwargs["v_proj_states"]
|
||||
q_conv_weights = kwargs["q_conv_weights"]
|
||||
k_conv_weights = kwargs["k_conv_weights"]
|
||||
v_conv_weights = kwargs["v_conv_weights"]
|
||||
assert isinstance(mixed_qkv, Tuple)
|
||||
(q_proj_states, k_proj_states, v_proj_states) = mixed_qkv
|
||||
(q_conv_weights, k_conv_weights, v_conv_weights) = layer.conv_weights
|
||||
(q_conv_bias, k_conv_bias, v_conv_bias) = layer.bias
|
||||
|
||||
q_conv_bias = kwargs["q_conv_bias"]
|
||||
k_conv_bias = kwargs["k_conv_bias"]
|
||||
v_conv_bias = kwargs["v_conv_bias"]
|
||||
head_dim = layer.head_qk_dim
|
||||
layer_id = layer.layer_id
|
||||
beta = b
|
||||
g = a
|
||||
|
||||
head_dim = kwargs["head_dim"]
|
||||
layer_id = kwargs["layer_id"]
|
||||
beta = kwargs["beta"]
|
||||
g = kwargs["gate"]
|
||||
A_log = layer.A_log
|
||||
dt_bias = layer.dt_bias
|
||||
|
||||
query_start_loc = self.forward_metadata.query_start_loc
|
||||
cache_indices = self.forward_metadata.mamba_cache_indices
|
||||
@@ -836,7 +824,7 @@ class GDNAttnBackend(MambaAttnBackendBase):
|
||||
self,
|
||||
layer: RadixLinearAttention,
|
||||
forward_batch: ForwardBatch,
|
||||
mixed_qkv: torch.Tensor,
|
||||
mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]],
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
**kwargs, # Unused, for compatibility with HybridLinearAttnBackend
|
||||
@@ -856,6 +844,7 @@ class GDNAttnBackend(MambaAttnBackendBase):
|
||||
query_start_loc = self.forward_metadata.query_start_loc
|
||||
cache_indices = self.forward_metadata.mamba_cache_indices
|
||||
|
||||
assert isinstance(mixed_qkv, torch.Tensor)
|
||||
mixed_qkv = causal_conv1d_update(
|
||||
mixed_qkv,
|
||||
conv_states,
|
||||
@@ -907,11 +896,12 @@ class GDNAttnBackend(MambaAttnBackendBase):
|
||||
self,
|
||||
layer: RadixLinearAttention,
|
||||
forward_batch: ForwardBatch,
|
||||
mixed_qkv: torch.Tensor,
|
||||
mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]],
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
**kwargs, # Unused, for compatibility with HybridLinearAttnBackend
|
||||
):
|
||||
assert isinstance(mixed_qkv, torch.Tensor)
|
||||
seq_len = mixed_qkv.shape[0]
|
||||
|
||||
conv_weights = layer.conv_weights
|
||||
@@ -1234,7 +1224,7 @@ class HybridLinearAttnBackend(AttentionBackend):
|
||||
q: Optional[torch.Tensor] = None, # For full attention
|
||||
k: Optional[torch.Tensor] = None, # For full attention
|
||||
v: Optional[torch.Tensor] = None, # For full attention
|
||||
mixed_qkv: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
mixed_qkv: Optional[Union[torch.Tensor, Tuple[torch.Tensor, ...]]] = None,
|
||||
a: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
b: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
**kwargs,
|
||||
@@ -1266,7 +1256,7 @@ class HybridLinearAttnBackend(AttentionBackend):
|
||||
q: Optional[torch.Tensor] = None, # For full attention
|
||||
k: Optional[torch.Tensor] = None, # For full attention
|
||||
v: Optional[torch.Tensor] = None, # For full attention
|
||||
mixed_qkv: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
mixed_qkv: Optional[Union[torch.Tensor, Tuple[torch.Tensor, ...]]] = None,
|
||||
a: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
b: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
**kwargs,
|
||||
@@ -1298,7 +1288,9 @@ class HybridLinearAttnBackend(AttentionBackend):
|
||||
layer: RadixAttention = None,
|
||||
forward_batch: ForwardBatch = None,
|
||||
save_kv_cache: bool = True,
|
||||
mixed_qkv: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
mixed_qkv: Optional[
|
||||
Union[torch.Tensor, Tuple[torch.Tensor, ...]]
|
||||
] = None, # For GDN linear attention
|
||||
a: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
b: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||
**kwargs,
|
||||
@@ -1308,9 +1300,15 @@ class HybridLinearAttnBackend(AttentionBackend):
|
||||
|
||||
if forward_batch.forward_mode.is_idle():
|
||||
if is_linear_attn:
|
||||
return mixed_qkv.new_empty(
|
||||
mixed_qkv.shape[0], layer.num_v_heads, layer.head_v_dim
|
||||
)
|
||||
# KDA:
|
||||
if isinstance(mixed_qkv, tuple):
|
||||
return mixed_qkv[0].new_empty(
|
||||
mixed_qkv[0].shape[0], layer.num_v_heads, layer.head_v_dim
|
||||
)
|
||||
else: # GDN:
|
||||
return mixed_qkv.new_empty(
|
||||
mixed_qkv.shape[0], layer.num_v_heads, layer.head_v_dim
|
||||
)
|
||||
return q.new_empty(q.shape[0], layer.tp_q_head_num * layer.v_head_dim)
|
||||
elif forward_batch.forward_mode.is_decode():
|
||||
return self.forward_decode(
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
"""Radix linear attention."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
from typing import TYPE_CHECKING, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -36,7 +36,8 @@ class RadixLinearAttention(nn.Module):
|
||||
head_qk_dim: int,
|
||||
head_v_dim: int,
|
||||
attention_tp_size: int = 1,
|
||||
conv_weights: Optional[torch.Tensor] = None,
|
||||
# GDN KDA Shared Weights
|
||||
conv_weights: Optional[Union[torch.Tensor, Tuple[torch.Tensor, ...]]] = None,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
activation: str = "silu",
|
||||
A_log: Optional[torch.Tensor] = None,
|
||||
@@ -64,13 +65,14 @@ class RadixLinearAttention(nn.Module):
|
||||
self.conv_weights = conv_weights
|
||||
self.bias = bias
|
||||
self.activation = activation
|
||||
|
||||
self.A_log = A_log
|
||||
self.dt_bias = dt_bias
|
||||
|
||||
def forward(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
mixed_qkv: torch.Tensor,
|
||||
mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]],
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
|
||||
Reference in New Issue
Block a user