十八、基于 Transformers 库调用 GPT2 中文诗歌模型实现文本续写实战
·
在 AI 创作领域,利用预训练语言模型实现文本续写是非常经典的应用场景。本文将以 GPT2 中文诗歌模型为例,详细讲解如何使用 Hugging Face 的 Transformers 库快速实现中文文本(诗歌风格)的智能续写,帮助新手理解从模型加载到推理配置的全流程。
一、环境准备
在运行代码前,需先配置好基础开发环境,确保依赖库安装完整:
1. 安装核心依赖
打开终端执行以下命令,安装 Transformers 库(核心)、PyTorch(模型运行框架):
bash
运行
# 安装Transformers库
pip install transformers
# 安装PyTorch(根据系统/显卡适配,此处为通用版)
pip install torch
2. 模型文件准备
本文使用的是gpt2-chinese-poem预训练模型(中文诗歌专用版),需提前将模型文件下载到本地指定路径(示例中为D:\pyprojecgt\flaskProject\langchainstudy\modelscope\gpt2-chinese-poem),模型文件应包含config.json、pytorch_model.bin、vocab.txt等核心文件。
二、完整代码与逐行解析
先展示完整可运行代码,再对核心部分逐一拆解:
完整代码
# 导入Transformers库的核心组件
from transformers import (
pipeline, AutoTokenizer, AutoModelForCausalLM,
BertTokenizer, GPT2LMHeadModel, TextGenerationPipeline
)
# 1. 定义本地模型路径
model_dir = r'D:\pyprojecgt\flaskProject\langchainstudy\modelscope\gpt2-chinese-poem'
# 2. 加载分词器(Tokenizer)
# GPT2中文诗歌模型适配BertTokenizer(中文分词更友好)
tokenizer = BertTokenizer.from_pretrained(model_dir)
# 3. 加载预训练模型
# weights_only=False:允许加载完整的模型权重(新手推荐默认值)
model = GPT2LMHeadModel.from_pretrained(model_dir, weights_only=False)
# 4. 创建文本生成推理管道
text_generator = TextGenerationPipeline(model, tokenizer)
# 5. 配置推理参数并执行续写
out = text_generator(
"[ECS]济南的冬天很冷,", # 续写的初始文本(prompt)
truncation=True, # 截断过长的输入,避免超出模型最大长度
max_new_tokens=80, # 生成新token的最大数量(控制续写长度)
do_sample=True, # 启用采样策略(非贪心生成,更具创造性)
temperature=0.7, # 温度参数(平衡创造性与逻辑性)
top_k=40, # 限制采样范围(仅选概率前40的token)
no_repeat_ngram_size=2 # 禁止2元语法重复(避免冗余)
)
# 6. 输出续写结果
print(out)
核心代码解析
1. 库导入:按需加载核心组件
from transformers import pipeline, AutoTokenizer, AutoModelForCausalLM, BertTokenizer, GPT2LMHeadModel, TextGenerationPipeline
BertTokenizer:中文场景下的分词器,适配 GPT2 中文模型的词汇表;GPT2LMHeadModel:GPT2 的因果语言模型头(专门用于文本生成任务);TextGenerationPipeline:封装好的文本生成管道,简化模型推理流程。
2. 模型与分词器加载
tokenizer = BertTokenizer.from_pretrained(model_dir)
model = GPT2LMHeadModel.from_pretrained(model_dir, weights_only=False)
- 分词器(Tokenizer)的作用:将人类可读的文本转换为模型能理解的数字 token;
weights_only=False:新手无需修改,该参数允许加载模型的完整权重文件,避免权重加载失败;- 本地路径加载:
from_pretrained支持本地路径,无需每次从网络下载模型。
3. 推理管道与核心参数配置
text_generator = TextGenerationPipeline(model, tokenizer)
out = text_generator(
"[ECS]济南的冬天很冷,",
truncation=True,
max_new_tokens=80,
do_sample=True,
temperature=0.7,
top_k=40,
no_repeat_ngram_size=2
)
这是整个代码的核心,重点讲解影响诗歌续写效果的关键参数:
| 参数 | 作用与取值逻辑(诗歌场景适配) |
|---|---|
max_new_tokens |
控制续写长度,80 个 token 适配现代诗 “短而精” 的特点,避免生成冗长内容;若写古体诗可设为 40-60 |
do_sample |
设为True启用采样生成(非贪心),让诗歌更有创造性;设为False则生成结果更固定,但缺乏灵气 |
temperature |
温度越高,生成的内容越跳脱(意象更丰富),越低则越写实;0.7 是诗歌创作的 “黄金值”,兼顾情感与逻辑 |
top_k |
仅从概率前 k 的 token 中采样,40 能过滤掉口语化、无意义的词汇,保证诗歌用词的文艺性 |
no_repeat_ngram_size=2 |
禁止 2 个连续词汇重复(如避免 “下雪的雪,吹风的风” 这类冗余表达),让诗歌更简洁 |
truncation=True |
若输入 prompt 过长,自动截断到模型支持的最大长度,避免报错 |
三、运行效果与输出解读
1. 典型输出示例
运行代码后,输出格式为列表(包含字典),典型结果如下:
[{'generated_text': '[ECS]济南的冬天很冷,风掠过护城河的岸,碎了薄冰,也碎了归人的盼。寒雾漫过老巷的檐,灯影摇着流年,一杯温酒,暖了指尖,却暖不透这北方的天。'}]
2. 输出解析
- 输出是一个列表,每个元素是包含
generated_text的字典; - 若想直接提取续写后的文本,可修改输出代码:
# 提取纯文本结果 result = out[0]['generated_text'] print("续写结果:", result)
四、优化建议(新手拓展)
- 参数调优:
- 想让诗歌更有想象力:将
temperature调至 0.8-0.9,top_k调至 50; - 想让诗歌更规整(如古体诗):将
temperature调至 0.5,no_repeat_ngram_size设为 3;
- 想让诗歌更有想象力:将
- 模型替换:若想生成散文 / 小说风格的文本,可将模型替换为通用版 GPT2 中文模型(如
gpt2-chinese); - 批量生成:通过循环传入不同 prompt,实现批量文本续写,适合创作素材积累。
总结
- 核心流程:加载分词器→加载 GPT2 模型→创建文本生成管道→配置推理参数→执行续写,是 Transformers 库实现文本生成的标准流程;
- 关键参数:
temperature(创造性)、top_k(词汇质量)、no_repeat_ngram_size(避免冗余)是影响诗歌续写效果的核心参数,需根据创作风格调整; - 适配性:GPT2 中文诗歌模型结合 BertTokenizer,能更好地适配中文语境,是新手入门 AI 文本创作的优质选择。
通过本文的实战教程,你可以快速掌握基于 Transformers 库调用预训练模型实现文本续写的方法,后续可结合自己的需求调整参数、替换模型,实现更多风格的 AI 创作。
更多推荐

所有评论(0)