PyMC 维度感知数学运算:pymc.dims.math 模块原理与实战指南
【免费下载链接】pymcBayesian Modeling and Probabilistic Programming in Python项目地址: https://gitcode.com/GitHub_Trending/py/pymc
导读:在 PyMC 中,
pymc.dims.math是维度感知(dimension-aware)建模体系pymc.dims的数学运算核心。本文基于 docs/source/api/dims/math.rst 展开,结合源码与测试用例,讲解该模块如何包装pytensor.xtensor.math、linalg子模块及模块级再导出机制,并通过真实样条回归示例对比传统pm.math与维度感知写法,帮助读者掌握用命名维度(named dims)编写更健壮、可读性更强的贝叶斯模型。
一、从 API 文档看 pymc.dims.math 的定位
math.rst 全文虽短,却精准界定了该模块的三层职责:
- 包装全部数学运算:
pymc.dims.math模块包装了pytensor.xtensor.math中定义的所有数学运算; - 提供 linalg 子模块:模块内含一个
linalg子模块,包装pytensor.xtensor.linalg中定义的全部线性代数运算; - 模块级再导出:在
pytensor.xtensor模块级定义的操作,可直接在pymc.dims模块级使用。
这三句话背后,是 PyMC 全新的"维度感知"建模范式:传统张量运算按位置轴(axis=0、axis=-1)操作,而pymc.dims体系中的所有变量都携带命名维度(如("batch", "time")),运算按维度名操作。pymc.dims.math正是让普通数学运算获得这种命名维度能力的关键入口。
二、模块实现:两行代码完成的全量包装
打开源码 pymc/dims/math.py,整个模块的实现简洁到极致:
from pytensor.xtensor import linalg from pytensor.xtensor.math import *from pytensor.xtensor.math import *:把pytensor.xtensor.math下定义的全部数学操作(softmax、log、exp、maximum、sqrt、reciprocal等)原样再导出到pymc.dims.math命名空间;from pytensor.xtensor import linalg:把线性代数子模块整体绑定为pymc.dims.math.linalg。
这种"薄包装"设计意味着:pymc.dims.math的函数签名、参数语义与pytensor.xtensor.math保持一致,所有操作都接受并返回携带命名维度的XTensorVariable(xarray 风格变量)。文档中使用的:ref:交叉引用(pytensor:libdoc_xtensor_math、pytensor:libdoc_xtensor_linalg)也印证了这一点——它们指向的是同一套 API 的官方文档。
三、pymc.dims 体系:math 模块的上下文
pymc.dims.math不是孤立存在的,它隶属于 pymc/dims/init.py 构建的完整维度感知生态。该包在导入时做了三件关键事情:
# pymc/dims/__init__.py(节选) from pytensor.tensor import TensorLike from pytensor.xtensor import as_xtensor, broadcast, concat, dot, full_like, ones_like, zeros_like from pytensor.xtensor.basic import tensor_from_xtensor from pytensor.xtensor.type import XTensorVariable from xarray import DataArray XTensorLike = TensorLike | XTensorVariable | DataArray from pymc.dims import math from pymc.dims.distributions import * from pymc.dims.model import Data, Deterministic, Potential- 再导出模块级操作:
as_xtensor、broadcast、concat、dot、full_like、ones_like、zeros_like等直接挂在pymc.dims顶层,即文档中"Operations defined at the module level in pytensor.xtensor are available at the pymc.dims module level"所指的内容; - 定义统一类型别名
XTensorLike = TensorLike | XTensorVariable | DataArray:维度感知 API 同时接受 PyTensor 张量、xtensor 变量与 xarrayDataArray,与pymc.dims.math的运算无缝衔接; - 注册维度感知分布
from pymc.dims.distributions import *(如Normal、MvNormal、Categorical等,详见 distributions.rst)与模型构造器Data、Deterministic、Potential(详见 model.rst)。
3.1 让 PyMC 感知 xtensor 功能
pymc/dims/__init__.py中的__init__()函数负责把 xtensor 能力接入 PyMC 的计算图优化与对数概率(logprob)系统:
MeasurableOp.register(XRV) logprob_rewrites_db.register( "pre_lower_xtensor", optdb.query("+lower_xtensor"), "basic", position=0.1 ) logprob_rewrites_db.register( "post_lower_xtensor", optdb.query("+lower_xtensor"), "cleanup", position=5.1 ) initial_point_rewrites_db.register( "lower_xtensor", optdb.query("+lower_xtensor"), "basic", position=0.1 )这意味着通过pymc.dims创建的带维度随机变量,其采样(XRV)、对数概率(logp/logcdf/icdf)与初始点生成都能被 PyMC 正确识别与计算——维度感知的数学运算也因此能无缝参与模型推断,而不仅是"看上去好看"。
四、从数学运算到带维度变量:as_xtensor 是入口
pymc.dims.math的运算对象是携带命名维度的 xtensor 变量,而创建这种变量的核心入口是as_xtensor(同时也是pymc.dims模块级再导出函数):
from pymc import dims as pmd # 把普通 NumPy 数组提升为带命名维度的 xtensor 变量 group_idx = pmd.math.as_xtensor(group_idx_np, dims=("obs",)) chol_xr = pmd.math.as_xtensor(chol, dims=("core1", "core2"))这一能力在源码中随处可见:
- tests/dims/distributions/test_core.py 中把常规 PyMC 分布的
chol结果用pmx.math.as_xtensor(chol, dims=("core1", "core2"))转成带维度变量后传给pmx.MvNormal; - pymc/dims/distributions/core.py 的
DimDistribution._as_xtensor则规定了反向规则:如果传入的参数不是 xtensor 且无法转换,会抛出明确错误,提示"Usepymc.dims.as_xtensor(..., dims=...)to specify the dims explicitly"。
由此形成完整闭环:as_xtensor创建带维度变量 →pymc.dims.math提供维度感知运算 →pymc.dims分布与模型构造器完成建模。
五、实战对比:传统 pm.math 与维度感知 pmd.math
test_model.py 中的分段线性样条回归(spline)模型是理解两者差异的最佳教材。两个模型数学上完全等价,仅书写方式不同。
5.1 传统写法:位置轴驱动
import pymc as pm delta_factors = pm.math.softmax(z, axis=-1) # (groups, knot) slope_factors = 1 - delta_factors[:, :-1].cumsum(axis=-1) # (groups, knot-1) spline_slopes = pm.math.concatenate([beta0[:, None], beta0[:, None] * slope_factors], axis=-1) beta = pm.math.concatenate([beta0[:, None], pm.math.diff(spline_slopes)], axis=-1) X = pm.math.maximum(0, x[:, None] - knots[None, :]) # (n, knot) mu = (X * beta[group_idx_np]).sum(-1)这里每一步都要手工维护axis、None扩维、索引切片,一旦维度顺序变化极易出错。
5.2 维度感知写法:按维度名驱动
from pymc import dims as pmd delta_factors = pmd.math.softmax(z, dim="knot") # 按名字 softmax slope_factors = 1 - delta_factors.isel(knot=slice(None, -1)).cumsum("knot") spline_slopes = pmd.concat([beta0, beta0 * slope_factors], dim="knot") beta = pmd.concat([beta0, spline_slopes.diff("knot")], dim="knot") X = pmd.math.maximum(0, x - knots) # 广播按维度对齐 mu = (X * beta.isel(group=group_idx)).sum("knot")要点解析:
| 维度感知写法 | 等价传统写法 | 含义 |
|---|---|---|
pmd.math.softmax(z, dim="knot") | pm.math.softmax(z, axis=-1) | 沿命名维度 softmax,无需数轴号 |
var.isel(knot=slice(None, -1)) | var[:, :-1] | 按维度名切片,无需记忆位置 |
var.cumsum("knot")/var.diff("knot") | var.cumsum(axis=-1)/pm.math.diff(var, axis=-1) | 命名维度上的累积和与差分 |
pmd.concat([...], dim="knot") | pm.math.concatenate([...], axis=-1) | 命名维度拼接 |
pmd.math.maximum(0, x - knots) | pm.math.maximum(0, x[:, None] - knots[None, :]) | 广播由维度名自动对齐,x的obs维与knots的knot维自动匹配 |
beta.isel(group=group_idx) | beta[group_idx_np] | 按维度名做高级索引 |
expr.sum("knot") | expr.sum(-1) | 命名维度求和 |
该测试随后验证了两者的初始点(initial_point)、对数概率(compile_logp)与随机抽样(pm.draw)结果完全一致,证明维度感知写法在数值上与传统写法等价,但代码更短、意图更清晰、对维度顺序变更更鲁棒。
5.3 分布定义中的维度语义
同样的维度语义贯穿到分布层。以Normal为例,tests/dims/distributions/test_core.py 展示了丰富的维度推导规则:
with pm.Model(coords=coords) as model: x = pmx.Data("x", np.random.randn(2, 3, 5), dims=("a", "b", "c")) y2 = pmx.Normal("y2", mu=x, dims=("a", "b", "c")) # 冗余但合法 y3 = pmx.Normal("y3", mu=x, dims=("b", "a", "c")) # 隐式转置 → shape (3, 2, 5) y4 = pmx.Normal("y4", mu=x, dims=("a", ...)) # 省略号补齐其余维度 y8 = pmx.Normal("y8", mu=x, dims=("d", "a", "b", "c")) # 新增额外维度 (7, 2, 3, 5)其底层实现位于 pymc/dims/distributions/core.py 的DimDistribution:分布输出经由rv.transpose(*dims)对齐到用户指定顺序;参数若缺少 dims 会被_as_xtensor拦截并报错,避免维度假设带来的隐性 bug。这正是文档中强调"维度感知"的核心动机——显式声明维度,杜绝隐式假设。
六、linalg 子模块:带维度的线性代数
按文档所述,pymc.dims.math.linalg包装了pytensor.xtensor.linalg的全部线性代数操作。在维度感知建模中,典型的应用场景包括:
- 协方差矩阵构造:把 Cholesky 分解结果、协方差矩阵等按命名维度(如
("core1", "core2"))组织,使其维度信息随变量一起传递; - 矩阵运算:
dot(同时作为模块级再导出存在)、求逆、行列式、特征分解等操作在矩阵的两个命名维度上语义化执行,不再依赖axis的"哪个是行、哪个是列"的记忆负担。
线性代数是多元分布(MvNormal、LKJ等)建模的基石,vector.py 中Categorical、MvNormal等向量分布正是通过core_dims机制(如core_dims=("core1", "core2"))把核心维度与批处理维度区分开,参见 tests/dims/distributions/test_core.py。pymc.dims.math.linalg保证这些核心维度上的运算全程可追踪。
七、模块级再导出:pymc.dims 顶层的快捷操作
文档第三条要点——"Operations defined at the module level in pytensor.xtensor are available at the pymc.dims module level"——对应 pymc/dims/init.py 中的直接导入:
from pytensor.xtensor import as_xtensor, broadcast, concat, dot, full_like, ones_like, zeros_like from pytensor.xtensor.basic import tensor_from_xtensor from pytensor.xtensor.type import XTensorVariable因此以下写法完全等价(以拼接为例):
from pymc import dims as pmd beta = pmd.concat([beta0, spline_slopes.diff("knot")], dim="knot") # 等价于: beta = pmd.math.concat([beta0, spline_slopes.diff("knot")], dim="knot")concat、dot属于日常高频操作,放在pymc.dims顶层(与Normal、Data、Deterministic、Potential并列)是为了书写便利;其余数学运算则统一收敛在pymc.dims.math命名空间,保持 API 的组织清晰。
八、常见陷阱与最佳实践
- 参数必须显式携带维度:
DimDistribution._as_xtensor(core.py)会在参数缺少 dims 时抛错,提示用as_xtensor(..., dims=...)显式声明。这是刻意设计:维度推导失败要尽早暴露,而不是静默产生错误广播。 - dims 必须属于模型 coords:在
DimDistribution.__new__(core.py)中,未在模型 coords 中注册的维度会抛出ValueError,提示先用model.add_coord或在Model(coords=...)中声明。 - dims 必须与分布输出完全匹配或用省略号:
test_core.py展示了缺维(dims=("a", "b")vs 输出("a","b","c"))、仅指定额外维(dims=("d",))等错误场景,正确做法是使用("a", ...)形式的省略号语法。 - 向量分布必须提供 core_dims:
VectorDimDistribution.dist(core.py)强制要求core_dims,因为涉及非标量输入/输出,缺失时会给出明确报错。 - 转换器是维度感知的:
pymc.dims使用DimTransform体系(LogTransform、LogOddsTransform、SimplexTransform、ZeroSumTransform等,见 transforms.rst 与 transforms.py),而非传统Transform,以确保变换的对数雅可比在命名维度上正确计算。
九、总结
pymc.dims.math是 PyMC 维度感知建模体系中的数学运算支柱,其设计可以概括为:
- 全量薄包装:两行源码把
pytensor.xtensor.math的全部运算与linalg子模块引入pymc.dims命名空间,API 与上游保持一致; - 模块级再导出:
as_xtensor、concat、dot、broadcast等高频操作在pymc.dims顶层可用; - 命名维度驱动:配合
as_xtensor、DimDistribution与Data/Deterministic/Potential包装器,实现"按维度名而非位置轴"编写模型; - 全链路集成:通过
MeasurableOp.register(XRV)与logprob_rewrites_db/initial_point_rewrites_db注册,维度感知变量可无缝参与采样、对数概率与初始点计算。
对于希望降低贝叶斯模型维护成本、提升代码可读性与鲁棒性的开发者,从pymc.dims.math入手逐步迁移现有模型,是一条低风险、高回报的路径。进一步的分布与模型构造器 API 可查阅 distributions.rst、model.rst 与 transforms.rst。
【免费下载链接】pymcBayesian Modeling and Probabilistic Programming in Python项目地址: https://gitcode.com/GitHub_Trending/py/pymc
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考