拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

字符级LSTM古诗词生成实战:从数据清洗到Flask部署全解析

字符级LSTM古诗词生成实战:从数据清洗到Flask部署全解析

简介:面向自然语言处理与诗歌生成方向的学习者,这份LSTM古诗自动生成系统覆盖建模、训练到Web展示全流程。针对RNN长序列梯度缺陷,采用LSTM优化,使用sparse_categorical_crossentropy损失函数和Adam(lr=0.002)训练,可生成五言律诗、七言绝句与藏头诗。资源共40个文件,压缩包81.76MB,包含Python源码、Keras模型文件(checkpoint、data、index、meta)、前端页面及配置说明,结构清晰便于复现。已有1874人学习下载。借助该资源可直接运行算法,也可参考模型设计、损失函数选择与Flask集成方式,适合NLP课程设计、毕业设计或古诗生成算法对比实验。

1. 古诗词生成不是随机拼字:LSTM 如何在字符级别上学会平仄与意象

我最初接到古诗词生成这个需求时,第一反应是拿现成的模板拼句子,但效果非常"打油诗"。换成 LSTM 之后才发现,古诗生成本质上是一个字符级序列预测问题:模型看到的不是词语,而是连续字符流。它从全唐诗里学到的不是词库,而是「平仄交替」「对仗工整」「意象搭配」这些藏在字符序列里的统计规律。这份资源是一个完整的 LSTM 古诗生成系统,包含数据处理、模型训练、采样生成和 Flask 网页展示。适合正在学序列建模、想做一个能跑起来的 NLP 实战项目的开发者,也适合想把模型快速包装成 demo 给客户看的人。读完你会明白为什么字符级 LSTM 比 n-gram 更合适,以及部署时那些让人头大的坑。

2. 数据准备与字符编码:从全唐诗到可训练的序列样本

2.1 语料来源与清洗:为什么只保留五言和七言

网上流传的《全唐诗》文本通常夹杂着作者、词牌名、注释、标点,甚至还有繁体异体字。直接丢给模型训练,字符集会膨胀到上万,而且噪音会让 LSTM 去学那些无意义的标点符号。我拿到语料后第一件事就是过滤出五言绝句和七言绝句。绝句本身结构严谨,每首只有四句,句子长度固定,非常适合作为定长序列的训练数据。五言和七言分开训练的话,生成时更容易控制格式。

清洗步骤其实很机械:先按行读取,去掉包含「诗」「卷」「作者」等关键字的行;然后保留正文行,用正则去掉所有非中文字符;遇到空缺的句子就跳过整首。下面是资源包里的清洗脚本核心部分:

import re def clean_poem(raw_text): # 只保留中文字符和句读符号 text = re.sub(r'[^\u4e00-\u9fa5,。!?]', '', raw_text) # 按句读切分,五言绝句应该是4句,每句5字;七言每句7字 lines = [l for l in re.split(r'[,。!?]', text) if l] if len(lines) != 4: return None if all(len(l) == 5 for l in lines): return ('五言', lines) if all(len(l) == 7 for l in lines): return ('七言', lines) return None

这段代码做的事情很简单:过滤后把一首诗切成四句,然后判断长度是否一致。之所以不直接按字数过滤,是因为古诗词里有「偷声」和「减字」的现象,同一词牌字数也可能不同。但我选用绝句是因为它的格式最规整,不用额外处理变体。

清洗完的语料保存为两个文件:poem_five.txt和poem_seven.txt,每行一首诗,句与句之间用空格分隔。这一步决定了下游字符集的大小。清洗后去重,我大概保留了 4 万首左右的五言和 3 万首七言,字符集在 3000 左右。如果语料太少,LSTM 学不到平仄模式;太多又会引入大量生僻字,导致最终生成的诗歌里频频出现「爨」「龘」这种字,交给用户完全是灾难。

2.2 构造训练样本:定长序列切分与字符映射表

LSTM 需要一个固定长度的时间步。古诗不是每个字都独立,后一个字依赖于前面的上下文,所以我把每首诗拼接成一个大字符串,然后用滑动窗口切出「输入序列」和「目标序列」。每个样本包含seq_len个输入字符,以及后移一位的相同长度目标字符,也就是input[i] -> target[i] = input[i+1]。

字符映射表构建要注意几个细节:必须包含三个特殊 token ——<PAD>用于填充、<S>用于句首、<E>用于句末。其中<S>和<E>很关键,因为生成时我们需要一个信号来启动和终止。我用word2idx和idx2word两个字典保存映射关系。

