summaryrefslogtreecommitdiff
path: root/experiments/transformer_crossover_t2_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/transformer_crossover_t2_smoke.py')
-rw-r--r--experiments/transformer_crossover_t2_smoke.py65
1 files changed, 65 insertions, 0 deletions
diff --git a/experiments/transformer_crossover_t2_smoke.py b/experiments/transformer_crossover_t2_smoke.py
new file mode 100644
index 0000000..febca40
--- /dev/null
+++ b/experiments/transformer_crossover_t2_smoke.py
@@ -0,0 +1,65 @@
+#!/usr/bin/env python3
+"""Deterministic audit of the complete Transformer T2 registry."""
+import os
+import sys
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from experiments.transformer_crossover_t2 import DEPTHS, METHODS, t2_jobs
+
+
+RATES = {
+ "bp": 0.001,
+ "fa": 0.001,
+ "dfa": 0.001,
+ "pepita": 0.0003,
+ "ff": 0.0003,
+ "ep": 0.0003,
+ "dualprop": 0.0003,
+ "clean_kp": 0.001,
+ "sdil": 0.001,
+}
+
+
+def value(command, flag):
+ index = command.index(flag)
+ return command[index + 1]
+
+
+def main():
+ jobs = t2_jobs(RATES)
+ assert len(jobs) == 27
+ assert {(job["method"], job["depth"]) for job in jobs} == {
+ (method, depth) for method in METHODS for depth in DEPTHS
+ }
+ assert len({job["experiment_name"] for job in jobs}) == 27
+ assert len({job["output"] for job in jobs}) == 27
+ for job in jobs:
+ command = job["command"]
+ rate = RATES[job["method"]]
+ assert value(command, "--depth") == str(job["depth"])
+ assert value(command, "--train_steps") == "5000"
+ assert value(command, "--batch_size") == "32"
+ assert value(command, "--context_length") == "64"
+ assert value(command, "--lr") == str(rate)
+ assert value(command, "--schedule") == "cosine"
+ assert value(command, "--min_lr") == str(rate * 0.1)
+ assert value(command, "--warmup_steps") == "100"
+ assert value(command, "--eval_every") == "0"
+ assert value(command, "--max_val_batches") == "0"
+ assert job["timeout_seconds"] == 48 * 60 * 60
+ assert "--test" not in command
+ assert value(command, "--traffic_ratio") == "4.0"
+ assert value(command, "--ep_free_steps") == "20"
+ assert value(command, "--ep_nudge_steps") == "4"
+ assert value(command, "--dp_inference_passes") == "16"
+ shard0 = jobs[0::2]
+ shard1 = jobs[1::2]
+ assert len(shard0) == 14 and len(shard1) == 13
+ for method in ("ep", "dualprop", "ff"):
+ assert any(job["method"] == method for job in shard0)
+ assert any(job["method"] == method for job in shard1)
+ print("Transformer T2 registry: 27/27 cells; schedules and shards exact")
+
+
+if __name__ == "__main__":
+ main()