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

TensorFlow中为TensorBoard混淆矩阵添加文本标签及自定义retrain.py问题

Hey there! Let me walk through how I'd approach adding a confusion matrix to TensorBoard while customizing the retrain.py script—since I've messed around with similar tweaks before, I feel your pain with those debugging hurdles!

Customizing retrain.py with Confusion Matrix in TensorBoard

Background & Context

First off, kudos on expanding the stock script with extra dense layers, Dropout, and momentum-based gradient descent—those are solid tweaks for a custom image dataset. Adding a confusion matrix to TensorBoard is a great call for visualizing class-wise performance, and I’ve also tested both top answers from that thread, so I know the second one can throw some weird debugging snags (looking at you, tensor shape mismatches!).

Step-by-Step Implementation in add_evaluation_step

Here’s how I modified the add_evaluation_step function to get the confusion matrix working reliably, building off Jerod’s approach:

1. Keep the Name Scope for Organization

Sticking with tf.name_scope ensures all confusion matrix-related operations are grouped neatly in TensorBoard, making it easier to navigate later.

2. Core Code for Confusion Matrix & Summaries

Inside the scope, you’ll need to derive predicted/true classes, compute the confusion matrix, and convert it to a visual summary TensorBoard can render. Here’s the full function:

def add_evaluation_step(result_tensor, ground_truth_tensor):
    with tf.name_scope('evaluation'):
        # Extract predicted and true class indices
        predicted_classes = tf.argmax(result_tensor, 1, name='predicted_classes')
        true_classes = tf.argmax(ground_truth_tensor, 1, name='true_classes')
        
        # Calculate confusion matrix
        confusion_mat = tf.math.confusion_matrix(
            labels=true_classes,
            predictions=predicted_classes,
            name='confusion_matrix'
        )
        
        # Convert matrix to an image summary for TensorBoard visualization
        # Normalize and reshape to fit image summary requirements
        cm_normalized = tf.cast(confusion_mat, tf.float32) / tf.reduce_sum(confusion_mat, axis=1, keepdims=True)
        cm_image = tf.reshape(cm_normalized, [1, -1, tf.shape(cm_normalized)[1], 1])
        cm_summary = tf.summary.image('normalized_confusion_matrix', cm_image, max_outputs=1)
        
        # Optional: Add scalar metrics for precision/recall to track alongside
        precision = tf.metrics.precision(labels=true_classes, predictions=predicted_classes)[1]
        recall = tf.metrics.recall(labels=true_classes, predictions=predicted_classes)[1]
        tf.summary.scalar('precision', precision)
        tf.summary.scalar('recall', recall)
        
        return cm_summary, precision, recall

Key Tips to Avoid Debugging Headaches

  • Shape Compatibility: Double-check that result_tensor and ground_truth_tensor have matching shapes. If using one-hot labels, both should be [batch_size, num_classes].
  • TF Version Adjustments: If you’re on TensorFlow 2.x, wrap this function in tf.function if you’re not using legacy graph mode—eager execution can throw off summary ops otherwise.
  • Merge Summaries: Don’t forget to add the new cm_summary to your existing summary merge operation in the training loop of retrain.py; otherwise, it won’t show up in TensorBoard.
  • Normalization: The normalized matrix (divided by true class counts) is far more useful for visualizing misclassification patterns than raw counts.

Why Jerod’s Approach Worked Better for Me

The second answer in that thread often relies on custom visualization functions or more complex summary ops, which can break when you’re already modifying other parts of the script (like adding dense layers or Dropout). Jerod’s approach sticks to core TensorFlow ops, making it more stable and easier to integrate with your existing customizations.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:01:56