☰
从零跑通SeqGAN:策略梯度与蒙特卡洛搜索实现序列对抗生成
2026/9/26 2:30:40 网站建设 项目流程

简介:面向深度学习与生成模型研究者的SeqGAN对抗神经网络Python实现,聚焦序列数据生成。项目将序列生成转化为策略优化问题,通过预言机监督预训练与对抗训练两阶段,帮助读者理解强化学习策略梯度、生成器与判别器博弈等核心机制,适合学习GAN进阶应用、复现序列生成实验的开发者。资源压缩包共13个文件,以6个py源码为主,覆盖生成器、判别器、数据加载、序列生成与rollout采样等模块,另含2个预训练模型参数文件、2张训练曲线图、实验日志、说明文档及示例zip,整体约5.75MB,结构清晰便于直接运行调试。目前已有393人学习浏览。通过阅读源码与运行实验,可掌握SeqGAN从数据预处理、模型搭建到训练调参的完整流程,同时理解对抗性学习与序列建模的交叉应用,为文本、音频等序列生成方向的创新提供可复现的实践基础。

1. 从零跑通SeqGAN对抗神经网络:策略梯度与蒙特卡洛搜索是核心

把GAN从图像搬到文本或序列上时,直接照搬必然会翻车:文本是离散的,每个词是一个token,而argmax操作让梯度无法回传,判别器给出的反馈根本指导不了生成器更新。SeqGAN的做法是把生成器当策略网络、判别器当奖励函数,用策略梯度绕开不可导这一环,再用蒙特卡洛搜索对没写完的句子打分。这套对抗神经网络的架构后来成了序列生成方向的基线,大量后续文本GAN都在它身上做文章。这份Python完整源码把数据预处理、生成器、判别器、rollout、训练循环都打进一个工程里,不用再从论文往代码里翻译。

打开之后你会看到典型的文本GAN工程结构,而不是零散的算法片段。它自带数据和词表构建脚本,意味着你把压缩包解压、装好依赖,就能直接跑通一次完整的训练。适合写过LSTM、却被生成文本重复或发散问题折磨的人,也适合想给对话回复、量化指标序列等场景引入对抗训练但还不知怎么下手的人。下面这五章,是我自己复现这类源码时走过的完整路径,每一步都对应实际会碰到的坑。

2. 原理拆解:为什么序列生成必须绕道强化学习

2.1 离散token让梯度无法回传

传统GAN的核心是生成器输出连续数据,比如图像像素值,判别器输出一个真伪概率,整条链路从判别器到生成器可以端到端反向传播。图像这个场景里,生成器的输出和判别器的输入之间是连续映射,梯度能沿着像素值一路传回生成器参数。

但语言不是这样的。生成器在每一步输出的是一个词表上的概率分布,真正要拿到下一个词,得从分布里采样或取argmax。采样和argmax都不可导,判别器对完整句子的打分传到离散token这一步就断了。就算用gumbel-softmax这类技巧强行让采样过程可导,连续松弛后的分布和你真正想要的分布之间仍有偏差,尤其在序列较长时,误差会累积。

SeqGAN的解法是把生成器看成一个策略网络,把判别器当成奖励函数。它把「生成一个完整句子」重新描述成强化学习里的序列决策问题:生成器在状态S_t(已经写出的前t-1个token)下采取动作a_t(选择下一个token),环境返回奖励R_t。生成器的目标变成最大化期望累积奖励,而期望奖励的梯度可以用策略梯度定理估计:

# 策略梯度的核心更新逻辑,SeqGAN的生成器更新就长这样 # log_probs: 每一步采样token对应的对数概率,形状 [batch, seq_len] # rewards: 每一步token拿到的奖励,形状 [batch, seq_len] # 用奖励加权对数概率,梯度方向让高奖励动作更容易被选中 g_loss = -torch.mean(log_probs * rewards)

代码里log_probs来自生成器在采样token处的分布取值,rewards来自判别器和蒙特卡洛搜索。这里有个容易误解的点:不是把整个句子当成一个奖励,而是让序列里每个token都拿到一个独立奖励,这样后期token的决策也能获得有效的梯度反馈。如果只给整句一个总奖励,前期token几乎学不到东西,梯度方差会非常大。

2.2 生成器与判别器:LSTM写诗,CNN判诗

