DROPS 类不平衡学习实战:分布鲁棒后处理在长尾分类中的应用(google-research/drops)
2026/9/20 15:47:22 网站建设 项目流程
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

导读:本文围绕 google-research 仓库中drops目录提供的实验代码,系统讲解 DROPS(Distribution RObust PoSthoc)方法在类不平衡(长尾分布)图像分类任务上的完整使用方案。你将掌握两步变体(先以 CE 损失训练基线,再对保存的 logits 做分布鲁棒后处理)与一步变体(端到端训练)两条训练路径的全部命令行参数,并从源码层面理解类级权重g_yalpha_y的更新机制、KL 散度约束下的最坏情况精度优化原理,以及配套的多种长尾损失函数族。

一、DROPS 是什么:将分布鲁棒优化引入类不平衡学习

真实场景中的数据往往呈现长尾分布(long-tailed distribution):少数"头类"拥有大量样本,而大量"尾类"样本稀少。直接用交叉熵(CE)损失训练的模型会严重偏向头类,导致尾类分类精度大幅下降。围绕这一经典问题,drops目录提供了 DROPS 项目的实验代码,核心思路是借助**分布鲁棒优化(Distributionally Robust Optimization, DRO)**来重新权衡类别权重,从而提升长尾场景下的整体与最坏类别表现。

从代码注释(drops/main_lt.py)可以看到,该目录定位为"Long-tail experiments",即长尾实验的实现,其方法名DROPS = Distribution RObust PoSthoc,含义为"带 DRO 约束权重的 posthoc 后处理":

  • 训练阶段先用普通 CE 损失得到基线模型;
  • 随后在已保存的模型预测 logits上,通过求解带散度约束的最优化问题,搜索一组"最坏情况"的类别权重g_y,并据此对 logits 做偏移(posthoc adjustment),使得模型对任何位于约束半径eps内的类别分布扰动都具有鲁棒性。

该方法以两种变体形式提供,二者的区别在于 DRO 权重搜索发生在何时:

变体训练流程对应命令
两步变体(Two-step)Step 1 用 CE 损失训练基线模型并保存 logits;Step 2 在保存的 logits 上离线执行 DROPS 后处理main_lt+drops_test_time
一步变体(One-step)训练过程中直接用 DROPS 损失交替更新模型与类权重main_lt --loss 'drops'

二、运行环境与依赖安装

drops目录的代码基于 TensorFlow 2(含 TensorFlow Datasets 与 Addons),并依赖cvxpy求解 DRO 约束优化问题。官方提供的环境初始化脚本 drops/run.sh 给出了完整流程:

set -e set -x virtualenv -p python3 . source ./bin/activate pip install tensorflow pip install -r drops/requirements.txt python -m drops.main_lt --dataset 'cifar10' --loss 'ce' --imb_ratio 0.1 --dro_div 'kl' --num_iters 3 --eval_freq 3

其中 drops/requirements.txt 声明的依赖版本约束为:

numpy>=1.16.0 tensorflow>=2.4.0 tensorflow_datasets tensorflow_addons edward2 cvxpy
  • tensorflow_addons用于模型中的 Group Normalization 层(见 drops/preact_resnet_models.py);
  • edward2preact_resnet_models.py导入,用于概率层扩展;
  • cvxpy在训练与测试时评估 DRO 精度时求解凸优化问题(见 drops/main_lt.py)。

run.sh中给出的验证命令特意使用了很小的num_iters 3eval_freq 3,方便快速跑通流程;实际实验请按 drops/README.md 的默认配置(5 万迭代)执行。

三、两步变体实战:先训 CE 基线,再做分布鲁棒后处理

两步变体是 README 首先介绍的训练方式,其设计动机是:DROPS 的核心计算只需要模型预测的 logits,无需重新训练模型,因此可以完全解耦"训练"与"后处理"两个阶段。

3.1 Step 1:训练 CE 损失基线模型

python -m drops.main_lt --dataset 'cifar10' --loss 'ce' --imb_ratio 0.1 --dro_div 'kl'

