| 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() |
|
|