在 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.jsonpytorch_model.binvocab.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)
    

四、优化建议(新手拓展)

  1. 参数调优
    • 想让诗歌更有想象力:将temperature调至 0.8-0.9,top_k调至 50;
    • 想让诗歌更规整(如古体诗):将temperature调至 0.5,no_repeat_ngram_size设为 3;
  2. 模型替换:若想生成散文 / 小说风格的文本,可将模型替换为通用版 GPT2 中文模型(如gpt2-chinese);
  3. 批量生成:通过循环传入不同 prompt,实现批量文本续写,适合创作素材积累。

总结

  1. 核心流程:加载分词器→加载 GPT2 模型→创建文本生成管道→配置推理参数→执行续写,是 Transformers 库实现文本生成的标准流程;
  2. 关键参数:temperature(创造性)、top_k(词汇质量)、no_repeat_ngram_size(避免冗余)是影响诗歌续写效果的核心参数,需根据创作风格调整;
  3. 适配性:GPT2 中文诗歌模型结合 BertTokenizer,能更好地适配中文语境,是新手入门 AI 文本创作的优质选择。

通过本文的实战教程,你可以快速掌握基于 Transformers 库调用预训练模型实现文本续写的方法,后续可结合自己的需求调整参数、替换模型,实现更多风格的 AI 创作。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