| from __future__ import annotations |
|
|
| import json |
| import shutil |
| from pathlib import Path |
|
|
| import numpy as np |
| import pandas as pd |
| import torch |
| import trackio |
| from model import HighwayNetwork, PlainDeepNetwork, parameter_count |
| from safetensors.torch import save_file |
| 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] |
| VISION_DATA = ROOT_DIR / "projects" / "tiny-vision-foundry" / "data" |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "highway-depth-pocket" |
| DATA_DIR = PROJECT_DIR / "data" |
| SEEDS = [2237, 2239, 2243, 2251, 2267] |
|
|
|
|
| def load_split(name: str, *, shuffle: bool, seed: int) -> DataLoader: |
| frame = pd.read_parquet(VISION_DATA / f"{name}.parquet") |
| pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 16 |
| labels = frame["label"].to_numpy(dtype=np.int64, copy=True) |
| return DataLoader( |
| TensorDataset(torch.from_numpy(pixels), torch.from_numpy(labels)), |
| batch_size=128, |
| shuffle=shuffle, |
| generator=torch.Generator().manual_seed(seed), |
| ) |
|
|
|
|
| @torch.inference_mode() |
| def evaluate( |
| model: torch.nn.Module, |
| loader: DataLoader, |
| *, |
| noise: float = 0.0, |
| seed: int, |
| ) -> dict: |
| model.eval() |
| correct = 0 |
| total = 0 |
| generator = torch.Generator().manual_seed(seed + 1) |
| for pixels, labels in loader: |
| if noise: |
| pixels = ( |
| pixels + torch.randn(pixels.shape, generator=generator) * noise |
| ).clamp(0, 1) |
| prediction = model(pixels).argmax(1) |
| correct += int((prediction == labels).sum()) |
| total += len(labels) |
| return {"accuracy": correct / total, "examples": total} |
|
|
|
|
| def train_one( |
| model: torch.nn.Module, |
| seed: int, |
| ) -> tuple[dict[str, torch.Tensor], int]: |
| train_loader = load_split("train", shuffle=True, seed=seed) |
| validation_loader = load_split("validation", shuffle=False, seed=seed) |
| optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3, weight_decay=1e-4) |
| best = -1.0 |
| best_epoch = 0 |
| best_state = None |
| for epoch in range(1, 121): |
| model.train() |
| for pixels, labels in train_loader: |
| loss = F.cross_entropy(model(pixels), labels) |
| optimizer.zero_grad(set_to_none=True) |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 5) |
| optimizer.step() |
| validation = evaluate(model, validation_loader, seed=seed) |
| if validation["accuracy"] > best: |
| best = validation["accuracy"] |
| best_epoch = epoch |
| best_state = { |
| name: value.detach().cpu().clone() |
| for name, value in model.state_dict().items() |
| } |
| assert best_state is not None |
| return best_state, best_epoch |
|
|
|
|
| def aggregate(runs: list[dict]) -> dict: |
| return { |
| "clean_accuracy_mean": float( |
| np.mean([run["clean_accuracy"] for run in runs]) |
| ), |
| "clean_accuracy_std": float( |
| np.std([run["clean_accuracy"] for run in runs]) |
| ), |
| "noise_0.20_accuracy_mean": float( |
| np.mean([run["noise_0.20_accuracy"] for run in runs]) |
| ), |
| "noise_0.20_accuracy_std": float( |
| np.std([run["noise_0.20_accuracy"] for run in runs]) |
| ), |
| "successful_seeds_above_95_percent": sum( |
| run["clean_accuracy"] >= 0.95 for run in runs |
| ), |
| "seeds": len(runs), |
| } |
|
|
|
|
| def main() -> None: |
| torch.set_num_threads(1) |
| assert parameter_count(HighwayNetwork()) == parameter_count( |
| PlainDeepNetwork() |
| ) == 19_306 |
| trackio.init( |
| project="highway-depth-pocket", |
| name="gated-versus-plain-depth-v1", |
| config={ |
| "parameters_per_model": 19_306, |
| "highway_layers": 8, |
| "plain_layers": 16, |
| "seeds": SEEDS, |
| }, |
| ) |
| all_runs = {"highway": [], "plain": []} |
| saved_states = {} |
| for seed in SEEDS: |
| test_loader = load_split("test", shuffle=False, seed=seed) |
| for name, constructor in [ |
| ("highway", HighwayNetwork), |
| ("plain", PlainDeepNetwork), |
| ]: |
| torch.manual_seed(seed) |
| model = constructor() |
| state, best_epoch = train_one(model, seed) |
| model.load_state_dict(state) |
| run = { |
| "seed": seed, |
| "best_epoch": best_epoch, |
| "clean_accuracy": evaluate(model, test_loader, seed=seed)["accuracy"], |
| "noise_0.20_accuracy": evaluate( |
| model, test_loader, noise=0.20, seed=seed |
| )["accuracy"], |
| } |
| all_runs[name].append(run) |
| trackio.log( |
| { |
| "seed": seed, |
| "variant": name, |
| **{key: value for key, value in run.items() if key != "seed"}, |
| } |
| ) |
| if seed == SEEDS[0]: |
| saved_states[name] = state |
| highway = HighwayNetwork() |
| highway.load_state_dict(saved_states["highway"]) |
| test_frame = pd.read_parquet(VISION_DATA / "test.parquet") |
| test_pixels = torch.from_numpy( |
| np.stack(test_frame["image"].to_numpy()).astype(np.float32) / 16 |
| ) |
| with torch.inference_mode(): |
| _, gates = highway(test_pixels, return_gates=True) |
| gate_profile = gates.mean(dim=(0, 2)).tolist() |
| results = { |
| "highway": { |
| "parameters": parameter_count(highway), |
| "aggregate": aggregate(all_runs["highway"]), |
| "runs": all_runs["highway"], |
| "mean_transform_gate_by_layer": gate_profile, |
| }, |
| "plain": { |
| "parameters": parameter_count(PlainDeepNetwork()), |
| "aggregate": aggregate(all_runs["plain"]), |
| "runs": all_runs["plain"], |
| }, |
| } |
| report = { |
| "experiment": "Highway gating versus parameter-matched plain depth", |
| "results": results, |
| } |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) |
| DATA_DIR.mkdir(parents=True, exist_ok=True) |
| save_file(saved_states["highway"], ARTIFACT_DIR / "highway.safetensors") |
| save_file(saved_states["plain"], ARTIFACT_DIR / "plain.safetensors") |
| (ARTIFACT_DIR / "evaluation.json").write_text( |
| json.dumps(report, indent=2), encoding="utf-8" |
| ) |
| shutil.copy2(VISION_DATA / "test.parquet", DATA_DIR / "test.parquet") |
| trackio.log( |
| { |
| "highway_clean_mean": results["highway"]["aggregate"][ |
| "clean_accuracy_mean" |
| ], |
| "plain_clean_mean": results["plain"]["aggregate"]["clean_accuracy_mean"], |
| } |
| ) |
| trackio.finish() |
| print(json.dumps(report, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|
|
|