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

TensorHub图像重训练示例:混淆矩阵与评估时间记录问题

Hey there! Let's tackle these two issues with the TensorFlow Hub retrain script one by one—I’ve tinkered with this exact script before, so I know the quirks pretty well.

1. Fixing Confusion Matrix Output (Tensor → Numpy Array)

The root issue here is that tf.confusion_matrix spits out a Tensor object by default, not a concrete value. You need to evaluate that tensor within an active TensorFlow session to get the actual numpy array. Here’s how to modify the run_final_eval function properly:

First, collect the real ground truth labels and model predictions by running them in the session first. Then compute the confusion matrix on those concrete values. Here’s a code snippet for the modified function:

def run_final_eval(sess, image_lists, label_names, class_count, jpeg_data_tensor,
                   decoded_image_tensor, resized_image_tensor, bottleneck_tensor,
                   bottleneck_input, ground_truth_input, final_tensor):
    # ... keep existing setup code for test batches ...

    ground_truth = []
    predictions = []
    
    # Pull actual label and prediction values from the session
    for test_batch in test_batches:
        batch_truth, batch_predictions = sess.run(
            [ground_truth_input, final_tensor],
            feed_dict={bottleneck_input: test_batch[0], ground_truth_input: test_batch[1]}
        )
        ground_truth.extend(np.argmax(batch_truth, axis=1))
        predictions.extend(np.argmax(batch_predictions, axis=1))
    
    # Option 1: Use TensorFlow's confusion matrix (evaluate it in the session)
    confusion_matrix_tensor = tf.confusion_matrix(labels=ground_truth, predictions=predictions, num_classes=class_count)
    confusion_matrix = sess.run(confusion_matrix_tensor)
    
    # Option 2: Use scikit-learn's confusion matrix (simpler, no TF tensor needed)
    # from sklearn.metrics import confusion_matrix
    # confusion_matrix = confusion_matrix(ground_truth, predictions)
    
    print("\nConfusion Matrix:")
    print(confusion_matrix)
    return confusion_matrix

By first fetching the actual label and prediction arrays via sess.run, you avoid dealing with unevaluated tensors. The confusion matrix will now be a numpy array you can save or analyze further.

2. Recording Evaluation Times Per Test Batch/Sample

To track how long each evaluation step takes, use Python’s time module to timestamp before and after each inference call. Insert this into the test batch loop in run_final_eval:

import time

def run_final_eval(sess, image_lists, label_names, class_count, jpeg_data_tensor,
                   decoded_image_tensor, resized_image_tensor, bottleneck_tensor,
                   bottleneck_input, ground_truth_input, final_tensor):
    # ... existing setup code ...
    
    evaluation_times = []
    ground_truth = []
    predictions = []
    
    for test_batch in test_batches:
        start_time = time.time()  # Start timer right before inference
        
        # Run the model evaluation for this batch
        batch_truth, batch_predictions = sess.run(
            [ground_truth_input, final_tensor],
            feed_dict={bottleneck_input: test_batch[0], ground_truth_input: test_batch[1]}
        )
        
        elapsed_time = time.time() - start_time  # Calculate time taken
        evaluation_times.append(elapsed_time)
        
        # Continue collecting labels/predictions as before
        ground_truth.extend(np.argmax(batch_truth, axis=1))
        predictions.extend(np.argmax(batch_predictions, axis=1))
    
    # Print or save the times
    print("\nEvaluation time per batch (seconds):", evaluation_times)
    # Save to a file if needed:
    # np.save("evaluation_times.npy", evaluation_times)
    
    # ... rest of the confusion matrix code ...

This tracks time per batch (the most common use case for performance tracking). If you need per-image times instead, you’ll have to modify the loop to process individual images (though this will slow down evaluation significantly). The code above only times the model inference step (sess.run), which is usually what you care about for evaluation speed.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 07:55:12