ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

继续预训练实战指南:让大模型真正理解行业知识

继续预训练实战指南:让大模型真正理解行业知识 上个月有位做电力行业系统的朋友找到我说他们手上有几十万份设备运行日志、检修工单和行业规程想让开源通用大模型更懂电力术语一上来就问“是不是该做指令微调SFT”。这个问题我最近被问了很多次答案其实要分场景如果你的目标是让模型真正吸收行业知识、读懂专业表述而不是单纯学会答题格式那么用的路子不应该是SFT而是继续预训练Continued Pre-Training简称CPT。这篇实战指南就是来聊这件事的。CPT说白了是在通用大模型的基础上用领域语料继续做自监督训练让模型把行业知识“内化”到参数里而不是像微调那样只学表层指令格式。文中会覆盖技术原理、算力与成本评估、数据清洗、训练参数配置、监控调优和常见坑点适合企业算法工程师、技术负责人以及准备往行业大模型方向走的开发者参考。1. 为什么企业要自己训练行业模型通用模型不是万能钥匙1.1 通用模型在行业场景里的三道坎市面上开源通用大模型的能力已经很强但在真正落到具体行业时普遍绕不开三个问题。第一是知识缺口。通用模型训练时用的是公开互联网数据很多行业内部资料、非公开文档、细颗粒度的专业规范根本不在训练集里。比如医疗场景的临床路径文本、法律场景的裁判文书细节、工业场景的设备故障案例这些内容模型没见过自然答不准。第二是术语理解能力弱。行业里的黑话、缩写、复合专业词通用模型经常理解不到位。同样是“卡脖子”在制造业和日常生活里意思差了十万八千里同一个“空腹血糖”在体检报告和临床诊断语境下含义也不同。通用模型能识别字面但抓不住行业语境下的精确含义。第三是专业推理链条薄弱。通用模型擅长泛化回答但面对需要多步骤专业推理的问题比如“这段设备日志里哪些参数组合预示着轴承故障”它缺乏足够的领域知识做支撑容易一本正经地胡说八道。这三道坎本质上都是知识层面的缺失不是指令格式层面的问题。所以解决思路必须落在继续预训练上。1.2 继续预训练与微调的本质区别先说清楚两个概念。指令微调SFT用的是“问题-答案”对目标是让模型学会按人类期望的格式做回答而继续预训练用的是纯文本语料目标是让模型通过自监督学习把行业知识嵌入模型权重。打个比方。SFT像是给一个刚毕业的全科医生培训班教他怎么写病历、怎么按规范开检查单CPT则是把他扔进专科病房轮转几个月让他大量读真实的专科病例、手术记录、文献综述真正建立起专科知识体系。前者管“形式上会做”后者管“知识上懂行”。还有一个关键点SFT的数据量通常很小几千到几万条指令样本就够CPT的数据量动辄上亿甚至几十亿token。小样本微调能让模型改变行为风格但塞不下大规模知识。想让模型真正“懂行业”CPT是绕不开的一步。当然完整的企业落地链路通常还要在CPT之后再做SFT、再做RLHF等对齐步骤。但本文的焦点是CPT这是行业化改造里最容易被忽略、也最需要深度实操的一环。1.3 什么项目适合做继续预训练先算好投入产出不是所有行业大模型项目都适合走CPT先掂量一下自己的场景。适合做CPT的场景通常有这几个特征领域知识密集且专业术语多通用模型明显“答非所问”行业语料有规模且可获取至少能攒到几亿token高质量文本业务场景对专业能力要求高于对话流畅度比如辅助诊断、智能审阅、故障诊断等。不适合一上来就做CPT的包括这些情况你要解决的只是任务格式变化通过几千条SFT样本就能搞定行业数据有限凑不出足够大规模的高质量语料团队没有GPU训练环境只有推理部署的机器或者你的需求用检索增强生成RAG就能基本满足没必要花大成本训练。这里还想多说一句RAG和CPT不是互斥的很多企业最后是两者结合用。RAG负责把实时更新的知识检索出来CPT负责把稳定的行业知识固化在模型里。先想清楚你的知识是“需要频繁更新”还是“长期稳定”再决定投入方向。2. 动手前必看算力、数据、成本的三笔账2.1 算力需求怎么估算从7B到70B的卡时账很多团队开聊CPT时最没概念的就是算力7B模型训练10B token到底要几张卡训练多久这里给出一个可以直接套用的估算方法。先说基本原理。一次训练的前向和反向传播算力消耗大约是 6 × 模型参数量 × 处理的token数也就是常说的6ND公式。对于7B模型、10B token数据总计算量约为6 × 7e9 × 1e10 ≈ 4.2e20 FLOPs再算单卡算力。以A100-80G为例BF16峰值算力约312 TFLOPS但实际训练中利用率很难跑满通常按40%~60%估算。取50%也就是单卡有效算力约156 TFLOPS即1.56e14 FLOPs/秒。如果用8张A100总算力约为1.25e15 FLOPs/秒。那么理想训练时间为4.2e20 ÷ 1.25e15 ≈ 336,000秒 ≈ 93小时约等于4天。这还没算通信开销、日志打印、checkpoint保存、中间评估等额外时间实际落地按5到7天估计比较稳妥。如果把模型换成70B、数据量加到100B token那总计算量直接变成约4.2e22 FLOPs是前面例子的100倍。除非有几十张H100级别的卡否则普通团队很难在合理时间内完成这时就需要认真考虑做部分参数CPT或者用LoRA这类低成本方案。显存这块也要提前算。7B模型用BF16存储权重本体占约14GB训练过程中还有梯度、优化器状态Adam通常要维护fp32的权重副本、一阶动量、二阶动量如果不开任何优化这几项加起来显存压力非常大。所以CPT训练中几乎离不开DeepSpeed ZeRO、gradient checkpointing、混合精度这类手段。后面第三章会详细讲配置。2.2 行业数据从哪来、怎么洗数据质量决定上限做CPT最常被低估的一步是数据。很多团队一上来就急着开训练结果练出来的模型效果很差回头才发现数据里全是重复段落、格式噪声和低质量内容。说句实在话CPT项目的上限由数据质量决定训练只是在把这个上限兑现出来。行业数据的来源常见的有几类企业内部文档包括工单日志、操作规程、技术方案、验收报告、知识库文章公开但偏行业化的语料比如行业白皮书、专利文本、专业论文、标准规范以及结构化数据转文本比如把数据库里的一些关键字段、统计报表经过模板拼接转成自然语言文本。数据清洗环节核心步骤通常包括去掉HTML标签、Markdown标记、XML标签等格式噪声过滤过长或过短的文本比如少于200字的内容通常信息量低剔除乱码、纯英文混杂、表情符号过多的段落用MinHash等方法做大规模去重防止同一段文字反复出现删除涉及个人隐私和敏感信息的片段这块对医疗、金融领域尤其重要数据规模方面我的经验是如果做垂直领域CPT高质量行业语料至少要有几亿token低于这个量级效果很容易被淹没。同时建议在训练数据里混入20%~50%的通用语料这能显著缓解“灾难性遗忘”问题后面4.2节再说。2.3 训练周期和成本的现实预期关于时间分布请提前做好心理准备数据处理通常占整个项目60%以上的时间真正的训练反而只是其中一部分。一个10B token级别的中型CPT项目从数据收集、清洗、去重、格式转换到试跑、调参、正式训练、评估全流程一个两三人小团队干两周到一个月是很正常的事。成本也要算细一点。以云GPU按小时计费为例8卡A100一小时约几十到上百元级别跑5天大致是一笔可观支出但又不是遥不可及。如果预算有限还有更轻量打法不训练全量参数用LoRA或QLoRA在冻结原模型的情况下做“轻量级继续预训练”。这个方案在7B模型上效果也不错尤其适合先用小成本验证数据有效性再决定要不要上全量训练。3. 跑通第一次继续预训练完整实战流程3.1 环境与框架选型DeepSpeed还是Megatron-LM实操之前先选好工具链。目前主流方案基本可以分三条路线。第一条是Hugging Face Transformers配合Trainer或Accelerate再加DeepSpeed做显存优化。这条路最成熟生态最全文档和报错信息丰富适合绝大多数第一次做CPT的团队。第二条是Megatron-LM或Megatron-DeepSpeed适合70B以上模型的训练。它支持张量并行、序列并行、流水线并行能把超大模型切到多机多卡上但配置复杂度高很多普通项目没必要一上来就碰。第三条是轻量级路线用PEFT库做LoRA/QLoRA继续预训练。优点是单卡24G显存就能跑7B模型成本极低缺点是学到的知识深度不如全量训练适合验证阶段。我在实际项目里最常用的组合是Transformers DeepSpeed ZeRO-3 BF16混合精度。环境安装直接看Hugging Face和DeepSpeed官方文档大致如下pip install torch transformers datasets accelerate deepspeed建议提前确认一下GPU驱动和CUDA版本PyTorch版本尽量选最新稳定版否则容易遇到算子不兼容的问题。DeepSpeed一般不需要单独装transformers里能直接调用。3.2 数据处理完整流程从原始文本到jsonlCPT的数据格式不用搞复杂最通用的就是jsonl文件每行一个JSON对象核心字段是text。下面是一个典型的清洗脚本片段可以直接改改拿来用import json import re from pathlib import Path def clean_text(text: str) - str: # 去掉HTML标签 text re.sub(r[^], , text) # 压缩连续空行 text re.sub(r\n{3,}, \n\n, text) # 压缩行内多余空白 text re.sub(r[ \t]{2,}, , text) # 去掉首尾空白字符 text text.strip() return text def doc_to_jsonl(text: str): cleaned clean_text(text) if len(cleaned) 200: return None return {text: cleaned} with open(raw_docs.txt, r, encodingutf-8) as f: raw_docs f.read().split(\n---\n) with open(data/train.jsonl, w, encodingutf-8) as out: for doc in raw_docs: item doc_to_jsonl(doc) if item: out.write(json.dumps(item, ensure_asciiFalse) \n)清洗完之后建议先抽样检查20~30条人工过一眼确认没有明显的格式问题或乱码再往下走。还可以统计一下token数总量确保数据和算力估算对得上。这里有一个细节要特别提醒CPT训练时用的是“文本续写”这种自监督方式所以不需要构造问答对。有些人习惯性地按SFT的思维去整理数据结果花大量时间做标注这是完全没必要的。你只需要高质量、多样化的行业文本让模型自己学。3.3 核心参数配置学习率、batch size、warmup参数配置是CPT实操里最见功力的一块。先说总原则继续预训练的学习率要远低于从零预训练也低于常见的SFT微调否则模型容易在领域数据上过拟合同时遗忘通用能力。我常用的配置大概是这样参数推荐值说明learning_rate1e-5 ~ 5e-5稳定优先太小收敛慢lr_scheduler_typecosine训练后期平滑衰减warmup_ratio0.03先把学习率从0缓慢升上去per_device_train_batch_size2~8视显存而定gradient_accumulation_steps8~32用于凑够global batch sizebf16true有A100/H100就用BF16weight_decay0.01常用默认值max_seq_length2048或4096视数据长度和显存调整关于global batch size这个概念要理解透。真正影响模型收敛的是global batch size它等于单卡batch size × 梯度累积步数 × 卡数。比如单卡batch size为4梯度累积16步8卡训练那么global batch size就是4×16×8512。这个量级在CPT里是比较常规的。Trainer的配置代码大致这样from transformers import TrainingArguments training_args TrainingArguments( output_dir./cpt_output, per_device_train_batch_size4, gradient_accumulation_steps16, learning_rate2e-5, warmup_ratio0.03, lr_scheduler_typecosine, bf16True, logging_steps10, save_steps500, save_total_limit5, deepspeedds_config.json, )对应的ds_config.json里核心是ZeRO stage和可选的CPU offload我用得最多的配置是stage 3显存不够时再把优化器offload到CPU{ train_batch_size: auto, train_micro_batch_size_per_gpu: auto, gradient_accumulation_steps: auto, bf16: { enabled: true }, zero_optimization: { stage: 3, offload_optimizer: { device: cpu } } }启动命令也很简单deepspeed --num_gpus8 train_cpt.py \ --model_name_or_path your_base_model \ --dataset_path data/train.jsonl第一次跑的时候建议先用小规模数据比如1万条做一次“冒烟测试”确认流程能通、显存不爆、loss在下降再放开到全量数据。这一步能省下大量排障时间。4. 训练中的监控与调优别让loss骗了你4.1 训练曲线怎么解读正常长什么样开始训练之后最常看的指标就是loss。但新手很容易被loss曲线误导。一个正常的CPT训练loss应该是缓慢平滑下降的中间有微小波动整体趋势向下。如果训练数据量足够跑了几个小时后loss曲线基本能看出形状前期快速下降中后期越来越平缓。这里要注意CPT的loss绝对值一般比从零预训练要低得多因为模型已经在通用数据上预训练过起点就低。如果你看到loss突然跳高第一反应不要是调学习率而是先怀疑数据是不是某个batch里混入了一段异常文本是不是某个文档里面有半个乱码词表把loss高的step对应的数据样本拉出来看一下比瞎调参效率高得多。还有一个经验多打印日志没有坏处。logging_steps可以设小一点比如10步打一次否则等半天看不到日志你根本不知道训练是不是卡死了。4.2 防止灾难性遗忘一份数据里混着通用和行业灾难性遗忘是CPT最核心的技术挑战。只管训练行业数据、不管通用数据模型会在专业测试上变好但通用对话能力、常识推理能力明显下滑最终变成一个“偏科生”下游很多应用都没法用。我用下来最有效的手段是数据混合。具体操作是给训练语料做两类混合一类是通用领域数据比如高质量网页、书籍、百科类文本另一类是行业领域数据。比例方面我通常从8:2的行业与通用比例开始如果发现通用能力掉得厉害就往7:3甚至6:4调整。此外小学习率本身就是防遗忘的手段。学习率太高模型参数会被剧烈拉扯牢牢记住行业数据的同时忘掉原本的通用知识学习率控制在1e-5到3e-5这个区间参数更新温和既能吸收新知识又不容易破坏原有能力。还有个容易被忽略的点保存多个checkpoint训练中每隔一段就单独评测一次通用能力。不要等到训练结束才想起来评估。如果模型在某个checkpoint时通用能力明显下降就回溯到之前的checkpoint把训练粒度放细找到平衡点。4.3 评估怎么做领域指标与通用能力双端验证CPT的效果评估不能只看loss。loss下降只是说模型“拟合”了训练文本不代表它学会了知识。我习惯在训练期间和结束后做双端评估。领域端评估主要看三个维度领域困惑度Perplexity在没参与训练的领域文本验证集上算PPL这个指标直观反映模型对领域文本的拟合程度。PPL下降越多说明模型在领域数据上的把握越强。领域问答测试集人工准备100到200条行业问答让模型生成答案后人工打分这个最贴近真实业务效果。行业任务指标比如实体识别场景看F1分类场景看准确率按你的实际业务来定。通用端评估比较常用的是MMLU、C-Eval、GSM8K等公开benchmark。虽然这些测试集有些老旧但用来对比训练前后通用能力是否退化依然是性价比最高的方式。评估建议做成表格直接对比基础模型、CPT后模型、CPT后再加SFT模型的效果。这样团队内部讨论时一目了然也方便向业务方交代“训练到底改变了什么”。5. 真实踩坑记录常见问题与排查速查表5.1 显存溢出OOM从根因到解法CPT训练里最常见的报错就是CUDA out of memory。这个问题的根源在于显存里同时要放模型权重、梯度、优化器状态、中间激活和K/V cache任意一项爆了都会报OOM。我的排查顺序是这样先把per_device_train_batch_size降到1看还爆不爆不爆的话就是激活值太大打开gradient_checkpointing能省一大截显存。如果batch size为1还爆那就是模型本身优化器状态超了显存这时候要把DeepSpeed开到stage 3甚至配合CPU offload。一张快速排查表现象可能原因解决手段batch_size1仍OOM模型权重/优化器状态超显存DeepSpeed stage 3 offloadbatch_size调大后OOM中间激活占显存开gradient_checkpointing训练一段时间后突然OOM显存碎片或checkpoint保存时临时占用降低save_steps频率预留8-10GB多机多卡训练报OOM通信内存与训练内存叠加给NCCL预留内存减少batch size5.2 loss变成NaN或突然飙升loss变成NaN是最让人头皮发麻的问题之一。根因无非这几类学习率太高导致梯度爆炸混合精度下用FP16发生溢出用BF16能明显改善数据里有异常文本比如超大整数、异常字符、截断产生的半截UTF-8字符偶尔也可能是优化器参数设置不当。这里分享一个独家技巧当训练中途出现NaN时先别急着改学习率去翻一下最近一个step的样本数据。我遇到过两次NaN最后都是因为数据清洗时漏了一个包含极端长数字串的文档模型在生成概率上出现log域异常直接崩了。5.3 模型“变笨”了通用能力下降怎么办训练完拿回通用benchmark一看MMLU降了好几个点这种挫败感我太熟悉了。面对这种情况先别慌按这个顺序排查和调整。首先检查行业数据与通用数据的混合比例。如果行业数据超过80%甚至接近100%那通用能力大概率会下降。建议把混合比例调整到7:3或6:4后再试。其次检查学习率和训练步数。CPT阶段的学习率如果超过5e-5或者训练步数过多模型可能在新数据上“背诵”过头把原有知识覆盖掉了。减少训练步数、降低学习率通常能挽回不少。最后如果通用能力仍然掉得很厉害可以考虑在CPT之后再补一轮高质量通用指令数据上的SFT。很多团队最终验证下来“CPT 通用SFT”的组合能很好地兼顾领域能力和通用能力。5.4 GPU利用率上不去训练跑起来了但一看GPU利用率只有30%心里肯定着急。这种情况常见原因有三个数据加载太慢、CPU预处理成为瓶颈、多卡通信开销过大。数据加载问题通常用DataLoader的num_workers和prefetch_factor解决数据量大时把num_workers调到8以上。如果是超大语料集建议直接换流式数据读取方案比如WebDataset避免一次性把全部数据塞进内存。多卡通信问题则要看是不是频繁打印日志、频繁做all-reduce导致通信卡顿。把logging_steps适当加大减少梯度同步次数对吞吐量也有帮助。我自己调试时还会用nvidia-smi和deepspeed的日志实时看单卡算力利用情况配合top查看CPU内存是否爆满基本几分钟就能定位瓶颈。最后再分享一个个人习惯CPT项目开工前我一定先用小批量数据跑通全流程、确认loss在降、存一个baseline checkpoint再启动正式训练。很多团队第一次跑CPT踩坑其实都不是技术本身有多难而是流程不规范、数据不干净、参数凭感觉拍脑袋。多做一点前期验证能帮你省掉后面好几天的返工时间。
返回列表