ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

手撕大模型训练全流程:从预训练到端侧部署的完整实践指南

2026/8/18 20:39:32 拓冰建站 浏览量
手撕大模型训练全流程:从预训练到端侧部署的完整实践指南 1. 先搞清楚“手撕”到底要解决什么问题看到“手撕大模型训练全流程”这个标题很多人第一反应是觉得门槛高、流程长、无从下手。其实这个主题的核心价值在于它把一个看似庞大的工程拆解成几个可以分步验证的独立阶段并且最终指向一个非常具体的目标得到一个能在普通设备比如手机上运行的、经过完整流程调优的模型。这解决了几个实际问题理解全貌很多人只接触过SFT监督微调或量化对预训练、RLHF人类反馈强化学习等环节只有模糊概念不清楚它们如何串联。验证可行性从零开始训练一个“大模型”听起来不现实但通过合理的流程设计可以在小规模数据、小参数量的模型上完整跑通所有环节验证每个阶段的效果。打通部署链路训练出的模型最终要能落地。流程中包含了量化、蒸馏等压缩技术目标就是让模型从“只能跑在服务器上”变成“能在资源受限的端侧如手机运行”。所以这篇文章适合两类人一是想系统性理解大模型训练全流程的开发者或学习者二是确实有需求想得到一个轻量级、可定制、能部署的模型但被复杂流程劝退的实践者。最关键的能力不是让你从零训练出一个GPT-4而是让你掌握一套可复现、可调试、分阶段验证的方法论。即使你只有单张消费级显卡甚至只有CPU也能通过调整模型规模和数据量走完整个流程获得直观感受。2. 环境与资源准备别在第一步卡住在开始“手撕”之前最实际的一步是准备好战场。这里的环境准备不是简单列个软件清单而是要明确每个阶段对资源的需求不同我们可以动态调整。2.1 硬件与基础软件环境核心硬件考量GPU最重要这是训练速度的瓶颈。对于全流程学习拥有一张显存8GB以上的NVIDIA显卡是比较理想的起点如RTX 3070/4060Ti等。这允许你运行参数量在1B10亿左右的模型进行全流程实验。如果只有4-6GB显存可以考虑更小的模型如300M参数或在预训练等显存消耗大的阶段使用更小的批次大小batch size。CPU与内存数据预处理、日志记录、部分评估任务会用到CPU。建议16GB以上内存。RLHF阶段需要同时加载多个模型训练中的模型、参考模型、奖励模型对内存和显存都有额外压力。磁盘准备至少100GB的可用空间。这用于存放原始语料、预处理后的数据、多个阶段的模型检查点、日志等。建议使用SSD以加速数据读取。基础软件栈操作系统LinuxUbuntu 20.04/22.04是首选对深度学习框架支持最完善。WindowsWSL2或macOSM系列芯片也可行但可能在某些环节遇到依赖问题。Python版本3.8-3.10。建议使用conda或venv创建独立的虚拟环境。深度学习框架PyTorch是当前大模型训练生态的事实标准。根据你的CUDA版本通过nvidia-smi查看安装对应版本的PyTorch。关键Python库transformers(Hugging Face)模型加载、训练、评估的核心。datasets(Hugging Face)数据集处理。accelerate简化分布式训练。peft参数高效微调如LoRA在资源有限时极其有用。bitsandbytes8-bit/4-bit量化用于降低推理和训练显存。trlTransformer Reinforcement Learning实现RLHF的核心库。wandb或tensorboard实验跟踪与可视化。安装命令示例以CUDA 11.8为例conda create -n llm-train python3.10 conda activate llm-train pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 pip install transformers datasets accelerate peft bitsandbytes trl wandb2.2 模型与数据准备这是最容易出问题的地方不要一上来就追求大数据、大模型。模型选择起点不要从随机初始化开始预训练那需要海量数据和算力。我们的起点是一个已有的、结构简单的中小型预训练模型。例如GPT-2 Small(124M参数)经典文档丰富非常适合学习全流程。BERT-base(110M参数)理解Encoder结构训练的好选择。Bloom-560m最近的开源模型性能不错。Qwen1.5-0.5B或Phi-2(2.7B)参数稍大能力更强但对显存要求也更高。从Hugging Face Hub下载这些模型非常简单from transformers import AutoModelForCausalLM; model AutoModelForCausalLM.from_pretrained(gpt2)。数据准备每个阶段需要不同的数据预训练/继续预训练需要大规模、高质量的通用文本语料。对于学习目的可以使用wikitext-2、openwebtext的子集或者自己准备一小部分纯文本文件如小说、技术文档。关键是文本要干净。SFT监督微调需要高质量的指令-回答对。可以使用开源数据集如Alpaca格式的数据、ShareGPT的清洗版本或者自己构造。格式通常是{instruction: ..., input: ..., output: ...}。RLHF需要偏好对比数据。格式如{prompt: ..., chosen: ..., rejected: ...}。可以基于SFT数据用一个大模型如GPT-4生成多个回答然后人工或基于规则进行排序来模拟。一个关键建议先只用极少量数据如几百条跑通每个阶段的训练循环确保代码、流程、保存加载没问题再逐步增加数据量。3. 分阶段实战预训练 - SFT - RLHF - 压缩现在进入核心实操。我们把流程分解为四个主要阶段每个阶段都关注输入、输出、核心代码和效果验证。3.1 阶段一预训练Pre-training—— 让模型“学会说话”目标让模型掌握语言的统计规律和世界知识。在我们的“手撕”场景中更可能是继续预训练Continue Pre-training即在已有模型基础上用新的领域数据让它学习新知识或风格。核心操作数据预处理将文本语料分词tokenize并处理成模型需要的输入格式如input_ids,attention_mask。对于因果语言模型如GPT任务是预测下一个token。配置训练参数from transformers import Trainer, TrainingArguments training_args TrainingArguments( output_dir./pt-checkpoints, # 检查点保存路径 overwrite_output_dirTrue, num_train_epochs1, # 对于继续预训练1-3个epoch通常足够 per_device_train_batch_size4, # 根据显存调整能设多大设多大 save_steps500, # 每500步保存一次 save_total_limit2, # 只保留最新的2个检查点 logging_steps100, learning_rate5e-5, # 继续预训练的学习率通常较小 fp16True, # 开启混合精度训练节省显存并加速 gradient_accumulation_steps4, # 模拟更大的batch size )开始训练from transformers import DataCollatorForLanguageModeling from datasets import load_dataset # 加载数据集和模型 dataset load_dataset(text, data_files{train: ./my_corpus.txt}) tokenizer AutoTokenizer.from_pretrained(gpt2) model AutoModelForCausalLM.from_pretrained(gpt2) # 对数据集进行分词 def tokenize_function(examples): return tokenizer(examples[text], truncationTrue, max_length512) tokenized_datasets dataset.map(tokenize_function, batchedTrue, remove_columns[text]) # 使用语言模型数据收集器MLM任务用DataCollatorForLanguageModeling这里是因果LM data_collator DataCollatorForLanguageModeling(tokenizertokenizer, mlmFalse) trainer Trainer( modelmodel, argstraining_args, train_datasettokenized_datasets[train], data_collatordata_collator, ) trainer.train()验证效果训练损失Loss观察损失是否平稳下降。这是最直接的指标。生成文本训练几个epoch后用model.generate()让模型续写一段话看语言是否通顺、是否符合领域特点。困惑度Perplexity, PPL在预留的验证集上计算PPL数值越低说明模型对文本的建模能力越强。避坑点显存溢出OOM首先降低per_device_train_batch_size其次开启gradient_checkpointing检查点技术再次尝试使用peft库的LoRA进行微调而非全参数训练。学习率继续预训练的学习率5e-5到1e-4通常比从头训练1e-3左右小得多。3.2 阶段二SFT监督微调—— 让模型“听从指令”目标将预训练好的通用语言模型微调成能理解并遵循人类指令的对话或任务模型。核心操作数据格式转换将指令数据转换成对话格式。例如将(instruction, input, output)转换成|system|You are a helpful assistant./s |user|{instruction}\n{input}/s |assistant|{output}/s具体模板需与模型的训练格式对齐。很多模型如Qwen, Llama有固定的对话模板。使用Seq2Seq训练SFT通常被建模为序列到序列的生成任务。配置训练参数与预训练类似但learning_rate可以稍大如1e-4到2e-5num_train_epochs通常为3-10取决于数据量。training_args TrainingArguments( output_dir./sft-checkpoints, per_device_train_batch_size4, gradient_accumulation_steps4, num_train_epochs3, learning_rate2e-5, logging_steps10, save_steps500, fp16True, remove_unused_columnsFalse, # 重要SFT数据通常有多列需要保留 )使用SFTTrainer推荐trl库提供了专为SFT优化的Trainer能更好地处理对话格式和打包序列。from trl import SFTTrainer from datasets import load_dataset dataset load_dataset(json, data_files./sft_data.jsonl) trainer SFTTrainer( modelmodel, argstraining_args, train_datasetdataset[train], dataset_text_fieldtext, # 你的数据中拼接好格式的文本字段名 max_seq_length1024, # 模型支持的最大长度 tokenizertokenizer, ) trainer.train()验证效果指令遵循给出训练数据中未见过的指令看模型是否能生成相关、合理的回答。格式一致性检查模型输出是否严格遵守了设定的对话格式如正确使用/s作为结束符。评估指标可以使用BLEU、ROUGE等自动指标但更重要的是人工评估回答的有用性、相关性和无害性。避坑点过拟合如果数据量少模型可能会机械记忆训练样本。通过早停early stopping、增加数据多样性、使用LoRA等PEFT方法缓解。格式混乱确保数据预处理时对话模板拼接正确否则模型学不会正确的对话轮次。3.3 阶段三RLHF人类反馈强化学习—— 让模型“输出更优”目标让模型的输出不仅正确而且更符合人类的偏好如更有帮助、更详细、更安全。核心流程简化版使用DPO RLHF传统流程训练奖励模型再用PPO优化非常复杂且不稳定。近年来DPO直接偏好优化方法变得流行它更简单、稳定适合我们“手撕”的场景。DPO的核心思想是直接利用偏好数据(prompt, chosen_response, rejected_response)来微调模型绕过训练独立奖励模型的步骤。核心操作准备偏好数据格式为{prompt: ..., chosen: ..., rejected: ...}。加载SFT模型DPO需要一个已经过SFT的模型作为起点。使用DPOTrainerfrom trl import DPOTrainer, DPOConfig from datasets import load_dataset # 加载模型和参考模型通常是同一个SFT模型的拷贝 model AutoModelForCausalLM.from_pretrained(./sft-final-checkpoint) model_ref AutoModelForCausalLM.from_pretrained(./sft-final-checkpoint) tokenizer AutoTokenizer.from_pretrained(./sft-final-checkpoint) # 加载偏好数据集 dataset load_dataset(json, data_files./preference_data.jsonl) # 配置DPO参数 dpo_args DPOConfig( output_dir./dpo-checkpoints, per_device_train_batch_size2, # DPO通常batch size更小 learning_rate1e-6, # DPO学习率非常小 num_train_epochs1, # 通常1-2个epoch就够了 beta0.1, # DPO温度参数控制偏离参考模型的程度 fp16True, ) trainer DPOTrainer( modelmodel, ref_modelmodel_ref, argsdpo_args, train_datasetdataset[train], tokenizertokenizer, ) trainer.train()验证效果偏好胜率在测试集上比较DPO微调后的模型和SFT模型看前者生成“chosen”类回答的胜率是否提高。人工对比评估准备一批新的prompt让SFT模型和DPO模型同时生成回答找人判断哪个回答更好。DPO模型应该在有用性、无害性、细节丰富度上表现更优。避坑点灾难性遗忘RLHF/DPO可能会损害模型原有的知识或能力。可以通过在损失函数中加入对原始SFT模型的KL散度惩罚来缓解DPOTrainer已内置。数据质量偏好数据的质量至关重要。“chosen”和“rejected”的差距必须清晰。噪声大的数据会导致优化方向错误。3.4 阶段四量化与蒸馏 —— 让模型“瘦身”并“提速”目标将训练好的大模型压缩降低其计算和存储开销使其能够部署在手机等资源受限的设备上。方案一量化Quantization量化将模型参数从高精度如FP32转换为低精度如INT8, INT4大幅减少模型体积和推理时的内存占用。动态量化推理时动态转换易于使用。静态量化需要校准数据精度更高。GPTQ/AWQ更高级的仅权重量化方法精度损失小。使用bitsandbytes进行加载时量化最简便from transformers import AutoModelForCausalLM, BitsAndBytesConfig import torch bnb_config BitsAndBytesConfig( load_in_4bitTrue, # 加载为4-bit bnb_4bit_quant_typenf4, # 使用NF4量化类型 bnb_4bit_compute_dtypetorch.float16, # 计算时使用fp16 ) model AutoModelForCausalLM.from_pretrained( ./dpo-final-checkpoint, quantization_configbnb_config, device_mapauto, # 自动将模型层分配到可用的GPU/CPU上 )量化后的模型可以直接用于推理速度更快显存占用极低。但量化模型通常无法直接继续训练如需训练需使用QLoRA等技术。方案二知识蒸馏Knowledge Distillation用一个庞大的“教师模型”来指导一个小的“学生模型”学习让学生模型模仿教师模型的输出或中间层特征。简易蒸馏流程准备数据收集一批输入文本。生成软标签用我们训练好的大模型教师模型对这些输入生成输出并获取输出概率分布软标签这比硬标签包含更多信息。训练学生模型用一个更小架构的模型学生模型在原始数据硬标签和教师模型的软标签上联合训练。损失函数通常结合学生预测与真实标签的交叉熵以及学生预测与教师软标签的KL散度。验证压缩效果模型大小对比量化/蒸馏前后模型的.bin文件大小。推理速度在相同硬件上测量生成固定长度文本的平均时间。资源占用使用nvidia-smi或内存监控工具观察推理时的GPU显存或CPU内存占用。性能保留在保留的测试集上评估压缩后模型的准确率、BLEU分数或进行人工评估看性能下降是否在可接受范围内例如精度损失5%。避坑点量化精度损失尝试不同的量化配置如load_in_8bitvsload_in_4bitnf4vsfp4找到精度和速度的平衡点。蒸馏难度如果教师模型和学生模型能力差距过大蒸馏可能失败。需要仔细设计学生模型架构和蒸馏损失函数。4. 端侧部署验证在手机上跑起来流程的终点是验证模型能否在目标环境手机中运行。这里以Android平台为例介绍核心思路。核心工具链模型转换将PyTorch模型转换为手机端推理框架支持的格式。ONNX通用的模型交换格式。可以使用torch.onnx.export将模型导出为ONNX。TFLiteTensorFlow Lite格式在移动端支持良好。可以通过ONNX - TFLite的转换工具实现。移动端推理引擎PyTorch Mobile直接支持TorchScript格式。TFLiteGoogle主推生态完善。MNN/NCNN优秀的国产移动端推理框架。前端封装使用AndroidJava/Kotlin或iOSSwift编写简单的App调用推理引擎完成输入文本的预处理、模型推理和输出文本的后处理。简化验证步骤导出模型将最终量化后的模型导出为ONNX格式。import torch model.eval() # 设置为评估模式 dummy_input torch.randint(0, tokenizer.vocab_size, (1, 10)).to(device) # 示例输入 torch.onnx.export(model, dummy_input, llm.onnx, opset_version14)简化模型对于移动端可能需要进一步优化如算子融合、常量折叠。可以使用onnxruntime的工具或onnx-simplifier。选择推理引擎集成例如在Android Studio项目中添加TFLite依赖将转换好的.tflite模型放入assets文件夹。编写核心推理代码Android示例片段// 加载模型和分词器需要事先将分词器词表等资源打包进App Interpreter tflite new Interpreter(loadModelFile(context)); // 将输入文本预处理成token ids int[][] inputIds preprocess(text); // 运行推理 tflite.run(inputIds, outputBuffer); // 将输出的token ids 后处理成文本 String result postprocess(outputBuffer);性能测试在真机上测试不同长度输入的推理延迟和内存占用。避坑点模型过大即使量化后模型也可能超过几十MB甚至上百MB。需要考虑App包体积或采用模型动态下载的方案。算子不支持某些Transformer层的特殊算子可能不被移动端推理引擎完全支持。导出ONNX时需注意算子版本或寻找替代实现。预处理/后处理分词逻辑Tokenizer需要在移动端复现这部分代码也需要移植和优化。5. 全流程串联与迭代心得走完一遍全流程后你会对每个环节的输入输出、资源消耗、调参敏感度有直观认识。这里分享几个串联时的关键心得1. 检查点管理是生命线每个阶段都会产生多个模型检查点。必须建立清晰的目录结构并记录每个检查点对应的训练配置、数据版本和评估结果。例如project/ ├── data/ ├── checkpoints/ │ ├── pt-epoch1/ │ ├── sft-lora/ │ ├── dpo-beta0.1/ │ └── quantized-int4/ ├── scripts/ # 训练和评估脚本 └── logs/ # 训练日志和评估结果使用wandb或tensorboard可以很好地跟踪实验。2. 评估必须贯穿始终不要等到最后才评估模型。每个阶段结束后都应进行快速评估预训练后计算验证集困惑度做文本生成。SFT后用一组标准指令测试指令遵循能力。RLHF/DPO后进行A/B测试对比SFT模型。量化/蒸馏后在测试集上评估性能保留率并测试端侧推理速度。 建立自动化的评估脚本让评估结果可量化、可比较。3. 资源有限下的优先级如果算力紧张可以这样取舍优先保证SFTSFT对模型能力提升最直接。可以使用LoRA等PEFT方法在单卡上微调7B/13B的模型。RLHF可以简化使用DPO替代传统的PPO流程稳定且省资源。甚至可以手动构造高质量偏好数据量少但质精。量化是部署的必选项bitsandbytes的加载时量化几乎无成本应优先使用。蒸馏则需要额外的训练成本可作为进阶优化。4. 迭代循环小步快跑不要试图用全部数据一次性跑完所有流程。建议采用迭代式开发用1%的数据跑通从数据预处理到SFT的完整pipeline。评估效果修复流程中的bug如数据格式错误、训练崩溃。增加数据量到10%加入DPO训练并尝试量化。评估各阶段输出调整超参数学习率、batch size等。用全量数据或你能承受的数据量进行最终训练。这个过程的核心不是产出世界顶尖的模型而是构建一个可理解、可控制、可迭代的模型开发工作流。当你掌握了这个工作流面对新的模型架构、新的训练技巧时你就能快速地将它们集成进来并进行有效的验证。这才是“手撕”全流程带来的真正价值。