From 98dddcfabb39c13f652b3aa44f0926ee04428c06 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 01:38:39 -0500 Subject: experiments: isolate frozen validation protocols --- experiments/run.py | 75 +++++++++++++++++++++++++++++++++++++++++------------- 1 file changed, 57 insertions(+), 18 deletions(-) (limited to 'experiments/run.py') diff --git a/experiments/run.py b/experiments/run.py index cc3e475..c66acca 100644 --- a/experiments/run.py +++ b/experiments/run.py @@ -24,7 +24,7 @@ from sdil.core import SDILNet, SDILConfig, sdil_step, neutral_p_update from sdil.baselines import BPNet, dfa_config, evaluate from sdil.local_baselines import FANet from sdil import probes -from sdil.data import (get_dataset, onehot, make_hierarchical, +from sdil.data import (get_dataset_splits, onehot, make_hierarchical, make_teacher_student, make_tentmap) @@ -95,9 +95,18 @@ def load_task(args, device): teacher between model seeds. """ if args.dataset in REAL_DATASETS: - return get_dataset( + train, validation, test, n_in, n_out, split = get_dataset_splits( args.dataset, batch_size=args.batch_size, device=device, - train_limit=args.train_examples or None) + train_limit=args.train_examples or None, + val_examples=args.val_examples, split_seed=args.split_seed) + if args.eval_split == "validation": + if validation is None: + raise ValueError("--eval_split validation requires --val_examples > 0") + evaluation = validation + else: + evaluation = test + split["evaluation_split"] = args.eval_split + return train, evaluation, n_in, n_out, split common = dict( n_train=args.task_train_examples, n_test=args.task_test_examples, @@ -106,32 +115,46 @@ def load_task(args, device): device=device, ) if args.dataset == "teacher": - return make_teacher_student( + task = make_teacher_student( n_in=args.task_n_in, n_classes=args.task_classes, t_depth=args.teacher_depth, t_width=args.teacher_width, residual=bool(args.teacher_residual), **common) - if args.dataset == "hierarchical": - return make_hierarchical( + elif args.dataset == "hierarchical": + task = make_hierarchical( levels=args.task_levels, n_classes=args.task_classes, **common) - if args.dataset == "tentmap": - return make_tentmap( + elif args.dataset == "tentmap": + task = make_tentmap( levels=args.task_levels, n_in=args.task_n_in, **common) - raise ValueError(f"unknown dataset: {args.dataset}") + else: + raise ValueError(f"unknown dataset: {args.dataset}") + train, test, n_in, n_out = task + split = { + "dataset": args.dataset, + "task_seed": args.task_seed, + "train_examples": args.task_train_examples, + "test_examples": args.task_test_examples, + "evaluation_split": "test", + "synthetic_generator": True, + } + if args.eval_split != "test" or args.val_examples: + raise ValueError("synthetic validation splits are not implemented yet") + return train, test, n_in, n_out, split def train(args): device = args.device torch.manual_seed(args.seed) - train_loader, test_loader, n_in, n_out = load_task(args, device) + train_loader, eval_loader, n_in, n_out, split = load_task(args, device) args.n_in, args.n_out = n_in, n_out net, cfg = build(args, device) # a fixed probe batch for stable alignment tracking - px, py = next(iter(test_loader)) + px, py = next(iter(eval_loader)) px, py = px[:args.probe_bs].to(device), py[:args.probe_bs].to(device) poh = onehot(py, n_out, device=device) - log = {"args": vars(args), "provenance": code_provenance(), "steps": [], "final": {}} + log = {"args": vars(args), "split": split, "provenance": code_provenance(), + "steps": [], "final": {}} step = 0 prev_error = None t0 = time.time() @@ -182,8 +205,9 @@ def train(args): if args.max_steps and step >= args.max_steps: break - acc, tloss = evaluate(net, test_loader) - msg = f"[{args.tag}] epoch {epoch} step {step} loss {loss:.4f} test_acc {acc:.4f}" + acc, tloss = evaluate(net, eval_loader) + metric = "val_acc" if args.eval_split == "validation" else "test_acc" + msg = f"[{args.tag}] epoch {epoch} step {step} loss {loss:.4f} {metric} {acc:.4f}" if args.mode == "fa": al = probes.fa_alignment_report(net, px, py, poh) meancos = sum(al["cos_fa_negg"]) / len(al["cos_fa_negg"]) @@ -193,10 +217,21 @@ def train(args): meancos = sum(al["cos_r_negg"]) / len(al["cos_r_negg"]) msg += f" mean_cos(r,-g) {meancos:+.3f} per-layer {['%.2f'%v for v in al['cos_r_negg']]}" print(msg, flush=True) - log["steps"].append({"epoch_end": epoch, "step": step, "test_acc": acc, "test_loss": tloss}) + record = {"epoch_end": epoch, "step": step, "eval_split": args.eval_split, + "eval_acc": acc, "eval_loss": tloss} + if args.eval_split == "test": + record.update({"test_acc": acc, "test_loss": tloss}) + else: + record.update({"val_acc": acc, "val_loss": tloss}) + log["steps"].append(record) - acc, tloss = evaluate(net, test_loader) - log["final"] = {"test_acc": acc, "test_loss": tloss, "wall_s": time.time() - t0} + acc, tloss = evaluate(net, eval_loader) + log["final"] = {"eval_split": args.eval_split, "eval_acc": acc, "eval_loss": tloss, + "wall_s": time.time() - t0} + if args.eval_split == "test": + log["final"].update({"test_acc": acc, "test_loss": tloss}) + else: + log["final"].update({"val_acc": acc, "val_loss": tloss}) if args.mode == "fa": log["final"].update(probes.fa_alignment_report(net, px, py, poh)) elif args.mode != "bp": @@ -212,7 +247,7 @@ def train(args): outpath = os.path.join(args.outdir, f"{args.tag}.json") with open(outpath, "w") as f: json.dump(log, f) - print(f"[{args.tag}] DONE test_acc={acc:.4f} -> {outpath}", flush=True) + print(f"[{args.tag}] DONE {args.eval_split}_acc={acc:.4f} -> {outpath}", flush=True) return log @@ -229,6 +264,10 @@ def get_args(): p.add_argument("--batch_size", type=int, default=128) p.add_argument("--train_examples", type=int, default=0, help="0 uses the full training split") + p.add_argument("--val_examples", type=int, default=0, + help="stratified validation examples held out from training") + p.add_argument("--split_seed", type=int, default=2027) + p.add_argument("--eval_split", default="test", choices=["validation", "test"]) p.add_argument("--task_seed", type=int, default=0, help="fixed target/data seed, separate from student --seed") p.add_argument("--task_train_examples", type=int, default=50000) -- cgit v1.2.3