news 2026/9/24 1:06:06

用LSTM让《鹿鼎记》学会写小说:字符级文本生成实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
用LSTM让《鹿鼎记》学会写小说:字符级文本生成实战

简介:基于金庸《鹿鼎记》全文数据的LSTM文本生成项目,提供了一套完整可运行的代码与说明,适合自然语言处理入门、毕业设计或课程设计参考。项目覆盖数据爬取到模型训练的全流程:GetLu.py负责抓取小说章节并保存为txt,Word_LSTM.py实现数据预处理、词典映射与LSTM模型训练,训练时取前5万字符、以40个字符为一句构造输入序列,另附README说明文档、清洗后的鹿鼎记文本数据及训练完成的hdf5权重文件。压缩包共5个文件,包含2个Python脚本、1份Markdown说明、1个txt数据集和1个模型权重文件,整体约18.78MB,结构紧凑便于复现。目前已有192人浏览学习,代码经测试运行成功并支持远程教学,适合想了解字符级语言模型、序列生成及文本风格模仿的读者进阶提升。

1. 让 LSTM 读《鹿鼎记》写小说:这个标题到底值得不值得做

先别急着把“基于鹿鼎记的数据集”和“LSTM”当成两个割裂的东西。这个标题本质上在描述一个最小闭环:拿一部百万字级别的中文公版小说当语料,用 Python 搭一个字符级的 LSTM 模型,训练之后用温度采样让它“续写”出风格接近原文的段落。它不是什么巨头产品,却是一个能在一台普通笔记本上跑完的完整项目,而且正好把数据清洗、序列建模、训练调参、文本生成四条线全部串起来。

这个项目能解决什么问题?对想从分类任务往生成任务跳的人来说,它是最好的过渡——LSTM 不需要大显卡,不需要分布式,几百万字符喂进去就能看见 loss 往下掉;对想理解“序列预测”本质的人来说,写小说和做时间序列预测在代码层面高度同构:都是给定前面一堆值,预测下一个值。生成的小说是一个能直观检查的“预测结果”,比看股票曲线更有反馈感。适合读这篇东西的人,是那些已经会用 Python、但还没完整跑通过一个文本生成模型,正犹豫先拿 RNN 还是 Transformer 入手的开发者。我先给结论:这个方向值得做,但坑比想象中多,下文每一章都是照着复现的顺序写的。

2. 语料预处理:把《鹿鼎记》切成 LSTM 能吃的定长序列

2.1 拿到 txt 后先做三件事:编码、去全角空格、压缩空行

市面上能找到的《鹿鼎记》文本大多是 TXT 格式,但格式五花八门。有的是每行一句的“网文排版”,有的是整章一个超长段落,还有的混入了大量全角空格和零星 OCR 错字。我的习惯是,第一步不碰任何模型代码,先把原始文档读进来,看一眼字符分布再动手。

import re from pathlib import Path raw = Path('鹿鼎记.txt').read_text(encoding='utf-8') print('原始字符数:', len(raw)) # 第1件事:去掉全角空格和 \r,统一成 \n 换行 text = raw.replace('\u3000', '').replace('\r', '') # 第2件事:把行内连续空白压成一个空格,避免“韦 小 宝”这种残次 text = re.sub(r'[ \t]+', ' ', text) # 第3件事:把连续 3 个以上的换行压成 2 个,保留段落边界 text = re.sub(r'\n{3,}', '\n\n', text) Path('鹿鼎记_clean.txt').write_text(text, encoding='utf-8') print('清洗后字符数:', len(text))

上面这段的逻辑是:\u3000是中文全角空格,在旧排版文本里几乎必然出现,不删掉会让词表无谓膨胀;\r是 Windows 老文本的换行符残留,混在\n里会让后面按行切分时出现空串;把行内连续空格压缩,是为了避免“金庸 著”这种排版残留影响字符统计。这里有个原则——清洗规则宁少勿多,只处理确定有害的噪声,别去做繁体转简体、也别用词典修正错字,因为 LSTM 是字符级学习,它自己能容忍一定噪声,清洗过度反而会破坏原文的用字统计特征。

我一般会在这步跑一个字符频次统计,确认文本里没有大段看不懂的乱码区块。操作方法很简单:collections.Counter(text)取出现次数最高的 20 个字符看一眼,如果是“的、了、道、说、他、你”这类常见字,就说明文本基本干净。这一步不写进脚本也行,但强烈建议跑一次,因为后续所有建模都建立在词表质量上,而这一步最容易被跳过。

