来源:
1. fla-org/flash-linear-attention (5538 stars)
→ upstream_ref/fla/ops/gated_delta_rule/naive.py (正确的纯 PyTorch GDN)
→ upstream_ref/fla/ops/gated_delta_rule/chunk.py (Triton chunk kernel)
→ upstream_ref/fla/layers/gated_deltanet.py (层集成)
2. vllm-project/vllm main (88717 stars)
→ upstream_ref/vllm_gdn/gdn/qwen_gdn_linear_attn.py (1751行, Qwen3.5 原生 GDN)
→ upstream_ref/vllm_gdn/ops/causal_conv1d.py (1289行, 正确的 Conv1d)
→ upstream_ref/vllm_gdn/third_party/ops/ (FLA Triton ops vendored)
→ upstream_ref/vllm_gdn/models/qwen3_5.py (vllm 最新 Qwen3.5 模型)
3. Deep-Spark/xllm (BI-V100 硬件厂商)
→ upstream_ref/xllm_latest/core/layers/npu_torch/qwen3_gated_delta_net_base.cpp (576行)
→ upstream_ref/xllm_latest/core/kernels/npu/npu_causal_conv1d.cpp
→ upstream_ref/xllm_latest/core/kernels/npu/npu_recurrent_gated_delta_rule.cpp
目的: 修复 corex_gdn.py Conv1d groups 接口不匹配问题
错误: conv1d_weight shape (2560,1,4) 被当成 (num_k_heads,1,4) 索引
conv_dim = key_dim*2 + value_dim = 10240, TP=4 后 2560
FLA naive.py 和 vllm qwen_gdn_linear_attn.py 有正确的实现可直接对接
59 lines
2.0 KiB
Python
59 lines
2.0 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
|
import torch
|
|
from transformers import PretrainedConfig
|
|
|
|
from vllm.config import (
|
|
VllmConfig,
|
|
)
|
|
from vllm.distributed import (
|
|
get_tensor_model_parallel_rank,
|
|
get_tensor_model_parallel_world_size,
|
|
)
|
|
from vllm.model_executor.custom_op import PluggableLayer
|
|
from vllm.model_executor.layers.mamba.abstract import MambaBase
|
|
from vllm.model_executor.layers.mamba.mamba_utils import (
|
|
MambaStateDtypeCalculator,
|
|
)
|
|
from vllm.model_executor.models.utils import extract_layer_index
|
|
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
|
|
|
|
|
|
class GatedDeltaNetAttention(PluggableLayer, MambaBase):
|
|
"""Base class for GatedDeltaNet attention layer."""
|
|
|
|
def __init__(
|
|
self,
|
|
config: PretrainedConfig,
|
|
vllm_config: VllmConfig,
|
|
prefix: str = "",
|
|
) -> None:
|
|
super().__init__()
|
|
self.prefix = prefix
|
|
self.tp_size = get_tensor_model_parallel_world_size()
|
|
self.tp_rank = get_tensor_model_parallel_rank()
|
|
self.layer_idx = extract_layer_index(prefix)
|
|
self.hidden_size = config.hidden_size
|
|
self.activation = config.hidden_act
|
|
self.layer_norm_epsilon = config.rms_norm_eps
|
|
self.model_config = vllm_config.model_config
|
|
self.cache_config = vllm_config.cache_config
|
|
self.quant_config = vllm_config.quant_config
|
|
self.speculative_config = vllm_config.speculative_config
|
|
self.num_spec = (
|
|
self.speculative_config.num_speculative_tokens
|
|
if self.speculative_config
|
|
else 0
|
|
)
|
|
|
|
@property
|
|
def mamba_type(self) -> MambaAttentionBackendEnum:
|
|
return MambaAttentionBackendEnum.GDN_ATTN
|
|
|
|
def get_state_dtype(self) -> tuple[torch.dtype, ...]:
|
|
return MambaStateDtypeCalculator.gated_delta_net_state_dtype(
|
|
self.model_config.dtype,
|
|
self.cache_config.mamba_cache_dtype,
|
|
self.cache_config.mamba_ssm_cache_dtype,
|
|
)
|