基于BERT微调的古诗生成器:从模型训练到Flask部署全解析
2026/9/23 17:08:53 网站建设 项目流程

简介:一份基于Python的古诗生成器完整工程源码,将后端算法与前端界面整合于一体,适合文学爱好者、编程学习者及AI技术实践者用来体验古诗自动创作并理解前后端协作流程。压缩包共43个文件,包含7个Python脚本、5个XML配置、5个CSS与5个JS前端资源,以及文本、图片、字体等辅助文件,整体约10.85MB;各脚本分别承担生成器初始化、数据加载、模型训练与评估等职责,前端文件则负责页面布局、交互与视觉呈现。目前已有323人学习,可作为功能完整的入门级NLP示例项目。读者可直接运行主入口脚本与页面,查看从数据集处理、模型调用到界面展示的完整链路,并借助配置模块与说明文档进行二次开发,体会Python在古诗词生成场景中的实际应用。

1. 这个古诗生成器,拆开看其实就是一套 BERT 微调样板间

如果你以为古诗生成器是什么高深的“AI 作诗黑匣子”,那这个项目会刷新你的认知:它的核心就是一套标准的中文 BERT 微调流程——加载预训练模型、准备数据集、训练、评估,然后通过 Flask 暴露 HTTP 接口,前端只负责把用户输入传给后端并把返回值渲染到页面上。整套源码 43 个文件,7 个 Python 脚本负责算法主体,前端用 layui + jQuery 搭建,适合想找一个「从模型训练到 Web 部署」完整闭环来练手的开发者。文学爱好者可以把它当成一个有趣的作诗玩具,编程学习者则能从这里看到 NLP 工程落地的常见套路:模型文件从哪来、数据怎么切、接口怎么设计、前端怎么对接。我会把这套项目的运行机制、关键脚本、前端联动和常见坑,按照实际拆解的顺序讲清楚。

2. 项目结构拆解:43 个文件里,真正干活的是哪几个

2.1 文件清单与职责划分:别被配置文件吓到

拿到源码包解压之后,第一眼看到 43 个文件可能会有点懵,但把这些文件按「运行时必需」和「开发期辅助」分成两类,思路立刻清晰。先看一组运行时必需的 Python 脚本:app.py是整个应用的入口,负责启动 Flask 服务;model.py定义模型结构与加载逻辑;utils.py提供工具函数,比如文本清洗、id 转换;dataset.py处理数据集的加载与批处理;train.py触发训练流程;eval.py做生成效果的评估;settings.py管理各类配置参数。这 7 个文件构成完整的「数据处理—模型定义—训练—评估—服务发布」流水线。

另一类文件是支撑模型推理的关键依赖:chinese_L-12_H-768_A-12是谷歌发布的中文 BERT 预训练模型目录,bert_config.json定义了 BERT 的 12 层 Transformer、768 维隐藏层、12 个注意力头这些结构参数,vocab.txt是 BERT 的分词词表。这三个加上模型权重文件,决定了生成器的基础能力,替换成其他预训练模型时,这三个文件必须一起换,否则加载直接报错。

前端部分集中在templates/index.htmlstatic目录下,CSS 样式统一在css/fishc.css,交互逻辑由js/run.js驱动,js/fishc.js里封装了调用后端的封装函数,js/layer.jsjs/jquery.min.js是第三方依赖库。这类项目里前端文件会被很多人忽略,但实际运行时前端渲染逻辑出了问题,后端模型再准用户也看不见结果。

2.2 启动流程一图流:从 Flask 到浏览器

把项目跑起来是第一步。通常做法是先把settings.py里的模型路径和数据路径改成你本机的绝对路径,然后命令行执行:

pip install -r requirements.txt python app.py

requirements.txt里一般会锁定 Flask、torch、transformers、pytorch-pretrained-bert 这几个核心库的版本。启动后 Flask 默认监听127.0.0.1:5000,浏览器访问这个地址就会加载templates/index.html。页面里的输入框接收上句古诗,点击生成按钮,run.js会发起一个 AJAX 请求到后端接口,后端调用加载好的模型执行推理,返回生成的下句与整首诗,前端再把结果渲染到页面上。

