资讯动态

LSTM中文诗歌生成实战:从数据管道到采样调参的完整指南

发布时间:2026/10/9 14:53:32 来源:尧图企业网站定制
简介这份资源是一套基于LSTM的中文诗歌生成Python项目适合计算机、人工智能等专业学生用于期末大作业、课程设计或毕设参考也适合想入门文本生成的小白进阶学习。项目在原开源代码基础上做了优化与bug修复重点针对中文诗歌生成训练与生成流程均已测试跑通答辩评审平均分达96分。压缩包共10个文件包含4个py源码文件、5个txt数据文件与1个md说明文档整体约4.76MB其中py文件负责模型、训练与采样逻辑txt提供诗歌语料md为使用说明。已有155人学习下载。读者可据此掌握LSTM文本生成的数据预处理、模型搭建、训练调参与采样输出全流程并可在源码基础上修改以扩展其他功能下载后建议先阅读README.md仅供学习参考切勿商用。1. 一份能跑通的 LSTM 写诗源码到底能省下多少折腾时间如果你正在找一份能直接跑起来的 LSTM 中文诗歌生成代码大概率已经翻过不少仓库有的依赖缺失、有的数据路径写死、有的训练两轮就报维度错误。这份LSTM-Generate_words-master属于少见的「拆开就能用」类型——它把汉语、英文、日文诗歌生成放在同一套框架里作者重点修了中文分支的 bug并调整了训练参数。目录里poetry.txt、jpn.txt、jay.txt、shakespeare.txt四份语料直接可用train.py、model.py、sample.py、read_utils.py四个核心文件分工明确。适合课程设计、期末大作业、毕设初期演示也适合想搞懂 LSTM 文本生成完整链路的入门者。下面按「数据怎么读 → 模型怎么搭 → 训练怎么调 → 生成怎么采样 → 坑在哪」的顺序拆一遍。2. 数据管道与词表构建read_utils.py 里的三个关键动作2.1 为什么中文诗歌不能直接按字喂进去LSTM 文本生成的本质是把字符序列映射成概率分布模型每一步预测下一个字符。中文和英文最大的差别在于英文有天然空格分词中文没有。如果按词切分需要额外分词工具和词表管理按字切分则词表规模可控poetry.txt里常用汉字大约几千个嵌入层压力小。这份代码走的是按字建模路线read_utils.py负责把整份语料转成整数 id 序列同时维护id2char和char2id两张映射表。常见做法是把所有字符去重后排序再分配索引特殊符号如空格、换行、句号也各占一个 id。这样模型学到的就是「给定前 N 个字下一个字最可能是什么」。2.2 读取语料与构建词表的具体步骤先看数据加载的核心逻辑下面这段代码对应read_utils.py中文本转 id 的部分# read_utils.py 核心逻辑示意 import codecs import numpy as np def read_data(file_path): with codecs.open(file_path, r, utf-8) as f: text f.read() # 去重后排序保证每次运行词表顺序一致 vocab sorted(set(text)) char2id {ch: i for i, ch in enumerate(vocab)} id2char {i: ch for i, ch in enumerate(vocab)} # 把整段文本转成整数序列 id_data np.array([char2id[ch] for ch in text], dtypenp.int32) return id_data, char2id, id2char逻辑说明set(text)去重得到所有出现过的字符sorted保证顺序稳定避免每次训练词表 id 漂移导致模型无法复用。char2id和id2char互为逆映射生成阶段靠id2char把预测 id 还原成汉字。参数方面file_path指向data/poetry.txt编码必须写utf-8否则中文会乱码。如果换用自己的语料只要保证纯文本、utf-8 编码即可不需要额外清洗。2.3 批次生成与序列截断训练时不能把整本诗集一次性喂进去需要按固定长度切段。常见做法是设定batch_size和seq_length从 id 序列里随机取起点截取连续片段。下面是一个批次生成函数示意def generate_batch(id_data, batch_size, seq_length): # 随机起点保证每个 batch 覆盖不同位置 starts np.random.randint(0, len(id_data) - seq_length - 1, batch_size) inputs np.array([id_data[s:s seq_length] for s in starts]) targets np.array([id_data[s 1:s seq_length 1] for s in starts]) return inputs, targetsinputs是当前序列targets是右移一位的序列模型学习的是「看到前一个字预测后一个字」。seq_length一般设 20 到 50太短学不到长距离依赖太长显存吃紧。batch_size根据显卡调整CPU 训练建议 32 或 64GPU 可以上 128。注意starts的上界要减去seq_length 1否则切片越界。3. 模型结构与训练脚本model.py 和 train.py 怎么配合3.1 嵌入层加 LSTM 加全连接的标准三段式model.py定义了一个典型的两层 LSTM 结构。输入先过嵌入层把整数 id 变成稠密向量再进 LSTM 提取时序特征最后接全连接层输出词表大小的 logits。下面是对应结构示意# model.py 结构示意 import tensorflow as tf class PoetryModel(tf.keras.Model): def __init__(self, vocab_size, embedding_dim, hidden_dim): super().__init__() self.embedding tf.keras.layers.Embedding(vocab_size, embedding_dim) self.lstm tf.keras.layers.LSTM(hidden_dim, return_sequencesTrue) self.fc tf.keras.layers.Dense(vocab_size) def call(self, inputs): x self.embedding(inputs) x self.lstm(x) logits self.fc(x) return logitsvocab_size等于词表大小由read_utils统计得出embedding_dim常见取 128 或 256hidden_dim取 256 或 512。return_sequencesTrue保证每个时间步都有输出才能逐字预测。损失函数用交叉熵优化器用 Adam学习率从 0.001 起步。如果训练 loss 下降很慢先检查词表大小是否和全连接层输出维度一致这是最常见的维度不匹配来源。3.2 训练命令与参数含义作者在说明里给出的训练命令是python train.py \ --use_embedding \ --input_file data/poetry.txt \ --name poetry \ --learning_rate 0.001 \ --num_steps 50 \ --hidden_size 128 \ --batch_size 64 \ --num_layers 2 \ --output_dir model/ \ --max_vocab 3500逐项说明--use_embedding开启嵌入层不加这个参数会走 one-hot 路线显存占用大且效果一般--input_file指定语料路径中文诗歌就指向data/poetry.txt--name是模型保存前缀--num_steps是截断长度对应前面的seq_length--hidden_size是 LSTM 隐藏单元数--num_layers是堆叠层数--max_vocab限制词表上限超出部分按频率截断防止词表过大拖慢训练。训练过程中会定期保存 checkpoint 到output_dir中断后可以从最新 checkpoint 恢复。3.3 训练过程要看哪些指标启动训练后终端会打印每一步的 loss 和 perplexity。loss 从 6 左右开始下降中文诗歌语料训练几千步后通常能降到 2 到 3 之间。如果 loss 卡在某个值不动常见原因是学习率太大导致震荡或者词表里低频字太多稀释了梯度。可以先把max_vocab调小到 2000 试试观察 loss 是否继续下降。另外注意训练集和验证集如果来自同一份语料loss 低不代表生成质量好最终还是要看sample.py的输出。4. 生成与采样sample.py 里温度参数怎么调4.1 从 checkpoint 恢复并逐字生成sample.py负责加载训练好的模型给定一个起始字逐字预测后续内容。核心逻辑如下# sample.py 生成逻辑示意 def generate(model, start_char, char2id, id2char, length100, temperature1.0): input_id char2id.get(start_char, 0) result start_char for _ in range(length): # 构造当前输入序列 input_seq np.array([[input_id]]) logits model(input_seq) # 温度调节温度越低分布越尖锐 logits logits / temperature probs tf.nn.softmax(logits) next_id np.random.choice(len(probs), pprobs.numpy().flatten()) result id2char[next_id] input_id next_id return resulttemperature是关键参数设为 1.0 时按原始分布采样输出多样性正常降到 0.5 以下会偏向高频字诗句更「通顺」但容易重复升到 1.5 以上会出现生僻字和乱码。我一般先用 0.8 试看输出是否成句再微调。length控制生成字数中文五言绝句 20 字、七言绝句 28 字设 50 到 100 能看出整体风格。4.2 起始字的选择影响很大起始字相当于给模型一个「主题提示」。用「春」「月」「风」这类高频意象字开头生成结果通常比较像诗用生僻字开头模型可能接不上。如果想让生成内容围绕某个主题可以把起始字设为主题字或者把起始序列设成一句已有的诗让模型续写。注意sample.py里如果直接加载 checkpoint要保证模型结构参数和训练时一致否则会报变量形状不匹配。4.3 生成结果的简单筛选模型输出不一定每首都像样常见做法是生成多条人工挑一条。可以写个循环批量生成再按重复字比例过滤def filter_poems(poems, max_repeat0.3): good [] for p in poems: repeat_ratio 1 - len(set(p)) / len(p) if repeat_ratio max_repeat: good.append(p) return goodmax_repeat设 0.3 表示重复字占比不超过三成超过就丢弃。这个阈值可以按语料风格调整唐诗重复率低宋词可能稍高。5. 避坑与排查训练和生成阶段最容易翻车的五个点5.1 现象训练启动即报维度不匹配原因max_vocab和模型输出层维度不一致或者换了语料后词表大小变了但模型没重建。解决先跑一遍read_utils打印词表大小把vocab_size和max_vocab对齐重新训练。5.2 现象loss 一直不降生成全是重复字原因学习率过大或过小也可能是seq_length太短导致模型学不到上下文。解决学习率从 0.001 调到 0.0005 试num_steps从 20 加到 50观察 loss 曲线是否改善。5.3 现象生成结果全是逗号句号原因语料里标点占比高模型学会了「偷懒」只输出高频标点。解决在数据预处理阶段过滤掉连续标点或者降低标点在词表中的权重也可以提高temperature增加多样性。5.4 现象GPU 显存溢出原因batch_size或hidden_size太大或者num_steps过长。解决先把batch_size减半再降hidden_size最后考虑缩短num_steps。CPU 训练时把batch_size设到 32 以下。5.5 现象换用 jay.txt 或 shakespeare.txt 后效果差原因不同语料的字符集差异大中文模型直接套英文语料词表会错乱。解决每种语料单独训练一个模型不要混用 checkpoint。作者也提到主要优化的是中文分支其他语料属于「能跑但没细调」状态。6. 进阶技巧把生成结果做成可复现的评估流程训练完一个 LSTM 写诗模型最怕的是「这次生成不错下次重启就复现不了」。我一般会固定三样东西随机种子、模型 checkpoint、采样参数。下面是一个可复现的生成脚本骨架import numpy as np import tensorflow as tf # 固定随机种子 np.random.seed(42) tf.random.set_seed(42) # 加载模型和词表 model tf.keras.models.load_model(model/poetry_final) id2char load_id2char(model/id2char.pkl) # 固定起始字和温度 start_chars [春, 月, 风, 花] for ch in start_chars: poem generate(model, ch, char2id, id2char, length56, temperature0.8) print(poem)固定种子后同样的起始字和温度每次输出一致方便对比不同 checkpoint 的效果。评估时不要只看一条至少生成 20 条统计重复率、平均句长、标点比例三个指标。重复率低于 0.2、平均句长在 5 到 7 字之间、标点占比 10% 左右基本算合格。如果要做课程设计答辩可以把这些指标做成表格比单纯贴几首诗更有说服力。还有一个容易忽略的点模型保存时最好把词表一起存下来。只存权重不存char2id换台机器加载时词表对不上生成结果就是乱码。我习惯在output_dir里同时放id2char.pkl和char2id.pkl加载时先读词表再建模型。从那以后我每次保存模型都强制走一遍「权重 词表 训练参数」三件套再也没出现过加载后输出乱码的情况。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