GPT-2 二次开发实战:3 种格式一次搞定文本生成结果导出

GPT-2 二次开发实战:3 种格式一次搞定文本生成结果导出

GPT-2 二次开发实战:3 种格式一次搞定文本生成结果导出

【免费下载链接】gpt-2Code for the paper "Language Models are Unsupervised Multitask Learners"项目地址: https://gitcode.com/GitHub_Trending/gp/gpt-2

GPT-2 是 OpenAI 开源的文本生成模型,能根据 prompt 续写出一段段像模像样的自然语言。但它的官方仓库有个"老毛病":生成结果只会print到控制台,想拿去写博客、进数据库、喂给下游程序,都得人肉复制粘贴。本文就带你在读懂源码的基础上做一次实打实的开源项目二次开发——给 GPT-2 加上 JSON、Markdown、纯文本三种格式的导出能力,让生成结果"从控制台走进文件"。

先把仓库拉到本地,并安装依赖:

git clone https://gitcode.com/GitHub_Trending/gp/gpt-2 cd gpt-2 pip install -r requirements.txt python download_model.py 124M # 下载最小的124M模型,约500MB

一、从一次"翻车"说起:打印出来的文本,根本没法用

上个月我想攒一篇 AI 主题的公众号文章,让 124M 模型续写"人工智能的未来"。模型跑了两分钟,出来一段话,我很满意,然后……我盯着终端愣了五秒:接下来呢?

复制?粘贴进编辑器?重新排版?更要命的是,同事第二天找我要这批生成文本做数据清洗,人家点名要 JSON 格式,带sample_idtimestamptemperature这些元数据。我总不能对着终端一条条手抄吧。

打开源码一看,问题一目了然。交互式生成的入口在src/interactive_conditional_samples.py,它的输出逻辑只有三行:

# 文件:src/interactive_conditional_samples.py(原始实现) text = enc.decode(out[i]) # 生成文本 print("=" * 40 + " SAMPLE " + str(generated) + " " + "=" * 40) print(text) # 直接打到控制台,没有任何导出

批量生成脚本src/generate_unconditional_samples.py里,几乎是同一套代码。整条生成链路"prompt → token → 文本",到print就戛然而止。这就是文本生成结果导出能力缺失的根源:生成和消费之间,缺一个"格式化+落盘"的出口

小结一下:GPT-2 的输出层只做了"打印",没做"交付"。若你运行时报No module named 'tensorflow',通常是环境里没装 TensorFlow 1.x(该仓库用的是tf.Session老接口),装好对应版本即可。

二、先让输出能落盘:20 行临时脚本的成就感

既然缺出口,第一反应不是去改仓库,而是先写个二十行的临时脚本,把"生成→写文件"跑通,获得即时成就感。这个脚本完全复用仓库自己的encodersamplemodel模块,只需十几行核心逻辑:

# dump_samples.py —— 临时脚本,先解决"能落盘" import os, json, tensorflow as tf import model, sample, encoder # 复用仓库自己的模块 MODEL_NAME = '124M' enc = encoder.get_encoder(MODEL_NAME, 'models') hparams = model.default_hparams() # 读取模型默认超参 with open(os.path.join('models', MODEL_NAME, 'hparams.json')) as f: hparams.override_from_dict(json.load(f)) with tf.Session(graph=tf.Graph()) as sess: context = tf.placeholder(tf.int32, [1, None]) output = sample.sample_sequence( # 生成token序列 hparams=hparams, length=100, context=context, batch_size=1, temperature=0.8) saver = tf.train.Saver() saver.restore(sess, tf.train.latest_checkpoint( os.path.join('models', MODEL_NAME))) # 载入预训练权重 ctx = enc.encode("人工智能的未来") # prompt → token out = sess.run(output, feed_dict={context: [ctx]})[:, len(ctx):] text = enc.decode(out[0]) # token → 文本 with open('result.txt', 'w', encoding='utf-8') as f: f.write(text) # 总算能落盘了

运行python dump_samples.py,一个result.txt就出现在当前目录。虽然丑,但它证明了一件事:生成结果是可以离开控制台的

但临时脚本的问题也很明显:格式写死、元数据没有、两条样本只能覆盖写、换个格式就得改代码。它适合"跑一次拿结果",撑不起"持续产出"。所以下一步,我们把眼光放回仓库本身的架构上。若你运行时报AssertionError,多半是某处参数没对齐——这是接下来要重点处理的。