这里要特别说明:chinese_L-12_H-768_A-12是 BERT 原始权重,不是 GPT 式的自回归生成模型,所以古诗生成的做法通常是设计一个生成策略——在 BERT 的 MLM(掩码语言模型)框架下,把待生成的位置设为[MASK],然后让模型预测该位置的 token,反复迭代得到完整诗句。这是理解整套代码的关键,也是你后续调参时的理论根基。

2.3 settings.py 里的关键参数:改哪里直接影响生成效果

命令行能启动不代表生成效果好,真正影响结果的是settings.py里的配置。我建议你打开这个文件仔细核对以下参数:

参数名典型值作用调整建议
model_path./chinese_L-12_H-768_A-12预训练模型目录,需包含bert_config.jsonvocab.txt、权重文件路径不能含中文和空格,否则加载报错
data_path./poetry.txt训练语料路径每行一首诗,格式不一致会导致预处理崩溃
max_seq_len128输入序列最大长度五言诗设 64 就够,七言诗设 128
batch_size32训练批次大小CPU 训练调到 8,否则内存撑不住
learning_rate2e-5微调学习率不改,BERT 微调的标准值
num_epochs5训练轮次数据量大可降到 3,防过拟合

settings.py里的device参数也值得注意——如果电脑没有 N 卡,写成cpu,否则 PyTorch 会在 CUDA 初始化时报错,别问我怎么知道的。

3. BERT 古诗生成的核心逻辑:model.py 和 utils.py 里藏着什么

3.1 模型加载与生成策略:为什么不是 GPT 那种逐字写

model.py的核心职责是加载 BERT 预训练模型并封装生成函数。BERT 本身是个双向编码器,它不像 GPT 那样从左到右逐字生成文本,因此古诗生成必须用“填空”的思路。常见做法是:给定上句如「床前明月光」,把整句构造成「床前明月光[MASK][MASK][MASK][MASK][MASK]」的形式,让 BERT 预测每个[MASK]位置上最可能的字,一次填完五个位置。这不代表生成质量一定高,但工程实现比自回归简单得多,而且中文古诗对仗工整,用 MLM 填空反而能利用双向上下文信息。

model.py的核心代码结构大致如下:

class AncientPoetryGenerator: def __init__(self, config_path, model_path, vocab_path, device="cpu"): # 加载 BERT 配置文件 self.config = BertConfig.from_pretrained(config_path) # 从本地目录加载预训练权重 self.model = BertForMaskedLM.from_pretrained(model_path, config=self.config) # 加载词表,用于 id 与 token 的相互转换 self.tokenizer = BertTokenizer.from_pretrained(vocab_path) self.device = torch.device(device) self.model.to(self.device) self.model.eval() # 切换到推理模式,关闭 dropout def predict_next_chars(self, text, mask_positions): # 将输入文本转换成 BERT 需要的 token id 序列 tokens = self.tokenizer.tokenize(text) indexed_tokens = self.tokenizer.convert_tokens_to_ids(tokens) # 构造输入张量,维度为 [1, seq_len] tokens_tensor = torch.tensor([indexed_tokens]) with torch.no_grad(): outputs = self.model(tokens_tensor) predictions = outputs[0] # 形状 [1, seq_len, vocab_size] # 取出每个 mask 位置 top-k 的候选字 result = [] for pos in mask_positions: probs = torch.softmax(predictions[0, pos], dim=-1) top_k = torch.topk(probs, k=10) result.append(top_k.indices.tolist()) return result

这段代码的逻辑路径很清晰:实例化时加载配置、权重和词表,推理时把文本转成 id 序列进模型,拿到每个位置的 logits 后做 softmax,再用topk取概率最高的前 10 个候选字。参数上需要注意的是torch.topk(probs, k=10)里的k代表候选字数量,如果你想让生成结果更可控,可以把k调小到 5 甚至 3,配合前端展示成多个可选方案。

3.2 utils.py 的数据处理:一个字错了整首诗就毁了

