欢迎光临
我们一直在努力

笔记二十:SFT(监督微调)实战完全指南

SFT(监督微调)实战完全指南:序列打包、掩码策略与灾难性遗忘诊断

本文是对《The Hitchhiker’s Guide to Agentic AI》第10章\”SFT Best Practices and Techniques\”的系统梳理。SFT是RLHF(基于人类反馈的强化学习)流程的基础——SFT模型的质量,直接决定了RL能达到的天花板高度【10†L1-L2】。


第1章:序列打包(Sequence Packing)——让GPU不再\”空转\”

1.1 问题:为什么训练时GPU算力被浪费了?

在深度学习训练中,一个批次(Batch)里的多个样本需要\”整齐划一\”地送入GPU计算。但问题是,不同样本的长度往往差异很大——有的问题只有十几个字(比如\”1+1等于几?\”),有的却是一篇长文档(比如几千字的论文摘要)。

为了能让它们在一个批次里一起运算,标准做法是:把所有序列都填充(Padding)到和本批次最长样本一样的长度。短样本后面会被塞满无意义的\”填充符\”(通常是0)。

大白话理解:就像学校排队做操,高个子站后面,矮个子站前面,但为了队伍整齐,矮个子脚下要垫砖头,直到和高个子一样高。这些\”砖头\”就是填充符,它们不包含任何有用信息,但GPU却要花同样的算力去处理它们。

浪费有多严重? 对于长短差异大的数据集(比如混合了短指令和长文档),这种填充方式会浪费**50%~80%**的算力【10†L6】。GPU花了大量时间在空数据上\”干算\”,非常不划算。

1.2 解决方案:序列打包——把短样本\”拼车\”在一起

序列打包的思路很简单:把多个短样本首尾相连,拼成一个接近长度上限的长序列,就像拼车一样把空座位填满。

