PyTorch中如何实现数值稳定的非中心卡方分布?
在PyTorch中实现数值稳定的非中心卡方分布
非中心卡方分布描述了多个独立非零均值正态变量的平方和的分布,PyTorch原生未提供该分布的实现。以下是针对概率密度函数(PDF)计算和采样场景的数值稳定实现方案。
一、数值稳定的PDF计算
非中心卡方的PDF可表示为泊松加权的中心卡方PDF混合:
$f(x; df, \lambda) = \sum_{k=0}^\infty \frac{e^{-\lambda/2} (\lambda/2)^k}{k!} \cdot f_{\chi^2}(x; df + 2k)$
其中$f_{\chi^2}$是中心卡方分布的PDF,$df$为自由度,$\lambda$为非中心参数。
直接求和易出现数值下溢/溢出,因此我们在对数空间中计算每一项,再用torch.logsumexp合并结果,避免精度损失:
import torch from torch.distributions import Chi2 def non_central_chi2_pdf(x, df, lam): # 统一转为张量格式 x = torch.as_tensor(x) df = torch.as_tensor(df) lam = torch.as_tensor(lam) # x<=0时PDF为0 pdf = torch.zeros_like(x) valid_mask = x > 0 x_valid = x[valid_mask] df_valid = df[valid_mask] if df.ndim > 0 else df lam_valid = lam[valid_mask] if lam.ndim > 0 else lam # 动态确定截断项数:基于泊松分布6σ原则,保证截断误差<1e-12 max_k = torch.ceil(lam_valid + 8 * torch.sqrt(lam_valid)).int() max_k = torch.max(max_k, torch.tensor(20, device=x.device)) # 兜底最小项数 # 生成k的范围(0到max_k) k = torch.arange(0, max_k.max().item() + 1, device=x.device) k = k.unsqueeze(0).expand(len(x_valid), -1) # 计算每一项的对数形式,避免直接计算导致的数值问题 log_poisson_term = -lam_valid.unsqueeze(1)/2 + k * torch.log(lam_valid.unsqueeze(1)/2) - torch.lgamma(k + 1) chi2_dist = Chi2(df_valid.unsqueeze(1) + 2 * k) log_chi2_pdf = chi2_dist.log_prob(x_valid.unsqueeze(1)) # 合并所有项的对数结果 log_total = torch.logsumexp(log_poisson_term + log_chi2_pdf, dim=1) pdf[valid_mask] = torch.exp(log_total) return pdf
核心稳定性优化:
- 用对数空间计算泊松项和中心卡方PDF,避免大λ时$e^{-\lambda/2}$的下溢问题
- 动态截断求和项,在精度和计算效率间取得平衡
- 用
torch.logsumexp替代普通求和,防止小数值相加时的精度丢失
二、数值稳定的采样实现
非中心卡方分布可通过泊松-中心卡方混合采样,天然避免数值问题:
- 采样泊松变量$Y \sim \text{Poisson}(\lambda/2)$
- 基于$Y$采样中心卡方变量$Z \sim \chi^2(df + 2Y)$
- $Z$即为非中心卡方分布的样本
def non_central_chi2_sample(df, lam, sample_shape=torch.Size()): df = torch.as_tensor(df) lam = torch.as_tensor(lam) # 采样泊松变量Y poisson_rate = lam / 2 y = torch.poisson(poisson_rate.expand(sample_shape)) # 采样对应自由度的中心卡方样本 chi2_df = df + 2 * y samples = torch.distributions.Chi2(chi2_df).sample() return samples
验证示例
对比PyTorch实现与scipy的结果,确保数值一致性:
import scipy.stats as stats # 测试PDF一致性 x = torch.tensor([1.0, 5.0, 10.0]) df = 3.0 lam = 2.0 torch_pdf = non_central_chi2_pdf(x, df, lam) scipy_pdf = torch.tensor(stats.ncx2.pdf(x.numpy(), df, lam)) print(torch.allclose(torch_pdf, scipy_pdf, rtol=1e-6)) # 应输出True # 测试采样均值一致性 samples = non_central_chi2_sample(df, lam, sample_shape=(10000,)) scipy_samples = stats.ncx2.rvs(df, lam, size=10000) print(torch.allclose(samples.mean(), torch.tensor(scipy_samples.mean()), rtol=1e-2)) # 均值接近
内容的提问来源于stack exchange,提问作者Kemsikov
相关产品推荐
相关产品推荐

