扩散式语言模型,简单说就是借用图像扩散的加噪-去噪思路来生成文本的一类模型。它和GPT这类自回归模型最大的区别在于生成方向:自回归只能从左到右逐字预测,扩散式模型则从一个满是噪声的序列开始,通过多轮去噪逐步还原出完整句子。这个方向的代表工作包括Diffusion-LM、D3PM以及后续各种改进,它们的核心都一样:给文本加噪,训练模型去噪,推理时从噪声反向走回文本。
这篇文章我会用PyTorch从零实现一个最小可运行的掩码扩散语言模型,把加噪、去噪、训练、采样和常见问题完整走一遍。适合有深度学习基础、想真正理解扩散语言模型内部机制的读者;如果你已经在看相关论文,这篇文章也可以帮你把论文里的抽象公式落到代码上。
在动手之前先给你一个判断:这个方向值得学,但不要指望一个简化版Demo就追上GPT的生成质量。扩散语言模型目前的优势更多体现在可控生成、并行解码和局部编辑,而不是单纯的语言流畅度。理解了这一点,后面的每一步才会更有方向。
1. 扩散式语言模型解决什么问题,和自回归模型差在哪
1.1 一个直观理解:不是“从左到右”,而是“从噪声里逐步还原”
自回归语言模型很好理解。你输入“今天天气”,模型预测下一个词最可能是“不错”“很好”“很差”。生成的时候用causal mask把注意力限制在左侧,让每个位置只能看到它前面的token。这种设计决定了它天生适合逐词生成,推理速度受序列长度限制,也没法回头修改已经生成的词。
扩散式语言模型的思路完全不一样。训练的时候,我们不预测下一个词,而是把一整句话拆开,随机遮掉一部分token,再让模型根据剩下的token把被遮掉的内容猜回来。更准确地说,我们会定义一条“加噪路径”:从一个完整句子出发,逐渐把越来越多的token替换成特殊符号,直到整个序列几乎全是噪声。模型需要学习的是这条路径的反方向:从带噪序列一步步还原出原始文本。
推理时,模型不是从左往右生成,而是先放上一整段mask,然后反复执行“预测所有位置、填上置信度最高的部分、保留已有内容”这个循环,多轮之后整段文本就浮现出来。
你可以这样类比:自回归像写作文,一行一行往下写,不能回头改;扩散式则像你把一篇作文用涂改液遮掉一部分字,再根据残留的字和上下文把内容恢复出来。遮掉得越多,恢复越难;模型要能做到在任意遮掉比例下都能猜。
1.2 它真正解决的是什么问题
扩散式语言模型不是来取代自回归模型的,它主要解决几个自回归模型不太好处理的问题。
第一,非自回归生成。自回归模型生成一个长度为L的句子,至少要执行L次前向推理。扩散模型只需要执行固定步数的去噪,比如50步,每一步同时预测所有位置,理论上可以做到并行生成,长句子的生成延迟不会线性增长。
第二,双向上下文利用。自回归模型天生只能用左侧信息,即使扩大上下文窗口,也没有改变这个方向限制。扩散模型在去噪过程中使用双向注意力,每个token都能同时参考左右两侧的信息。这让它更适合做局部文本编辑、文本修复这类任务。
第三,可控生成。在扩散采样过程中,你可以对每一轮去噪结果施加约束,比如指定某些位置必须是某个词、整体情感倾向要偏正面、生成内容要包含某些关键词。这种“在生成过程中引导”的方式比自回归模型重新采样或者调prompt更直接。
1.3 适合谁学、需要什么基础
如果你对Transformer和PyTorch已经有实际使用经验,能看懂attention、embedding、交叉熵这些概念,那这篇文章的代码部分对你没有障碍。如果你只是听说过“扩散模型”但没写过图像生成,也可以学,我会把加噪和去噪的逻辑一步步拆清楚。
如果完全没写过Transformer,建议先跑一个最小Transformer分类器再回来。扩散语言模型的代码结构本身不复杂,但一旦出问题,排查的时候会同时涉及数据、网络、损失函数、采样策略好几层,没有基础容易一头雾水。
2. 核心原理拆解:加噪、去噪、损失函数
2.1 文本是离散的,所以不能直接照搬图像扩散
图像扩散模型在连续像素上加高斯噪声:像素从清晰变成模糊,再从模糊变回清晰。这个过程数学上很干净,因为连续空间可以用正态分布描述,去噪的每一步也有解析解。
但文本token是离散的。“加噪”在文本里不是加一个随机浮点数,而是把某个token替换成另一个东西。怎么替换,替换成什么,就成了文本扩散设计时最先要回答的问题。
目前主要有三类做法。
第一类叫掩码扩散。把部分token替换成一个特殊的[MASK]标记。模型看到带掩码的序列,目标是预测被遮住的原始token。这类方法直观,效果也够用,本文的Demo就按这个思路写。
第二类叫转移矩阵扩散。不只用mask,还允许token被替换成其他词汇,比如均匀随机取一个词,或者按语言相似度转移。代表工作是D3PM。
第三类是在连续空间做扩散。先把token映射成embedding向量,在向量空间加高斯噪声,最后通过一个rounding步骤把向量映射回离散token。代表工作是Diffusion-LM。
入门阶段,掩码扩散最容易理解,也是理解其他所有变体的基础。
2.2 前向过程:按时间步逐步掩码
先定义总步数T,比如T=100。对于一条原始文本x0,我们随机采样一个时间步t,t越大表示噪声越重。
前向加噪规则是:按照比例 t/T 随机选择一部分位置,把这些位置的token替换成[MASK]。t=0时不遮任何token,t接近T时几乎把整句话都遮掉。
这个设计有一个关键点:每个样本在训练时只会被加噪一次,而不是在同一个batch里展示所有噪声程度。因为t是随机采样的,所以一个batch里有的样本噪声轻,有的样本噪声重。模型必须在同一套参数下处理各种噪声程度,这就逼着它学会“在模糊信息中还原”。
加噪时还要注意一个问题:不要把所有token都遮住,至少要保留一个可见token。否则模型拿到的是纯噪声,没有任何上下文可以依赖,预测就变成了瞎猜。虽然理论上模型可以从词频先验去猜,但实验里通常会让模型至少看到一个token。
2.3 反向过程:用双向Transformer去噪
去噪模型的输入是带噪序列和当前时间步t,输出是每个位置对所有词表的概率分布。模型的结构用双向Transformer Encoder,不能带causal mask。
原因很简单:mask位置需要同时看左右两边的可见token。比如“今天[天气]很好”,要猜“天气”这个词,既需要看左边的“今天”,也需要看右边的“很好”。自回归的causal mask只让看左边,信息不完整。
时间步t怎么融入模型?常见的做法是把t编码成一个向量,然后用MLP映射到模型维度,加到每个token的embedding上。这样模型知道当前处于哪个噪声级别,从而调整去噪策略:噪声轻时可以大胆填,噪声重时必须保守。
2.4 训练目标与时间步采样
训练目标是在被mask的位置上计算交叉熵损失,让模型预测的分布尽量接近真实token。
这里要注意:未被mask的位置不参与loss计算。因为这些位置是模型能看到的输入信息,如果让模型去预测它们,模型只需要记住输入就行,学不到任何去噪能力。
时间步t的采样必须覆盖0到T-1的整个范围。如果只固定一个t训练,模型只能处理固定噪声程度。均匀采样的目的是让模型在任意噪声比例下都能工作,推理时我们才能从高噪声逐步走到低噪声。
训练时还有一个细节:同一个batch里,各条样本的t可以不同。因为每一条样本是独立的,模型会通过t的embedding区分当前噪声程度。这比让整个batch共享同一个t要高效。
3. 环境准备:用最小的代价跑起来
3.1 软件依赖与版本建议
本文代码只需要Python和PyTorch,不需要额外安装Transformer库。至少需要torch、numpy、tqdm这三个包。
版本方面没有非常严格的要求,建议使用Python 3.9以上,PyTorch 2.0以上。新版本PyTorch对Transformer Encoder的封装更完善,代码也更省事。
如果你本地没有GPU,用CPU也能够完成这个实验,只是速度会慢一些。这个Demo设计的参数量很小,CPU上训练几千步是可行的。
3.2 硬件条件:CPU/GPU/显存
硬件条件主要看你要训练多久。如果只是跑通流程、看生成效果,CPU完全够。但如果你想训练一个看起来还行的模型,建议还是用GPU。
我的建议配置是:
- 最小编译要求:内存8GB,磁盘剩余10GB,CPU环境下能跑通训练循环。
- 入门GPU训练:显存4GB以上,序列长度128,batch_size 32,d_model 256,可正常训练。
- 更充分的训练:显存8GB以上,可以加大模型和batch_size。
如果你的机器只有CPU,不要开太大的batch_size和序列长度,否则一个batch要等很久。先用小参数把流程跑通,比追求训练速度更重要。
3.3 数据集选择
入门实验用不着几GB的大语料。WikiText-2是常见选择,但下载链接偶尔会变动;你也可以直接用自己手头的纯文本文件。
最省事的做法是把几篇技术文档、新闻文章或者小说章节拼成一个纯文本文件,按行切分,每行作为一条样本。关键是文本本身要有一定自然语言结构,不能只是无意义的字符堆积。
数据量也不需要太大。10MB到50MB的纯文本足以让这个Demo产生有意义的结果。如果你的数据集比较小,训练时可以把epoch设多一些,或者用小一些的模型容量来避免过拟合。
4. 最小可运行实现:一个掩码扩散语言模型
这一节是全文核心。我会按照“tokenizer -> 加噪 -> 去噪网络 -> 训练 -> 采样”的顺序给出可运行代码。为了让流程最短,我使用字符级tokenizer,不引入外部语料库。
4.1 整体结构
整个实现由五个部分组成:
- CharTokenizer:把文本转成token序列,并定义
[PAD]和[MASK]两个特殊token。 - add_noise:对给定token序列执行掩码加噪。
- DiffusionLM:双向Transformer去噪网络,接收带噪序列和时间步t。
- train_step:一个batch的训练逻辑。
- sample:从全mask序列开始逐步去噪生成。
先写tokenizer。
class CharTokenizer: def __init__(self, texts): chars = set("".join(texts)) self.vocab = sorted(chars) self.stoi = {c: i for i, c in enumerate(self.vocab)} self.itos = {i: c for i, c in enumerate(self.vocab)} self.pad_token_id = len(self.vocab) self.mask_token_id = len(self.vocab) + 1 self.vocab_size = len(self.vocab) + 2 def encode(self, text): return [self.stoi[c] for c in text] def decode(self, ids): return "".join( self.itos[i] for i in ids if i < len(self.vocab) )字符级tokenizer的好处是词表很小、实现简单,不需要下载预训练词表。缺点是每个token的信息量低,生成的东西读起来会有“字符感”。如果要更好的效果,可以换成BPE tokenizer,但那是工程优化,不影响理解原理。
4.2 加噪函数
加噪函数输入原始序列x0和时间步t,输出带噪序列和mask标记矩阵。
def add_noise(x0, t, mask_token_id, T): # x0: [B, L] 原始token序列 # t: [B] 当前时间步,范围 1 ~ T-1 B, L = x0.shape mask = torch.zeros_like(x0, dtype=torch.bool) for i in range(B): ratio = t[i].item() / T num_mask = int(ratio * L) # 至少保留一个可见token num_mask = min(num_mask, L - 1) perm = torch.randperm(L) mask[i, perm[:num_mask]] = True xt = x0.clone() xt[mask] = mask_token_id return xt, mask这里为什么要限制num_mask最多为L-1?前面说过,如果整句话被遮完,模型没有任何上下文,只能靠词频瞎猜。保留至少一个token让模型有机会学“基于局部上下文重构”。
另外注意,mask位置是随机选的,不是固定选前几个或后几个。这样才能保证模型学习到不同位置、不同比例下的去噪能力。
4.3 去噪网络
去噪网络的主体是Transformer Encoder,加时间步嵌入。
import math import torch import torch.nn as nn import torch.nn.functional as F def get_sinusoidal_position_embedding(max_len, d_model): pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len).unsqueeze(1).float() div_term = torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) return pe class DiffusionLM(nn.Module): def __init__(self, vocab_size, d_model=256, nhead=4, num_layers=4, max_len=512): super().__init__() self.vocab_size = vocab_size self.d_model = d_model self.token_emb = nn.Embedding(vocab_size, d_model) self.pos_emb = get_sinusoidal_position_embedding(max_len, d_model) self.pos_emb = nn.Parameter(self.pos_emb, requires_grad=False) self.time_mlp = nn.Sequential( nn.Linear(d_model, d_model), nn.SiLU(), nn.Linear(d_model, d_model), ) encoder_layer = nn.TransformerEncoderLayer( d_model=d_model, nhead=nhead, batch_first=True, ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) self.out_proj = nn.Linear(d_model, vocab_size) def forward(self, x, t): # x: [B, L], t: [B] B, L = x.shape h = self.token_emb(x) # [B, L, D] h = h + self.pos_emb[:L].unsqueeze(0) # 加位置编码 t_emb = self._time_embedding(t) # [B, D] t_emb = self.time_mlp(t_emb).unsqueeze(1) # [B, 1, D] h = h + t_emb h = self.encoder(h) # [B, L, D] logits = self.out_proj(h) # [B, L, V] return logits def _time_embedding(self, t): # t: [B] device = t.device half_dim = self.d_model // 2 emb = math.log(10000.0) / (half_dim - 1) emb = torch.exp(torch.arange(half_dim, device=device) * -emb) emb = t.unsqueeze(1) * emb.unsqueeze(0) emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) return emb三个关键设计:
第一,位置编码用sinusoidal固定编码,不学习。因为文本长度是动态的,固定编码对未出现过的长度更友好。
第二,时间步t通过sinusoidal编码后进入两层MLP,再与token序列相加。这比直接拼一个整数进embedding更科学,因为sinusoidal编码能让相近的t有相近的表示,模型更容易理解噪声级别的连续性。
第三,Transformer Encoder默认是双向注意力,没有causal mask。这一点对扩散语言模型是必需的,不要改成causal。
4.4 训练循环
训练逻辑是:取一个batch -> 采样t -> 加噪 -> 预测 -> 只计算mask位置的交叉熵 -> 反向传播。
def train_step(model, optimizer, batch, mask_token_id, T): model.train() x0 = batch # [B, L] B = x0.shape[0] t = torch.randint(1, T, size=(B,), device=x0.device) xt, mask = add_noise(x0, t, mask_token_id, T) logits = model(xt, t) # [B, L, V] # 只在被mask的位置计算loss loss = F.cross_entropy( logits[mask], x0[mask], reduction="mean", ) optimizer.zero_grad() loss.backward() optimizer.step() return loss.item()这里logits[mask]是利用PyTorch的布尔索引,把每个样本中被mask位置的logits取出来,形状是[总mask数, vocab_size],再与对应位置的原始token计算交叉熵。
为什么要用mean而不是sum?因为每个batch里mask数量不同,如果用sum,batch之间的loss量级会随mask数量变化,不好对比。用mean之后,不管mask多少,loss都代表平均每个被遮token的预测误差。
主训练循环如下:
def main(): # 准备数据 with open("data.txt", "r", encoding="utf-8") as f: lines = [line.strip() for line in f if line.strip()] tokenizer = CharTokenizer(lines) seq_len = 128 samples = [] for line in lines: ids = tokenizer.encode(line) if len(ids) >= seq_len: for i in range(0, len(ids) - seq_len + 1, 64): samples.append(ids[i:i+seq_len]) else: padded = ids + [tokenizer.pad_token_id] * (seq_len - len(ids)) samples.append(padded) T = 100 model = DiffusionLM( vocab_size=tokenizer.vocab_size, d_model=256, nhead=4, num_layers=4, max_len=seq_len, ) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3) batch_size = 32 for step in range(5000): idxs = torch.randint(0, len(samples), size=(batch_size,)) batch = torch.tensor( [samples[i] for i in idxs], dtype=torch.long ) loss = train_step( model, optimizer, batch, tokenizer.mask_token_id, T, ) if step % 500 == 0: print(f"step {step}, loss {loss:.4f}") # 每个阶段生成一个样例 sample_text = sample( model, tokenizer, seq_len=64, mask_token_id=tokenizer.mask_token_id, T=T, ) print("sample:", sample_text)我一般建议先训练几百步看loss有没有下降趋势,再决定是否继续。不要一上来就写5000步,万一数据或代码有bug,浪费时间。
4.5 采样生成
采样是扩散语言模型最关键的环节。我的简化实现采用“逐步填充”策略:
- 初始化一个全mask序列。
- 从高噪声时间步开始,让模型预测所有mask位置的分布。
- 选出置信度最高的位置,用概率最高的token填充。
- 降低噪声级别,重复上述过程,直到填满所有位置。
@torch.no_grad() def sample(model, tokenizer, seq_len, mask_token_id, T=100, steps=50, temperature=1.0): model.eval() device = next(model.parameters()).device x = torch.full((1, seq_len), mask_token_id, device=device) B = 1 for step in range(steps): cur_t = T - 1 - int((T / steps) * step) cur_t = max(cur_t, 0) t = torch.full((B,), cur_t, device=device) logits = model(x, t) # [1, L, V] # 压低mask和pad的输出概率,防止模型生成特殊token logits[:, :, mask_token_id] = -1e9 logits[:, :, tokenizer.pad_token_id] = -1e9 probs = F.softmax(logits / temperature, dim=-1) token_probs, token_ids = probs.max(dim=-1) # [1, L] mask_positions = (x == mask_token_id) ratio = cur_t / T target_mask_count = min(int(ratio * seq_len), seq_len - 1) current_mask_count = mask_positions.sum().item() fill_count = min( max(current_mask_count - target_mask_count, 0), current_mask_count, ) if fill_count == 0: continue scores = token_probs.clone() scores[~mask_positions] = -1e9 topk_idx = torch.topk(scores, k=fill_count).indices for pos in topk_idx: x[0, pos] = token_ids[0, pos] return tokenizer.decode(x[0].tolist())这个采样策略有一个重要细节:每一轮不是把所有mask位置都填上,而是只填“置信度最高”的一部分。为什么要这样?
因为模型在单次前向里给出的预测不一定准确。如果一次性把所有位置都填死,低置信度位置的错误会保留到最后。分批填充时,先填最有把握的位置,这些位置的信息可以给下一轮去噪提供上下文,帮助模型修正对剩余位置的判断。这有点类似人在做填空时先填有把握的,再回头推理难的。
还有个细节:把mask_token_id和pad_token_id的logits压低。虽然训练时模型几乎不会把原始token预测成mask,但采样过程中为了避免小概率输出特殊token,直接封掉更稳妥。
temperature参数控制概率分布的尖锐程度。temperature越低,分布越尖锐,生成越保守;temperature越高,分布越平坦,生成越多样。一般取值在0.8到1.5之间。训练质量不太好的时候,先用低temperature看稳定输出。
5. 训练与验证:怎么判断模型真的学会了
5.1 训练时盯哪些指标
训练loss是最直接的指标。掩码扩散模型的loss下降曲线通常不像自回归模型那么平滑,因为每次随机采样了不同的t,噪声程度差异大。
我更推荐同时在固定样本上做生成验证。每500步生成一条样例,观察内容从完全乱码逐渐变成有意义的单词或短句。这个变化比loss数字更能说明问题。
如果想更量化,可以准备一条验证集,固定一批样本和固定的mask位置,计算模型在这些固定mask位置上的预测准确率。固定验证的好处是不同训练阶段之间可以公平对比;如果不固定,mask位置每次都不一样,准确率波动会很大。
5.2 生成质量的三个检查
看到生成结果后,不要只凭直觉判断“好不好”,按三个角度检查。
第一,语言是否通顺。模型生成的字符能否组成正常单词,单词之间是否符合基本语法。对于字符级模型,刚开始只能生成词汇碎片,这是正常的;训练充分后应该能拼出完整单词。
第二,是否复读了训练集。如果模型把训练样本原封不动输出,说明它记住了数据,但没有学会泛化。这种情况在小数据集上很常见。判断方法是看生成结果里有没有出现训练集中少见的连续片段。
第三,长度和语义是否合理。扩散模型没有“生成终止”机制,所以输出的长度是预设的。如果生成到后半段开始乱码,说明模型的语言连贯性只维持了短距离。
5.3 参数调节:温度、去噪步数、训练步数
temperature和生成多样性直接相关。模型训练不充分时,用低temperature更容易看到可读输出,比如0.8;模型训练充分后,可以调到1.0或更高来增加多样性。
去噪步数steps影响生成质量和速度。steps太少,比如5步,每轮填充大量位置,预测错误率高;steps