Kev 自训练 loss 不降 / 是数据格式还是 batch 设置有问题?

文章导读
Kev 遇到的这类现象,通常不是单一变量造成的:数据没对齐、标签错位、字段名写错,和 batch size、学习率、序列长度设置不当,会给出非常接近的 loss 曲线。可操作的判断顺序是先把数据层排除掉——能不能正确读到样本、输入和标签是否对得上、模型能不能在极少样本上把 loss 压下去;确认模型本身有学习能力之后,再逐项调超参,最后回头核对损失函数与任务是否匹配。这些环节都可以通过打印样本、翻
📋 目录
  1. Ⅰ 用一小批数据检查输入和标签是否对齐
  2. Ⅱ 在训练日志里看 loss 是平、震荡还是 NaN
  3. Ⅲ 用极小数据集做一次过拟合测试
  4. Ⅳ 逐项调整 batch size、学习率和序列长度
  5. Ⅴ 确认自训练配置里的损失函数与任务匹配
A A

Kev 遇到的这类现象,通常不是单一变量造成的:数据没对齐、标签错位、字段名写错,和 batch size、学习率、序列长度设置不当,会给出非常接近的 loss 曲线。可操作的判断顺序是先把数据层排除掉——能不能正确读到样本、输入和标签是否对得上、模型能不能在极少样本上把 loss 压下去;确认模型本身有学习能力之后,再逐项调超参,最后回头核对损失函数与任务是否匹配。这些环节都可以通过打印样本、翻训练日志、跑小实验来验证,不需要靠猜。

loss 不降优先怀疑数据:输入与标签错位、padding 与 labels 不对应、字段名写错,都会让模型学不到东西。建议先打印一条样本人工比对,再用几条样本做过拟合测试;如果连这几条都记不住,先别动 batch size 和学习率。只有在过拟合测试能降 loss、而全量训练不降时,才把注意力转到学习率、batch size、序列长度和损失函数配置上。

用一小批数据检查输入和标签是否对齐

先确认数据格式问题是不是第一原因。做法很简单:绕开训练循环,直接从 Dataset 或数据文件里取一条样本打印出来,人工看输入和标签是否语义对应。字段名以实际仓库为准,下面的骨架只是说明要打印哪些内容。

# 读一条样本,打印输入和标签的形状与内容
sample = train_dataset[0]          # 换成实际的 Dataset / 读数据函数
tokens = sample["input_ids"]        # 字段名以仓库实现为准
labels = sample["labels"]
print("keys:", list(sample.keys()))
print("tokens shape:", getattr(tokens, "shape", None), "head:", tokens[:16])
print("labels shape:", getattr(labels, "shape", None), "head:", labels[:16])
print("tokens dtype:", getattr(tokens, "dtype", None))
print("labels dtype:", getattr(labels, "dtype", None))
print("decoded:", tokenizer.decode(tokens[:64]))   # 有 tokenizer 时看文本
print("label text:", tokenizer.decode(labels[:64]))

打印完之后逐项比对,比看 shape 更重要:

  • tokens 和 labels 的长度是否一致,差一位往往意味着错位一格。
  • labels 里的忽略位(常见为 -100 或仓库自定义的 ignore_index)是否只出现在该忽略的位置。
  • 分类任务的 label 是否落在 [0, num_classes) 范围内,越界会让 loss 直接失控或报错。
  • 回归任务的 label 量纲是否和输出层一致,例如标签是 0-1、模型输出却是原始尺度。
  • padding 的位置在 loss 计算里是否被 mask 掉,pad token 被当成有效标签时,loss 会被稀释得很难下降。

人工比对的方式是:把 decoded 出来的输入和对应标签读一遍,判断“这条样本的标签是不是这段输入该有的答案”。如果肉眼都看不出对应关系,训练侧再怎么调参也很难救回来。

在训练日志里看 loss 是平、震荡还是 NaN

“不降”其实是三种不同曲线的统称,排查方向不一样。先把日志里的 loss 抽出来画一条简单曲线,形态比单个数值更有信息量。

# 从日志抽 loss 并画曲线(字段名按自己日志格式调整)
grep -oE "loss[=: ]+[0-9.eE+-]+" train.log | awk '{print NR, $NF}' > loss.txt
python - <<'PY'
xs, ys = [], []
for line in open("loss.txt"):
    a, b = line.split()
    xs.append(int(a)); ys.append(float(b))

