EBP 能量基过程(Energy-Based Processes for Exchangeable Data):原理、安装、训练与评估实战指南
2026/9/20 6:36:06 网站建设 项目流程
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

本文以 ebp/README.md 为主体,结合 ebp 仓库源码,系统讲解 Google Research 开源的 EBP(Energy-Based Processes)实现:它如何用能量函数为可交换数据(exchangeable data)建模、如何安装运行、如何完成训练与能量热力图评估,并深入剖析其对抗式训练与 MCMC 采样的源码级原理。读完本文,你将能独立复现多模态合成数据、MNIST 图像补全与点云生成/去噪三类实验,并掌握通过命令行参数调控训练行为的完整方法。

图 1:多模态合成数据实验。左起为 Ground Truth 与 GP、NP、VIP、EBP 学习的能量分布可视化,EBP 成功捕获了 toy 数据的多模态特性(来源:ebp/figures/exp_syn.png)。

一、项目背景:什么是 Energy-Based Processes

EBP(Energy-Based Processes for Exchangeable Data)是 Mengjiao Yang、Bo Dai、Hanjun Dai 与 Dale Schuurmans 发表于 2020 年的工作,论文见 arXiv:2003.07521。其核心思想是:对可交换数据(即数据的排列不影响联合分布的一类随机过程,如高斯过程 GP)用显式能量函数建模,而非直接刻画概率密度本身。

仓库中score_func的命名与设计揭示了这一思路的关键:模型学到的不是显式密度,而是数据的能量/得分函数,配合 MCMC 采样(如 HMC、SGLD)即可从能量分布中抽取样本。与高斯过程(GP)、神经过程(NP)、变分隐式过程(VIP)相比,EBP 的优势在于能够捕捉数据分布的多样性与多模态——这在 ebp/figures/exp_syn.png 的能量热力图中得到直观体现。

若在科研工作中使用本代码库,请按仓库 README 给出的 BibTeX 引用该论文:

@article{yang2020energy, title={Energy-Based Processes for Exchangeable Data}, author={Yang, Mengjiao and Dai, Bo and Dai, Hanjun and Schuurmans, Dale}, journal={arXiv preprint arXiv:2003.07521}, year={2020} }

二、仓库结构与核心模块

从仓库目录结构看,项目分为三个层次:

路径职责
ebp/ebp/experiments/训练/测试入口main.py与启动脚本run_ebp.sh
ebp/ebp/common/公共组件:命令行参数、能量函数族、流模型、生成器、数据读取、绘图工具
ebp/figures/README 展示的实验结果图

公共组件中几个关键文件的角色如下(可从源码结构推断其依赖关系):

  • ebp/ebp/common/cmd_args.py:统一使用argparse注册全部超参数,parse_known_args保证未知参数(如启动脚本追加的$@)不会报错,并自动创建save_dir目录;
  • ebp/ebp/common/f_family.py:定义DeepsetEncoder(DeepSets 集合编码器,对集合元素做置换不变聚合)与VAE(将集合编码为隐变量 z 的变分后验),为能量函数提供隐变量条件;
  • ebp/ebp/common/generator.py:定义HyperGen条件生成器,负责从隐变量与上下文生成目标样本;
  • ebp/ebp/common/data_utils/curve_reader.py:get_reader-data_name提供合成数据读取器,generate_curves()生成训练曲线数据;
  • ebp/ebp/common/plot_utils/plot_2d.py:plot_samples绘制 2D 样本散点图。

依赖关系上,ebp/ebp/experiments/main.py 将curve_readercmd_argsScoreFunc/VAEHyperGenplot_samples全部串起来,构成一条完整的“数据读取 → 能量函数/后验 → 条件生成器 → 对抗训练 → 可视化/评估”流水线。

三、安装指南

仓库通过setuptools打包,安装只需在项目根目录执行:

pip3 install -e .

根据 ebp/setup.py 中的install_requires,核心依赖为:

  • dm-sonnet==1.23(Sonnet 模块化网络库)
  • tensorflow==1.13.1(TensorFlow 1.x)
  • numpytqdmscipymatplotlib

