批量计算与全量计算模型指标差异异常问题求助
使用预训练MobileNet训练图像分类模型,训练阶段训练集、验证集的Accuracy、Precision、Recall、F1-score均达70%以上,但采用全量数据集(无批次,单次计算全量指标)评估时,所有指标均不足1%。编写了两种评估函数(代码如下),结果一致,无法定位问题,需分析差异成因并给出解决办法。
旧评估代码
def test_model(model, data, CLASSES, label_one_hot=True, average="micro", threshold_analysis=False, thres_analysis_start_point=0.0, thres_analysis_end_point=0.95, thres_step=0.05, classwise_analysis=False, produce_confusion_matrix=False): images_ds = data.map(lambda image, label: image) labels_ds = data.map(lambda image, label: label).unbatch() NUM_VALIDATION_IMAGES = count_data_items(tf_records_filenames=data) cm_correct_labels = next(iter(labels_ds.batch(NUM_VALIDATION_IMAGES))).numpy() # get everything as one batch if label_one_hot is True: cm_correct_labels = np.argmax(cm_correct_labels, axis=-1) cm_probabilities = model.predict(images_ds) cm_predictions = np.argmax(cm_probabilities, axis=-1) warnings.filterwarnings('ignore') overall_score = f1_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average=average) overall_precision = precision_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average=average) overall_recall = recall_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average=average) overall_test_results = {'overall_f1_score': overall_score, 'overall_precision':overall_precision, 'overall_recall':overall_recall} if classwise_analysis is True: label_index_dict = get_index_label_from_tf_record(dataset=data) label_index_dict = {k:v for k, v in sorted(list(label_index_dict.items()))} label_index_df = pd.DataFrame(label_index_dict, index=[0]).T.reset_index().rename(columns={'index':'class_ind', 0:'class_names'}) classwise_score = f1_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average=None) classwise_precision = precision_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average=None) classwise_recall = recall_score(cm_correct_labels, cm_predictions, labels=range(len(CLASSES)), average=None) ind_class_count_df = class_ind_counter_from_tfrecord(data) ind_class_count_df = ind_class_count_df.merge(label_index_df, how='left', left_on='class_names', right_on='class_names') classwise_test_results = {'classwise_f1_score':classwise_score, 'classwise_precision':classwise_precision, 'classwise_recall':classwise_recall, 'class_names':CLASSES} classwise_test_results_df = pd.DataFrame(classwise_test_results) if produce_confusion_matrix is True: cmat = confusion_matrix(cm_correct_labels, cm_predictions, labels=range(len(CLASSES))) return overall_test_results, classwise_test_results, cmat return overall_test_results, classwise_test_results if produce_confusion_matrix is True: cmat = confusion_matrix(cm_correct_labels, cm_predictions, labels=range(len(CLASSES))) return overall_test_results, cmat warnings.filterwarnings('always') return overall_test_results
TensorFlow新评估代码
def eval_model(y_true, y_pred): eval_results = {} unbatch_accuracy = tf.keras.metrics.CategoricalAccuracy(name='unbatch_accuracy') unbatch_recall = tf.keras.metrics.Recall(name='unbatch_recall') unbatch_precision = tf.keras.metrics.Precision(name='unbatch_precision') unbatch_f1_micro = tfa.metrics.F1Score(name='unbatch_f1_micro', num_classes=n_labels, average='micro') unbatch_f1_macro = tfa.metrics.F1Score(name='unbatch_f1_macro', num_classes=n_labels, average='macro') unbatch_accuracy.update_state(y_true, y_pred) unbatch_recall.update_state(y_true, y_pred) unbatch_precision.update_state(y_true, y_pred) unbatch_f1_micro.update_state(y_true, y_pred) unbatch_f1_macro.update_state(y_true, y_pred) eval_results['unbatch_accuracy'] = unbatch_accuracy.result().numpy() eval_results['unbatch_recall'] = unbatch_recall.result().numpy() eval_results['unbatch_precision'] = unbatch_precision.result().numpy() eval_results['unbatch_f1_micro'] = unbatch_f1_micro.result().numpy() eval_results['unbatch_f1_macro'] = unbatch_f1_macro.result().numpy() unbatch_accuracy.reset_states() unbatch_recall.reset_states() unbatch_precision.reset_states() unbatch_f1_micro.reset_states() unbatch_f1_macro.reset_states() return eval_results
注:两种函数输出结果基本一致。
核心成因推测
样本顺序错位
旧评估代码中,images_ds保留原数据集的批次结构,而labels_ds做了unbatch()操作,导致图像和标签的全局顺序完全不匹配。训练阶段按批次计算时,每个批次内的图像和标签一一对应,指标正常;但全量评估时,预测结果和标签完全错位,指标趋近于随机猜测水平(类别数较多时不足1%)。标签处理逻辑不一致
训练时的标签格式(如是否one-hot、类别索引映射)可能与评估时的处理逻辑冲突。例如训练时标签是整数索引,但评估时错误执行np.argmax,导致标签全部变为无效值。全量数据加载截断
当数据集过大时,labels_ds.batch(NUM_VALIDATION_IMAGES)可能无法一次性加载所有数据,导致实际获取的标签数量少于预测结果数量,引发对应关系错乱。
解决办法
修正样本对应关系
避免分别提取图像和标签,直接遍历完整数据集同时获取两者,确保顺序完全对齐:
# 替换原images_ds、labels_ds的提取逻辑 all_images = [] all_labels = [] for img, lbl in data.unbatch(): all_images.append(img.numpy()) all_labels.append(lbl.numpy()) all_images = np.array(all_images) all_labels = np.array(all_labels) # 后续预测和标签处理 cm_probabilities = model.predict(all_images) cm_predictions = np.argmax(cm_probabilities, axis=-1) if label_one_hot is True: cm_correct_labels = np.argmax(all_labels, axis=-1) else: cm_correct_labels = all_labels
验证标签处理逻辑
打印cm_correct_labels和cm_predictions的前10个样本,检查类别分布是否合理。若标签全为同一值,说明标签转换逻辑错误,需与训练时的标签格式对齐。
优化全量数据加载
若数据集过大,分批次遍历累积结果,避免内存问题同时保证顺序一致:
cm_correct_labels = [] cm_probabilities = [] for img_batch, lbl_batch in data: cm_probabilities.append(model.predict(img_batch)) cm_correct_labels.append(lbl_batch.numpy()) # 合并所有批次 cm_probabilities = np.concatenate(cm_probabilities, axis=0) cm_correct_labels = np.concatenate(cm_correct_labels, axis=0)
规范新评估代码使用
确保传入eval_model的y_true和y_pred样本数量、顺序完全一致,且格式匹配(如均为one-hot编码或整数索引)。
内容的提问来源于stack exchange,提问作者learner

