TensorFlow 2自定义多元正态损失函数报InvalidArgumentError错误求解
错误原因
- 维度拼接逻辑错误:你使用Python列表嵌套的方式传入
loc和covariance_matrix参数时,TF会将列表维度添加到张量的最前面,而非事件维度所在的最后面。以batch_size=32、序列长度为10的场景为例,你的mu1、mu2形状为(32, 10, 1),用[mu1, mu2]得到的loc形状为(2, 32, 10, 1),完全不符合MultivariateNormalFullCovariance要求的loc最后一维为事件维度(此处为2)、covariance_matrix最后两维为(2,2)方阵的格式,因此触发"Input matrix must be square"报错。 - 隐含的协方差矩阵合法性问题:当前代码未对协方差项sigma12做约束,很容易出现
sigma12^2 > sigma11 * sigma22的情况,导致协方差矩阵非半正定,后续Cholesky分解依然会报错。
解决方案
直接修改损失函数中分布的构造逻辑,正确堆叠张量维度,并增加协方差矩阵半正定约束,修改后的核心代码如下:
def negative_normdist_loss_2(y_true, y_pred): # Separate the parameters mu1, mu2, sigma11, sigma12, sigma22 = tf.unstack(y_pred, num=5, axis=-1) # 对对角线方差项做softplus保证为正 sigma11 = tf.keras.activations.softplus(tf.expand_dims(sigma11, -1)) sigma22 = tf.keras.activations.softplus(tf.expand_dims(sigma22, -1)) # 对相关系数做tanh约束到[-1,1]区间,保证协方差矩阵半正定 rho = tf.keras.activations.tanh(tf.expand_dims(sigma12, -1)) sigma12 = rho * tf.sqrt(sigma11 * sigma22) # 正确构造loc:最后一维为事件维度2,形状为(..., 2) loc = tf.concat([tf.expand_dims(mu1, -1), tf.expand_dims(mu2, -1)], axis=-1) # 正确构造协方差矩阵:最后两维为2x2方阵,形状为(..., 2, 2) cov_row1 = tf.concat([sigma11, sigma12], axis=-1) cov_row2 = tf.concat([sigma12, sigma22], axis=-1) covariance_matrix = tf.stack([cov_row1, cov_row2], axis=-2) # 计算负对数似然 dist = tfp.distributions.MultivariateNormalFullCovariance( loc=loc, covariance_matrix=covariance_matrix ) nll = tf.reduce_mean(-dist.log_prob(y_true)) return nll
另外请额外确认:你的y_true最后一维必须为2,对应两条时间序列的真实值,否则会出现维度不匹配报错。
内容的提问来源于stack exchange,提问作者user10176535
相关产品推荐
相关产品推荐

