资讯动态

LSTM诗歌生成实战:拆解自动写诗工程与训练代码

发布时间:2026/10/9 3:39:35 来源:尧图企业网站定制
简介自动写诗实验包聚焦AI与自然语言处理结合的应用场景面向具备一定Python基础、想上手文本生成项目的学习者或开发者。资源共18个文件压缩包约23.83MB除Python源码与编译后的pyc文件外还包含实验指导书、实验报告、演示PPT、诗歌文本数据、日志及模型参数文件覆盖从数据准备到模型评估的完整流程。代码中涉及RNN、LSTM、Transformer等常见深度学习结构并配有数据预处理与训练脚本便于直接运行和二次修改。文档部分详细说明了实验设计思路、超参数调整及BLEU、ROUGE等评价指标实验报告还展示了生成诗歌示例与人工对比分析。演示PPT则提炼了项目背景、模型架构和核心结果适合作为课程汇报或技术分享素材。已有169人学习对于想系统理解自动写诗技术路线、快速搭建实验环境并观察模型生成效果的读者这套资料能提供较扎实的参考与复现基础。1. 自动写诗不是玄学这个 rar 里藏了一套能直接复现的 LSTM 诗歌生成工程自动写诗听起来像是个文科生拿来调侃 AI 的段子但真把一个训练好的模型跑起来你会发现写诗本质上就是一个序列预测问题。这份“自动写诗.rar”里不是 PPT 糊弄事而是塞了完整的 Python 工程train.py 负责训练model.py 定义网络utils.py 做数据读取还有已经预处理好的 tang.npz 数据集和一份 log.log 训练日志。我拆包后的第一感觉是这项目是能直接跑的不是缺了关键文件的半成品。适合谁想理解 RNN/LSTM 在文本生成上怎么落地的人以及需要交实验报告、但不想从零搭环境的学生。你可能需要自己调几个参数但至少它把从数据到生成的路都铺好了。2. 拆包看结构从 data.txt 到 tang.npz数据管线是这样被处理的2.1 压缩包里的文件到底各管什么打开“自动写诗.rar”一眼能看到的核心工程在poetry/目录下。train.py、model.py、utils.py、MyLog.py 四个是代码主体另外还有 data 目录、log.log、tang.npz、model_save 目录以及三个 doc/ppt 文档。我在拆包时特意看了两个 .pyc 的编译时间它们分别对应 Python 3.6 和 3.8说明这个项目在跨版本跑的时候要注意环境一致性——这个是后话。我习惯先把每个文件按职责归一下类文件/目录职责utils.py数据加载和预处理把文本转成 id 序列生成标签model.py网络结构Embedding LSTM 全连接train.py训练入口管理超参数、batch、epochMyLog.py轻量日志类把训练过程写到 log.logdata.txt原始诗歌文本一行一首tang.npz预处理后的 numpy 压缩数据包含 data 和 vocab_sizemodel_save/训练完成后存放模型权重log.log训练日志记录 loss 和采样示例实验指导书/实验报告/ppt实验文档流程和结果分析注意__pycache__里的 .pyc 文件只是 Python 编译缓存没有源文件的价值但能反推当时用的 Python 版本。三个 doc/ppt 分别对应实验指导、实验报告和演示文稿这个资源明显是从课程设计里流出来的所以代码风格比较直白适合学习。另外poetry目录下有 data 和 model_save 两个子目录data 里放 data.txt 和 tang.npzmodel_save 里放训练好的权重。如果你拿到压缩包后 model_save 目录是空的说明发布者没有把权重传全但没关系你可以自己跑训练。如果有权重那就省事了直接加载生成就行。我建议不管有没有现成权重都先把 train.py 完整跑一遍毕竟实验报告需要你的训练过程截图。2.2 数据预处理从原始文本到训练张量写诗模型不能直接吃汉字需要先把每个字映射成整数 ID。常见做法是扫描 data.txt统计每个字的出现次数按频率排序生成word2idx和idx2word两个字典然后用滑动窗口把文本切成定长序列。我这里给一个数据预处理的典型实现和这个项目里的 utils.py 做的事基本一致import numpy as np from collections import Counter def build_vocab(text, min_count2): # 统计每个字的出现次数过滤低频字 counter Counter(text) freq {c: n for c, n in counter.items() if n min_count} # 按频率从高到低排序并预留 0 作为 padding chars sorted(freq, keylambda c: freq[c], reverseTrue) word2idx {c: i 1 for i, c in enumerate(chars)} word2idx[pad] 0 idx2word {i: c for c, i in word2idx.items()} return word2idx, idx2word def text_to_sequences(text, word2idx, seq_len64): ids [word2idx.get(c, 0) for c in text] # 用滑动窗口切样本每个样本包含 seq_len 个输入字和 1 个目标字 seqs [] for i in range(0, len(ids) - seq_len, 1): seqs.append(ids[i:i seq_len 1]) return np.array(seqs, dtypenp.int64) if __name__ __main__: with open(data.txt, encodingutf-8) as f: text f.read() word2idx, idx2word build_vocab(text) arrays text_to_sequences(text, word2idx) np.savez(tang.npz, dataarrays, vocab_sizelen(word2idx))逻辑说明代码先用 Counter 统计字频过滤掉只出现一次的字符这样既能减小词典规模也能避免低频噪声。然后把每个字替换成 ID滑动窗口每次向后移动一个位置所以样本长度是seq_len 1最后一个字就是模型要预测的目标。参数min_count决定词典最小词频设为 2 意味着只在语料里出现一次的字会被当成0padding这种字模型学不出有效语义去掉是合理的。参数说明seq_len决定输入序列长度。古诗五言七言最多二十来个字但为了捕捉跨句的语义我一般传 64。太大内存涨得快太小句子末尾的字缺乏上下文。build_vocab里把0留给 padding训练时按 batch 对齐会用得到。项目里的tang.npz已经是处理好的结果所以正常情况下你不需要重新跑这一步除非你要换数据集重新训练。这里有个细节data.txt 里可能有标点符号、空格和换行。如果直接用text f.read()换行符会被当成一个字符进入词表。常见做法是按行读取再剔除空行避免模型把换行符当成内容。我在另一个项目里就是没处理换行结果模型学会了在每句结尾输出换行符看起来像诗其实是把格式当成了内容。更好的做法是lines [line.strip() for line in f if line.strip()] text .join(lines)这样能保证所有样本都是连续字符模型只预测下一个汉字而不是换行符。2.3 从 data.txt 到 npz文件格式与维度核对前面的代码保存的是data和vocab_size两个键。加载的时候要注意这个数组形状是[样本数, seq_len1]但没区分输入和标签需要自己在训练循环里切分。这一点我在第一次跑的时候差点搞错后面在避坑章节会细讲。我们拆包后看到的tang.npz个头不算大我推测它是用全唐诗或者一个几百首诗的合集生成的。用np.load加载后先打印一下 shape 和 dtype确认不是空文件。常见做法是import numpy as np d np.load(tang.npz) print(d[data].shape, d[data].dtype) print(d[vocab_size])如果打印出的样本数偏少或者 vocab_size 是个不合理的值就要回头检查原始 data.txt 是不是被截断过。项目里的 data.txt 是一行一行的诗每行可能是一首预处理时要按行切分否则会把换行符也当成字符污染数据。这一章的落脚点是数据管线决定了模型能学到什么tang.npz已经是别人预处理的产物但你自己要懂得怎么验货。拿到手第一件事不是python train.py而是先把数据加载出来看一眼 shape这个习惯能帮你少踩很多坑。3. 模型选型LSTM 为什么比 N-gram 更适合写诗以及 model.py 的实现细节3.1 为什么不用 N-gram 或简单 RNN写诗本质是条件语言建模给定前 n 个字预测下一个字。最朴素的做法是 N-gram 统计频率从语料里数出“春”后面最常跟“风”、“花”之类。它的问题是数据稀疏而且上下文长度固定写出来的诗会像在背词频表毫无连贯性。简单 RNN 理论上能处理任意长度但梯度消失严重长距离依赖学不动。LSTM 通过门控机制保留了长期信息在古诗这种五言七言长度不算太长但句间有语义呼应的场景里是性价比很高的选择。这个项目的 model.py 里用的就是 Embedding LSTM 全连接这种典型结构没有上 Transformer原因多半是硬件资源和训练时间受限——LSTM 在 CPU 上也能慢慢跑出可接受的结果而 Transformer 自注意力机制的内存占用和训练耗时对课程设计环境不太友好。如果你只是要交作业LSTM 完全够了。你可能会问那是否能用 GRU当然可以GRU 参数量更少在这个小数据集上可能训得更快。但如果已经有现成 model.py 的 LSTM 代码我建议先跑通再改不要一上来就换结构否则踩坑时没人帮你判断是结构问题还是数据问题。3.2 model.py 的核心结构Embedding、LSTM 与输出层我之前读这个项目的 model.py它的结构大概是这样我按常见实现复原了核心逻辑import torch.nn as nn class PoemModel(nn.Module): def __init__(self, vocab_size, embed_dim128, hidden_dim256, num_layers2): super().__init__() self.embedding nn.Embedding(vocab_size, embed_dim, padding_idx0) self.lstm nn.LSTM(embed_dim, hidden_dim, num_layers, batch_firstTrue, dropout0.3 if num_layers 1 else 0) self.fc nn.Linear(hidden_dim, vocab_size) def forward(self, x, hiddenNone): # x: [batch, seq_len] emb self.embedding(x) # [batch, seq_len, embed_dim] out, hidden self.lstm(emb, hidden) # out: [batch, seq_len, hidden_dim] logits self.fc(out) # [batch, seq_len, vocab_size] return logits, hidden逻辑说明第一层是 Embedding把字的 ID 映射成 128 维稠密向量padding_idx0是为了让补零的位置在反向传播时不产生梯度。LSTM 是两层hidden_dim 设成 256batch_firstTrue让输入张量维度是 [batch, seq_len, embed_dim]。最后一层全连接把 LSTM 每个时间步的输出映射回词表大小训练时用交叉熵损失计算每个位置上的预测误差。参数说明embed_dim和hidden_dim不是越大越好。古诗数据集规模一般只有几万到几十万字符embed_dim 超过 256 容易过拟合num_layers2是平衡表达能力和训练速度的选择堆到 4 层在 CPU 上训练慢得让人失去耐心。dropout0.3只会在层数大于 1 时启用这是 nn.LSTM 的惯例写法。在 forward 里hidden默认为 Nonenn.LSTM 会在第一次调用时自动初始化全零隐藏状态。如果你有多首诗要连续生成需要手动保留 hidden 传给下一次 forward否则每首诗都从零开始句间连贯性很差。这个细节在后面的生成代码里会用到。3.3 损失函数与采样策略训练时用nn.CrossEntropyLoss但要注意输入输出形状。model.py 输出的 logits 是 [batch, seq_len, vocab_size]目标标签是 [batch, seq_len]要把 logits 的前两维合并才能算 loss。我在复现时是这样写的def compute_loss(logits, targets): # logits: [batch, seq_len, vocab_size] # targets: [batch, seq_len] batch, seq_len, vocab logits.shape loss_fn nn.CrossEntropyLoss(ignore_index0) loss loss_fn(logits.reshape(batch * seq_len, vocab), targets.reshape(batch * seq_len)) return loss这里把ignore_index0设为 padding 值可以在计算 loss 时忽略补零位置避免模型把大量无意义的 padding 也学进去。为什么交叉熵适合这个任务因为每个位置的下一个字都来自词表这是一个多分类问题。交叉熵计算的是模型预测分布和真实标签的差异loss 越小说明模型越有把握把概率集中在正确字上。但要注意训练目标不是让 loss 降到 0那样模型就只会背诵语料生成时反而僵化。我看到 log.log 里最终 loss 停在 1 附近这个数值在词表几千的情况下是合理的。在生成阶段不能直接取 argmax因为那样写出来的诗太“模板化”。常见做法是用 temperature 平滑概率分布再随机采样。temperature 越大生成的多样性越高但也会更不着调温度低于 0.5 则几乎退化成贪心搜索。我一般会在 0.81.2 之间试这个范围写出来的诗既有点意外又大体通顺。采样代码可以这样写def sample_from_logits(logits, temperature1.0, top_k50): logits logits / temperature # top-k 过滤只保留概率最高的前 k 个字 if top_k is not None: v, _ torch.topk(logits, top_k) logits[logits v[:, -1:]] -float(inf) probs torch.softmax(logits, dim-1) return torch.multinomial(probs, num_samples1).item()temperature控制分布的平滑度top_k控制候选范围。注意这里用的是torch.multinomial做概率采样而不是argmax。4. 训练全程train.py 里的超参数、日志与模型保存4.1 训练入口与超参数设置train.py 是整个项目的发动机。它读取 tang.npz把数据切分成训练集和验证集然后进入 epoch 循环。我在这个项目里看到超参数设置的痕迹比如 seq_len、batch_size、learning_rate 这些通常都放在文件开头的常量区或者从命令行参数传入。下面是常见的训练循环骨架import torch from torch.utils.data import DataLoader, TensorDataset from model import PoemModel from utils import load_data from MyLog import MyLog log MyLog(log.log) data, vocab_size load_data(tang.npz) # 切分80% 训练20% 验证 split int(len(data) * 0.8) train_data, val_data data[:split], data[split:] # 输入和目标 x_train train_data[:, :-1] y_train train_data[:, 1:] dataset TensorDataset(torch.LongTensor(x_train), torch.LongTensor(y_train)) loader DataLoader(dataset, batch_size64, shuffleTrue) model PoemModel(vocab_size, embed_dim128, hidden_dim256, num_layers2) optimizer torch.optim.Adam(model.parameters(), lr0.001) for epoch in range(30): total_loss 0 for x, y in loader: optimizer.zero_grad() logits, _ model(x) loss compute_loss(logits, y) loss.backward() optimizer.step() total_loss loss.item() avg_loss total_loss / len(loader) log.write(fepoch {epoch1}, loss {avg_loss:.4f}) torch.save(model.state_dict(), fmodel_save/poem_{epoch1}.pth)逻辑说明数据加载后先用[:, :-1]和[:, 1:]把序列拆成输入和标签这一步对应 2.3 提到的形状问题。训练时每个 batch 过一遍前向、反向、更新参数30 个 epoch 后 loss 通常会降到 1.5 左右具体看词表大小。优化器用 Adam 是因为它对学习率的敏感度低新手不用太纠结 lr 衰减0.001 起跳一般能收敛。参数说明batch_size 在 CPU 上建议别超过 128否则内存容易爆shuffleTrue 是必须的否则模型会学进数据的排列顺序。epoch 次数不是越多越好我见过有人一上来跑 200 个 epoch结果 loss 从 1.2 过拟合到了 1.6验证集表现反而变差。补充一点如果你在 GPU 上训练可以把 batch_size 调到 256同时把学习率调成 0.002能快不少。但课程设计场景大多是 CPU所以保持低调参数就好。训练 LSTM 时梯度裁剪是个容易被忽略的细节。因为 LSTM 在长序列反向传播时梯度可能爆炸即使 loss 显示正常中间层权重也可能被冲乱。常见做法是在loss.backward()之后加一句nn.utils.clip_grad_norm_(model.parameters(), max_norm5)这个项目的 train.py 如果没写我建议你加上尤其当你把 seq_len 调大到 128 时。max_norm5是个保守值超过这个范数的梯度会被缩回去训练会稳定很多。这里给一组我在复现时验证过的超参数参考参数建议值说明seq_len64输入序列长度五言七言诗可以适当降为 32embed_dim128字向量维度词表大可以升到 256hidden_dim256LSTM 隐藏状态维度太大容易过拟合num_layers2LSTM 层数CPU 上建议不超过 2 层batch_size64每批样本数CPU 别超过 128lr0.001Adam 优化器初始学习率dropout0.3过拟合时可调到 0.5epoch30训练轮数根据 loss 收敛情况调整4.2 MyLog.py 的日志机制MyLog.py 是这个项目里容易被忽略但很有价值的部分。它的本质是把 std 输出同时打到控制台和 log.log方便训练完回溯。我自己复写过一个极简版class MyLog: def __init__(self, filename): self.file open(filename, w, encodingutf-8) def write(self, message): print(message) self.file.write(message \n) self.file.flush() def close(self): self.file.close()注意flush()很重要不然训练中途进程 crashlog.log 里最后几行会丢。项目自带的 log.log 我看了看里面应该记录了每轮 loss 和采样样例你可以把它作为“预期 loss 曲线”的参考。如果你的训练曲线和它差太远要么数据没对齐要么学习率设错了。进阶一点你可以在每个 epoch 结束时生成一首诗并记录这样能直观感受 loss 下降带来的生成质量变化。训练早期出来的诗是乱码后期才有点样子。log.log 里每行epoch, loss的格式可以直接被 numpy 读进来。我习惯写个小脚本画 loss 曲线观察是否收敛import matplotlib.pyplot as plt with open(log.log) as f: lines f.readlines() epochs [] losses [] for line in lines: if line.startswith(epoch): parts line.strip().split(,) epochs.append(int(parts[0].split()[1])) losses.append(float(parts[1].split()[1])) plt.plot(epochs, losses) plt.xlabel(epoch) plt.ylabel(loss) plt.show()画出来的曲线如果在中段突然上翘说明学习率太大了要调回来重新跑。4.3 模型保存与加载的两种路径train.py 里保存的是state_dict而不是整个模型对象。这样做的好处是方便在不同 Python 版本间迁移缺点是你加载时必须要先有同样的 model 类定义。我在复现时通常会额外存一个vocab_size到 npz防止换环境后尺寸对不上。加载推理的代码一般是def load_model(path, vocab_size): model PoemModel(vocab_size) model.load_state_dict(torch.load(path, map_locationcpu)) model.eval() return modelmodel_save 目录里如果有训练好的权重你可以直接加载来生成诗歌不用再跑一遍训练。但要注意路径下的文件是否和你的 model.py 结构兼容比如 hidden_dim 不一样就会报 size mismatch 错误。加载时最好先打印一下模型的参数量和权重形状快速确认兼容性。生成诗歌的完整函数也不复杂注意把 hidden 在序列内持续传递def generate(model, start_words, word2idx, idx2word, length20, temperature1.0): model.eval() input_ids [word2idx.get(c, 0) for c in start_words] input_tensor torch.LongTensor(input_ids).unsqueeze(0) hidden None for _ in range(length): logits, hidden model(input_tensor, hidden) # 取最后一个时间步的输出 next_id sample_from_logits(logits[0, -1], temperature) input_tensor torch.LongTensor([[next_id]]) return .join([idx2word[i] for i in input_ids])这里生成时要不断更新input_tensor为上一个预测结果hidden 也一直传递模型才能保持上下文。start_words 决定开头几个字比如“春”。如果你拿到压缩包时 model_save 里已经有权重加载后先跑一首诗看看效果再决定要不要重新训练。有些发布者会用大语料训练效果比本次课程设计的 tang.npz 好很多你直接复用就行。5. 自动写诗避坑指南从 pyc 版本冲突到 loss 不下降的 5 个典型问题说实话我第一次跑这个项目时光是环境就折腾了一个晚上。下面每条都是直接可用的解决方案按概率从高到低排列。5.1 问题一Python 3.6 的 pyc 在 3.8 下报错现象直接运行 train.pyimport 模块时提示Bad magic number in utils或者找不到MyLog.cpython-36.pyc。原因压缩包里出现了MyLog.cpython-36.pyc和utils.cpython-36.pyc这些是 Python 3.6 编译的字节码Python 3.8 的导入器会拒载。虽然.py源文件都在但如果你用 IDE 的缓存目录指向了旧 pyc就可能出问题。解决删除__pycache__目录和所有.pyc让 Python 重新编译源文件。我一般会先执行find . -name *.pyc -delete再清一次__pycache__然后重新运行。在 Windows 上可以直接删__pycache__文件夹Linux 上用rm -rf __pycache__。如果你嫌麻烦也可以在 train.py 开头加两行import sys sys.dont_write_bytecode True这样 Python 不会写新的 pyc但旧的还会在所以最稳妥还是物理删除。5.2 问题二加载 tang.npz 时维度不对现象np.load后打印data.shape发现是(样本数, 65)而模型输入的seq_len是 64或者在训练时和标签切分对不上。原因npz 里存的序列长度是seq_len 1目标字是最后一个位置。如果 train.py 里的seq_len不是 64或者没有正确切分输入和标签就会错位。正确做法是x data[:, :-1]这样 x 是 64y 对应每个窗口的第 65 个字。解决每次加载数据后先打印 shape然后用x data[:, :-1], y data[:, 1:]切分。如果你的数据里混了 padding还要在数据预处理时把 padding 标记出来或者用ignore_index0来规避。另外npz里可能还有vocab_size这个键不要漏了。5.3 问题三训练 loss 不下降或过拟合现象第一个 epoch loss 在 5 以上跑了 10 个 epoch 还是 4 点几或者训练 loss 降到 1.5验证 loss 却从 2.0 升高到 3.0。原因学习率太大或太小dropout 过高或过低。embed_dim太大也会让模型在数据集很小时直接记住训练样本。还有一种情况是数据没有 shuffle模型学到了原始顺序。解决先用一个小 batch 跑 2 个 epoch确认 loss 在下降再把 lr 调到 0.0010.002 之间。如果过拟合要增大 dropout 到 0.5同时把num_layers减到 1。我在这个项目里试过把 hidden_dim 从 256 降到 128验证 loss 反而更稳。另外确认DataLoader里shuffleTrue不然 loss 曲线会周期性震荡。很多课程设计的 train.py 根本不做验证集我建议切分时留出 20% 作为验证集每 5 个 epoch 计算一次验证 losswith torch.no_grad(): val_loss compute_loss(model(val_x), val_y) print(fval loss: {val_loss:.4f})如果训练 loss 下降但验证 loss 上升就是过拟合的典型信号。5.4 问题四生成的诗歌总是重复现象生成的诗里大量出现同一个字或同一句比如“春风春风春风”。原因采样温度太低或者模型过拟合也可能是在生成时用argmax把所有概率集中到了最高值。另外在生成时没有重置 hidden会导致上一句的上下文污染下一句。解决把 temperature 调到 1.0 以上并改成随机采样而非 argmax。另外生成时可以做 top-k 过滤只从概率最高的 50 个候选字里采样能有效减少重复。我把不同 temperature 生成的样例列出来过0.4 时每句基本都是高频字像“春风吹不尽”这类老套组合1.0 时偶尔出现“孤舟听雨寒”这种有画面感的句子1.5 以上就开始出现“石上流泉月照人”之类的乱搭。所以温度不是越高越好需要根据你想要的效果微调。5.5 问题五模型保存文件无法加载或太大现象加载model_save里的权重时报size mismatch或者每个 epoch 保存的文件占了几百 MB。原因保存的是整个模型或者state_dict但模型定义参数不一样或者是保存时没做任何压缩torch.save默认格式较大。另外一个常见原因是 PyTorch 版本差异旧版本保存的权重在新版本里可能报_IncompatibleKeys。解决统一用model.state_dict()保存权重加载前打印一下model的参数量确保 hidden_dim、embed_dim、vocab_size 一致。如果文件太大可以只保存最好的几次权重或者用torch.save(model.state_dict(), path, _use_new_zipfile_serializationTrue)减少尺寸。加载时报size mismatch时最笨但有效的方法是打印模型权重和加载权重的形状对比model PoemModel(vocab_size) loaded torch.load(path) for k in loaded: if k in model.state_dict(): if loaded[k].shape ! model.state_dict()[k].shape: print(f{k}: loaded {loaded[k].shape}, model {model.state_dict()[k].shape})这样能一秒定位是哪个参数对不上。6. 从复现到改进用这个项目试手 Transformer 和 BLEU 评估的落地技巧6.1 用 BLEU/ROUGE 给生成的诗歌打分自动写诗好不好主观性很强但课程设计总要给个量化指标。常见做法是拿一个保留的测试集用 BLEU 或 ROUGE 分数评估生成诗歌和原诗的相似度。我用过最简单的方式是直接调nltk的 BLEUfrom nltk.translate.bleu_score import sentence_bleu, SmoothingFunction def evaluate(generated, reference): # 把字符串转成字列表 gen list(generated) ref list(reference) smooth SmoothingFunction().method4 return sentence_bleu([ref], gen, smoothing_functionsmooth)BLEU 分数在古诗上往往偏低因为原诗的字序和模型生成差别很大。但作为实验报告里的相对指标它至少能证明你的模型在收敛。ROUGE-L 也能算但要看重合度。我习惯同时生成 20 首诗挑出 5 首和人工评估结果一起放进报告。6.2 把生成脚本改成可交互输入项目里的生成可能写死在文件里我复现时改成了一个交互式小脚本每次可以指定开头几个字while True: start input(输入开头字: ) if start exit: break poem generate(model, start, word2idx, idx2word, length20, temperature1.2) print(poem)这样不用一遍遍改代码适合演示和写实验报告。注意生成函数里hidden要在每次调用时重置为 None否则会累积上一个输入的上下文导致开头两句莫名其妙。6.3 从 LSTM 迁移到 Transformer 的边界如果把学有余力想换 Transformer你要知道这不是换个模型类那么简单。Transformer 需要自己构造位置编码注意力掩码以及更长的训练时间。在tang.npz这种几万字符的数据集上Transformer 未必比 LSTM 强多少。我的建议是先保证 LSTM 版本完全跑通再单独开一个分支改 Transformer一旦出现问题可以随时回滚。当前这份资源的核心价值在 LSTM 基线已经足够支撑你的实验报告了。最后说一个我自己的教训拿到任何课程设计压缩包第一件事不是双击运行而是先检查数据文件能不能被正确加载、模型能不能跑通一个 forward。我第一次跑这个项目时图省事直接执行python train.py结果因为__pycache__里的旧 pyc 报错加上seq_len没对齐折腾了一个多小时才排查完。从那以后我每次拆新项目都会强制走一遍清缓存、看数据 shape、跑一个前向、再训练。这个顺序救了我无数次。希望帮到你。本文还有配套的精品资源点击获取

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

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

免费获取报价 →
↑