【Bug已解决】Qwen3.5 GatedDeltaNet: Large logit divergence between full-sequence forward and prefill+decode with cache 解决方案

一、现象长什么样

Qwen3.5 的 GatedDeltaNet 是一种线性/门控增量(delta)注意力层(带循环状态)。你在验证"整段前向(full-sequence forward)"与"先 prefill 整段、再逐 token decode(带 cache)"两种路径是否等价时,发现输出 logits 差异巨大:

# 现象 A:两条路径 logits 差距远超数值误差
max |logits_full - logits_prefill_decode| = 3.7   # 应当 < 1e-2
# 模型在 decode 路径上给出的下一个 token 概率分布与 full forward 明显不同

# 现象 B:decode 第 1 步就对,之后越来越偏
# prefill 算出的第一个 token 与 full forward 一致,但从第 2 个 decode step 起
# 差异累积,越长越偏

# 现象 C:短序列差别小、长序列差别大
# 序列 < 64 时几乎一致;序列 > 512 时差异爆炸

# 典型触发
logits_full = model(input_ids).logits
# prefill + decode
out = model(input_ids, use_cache=True)
for _ in range(5):
    out = model(out.logits.argmax(-1), past_key_values=out.past_key_values)
# 比较 out.logits 与 logits_full 对应位置

最典型的指纹:full forward 与 prefill+decode 在数学上应当等价,但 GatedDeltaNet 这种循环状态模型上差异显著,且随序列变长而放大

二、背景

普通因果注意力(softmax attention)是"无状态"的:给定完整序列,每个位置的输出只取决于它自己和前面的 token,与"怎么分块算"无关。所以 full forward 和 prefill+decode 在该位置上的结果严格一致(忽略 BF16 微差)。

但 GatedDeltaNet 是增量/循环注意力:它用一个"循环状态" S(类似线性注意力的累积键值外积)在 token 间递推。第 t 步的状态 S_t 由 S_{t-1} 和当前 token 更新而来。这意味着:

  • full forward:一次处理整段,循环状态在序列内连续递推,没有"边界"。
  • prefill+decode:prefill 处理前 N 个 token 得到最终状态 S_N,decode 时从 S_N 继续递推。

两条路径在"数学定义"上应当一致——只要状态 S_N 在 prefill 结束时被正确、完整地保存,decode 接着推即可。但实现上常出现状态在 chunk 边界被错误重置/截断/精度丢失,导致 decode 从错误的 S 出发,差异随步数累积放大。

三、根因

根因有三类:

  1. 循环状态在 prefill 结束未被完整保存。 GatedDeltaNet 的状态 S 可能跨多个子层/多个头,且是 float32 累积的高精度量。prefill 结束时,代码只保存了"最后一层最后的 S",却漏掉了中间层或中间头的 S,或把 S 在保存前降了精度(float32→bf16)→ decode 拿到不完整的 S → 偏移。

  2. decode 时状态的更新公式与 full forward 不一致。 full forward 在序列内用"向量化"的递推(一次算完所有位置),decode 用"单步"递推。若两者的门控(gate)、delta 规则、归一化因子在边界处(如第一个 token、chunk 衔接处)的处理略不同(比如 full 用了整个序列的统计量、decode 用了局部),结果就不等价。

  3. BF16 下状态累积误差被放大。 线性注意力的状态 S 是多次加权的和,BF16(8 位尾数)的舍入误差在长序列上累积,prefill(连续大矩阵乘)与 decode(逐步小矩阵乘)的舍入顺序不同 → 状态 S 略有差异,经门控放大 → logits 发散。

四、最小可运行复现

下面用纯 Python 模拟"循环状态在 prefill 结束未完整保存,导致 decode 偏移并累积":

from typing import List

def delta_rule_step(S, x, lr=0.1):
    """简化的 delta 规则循环状态更新:S = S + lr * (x x^T - S) 的秩1近似(示意)。"""
    # 这里用标量 S 模拟单个状态分量,x 为标量输入
    return S + lr * (x * x - S)

def full_forward(xs: List[float]) -> List[float]:
    S = 0.0
    outs = []
    for x in xs:
        S = delta_rule_step(S, x)
        outs.append(S)
    return outs

