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

批量计算与全量计算模型指标差异异常问题求助

问题描述

使用预训练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

注:两种函数输出结果基本一致。

差异成因分析与解决办法

核心成因推测

  1. 样本顺序错位
    旧评估代码中,images_ds保留原数据集的批次结构,而labels_ds做了unbatch()操作,导致图像和标签的全局顺序完全不匹配。训练阶段按批次计算时,每个批次内的图像和标签一一对应,指标正常;但全量评估时,预测结果和标签完全错位,指标趋近于随机猜测水平(类别数较多时不足1%)。

  2. 标签处理逻辑不一致
    训练时的标签格式(如是否one-hot、类别索引映射)可能与评估时的处理逻辑冲突。例如训练时标签是整数索引,但评估时错误执行np.argmax,导致标签全部变为无效值。

  3. 全量数据加载截断
    当数据集过大时,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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 21:57:41