**作者**:昇腾实战派

# 一、背景介绍

2026年,千问发布了Qwen3.5系列模型,模型结构与Qwen3-Next类似,出现了Gated Attention模块,该系列模型可通过vllm-ascend框架在昇腾平台上部署。本文基于transform和vllm-ascend框架中的代码分析,提取了Qwen3.5-27B的Profiling,将对Qwen3.5-27B模型代码展开讲解,详细介绍每个API的内容和作用。

vllm-ascend相关代码链接:[vllm-ascend/vllm\_ascend/ops/gdn.py at e14b89cf3021639ae2f4093a3052ad96c13e8b52 · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/e14b89cf3021639ae2f4093a3052ad96c13e8b52/vllm_ascend/ops/gdn.py#L253)

# 二、Qwen3.5-27B模型结构

## 2.1 Qwen3.5-27B模型结构介绍

Qwen3.5-27B模型结构如图1所示,主要包含几个模块:

**(1)3层Linear Attention(主要包含Gated DeltaNet和FNN模块)**

**(2)1层Full Attention(主要包含Gated Attention和FNN模型,其中Gated Attention与Qwen3系列模型Attention基本一致,仍然沿用GQA的结构)**

**(3)MTP**

![](./12734313-df87-4549-ae82-9291f82eb944.png)

图1 Qwen3.5-27B模型结构示意图

## 2.2 GDN简介

### 2.2.1 GDN背景介绍

通过注意力核线性化(linear attention)改造传统自注意力,将复杂度从O(n2)降至O(n);采用业界顶尖的**GDN(Gated Delta Rule)线性注意力架构,继承 Delta Net 原生动态 Delta 离散化规则,并吸收 Mamba 选择性 SSM 的全局遗忘与门控设计,在实现跨序列全局平滑衰减的同时,以单键值隐状态累积替换**机制完成序列建模,无需全局 token 两两交互。Linear Attention主要优势有如下3点:

* Linear Attention 将标准 Attention 的 softmax 替换为近似核函数,并改变 Q/K/V 的计算顺序(先算![](./6cf8fe1b-94f1-4ffe-8721-898ebc713ce5.png)再乘 Q),将复杂度从O(L^2d)降为O(Ld^2);
* 在自回归场景下,Linear Attention 可以转化为 RNN 递推形式(有状态 h、时序递归),每一步仅需维护一个​**固定尺寸的隐状态矩阵 h**​,无需保存历史全部 KV。
* 递推形式每步串行且无法利用 tensor core,因此采用**​ chunkwise 分块并行算法**兼顾效率。随着序列长度增加,Flash Attention 的O(L^2d)计算量导致 kernel 逐渐进入 compute-bound,而 Linear Attention 的O(Ld^2)复杂度优势开始显现。

### 2.2.2 GDN算法简介

Linear attention: 每个时间步 t 通过线性投影产生三个列向量: qt, vt, kt,则当前的状态St可以表示为:

![](./058be2cc-c236-4c93-bdf1-7acffd30edd5.png),其中![](./769b29ed-da4d-4e06-95ab-21cc04437cdd.png)

通过交换律可以推导到原来的公式![](./02339eed-9ce1-43a6-90f6-48317058734e.png):

![](./3557405e-4bf1-4037-a39b-6a94d0159446.png)

原始 Linear Attention 在模型效果上远弱于标准 Transformer,根本原因是​**固定大小的隐状态容量有限,导致”记忆碰撞”,且缺乏遗忘机制**​。为解决这个问题,有以下两种改进方向:

* ​**Mamba2/GLA**​:引入标量门控衰减 ![](./9fa209e8-bcbb-421e-9a49-fc2705041eff.png) ,可以快速擦除全局记忆,但不能选择性更新单个键值对。
* ​**DeltaNet**​:使用 delta 规则精确替换特定键值对,但缺乏全局记忆的快速清除机制。

**Gated Delta Networks (GDN)** (GDN结构如图2所示)将两者结合,在线性递推的基础上同时引入门控衰减![](./cea2d860-572f-484d-a154-e52d346676ab.png)和 delta 更新强度 ![](./34dd0f39-6190-4b24-8ef2-5508b3ea382d.png),兼具全局遗忘和定向更新能力,但需要用chunkwise算法实现。

![](./f0458176-62d9-46e4-b784-2165298b35c3.png)

图2 Gated DeltaNet网络结构示意图

### 2.2.3 GDN关键步骤总览

**主要流程参考如下:**

1. 输入投影(in_proj_qkv / in_proj_z / in_proj_ba)
   
   → 得到 QKV、z、b、a
2. 因果卷积(causal_conv1d)
   
   → 局部特征提取
3. 生成 g(alpha)、beta
   
   → 门控参数
4. Gated Delta Rule 核心计算(循环状态机)
   
   → 长序列建模
5. 门控归一化 + 输出投影
   
   → norm(z * out) → out_proj

# 三、Qwen3.5-27B模型结构代码走读

## 3.1 Qwen3.5-27B模型代码结构总览

![](./777d19ae-0bc6-4d62-b83e-7dca70c3c992.png)

## 3.2 GatedDeltaNetAttention分析

### 3.2.1 GatedDeltaNetAttention初始化

```
self.hidden_size      = 模型维度 (如 2048 / 4096)
self.num_k_heads      = Q/K 头数
self.num_v_heads      = V 头数 (通常是 K 的整数倍)
self.head_k_dim       = Q/K 每头维度
self.head_v_dim       = V 每头维度
self.key_dim          = num_k_heads * head_k_dim
self.value_dim        = num_v_heads * head_v_dim
self.conv_kernel_size = 因果卷积核大小
```

### 3.2.2 QKV、z、b、a输入映射

```
if hasattr(self, "in_proj_qkv"):
           # LoRA path (Qwen3.5 only): separate in_proj_qkv and in_proj_z
           mixed_qkv, _ = self.in_proj_qkv(hidden_states)
           ba, _ = self.in_proj_ba(hidden_states)
           z, _ = self.in_proj_z(hidden_states)
           z = z.reshape(z.size(0), -1, self.head_v_dim)
           b, a = ba.chunk(2, dim=-1)
           b = b.contiguous()
           a = a.contiguous()
       else:
           mixed_qkvz, _ = self.in_proj_qkvz(hidden_states)
           ba, _ = self.in_proj_ba(hidden_states)

           if self.gqa_interleaved_layout:
               # Qwen3-Next: unpack the interleaved GQA layout
               query, key, value, z, b, a = self.fix_query_key_value_ordering(
                   mixed_qkvz, ba
               )
               query, key, value = map(
                   lambda x: rearrange(x, "l p d -> l (p d)"), (query, key, value)
               )
               mixed_qkv = torch.cat((query, key, value), dim=-1)
           else:
               # Qwen3.5: weights are already in [q, k, v, z] and [b, a] order
               qkv_size = (self.key_dim * 2 + self.value_dim) // self.tp_size
               z_size = self.value_dim // self.tp_size
               mixed_qkv, z = mixed_qkvz.split([qkv_size, z_size], dim=-1)
               z = z.reshape(z.size(0), -1, self.head_v_dim)
               b, a = ba.chunk(2, dim=-1)
               b = b.contiguous()
               a = a.contiguous()
```

* ​**in_proj_qkv**​: 把输入映射成 QKV shape: `hidden_size → key_dim + key_dim + value_dim`

输入shape:[batch_size,seq_len, hidden_size//tp]

输出shape:[batch_size,seq_len, (key_dim+key_dim+value_dim)//tp]

* ​**in_proj_z**​: 门控信号 shape: `hidden_size → value_dim `

输入shape:[batch_size,seq_len, hidden_size//tp]

输出shape:[batch_size,seq_len, value_dim//tp]--- before reshape

输出shape:[batch_size,seq_len, num_v_heads//tp, head_v_dim]--- after reshape

* ​**in_proj_ba**​: 循环门控参数

输入shape:[batch_size,seq_len, hidden_size//tp]

输出shape:[batch_size,seq_len, 2*num_v_heads//tp]

后续每个输入会根据上面输出进行拆解:

q shape:[batch_size,seq_len, key_dim//tp]

k shape:[batch_size,seq_len, key_dim//tp]

v shape:[batch_size,seq_len, value_dim//tp]

b shape: [batch_size,seq_len, num_v_heads//tp]

a shape: [batch_size,seq_len, num_v_heads//tp]

### 3.2.3 gdn_attention_core计算

整个GatedDeltaNetAttention网络最核心的部分是调用了 gdn_attention_core :

```
core_attn_out = torch.zeros(
            (num_tokens, self.num_v_heads // self.tp_size, self.head_v_dim),
            dtype=hidden_states.dtype,
            device=hidden_states.device,
        )

torch.ops.vllm.gdn_attention_core(
            mixed_qkv,
            b,
            a,
            core_attn_out,
            _encode_layer_name(self.prefix),
        )
```

**关键步骤**

1. 申请一块全是0的输入内存 (vLLM 算子是**in-place 写入**模式,必须提前给空间。)
2. 跑完整的GatedDeltaNet计算,结果写入上面空内存

**【gdn_attention_core计算核心逻辑】**

![](./10102a14-e62b-4f20-b847-84c05c5084e4.png)

#### 3.2.3.1 causal_conv1d_update_npu

**(1)API功能介绍**

![](./0c705415-fc8b-4f3c-9310-003e1c6348ce.png)

**(2)API入参介绍**

decode阶段对应的causal_conv1d(备注:prefill阶段是调用到了Ascend C,固在此不展开描述)

* `KERNEL_WIDTH = w` 卷积核长度(qwen3.5固定width=4,卷积核大小较小,属于small kernel)
* `d`:通道维度 dim
* `t`:时间步 token 位置
* xt​[d]:第 t 个 token、第 d 通道输入
* w[d, k]:第 d 通道、卷积核第 k 个权重(代码 `weight: (dim, width)`)
* b[d]:偏置 bias
* yt​[d]:卷积原始输出
* ot​[d]:最终输出(代码里的 o)
* 因果规则:**只能用 t, t-1, t-2,... 历史,不能用未来 t+1**

**(3)API核心公式**

![](./ef77e859-8673-4b1d-9408-34ea13cf49a1.png)

拆开逐段对应代码:

1. 卷积求和(内核循环 j)

![](./b6dae29b-3ede-41de-8440-4cac69ef6ab0.png)

2. SiLU 激活(源码 `SILU_ACTIVATION`)

![](./e4850883-768c-472b-bce0-13dee1edfd78.png)

最终

![](./0bc5ad5b-19a3-4f2a-a722-5f78a9aeed73.png)

* ​**xₜ**​:当前 token(你有)
* ​**xₜ₋₁, xₜ₋₂ ...**​:过去的 token(​**必须存在 cache 里**​)

decode推理的时候是**一个 token 一个 token 生成**的,不可能每次都把整个历史序列重新传一遍。

所以:

**cache_t = 保存最近的w个历史 x,供下一次卷积直接用**

输出是mixed_qkv,后续再拆分成q,k,v作为gated delta rule的输入

公式如下:

```
query_spec, key_spec, value_spec = self.rearrange_mixed_qkv(mixed_qkv_spec)
query_non_spec, key_non_spec, value_non_spec = self.rearrange_mixed_qkv(mixed_qkv_non_spec)
```

#### 3.2.3.2 fused_gdn_gating_patch

(1)API 功能介绍

什么时候调用:

在prefill阶段和decode开启mtp阶段会用到

(2)API 核心公式

![](./a923276b-3056-4e22-93e5-1066437736db.png)

这里输出是g和![](./8dd7cc4a-7107-4d26-89ca-b0f23a302273.png),这里的![](./7ee1b32a-898f-4106-969c-917b3f3f799a.png)就是后面的![](./3c63ec93-e0e6-4d35-b456-9c1523e3761a.png)

关于g的操作:

* Softplus 保证输出​**恒正**​:softplus(x)≥0,β是温度系数,缩放时间步;`threshold`是溢出保护:当βx过大时直接线性截断,避免 ![](./9c04ae87-3f81-438c-b32c-5d46efadb8eb.png) 爆炸。
* 负号 `-` → 保证 `exp(g)` 永远 ​**< 1**​,对于后续GDN作用是让历史状态不断遗忘;
* `A_log` → 每个头固定衰减率

因为 beta_output 是​**输入门控**​,必须满足:

* 范围 **(0,1)**
* 对于后续GDN作用:**​控制当前 token 对状态 h 的贡献强度,​**也就是**控制**新信息注入强弱****

![](./6fb00c9e-9338-4054-924b-a0c9c137bcfd.png)

#### **3.2.2.3 chunk_gated_delta_rule**

**(1)为什么在prefill阶段要用chunk**

把长序列切成大小固定的块 BT=64

**长序列(尤其 Prefill 阶段)在 NPU/GPU 上**直接算会极慢、显存爆炸、无法并行,且NPU/GPU 不擅长逐 token 循环,对于Prefill 长序列,无法一次性载入算子

​**串行依赖**​:对于每个时间步骤St依赖St-1,序列方向完全串行,无法利用序列维度的并行性。

1. ​**块内**​:64 个 token 一次性并行算完内部递归
2. ​**块间**​:只传递块首尾的状态 S
   * 块开始时载入上一块的最终状态 S_prev​
   * 块内跑完所有 token 后输出本块最终状态 S_curr​
   * 传给下一个块当初始状态

**计算复杂度的下降**

对于chunk计算可以写成如下公式:

![](./9f539215-79b9-4e08-bc16-5a7e0753ed7a.png)

这里N代表新项,o代表旧项的门控衰减系数,均可看成常量

进行移项可以得到如下公式:

![](./90b5662f-58fa-4386-ba20-19411d12589f.png)

得到的因果下三角核如下

![](./0636117a-5b47-42b3-9c2b-b3b0394c43b1.png)

也就是对应的公式:

![](./9eddc73f-01f1-4201-87c2-16edfde37710.png)

S计算可以通过计算![](./a2d63b44-0b3f-41a6-9bec-0b7c2aadf66b.png)得到。由于这里的A是对应一个chunk大小的矩阵,通过等比求和公式,可以得到如下公式

![](./6afccbf9-cf43-4c50-8899-de2468cae966.png)

假设chunk大小是K,对于下三角矩阵,A^K=0,固对于k做截断只需要计算I+A+...+A^k-1即可,复杂度降到了O(K^2),总的计算复杂度从(如果序列长度是L)O(L^2)降到了O(LK)

**为了保持全文的连贯性与可读性,后续隐状态Ht均用St来表示**

**(2)输入shape介绍**

q shape:[batch_size,seq_len, num_k_heads//tp, head_k_dim]

k shape:[batch_size,seq_len, num_k_heads//tp, head_k_dim]

v shape:[batch_size,seq_len, num_v_heads//tp, head_v_dim]

beta shape: [batch_size,seq_len, num_v_heads//tp]

g shape: [batch_size,seq_len, num_v_heads//tp]

**(3)调用逻辑**

```
chunk_gated_delta_rule()  [入口函数:参数校验 + 格式转换]
   ↓
ChunkGatedDeltaRuleFunction.apply()  
   ↓
   forward()  [自动求导的前向传播]
      ↓ (可选)
      l2norm_fwd(q)  L2归一化
      l2norm_fwd(k)
      ↓
   chunk_gated_delta_rule_fwd()  [真正的核心前向计算]
      ↓
      ┌─────────────────────────────────────────┐
      │ 核心计算流水线(分块并行计算)           │
      │ 1. chunk_local_cumsum(g)       门控累积和 │
      │ 2. chunk_scaled_dot_kkt_fwd()  计算A矩阵 │
      │ 3. solve_tril()                下三角求解 │
      │ 4. recompute_w_u_fwd()         计算w, u  │
      │ 5. chunk_gated_delta_rule_fwd_h() 计算h  │
      │ 6. chunk_fwd_o()               计算最终输出o │
      └─────────────────────────────────────────┘
      ↓
返回 o, final_state
```

##### 3.2.2.3.1 l2norm_fwd

**(1)功能介绍**

这段代码中主要是通过triton实现L2 归一化

实现的公式可以写成如下

```
y = x / sqrt( sum(x²) + eps )
```

目标:

**让![](./778d7221-0276-4bbf-aa87-ad2d2dca383c.png),变成单位向量**

作用:

对于fwd_h里的公式(后续会提到)

![](./17f69108-f6e2-46ea-a865-76492bf26c15.png)

有大量的指数衰减项(exp)

* 如果k很大,每次更新的S就会很大,导致S指数膨胀
* 如果k太大就会导致chunk跨块传递S导致数值爆炸

**L2Norm 把 K 固定模长 = 1**

→ K 大小永远不变

→ S 只随遗忘 g、递归更新,**不会随向量长度膨胀**

→ Chunk 分块递归数值完全稳定

**(2)代码介绍**

```
q = l2norm_fwd(q)
k = l2norm_fwd(k)
```

![](./ced0ce2f-91c6-4fc2-a3ea-15f3ebbe899c.png)

##### 3.2.1.3.2 chunk_gated_delta_rule_fwd

![](./14727bee-9539-4ae0-8722-5a1257751ad5.png)

**(1)代码介绍**

```
def chunk_gated_delta_rule_fwd(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    g: torch.Tensor,
    beta: torch.Tensor,
    scale: float,
    initial_state: torch.Tensor,
    output_final_state: bool,
    cu_seqlens: torch.LongTensor | None = None,
    prebuilt_meta=None,
):
    ......
    g = chunk_local_cumsum(
        g,
        chunk_size=chunk_size,
        cu_seqlens=cu_seqlens,
        block_indices=block_indices_cumsum,
    )
    # obtain WY representation. u is actually the new v.
    A = chunk_scaled_dot_kkt_fwd(
        k=k,
        beta=beta,
        g_cumsum=g,
        cu_seqlens=cu_seqlens,
        chunk_indices=chunk_indices_chunk64,
        output_dtype=torch.float32,
    )
    A = solve_tril(
        A=A,
        cu_seqlens=cu_seqlens,
        chunk_indices_large_block=chunk_indices_large_block,
        chunk_indices_bt=chunk_indices_chunk64,
        output_dtype=k.dtype,
    )
    w, u = recompute_w_u_fwd(
        k=k,
        v=v,
        beta=beta,
        A=A,
        g_cumsum=g,
        cu_seqlens=cu_seqlens,
        chunk_indices=chunk_indices_chunk64,
    )
    h, v_new, final_state = chunk_gated_delta_rule_fwd_h(
        k=k,
        w=w,
        u=u,
        g=g,
        initial_state=initial_state,
        output_final_state=output_final_state,
        cu_seqlens=cu_seqlens,
        chunk_indices=chunk_indices_chunk64,
        chunk_offsets=chunk_offsets_chunk64,
    )

    ......
    o = chunk_fwd_o(
        q=q,
        k=k,
        v=v_new,
        h=h,
        g=g,
        scale=scale,
        cu_seqlens=cu_seqlens,
        chunk_offsets=chunk_offsets_chunk64,
    )

   ......
```

**(2)入参介绍**

这里入参就是上面经过因果卷积,ab转换成g和beta还有l2 norm之后的输出

q shape:[batch_size,seq_len, key_dim//tp]

k shape:[batch_size,seq_len, key_dim//tp]

v shape:[batch_size,seq_len, value_dim//tp]

beta shape: [batch_size,seq_len, num_v_heads//tp]

g shape: [batch_size,seq_len, num_v_heads//tp]

**(3)API调用逻辑**

###### 1)准备工作

```
forward_context = get_forward_context()
num_decodes = 0
chunk_size = 64  # 固定分块大小:64 token 一块,根据seq_len切分
```

###### 2) chunk_local_cumsum

```
g = chunk_local_cumsum(g, chunk_size=64, ...)
```

对每个 chunk 内的 g 做累计和,用于快速计算历史状态衰减

这里的 g 必须做 cumsum,是为了在下面公式中(后面的fwd_h)让chunk内的递归公式变成​**可并行计算**​!

![](./bc488000-e74c-4e24-a251-09e9ea942daf.png)

更详细地:

原本的GDN状态更新是:

![](./20b8f8ab-9ee1-4b24-a1db-1ac21a58ad55.png)

其中

![](./acc0196f-592f-49b8-990c-1e720ff973b0.png)

ht必须依赖ht-1,必须串行

现在对于g做累计和:

![](./1b9c4bb0-9b60-4012-9e5b-c397e1a20a0e.png)

代入原公式可以得到:

![](./83c9edb8-a150-4379-8708-776cb60daee9.png)

**展开后,St 不再依赖前一个 S_(t-1)了!**

只依赖:

* 初始状态 S0
* 累积和 G_t、G_i

输入输出shape不变:g shape: [batch_size,seq_len, num_v_heads//tp]

###### 3) chunk_scaled_dot_kkt_fwd

```
A = chunk_scaled_dot_kkt_fwd(k, beta, g_cumsum=g, ...)
```

* **chunk_scaled_dot_kkt_fwd_kernel 做一件事:**
  在每个 chunk 内,计算带门控、带缩放、带因果掩码的 k·k.T 矩阵
  公式:
  A[i,j] = β[i] * exp(g[i]-g[j]) * k[i]·k[j].T
  且 i > j 才有效

A shape:[batch_size,seq_len, num_v_heads//tp, chunk_size]

###### 4) solve_tril

三角求解 (I + A)⁻¹

```
A = solve_tril(A, ...)
```

​**目的**​:让后面 state 递推变成 **O (1) 计算**

###### 5)recompute_w_u_fwd

重新计算 w, u(新 K, V)

```
w, u = recompute_w_u_fwd(k, v, beta, A, g_cumsum=g, ...)
```

* ​**输入**​:`k, v, beta, A`
* ​**输出**​:
  * `w`:新 key
  * `u`:新 value
* **u = A · (v · β)** → 新 value
* **w = A **​ ·** (k · β · exp(g))** → 新 key
* ​**目的**​:让记忆 state 更新更快、更稳定

v 原始 shape = [batch_size,seq_len, num_v_heads//tp, head_v_dim]  beta 进 kernel 前被转置成 [num_v_heads//tp,batch_size,seq_len] = [16, 1, 8192]

对于每个batch_size和每个head_num的chunk的v和k进行recompute

```
┌─────────────┐     ┌─────────────┐
│  A [64,64]  │     │  V [64,128] │
└──────┬──────┘     └──────┬──────┘
       │                   │
       │             beta [64] → [64,1]
       │                   │
       │                   ▼
       │           V * beta = [64,128]
       │                   │
       └─────────┬─────────┘
                 │
                 ▼
          U = A @ (V*beta)
              [64,128]
```

chunk A和chunk V和chunk beta做运算得到chunk u

最终按照batch size和head num放到最终u不同的位置上

| 数据            | shape                        | 含义                                        |
| ----------------- | ------------------------------ | --------------------------------------------- |
| 单个 chunk 输出 | **[64, 128]**          | 一个 chunk + 一个 head 的结果               |
| 最终完整 U      | **[1, 8192, 48, 128]** | 全部 batch、全部 token、全部 head、全部维度 |

对于k和beta和g的num_k_head和num_v_head不一致的情况,这里用了如下方法

```
ptr_k = k + (bos * Hg + i_h // (H // Hg)) * K + offs_t_2d * (Hg * K) + offs_k * 1
```

举例:

num_k_heads = 48

num_v_heads = 16

这里head的映射关系如下(每 3 个 V head 共享 1 个 K head)

```
i_h (V head) | 对应的 K head
0            → 0 //3 = 0
1            → 1 //3 = 0
2            → 2 //3 = 0
3            → 3 //3 = 1
4            → 4 //3 = 1
5            → 5 //3 = 1
...
45,46,47     → 15
```

类似的,w输出的shape是:[batch_size,seq_len, num_v_heads//tp, head_v_dim]

###### 6)chunk_gated_delta_rule_fwd_h

这里是在计算隐状态S,是一个串行运行各个chunk的函数,每个后面的St都和前面的St-1相关,chunk内是并行操作

最重要的步骤

```
h, v_new, final_state = chunk_gated_delta_rule_fwd_h(
        k=k,
        w=w,
        u=u,
        g=g,
        initial_state=initial_state,
        output_final_state=output_final_state,
        cu_seqlens=cu_seqlens,
        chunk_indices=chunk_indices_chunk64,
        chunk_offsets=chunk_offsets_chunk64,
    )
```

**输入**

* `k`:key
* `w`:优化后的 key
* `u`:优化后的 value
* `g`:门控累计和(控制更新强度)
* `initial_state`:初始记忆 S₀(上文已讲,后续h全部用s代替来增加可读性,本质不变),可以是none:这时就会跳过内部如下命令 如果有 shape=[real_batch_size, num_v_heads//tp, head_k_dim, head_v_dim],后续每次都有前一次chunk做完的隐状态作为这次的initial state

```
if USE_INITIAL_STATE:
    h0_ptr = h0 + i_nh * K * V
    ptr_h0_bv1 = h0_ptr + offs_k * V + offs_v1 * 1
    b_h1_bv1 += tl.load(ptr_h0_bv1, mask=mask_kv1, other=0.0).to(tl.float32)

    ptr_h0_bv2 = h0_ptr + offs_k * V + offs_v2 * 1
    b_h1_bv2 += tl.load(ptr_h0_bv2, mask=mask_kv2, other=0.0).to(tl.float32)
```

**输出**

* **S** (每个 chunk 的状态)

shape: [batch_size,chunk_num, num_v_heads//tp, head_k_dim, head_v_dim]

**chunk_num = 序列被切成多少个 chunk**

* **St (final_state)** (最终状态)

shape: [real_batch_size, num_v_heads//tp, head_k_dim, head_v_dim]

真正在运行的代码是这一段

```
b_v_new1 = b_v1 - tl.dot(b_w, b_h1_bv1)
b_v_new1 = b_v_new1 * b_g
b_h1_bv1 = b_h1_bv1 * b_g_last
b_h1_bv1 += tl.dot(b_k, b_v_new1)
```

固,对于第t个chunk写成数学公式为:

![](./924daab7-be7e-48e2-9052-d3917396ccdf.png)

* St-1:第 t−1 块处理完毕后,整条序列到目前为止所有历史信息压缩状态
* St:第 t 块处理完之后,新的全局累积状态
* Kt:当前分块的键
* Vt:当前分块的值,这个v是前一步骤转换后的u
* Wt:可以当成全局固定化权重,是前一步k转换后的w
* ​增量(Vt-W*Ht-1)​,这就是 “delta rule” 名字来源
* gt​:第t个chunk的门控衰减系数(指数衰减,控制历史信息遗忘)
* e^gt​:对历史状态Ht-1做指数衰减 (e^gt*Ht-1:遗忘旧记忆)
* e^(gL−gt)​:对当前增量做衰减
* gL:当前 chunk 最后一个 token 的门控,也就是chunk内原始g的累加和

更准确地公式可以如下

![](./76802b18-8833-41a8-912a-cb70f00c5299.png)

###### 7) chunk_fwd_o

算最终输出 o。这个函数对于所有chunk可以并行运行,根据如下公式可以看到,每个chunk直接选取H对应的位置即可

公式如下:

![](./dd5e1b3d-bf29-4d0e-8c55-76eb62ff6e8b.png)

输出由两部分组成:​**inter-chunk**​(当前 query 与历史累积状态的交互)和 ​**intra-chunk**​(当前 chunk 内部 query 和 key 的注意力交互)。

**​inter-chunk:​**当前 chunk 的 query 与前面所有 chunk 的累积隐状态的交互。

对应代码:

```
b_o += tl.dot(b_q, b_h)
b_o = b_o * tl.exp(b_g)[:, None]
```

**​intra-chunk:​**当前 chunk 内部注意力

对应代码:

```
b_A += tl.dot(b_q, b_k)
b_A = b_A * safe_exp(b_g[:, None] - b_g[None, :])
b_A = tl.where(m_A, b_A, 0)
tl.dot(b_A, b_v)
```

![](./f0e8c88f-0eac-49f9-b8e1-838ba7eb201f.png)

mask:只保留下三角,上三角全部置 0,禁止看到未来 token

```
# 门控衰减矩阵 Gamma[i,j] = exp(g[i] - g[j]),g 是 log-space 累积和
# 举例:g_cumsum = [-0.1, -0.3, -0.5, -0.7]
#
# Gamma =
#   j=0        j=1        j=2        j=3
# i=0 [exp(0)     0          0          0    ]     <- 对角线无衰减
# i=1 [exp(-0.2)  exp(0)     0          0    ]     <- token 0 对 token 1 衰减 exp(-0.2)
# i=2 [exp(-0.4)  exp(-0.2)  exp(0)     0    ]
# i=3 [exp(-0.6)  exp(-0.4)  exp(-0.2)  exp(0)]    <- 越远的过去衰减越多
```

为了更清楚地解释上面3-5步骤对于6-7步骤的逻辑,这里引入WY表示:

这里通过deltanet的公式进行解释(gated deltanet是一样的)deltanet比较容易理解:

![](./3288d665-3df5-4c98-af21-9ff9c3c3374c.png)

该公式可以用简单数学归纳法证明。首先定义![](./1723d31e-0fbf-4301-ac8a-f89dcd82a34d.png)n表示第n个token。当n=1时,公式固然成立。

假设对于n-1也成立,那么我们证明对于n也成立,证明过程如下:

![](./12189e30-1cae-4aa3-8d32-f4396d4665dc.png)

该证明不仅证明了公式的正确性,也提供了w_n的计算公式!

通过Sn的递推公式,我们可以证明![](./47289d9c-fab9-4e50-b005-f064c8e5543b.png)通过归纳法:

![](./c8ffb33b-8e6c-466b-a7e9-b3650258f99a.png)

那么,对于在一个chunk i内部第r个位置,我们通过递归到S[i]会有如下公式:

![](./8b17045d-eb00-4c3e-8175-20e1f3f77a4a.png)

w和u在这里都是通过WY表示计算的,但是从每一个块的第一个位置开始,不是从序列起始位置开始,从而第r个位置有如下w和u:

![](./bbd27e13-23f7-4506-86b9-5b22cc1941dc.png)

对于输出计算

![](./b165dc11-a141-4631-b27f-29a653f8e813.png)

结合矩阵乘法形式,可以得到最终公式:

![](./ab407278-129a-4840-8bad-a83b3a387eaa.png)

![](./53028b87-468d-4eae-8413-19649c629bbc.png)

#### 3.2.1.4 recurrent gated delta rule

开启投机推理会运行该算子,替代原先的fused_sigmoid_gating_delta_rule

Recurrent 模式是直接按时间步逐个计算递推过程,适用于推理解码场景(逐 token 生成)。

和prefill阶段的区别是:把chunk_fwd_h和chunk_fwd_o都合并在一步里

公式如下:是对于逐个token递归

![](./c3205da4-e66f-41c8-bdac-a3cd2eed09f3.png)

![](./6ce0a7d9-c8e5-4db5-8f4a-2932ecb58723.png)

这个和原生的gated delta rule是等价的

![](./f66b4318-b915-4b1b-9276-e413687632a7.png)

门控衰减:![](./032543e4-2db9-4213-a005-87b3a0e63ec1.png)对应![](./2c2b4cb7-9eb0-4144-9ba4-532dc780fb4e.png),保证了 delta 项中擦除的是衰减后的旧值。

代码实现如下

```
b_h *= exp(b_g)                                        # h = alpha * h
b_v = b_beta * (b_v - tl.sum(b_h * b_k[:, None], 0))   # v_new = beta * (v - alpha*S @ k)
b_h += b_k[:, None] * b_v                              # h += k * v_new^T
```

另外,在 kernel 中也做了GVA 的支持,Q/K 每时间步前进 num_k_heads*head_k_dim//tp,V/O 每时间步前进 num_v_heads*head_v_dim//tp。g 的步长也是 num_v_heads//tp或者是num_v_heads*head_k_dim//tp取决于是否是KDA(KDA 模式:​**一个头,K 个 g(每个 key 维度一个)**​),beta 的步长取决于是否为 headwise 模式。

```
p_q += H * K
p_k += H * K
p_o += HV * V
p_v += HV * V
if not IS_KDA:
    p_g += HV
else:
    p_gk += HV * K
p_beta += HV * (V if IS_BETA_HEADWISE else 1)
```

#### 3.2.1.5 fused sigmoid gating delta rule

和fused recurrent gated delta rule的区别是输入sigmoid输入是a,b recurrent输入是g和beta

公式和fused recurrent gated delta rule一样,只是在运行这段代码前又做了a,b到g,beta的转换,可参考fused_gdn_gating_patch

## 3.3 Qwen3.5-27B Profiling算子分析

| **profiling算子名称**                            | **对应代码调用位置**                                                                                                                                                                                                                                                                                                                                                                                                                                                                       | **代码路径**                                                                                                                                                                                                               |
| -------------------------------------------------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
| _causal_conv1d_update_kernel_npu_tiled_5        | _causal_conv1d_update_kernel_npu_tiled[grid]](x,weight,bias,conv_state,conv_state_indices,num_accepted_tokens,query_start_loc,block_idx_last_scheduled_token,initial_state_idx,out,batch,dim,seqlen,...)                                                                                                                                                                                                                                                                       | [vllm-ascend/vllm\_ascend/ops/triton/mamba/causal\_conv1d.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/mamba/causal_conv1d.py#L527)            |
| l2norm_fwd_kernel2_loop                             | q = l2norm_fwd(q)k = l2norm_fwd(k)                                                                                                                                                                                                                                                                                                                                                                                                                                                             | [vllm-ascend/vllm\_ascend/ops/triton/fla/l2norm.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/fla/l2norm.py#L34)                                |
| chunk_local_cumsum_scalar_kernel                   | chunk_local_cumsum(g,chunk_size=chunk_size,cu_seqlens=cu_seqlens,block_indices=block_indices_cumsum,)                                                                                                                                                                                                                                                                                                                                                                                   | [vllm-ascend/vllm\_ascend/ops/triton/fla/cumsum.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/fla/cumsum.py#L116)                               |
| chunk_scaled_dot_kkt_fwd_kernel                   | chunk_scaled_dot_kkt_fwd(k=k,beta=beta,g_cumsum=g,cu_seqlens=cu_seqlens,chunk_indices=chunk_indices_chunk64,output_dtype=torch.float32,)                                                                                                                                                                                                                                                                                                                                              | [vllm-ascend/vllm\_ascend/ops/triton/fla/chunk\_scaled\_dot\_kkt.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/fla/chunk_scaled_dot_kkt.py#L83) |
| solve_tril_16x16_kernel                             | A = solve_tril(A=A,cu_seqlens=cu_seqlens,chunk_indices_large_block=chunk_indices_large_block,chunk_indices_bt=chunk_indices_chunk64,output_dtype=k.dtype,)                                                                                                                                                                                                                                                                                                                         | [vllm-ascend/vllm\_ascend/ops/triton/fla/solve\_tril.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/fla/solve_tril.py#L330)                      |
| recompute_w_u_fwd_kernel                           | w, u = recompute_w_u_fwd(k=k,v=v,beta=beta,A=A,g_cumsum=g,cu_seqlens=cu_seqlens,chunk_indices=chunk_indices_chunk64,)                                                                                                                                                                                                                                                                                                                                                                   | [vllm-ascend/vllm\_ascend/ops/triton/fla/wy\_fast.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/fla/wy_fast.py#L98)                             |
| chunk_gated_delta_rule_fwd_kernel_h_blockdim64  | h, v_new, final_state = chunk_gated_delta_rule_fwd_h(k=k,w=w,u=u,g=g,initial_state=initial_state,output_final_state=output_final_state,cu_seqlens=cu_seqlens,chunk_indices=chunk_indices_chunk64,chunk_offsets=chunk_offsets_chunk64,)                                                                                                                                                                                                                                      | [vllm-ascend/vllm\_ascend/ops/triton/fla/chunk\_delta\_h.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/fla/chunk_delta_h.py#L179)               |
| chunk_fwd_kernel_o                                  | o = chunk_fwd_o(q=q,k=k,v=v_new,h=h,g=g,scale=scale,cu_seqlens=cu_seqlens,chunk_offsets=chunk_offsets_chunk64,)                                                                                                                                                                                                                                                                                                                                                                          | [vllm-ascend/vllm\_ascend/ops/triton/fla/chunk\_o.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/fla/chunk_o.py#L112)                            |
| fused_recurrent_gated_delta_rule_fwd_kernel_11  | core_attn_out_non_spec, last_recurrent_state = fused_recurrent_gated_delta_rule(q=query_non_spec,k=key_non_spec,v=value_non_spec,g=g_non_spec,beta=beta_non_spec,initial_state=ssm_state,inplace_final_state=True,cu_seqlens=non_spec_query_start_loc[: attn_metadata.num_decodes + 1],ssm_state_indices=non_spec_state_indices_tensor,use_qk_l2norm_in_kernel=True,)                                                                                   | [vllm/vllm/model\_executor/layers/fla/ops/fused\_recurrent.py at main · vllm-project/vllm (github.com)](https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/fla/ops/fused_recurrent.py#L481)                 |
| fused_sigmoid_gating_delta_rule_update_kernel_0 | core_attn_out_non_spec = fused_sigmoid_gating_delta_rule_update(A_log=self.A_log.contiguous(),dt_bias=self.dt_bias.contiguous(),q=query_non_spec.contiguous(),k=key_non_spec.contiguous(),v=value_non_spec.contiguous(),a=a.contiguous(),b=b.contiguous(),initial_state_source=ssm_state,initial_state_indices=non_spec_state_indices_tensor,cu_seqlens=non_spec_query_start_loc,use_qk_l2norm_in_kernel=True,softplus_beta=1.0,softplus_threshold=20.0,) | [vllm-ascend/vllm\_ascend/ops/triton/fla/sigmoid\_gating.py at main · vllm-project/vllm-ascend (github.com)](https://github.com/vllm-project/vllm-ascend/blob/main/vllm_ascend/ops/triton/fla/sigmoid_gating.py#L180)              |

# 四、总结

本文总结了 Qwen3.5-27B 模型结构,并针对 GDN 网络结构做了完整拆解与原理分析。

GDN 模块预填充(Prefill)阶段最核心算子为chunk_gated_delta_rule_fwd_kernel_h和chunk_fwd_kernel_o;自回归解码(Decode)阶段,开启 MTP 投机解码时核心算子为fused_recurrent_gated_delta_rule,关闭 MTP 常规单 Token 解码时核心算子为fused_sigmoid_gating_delta_rule。

在 Prefill 分块 Chunk 计算流程中,矩阵 A 构建、门控 g/beta 计算、Q/K 归一化、初始状态加载等全部前置步骤,均是为后续**fwd_h 隐状态计算**与**fwd_o 输出计算**两个核心步骤提供输入支撑。其中chunk_gated_delta_rule_fwd_kernel_h负责完成 GDN 内部时序隐状态的分块并行求解,chunk_fwd_kernel_o基于求解完成的隐状态进一步计算注意力最终输出。

同时本文梳理了投机解码模式下的 GDN 执行逻辑:投机草稿 Token 分支与已验证正常 Token 分支分开计算、分开执行因果卷积与时序递归,仅在输出端按照位置索引合并结果,既保证投机解码加速能力,又保证模型因果正确性。

整体 GDN 架构采用**Causal Conv1D 局部时序建模 + 线性注意力全局长依赖**的互补设计,通过 Chunk 并行预填充、Recurrent 串行解码实现线性复杂度 O (N) 超长上下文推理能力,相比传统 Transformer 注意力具备显著效率优势。


 

更多推荐