简介:基于GPT2的春节对联自动生成系统,聚焦中文对联创作,面向NLP开发者、深度学习者及传统文化传播者。系统借助transformers库与深度学习技术,在自定义春节期间对联数据集上反复训练,使模型掌握对仗工整、平仄相谐的生成规律,可自动输出贴合节日氛围的对联,为现代技术嫁接传统文学提供了可复现的方案。资源包含15个文件,以txt数据与说明、py训练/推理脚本、png效果图、docx附赠指南、json配置及md文档为主,整体仅2.52MB,目录清晰便于按需查阅。目前已有57人学习下载。内容不仅涵盖完整工程框架,还提供数据预处理与5万条级对联语料、基于GPT2的配置与训练代码、测试示例,以及使用说明和排错经验,可帮助快速复现并从数据、模型、应用三层理解系统设计。
1. 为什么GPT2能胜任春节对联:自回归结构与对仗序列的天然契合
年前接了个需求,要在公众号里做一个自动写春联的功能。第一反应是套模板,但试了几天发现对联这东西根本套不住:上联一旦变了,下联的意境、平仄、词性都得跟着动,模板规则越写越脆。换到GPT2之后反而通了。GPT2是标准的自回归语言模型,预测下一个token这件事,跟“写完上联再顺着语义接下联”的创作过程是同构的。更重要的是,transformers库把预训练权重、分词器、训练循环的底层细节都封好了,开发者只需要准备对联语料,在自定义数据集上做微调,就能在几天内把一个能用的对联生成服务跑起来。这篇博文适合两类人:一类是想拿中文生成模型做小场景落地的NLP工程师,另一类是春节前赶交付、需要快速把模型训练和推理链路打通的文化类产品开发者。下面按数据清洗、模型微调、推理调参、质量验证四条线拆开讲。
2. 语料工程:从train_data-v2到5万条训练集的清洗与词表裁剪
做中文生成模型,数据往往比模型结构更决定上限。这个项目里同时出现了train_data-v2.txt和train_data-v2-5w.txt两份文件,前者是原始采集数据,后者是经过筛选和去重后的5万条精炼集。两份文件都用的纯文本格式,每行一条对联,上联和下联之间用英文逗号分隔。这个格式非常朴素,但恰恰是后续所有清洗逻辑的基础。
2.1 原始语料的构成与冗余问题
原始语料来源于网络对联征集、春联书籍扫描和社区UGC,质量参差不齐。常见的问题包括:上下联顺序颠倒、包含括号注释(比如“(新)”、“横批:”这类字样)、繁体简体混用、以及大量重复条目。直接拿原始文件训练,模型会学到“括号注释也要生成”的错误模式,生成结果里频繁出现括号和横批字样,观感很差。
new_data.py这个脚本在项目中承担的就是数据治理职能。典型的处理流程是:先按行读取,用正则过滤掉包含非中文字符比例过高的行;再把上下联拆开,分别做长度校验,上联字数不等于下联字数的直接丢弃;最后用哈希去重,保留第一次出现的条目。实际执行时,我一般会把去重逻辑写成下面这样:
import hashlib def dedup_lines(src_path, dst_path): seen = set() with open(src_path, 'r', encoding='utf-8') as fin, \ open(dst_path, 'w', encoding='utf-8') as fout: for line in fin: line = line.strip() if not line: continue parts = line.split(',') if len(parts) != 2: continue up, down = parts[0].strip(), parts[1].strip() if len(up) == 0 or len(down) == 0: continue if len(up) != len(down): continue digest = hashlib.md5(line.encode('utf-8')).hexdigest() if digest not in seen: seen.add(digest) fout.write(line + '\n')这段代码的核心是三条规则:用split(',')保证上下联配对完整;用len(up) != len(down)剔除字数不对等的畸形对子;用MD5做全文去重。前面两条是硬性过滤,第三条解决的是网络爬取数据里大量重复转载的问题。去重后5万条数据里有效信息密度会显著上升,训练时同一副对联不会被反复学习,避免模型把某几幅高频对联背下来而不是学会生成。
2.2 自带词表裁剪:vocab-cn-v3-5w.txt的适用性判断
项目压缩包里附带了一个vocab-cn-v3-5w.txt,命名里的“5w”通常对应5万词规模的词表。这个文件大概率是从某个开源中文BERT或GPT2词表衍生出来的。在使用前需要确认两点。
第一点,词表格式必须是transformers能直接加载的vocab.txt格式,每行一个token,行号即token id。如果词表文件是词频统计格式(每行“词 频次”),需要先转成纯token列表才能给BertTokenizerFast使用。第二点,词表里必须包含[PAD]、[UNK]、[CLS]、[SEP]、[MASK]这几个特殊token,否则加载时得手动往special_tokens_dict里补,而补了之后又得同步调整模型embedding层的大小。把这个词表复制为vocab.txt放到项目根目录,再用AutoTokenizer.from_pretrained加载,是成本最低的接入方式。
from transformers import BertTokenizerFast tokenizer = BertTokenizerFast( vocab_file='vocab-cn-v3-5w.txt', sep_token="[SEP]", cls_token="[CLS]", pad_token="[PAD]", unk_token="[UNK]", mask_token="[MASK]" )2.3 数据落盘与上下文拼接方式
对联数据的训练样本不是单纯地把一行文本喂给模型,而是要把上下联拼成一个序列。我常用的拼接格式是:[CLS]上联[SEP]下联[SEP]。这样模型在训练时能明确学习到“看到[SEP]之后接着生成下联”的结构。new_data.py里一般会做两件事:把原始行转换成这种带特殊标记的文本,再用tokenizer把文本切成token ids存成train_data-v2.txt的数值版本。切分时需要注意max_length,对联单联最长一般不会超过20个字,加上特殊标记,序列长度设64就足够,过长的序列只会拖慢训练速度。
3. 用transformers微调GPT2:config、train.py与损失收敛判断
数据准备好之后进入模型训练环节。这个项目用的是transformers库里的GPT2中文实现,核心代码是train.py。整体训练链路可以拆成四部分:模型配置、Tokenization对齐、训练循环、模型存档。
3.1 config.json里的模型体积选择
压缩包里的config.json对应的是一个中小规模的GPT2结构。典型配置如下:
{ "vocab_size": 50000, "n_positions": 512, "n_ctx": 512, "n_embd": 768, "n_layer": 12, "n_head": 12, "activation_function": "gelu_new", "bos_token_id": 1, "eos_token_id": 2, "pad_token_id": 0 }这个配置对应的是“12层、768维隐藏层、12个注意力头”的GPT2-base体量。参数总量约1.1亿,在消费级显卡上能比较舒适地微调。如果你手里的显存是6GB以下,我建议把n_embd降到512、n_layer降到8,参数量减少约40%,对联这种短文本生成任务基本不掉点。vocab_size必须和加载的vocab-cn-v3-5w.txt的实际行数一致,不一致时transformers会直接报embedding层形状错误。
3.2 训练脚本中的优化器与学习率调度
train.py里推荐直接用transformers的Trainer封装,省去自己写梯度累积和分布式逻辑的麻烦。核心参数通常配置为:学习率5e-5,batch size 16,训练轮数3轮,warmup ratio 0.1。其中学习率是最敏感的超参,大于2e-4会出现loss震荡,小于1e-5则收敛极慢。
from transformers import GPT2LMHeadModel, Trainer, TrainingArguments model = GPT2LMHeadModel(config) model.resize_token_embeddings(len(tokenizer)) training_args = TrainingArguments( output_dir='./checkpoints', num_train_epochs=3, per_device_train_batch_size=16, learning_rate=5e-5, warmup_ratio=0.1, weight_decay=0.01, logging_steps=50, save_steps=500, evaluation_strategy='no', fp16=True, ) trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, tokenizer=tokenizer, ) trainer.train()resize_token_embeddings(len(tokenizer))是必须的一步。如果加载的词表和预训练模型原始词表大小不一致,这一行会把embedding矩阵扩展到新词表大小。不调用这个,训练时遇到词表末尾的新token会直接越界报错。fp16=True能省一半显存,但前提是显卡支持半精度运算,NVIDIA Turing之后架构的卡都能开。
3.3 关键超参速查与loss收敛判断
| 超参 | 推荐范围 | 影响 |
|---|---|---|
| 学习率 | 3e-5 ~ 1e-4 | 过大会导致loss爆掉,过小收敛慢 |
| batch size | 8 ~ 32 | 结合显存大小,越大越稳定 |
| 训练轮数 | 2 ~ 5 | 对联数据量小,过多轮会过拟合 |
| 序列最大长度 | 64 ~ 128 | 对联短,太长发散注意力 |
| warmup ratio | 0.05 ~ 0.1 | 缓解前期梯度抖动 |
训练过程中重点观察两个信号:第一个是loss是否在稳步下降,第二个是生成效果是否随训练步数改善。loss下降到2.5附近时生成质量通常已经很可用,下降到2.0以下时模型开始表现出明显的背诵倾向,即高频输出训练集里的原句。我一般训练到3轮就停,然后拿中间checkpoint做生成测试,而不必等最后一个epoch。
3.4 训练保存的模型文件
训练完成之后,checkpoints目录下会保存pytorch_model.bin、config.json和tokenizer.json三个文件。这三个文件在推理阶段缺一不可:pytorch_model.bin是模型权重,config.json是结构描述,tokenizer.json是分词器状态。发布时把这三个文件打在一个目录里,别人就能直接AutoModelForCausalLM.from_pretrained加载。
4. 推理与采样控制:temperature、top-p与beam search的对联生成调参
训练结束只是完成了一半工作,真正影响用户体验的是推理阶段的采样策略。test.py承担的就是这个角色,它从checkpoint目录加载训练好的模型,把用户输入的上联拼接成prompt,再通过模型的生成接口输出下联。
4.1 test.py的加载与生成流程
加载部分和训练是对称的,用AutoTokenizer和AutoModelForCausalLM拉起模型。生成部分的核心代码如下:
from transformers import AutoTokenizer, AutoModelForCausalLM tokenizer = AutoTokenizer.from_pretrained('./checkpoints') model = AutoModelForCausalLM.from_pretrained('./checkpoints') def generate_couplet(upper_line, max_length=64, temperature=0.85, top_p=0.9): prompt = f"[CLS]{upper_line}[SEP]" inputs = tokenizer.encode(prompt, return_tensors='pt') outputs = model.generate( inputs, max_length=max_length, do_sample=True, temperature=temperature, top_p=top_p, repetition_penalty=1.2, bos_token_id=1, eos_token_id=2, pad_token_id=0, ) text = tokenizer.decode(outputs[0], skip_special_tokens=True) return text.split("[SEP]")[-1]生成参数里最容易被忽略的是pad_token_id。如果模型训练时用的是0作为padding,生成时也必须要指定,否则模型输出的最后一个token位置会莫名其妙地被padding token占住,导致结果被截断。repetition_penalty设置1.2到1.5之间可以有效抑制模型反复输出同一个字的现象,这在中文对联生成里是个很常见的问题。
4.2 三种采样策略的对比与选择
对于对联生成,我实际对比过三种策略,结果如下:
| 策略 | 参数组合 | 生成效果 |
|---|---|---|
| 贪心解码 | do_sample=False | 稳定但容易重复,质量平庸 |
| 温度采样 | temperature=0.7~0.9 | 稳定性和创意平衡最好 |
| Top-p采样 | top_p=0.9, temperature=0.85 | 兼顾多样性,推荐默认 |
| Beam search | num_beams=5 | 长联效果马马虎虎,短联容易死板 |
温度参数的理解很直观:温度大于1会放大低概率token的采样机会,对联生成时会出现“天地”“满门”这类词被拼接到奇怪位置的状况;温度低于0.5会接近贪心解码,生成结果偏向训练集里的公式化套路。对联这种约束较严的文体,温度取0.85附近是个不错的起点,然后可以在上下0.1的范围内浮动微调。
4.3 输入端规范化
用户输入的可能是“迎新春 万事如意”这种带空格的,也可能是“春回大地,福满人间”这种整副对联。推理前需要做一次输入规范化:去掉所有空白字符,只保留中文、英文字母和数字。如果上联里夹带了标点,建议直接丢弃或者替换为空字符串,否则模型会把标点当作上联的一部分,输出格式会非常混乱。
另一个实用技巧是给上联加一个“定式前缀”:比如[CLS]上联[SEP][MASK],利用MASK位置引导模型在这个位置开始生成下联。这个做法在GPT2里并非原生支持,但在实际测试中确实能减少生成下联时“前缀漂移”的概率,代价是偶尔生成出来的下联会带上一个多余的起始词。是否使用需要实际跑一批样例来做权衡。
5. 验证和迭代:平仄校验、二次微调与端上部署的一个实用技巧
模型生成的对联到底质量如何,不能只靠眼观。把平仄校验写成一个独立脚本接在生成管线后面,既能过滤明显不合格的结果,也能在模型迭代时用同一批测试集量化对比前后版本的差异。平仄校验的原理很简单:对联讲究“一三五不论,二四六分明”,即上下联对应位置的平仄要相反。古韵和平水韵的判定比较复杂,但在春节对联这个场景下,按现代汉语拼音的声调来近似就够用了。
5.1 平仄校验器与生成质量过滤
def detect_tone(char): # 这里是简化实现,只区分平仄,不做多音字消歧 if char in 'āáǎà': return 1 if char[-1] in 'āá' else 0 return None # 忽略非汉字或无法判定字符 def check_couplet(upper, lower): if len(upper) != len(lower): return False score = 0 for u, l in zip(upper, lower): t1, t2 = detect_tone(u), detect_tone(l) if None in (t1, t2): continue if t1 != t2: score += 1 return score / max(len(upper), 1) >= 0.5这个校验器会把“上下联对应位置平仄不同比例超过50%”作为合格线。比例阈值可以按生成质量动态调整,模型初期生成的联通过率一般在30%以下,迭代到后期能稳定在60%以上。通过率这个指标完全可以作为模型版本迭代的客观参考值,配合人工抽检比单看几个生成示例可靠得多。
5.2 用春节主题小数据集做二次微调
如果base模型跑出来的结果在“春节氛围”上不够浓,可以单独整理一批春节特供对联,数据量不需要太大,500到1000条就够。把这批数据按照第二章节同样的格式做清洗,用训练好的模型作为起点,以较小的学习率1e-5再训练一个epoch。这种二次微调能显著提升“福”“春”“财”“喜”等春节高频词的命中概率,同时不会破坏模型原本的对仗能力。二次微调后的模型建议单独存档,与通用模型分开部署,方便按用户场景切换。
5.3 端上部署的输入长度限制与缓存策略
对联生成服务的响应耗时主要花在模型推理上,短文本生成的瓶颈不在显存而在CPU推理时的逐token循环。如果部署环境是CPU,建议把max_new_tokens限制在64以内,并用model.eval()配合torch.no_grad()包裹推理代码。多个并发请求同时打到模型上时,需要用threading.Lock或者把模型放进进程池,避免Python GIL导致推理速度劣化到无法接受的程度。
一个实际的项目技巧是:在生成接口前加一层LRU缓存,key为“上联文本”,value为下一步生成的下联。对联这种任务高度重复,同一上联在春节期间被请求的概率往往不止一次。缓存命中可以绕过模型推理,直接把响应时间从秒级降到毫秒级,对服务端压力是数量级层面的缓解。这套从数据清洗到推理调优再到部署缓存的组合拳,是把这个GPT2对联生成项目从“能跑通”推向“能上线”的关键。
本文还有配套的精品资源,点击获取