summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:23:50 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 14:23:50 -0500
commitbf0243d4aa10221f53ff14acbb2ebfbf32e30b80 (patch)
treeb686bebbe010595ef66d81db98c552386442fe26
parente1e2c2eb9011eb75ff2108b06c9ac3fefa596767 (diff)
experiment: add complete Transformer crossover runner
-rw-r--r--experiments/transformer_crossover_native.py502
1 files changed, 502 insertions, 0 deletions
diff --git a/experiments/transformer_crossover_native.py b/experiments/transformer_crossover_native.py
new file mode 100644
index 0000000..43d3e62
--- /dev/null
+++ b/experiments/transformer_crossover_native.py
@@ -0,0 +1,502 @@
+#!/usr/bin/env python3
+"""Validation-only runner for the matched Transformer crossover."""
+import argparse
+import hashlib
+import json
+import math
+import os
+import platform
+import subprocess
+import sys
+import time
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from sdil.transformer import ( # noqa: E402
+ LocalDecoderTransformer,
+ LocalTransformerConfig,
+)
+
+
+METHODS = (
+ "bp", "fa", "dfa", "pepita", "ff", "ep", "dualprop",
+ "clean_kp", "sdil",
+)
+DEFAULT_DATA_DIR = "/home/yurenh2/local-transformer/data/shakespeare_char"
+
+
+def file_sha256(path):
+ digest = hashlib.sha256()
+ with open(path, "rb") as handle:
+ while True:
+ block = handle.read(1024 * 1024)
+ if not block:
+ break
+ digest.update(block)
+ return digest.hexdigest()
+
+
+def git_output(*args):
+ return subprocess.check_output(
+ ["git", *args], text=True,
+ stderr=subprocess.DEVNULL).strip()
+
+
+def sync(device):
+ if str(device).startswith("cuda"):
+ torch.cuda.synchronize(torch.device(device))
+
+
+class TokenData:
+
+ def __init__(self, data_dir, context_length):
+ self.context_length = int(context_length)
+ self.paths = {
+ split: os.path.join(data_dir, f"{split}.bin")
+ for split in ("train", "val")}
+ for path in self.paths.values():
+ if not os.path.isfile(path):
+ raise FileNotFoundError(path)
+ self.arrays = {
+ split: torch.from_numpy(np.array(
+ np.memmap(path, dtype=np.uint16, mode="r"),
+ dtype=np.int64, copy=True))
+ for split, path in self.paths.items()}
+ self.hashes = {
+ split: file_sha256(path)
+ for split, path in self.paths.items()}
+
+ def random_batch(self, batch_size, generator, device):
+ data = self.arrays["train"]
+ upper = data.numel() - self.context_length - 1
+ starts = torch.randint(
+ 0, upper, (batch_size,), generator=generator)
+ offsets = torch.arange(self.context_length + 1)
+ windows = data[starts[:, None] + offsets[None, :]]
+ return (
+ windows[:, :-1].to(device, non_blocking=True),
+ windows[:, 1:].to(device, non_blocking=True))
+
+ def validation_batches(self, batch_size, device, max_batches=0):
+ data = self.arrays["val"]
+ starts = torch.arange(
+ 0, data.numel() - self.context_length,
+ self.context_length)
+ if max_batches:
+ starts = starts[:batch_size * max_batches]
+ offsets = torch.arange(self.context_length + 1)
+ for begin in range(0, starts.numel(), batch_size):
+ selected = starts[begin:begin + batch_size]
+ windows = data[selected[:, None] + offsets[None, :]]
+ yield (
+ windows[:, :-1].to(device, non_blocking=True),
+ windows[:, 1:].to(device, non_blocking=True))
+
+ @property
+ def train_tokens(self):
+ return self.arrays["train"].numel()
+
+ @property
+ def validation_tokens(self):
+ return self.arrays["val"].numel()
+
+
+def build(args):
+ config = LocalTransformerConfig(
+ vocab_size=65,
+ context_length=args.context_length,
+ depth=args.depth,
+ width=args.width,
+ heads=args.heads,
+ mlp_ratio=args.mlp_ratio,
+ dropout=0.0,
+ bias=False,
+ traffic_ratio=args.traffic_ratio,
+ seed=args.model_seed,
+ )
+ return LocalDecoderTransformer(config, args.method).to(args.device)
+
+
+def scheduled_rate(args, step):
+ if args.schedule == "constant":
+ return args.lr
+ if step < args.warmup_steps:
+ return args.lr * (step + 1) / max(args.warmup_steps, 1)
+ progress = (
+ (step - args.warmup_steps)
+ / max(args.train_steps - args.warmup_steps - 1, 1))
+ return args.min_lr + 0.5 * (args.lr - args.min_lr) * (
+ 1.0 + math.cos(math.pi * min(max(progress, 0.0), 1.0)))
+
+
+def train_step(net, optimizer, tokens, targets, args, step, generators):
+ optimizer.zero_grad(set_to_none=True)
+ if args.method in {"bp", "fa", "clean_kp", "sdil"}:
+ output = net(tokens, targets)
+ loss = output["loss"]
+ loss.backward()
+ metrics = {"loss": float(loss.detach())}
+ elif args.method == "dfa":
+ with torch.no_grad():
+ loss = net(tokens, targets)["loss"]
+ metrics = {
+ "loss": float(loss),
+ **net.dfa_gradients(tokens, targets),
+ }
+ elif args.method == "pepita":
+ local = net.pepita_gradients(tokens, targets)
+ metrics = {"loss": local.pop("clean_loss"), **local}
+ elif args.method == "ff":
+ layers = net.ff_num_layers
+ steps_per_layer = math.ceil(args.train_steps / layers)
+ layer = min(step // steps_per_layer, layers - 1)
+ offsets = torch.randint(
+ 1, net.config.vocab_size, targets.shape,
+ generator=generators["negative"], device=targets.device)
+ negative = (targets + offsets) % net.config.vocab_size
+ loss, local = net.ff_local_loss(
+ layer, tokens, targets, negative,
+ threshold=args.ff_threshold)
+ loss.backward()
+ metrics = {
+ "loss": float(loss.detach()),
+ "ff_layer": layer,
+ **local,
+ }
+ elif args.method == "dualprop":
+ local = net.dualprop_gradients(
+ tokens, targets, alpha=args.dp_alpha, beta=args.dp_beta,
+ inference_passes=args.dp_inference_passes)
+ metrics = {"loss": local.pop("clean_loss"), **local}
+ elif args.method == "ep":
+ sign_draw = torch.randint(
+ 0, 2, (), generator=generators["ep"],
+ device=targets.device)
+ sign = 1.0 if int(sign_draw) else -1.0
+ local, _, _ = net.ep_gradients(
+ tokens, targets, ep_beta=args.ep_beta, dt=args.ep_dt,
+ free_steps=args.ep_free_steps,
+ nudge_steps=args.ep_nudge_steps, beta_sign=sign)
+ metrics = {"loss": local.pop("free_mse"), **local}
+ else:
+ raise ValueError(args.method)
+
+ gradients = [
+ parameter.grad for parameter in net.parameters()
+ if parameter.grad is not None]
+ metrics["gradient_tensors"] = len(gradients)
+ metrics["gradients_finite"] = bool(
+ gradients and all(torch.isfinite(value).all() for value in gradients))
+ if not math.isfinite(metrics["loss"]) or not metrics["gradients_finite"]:
+ return metrics
+ for group in optimizer.param_groups:
+ group["lr"] = scheduled_rate(args, step)
+ optimizer.step()
+ return metrics
+
+
+def evaluate(net, data, args):
+ was_training = net.training
+ net.eval()
+ loss_sum = 0.0
+ correct = 0
+ tokens_seen = 0
+ batches = 0
+ candidate_presentations = 0
+ relaxation_passes = 0
+ for token_batch, target_batch in data.validation_batches(
+ args.eval_batch_size, args.device, args.max_val_batches):
+ if args.method == "ff":
+ with torch.no_grad():
+ scores = net.ff_candidate_scores(
+ token_batch,
+ score_from_layer=args.ff_score_from_layer)
+ candidate_presentations += (
+ target_batch.numel() * net.config.vocab_size)
+ elif args.method == "ep":
+ states = net.ep_settle(
+ token_batch, target_batch, beta=0.0,
+ steps=args.ep_free_steps, dt=args.ep_dt)
+ scores = states[-1]
+ relaxation_passes += (
+ target_batch.numel() * args.ep_free_steps)
+ else:
+ with torch.no_grad():
+ scores = net(token_batch)["logits"]
+ with torch.no_grad():
+ loss_sum += float(F.cross_entropy(
+ scores.reshape(-1, net.config.vocab_size),
+ target_batch.reshape(-1), reduction="sum"))
+ correct += int(
+ (scores.argmax(dim=-1) == target_batch).sum())
+ tokens_seen += target_batch.numel()
+ batches += 1
+ if was_training:
+ net.train()
+ mean_loss = loss_sum / tokens_seen
+ return {
+ "nll": mean_loss,
+ "perplexity": math.exp(min(mean_loss, 80.0)),
+ "accuracy": correct / tokens_seen,
+ "tokens": tokens_seen,
+ "batches": batches,
+ "candidate_token_presentations": candidate_presentations,
+ "relaxation_token_passes": relaxation_passes,
+ }
+
+
+def hardware_report(device):
+ report = {
+ "platform": platform.platform(),
+ "python": sys.version,
+ "torch": torch.__version__,
+ "numpy": np.__version__,
+ "requested_device": str(device),
+ "cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES"),
+ }
+ if str(device).startswith("cuda"):
+ index = torch.device(device).index
+ index = torch.cuda.current_device() if index is None else index
+ properties = torch.cuda.get_device_properties(index)
+ report.update({
+ "cuda_runtime": torch.version.cuda,
+ "visible_cuda_index": index,
+ "device_name": properties.name,
+ "device_total_memory": properties.total_memory,
+ })
+ return report
+
+
+def work_report(net, args, train_tokens, validation_evaluations):
+ ordinary_validation = sum(
+ row["tokens"] for row in validation_evaluations)
+ presentations = train_tokens
+ if args.method in {"pepita", "ff"}:
+ presentations *= 2
+ local_vjps = 0
+ relaxation = 0
+ if args.method == "dualprop":
+ local_vjps = (
+ train_tokens * (args.depth + 1)
+ * args.dp_inference_passes)
+ relaxation = train_tokens * args.dp_inference_passes
+ elif args.method == "ep":
+ local_vjps = (
+ train_tokens * (args.depth + 1)
+ * (args.ep_free_steps + args.ep_nudge_steps))
+ relaxation = (
+ train_tokens * (args.ep_free_steps + args.ep_nudge_steps)
+ + sum(row["relaxation_token_passes"]
+ for row in validation_evaluations))
+ return {
+ "ordinary_training_tokens": train_tokens,
+ "ordinary_validation_tokens": ordinary_validation,
+ "training_token_presentations": presentations,
+ "candidate_token_presentations": sum(
+ row["candidate_token_presentations"]
+ for row in validation_evaluations),
+ "relaxation_token_passes": relaxation,
+ "local_vjp_token_evaluations": local_vjps,
+ "forward_parameter_count": net.n_forward_parameters,
+ "feedback_parameter_count": net.n_feedback_parameters,
+ "completed_optimizer_steps": len(
+ [row for row in validation_evaluations
+ if row.get("kind") == "train_marker"]),
+ }
+
+
+def run(args):
+ if os.path.exists(args.out):
+ raise FileExistsError(f"refusing to overwrite {args.out}")
+ torch.manual_seed(args.run_seed)
+ if str(args.device).startswith("cuda"):
+ if not torch.cuda.is_available():
+ raise RuntimeError("CUDA requested but unavailable")
+ torch.cuda.manual_seed_all(args.run_seed)
+ torch.cuda.reset_peak_memory_stats(torch.device(args.device))
+ data = TokenData(args.data_dir, args.context_length)
+ net = build(args)
+ net.train()
+ optimizer = torch.optim.AdamW(
+ net.parameters(), lr=args.lr, betas=(0.9, 0.95),
+ weight_decay=args.weight_decay)
+ batch_generator = torch.Generator().manual_seed(args.loader_seed)
+ generator_device = (
+ torch.device(args.device)
+ if str(args.device).startswith("cuda") else torch.device("cpu"))
+ generators = {
+ name: torch.Generator(device=generator_device).manual_seed(seed)
+ for name, seed in (
+ ("negative", args.negative_seed),
+ ("ep", args.ep_sign_seed))}
+ provenance = {
+ "git_commit": git_output("rev-parse", "HEAD"),
+ "git_tracked_dirty": bool(git_output(
+ "status", "--porcelain", "--untracked-files=no")),
+ }
+ started = time.time()
+ train_tokens = 0
+ first_nonfinite_step = None
+ history = []
+ validation_evaluations = []
+
+ for step in range(args.train_steps):
+ tokens, targets = data.random_batch(
+ args.batch_size, batch_generator, args.device)
+ sync(args.device)
+ step_started = time.time()
+ metrics = train_step(
+ net, optimizer, tokens, targets, args, step, generators)
+ sync(args.device)
+ train_tokens += targets.numel()
+ row = {
+ "step": step + 1,
+ "lr": scheduled_rate(args, step),
+ "train_seconds": time.time() - step_started,
+ **metrics,
+ }
+ if (step + 1) % args.log_every == 0 or step == 0:
+ history.append(row)
+ print(
+ f"step={step + 1}/{args.train_steps} "
+ f"loss={metrics['loss']:.6g} "
+ f"seconds={row['train_seconds']:.3f}", flush=True)
+ if not metrics["gradients_finite"] or not math.isfinite(
+ metrics["loss"]):
+ first_nonfinite_step = step + 1
+ break
+ if args.eval_every and (step + 1) % args.eval_every == 0:
+ sync(args.device)
+ eval_started = time.time()
+ evaluation = evaluate(net, data, args)
+ sync(args.device)
+ evaluation.update({
+ "step": step + 1,
+ "evaluation_seconds": time.time() - eval_started,
+ })
+ validation_evaluations.append(evaluation)
+ print(
+ f"validation step={step + 1} "
+ f"nll={evaluation['nll']:.6g} "
+ f"ppl={evaluation['perplexity']:.6g}", flush=True)
+
+ sync(args.device)
+ final_started = time.time()
+ final = evaluate(net, data, args)
+ sync(args.device)
+ final["step"] = min(args.train_steps, (
+ first_nonfinite_step or args.train_steps))
+ final["evaluation_seconds"] = time.time() - final_started
+ validation_evaluations.append(final)
+ peak_allocated = None
+ peak_reserved = None
+ if str(args.device).startswith("cuda"):
+ peak_allocated = torch.cuda.max_memory_allocated(
+ torch.device(args.device))
+ peak_reserved = torch.cuda.max_memory_reserved(
+ torch.device(args.device))
+ record = {
+ "schema_version": 1,
+ "protocol_family": "transformer_local_learning_crossover",
+ "args": vars(args),
+ "provenance": provenance,
+ "dataset": {
+ "name": "tiny Shakespeare character language modeling",
+ "train_path": data.paths["train"],
+ "validation_path": data.paths["val"],
+ "train_sha256": data.hashes["train"],
+ "validation_sha256": data.hashes["val"],
+ "train_file_tokens": data.train_tokens,
+ "validation_file_tokens": data.validation_tokens,
+ "vocab_size": 65,
+ },
+ "architecture": {
+ "family": "pre-LN causal decoder Transformer",
+ "depth": args.depth,
+ "width": args.width,
+ "heads": args.heads,
+ "context_length": args.context_length,
+ "mlp_ratio": args.mlp_ratio,
+ "dropout": 0.0,
+ "forward_parameter_count": net.n_forward_parameters,
+ "feedback_parameter_count": net.n_feedback_parameters,
+ },
+ "history": history,
+ "validation": validation_evaluations,
+ "first_nonfinite_step": first_nonfinite_step,
+ "evaluation_protocol": {
+ "split": "validation",
+ "test_evaluations": 0,
+ "test_used_for_selection": False,
+ },
+ "final": final,
+ "work": work_report(
+ net, args, train_tokens, validation_evaluations),
+ "hardware": {
+ **hardware_report(args.device),
+ "peak_memory_allocated_bytes": peak_allocated,
+ "peak_memory_reserved_bytes": peak_reserved,
+ },
+ "total_wall_seconds": time.time() - started,
+ }
+ # completed_optimizer_steps is easier and less error-prone to state here.
+ record["work"]["completed_optimizer_steps"] = final["step"]
+ os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
+ with open(args.out, "w", encoding="utf-8") as handle:
+ json.dump(record, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print(json.dumps({
+ "out": args.out,
+ "final": final,
+ "work": record["work"],
+ }, indent=2, sort_keys=True))
+
+
+def parse_args():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--method", choices=METHODS, required=True)
+ parser.add_argument("--out", required=True)
+ parser.add_argument("--device", default="cpu")
+ parser.add_argument("--data_dir", default=DEFAULT_DATA_DIR)
+ parser.add_argument("--depth", type=int, choices=(4, 8, 12), default=4)
+ parser.add_argument("--width", type=int, default=128)
+ parser.add_argument("--heads", type=int, default=4)
+ parser.add_argument("--mlp_ratio", type=int, default=4)
+ parser.add_argument("--context_length", type=int, default=64)
+ parser.add_argument("--batch_size", type=int, default=32)
+ parser.add_argument("--eval_batch_size", type=int, default=32)
+ parser.add_argument("--train_steps", type=int, default=1000)
+ parser.add_argument("--lr", type=float, required=True)
+ parser.add_argument(
+ "--schedule", choices=("constant", "cosine"), default="constant")
+ parser.add_argument("--min_lr", type=float, default=1e-4)
+ parser.add_argument("--warmup_steps", type=int, default=100)
+ parser.add_argument("--weight_decay", type=float, default=0.1)
+ parser.add_argument("--run_seed", type=int, default=0)
+ parser.add_argument("--model_seed", type=int, default=2027)
+ parser.add_argument("--loader_seed", type=int, default=0)
+ parser.add_argument("--negative_seed", type=int, default=5001)
+ parser.add_argument("--ep_sign_seed", type=int, default=5002)
+ parser.add_argument("--traffic_ratio", type=float, default=4.0)
+ parser.add_argument("--ff_threshold", type=float, default=2.0)
+ parser.add_argument("--ff_score_from_layer", type=int, default=1)
+ parser.add_argument("--ep_beta", type=float, default=0.5)
+ parser.add_argument("--ep_dt", type=float, default=0.5)
+ parser.add_argument("--ep_free_steps", type=int, default=20)
+ parser.add_argument("--ep_nudge_steps", type=int, default=4)
+ parser.add_argument("--dp_alpha", type=float, default=0.0)
+ parser.add_argument("--dp_beta", type=float, default=0.1)
+ parser.add_argument("--dp_inference_passes", type=int, default=16)
+ parser.add_argument("--eval_every", type=int, default=0)
+ parser.add_argument("--log_every", type=int, default=50)
+ parser.add_argument(
+ "--max_val_batches", type=int, default=0,
+ help="smoke-only cap; formal runs use zero for the full validation set")
+ return parser.parse_args()
+
+
+if __name__ == "__main__":
+ run(parse_args())