news 2026/9/15 11:32:51

PyMC 模型与 PyTensor FunctionGraph 互转指南:fgraph_from_model、model_from_fgraph 与 clone_model 全解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyMC 模型与 PyTensor FunctionGraph 互转指南:fgraph_from_model、model_from_fgraph 与 clone_model 全解析

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_modelModel编译成 PyTensor 的FunctionGraphmodel_from_fgraph再把计算图还原成新的Model,而clone_model则是二者的便捷组合。这套机制是模型重写(rewriting)、模型克隆与多种模型变换(如剪枝、去 Minibatch、条件化)的底层基石,读完本文你将掌握这三个 API 的完整签名、底层ModelVar占位算子机制、约束条件与实战改造流程。

一、为什么需要 Model 与 FunctionGraph 互转?

Model是用户友好的高层容器:它记录free_RVsobserved_RVspotentialsdeterministicsdata_vars等结构化信息,并提供dims、坐标(coords)等语义元数据。但要做图级别的自动变换(例如把居中式参数化改为非居中式、剪掉与观测无关的变量、移除 Minibatch 节点),最自然的载体是 PyTensor 的FunctionGraph——一个可以执行node_rewriterreplace_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_model

1.fgraph_from_model(model, inlined_views=False)

将 PyMC 模型转换为 PyTensorFunctionGraph,返回二元组(fgraph, memo)

  • modelpymc.model.core.Model实例。
  • inlined_viewsbool,默认False。决定 Deterministic 与 Data 等 "view" 变量是作为独立的图分支出现,还是被内联(inline)进随机变量之间。
  • 返回值
    • fgraph:包含模型变量副本的FunctionGraph,每个变量被包装在哑ModelVar算子中,保证可以用model_from_fgraph还原出合法模型;
    • memo:从原始模型变量到 fgraph 中等价节点的映射字典。

函数在转换前会做三类前置校验(fgraph.py):

  1. 存在非默认initial_values时抛出NotImplementedError("Cannot convert models with non-default initial_values");
  2. 模型是嵌套子模型(model.parent is not None)时抛出ValueError("Nested sub-models cannot be converted..."),因为子模型必须通过父模型转换;
  3. 检测到名称含_rotated__hsgp_coeffs_的变量时发出UserWarning,提示这些变量可能来自旧的 GP 对象,继续使用旧 GP 对象可能把旧模型变量重新引入。

2.model_from_fgraph(fgraph, mutate_fgraph=False)

将带有哑ModelVar算子的FunctionGraph还原为 PyMC 模型:

  • fgraphfgraph_from_model产出的计算图。
  • mutate_fgraphbool,默认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。它把变量的namedims存为 Op 属性(重建模型所需的元信息),而被包装的变量(以及取值变量)作为输入。

perform方法直接抛出RuntimeError("ModelVars should never be in a final graph!"),确保这些占位算子绝不可能出现在最终的计算图中。do_constant_folding返回False,防止折叠破坏结构。

在此基础上派生了完整的算子家族(fgraph.py):

算子类构造函数属性对应模型角色
ModelVarname, dims基类
ModelValuedVarname, dims, transform带取值变量(value var)的变量基类
ModelFreeRV继承ModelValuedVar自由随机变量(free RV),额外携带transform
ModelObservedRV继承ModelValuedVar观测随机变量
ModelPotentialname, dims势能项(pm.Potential
ModelDeterministicname, dims确定性变量(pm.Deterministic
ModelNamedname, 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完成替换。

③ 构造 FunctionGraphFunctionGraph(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 的底层变量;
  • 按算子类型填充新模型的各映射表:ModelFreeRVcreate_value_var(var, transform=..., value_var=value)set_initvalModelObservedRV走无 transform 的create_value_varModelPotential追加到potentialsModelDeterministic追加到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_lengthsnamed_vars_to_dims、六类变量归属、rvs_to_transforms(如HalfNormallogtransform)全部保留;且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 表明:模型参数(如musigma)虽被克隆为新对象(mu_new is not mu),但仍指向同一底层容器same_storage(mu, mu_new)为 True),即"共享同一份数据、对象相互独立";
  • Deterministic 的链式视图:test_deterministics 覆盖了 Deterministic 直接复制 RV 的特殊情形,重建后y_y__都直接指向y,不会重复插入多余的确定性节点;
  • 多元变换保留:test_multivariate_transform 验证Dirichletsimplex变换与LKJCholeskyCovcholesky-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=Falseclone_replace后手工恢复_coords_dim_lengths
  • model/core.py 的Model.clone()
  • model/transform/conditioning.py、model/transform/optimization.py 等模块也大量复用该互转管线;
  • testing.py 在测试工具中用它做模型等价性比较。

十、使用注意事项小结

  1. 转换会丢弃(并拒绝)非默认初始值,见 test_core.py 的注释;
  2. model_from_fgraph不能运行在活动的Model上下文内;
  3. 嵌套子模型只能通过父模型整体转换;
  4. 克隆后 RNG 独立、Data 与 dim_lengths 的共享变量也独立,但用户自定义 shared variable 共享底层内存——若需要彻底隔离请自行重建;
  5. 大体积 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/15 11:31:52

考研复试Python备考指南与科研应用实战

1. 为什么考研复试要考Python?作为一名经历过考研复试的过来人,我清楚地记得当时看到复试要求"Python基础"时的困惑。直到后来读研期间参与多个科研项目,我才真正理解Python在学术研究中的重要性。Python在科研领域的应用场景远超想…

作者头像 李华
网站建设 2026/9/15 11:29:51

智能生产系统演进:从自动化到AI决策的十年实践

1. 维他动力十年演进概述(2015-2025)维他动力作为一家专注于健康饮品研发的企业,在过去十年间经历了从传统配方到智能化生产的完整转型周期。2015年我们推出首款含电解质运动饮料时,生产线还需要人工调配基础溶液;而到…

作者头像 李华
网站建设 2026/9/15 11:29:41

CloudCompare实操:三维模型转点云全流程与避坑指南

大概两年前我接手了一个古建筑数字化项目,对方给了一批高精度三维模型,但验收时却要求提交点云数据用于后续的形变分析。当时我第一反应是“这不就是导一下格式的事吗”,结果真上手才发现,三维模型(Mesh)和…

作者头像 李华