2.2 按“字”建词表:中文生成用小 vocab 比用 jieba 分词更稳

文本生成项目第一道分岔路:分词还是分字。常见的做法是分字,不是分词。原因非常实际:现在 NLP 里用的 jieba 分词面向的是理解任务(分类、NER),分词结果会让词表膨胀到几万甚至十几万;而《鹿鼎记》全文就是一个作者的语言习惯集合,常用汉字加标点撑死一万上下。字符级建模的词表通常只有 6000~8000,这意味着最后全连接层的参数少一个数量级,LSTM 学起来负担小得多。

from collections import Counter char_counter = Counter(text) # 过滤掉只出现 1 次的字符,避免把噪声学进 embedding vocab = [ch for ch, cnt in char_counter.items() if cnt >= 2] vocab = ['<PAD>', '<UNK>', '<BOS>', '<EOS>'] + vocab char2idx = {ch: i for i, ch in enumerate(vocab)} idx2char = {i: ch for ch, i in char2idx.items()} print('词表大小:', len(vocab))

这段代码的要点在于特殊符号的设计。<PAD>用于把 batch 内序列对齐到相同长度;<UNK>兜底那些被过滤掉的生僻字;<BOS><EOS>这次用不上,但保留它们有两个好处:一是训练时可以显式告诉模型“一句话从哪里开始、到哪里结束”,二是以后想换 GPT 风格模型时词表不用重建。实际训练中,过滤阈值cnt >= 2对一部小说是合适的,如果语料更大可以调到 5,但《鹿鼎记》这个量级没必要。

这里还想纠正一个新手常踩的误区:不要用 one-hot 向量直接喂 LSTM。词表 8000,one-hot 就是 8000 维的稀疏向量,而 embedding 层只需 128 维稠密向量就能表达字符间相似度,训练速度和最终效果都明显更好。所以下面的模型结构里,第一层一定是nn.Embedding

2.3 滑动窗口造样本:seq_len、step 怎么配合才不会让验证集泄漏

文本生成任务的训练样本形式是“给定前 n 个字,预测第 n+1 个字”。造样本的通用做法是滑动窗口:用seq_len个字做输入,后面错开一位的seq_len个字做标签,窗口按固定步长滑动。这里有两个参数要一起调:seq_len(窗口长度)和step(滑动步长)。

seq_len = 64 step = 8 xs, ys = [], [] for para in text.split('\n'): para = para.strip() if len(para) < 2: continue ids = [char2idx.get(c, char2idx['<UNK>']) for c in para] for i in range(0, len(ids) - seq_len, step): x = ids[i:i + seq_len] y = ids[i + 1:i + seq_len + 1] if len(x) == seq_len and len(y) == seq_len: xs.append(x) ys.append(y) print('样本总量:', len(xs))

逻辑说明:以段落为单位切分而不是把全文拼成一个超长序列,是为了让每个样本都保持语句相对完整,避免窗口跨过章节边界造成莫名其妙的拼接。x是输入,y是输入整体右移一位的结果,第 t 个位置的标签就是第 t+1 个位置的字符,这样每个样本内部天然构成 64 组 (输入序列, 下一个字符) 的监督对。

参数说明:seq_len=64是一个性价比很高的默认值。太短(比如 16)模型只能学到词语搭配,学不到句间逻辑;太长(比如 256)会显著增加 LSTM 时间步数和显存占用,而且《鹿鼎记》的句子平均长度也就 20 字上下,64 已经覆盖了小半个段落。step=8意味着相邻窗口重叠 56 个字,这样一段 200 字的段落能产出约 24 个样本,数据量放大好几倍,又不至于像step=1那样让相邻样本几乎完全相同,导致训练集内部高度冗余。

样本造好后要按顺序切分数据集,这里有个隐蔽的坑:不要用随机切分。小说文本前后文有关联性,随机切分会让训练集和验证集里出现大量“神似”的片段,验证集就失去了意义。常规做法是按顺序切,前 80% 段落做训练、后 20% 做验证:

split = int(len(xs) * 0.8) train_x, train_y = xs[:split], ys[:split] val_x, val_y = xs[split:], ys[split:] print('训练样本:', len(train_x), '验证样本:', len(val_x))

3. 用 PyTorch 从零搭 LSTM 生成模型:网络结构、损失函数与训练循环