SeqGAN里的生成器和判别器都不是花哨结构。生成器一般就是一个单层LSTM,输入是前一个token的embedding,输出是词表上的logits;判别器一般是一个text-CNN,用多个不同尺寸的卷积核提取句子特征,然后拼起来做二分类。论文里和大多数开源实现都采用这个配置,因为它在效果和训练速度之间最平衡。

class Generator(nn.Module): def __init__(self, vocab_size, embed_dim, hidden_dim): super().__init__() self.embed = nn.Embedding(vocab_size, embed_dim) self.lstm = nn.LSTM(embed_dim, hidden_dim, batch_first=True) self.fc = nn.Linear(hidden_dim, vocab_size) def forward(self, x, hidden): emb = self.embed(x) # [batch, seq_len, embed_dim] out, hidden = self.lstm(emb, hidden) logits = self.fc(out) # [batch, seq_len, vocab_size] return logits, hidden

这里embed_dim和hidden_dim就是最常见的可调参数。embed_dim控制在64到128之间,hidden_dim在128到256之间,效果都不错;如果数据量不大,把hidden_dim堆到512反而容易让生成器记住训练集里的原句而不是学会泛化。

判别器我一般会写成几组并行的Conv1d,分别提取unigram、bigram、trigram特征,再走一个全连接输出标量。要注意的是判别器每轮训练用的负样本必须来自当前版本的生成器,不能拿一开始预训练采样出来的旧样本反复用——生成器策略一直在变,旧样本会让判别器学到一个过时的决策边界,SeqGAN的判别器就变成在跟历史版本对抗,而不是跟当前版本对抗。

