Diffusion LLM 为什么能并行生成:从逐 Token 解码到多位置去噪实战

大语言模型生成一句话时,我们已经习惯看到文字从左到右逐个出现。这个交互形式并不只是界面动画,而是主流自回归模型的概率分解方式:只有第一个新 Token 确定以后,第二个 Token 的条件才完整;第二个确定以后,第三个才能继续。模型可以把一段输入并行编码,却不能在严格保持这套生成分布的前提下,一次独立确定所有未来 Token。

Diffusion Language Model,简称 Diffusion LM 或 dLLM,试图改写这个前提。以离散掩码扩散为例,生成可以从一串 [MASK] 开始。模型在一次前向计算里同时预测许多位置,再保留置信度较高的部分,把不确定的位置留到下一轮,甚至重新遮住已经填过但相互冲突的 Token。于是,句子的多个区域可能一起成形,而不是永远服从从左到右的唯一顺序。

这很容易被包装成一句过强的宣传:"并行生成,所以一定比自回归模型快。"事实远没有这么简单。一次预测多个位置,不等于一次就能提交所有位置;双向注意力的每轮成本、去噪步数、输出长度处理、缓存复用、采样器设计和硬件利用率,都会决定真实时延。某个论文设置中的步数减少,也不能直接换算成线上接口的固定倍数。

本文不预设 Diffusion LLM 会取代所有自回归模型,而是回答六个更实际的问题:它与传统逐 Token 解码究竟差在哪里;前向加噪和反向去噪怎样落到离散文本;置信度提交与重掩码为何有效;并行性为什么不保证加速;怎样与 Speculative Decoding 公平比较;以及如何构建一套不会被演示样例误导的评测与上线流程。

1. 自回归模型为什么天然按顺序生成

给定 Token 序列 (x_1,x_2,\ldots,x_L),自回归模型把联合概率写成条件概率的乘积:

