ARTICLE DETAIL

资讯详情

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

用LSTM训练《鹿鼎记》生成中文小说:从数据处理到参数调优

用LSTM训练《鹿鼎记》生成中文小说:从数据处理到参数调优 简介一套以金庸《鹿鼎记》为语料的LSTM小说文本生成项目适合具备基础Python与深度学习概念的在校学生、毕设开发者及NLP爱好者用于学习循环神经网络与字符级文本生成。项目先用requests爬取金庸网鹿鼎记目录及正文保存为纯文本语料受硬件限制取前5万字符经过去重、排序并构建词典到整数的映射按每句40个字符组织训练序列最终训练字符级LSTM模型完整链路覆盖数据采集、清洗、编码、训练与文本采样代码均可直接运行。压缩包共5个文件包含2个Python脚本分别负责爬取语料与训练生成另有1份说明文档、1个txt小说语料和1个hdf5训练权重文件整体仅18.78MB结构紧凑、便于快速上手。训练权重可直接加载测试生成效果也可自行修改语料或调整超参数继续实验既能用于毕业设计、课程设计也可作为自然语言处理入门练手。目前已有192人浏览学习代码经上传者验证可运行下载后可按说明文档依序执行。1. 用《鹿鼎记》教 LSTM 写小说这条路到底通向哪里基于鹿鼎记的数据集用 LSTM 学写小说是一个把 python 源码、文档说明和数据整理齐的 NLP 练手项目。它要解决的问题很具体给 LSTM 一部百万字级的中文长篇小说让模型按字符一个接一个地往外吐文本最终读出金庸的用词习惯和句子节奏。这个任务没有想象中那么吃算力把全文按字符拆开后一张普通显卡甚至强一点的 CPU 都能完成训练。适合刚入门的深度学习者建立「序列数据怎么进、loss 怎么降、生成结果为什么总是重复」这些直觉也适合想评估传统 RNN 在文本生成上还行不行的研发人员拿来做基线。下面这套方案不是流水账每一步都能直接复现。2. 把《鹿鼎记》变成训练样本清洗、切块与字符编码2.1 从 txt 到干净正文编码探测与空白清理先处理数据。网络上流传的《鹿鼎记》txt 版本编码不统一常见的是 utf-8 和 gb18030 两种直接按 utf-8 硬读很容易报UnicodeDecodeError。我一般先做编码探测再处理空白字符。import re def load_and_clean(path): text None for encoding in (utf-8, gb18030): try: with open(path, r, encodingencoding) as f: text f.read() break except UnicodeDecodeError: continue if text is None: raise ValueError(无法识别文件编码请手动确认源文件格式) # 去掉全角空格与普通空格、Tab保留换行作为段落边界 text re.sub(r[ \t\u3000], , text) # 把连续换行压缩成单一换行避免空白页进入训练语料 text re.sub(r\n{2,}, \n, text) return text逻辑说明全角空格\u3000在中文排版里极其常见如果不清掉它会在字符表里占一个位置白白增加输出维度。连续换行往往对应原书的空白页或章节间隔压缩成单个\n后模型仍然能看到「换行 段落结束」这个信号但不会被多余空行干扰。章节标题怎么处理取决于你的目标。如果只是想学叙事语言删掉「第X回」这类行可以让训练数据更干净如果希望模型学会回目结构和「且听下回分解」这种固定表达就保留。注意这里直接对文件做内存级读取鹿鼎记全文百万字级别完全放得下不需要流式读取。2.2 固定序列切块seq_len 与 stride 的配合LSTM 不接收整本书它只看固定长度的字符窗口。切块时有两个参数seq_len是窗口长度stride是每次移动的步长。下面这段代码把所有文本变成「输入-目标」样本对seq_len 64 stride 32 text load_and_clean(ludingji.txt) samples, targets [], [] for i in range(0, len(text) - seq_len, stride): chunk text[i:i seq_len] samples.append(chunk) targets.append(text[i 1:i seq_len 1]) # 右移一位逐字符预测逻辑说明第 i 个样本的输入是text[i:iseq_len]目标是把整个窗口右移一个字符让模型在每个时间步预测下一个字符。当 stride 小于 seq_len 时相邻样本之间有重叠相当于把百万字语料复制扩展了好几份。参数说明seq_len64是中文小说生成常用的起步值大约能覆盖两到三个短句。stride 一般取 seq_len 的一半这样样本量约为(总字数 - seq_len) / stride百万字语料能切出三万个左右的样本stride 设得太小会让相邻样本高度相似训练变慢还容易过拟合。2.3 字符级编码中文小说为什么先走 char-level中文文本建模有两种路子按词切分或按字符切分。按词需要分词工具词表通常几万到几十万还会遇到未登录词字符级直接把每个汉字和标点当作一个单位整部《鹿鼎记》去重后也就几千个字符。对「学语言风格」这件事字符级天然能捕捉字与字的搭配输出层也小得多。chars sorted(set(text)) char2idx {c: i for i, c in enumerate(chars)} idx2char {i: c for c, i in char2idx.items()} vocab_size len(chars) import torch from torch.utils.data import TensorDataset, DataLoader X torch.tensor([[char2idx[c] for c in s] for s in samples], dtypetorch.long) Y torch.tensor([[char2idx[c] for c in t] for t in targets], dtypetorch.long) dataset TensorDataset(X, Y) loader DataLoader(dataset, batch_size128, shuffleTrue)逻辑说明chars sorted(set(text))会把所有出现过的字符收集起来包括汉字、中英文标点和换行符。映射装进两个 dict训练时用char2idx把字符转索引生成时用idx2char把索引还原成字符。这里有个容易忽略的点vocab_size不要手动指定直接用len(chars)。某次我偷懒写死 5000结果实际字符表比这大一点索引越界直接崩了。字符级模型的词表就是字符集合本身动态生成最稳。3. 用 PyTorch 搭一个能跑通的 LSTM 小说生成模型3.1 网络结构Embedding、双层 LSTM 与输出投影文本不能直接进 LSTM先要通过 Embedding 把每个字符索引映射成一个稠密向量再用两层 LSTM 做序列建模最后接一个全连接层输出每个字符的分数。结构如下import torch.nn as nn class NovelLSTM(nn.Module): def __init__(self, vocab_size, embed_size256, hidden_size512, num_layers2, dropout0.3): super().__init__() self.embedding nn.Embedding(vocab_size, embed_size) self.lstm nn.LSTM(embed_size, hidden_size, num_layers, batch_firstTrue, dropoutdropout) self.fc nn.Linear(hidden_size, vocab_size) def forward(self, x, hiddenNone): emb self.embedding(x) out, hidden self.lstm(emb, hidden) logits self.fc(out) # shape: [batch, seq_len, vocab_size] return logits, hidden逻辑说明Embedding 把索引转成 256 维向量LSTM 接收[batch, seq_len, embed_size]输出每个时间步的隐层状态全连接层把 512 维隐层投影回字符表大小得到每个位置对所有字符的分数。hidden在训练时用不到但生成时需要逐段传递。参数说明embed_size256、hidden_size512是中等配置显存吃紧可以都降到 256。num_layers2表示堆叠两层 LSTM第一层能捕捉字词级搭配第二层在更高抽象级别建模句式。nn.LSTM的 dropout 参数只在层数大于 1 时对中间层生效单层时设了也被忽略。实例化时一行搞定model NovelLSTM(vocab_sizelen(chars))3.2 训练循环交叉熵、Adam 与梯度裁剪训练目标是用前 63 个字符预测第 64 个字符整个序列每个位置都做一次预测。交叉熵损失会对每个时间步求平均数值大小跟字符表规模有关。鹿鼎记字符表五六千字随机初始化的 loss 大约在ln(5000) ≈ 8.5左右训练后能掉到 2 以下说明模型已经记住了常用搭配。criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) model.train() for epoch in range(10): for x, y in loader: logits, _ model(x) loss criterion(logits.reshape(-1, vocab_size), y.reshape(-1)) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() print(fepoch {epoch 1}, loss {loss.item():.4f})逻辑说明logits形状是[batch, seq_len, vocab_size]reshape 成[batch*seq_len, vocab_size]才能和展平后的目标索引算交叉熵。梯度裁剪是长序列训练的关键LSTM 在反向传播时梯度容易指数级增长裁剪到 5.0 相当于给梯度上了保险丝。参数说明Adam 初始学习率 1e-3 对字符级模型来说足够如果 loss 震荡得厉害降到 3e-4 或 5e-4。看到 loss 数值不再变化时优先考虑调序列长度和模型宽度而不是继续堆训练轮数。3.3 生成函数温度参数与 top-k 采样训练完成后生成是逐个字符往外蹦的。最粗暴的写法是每次取分数最高的字符但这会让文本陷入无限重复。更好的做法是用温度缩放分布再按概率采样。def generate(model, seed, length200, temperature0.8, top_k30): model.eval() chars list(seed[-seq_len:]) if len(seed) seq_len else list(seed) # 如果 seed 太短用换行补齐到 seq_len if len(chars) seq_len: chars [\n] * (seq_len - len(chars)) chars with torch.no_grad(): for _ in range(length): x torch.tensor([[char2idx[c] for c in chars[-seq_len:]]], dtypetorch.long) logits, _ model(x) next_logits logits[0, -1] / temperature values, indices next_logits.topk(top_k) mask next_logits values[-1] next_logits[mask] -float(inf) prob torch.softmax(next_logits, dim-1) next_idx torch.multinomial(prob, 1).item() chars.append(idx2char[next_idx]) return .join(chars)逻辑说明logits[0, -1]取出最后一个时间步的输出只预测下一个字符。除以 temperature 会改变分布的尖锐程度温度低时概率集中在高分字符上输出保守温度高时分布变平更容易选出冷门字。top-k 过滤把分数低于第 k 名的候选全部屏蔽避免生成「生僻错字」。最后用torch.multinomial采样而不是argmax这是保留生成随机性的关键。参数说明temperature 0.8、top_k 30 是中文小说场景的稳妥组合想要更稳就调成 0.6 和 20想探险可以试 1.0 和 50。seed 我习惯直接抄原文一句话比如「韦小宝笑道」。4. 参数怎么设序列长度、batch、学习率与 dropout 的取舍4.1 序列长度先定上下文窗口决定连贯上限代码能跑通后第一件事是定 seq_len。它对生成质量的影响比网络宽度更直接。下表是我在不同 seq_len 下观察到的典型现象seq_len生成现象观察显存/训练成本16三五个字内语法正确词序随意飘低CPU 可跑32短语通顺高频虚词开始占主导低64短句内部基本通顺局部有呼应中建议用 GPU128能记住前一句的主语跨句略有关联较高训练明显变慢seq_len64 是最划算的起点。上升到 128 时模型能利用的上下文更长但 LSTM 的反向传播时间步也翻倍训练耗时会显著增加。如果显存只有 4G老老实实从 32 或 64 开始先把流程跑通再谈质量。4.2 batch、学习率与梯度裁剪稳定训练的三角组合字符级模型的 batch 影响的是 loss 曲线的平滑程度。batch 太大收敛变慢太小则 loss 噪声大。128 是常用的平衡点CPU 训练时用 64 也能接受速度差别不大。accum_steps 4 for step, (x, y) in enumerate(loader): logits, _ model(x) loss criterion(logits.reshape(-1, vocab_size), y.reshape(-1)) loss / accum_steps loss.backward() if (step 1) % accum_steps 0: nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() optimizer.zero_grad()逻辑说明这段代码实现了梯度累积每 4 个 batch 更新一次参数。loss 除以 accum_steps 是为了让累积梯度的量级和单次更新等价。显存不够但不想降 batch 时这是最直接的补救手段。参数说明学习率 1e-3 跑 2000 步后如果 loss 下降变缓把 lr 降到 3e-4 再继续。梯度裁剪固定 5.0 即可不用频繁调整。4.3 网络深度与 dropout数据量撑不起太深的 LSTM很多人一上来就想堆 4 层 LSTM觉得越深越强。但百万字级字符语料对 LSTM 来说并不算多层数加深容易过拟合训练还慢。2 层 512 维已经是这个数据量级的上限附近。dropout 设 0.3 到 0.5 能有效缓解过拟合但注意推理时必须调用model.eval()否则 dropout 还在随机丢神经元生成结果会不稳定。5. 避坑用《鹿鼎记》训练 LSTM 的 5 条血泪经验5.1 全角与半角标点混排模型学到两种逗号现象生成的文本里同时出现「」和「,」甚至一个句子混用两种标点。原因原始 txt 里全角半角标点混排字符集合里同时存在中文逗号和英文逗号两个位置。模型把它们当两个独立字符处理输出时随机选用。解决清洗阶段做一次标点统一。text text.replace(,, ).replace(!, ).replace(?, )这样字符表里只保留全角版本生成的标点风格立刻统一。5.2 固定长度切块切断引号对话格式永远学不会现象生成文本里的双引号经常只开不闭对话片段会突然卡住。原因按固定字符数硬切时序列可能从一句对话的中间断开。模型看到的是「前引号没有后引号」的残缺样本学不到引号配对规律。解决切块时优先对齐句子边界。简单做法是先按。\n把文本切成句子再按句号边界拼回 seq_len 长度的块句子太长时回退到暴力切块避免样本丢失。这个策略会显著提升对话生成质量。5.3 loss 卡在 1.8 附近生成全是高频虚词现象训练后期 loss 不再下降生成结果满屏「的」「了」「是」。原因字符级模型天然偏向高频字符序列长度过短时低频实词缺少足够上下文很难被正确激活。解决先把 seq_len 提到 128 试一轮再看生成情况。如果仍旧高频虚词泛滥生成时把 temperature 调到 1.0 到 1.2提高冷门候选被选中的概率。也可以对 top_k 采样做微调把 top_k 从 30 提到 50给低频字更多机会。5.4 显存随序列长度暴增batch 调小照样崩现象seq_len 从 64 调到 128 后显存直接 OOM把 batch 降到 16 还是崩。原因LSTM 的 BPTT 需要保存所有时间步的中间状态显存占用随 seq_len 线性增长。batch 对显存的影响远不如 seq_len 大。解决回到 seq_len64用梯度累积模拟大 batch。训练脚本里加一句torch.cuda.empty_cache()在每轮循环结束后清缓存也能缓解碎片化问题。5.5 CPU 训练半天 loss 不动先跑 500 步冒烟测试现象挂机训练两小时loss 几乎没掉。原因字符级 LSTM 收敛慢但更常见的问题是数据没进对或者学习率设得太低。解决第一次跑通流程时只取前 10 万字做冒烟测试训练 500 步看 loss 是否从 8 字头往下走。确认流程没问题后再全量挂机。这个习惯帮我避开了好几次「训练一晚上第二天发现字符映射错了」的尴尬。6. 生成结果更像金庸温度微调、标点修正与验证方法温度是生成阶段最值得花时间调的一个旋钮。0.6 附近生成的文本最稳句子基本通顺但略显保守适合拿来做「像不像小说」的初版测试0.8 到 1.0 之间是最好的创作区既有可读性又不乏意外搭配超过 1.2 后句子结构开始松散偶尔能蹦出惊艳的词组但整体不可控。我一般固定 top_k30只改 temperature用同一个 seed 生成 5 段对比挑流畅度最稳定的区间。生成之后还有一道后处理工序。LSTM 不擅长处理排版细节常见问题是连续换行过多、半角标点混入、引号不成对。这时候靠一个小函数兜底def postprocess(text): text re.sub(r\n{3,}, \n\n, text) text text.replace(,, ).replace(!, ).replace(?, ) if text.count(”) % 2 1: text ” return text引号计数配对只是近似处理真正严谨的引号修复需要状态机但生成场景下能保证每段文本不出现肉眼可见的标点残缺就够了。验证模型有没有过拟合我的习惯是留出最后 5% 的样本不参与训练单独算 validation loss。训练 loss 一路下降、validation loss 却在 3000 步后掉头上升说明模型开始背数据而不是学语言此时把 hidden_size 从 512 降到 256或把 dropout 从 0.3 提到 0.5重新训练。还有一个我很喜欢的小验证法把「且听下回分解」作为 seed观察模型接出来的下一回开头是否接近原文回目的语气这一步能快速看出模型学的到底是「文本结构」还是「碎词拼接」。我自己第一次跑这个方案时贪心把 seq_len 拉到 256显存崩了两次后来才明白字符级 LSTM 的配置要克制与其堆长度不如把温度调明白。希望帮到你。本文还有配套的精品资源点击获取
返回列表