CoolFace
Modelpublic

iskhare/iclr-debug

sourceHugging Faceupdated 1d agoView on Hugging Face
0likes
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