pixel-shield / source /train.py
ARotting's picture
Publish Adversarially trained 2.2K parameter vision model
a4d55b6 verified
Raw
History Blame Contribute Delete
6.04 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 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()