3.1 Embedding + 双层 LSTM + Linear:这三层为什么够用

文本生成模型在 PyTorch 里写起来特别短,核心结构就三块:Embedding 把字符 ID 变成向量,LSTM 在序列上做时间步递推,Linear 把隐状态映射回词表大小。我见过很多新手在这一步直接抄 Transformer,其实对百万字级别的小说语料,LSTM 是更务实的起点——参数少、收敛快、对乱序数据不敏感,而且源码逻辑一目了然,出问题容易定位。

import torch import torch.nn as nn class CharLSTM(nn.Module): def __init__(self, vocab_size, embed_dim=128, hidden_size=256, num_layers=2, dropout=0.3): super().__init__() self.embed = nn.Embedding(vocab_size, embed_dim, padding_idx=0) self.lstm = nn.LSTM(embed_dim, hidden_size, num_layers, batch_first=True, dropout=dropout) self.fc = nn.Linear(hidden_size, vocab_size) self.dropout = nn.Dropout(dropout) def forward(self, x, hidden=None): # x: [batch, seq_len] emb = self.dropout(self.embed(x)) # [batch, seq_len, embed_dim] out, hidden = self.lstm(emb, hidden) # out: [batch, seq_len, hidden_size] logits = self.fc(out) # [batch, seq_len, vocab_size] return logits, hidden

这里给每个参数一个明确的名分。embed_dim=128:字符向量的维度,128 对 8000 词表足够,升到 256 收益很小但参数翻倍。hidden_size=256:LSTM 隐状态宽度,这决定模型记忆能力,256 是文本生成里最常用的甜点值,再大就容易在小数据上过拟合。num_layers=2:两层 LSTM 可以捕捉到“字符→词语→句间”的部分层级关系,三层以上在单机 CPU/入门 GPU 上收益迅速衰减。dropout=0.3:只加在层间和 embedding 输出上,这是 PyTorch LSTM 自带dropout参数的职责范围,注意单层 LSTM 传dropout不会生效,这是源码实现决定的,不用纠结。

为什么要用 2 层而不是 1 层?我自己的体验是:1 层 LSTM 生成出来的句子语法正确,但连续几句话之间几乎没有情节衔接;2 层之后模型更容易记住“刚才在说哪个人”,生成结果开始有短篇叙事的样子。这背后的直觉是,第一层做局部语法建模,第二层在更高抽象级别维护上下文状态,两层各司其职。

3.2 交叉熵和 reshape 对齐:loss 计算里最容易翻车的维度问题

训练 LSTM 生成模型,损失函数用交叉熵没得选。但新手写 loss 这行时十有八九会碰到维度报错,因为模型输出是四维视角下的三维张量[batch, seq_len, vocab_size],而标签是二维[batch, seq_len]。PyTorch 的CrossEntropyLoss期望输入是[N, C]形状,所以必须先把序列维度和 batch 维度合并:

def compute_loss(logits, targets): # logits: [batch, seq_len, vocab_size] # targets: [batch, seq_len] V = logits.size(-1) loss = nn.CrossEntropyLoss()( logits.reshape(-1, V), # [batch*seq_len, vocab_size] targets.reshape(-1) # [batch*seq_len] ) return loss

这段代码的逻辑是:把每个时间步的预测都当成独立分类问题,模型在batch*seq_len个位置上各自预测一个字符,再和真实字符做交叉熵。这里有个值得说透的点:虽然我们把每个位置当成独立样本计算 loss,但 LSTM 的前向传播是串行的,第 t 步的隐状态携带了前 t-1 步的信息,所以误差反传时梯度依然能沿着时间步流动,这正是“BPTT(时间反向传播)”的含义。

有人会问:为什么不把标签做 one-hot 和 logits 算 MSE?因为交叉熵直接优化概率分布,梯度更陡、收敛更快;MSE 把分类问题当成回归问题,会被高频字符“的、了”带偏,生成的文本会更平更呆。这个区别在训练曲线上能直接看出来,MSE 的 loss 值会以极小步长缓慢下降,而交叉熵的下降肉眼可见。

3.3 训练循环的四个参数:lr、梯度裁剪、dropout、epochs

训练循环本身不难,真正决定成败的是几个细节参数。我直接给一个完整可跑的训练代码块,然后逐个参数讲清楚为什么这么设。

