summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/run.py63
-rw-r--r--experiments/synthetic_smoke.py46
2 files changed, 101 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)
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")