news 2026/8/9 11:18:09

【Bug已解决】Understanding loss in Training LLM 解决方案

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
【Bug已解决】Understanding loss in Training LLM 解决方案

【Bug已解决】Understanding loss in Training LLM 解决方案

一、现象长什么样

训练自己的 LLM(用transformersTrainer或自己写的训练循环)时,遇到一类「看不懂 loss」的问题:

  • loss 数值异常大(比如 10+、20+),怎么调学习率都下不来;
  • loss 看上去在降,但模型生成全是乱码、复读;
  • 验证集 loss 比训练集还低,或者两边都不正常;
  • 切换到不同tokenizer/ 不同 padding 策略后,loss 量级突然变了,但模型结构没动。

最典型的复现:把一批长短不一的样本 pad 到同一长度送进模型,直接用input_idslabels算交叉熵,发现 loss 被 padding 位置严重拉高,训练目标其实是「学会预测padding token」,而不是「学会预测下一个真实 token」。

这类问题不报错,但训练出来的模型就是「不懂人话」——因为 loss 的含义从一开始就算错了。

二、背景

自回归 LLM 的训练目标是「给定前 i 个 token,预测第 i+1 个 token」。交叉熵 loss 对每个位置算一次,再平均。关键点:padding 位置不该参与 loss

transformersmodel(**inputs)在传入labels时,会自动对labels == -100的位置跳过(用ignore_index)。但很多人这么写:

# 错误写法:直接把 input_ids 当 labels outputs = model(input_ids=batch, labels=input_ids) loss = outputs.loss

如果 batch 里有 padding(input_idspad_token_id),那么 padding 位置也被当成「要预测的真实标签」,模型被迫去预测pad_token_id,这些位置的 loss 被算进平均。后果:

  1. loss 被 padding 稀释/拉高:短样本多的 batch,padding 占比大,loss 大部分在「学 padding」,真实语义信号被淹没。
  2. 训练目标错位:模型花大量精力拟合 padding,生成时容易吐 padding 或复读。
  3. 跨 tokenizer 不可比:不同 tokenizer 的 pad 比例不同,loss 量级跟着变,你以为换了模型,其实只是 pad 多了。

正确理解 loss 的前提,就是「让 padding 不参与 loss」。

三、根因

根因一句话:训练时labels没有把 padding 位置标成ignore_index=-100,导致交叉熵把 padding 也当成要预测的目标,loss 含义错误、训练目标错位。

三点展开:

  1. padding 参与计算labels = input_ids让 pad 位置进入 loss,ignore_index没生效。
  2. 平均基准错:loss 平均的分母包含 padding 位置数,真实 token 的梯度被稀释。3.缺校验:没有在送入模型前断言labels里 padding 已被-100覆盖,于是错误静默存在。

不是模型不会学,是「学什么」被 padding 污染了。

四、最小可运行复现

不依赖真实大模型,用一个最小交叉熵演示 padding 如何污染 loss:

import torch import torch.nn.functional as F vocab, seq = 10, 6 pad_id = 0 logits = torch.randn(1, seq, vocab) # 模型输出(未归一化) targets_raw = torch.tensor([[1, 2, 3, pad_id, pad_id, pad_id]]) # 含 padding # 错误:直接拿含 pad 的 target 算 loss loss_with_pad = F.cross_entropy( logits.view(-1, vocab), targets_raw.view(-1) ) # 正确:padding 标成 ignore_index=-100 targets_masked = targets_raw.clone() targets_masked[targets_raw == pad_id] = -100 loss_no_pad = F.cross_entropy( logits.view(-1, vocab), targets_masked.view(-1), ignore_index=-100 ) print("含 padding 的 loss:", round(loss_with_pad.item(), 4)) print("忽略 padding 的 loss:", round(loss_no_pad.item(), 4)) print("两者是否相同:", torch.isclose(loss_with_pad, loss_no_pad))

跑出来:含 padding 的 loss 把 3 个 pad 位置也学进去了,数值和「只看真实 3 个 token」的 loss 明显不同(pad 多时差异更大)。这就是「loss 算错」的精确复现。

五、解决方案(第一层:最小直接修复)

最小修复:构造labels时,把所有pad_token_id位置替换成-100,再送进模型。

import torch from transformers import AutoModelForCausalLM, AutoTokenizer tokenizer = AutoTokenizer.from_pretrained("your-model") model = AutoModelForCausalLM.from_pretrained("your-model") def make_labels(input_ids: torch.Tensor) -> torch.Tensor: labels = input_ids.clone() # 关键:padding 位置标成 -100,交叉熵忽略它 labels[labels == tokenizer.pad_token_id] = -100 return labels # 训练循环 for batch_input_ids in dataloader: labels = make_labels(batch_input_ids) outputs = model(input_ids=batch_input_ids, labels=labels) loss = outputs.loss # 现在只统计真实 token loss.backward() optimizer.step() optimizer.zero_grad()

