#!/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 = os.environ.get( "SDIL_TRANSFORMER_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(): cache = net(tokens, targets, return_cache=True) loss = cache["loss"] metrics = { "loss": float(loss), **net.dfa_gradients(tokens, targets, cache=cache), } 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, completed_steps, 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)) width = args.width block_affine_macs_per_token = ( (4 + 2 * args.mlp_ratio) * width * width) block_attention_macs_per_token = ( 2 * args.context_length * width) forward_affine_macs_per_token = ( args.depth * block_affine_macs_per_token + width * net.config.vocab_size) forward_attention_macs_per_token = ( args.depth * block_attention_macs_per_token) training_forward_multiplier = { "bp": 1, "fa": 1, "dfa": 1, "pepita": 2, "ff": 2, "ep": 2, "dualprop": 1, "clean_kp": 1, "sdil": 1, }[args.method] if args.method == "ff": validation_forward_tokens = sum( row["candidate_token_presentations"] for row in validation_evaluations) else: validation_forward_tokens = ordinary_validation full_forward_token_passes = ( training_forward_multiplier * train_tokens + validation_forward_tokens) local_block_token_evaluations = 0 local_head_token_evaluations = 0 if args.method in {"dfa", "pepita"}: local_block_token_evaluations = args.depth * train_tokens local_head_token_evaluations = train_tokens elif args.method == "ff": steps_per_layer = math.ceil( args.train_steps / net.ff_num_layers) block_steps = 0 head_steps = 0 for step in range(completed_steps): layer = min( step // steps_per_layer, net.ff_num_layers - 1) block_steps += int(1 <= layer <= args.depth) head_steps += int(layer == args.depth + 1) tokens_per_step = args.batch_size * args.context_length local_block_token_evaluations = block_steps * tokens_per_step local_head_token_evaluations = head_steps * tokens_per_step elif args.method == "dualprop": multiplier = 2 * args.dp_inference_passes + 1 local_block_token_evaluations = ( args.depth * train_tokens * multiplier) local_head_token_evaluations = train_tokens * multiplier elif args.method == "ep": train_multiplier = ( 2 * (args.ep_free_steps + args.ep_nudge_steps) + 2) validation_multiplier = 2 * args.ep_free_steps local_block_token_evaluations = args.depth * ( train_tokens * train_multiplier + ordinary_validation * validation_multiplier) local_head_token_evaluations = ( train_tokens * train_multiplier + ordinary_validation * validation_multiplier) 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), "logical_task_loss_queries": 0, "relaxation_token_passes": relaxation, "local_vjp_token_evaluations": local_vjps, "full_forward_token_passes": full_forward_token_passes, "local_block_token_evaluations": local_block_token_evaluations, "local_head_token_evaluations": local_head_token_evaluations, "block_affine_macs_per_token": block_affine_macs_per_token, "block_attention_macs_per_token": block_attention_macs_per_token, "forward_affine_macs_per_token": forward_affine_macs_per_token, "forward_attention_macs_per_token": forward_attention_macs_per_token, "enumerated_full_forward_macs": ( full_forward_token_passes * (forward_affine_macs_per_token + forward_attention_macs_per_token)), "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, final["step"], 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())