Download Classification/inference.py from ODELIA-AI/SwinUNETR: direct link, hf CLI and curl.
- Browser
- Download file 9.26 kB
-
https://huggingface.co/ODELIA-AI/SwinUNETR/resolve/main/Classification/inference.py
- Command line
-
hf download hf://ODELIA-AI/SwinUNETR/Classification/inference.py
-
curl -L -o inference.py https://huggingface.co/ODELIA-AI/SwinUNETR/resolve/main/Classification/inference.py
9.26 kB
| import torch | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| from torch.utils.data import DataLoader | |
| from tqdm import tqdm | |
| from dataloading.dataloader2D import NiftiSegmentationDataset | |
| import yaml | |
| from pathlib import Path | |
| import matplotlib.pyplot as plt | |
| from sklearn.calibration import calibration_curve | |
| from sklearn.metrics import roc_auc_score, roc_curve | |
| from sklearn.preprocessing import label_binarize | |
| from odelia_breast_mri.scripts.main_predict import evaluate | |
| import torchvision.transforms.functional as TF | |
| from dataloading.collate_function import custom_collate | |
| from models.swinunetr import SwinUNETRMultiTask | |
| def tta_transforms(x): | |
| """ | |
| Apply test-time augmentations to a single batch tensor. | |
| Input: x (B, C, H, W) | |
| Returns a list of augmented versions. | |
| """ | |
| return [ | |
| x, # original | |
| TF.hflip(x), | |
| TF.vflip(x), | |
| TF.rotate(x, 90), | |
| TF.rotate(x, 180), | |
| TF.rotate(x, 270), | |
| TF.hflip(TF.vflip(x)), # h+v flip | |
| ] | |
| def plot_multiclass_reliability(probs, labels, class_names=None, n_bins=10): | |
| """ | |
| Plot calibration curves for each class in a multiclass problem. | |
| Args: | |
| probs: numpy array of shape (n_samples, n_classes) with predicted probabilities. | |
| labels: numpy array of shape (n_samples,) with integer class labels. | |
| class_names: list of class names (optional). | |
| n_bins: number of bins for calibration curve. | |
| """ | |
| n_classes = probs.shape[1] | |
| if class_names is None: | |
| class_names = [f"Class {i}" for i in range(n_classes)] | |
| plt.figure(figsize=(8, 8)) | |
| for i in range(n_classes): | |
| # Binary labels: is this class or not | |
| binarized_labels = (labels == i).astype(int) | |
| fraction_of_positives, mean_predicted_value = calibration_curve( | |
| binarized_labels, probs[:, i], n_bins=n_bins, strategy='uniform' | |
| ) | |
| plt.plot( | |
| mean_predicted_value, fraction_of_positives, "s-", | |
| label=f"{class_names[i]}" | |
| ) | |
| plt.plot([0, 1], [0, 1], "k:", label="Perfectly calibrated") | |
| plt.title("Multiclass Reliability Diagram") | |
| plt.xlabel("Mean Predicted Probability") | |
| plt.ylabel("Fraction of Positives") | |
| plt.legend() | |
| plt.grid() | |
| # plt.show() | |
| plt.savefig('/workspace/ClassifierSegmenter/results/multiclass_reliability_diagram.png') | |
| # ---------------------- | |
| # Load Config & Set Up | |
| # ---------------------- | |
| with open("/workspace/ClassifierSegmenter/config2d.yaml", "r") as f: | |
| config = yaml.safe_load(f) | |
| device = torch.device(config["device"] if torch.cuda.is_available() else "cpu") | |
| # ---------------------- | |
| # Load Model | |
| # ---------------------- | |
| checkpoint_paths = [ | |
| "/workspace/Classifier/checkpoints/final/model1/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model2/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model3/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model4/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model5/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model6/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model7/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model8/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model9/best_score_model.pth", | |
| "/workspace/Classifier/checkpoints/final/model10/best_score_model.pth" | |
| ] | |
| models = [] | |
| for i, ckpt_path in enumerate(checkpoint_paths, start=1): | |
| model = SwinUNETRMultiTask( | |
| img_size=(256, 256), | |
| in_channels=4, | |
| out_seg_channels=2, | |
| out_cls_classes=3 | |
| ).to(device) | |
| if ckpt_path is not None: | |
| state = torch.load(ckpt_path, map_location=device, weights_only=True) | |
| model.load_state_dict(state) | |
| model.to(device) | |
| model.eval() | |
| models.append(model) | |
| # ---------------------- | |
| # Load Test Dataset | |
| # ---------------------- | |
| test_dataset = NiftiSegmentationDataset(config["csv_file_test"], channel_keys=config["channel_keys"], augment=False) | |
| test_loader = DataLoader(test_dataset, batch_size=config["batch_size"], collate_fn=custom_collate, shuffle=False) | |
| metrics = {} | |
| # ---------------------- | |
| # Inference + Metrics | |
| # ---------------------- | |
| all_preds, all_probs, all_targets = [], [], [] | |
| all_patient_ids = [] | |
| with torch.no_grad(): | |
| if not config["tta"]: | |
| for batch in tqdm(test_loader, desc="Running Test Inference"): | |
| x = batch['image'].to(device) | |
| y = batch['cls_label'].to(device) | |
| mask = batch['mask'].to(device) if batch['mask'] is not None else None | |
| has_mask = batch['has_mask'].to(device) if batch['has_mask'] is not None else None | |
| with torch.autocast(device_type='cuda'): | |
| # Collect logits from all models | |
| logits_list = [] | |
| for model in models: | |
| _, out, _ = model(x) | |
| logits_list.append(out) | |
| # Average probabilities across models | |
| probs_list = [torch.softmax(logits, dim=1) for logits in logits_list] | |
| probs = torch.stack(probs_list, dim=0).mean(dim=0) | |
| preds = torch.argmax(probs, dim=1) | |
| all_preds.append(preds.cpu()) | |
| all_probs.append(probs.cpu()) | |
| all_targets.append(y.cpu()) | |
| all_patient_ids.extend(batch['patient_id']) | |
| else: | |
| for batch in tqdm(test_loader, desc="Running Test Inference w/ TTA"): | |
| x = batch['image'].to(device) | |
| y = batch['cls_label'].to(device) | |
| mask = batch['mask'].to(device) if batch['mask'] is not None else None | |
| has_mask = batch['has_mask'].to(device) if batch['has_mask'] is not None else None | |
| tta_versions = tta_transforms(x) # list of augmented batches | |
| # Combine TTA with ensemble: collect logits for each model and each augmentation | |
| model_logits = [] | |
| for model in models: | |
| all_logits = [] | |
| for aug_x in tta_versions: | |
| with torch.autocast(device_type='cuda'): | |
| _, logits, _ = model(aug_x) | |
| all_logits.append(logits) | |
| stacked_logits = torch.stack(all_logits, dim=0).mean(dim=0) | |
| model_logits.append(stacked_logits) | |
| # Average across models | |
| mean_logits = torch.stack(model_logits, dim=0).mean(dim=0) | |
| probs = torch.softmax(mean_logits, dim=1) | |
| preds = torch.argmax(probs, dim=1) | |
| all_preds.append(preds.cpu()) | |
| all_probs.append(probs.cpu()) | |
| all_targets.append(y.cpu()) | |
| all_patient_ids.extend(batch['patient_id']) | |
| # Concatenate everything | |
| all_preds = torch.cat(all_preds) | |
| all_probs = torch.cat(all_probs) | |
| all_targets = torch.cat(all_targets) | |
| # ---------------------- | |
| # Evaluate MACRO metrics | |
| # ---------------------- | |
| accuracy = (all_preds == all_targets).sum().item() / len(all_targets) | |
| auc, sensitivity, specificity = evaluate( | |
| all_targets, | |
| all_preds, | |
| all_probs, | |
| path_out=Path('/workspace/ClassifierSegmenter/results') | |
| ) | |
| print("\n✅ MACRO Results") | |
| print(f"Accuracy: {accuracy:.4f}") | |
| print(f"AUC: {auc:.4f}") | |
| print(f"Sensitivity: {sensitivity:.4f}") | |
| print(f"Specificity: {specificity:.4f}") | |
| plot_multiclass_reliability( | |
| all_probs.numpy(), | |
| all_targets.numpy(), | |
| class_names=['no lesion', 'benign', 'malignant'], | |
| n_bins=10 | |
| ) | |
| # ---------------------- | |
| # Evaluate MICRO metrics | |
| # ---------------------- | |
| y_true_hot = label_binarize(all_targets.numpy().astype(str), classes=['0', '1', '2']) | |
| fpr, tpr, thresholds = roc_curve(y_true_hot.ravel(), all_probs.numpy().ravel(), drop_intermediate=False) | |
| roc_auc = roc_auc_score(y_true_hot, all_probs.numpy(), average="micro") | |
| roc_auc_macro = roc_auc_score(y_true_hot, all_probs.numpy(), average="macro") | |
| # Sensitivity at 90% specificity | |
| specificity_threshold = 0.90 | |
| fpr_threshold = 1 - specificity_threshold | |
| sensitivity_at_90_specificity = np.interp(fpr_threshold, fpr, tpr) | |
| # Specificity at 90% sensitivity | |
| sensitivity_threshold = 0.90 | |
| fpr_at_90_sensitivity = np.interp(sensitivity_threshold, tpr, fpr) | |
| specificity_at_90_sensitivity = 1 - fpr_at_90_sensitivity | |
| amalgamated_results = [roc_auc, specificity_at_90_sensitivity, sensitivity_at_90_specificity] | |
| averaged_results = np.mean(amalgamated_results) | |
| metrics["results"] = { | |
| "AUC": roc_auc, | |
| "AUC macro": roc_auc_macro, | |
| "Specificity": specificity_at_90_sensitivity, | |
| "Sensitivity": sensitivity_at_90_specificity, | |
| "Score": averaged_results | |
| } | |
| print("\n✅ MICRO Results") | |
| print(f"AUC: {roc_auc:.4f}") | |
| print(f"Sensitivity: {sensitivity_at_90_specificity:.4f}") | |
| print(f"Specificity: {specificity_at_90_sensitivity:.4f}") | |
| print(f"Score: {averaged_results:.4f}") | |
| # Save results | |
| results_df = pd.DataFrame({ | |
| "PatientID": all_patient_ids, | |
| "TrueLabel": all_targets.numpy(), | |
| "PredLabel": all_preds.numpy(), | |
| **{f"Prob_Class_{i}": all_probs[:, i].numpy() for i in range(3)} | |
| }) | |
| results_df.to_csv("/workspace/ClassifierSegmenter/results/test_predictions.csv", index=False) | |
| print("\n📁 Saved test predictions to 'results/test_predictions.csv'") | |