如果做「下一 token 预测」且输入已经是「输入+标签移位」的格式,注意:自回归模型内部会自己处理移位,你只需保证labels里 padding 是-100,不要把labels再做一次[:, 1:]移位(那会和模型内部的 shift 重复)。

要点:

  • labels[labels == pad_token_id] = -100一行解决 padding 污染。
  • model(..., labels=labels)内部用ignore_index=-100自动跳过。
  • loss 现在只反映「真实 token 的预测质量」,量级和训练目标都正确。

这一步单独就让 loss 回归正确含义。

六、解决方案(第二层:结构性改进)

第一层是「在循环里加一行」。但训练脚本里多个数据路径(SFT、预训练、带 mask 的指令数据)都构造 labels,容易漏。更稳的做法把「labels 如何正确屏蔽 padding / 特殊 token」收敛成单一策略对象。

from dataclasses import dataclass, field from typing import List, Optional import torch @dataclass class LlmLossAuditor: """LLM 训练 loss 标签屏蔽的单一策略。""" # 需要忽略的 token id 集合(padding、特殊 token 等) ignore_ids: List[int] = field(default_factory=list) # 是否同时忽略序列左侧(prompt)只学回答(SFT 常用) train_on_completion_only: bool = False # completion 起始标记(SFT 用) response_start_id: Optional[int] = None def build_labels(self, input_ids: torch.Tensor) -> torch.Tensor: labels = input_ids.clone() for ig in self.ignore_ids: labels[labels == ig] = -100 if self.train_on_completion_only and self.response_start_id is not None: # 找到每个样本里 response_start 的位置,其之前全标 -100 mask = (input_ids == self.response_start_id) # 用 cumsum:start 之前为 0,之后为 1 pos = mask.cumsum(dim=-1) labels[pos == 0] = -100 return labels def check(self, labels: torch.Tensor): # 防御:整行全 -100 意味着该样本无监督信号 all_ignored = (labels == -100).all(dim=-1) if all_ignored.any(): print(f"[LlmLossAuditor] 警告: {int(all_ignored.sum())} 条样本整行被忽略") # 用法 auditor = LlmLossAuditor(ignore_ids=[tokenizer.pad_token_id, tokenizer.bos_token_id]) for ids in dataloader: labels = auditor.build_labels(ids) auditor.check(labels) loss = model(input_ids=ids, labels=labels).loss loss.backward(); optimizer.step(); optimizer.zero_grad()

结构收益:

  • 单一策略:padding、特殊 token、SFT「只学回答」的屏蔽都集中在一处。-可校验check抓出「整行无监督」的废样本。
  • 可扩展:加新的忽略规则只改LlmLossAuditor,不动训练循环。

七、解决方案(第三层:断言 / CI 守护)

写 pytest 守三条:(1) padding 被标-100;(2) 真实 token 不被误标;(3) 计算出的 loss 与「仅真实 token」一致。

import torch import torch.nn.functional as F import pytest from your_lib import LlmLossAuditor @pytest.fixture def auditor(): return LlmLossAuditor(ignore_ids=[0]) # 假设 pad_id=0 def test_pad_masked_to_neg100(auditor): ids = torch.tensor([[1, 2, 0, 0]]) labels = auditor.build_labels(ids) assert labels[0, 2].item() == -100 assert labels[0, 3].item() == -100 def test_real_tokens_kept(auditor): ids = torch.tensor([[1, 2, 0, 0]]) labels = auditor.build_labels(ids) assert labels[0, 0].item() == 1 assert labels[0, 1].item() == 2 def test_loss_ignores_pad(): vocab, seq = 10, 4 pad_id = 0 logits = torch.randn(1, seq, vocab) raw = torch.tensor([[1, 2, pad_id, pad_id]]) masked = raw.clone(); masked[raw == pad_id] = -100 l_pad = F.cross_entropy(logits.view(-1, vocab), raw.view(-1)) l_mask = F.cross_entropy(logits.view(-1, vocab), masked.view(-1), ignore_index=-100) assert not torch.isclose(l_pad, l_mask), "含 pad 的 loss 应与忽略 pad 的不同" def test_completion_only_mode(): a = LlmLossAuditor(ignore_ids=[0], train_on_completion_only=True, response_start_id=5) ids = torch.tensor([[1, 5, 6, 7]]) # 5 之后才是回答 labels = a.build_labels(ids) assert labels[0, 0].item() == -100 # prompt 部分忽略 assert labels[0, 1].item() == -100 # response_start 本身可忽略 assert labels[0, 2].item() == 6 # 回答部分保留

