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

Scipy/Python中4维多元正态分布的维度顺序问题排查

4维高斯分布可视化中心偏移问题的解决方案

问题背景

需要在4维网格(变量为x1,y1,x2,y2)上计算4维高斯分布,设定均值为(x1=1, y1=0, x2=2, y2=0)。目标是绘制y1=y2=0时x1-x2方向的2D等高线图,期望看到中心位于(1,2)的高斯分布,但实际绘图显示中心为(2,0)。手动重塑点矩阵后结果仍一致,问题根源在于网格定义的索引方式。

初始错误代码

import numpy as np
from matplotlib import pyplot as plt
from scipy.stats import multivariate_normal


xy_min = -5
xy_max = 5
npoints = 50
x = np.linspace(xy_min, xy_max, npoints)
dim = 4
xx1,yy1,xx2,yy2 = np.meshgrid(x, x,x,x)
points = np.concatenate([xx1[:, :,:, :,None], yy1[:, :, :,:,None],xx2[:, :, :,:,None],yy2[:, :, :,:,None]], axis=-1)

cov = np.diag(np.ones(4))
mean=np.array([1,0,2,0])
rv = multivariate_normal.pdf(points , mean=mean, cov=cov)

plt.figure()
plt.contourf(x, x, rv[:,0,:,0])

尝试的重塑代码(未解决问题)

points_resh = np.reshape(points,[npoints**4,dim],order='C')
rv_resh = multivariate_normal.pdf(points_resh , mean=mean, cov=cov)
rv2 = np.reshape(rv_resh,[npoints,npoints,npoints,npoints],order='C')

plt.figure()
plt.contourf(x, x, rv2[:,0,:,0])

错误原因

numpy.meshgrid默认使用indexing='xy'(笛卡尔索引),这种索引方式会将第一个输入数组映射到最后一个维度,第二个输入数组映射到倒数第二个维度,以此类推。在4维场景下,这种轴顺序会导致xx1、xx2等变量的维度映射和预期不符,提取y1=y2=0的切片时,实际对应的变量维度错位,最终导致可视化的中心位置偏移。

解决方法

使用meshgrid的indexing='ij'(矩阵索引)模式,该模式会让输入数组的顺序和输出网格的轴顺序一一对应(第一个输入对应第一个轴,第二个输入对应第二个轴,以此类推)。同时,由于matplotlib的contourf默认使用笛卡尔坐标,需要对提取的4维分布切片进行转置,以匹配绘图的坐标逻辑。

正确示例代码

import numpy as np
from matplotlib import pyplot as plt
from scipy.stats import multivariate_normal

# 定义各维度的取值范围和点数
x = np.linspace(-5, 5, 50)    # x1维度
y = np.linspace(-3, 3, 30)    # y1维度
z = np.linspace(-2, 2, 20)    # x2维度
w = np.linspace(-1, 1, 10)    # y2维度

# 使用ij索引创建4维网格
x4d, y4d, z4d, w4d = np.meshgrid(x, y, z, w, indexing='ij')
# 拼接成形状为(nx, ny, nz, nw, 4)的点矩阵
points4d = np.concatenate([x4d[..., None], y4d[..., None], z4d[..., None], w4d[..., None]], axis=-1)

# 计算4维高斯分布概率密度
rv4d = multivariate_normal.pdf(points4d, mean=[1.0, 0.0, 2.0, 0.0], cov=np.diag([0.1, 0.1, 0.1, 0.1]))

# 绘制y1=0、y2=0时x1-x2的等高线图,注意转置切片
fig, ax = plt.subplots()
ax.contourf(x, z, rv4d[:, 0, :, 0].T)
ax.set(xlabel='x1', ylabel='x2')
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 12:15:40