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

在Numpyro中估计部分观测参数向量的后验分布

解决方案:用解析后验实现精准高效的后验采样

核心思路

你的问题本质是带线性硬约束的高斯先验后验推断,因为先验是独立高斯,观测是线性变换的部分维度,后验属于条件高斯分布,有精确的解析解。直接利用这个性质建模,比numpyro.factor()或虚构噪声的方案更精准、高效,完全避免近似误差。

具体实现方案

方案1:直接用解析后验的高斯分布采样

数学推导

设:

  • M维参数 $P \sim \mathcal{N}(0, I_M)$($I_M$是M维单位矩阵)
  • 观测变换矩阵 $R \in \mathbb{R}^{K \times M}$(K是完整观测空间维度)
  • 取R中对应观测维度的行组成子矩阵 $R_{\text{obs}}$(比如只取第2行,对应D=1维观测),观测值为 $y_{\text{obs}}$

根据条件高斯分布公式,后验 $P(P | R_{\text{obs}}P = y_{\text{obs}})$ 的均值和协方差可直接计算:

  • 后验均值:$\mu_{\text{post}} = R_{\text{obs}}^T (R_{\text{obs}} R_{\text{obs}}T){-1} y_{\text{obs}}$
  • 后验协方差:$\Sigma_{\text{post}} = I_M - R_{\text{obs}}^T (R_{\text{obs}} R_{\text{obs}}T){-1} R_{\text{obs}}$

Numpyro代码实现

import numpyro
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS
import jax.numpy as jnp
import jax

def model(R_obs, y_obs):
    M = R_obs.shape[1]
    # 计算后验的均值和协方差(用solve替代直接求逆,提升数值稳定性)
    RRT = jnp.dot(R_obs, R_obs.T)
    inv_RRT_times_y = jax.scipy.linalg.solve(RRT, y_obs)
    mu_post = jnp.dot(R_obs.T, inv_RRT_times_y)
    
    RTR = jnp.dot(R_obs.T, R_obs)
    sigma_post = jnp.eye(M) - jnp.dot(R_obs.T, jax.scipy.linalg.solve(RRT, R_obs))
    
    # 采样后验参数
    P = numpyro.sample("P", dist.MultivariateNormal(loc=mu_post, covariance_matrix=sigma_post))
    # 生成完整观测空间的样本(替换成你的完整R矩阵即可)
    full_R = jnp.array([[1, 0, 0], [0, 1, 0], [0, 0, 1]])
    numpyro.deterministic("full_observation", jnp.dot(full_R, P))

# 示例:M=3,观测第2维,观测值为2.0
R_obs = jnp.array([[0, 1, 0]])
y_obs = jnp.array([2.0])

# 运行MCMC(解析高斯分布采样效率极高)
nuts_kernel = NUTS(model)
mcmc = MCMC(nuts_kernel, num_warmup=300, num_samples=1000)
mcmc.run(jax.random.PRNGKey(42), R_obs=R_obs, y_obs=y_obs)
mcmc.print_summary()

# 获取结果
posterior_P = mcmc.get_samples()["P"]
full_obs_samples = mcmc.get_samples()["full_observation"]

方案2:参数分解(适合高维场景,避免协方差求逆)

如果M很大,直接计算协方差矩阵可能有数值稳定性问题,可以把参数P拆成受约束的固定部分和自由采样部分:

  1. 对$R_{\text{obs}}$做SVD分解,得到行空间正交基$V$和零空间正交基$Q$(满足$R_{\text{obs}}Q=0$)
  2. 令$P = V\alpha + Q\beta$,其中$\alpha$是受约束的D维变量,$\beta$是M-D维自由变量
  3. 约束$R_{\text{obs}}P=y_{\text{obs}}$等价于$\alpha = (R_{\text{obs}}V)^{-1}y_{\text{obs}}$,$\beta$仍服从先验的$\mathcal{N}(0, I_{M-D})$

Numpyro代码实现

def model_param_decompose(R_obs, y_obs):
    M = R_obs.shape[1]
    D = R_obs.shape[0]
    # SVD分解获取行空间和零空间基
    _, _, Vh = jnp.linalg.svd(R_obs, full_matrices=True)
    V = Vh[:D, :].T  # 行空间基,M×D
    Q = Vh[D:, :].T  # 零空间基,M×(M-D)
    
    # 计算受约束的alpha(固定值)
    alpha = jax.scipy.linalg.solve(jnp.dot(R_obs, V), y_obs)
    # 自由采样beta
    beta = numpyro.sample("beta", dist.MultivariateNormal(loc=jnp.zeros(M-D), covariance_matrix=jnp.eye(M-D)))
    
    # 重构参数P
    P = jnp.dot(V, alpha) + jnp.dot(Q, beta)
    numpyro.deterministic("P", P)
    # 生成完整观测
    full_R = jnp.array([[1,0,0],[0,1,0],[0,0,1]])
    numpyro.deterministic("full_observation", jnp.dot(full_R, P))

方案优势对比

  • 对比虚构噪声:完全没有近似偏差,不需要调整噪声方差(哪怕方差设极小也会引入数值不稳定)
  • 对比numpyro.factor():factor()是通过添加极大负对数密度模拟硬约束,本质是软约束,采样时容易出现数值问题且效率低;而解析方法直接生成符合约束的样本,采样效率极高,结果完全精准

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 12:05:03