CI 常驻跑这四条后,任何「padding 又混进 loss」「真实 token 被误标」的回归都会立刻爆红。

八、排查清单

训练 LLM「loss 看不懂」时按顺序查:

  1. 先打印labelspad_token_id是否还在——在就说明 padding 参与了 loss。
  2. 确认labelsinput_ids的克隆并做了-100替换,而不是直接用input_ids
  3. 确认用的loss来自model(..., labels=labels).loss,而不是自己手写的、没传ignore_indexF.cross_entropy
  4. SFT 场景确认是否「只学回答」:prompt 部分应标-100,否则模型在学复述问题。
  5. 确认ignore_index=-100与模型内部一致(transformers 默认就是 -100,别改成别的)。
  6. 换 tokenizer 后重新核对pad_token_id,不同 tokenizer 的 pad id 可能不同。
  7. 批量打印几个样本的labels肉眼确认:除 pad/特殊位外,真实 token 都保留。

九、小结

「训练 LLM 但 loss 看不懂 / 模型学不会」的常见根子是labels没把 padding 标成ignore_index=-100,交叉熵把 padding 也当成学习目标,loss 含义错位、训练目标被污染。修复三层次:第一层构造 labels 时labels[labels==pad_token_id] = -100;第二层用LlmLossAuditordataclass 把 padding/特殊 token/SFT「只学回答」的屏蔽收敛为单一策略并加整行忽略校验;第三层用 pytest 守「padding 被标 -100」「真实 token 保留」「含 pad 与忽略 pad 的 loss 不同」「completion-only 正确」。

工程启示:自回归训练里,「labels 怎么构造」决定了「模型学什么」。padding 必须-100、prompt(SFT 时)必须-100、特殊 token 通常也要-100。任何训练脚本上线前,先肉眼看一眼labels再训,比训完发现模型废了再回头查省事得多。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/9 11:16:54

Apigee Hybrid与Cassandra的API管理数据存储优化实践

1. 项目概述:当API管理遇上分布式数据库在混合云架构成为企业标配的今天,Apigee Hybrid作为API管理平台中的"瑞士军刀",其底层数据存储机制直接决定了API流量分析、策略执行和监控告警的可靠性。而Cassandra这个高度可扩展的NoSQL数…

作者头像 李华
网站建设 2026/8/9 11:16:46

商用饮水机选购指南:核心考量与避坑要点

1. 商用饮水机选购核心考量因素 商用饮水机与家用产品存在本质区别,需要从以下几个维度进行专业评估: 1.1 日均供水量测算 根据我服务过30企业的经验,建议按照"员工人数1.5L1.2安全系数"计算基础需求。例如100人团队:…

作者头像 李华
网站建设 2026/8/9 11:16:45

电力系统集群规划中的空间约束优化方法与实践

1. 电力系统集群规划的背景与挑战现代电力系统正朝着分布式、智能化的方向快速发展,集群化运行已成为提升系统可靠性和经济性的重要手段。但在实际规划中,我们常常面临一个关键问题:如何合理划分电力设备集群,使其既符合电气连接特…

作者头像 李华
网站建设 2026/8/9 11:15:53

Project Deskless:本地部署语音驱动AI智能体Viktor的实践指南

这次我们来看一个名为 Project Deskless 的开源项目,它主打一个非常直接的概念:通过语音指令,一键指挥一个名为 Viktor 的 AI 员工为你工作。想象一下,你只需要对着麦克风说出任务,比如“帮我写一份周报”或“分析…

作者头像 李华
网站建设 2026/8/9 11:13:23

DC-3靶机渗透实战:从Web漏洞到Root提权

1. DC-3靶机实战概述 DC-3是Vulnhub平台上经典的渗透测试训练靶机,基于Debian系统构建,包含多个故意设置的安全漏洞。这个靶机特别适合中阶渗透测试人员练习综合渗透技巧,从信息收集到权限提升的全流程都能得到充分锻炼。我花了三周时间反复测…

作者头像 李华
网站建设 2026/8/9 11:13:03

3步让老旧游戏手柄重获新生:XOutput完整指南

3步让老旧游戏手柄重获新生:XOutput完整指南 【免费下载链接】XOutput DirectInput to XInput wrapper 项目地址: https://gitcode.com/gh_mirrors/xo/XOutput 你是否曾为那些功能完好却被现代游戏"抛弃"的老旧游戏手柄感到惋惜?那些陪…

作者头像 李华