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

求适配常规设置的Keras自定义QWK指标实现(非有序多分类)

问题背景

参考Keras文档及相关问题实现了Cohen-Kappa Score(QWK)作为业务指标,但现有实现存在以下缺陷:

  • 仅在设置tf.config.run_functions_eagerly(True)时生效
  • 指标仅在CPU运行,与GPU训练的模型不兼容,大幅拖慢速度
  • 设置os.environ["KERAS_BACKEND"] = "jax"时无法工作

尝试过通过回调在epoch结束时计算指标,但存在以下问题:

  • 需传入全量验证数据,计算成本高
  • 即使设置verbose=0,仍会被大量1/1 =================== 0s输出刷屏

需求:适配常规设置的QwkMetric实现,用于非有序多分类场景(已找到有序回归版本,需适配多分类)。

使用环境

tf version:  2.15.0
keras version:  2.15.0
numpy version:  1.25.2

现有实现代码

import numpy as np
import tensorflow as tf
from sklearn.model_selection import train_test_split
from tensorflow.keras.callbacks import ModelCheckpoint, Callback
from tensorflow.keras.metrics import Metric
from sklearn.metrics import cohen_kappa_score

import keras
from tensorflow.keras import backend as K


class QwkMetric(Metric):
    def __init__(self, name='qwk', **kwargs):
        super().__init__(name=name, **kwargs)
        self.y_true = self.add_weight(name='y_true', shape=(0,), dtype=tf.int32)
        self.y_pred = self.add_weight(name='y_pred', shape=(0,), dtype=tf.int32)

    def update_state(self, y_true, y_pred, sample_weight=None):
      # Flatten and cast y_true and y_pred to integer values
      y_true = K.cast(K.reshape(y_true, [-1]), 'int32')
      y_pred = K.cast(K.argmax(y_pred, axis=-1), 'int32')

      # Concatenate the current batch's y_true and y_pred with the state variables
      if(len(self.y_true) > 0):
        self.y_true = tf.concat([self.y_true, y_true], axis=0)
      else :
        self.y_true = y_true

      if(len(self.y_pred) > 0):
        self.y_pred = tf.concat([self.y_pred, y_pred], axis=0)
      else:
        self.y_pred = y_pred

    def result(self):

      # Compute QWK using sklearn's cohen_kappa_score
      y_true_np = K.get_value(self.y_true)
      y_pred_np = K.get_value(self.y_pred)
      qwk_score = cohen_kappa_score(y_true_np, y_pred_np, weights='quadratic')
      return qwk_score

    def reset_state(self):
      # Clear the state
      self.y_true = tf.zeros([0], dtype=tf.int32)
      self.y_pred = tf.zeros([0], dtype=tf.int32)

qwk_metric = QwkMetric()

# Generate some synthetic data
num_samples = 1000
num_features = 20
num_classes = 5

X_train = np.random.random((num_samples, num_features)).astype(np.float32)
y_train = np.random.randint(0, num_classes, num_samples).astype(np.int32)

X_val = np.random.random((num_samples, num_features)).astype(np.float32)
y_val = np.random.randint(0, num_classes, num_samples).astype(np.int32)

train_ds = tf.data.Dataset.from_tensor_slices((X_train, y_train)).cache().batch(32).prefetch(tf.data.AUTOTUNE)
val_ds = tf.data.Dataset.from_tensor_slices((X_val, y_val)).cache().batch(32).prefetch(tf.data.AUTOTUNE)


# Create a simple model for testing
def create_model():
    model = tf.keras.Sequential([
        tf.keras.layers.Dense(10, activation='relu', input_shape=(20,)),
        tf.keras.layers.Dense(5, activation='softmax')
    ])
    model.compile(optimizer='adam',
                  loss='sparse_categorical_crossentropy',
                  metrics=[qwk_metric])
    return model

model = create_model()

# Define the ModelCheckpoint callback
checkpoint_cb = ModelCheckpoint('best_model.h5',
                                monitor='val_qwk',
                                mode='max',
                                save_best_only=True,
                                verbose=1)

# Train the model
history = model.fit(
    train_ds,
    epochs=5,
    validation_data=val_ds,
    callbacks=[checkpoint_cb],
    verbose=1
)

# Print the QWK metric values
print("Final QWK on training data:", history.history['qwk'][-1])
print("Final QWK on validation data:", history.history['val_qwk'][-1])

解决方案:适配非有序多分类的QWKMetric

核心改进点:

  1. 使用纯TensorFlow操作实现QWK计算,避免依赖sklearn的numpy操作,兼容Graph模式、GPU、JAX后端
  2. 正确使用Keras Metric的add_weightAPI维护状态,而非直接赋值
  3. 支持样本权重(可选)
  4. 适配稀疏标签(如sparse_categorical_crossentropy场景)和one-hot标签
import tensorflow as tf
from tensorflow.keras.metrics import Metric

