You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用Einstein Sum时PyMC3模型触发ValueError错误如何解决

问题原因

PyMC3 3.x版本底层依赖Theano张量计算框架,模型上下文内定义的变量都是Theano张量类型,而非Numpy数组。np.einsum是Numpy专属接口,只能处理Numpy数组,无法正确解析Theano张量的维度结构,因此会抛出维度匹配错误;你在模型上下文外执行时输入X是Numpy数组,所以可以正常运行。

解决方法

方法1:替换为Theano自带的einsum接口

theano.tensor.einsum的语法规则和Numpy的einsum完全一致,直接替换即可:
首先导入依赖:

import theano.tensor as tt

修改后的模型代码:

with pm.Model() as model:
    gamma= pm.Normal("gamma", mu=0, sigma=100, shape=())
    beta= pm.Normal("beta", mu=0, sigma=100, shape=(X.shape[2], 1))
    u= pm.Normal("u", mu=0, sigma=1, shape=(1, X.shape[0]))
    r = pm.Gamma("r", alpha=9, beta=4, shape=())

    # 将np.einsum替换为tt.einsum
    y_hat = gamma + tt.einsum("jik,kl->ij", X, beta) + u

    y_like = pm.Normal("y_like", mu=y_hat, sigma=r, observed=y)

方法2:用tensordot替代einsum实现相同逻辑

如果担心einsum的兼容性,也可以用tt.tensordot配合维度调整实现相同计算效果,计算效率通常更高:

with pm.Model() as model:
    gamma= pm.Normal("gamma", mu=0, sigma=100, shape=())
    beta= pm.Normal("beta", mu=0, sigma=100, shape=(X.shape[2], 1))
    u= pm.Normal("u", mu=0, sigma=1, shape=(1, X.shape[0]))
    r = pm.Gamma("r", alpha=9, beta=4, shape=())

    # 按K维度相乘后调整维度
    dot_res = tt.tensordot(X, beta, axes=([2], [0])) # 输出形状为(N, T, 1)
    dot_res = tt.squeeze(dot_res).T # 压缩多余维度后转置,得到(T, N)形状
    y_hat = gamma + dot_res + u

    y_like = pm.Normal("y_like", mu=y_hat, sigma=r, observed=y)

方法3:提前对输入做维度预处理

可以在进入模型上下文前把三维的X转为二维数组,直接用矩阵乘法计算,逻辑更直观:

# 预处理:X原本形状(N, T, K),转置后拉平为(T*N, K)
X_reshaped = X.transpose(1, 0, 2).reshape(-1, X.shape[2])

with pm.Model() as model:
    gamma= pm.Normal("gamma", mu=0, sigma=100, shape=())
    beta= pm.Normal("beta", mu=0, sigma=100, shape=(X.shape[2],))
    u= pm.Normal("u", mu=0, sigma=1, shape=(X.shape[0],))
    r = pm.Gamma("r", alpha=9, beta=4, shape=())

    # 计算后调整回(T, N)形状
    dot_res = tt.dot(X_reshaped, beta).reshape(y.shape)
    y_hat = gamma + dot_res + u

    y_like = pm.Normal("y_like", mu=y_hat, sigma=r, observed=y)

内容的提问来源于stack exchange,提问作者Michael

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.29 14:57:03