class PoetryDataset(Dataset): def __init__(self, poems, seq_len=64, char2idx=None, idx2char=None): self.seq_len = seq_len self.char2idx = char2idx or {} self.idx2char = idx2char or [] # 把所有诗句合成一个长序列 all_text = '<S>' + '<S>'.join(poems) + '<E>' self.indices = [self.char2idx.get(c, self.char2idx['<PAD>']) for c in all_text] def __len__(self): return len(self.indices) - self.seq_len - 1 def __getitem__(self, i): x = torch.tensor(self.indices[i : i + self.seq_len], dtype=torch.long) y = torch.tensor(self.indices[i + 1 : i + self.seq_len + 1], dtype=torch.long) return x, y

seq_len我取 64,足够覆盖一首七言绝句加上标点(28 个字)的两倍长度。这样模型在训练时能看到超过一首诗的上文,有助于学习跨句的承接关系。如果你发现生成的诗句频繁出现「上句不接下句」,可以尝试把seq_len增加到 128,代价是训练时间变长。

这里有个容易忽视的问题:字符<S>和<E>如果在多个样本中频繁出现,模型可能学会「看到<S>就输出<S>」的偷懒策略。所以我统计字符频率后,把<S>和<E>的频率权重调低了一些,具体做法是在损失函数里加一个 class weight,后面训练章节会讲到。

2.3 数据加载器实现:PyTorch Dataset 与批处理细节

直接返回长度不等的序列会拖慢训练,所以我在__getitem__里固定返回seq_len长度的片段。PyTorch 的DataLoader会自动把 batch 里的样本堆叠成(batch, seq_len)的张量,但需要确保collate_fn不额外做 padding,因为我们每个样本长度已经一致。

资源里提供了一个padding_collate,它唯一的职责是检查输入输出长度一致。其实默认的 collate 就够用,但如果你的seq_len不是固定的,就要自己写。这里我习惯设置pin_memory=True,在 GPU 训练时能明显减少 CPU 到 GPU 的拷贝耗时。

train_loader = DataLoader(dataset, batch_size=64, shuffle=True, pin_memory=True, num_workers=4)

num_workers我设为 4,在 Windows 上如果报错就改为 0。注意多进程加载时,word2idx字典必须作为全局变量,或者通过dataset构造时传入,不然每个 worker 会重新构建一份映射,导致索引错位。

数据准备的最后一步是把char2idx和idx2char保存成 json,模型训练完还要用它们把输出转回汉字。忘了保存这一步,后面一切生成都无法进行。

3. 模型设计与训练参数:三层 LSTM 加 dropout 的效果差异

3.1 网络结构选型:为什么用字符级 LSTM 而不是 word2vec

有人问为什么不用预训练的词向量。古诗词的字义高度依赖语境,比如「春」在不同诗里可能代表生机也可能代表愁绪,用静态词向量很容易丢失这种多义性。字符级 LSTM 不依赖分词质量,每个汉字是一个独立输入,模型自己学习「春」和「秋」的搭配关系。另外,汉字本身就有平仄属性,LSTM 的隐状态可以把这个信息编码到序列的长期依赖里。

网络结构我采用三层 LSTM。单层 LSTM 对「前一句的末尾平仄」影响「后一句开头」这种远距离依赖无能为力。三层堆叠后,第一层捕捉基本字词搭配,第二层学习短语节奏,第三层整合成句子的语义和格律。hidden_size设 256,过大会导致参数量膨胀且容易过拟合。每层之间加dropout=0.3,只在层间生效,不在时间步之间生效。

class PoetryLSTM(nn.Module): def __init__(self, vocab_size, embedding_dim=128, hidden_size=256, num_layers=3, dropout=0.3, seq_len=64): super().__init__() self.embedding = nn.Embedding(vocab_size, embedding_dim) self.lstm = nn.LSTM(embedding_dim, hidden_size, num_layers, dropout=dropout, batch_first=True) self.fc = nn.Linear(hidden_size, vocab_size) def forward(self, x, hidden=None): emb = self.embedding(x) # (batch, seq_len, emb) out, hidden = self.lstm(emb, hidden) # out: (batch, seq_len, hidden) out = self.fc(out) # (batch, seq_len, vocab_size) return out, hidden

这里batch_first=True意味着输入维度是(batch, seq_len, embedding_dim),Dataloader 出来的张量正好是这个形状。最后一层fc使用权重共享:我让fc.weight与embedding.weight共享参数,这样可以减少参数量,而且效果表明它能在输出层复用输入的语义表征,让生成的句子更「像」训练集里的用词。代码里加一行self.fc.weight = self.embedding.weight即可,但注意 PyTorch 要求两个 weight 形状一致,所以vocab_size必须与 embedding 的最后一维一致。

