iskhare/iclr-debug
0
1diff --git a/evaluation/eval_random_gen_ppl_checkpoint.py b/evaluation/eval_random_gen_ppl_checkpoint.py2index 2d5c741..eb4eb89 1006443--- a/evaluation/eval_random_gen_ppl_checkpoint.py4+++ b/evaluation/eval_random_gen_ppl_checkpoint.py5@@ -5,6 +5,7 @@ import json6 import math7 import sys8 from collections import Counter9+from contextlib import contextmanager10 from pathlib import Path11 12 REPO_ROOT = Path(__file__).resolve().parents[1]13@@ -22,14 +23,35 @@ from utils.misc import set_manual_seed14 from utils.registry import trainers15 16 17+def seed_evaluation(seed, reference_dtype):18+ # The shared seed helper enables TF32, so precision must be set afterward.19+ set_manual_seed(seed)20+ if reference_dtype in ("float32", "float64"):21+ torch.backends.cuda.matmul.allow_tf32 = False22+ torch.backends.cudnn.allow_tf32 = False23+24+25+@contextmanager26+def reference_precision_context(reference_dtype):27+ previous_dtype = torch.get_default_dtype()28+ try:29+ # Transformers' causal-mask helper creates a tensor without a dtype30+ # before converting it. Its FP64 minimum overflows a default FP32 tensor.31+ if reference_dtype == "float64":32+ torch.set_default_dtype(torch.float64)33+ yield34+ finally:35+ torch.set_default_dtype(previous_dtype)36+37+38 def main():39 parser = argparse.ArgumentParser()40 parser.add_argument("--ckpt", required=True)41 parser.add_argument("--codec")42 parser.add_argument("--validation-source")43 parser.add_argument("--ref-model", default="gpt2-large")44- parser.add_argument("--ref-dtype", choices=("bfloat16", "float32"), default="bfloat16",45- help="Legacy default retained; new ICLR runs explicitly use float32")46+ parser.add_argument("--ref-dtype", choices=("bfloat16", "float32", "float64"), default="bfloat16",47+ help="Reference model and scoring dtype; float32/float64 disable TF32")48 parser.add_argument("--seed", type=int, default=0)49 parser.add_argument("--n-samples", type=int, default=128)50 parser.add_argument("--batch-size", type=int, default=8)51@@ -40,10 +62,6 @@ def main():52 parser.add_argument("--save-all-generations", action="store_true")53 args = parser.parse_args()54 55- if args.ref_dtype == "float32":56- torch.backends.cuda.matmul.allow_tf32 = False57- torch.backends.cudnn.allow_tf32 = False58-59 device = "cuda"60 checkpoint = torch.load(args.ckpt, map_location="cpu", weights_only=False)61 cfg = checkpoint["args"]62@@ -58,7 +76,7 @@ def main():63 else:64 cfg.optimization.stochastic_encode = {"enabled": False}65 66- set_manual_seed(args.seed)67+ seed_evaluation(args.seed, args.ref_dtype)68 Trainer = hydra.utils.get_class(trainers[cfg.task])69 trainer = Trainer(cfg, device)70 trainer.model.load_state_dict(checkpoint["model"])71@@ -129,18 +147,19 @@ def main():72 torch.cuda.empty_cache()73 74 ref_tokenizer = AutoTokenizer.from_pretrained(args.ref_model)75- ref_model = AutoModelForCausalLM.from_pretrained(76- args.ref_model, torch_dtype=getattr(torch, args.ref_dtype)77- ).to(device).eval()78- result = generation_perplexity(79- ref_model,80- ref_tokenizer,81- prompts,82- generations,83- score_only_generated=True,84- batch_size=args.batch_size,85- device=device,86- )87+ with reference_precision_context(args.ref_dtype):88+ ref_model = AutoModelForCausalLM.from_pretrained(89+ args.ref_model, torch_dtype=getattr(torch, args.ref_dtype)90+ ).to(device).eval()91+ result = generation_perplexity(92+ ref_model,93+ ref_tokenizer,94+ prompts,95+ generations,96+ score_only_generated=True,97+ batch_size=args.batch_size,98+ device=device,99+ )100 101 payload = {102 "checkpoint": str(Path(args.ckpt).resolve()),103diff --git a/scripts/run_iclr_debug.py b/scripts/run_iclr_debug.py104index da25b0c..ce5ae4e 100644105--- a/scripts/run_iclr_debug.py106+++ b/scripts/run_iclr_debug.py107@@ -23,7 +23,15 @@ from utils.iclr_training import stage_config, export_training_checkpoint108 109 def make_plan(args):110 root = args.output.resolve()111- stages = ["tokenizer", "generator"] if args.variant == "two-stage" else ["decoder-mstok"]112+ stages = {"two-stage": ["tokenizer", "generator"], "generator": ["generator"],113+ "decoder-mstok": ["decoder-mstok"]}[args.variant]114+ external_tokenizer = getattr(args, "tokenizer_checkpoint", None)115+ external_steps = getattr(args, "tokenizer_steps", None)116+ if args.variant == "generator":117+ if external_tokenizer is None or external_steps is None or external_steps <= 0:118+ raise ValueError("Generator-only training requires --tokenizer-checkpoint and positive --tokenizer-steps")119+ elif external_tokenizer is not None or external_steps is not None:120+ raise ValueError("External tokenizer options require --variant generator")121 if args.tokenizer_only and args.variant != "two-stage":122 raise ValueError("--tokenizer-only requires --variant two-stage")123 if args.pilot_steps is not None and not 32 <= args.pilot_steps <= 500:124@@ -32,13 +40,18 @@ def make_plan(args):125 for stage in stages:126 checkpoint, tokenizer_steps = None, None127 if stage == "generator":128- tokenizer_steps = configs["tokenizer"].iclr_debug.stop_step129- checkpoint = root / "tokenizer" / f"milestone-iter-{tokenizer_steps}.pt"130+ if args.variant == "generator":131+ tokenizer_steps = external_steps132+ checkpoint = Path(external_tokenizer).resolve()133+ else:134+ tokenizer_steps = configs["tokenizer"].iclr_debug.stop_step135+ checkpoint = root / "tokenizer" / f"milestone-iter-{tokenizer_steps}.pt"136 budget = args.generator_positions if stage == "generator" else args.tokenizer_positions137 if stage == "decoder-mstok":138 budget = args.joint_positions139 configs[stage] = stage_config(stage, args.data_dir, root, target_positions=budget,140- horizon_positions=args.joint_lr_horizon_positions if stage == "decoder-mstok" else None,141+ horizon_positions=(args.joint_lr_horizon_positions if stage == "decoder-mstok" else142+ getattr(args, "generator_lr_horizon_positions", None) if stage == "generator" else None),143 pilot_steps=args.pilot_steps, wandb_enabled=not args.no_wandb,144 tokenizer_checkpoint=checkpoint, tokenizer_steps=tokenizer_steps)145 if stage == "generator":146@@ -65,7 +78,7 @@ def latest(root):147 return max(candidates, key=lambda x: x[0], default=(0, None))148 149 150-def generation_evaluation(controller, config, step, directory):151+def generation_evaluation(controller, config, step, directory, reference_dtype="float32"):152 stage_root = Path(config.experiment_dir)153 output = stage_root / "evaluation" / f"step-{step}"154 output.mkdir(parents=True, exist_ok=True)155@@ -78,19 +91,25 @@ def generation_evaluation(controller, config, step, directory):156 controller.run(f"{config.iclr_debug.stage}-eval-{step}-seed{seed}", [sys.executable,157 "evaluation/eval_random_gen_ppl_checkpoint.py", "--ckpt", directory / "ncp.pt",158 "--codec", directory / "vqvae.pt", "--validation-source", config.dataset.validate_source,159- "--ref-model", "gpt2-large", "--ref-dtype", "float32", "--seed", seed,160+ "--ref-model", "gpt2-large", "--ref-dtype", reference_dtype, "--seed", seed,161 "--n-samples", 128, "--batch-size", 4, "--temperature", 1,162 "--top-k", 0, "--top-p", 1, "--save-all-generations", "--out", path],163 path.with_suffix(".log"), env=dict(controller.env, CUDA_VISIBLE_DEVICES="0"), timeout=3600)164 row = json.loads(path.read_text())165- if (row["reference_dtype"] != "float32" or row["protocol"]["n_levels"] != 16166+ if (row["reference_dtype"] != reference_dtype or row.get("tf32_allowed") is not False167+ or row["protocol"]["n_levels"] != 16168 or row["checkpoint_step"] != step or row["seed"] != seed169 or row["requested_samples"] != 128 or row["scored_samples"] != 128):170 raise ValueError(f"Incomplete/mismatched evaluation: {path}")171 controller.run(f"summarize-{step}", [sys.executable, "evaluation/summarize_gen_ppl.py",172 "--input-glob", output / "random-seed-[0-4].json", "--out", output / "summary.json"],173 output / "summarize.log", timeout=120)174- write_json(output / "EVAL_DONE.json", dict(step=step, reference_dtype="float32", sampling="supplied-level0-untruncated"))175+ summary_path = output / "summary.json"176+ summary = json.loads(summary_path.read_text())177+ summary.update(reference_dtype=reference_dtype, tf32_allowed=False)178+ write_json(summary_path, summary)179+ write_json(output / "EVAL_DONE.json", dict(step=step, reference_dtype=reference_dtype,180+ tf32_allowed=False, sampling="supplied-level0-untruncated"))181 (output / "EVAL_FAILED.json").unlink(missing_ok=True)182 183 184@@ -130,6 +149,19 @@ def run(args):185 raise FileNotFoundError("Resume requires this launcher's manifest")186 expected = json.loads((ROOT / "config/iclr-debug/full-owt-audit.json").read_text())187 data_preflight(args.data_dir, expected)188+ tokenizer_provenance = None189+ if args.variant == "generator":190+ import torch191+ checkpoint = Path(configs["generator"].iclr_debug.tokenizer_checkpoint)192+ payload = torch.load(checkpoint, map_location="cpu", mmap=True, weights_only=False)193+ if (OmegaConf.select(payload["args"], "iclr_debug.stage") != "tokenizer"194+ or int(payload["step"]) != configs["generator"].iclr_debug.tokenizer_steps):195+ raise ValueError("External tokenizer checkpoint stage/step mismatch")196+ for key in ("codec", "dataset", "pre_tokenizer"):197+ if OmegaConf.to_container(payload["args"][key], resolve=True) != OmegaConf.to_container(configs["generator"][key], resolve=True):198+ raise ValueError(f"External tokenizer changes {key}")199+ tokenizer_provenance = dict(path=str(checkpoint), step=int(payload["step"]), sha256=sha256(checkpoint))200+ del payload201 root.mkdir(parents=True, exist_ok=True)202 lock = (root / "controller.lock").open("a")203 fcntl.flock(lock, fcntl.LOCK_EX | fcntl.LOCK_NB)204@@ -141,6 +173,8 @@ def run(args):205 configs={s: OmegaConf.to_container(c, resolve=True) for s, c in configs.items()},206 sources=sources, dataset=expected, hf_repo=args.hf_repo,207 generation_evaluation=not args.no_generation_eval,208+ evaluation_reference_dtype=args.eval_ref_dtype,209+ external_tokenizer=tokenizer_provenance,210 versions={name: version(name) for name in ("torch", "transformers", "numpy", "hydra-core")})211 if args.resume:212 if json.loads((root / "manifest.json").read_text()) != manifest:213@@ -170,7 +204,7 @@ def run(args):214 if args.pilot_steps:215 # Force a real save/reload/continuation in every stage.216 targets = [args.pilot_steps // 2, args.pilot_steps]217- if stage == "generator" and not (root / "tokenizer" / "STAGE_DONE.json").is_file():218+ if stage == "generator" and args.variant == "two-stage" and not (root / "tokenizer" / "STAGE_DONE.json").is_file():219 raise RuntimeError("Tokenizer stage must finish before generator handoff")220 for target in targets:221 done, checkpoint = latest(stage_root)222@@ -195,7 +229,7 @@ def run(args):223 export_training_checkpoint(checkpoint, exports)224 if not args.no_generation_eval and target in config.training.generation_eval_steps:225 try:226- generation_evaluation(controller, config, target, exports)227+ generation_evaluation(controller, config, target, exports, args.eval_ref_dtype)228 except InterruptedError:229 raise230 except Exception as exc:231@@ -223,11 +257,16 @@ def run(args):232 233 def main():234 p = argparse.ArgumentParser(description=__doc__)235- p.add_argument("--variant", choices=("two-stage", "decoder-mstok"), required=True)236+ p.add_argument("--variant", choices=("two-stage", "generator", "decoder-mstok"), required=True)237 p.add_argument("--data-dir", type=Path, required=True)238 p.add_argument("--output", type=Path, required=True)239 p.add_argument("--tokenizer-positions", type=int, default=135_000_000_000)240 p.add_argument("--generator-positions", type=int, default=135_000_000_000)241+ p.add_argument("--generator-lr-horizon-positions", type=int,242+ help="Generator cosine-decay horizon; hold minimum LR after this many input positions")243+ p.add_argument("--tokenizer-checkpoint", type=Path, help="Frozen tokenizer training checkpoint for --variant generator")244+ p.add_argument("--tokenizer-steps", type=int, help="Expected step of the external tokenizer checkpoint")245+ p.add_argument("--eval-ref-dtype", choices=("float32", "float64"), default="float32")246 p.add_argument("--generator-microbatch", type=int, choices=(32, 64), default=32,247 help="32x16 (measured fastest) or 64x8; both keep global batch 4096")248 p.add_argument("--joint-positions", type=int, default=135_000_000_000)249diff --git a/tests/test_iclr_training.py b/tests/test_iclr_training.py250index 2bc7aee..d19b507 100644251--- a/tests/test_iclr_training.py252+++ b/tests/test_iclr_training.py253@@ -90,6 +90,51 @@ def test_plan_handoff_and_pilot(tmp_path):254 assert pilot["generator"].training.generation_eval_steps == []255 256 257+def test_generator_only_short_decay_plan(tmp_path):258+ args = SimpleNamespace(output=tmp_path / "fresh-generator", data_dir=tmp_path,259+ variant="generator", tokenizer_only=False, pilot_steps=None, no_wandb=False,260+ tokenizer_checkpoint=tmp_path / "tokenizer-128747.pt", tokenizer_steps=128747,261+ generator_positions=135_000_000_000, generator_lr_horizon_positions=27_000_000_000,262+ joint_lr_horizon_positions=None)263+ configs = make_plan(args)264+ assert set(configs) == {"generator"}265+ cfg = configs["generator"]266+ assert cfg.iclr_debug.tokenizer_checkpoint == str(args.tokenizer_checkpoint.resolve())267+ assert cfg.iclr_debug.tokenizer_steps == 128747268+ assert cfg.training.total_iters == 128747269+ assert cfg.optimization.generator_lr_decay_iters == 25750270+ assert cfg.optimization.generator_warmup_iters == 225271+ assert cfg.optimization.generator_min_lr == 1e-5272+ assert cfg.training.resume_checkpoint is None273+ args.tokenizer_steps = None274+ with pytest.raises(ValueError, match="requires"):275+ make_plan(args)276+277+278+def test_generator_lr_stays_at_minimum_after_short_horizon_and_resume(tmp_path):279+ cfg = stage_config("generator", tmp_path, tmp_path, horizon_positions=27_000_000_000)280+ model = SimpleNamespace(ncp=torch.nn.Linear(1, 1), vqvae=torch.nn.Linear(1, 1).requires_grad_(False))281+ optimizer, scheduler = make_optimizer_scheduler(model, cfg)282+ horizon = cfg.optimization.generator_lr_decay_iters283+ # Exercise the actual scheduler on both sides of the boundary and after reload.284+ scheduler.last_epoch = horizon - 2285+ rates = []286+ for _ in range(4):287+ optimizer.step()288+ scheduler.step()289+ rates.append(optimizer.param_groups[0]["lr"])290+ assert rates[0] > cfg.optimization.generator_min_lr291+ assert rates[1:] == [cfg.optimization.generator_min_lr] * 3292+ opt2, sched2 = make_optimizer_scheduler(model, cfg)293+ opt2.load_state_dict(optimizer.state_dict())294+ sched2.load_state_dict(scheduler.state_dict())295+ for step in (horizon + 100, 51499, 128747):296+ sched2.last_epoch = step - 1297+ opt2.step()298+ sched2.step()299+ assert opt2.param_groups[0]["lr"] == cfg.optimization.generator_min_lr300+301+302 @pytest.mark.parametrize("stage", ["tokenizer", "decoder-mstok"])303 def test_strict_checkpoint_exact_continuation(stage, tmp_path):304 torch.set_num_threads(1)305 