ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

神经对话生成对抗训练复现指南:从Seq2Seq到Policy Gradient的完整实现

神经对话生成对抗训练复现指南:从Seq2Seq到Policy Gradient的完整实现 简介面向机器学习课程设计的高分复现项目以神经对话生成对抗性学习论文为蓝本帮助高校学生快速搭建可运行的期末大作业。资源共二十个文件包含十二个源码文件分别覆盖数据生成、预训练、生成器、判别器、训练与测试等模块另有五个工程配置文件、一份说明文档、一份项目简介以及一个项目结构描述文件压缩包仅约五百七十KB轻量易部署。已有五百五十五人学习下载。代码注释完整并附文档讲解复现思路新手也能较快上手。项目涵盖对话生成、对抗训练、预训练与判别模型等核心环节可直接运行或修改也可作为实验报告与答辩演示的参考模板适合需要完成机器学习大作业、课程设计或论文复现实验的读者。1. 神经对话生成对抗性学习复现这份源代码包先解决三个离线问题机器学习大作业选「复现论文神经对话生成对抗性学习」的人多真正能从源码包跑出像样对话指标的少。这个项目解决的是对话生成里最经典的问题MLE 训练出的对话模型只会说「I dont know」「Me too」这类安全回复对抗训练能把多样性拉回来。实现上把 GAN 思路搬进对话任务生成器是 Seq2Seq Attention判别器判断一条回复像真人还是机器生成两边交替训练互相较劲。三类人需要它做大作业和毕设的学生、想复现 SeqGAN 类论文的读者、准备在对话生成方向做 baseline 的从业者前提是你已经能独立用 Python 和 PyTorch 写基本模型。如果只把源码包当黑匣子跑一遍就交差这题目基本白做。对抗训练给的 reward 信号稀疏不搞懂 loss 和 reward 怎么构造换数据、换参数时每一步都是盲区。下面从数据预处理开始带你把这套代码拆开看清楚。2. 从论文到代码数据预处理与 Seq2Seq 生成器的搭建细节复现类项目翻车多半不是模型写错而是数据没洗干净。对抗训练对数据质量比普通监督学习更敏感判别器要区分真实回复和机器回复语料里混了噪声样本判别器随便抓个特征就能赢生成器什么都学不到。所以第一步不是读模型是处理语料和词表。2.1 语料格式与词表构建min_count 和 max_vocab 是第一道坎常见做法是把对话语料整理成 tsv每行一对context \t response。Cornell Movie Dialogs 和 DailyDialog 预处理后都能转成这个格式关键在词表怎么建# data/build_vocab.py from collections import Counter def build_vocab(data_path, min_count3, max_vocab20000): counter Counter() with open(data_path, r, encodingutf-8) as f: for line in f: parts line.rstrip(\n).split(\t) if len(parts) 2: continue counter.update(parts[0].split()) counter.update(parts[1].split()) vocab {pad: 0, bos: 1, eos: 2, unk: 3} for w, c in counter.most_common(max_vocab - 4): if c min_count: vocab[w] len(vocab) return vocab这里有两个参数直接决定后面模型能不能收敛。min_count设 3 或 5把只出现过一两次的噪声词全过滤掉否则解码时模型总在 OOV 边界试探loss 居高不下。max_vocab限制在 2 万左右别贪大embedding 层占显存不说稀疏词对对抗训练没有正向帮助。四个特殊 token 的位置固定写死后面所有 mask、padding 逻辑都依赖这份约定换位置会出很难查的错。2.2 数据加载与长度分桶batch 里混进长尾句子会拖垮训练对话句子长度差异极大短的 3 个 token长的上百。直接把一个 batch 里所有句子 pad 到最长每一轮都有大量无效计算显存也被无谓占用。我一般会按长度分桶让每个 batch 内部句子长度尽量接近# data/bucket_dataset.py import torch from torch.nn.utils.rnn import pad_sequence def collate_with_bucketing(batch): # 按最长句子排序pad 长度由 batch 内最大值决定 batch sorted(batch, keylambda x: max(len(x[0]), len(x[1]))) src pad_sequence([torch.tensor(x[0]) for x in batch], padding_value0) tgt pad_sequence([torch.tensor(x[1]) for x in batch], padding_value0) return src, tgt代码逻辑是先把 batch 按最长句排序再用 pad_sequence 填充到该 batch 的最大长度。效果是短句多的 batch 计算量明显变小训练速度能提升 30% 左右。注意padding_value0必须和词表里pad的 id 一致否则 mask 逻辑全乱。这个函数配合 DataLoader 的collate_fn参数使用即可。2.3 生成器结构双向编码器 Attention 解码器对抗训练里的生成器一般沿用标准 Seq2Seq 加注意力没有 Attention 的生成器在长句上撑不住判别器轻易就能抓出机器味道。核心结构如下# models/generator.py import torch.nn as nn class Seq2SeqGenerator(nn.Module): def __init__(self, vocab_size, emb_dim256, hidden_dim512, dropout0.3): super().__init__() self.embedding nn.Embedding(vocab_size, emb_dim, padding_idx0) self.encoder nn.LSTM( emb_dim, hidden_dim, batch_firstTrue, bidirectionalTrue ) self.decoder nn.LSTM(emb_dim, hidden_dim * 2, batch_firstTrue) self.attention nn.Linear(hidden_dim * 4, 1, biasFalse) self.dropout nn.Dropout(dropout) self.out nn.Linear(hidden_dim * 2, vocab_size) def forward(self, src, tgt): # src: [B, T_s], tgt: [B, T_t] src_mask (src ! 0).unsqueeze(1) enc_out, (h, c) self.encoder(self.embedding(src)) dec_out, _ self.decoder(self.embedding(tgt), (h, c)) # 拼接 dec_out 和 enc_out计算 attention 权重 dec_expand dec_out.unsqueeze(2).expand(-1, -1, src.size(1), -1) enc_expand enc_out.unsqueeze(1).expand(-1, dec_out.size(1), -1, -1) attn_logits self.attention(torch.cat([dec_expand, enc_expand], dim-1)).squeeze(-1) attn_logits attn_logits.masked_fill(src_mask 0, -1e9) attn_w torch.softmax(attn_logits, dim-1) ctx torch.bmm(attn_w, enc_out) logits self.out(self.dropout(torch.cat([dec_out, ctx], dim-1))) return logits这里有几个关键点。第一编码器双向隐藏层维度翻倍解码器 LSTM 的 hidden 尺寸也要传hidden_dim * 2否则初始化时 shape 对不上直接报错。第二attention 打分前必须 mask 掉 padding 位置-1e9让 softmax 结果趋近于零否则注意力会分给填充符生成结果出现奇怪的重复 token。第三output 层可以分享 embedding 的权重词表大时节省参数但要注意 padding_idx 的梯度保护。这个类里还需要补sample和log_prob方法后面对抗训练主循环会用到。2.4 MLE 预训练没学会说人话就进对抗必废生成器不能从随机权重直接开始对抗训练。对抗训练的 reward 信号太稀疏随机模型采样的句子几乎没有完整词法结构判别器一步到位学成「一句话全是狗屁」的分类器整个训练就废了。所以先做 MLE 预训练# train_mle.py import torch import torch.nn.functional as F from torch.nn.utils import clip_grad_norm_ def train_mle(model, loader, optimizer, epochs10, max_norm5.0): model.train() for epoch in range(epochs): total 0.0 for src, tgt in loader: optimizer.zero_grad() logits model(src, tgt[:, :-1]) # teacher forcing loss F.cross_entropy( logits.reshape(-1, logits.size(-1)), tgt[:, 1:].reshape(-1), ignore_index0, ) loss.backward() clip_grad_norm_(model.parameters(), max_norm) optimizer.step() total loss.item() * src.size(0) print(fepoch {epoch:02d} loss {total / len(loader.dataset):.4f})teacher forcing 是必须的输入真实的上文词而不是上一个预测词否则早期误差传播太慢。ignore_index0让 padding 位置不参与 loss 计算这是和普通分类任务不一样的地方。max_norm5.0的梯度裁剪放在 step 前LSTM 训练最容易梯度爆炸这个值是我试下来比较稳的。MLE 阶段一般训到 loss 曲线不再明显下降大约 5~8 个 epoch。这里特别提醒别训到过拟合MLE 过拟合会让生成分布集中在少数高频词上对抗训练需要花更长的时间才能把分布拉开。3. 对抗训练核心CNN 判别器与 Policy Gradient 的联合更新生成器就绪后进入对抗阶段。这一章是整个复现项目的灵魂生成器输出的是离散 token梯度没法从判别器直接反传所以要用策略梯度Policy Gradient把判别器的打分变成 reward。工程上主要拆成三件事判别器怎么搭、reward 怎么算、两个网络怎么交替更新。3.1 判别器选型TextCNN 对上下文拼接做二分类原始论文里的判别器结构比较复杂涉及分层编码。开源复现里最常见的替代方案是 TextCNN把 context 和 response 拼接成一条序列用多尺寸卷积核抓 n-gram 特征速度快且效果稳定。实际训练下来这个方案对短对话足够用了# models/discriminator.py import torch import torch.nn as nn import torch.nn.functional as F class TextCNNDiscriminator(nn.Module): def __init__(self, vocab_size, emb_dim256, num_filters128): super().__init__() self.embedding nn.Embedding(vocab_size, emb_dim, padding_idx0) self.convs nn.ModuleList([ nn.Conv2d(1, num_filters, (k, emb_dim), padding(k - 1, 0)) for k in (2, 3, 4) ]) self.dropout nn.Dropout(0.5) self.fc nn.Linear(num_filters * 3, 2) def forward(self, context, response): # context [B, T_c], response [B, T_r] x torch.cat([context, response], dim1) # [B, T_c T_r] emb self.embedding(x).unsqueeze(1) # [B, 1, L, D] pooled [] for conv in self.convs: feat torch.relu(conv(emb)).squeeze(3) # [B, F, L] pooled.append(F.max_pool1d(feat, feat.size(2)).squeeze(2)) cat torch.cat(pooled, dim1) return self.fc(self.dropout(cat))为什么卷积核尺寸选 2、3、4对话里的判别性特征主要是 bigram 和 trigram比如「I dont」「I dont know」这种固定搭配再长的 n-gram 特征稀疏且容易被 max pooling 丢掉。padding 设置成(k-1, 0)是让卷积输出长度和输入保持一致方便直接 max pooling。两个关键约定一是判别器和生成器各自维护自己的 embedding不要共享权重否则对抗梯度会直接泄漏到生成器词向量里训练不稳二是判别器输出 2 分类 logits后续通过 softmax 取「真实类别」的概率作为 reward。3.2 Reward 构造完整序列打分与 MC rollout 补全策略梯度的核心是把离散采样和连续梯度连接起来。生成器对每个 context 采样一条 response然后把这条 response 交给判别器打分分数就是 reward。但这里有个细节对话生成是逐 token 产生序列的如果整条序列只给一个端到端 reward中间每个 token 都不知道自己到底做得好不好。常见的解法是 MC rollout用生成器把部分序列补全到结尾再打分多次采样取平均# rl_reward.py def mc_reward(discriminator, generator, contexts, responses, rollout_num5): rewards [] for c, r in zip(contexts, responses): # c: [T_c], r: [T_r] token_rews [] for t in range(1, len(r)): avg 0.0 for _ in range(rollout_num): rollout generator.rollout(c, r[:t], max_lenlen(r)) avg discriminator.prob_real(c, rollout) avg / rollout_num token_rews.append(avg) rewards.append(token_rews) return torch.tensor(rewards)逐 token 计算 reward 的逻辑是对于 response 里的第 t 个位置用已经生成的前 t 个 token 作为前缀让生成器自己 rollout 出后续内容再由判别器对完整序列打分。这样中间每个 token 都能获得一个「做完这件事值多少分」的估计。rollout_num大一点方差小但计算量成倍上涨5 是性价比比较高的值。max_len建议和训练时的最大长度一致避免 rollout 出过短或过长的序列导致打分偏差。这里没有 baseline 处理后面主循环里会补上。3.3 对抗更新主循环d_step 与 g_step 的节奏怎么定判别器和生成器的更新必须交替进行而且节奏比例很关键。判别器更新太少reward 区分度低梯度信号弱更新太多判别器过强生成器学到的只有「被惩罚」loss 震荡。常见的训练节奏是判别器每轮更新 4 次生成器只更新 1 次# train_adv.py def adversarial_step(gen, disc, ctx_loader, gen_opt, disc_opt, d_step4, g_step1, top_k10): # 1) 更新判别器 for _ in range(d_step): contexts, _ next(ctx_loader) true_resp sample_real_responses(contexts) # 从语料取真实回复 fake_resp gen.sample(contexts, top_ktop_k) # 当前生成器采样 d_logits disc(contexts, torch.cat([true_resp, fake_resp], dim0)) labels torch.cat([ torch.ones_like(true_resp[:, 0]), torch.zeros_like(fake_resp[:, 0]) ]).long() d_loss F.cross_entropy(d_logits, labels) disc_opt.zero_grad(); d_loss.backward(); disc_opt.step() # 2) 更新生成器 for _ in range(g_step): contexts, _ next(ctx_loader) responses gen.sample(contexts, top_ktop_k) rewards mc_reward(disc, gen, contexts, responses) rewards (rewards - rewards.mean()) / (rewards.std() 1e-8) log_probs gen.log_prob(contexts, responses) g_loss -(log_probs * rewards).sum() / responses.size(0) gen_opt.zero_grad(); g_loss.backward(); gen_opt.step()判别器部分没什么好说的正负样本各占一半交叉熵训练。重点在生成器部分。log_probs是采样出的这堆 token 在生成器分布下的对数概率reward 乘上去再取负就是策略梯度的标准形式奖励高的序列对应 token 的概率会被推高。reward 归一化用「减均值除标准差」替代简单的减 baseline这是我从强化学习实践里学来的习惯比减常数收敛稳定得多。gen.sample必须用带随机性的采样greedy decoding 产生的序列分布太尖奖励信号对生成器没有区分度。top_k10的意思是采样候选只保留概率最高的 10 个 token既保证多样性又避免采样到语法噪声。4. 评估指标与训练调参BLEU、distinct-1/2 和超参表对抗训练跑起来之后最常遇到的问题不是代码报错而是「看着 loss 在降但生成结果怎么还是像复读机」。这说明评估指标选错了。对话生成是开放域的一对多任务一个 context 可以有无数种合理回复传统机器翻译指标在这里会给出误导性的结论。4.1 BLEU 在对话生成里只能当参考BLEU 衡量生成序列和参考序列的 n-gram 重合度思想是「翻译结果越接近标准答案越好」。但对话生成里没有标准答案同一个 context 换一种说法、换一个角度回答语义完全合理BLEU 照样给低分。复现论文里一般会报告 BLEU但那是为了和原文对比不是用来判断训练好坏的。如果你发现对抗训练后 BLEU 掉了一些先别慌看下一个指标。4.2 distinct-1/distinct-2多样性才是对抗训练的主目标对抗训练本身不直接优化多样性它通过「骗过判别器」间接让生成器远离安全回复。所以评估多样性要看 distinct 指标计算生成语料里不重复 n-gram 的比例# eval/distinct.py def distinct_ngram(responses, n2): ngrams set() total 0 for text in responses: toks text.split() for i in range(len(toks) - n 1): ngrams.add(tuple(toks[i:i n])) total 1 return len(ngrams) / total if total else 0.0n1是词级别的多样性n2是短语级别的多样性。实战中主要看 distinct-2因为单个词重复容易被停用词刷高双词组合更能反映真实表达多样性。对抗训练前后对比如果 distinct-2 从 0.03 左右涨到 0.10 上下说明生成器确实在摆脱安全回复如果原地踏步多半是判别器太弱或 reward 没有区分度回头查第 5 章。4.3 超参表每个参数在什么范围改坏了是什么表现参数建议值失败特征emb_dim256过小模型容量不足MLE loss 难下降hidden_dim512过大显存翻倍收益有限MLE 学习率1e-3过大 loss 在 3 个 epoch 内变 NaN对抗生成器学习率5e-5 ~ 1e-4过大会导致采样分布剧烈震荡对抗判别器学习率1e-4过大会让判别器 loss 瞬间归零d_step / g_step4 / 1判别器更新太少reward 区分度低top_k 解码10 ~ 20太小多样性差太大出现语法噪声max_len32过长 OOM过短学不到长句结构MLE 阶段学习率比普通分类任务要保守因为对话序列误差会沿着 LSTM 累积学习率稍大就梯度爆炸。对抗阶段生成器学习率要再降一个量级策略梯度的方差远大于交叉熵学习率大了训练曲线像心电图。max_len32对短对话足够超过这个长度的句子可以直接在预处理阶段截断省掉的显存足够把 batch size 翻倍。训练日志的形态也要心里有数。MLE 阶段 loss 应该平滑下降如果出现突然跳高再回落说明学习率偏大或某个 batch 里混了超长句子。对抗阶段判别器 loss 会在 0.6~0.7 附近震荡这是正常的说明它在真实和生成样本之间摇摆如果稳定在 0.0~0.1 附近说明判别器过强生成器已经被完全碾压需要减小 d_step 或给判别器加 dropout。5. 论文复现避坑指南五个跑通后才知道的坑对抗训练是出了名的难调很多问题从 loss 数字上看不出来只有生成结果暴露真实情况。这一章把最常见的五个坑按「现象 → 原因 → 解决」梳理清楚每条都是血泪经验。5.1 判别器 loss 瞬间归零或者 reward 全在 0.5 附近抖动现象判别器训练几轮后 loss 变成 0.0x生成器采样的句子全被判别器判为假。另一种情况是判别器 loss 一直在 0.69 左右不动reward 完全没有区分度。原因第一种是判别器学习率太大或正负样本分布失衡真实样本太少判别器一步就记住了所有真实回复的特征。第二种是判别器被先前的强负样本击穿所有输入都输出接近 0.5 的概率。解决判别器学习率统一降到 1e-4正负样本严格 1:1 构造。负样本不要用 greedy decoding 的输出那种序列高度雷同要用top_k10的随机采样保证负样本分布尽量接近真实数据。如果 reward 一直在 0.5 抖先把生成器冻住单独训练判别器 20 步直到它在验证集上准确率超过 80%再打开生成器的梯度。这一步是最容易忽略的判别器没训好之前生成器怎么更新都是原地转圈。5.2 生成器采样全是「I dont know」和「 」现象对抗训练跑了 2000 步采样出来的 20 句回复里有 15 句是「I dont know」或直接输出结束符。原因这不是对抗训练的问题是 MLE 预训练阶段就埋下的雷。安全回复在语料里出现频率高MLE 把这些高频回复的概率推得很尖对抗训练的 reward 信号再强也很难撼动已经固化的分布。另外解码温度太低也会让采样集中到概率最高的几个 token。解决回退到 MLE 阶段的 checkpoint确认 loss 是否已经收敛到正常范围比如 5.0 以下。如果 MLE 阶段采样出来的句子就已经是复读机说明训练不够或数据里有大量重复对话对。分布固化严重时把解码 temperature 提高到 1.2~1.5让概率分布变平缓再配合 top_k 采样能明显改善多样性。还有一个取巧做法对抗训练的前 100 步不更新生成器只让判别器见识足够多的真实回复这时生成器的旧分布不会进一步固化。5.3 OOM 和训练速度玄学小 batch 也炸显存现象batch size 已经降到 16max_len 也限制在 32还是 OOM训练速度慢到一小时跑不完一个 epoch。原因大概率是 batch 内句子长度差距太大padding 后有效计算量极低。另一个隐藏开销是 MC rolloutrollout_num5意味着每条采样回复要额外做 5 次生成和 5 次判别器前向显存和算力都翻了好几倍。解决一定要分桶按 2.2 节的方式让 batch 内长度接近。rollout 不是每一步都要做满实践里可以只在后半段 token 做 rollout前半段直接用完整序列的 reward 代替效果几乎不变。再就是生成器的 sample 阶段要开torch.no_grad()判别器的 rollout 打分也同理这两个阶段都不需要保留计算图。把这两处 grad 关掉显存占用能降一半。5.4 评估指标虚高distinct 涨了但生成结果还是复读机现象distinct-1 从 0.02 涨到 0.30看着很漂亮人工一看生成结果全是「Yes」「No」「OK」这类单字回复连不成一句话。原因distinct-1 统计的是不同单词的数量比例短回复天然占优势。「Yes」「No」这种单字回复每个都是不同的特殊 tokendistinct 自然虚高。这是评估指标的盲区不是模型真的变好了。解决主要看 distinct-2单字回复的组合数少刷不动这个指标。更稳的做法是统计平均回复长度低于 6 个 token 说明生成器在走捷径。评估时还可以把newline、标点等无意义 token 从统计里剔除。我自己的习惯是每个 epoch 采样 200 条回复同时打印 distinct-2 和平均长度两个数一起看才有参考价值。6. 进阶三阶段验证法把复现实验从玄学变成流程对抗训练最大的问题是「不知道自己训得对不对」。loss 降了不代表对话质量好reward 涨了也可能只是判别器过拟合。我踩过几次坑之后养成了一个习惯每次拿到新数据、改完模型结构都强制走一遍三阶段验证。这套流程三分钟跑完能挡住 80% 的无效训练。Phase 0 是数据体检检查三件事语料有多少行、平均句子长度多少、词表过滤后的 OOV 比例。OOV 比例超过 5% 就直接回头调min_count别带病上场。Phase 1 是 MLE 快速收敛测试只训 3 个 epoch 看 loss 曲线必须平滑下降且最终值低于正常训练的 60% 才有资格进对抗阶段。Phase 2 才是对抗训练但前 500 步只看两个信号判别器验证准确率有没有先冲到 80%distinct-2 有没有开始抬头。两个信号都出现才继续往下训。# scripts/verify.sh —— 每次改数据或模型后先跑一遍 python data/check_corpus.py --path data/train.tsv --min_count 3 python train_mle.py --epochs 3 --save checkpoints/mle_quick.pt python eval/generate.py --ckpt checkpoints/mle_quick.pt --show 20 python train_adv.py --resume checkpoints/mle_quick.pt \ --steps 500 --log_every 100 --eval_every 100check_corpus.py对应 Phase 0输出语料行数、平均长度和 OOV 率train_mle.py --epochs 3对应 Phase 1快速验证梯度链路没断generate.py --show 20打印 20 条采样结果人工扫一眼就知道有没有复读机倾向最后的train_adv.py --steps 500是 Phase 2 的预热看判别器准确率和 distinct-2 有没有按预期走。每跑完一个命令肉眼核对一步不对就止血绝不带着疑问继续往前。从那以后我每次配新数据、调模型结构都强制走一遍这三步三分钟换掉一晚上的无效训练这个习惯一直留到现在。希望帮到你。本文还有配套的精品资源点击获取
返回列表