ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

bnn_hmc 复现指南:用全批量 Hamiltonian Monte Carlo 探究贝叶斯神经网络后验的真实形态

bnn_hmc 复现指南:用全批量 Hamiltonian Monte Carlo 探究贝叶斯神经网络后验的真实形态 bnn_hmc 复现指南用全批量 Hamiltonian Monte Carlo 探究贝叶斯神经网络后验的真实形态【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research本文以 Google Research 仓库中的 bnn_hmc/README.md 为骨架结合其 JAX 源码完整讲解如何复现论文What Are Bayesian Neural Network Posteriors Really Like?arXiv:2104.14421作者 Pavel Izmailov、Sharad Vikram、Matthew D. Hoffman 与 Andrew Gordon Wilson中的全部实验。读完本文你将掌握bnn_hmc项目的环境搭建、HMC / SGD / Deep Ensembles / SGMCMC / MFVI 五类训练脚本的完整命令行用法、关键超参数的含义与推荐取值以及后验密度可视化工具的使用方式并能对照源码理解每一步计算的底层实现。一、项目背景论文的核心问题与主要发现该项目使用**全批量full-batchHamiltonian Monte CarloHMC**作为近乎精确的贝叶斯推理工具去回答贝叶斯深度学习中的一系列基础问题贝叶斯神经网络BNN的后验到底是什么样的它能否带来实际性能收益常用的近似方法深度集成、SGMCMC、变分推断与真实后验有多接近论文在实验中得到的主要结论如下也是本项目所有脚本设计意图的集中体现BNN 相比标准训练和深度集成可以获得显著的性能提升单条长 HMC 链即可提供与多条短链相当的后验表示即长链采样效率优于多条短链并联与近期一些研究相反论文发现后验降温posterior tempering并非达到近似最优性能所必需几乎没有“冷后验cold posterior”效应的证据——论文进一步指出该效应在很大程度上是数据增强data augmentation的产物贝叶斯模型平均BMA的性能对先验尺度prior scale的选择比较鲁棒并且对角高斯先验、高斯混合先验与逻辑斯蒂先验下的结果较为接近BNN 在分布漂移domain shift下泛化表现令人意外地差深度集成与 SGMCMC 等更廉价的替代方案虽然能取得不错的泛化但它们给出的预测分布与 HMC 明显不同值得注意的是深度集成的预测分布与 HMC 的接近程度与标准 SGLD 相当且明显优于标准变分推断。仓库以 JAX 实现了上述全部实验代码并附带 bibtex 引用信息见 bnn_hmc/README.md引用该论文时可直接使用article{izmailov2021bayesian, title{What Are Bayesian Neural Network Posteriors Really Like?}, author{Izmailov, Pavel and Vikram, Sharad and Hoffman, Matthew D and Wilson, Andrew Gordon}, journal{arXiv preprint arXiv:2104.14421}, year{2021} }二、环境搭建requirements.txt 与 pip 安装仓库提供了 bnn_hmc/requirements.txt其头部注释写明可以直接用来创建 conda 环境conda create --name env --file requirements.txt从该文件内容看这是一个基于Python 3.8.8、Linux-64平台锁定的环境清单核心依赖包括jax0.2.12、jaxlib0.1.65cuda112CUDA 11.2 版本、dm-haiku0.0.5.dev0、optax0.0.6、tensorflow2.4.1、tensorflow-datasets4.2.0、numpy1.19.5、scipy1.6.3、tabulate0.8.9、tensorboard2.5.0等。可见代码运行同时依赖 JAX 生态数值计算、Haiku网络定义与 TensorFlow数据加载与日志。如果使用 pip 手动安装README 给出的步骤为pip install tensorflow pip install --upgrade pip pip install --upgrade jax jaxlib0.1.65cuda112 -f \ https://storage.googleapis.com/jax-releases/jax_releases.html pip install githttps://github.com/deepmind/dm-haiku pip install tensorflow_datasets pip install tabulate pip install optax其中 JAX 的安装方式与硬件强相关示例中的jaxlib0.1.65cuda112对应 CUDA 11.2 环境安装其他硬件上的 JAX 请以 JAX 官方仓库的最新安装说明为准。三、文件结构总览仓库目录结构在 README 中给出了清晰说明详见 bnn_hmc/README.md实际目录与 README 描述一致可归纳为三层核心算法层bnn_hmc/corehmc.pyHamiltonian Monte Carlo 算法实现sgmcmc.py以 optax optimizer 形式实现的 SGMCMC 方法SGLD / SGHMCvi.py均值场变分推断Mean Field VI实现。工具层bnn_hmc/utils训练脚本共用的功能模块包括train_utils.py训练轮次与更新规则、models.py实验用网络结构、losses.py先验与似然函数、data_utils.py数据加载与预处理、optim_utils.py优化器与学习率调度、ensemble_utils.py预测集成、metrics.py评估指标、cmd_args_utils.py公共命令行参数、script_utils.py训练脚本公共功能、checkpoint_utils.py检查点保存/加载、logging_utils.py日志打印、precision_utils.py数值精度控制、tree_utils.pypytree 通用操作。训练脚本层bnn_hmc 根目录run_hmc.pyHMC 训练脚本run_sgd.pySGD 训练脚本深度集成即多次以不同种子运行它run_sgmcmc.pySGMCMC 训练脚本run_vi.pyMFVI 训练脚本make_posterior_surface_plot.py后验密度可视化脚本。此外 bnn_hmc/notebooks 目录还提供了若干交互式 Notebook如synthetic_regression_inference.ipynb、mcmc_gaussian_test.ipynb可用于小规模验证 MCMC 推断流程。四、公共命令行参数所有训练脚本共用的基础配置所有训练脚本run_hmc.py、run_sgd.py、run_sgmcmc.py、run_vi.py共享由 utils/cmd_args_utils.py 中add_common_flags()定义的一组参数参数说明seed随机种子源码默认0dir训练目录用于保存检查点与 TensorBoard 日志必填dataset_name数据集名称例如cifar10、cifar100、imdbUCI 数据集写作UCI 数据集名_随机种子的形式例如yacht_2其中种子决定训练/测试划分subset_train_to从数据集使用的样本数默认使用全部数据model_name神经网络结构名称例如lenet、resnet20_frn_swish、cnn_lstm、mlp_regression_smallweight_decay权重衰减对贝叶斯方法权重衰减决定了先验方差prior_var 1 / weight_decaytemperature后验温度默认1init_checkpoint用于初始化链/模型的检查点路径可选tabulate_freq表格表头打印频率use_float64使用 float64 精度在 TPU 及部分 GPU 上不可用默认使用float32另外源码中还包含--tpu_ip参数默认None用于指定 Cloud TPU 的内部 IP配合 utils/train_utils.py 中set_up_jax()的jax_backend_target配置可将计算运行在 TPU 上use_float64则通过jax_enable_x64开关开启。需要特别说明的是weight_decay与先验的关系在 utils/losses.py 的make_gaussian_log_prior()中高斯先验的对数密度为-0.5 * weight_decay * ||params||² 常数项再除以温度temperature因此weight_decay越大意味着先验方差越小、先验越强。这也是论文中讨论“先验尺度鲁棒性”实验的操作入口。五、运行 HMCrun_hmc.py5.1 HMC 专属参数除公共参数外HMC 训练脚本 run_hmc.py 还定义以下参数参数说明源码默认值step_sizeHMC 步长1e-4trajectory_lenHMC 轨迹长度1e-3num_iterationsHMC 迭代总次数1000max_num_leapfrog_steps允许的最大 leapfrog 步数作为 sanity check应大于trajectory_len / step_size10000num_burn_in_iterationsburn-in 迭代次数0--no_mh若设置忽略 Metropolis–Hastings 校正False其中 leapfrog 步数在训练循环中由ceil(trajectory_len / step_size)计算得出见 utils/train_utils.py 的make_hmc_update()若超过max_num_leapfrog_steps会直接断言报错。5.2 示例命令IMDB 上的 CNN-LSTM论文在 IMDB 情感分类任务上使用 CNN-LSTM 模型cnn_lstm并以不同温度运行 HMC 验证“冷后验”问题。README 给出的四组命令完整复刻如下运行环境为 8 张 NVIDIA Tesla V-100 GPU# Temperature 1 python3 run_hmc.py --seed1 --weight_decay40. --temperature1. \ --dirruns/hmc/imdb/ --dataset_nameimdb --model_namecnn_lstm \ --use_float64 --step_size1e-5 --trajectory_len0.24 \ --max_num_leapfrog_steps30000 # Temperature 0.3 python3 run_hmc.py --seed1 --weight_decay40. --temperature0.3 \ --dirruns/hmc/imdb/ --dataset_nameimdb --model_namecnn_lstm \ --use_float64 --step_size3e-6 --trajectory_len0.136 \ --max_num_leapfrog_steps46000 # Temperature 0.1 python3 run_hmc.py --seed1 --weight_decay40. --temperature0.1 \ --dirruns/hmc/imdb/ --dataset_nameimdb --model_namecnn_lstm \ --use_float64 --step_size1e-6 --trajectory_len0.078 \ --max_num_leapfrog_steps90000 # Temperature 0.03 python3 run_hmc.py --seed1 --weight_decay40. --temperature0.03 \ --dirruns/hmc/imdb/ --dataset_nameimdb --model_namecnn_lstm \ --use_float64 --step_size1e-6 --trajectory_len0.043 \ --max_num_leapfrog_steps45000可以观察到温度越低所需步长越小、leapfrog 步数越多——这是因为降温使得后验更“尖”需要更精细的积分步长。IMDB 任务使用--use_float64以获得数值稳定的接受率计算。5.3 示例命令MNIST 160 样本子集上的 MLP论文中还有一个轻量示例——在 MNIST 随机取 160 个样本的子集上训练 MLP 分类器mlp_classification单张 GPU 即可运行适合快速验证整个 HMC 流程python3 run_hmc.py --seed0 --weight_decay1. --temperature1. \ --dirruns/hmc/mnist_subset160 --dataset_namemnist \ --model_namemlp_classification --step_size3.e-5 --trajectory_len1.5 \ --num_iterations100 --max_num_leapfrog_steps50000 \ --num_burn_in_iterations10 --subset_train_to160注意这里通过--subset_train_to160限制训练样本数并设置了 10 次 burn-in 迭代。5.4 HMC 的源码级原理HMC 更新逻辑在 core/hmc.py 中实现核心流程清晰可循leapfrog 积分make_leapfrog对动量做半步更新 → 更新参数 → 重新计算梯度 → 再对动量做半步更新整个过程用jax.lax.fori_loop串起n_leapfrog步接受概率make_accept_prob基于动能差、对数似然差与对数先验差计算能量差得到min(1, exp(energy_diff))代码注释特别指出将似然与先验分开返回是为了在 float32 下更精确地计算接受/拒绝步骤中的先验密度比自适应步长adapt_step_size按step_size * exp(speed * (accept_prob - target_accept_rate))调整步长使接受率逼近目标值MH 校正make_adaptive_hmc_updateburn-in 之外默认开启 Metropolis–Hastings 校正--no_mh可关闭。在 run_hmc.py 的训练主循环中每轮迭代都会评估 train/test 指标、将被接受的样本或--no_mh时的所有样本累加进ensemble_predictions即 BMA 集成、保存检查点并把accept_prob、num_ensembled、当前step_size等写入 TensorBoard。训练目录名由模型、权重衰减、步长、轨迹长度、burn-in、MH 开关、温度与种子自动拼接而成便于对比实验归档。需要说明的是README 明确指出论文中 CIFAR-10 上的 HMC 实验是在包含 512 个 TPU 设备的 TPU pod 上、使用后续将发布的修改版代码运行的——即本仓库开箱即用的代码主要面向 GPU 场景CIFAR 级全量 HMC 需要特殊规模的硬件。六、运行 SGD 与深度集成run_sgd.pySGD 脚本 run_sgd.py 使用cosine 学习率调度其专属参数如下由 utils/cmd_args_utils.py 的add_sgd_flags()定义参数说明源码默认值init_step_size初始 SGD 步长1e-6num_epochsSGD 总 epoch 数300batch_size批大小80eval_freq评估频率epoch10save_freq检查点保存频率epoch50momentum_decaySGD 动量衰减0.9README 给出的示例ResNet-20-FRNresnet20_frn_swish在 CIFAR-10 / CIFAR-100 上训练 500 个 epoch并使用--subset_train_to40960限制样本数python3 run_sgd.py --seed1 --weight_decay10 --dirruns/sgd/cifar10/ \ --dataset_namecifar10 --model_nameresnet20_frn_swish \ --init_step_size3e-7 --num_epochs500 --eval_freq10 --batch_size80 \ --save_freq500 --subset_train_to40960 python3 run_sgd.py --seed1 --weight_decay10 --dirruns/sgd/cifar100/ \ --dataset_namecifar100 --model_nameresnet20_frn_swish \ --init_step_size1e-6 --num_epochs500 --eval_freq10 --batch_size80 \ --save_freq500 --subset_train_to40960IMDB 上的 CNN-LSTMpython3 run_sgd.py --seed1 --weight_decay3. --dirruns/sgd/imdb/ \ --dataset_nameimdb --model_namecnn_lstm --init_step_size3e-7 \ --num_epochs500 --eval_freq10 --batch_size80 --save_freq500深度集成的训练方式README 明确说明只需用不同随机种子多次运行run_sgd.py将得到的多个 SGD 解进行预测集成即可。论文比较了深度集成与 HMC 的预测分布差异得出“深度集成与 HMC 的接近程度和标准 SGLD 相当、且优于标准变分推断”的结论。从实现上看SGD 训练过程位于 utils/train_utils.py 的make_sgd_train_epoch()数据被切分为num_batches份每个 epoch 内用jax.lax.scan串行执行各 batch 的梯度更新并由jax.lax.psum完成跨设备梯度聚合体现 JAX 多设备如多 GPU并行训练能力。七、运行 SGMCMCrun_sgmcmc.pySGMCMC 脚本 run_sgmcmc.py 共享 SGD 的全部参数并额外引入以下参数参数说明源码默认值preconditioner预处理器选择None或RMSpropNonestep_size_schedule步长调度constant或cyclical。constant表示先做num_burnin_epochs个 epoch 的 cosine burn-in随后步长固定为final_step_sizecyclical表示先做固定步长的 burn-in再进入 cosine 循环调度constantnum_burnin_epochs达到最终学习率前的 epoch 数300final_step_size最终步长仅用于constant调度默认取init_step_sizestep_size_cycle_length_epochs循环周期长度epoch仅用于cyclical调度50save_all_ensembled保存所有被集成的网络Falseensemble_freq集成 iterates 的频率epoch10CIFAR-10 上 ResNet-20-FRN 的四种 SGMCMC 变体示例# SGLD python3 run_sgmcmc.py --seed1 --weight_decay5. --dirruns/sgmcmc/cifar10/ \ --dataset_namecifar10 --model_nameresnet20_frn_swish --init_step_size1e-6 \ --final_step_size1e-6 --num_epochs10000 --num_burnin_epochs1000 \ --eval_freq10 --batch_size80 --save_freq10 --momentum0. \ --subset_train_to40960 # SGHMC python3 run_sgmcmc.py --seed1 --weight_decay5 --dirruns/sgmcmc/cifar10/ \ --dataset_namecifar10 --model_nameresnet20_frn_swish --init_step_size3e-7 \ --final_step_size3e-7 --num_epochs10000 --num_burnin_epochs1000 \ --eval_freq10 --batch_size80 --save_freq10 --subset_train_to40960 \ --momentum0.9 # SGHMC-CLR python3 run_sgmcmc.py --seed1 --weight_decay5 --dirruns/sgmcmc/cifar10/ \ --dataset_namecifar10 --model_nameresnet20_frn_swish --init_step_size3e-7 \ --num_epochs10000 --num_burnin_epochs1000 --step_size_schedulecyclical \ --step_size_cycle_length_epochs50 --ensemble_freq50 --eval_freq10 \ --batch_size80 --save_freq1000 --subset_train_to40960 \ --preconditionerNone --momentum0.95 --eval_freq10 --save_all_ensembled # SGHMC-CLR-Prec python3 run_sgmcmc.py --seed1 --weight_decay5 --dirruns/sghmc/cifar10/ \ --dataset_namecifar10 --model_nameresnet20_frn_swish --init_step_size3e-5 \ --num_epochs10000 --num_burnin_epochs1000 --step_size_schedulecyclical \ --step_size_cycle_length_epochs50 --ensemble_freq50 --eval_freq10 \ --batch_size80 --save_freq50 --subset_train_to40960 \ --preconditionerRMSprop --momentum0.95 --eval_freq10 --save_all_ensembled四个变体对应论文中比较的 SGMCMC 方法族标准 SGLDmomentum0、SGHMC动量0.9、SGHMC 循环学习率CLRstep_size_schedulecyclical且momentum0.95以及在此基础上叠加 RMSprop 预处理preconditionerRMSprop。注意 SGLD 与 SGHMC 使用了不同的初始/最终步长1e-6与3e-7这与各自方法的噪声特性有关。SGMCMC 的实现位于 core/sgmcmc.py以 optaxGradientTransformation形式提供sgld_gradient_update()实现 SGLD当momentum_decay0时为经典 SGLDWelling Teh, ICML 2011否则退化为欠阻尼 SGLD即 SGHMCChen et al., ICML 2014。更新规则中梯度项乘lr_sqrt高斯噪声项的标准差为sqrt(2 * (1 - momentum_decay))再叠加动量项momentum_decay * m预处理通过Preconditioner抽象实现None对应IdentityPreconditionerState恒等变换RMSprop对应get_rmsprop_preconditioner()其基于梯度二阶矩估计默认衰减因子0.99、eps1e-7对噪声与更新量分别做M^{-1/2}与M^{-1}变换即论文中 SGHMC-CLR-Prec 的“Prec”来源。步长调度在 run_sgmcmc.py 的get_lr_schedule()中构建constant调度调用make_constant_lr_schedule_with_cosine_burnin(init_step_size, final_step_size, burnin_steps)cyclical调度调用make_cyclical_cosine_lr_schedule_with_const_burnin(...)周期由step_size_cycle_length_epochs换算为步数。训练过程中ensemble_freq决定每隔多少个 epoch 将当前 iterate 的预测累加进集成save_all_ensembled则同时把被集成的网络参数落盘。八、运行 MFVIrun_vi.py均值场变分推断脚本 run_vi.py 同样共享 SGD 参数并额外定义参数说明源码默认值optimizer优化器选择SGD或Adam源码默认AdamREADME 描述为默认 SGD请以仓库代码为准vi_sigma_initMFVI 中权重标准差σ的初始值1e-3vi_ensemble_sizeVI 评估时采样的集成规模20mean_init_checkpoint用于初始化 MFVI 均值mean的 SGD 检查点None示例命令# ResNet-20-FRN on CIFAR-10 or CIFAR-100 python3 run_vi.py --seed11 --weight_decay5. --dirruns/vi/cifar100/ \ --dataset_name[cifar10 | cifar100] --model_nameresnet20_frn_swish \ --init_step_size1e-4 --num_epochs300 --eval_freq10 --batch_size80 \ --save_freq300 --subset_train_to40960 --optimizerAdam \ --vi_sigma_init0.01 --temperature1. --vi_ensemble_size20 \ --mean_init_checkpointpath-to-sgd-solution # CNN-LSTM on IMDB python3 run_vi.py --seed11 --weight_decay5. --dirruns/vi/imdb/ \ --dataset_nameimdb --model_namecnn_lstm --init_step_size1e-4 \ --num_epochs500 --eval_freq10 --batch_size80 --save_freq200 \ --optimizerAdam --vi_sigma_init0.01 --temperature1. --vi_ensemble_size20 \ --mean_init_checkpointpath-to-sgd-solution其中--mean_init_checkpoint通常指向前面用run_sgd.py训练出的解即论文中的“用 SGD 解初始化 MFVI 均值”策略。MFVI 的实现在 core/vi.pyget_mfvi_model_fn()将每个权重参数化为{mean, inv_softplus_std}两个量均值直接用原参数复制标准差则存其 inverse-softplus 形式sigma softplus(inv_softplus_std)保证标准差恒为正sample_parms_fn用重参数化技巧m n * sn 为标准正态噪声从变分分布采样make_kl_with_gaussian_prior(weight_decay, temperature)实现 ELBO 中的先验 KL 项对每个参数计算log(σ_prior/σ_vi) (σ_vi² μ_vi²)/(2σ_prior²) - 1/2其中σ_prior sqrt(1/weight_decay)并以temperature作为 KL 项权重评估时vi_ensemble_predict_fn从变分后验采样vi_ensemble_size次并集成预测同时run_vi.py还会将每个参数的 softplus 标准差以直方图形式写入 TensorBoardMFVI/param_stds便于观察各层不确定性的分布。九、可视化后验密度make_posterior_surface_plot.py脚本 make_posterior_surface_plot.pyREADME 中笔误写作makemake_posterior_surface_plot.py实际文件名以此处为准可以在由三个检查点张成的二维平面上可视化后验的 log 密度、log 似然与 log 先验复现论文中的后验曲面图。参数如下参数说明源码默认值limit_bottom曲面可视化在垂直方向底部的范围-0.25limit_top曲面可视化在垂直方向顶部的范围1.25limit_left曲面可视化在水平方向左侧的范围-0.25limit_right曲面可视化在水平方向右侧的范围1.25grid_size每个方向上的网格点数20checkpoint1第一个检查点路径—checkpoint2第二个检查点路径—checkpoint3第三个检查点路径—示例IMDB 上的 CNN-LSTMpython3 make_posterior_surface_plot.py --weight_decay40 --temperature1. \ --dirruns/surface_plots/imdb/ --model_namecnn_lstm --dataset_nameimdb \ --checkpoint1ckpt1 --checkpoint2ckpt2 --checkpoint3ckpt3 --limit_bottom-0.75 --limit_left-0.75 --limit_right1.75 --limit_top1.75 \ --grid_size50从源码看该脚本以三个检查点构造平面坐标系每点对应平面上的一个二维坐标在grid_size × grid_size网格上逐点计算 log 后验密度、log 似然与 log 先验依赖 utils/losses.py 与 utils/models.py 提供的前向计算并用 matplotlib 输出可视化结果是理解“BNN 后验为何不像高斯”“冷后验现象由何而来”等论文论点的直观工具。十、小结如何组织一次完整的复现实验综合以上内容一次完整的论文实验复现可按如下流程组织按 第二节 搭建 conda/pip 环境用run_sgd.py在不同种子下训练若干 SGD 解得到基线精度与深度集成用run_hmc.py在对应任务上运行全批量 HMC建议先用 5.3 节 的 MNIST 160 样本示例在小规模上验证流程得到近乎精确的后验与 BMA 集成结果用run_sgmcmc.py分别运行 SGLD / SGHMC / SGHMC-CLR / SGHMC-CLR-Prec与 HMC 对比预测分布的接近程度用run_vi.py以 SGD 解初始化均值运行 MFVI作为变分推断基线用make_posterior_surface_plot.py对若干代表性检查点生成后验曲面图直观检验论文关于先验尺度、温度与后验形态的结论。所有脚本都将检查点与 TensorBoard 日志写入--dir指定目录可通过 TensorBoard 对比各方法的 accuracy、NLL、ECE分类任务或 RMSE / NLL回归任务等指标以及test/ens_*前缀的集成指标。值得留意的是论文涉及温度对比实验时如 HMC 的四个温度档位需保持weight_decay一致而仅改变temperature并在解读结果时记住该仓库代码默认不使用数据增强而论文指出的“冷后验效应很大程度上是数据增强的产物”这一结论正是建立在这样的实验设置之上的。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进