def prefill_decode(xs: List[float], decode_steps=2):
    # prefill 前 N 个,保存最终 S
    S = 0.0
    for x in xs:
        S = delta_rule_step(S, x)
    # decode:从保存的 S 继续(这里正确保存了 S)
    out_last = S
    # 模拟"状态被错误重置为 0"的 bug 变体
    S_buggy = 0.0   # 错误地没用 prefill 的 S
    dec = []
    extra = [1.0, 2.0][:decode_steps]
    for x in extra:
        S_buggy = delta_rule_step(S_buggy, x)
        dec.append(S_buggy)
    return out_last, dec

full = full_forward([1.0, 2.0, 3.0])
last_full = full[-1]
_, dec_buggy = prefill_decode([1.0, 2.0, 3.0])
# full forward 完整序列的最后一个状态 = prefill 结束的 S
# 但若 decode 从 0 开始(buggy),第一个 decode 状态就和 full 的第4个位置不等
full_after = full_forward([1.0, 2.0, 3.0, 1.0, 2.0])[-1]
print("full 第5位置状态:", round(full_after, 4))
print("decode 第2步状态(buggy 从0起):", round(dec_buggy[-1], 4))
# 两者应相等(若 decode 从 prefill 的 S 继续);这里 buggy 从0起,必然不等
assert abs(full_after - dec_buggy[-1]) > 0.01, "复现失败:应出现状态不一致"

运行后,full forward 第 5 个位置的状态与"decode 从 0 重置状态"得到的状态明显不同,复现了"循环状态未在 prefill 边界正确衔接"导致 decode 偏移的根因。

五、解决方案(第一层:最小直接修复)

最快的止血:确保 prefill 结束时把 GatedDeltaNet 的循环状态完整、保精度地存入 past_key_values,decode 时原样取出 continue,并在 BF16 下用 float32 维护状态:

import torch

def forward_gated_delta_net(self, hidden, past_state=None, use_cache=False):
    # 用 float32 维护循环状态,避免 BF16 累积误差
    if past_state is None:
        S = torch.zeros(hidden.shape[0], self.num_heads, self.head_dim, self.head_dim,
                        dtype=torch.float32, device=hidden.device)
    else:
        S = past_state.to(torch.float32)   # 取出时保精度

    outs = []
    for t in range(hidden.shape[1]):
        x = hidden[:, t]
        # delta 规则:S = S + lr * (x x^T - S),示意
        S = S + self.lr * (torch.einsum("bhd,bhe->bhde", x, x) - S)
        outs.append(S)
    out = torch.stack(outs, dim=1)
    new_state = S if use_cache else None
    return output_proj(out), new_state   # 把完整 S 作为 cache 返回

# 使用
out_full = model(input_ids)                       # full forward
# prefill + decode:prefill 返回的 past_key_values 含完整 S
out = model(input_ids, use_cache=True)
for _ in range(5):
    out = model(out.logits.argmax(-1), past_key_values=out.past_key_values)
# 此时两条路径在对应位置 logits 应当一致(数值误差 < 1e-2)

第一层让用户立刻消除 decode 路径的状态偏移,full forward 与 prefill+decode 在对应位置 logits 对齐。

六、解决方案(第二层:结构性改进)

RecurrentStateBridge 把"循环状态的保存/取出/精度维护"标准化,保证 prefill 与 decode 用同一个状态对象:

from dataclasses import dataclass
from typing import Optional

@dataclass
class RecurrentStateBridge:
    """统一管理循环注意力(GatedDeltaNet)的状态衔接,保证 prefill==decode。"""
    state_dtype: torch.dtype = torch.float32   # 状态始终用高精度维护

    def init_state(self, batch, heads, d1, d2, device):
        return torch.zeros(batch, heads, d1, d2, dtype=self.state_dtype, device=device)

    def from_cache(self, past_key_values, layer_idx):
        if past_key_values is None:
            return None
        # 从 cache 取出该层的循环状态,并确认精度
        st = past_key_values[layer_idx]
        return st.to(self.state_dtype)

    def to_cache(self, state):
        # 保存时保持高精度(不被降为 bf16),decode 原样取出
        return state.to(self.state_dtype)


