diff options
Diffstat (limited to 'experiments/transformer_crossover_native.py')
| -rw-r--r-- | experiments/transformer_crossover_native.py | 502 |
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()) |