utils.py通常承担文本清洗和格式转换工作。古诗数据集的质量参差不齐,有的带标点,有的不带,有的混入作者和题目信息,所以清洗环节决定了模型能不能学到真正有效的韵律特征。常见的处理函数包括:去掉全角空格和特殊符号、统一为简体中文、过滤掉长度异常的句子、把诗句按「上句—下句」配对。

def clean_poem_line(line): """ 清洗单行古诗文本,返回纯诗句字符串。 常见脏数据包括:标题、作者、标点符号、空白字符。 """ # 只保留中文字符和常见标点 cleaned = re.sub(r'[^\u4e00-\u9fa5,。!?、]', '', line) # 去除全角空格 cleaned = cleaned.replace('\u3000', '') # 去除首尾空白 return cleaned.strip() def build_pairs(poem_lines, max_len=64): """ 将清洗后的诗句按两句一组拆成 (上句, 下句) 训练对。 max_len 超过该长度的诗句会被丢弃,防止模型学到过长噪声。 """ pairs = [] for i in range(0, len(poem_lines) - 1, 2): upper = clean_poem_line(poem_lines[i]) lower = clean_poem_line(poem_lines[i + 1]) if len(upper) <= max_len and len(lower) <= max_len: pairs.append((upper, lower)) return pairs

clean_poem_line里的正则只保留\u4e00-\u9fa5这个 Unicode 范围内的中文字符,这会过滤掉日文假名和生僻扩展区汉字,对标准古诗够用但遇到生僻字会被误删。build_pairs按相邻两行配对的前提是数据集中每行就是一句诗,如果你的语料是一整首诗占一行,这个函数就完全不适用,需要先按逗号或句号拆句。

3.3 dataset.py 的数据流:训练数据是怎么喂给模型的

dataset.py在训练阶段负责把清洗后的诗句对转换成模型能消费的张量格式。它的工作流程是:读取poetry.txt→ 对每首诗做上句和下句的分割 → 把上下句拼接成「上句[MASK][MASK]…下句」的输入格式 → 生成input_idstoken_type_idsattention_mask三个张量。其中token_type_ids用来区分前后句,attention_mask用来标记哪些位置是真实 token、哪些是 padding。这部分代码不需要你逐行读懂,但你要知道训练数据长什么样、模型根据什么学习律诗的对仗关系,否则后面排查生成质量问题时完全没有方向。

dataset.py里还可能包含一个create_mask_input函数,它决定了下句的哪些位置被挖掉——是每个位置都挖,还是随机挖一半。这个设计直接决定了训练时模型看到的「残缺程度」,如果每句只挖掉最后几个字,那训练目标只关注结尾的韵脚,中间部分的对仗和意境完全学不到。

4. train.py 训练细节与 eval.py 评估指标:生成质量靠什么保证

4.1 训练流程与损失函数:微调 BERT 不是从头训练

train.py做的事情是在预训练 BERT 的基础上做下游任务微调,损失函数用的是交叉熵损失,目标位置是训练数据中被[MASK]覆盖的 token。这一步不需要更新整个模型的全部参数——最省资源的做法是冻结 BERT 底层参数,只微调顶层和输出层,但这样做古诗生成效果通常一般。我一般会全参数微调,前提是 GPU 显存足够。

def train_epoch(model, dataloader, optimizer, device, clip_grad=1.0): """ 单轮训练:遍历 dataloader,对每个 batch 计算 loss 并回传梯度。 clip_grad 是梯度裁剪阈值,防止梯度爆炸导致 loss 变成 NaN。 """ model.train() total_loss = 0.0 for step, batch in enumerate(dataloader): input_ids = batch['input_ids'].to(device) token_type_ids = batch['token_type_ids'].to(device) attention_mask = batch['attention_mask'].to(device) labels = batch['labels'].to(device) outputs = model(input_ids, token_type_ids=token_type_ids, attention_mask=attention_mask, labels=labels) loss = outputs.loss optimizer.zero_grad() loss.backward() # 梯度裁剪,防止梯度过大 torch.nn.utils.clip_grad_norm_(model.parameters(), clip_grad) optimizer.step() total_loss += loss.item() if step % 100 == 0: print(f"Step {step}, Loss: {loss.item():.4f}") return total_loss / len(dataloader)

