- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
导读:本文围绕 google-research 仓库中drops目录提供的实验代码,系统讲解 DROPS(Distribution RObust PoSthoc)方法在类不平衡(长尾分布)图像分类任务上的完整使用方案。你将掌握两步变体(先以 CE 损失训练基线,再对保存的 logits 做分布鲁棒后处理)与一步变体(端到端训练)两条训练路径的全部命令行参数,并从源码层面理解类级权重g_y、alpha_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 cvxpytensorflow_addons用于模型中的 Group Normalization 层(见 drops/preact_resnet_models.py);edward2被preact_resnet_models.py导入,用于概率层扩展;cvxpy在训练与测试时评估 DRO 精度时求解凸优化问题(见 drops/main_lt.py)。
run.sh中给出的验证命令特意使用了很小的num_iters 3与eval_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内部做了以下几件事:
- 构造长尾数据集:通过
get_cls_num(drops/main_lt.py)按指数衰减规则为每个类分配样本数——img_num_per_cls = img_max * (imb_factor ** (cls_idx / (cls_num - 1))),其中imb_ratio即img_min / img_max。随后get_cls_idx(drops/main_lt.py)按每类配额从原始数据中随机抽取索引,形成不平衡训练子集。 - 划分数据集:
get_dataset(drops/main_lt.py)使用train[:90%]作为训练集、train[90%:]作为验证集、test作为测试集,并施加 CIFAR 标准预处理(32×32 随机裁剪、随机水平翻转、按通道均值/方差归一化)。 - 构建模型:调用
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)。 - 训练:使用 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.h5与val/test_mean_best_{logits,labels}.txt - 最坏类精度最优 →
worst_best.h5与val/test_worst_best_{logits,labels}.txt - DRO 精度最优 →
dro_best.h5与val/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.9drops_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)。执行流程如下:
- 根据
model_dir、dataset、imb_ratio、run_id自动拼出基线模型目录,并固定加载_loss_ce_{run_id}的 CE 模型(drops/drops_test_time.py),即它只对 CE 训练的基线做后处理; - 通过
get_pre_train(drops/drops_test_time.py)读取val_{req}_best_logits.txt与test_{req}_best_logits.txt(req由--req指定,默认dro,可选mean/dro/worst,即用哪种指标选出的基线模型); - 以验证集 logits 为优化数据,迭代更新类权重
g_y、拉格朗日乘子lambd与后处理偏移alpha_y; - 用更新后的权重对测试 logits 施加偏移并评估平均精度、最坏类精度与 DRO 精度。
3.3 参数详解
main_lt.py与drops_test_time.py均使用absl.flags解析命令行参数。main_lt.py的完整参数如下(默认值取自源码定义,drops/main_lt.py):
| 参数 | 默认值 | 说明 |
|---|---|---|
--model_dir | ./ | 模型与日志存储目录,实际会追加{dataset}_{imb_ratio}/_loss_{loss}_{run_id}/子路径 |
--lr | 0.1 | 初始学习率 |
--num_iters | 50000 | 训练迭代次数 |
--eval_freq | 1000 | 评估频率 |
--batch_size | 128 | 批大小 |
--train_on_full | False | 是否在完整训练集(而非长尾子集)上训练 |
--run_id | 0 | 实验编号,用于区分同配置多次运行 |
--dataset | cifar10 | 数据集,支持cifar10/cifar100 |
--num_classes | 10 | 类别数(cifar100 时需改为 100) |
--imb_ratio | 0.1 | 不平衡因子n_min / n_max |
--loss | drops | 损失类型:ce/focal/bsm/cb/cb_focal/ldam/logit_adj/posthoc/posthoc_ce/drops |
--alpha | 1.0 | focal 损失的 alpha 参数 |
--gamma | 0.5 | focal / LDAM(C) 的 gamma 参数 |
--beta | 0.9999 | class-balanced 损失的 beta 参数 |
--s | 1 | LDAM 中 logits 的缩放参数 |
--re_weight_type | prior | 类别重加权方式:none/prior/sqrt |
--warmup | 0 | 热身迭代数,期间使用 CE 损失 |
--is_upsampling | False | 是否对尾类做上采样平衡 |
--dro_div | kl | DRO 度量使用的散度:kl/l2/l1(也支持reverse-kl) |
--eps | 0.9 | DRO 度量的扰动半径(散度约束上界 delta) |
--metric_base | uniform | DRO 度量的类权重基准:prior/recip_prior/uniform |
--eta_g | 0.01 | EG(指数梯度)更新g_y的步长 |
--eta_lambda | 0.01 | EG 更新lambd的步长 |
--n_it_update | 1 | g、lambda的更新频率(每多少步更新一次) |
--tau | 1.0 | logit adjustment 项的温度常数 |
--weight_type | ce | 计算类损失时使用的权重方式:0_1(0-1 损失)或ce(交叉熵) |
--g_type | not-eg | 是否使用 EG(指数梯度)风格更新g_y:eg或简化式not-eg |
drops_test_time.py的参数在main_lt.py基础上增加了三个后处理专属选项(drops/drops_test_time.py):
| 参数 | 默认值 | 说明 |
|---|---|---|
--prior_type | train | 类别先验来源:train(用训练集指数分布先验)或val |
--cal_type | none | 校准方式:temp(温度)/scale(缩放)/shift(偏移)/shift_scale/none |
--eta_lambda_mult | 1.0 | 每轮迭代对eta_lambda的衰减因子(README 示例取0.95,即逐步缩小拉格朗日乘子步长) |
--req | dro | 选择基线模型的标准:mean/dro/worst |
drops_test_time.py中eps的默认值为0.2、num_iters默认25、eta_lambda默认0.1;README 示例将其调整为eps 0.9、num_iters 25、eta_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.py每n_it_update(默认 1)步就会调用updateg在验证集上更新g_y、lambd与alpha_y,随后用更新后的alpha_y重建 DROPS 损失函数继续训练(drops/main_lt.py)。因此一步变体的训练开销高于两步变体——DRO 权重搜索发生在每一次参数更新之间,而不是只做一次。
五、源码级原理:DROPS 的优化机制
5.1 类权重g_y、alpha_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)":
- 计算每类损失
L_y:对验证集每个 batch,先取 softmax 预测,再用alpha_y加权得到"后偏移预测"y_weighted_pred;按weight_type分别用 CE 损失或 0-1 损失累计每个类别的平均损失loss_y_list; - 计算散度约束项:按
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); - 构造 Lagrangian:
L = Σ g_y·L_y - lambd·(D - delta); - 更新
g_y:g_type='eg'时按指数梯度(EG)闭式更新,否则用简化式g_y ← r_y·exp(L_y / lambd)后归一化; - 更新
lambd:lambd ← lambd + eta_lambda·cons; - 更新
alpha_y:alpha_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,即在允许的分布扰动范围内最坏可能的加权精度。在kl与reverse-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 | 实现函数 | 变换形式 |
|---|---|---|
temp | calibrate_temp | logits * temp(温度校准) |
shift | calibrate_shifts | logits - shifts |
scale | calibrate_scales | logits * scales |
shift_scale | calibrate_shifts_scales | logits * scales - shifts |
none | — | 不校准 |
校准参数(temp、shifts、scales)均在验证集 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 名称 | 实现类/函数 | 关键超参 |
|---|---|---|
ce | CELoss | — |
up_ce | CELoss(配合上采样) | --is_upsampling |
ldam | LDAMLoss | --gamma、--s |
focal | FocalLoss | --gamma |
cb | CBLoss | --beta |
cb_focal | CBFocal | --beta、--gamma |
bsm | BalancedSoftmax | — |
logit_adj | LogitAdjust | --tau |
posthoc | LogitAdjust(仅评估时偏移) | --tau |
posthoc_ce | CELoss+ 评估时 posthoc 偏移 | --tau |
drops | CELoss+ 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 优化动态求解——这正是它与posthoc、logit_adj的本质区别。
八、运行结果与日志解读
两类入口都会在model_dir下输出文本日志与 TensorBoard 摘要:
results_log.txt(main_lt.py)或results_loss=...,delta_train=...命名的日志文件(drops_test_time.py),逐条记录验证/测试损失、平均精度、最坏类精度以及各eps网格点上的 DRO 精度;summaries/train与summaries/eval目录供 TensorBoard 可视化(drops/main_lt.py);drops_test_time.py的日志末尾还会以机器可读的delta_evals=、DRO-Acc=、kldiv_cons=三段式输出,便于直接绘制 DRO 精度–扰动半径曲线(drops/drops_test_time.py)。
九、适用前提与注意事项
- 数据集支持:源码只实现了 CIFAR-10 / CIFAR-100 的合成长尾构造(
get_cls_num中按类别索引指数分配样本数);drops_test_time.py的--dataset注释虽提到 imagenet-lt,但当前仓库未包含对应数据加载实现,实际可用范围为 CIFAR 系列。 - 两步变体的耦合约束:
drops_test_time.py固定读取 CE 基线的 logits 文件,因此必须先以--loss 'ce'完整跑完main_lt并确保model/下存在val/test_{req}_best_{logits,labels}.txt,否则后处理会因文件缺失而失败。 num_classes需与数据集匹配:cifar100 时请将--num_classes 100与--dataset 'cifar100'一并设置,否则 logits 的形状解析会出错。- 实验定位:README 明确标注该目录为Experimental code,适用于复现与学术研究,生产环境部署前需结合自身数据做充分验证。
通过本文提供的两条命令路径与参数对照表,配合对updateg、eval_dro_metrics等核心函数的源码理解,你可以直接在 drops 目录中复现 DROPS 的两步/一步训练流程,并进一步修改eps、dro_div、eta_lambda等关键超参,探索分布鲁棒优化在长尾分类中的行为特性。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
CIW/CICW 标签噪声鲁棒学习实战指南:约束实例与类别重加权在 Google Research 开源代码中的应用
CIW/CICW 标签噪声鲁棒学习实战指南:约束实例与类别重加权在 Google Research 开源代码中的应用 导读 本文围绕 ciw_label_noi
人工智能深度学习NLP计算机视觉强化学习多语言支持详解:Plus Jakarta Sans如何实现越南语等复杂文字的完美呈现
多语言支持详解:Plus Jakarta Sans如何实现越南语等复杂文字的完美呈现 Plus Jakarta Sans是一款由Tokotype的Gumpita
人工智能深度学习NLP计算机视觉强化学习Google Authenticator深度解析:开源双因素认证技术实现
Google Authenticator深度解析:开源双因素认证技术实现 Google Authenticator作为谷歌开源的动态验证码生成器,基于OATH开
网络安全应用安全
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考