Keras的weighted_metrics未纳入样本权重计算的解决方法咨询
问题根源
你遇到的指标不生效的问题来自两个核心原因:
- 你直接传入
tf.keras.losses.categorical_crossentropy作为指标,而非tf.keras.metrics下的指标类实例,Keras的自动包装逻辑不会处理样本权重与输出维度的匹配。 - 你的模型输出是形状为
(None, 400, 22)的3D序列张量,而传入的样本权重是(None,)的1D样本级权重,维度不匹配导致权重无法被自动应用到每个时间步的损失计算中。
解决方案
方案1:自定义加权分类交叉熵指标(推荐)
手动实现指标类处理权重广播逻辑,适配你的序列输出场景:
import tensorflow as tf class WeightedCategoricalCrossentropy(tf.keras.metrics.Metric): def __init__(self, name='weighted_categorical_crossentropy', **kwargs): super().__init__(name=name, **kwargs) # 定义累加变量:总加权损失、总权重和 self.total_loss = self.add_weight(name='total_loss', initializer='zeros') self.total_weight = self.add_weight(name='total_weight', initializer='zeros') def update_state(self, y_true, y_pred, sample_weight=None): # 计算每个时间步的交叉熵,输出形状为 (None, 400) base_cce = tf.keras.losses.categorical_crossentropy(y_true, y_pred) # 对每个样本的所有时间步求平均,得到每个样本的平均交叉熵,形状为 (None,) sample_cce = tf.reduce_mean(base_cce, axis=-1) # 应用样本权重 if sample_weight is not None: sample_weight = tf.cast(sample_weight, dtype=sample_cce.dtype) weighted_cce = sample_cce * sample_weight self.total_weight.assign_add(tf.reduce_sum(sample_weight)) else: weighted_cce = sample_cce self.total_weight.assign_add(tf.cast(tf.shape(y_true)[0], dtype=self.total_weight.dtype)) self.total_loss.assign_add(tf.reduce_sum(weighted_cce)) def result(self): # 返回加权平均后的交叉熵 return self.total_loss / self.total_weight def reset_state(self): # 每个epoch重置累加变量 self.total_loss.assign(0.0) self.total_weight.assign(0.0)
编译模型时将自定义指标传入weighted_metrics即可:
model.compile( optimizer=tf.keras.optimizers.Adam(1e-3), loss=tf.keras.losses.categorical_crossentropy, weighted_metrics=[WeightedCategoricalCrossentropy()] )
注意:最终loss和指标的微小差异来自L2正则项,因为loss会累加正则损失,而该指标仅计算加权交叉熵,若要完全对齐可去掉正则项验证。
方案2:预广播样本权重
如果不想自定义指标,可提前把样本权重广播到和输出时间步维度匹配,让Keras能自动对应权重:
# 原样本权重形状为 (20,) sample_weight = np.random.choice([0.01, 0.1, 1], size=20) # 广播为 (20, 400),每个样本的400个时间步共享同一个权重 sample_weight = np.repeat(sample_weight[:, np.newaxis], 400, axis=1)
编译时直接使用内置指标类即可:
model.compile( optimizer=tf.keras.optimizers.Adam(1e-3), loss=tf.keras.losses.categorical_crossentropy, metrics=[tf.keras.metrics.CategoricalCrossentropy()] )
内容的提问来源于stack exchange,提问作者ProteinGuy
相关产品推荐
相关产品推荐