这里有个容易踩坑的地方:labels张量中,非[MASK]位置的标签值通常设为-100,因为 PyTorch 的交叉熵损失会自动忽略-100所在的索引。有些初学者会把这些位置设为 0,结果模型被引导去“预测”原始 token 的位置,训练出来的效果一塌糊涂。训练轮次和 batch size 的搭配也需要小心:settings.py里如果num_epochs=5batch_size=32,在 8G 显存上很容易 OOM。常见处理是安装梯度累积插件或者在train.py里手动实现梯度累积——每 4 个小 batch 更新一次参数,等效于 128 的 batch size 但显存压力小得多。

4.2 eval.py 的评估逻辑:别只盯着 loss 数字

评估模块最容易被人忽略,但它才是判断「生成的诗句像不像诗」的直接依据。eval.py通常会实现三种指标:loss 数值、top-k 准确率、人工抽检的对比队列。loss 数值只能反映模型在验证集上的拟合程度,而 top-k 准确率衡量的是「正确答案是否出现在模型预测的前 k 个候选字里」。这两个指标结合起来,才能判断模型到底是真学到了韵律规则,还是只是把训练数据背了下来。

我在实际评估这个项目时,遇到过一个很典型的现象:loss 降到 1.2 左右就再也不动,看起来是收敛了,但生成的句子读起来毫无诗意。原因在于训练数据的 poem 对如果清洗得太狠,把标点全删了,模型学不到句读的停顿节奏,生成出来的句子就是一堆语义通顺但没有平仄韵律的字。所以eval.py里最好加一个「标点重现率」的检查——统计模型预测结果中逗号句号的分布是否符合五言、七言诗的断句规律。如果有条件,保留一份原始带标点的语料和清洗后的语料做对比,能快速定位是数据问题还是模型问题。

4.3 训练后的产物保存与模型加载:除了 PyTorch 原版格式还要导出什么

训练完成后,torch.save(model.state_dict(), 'model.pt')是最直接的保存方式,但这只保存了模型参数,没保存配置和词表。为了在app.py里快速加载,我一般会同时保存三样东西:

# 保存模型参数 torch.save(model.state_dict(), 'trained_model.pt') # 保存模型的配置文件副本,防止后续加载时路径找不到 model.config.to_json_file('trained_config.json') # 保存词表映射,方便推理时直接使用 tokenizer.save_vocabulary('./')

这段代码看起来平平无奇,但实际部署时很多新手只保存了权重文件,然后在app.py里用BertForMaskedLM.from_pretrained('./trained_model.pt')加载,直接报错说缺少配置文件。正确做法是确保trained_model.pttrained_config.jsonvocab.txt三个文件在同一个目录下,推理时用model_path指向这个目录。

5. 前端集成与 Flask 接口设计:run.js 是怎么把诗句画到页面上的

5.1 app.py 的路由设计:一个接口撑起整个交互

app.py是这个项目的前后端连接器,它用 Flask 定义了两个路由:GET /返回index.html页面,POST /generate接收用户输入的上句,调用模型生成下句并返回 JSON。前端不直接操作模型,所有计算都在后端完成,这是这类项目最基本的架构约束。

@app.route('/generate', methods=['POST']) def generate(): """ 接收前端传来的 JSON 请求,格式为 {"prompt": "床前明月光", "top_k": 10} 返回格式为 {"poem": "疑是地上霜", "candidates": ["疑是地上霜", ...]} """ data = request.get_json() prompt = data.get('prompt', '') top_k = data.get('top_k', 10) if not prompt or len(prompt) > 10: return jsonify({'error': '上句长度需在 1-10 个字之间'}), 400 # 调用模型生成下句 candidates = generator.predict_next_chars(prompt, mask_positions=range(len(prompt), len(prompt) + 5), top_k=top_k) return jsonify({'poem': candidates[0], 'candidates': candidates})