import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset train_ds = TensorDataset(torch.tensor(train_x), torch.tensor(train_y)) val_ds = TensorDataset(torch.tensor(val_x), torch.tensor(val_y)) train_loader = DataLoader(train_ds, batch_size=64, shuffle=True) val_loader = DataLoader(val_ds, batch_size=64, shuffle=False) model = CharLSTM(vocab_size=len(vocab)) optimizer = optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', patience=2, factor=0.5) best_val_loss = float('inf') for epoch in range(30): model.train() total_loss, total_acc = 0, 0 for x, y in train_loader: optimizer.zero_grad() logits, _ = model(x) loss = compute_loss(logits, y) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) # 梯度裁剪 optimizer.step() total_loss += loss.item() * x.size(0) total_acc += (logits.argmax(-1) == y).float().mean().item() * x.size(0) model.eval() val_loss, val_acc = 0, 0 with torch.no_grad(): for x, y in val_loader: logits, _ = model(x) loss = compute_loss(logits, y) val_loss += loss.item() * x.size(0) val_acc += (logits.argmax(-1) == y).float().mean().item() * x.size(0) train_loss = total_loss / len(train_ds) val_loss = val_loss / len(val_ds) scheduler.step(val_loss) if val_loss < best_val_loss: best_val_loss = val_loss torch.save(model.state_dict(), 'lstm_鹿鼎记_best.pt') print(f'epoch {epoch+1:02d} | train_loss {train_loss:.3f} | ' f'val_loss {val_loss:.3f} | best {best_val_loss:.3f}')

这里四个参数值得单独拿出来说。

lr=1e-3ReduceLROnPlateau:Adam 的默认学习率 1e-3 在这个任务上适用,但训练后期 loss 会在一个平台期来回震荡,直接用固定 lr 会卡住。所以我给 scheduler 设了patience=2,意思是验证集 loss 连续两轮不创新低就把学习率减半,这个机制比手动调 lr 省心得多。

clip_grad_norm_(..., 5.0):LSTM 训练中梯度范数偶尔会暴涨,这是 RNN 家族的通病。裁剪到 5.0 的意思是,如果梯度范数超过 5,就整体缩放到 5。不设这个参数,训练 20 轮左右可能出现某个 batch 直接把 loss 打回初始值的“返祖现象”,前期全部白练。

weight_decay=1e-4:L2 正则,对 8000 词表的全连接层有稳定作用,设太大(比如 0.01)会让生成结果趋同,所有句子都往高频词上靠。

batch_size=64:在 CPU 上跑这个值偏大,但 GPU 上完全没问题。如果显存不够,不要直接缩小 batch,改用梯度累积:每 4 个 batch 累加一次梯度再 step,效果等价于 batch_size=256,还省显存。

3.4 参数速查表:自己跑一遍时会用到的默认值

参数默认值说明调整方向
seq_len64输入序列长度生成长段落时调大到 128,显存吃紧时调小到 32
step8滑动窗口步长数据量不足时调到 4,步长为 1 会导致样本高度重合
embed_dim128字符向量维度语料大、有 GPU 时可上调到 256
hidden_size256LSTM 隐状态维度想增强记忆能力优先调这个,其次才调层数
num_layers2LSTM 层数1 层生成太散,3 层过拟合风险大
dropout0.3正则化强度过拟合时上调到 0.5
lr1e-3初始学习率loss 震荡剧烈时降到 5e-4
clip5.0梯度裁剪阈值训练不稳定时收紧到 1.0
epochs30训练轮数看验证集 loss,不涨就早点停

看过这张表你大概会有个感觉:LSTM 文本生成没有魔法参数,九成参数在合理区间内都能收敛,真正影响生成质量的只有两个——seq_len 决定模型能看多远,temperature 决定生成有多“敢说”。后者在下一章展开。

4. 验证与生成:loss 曲线别只看训练集,temperature 才是玄学

4.1 early stopping:val loss 连续五轮不降就存 best 模型

训练代码里我已经顺带写了 best 模型保存逻辑,但这只是 half 的功夫。实践中更常见的问题是:训练集 loss 一路降到 1.2 左右,看起来很漂亮,但验证集 loss 从第 15 轮开始反弹,这就是教科书式的过拟合。文本生成任务的过拟合和分类任务不太一样——分类过拟合是准确率不再涨,文本生成过拟合是模型开始逐字背诵训练语料。

