尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

EBP 能量基过程(Energy-Based Processes for Exchangeable Data):原理、安装、训练与评估实战指南

EBP 能量基过程(Energy-Based Processes for Exchangeable Data):原理、安装、训练与评估实战指南 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载本文以 ebp/README.md 为主体结合 ebp 仓库源码系统讲解 Google Research 开源的 EBPEnergy-Based Processes实现它如何用能量函数为可交换数据exchangeable data建模、如何安装运行、如何完成训练与能量热力图评估并深入剖析其对抗式训练与 MCMC 采样的源码级原理。读完本文你将能独立复现多模态合成数据、MNIST 图像补全与点云生成/去噪三类实验并掌握通过命令行参数调控训练行为的完整方法。图 1多模态合成数据实验。左起为 Ground Truth 与 GP、NP、VIP、EBP 学习的能量分布可视化EBP 成功捕获了 toy 数据的多模态特性来源ebp/figures/exp_syn.png。一、项目背景什么是 Energy-Based ProcessesEBPEnergy-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.shebp/ebp/common/公共组件命令行参数、能量函数族、流模型、生成器、数据读取、绘图工具ebp/figures/README 展示的实验结果图公共组件中几个关键文件的角色如下可从源码结构推断其依赖关系ebp/ebp/common/cmd_args.py统一使用argparse注册全部超参数parse_known_args保证未知参数如启动脚本追加的$不会报错并自动创建save_dir目录ebp/ebp/common/f_family.py定义DeepsetEncoderDeepSets 集合编码器对集合元素做置换不变聚合与VAE将集合编码为隐变量 z 的变分后验为能量函数提供隐变量条件ebp/ebp/common/generator.py定义HyperGen条件生成器负责从隐变量与上下文生成目标样本ebp/ebp/common/data_utils/curve_reader.pyget_reader按-data_name提供合成数据读取器generate_curves()生成训练曲线数据ebp/ebp/common/plot_utils/plot_2d.pyplot_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-sonnet1.23Sonnet 模块化网络库tensorflow1.13.1TensorFlow 1.xnumpy、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其完整内容如下datamix_line bsize6 ctx15 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_epochs5适合快速冒烟验证。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随机种子脚本默认10086main.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采样器类型可选HMC、GeneralHmc、ResGeneralHmc、SGLD配合-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 可以还原每一轮训练的具体步骤构造能量函数与后验在score_func变量作用域内用VAE将真实数据(query, target_y)编码为隐变量z_outer含mu/sigma/neg_klScoreFunc(embed_dim32)以集合与隐变量为输入输出能量值main.py构造条件生成器在generator作用域内HyperGen(dim1, condx_dim1, condz_dim32, num_layers10)根据上下文与隐变量生成伪造样本x_fake及其对数似然ll_fakemain.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_lossmain.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, axis0)取平均得到heat-{epoch}.pdfmain.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.shnum_epochs5快速确认环境与数据管线正常再切换 ebp/ebp/experiments/run_ebp.sh 进行完整训练结果目录结构训练产物model-{epoch}.ckpt、plot-{epoch}.pdf、heat-{epoch}.pdf等统一写入-save_dircmd_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-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/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人工智能计算机视觉深度学习预训练上一篇告别后端依赖在浏览器中玩转PL/pgSQL存储过程——PGlite全攻略下一篇Zotero Style插件终极指南如何用智能标签和进度追踪提升文献管理效率创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表