from __future__ import annotations import json import random from pathlib import Path import numpy as np import pandas as pd import torch import trackio from metrics import expected_calibration_error from model import ProbabilisticTinyCNN, parameter_count from safetensors.torch import save_file from sklearn.metrics import accuracy_score, log_loss, roc_auc_score from torch.nn import functional as F from torch.utils.data import DataLoader, TensorDataset PROJECT_DIR = Path(__file__).resolve().parent ROOT_DIR = PROJECT_DIR.parents[1] DATA_DIR = ROOT_DIR / "projects" / "tiny-vision-foundry" / "data" ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "uncertainty-lens" def seed_everything(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) def load_split(name: str) -> tuple[torch.Tensor, torch.Tensor]: frame = pd.read_parquet(DATA_DIR / f"{name}.parquet") pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 16.0 labels = frame["label"].to_numpy(dtype=np.int64, copy=True) return ( torch.from_numpy(pixels.reshape(-1, 1, 8, 8)), torch.from_numpy(labels), ) def train_member( train_pixels: torch.Tensor, train_labels: torch.Tensor, validation_pixels: torch.Tensor, validation_labels: torch.Tensor, seed: int, ) -> tuple[ProbabilisticTinyCNN, float, int]: seed_everything(seed) model = ProbabilisticTinyCNN() loader = DataLoader( TensorDataset(train_pixels, train_labels), batch_size=64, shuffle=True, generator=torch.Generator().manual_seed(seed), ) optimizer = torch.optim.AdamW(model.parameters(), lr=0.003, weight_decay=0.002) epochs = 70 scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) best_accuracy = -1.0 best_epoch = 0 best_state = None for epoch in range(1, epochs + 1): model.train() for pixels, labels in loader: loss = F.cross_entropy(model(pixels), labels) optimizer.zero_grad(set_to_none=True) loss.backward() optimizer.step() scheduler.step() model.eval() with torch.no_grad(): predictions = model(validation_pixels).argmax(dim=1) accuracy = float((predictions == validation_labels).float().mean()) if accuracy > best_accuracy: best_accuracy = accuracy best_epoch = epoch best_state = { key: value.detach().cpu().clone() for key, value in model.state_dict().items() } assert best_state is not None model.load_state_dict(best_state) return model, best_accuracy, best_epoch @torch.inference_mode() def logits_for(model: ProbabilisticTinyCNN, pixels: torch.Tensor) -> torch.Tensor: model.eval() outputs = [] for start in range(0, len(pixels), 256): outputs.append(model(pixels[start : start + 256])) return torch.cat(outputs) def classification_metrics(probabilities: np.ndarray, labels: np.ndarray) -> dict: one_hot = np.eye(10)[labels] return { "accuracy": float(accuracy_score(labels, probabilities.argmax(axis=1))), "negative_log_likelihood": float(log_loss(labels, probabilities, labels=range(10))), "brier_score": float(np.mean(np.sum((probabilities - one_hot) ** 2, axis=1))), "expected_calibration_error": expected_calibration_error( probabilities, labels, ), } def fit_temperature(logits: torch.Tensor, labels: torch.Tensor) -> float: candidates = np.linspace(0.5, 3.0, 251) losses = [ F.cross_entropy(logits / float(temperature), labels).item() for temperature in candidates ] return float(candidates[int(np.argmin(losses))]) def predictive_entropy(probabilities: np.ndarray) -> np.ndarray: clipped = np.clip(probabilities, 1e-9, 1) return -np.sum(clipped * np.log(clipped), axis=1) def ensemble_uncertainty(member_logits: torch.Tensor, temperature: float) -> dict: member_probabilities = torch.softmax(member_logits / temperature, dim=-1).numpy() mean_probabilities = member_probabilities.mean(axis=0) predictive = predictive_entropy(mean_probabilities) member_entropy = np.mean( -np.sum( np.clip(member_probabilities, 1e-9, 1) * np.log(np.clip(member_probabilities, 1e-9, 1)), axis=2, ), axis=0, ) return { "probabilities": mean_probabilities, "predictive_entropy": predictive, "mutual_information": predictive - member_entropy, } def make_ood(test_pixels: torch.Tensor) -> tuple[torch.Tensor, dict[str, int]]: generator = torch.Generator().manual_seed(2036) noise = torch.rand(test_pixels.shape, generator=generator) permutation = torch.randperm(64, generator=generator) scrambled = test_pixels.reshape(len(test_pixels), 64)[:, permutation].reshape( -1, 1, 8, 8, ) return torch.cat([noise, scrambled]), { "uniform_noise": len(noise), "pixel_scrambled": len(scrambled), } def main() -> None: seed_everything(2036) train_pixels, train_labels = load_split("train") validation_pixels, validation_labels = load_split("validation") test_pixels, test_labels = load_split("test") trackio.init( project="uncertainty-lens", name="five-member-deep-ensemble-v1", config={ "members": 5, "parameters_per_member": parameter_count(ProbabilisticTinyCNN()), "calibration_split": "validation", "ood_sets": ["uniform_noise", "pixel_scrambled"], }, ) members = [] member_training = [] for index, seed in enumerate(range(2036, 2041)): model, accuracy, best_epoch = train_member( train_pixels, train_labels, validation_pixels, validation_labels, seed, ) members.append(model) member_training.append( { "member": index, "seed": seed, "best_validation_accuracy": accuracy, "best_epoch": best_epoch, } ) trackio.log( { "member": index, "best_validation_accuracy": accuracy, "best_epoch": best_epoch, } ) validation_logits = torch.stack( [logits_for(model, validation_pixels) for model in members] ) test_logits = torch.stack([logits_for(model, test_pixels) for model in members]) mean_validation_logits = validation_logits.mean(dim=0) temperature = fit_temperature(mean_validation_logits, validation_labels) single_probabilities = torch.softmax(test_logits[0], dim=1).numpy() ensemble_uncalibrated = ensemble_uncertainty(test_logits, temperature=1.0) ensemble_calibrated = ensemble_uncertainty(test_logits, temperature=temperature) ood_pixels, ood_composition = make_ood(test_pixels) ood_logits = torch.stack([logits_for(model, ood_pixels) for model in members]) ood_uncertainty = ensemble_uncertainty(ood_logits, temperature=temperature) clean_entropy = ensemble_calibrated["predictive_entropy"] ood_entropy = ood_uncertainty["predictive_entropy"] detection_labels = np.concatenate( [np.zeros(len(clean_entropy)), np.ones(len(ood_entropy))] ) detection_scores = np.concatenate([clean_entropy, ood_entropy]) ood_roc_auc = roc_auc_score(detection_labels, detection_scores) results = { "model": "Uncertainty Lens Deep Ensemble", "members": len(members), "parameters_per_member": parameter_count(members[0]), "total_parameters": sum(parameter_count(model) for model in members), "member_training": member_training, "temperature": temperature, "single_member_test": classification_metrics( single_probabilities, test_labels.numpy(), ), "ensemble_uncalibrated_test": classification_metrics( ensemble_uncalibrated["probabilities"], test_labels.numpy(), ), "ensemble_calibrated_test": classification_metrics( ensemble_calibrated["probabilities"], test_labels.numpy(), ), "ood_detection": { "composition": ood_composition, "entropy_roc_auc": float(ood_roc_auc), "clean_mean_predictive_entropy": float(clean_entropy.mean()), "ood_mean_predictive_entropy": float(ood_entropy.mean()), "clean_mean_mutual_information": float( ensemble_calibrated["mutual_information"].mean() ), "ood_mean_mutual_information": float( ood_uncertainty["mutual_information"].mean() ), }, } trackio.log( { "test_ensemble_accuracy": results["ensemble_calibrated_test"]["accuracy"], "test_calibrated_ece": results["ensemble_calibrated_test"][ "expected_calibration_error" ], "ood_entropy_roc_auc": ood_roc_auc, } ) trackio.finish() ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) for index, model in enumerate(members): save_file(model.state_dict(), ARTIFACT_DIR / f"member_{index}.safetensors") (ARTIFACT_DIR / "evaluation.json").write_text( json.dumps(results, indent=2), encoding="utf-8", ) print(json.dumps(results, indent=2)) if __name__ == "__main__": main()