- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
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_reader、cmd_args、ScoreFunc/VAE、HyperGen、plot_samples全部串起来,构成一条完整的“数据读取 → 能量函数/后验 → 条件生成器 → 对抗训练 → 可视化/评估”流水线。
三、安装指南
仓库通过setuptools打包,安装只需在项目根目录执行:
pip3 install -e .根据 ebp/setup.py 中的install_requires,核心依赖为:
dm-sonnet==1.23(Sonnet 模块化网络库)tensorflow==1.13.1(TensorFlow 1.x)numpy、tqdm、scipy、matplotlib
注意两点硬性前提:
- 安装过程需要 gcc 编译器(部分扩展需本地编译);
- 若启用 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_name | None | 合成数据名,脚本默认mix_line;由 curve_reader.py 的get_reader分发 |
-batch_size | 100 | 小批量大小,脚本默认6(小 batch 便于曲线级建模) |
-num_ctx | 10 | 上下文(context)点数量,脚本默认15 |
-gp_lambda | 0 | 梯度惩罚(gradient penalty)系数,脚本默认1。见 main.py:在真实/伪造样本之间插值并对能量函数施加 Lipschitz 约束 |
-ent_lam | 1.0 | 生成器熵正则系数,脚本默认0.01,乘在伪造样本对数似然项上(loss = -mean(f) + ent_lam * mean(ll_fake)) |
-num_epochs | 50000 | 训练轮数,脚本默认50 |
-seed | 1 | 随机种子,脚本默认10086;main.py 中同步设置random/numpy/tf三处种子 |
-sigma_eps | 1e-1 | 重参数化的标准差尺度;f_family.py 中sigma = sigmoid(logit_sigma) * sigma_eps,即用它约束后验标准差上限 |
-beta1 | 0.9 | Adam 第一矩衰减系数,脚本设为0 |
-learning_rate | 0.001 | Adam 学习率 |
-energy_type | mlp | 能量函数类型 |
-mcmc_type | None | 采样器类型,可选HMC、GeneralHmc、ResGeneralHmc、SGLD,配合-mcmc_steps、-hmc_step_size等使用 |
-score_type/-score_func | agg/single | 得分聚合方式(agg/prod)与得分函数形式(single/mixture) |
-flow_type/-num_flows | planar/1 | 生成器使用的流类型与流数量 |
-epoch_load | -1 | 加载指定 epoch 的 checkpoint 进行评估;-1表示从零开始训练 |
4.3 训练循环内部机制
从 main.py 可以还原每一轮训练的具体步骤:
- 构造能量函数与后验:在
score_func变量作用域内,用VAE将真实数据(query, target_y)编码为隐变量z_outer(含mu/sigma/neg_kl),ScoreFunc(embed_dim=32)以集合与隐变量为输入输出能量值(main.py); - 构造条件生成器:在
generator作用域内,HyperGen(dim=1, condx_dim=1, condz_dim=32, num_layers=10)根据上下文与隐变量生成伪造样本x_fake及其对数似然ll_fake(main.py); - 判别器(能量)更新:
get_disc_loss最小化mean(-f(x)) + mean(f(x_fake)) - mean(neg_kl),即拉低真实样本能量、抬高伪造样本能量,并对梯度做 NaN 防护(main.py); - 生成器更新:
get_gen_loss最大化伪造样本的期望能量并加入熵正则(main.py); - 交替优化:每个 batch 内判别器更新 1 次、生成器更新 3 次(
for i in range(3)),并以 tqdm 实时打印disc_loss/gen_loss(main.py); - 周期存档:每个 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恢复模型,随后依次执行:
- 条件样本生成:用测试数据跑生成器,绘制伪造样本散点图,输出
fig-{epoch}.pdf; - 上下文可视化:输出
observe-{epoch}.pdf展示测试上下文点; - 能量热力图:在
[-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 的阅读器中查看动态过程。
六、实践建议与注意事项
结合源码,给出几条实操建议:
- 首次运行先小规模验证:可使用 ebp/run.sh(
num_epochs=5)快速确认环境与数据管线正常,再切换 ebp/ebp/experiments/run_ebp.sh 进行完整训练; - 结果目录结构:训练产物(
model-{epoch}.ckpt、plot-{epoch}.pdf、heat-{epoch}.pdf等)统一写入-save_dir;cmd_args.py会自动创建该目录,无需手动 mkdir; - 环境约束:依赖锁定 TensorFlow 1.x 与 dm-sonnet 1.23,需在兼容的 Python 环境中运行;GPU 训练需提前配好 CUDA,且显存采用按需增长策略;
- 采样器扩展:如需在评估阶段使用 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
相关推荐
Yi 模型 GPTQ 量化实战指南:基于 AutoGPTQ 的后训练量化与推理评估
Yi 模型 GPTQ 量化实战指南:基于 AutoGPTQ 的后训练量化与推理评估 导读 本文是 Yi 开源模型仓库(GitHub_Trending/yi/Yi
人工智能大模型基础模型微调模型量化多模态SkyPilot 并行训练与评估 Job Group 实战:基于共享卷实现训练-评估流水线
SkyPilot 并行训练与评估 Job Group 实战:基于共享卷实现训练 评估流水线 本教程以 SkyPilot 仓库中的 train eval jobg
后端任务调度MLOps集群管理GroundingDINO日志分析:训练过程监控与性能评估全指南
GroundingDINO日志分析:训练过程监控与性能评估全指南 引言:你还在为目标检测训练调参焦头烂额? 开放式目标检测(Open vocabulary Ob
人工智能计算机视觉深度学习预训练
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考