3.2 损失函数与优化器配置:交叉熵、Adam 与学习率衰减

损失函数用CrossEntropyLoss,它的输入是(batch * seq_len, vocab_size)的 logits,target 是(batch * seq_len)的索引。我们需要把模型输出的前两维合并。另一个关键点是加入字符频率权重:高频字「之」「不」出现的概率大,如果给它们过高的权重,生成的诗会趋于平淡。我使用sklearn的compute_class_weight来算每个字符的逆频率,然后传给CrossEntropyLoss。

loss_fn = nn.CrossEntropyLoss(ignore_index=char2idx['<PAD>'], weight=class_weight_tensor)

优化器用 Adam,初始学习率lr=1e-3,每 5 个 epoch 按lr *= 0.5衰减。LSTM 对学习率很敏感,过大会导致 loss 震荡,过小则训练缓慢。我观察到 1e-3 对三层的结构是安全的,如果 loss 在最后几个 epoch 突然升高,就说明学习率没降下来。

训练循环里有一个很重要的习惯:每 10 个 epoch 就用当前模型生成一首诗,对比不同阶段生成质量。只看 loss 下降是不够的,loss 低可能只是模型学会了输出高频词。生成样例能直观反映模型是否开始形成格律,这个观察比任何指标都靠谱。

3.3 训练循环与模型保存:每隔 N 个 epoch 生成一首诗来观察

def train(model, loader, optimizer, loss_fn, epochs, device): model.train() for epoch in range(epochs): total_loss = 0 for x, y in loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() out, _ = model(x) loss = loss_fn(out.view(-1, out.size(-1)), y.view(-1)) loss.backward() # 梯度裁剪,防止长序列训练时梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) optimizer.step() total_loss += loss.item() if epoch % 10 == 0: print(f"epoch {epoch}, loss {total_loss/len(loader):.4f}") generate_sample(model, char2idx, idx2char, device)

梯度裁剪的max_norm=5.0是我调出来的一个折中。太小会让模型训练变慢,太大则失效。在 LSTM 序列任务里,这个值几乎必设,否则训练到第 30 个 epoch 时 loss 可能突然变成nan。

保存模型时不要只保存model.state_dict(),一定把char2idx、idx2char、seq_len这些配置打包进同一个 checkpoint,这样才能在部署时无缝恢复。我用.pth.tar格式,本质是一个字典。资源包里的checkpoint/目录给出了训练 50 个 epoch 后的完整文件和日志,你可以直接加载来做推理,也可以自己从头训练。

4. 生成策略与后处理:温度采样、平仄校验和意象过滤

4.1 温度参数对随机性的控制:从贪心到 top-p

训练好模型后,生成有两种常用策略:贪心解码和采样。贪心每次选概率最大的字符,结果可能陷入重复循环。我采用带温度系数的多项式采样,公式是P(w) = exp(z_i / T) / sum(exp(z_j / T))。温度T=1时保持原分布,T越低越保守,越高越随机。实践下来,T=0.8生成的句子在通顺和创意之间最平衡。

仅仅调温度还不够,有时会出现「的」「之」这类停用词被反复采样。我加了一个top_p累积概率过滤:只从累积概率超过p=0.9的最小候选集里采样,其余低概率字符直接丢弃。这个操作能减少生僻字和重复字。

def sample_from_logits(logits, temperature=0.8, top_p=0.9): logits = logits / temperature probs = torch.softmax(logits, dim=-1) sorted_probs, sorted_idx = torch.sort(probs, descending=True) cumsum = torch.cumsum(sorted_probs, dim=-1) mask = cumsum - sorted_probs > top_p sorted_probs[mask] = 0 normalized = sorted_probs / sorted_probs.sum() idx = torch.multinomial(normalized, 1).item() return sorted_idx[idx].item()

注意top_p过滤要在温度缩放之后做,顺序反了会影响概率分布形状。我在代码里加了这个注释,因为很多开源实现把顺序搞反,导致top_p形同虚设。

4.2 平仄与押韵的规则校验:生成后的硬约束修正

LSTM 学到的平仄是概率性的,并不能保证完全合规。我准备了一份平水韵表,把每个汉字标成平声或仄声。生成四句之后,先检查每句内部的平仄交替是否符合基本规则。比如五言绝句的常见格式是「仄仄平平仄,平平仄仄平」,如果检测到连续三个平声或三个仄声,就标记为不合格。