注意两点硬性前提:

  1. 安装过程需要 gcc 编译器(部分扩展需本地编译);
  2. 若启用 GPU 加速,需要 CUDA 环境(README 明确说明 "if gpu is enabled");同时 ebp/requirements.txt 直接声明了tensorflow-gpu,这与 README 的说明相互印证。

主程序 ebp/ebp/experiments/main.py 在ConfigProto中设置gpu_options.allow_growth = True,即 GPU 显存按需增长,避免一次性占满显存。

四、训练:启动脚本与参数解读

4.1 标准训练流程

README 给出的训练命令是:

cd ebp/experiments/ ./run_ebp.sh

由于本仓库实际目录结构为ebp/ebp/experiments/,对应脚本位于 ebp/ebp/experiments/run_ebp.sh,其完整内容如下:

data=mix_line bsize=6 ctx=15 save_dir=$HOME/scratch/results/ebp/$data-$bsize-$ctx python3 main.py \ -save_dir $save_dir \ -data_name $data \ -batch_size $bsize \ -num_ctx $ctx \ -gp_lambda 1 \ -ent_lam 0.01 \ -num_epochs 50 \ -seed 10086 \ -sigma_eps 1e-1 \ -beta1 0 \ $@

脚本末尾的$@会将你追加的任意命令行参数透传给main.py,这正是 README 中./run_ebp.sh -epoch_load 99能工作的机制——也便于在不改脚本的情况下覆盖默认超参数。此外,仓库根目录的 ebp/run.sh 提供了等价的模块化启动方式(python3 -m ebp.experiments.main,默认num_epochs=5,适合快速冒烟验证)。

4.2 核心超参数语义

以下参数均在 ebp/ebp/common/cmd_args.py 中注册,结合源码注释与主程序使用方式说明其作用:

