如何使用PyTorch从多元正态分布指定概率区域采样
关键认知澄清
你觉得log_prob()返回值不符合预期,核心是混淆了连续分布的概率密度和离散概率:
- 多元正态是连续分布,不存在单点概率
P(x),dist.log_prob(x)返回的是对数概率密度值,数值大小和特征维度、协方差矩阵的尺度直接绑定,不能直接和0.05这类概率阈值比较。 - 连续分布的低概率区域不需要直接用概率密度判定:d维多元正态分布中,样本到均值的马氏距离平方
(x-μ)^TΣ^{-1}(x-μ)服从自由度为d的卡方分布χ²(d),概率密度越低的样本,马氏距离越大,用这个指标做筛选数值稳定性远高于直接用概率密度。
具体实现方法
整个流程不需要复杂操作,分三步即可:
- 协方差矩阵正则:从ResNet表征算出来的经验协方差通常是半正定的,直接传入
MultivariateNormal很容易触发数值错误,采样前加一个极小的单位矩阵做扰动即可。 - 计算分位阈值:自由度为特征维度的卡方分布的95%分位点,就是高低概率区域的分界——95%的正常样本马氏距离平方小于该值,剩下5%就是你需要的低概率尾部样本。
- 批量计算样本马氏距离,筛选符合阈值要求的样本即可。
可直接复用的代码
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
相关产品推荐
相关产品推荐

