PyMC 模型与 PyTensor FunctionGraph 互转指南:fgraph_from_model、model_from_fgraph 与 clone_model 全解析
【免费下载链接】pymcBayesian Modeling and Probabilistic Programming in Python项目地址: https://gitcode.com/GitHub_Trending/py/pymc
导读
在 PyMC 中,pymc.model.fgraph模块提供了一套"模型 ↔ 计算图"双向转换机制:fgraph_from_model把Model编译成 PyTensor 的FunctionGraph,model_from_fgraph再把计算图还原成新的Model,而clone_model则是二者的便捷组合。这套机制是模型重写(rewriting)、模型克隆与多种模型变换(如剪枝、去 Minibatch、条件化)的底层基石,读完本文你将掌握这三个 API 的完整签名、底层ModelVar占位算子机制、约束条件与实战改造流程。
一、为什么需要 Model 与 FunctionGraph 互转?
Model是用户友好的高层容器:它记录free_RVs、observed_RVs、potentials、deterministics、data_vars等结构化信息,并提供dims、坐标(coords)等语义元数据。但要做图级别的自动变换(例如把居中式参数化改为非居中式、剪掉与观测无关的变量、移除 Minibatch 节点),最自然的载体是 PyTensor 的FunctionGraph——一个可以执行node_rewriter、replace_all、拓扑替换等重写操作的低层数据结构。
pymc.model.fgraph就是这两者之间的桥梁,其 API 参考页面见 fgraph.rst,它被收录在 model.rst 的文档体系之下。
二、核心 API 一览
模块导出的公共接口定义在 fgraph.py 的__all__中,共三个函数:
from pymc.model.fgraph import fgraph_from_model, model_from_fgraph, clone_model1.fgraph_from_model(model, inlined_views=False)
将 PyMC 模型转换为 PyTensorFunctionGraph,返回二元组(fgraph, memo):
- model:
pymc.model.core.Model实例。 - inlined_views:
bool,默认False。决定 Deterministic 与 Data 等 "view" 变量是作为独立的图分支出现,还是被内联(inline)进随机变量之间。 - 返回值:
fgraph:包含模型变量副本的FunctionGraph,每个变量被包装在哑ModelVar算子中,保证可以用model_from_fgraph还原出合法模型;memo:从原始模型变量到 fgraph 中等价节点的映射字典。
函数在转换前会做三类前置校验(fgraph.py):
- 存在非默认
initial_values时抛出NotImplementedError("Cannot convert models with non-default initial_values"); - 模型是嵌套子模型(
model.parent is not None)时抛出ValueError("Nested sub-models cannot be converted..."),因为子模型必须通过父模型转换; - 检测到名称含
_rotated_或_hsgp_coeffs_的变量时发出UserWarning,提示这些变量可能来自旧的 GP 对象,继续使用旧 GP 对象可能把旧模型变量重新引入。
2.model_from_fgraph(fgraph, mutate_fgraph=False)
将带有哑ModelVar算子的FunctionGraph还原为 PyMC 模型:
- fgraph:
fgraph_from_model产出的计算图。 - mutate_fgraph:
bool,默认False。为True时允许函数就地修改 fgraph 及其变量(适合该 fgraph 之后不再使用的情形);为False时会先克隆一份 fgraph 再重建,避免副作用。
3.clone_model(model)
一键克隆模型,等价于:
model_from_fgraph(fgraph_from_model(model)[0], mutate_fgraph=True)克隆后的模型拥有原模型全部变量的新对象(名称、dims、coords 保持一致),但共享变量(如pm.Data)指向同一份底层内存容器;常量(Constant)不克隆。
三、底层机制:ModelVar哑算子家族
整个互转机制的精髓,是 fgraph.py 中定义的ModelVar——一个用于描述"模型变量用途"的哑Op。它把变量的name与dims存为 Op 属性(重建模型所需的元信息),而被包装的变量(以及取值变量)作为输入。
其perform方法直接抛出RuntimeError("ModelVars should never be in a final graph!"),确保这些占位算子绝不可能出现在最终的计算图中。do_constant_folding返回False,防止折叠破坏结构。
在此基础上派生了完整的算子家族(fgraph.py):
| 算子类 | 构造函数属性 | 对应模型角色 |
|---|---|---|
ModelVar | name, dims | 基类 |
ModelValuedVar | name, dims, transform | 带取值变量(value var)的变量基类 |
ModelFreeRV | 继承ModelValuedVar | 自由随机变量(free RV),额外携带transform |
ModelObservedRV | 继承ModelValuedVar | 观测随机变量 |
ModelPotential | name, dims | 势能项(pm.Potential) |
ModelDeterministic | name, dims | 确定性变量(pm.Deterministic) |
ModelNamed | name, dims | 命名变量(主要是 Data 等) |
对应的便捷工厂函数:model_free_rv(rv, value, transform, name, *dims)、model_observed_rv(rv, value, name, *dims)、model_potential(rv, name, *dims)、model_deterministic(rv, name, *dims)、model_named(rv, name, *dims)。
四、fgraph_from_model转换流程详解
整体流程在 fgraph.py,大致分为五步:
① 收集变量并处理 View:以model.rvs_to_values为线索收集所有 RV,把 Deterministic 与具名取值变量用view_op包装成 View(inlined_views=False时),这样它们不会穿插在"主变量"之间;随后通过local_remove_view重写(fgraph.py)把多余的 View 移除。
② 深拷贝共享变量:deepcopy_shared_variable(fgraph.py)手动重建三类共享变量——RNG 节点(通过 pytensorf.py 的find_rng_nodes定位)、dim_lengths中的共享变量、named_vars中的共享变量(Data)。注释特别说明:"Data(可能显著增加内存)"。共享变量没有 deepcopy 方法,因此这里通过type(var)(type=..., value=None, strict=None, container=deepcopy(var.container), name=...)手工重建,并把新变量放进memo完成替换。
③ 构造 FunctionGraph:FunctionGraph(outputs=model_vars, clone=True, memo=memo, copy_orphans=True, copy_inputs=True),并把模型的_coords与_dim_lengths(经 memo 映射后)复制到 fgraph 上。
④ 引入哑 ModelVar 算子:按"自由 RV / 观测 RV / Potential / Deterministic / 具名变量 / 未命名取值变量"六类逐一替换,并用toposort_replace(pytensorf.py,按拓扑序就地批量替换)把变量替换为包装后的节点;同时更新 memo 的反向映射。
⑤ 收尾清理:移除对应未命名取值变量的多余输出,最后应用remove_view_rewrite清掉噪音 View。
五、model_from_fgraph重建流程详解
重建逻辑在 fgraph.py:
- 关键约束:不能在
with pm.Model():上下文内调用。实现通过Model(model=None)显式不继承上下文模型,测试 test_context_error 验证了返回模型的parent is None; - 非
mutate_fgraph模式先克隆 fgraph(fgraph.clone_get_equiv),并同步更新_dim_lengths的 memo 映射; - 遍历
fgraph.toposort()收集所有ModelVar节点,用first_non_model_var递归解包到第一个非 ModelVar 的底层变量; - 按算子类型填充新模型的各映射表:
ModelFreeRV走create_value_var(var, transform=..., value_var=value)并set_initval;ModelObservedRV走无 transform 的create_value_var;ModelPotential追加到potentials;ModelDeterministic追加到deterministics(若它只是某个 RV 的直接视图,则用view_op包一层);ModelNamed追加到data_vars; - 最后用
op.name恢复变量名、op.dims恢复维度,调用add_named_variable注册。
六、clone_model实战
clone_model 的 docstring 给出了完整示例——克隆后可以在不影响原模型的前提下继续扩展图:
import pymc as pm from pymc.model.fgraph import clone_model with pm.Model() as m: p = pm.Beta("p", 1, 1) x = pm.Bernoulli("x", p=p, shape=(3,)) with clone_model(m) as clone_m: # 按名字访问克隆变量 clone_x = clone_m["x"] # z 只属于 clone_m,不属于 m z = pm.Deterministic("z", clone_x + 1)测试 test_basic 验证了往返转换的完整性:coords({"test_dim": tuple(range(3))})、_dim_lengths、named_vars_to_dims、六类变量归属、rvs_to_transforms(如HalfNormal的logtransform)全部保留;且pm.draw得到的随机样本与compile_logp计算的对数概率在克隆前后完全一致。
值得注意:Model.clone()方法在 core.py 中直接委托给clone_model(self),因此用户可以直接调用m.clone()。
七、图重写(Rewrite)实战:非居中式参数化
互转机制最大的价值在于"改图"。测试 test_fgraph_rewrite 与夹具non_centered_rewrite(test_fgraph.py)演示了完整流程:定义一个node_rewriter(tracksModelFreeRV),在 fgraph 上把居中式 Normal 替换为raw_标准正态 + 确定性变换:
@node_rewriter(tracks=[ModelFreeRV]) def non_centered_param(fgraph, node): rv, value = node.inputs name, dims = node.op.name, node.op.dims if not isinstance(rv.owner.op, pm.Normal): return rng, size, loc, scale = rv.owner.inputs if rv_size_is_none(size): return None # 构造 raw 标准正态并注册为新自由 RV raw_name = f"{name}_raw_" raw_norm = pm.Normal.dist(0, 1, size=size, rng=rng) raw_norm_value = raw_norm.clone() raw_norm_value.name = raw_name fgraph.add_input(raw_norm_value) raw_norm = model_free_rv(raw_norm, raw_norm_value, node.op.transform, raw_name, *dims) # 重建原变量为确定性变量 new_norm = loc + raw_norm * scale fgraph.add_output(model_deterministic(new_norm, name, *dims)) return [new_norm]随后non_centered_rewrite.apply(fg)作用于fgraph_from_model产出的图,再model_from_fgraph(fg)重建。测试断言:free_RVs变为{group_mean, group_std, subject_mean_raw_},subject_mean降级为deterministics,且新模型与手写参考模型在随机采样和对数概率上完全等价。
八、边界情况与数据独立性
- Data 与共享变量独立:test_data 验证 RNG、Data、dim_lengths 三类共享变量在克隆后不再共享同一块存储(
same_storage返回 False),在新模型中pm.set_data修改x不会影响原模型; - 用户自定义共享变量:test_shared_variable 表明:模型参数(如
mu、sigma)虽被克隆为新对象(mu_new is not mu),但仍指向同一底层容器(same_storage(mu, mu_new)为 True),即"共享同一份数据、对象相互独立"; - Deterministic 的链式视图:test_deterministics 覆盖了 Deterministic 直接复制 RV 的特殊情形,重建后
y_、y__都直接指向y,不会重复插入多余的确定性节点; - 多元变换保留:test_multivariate_transform 验证
Dirichlet的simplex变换与LKJCholeskyCov的cholesky-cov-packed变换在克隆前后产生一致的初始点; - 子模型拒绝:test_sub_model_error 验证嵌套子模型抛出
ValueError。
九、在仓库内部的真实应用
这套机制并非孤立 API,而是多个模型变换功能的地基(均以fgraph_from_model(model, inlined_views=True)+model_from_fgraph(fgraph, mutate_fgraph=True)的模式出现):
- transform/basic.py 的
prune_vars_detached_from_observed:利用 fgraph 的祖先分析剪掉与观测无关的变量(含 Potentials 时直接NotImplementedError); - transform/basic.py 的
remove_minibatched_nodes:把pm.Minibatch节点替换为原始输入,注意通过rebuild_strict=False的clone_replace后手工恢复_coords与_dim_lengths; - model/core.py 的
Model.clone(); - model/transform/conditioning.py、model/transform/optimization.py 等模块也大量复用该互转管线;
- testing.py 在测试工具中用它做模型等价性比较。
十、使用注意事项小结
- 转换会丢弃(并拒绝)非默认初始值,见 test_core.py 的注释;
model_from_fgraph不能运行在活动的Model上下文内;- 嵌套子模型只能通过父模型整体转换;
- 克隆后 RNG 独立、Data 与 dim_lengths 的共享变量也独立,但用户自定义 shared variable 共享底层内存——若需要彻底隔离请自行重建;
- 大体积 Data 被深拷贝,可能显著增加内存,可考虑
inlined_views与剪枝策略配合。
掌握fgraph_from_model→ 图重写 →model_from_fgraph这条管线,你就能以源码级的控制力实现自定义模型变换,这也是 PyMC 中许多高级功能(如自动非居中式化、模型剪枝)的通用底层范式。
【免费下载链接】pymcBayesian Modeling and Probabilistic Programming in Python项目地址: https://gitcode.com/GitHub_Trending/py/pymc
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考