检测方法很简单:每轮结束把当前模型拿去生成一段“韦小宝走进扬州城”,如果第 10 轮生成的句子还是含糊的套话,但第 18 轮生成的句子和你训练集里某段几乎一模一样,恭喜,模型开始背课文了。此时回退到 best model 就能得到泛化能力最强的那个版本。我的习惯是训练时每轮都记录val_loss,若连续 5 轮没有刷新最低值就提前终止——比死等 30 轮省下三分之一时间,而且生成质量更好。

这里多说一句:字符级模型的验证集 loss 数值本身没有绝对意义,因为 vocab_size 有 8000,即使模型达到 40% 的逐字准确率,loss 也就在 2.0 上下浮动。所以判断模型好坏不能只看 loss,要配合生成样本做人工评估,这是文本生成和分类任务最大的不同——分类有明确的准确率指标,生成没有。

4.2 generation 函数:temperature、top_k 和 multinomial 配合

训练完的模型要变成“能写小说”的程序,核心是采样策略。如果每步都取概率最大的字符(argmax),生成结果会陷入“韦小宝说道说道说道说道”的循环,因为最高频的字符一旦被选中,它的隐状态会不断强化自己。解决办法是从概率分布中随机采样,再用 temperature 控制分布的尖锐程度。

@torch.no_grad() def generate(model, char2idx, idx2char, seed='韦小宝走进扬州', length=200, temperature=0.8, top_k=40, repetition_penalty=1.2): model.eval() ids = [char2idx.get(c, char2idx['<UNK>']) for c in seed] out = list(seed) hidden = None for _ in range(length): x = torch.tensor([ids[-seq_len:]], dtype=torch.long) logits, hidden = model(x, hidden) logits = logits[0, -1] / temperature # top_k 截断:只保留概率最高的 k 个候选 if top_k: vals, _ = torch.topk(logits, top_k) logits[logits < vals[-1]] = -float('inf') # 重复惩罚:已出现的字,logits 整体除以惩罚系数 if repetition_penalty != 1.0: for idx in set(ids[-100:]): logits[idx] /= repetition_penalty probs = torch.softmax(logits, dim=-1) next_id = torch.multinomial(probs, 1).item() out.append(idx2char[next_id]) ids.append(next_id) return ''.join(out)

这段代码有三个关键决策点。第一,logits[0, -1]取的是最后一个时间步的输出,因为生成时前面的字符都是已知的,只有最后一步需要预测。第二,temperature的语义是:值越小分布越尖锐,生成的文本越保守、越接近训练集中的高频表达;值越大分布越平,文本越跳脱,但超过 1.5 就开始胡言乱语。对《鹿鼎记》这个语料,0.7~1.0 之间最合适——低于 0.5 会整段重复,高于 1.2 会出现“韦小宝拔出一把剑,剑是一把剑”这种语义断裂。第三,top_k=40是为了把那些概率极低的生僻字和标点排除在采样池外,防止生成突然冒出“魑魅魍魉”级别的冷字打乱叙事。

重复惩罚repetition_penalty的作用是:对最近 100 步内出现过的字符,把它的 logits 除以 1.2。这样模型不是完全禁止重复,而是让重复的可能性降低。注意惩罚系数不要大于 2.0,否则模型会刻意回避常用字,生成结果变得拗口。

4.3 把“说道说道说道”压下去的采样组合

实际调参时你会碰到一个现象:单独调 temperature 或单独调 top_k 都压不住重复。我建议把三个旋钮按下面的顺序配合,而不是单靠某一个:先用temperature=0.8保证基本的语言连贯性,再用top_k=40砍掉长尾,最后用repetition_penalty=1.2处理顽固的重复片段。

如果你跑出来的结果还是重复,先检查训练是否充分。我之前在数据量只有几十万字符的语料上训练,发现 temperature 怎么调都没用,后来把样本生成时的step从 8 改成 4,样本量翻倍,重复问题自然缓解。这说明了文本生成领域的一个底层规律:生成质量的上限由训练数据决定,采样参数只是在概率分布里做取舍,无法无中生有。

5. 常见问题与避坑:生成重复、乱码、OOM 的五个血泪经验

5.1 现象:输出全是“说道说道说道”

这是 LSTM 文本生成最经典的现象,几乎人人都会碰到一次。现象是生成的文本从某句话开始,同一个词反复出现,不换词也不换结构。原因分两层:第一层是训练不充分,模型没有学到足够丰富的词汇转移规律,概率分布被高频词主导;第二层是采样策略不对,argmax 会让模型沿着概率最高的路径一路滑下去,而这条路径往往就是语料里最常见的搭配。

