求适配常规设置的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
核心改进点:
- 使用纯TensorFlow操作实现QWK计算,避免依赖sklearn的numpy操作,兼容Graph模式、GPU、JAX后端
- 正确使用Keras Metric的
add_weightAPI维护状态,而非直接赋值 - 支持样本权重(可选)
- 适配稀疏标签(如
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
相关产品推荐
相关产品推荐