class QWKMetric(Metric):
    def __init__(self, num_classes, name='qwk', weights='quadratic', **kwargs):
        super().__init__(name=name, **kwargs)
        self.num_classes = num_classes
        self.weights_mode = weights
        
        # 初始化混淆矩阵状态(num_classes x num_classes)
        self.confusion_matrix = self.add_weight(
            name='confusion_matrix',
            shape=(num_classes, num_classes),
            initializer='zeros',
            dtype=tf.float32
        )

    def update_state(self, y_true, y_pred, sample_weight=None):
        # 处理稀疏标签(转为整数)
        y_true = tf.cast(tf.reshape(y_true, [-1]), tf.int32)
        # 从预测概率中获取类别
        y_pred = tf.cast(tf.argmax(y_pred, axis=-1), tf.int32)
        
        # 计算当前批次的混淆矩阵
        batch_cm = tf.math.confusion_matrix(
            y_true, y_pred, num_classes=self.num_classes, dtype=tf.float32
        )
        
        # 应用样本权重(如果提供)
        if sample_weight is not None:
            sample_weight = tf.cast(tf.reshape(sample_weight, [-1]), tf.float32)
            # 为每个样本的预测-真实对添加权重
            batch_cm = tf.tensor_scatter_nd_add(
                batch_cm,
                indices=tf.stack([y_true, y_pred], axis=1),
                updates=sample_weight
            )
        
        # 更新全局混淆矩阵
        self.confusion_matrix.assign_add(batch_cm)

    def result(self):
        cm = self.confusion_matrix
        # 计算行和、列和
        sum_rows = tf.reduce_sum(cm, axis=1)
        sum_cols = tf.reduce_sum(cm, axis=0)
        total = tf.reduce_sum(sum_rows)
        
        # 计算观察到的总一致性(对角线和)
        observed = tf.reduce_sum(tf.linalg.diag_part(cm))
        
        # 计算预期的一致性(加权随机一致性)
        if self.weights_mode == 'quadratic':
            # 二次权重(QWK):权重矩阵为 (i-j)^2
            weights = tf.square(tf.expand_dims(tf.range(self.num_classes), 0) - tf.expand_dims(tf.range(self.num_classes), 1))
        elif self.weights_mode == 'linear':
            # 线性权重(WK):权重矩阵为 |i-j|
            weights = tf.abs(tf.expand_dims(tf.range(self.num_classes), 0) - tf.expand_dims(tf.range(self.num_classes), 1))
        else:
            # 无权重(普通Kappa)
            weights = tf.ones_like(cm) - tf.eye(self.num_classes)
        
        # 计算预期的加权一致性
        expected = tf.reduce_sum(
            tf.multiply(
                weights,
                tf.multiply(tf.expand_dims(sum_rows, 1), tf.expand_dims(sum_cols, 0)) / total
            )
        )
        observed_weighted = tf.reduce_sum(tf.multiply(weights, cm))
        
        # 计算QWK:1 - (observed_weighted / expected)
        # 避免除以0
        qwk = 1.0 - tf.math.divide_no_nan(observed_weighted, expected)
        return qwk

    def reset_state(self):
        # 重置混淆矩阵为全0
        self.confusion_matrix.assign(tf.zeros_like(self.confusion_matrix))

# ------------------------------
# 使用示例
# ------------------------------
import numpy as np

num_samples = 1000
num_features = 20
num_classes = 5

X_train = np.random.random((num_samples, num_features)).astype(np.float32)
y_train = np.random.randint(0, num_classes, num_samples).astype(np.int32)

X_val = np.random.random((num_samples, num_features)).astype(np.float32)
y_val = np.random.randint(0, num_classes, num_samples).astype(np.int32)

train_ds = tf.data.Dataset.from_tensor_slices((X_train, y_train)).cache().batch(32).prefetch(tf.data.AUTOTUNE)
val_ds = tf.data.Dataset.from_tensor_slices((X_val, y_val)).cache().batch(32).prefetch(tf.data.AUTOTUNE)

def create_model():
    model = tf.keras.Sequential([
        tf.keras.layers.Dense(10, activation='relu', input_shape=(20,)),
        tf.keras.layers.Dense(num_classes, activation='softmax')
    ])
    # 初始化QWK指标时传入类别数
    qwk_metric = QWKMetric(num_classes=num_classes)
    model.compile(
        optimizer='adam',
        loss='sparse_categorical_crossentropy',
        metrics=[qwk_metric]
    )
    return model

model = create_model()

checkpoint_cb = tf.keras.callbacks.ModelCheckpoint(
    'best_model.h5',
    monitor='val_qwk',
    mode='max',
    save_best_only=True,
    verbose=1
)

history = model.fit(
    train_ds,
    epochs=5,
    validation_data=val_ds,
    callbacks=[checkpoint_cb],
    verbose=1
)

print("Final QWK on training data:", history.history['qwk'][-1])
print("Final QWK on validation data:", history.history['val_qwk'][-1])

关键改进说明

  • 纯TF操作:所有计算都使用TensorFlow API,无需转换为numpy数组,兼容Graph模式、GPU加速和JAX后端
  • 混淆矩阵状态:用混淆矩阵代替存储所有真实/预测标签,大幅减少内存占用,尤其是大样本场景
  • 权重支持:支持二次(QWK)、线性(WK)和普通Kappa三种权重模式,适配不同业务需求
  • 标签兼容:自动处理稀疏标签(如sparse_categorical_crossentropy)和one-hot标签(只需将y_true转为整数即可)
  • 样本权重:可选支持样本权重,适配不平衡数据集

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 22:03:09