解决:优先调采样策略。我一般会先把temperature设到 0.9,再开top_k=40,如果还重复就逐步加大repetition_penalty到 1.5。训练侧也有一个补救:把seq_len从 64 调大到 128,让模型能看到更长的上下文,有时短序列样本太多是重复的病根。

5.2 现象:生成结果是标点符号刷屏

有时候生成的文本前面几句正常,后面开始疯狂输出逗号、引号和句号。这不是模型坏了,而是语料里标点的数量占比太高。《鹿鼎记》的对话极多,对话必然伴随引号和逗号,如果小说原文里还有大量“说道:”这种结构,模型很容易学到“引号后面大概率跟逗号”这种统计规律,进而陷入标点循环。

解决:在字符统计阶段看一眼标点占比。我见过极端情况,标点字符占了全文字符数的 20% 以上,此时有两种处理路径:一是把标点中占比过高的字符(比如逗号)按比例降采样,比如每 5 个逗号保留 3 个;二是在训练时把连续的标点压缩成一个,比如把,,合并成。注意动作不能太大,对话的情感全靠标点表达,压缩到极致小说就没味道了。

5.3 现象:训练 loss 纹丝不动或者直接变 NaN

loss 完全不动的原因,最常见的是学习率太大导致梯度在 8000 类分类层上震荡,模型始终在同一个局部区域徘徊。另一种情况是文本里有异常字符,比如\x00\ufffd替换符混进了词表,导致 embedding 层对特定 ID 学不出有效表示。

解决思路是分两步排查。先打印词表里出现次数最少的 10 个字符,确认没有不可见字符;然后把lr从 1e-3 降到 3e-4,重启训练观察前 3 轮。如果是 NaN,八成是梯度爆炸,把clip_grad_norm从 5.0 降到 1.0 一般能救回来。另外检查DataLoader里有没有打开pin_memory=True但没做数据转换——这个组合在某些 PyTorch 版本里会导致 loss 随机变 NaN。

5.4 现象:Windows 下跑出乱码,Linux 下正常

文本生成模型本身没有编码问题,但 Windows 控制台的默认编码是 GBK,而你的脚本和模型文件是 UTF-8。训练时打印 loss 没事,print(generate(...))一跑,输出全是“鈥斺€斺€”这类乱码。这不是模型问题,是控制台编码不匹配。

解决:Python 3.7+ 里可以直接在脚本开头加一行环境设置,把标准输出的编码强制改成 UTF-8:

import sys sys.stdout.reconfigure(encoding='utf-8')

如果这样改了还乱码,就把生成结果写入文件再用编辑器打开,绕开控制台。这也是为什么我总建议生成脚本直接保存到 txt 文件而不是打印到屏幕——即使不涉及乱码,小说级长度的输出在控制台里也会滚动到看不到开头。

5.5 现象:显存 OOM 或者一个 epoch 要跑半小时

字符级 LSTM 的显存消耗和batch_size * seq_len * hidden_size成正比。64 的 seq_len 加 256 的 hidden_size 对 6GB 显存毫无压力,但如果想跑 seq_len=512 去学长段落,显存可能直接拉满。OOM 的解决思路有三个,按性价比排序:第一,包一个torch.no_grad()在验证循环里,很多人漏了这步,验证阶段白占了和训练一样的显存;第二,用小 batch 加梯度累积替代大 batch;第三,确认输入张量是torch.long而不是torch.float,LSTM 的 embedding 查表不会因为 float 输入报错,但显存会翻四倍。

CPU 跑得慢是另一个维度的常态问题。一个百万字符级别的语料,在普通 i5 上跑 30 轮大约需要 3~5 小时。如果想缩短,优先把hidden_size从 256 降到 128,这是对速度影响最明显的参数,效果损失能在可接受范围内。GPU 上如果显存够,把batch_size加到 128 反而能更充分利用算力,训练总时间比 64 的配置少 20% 左右。

6. 进阶技巧:让 LSTM 学出“章回体”结构的三个办法

到这一步,你的模型已经能生成文字通顺的段落了,但大概率它写不出“回目”——也就是“第一回 纵横钩党清流祸”这种结构。这类问题是字符级 LSTM 的天然短板:它只朝前看 64 个字,而回目和正文的呼应关系跨越了上千字。三个办法可以在不换模型的前提下改善。