class Discriminator(nn.Module): def __init__(self, vocab_size, embed_dim, seq_len, filter_sizes=(1, 2, 3)): super().__init__() self.embed = nn.Embedding(vocab_size, embed_dim) self.convs = nn.ModuleList([ nn.Conv1d(embed_dim, 64, k, padding=k // 2) for k in filter_sizes ]) self.fc = nn.Linear(len(filter_sizes) * 64, 2) def forward(self, x): emb = self.embed(x).transpose(1, 2) # [batch, embed_dim, seq_len] feats = [torch.relu(conv(emb)).max(dim=-1)[0] for conv in self.convs] feat = torch.cat(feats, dim=-1) return self.fc(feat)

这里filter_sizes里的(1,2,3)就是channel对不同长度n-gram的覆盖,中文按字切分时,1/2/3分别对应一个字、两个字的词、三个字的片段。对英文分词后的序列,也可以把filter_sizes调成(2,3,4),因为英文词本身比汉字更短,太小的卷积核感受野不足。

2.3 蒙特卡洛搜索:给没写完的句子打分

判别器只能给完整序列打分,但生成器每一步都要一个奖励值,这个矛盾怎么解决?SeqGAN选择用蒙特卡洛搜索,又叫rollout策略。假设当前已经生成了前t个token,把这t个token当作固定前缀,让生成器继续按当前策略随机补全到完整长度,重复rollout_num次,然后把每次补全得到的完整序列都丢给判别器打分,最后取平均,作为这前t个token的奖励估计。

def mc_search(gen, disc, prefix_tokens, seq_len, rollout_num): """对已生成前缀做rollout补全,返回平均奖励作为每个token的reward""" rewards = [] for _ in range(rollout_num): seq = prefix_tokens[:] # 复制已生成的token序列 hidden = gen.init_hidden(1) for t in range(len(prefix_tokens), seq_len): logits, hidden = gen(seq[-1:], hidden) # 只输入最后一个token prob = torch.softmax(logits, dim=-1) token = torch.multinomial(prob, 1) # 多项式采样而不是argmax seq = torch.cat([seq, token], dim=-1) score = disc(seq) # 完整序列交给判别器 score = torch.softmax(score, dim=-1)[0, 1] # 取“真样本”概率 rewards.append(score.item()) return torch.mean(rewards)

rollout_num这个参数很有讲究。我在复现时通常取16到32,值越大奖励估计越平滑,但训练时间线性上升。如果取值太小比如4,奖励方差太大,生成器loss会像心电图一样跳,早停判断基本没法做。还要注意一个细节:MC搜索时生成器的forward不需要保留梯度,因为采样出来的token本身不是通过反向传播挑的,这里的目的是估算奖励,不是更新生成器。等拿到奖励之后,再重新走一次带梯度的生成器forward,把对应token的log概率乘以奖励。

这种「先用无梯度采样估算奖励,再回传梯度」的做法,是SeqGAN速度最快的版本。我在第一次复现时没注意,把MC搜索整个放进了autograd计算图,结果显存直接翻了三倍,batch_size只能降到原来的四分之一。

3. 把源码跑起来:环境配置、数据预处理与训练循环

3.1 环境准备:先确认Python版本与依赖

拿到源码包后第一件事不是急着跑main.py,而是先看依赖清单。网上流传的SeqGAN复现有两个血统:TensorFlow 1.x版和PyTorch版。TensorFlow 1.x的版本在Python 3.8以上几乎装不上,如果你下载到的是TF版,建议直接用conda新建一个Python 3.6环境。PyTorch版就好办得多,Python 3.8到3.10都能跑,主要依赖就是torch、numpy、tqdm。

# 建议先建独立环境,避免跟系统Python打架 python -m venv seqgan_env source seqgan_env/bin/activate pip install -r requirements.txt

如果你是在Windows上复现,用conda更稳一点,因为LSTM训练在Windows上对gcc版本不敏感,但有些老代码会偷偷调用nltk的perl脚本,缺perl环境会直接报错。装好之后用python -m pip list检查torch版本,CPU版本也能跑但速度会慢很多,显存不足后面会专门讲。顺便说一句,现在很多python安装教程和vscode python环境配置都会默认帮你把pip源改成清华源,如果装torch比较慢,可以先确认一下源是不是被改过,torch的安装包比较大,用官方源经常卡到超时。

源码的目录结构一般是这样的:

文件/目录作用
data/原始语料和预处理后的token序列
generator.py生成器实现
discriminator.py判别器实现
rollout.py蒙特卡洛搜索实现
data_utils.py词表构建、token映射、batch切分
main.py训练入口,包含预训练和对抗训练
save/模型断点保存目录

3.2 数据预处理:从原始文本到token序列

SeqGAN的训练数据必须满足一个要求:所有样本长度一致。判别器是CNN结构,CNN在PyTorch里处理变长序列要么用padding加mask,要么干脆把序列长度统一。原始语料是每行一句话或一首诗的纯文本,长度参差不齐,必须先做定长处理。

中文语料我一般按字符切,而不是按分词之后的词切。原因有两个:一是字符级词表小,一般也就几千个字符,训练更容易收敛;二是古诗、短句这类文本本身没有明显的分词边界,按字切不会丢信息。英文语料则按空格分词,再建立word2idx。

def build_vocab(lines): chars = set() for line in lines: chars.update(list(line.strip())) chars = sorted(chars) word2idx = {c: i + 2 for i, c in enumerate(chars)} word2idx['<pad>'] = 0 word2idx['<unk>'] = 1 return word2idx

word2idx里pad占0、unk占1,是BERT系词表的常用做法,主要为了后面做embedding时能统一处理。做完词表之后还要把所有句子转成长度一致:超过seq_len的直接截断,不足的用pad补齐。

def pad_sequence(seq, max_len): if len(seq) > max_len: return seq[:max_len] return seq + ['<pad>'] * (max_len - len(seq))

这里max_len直接对应后面的seq_length参数。古诗一句通常五言或七言,如果切的是整首诗,一般设置32到64就够。如果切的是单句,16就够。序列太长不是好事,MC搜索的时间复杂度是rollout_num乘以seq_len,每多一个token,补全和打分就多一轮,训练成本翻倍涨。

3.3 训练循环:预训练、对抗训练与采样参数

SeqGAN的训练不能直接进对抗阶段,否则生成器一开始连一个像样的词都吐不出来,判别器瞬间就能分辨真假,梯度直接消失。标准流程分三步走。

第一步,用最大似然估计MLE预训练生成器。这时的生成器退化成普通语言模型,拿真实语料当监督信号,学习下一token的条件分布。这一步很重要,它让生成器先学会语法的基本骨架。我一般把MLE预训练跑80到120个epoch,直到采样输出看起来是通顺的句子为止。

python main.py \ --mode pretrain_gen \ --epochs 100 \ --batch_size 64 \ --seq_length 32 \ --embed_dim 64 \ --hidden_dim 128

第二步,预训练判别器。用预训练好的生成器采样一批负样本,混入真实样本,训练判别器做一个简单的真假二分类。这一步让判别器有一个说得过去的初始能力,而不是从随机权重开始跟生成器大眼瞪小眼。

python main.py --mode pretrain_dis --epochs 50

第三步,正式对抗训练。每个epoch里,先用当前生成器采样负样本来更新判别器,再用MC搜索估算奖励来更新生成器。两个网络交替更新,生成器每更新一次,判别器也要跟上,否则生成器很容易骗过一个陈旧的判别器。

python main.py --mode train --epochs 200 --rollout_num 16

核心参数我一般这样设置:

参数名常用范围影响
embed_dim32~128太小词义表达不足,太大会过拟合
hidden_dim64~256LSTM容量,决定生成器表达能力
seq_length16~64越长MC搜索越贵,越容易发散
rollout_num8~32奖励估计方差,越大越稳但越慢
temperature0.8~1.0采样多样性,越低越保守,越高越散
pretrain_epoch80~120生成器初始质量,短了后面全崩

其中temperature是最容易忽视的一个。很多源码实现里torch.multinomial的输入softmax概率没有做temperature缩放,导致生成结果千篇一律。你可以在softmax之前把logits除以temperature,大于1时分布变平、生成更发散,小于1时分布变尖、生成更保守。SeqGAN对抗训练阶段我常用0.9左右,太低会诱发复读机问题。

4. 常见问题排查:SeqGAN训练中的五个真实翻车现场

4.1 现象一:对抗训练的loss不降反升

我复现时第一次遇到的是这个:MLE预训练跑得挺好,一旦切到对抗训练,生成器的loss直接从前几个epoch的0.4跳到3.8,然后一路爬升不收住。看生成的文本,句子结构开始崩坏,逐渐出现不通顺的token组合。

原因是MC搜索估算的奖励噪声很大,尤其rollout_num较小时方差更高,策略梯度的更新方向几乎被噪声主导。另一个关键因素是判别器刚切换训练模式时还没稳定,给出的奖励信号本身就是乱的。

解决方法是分两步调整。先把rollout_num从8提到24,reward变得更平滑。然后在策略梯度损失里减去一个baseline,我一般用最近一个batch的奖励均值做baseline,这相当于中心化奖励,能显著压低梯度方差。调整之后loss虽然还是上下浮动,但整体趋势会稳定在一个区间里,而不是一路狂奔。

4.2 现象二:生成器只会复读机,来来回回就那几句

训练跑了一百多个epoch,每次采样出来的诗句都差不多,甚至完全相同。这是明显的模式崩溃。判别器的判断能力跟不上生成器,或者判别器被生成器的高频模板骗过了,生成器发现只要输出那几个高概率模板,就能稳定骗到0.8以上的真样本分数,于是不再探索其他句式。

解决得从判别器下手。第一,降低temperature让采样更多样化,把0.8调到0.95以上;第二,检查判别器训练频率,我遇到过判别器每个epoch只更新一次但生成器更新了三次的情况,对抗严重失衡,调整为每轮生成器更新一次、判别器更新两次后明显改善;第三,给策略梯度加一个最大熵正则项,相当于在损失里加一项对log-prob的熵惩罚,让生成器别把所有概率质量压在同一批token上。

# 带熵正则的策略梯度损失,熵项系数entropy_coef一般取0.05~0.2 policy_loss = -torch.mean(log_probs * rewards) entropy = -(prob * torch.log(prob + 1e-8)).sum(dim=-1).mean() total_loss = policy_loss - entropy_coef * entropy

4.3 现象三:GPU显存爆得比图像GAN还快

这是个很容易懵的坑。batch_size设成64,seq_len设成32,这个规模在图像GAN里根本算不了什么,但SeqGAN却直接OOM。因为MC搜索会把batch按rollout_num扩写:假设batch_size为64,rollout_num为16,实际输入判别器的补全序列数量高达1024条,每条序列长度32。相当于把batch_size翻到1024,显存当然扛不住。

解决方法是把MC搜索的补全过程整体包在torch.no_grad()里,让补全序列的处理不需要梯度图。奖励拿到后,生成器的更新重新走一个小batch的前向计算,提取对应token的log概率。经过这个改动,显存峰值能降到原来的三分之一到四分之一。如果还爆,就把batch_size降到16或32,graph编译也会省不少显存。

4.4 现象四:训练速度慢到怀疑人生

rollout_num设为32、seq_len设成64时,在1080Ti上跑一个epoch要四十多分钟,完全没法调参。原因是MC搜索是逐token采样补全的,每个token都要过一遍生成器的forward,本身是串行的。

我的做法是先做小规模验证,把seq_len压到16,rollout_num压到8,数据量裁到原来的五分之一,确认代码逻辑没有bug、loss有下降趋势,再调大参数上全量数据。另一个很实用的技巧是,判别器打分阶段用一个小batch循环代替一次性forward全量补全序列,虽然总计算量没变,但能让GPU显存占用变小,避免内存换页拖慢整机速度。

4.5 现象五:判别器acc稳在100%,生成器完全躺平

训练到中期,判别器准确率刷到100%,生成器的loss冻结在一个常数附近,再多个epoch都纹丝不动。这是判别器过强的典型表现。判别器已经把真实样本和生成样本的特征彻底分开,生成器再怎么改输出,判别器都能一眼识破,等于梯度消失了。

解决思路是按顺序做三件事。先减少判别器预训练轮数,从50砍到20;再把生成器的学习率适当调大,让它在对抗中更有竞争力;最后也可以尝试在判别器输入上加一点dropout或噪声,人为降低判别器的绝对判断能力,给生成器留一点生存空间。

5. 换到自己的数据上:格式改造与生成质量验证

5.1 数据格式改造:把任意序列转成SeqGAN能吃的token

SeqGAN不挑数据,只要是序列化的符号都能训练,但格式需要你自己换。古诗是字符级序列,如果换成量化数据,面对的是一串连续浮点数,本站的关键问题就变成了「怎么把连续数变成离散token」。我一般先把原始序列做一阶差分,再把差分值按分位数分箱。分箱数量取50到100之间,太少丢失趋势细节,太多让词表过大、样本过于稀疏。

def series_to_tokens(series, n_bins=50): diff = np.diff(series) / series[:-1] # 计算收益率序列 q = np.percentile(diff, np.linspace(0, 100, n_bins + 1)) tokens = np.digitize(diff, bins=q[1:-1]) # 每个差分映射成bin索引 return tokens.tolist()

这段代码把连续收益率离散化成symbol序列,SeqGAN就可以把历史走势当作文本一样学习分布。生成的token序列再映射回收益率分箱中心值,就能得到合成走势。但要提醒自己:这生成的是「分布上相似的样本」,不是预测下一天的行情,把它当数据增强工具可以,直接拿去做量化信号会后悔的。我见过不少人拿这类对抗生成序列做策略回测增强,前提是你严格控制样本泄露,否则结果全是幻觉。

5.2 验证生成质量:别只用loss当标准

SeqGAN训练过程中,loss下降不代表生成质量变好,我只看三个指标。首先看人工采样文本,每隔固定步数把生成器采样出的句子打印出来,通顺度是最直观的;其次看n-gram多样性,重复率高的模型多半已经mode collapse;最后跑一遍留存集上的困惑度,数值陡增说明生成器在捏造没学过的模式。

def distinct_n(sentences, n=2): all_ngrams = set() total = 0 for s in sentences: tokens = s.split() ngrams = [tuple(tokens[i:i + n]) for i in range(len(tokens) - n + 1)] all_ngrams.update(ngrams) total += len(ngrams) return len(all_ngrams) / max(total, 1)

distinct_n接近1说明句子之间差异大,接近0说明高度重复。对古诗生成任务,我要求distinct_2大于0.6才算合格。还可以把生成样本做python数据分析与可视化,比如统计每个token位置上的分布熵,如果你发现后半个序列的熵比前半段低很多,说明生成器越往后越保守,经常提前锁死后续内容。

从那以后我每次换数据集,都会强制走一遍同样的验证清单:先把MLE预训练跑通、看一眼采样输出,再做对抗训练;对抗阶段每隔固定步数保存一组采样样本回看,distinct_n低于阈值立刻停。毕竟生成模型看的是样本,不是训练曲线上的数字。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询