ARTICLE DETAIL

资讯详情

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

SeqGAN对抗神经网络:离散文本生成的工程解法与Python实战

SeqGAN对抗神经网络:离散文本生成的工程解法与Python实战 简介这份资源是SeqGAN对抗神经网络的Python完整源码与配套数据面向希望深入理解序列生成与强化学习交叉应用的开发者及研究人员。SeqGAN将序列生成转化为策略优化问题通过策略梯度更新生成器弥补传统GAN在时间依赖与顺序性数据上的不足。资源包共13个文件约5.75MB以6个py源码文件为核心涵盖生成器、判别器、rollout、序列GAN主流程与数据加载等模块另含2个pkl模型参数文件、2张png训练曲线图、1个txt实验日志、1个md说明文档及1个zip压缩包结构清晰便于按模块研读。项目分两阶段展开先以预言机正样本和最大似然估计进行监督学习再通过生成器与判别器的对抗博弈提升生成质量。已有393人学习读者可借此掌握损失函数设计、训练循环与策略梯度更新的实际写法并理解序列奖励函数在文本或音频等序列数据生成中的落地方式。1. SeqGAN 对抗神经网络把离散文本生成塞进 GAN 的工程解法做文本生成的人大多踩过同一个坑用交叉熵逐词训练出来的模型推理时一旦换成 BLEU、ROUGE 或者人工评分指标就往下掉。原因是训练目标和评估目标根本不一致——模型学的是「给定前文下一个词概率最大」而不是「整句话读起来好不好」。SeqGAN 就是冲着这个断层来的它把生成器当成强化学习里的策略网络用判别器给出的分数当奖励直接优化整句的生成质量。标题里的「对抗神经网络」不是把图像那套 DCGAN 搬过来而是用 GAN 的对抗思想去驱动一个序列决策过程。配套的 Python 完整源码和数据解决的正是复现门槛——这套东西自己从零写光是策略梯度那几行就能卡住大半天。2. 为什么文本生成不能直接套用图像 GAN离散采样的梯度断点2.1 图像 GAN 能端到端反传文本 GAN 卡在哪图像生成里生成器输出一张图判别器给一个标量梯度顺着判别器一路回传到生成器参数链路是通的。文本不一样生成器输出的是每个位置上词表上的概率分布要变成真实 token 得做采样或者 argmax这一步是离散的、不可导的。梯度到采样这一步就断了判别器的分数传不回生成器。早期有人试过 Gumbel-Softmax 做松弛把离散采样近似成可导的连续分布。这个思路在短序列、小词表上能跑但词表一大、序列一长松弛误差累积得厉害生成结果会退化成重复词或者乱码。SeqGAN 的作者选了另一条路既然梯度传不过去那就不传了改用强化学习的策略梯度把判别器输出当成 reward用 REINFORCE 更新生成器。这个选择决定了整个框架的形态。2.2 判别器只给整句打分中间步骤的奖励从哪来强化学习里有个经典难题叫信用分配一盘棋下完赢了到底是哪一步走得好文本生成同理判别器只能对一句完整的话给一个「真/假」概率但生成是逐词进行的每个词该分到多少奖励SeqGAN 的解法是蒙特卡洛搜索。对当前已经生成的前缀用生成器自己往后补完若干条完整序列把这些补完的序列丢给判别器打分取平均作为当前这一步的奖励估计。这样每个中间状态都能拿到一个带噪声但可用的奖励信号。补完的条数是个关键参数太少方差大太多计算量爆炸源码里一般设 16 或 32。2.3 生成器和判别器的训练节奏怎么错开GAN 训练最怕的就是一方碾压另一方。SeqGAN 里这个矛盾更尖锐判别器太强生成器拿到的奖励全是接近 0 的负数策略梯度没有有效信号判别器太弱奖励区分度不够生成器学不到东西。常见做法是每轮先训 k 次判别器再训 1 次生成器k 取 1 到 5 之间。源码里通常还会对判别器的输出做一点平滑避免它输出绝对的 0 或 1。另外生成器的预训练很关键——先用最大似然把生成器训到一个「能说出人话」的水平再接入对抗训练。如果一上来就对抗生成器输出的是随机词判别器一眼识破奖励信号毫无梯度可言。3. 用 Python 把 SeqGAN 跑起来从环境到第一个训练轮次3.1 环境依赖与数据准备这套源码依赖比较轻核心就是 PyTorch 和 NumPy。建议用 Python 3.8 以上PyTorch 1.10 以上。数据方面标题里提到的数据集通常是新闻标题或者诗句这类短文本每行一句UTF-8 编码。# 创建虚拟环境并安装依赖 python -m venv seqgan_env source seqgan_env/bin/activate # Windows 用 seqgan_env\Scripts\activate pip install torch numpy tqdm数据文件放在data/目录下命名为real.txt每行一条样本。如果用自己的数据注意两点一是去掉空行和超长行二是控制词表大小超过 10000 的词表会让判别器的 embedding 层变得很重训练速度明显下降。# 数据预处理构建词表并把句子转成 id 序列 from collections import Counter def build_vocab(path, max_vocab10000): counter Counter() with open(path, encodingutf-8) as f: for line in f: line line.strip() if line: counter.update(list(line)) # 按字符切分中文场景常用 # 保留最高频的 max_vocab 个词其余归入 unk vocab {w: i2 for i, (w, _) in enumerate(counter.most_common(max_vocab))} vocab[pad] 0 vocab[unk] 1 return vocab这里按字符切分是中文文本生成的常见做法英文场景可以换成按空格分词。pad和unk占掉 0 和 1 两个位置后面 embedding 的num_embeddings要设成len(vocab)。词表建好后存成 json训练和推理共用同一份避免两次运行词表不一致导致 id 对不上。3.2 生成器与判别器的网络结构生成器用 LSTM 加全连接输出层判别器用 CNN 做文本分类。这个组合是 SeqGAN 原论文的配置CNN 判别器比 LSTM 判别器训练更稳不容易出现梯度消失。import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, vocab_size, embed_dim32, hidden_dim32): super().__init__() self.embed nn.Embedding(vocab_size, embed_dim, padding_idx0) self.lstm nn.LSTM(embed_dim, hidden_dim, batch_firstTrue) self.fc nn.Linear(hidden_dim, vocab_size) def forward(self, x, hiddenNone): # x: (batch, seq_len) 的 token id emb self.embed(x) out, hidden self.lstm(emb, hidden) logits self.fc(out) # (batch, seq_len, vocab_size) return logits, hidden class Discriminator(nn.Module): def __init__(self, vocab_size, embed_dim32, num_filters64, kernel_sizes(2,3,4)): super().__init__() self.embed nn.Embedding(vocab_size, embed_dim, padding_idx0) self.convs nn.ModuleList([ nn.Conv2d(1, num_filters, (k, embed_dim)) for k in kernel_sizes ]) self.fc nn.Linear(num_filters * len(kernel_sizes), 1) def forward(self, x): emb self.embed(x).unsqueeze(1) # (batch, 1, seq_len, embed_dim) feats [torch.relu(conv(emb)).squeeze(3) for conv in self.convs] pooled [torch.max(f, dim2)[0] for f in feats] cat torch.cat(pooled, dim1) return torch.sigmoid(self.fc(cat))生成器的hidden_dim设 32 是原论文的配置实际用的时候如果数据量大可以加到 64 或 128。判别器的kernel_sizes覆盖 2、3、4 三种窗口分别捕捉二元、三元、四元短语特征这是文本分类里的经典 TextCNN 结构。num_filters每个窗口 64 个卷积核太小判别器学不动太大容易过拟合。3.3 蒙特卡洛搜索与策略梯度更新这是整个 SeqGAN 最核心也最容易写错的部分。生成器每生成一个词就要用蒙特卡洛补完来估计这一步的奖励。def get_reward(gen, dis, prefix, rollout_num16, max_len20): 对当前前缀做 rollout返回平均判别器分数作为奖励 rewards [] for _ in range(rollout_num): seq prefix.clone() hidden None # 用生成器补完剩余位置 for _ in range(max_len - prefix.size(1)): logits, hidden gen(seq, hidden) next_token torch.multinomial( torch.softmax(logits[:, -1, :], dim-1), 1 ) seq torch.cat([seq, next_token], dim1) # 补完的完整序列丢给判别器打分 score dis(seq) rewards.append(score) return torch.mean(torch.stack(rewards), dim0)rollout_num设 16 是精度和速度的折中。补完时用multinomial采样而不是 argmax是为了保持探索性argmax 会让所有 rollout 结果一样奖励估计失去意义。max_len要和训练数据的最大长度对齐短了截断信息长了浪费计算。拿到奖励后用策略梯度更新生成器def update_generator(gen, dis, real_data, gen_optim, max_len20): gen_optim.zero_grad() # 从起始符开始逐词生成 start torch.full((real_data.size(0), 1), 2, dtypetorch.long) # 2 是起始符 id seq start log_probs [] rewards [] hidden None for step in range(max_len): logits, hidden gen(seq, hidden) prob torch.softmax(logits[:, -1, :], dim-1) next_token torch.multinomial(prob, 1) log_prob torch.log(prob.gather(1, next_token) 1e-8) log_probs.append(log_prob) # 对当前前缀做 rollout 拿奖励 r get_reward(gen, dis, torch.cat([seq, next_token], dim1)) rewards.append(r) seq torch.cat([seq, next_token], dim1) # 策略梯度reward 作为权重乘在 log_prob 上 log_probs torch.cat(log_probs, dim1) rewards torch.cat(rewards, dim1) # 对奖励做标准化降低方差 rewards (rewards - rewards.mean()) / (rewards.std() 1e-8) loss -(log_probs * rewards).mean() loss.backward() gen_optim.step() return loss.item()奖励标准化这一步很关键不做的话策略梯度方差极大训练会剧烈震荡。1e-8是防止 log 零的兜底。起始符 id 设成 2是因为 0 和 1 被 pad 和 unk 占了实际用的时候要跟词表构建逻辑对齐。4. SeqGAN 训练避坑从奖励塌缩到模式崩溃的排查清单4.1 判别器分数一直停在 0.5 附近现象训练几十轮后判别器对真实数据和生成数据的输出都在 0.5 上下生成器奖励没有区分度loss 不降。原因判别器容量不够或者学习率太低学不到真假数据的边界。也可能是生成器预训练太充分生成的数据已经和真实数据分布很接近判别器确实分不出来。解决先把判别器的学习率调大一到两倍观察分数是否拉开。如果还不行检查判别器的 embedding 维度是不是太小32 维在词表较大时确实不够用可以加到 64。另外确认预训练轮次没有过多一般生成器预训练 50 到 100 轮就够了再多对抗阶段就没有提升空间。4.2 生成结果全是重复词或者固定句式现象生成器输出「的的的的的」或者每次都生成同一句话多样性极低。原因典型的模式崩溃。奖励标准化之后如果某个模式偶然拿到高分策略梯度会拼命强化这个模式其他模式被压制。另外蒙特卡洛补完的方差太大奖励信号噪声高生成器会倾向于收敛到最安全的输出。解决把rollout_num从 16 加到 32降低奖励估计的方差。同时在奖励里加一个长度惩罚或者重复惩罚项对连续重复的词扣分。还有一个工程上的技巧是给判别器的输出加标签平滑把真实样本的标签从 1 改成 0.9假样本从 0 改成 0.1防止判别器输出过于极端导致奖励信号饱和。4.3 训练到一半 loss 突然变成 NaN现象前几十轮正常突然某一轮 loss 爆成 NaN参数全废。原因策略梯度里的 log_prob 在概率接近 0 时数值不稳定加上奖励标准化时分母的 std 可能接近 0除法直接炸掉。另外 LSTM 在长序列上梯度累积也容易溢出。解决在 log 里加的1e-8可以适当放大到1e-6。奖励标准化的分母加1e-8不够改成1e-5更稳。如果还炸加梯度裁剪torch.nn.utils.clip_grad_norm_(gen.parameters(), max_norm5.0)max_norm设 5.0 是经验值太小会拖慢训练太大起不到保护作用。另外检查输入序列长度超过 30 的序列在 LSTM 上容易出问题建议截断到 20 到 25。4.4 生成器和判别器 loss 同步下降然后一起卡住现象两个 loss 一起降到一个平台期之后都不动生成质量也不再提升。原因双方进入了平衡态判别器分不出真假生成器也没有梯度信号。这在 GAN 训练里很常见不一定是 bug但意味着训练到此为止了。解决如果当前生成质量已经可用直接停掉取模型。如果还不够好可以尝试给判别器加一点噪声或者调整两者的学习率比例打破平衡。另一个思路是换更难的判别任务比如把二分类改成多分类让判别器判断样本属于真实数据、生成器 A、生成器 B 中的哪一类增加判别器的学习压力。4.5 用自己的数据训练时效果远差于示例数据现象示例数据上跑得好好的换成自己的数据后生成结果一塌糊涂。原因数据分布差异。示例数据通常是短句、格式规整自己的数据可能有长尾词、特殊符号、长度差异大。词表构建时如果按字符切分英文和数字会被拆得七零八落。解决先做数据清洗去掉 HTML 标签、特殊符号、超长行。英文场景改成按空格分词中文场景可以试试按 jieba 分词而不是按字符。词表大小根据数据量调整数据少于 1 万条时词表控制在 5000 以内否则 embedding 层参数太多学不好。最后确认训练数据的长度分布如果大部分句子在 10 个词以内max_len设 20 就够设太大反而引入噪声。5. 把 SeqGAN 用出效果奖励设计和评估的两个进阶技巧5.1 用 BLEU 做混合奖励比纯判别器分数更稳纯用判别器分数做奖励训练前期信号很弱因为判别器还没学好。一个实用的改进是把 BLEU 分数混进去。具体做法是对每个 rollout 补完的序列除了判别器分数再算一个和真实数据的 BLEU两者加权求和作为最终奖励。from nltk.translate.bleu_score import sentence_bleu def mixed_reward(gen_seq, real_seq, dis_score, alpha0.3): 判别器分数和 BLEU 的加权组合 gen_tokens [str(t) for t in gen_seq.tolist()] real_tokens [str(t) for t in real_seq.tolist()] bleu sentence_bleu([real_tokens], gen_tokens) return (1 - alpha) * dis_score alpha * bleualpha控制 BLEU 的权重0.3 是个保守的起点。训练前期判别器不准可以把 alpha 调到 0.5后期再降回 0.2。注意 BLEU 计算本身有开销rollout_num 大的时候会明显拖慢训练可以每几个 step 才算一次 BLEU中间步骤只用判别器分数。5.2 评估不能只看 loss要盯生成样本的多样性和合理性GAN 的 loss 和生成质量没有直接对应关系这是血泪经验。评估 SeqGAN 至少要看三个指标一是生成样本的 distinct-1 和 distinct-2衡量用词多样性二是人工抽检 50 条生成结果看有没有语法错误和逻辑断裂三是如果下游有具体任务比如对话或者摘要直接拿生成结果去跑下游指标。我一般会在训练脚本里每 10 轮存一次生成样本到文件训练结束后统一看。这样比盯着 loss 曲线有用得多。另外记得固定随机种子做对比实验不然两次运行的生成结果差异可能比模型改进带来的差异还大。# 每 10 轮保存生成样本 if epoch % 10 0: gen.eval() with torch.no_grad(): samples generate_samples(gen, num20, max_len20) with open(fsamples_epoch_{epoch}.txt, w, encodingutf-8) as f: for s in samples: f.write(s \n) gen.train()生成样本时记得把生成器切到 eval 模式关掉 dropout。生成完再切回 train不然下一轮训练会受影响。这个细节不注意的话训练后期会莫名其妙变差排查半天才发现是模式没切回来。这套东西值不值得投入取决于你的场景对生成质量的要求。如果只是做个 demo 或者短文本生成SeqGAN 的训练成本和调参难度确实比直接微调一个预训练模型高。但如果你需要的是可控的、能针对特定指标优化的生成器而且数据量不大、不想依赖大模型SeqGAN 这套框架仍然是个扎实的起点。源码跑通之后把奖励函数改成你关心的指标往往比换模型结构带来的提升更直接。希望帮到你。本文还有配套的精品资源点击获取
返回列表