import matplotlib.pyplot as plt
plt.plot(xs, ys)
plt.xlabel("step"); plt.ylabel("loss")
plt.savefig("loss_curve.png")
print("first:", ys[:5], "last:", ys[-5:])
PY
  • 平:几乎一条水平线。优先查数据与标签对齐、loss 是否被 padding 稀释、输出层是否被冻结,超参放在这一步之后再看。
  • 震荡:上下跳动但没有下降趋势。优先查学习率是否偏大、batch size 是否过小导致梯度噪声大、数据是否被 shuffle 得过于无序。
  • NaN:跳跃出现或从某步开始。优先查数值稳定性——学习率过大、梯度爆炸、混合精度下的溢出、标签越界或含非法值。

日志里的 loss 如果是按 step 打印的,建议同时记录对应的学习率(有些框架会在日志里输出当前 lr),便于判断震荡是不是由 warmup 或调度器引起。曲线形态确认之后,再决定下一步动数据还是动超参。

用极小数据集做一次过拟合测试

这是把“模型有没有学习能力”和“超参调得好不好”分开的关键一步。思路是拿几条样本反复训练,如果模型结构、损失函数、数据管道都没问题,loss 应该能明显下降,甚至接近记住这几条样本。

Kev 自训练 loss 不降 / 是数据格式还是 batch 设置有问题?
# 过拟合测试:样本数极少、步数很少
from torch.utils.data import DataLoader, Subset

tiny = Subset(train_dataset, list(range(0, 8)))   # 只取几条,按需调整
loader = DataLoader(tiny, batch_size=<小到能放下>, shuffle=True)

model.train()
for step, batch in enumerate(loader):
    out = model(**batch)
    loss = out.loss            # 具体取哪个字段以仓库实现为准
    loss.backward()
    optim.step()
    optim.zero_grad()
    if step % 5 == 0:
        print(step, float(loss))
    if step >= 50:             # 只跑几十步
        break

验证标准是 loss 是否明显下降。明显下降说明模型和数据管道具备学习能力,问题更可能在超参或全量数据分布上;如果 loss 在这几条样本上仍然不降,优先回到第 1 节检查标签对齐和损失函数配置,而不是继续调 batch size。

逐项调整 batch size、学习率和序列长度

确认模型能过拟合之后,再用控制变量法把超参影响从数据问题里分离出来:每次只改一个参数,其他参数保持不变,逐个实验记录 loss 变化。下面是可以直接套用的实验表格模板,参数值按自己当前配置填写。

实验编号改动项其他参数记录内容观察到什么
E0基线,不改动lr / batch / seq_len 保持当前值前若干步与后若干步的 loss曲线形态:平 / 震荡 / NaN
E1只改 batch sizelr、seq_len 不变同样步数下的 loss 序列是否比 E0 更容易下降
E2只改 lrbatch、seq_len 不变同样步数下的 loss 序列是否震荡加剧或变成 NaN
E3只改序列长度 / 截断长度batch、lr 不变loss 与显存占用是否因截断丢失关键标签

几点需要结合环境确认的地方:batch size 变化时,有些框架要求同步调整学习率或 warmup 步数,否则两个变量被一起改了,实验结果就不可比;序列长度变大通常显存占用上升,可能触发框架自动换用更小的有效 batch,这一点要在日志或显存监控里确认。表里只记录 loss 变化本身,不预设哪个参数一定有效。

确认自训练配置里的损失函数与任务匹配

最后一步是排除“假性不收敛”:目标函数和输出格式不匹配时,loss 也可能长期不降。先在配置文件或训练脚本里找到损失函数字段,确认它和任务类型是否对应。

# 通用检查方式:打印实际生效的配置与 loss 对象
print(cfg.loss)                       # 配置里的 loss 字段名以仓库为准
print(type(criterion))                # criterion 是框架内置还是自定义
out = model(**batch)
print("model output keys:", out.keys())
print("logits shape:", out.logits.shape)   # 与 labels 形状对照
  • 单标签分类:交叉熵类损失,标签是类别索引;若输出已经过 softmax,再套一层交叉熵会重复归一化。
  • 多标签分类:二分类交叉熵类损失,标签通常是 0/1 向量,形状要和 logits 对齐。
  • 回归:L1/L2 或平滑 L1 类损失,输出维度与标签维度必须一致。
  • 序列到序列:交叉熵作用于每个位置的 logits,labels 需要做位移(shift)与忽略位处理。

具体用哪一种、字段叫什么名字、是否内部自动 shift,都要以实际仓库实现为准。确认任务与损失函数对应之后,如果前面的过拟合测试是通过的,就可以回到数据与超参两侧继续收窄范围,而不必同时改动多个地方。