基于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.
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)
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])
- Multi-class Adjustments: For multi-class tasks, update
roc_auc_scorewithmulti_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_endif it supports thereset()method (e.g., generators fromImageDataGenerator.flow()).
内容的提问来源于stack exchange,提问作者snailbee

