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

基于Keras Model API实现每轮AUROC计算及批次聚合问询

Got it, let's work through your problem step by step. Since you're using Keras' Model API and can't rely on self.model, self.validation_data, or self.model.predict directly in callbacks, we'll build a custom callback that works around these constraints, plus address your questions about batch-level AUROC calculation and aggregation.


1. Custom Callback for Per-Epoch AUROC Calculation & Storage

We'll create a callback that accepts your model and test data generator as parameters (instead of relying on self.model), then computes AUROC at the end of each epoch and stores the results.

First, import the required libraries:

import numpy as np
from sklearn.metrics import roc_auc_score
from tensorflow.keras.callbacks import Callback

Then define the callback class:

class AUROCCallback(Callback):
    def __init__(self, target_model, test_generator, total_test_samples):
        super().__init__()
        self.target_model = target_model  # Pass your Model instance here
        self.test_gen = test_generator    # Pass your test data generator
        self.total_test_samples = total_test_samples  # Total samples in test set
        self.epoch_aurocs = []  # Store AUROC for each epoch

    def on_epoch_end(self, epoch, logs=None):
        # Collect all predictions and true labels from the test set
        all_predictions = []
        all_true_labels = []

        # Reset generator to avoid repeating data across epochs (if supported)
        if hasattr(self.test_gen, 'reset'):
            self.test_gen.reset()

        # Iterate through all batches in the test generator
        for _ in range(len(self.test_gen)):
            x_batch, y_batch = next(self.test_gen)
            batch_preds = self.target_model.predict(x_batch, verbose=0)
            
            # Adjust flattening based on your task:
            # - For binary classification: flatten to 1D array of probabilities
            # - For multi-class: use appropriate formatting (e.g., keep class probabilities)
            all_predictions.extend(batch_preds.flatten())
            all_true_labels.extend(y_batch.flatten())

        # Calculate full-test-set AUROC
        epoch_auc = roc_auc_score(all_true_labels, all_predictions)
        self.epoch_aurocs.append(epoch_auc)

        # Print results for visibility
        print(f"\nEpoch {epoch+1} - Test AUROC: {epoch_auc:.4f}")

Usage in Your Main Code

# Assume you've already defined your Model instance `my_model` and test generator `test_gen`
total_test_samples = test_gen.n  # Use this if using ImageDataGenerator; adjust otherwise

# Initialize the callback
auroc_callback = AUROCCallback(target_model=my_model, test_generator=test_gen, total_test_samples=total_test_samples)

# Start training with the callback
my_model.fit(
    train_generator,
    epochs=30,
    callbacks=[auroc_callback]
)

# After training, access stored AUROCs for visualization/analysis
print("All epoch AUROCs:", auroc_callback.epoch_aurocs)

2. Batch-Level AUROC Calculation & Aggregation

Quick Answer to Your Question

By default, the code above calculates AUROC using all test samples at once, not per batch. If you want to compute AUROC for each test batch first, then aggregate the results, here's how to do it:

Modified Callback for Batch-Level Calculation

We'll compute AUROC for each batch, then aggregate using either simple average or weighted average (weighted by batch size, which is more accurate if batch sizes vary):

class BatchAUROCCallback(Callback):
    def __init__(self, target_model, test_generator, total_test_samples):
        super().__init__()
        self.target_model = target_model
        self.test_gen = test_generator
        self.total_test_samples = total_test_samples
        self.epoch_batch_aurocs = []  # Store AUROC for each batch per epoch
        self.epoch_aggregated_aucs = []  # Store aggregated results per epoch

    def on_epoch_end(self, epoch, logs=None):
        batch_aucs = []
        batch_sizes = []

        if hasattr(self.test_gen, 'reset'):
            self.test_gen.reset()

        for _ in range(len(self.test_gen)):
            x_batch, y_batch = next(self.test_gen)
            batch_preds = self.target_model.predict(x_batch, verbose=0)
            preds_flat = batch_preds.flatten()
            labels_flat = y_batch.flatten()

            # Skip batches with only one class (avoids roc_auc_score errors)
            try:
                batch_auc = roc_auc_score(labels_flat, preds_flat)
                batch_aucs.append(batch_auc)
                batch_sizes.append(len(labels_flat))
            except ValueError:
                print(f"Skipping batch (only one class present)")
                continue

        # Aggregate results
        simple_avg_auc = np.mean(batch_aucs)
        weighted_avg_auc = np.average(batch_aucs, weights=batch_sizes)

        self.epoch_batch_aurocs.append(batch_aucs)
        self.epoch_aggregated_aucs.append({
            'simple_average': simple_avg_auc,
            'weighted_average': weighted_avg_auc
        })

        # Print aggregated results
        print(f"\nEpoch {epoch+1} - Simple Avg AUROC: {simple_avg_auc:.4f}, Weighted Avg AUROC: {weighted_avg_auc:.4f}")

Usage

batch_auroc_callback = BatchAUROCCallback(target_model=my_model, test_generator=test_gen, total_test_samples=total_test_samples)
my_model.fit(train_generator, epochs=30, callbacks=[batch_auroc_callback])

# Access batch-level results for a specific epoch (e.g., 3rd epoch, index 2)
print("3rd epoch batch AUROCs:", batch_auroc_callback.epoch_batch_aurocs[2])
print("3rd epoch aggregated AUROCs:", batch_auroc_callback.epoch_aggregated_aucs[2])

Key Notes
  • Multi-class Adjustments: For multi-class tasks, update roc_auc_score with multi_class='ovr' or 'ovo', and ensure predictions/labels are formatted correctly (e.g., one-hot labels paired with class probabilities).
  • Memory Efficiency: Batch-level aggregation uses less memory than storing all predictions/labels at once, which is helpful for large test sets.
  • Generator Reset: Always reset your test generator at the start of on_epoch_end if it supports the reset() method (e.g., generators from ImageDataGenerator.flow()).

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 08:13:00