具体操作分三步【10†L8-L10】:

  • 排序(可选):先把样本按长度排个序,方便后续\”装箱\”,提高打包效率。
  • 贪婪装箱:把短样本一个个塞进长度为 max_seq_length 的\”箱子\”里,直到塞不下为止。
  • 注意力掩码:这是最关键的一步——用块对角注意力掩码,确保不同样本之间的token不会互相\”偷看\”,防止前后文混淆【10†L10】。
  • 大白话理解:就像几拨陌生人拼一辆大巴车,虽然坐在同一辆车里,但每拨人只跟自己人聊天,不去听别人的对话。注意力掩码就是给每个样本画了个\”隔间\”。

    1.3 效果有多好?(数据说话)

    指标
    传统填充法
    序列打包法
    GPU利用率 20%~50% 85%~95%【10†L13】
    训练速度 基准 快2~4倍(长短差异大的数据集)【10†L14】
    显存占用 基准 差不多(因为每批处理的总有效token数没变)【10†L15】

    需要注意的坑:必须小心处理注意力掩码,否则不同样本的内容会互相污染,模型就学乱了【10†L16】。

    1.4 代码怎么实现?(TRL库一行搞定)

    在Hugging Face的TRL库中,只需要在 SFTConfig 里设置 packing=True,一切自动完成【10†L20-L31】:

    from trl import SFTConfig, SFTTrainer

    config = SFTConfig(
    max_seq_length=4096, # 最大序列长度
    packing=True, # 开启序列打包 —— 就这一行!
    output_dir=\”sft_model\”,
    per_device_train_batch_size=4,
    gradient_accumulation_steps=4,
    learning_rate=2e-5,
    num_train_epochs=3,
    )

    trainer = SFTTrainer(
    model=model,
    args=config,
    train_dataset=dataset,
    )
    trainer.train()

    小结:序列打包解决的是\”算力效率\”问题——让GPU每一分算力都用在刀刃上,而不是浪费在填充符上。


    第2章:聊天模板(Chat Templates)——让模型分清谁在说话

    2.1 问题:为什么聊天模板如此重要?

    底层模型(基座模型)只认识\”裸文本\”——就是一串连续的字符,分不清哪句是系统指令、哪句是用户提问、哪句是助手回答【10†L35】。

    但指令微调的目标是让模型学会\”遵循指令\”,所以必须在训练数据中明确标记出角色边界:System(系统)、User(用户)、Assistant(助手)【10†L36】。

    后果有多严重? 如果推理时用了错误的模板,或者干脆不用模板,模型会把\”用户问题\”和\”系统提示\”混为一谈,导致回答风格崩坏、指令遵循能力大幅下降【10†L37】。

    大白话理解:就像看一部剧本,如果不标明哪句话是哪个角色说的,演员就不知道谁该念哪句词。聊天模板就是剧本里的\”角色标签\”。

    2.2 主流模板长什么样?(以ChatML为例)

    ChatML是目前最通用的聊天模板格式之一【10†L40-L47】:

    <|im_start|>system
    你是一个乐于助人的助手。
    <|im_end|>
    <|im_start|>user
    什么是机器学习?
    <|im_end|>
    <|im_start|>assistant
    机器学习是人工智能的一个分支…
    <|im_end|>

    格式解析:

    • <|im_start|> 表示某个角色的开始
    • <|im_end|> 表示该角色的结束
    • 按\”系统→用户→助手\”的顺序拼接

    注意:不同模型家族的模板格式不同(比如Llama用的是 <|start_header_id|>,Qwen用的是 <|im_start|>),必须使用和模型配套的模板。

    2.3 在TRL中如何正确应用?

    标准做法是三步走【10†L52-L70】:

    第一步:加载分词器(Tokenizer)

    from transformers import AutoTokenizer
    from trl import SFTConfig, SFTTrainer

    tokenizer = AutoTokenizer.from_pretrained(\”meta-llama/Llama-3.1-8B-Instruct\”)

    为什么是分词器来管模板?因为模板的\”语法规则\”是写在分词器的配置文件(tokenizer_config.json)里的,不同模型的模板不一样。

    第二步:定义格式化函数

    def formatting_func(example):
    \”\”\”把数据集中的一条数据转换成符合模板规范的格式\”\”\”
    messages = [
    {

    \”role\”: \”system\”, \”content\”: \”你是一个乐于助人的助手。\”},
    {

    \”role\”: \”user\”, \”content\”: example[\”instruction\”]}, # 数据集里的\”问题\”字段
    {

    \”role\”: \”assistant\”, \”content\”: example[\”response\”]}, # 数据集里的\”回答\”字段
    ]

    return tokenizer.apply_chat_template(
    messages,
    tokenize=False, # 先不转成token ID,先返回字符串
    add_generation_prompt=False, # 训练时设为False
    )

    第三步:传给Trainer

    config = SFTConfig(
    max_seq_length=2048,
    output_dir=\”sft_model\”,
    )

    trainer = SFTTrainer(
    model=model,
    tokenizer=tokenizer,
    args=config,
    train_dataset=dataset,
    formatting_func=formatting_func, # 传入格式化函数
    )
    trainer.train()

    2.4 一个重要细节:add_generation_prompt 参数

    • 训练时(add_generation_prompt=False):把完整的\”用户问+助手答\”作为训练目标,让模型学习完整的对话模式。
    • 推理时(add_generation_prompt=True):只给\”用户问\”,在助手回答位置留空,让模型自动生成后续内容【10†L67】。

    大白话理解:训练时相当于给模型看\”例题+答案\”,推理时只给\”例题\”,让模型自己写答案。

    小结:聊天模板解决的是\”角色格式\”问题——告诉模型谁是系统、谁是用户、谁是助手,让模型知道该以什么身份、什么风格来回答。


    第3章:仅补全掩码(Completion-Only Masking)——让模型只学\”怎么回答\”,不学\”怎么提问\”

    3.1 问题:为什么不能让模型学习全部文本?

    在指令微调的数据里,一条样本通常包含两个部分:

    • 用户问(Prompt):比如\”请解释一下什么是梯度下降?\”
    • 助手答(Response):比如\”梯度下降是一种优化算法…\”

    如果我们让模型对整段文本(包括用户问题)都计算损失(Loss),会出现两个严重问题:

    问题一:浪费算力 模型会浪费宝贵的梯度信号去学习\”如何预测用户的下一个问题\”,但这根本不是我们想要的生成能力。我们想要的是模型根据问题生成答案,而不是续写问题本身。

    问题二:死记硬背(灾难性\”作弊\”) 模型可能会\”记住\”特定问题的措辞,导致它过于依赖训练数据中的固定表述。测试时换一种问法,模型就懵了。更糟的是,模型学会了\”背答案\”而不是\”理解问题-生成答案\”的映射关系。

    大白话理解:就像老师给学生一张卷子,上面既有题目又有答案。如果学生把题目和答案一起背下来,考试时题目稍微变一下他就不会了。我们应该只让学生学\”看到题目→想到答案\”这个推理过程,而不是把题目也背下来。

    3.2 解决方案:仅补全掩码——给\”问题\”部分打个×,不算分

    仅补全掩码的核心思想极其简单粗暴:在计算损失时,把所有非助手角色的token(即系统提示和用户问题)的Loss权重设为0;只保留助手回答部分的Loss来计算梯度并更新模型。

    大白话理解:就像老师在批改作业时,用红笔把\”题目部分\”全部划掉,只批改\”答案部分\”。模型只会因为\”答案写得好不好\”而受到奖励或惩罚,不会因为\”题目背得熟不熟\”而得到任何信号。

    3.3 在TRL中如何实现?

    在TRL中,通过 DataCollatorForCompletionOnlyLM 这个数据整理器来实现:

    from trl import SFTConfig, SFTTrainer, DataCollatorForCompletionOnlyLM
    from transformers import AutoTokenizer

    tokenizer = AutoTokenizer.from_pretrained(\”meta-llama/Llama-3.1-8B-Instruct\”)

    # 关键:告诉数据整理器\”回答从哪里开始\”
    # 对于Llama-3,助手回答的起始标记是 <|start_header_id|>assistant<|end_header_id|>
    response_template = \”<|start_header_id|>assistant<|end_header_id|>\”

    # 创建数据整理器
    collator = DataCollatorForCompletionOnlyLM(
    response_template=response_template,
    tokenizer=tokenizer,
    )

    config = SFTConfig(
    max_seq_length=2048,
    output_dir=\”sft_model\”,
    )

    trainer = SFTTrainer(
    model=model,
    tokenizer=tokenizer,
    args=config,
    train_dataset=dataset,
    data_collator=collator, # 传入带掩码功能的数据整理器
    formatting_func=formatting_func, # 前面我们定义的聊天模板函数
    )
    trainer.train()

    代码背后的逻辑:

  • 数据整理器拿到一个批次的数据(已经是按聊天模板格式化好的字符串)。
  • 它把这些字符串转成token ID序列。
  • 它在token序列里搜索 response_template 对应的那串token ID。
  • 找到之后,把这个位置之前的所有token(系统提示+用户问题)都标记为\”不计算损失\”。
  • 这个位置及之后的token(助手回答)正常计算损失。
  • 3.4 三大致命陷阱(实操必看!)

    文末特别列出了三个极易踩坑的地方,每一个都能让你的训练白费:

    陷阱一:模板必须和分词器完全匹配(精确匹配)

    response_template 这个字符串在传入后,会被分词器转换成token ID序列。如果你的模板字符串里多了一个空格、少了一个换行符、或者标点符号全角半角不对,那么分词后得到的token ID序列就和实际数据里的序列对不上。

    赞(0)
    未经允许不得转载:171主机测评 » 笔记二十:SFT(监督微调)实战完全指南
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址