拓冰建站拓冰建站
首页 / 资讯中心 / 正文

Active Selective Prediction 实战指南:基于 ASPEST 的主动学习与选择性预测框架

Active Selective Prediction 实战指南基于 ASPEST 的主动学习与选择性预测框架【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research导读本文是 google-research 仓库中 active_selective_prediction 模块的完整技术指南。该模块是论文《ASPEST: Bridging the Gap Between Active Learning and Selective Prediction》的官方实现目标是在分布偏移dataset shift场景下把主动学习Active Learning与选择性预测Selective Prediction统一到同一套主动选择预测流水线中先用源域数据训练模型再在目标域上通过采样方法挑选少量样本标注并利用集成与自训练提升目标域性能。读完本文你将掌握六个基准数据集的构建方法、三类选择性预测方法SR / DE / ASPEST的原理、九种主动学习采样策略的取舍以及从训练到端到端评估的完整命令行操作。一、项目定位为什么要桥接主动学习与选择性预测选择性预测允许模型对不确定的样本弃权abstain主动学习则通过挑选最有价值的目标域样本进行标注来降低标注成本。传统研究通常将两者割裂对待而 ASPESTActive Selective Prediction using Ensembles and Self-Training将两者结合在目标域分布偏移的情况下利用集成的预测置信度来决定该标注哪些样本主动学习以及该预测哪些样本选择性预测。从源码结构看该模块采用清晰的五层组织methods/选择性预测方法SR、DE、ASPEST其中 aspest.py 定义了核心类ASPEST继承自 base_sp.py 中的SelectivePredictionMethodsampling_methods/九种主动学习采样策略models/自定义模型架构如 custom_model.pytfds_generators/非 TFDS 内置数据集的GeneratorBasedBuilder实现utils/数据、模型、TensorFlow 通用工具。二、环境要求与依赖安装README 明确给出验证过的运行环境操作系统Debian 4.19.260-1 (2022-09-29) x86_64 GNU/LinuxPython 版本3.7.12。安装依赖只需一条命令pip install -r active_selective_prediction/requirements.txt从 requirements.txt 可以看到关键依赖及其版本下限依赖版本要求用途tensorflow / keras2.10.0模型训练与推理主框架tensorflow-datasets4.7.0内置数据集MNIST、SVHN、CIFAR-10、DomainNet加载transformers4.24.0amazon_review 数据集 RoBERTa 特征提取wilds2.0.0FMoW、Amazon Review 数据获取scikit-learn / scipy / numpy见文件采样方法与距离计算pandas / tqdm / tensorboard见文件数据处理与训练监控此外项目还依赖keras-nightly与tf-estimator-nightly系列版本说明该实现针对 TensorFlow 2.11 时代的 API 编写升级大版本时需注意兼容性。三、六个基准数据集与构建流程项目创建了六个分布偏移基准数据集均以源域训练 → 目标域测试的形式组织基准源域 → 目标域数据来源mnist-svhnMNIST → SVHNTFDS 内置cifar10-cinic10CIFAR-10 → CINIC-10CINIC-10 手动下载fmowFMoW 分布内 → FMoW-OODWILDS 下载amazon_review分布内 → 分布外子集WILDS 下载domainnetreal → painting/clipart/infograph/sketchTFDS 内置otto训练/验证 → 测试Kaggle 竞赛数据其中 TFDS 已内置mnist、svhn、cifar10、domainnet可直接加载其余数据集需手动下载cinic10从 Edinburgh DataShare 下载后放入$RAW_DATASET_DIR/cinic10fmow从 WILDS 获取后放入$RAW_DATASET_DIR/wilds_data/fmow_v1.1amazon_review从 WILDS 获取后放入$RAW_DATASET_DIR/wilds_data/amazon_v2.1otto从 Kaggle 的 Otto Group Product Classification 竞赛下载后放入$RAW_DATASET_DIR/otto-group-product-classification。3.1 构建 TensorFlow 数据集执行一条命令即可将全部原始数据转换为 TFDS 数据集./active_selective_prediction/build_datasets.sh $RAW_DATASET_DIR $DATA_DIR其中$RAW_DATASET_DIR存放原始数据$DATA_DIR存放生成后的 TFDS 数据集。从 build_datasets.sh 的源码可以看到脚本内部会依次对mnist、cifar10、domainnet、fmow、amazon_review、otto六个数据集调用python -m active_selective_prediction.build_datasets --gpu 0 --dataset $dataset --data-dir $data_dir --raw-dataset-dir $raw_dataset_dir即真正的构建逻辑在 build_datasets.py 中脚本只是循环调用的封装。生成的数据集通过 utils/data_util.py 中的函数以tf.data.Dataset形式加载。注意该文件顶部定义了全局变量DATA_DIR ~/tensorflow_datasets/data_util.py所有tfds.builder(..., data_dirDATA_DIR)调用都依赖它因此必须把DATA_DIR指向你实际的$DATA_DIR否则加载会失败。3.2 特殊说明amazon_review 需要 GPU对于amazon_review数据集构建过程需要 GPU 来用 RoBERTa 提取文本嵌入对应 tfds_generators/tfds_amazon_review.py 的生成逻辑因此build_datasets.sh中以--gpu 0调用。3.3 向代码库新增数据集的完整流程README 给出了可操作的五步扩展指南若新数据集不在 TFDS 目录中实现一个tfds.core.GeneratorBasedBuilder子类放入./tfds_generators目录仓库中已有 tfds_cinic10.py、tfds_fmow.py、tfds_amazon_review.py、tfds_otto.py 四个范例在 build_datasets.py 中为新数据集添加 builder 函数修改 build_datasets.sh 并运行构建在 utils/data_util.py 中添加以tf.data.Dataset加载新数据集的函数若新数据集需要新的模型架构在 models/custom_model.py 中实现并在 utils/model_util.py 添加加载函数同时在 train.py 添加源域训练代码在 eval_model.py 与 eval_pipeline.py 添加评估代码修改./configs/下的配置文件以纳入新数据集。四、选择性预测方法SR / DE / ASPEST4.1 SRSoftmax Response最简单直接的置信度方法直接用模型输出的 softmax 最大概率作为预测置信度低于阈值即弃权。实现位于 methods/sr.py。4.2 DEDeep Ensembles深度集成训练多个模型默认num_models 5组成集成以集成后的预测一致性/置信度作为选择性预测依据通常比单个模型更可靠。实现位于 methods/de.py。4.3 ASPEST本文提出的方法ASPEST 全称 Active Selective Prediction using Ensembles and Self-Training基于集成与自训练的主动选择性预测实现位于 methods/aspest.py。从源码结构看ASPEST类aspest.py继承SelectivePredictionMethod关键设计包括强制使用平均 margin 采样构造函数将sampling_method硬编码为average_marginaspest.py即基于集成输出的平均 margin 挑选样本软集成soft ensembleself.ensemble_method soft配合reset_ensemble_state()维护ensemble_model_outputs与counts状态aspest.py两阶段训练train_init_model先在源域上以finetune_kwargs[init_steps]步默认 1000见配置文件初始化模型aspest.py随后进入采样 → 微调 → 自训练的迭代循环。其中AverageMarginSamplingaverage_margin_sampling.py的评分逻辑为对集成输出的每行排序后取最大与次大概率之差作为 scoresorted_outputs[:, -1] - sorted_outputs[:, -2]分数越低代表模型越不确定越值得被采样标注。五、主动学习采样方法一览README 列出了当前支持的九种采样方法全部实现在 sampling_methods/ 下统一继承 base_sampler.py 中的SamplingMethod方法策略思想实现文件UniformSampling均匀随机采样基线uniform_sampling.pyConfidenceSampling基于模型置信度confidence_sampling.pyEntropySampling基于模型熵entropy_sampling.pyMarginSampling基于预测 marginmargin_sampling.pyKCenterGreedySamplingK-center 贪心覆盖性kcenter_greedy_sampling.pyCLUESamplingCLUE 采样clue_sampling.pyBADGESamplingBADGE 采样badge_sampling.pyAverageKLDivergenceSampling平均 KL 散度average_kl_divergence_sampling.pyAverageMarginSampling平均 marginASPEST 专用average_margin_sampling.py前四个为经典的基于不确定性/多样性的策略K-center、CLUE、BADGE 属于面向批量选择的代表性方法后两个Average系列专为集成式方法设计。可通过eval_pipeline.py的--method与配置文件中的sampling_method字段自由组合。六、训练在源域上训练标准监督模型6.1 命令格式python -m active_selective_prediction.train --gpu $gpu --dataset $dataset6.2 参数详解来自 train.py参数默认值说明--gpu0使用的 GPU 编号写入CUDA_VISIBLE_DEVICES--seed100固定随机种子--datasetcolor_mnist数据集可选cifar10、domainnet、color_mnist、fmow、amazon_review、otto--save-dir./checkpoints/standard_supervised/模型权重保存目录最终权重写入{save_dir}/{dataset}/checkpoint注意color_mnist是项目中一个轻量级调试数据集基于 MNIST 灰度转 RGB 并 padding 到 32×32见 data_util.py默认配置用它作为快速验证目标。6.3 各数据集的训练配置源码事实从 train.py 可以提取每个数据集的训练细节数据集模型架构优化器学习率epochs类别数color_mnistsimple_convnetAdam1e-32010cifar10cifar_resnetSGD(momentum0.9)1e-120010domainnetresnet50imagenet 预训练Adam1e-450345fmowdensenet121imagenet 预训练Adam1e-45062amazon_reviewroberta_mlpAdam1e-32005ottosimple_mlpAdam1e-32009其中 cifar10 还配有分段学习率调度epoch 80/120/160 处衰减 0.1epoch 180 处乘 0.5train.pybatch size 统一为 128验证集 128 或 200。训练完成后调用model.save_weights()将权重保存为 checkpointtrain.py。七、评估源模型精度与主动选择性预测流水线7.1 评估源模型在源验证集与目标测试集上的精度python -m active_selective_prediction.eval_model --gpu $gpu --source-dataset $source --model-path $path其中$path指向存放源训练模型 checkpoint 的目录例如color_mnist的默认路径为./checkpoints/standard_supervised/color_mnist。从 eval_model.py 源码可以看到load_pretrained_model根据源数据集选择对应架构simple_convnet / cifar_resnet / resnet50 / densenet121 / roberta_mlp / simple_mlp用源域 batch 做一次前向传播创建变量后load_weights(...).expect_partial()载入权重eval_model.py评估时在源验证集上计算Source accuracy再对每个目标数据集如 SVHN、CINIC-10、DomainNet 的 painting/clipart/infograph/sketch、FMoW-OOD、Amazon-Review-OOD、otto-test计算Target accuracyeval_model.py。7.2 评估主动选择性预测方法python -m active_selective_prediction.eval_pipeline --gpu $gpu --source-dataset $source --method $method --method-config-file $config$config指向 configs/ 目录下的方法配置文件仓库为 SR、DE、ASPEST 各提供了默认配置sr.json、de.json、aspest.json。--method参数与配置文件中的name字段sr/de/aspest对应。7.3 配置文件字段解析三个配置文件的结构基本一致字段含义如下以 aspest.json 为例字段SR 默认值DE/ASPEST 默认值含义model_path./checkpoints/standard_supervised/{dataset}同左源模型 checkpoint 目录namesr无方法名仅 sr.json 含num_models无5集成模型数量DE/ASPEST 使用label_budget100500主动学习总标注预算domainnet 在 de.json 中为 100batch_size128128训练 batch 大小sampling_rounds1010采样轮数max_epochs200200微调最大 epoch 数patience_epochs1010早停 patiencemin_epochs5050最小 epoch 数optimizer_nameAdam/SGD同左优化器cifar10 为 SGDoptimizer_kargs如learning_rate、momentum同左优化器参数sampling_methodmargin无ASPEST 硬编码 average_marginde.json 为margin采样策略sampling_kwargs{}{}采样附加参数self_train_kwargs无pseudo_train_epochs20、pseudo_ckpt_epoch5、frac0.1、use_checkpoint_ensembletrue、lower_threshold0.9、upper_threshold1.0ASPEST 自训练参数finetune_methodjoint_trainjoint_train微调方式finetune_kwargslambda1.0de 另含init_steps1000ckpt_epoch5、lambda1.0、init_steps1000、init_ckpt_step200微调参数debug_infofalsefalse是否输出调试信息print_freq100sr500打印频率其中 ASPEST 的self_train_kwargs语义结合 aspest.py 结构可推断pseudo_train_epochs自训练伪标签阶段训练的 epoch 数pseudo_ckpt_epoch自训练阶段保存 checkpoint 的间隔 epochfrac从伪标签样本中选取的比例use_checkpoint_ensemble是否用训练过程中的 checkpoint 组成集成lower_threshold/upper_threshold伪标签置信度阈值区间仅信任落在区间内的伪标签样本。八、端到端运行示例MNIST → SVHN仓库提供了现成脚本 run_mnist_to_svhn_exp.sh一键跑通完整实验链路./active_selective_prediction/run_mnist_to_svhn_exp.sh从脚本内容看它依次完成创建虚拟环境并安装依赖virtualenv -p python3 .pip install -r active_selective_prediction/requirements.txt在 MNIST 上训练标准模型python -m active_selective_prediction.train --gpu 0 --dataset color_mnist注意脚本实际以color_mnist作为源任务其目标域正是 SVHN评估源模型在 MNIST 测试集与 SVHN 测试集上的精度python -m active_selective_prediction.eval_model --gpu 0 --source-dataset color_mnist --model-path ./checkpoints/standard_supervised/color_mnist依次评估三种方法sr、de、aspest分别读取configs/sr.json、configs/de.json、configs/aspest.json。该脚本以set -e与set -x开启错误即停与执行回显便于排错也演示了本项目训练 → 模型评估 → 方法评估的标准工作流。README 同时说明仓库还提供通用脚本 run.sh 供参考使用。九、小结与扩展建议ASPEST 模块为在分布偏移下以最少标注获得可靠预测提供了完整可复现的工程实现六个跨视觉、NLP、表格数据域的基准三个选择性预测方法其中 ASPEST 以软集成 平均 margin 采样 自训练为核心九个可插拔的采样策略以及配置驱动的一键训练评估流水线。若要在自己的数据上使用建议路径为先按第三节的步骤把新数据接入 TFDS 生成器与data_util再仿照 configs/aspest.json 添加配置项最后分别运行train.py训练源模型、eval_pipeline.py对比 SR/DE/ASPEST 在目标域上的精度与覆盖度coverage表现。由于配置中label_budget100/500与sampling_rounds10共同决定了标注成本实际使用时应结合自身预算调整这两个核心超参数。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门