TF2多分类场景下多元正态分布样本概率计算报错如何解决
核心问题疏漏
- 多元正态分布参数维度匹配错误:x维多元正态分布的均值
loc应为长度为x的向量,对角协方差对应的scale_diag也应为长度为x的向量。你当前对每个类别的[N, x]样本全局求均值/标准差,得到的是单个标量,相当于把所有特征的统计量合并为一个值,构建的本质是c个单变量正态分布,自然不支持输入x维样本计算概率。 - 批分布维度逻辑错误:你需要为每个类别构建一个x维多元正态,对应的
loc参数形状应为[c, x],scale_diag参数形状也应为[c, x],这样TFP会自动生成c个独立的x维多元正态分布批,后续可直接计算样本属于每个类别的概率。
正确实现方案
实现逻辑
- 按类别计算每个特征维度的独立均值和标准差,得到形状为
[c, x]的均值矩阵和标准差矩阵 - 用上述参数构建
MultivariateNormalDiag分布批 - 输入x维样本计算概率,输出结果为长度为c的向量,对应样本属于每个类别对应多元正态分布的概率
修正后代码
import tensorflow as tf import tensorflow_probability as tfp tfd = tfp.distributions # 假设原始数据为data,形状为[c, N, x],c为类别数,N为每类样本数,x为特征维度 c, N, x = data.shape # 1. 按类别对样本维度求平均/标准差,保留类别和特征维度,输出形状均为[c, x] mean_vec = tf.reduce_mean(data, axis=1) # 加极小值避免某特征所有样本取值相同导致标准差为0的数值问题 std_vec = tf.math.reduce_std(data, axis=1) + 1e-6 # 2. 构建c个独立的x维对角协方差多元正态分布批 distr = tfd.MultivariateNormalDiag( loc=mean_vec, scale_diag=std_vec ) # 3. 计算x维样本属于每个类别的概率,y形状为[x]时输出prob_per_class形状为[c] y = tf.random.normal(shape=(x,)) # 随机生成的x维样本 prob_per_class = distr.prob(y)
如果需要批量处理多个样本,输入y的形状为[batch_size, x]时,输出概率形状为[batch_size, c],对应每个样本属于每个类别的概率。
提示:如果是做分类任务,可将上述得到的每个类别的概率乘以对应类别的先验概率,再取最大值作为分类结果。
内容的提问来源于stack exchange,提问作者TheEnigmist
相关产品推荐
相关产品推荐

