summaryrefslogtreecommitdiff
path: root/experiments/run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-21 09:49:13 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-21 09:49:13 -0500
commite07ee19312f462d204aa48e77e739abd102bcabb (patch)
tree9febb8456dec0c31ee87665857d5d077e9023b31 /experiments/run.py
parentda2c397985e8ee8e05fef065220481328804c2bf (diff)
feat: connect compositional scaling tasks
Diffstat (limited to 'experiments/run.py')
-rw-r--r--experiments/run.py63
1 files changed, 55 insertions, 8 deletions
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)