这个接口的参数设计有几个细节值得注意:len(prompt) > 10的限制是因为模型输入长度太大时,推理时间会线性增长,而且古诗通常是五言或七言,超长上句本身就不符合生成场景。top_k从请求体里读取而不是写死在代码里,给前端留了调参空间,用户可以在页面上选择“严谨”或“创意”模式,对应不同 top_k 值。返回的candidates是一个列表,前端拿到后可以展示一个下拉框或者点击换一组,这是一个加分交互。

5.2 index.html 的页面骨架与 run.js 的交互逻辑

前端页面的核心不复杂,但它的文件组织是很多初学者容易搞混的地方。index.html引用了css/fishc.css做整体样式,js/jquery.min.jsjs/layer.js是基础库,js/run.js是业务逻辑。run.js中封装了一个generatePoem函数,核心流程是:读取输入框的值 → 组装 JSON 数据 → 发送 AJAX 请求 → 把返回结果填充到两个 DOM 节点——一个显示完整诗句,另一个显示候选列表。

function generatePoem() { let prompt = $('#prompt-input').val().trim(); if (prompt.length === 0) { layer.msg('请输入一句上联或上句', {icon: 0}); return; } $.ajax({ url: '/generate', type: 'POST', contentType: 'application/json', data: JSON.stringify({prompt: prompt, top_k: 5}), success: function(res) { if (res.error) { layer.msg(res.error, {icon: 2}); } else { $('#result-poem').text(prompt + res.poem); renderCandidates(res.candidates); } }, error: function() { layer.msg('服务器开小差了,请检查 app.py 是否在运行', {icon: 2}); } }); }

这里有一个前端研发常踩的坑:contentType写成application/json时,后端必须用request.get_json()解析,如果后端用的是request.form大概率拿到空值。反过来前端用表单格式提交、后端用 JSON 解析也一样报错。layer.msg是 layui 的弹窗组件,它的样式依赖css/layui.cssjs/layer.js两个文件,只引用了 JS 没引用 CSS 会导致弹窗乱成一行文字。

5.3 前后端联调时的接口规范和异常处理