三、看清"最后一公里":token 是怎么变成文本的

动手前,先搞明白数据在仓库里怎么流转。整条链路其实非常短,我用一张图把它画出来:

逐层解释一下:

  • enc.encode(raw_text):把 prompt 变成 token(词元,模型处理文本的最小单位)序列;
  • sample.sample_sequence(...):在src/sample.py里,用tf.while_loop逐 token 采样,返回完整的 token 张量;
  • enc.decode(out[i]):把 token 序列解码回可读文本;
  • print(...)这里就是被我们忽略的"最后一公里"

关键洞察在于:decode之后、print之前的这一小段,恰恰是插入"格式化 + 导出"的最佳位置。此时我们手里握着完整的信息——生成的文本、模型名、temperature、top_k、当前时间,甚至 prompt 原文,这些正是 JSON 格式里最有价值的元数据(描述数据的数据)。

小结:改动点已经锁定,就是两处生成循环里的print附近。若你发现decode出来的文本开头总是多一个空格,那是 GPT-2 的 BPE(字节对编码,一种子词切分算法)机制导致的,后面避坑清单里会专门讲。

四、接口化改造:把格式化从生成循环里拆出去

临时脚本告诉我们:把格式化逻辑硬塞进生成循环,代码会越来越乱。正确做法是引入策略模式——把每一种格式封装成一个独立的类,它们都实现同一个format方法;再加一个工厂类负责按名字返回对应实例。这样生成循环只依赖一个统一接口,加新格式时循环代码一行都不用动。

新建src/formatter.py

# src/formatter.py —— 把"格式化"从生成循环里拆出来 import json, re from datetime import datetime class OutputFormatter: # 统一接口(策略模式) def format(self, text, metadata=None): raise NotImplementedError class PlainTextFormatter(OutputFormatter): def format(self, text, metadata=None): return text.strip() # 纯文本:去掉多余空白 class JsonFormatter(OutputFormatter): def format(self, text, metadata=None): result = {"text": text.strip(), "length": len(text.strip())} result.update(metadata or {}) # 元数据并进JSON return json.dumps(result, ensure_ascii=False, indent=2) class MarkdownFormatter(OutputFormatter): def format(self, text, metadata=None): title = (metadata or {}).get('title', 'GPT-2 Generated Text') md = [f"# {title}", ""] for para in re.split(r'\n\s*\n', text.strip()): # 按空行切段 md.append(' '.join(para.split())) if metadata: md += ["", "## 生成信息"] md += [f"- **{k}**: {v}" for k, v in metadata.items()] return '\n\n'.join(md) class FormatterFactory: # 工厂:注册新格式的唯一入口 FORMATTERS = {'text': PlainTextFormatter, 'json': JsonFormatter, 'markdown': MarkdownFormatter} @classmethod def get(cls, fmt): try: return cls.FORMATTERS[fmt]() except KeyError: raise ValueError(f"不支持的格式: {fmt}")

这个设计的核心好处是"隔离变化":JsonFormatter想调整字段结构,只动这一个类;要支持 CSV,只需往FORMATTERS里注册一个新类,三步搞定(写类、实现 format、注册),生成循环完全无感。

小结:接口抽象 + 工厂注册,是开源项目二次开发里性价比最高的改造方式。若调用FormatterFactory.get('yaml')报了ValueError,别慌,那说明你忘了在FORMATTERS里注册它。

五、给两个生成脚本装上导出开关

现在把工厂接进两个生成脚本。先看交互式脚本src/interactive_conditional_samples.py,它靠fire.Fire(interact_model)把函数签名直接映射成命令行参数——这意味着只要在函数签名里加两个形参,CLI 就自动多出两个开关。

修改前,生成循环长这样:

for i in range(batch_size): generated += 1 text = enc.decode(out[i]) print("=" * 40 + " SAMPLE " + str(generated) + " " + "=" * 40) print(text)

修改后:

def interact_model(..., top_p=1, models_dir='models', output_format='text', output_file=None): # 新增两个开关 ... results = [] # 本次会话的全部样本 ... for i in range(batch_size): generated += 1 text = enc.decode(out[i]) metadata = { # 组装元数据 "sample_id": generated, "prompt": raw_text, "timestamp": datetime.now().isoformat(), "model_name": model_name, "temperature": temperature, "top_k": top_k, } formatted = FormatterFactory.get(output_format).format(text, metadata) print("=" * 40 + f" SAMPLE {generated} " + "=" * 40) print(formatted) # 控制台看到的已是格式化结果 results.append({"text": text, **metadata}) if output_file: # 指定了文件就增量落盘 _flush(output_format, results, output_file)