解决方式不是重新生成整首,而是做局部替换。我提取不合格位置的上下文,把该位置作为待选字,重新用模型预测该位置的字符,但强制候选字符满足平仄要求。这一步相当于把模型当成一个填空器,而不是从头采样。资源里postprocess.py实现了这个功能,替换时还考虑了韵脚:第二句和第四句的最后一个字必须在同一韵部。

这里有一个细节:平仄校验必须在生成完整四句后进行,而不是逐句生成时进行。因为模型在生成第二句时并不知道第四句的韵脚,如果你逐句硬控,最后可能韵脚冲突。我的做法是先快速生成一批候选诗,然后按「平仄正确率 + 押韵数量」排序,选综合分最高的那首。

4.3 标题生成与五言/七言格式控制

用户输入一个关键词,比如「春」,系统需要生成一首以「春」为主题的绝句。我的做法是输入序列用<S>春作为起始,让模型从主题字开始扩散。但这样直接生成容易让第一句的主题词被遗忘,所以我把它变成条件:在生成第一句时,强制第一个字符为「春」,后续字符从模型中采样。其后每句的起始字符由上一句的语义推断,不做额外限制。

格式控制则靠生成时的长度约束。五言诗每句必须 5 字,七言每句 7 字。我在逐字采样时维护一个计数器,当句子长度达到目标时强制输出句读符号「,」或「。」,然后开始下一句。为了避免模型反复输出句读,我把句读符号从采样候选集里临时剔除,只在强制位置添加。

def generate_poem(model, prefix, seq_len, max_len, device): model.eval() with torch.no_grad(): input_ids = [char2idx[c] for c in ('<S>' + prefix)] for _ in range(max_len): x = torch.tensor([input_ids[-seq_len:]], device=device) out, _ = model(x) logits = out[0, -1, :] next_id = sample_from_logits(logits) if idx2char[next_id] in ',。!?': next_id = sample_from_logits(logits, top_p=0.7) input_ids.append(next_id) return ''.join(idx2char[i] for i in input_ids)

这个循环里,当采样到标点时,我会重新采样一次,并把top_p调低到 0.7。这样能避免模型在句子中间过早结束。真正的句读由外部强制插入,也就是说整个生成过程不依赖模型输出句读符号。这样做让生成结果更规整。

5. Flask 系统实现与部署:接口设计、线程安全与五个常见坑

5.1 后端接口与前端交互:一次请求生成一首诗

把训练好的模型包装成 Flask 服务,本质上就是加载 checkpoint,然后在 POST 请求里调用生成函数。我用一个全局变量_MODEL保存模型实例,在应用启动时预热。接口设计如下:

from flask import Flask, request, jsonify import torch app = Flask(__name__) _model = None def load_model(): global _model checkpoint = torch.load('checkpoint/20240501.pth.tar', map_location='cpu') char2idx = checkpoint['char2idx'] idx2char = checkpoint['idx2char'] _model = PoetryLSTM(vocab_size=len(char2idx), seq_len=checkpoint['seq_len']) _model.load_state_dict(checkpoint['model_state_dict']) _model.eval() return char2idx, idx2char, _model @app.route('/api/generate', methods=['POST']) def generate(): data = request.get_json() theme = data.get('theme', '春') style = data.get('style', '五言') # 五言或七言 # 生成主逻辑省略 poem = do_generate(theme, style) return jsonify({'poem': poem, 'theme': theme, 'style': style}) if __name__ == '__main__': load_model() app.run(host='0.0.0.0', port=5000)

load_model在app.run之前执行,确保第一个请求到达时模型已在内存中。实际部署时,我一般用 gunicorn 启动,设置 4 个 worker。但要注意每个 worker 都会加载一份模型副本,内存消耗翻倍,如果服务器只有 2G 内存,最好改成单 worker + 多线程模式,或者用torch.jit.script把模型序列化为 TorchScript,推理速度也能提升 20% 左右。

前端那边我提供了一个非常简单的index.html,只有一个输入框、一个下拉框和结果区域,用 fetch 调用接口。没有用 Vue 或 React,因为目标用户只是需要看个效果,没必要增加打包构建的复杂度。前端代码里没有坑,主要问题都出在后端的并发和兼容性上。

5.2 模型加载与线程安全:不要在 request 里初始化模型

最容易犯的错误是在请求处理函数里加载模型。这样每次请求都会读取磁盘、重建图结构,响应时间会飙到好几秒,而且高并发时内存不断增长。正确的做法是把模型加载到全局变量,并且只加载一次。