前后端联调是这类项目最容易翻车的阶段,问题集中在三个方面:跨域、请求格式、异常反馈。如果一个 Dev 把app.py跑在 5000 端口、前端静态页面直接从文件系统打开(file://协议),浏览器会直接拦截跨域请求,这时需要后端加上CORS中间件或者让前端也通过 Flask 的 5000 端口访问。请求格式不匹配的问题上面说了,解决方法是前后端约定死一个JSONschema,并在eval.pyapp.py里打印请求体做日志。

异常反馈的设计也很重要——很多初学者在模型推理出错时,后端直接抛 500,前端只会看到一堆 ChunkedEncodingError。我一般会在app.py里加一个全局异常捕获,把模型的报错信息转成 JSON 返回给前端,这样至少知道是 GPU 显存不足还是 key 拼写错误。

6. 避坑指南:古诗生成器部署中的五个高频故障

6.1 坑一:BERT 模型加载报错「路径不存在」

现象:执行app.py时提示Model name 'chinese_L-12_H-768_A-12' was not found
原因:transformers库的from_pretrained会自动检查传入路径是否是本地目录,有时会误以为chinese_L-12_H-768_A-12是 Hugging Face 模型库中的模型名,而且合成路径错误。
解决:确认settings.py里的model_path是绝对路径,并且模型目录下确实存在bert_config.jsonvocab.txt和权重文件。不需要使用 os.path.abspath 也能跑通,但用绝对路径是排查这个问题最快的办法。

6.2 坑二:Windows 环境下编码报错

现象:读取poetry.txt时报UnicodeDecodeError: 'gbk' codec can't decode byte,或者生成结果全是乱码。
原因:Windows 环境中 Python 的默认读写编码是 GBK,而poetry.txt通常是 UTF-8 编码。
解决:所有涉及文件读写的操作,显示指定编码:

with open('poetry.txt', 'r', encoding='utf-8') as f: lines = f.readlines()

这一步虽然简单,但几乎每次部署到新环境都会遇到,我已经把它写进项目部署 checklist 的第一条了。

6.3 坑三:训练的时候 Loss 变成 NaN

现象:训练到某个 step 时 loss 突变为nan,之后所有数值都是nan
原因:最常见的原因是学习率过大导致梯度爆炸,其次是数据里有长度为 0 的诗句,导致 masked 位置没有有效的 label。
解决:先检查poetry.txt里是否有空行和单字行,清洗时过滤掉len(text) < 4的行。如果数据没问题,把settings.py里的learning_rate2e-5降到1e-5,或者在train.py里加上梯度裁剪。我的习惯是用clip_grad_norm_(model.parameters(), 1.0)兜底,它能解决绝大多数 nan 问题。

6.4 坑四:生成的句子「驴唇不对马嘴」——语义不连贯

现象:训练完跑出结果,上句「白日依山尽」,下句是「青山横北郭」,单独看每个字都正常,但两句之间完全没有对仗关系。
原因:数据集清洗时把标点符号和断句信息全删了,模型学到的是「字级别的组合概率」,而不是「诗句级的对仗结构」。
解决:检查utils.py里的clean_poem_line,不要删掉逗号和句号,BERT 词表里本身就有这些标点的 token。训练数据保留标点后重新训练一下,生成结果的质量会有明显提升——这是我从 NER 任务迁移过来的经验,标点在中文 NLP 里从来不是噪声。

6.5 坑五:GPU 训练时显存溢出(OOM)

现象:训练刚开始第一个 batch 就报CUDA out of memory
原因:batch_size设置得太大,或者max_seq_len太长,导致中间激活值占用过多显存。
解决:先把max_seq_len从 128 降到 64,再降batch_size从 32 到 8。如果还是不行,在train.py里增加梯度累积逻辑——每 4 个 batch 累积一次梯度,等效 batch size 不变但显存峰值降到原来的四分之一。还能做的就是把模型从 fp32 转成 fp16 混合精度,但 BERT 微调场景下 fp16 容易掉点,不如前两个方案保险。

7. 把生成器往实用方向推:换数据集、调温度、部署到公网

这个项目给人最大的发挥空间是「换数据」。现在内置的poetry.txt可能只有几百首常见古诗,生成结果容易撞车——同一个上句生成的下句永远是那几个高频组合。解决办法是换上更全的《全唐诗》数据集,大概 5 万首以上,每行格式保持「诗句,下句」的划分,清洗逻辑就可以完全复用。换数据之后需要重新跑train.py,训练时间会拉长,但生成结果的多样性和新颖度会好很多。

第二个可调的位置是采样温度temperature。当前model.py里用的是topk截断,没有温度系数。加入温度很简单,在torch.softmax计算前把 logits 除以 temperature:

temperature = 0.8 # 值越小生成的句子越保守,越大越有创意,0.8 是古诗场景的折中值 probs = torch.softmax(predictions[0, pos] / temperature, dim=-1)

temperature 大于 1 时会增加低概率词被选中的机会,生成结果更「出格」;小于 1 时会集中在高概率词上,结果更工整。我的经验是七言绝句用 0.8、五言用 0.7 比较合适,太高了容易生成不通顺的生僻词。

如果想让局域网里的其他设备访问,把app.py的启动参数从app.run()改成:

app.run(host='0.0.0.0', port=5000, debug=False)

这样同一局域网内的手机和电脑就能通过http://你的IP:5000访问生成器了,不用每次都挤在同一台电脑前演示。调试的时候可以开着 debug=True,能自动重载代码,但真正常时间跑服务必须关掉,否则接口异常信息会直接泄露在页面上。

以前接手这类 NLP 小项目总喜欢直接改代码,后来被模型路径和编码问题连续坑了几次,从那以后每次部署都强制走一遍「路径检查 → 编码检查 → batch_size 试探」的固定流程,这份源码本身写得不算花哨,但作为 BERT 微调和前后端联动的参照物,值得花一个晚上把每个脚本的输入输出理一遍。希望这篇拆解能帮你在复现和二次开发时少走几步弯路。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询