ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

non_decomp 实战指南:用 logit-adjusted 成本敏感损失训练过参数化模型,优化不可分解目标(NeurIPS 2021 代码复现)

non_decomp 实战指南:用 logit-adjusted 成本敏感损失训练过参数化模型,优化不可分解目标(NeurIPS 2021 代码复现) non_decomp 实战指南用 logit-adjusted 成本敏感损失训练过参数化模型优化不可分解目标NeurIPS 2021 代码复现【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research本指南基于 non_decomp/README.md 及同目录源码系统讲解如何复现 NeurIPS 2021 论文《Training Over-parameterized Models with Non-decomposable Objectives》所提出的最坏类别召回率worst-case recall最大化方法论文 Algorithm 1。读者将掌握四种训练模式ERM / balanced / reweighted / proposed的含义与命令行用法、CIFAR-10/100 长尾数据集的获取与放置、核心参数tau、eg_lr、update_freq 等的调优思路以及基于tensorflow_constrained_optimizationTFCO的指数梯度类权重更新机制的底层原理。一、背景不可分解目标与长尾学习在类别不平衡long-tail场景下常见的全局指标如总体准确率是可分解decomposable的——它等于逐样本损失的均值。而像最坏类别召回率minimum per-class recall这类指标则不可分解它取决于每一类的召回率中的最小值无法写成逐样本损失求和的简单形式。直接优化这类目标在过参数化模型上十分困难。本仓库正是针对这一问题的官方实现对应论文Training Over-parameterized Models with Non-decomposable ObjectivesHarikrishna Narasimhan, Aditya Krishna MenonNeurIPS 2021。该仓库从同项目下的 logit_adjustmentMenon 等人 ICLR 2021 的 Logit Adjustment 长尾学习代码分支而来。核心思路是通过一个带有 logit 调整的成本敏感损失logit-adjusted cost-sensitive loss配合同步更新的类别权重在训练过程中逐步逼近最坏类别召回率最大化目标。二、仓库结构一览non_decomp/ ├── README.md # 官方使用说明本文主体 ├── main.py # 训练/评估主程序absl flags 入口 ├── models.py # CIFAR ResNet-32 模型定义 ├── utils.py # 数据集、损失函数、指标、学习率调度等工具 ├── requirements.txt # Python 依赖 ├── run.sh # 一键虚拟环境 冒烟测试脚本 └── data/ # base_probs 文件与冒烟测试数据 ├── cifar10-lt_base_probs.txt ├── cifar100-lt_base_probs.txt └── test.tfrecord其中data/目录预置了两份关键的基类概率base probabilities文件详见第四节以及一个用于快速验证的哑数据集test.tfrecord。三、环境搭建与依赖官方声明代码在Python 3.7.2下测试通过。在仓库根目录执行# from google-research/ pip install -r non_decomp/requirements.txtnon_decomp/requirements.txt 中的依赖为numpy1.16.0tensorflow2.4.0tensorflow_constrained_optimizationTFCO用于计算各类别的假阴性率并驱动约束优化仓库还提供了 non_decomp/run.sh 一键脚本它使用virtualenv -p python3 .创建本地虚拟环境、激活后安装依赖并自动执行一次冒烟测试--datasettest --modeerm --train_batch_size2 --test_batch_size2。四、快速冒烟测试安装依赖后先用内置哑数据集验证代码链路是否畅通# from google-research/ python -m non_decomp.main --datasettest --modebaseline --train_batch_size2 --test_batch_size2该命令应能快速无错完成。注意--datasettest对应 utils.py 中的Dataset(test, 10, test.tfrecord, test.tfrecord, 4, 4, 2, ...)即 4 个训练样本、4 个测试样本、共训练 2 个 epoch此模式下 README 使用了--modebaseline而在 main.py 中mode的合法取值为[erm, balanced, reweighted, proposed]因此更准确的做法是--modeermrun.sh 中正是如此。五、数据准备CIFAR-10/100 长尾数据集训练真实模型需要下载 4 个.tfrecord文件下载链接见 README托管于storage.googleapis.com的gresearch/logit_adjustment/路径下数据集训练文件测试文件CIFAR-10 长尾cifar10-lt_train.tfrecordcifar10_test.tfrecordCIFAR-100 长尾cifar100-lt_train.tfrecordcifar100_test.tfrecord要点训练集按论文中的EXP-100 剖面指数衰减的类别频率构造测试集即标准 CIFAR-10/100 测试集。下载后将文件放入data_home指定的目录。README 沿用了上游 logit_adjustment 的说明放入logit_adjustment/data/但本仓库 main.py 中data_home的默认值为non_decomp/data因此实际运行时应将训练/测试.tfrecord放入 non_decomp/data/或通过--data_home显式指定目录。仓库已预置两份base_probs.txt与数据集配套使用non_decomp/data/cifar10-lt_base_probs.txt10 行概率从约 0.403 指数衰减到约 0.004non_decomp/data/cifar100-lt_base_probs.txt100 行概率。这些文件对所有非 ERM 模式是必需的main.py 会尝试读取{data_home}/{dataset}_base_probs.txt若文件缺失且mode ! erm会抛出app.UsageError明确提示该文件必须存在只有erm模式允许缺失此时base_probs None。六、运行四种训练模式在仓库根目录对 CIFAR-10 长尾数据集执行# from google-research/ # 1) ERM 基线标准经验风险最小化 python -m non_decomp.main --datasetcifar10-lt --modeerm # 2) balanced 基线Menon 等 2021 的 logit-adjusted 损失即balanced版本 python -m non_decomp.main --datasetcifar10-lt --modebalanced # 3) reweighted 基线重加权成本敏感损失 python -m non_decomp.main --datasetcifar10-lt --modereweighted # 4) proposed论文提出的 logit-adjusted 成本敏感方法Algorithm 1 python -m non_decomp.main --datasetcifar10-lt --modeproposed将cifar10-lt替换为cifar100-lt即可复现 CIFAR-100 长尾实验。四种模式的损失函数映射关系在 main.py 中一目了然erm→loss_type standardreweighted→loss_type reweightedbalanced与proposed→loss_type logit_adjusted也就是说balanced和proposed使用同一个损失函数二者的差异在于proposed额外维护一组可学习的类别权重见第七节。七、核心机制源码级解析7.1 三类损失函数的统一实现utils.py 的build_loss_fn把三种损失统一为外层权重 × 逐样本交叉熵的框架standardERMouter_weights全为 1即普通交叉熵reweightedscaled_class_weights 1 / base_probs^tau作为逐样本外层权重logit_adjusted同样计算scaled_class_weights但不做加权而是从 logits 中减去log(scaled_class_weights)等价于在 softmax 前给每个类别加上先验补偿项。其中tau是对基类概率做温度缩放temperature scaling的超参数默认1.0控制调整强度实现中加入了1e-12的微小扰动以避免除零。最终损失为sum(loss * per_sample_weights) / sum(per_sample_weights)的加权平均形式。7.2 proposed 模式指数梯度Exponentiated Gradient更新类别权重proposed模式的核心循环在 main.py初始化class_weights为均匀分布1/num_classes一个可训练的tf.Variable每update_freq默认 32个梯度步取一个验证批次用 TFCO 库计算每个类别的假阴性率FNR 1 - recallexp_class_weights class_weights * tf.math.exp(FLAGS.eg_lr * fnrs.result()) class_weights.assign(exp_class_weights / tf.reduce_sum(exp_class_weights))即把假阴性率高的类别权重按eg_lr默认 0.1指数放大后归一化——这正是指数梯度更新名称的来源。计算 FNR 依赖 utils.py 的FalseNegativeRates类它借助 TFCO 的tfco.multiclass_rate_context构造了num_classes个假阴性率约束通过problem.update_ops()与problem.constraints()读取每类 FNR。这就是tensorflow_constrained_optimization在仓库中的实际用途。7.3 模型与优化器模型models.py 中的cifar_resnet32(num_classes)输入(32, 32, 3)配置[(5,16,1), (5,32,2), (5,64,2)]即 3 个阶段 × 5 个 block通道数 16/32/64后两阶段步长 2共 32 层卷积层与全连接层均施加 L2 正则权重衰减1e-4BatchNorm 动量 0.9、epsilon 1e-5。优化器main.py 使用 SGD momentum 0.9 Nesterov基学习率 0.1学习率调度utils.py 的LearningRateSchedule按 epoch 分段衰减先线性 warmup再按倍率逐步降低各数据集的调度表定义在 utils.py 的dataset_mappings()中。7.4 数据集定义与验证/测试划分dataset_mappings()定义了三个数据集dataset类别数训练样本测试样本epoch 数cifar10-lt1012406100001241cifar100-lt10010847100001419test10442由 utils.py 的create_tf_dataset可知测试文件的前 5000 条作为 test 集其余部分划作验证集vali验证集被repeat()无限重复用于训练过程中的间隔式类权重更新。训练数据增强包括 4 像素填充后随机裁剪与随机水平翻转见_process_image。八、全部命令行参数速查main.py 中通过 absl flags 定义的参数如下参数默认值说明--datasetcifar10-lt数据集cifar10-lt/cifar100-lt/test--data_homenon_decomp/data.tfrecord与base_probs.txt所在目录--train_batch_size128训练批大小--vali_batch_size4096验证批大小--test_batch_size100测试批大小--modeproposederm/balanced/reweighted/proposed--tau1.0logit 调整中基类概率的温度缩放参数--eg_lr0.1类权重指数梯度更新的学习率--update_freq32每多少个梯度步更新一次类权重--tb_log_dirnon_decomp/logTensorBoard 日志输出目录调参建议基于源码语义推断tau控制先验补偿强度balanced模式可直接调整它proposed模式下eg_lr决定类权重更新步长、update_freq决定更新频率二者共同影响收敛稳定性——权重更新过快可能震荡过慢则拖慢对尾部类别的适应。九、训练进度监控与日志每次运行都会向控制台打印训练损失、验证最坏召回率等日志当处于proposed模式时还会打印每类 FNR 与类权重数值。最终测试准确率也会出现在日志中。可用 TensorBoard 实时监控# from google-research/ tensorboard --logdir./non_decomp/logTensorBoard 摘要写入{tb_log_dir}/train与{tb_log_dir}/test两个子目录main.py记录批量损失、最坏召回率、训练/验证/测试指标等标量。注意每次重新训练前建议删除旧的non_decomp/log目录避免不同实验的曲线混在一起。十、常见问题与注意事项base_probs缺失报错非 ERM 模式要求non_decomp/data/cifar10-lt_base_probs.txt或cifar100-lt存在缺失时main.py会直接以UsageError终止仓库已预置这两份文件。数据集放置位置以--data_home实际生效目录为准默认non_decomp/dataREADME 中提及的logit_adjustment/data/是上游分支遗留说明。TensorBoard 目录污染多次运行间清空non_decomp/log避免曲线串扰。冒烟测试模式名README 示例中的--modebaseline对应本仓库的--modeerm二者等价。依赖版本代码基于 TF 2.x TFCO 编写升级或降级 TensorFlow 版本可能导致 API 兼容性问题。若运行中遇到问题可联系 README 中给出的作者邮箱hnarasimhan {at} google.com。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表