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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.23 21:24:05