p(x_{1:L})=\\prod_{i=1}\^{L}p(x_i\\mid x_{\

训练时答案是已知的,可以使用因果注意力掩码并行计算每个位置的下一个 Token 损失,这叫 Teacher Forcing。推理时答案未知,位置 (i) 的输入必须包含刚刚采样出的 (x_{i-1}),所以生成依赖链不能凭空消失。批量处理多个请求能够提高 GPU 利用率,却没有让单条序列内部的未来 Token 彼此独立。

自回归路线并非只有缺点。它的最大工程优势之一是 KV Cache:前缀中已经计算过的 Key、Value 可以缓存,下一步主要处理一个新 Token,而不必重复完成整段前缀的所有投影。连续批处理、PagedAttention、量化和融合算子也围绕这种工作负载积累了成熟生态。线上服务看到的"逐个生成",背后其实有大量经过优化的缓存与调度。

它的结构性限制同样明确。单请求输出 (L) 个 Token,至少要经历约 (L) 次有前后依赖的决策。每一步计算规模可能不大,却需要一次模型执行、采样和调度。低并发时,GPU 可能没有被充分利用;长输出的 Token 间时延会不断累积;模型较早产生的错误还会成为后续位置的条件,使整个后缀沿错误方向继续展开。

Diffusion LLM 并不是简单删除因果掩码,然后让普通聊天模型同时猜完答案。它改变了训练目标和生成过程:模型必须学会在不同噪声强度下,根据仍然可见的上下文恢复缺失 Token。只有经过这种训练,任意顺序补全、多位置同时更新和反复修订才成为模型本身的能力。

2. 文本扩散和图像扩散不是同一套噪声

图像像素或潜变量通常是连续值,可以逐步添加高斯噪声,再学习预测噪声、速度或干净样本。Token 是词表中的离散编号,对编号直接加一个高斯小数没有语言学意义:"数据库"的编号加 0.1 不会变成语义相近的"存储"。文本扩散因此出现了多条路线。

早期 Diffusion-LM 将文本映射到连续词向量,在连续空间中加噪和去噪,再映射回离散词。D3PM 则为离散状态设计转移矩阵,可以把 Token 随机替换成其他 Token,也可以转移到一个吸收态。掩码扩散采用的 [MASK] 就是一种直观吸收态:随着时间增大,越来越多原始 Token 被替换为掩码。

设 (x_0) 是干净文本,(t\in0,1) 表示噪声强度,(\alpha(t)) 表示保留原 Token 的概率。一种简化前向过程可以写成:

q(x_t\^i\\mid x_0\^i)= \\begin{cases} x_0\^i, \& \\text{概率 }\\alpha(t)\\ \[MASK\], \& \\text{概率 }1-\\alpha(t) \\end{cases}

当 (t) 接近 0,序列大部分可见;当 (t) 接近 1,序列接近全掩码。训练时随机采样噪声等级,构造不同破坏程度的输入,让 Transformer 预测被遮住位置的原始 Token。模型既会处理只缺几个词的简单修复,也会处理线索极少的高噪声恢复。

真正的论文目标通常会包含由噪声日程导出的权重、似然下界或连续时间形式,并不等同于"随机做一次 BERT 掩码训练"。MDLM 对掩码扩散目标进行了简化和 Rao-Blackwellization;SEDD 从离散分布比率与 Score Entropy 出发;LLaDA 则展示了在预训练和指令微调规模上构建掩码扩散大模型的路线。它们共享迭代去噪的家族特征,但损失、参数化和采样器不能随意互换。

3. 从全掩码到答案:反向去噪如何工作

推理可以从提示词加一段固定长度的 [MASK] 开始。提示词属于条件区域,不被修改;答案区域属于可编辑区域。每一轮把当前序列和噪声等级送入模型,得到所有可编辑位置上的词表分布,然后按照采样器决定本轮提交哪些 Token、哪些位置继续保持掩码。

最朴素的方法是每轮均匀解开一部分位置。更常见的思路是使用模型置信度:先为所有候选位置取最大概率或采样概率,再优先提交高置信位置。高置信 Token 会成为下一轮的上下文,帮助模型判断其他位置。MaskGIT 在离散图像 Token 生成中系统展示了这种并行预测与置信度调度思路,后来的掩码语言扩散采样也广泛讨论类似机制。
#mermaid-svg-C16vzOzANEAZc1Z4{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-C16vzOzANEAZc1Z4 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-C16vzOzANEAZc1Z4 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-C16vzOzANEAZc1Z4 .error-icon{fill:#552222;}#mermaid-svg-C16vzOzANEAZc1Z4 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-C16vzOzANEAZc1Z4 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-C16vzOzANEAZc1Z4 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-C16vzOzANEAZc1Z4 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-C16vzOzANEAZc1Z4 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-C16vzOzANEAZc1Z4 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-C16vzOzANEAZc1Z4 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-C16vzOzANEAZc1Z4 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-C16vzOzANEAZc1Z4 .marker.cross{stroke:#333333;}#mermaid-svg-C16vzOzANEAZc1Z4 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-C16vzOzANEAZc1Z4 p{margin:0;}#mermaid-svg-C16vzOzANEAZc1Z4 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-C16vzOzANEAZc1Z4 .cluster-label text{fill:#333;}#mermaid-svg-C16vzOzANEAZc1Z4 .cluster-label span{color:#333;}#mermaid-svg-C16vzOzANEAZc1Z4 .cluster-label span p{background-color:transparent;}#mermaid-svg-C16vzOzANEAZc1Z4 .label text,#mermaid-svg-C16vzOzANEAZc1Z4 span{fill:#333;color:#333;}#mermaid-svg-C16vzOzANEAZc1Z4 .node rect,#mermaid-svg-C16vzOzANEAZc1Z4 .node circle,#mermaid-svg-C16vzOzANEAZc1Z4 .node ellipse,#mermaid-svg-C16vzOzANEAZc1Z4 .node polygon,#mermaid-svg-C16vzOzANEAZc1Z4 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-C16vzOzANEAZc1Z4 .rough-node .label text,#mermaid-svg-C16vzOzANEAZc1Z4 .node .label text,#mermaid-svg-C16vzOzANEAZc1Z4 .image-shape .label,#mermaid-svg-C16vzOzANEAZc1Z4 .icon-shape .label{text-anchor:middle;}#mermaid-svg-C16vzOzANEAZc1Z4 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-C16vzOzANEAZc1Z4 .rough-node .label,#mermaid-svg-C16vzOzANEAZc1Z4 .node .label,#mermaid-svg-C16vzOzANEAZc1Z4 .image-shape .label,#mermaid-svg-C16vzOzANEAZc1Z4 .icon-shape .label{text-align:center;}#mermaid-svg-C16vzOzANEAZc1Z4 .node.clickable{cursor:pointer;}#mermaid-svg-C16vzOzANEAZc1Z4 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-C16vzOzANEAZc1Z4 .arrowheadPath{fill:#333333;}#mermaid-svg-C16vzOzANEAZc1Z4 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-C16vzOzANEAZc1Z4 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-C16vzOzANEAZc1Z4 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-C16vzOzANEAZc1Z4 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-C16vzOzANEAZc1Z4 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-C16vzOzANEAZc1Z4 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-C16vzOzANEAZc1Z4 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-C16vzOzANEAZc1Z4 .cluster text{fill:#333;}#mermaid-svg-C16vzOzANEAZc1Z4 .cluster span{color:#333;}#mermaid-svg-C16vzOzANEAZc1Z4 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-C16vzOzANEAZc1Z4 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-C16vzOzANEAZc1Z4 rect.text{fill:none;stroke-width:0;}#mermaid-svg-C16vzOzANEAZc1Z4 .icon-shape,#mermaid-svg-C16vzOzANEAZc1Z4 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-C16vzOzANEAZc1Z4 .icon-shape p,#mermaid-svg-C16vzOzANEAZc1Z4 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-C16vzOzANEAZc1Z4 .icon-shape rect,#mermaid-svg-C16vzOzANEAZc1Z4 .image-shape rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-C16vzOzANEAZc1Z4 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-C16vzOzANEAZc1Z4 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-C16vzOzANEAZc1Z4 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 有
无
提示词 + 全部 MASK
模型并行预测所有答案位置
计算每个候选的置信度
提交本轮高置信位置
低置信位置继续或重新 MASK
是否还有未决位置
长度检查与最终答案

假设答案区域有 64 个位置,采样器计划运行 16 轮。它不一定每轮固定提交 4 个 Token,也不一定从左到右选择。第一轮可能先出现标点、关键词或结构词;之后主语和结论在不同区域同时稳定。调度函数可以是线性、余弦或模型专用日程,核心是在"更早获得上下文"和"不要过早锁死错误"之间折中。

这里必须区分两种实现。单调解掩码一旦提交某个位置,就不再改变它,实现简单且便于追踪;可重掩码解码会重新评估已填位置,把低置信或与新上下文冲突的位置再次遮住。后者更符合"全局修改"的直觉,却需要额外策略防止 Token 反复震荡。某个模型训练时采用什么转移过程、官方推理代码采用什么采样器,应以该模型论文和固定版本实现为准。

4. 一个教学级离散 Mask 扩散训练器

下面用 PyTorch 写一个最小模型,目的不是训练可用聊天模型,而是把数据流说明白:随机选择噪声强度,只掩盖允许预测的位置,让双向 Transformer 根据其余 Token 恢复原词。代码使用线性掩码率,并只计算被掩盖位置的交叉熵。

python 复制代码
# toy_mask_diffusion.py
from __future__ import annotations

import torch
from torch import Tensor, nn
from torch.nn import functional as F


PAD_ID = 0
MASK_ID = 1
BOS_ID = 2
EOS_ID = 3


class TinyMaskDiffusion(nn.Module):
    def __init__(
        self,
        vocab_size: int,
        max_length: int = 64,
        hidden_size: int = 128,
        layers: int = 3,
        steps: int = 16,
    ) -> None:
        super().__init__()
        self.steps = steps
        self.token_embedding = nn.Embedding(vocab_size, hidden_size)
        self.position_embedding = nn.Embedding(max_length, hidden_size)
        self.time_embedding = nn.Embedding(steps + 1, hidden_size)
        block = nn.TransformerEncoderLayer(
            d_model=hidden_size,
            nhead=4,
            dim_feedforward=hidden_size * 4,
            dropout=0.0,
            batch_first=True,
            norm_first=True,
        )
        self.encoder = nn.TransformerEncoder(block, num_layers=layers)
        self.output = nn.Linear(hidden_size, vocab_size)

    def forward(self, tokens: Tensor, time_step: Tensor) -> Tensor:
        batch, length = tokens.shape
        positions = torch.arange(length, device=tokens.device)
        hidden = self.token_embedding(tokens)
        hidden = hidden + self.position_embedding(positions)[None, :, :]
        hidden = hidden + self.time_embedding(time_step)[:, None, :]
        padding_mask = tokens.eq(PAD_ID)
        hidden = self.encoder(hidden, src_key_padding_mask=padding_mask)
        return self.output(hidden)


def corrupt(tokens: Tensor, steps: int) -> tuple[Tensor, Tensor, Tensor]:
    """线性掩码日程:t 越大,被 MASK 的概率越高。"""
    batch = tokens.size(0)
    time_step = torch.randint(1, steps + 1, (batch,), device=tokens.device)
    mask_probability = time_step.float() / steps
    editable = tokens.ne(PAD_ID) & tokens.ne(BOS_ID) & tokens.ne(EOS_ID)
    random_values = torch.rand(tokens.shape, device=tokens.device)
    selected = editable & (random_values < mask_probability[:, None])

    # 极短样本也至少制造一个监督位置,避免整批标签都是 -100。
    for row in range(batch):
        if not selected[row].any() and editable[row].any():
            first = editable[row].nonzero(as_tuple=False)[0, 0]
            selected[row, first] = True

    noisy = tokens.clone()
    noisy[selected] = MASK_ID
    labels = tokens.masked_fill(~selected, -100)
    return noisy, labels, time_step


def train_step(model: TinyMaskDiffusion, batch: Tensor, optimizer: torch.optim.Optimizer) -> float:
    model.train()
    noisy, labels, time_step = corrupt(batch, model.steps)
    logits = model(noisy, time_step)
    loss = F.cross_entropy(logits.flatten(0, 1), labels.flatten())
    optimizer.zero_grad(set_to_none=True)
    loss.backward()
    optimizer.step()
    return float(loss.detach())


def demo() -> None:
    torch.manual_seed(7)
    model = TinyMaskDiffusion(vocab_size=32)
    optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
    batch = torch.tensor(
        [
            [BOS_ID, 4, 5, 6, EOS_ID, PAD_ID],
            [BOS_ID, 7, 8, 9, 10, EOS_ID],
        ]
    )
    loss = train_step(model, batch, optimizer)
    assert loss > 0.0
    print(f"one teaching step, loss={loss:.4f}")


if __name__ == "__main__":
    demo()

这段代码刻意省略了多项生产要素:没有论文级连续时间损失权重,没有大规模数据管线,没有混合精度、分布式训练、FlashAttention、检查点恢复和指令微调,也没有处理词表中真实 [MASK] 的配置差异。它只能验证"按噪声强度破坏,再预测干净 Token"的最小闭环,不能据此推导大型模型的质量或速度。

另一个容易忽略的差别是注意力方向。代码中的 TransformerEncoder 使用双向注意力,被遮住位置能查看左右两侧已经可见的 Token;自回归模型使用因果掩码,只能查看左侧前缀。直接拿一个普通因果语言模型套上 [MASK] 输入,通常不具备这里需要的训练分布和注意力结构。

5. 置信度提交与重掩码解码

训练完成后还需要采样器。下面的教学实现把提示区域锁定,把答案区域视为可编辑集合。每轮都并行预测这些位置,根据余弦日程增加已提交数量;排名落后的候选重新变成 [MASK]。因此某个位置前一轮被保留,下一轮也可能因为新上下文改变而被修订。

python 复制代码
# toy_parallel_decode.py
from __future__ import annotations

import math

import torch
from torch import Tensor

from toy_mask_diffusion import MASK_ID, TinyMaskDiffusion


@torch.inference_mode()
def parallel_decode(
    model: TinyMaskDiffusion,
    prompt: Tensor,
    answer_length: int,
    steps: int = 16,
    temperature: float = 0.0,
) -> Tensor:
    if prompt.ndim != 1:
        raise ValueError("prompt must be a one-dimensional token tensor")
    if answer_length <= 0 or steps <= 0:
        raise ValueError("answer_length and steps must be positive")

    device = next(model.parameters()).device
    prompt = prompt.to(device)
    sequence = torch.cat(
        [prompt, torch.full((answer_length,), MASK_ID, device=device)]
    ).unsqueeze(0)
    editable = torch.zeros_like(sequence, dtype=torch.bool)
    editable[:, prompt.numel() :] = True

    model.eval()
    for index in range(steps):
        time_value = max(1, model.steps - index * model.steps // steps)
        time_step = torch.tensor([time_value], device=device)
        logits = model(sequence, time_step)
        probabilities = logits.softmax(dim=-1)

        if temperature > 0.0:
            scaled = (logits / temperature).softmax(dim=-1)
            candidate = torch.multinomial(
                scaled[0, editable[0]], num_samples=1
            ).squeeze(-1)
            proposal = sequence.clone()
            proposal[0, editable[0]] = candidate
        else:
            proposal = logits.argmax(dim=-1)

        confidence = probabilities.gather(-1, proposal.unsqueeze(-1)).squeeze(-1)
        progress = 1.0 - math.cos((index + 1) / steps * math.pi / 2.0)
        keep_count = max(1, math.ceil(answer_length * progress))
        editable_confidence = confidence.masked_fill(~editable, float("-inf"))
        keep = editable_confidence.topk(keep_count, dim=1).indices

        next_sequence = sequence.masked_fill(editable, MASK_ID)
        next_sequence.scatter_(1, keep, proposal.gather(1, keep))
        sequence = next_sequence

    # 最后一轮理论上已保留全部位置;断言用于抓住调度实现错误。
    assert not sequence[editable].eq(MASK_ID).any()
    return sequence.squeeze(0)


def demo() -> None:
    torch.manual_seed(11)
    model = TinyMaskDiffusion(vocab_size=32, max_length=32)
    result = parallel_decode(
        model=model,
        prompt=torch.tensor([2, 4, 5]),
        answer_length=6,
        steps=8,
    )
    assert result.shape == (9,)
    print(result.tolist())


if __name__ == "__main__":
    demo()

未训练模型只会输出随机结果,这个 demo 验证的是张量形状、日程终止和多位置更新,不验证语言质量。生产采样还必须屏蔽非法特殊 Token、处理温度和随机种子、支持批次内不同长度、定义 EOS 之后的 Padding,并与官方模型的噪声日程完全对齐。

置信度也不是事实正确率。Softmax 最大概率高,只说明模型当前分布尖锐,模型可能对错误答案同样自信。跨轮比较置信度时,噪声等级改变了分布;不同位置的词频和词表竞争也不相同。因此"每轮留下概率最高的若干位置"是一种解码启发式,不应被表述为数学上保证留下正确 Token。

重掩码能修复一部分早期冲突,也可能制造震荡。例如模型在"使用 PostgreSQL"与"使用 MySQL"之间来回切换,相关配置项和 SQL 方言会一起变化。可以限制单个位置的最大修改次数、对已稳定 Token 添加保留偏置,或者使用块级解码减少全局波动。但每种稳定化策略都可能牺牲纠错能力,需要由验证集决定,而不是只看演示动画是否"像在思考"。

6. 并行生成为什么不等于固定加速

理解性能时,首先区分"单轮并行预测多少位置"和"生成整段文本需要多少模型计算"。自回归模型一次通常提交一个 Token,但能够复用前缀 KV Cache;掩码扩散模型一次能评估多个位置,却可能需要对完整答案区域运行双向注意力,而且要重复多轮。

可以用一个不带具体数字的近似式思考总时延:

T_{diff}\\approx S\\times(T_{full\\ forward}+T_{rank}+T_{sync})

其中 (S) 是去噪步数,full forward 是当前长度上的模型执行,rank 是置信度排序与重掩码,sync 是调度和设备同步。自回归时延则近似为:

T_{ar}\\approx L\\times(T_{cached\\ step}+T_{sample}+T_{schedule})

如果扩散模型用远少于 (L) 的去噪轮数完成同等质量,且每轮全序列计算在硬件上高效,它有机会改善单请求生成时延。反之,若为了质量需要很多轮、答案很短、双向全序列计算昂贵,理论并行度不会自动变成墙钟时间收益。

硬件也决定结论。大矩阵计算通常比大量很小的逐 Token 内核更容易吃满 GPU,扩散解码可能提高单请求算术强度;但显存带宽、注意力长度、批次大小、张量并行通信和置信度排序都可能成为新瓶颈。GPU 上的优势不能直接外推到 CPU、移动端或不同加速器。

服务并发会进一步改变结果。低并发自回归解码可能利用不足,而高并发连续批处理已经能把许多请求的单 Token 步骤拼成大批次。扩散模型的请求处于不同去噪轮次、具有不同答案长度时,也面临动态批处理和尾部拖延。只在并发 1 上展示单条样例,不能证明在线吞吐更高。

因此不要写"并行生成带来固定 N 倍加速",除非 N 来自明确模型、硬件、精度、输入输出分布、采样器和并发下的可复现实测。论文报告的数据只说明其特定设置,不能替代自己的基线。本文不给出一个装饰性的倍数,正是因为不存在脱离条件仍然成立的倍数。

7. 长度不是免费变量

自回归模型生成 EOS 后自然停止,长度由逐步条件概率决定。全掩码扩散常常需要先分配答案槽位:给 32 个 [MASK] 还是 512 个,会直接影响计算量和输出结构。预估太短会截断内容,预估太长则可能产生冗余、重复或大量需要清理的特殊位置。

一种做法是训练长度预测器,先预测答案长度再去噪;另一种做法是在词表中生成 EOS,并规定 EOS 后全部转成 Padding;还可以使用块级半自回归,在需要时追加一块掩码区域。MDLM 论文讨论了以半自回归方式生成任意长度文本,这也说明"大段一次性全并行"并不是唯一方案。

块级方法把输出分成若干区块:块内并行去噪,块间按顺序扩展。它牺牲部分全局并行性,换来可控显存、流式输出和更自然的长度增长。对聊天产品而言,首屏响应和持续输出体验可能比整段完成时间更重要;等待整个 512 Token 区域去噪完毕再一次返回,即使总耗时下降,也可能让用户觉得系统卡住。

长度策略还会影响质量评测。固定参考答案长 80 Token,模型生成 160 Token 时,后半段重复会拉低匹配指标;粗暴截到 80 又掩盖停止能力。应同时报告内容质量、长度偏差、EOS 合法率、截断率和重复率,并把长度预测模块的时间计算进端到端时延。

8. KV Cache 与服务化限制

自回归 KV Cache 成立的关键是旧前缀表示在增加一个末尾 Token 后不需要重算。双向掩码扩散中,某个位置从 [MASK] 变成实词,会影响其他位置对它的注意力;下一轮多个位置都可能变化,旧状态不能像固定因果前缀那样完整复用。研究和工程实现可以探索局部缓存、块级缓存或只更新变化区域,但不能默认获得与自回归相同的缓存收益。

可观测性也要改变。除了 TTFT、Token 间时延和吞吐,还应记录每轮掩码数、提交数、重掩码数、Token 修改次数、平均置信度、答案长度预测误差,以及每次前向和排序耗时。只保留最终文本会丢失最有价值的故障现场:究竟是模型一直不确定,还是采样器过早锁定了错误。

接口层不要假装二者完全等价。自回归流式接口可以把每个新 Token 视作不可撤销增量;可重掩码扩散的中间结果可能改变。如果把草稿直接展示给用户,就需要支持文本替换事件,而不只是追加事件。更稳妥的初期方案是按稳定块输出,或只在整段完成后返回,并明确测量用户感知延迟。

9. Diffusion LLM 与 Speculative Decoding 的本质区别

两者都希望减少生成过程中的串行等待,也都可能一次处理多个候选 Token,但概率目标完全不同。Speculative Decoding 通常保留一个自回归目标模型,用小草稿模型或其他提议器生成候选,再由目标模型批量验证。严格的推测采样通过接受与修正步骤保持目标模型分布,目标模型仍然定义最终答案。

Diffusion LLM 则从训练阶段就学习加噪与去噪。它不是给某个现成自回归目标模型打补丁,也没有"草稿错了就由目标 AR 模型无损纠正"的默认保证。它生成的是自身扩散模型和采样器定义的分布,改变步数、提交日程或重掩码规则可能改变输出质量。

两者的缓存和成本结构也不同。推测解码仍可利用目标模型的因果 KV Cache,但要承担草稿计算、候选验证和低接受率浪费;扩散解码用双向多位置预测换取更少轮次,却要承担全序列迭代、长度管理和缓存困难。前者的关键指标包括接受率和每轮平均接受长度,后者则需要关注去噪步数、每轮提交比例、重掩码率和最终质量。

选择时先看约束。如果必须复用现有自回归模型权重、保持其目标分布,并且有兼容草稿器,优先测试 Speculative Decoding。如果任务天然包含任意位置补全、全局约束或反复修订,而且能够接受专门模型与推理栈,Diffusion LLM 值得实验。二者甚至可以在研究系统中组合,但在没有单独基线前叠加优化只会让性能归因更困难。

对比维度 自回归模型 Speculative Decoding Masked Diffusion LLM
训练目标 下一个 Token 目标模型仍为下一个 Token 多噪声等级下恢复被掩码 Token
生成顺序 左到右 草稿多步、目标批量验证 多位置迭代去噪
是否需新模型 基线模型 常需草稿器或专用头 需要扩散训练的模型
分布保证 模型自身分布 严格算法可保持目标分布 由扩散模型与采样器定义
主要风险 串行步数 接受率低、额外草稿成本 去噪轮数、全序列计算、长度与缓存

10. 如何做一次不自欺的评测

评测必须把质量和性能绑定起来。把去噪步数从 64 降到 4,速度大概率改变,但如果答案准确率、代码可执行率或指令遵循显著下降,就不能把它记作无条件加速。应画出"质量---时延"曲线,而不是只挑一个最快配置与自回归最高质量配置比较。

首先固定可比条件:同一硬件、精度、最大显存预算、请求集、提示模板和并发档位;记录模型参数量、上下文限制和服务版本。不同架构很难做到逐项完全相同,因此报告中要保留差异,而不是用"都是 7B"掩盖训练数据和模型能力差距。

请求集至少按输入长度、目标输出长度、任务类型和约束类型分桶。短问答、长摘要、代码生成、JSON 输出、文本填空和双向补全对两类模型的友好程度不同。扩散模型擅长填空,不代表开放式长生成同样占优;代码 benchmark 得分也不能替代中文客服质量。

性能至少报告首个稳定块时延、完整答案时延、P50、P95、请求吞吐、有效 Token 吞吐、峰值显存和每请求计算轮数。若中间 Token 会重写,"首 Token 时间"需要重新定义:是第一次出现任何候选,还是首次出现不会再修改的稳定内容。定义不清的 TTFT 没有可比性。

质量指标依任务选择。知识问答要检查事实与引用,代码要运行测试,结构化输出要验证 Schema,摘要要检测遗漏和幻觉。通用困惑度对不同建模方式未必可以直接横比,最好增加任务成功率、人工盲评和成对偏好,并给出置信区间或至少多次随机种子的离散程度。

下面的脚本读取两个系统的 JSONL 结果,按任务桶计算完整时延分位数、长度误差、完全匹配率和 JSON 合法率。它不自动宣布胜者,因为上线阈值应由业务定义。每行应包含 system、group、latency_s、prediction、reference 和 target_tokens。

python 复制代码
# summarize_diffusion_eval.py
from __future__ import annotations

import argparse
import json
from collections import defaultdict
from pathlib import Path
from statistics import mean
from typing import Iterable


def percentile(values: Iterable[float], q: float) -> float:
    ordered = sorted(values)
    if not ordered:
        raise ValueError("values must not be empty")
    position = (len(ordered) - 1) * q
    lower = int(position)
    upper = min(lower + 1, len(ordered) - 1)
    weight = position - lower
    return ordered[lower] * (1.0 - weight) + ordered[upper] * weight


def is_valid_json(text: str) -> bool:
    try:
        json.loads(text)
        return True
    except json.JSONDecodeError:
        return False


def load_rows(path: Path) -> list[dict]:
    rows = []
    for line_number, line in enumerate(path.read_text(encoding="utf-8").splitlines(), 1):
        if not line.strip():
            continue
        row = json.loads(line)
        required = {"system", "group", "latency_s", "prediction", "reference", "target_tokens"}
        missing = required.difference(row)
        if missing:
            raise ValueError(f"line {line_number} missing fields: {sorted(missing)}")
        rows.append(row)
    return rows


def summarize(rows: list[dict]) -> None:
    buckets: dict[tuple[str, str], list[dict]] = defaultdict(list)
    for row in rows:
        buckets[(row["system"], row["group"])].append(row)

    for (system, group), items in sorted(buckets.items()):
        latencies = [float(item["latency_s"]) for item in items]
        exact = [item["prediction"].strip() == item["reference"].strip() for item in items]
        length_errors = [
            abs(len(item["prediction"].split()) - int(item["target_tokens"]))
            for item in items
        ]
        json_valid = [is_valid_json(item["prediction"]) for item in items]
        print(
            json.dumps(
                {
                    "system": system,
                    "group": group,
                    "samples": len(items),
                    "latency_p50_s": percentile(latencies, 0.50),
                    "latency_p95_s": percentile(latencies, 0.95),
                    "exact_match": mean(exact),
                    "mean_length_error": mean(length_errors),
                    "json_valid_rate": mean(json_valid),
                },
                ensure_ascii=False,
            )
        )


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("results", type=Path)
    args = parser.parse_args()
    summarize(load_rows(args.results))


if __name__ == "__main__":
    main()

完全匹配只适合答案唯一的任务,开放式写作不能用它代表整体质量;脚本用空格近似 Token 数,也不能替代真实 tokenizer。生产评测应把服务实际返回的 Token 数写入结果,并对 JSON 任务和普通文本任务分别统计。示例的价值是固定原始字段和分桶方式,不是制造一个综合分数掩盖取舍。

11. 常见失败模式与排查顺序

11.1 早期高置信错误被锁死

模型可能先确定一个错误实体,其他位置随后围绕它补全,最终句子语法流畅却事实错误。排查时观察首轮被提交的位置和后续修改轨迹,而不是只看最终概率。可以增加重掩码机会、放慢早期提交日程,或在外部加入检索与事实验证,但不能把"语句整体更协调"误认为事实更可靠。

11.2 Token 在多轮之间震荡

低置信位置反复切换,去噪轮数增加但答案没有稳定。先确认噪声时间步是否与训练一致,再检查温度、随机种子和重掩码比例。盲目增加轮数可能只会重复震荡;需要统计每个位置的修改次数,并在达到上限时触发保守提交或重新生成。

11.3 固定长度导致截断和填充废话

答案区域太短时,代码括号、JSON 结构或论证结尾被切断;区域太长时,模型为了填满槽位生成重复句。应将长度预测误差单独监控,为结构化任务预留闭合空间,并明确 EOS 后处理。不要通过事后裁剪隐藏模型没有正确停止的问题。

11.4 并行候选局部合理、全局冲突

多个位置在同一轮依据相同残缺上下文预测,单独看都合理,合在一起却重复主语、单位不一致或变量未定义。下一轮双向上下文理论上能修复冲突,但单调提交策略可能已经锁死。对代码和结构化输出,应在每轮或最终阶段运行解析器、类型检查与约束验证,失败时保留现场而不是静默修补。

11.5 低步数看起来快,任务成功率下降

只测延迟很容易把不完整、格式错误的输出当成性能胜利。所有步数实验必须绑定质量阈值;若低步数未达到阈值,它只是不同质量档位,不是同质量加速。可以为不同任务设置自适应步数,但自适应控制器本身也要计入延迟并接受回归测试。

11.6 训练噪声与推理日程错位

模型训练时很少见到接近全掩码输入,推理却从全掩码开始;或者训练采用某个连续时间权重,服务端错误映射成整数步。结果常表现为第一轮概率异常尖锐或长期接近均匀。部署时要固定模型 revision、采样器代码和噪声配置,把这些信息写进输出 Trace。

11.7 把普通 BERT 当生成式 Diffusion LLM

BERT 能恢复掩码不代表它已经学习完整生成分布和迭代采样。训练噪声范围、损失权重、长度处理和指令对齐都不同。用 BERT 做演示可以帮助理解 Mask Prediction,但不能据此宣称获得了可替代聊天模型的 Diffusion LLM。

11.8 在不支持的硬件上照搬论文结论

不同 GPU、驱动、编译器和注意力实现对全序列前向的利用率差异很大。论文中的批次、精度和算子可能无法复现。先用性能分析器拆出模型前向、排序、采样、通信和数据搬运时间,再判断瓶颈;不要仅凭网络前后总耗时猜测算法原因。

12. 从实验到上线的最小闭环

第一步是选一个有明确官方推理代码和许可证的模型,固定权重 revision、依赖镜像和采样配置。不要先重写服务框架,也不要同时加入量化、编译、张量并行和自定义采样器。先让官方最小示例在一张卡上生成可重复结果,并保存完整配置。

第二步构建当前业务可用的自回归基线。固定提示词、输出上限、精度和硬件,按真实流量分桶;若训练数据和能力不同,就分别展示任务质量,不把架构差异包装成纯推理差异。

第三步扫描去噪步数、提交日程和答案长度策略。每一个配置都同时记录质量、完整时延、稳定首块时延、显存和失败率。先删掉未达到质量门槛的配置,再在剩余配置中比较性能。这样可以避免"最快配置"其实只是少做工作。

第四步增加并发和长上下文。观察批次能否有效合并、尾延迟是否抬升、显存是否随答案槽位快速增长。若目标是离线批处理,重点看总样本吞吐和单位任务成本;若目标是交互聊天,重点看稳定首块、P95 和流式体验。两个目标不应共用一句"更快"。

第五步做失败注入。故意测试错误长度预测、极长提示、空提示、全特殊字符、JSON Schema、代码括号、重复文本和设备内存压力。确认请求超时、取消和 OOM 不会留下失控任务;采样器发生异常时返回明确错误,不能用半成品文本冒充成功。

最后才是小流量灰度。日志中关联请求 ID、模型 revision、去噪步数、长度预测、每轮掩码数和最终校验结果。设置质量失败率、P95、显存和队列长度告警,并保留快速回退到原自回归服务的路由。新架构的价值必须来自长期业务指标,而不是一次漂亮的去噪动画。

13. 哪些场景更值得优先试验

Diffusion LLM 的任意顺序生成适合文本填空、代码局部修复、结构中多个字段协同补全,以及需要全局反复修改的任务。已知左右两侧上下文时,自回归模型需要特殊填充策略,而掩码模型天然可以把中间区域作为未知变量。固定长度、格式明确的任务也更容易控制答案槽位。

相反,输出很短、现有自回归栈高度优化、必须逐 Token 实时流式、或必须保持某个目标模型精确分布的场景,不应仅为追逐热点迁移。超长开放式生成还要认真处理答案长度和全序列计算成本。团队如果没有能力重新评测模型质量、维护专用服务与采样器,采用成熟自回归方案通常更稳妥。

最终判断不是"Diffusion 与 Autoregressive 谁赢",而是哪个系统在指定任务、硬件和服务目标下形成更好的质量---延迟---成本边界。架构名称不会替你完成这次测量。

14. 总结

Diffusion LLM 的关键变化不是把同一个自回归循环简单并行化,而是重新定义生成:训练阶段对离散 Token 加掩码噪声,模型学习在不同噪声等级下恢复干净序列;推理阶段从全掩码或部分掩码开始,多位置预测、按置信度提交,并在需要时重掩码修订。

这种机制带来了任意顺序补全和较少串行轮次的可能,却没有免费消除计算。每轮全序列双向注意力、去噪步数、长度预测、KV Cache 复用困难、动态批处理和硬件算子效率,都可能吞掉并行收益。并行是一种执行机会,不是固定加速承诺。

它也不同于 Speculative Decoding:后者以自回归目标模型为裁判,严格算法可以保持目标分布;前者是独立训练的生成模型,质量由扩散目标与采样器共同决定。评测时必须在相同质量门槛下比较完整时延、吞吐、显存和失败率,并保留中间去噪轨迹解释问题。

最合理的工程态度是把 Diffusion LLM 当作一条有明确优势、也有明确约束的新路线。先用教学实现理解机制,再用官方模型做固定版本实验,最后让真实请求和业务阈值决定是否上线。它可能改变语言模型的生成方式,但目前没有证据允许我们跳过验证,直接宣布所有自回归模型都会被替代。

参考资料

  1. Vaswani et al., Attention Is All You Need: https://arxiv.org/abs/1706.03762
  2. Austin et al., Structured Denoising Diffusion Models in Discrete State-Spaces: https://arxiv.org/abs/2107.03006
  3. Li et al., Diffusion-LM Improves Controllable Text Generation: https://arxiv.org/abs/2205.14217
  4. Chang et al., MaskGIT: Masked Generative Image Transformer: https://arxiv.org/abs/2202.04200
  5. Lou et al., Discrete Diffusion Modeling by Estimating the Ratios of the Data Distribution: https://arxiv.org/abs/2310.16834
  6. Sahoo et al., Simple and Effective Masked Diffusion Language Models: https://arxiv.org/abs/2406.07524
  7. Nie et al., Large Language Diffusion Models (LLaDA): https://arxiv.org/abs/2502.09992
  8. Ye et al., Dream 7B: Diffusion Large Language Models: https://arxiv.org/abs/2508.15487
  9. Leviathan et al., Fast Inference from Transformers via Speculative Decoding: https://arxiv.org/abs/2211.17192
相关推荐
znx9391 小时前
机器学习赋能量化交易:解锁新时代投资的核心竞争优势
人工智能·python·机器学习·期魔方
atwednesday1 小时前
理解 Transformer
人工智能
Timmy1 小时前
前端转全栈笔记:讲框架之前,先把 TypeScript 这关过了
前端·人工智能·node.js
_Meilinger_1 小时前
碎片笔记|生成图像检测与溯源领域近期审稿意见锦囊
人工智能·生成图像检测·投稿经验·生成图像取证·ai治理·生成图像溯源·科研经验
XUHUOJUN1 小时前
AI-Native Software Engineering:氛围编程的工程化与 Software Change Management
人工智能·ai-native
AI日报派送佬1 小时前
2026年9月30日AI行业日报|OpenAI DevDay连发25项更新,GPT-6.1 Sol&Dots智能体落地,全球AI监管收紧
人工智能·gpt·openai·ai智能体·大模型技术·人工智能前沿·ai监管
Frag0ut1 小时前
Chrome Canary 最新实验功能:AI Mode、按设备筛选历史记录与 Gemini 自动改密
前端·人工智能·chrome·dev·beta·canary·最新实验功能
minji...1 小时前
LangGraph-AI智能体开发框架 - LangGraph 入门案例1 : 智能快递配送系统
人工智能·python·ai·langchain·大语言模型·agent·langgraph
xiangzhihong81 小时前
随心陪玩游戏全栈系统
人工智能