| 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 model import RobustTinyCNN, parameter_count |
| from safetensors.torch import load_file, save_file |
| from sklearn.metrics import accuracy_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] |
| SOURCE_DIR = ROOT_DIR / "projects" / "tiny-vision-foundry" |
| DATA_DIR = SOURCE_DIR / "data" |
| BASE_WEIGHTS = SOURCE_DIR / "artifacts" / "tiny-student-scratch" / "model.safetensors" |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "pixel-shield-robust" |
|
|
|
|
| def seed_everything(seed: int) -> None: |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
|
|
|
|
| def load_split(name: str, *, shuffle: bool, batch_size: int) -> DataLoader: |
| frame = pd.read_parquet(DATA_DIR / f"{name}.parquet") |
| pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 16.0 |
| images = torch.from_numpy(pixels.reshape(-1, 1, 8, 8)) |
| labels = torch.from_numpy(frame["label"].to_numpy(dtype=np.int64, copy=True)) |
| return DataLoader( |
| TensorDataset(images, labels), |
| batch_size=batch_size, |
| shuffle=shuffle, |
| generator=torch.Generator().manual_seed(2026), |
| ) |
|
|
|
|
| def fgsm( |
| model: RobustTinyCNN, |
| pixels: torch.Tensor, |
| labels: torch.Tensor, |
| epsilon: float, |
| ) -> torch.Tensor: |
| attacked = pixels.detach().clone().requires_grad_(True) |
| loss = F.cross_entropy(model(attacked), labels) |
| gradient = torch.autograd.grad(loss, attacked)[0] |
| return torch.clamp(attacked + epsilon * gradient.sign(), 0, 1).detach() |
|
|
|
|
| def evaluate(model: RobustTinyCNN, loader: DataLoader, epsilon: float) -> float: |
| model.eval() |
| labels, predictions = [], [] |
| for pixels, targets in loader: |
| if epsilon: |
| pixels = fgsm(model, pixels, targets, epsilon) |
| with torch.no_grad(): |
| logits = model(pixels) |
| labels.extend(targets.tolist()) |
| predictions.extend(logits.argmax(dim=1).tolist()) |
| return float(accuracy_score(labels, predictions)) |
|
|
|
|
| def benchmark(model: RobustTinyCNN, loader: DataLoader) -> dict[str, float]: |
| return { |
| f"epsilon_{epsilon:.2f}": evaluate(model, loader, epsilon) |
| for epsilon in [0.0, 0.05, 0.10, 0.15, 0.20, 0.25] |
| } |
|
|
|
|
| def main() -> None: |
| seed_everything(2030) |
| if not BASE_WEIGHTS.exists(): |
| raise FileNotFoundError("Train Tiny Vision Foundry before running PixelShield.") |
| train_loader = load_split("train", shuffle=True, batch_size=64) |
| validation_loader = load_split("validation", shuffle=False, batch_size=256) |
| test_loader = load_split("test", shuffle=False, batch_size=256) |
| baseline = RobustTinyCNN() |
| baseline.load_state_dict(load_file(BASE_WEIGHTS)) |
| baseline_metrics = benchmark(baseline, test_loader) |
|
|
| robust = RobustTinyCNN() |
| robust.load_state_dict(load_file(BASE_WEIGHTS)) |
| optimizer = torch.optim.AdamW(robust.parameters(), lr=0.0015, weight_decay=0.002) |
| epochs = 50 |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) |
| best_score = -1.0 |
| best_epoch = 0 |
| best_state = None |
| trackio.init( |
| project="pixel-shield", |
| name="tiny-student-fgsm-training-v1", |
| config={ |
| "parameters": parameter_count(robust), |
| "epochs": epochs, |
| "training_epsilon": 0.15, |
| "clean_adversarial_mix": "50/50", |
| }, |
| ) |
| for epoch in range(1, epochs + 1): |
| robust.train() |
| running_loss = 0.0 |
| examples = 0 |
| for pixels, labels in train_loader: |
| adversarial = fgsm(robust, pixels, labels, epsilon=0.15) |
| combined_pixels = torch.cat([pixels, adversarial]) |
| combined_labels = torch.cat([labels, labels]) |
| logits = robust(combined_pixels) |
| loss = F.cross_entropy(logits, combined_labels) |
| optimizer.zero_grad(set_to_none=True) |
| loss.backward() |
| optimizer.step() |
| running_loss += loss.item() * len(combined_labels) |
| examples += len(combined_labels) |
| scheduler.step() |
| clean_accuracy = evaluate(robust, validation_loader, epsilon=0.0) |
| robust_accuracy = evaluate(robust, validation_loader, epsilon=0.15) |
| selection_score = 0.35 * clean_accuracy + 0.65 * robust_accuracy |
| trackio.log( |
| { |
| "epoch": epoch, |
| "train_loss": running_loss / examples, |
| "validation_clean_accuracy": clean_accuracy, |
| "validation_fgsm_0.15_accuracy": robust_accuracy, |
| "selection_score": selection_score, |
| "learning_rate": scheduler.get_last_lr()[0], |
| } |
| ) |
| if selection_score > best_score: |
| best_score = selection_score |
| best_epoch = epoch |
| best_state = { |
| key: value.detach().cpu().clone() |
| for key, value in robust.state_dict().items() |
| } |
| trackio.finish() |
| assert best_state is not None |
| robust.load_state_dict(best_state) |
| robust_metrics = benchmark(robust, test_loader) |
| results = { |
| "model": "PixelShield Robust Tiny CNN", |
| "parameters": parameter_count(robust), |
| "best_epoch": best_epoch, |
| "attack": "white-box FGSM over normalized [0,1] pixels", |
| "baseline": baseline_metrics, |
| "adversarially_trained": robust_metrics, |
| "accuracy_delta": { |
| key: robust_metrics[key] - baseline_metrics[key] for key in baseline_metrics |
| }, |
| } |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) |
| save_file(robust.state_dict(), ARTIFACT_DIR / "model.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() |
|
|