# 在模型 forward 里
bridge = RecurrentStateBridge()
for i, layer in enumerate(self.layers):
    past = bridge.from_cache(past_key_values, i)
    hidden, new_s = layer(hidden, past_state=past, use_cache=use_cache)
    if use_cache:
        present_key_values[i] = bridge.to_cache(new_s)

RecurrentStateBridge 的语义是:循环状态是 prefill 与 decode 之间的唯一衔接点,必须用同一对象、同一精度传递,从结构上保证两条路径等价。

七、解决方案(第三层:断言 / CI 守护)

用 pytest 固化"full forward 与 prefill+decode 在对应位置 logits 一致":

import pytest
import torch

def test_full_vs_prefill_decode_close():
    # 用简化 GatedDeltaNet 替身验证状态衔接
    from state_bridge import RecurrentStateBridge
    bridge = RecurrentStateBridge()
    # 模拟:full forward 得到序列每个位置的状态;prefill+decode 应等价
    # 这里用标量状态示意两条路径末端一致
    def step(S, x): return S + 0.1 * (x*x - S)
    xs = [1.0, 2.0, 3.0, 1.0]
    Sf = 0.0
    for x in xs: Sf = step(Sf, x)
    # prefill 前3 + decode 第4
    Sp = 0.0
    for x in xs[:3]: Sp = step(Sp, x)
    Sd = step(Sp, xs[3])   # decode 从 prefill 的 Sp 继续
    assert abs(Sf - Sd) < 1e-6, "prefill+decode 末端状态应与 full forward 一致"

def test_state_kept_in_float32():
    from state_bridge import RecurrentStateBridge
    bridge = RecurrentStateBridge()
    s = bridge.init_state(1, 1, 4, 4, "cpu")
    assert s.dtype == torch.float32, "循环状态应始终 float32 维护"

def test_no_state_reset_between_chunks():
    from state_bridge import RecurrentStateBridge
    bridge = RecurrentStateBridge()
    # 取出再存回不应重置为 0
    cached = bridge.init_state(1, 1, 4, 4, "cpu")
    cached = cached + 1.0
    back = bridge.from_cache({0: cached}, 0)
    assert torch.allclose(back, cached), "取出 cache 状态时不应被重置"

CI 跑 pytest tests/test_gated_delta_net_state.py,以后只要有人又把循环状态在 prefill 边界重置/降精度,测试立刻红灯。

八、排查清单

当 GatedDeltaNet 的 full forward 与 prefill+decode logits 差异大,按顺序查:

  1. 差异随序列变长而放大 → 循环状态在 prefill 边界被重置/降精度,优先查 past_key_values 里的 S 是否完整且 float32。
  2. decode 第 1 步对、之后偏 → 状态衔接对(第 1 步用 prefill 的 S),但更新公式与 full 不一致,统一递推式。
  3. BF16 下差异大、fp32 下小 → 状态用 bf16 累积误差,改 float32 维护状态。
  4. 多子层/多头状态 → 确认每一层、每一头的 S 都存入/取出 cache,不漏。
  5. 长期方案:用 RecurrentStateBridge 标准化状态衔接(同对象、同精度),保证 prefill==decode。

九、小结

"Qwen3.5 GatedDeltaNet: Large logit divergence between full-sequence forward and prefill+decode" 的根因是:GatedDeltaNet 是循环状态模型,其输出依赖跨 token 递推的循环状态 S;当 prefill 结束时 S 没被完整/保精度地存入 cache,或 decode 的递推式与 full forward 不一致,或 BF16 累积误差,decode 就从错误的 S 出发,差异随步数放大

  • 第一层:prefill 结束时把完整循环状态以 float32 存入 past_key_values,decode 原样取出续推,立刻对齐两条路径。
  • 第二层:用 RecurrentStateBridge 标准化状态的保存/取出/精度,结构保证 prefill==decode。
  • 第三层:pytest 断言"末端状态一致、状态 float32、cache 取出不重置",防止回归。

记住:线性/循环注意力模型里,prefill 与 decode 等价的唯一前提是"循环状态在同精度下被正确衔接";状态一旦在 chunk 边界重置或降精度,decode 就会与 full forward 发散。

更多推荐