ARotting's picture
Publish Five-member calibrated tiny vision ensemble
7671040 verified
Raw
History Blame Contribute Delete
9.65 kB
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()