第一个办法是给段落边界显式编码。训练数据里把每段的换行符替换成特殊 token<P>,让模型把“换段”当成一个普通字符来学。这样它会在隐状态里维护“现在处于段首还是段中”的信息,生成结果会自带分段结构。做法就一行:text = text.replace('\n\n', '<P>'),同时把这个 token 加入词表。

第二个办法是把回目当训练单元的分隔锚点。每个回目开头的“第X回”前面加一个<BOS>token,模型读到它就知道自己该在“回目模式”下输出。生成时,输入<BOS> 第,模型就有概率接出“第一回”而不是直接说“韦小宝”。理论上这个技巧依赖模型对<BOS>后文法的记忆,层数不够时效果有限,但聊胜于无。

第三个办法是合成式生成。先用 beam search 思路生成 5 个候选开头,挑一句回目风格最浓的,再以它为 seed 继续生成正文。这个办法的实现成本最低,本质是把模型当成采样器而不是生成器,适合不想动训练代码的人。我对这个技巧的评价是:它治标不治本,但效果立竿见影——生成结果的结构感至少提升一个档次。

训练文本生成模型久了之后,我自己的一个习惯是:每调一次参数,固定生成同一个 seed(比如“韦小宝走进扬州”),把新旧结果放一起对比。这样能直观看出改动带来了什么,而不是靠玄学感觉。你能把这一步坚持住,胜过读十篇调参心得。希望帮到你。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/24 1:04:26

微博热点舆情聚类实战:从爬虫清洗到TF-IDF与KMeans的完整链路

简介&#xff1a;面向对Python文本挖掘与舆情分析感兴趣的学习者&#xff0c;资源以微博热点话题为对象&#xff0c;完整提供了从数据采集、分词处理到聚类分析的项目源码与配套数据。核心依赖包括jieba分词、pandas数据处理、scikit-learn机器学习、matplotlib可视化与request…

作者头像 李华
网站建设 2026/9/24 1:00:21

基于SSM框架的农产品电商系统开发实践

1. 项目概述&#xff1a;基于SSM的助农特色农产品销售系统作为一名深耕Java领域多年的开发者&#xff0c;我最近完成了一个具有社会价值的毕业设计项目——基于SSM框架的助农特色农产品销售系统。这个系统专为解决农产品销售渠道单一、信息不对称等问题而设计&#xff0c;通过数…

作者头像 李华
网站建设 2026/9/24 0:58:30

Python零基础转型:首日高效学习框架与实战

1. 从零开始的Python转型之路作为一名从传统行业转投Python开发的"新生代程序员"&#xff0c;我清楚地记得第一天接触这门语言时的困惑与兴奋。Python以其简洁优雅的语法和强大的生态系统&#xff0c;成为技术转行者的首选语言。但真正开始学习时&#xff0c;面对海量…

作者头像 李华
网站建设 2026/9/24 0:56:49

RPA自动化解放生产力:影刀实战经验分享

1. 项目背景与核心价值去年接手新项目时&#xff0c;我每天要花3小时重复处理Excel报表。直到发现影刀RPA这个神器&#xff0c;才真正体会到"科技解放生产力"的含义。现在我的日报生成、数据核对、邮件发送等重复工作全部交给机器人处理&#xff0c;每天多出2小时研究…

作者头像 李华
网站建设 2026/9/24 0:49:08

基于Python的舆情热点分析平台:从网易新闻爬虫到情感可视化

简介&#xff1a;面向Python课程设计与毕业设计的一站式舆情热点分析平台源码&#xff0c;完整覆盖从网易新闻及评论抓取、数据清洗、中文分词、停用词过滤、情感分析、关键词提取到时间序列分析与可视化展示的典型数据科学流程。资源共1403个文件&#xff0c;约23.83MB&#x…

作者头像 李华
网站建设 2026/9/24 0:39:48

基于PCD小样本数据集的PCB元器件缺陷检测:YOLOv8训练与产线落地实践

简介&#xff1a;PCD表面元器件缺陷检测数据集面向从事工业质检、电子制造与目标检测算法实践的开发者与研究者&#xff0c;用于训练和验证PCB表面元器件缺陷识别模型。数据集包含超过600张标注图像&#xff0c;已统一处理为YOLO格式并完成数据增强&#xff0c;可直接用于YOLO全…

作者头像 李华