该命令会在合成长尾 CIFAR-10 数据集上以 CE 损失训练基线模型。main_lt.py内部做了以下几件事:

  1. 构造长尾数据集:通过get_cls_num(drops/main_lt.py)按指数衰减规则为每个类分配样本数——img_num_per_cls = img_max * (imb_factor ** (cls_idx / (cls_num - 1))),其中imb_ratioimg_min / img_max。随后get_cls_idx(drops/main_lt.py)按每类配额从原始数据中随机抽取索引,形成不平衡训练子集。
  2. 划分数据集get_dataset(drops/main_lt.py)使用train[:90%]作为训练集、train[90%:]作为验证集、test作为测试集,并施加 CIFAR 标准预处理(32×32 随机裁剪、随机水平翻转、按通道均值/方差归一化)。
  3. 构建模型:调用resnet_models.create_resnet18(input_shape=(32, 32, 3), num_classes=FLAGS.num_classes, norm='batch')(drops/main_lt.py)创建 ResNet-18 预激活网络;该网络的 BatchNorm decay 为 0.9、L2 weight decay 为 1e-4(drops/preact_resnet_models.py)。
  4. 训练:使用 Nesterov SGD(momentum 0.9)配合分段常数学习率计划——在 30 / 80 / 110 epoch 处分别衰减至 0.01 / 0.001 / 0.0001(drops/main_lt.py)。

训练过程中,模型会按验证集上的三类指标分别保存最优 checkpoint 及对应 logits(drops/main_lt.py):

  • 平均精度最优 →mean_best.h5val/test_mean_best_{logits,labels}.txt
  • 最坏类精度最优 →worst_best.h5val/test_worst_best_{logits,labels}.txt
  • DRO 精度最优 →dro_best.h5val/test_dro_best_{logits,labels}.txt

保存 logits 的功能由save_logits(drops/main_lt.py)完成,它会将全部验证/测试样本的 logits 与标签以文本格式写入model/子目录,这正是 Step 2 后处理的输入。

3.2 Step 2:在保存的 logits 上执行 DROPS 后处理

python -m drops.drops_test_time --dataset 'cifar10' --loss 'drops' --imb_ratio 0.1 --eta_lambda 10 --num_iters 25 --eta_lambda_mult 0.95 --prior_type 'train' --cal_type 'none' --eps 0.9

drops_test_time.py纯测试时(test-time)后处理入口,其 docstring 明确说明"Train and evaluate models of long-tail experiments with DROPS only at test time"(drops/drops_test_time.py)。执行流程如下:

  1. 根据model_dirdatasetimb_ratiorun_id自动拼出基线模型目录,并固定加载_loss_ce_{run_id}的 CE 模型(drops/drops_test_time.py),即它只对 CE 训练的基线做后处理;
  2. 通过get_pre_train(drops/drops_test_time.py)读取val_{req}_best_logits.txttest_{req}_best_logits.txtreq--req指定,默认dro,可选mean/dro/worst,即用哪种指标选出的基线模型);
  3. 以验证集 logits 为优化数据,迭代更新类权重g_y、拉格朗日乘子lambd与后处理偏移alpha_y
  4. 用更新后的权重对测试 logits 施加偏移并评估平均精度、最坏类精度与 DRO 精度。

3.3 参数详解

main_lt.pydrops_test_time.py均使用absl.flags解析命令行参数。main_lt.py的完整参数如下(默认值取自源码定义,drops/main_lt.py):

