From e07ee19312f462d204aa48e77e739abd102bcabb Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Tue, 21 Jul 2026 09:49:13 -0500 Subject: feat: connect compositional scaling tasks --- experiments/run.py | 63 ++++++++++++++++++++++++++++++++++++------ experiments/synthetic_smoke.py | 46 ++++++++++++++++++++++++++++++ 2 files changed, 101 insertions(+), 8 deletions(-) create mode 100644 experiments/synthetic_smoke.py diff --git a/experiments/run.py b/experiments/run.py index 09e5a30..cb5de4e 100644 --- a/experiments/run.py +++ b/experiments/run.py @@ -23,7 +23,12 @@ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from sdil.core import SDILNet, SDILConfig, sdil_step, neutral_p_update from sdil.baselines import BPNet, dfa_config, evaluate from sdil import probes -from sdil.data import get_dataset, onehot, make_hierarchical, make_teacher_student +from sdil.data import (get_dataset, onehot, make_hierarchical, + make_teacher_student, make_tentmap) + + +REAL_DATASETS = ("mnist", "fmnist", "cifar10") +SYNTHETIC_DATASETS = ("teacher", "hierarchical", "tentmap") def code_provenance(): @@ -42,7 +47,7 @@ def code_provenance(): def build(args, device): - sizes = [args.n_in] + [args.width] * args.depth + [10] + sizes = [args.n_in] + [args.width] * args.depth + [args.n_out] if args.mode == "bp": net = BPNet(sizes, act=args.act, device=device, seed=args.seed, w_scale=args.w_scale, nuis_rho=0.0, residual=bool(args.residual), @@ -72,13 +77,44 @@ def build(args, device): return net, cfg +def load_task(args, device): + """Load a real dataset or construct one fixed synthetic task. + + ``task_seed`` controls examples and the target function; ``seed`` controls + only the student initialization. A depth sweep therefore compares models + on exactly the same compositional problem instead of silently changing the + teacher between model seeds. + """ + if args.dataset in REAL_DATASETS: + return get_dataset( + args.dataset, batch_size=args.batch_size, device=device, + train_limit=args.train_examples or None) + common = dict( + n_train=args.task_train_examples, + n_test=args.task_test_examples, + seed=args.task_seed, + batch_size=args.batch_size, + device=device, + ) + if args.dataset == "teacher": + return 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( + levels=args.task_levels, n_classes=args.task_classes, **common) + if args.dataset == "tentmap": + return make_tentmap( + levels=args.task_levels, n_in=args.task_n_in, **common) + raise ValueError(f"unknown dataset: {args.dataset}") + + def train(args): device = args.device torch.manual_seed(args.seed) - train_loader, test_loader, n_in, n_out = get_dataset( - args.dataset, batch_size=args.batch_size, device=device, - train_limit=args.train_examples or None) - args.n_in = n_in + train_loader, test_loader, n_in, n_out = 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 @@ -158,15 +194,26 @@ def train(args): def get_args(): p = argparse.ArgumentParser() p.add_argument("--mode", default="sdil", choices=["bp", "dfa", "sdil"]) - p.add_argument("--dataset", default="mnist", choices=["mnist", "fmnist", "cifar10"]) + p.add_argument("--dataset", default="mnist", + choices=list(REAL_DATASETS + SYNTHETIC_DATASETS)) p.add_argument("--depth", type=int, default=3) # hidden layers p.add_argument("--width", type=int, default=256) - p.add_argument("--act", default="tanh", choices=["tanh", "gelu", "silu"]) + p.add_argument("--act", default="tanh", choices=["tanh", "gelu", "silu", "relu"]) p.add_argument("--residual", type=int, default=0) # skip connections (deep no-BN) p.add_argument("--epochs", type=int, default=15) 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("--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) + p.add_argument("--task_test_examples", type=int, default=10000) + p.add_argument("--task_levels", type=int, default=8) + p.add_argument("--task_n_in", type=int, default=128) + p.add_argument("--task_classes", type=int, default=10) + p.add_argument("--teacher_depth", type=int, default=8) + p.add_argument("--teacher_width", type=int, default=64) + p.add_argument("--teacher_residual", type=int, default=1) p.add_argument("--eta", type=float, default=0.05) p.add_argument("--eta_A", type=float, default=0.02) p.add_argument("--eta_P", type=float, default=0.002) diff --git a/experiments/synthetic_smoke.py b/experiments/synthetic_smoke.py new file mode 100644 index 0000000..c118813 --- /dev/null +++ b/experiments/synthetic_smoke.py @@ -0,0 +1,46 @@ +"""Fast checks for the synthetic depth-scaling task plumbing.""" +import os +import sys + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from experiments.run import build, get_args, load_task + + +def check_task(dataset, expected_in, expected_out, extra): + argv = [ + "run.py", "--dataset", dataset, "--device", "cpu", + "--task_train_examples", "64", "--task_test_examples", "32", + "--batch_size", "16", "--depth", "2", "--width", "12", + ] + extra + old = sys.argv + try: + sys.argv = argv + args = get_args() + finally: + sys.argv = old + train, test, n_in, n_out = load_task(args, "cpu") + assert (n_in, n_out) == (expected_in, expected_out) + args.n_in, args.n_out = n_in, n_out + x, y = next(iter(train)) + assert x.shape == (16, n_in) + assert y.min() >= 0 and y.max() < n_out + for mode in ("bp", "dfa", "sdil"): + args.mode = mode + net, _ = build(args, "cpu") + assert net.logits(x).shape == (16, n_out) + assert sum(batch_y.numel() for _, batch_y in test) == 32 + print(f"{dataset}: input={n_in}, classes={n_out}, fixed task seed={args.task_seed}") + + +if __name__ == "__main__": + torch.manual_seed(0) + check_task("teacher", 20, 4, + ["--task_n_in", "20", "--task_classes", "4", + "--teacher_depth", "3", "--teacher_width", "10"]) + check_task("hierarchical", 16, 4, + ["--task_levels", "4", "--task_classes", "4"]) + check_task("tentmap", 5, 2, + ["--task_levels", "4", "--task_n_in", "5"]) + print("ALL SYNTHETIC TASK CHECKS PASSED") -- cgit v1.2.3