参数默认值语义与源码佐证
-data_nameNone合成数据名,脚本默认mix_line;由 curve_reader.py 的get_reader分发
-batch_size100小批量大小,脚本默认6(小 batch 便于曲线级建模)
-num_ctx10上下文(context)点数量,脚本默认15
-gp_lambda0梯度惩罚(gradient penalty)系数,脚本默认1。见 main.py:在真实/伪造样本之间插值并对能量函数施加 Lipschitz 约束
-ent_lam1.0生成器熵正则系数,脚本默认0.01,乘在伪造样本对数似然项上(loss = -mean(f) + ent_lam * mean(ll_fake)
-num_epochs50000训练轮数,脚本默认50
-seed1随机种子,脚本默认10086;main.py 中同步设置random/numpy/tf三处种子
-sigma_eps1e-1重参数化的标准差尺度;f_family.py 中sigma = sigmoid(logit_sigma) * sigma_eps,即用它约束后验标准差上限
-beta10.9Adam 第一矩衰减系数,脚本设为0
-learning_rate0.001Adam 学习率
-energy_typemlp能量函数类型
-mcmc_typeNone采样器类型,可选HMCGeneralHmcResGeneralHmcSGLD,配合-mcmc_steps-hmc_step_size等使用
-score_type/-score_funcagg/single得分聚合方式(agg/prod)与得分函数形式(single/mixture
-flow_type/-num_flowsplanar/1生成器使用的流类型与流数量
-epoch_load-1加载指定 epoch 的 checkpoint 进行评估;-1表示从零开始训练

4.3 训练循环内部机制

从 main.py 可以还原每一轮训练的具体步骤:

  1. 构造能量函数与后验:在score_func变量作用域内,用VAE将真实数据(query, target_y)编码为隐变量z_outer(含mu/sigma/neg_kl),ScoreFunc(embed_dim=32)以集合与隐变量为输入输出能量值(main.py);
  2. 构造条件生成器:在generator作用域内,HyperGen(dim=1, condx_dim=1, condz_dim=32, num_layers=10)根据上下文与隐变量生成伪造样本x_fake及其对数似然ll_fake(main.py);
  3. 判别器(能量)更新get_disc_loss最小化mean(-f(x)) + mean(f(x_fake)) - mean(neg_kl),即拉低真实样本能量、抬高伪造样本能量,并对梯度做 NaN 防护(main.py);
  4. 生成器更新get_gen_loss最大化伪造样本的期望能量并加入熵正则(main.py);
  5. 交替优化:每个 batch 内判别器更新 1 次、生成器更新 3 次(for i in range(3)),并以 tqdm 实时打印disc_loss/gen_loss(main.py);
  6. 周期存档:每个 epoch 用tf.train.Saver保存model-{epoch}.ckpt,同时输出plot-{epoch}.pdf采样可视化图(main.py)。

五、测试与评估:能量热力图与条件补全

5.1 加载 checkpoint 进行测试

README 给出的测试方式是先训练(或复用已有 checkpoint),再指定加载的 epoch:

./run_ebp.sh ./run_ebp.sh -epoch_load 99

-epoch_load >= 0时,main.py 会从{save_dir}/model/model-{epoch}.ckpt恢复模型,随后依次执行:

  1. 条件样本生成:用测试数据跑生成器,绘制伪造样本散点图,输出fig-{epoch}.pdf
  2. 上下文可视化:输出observe-{epoch}.pdf展示测试上下文点;
  3. 能量热力图:在[-2, 2] × [-2, 2]的 50×50 网格上,对每个网格点计算能量分数,重复 100 次采样后softmax(score * 10, axis=0)取平均,得到heat-{epoch}.pdf(main.py)。

这就是 README 中“plot the energy heatmap, pass the latest check-pointed epoch number”的完整实现逻辑:热力图反映的是当前能量函数在二维平面上的概率分布倾向,可用于直观检查模型是否捕获了真实数据的结构与多模态。

5.2 三组官方实验结果

README 展示了四张实验结果图,对应三类核心能力:

  • 多模态合成数据:ebp/figures/exp_syn.png 对比 GP、NP、VIP、EBP 学习的能量分布,EBP 的曲线结构与 Ground Truth 最为吻合,成功捕获 toy 数据的多模态特性;
  • 图像补全—— 给定部分像素(左半部为带噪/遮挡输入),EBP 通过能量函数条件采样补全出完整的 MNIST 数字(图见 ebp/figures/exp_mnist.png),这一能力对应仓库中的-binary(二值图像)、-img_size(图像尺寸)等图像相关参数;
  • 点云生成与去噪:ebp/figures/generation.gif 展示学习到的 RNN 采样器逐步生成三维点云;ebp/figures/denoising.gif 展示利用学习到的能量函数对带噪点云做去噪(能量越低越接近真实流形)。两个 GIF 均位于 ebp/figures/ 目录,可在支持 GIF 的阅读器中查看动态过程。

六、实践建议与注意事项

结合源码,给出几条实操建议:

  1. 首次运行先小规模验证:可使用 ebp/run.sh(num_epochs=5)快速确认环境与数据管线正常,再切换 ebp/ebp/experiments/run_ebp.sh 进行完整训练;
  2. 结果目录结构:训练产物(model-{epoch}.ckptplot-{epoch}.pdfheat-{epoch}.pdf等)统一写入-save_dircmd_args.py会自动创建该目录,无需手动 mkdir;
  3. 环境约束:依赖锁定 TensorFlow 1.x 与 dm-sonnet 1.23,需在兼容的 Python 环境中运行;GPU 训练需提前配好 CUDA,且显存采用按需增长策略;
  4. 采样器扩展:如需在评估阶段使用 HMC/SGLD 等 MCMC 采样器,通过-mcmc_type-mcmc_steps-hmc_step_size等参数即可切换,无需修改代码。

七、总结

EBP 将能量基建模与可交换数据结合,为随机过程学习提供了区别于 GP/NP/VIP 的显式能量视角。本仓库给出了从合成数据到图像补全、再到点云生成/去噪的完整可复现实现:run_ebp.sh一键训练,-epoch_load加载 checkpoint 输出能量热力图与条件样本,cmd_args.py暴露全部可调超参数。读者可以在此基础上替换-data_name对应的数据读取器,或通过-mcmc_type-flow_type等参数定制自己的能量过程模型。

  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

相关推荐

上一篇:告别后端依赖:在浏览器中玩转PL/pgSQL存储过程——PGlite全攻略
下一篇:Zotero Style插件终极指南:如何用智能标签和进度追踪提升文献管理效率

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询