|
| 1 | +"""Run Knowledge Distillation experiments across multiple seeds.""" |
| 2 | + |
| 3 | +import argparse |
| 4 | +import signal |
| 5 | +import subprocess |
| 6 | +import sys |
| 7 | +from pathlib import Path |
| 8 | + |
| 9 | +from experiments.config import DATASETS, DEFAULTS, MODELS, SEEDS |
| 10 | + |
| 11 | +# Global state for interrupt handler |
| 12 | +_sweep_state = {"completed": 0, "failed": 0, "skipped": 0, "total": 0} |
| 13 | + |
| 14 | + |
| 15 | +def find_teacher_path(model: str, dataset: str, results_dir: Path, teacher_seed: int) -> Path | None: |
| 16 | + """Find the FP32 teacher checkpoint for a given model/dataset.""" |
| 17 | + teacher_path = results_dir / dataset / model / f"std_s{teacher_seed}" / "best_model.pth" |
| 18 | + return teacher_path if teacher_path.exists() else None |
| 19 | + |
| 20 | + |
| 21 | +def run_kd_experiment( |
| 22 | + model: str, |
| 23 | + dataset: str, |
| 24 | + seed: int, |
| 25 | + teacher_path: Path, |
| 26 | + output_dir: Path, |
| 27 | + epochs: int, |
| 28 | + temperature: float, |
| 29 | + alpha: float, |
| 30 | + index: int, |
| 31 | + total: int, |
| 32 | + dry_run: bool = False, |
| 33 | +) -> str: |
| 34 | + """Run a single KD experiment. Returns 'completed', 'skipped', or 'failed'.""" |
| 35 | + run_name = f"{model}_kd_{dataset}_s{seed}" |
| 36 | + run_dir = output_dir / dataset / model / f"bit_kd_s{seed}" |
| 37 | + prefix = f"[{index}/{total}]" |
| 38 | + |
| 39 | + if (run_dir / "results.json").exists(): |
| 40 | + print(f"{prefix} Skipping {run_name} (already completed)") |
| 41 | + return "skipped" |
| 42 | + |
| 43 | + cmd = [ |
| 44 | + sys.executable, |
| 45 | + "-m", |
| 46 | + "experiments.train_kd", |
| 47 | + "--model", |
| 48 | + model, |
| 49 | + "--dataset", |
| 50 | + dataset, |
| 51 | + "--teacher-path", |
| 52 | + str(teacher_path), |
| 53 | + "--seed", |
| 54 | + str(seed), |
| 55 | + "--epochs", |
| 56 | + str(epochs), |
| 57 | + "--temperature", |
| 58 | + str(temperature), |
| 59 | + "--alpha", |
| 60 | + str(alpha), |
| 61 | + "--output-dir", |
| 62 | + str(run_dir), |
| 63 | + "--quiet", |
| 64 | + ] |
| 65 | + |
| 66 | + print(f"{prefix} Running {run_name} (teacher: {teacher_path.name})...", flush=True) |
| 67 | + if dry_run: |
| 68 | + print(f" Command: {' '.join(cmd)}") |
| 69 | + return "completed" |
| 70 | + |
| 71 | + best_acc_line = None |
| 72 | + error_line = None |
| 73 | + with subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True) as proc: |
| 74 | + for line in proc.stdout: # type: ignore[union-attr] |
| 75 | + line = line.rstrip() |
| 76 | + if "Progress:" in line or "Training complete" in line or "Skipping" in line: |
| 77 | + print(f" {line}") |
| 78 | + if "Best accuracy" in line: |
| 79 | + best_acc_line = line |
| 80 | + proc.wait() |
| 81 | + if proc.returncode != 0 and proc.stderr: |
| 82 | + error_line = proc.stderr.read().strip().splitlines()[-1] if proc.stderr else None |
| 83 | + |
| 84 | + if proc.returncode == 0: |
| 85 | + if best_acc_line: |
| 86 | + print(f" Done: {best_acc_line.split('Best accuracy:')[-1].strip()}") |
| 87 | + else: |
| 88 | + print(" Done") |
| 89 | + return "completed" |
| 90 | + else: |
| 91 | + print(" FAILED") |
| 92 | + if error_line: |
| 93 | + print(f" Error: {error_line}") |
| 94 | + return "failed" |
| 95 | + |
| 96 | + |
| 97 | +def print_summary() -> None: |
| 98 | + """Print sweep summary.""" |
| 99 | + state = _sweep_state |
| 100 | + print(f"\n{'=' * 40}") |
| 101 | + print(f"KD Sweep: {state['completed']} completed, {state['skipped']} skipped, {state['failed']} failed") |
| 102 | + print(f"{'=' * 40}") |
| 103 | + |
| 104 | + |
| 105 | +def handle_interrupt(_signum: int, _frame: object) -> None: |
| 106 | + """Handle Ctrl+C gracefully.""" |
| 107 | + print("\n\nInterrupted by user.") |
| 108 | + print_summary() |
| 109 | + sys.exit(1) |
| 110 | + |
| 111 | + |
| 112 | +def main() -> None: |
| 113 | + parser = argparse.ArgumentParser(description="Run KD experiments across seeds") |
| 114 | + parser.add_argument("--output-dir", default="results/raw_kd") |
| 115 | + parser.add_argument("--results-dir", default="results/raw", help="Dir with FP32 teacher checkpoints") |
| 116 | + parser.add_argument("--epochs", type=int, default=DEFAULTS.epochs) |
| 117 | + parser.add_argument("--models", nargs="+", default=MODELS, choices=MODELS) |
| 118 | + parser.add_argument("--datasets", nargs="+", default=DATASETS, choices=DATASETS) |
| 119 | + parser.add_argument("--seeds", nargs="+", type=int, default=SEEDS, help="Student seeds") |
| 120 | + parser.add_argument("--teacher-seed", type=int, default=42, help="Teacher seed (default: 42)") |
| 121 | + parser.add_argument("--temperature", type=float, default=4.0) |
| 122 | + parser.add_argument("--alpha", type=float, default=0.9) |
| 123 | + parser.add_argument("--dry-run", action="store_true", help="Print commands without running") |
| 124 | + args = parser.parse_args() |
| 125 | + |
| 126 | + signal.signal(signal.SIGINT, handle_interrupt) |
| 127 | + |
| 128 | + results_dir = Path(args.results_dir) |
| 129 | + output_dir = Path(args.output_dir) |
| 130 | + |
| 131 | + # Build experiment list |
| 132 | + experiments = [] |
| 133 | + for model in args.models: |
| 134 | + for dataset in args.datasets: |
| 135 | + teacher_path = find_teacher_path(model, dataset, results_dir, args.teacher_seed) |
| 136 | + if teacher_path is None: |
| 137 | + print(f"Warning: No teacher found for {model}/{dataset} (seed {args.teacher_seed}), skipping") |
| 138 | + continue |
| 139 | + for seed in args.seeds: |
| 140 | + experiments.append((model, dataset, seed, teacher_path)) |
| 141 | + |
| 142 | + total = len(experiments) |
| 143 | + _sweep_state["total"] = total |
| 144 | + |
| 145 | + print(f"Total KD experiments: {total}") |
| 146 | + print(f"Models: {args.models}") |
| 147 | + print(f"Datasets: {args.datasets}") |
| 148 | + print(f"Student seeds: {args.seeds}") |
| 149 | + print(f"Teacher seed: {args.teacher_seed}") |
| 150 | + print(f"KD params: T={args.temperature}, alpha={args.alpha}") |
| 151 | + print(f"Epochs: {args.epochs}") |
| 152 | + print() |
| 153 | + |
| 154 | + for i, (model, dataset, seed, teacher_path) in enumerate(experiments, 1): |
| 155 | + result = run_kd_experiment( |
| 156 | + model, |
| 157 | + dataset, |
| 158 | + seed, |
| 159 | + teacher_path, |
| 160 | + output_dir, |
| 161 | + args.epochs, |
| 162 | + args.temperature, |
| 163 | + args.alpha, |
| 164 | + i, |
| 165 | + total, |
| 166 | + args.dry_run, |
| 167 | + ) |
| 168 | + _sweep_state[result] += 1 |
| 169 | + |
| 170 | + print_summary() |
| 171 | + |
| 172 | + |
| 173 | +if __name__ == "__main__": |
| 174 | + main() |
0 commit comments