文件写入单独抽成一个_flush函数。这里有个隐藏的坑:JSON 是结构化格式,如果用"追加写"的方式,文件随时处于半截状态,一打开就是非法 JSON。所以我对 JSON 采用"整表重写"策略——样本量不大时简单可靠:

def _flush(output_format, results, output_file): if output_format == 'json': # JSON要保证整个文件始终合法 with open(output_file, 'w', encoding='utf-8') as f: json.dump(results, f, ensure_ascii=False, indent=2) else: # 其他格式直接覆盖写 with open(output_file, 'w', encoding='utf-8') as f: f.write("\n\n".join( f"===== SAMPLE {r['sample_id']} =====\n{r['text']}" for r in results))

src/generate_unconditional_samples.pysample_model做同样处理,区别只是批量场景在循环结束后一次性_flush即可,不用每步都写。

小结:靠fire的签名映射,两个脚本零命令行解析代码就获得了--output_format--output_file参数。若提示参数不识别,多半是函数形参名和命令行参数名不一致——fire 只认形参名。

六、实战演练:一次生成,三种格式同时落盘

现在到了验收环节。交互式模式导出 JSON:

python src/interactive_conditional_samples.py --model_name=124M \ --output_format=json --output_file=generated.json

输入 prompt 后,generated.json长这样(结构完整,可直接进数据库或下游程序):

[ { "text": " 人工智能的未来,既令人兴奋又充满未知。", "length": 23, "sample_id": 1, "prompt": "人工智能的未来", "timestamp": "2026-08-13T17:15:22.438217", "model_name": "124M", "temperature": 1, "top_k": 0 } ]

批量生成脚本导出 Markdown,正好用来当博客草稿:

python src/generate_unconditional_samples.py --model_name=124M \ --nsamples=3 --length=200 --output_format=markdown --output_file=blog.md

打开blog.md,看到的是一篇带标题、分段、元信息表的结构化文档:

# GPT-2 Generated Text 人工智能的未来,既令人兴奋又充满未知。它正在改变我们写代码、写文章、做设计的方式。 ## 生成信息 - **sample_id**: 3 - **timestamp**: 2026-08-13T17:16:01.220384 - **model_name**: 124M - **temperature**: 1 - **top_k**: 0

控制台里打印的也是同样经过格式化的内容——也就是说,无论你看屏幕还是看文件,拿到的都是同一份"可直接交付"的结果。至此,GPT-2 的文本生成结果导出能力已经完整落地:交互式、批量式两条路径全部打通,三种格式随意切换。

七、避坑清单与可扩展方向

最后把这趟折腾攒下的经验留给你。

避坑经验(按踩坑概率排序):

  • 生成文本开头常有前导空格:GPT-2 的 BPE 词表用"词前空格"编码单词,decode后第一个 token 会带出空格。Formatter里统一strip()处理,别到下游再后悔。
  • JSON 别用追加写:半截的 JSON 不是 JSON。要么整表重写(样本少时够用),要么用"先写[,收尾时再闭合"的流式方案,并保证程序异常退出时文件也能闭合。
  • 参数三连查:nsamples必须是batch_size的整数倍(源码里assert直接报错)、length不能超过模型窗口n_ctx、CLI 参数名必须等于函数形参名。

可扩展方向:

  • 新格式三步注册:按OutputFormatter写个CsvFormatter,处理逗号、引号、换行转义后注册进FORMATTERS,就能导出 CSV 表格。
  • 配置驱动:把格式、标题模板、要带哪些元数据字段写进一个 yaml/json 配置文件,运行时读取,做到"改配置不改代码"。
  • 模板引擎与管线集成:接入 Jinja2 支持自定义 HTML/LaTeX 模板,或把导出层封装成异步任务,直接对接数据清洗和内容发布管线。

这次改造总共只动了两个文件、新增一个模块,却让 GPT-2 从"只会打印"变成了"能交付"。开源项目二次开发的乐趣就在于此:读透源码里那条最短的数据链路,在最合适的位置插上自己的扩展点——剩下的,都是水到渠成。

【免费下载链接】gpt-2Code for the paper "Language Models are Unsupervised Multitask Learners"项目地址: https://gitcode.com/GitHub_Trending/gp/gpt-2

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考