SciPy.stats分布类如何灵活处理NumPy数组?如何实现同款兼容逻辑?
SciPy数组适配逻辑与NormalGamma分布实现示例
一、SciPy灵活适配任意形状数组的核心逻辑
核心是固定分量轴位置+批量维度自动对齐,一共4个步骤:
- 约定分量轴位置:要求输入数组的最后一个轴对应分布的分量维度,比如k元正态分布要求最后一个轴长度为k;NormalGamma是正态变量x和精度τ的联合分布,要求最后一个轴长度为2。除最后一个轴之外的所有维度都属于批量维度,可以是任意形状。
- 维度拉平:把输入数组除最后一个轴外的所有维度合并成一个一维的批量维度,不管原来的批量维度是(10,)还是(20,20),拉平后都是
(批量大小, 分量维度)的二维结构,统一计算逻辑。 - 向量化计算:所有PDF计算用NumPy元素级运算实现,天然支持批量维度,不需要手动写循环。
- 形状恢复:计算完成后把结果的批量维度恢复为输入时的原始批量形状,保证输出形状和输入的批量维度完全匹配。
二、NormalGamma分布简化实现示例
公式说明
单变量NormalGamma是随机变量$x$(正态分布)和$\tau$(Gamma分布,代表正态分布的精度)的联合分布,参数为$\mu$(正态均值)、$\lambda$(正态精度权重)、$\alpha$(Gamma形状参数)、$\beta$(Gamma速率参数),PDF公式为:
$$
p(x,\tau) = \frac{\beta^\alpha \sqrt{\lambda}}{\Gamma(\alpha)\sqrt{2\pi}} \tau^{\alpha-1/2} \exp\left(-\beta\tau - \frac{\lambda\tau(x-\mu)^2}{2}\right)
$$
代码实现
import numpy as np from scipy.special import gamma class NormalGamma: def __init__(self, mu, lam, alpha, beta): # 参数转为numpy数组,支持广播适配 self.mu = np.asarray(mu) self.lam = np.asarray(lam) self.alpha = np.asarray(alpha) self.beta = np.asarray(beta) # 联合分布分量维度为2:[x, tau] self.component_dim = 2 def pdf(self, x): x = np.asarray(x) # 兼容未显式加分量轴的1维输入 if x.ndim == 1 and x.shape[0] != self.component_dim: x = x[..., np.newaxis] # 校验最后一个轴是否匹配分量维度 if x.shape[-1] != self.component_dim: raise ValueError(f"输入最后一个轴长度必须为{self.component_dim}") # 保存原始批量形状 batch_shape = x.shape[:-1] # 拉平为[批量大小, 分量维度]的统一结构 x_flat = x.reshape(-1, self.component_dim) # 拆分分量 x_val = x_flat[:, 0] tau_val = x_flat[:, 1] # 向量化计算PDF norm_const = (self.beta ** self.alpha * np.sqrt(self.lam)) / (gamma(self.alpha) * np.sqrt(2 * np.pi)) power_term = tau_val ** (self.alpha - 0.5) exp_term = np.exp(-self.beta * tau_val - 0.5 * self.lam * tau_val * (x_val - self.mu) ** 2) pdf_flat = norm_const * power_term * exp_term # 恢复原始批量形状返回 return pdf_flat.reshape(batch_shape)
兼容测试用例
用例1:1维批量输入
# 初始化分布 rv = NormalGamma(mu=0, lam=1, alpha=2, beta=2) # 生成10个样本,x固定为1,tau从0到5取值,输入形状为(10, 2) tau_arr = np.linspace(0, 5, 10, endpoint=False) x_input = np.column_stack([np.full_like(tau_arr, 1), tau_arr]) # 输出形状为(10,),和输入批量维度匹配 pdf_1d = rv.pdf(x_input) print(pdf_1d.shape) # 输出 (10,)
用例2:3维网格输入
# 生成x和tau的二维网格 x_grid, tau_grid = np.mgrid[-3:3:0.1, 0.1:5:0.1] # 合并为3维输入,形状为(60, 49, 2),最后一个轴为分量轴 pos = np.dstack((x_grid, tau_grid)) # 输出形状为(60, 49),和网格批量维度匹配,可直接用于等高线绘制 pdf_grid = rv.pdf(pos) print(pdf_grid.shape) # 输出 (60, 49)
内容的提问来源于stack exchange,提问作者lpnorm
相关产品推荐
相关产品推荐

