ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

DeepMind 对抗鲁棒性模型评估指南:基于 Data Augmentation Can Improve Robustness 的预训练模型在 JAX 与 PyTorch 中的复现与评测

DeepMind 对抗鲁棒性模型评估指南:基于 Data Augmentation Can Improve Robustness 的预训练模型在 JAX 与 PyTorch 中的复现与评测 人工智能深度学习机器学习计算机视觉NLP强化学习【免费下载链接】deepmind-researchThis repository contains implementations and illustrative code to accompany DeepMind publications项目地址https://gitcode.com/gh_mirrors/de/deepmind-research点击查看免费下载本指南聚焦adversarial_robustness/iclrw2021data目录所配套的 ICLR 2021 论文Data Augmentation Can Improve RobustnessRebuffi et al., 2021介绍其发布的五组高精度对抗鲁棒模型CIFAR-10/CIFAR-100WideResNet 架构的下载方式、评估命令与源码实现细节。读完本文你将能够独立下载预训练权重在 JAX 或 PyTorch 环境中复现干净准确率clean accuracy与对抗鲁棒准确率robust accuracy的评测流程并理解评估脚本背后 PGD 攻击与模型归一化的具体实现。背景数据增强如何改善对抗鲁棒性Data Augmentation Can Improve Robustness是 DeepMind 于 ICLR 2021 Security and Safety in Machine Learning Systems WorkshopAML-ICLR2021发表的工作核心结论是在不借助外部数据的前提下通过精心设计的数据增强策略如 CutMix并配合对抗训练可以在标准数据集上获得显著的对抗鲁棒性提升。与其姊妹篇Fixing Data Augmentation to Improve Adversarial RobustnessRebuffi et al., 2021一脉相承。adversarial_robustness/iclrw2021data/README.md正是这篇论文配套的模型发布说明仓库中提供了与 JAX 和 PyTorch 两种框架兼容的顶级模型权重以及对应的模型定义源码方便研究者直接评估或二次开发。仓库结构与配套资源围绕该文档仓库adversarial_robustness/目录提供了完整的评估与训练体系路径作用jax/JAX 生态Haiku Optax下的模型定义、数据集处理、攻击实现与评估/训练入口pytorch/PyTorch 生态下的模型定义WideResNet 与 PreActResNet与评估入口run.sh一键创建虚拟环境、安装依赖并跑通测试程序requirements.txt锁定版本的依赖清单其中 JAX 目录下各文件职责清晰eval.py负责加载权重并评测attacks.py提供完整的对抗攻击组件库model_zoo.py给出 WideResNet 的 Haiku 实现datasets.py处理 CIFAR/MNIST 的加载与归一化experiment.py与train.py则承载复现论文的训练流程。预训练模型清单与性能基准文档给出了论文发布的五组模型涵盖 CIFAR-10 与 CIFAR-100 两个数据集、ℓ∞范数下 8/255 扰动半径的对抗训练设置。干净准确率clean在完整测试集上测得鲁棒准确率robust使用 AutoAttack 评测datasetnormradiusarchitectureextra datacleanrobustlinkCIFAR-10ℓ∞8 / 255WRN-70-16✓92.23%66.58%jax, ptCIFAR-10ℓ∞8 / 255WRN-70-16✗87.25%60.07%jax, ptCIFAR-10ℓ∞8 / 255WRN-28-10✗86.09%57.61%jax, ptCIFAR-100ℓ∞8 / 255WRN-70-16✗65.76%32.43%jax, ptCIFAR-100ℓ∞8 / 255WRN-28-10✗62.97%29.80%jax, pt两点关键观察额外数据extra data标记为 ✓ 的仅有一组cifar10_linf_wrn70-16_cutmix_external其 CIFAR-10 干净准确率 92.23%、鲁棒准确率 66.58%是五组中鲁棒性最高的其余模型均不依赖额外数据验证了论文纯数据增强即可改善鲁棒性的主张。文件名中的cutmix表明训练阶段使用了 CutMix 数据增强与论文核心方法对应。每个模型均提供两种格式.npyJAX/Haiku 参数可直接被np.load解析与.ptPyTorchtorch.load加载的 state_dict按需下载即可。使用预训练模型评估命令与参数详解模型下载完成后在jax或pytorch目录下运行eval.py即可完成评测。文档给出的标准命令如下cd jax python3 eval.py \ --ckpt${PATH_TO_CHECKPOINT} --depth70 --width16 --datasetcifar10JAX 版本参数说明对照 jax/eval.py 的 flag 定义各参数含义与取值范围如下参数默认值说明--ckpt必填权重文件路径传入dummy可随机初始化参数用于冒烟测试--datasetcifar10枚举值cifar10/cifar100/mnist--width16WideResNet 宽度因子宽度乘数--depth70WideResNet 深度源码要求满足depth 6n 4否则抛出ValueError--batch_size100评估批大小--num_batches0评估批数上限0表示跑完整测试集10000 张评估流程的源码级拆解从 jax/eval.py 可以看到完整的评测管线加载数据集根据--dataset选择 CIFAR-10/100 或 MNIST 测试集图像归一化到[0,1]浮点区间images.astype(np.float32) / 255.。构建模型使用model_zoo.WideResNet(num_classes10, depth, width, activationswish)注意官方模型激活函数为Swish而非默认 ReLU输入先经过datasets.py中的均值/方差归一化CIFAR-10 均值为(0.4914, 0.4822, 0.4465)标准差(0.2471, 0.2435, 0.2616)见 jax/datasets.py。加载权重--ckptdummy时用随机初始化参数否则np.load(ckpt, allow_pickleTrue)直接解析.npy。构造攻击评测鲁棒准确率时运行PGD-40 Adam 优化器 margin loss的不定目标攻击扰动半径epsilon 8/255优化器为 Adam学习率采用分段常数调度初始0.1第 20 步后衰减0.1倍第 30 步后再衰减0.01倍linf_initialize_fn在[-epsilon, epsilon]内均匀采样初始化扰动linf_project_fn将对抗样本裁剪回[0,1]像素范围见 jax/attacks.py。输出指标脚本最终打印Accuracy on the N test images: xx.xx%干净准确率与Robust accuracy: xx.xx%对抗鲁棒准确率。PyTorch 版本差异pytorch/eval.py 的参数大体一致但有两点区别--width0时自动切换到 PreActResNet预激活残差网络此时--depth仅支持18或34由 pytorch/model_zoo.py 中的分支决定块数量新增--use_cuda默认True控制是否使用 GPUMNIST 场景下会自动设置mean.5, std.5, padding2, num_input_channels1。PyTorch 版eval.py只计算干净准确率不内置攻击若要评测鲁棒准确率可配合 RobustBench 模型库该系列模型已收录其中或自行叠加攻击脚本。环境搭建与依赖说明评估代码的运行环境由 run.sh 一键搭建它在/tmp下创建名为adversarial_robustness_venv的 virtualenv安装 requirements.txt 中锁定的依赖并以小规模冒烟测试--ckptdummy --width1 --depth10 --batch_size1 --num_batches1验证 JAX 与 PyTorch 两套评估入口均可正常导入与运行。激活环境使用source /tmp/adversarial_robustness_venv/bin/activate依赖清单中的关键版本组合均为论文发布时的锁定版本包括jax0.2.16、jaxlib0.1.68、dm-haiku0.0.4、optax0.0.8、jaxline0.0.3、tensorflow2.5.0、torch1.9.0、torchvision0.10.0、numpy1.19.5、absl-py0.12.0。若需 GPU 支持应在执行run.sh前自行调整 JAX 相关依赖例如将jaxline替换为带cuda后缀的构建版本并遵循 JAX 官方的安装指引。需要强调的是这些版本号较老在新环境尤其是新版本 CUDA/Python中直接复现时可能需要适配。模型定义细节从源码看 WideResNet 结构为理解权重与架构的对应关系可对照两份模型定义JAX 版jax/model_zoo.pyWideResNet以(depth - 4) % 6 ! 0校验深度合法性每层包含(depth - 4) // 6个基本块三个阶段的滤波器数量分别为16×width、32×width、64×width每个块由 BatchNorm → 激活 → 3×3 卷积组成首块带 1×1 投影捷径全局使用 3×3 初始卷积与均值池化收尾。例如WRN-70-16即depth70, width16。PyTorch 版pytorch/model_zoo.py结构对齐 JAX 版其中 Swish 激活以自定义 autograd Function 实现见_Swish卷积手动补齐 padding 以等价 TensorFlow 的SAME语义并在前向中内建了(x - mean) / std的归一化。两版模型定义在评估时必须与--depth/--width严格匹配否则权重加载或网络结构将不兼容。引用与免责声明若在研究中使用了本文档对应的代码、数据或模型请按论文引用规范引用Data Augmentation Can Improve Robustness的完整版本即结合生成样本的版本article{rebuffi2021fixing, title{Fixing Data Augmentation to Improve Adversarial Robustness}, author{Rebuffi, Sylvestre-Alvise and Gowal, Sven and Calian, Dan A. and Stimberg, Florian and Wiles, Olivia and Mann, Timothy}, journal{arXiv preprint arXiv:2103.01946}, year{2021}, url{https://arxiv.org/pdf/2103.01946} }需要说明的是该代码与模型由 DeepMind 团队发布并非 Google 官方产品见文档末尾 Disclaimer。本指南所描述的下载链接、评估命令与参数均以当前仓库adversarial_robustness/目录下的文档与源码为准若你希望深入训练侧复现可进一步研读 jax/experiment.py 与 jax/train.py 中基于 Jaxline 的训练管线并结合 adversarial_robustness/README.md 中关于生成数据集DDPM 生成的 1M 样本npz与更多模型WRN-106-16、ResNet-18 等的补充说明进行整体把握。赞分享人工智能深度学习机器学习计算机视觉NLP强化学习【免费下载链接】deepmind-researchThis repository contains implementations and illustrative code to accompany DeepMind publications项目地址https://gitcode.com/gh_mirrors/de/deepmind-research点击查看免费下载相关推荐deepmind-research 对抗鲁棒性 PyTorch 评估与训练复现指南deepmind research 对抗鲁棒性 PyTorch 评估与训练复现指南 本指南围绕 adversarial_robustness https://l人工智能深度学习机器学习计算机视觉NLP强化学习Rust程序设计语言中文版如何构建并发安全的Rust应用程序Rust程序设计语言中文版如何构建并发安全的Rust应用程序 Rust程序设计语言中文版提供了强大的并发安全特性帮助开发者构建高效且安全的多线程应用程序。通上一篇Aspire 仓库 PR 代码审查技能code-review全指南从分支准备到问题上报的完整工作流下一篇如何使用Engauge Digitizer快速完成图表数字化从图像到数据的完整指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表