欢迎光临
我们一直在努力

损失函数 中的高损失样本筛查

损失函数 中的高损失样本筛查

flyfish

用当前已经训练出一定能力的模型,去推理整个训练集,逐张计算每张图片的分类损失,再按损失从高到低排序。
损失越高,代表模型越搞不定这张图——大概率是标注错误、图片模糊、目标极小、属于极难分样本。
它是做数据清洗、排查错标、定位猫狗互判难样本的手段。

@torch.no_grad()
def check_high_loss_samples(model, dataloader, device, criterion=None):
"""
逐样本计算交叉熵损失,按损失降序排列,用于排查标注错误
:param model: 训练好的基线模型
:param dataloader: 训练集dataloader
:param device: 设备
:param criterion: 损失函数,默认用标准交叉熵(客观无加权)
:return: 按损失排序的DataFrame
"""

if criterion is None:
criterion = nn.CrossEntropyLoss(reduction='none').to(device)

model.eval()
sample_records = []

for batch_idx, (images, labels, img_paths) in enumerate(tqdm(dataloader, desc="筛查高损失样本")):
images = images.to(device)
labels = labels.to(device)

# 前向推理
logits = model(images)
# 逐样本计算损失(不取平均)
loss_per_sample = criterion(logits, labels)
# 获取模型预测类别
preds = logits.argmax(dim=1)

# 记录每个样本信息
for i in range(len(images)):
sample_records.append({
'img_path': img_paths[i],
'true_label': labels[i].item(),
'true_label_name': CLASS_NAMES[labels[i].item()],
'pred_label': preds[i].item(),
'pred_label_name': CLASS_NAMES[preds[i].item()],
'loss': round(loss_per_sample[i].item(), 6)
})

# 按损失降序排列
df = pd.DataFrame(sample_records)
df_sorted = df.sort_values(by='loss', ascending=False).reset_index(drop=True)
return df_sorted

解释

1. 函数头与装饰器

@torch.no_grad()
def check_high_loss_samples(model, dataloader, device, criterion=None):

@torch.no_grad():关闭梯度计算,全程只做推理、不反向传播、不更新参数。作用是大幅节省显存、提升计算速度,是所有推理类函数的标准写法。
参数说明:
model:已经训练好的模型(一般用阶段1训练完的模型)
dataloader:要筛查的数据集加载器(一般传训练集train_loader)
device:CPU/GPU设备
criterion:损失函数,可选;不传就默认用标准交叉熵

2. 损失函数初始化

if criterion is None:
criterion = nn.CrossEntropyLoss(reduction='none').to(device)

默认使用标准交叉熵损失,不用训练用的FocalLoss,原因是:
FocalLoss加了gamma聚焦、alpha类别加权,损失数值是被人工调制过的,不能客观反映样本本身的难易程度;标准交叉熵无加权、无修改,能最公平地反映模型预测和真实标签的差距,筛查结果更客观。
关键参数 reduction='none':
不做取平均/求和,输出和batch长度完全一致的张量,每个位置对应一张图的单独损失值。这是实现逐样本筛查的基础。

3. 主循环:逐批次推理计算

model.eval()
sample_records = []

for batch_idx, (images, labels, img_paths) in enumerate(tqdm(dataloader, desc="筛查高损失样本")):
images = images.to(device)
labels = labels.to(device)

logits = model(images) # 前向推理,得到模型原始输出
loss_per_sample = criterion(logits, labels) # 逐样本计算损失
preds = logits.argmax(dim=1) # 得到模型预测的类别

model.eval():切换到评估模式,关闭Dropout、BatchNorm等训练专属层,保证推理结果稳定。
logits:模型的原始输出(还没经过softmax转概率),交叉熵内部会自动做softmax。
loss_per_sample:长度等于batch size的一维张量,每个数对应一张图的损失值。
preds:每张图模型预测的类别编号(0/1/2)。

4. 逐样本记录信息

for i in range(len(images)):
sample_records.append({
'img_path': img_paths[i], # 图片文件路径,用于定位原图
'true_label': labels[i].item(), # 真实标签(数字)
'true_label_name': CLASS_NAMES[labels[i].item()], # 真实标签(类别名)
'pred_label': preds[i].item(), # 预测标签(数字)
'pred_label_name': CLASS_NAMES[preds[i].item()], # 预测标签(类别名)
'loss': round(loss_per_sample[i].item(), 6) # 该样本的损失值
})

把每张图的完整信息存成字典,最后汇总成列表。
价值是 img_path:能直接定位到具体图片文件,方便人工打开核对标注是否正确。
true_label_name / pred_label_name:不用自己对照编号,一眼就能看出真实是猫、预测成狗这类互判情况。

5. 结果排序与返回

df = pd.DataFrame(sample_records)
df_sorted = df.sort_values(by='loss', ascending=False).reset_index(drop=True)
return df_sorted

把所有记录转成 pandas DataFrame 表格格式,方便筛选、排序、保存。
按 loss 列降序排列:损失越高、问题越大的样本排在最前面,优先排查Top样本收益最高。
重置索引,返回排序后的表格。

调用方法与使用流程

1. 最佳调用时机

在阶段1冻结训练结束后调用,收益最高:
此时模型已经学懂了基础分类规律,不会因为模型太菜导致所有样本损失都很高;
筛查出来的高损失样本,才是真的标注有问题、或者特别难分的样本。
不要在模型随机初始化的时候调用,没有参考意义。

2. 基础调用示例(主函数里的标准用法)

# 调用函数,得到排序后的结果表格
high_loss_df = check_high_loss_samples(model, train_loader, DEVICE)

# 打印损失最高的30张样本,快速浏览
print("\\n损失最高的30个样本:")
print(high_loss_df[['img_path', 'true_label_name', 'pred_label_name', 'loss']].head(30))

# 保存为CSV文件,方便人工逐张打开核对
high_loss_df.to_csv("train_high_loss_samples.csv", index=False, encoding="utf-8-sig")

3. 进阶用法:筛选特定互判样本

比如现在最关心猫被误判成狗的样本,可以直接从结果里筛选:

# 筛选:真实是cat,预测是dog的样本
cat_dog_confuse = high_loss_df[
(high_loss_df['true_label_name'] == 'cat') &
(high_loss_df['pred_label_name'] == 'dog')
]
print("猫误判成狗的难样本数量:", len(cat_dog_confuse))
print(cat_dog_confuse.head(20))

可以快速定位最难分的猫狗互判样本,针对性补充数据或修正标注。

赞(0)
未经允许不得转载:171主机测评 » 损失函数 中的高损失样本筛查
分享到: 更多 (0)

评论 抢沙发

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