大模型SFT训练流程代码逐行解析
大模型 SFT(监督式微调,Supervised Fine-Tuning)是指:先拿一个已经在海量无标签文本上完成自监督预训练(Pretraining)的语言模型(例如 GPT、LLaMA、ChatGLM 等),然后用人工标注的「输入→期望输出」(prompt→response)对,在有监督的条件下进一步训练模型,使其在特定任务或领域上产出更准确、更符合预期的回答。之所以称作「大模型 SFT」,通常是指模型规模(参数量)非常大,例如几亿、几十亿乃至上千亿参数的 Transformer 结构。与之对应的,SFT 的目标就是在保持原始预训练能力的前提下,让模型更擅长特定场景的「生成式任务」或「问答任务」。SFT 与预训练(PT)、RLHF 等的区别预训练(PT):使用无监督的语言建模目标(如最大化下一个 token 的概率、对比学习等),在海量通用文本语料上训练,让模型学到通用语言知识。SFT:在已经预训练好的大模型基础上,用「人类人工撰写的示例对」做有监督学习(常见格式是 |prompt| 用户输入 |response| 期望回答),只更新模型参数中的一部分或全部参数,使得模型“学会”按照示例里期望的风格、格式或知识领域来生成文本。RLHF(Reinforcement Learning from Human Feedback):在完成 SFT 之后,再用人类打分模型输出的好坏,构造奖励函数;通过策略梯度(PPO 等)进一步调整生成策略,最终得到对齐更好、更安全的输出。为什么要做大模型 SFT?预训练阶段的目标是「通用性」和「覆盖各种写作风格」,但在某些垂直场景(如客服对话、医学问诊、编程辅助)需要更精准、更符合规范的答案。SFT 可以让模型迅速掌握某一领域的风格、格式和专有知识,减少对错误、不相关内容的生成。相较于从头训练大模型,先预训练再 SFT 大幅度节省标注成本和算力开销,同时保留原有的大规模语言理解能力。0 训练流程简介参数与日志处理:开始程序后,首先解析命令行参数,包含模型、脚本和训练相关的参数。接着设置日志输出规则,确定日志的格式、级别和输出位置。设备与模型准备:设置随机种子以保证结果的可重复性。判断是否有可用的 GPU,若有则使用cuda,否则使用cpu。初始化分词器,加载模型配置并设置不使用kv cache。加载模型并将其移动到指定的设备上。计算并打印模型的总参数数量和可训练参数数量。数据加载与训练:加载 SFT 数据集。初始化Trainer,用于管理训练过程。根据是否需要恢复训练的设置,决定是从checkpoint恢复训练还是直接开始训练。模型保存:训练完成后,创建保存模型的目录。保存分词器和模型的权重及配置。结束:整个流程结束。各部分具体介绍如下1.主函数defmain():# 1.解析参数# 创建解析器,自动从命令行读取并填充三个dataclass类型的参数结构parser=HfArgumentParser((ModelArguments,ScriptArguments,TrainingArguments))# 解析命令行参数,将结果分别赋值给 model_args、script_args、training_argsmodel_args,script_args,training_args=parser.parse_args_into_dataclasses()# 2.设置日志输出规则# logger formatlogging.basicConfig(format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",datefmt="%m/%d/%Y %H:%M:%S",# 日志格式:时间戳、日志级别、模块名、消息内容level=logging.WARN,# 如果不是主进程,则使用 WARNING 及以上级别;否则也用 WARNINGhandlers=[logging.StreamHandler(sys.stdout)],)# 将日志输出到控制台(标准输出)iftraining_args.should_log:# 如果满足 should_log 条件,将 Transformers 库的日志级别设为 INFO,方便输出更多细节transformers.utils.logging.set_verbosity_info()# 根据当前进程是主进程还是工作进程,自动获取合适的日志级别(如 INFO 或 WARNING)log_level=training_args.get_process_log_level()# 设置全局 logger 的日志级别logger.setLevel(log_level)# 将 datasets 库的日志级别与主程序统一datasets.utils.logging.set_verbosity(log_level)# 设置 Transformers 库的日志等级(模型加载、Trainer 进度等信息)transformers.utils.logging.set_verbosity(log_level)# 启用默认的日志处理器,保证日志按照之前定义的格式输出transformers.utils.logging.enable_default_handler()# 使用显式的日志格式(即上面 logging.basicConfig 定义的格式)transformers.utils.logging.enable_explicit_format()# 在每个进程启动后,打印一条小结信息,包含进程 rank、设备信息、GPU 数量、是否分布式训练、是否使用 16 位精度训练logger.warning(f"Process rank:{training_args.local_rank}, device:{training_args.device}, n_gpu:{training_args.n_gpu}"+f"distributed training:{bool(training_args.local_rank!=-1)}, 16-bits training:{training_args.fp16}")# 3.设置随机种子set_seed(training_args.seed)# 4.判断是否有可用的 GPU,有则使用 "cuda",否则使用 "cpu"device="cuda"iftorch.cuda.is_available()else"cpu"# 5.初始化分词器tokenizer=transformers.AutoTokenizer.from_pretrained(script_args.base_model_path,# 基础模型所在路径或名称use_fast=False,# 不使用 fast tokenizer(避免某些自定义模型不兼容)trust_remote_code=True,# 允许加载模型作者上传的自定义代码model_max_length=model_args.max_position_embeddings# 支持的最大序列长度)# 6.加载模型配置文件(config.json),用于创建模型时指定参数config=transformers.AutoConfig.from_pretrained(script_args.base_model_path,# 与分词器相同的模型路径trust_remote_code=True# 允许加载自定义代码中的配置)config.use_cache=False# 不使用kv cache,即每步生成都重新计算所有 attention;更慢,但对某些训练/微调场景有用。推理时建议打开,训练时关闭# 7.加载模型model=transformers.AutoModelForCausalLM.from_pretrained(script_args.base_model_path,# 基础模型路径config=config,# 使用上面加载的配置trust_remote_code=True# 允许加载自定义代码实现的模型结构)# 将模型移动到指定设备(GPU 或 CPU)model.to(device)# 打印模型参数细节:总参数量和可训练参数量total_params=sum(p.numel()forpinmodel.parameters())# 计算模型所有参数的元素数量trainable_params=sum(p.numel()forpinmodel.parameters()ifp.requires_grad)#可训练参数量,即梯度参数 (requires_grad=True)logger.info(f"总参数:{total_params},{total_params/2**20:.2f}M params")# 参数量以百万M表示,二进制百万是2^20logger.info(f"可训练参数:{trainable_params}")# 输出可训练参数总数############### 8.加载SFT (Supervised Fine-Tuning) 数据集sft_dataset=SFTDataset(script_args.dataset_dir_or_p