1. 项目背景与核心目标
在自然语言处理领域,故事生成与问答能力一直是极具挑战性的任务组合。传统方法通常将这两个任务分开处理,导致生成的故事缺乏可解释性,而问答模型又难以理解复杂叙事逻辑。Qwen3-8B作为通义千问系列的最新开源模型,其80亿参数的规模在单卡环境下展现出独特的性价比优势。
这个项目的核心创新点在于:通过精心设计的微调方案,让同一个模型同时掌握"会讲故事"和"能回答故事问题"两种能力。具体来说,我们需要模型能够:
- 根据简短的提示生成情节连贯、风格统一的故事文本
- 准确理解故事内容并回答基于文本细节的各类问题
- 在长程叙事中保持角色和逻辑的一致性
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 数据集选型与特性分析
2.1 TinyStories数据集
这个由微软研究院发布的数据集包含近500万条样本,全部使用3-4岁儿童能理解的词汇(约1500个基础单词)。其核心价值在于:
- 叙事模板丰富:包含日常生活、童话、科普等多种题材
- 语言风格统一:所有样本都采用规范的JSON格式
- 长度分布合理:平均token数在200-300之间
实际使用中我们发现,直接微调全量数据会导致严重的过拟合。我们的解决方案是:
- 按主题进行分层抽样,保留各题材的代表性样本
- 对长尾分布进行截断,去除极端长度的样本
- 添加5%的噪声数据增强鲁棒性
2.2 WritingPrompts数据集
这个经典数据集包含30万条"短提示→长故事"的配对样本,其独特价值体现在:
- 高压缩比:平均23个token的提示能生成544个token的故事
- 复杂叙事:包含多角色互动、情节转折等高级叙事技巧
- 风格多样:从现实主义到奇幻题材应有尽有
我们特别注意到数据中存在[WP]前缀污染问题。清洗策略包括:
- 统一去除所有样本中的[WP]标记
- 对提示词进行词干提取和停用词过滤
- 建立基于TF-IDF的重复样本检测机制
2.3 FairytaleQA数据集
这个教育领域标杆数据集的特点在于:
- 结构化监督:问题与故事片段通过cor_section精确关联
- 问题类型丰富:包含显式/隐式、局部/总结等多种题型
- 教育学标签:每个问题都标注了认知难度和考察重点
在实际处理时,我们将其转换为标准的QA格式:
code复制{
"instruction": "为什么小红帽要穿过森林?",
"input": "[相关故事段落文本]",
"output": "为了给生病的外婆送食物"
}
3. 微调方案设计
3.1 硬件配置优化
在单张RTX 3090/4090上的关键配置策略:
- 采用LoRA适配器技术,将可训练参数控制在1%以内
- 动态batch策略:根据序列长度自动调整micro batch大小
- 梯度累积步数设为8,确保有效batch size达到32
显存占用实测数据:
| 配置项 | 2048长度 | 1536长度 |
|---|---|---|
| 全参数 | OOM | OOM |
| LoRA32 | 18GB | 14GB |
| LoRA64 | 21GB | 16GB |
3.2 三阶段训练流程
阶段A:基础叙事能力构建
- 数据:纯TinyStories采样数据(约50万条)
- 关键参数:
bash复制
lora_rank=32 lr=1e-4 max_length=1024 - 训练目标:最小化生成文本的困惑度(PPL)
阶段B:复杂叙事增强
- 数据:TinyStories + WritingPrompts按4:6混合
- 调整策略:
- 逐步提高max_length到1536
- 添加长度惩罚项:
length_penalty=1.2 - 引入课程学习,先训练中等长度样本
阶段C:问答能力对齐
- 数据:三数据集按3:5:2比例混合
- 特殊处理:
- 对QA样本添加10%的负例(错误答案)
- 使用Focal Loss处理类别不平衡
- 学习率降至5e-5防止灾难性遗忘
4. 关键实现细节
4.1 损失函数设计
除了标准的交叉熵损失,我们还引入了:
- 连贯性奖励:通过预训练模型计算段落间语义相似度
python复制def coherence_loss(outputs): segments = split_into_paragraphs(outputs) embeds = [model.encode(s) for s in segments] return -cosine_similarity(embeds[:-1], embeds[1:]).mean() - 问答一致性惩罚:确保生成答案与问题类型匹配
- 多样性正则项:基于n-gram重复率的动态权重调整
4.2 动态掩码策略
针对不同任务类型采用差异化的label masking:
- 生成任务:仅计算output部分的loss
- 问答任务:同时计算问题和答案的loss
- 混合样本:根据样本类型自动切换mask模式
5. 评估与优化
5.1 自动化评估指标
我们构建了多维度的评估体系:
| 维度 | 指标 | 目标值 |
|---|---|---|
| 生成质量 | PPL | <15 |
| Distinct-2 | >0.4 | |
| 问答准确 | ROUGE-L | >0.6 |
| BLEU-4 | >0.5 | |
| 一致性 | Hallucination Rate | <5% |
5.2 人工评估方案
设计了三层评估机制:
- 基础检查:故事是否通顺、问答是否相关
- 深度分析:角色行为是否符合逻辑、答案是否有文本依据
- 压力测试:故意提供矛盾提示检验模型鲁棒性
6. 部署实践
6.1 推理加速技巧
在实际部署中发现的关键优化点:
- 使用vLLM的PagedAttention技术,将吞吐量提升3倍
- 对长故事生成采用分块处理,避免OOM
- 问答任务启用精确的attention mask,减少计算量
6.2 实用调试命令
记录几个常用诊断命令:
bash复制# 显存监控
nvidia-smi -l 1
# 性能剖析
nsys profile --stats=true python infer.py
# 精度检查
torch.autograd.set_detect_anomaly(True)
7. 典型问题排查
在实际运行中遇到的几个典型问题及解决方案:
-
显存溢出
- 现象:训练中途突然崩溃
- 排查:发现是max_length设置不统一
- 修复:对所有数据集进行长度标准化
-
模式坍塌
- 现象:生成故事重复率升高
- 排查:温度参数设置过高
- 修复:采用动态温度调度策略
-
问答偏差
- 现象:答案总是偏向某种类型
- 排查:数据分布不均衡
- 修复:引入样本加权采样
经过三个月的迭代优化,最终模型在保持基础叙事能力的同时,问答准确率提升了42%。特别是在处理隐式问题(如"为什么角色会这样做?")时,模型展现出了令人惊喜的推理能力。这个项目证实了通过精心设计的微调方案,中等规模模型也能完成复杂的多任务学习。
