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

Keras模型自定义损失函数:惩罚特定误分类的实现疑问

实现带自定义惩罚矩阵的多分类损失函数

你需要用Keras后端(K)的张量操作来处理标签索引与惩罚权重的匹配,因为y_true和y_pred都是张量,无法直接用普通Python索引操作。以下是完整实现步骤与代码:

1. 定义惩罚矩阵

先创建12×12的惩罚矩阵,其中penalty_matrix[真实标签][预测标签]对应该分类组合的惩罚权重,再将其转为Keras张量(确保数据类型与模型一致):

import tensorflow as tf
from tensorflow.keras import backend as K

# 示例:构造12×12惩罚矩阵,对角线设为1(正确分类无额外惩罚),误分类按需设置权重
penalty_matrix = K.constant([
    [1.0, 3.0, 1.0, ...],  # 真实类别0时,各预测类别的惩罚权重
    [1.0, 1.0, 5.0, ...],  # 真实类别1时,各预测类别的惩罚权重
    # 补全剩余10行的权重设置
], dtype=K.floatx())

2. 完整自定义损失函数实现

def custom_loss(y_true, y_pred):
    # 计算基础交叉熵损失
    ce_loss = K.categorical_crossentropy(y_true, y_pred)
    
    # 从one-hot编码中提取真实/预测标签的类别索引
    true_labels = K.argmax(y_true, axis=-1)
    pred_labels = K.argmax(y_pred, axis=-1)
    
    # 构造(真实标签, 预测标签)的二维索引对,形状为(batch_size, 2)
    indices = K.stack([true_labels, pred_labels], axis=1)
    # 从惩罚矩阵中批量取出每个样本对应的惩罚权重
    penalty_weights = K.gather_nd(penalty_matrix, indices)
    
    # 将交叉熵损失与对应惩罚权重相乘
    weighted_loss = ce_loss * penalty_weights
    
    return weighted_loss

3. 模型编译时使用该损失

model.compile(optimizer='adam', loss=custom_loss, metrics=['accuracy'])

关键细节说明

  • K.argmax:从one-hot编码张量中提取类别索引,axis=-1表示沿最后一个维度(类别维度)取最大值索引。
  • K.stack:将真实标签与预测标签的索引张量堆叠为二维索引对,确保每个样本的分类组合能对应到惩罚矩阵的位置。
  • K.gather_nd:根据二维索引批量提取惩罚权重,实现每个样本损失与对应惩罚值的匹配。
  • 惩罚矩阵的对角线元素建议设为1,确保正确分类的损失不被调整,仅误分类应用自定义惩罚。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 19:15:35