参数默认值说明
--model_dir./模型与日志存储目录,实际会追加{dataset}_{imb_ratio}/_loss_{loss}_{run_id}/子路径
--lr0.1初始学习率
--num_iters50000训练迭代次数
--eval_freq1000评估频率
--batch_size128批大小
--train_on_fullFalse是否在完整训练集(而非长尾子集)上训练
--run_id0实验编号,用于区分同配置多次运行
--datasetcifar10数据集,支持cifar10/cifar100
--num_classes10类别数(cifar100 时需改为 100)
--imb_ratio0.1不平衡因子n_min / n_max
--lossdrops损失类型:ce/focal/bsm/cb/cb_focal/ldam/logit_adj/posthoc/posthoc_ce/drops
--alpha1.0focal 损失的 alpha 参数
--gamma0.5focal / LDAM(C) 的 gamma 参数
--beta0.9999class-balanced 损失的 beta 参数
--s1LDAM 中 logits 的缩放参数
--re_weight_typeprior类别重加权方式:none/prior/sqrt
--warmup0热身迭代数,期间使用 CE 损失
--is_upsamplingFalse是否对尾类做上采样平衡
--dro_divklDRO 度量使用的散度:kl/l2/l1(也支持reverse-kl
--eps0.9DRO 度量的扰动半径(散度约束上界 delta)
--metric_baseuniformDRO 度量的类权重基准:prior/recip_prior/uniform
--eta_g0.01EG(指数梯度)更新g_y的步长
--eta_lambda0.01EG 更新lambd的步长
--n_it_update1glambda的更新频率(每多少步更新一次)
--tau1.0logit adjustment 项的温度常数
--weight_typece计算类损失时使用的权重方式:0_1(0-1 损失)或ce(交叉熵)
--g_typenot-eg是否使用 EG(指数梯度)风格更新g_yeg或简化式not-eg

drops_test_time.py的参数在main_lt.py基础上增加了三个后处理专属选项(drops/drops_test_time.py):

参数默认值说明
--prior_typetrain类别先验来源:train(用训练集指数分布先验)或val
--cal_typenone校准方式:temp(温度)/scale(缩放)/shift(偏移)/shift_scale/none
--eta_lambda_mult1.0每轮迭代对eta_lambda的衰减因子(README 示例取0.95,即逐步缩小拉格朗日乘子步长)
--reqdro选择基线模型的标准:mean/dro/worst

drops_test_time.pyeps的默认值为0.2num_iters默认25eta_lambda默认0.1;README 示例将其调整为eps 0.9num_iters 25eta_lambda 10并配合eta_lambda_mult 0.95,属于 DRO 约束半径较大的典型配置。

四、一步变体实战:端到端训练 DROPS

一步变体将 DRO 权重更新融入训练循环,无需事后处理:

python -m drops.main_lt --dataset 'cifar10' --loss 'drops' --imb_ratio 0.1 --dro_div 'kl'

与 CE 训练的唯一区别是--loss 'drops'。在训练过程中,main_lt.pyn_it_update(默认 1)步就会调用updateg验证集上更新g_ylambdalpha_y,随后用更新后的alpha_y重建 DROPS 损失函数继续训练(drops/main_lt.py)。因此一步变体的训练开销高于两步变体——DRO 权重搜索发生在每一次参数更新之间,而不是只做一次。

五、源码级原理:DROPS 的优化机制

5.1 类权重g_yalpha_y的初始化

在 drops/main_lt.py 中,DROPS 的核心变量按以下规则初始化:

  • g_y(类权重,对应 DRO 中的分布g):默认以metric_base='uniform'初始化为均匀权重;若metric_base='prior'则取1/样本数,若metric_base='recip_prior'则取样本数本身,最后归一化为概率分布;
  • alpha_y(后处理偏移系数):初始化为1/样本数再乘以样本总数,即alpha_y ∝ 1/π_y
  • r_list:约束D(u, g) < delta中的基准分布u,默认与g_y相同;
  • lambd:拉格朗日乘子,初始化为 1.0。

5.2updateg:伪代码第 1 步的完整实现

updateg(drops/main_lt.py)实现了 DRO 权重更新的全部 8 个步骤,代码注释将其描述为"Pesuedo code 1: (Updating g, lambda)":

  1. 计算每类损失L_y:对验证集每个 batch,先取 softmax 预测,再用alpha_y加权得到"后偏移预测"y_weighted_pred;按weight_type分别用 CE 损失或 0-1 损失累计每个类别的平均损失loss_y_list
  2. 计算散度约束项:按dro_div计算D(r_list, g_y)——l2为平方和、l1为绝对值和、klΣ g_y·log(g_y/r_y)reverse-klΣ r_y·log(r_y/g_y)
  3. 构造 LagrangianL = Σ g_y·L_y - lambd·(D - delta)
  4. 更新g_yg_type='eg'时按指数梯度(EG)闭式更新,否则用简化式g_y ← r_y·exp(L_y / lambd)后归一化;
  5. 更新lambdlambd ← lambd + eta_lambda·cons
  6. 更新alpha_yalpha_y = g_y / π_y(乘以样本总数归一)。

5.3 评估指标:DRO 最坏情况精度

eval_dro_metrics(drops/main_lt.py)通过cvxpy求解一个凸优化问题来度量模型在分布扰动下的最坏情况精度:

minimize Σ v_y · acc_y subject to v >= 0, Σ v = 1, D(v, u) <= eps

其中v是待求解的最坏情况类权重,u是均匀基准分布,散度约束D(v, u) <= eps依据dro_div分别实现为 L2、L1、KL 或 reverse-KL 约束(drops/main_lt.py)。求解失败时自动回退到 SCS 求解器(drops/main_lt.py)。最终 DRO 精度为Σ v*·acc_y,即在允许的分布扰动范围内最坏可能的加权精度。在klreverse-kl模式下,训练与测试评估还会遍历一组预定义的eps网格(eps_list),输出 DRO 精度随扰动半径变化的曲线数据(drops/main_lt.py)。

5.4 测试时后处理的差异

drops_test_time.py中的updateg(drops/drops_test_time.py)在实现上稍有不同:它直接操作已加载的 logits 数组而非数据迭代器,且拉格朗日乘子更新采用eta_lambda *= eta_lambda_mult的逐步衰减策略(drops/drops_test_time.py),迭代在约束残差|cons| < 1e-5时提前终止。这使得它非常适合在单个 GPU 上对预训练 logits 做快速离线优化。

六、测试时校准:cal_type的作用

drops_test_time.py内置了四类 logits 校准工具(drops/drops_test_time.py),在 DROPS 后处理之前可选执行,用于修正基线模型 logits 的温标/偏移/缩放失配:

cal_type实现函数变换形式
tempcalibrate_templogits * temp(温度校准)
shiftcalibrate_shiftslogits - shifts
scalecalibrate_scaleslogits * scales
shift_scalecalibrate_shifts_scaleslogits * scales - shifts
none不校准

校准参数(tempshiftsscales)均在验证集 logits 上通过 100 轮 SGD 最小化 CE 损失学习得到,初始学习率 0.1 并以多项式衰减至 0.01(如calibrate_temp,drops/drops_test_time.py)。README 的示例默认使用--cal_type 'none',即直接对原始 logits 做 DROPS 后处理。

七、损失函数族:losses_lt.py的完整支持清单

MakeLossFunc(drops/losses_lt.py)是损失工厂函数,根据--loss参数返回对应的长尾损失实例,覆盖了长尾分类领域的主要方法:

loss 名称实现类/函数关键超参
ceCELoss
up_ceCELoss(配合上采样)--is_upsampling
ldamLDAMLoss--gamma--s
focalFocalLoss--gamma
cbCBLoss--beta
cb_focalCBFocal--beta--gamma
bsmBalancedSoftmax
logit_adjLogitAdjust--tau
posthocLogitAdjust(仅评估时偏移)--tau
posthoc_ceCELoss+ 评估时 posthoc 偏移--tau
dropsCELoss+ DRO 约束类权重偏移--eps--dro_div

其中LogitAdjust实现了 logit adjustment 损失:logits + tau·log(π_y)(drops/losses_lt.py),BalancedSoftmax则对 softmax 前指数用类先验加权(drops/losses_lt.py)。DROPS 的"后偏移"与 logit adjustment 形式一致,但关键在于偏移量alpha_y不再由静态先验决定,而是由 DRO 优化动态求解——这正是它与posthoclogit_adj的本质区别。

八、运行结果与日志解读

两类入口都会在model_dir下输出文本日志与 TensorBoard 摘要:

  • results_log.txtmain_lt.py)或results_loss=...,delta_train=...命名的日志文件(drops_test_time.py),逐条记录验证/测试损失、平均精度、最坏类精度以及各eps网格点上的 DRO 精度;
  • summaries/trainsummaries/eval目录供 TensorBoard 可视化(drops/main_lt.py);
  • drops_test_time.py的日志末尾还会以机器可读的delta_evals=DRO-Acc=kldiv_cons=三段式输出,便于直接绘制 DRO 精度–扰动半径曲线(drops/drops_test_time.py)。

九、适用前提与注意事项

  1. 数据集支持:源码只实现了 CIFAR-10 / CIFAR-100 的合成长尾构造(get_cls_num中按类别索引指数分配样本数);drops_test_time.py--dataset注释虽提到 imagenet-lt,但当前仓库未包含对应数据加载实现,实际可用范围为 CIFAR 系列。
  2. 两步变体的耦合约束drops_test_time.py固定读取 CE 基线的 logits 文件,因此必须先以--loss 'ce'完整跑完main_lt并确保model/下存在val/test_{req}_best_{logits,labels}.txt,否则后处理会因文件缺失而失败。
  3. num_classes需与数据集匹配:cifar100 时请将--num_classes 100--dataset 'cifar100'一并设置,否则 logits 的形状解析会出错。
  4. 实验定位:README 明确标注该目录为Experimental code,适用于复现与学术研究,生产环境部署前需结合自身数据做充分验证。

通过本文提供的两条命令路径与参数对照表,配合对updategeval_dro_metrics等核心函数的源码理解,你可以直接在 drops 目录中复现 DROPS 的两步/一步训练流程,并进一步修改epsdro_diveta_lambda等关键超参,探索分布鲁棒优化在长尾分类中的行为特性。

  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询