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

如何使用PyTorch从多元正态分布指定概率区域采样

关键认知澄清

你觉得log_prob()返回值不符合预期,核心是混淆了连续分布的概率密度和离散概率:

  • 多元正态是连续分布,不存在单点概率P(x),dist.log_prob(x)返回的是对数概率密度值,数值大小和特征维度、协方差矩阵的尺度直接绑定,不能直接和0.05这类概率阈值比较。
  • 连续分布的低概率区域不需要直接用概率密度判定:d维多元正态分布中,样本到均值的马氏距离平方(x-μ)^TΣ^{-1}(x-μ)服从自由度为d的卡方分布χ²(d),概率密度越低的样本,马氏距离越大,用这个指标做筛选数值稳定性远高于直接用概率密度。
具体实现方法

整个流程不需要复杂操作,分三步即可:

  1. 协方差矩阵正则:从ResNet表征算出来的经验协方差通常是半正定的,直接传入MultivariateNormal很容易触发数值错误,采样前加一个极小的单位矩阵做扰动即可。
  2. 计算分位阈值:自由度为特征维度的卡方分布的95%分位点,就是高低概率区域的分界——95%的正常样本马氏距离平方小于该值,剩下5%就是你需要的低概率尾部样本。
  3. 批量计算样本马氏距离,筛选符合阈值要求的样本即可。
可直接复用的代码
import torch
from torch.distributions.multivariate_normal import MultivariateNormal
from torch.distributions.chi2 import Chi2

# 输入参数:假设prototype形状为(1, feature_size),cov形状为(feature_size, feature_size)
feature_size = cov.shape[-1]
device = cov.device
dtype = cov.dtype

# 协方差正则,避免非正定导致的采样错误
cov_reg = cov + 1e-6 * torch.eye(feature_size, device=device, dtype=dtype)
# prototype压缩为一维,避免分布维度和batch维度混淆
proto_1d = prototype.squeeze()

# 初始化多元正态分布
dist = MultivariateNormal(proto_1d, covariance_matrix=cov_reg)

# 按需求采样,示例为采10000个样本
n_sample = 10000
samples = dist.sample(torch.Size([n_sample]))

# 计算5%低概率区域对应的马氏距离平方阈值
chi2 = Chi2(df=feature_size)
threshold = chi2.icdf(torch.tensor(0.95, device=device, dtype=dtype)).item()

# 批量计算所有样本的马氏距离平方,无循环
diff = samples - proto_1d
mahal_sq = torch.sum(diff @ dist.precision_matrix * diff, dim=-1)

# 筛选得到低概率样本
low_prob_samples = samples[mahal_sq > threshold]
补充说明

如果你的prototype保留(1, feature_size)的形状不做压缩,MultivariateNormal会默认把第一维识别为分布的batch维度,返回的log_prob、协方差逆矩阵的维度都会和预期不符,这也是很多人觉得接口返回值异常的常见诱因。

如果你坚持要用log_prob做筛选也完全可行:多元正态的对数概率密度和马氏距离平方是线性负相关关系,公式为:
log_prob(x) = -0.5 * (d*log(2π) + logdet(Σ) + mahal_sq(x))
你只需要计算所有样本的log_prob,保留数值最小的5%样本即可,结果和马氏距离筛选完全一致,但马氏距离方法不需要计算协方差行列式,在特征维度较高时不会出现数值溢出问题,稳定性更好。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 06:33:29