还有一个隐性问题:PyTorch 模型在eval模式下,多次调用forward是线程安全的吗?严格说,如果没有任何共享的可变状态,是安全的。但我在采样函数里用了torch.multinomial,它会维护一个全局的随机数生成器。多线程同时调用时,理论上会争夺全局 RNG 的状态,导致生成结果不稳定甚至报错。

解决方法是每个请求使用独立的 RNG 状态:

from torch import manual_seed import random, time def do_generate(*args): seed = int(time.time() * 1000) % (2**32) torch.manual_seed(seed) random.seed(seed) # 后续采样操作都是线程独立的了 poem = generate_poem(...) return poem

因为我在生成函数里没有用torch.Generator来显式控制抽样,所以通过手动设置全局种子来隔离线程间的 RNG 冲突。这是我在压测时翻车后总结出来的。

5.3 常见问题与避坑:从版本兼容到路径编码的五个记录

这一节记录了我在实际部署和用户反馈中遇到频率最高的 5 个问题,每条都按「现象 → 原因 → 解决」列出。

第一个坑:加载 checkpoint 时RuntimeError: Unsupported weight type。原因是训练时用了torch.save(model, ...)保存的整个对象,而部署环境的 PyTorch 版本不一致。解决:用state_dict保存,加载时用load_state_dict,并且不要包含优化器状态,除非你想断点续训。

第二个坑:Requests 并发一多,返回的诗只有一两行。原因是 Flask 自带的单 worker 是串行处理,但 gunicorn 多 worker 时,每个 worker 都有独立的模型副本,char2idx却可能在加载时被 python 的copy-on-write机制共享,某些 worker 的映射表不完整。解决:确保每个 worker 启动时都执行完整的load_model(),不要依赖父进程的全局变量。

第三个坑:中文乱码。Flask 返回 JSON 默认 ASCII 编码,汉字会变成\u6625,前端不好展示。解决:app.config['JSON_AS_ASCII'] = False。

第四个坑:生成的诗歌中出现「」这种空格或不可见字符。原因是原始语料清洗不干净,留下了全角空格。解决:在数据清洗阶段增加re.sub(r'\s', '', text),并在字符映射表里排除空字符。

第五个坑:用户输入的主题词超出词汇表。比如输入「火星」,「火」在古汉语有但「星」不在,模型无法处理。解决:在接口层做字符级过滤,对不在char2idx里的字符用<PAD>替换,或者提示用户更换关键词。我选择提示,避免生成无意义内容。

6. 进阶验证与调优技巧:从困惑度到人工评分

6.1 用困惑度判断模型是否过拟合

模型训练完后,除了看 loss,我还计算验证集的困惑度(perplexity)。困惑度是exp(loss),表示模型对下一个字符的平均不确定性。如果训练 loss 持续下降但验证的困惑度上升,基本可以判断过拟合。字符级模型在词汇只有几千的情况下,困惑度降到 2.5 左右就比较理想,代表平均候选字符只有 2.5 个,这已达到常用词的确定性输出水平。例如,在「春眠不觉晓」后面,模型给「晓」分配的候选概率集中在「明」「春」等字,困惑度很低。

6.2 人工评分表:从格律、意境、通顺度三个维度打分

自动指标永远不能替代人的审美。我设计了一个简单的人工评分表,邀请 10 位读者对生成的 20 首诗打分,每个维度 5 分,最后取平均。格律分看平仄和押韵,意境分看意象是否统一,通顺度分看句子是否像人话。这个表的优点是评分者不需要是中文专业,每条都有具体说明。资源里附带了scoring_template.xlsx,可以自行复用。

6.3 调优实验:温度与 seq_len 的影响

我用控制变量法对比了几组参数。seq_len=32时生成的句子前后关联弱,经常出现「前半句是山,后半句是水」的断裂感;seq_len=128时诗句变得流畅,但训练时间增加 40%,且重复度略高。最终折中取 64 效果最好。温度方面,T=0.6时生成的诗过于保守,每次都是「风吹柳絮飞」这类常见搭配;T=0.9时出现「石上流泉咽」这种稍微新奇的表达,但偶尔会有不通顺的句子。我现在固定T=0.8,再配合top_p过滤。

从那次之后,我每次调参都会在训练日志里记录 temp、top_p 和人工评分的均值,而不是只看 loss。现在我把这套验证流程固化到项目里,每次新数据进来都会强制走一遍「训练 → 困惑度检查 → 人工抽样 → 调整温度」的闭环。如果你要用这份资源做二次开发,建议保留这个习惯。希望帮到你。

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

返回列表