diff --git a/docker/patch/latest/megatron.patch b/docker/patch/latest/megatron.patch index 3be8152b..6a4b200c 100644 --- a/docker/patch/latest/megatron.patch +++ b/docker/patch/latest/megatron.patch @@ -639,10 +639,22 @@ index 5b31ddedf..ead60f2dd 100644 ) if self.fuse_linear_cross_entropy: diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py -index a4364f5e9..c76f6daac 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py -@@ -686,6 +686,8 @@ class DistributedOptimizer(MixedPrecisionOptimizer): +@@ -645,8 +645,10 @@ + # Extract 'step', for non-Apex/TE support. + if not HAVE_APEX_OR_TE: + steps = list(set([s["step"].item() for s in inner_state_dict["state"].values()])) +- assert len(steps) == 1 +- step = steps[0] ++ # A fresh Torch optimizer has no per-parameter state yet. The ++ # checkpoint loader queries this template before allocating state. ++ assert len(steps) <= 1, f"steps: {steps}" ++ step = steps[0] if steps else 0 + elif isinstance(self.optimizer, HybridDeviceOptimizer): + step = None + for optimizer in self.optimizer.sub_optimizers: +@@ -686,6 +688,8 @@ # TE FusedAdam will not accumulate step for empty param groups, so we need to # align the step across param groups. param_group["step"] = int(step) @@ -651,8 +663,31 @@ index a4364f5e9..c76f6daac 100644 # Grad scaler state. if self.grad_scaler: -@@ -969,7 +971,12 @@ class DistributedOptimizer(MixedPrecisionOptimizer): - for bucket_idx, gbuf_range_map in enumerate(gbuf_range_map_for_all_buckets): +@@ -828,7 +832,12 @@ + + for s in state_dict_state.values(): + # Native PyTorch state dict requires step (i.e., iteration). +- s["step"] = step ++ # Torch Adam increments each parameter's step in-place. A ++ # shared scalar would advance once per parameter, not per update. ++ s["step"] = step.detach().clone() ++ for group in state_dict_param_groups: ++ # This group field is checkpoint metadata, not Torch Adam state. ++ group.pop("step", None) + elif isinstance(self.optimizer, HybridDeviceOptimizer): + # Handle Torch AdamW special case, which, unlike FusedAdam, Torch AdamW + # has an extra optimizer state "step". +@@ -943,6 +952,9 @@ + optim_state = self.optimizer.state[main_param] + dst_tensors = {"param": main_param, **optim_state} + for key in dst_tensors: ++ if key == "step": ++ # Restored from common optimizer state, not parameter shards. ++ continue + dst_tensors[key].copy_(tensors[key]) + + def get_parameter_state_dp_reshardable(self): +@@ -968,6 +980,11 @@ bucket_state = [] for model_param, param_range_map in gbuf_range_map["param_map"].items(): tensors = self._get_main_param_and_optimizer_states(model_param) @@ -664,7 +699,7 @@ index a4364f5e9..c76f6daac 100644 tensors.update( { "gbuf_local_start": param_range_map["gbuf_local"].start, -@@ -1667,6 +1669,11 @@ class DistributedOptimizer(MixedPrecisionOptimizer): +@@ -1667,6 +1684,11 @@ if key == 'padding': tensors[key] = LocalNonpersistentObject(tensors[key]) continue @@ -676,7 +711,7 @@ index a4364f5e9..c76f6daac 100644 assert tensors[key].shape == (gbuf_local_end - gbuf_local_start,), ( tensors[key].shape, gbuf_local_start, -@@ -1808,6 +1815,11 @@ class DistributedOptimizer(MixedPrecisionOptimizer): +@@ -1801,6 +1823,11 @@ for src_tensors, (model_param, param_range_map) in zip( bucket_state, gbuf_range_map["param_map"].items() ): diff --git a/tests/checkpoint_resume_audit.py b/tests/checkpoint_resume_audit.py new file mode 100644 index 00000000..1d974f38 --- /dev/null +++ b/tests/checkpoint_resume_audit.py @@ -0,0 +1,226 @@ +"""Observation-only fixed-batch resume audit, loaded through Vime's public hook.""" + +import ast +import hashlib +import json +import math +import os +import pickle +import random +import re +import time +from pathlib import Path + + +def tensor_record(tensor): + import torch + + value = tensor.detach().cpu().contiguous() + return { + "shape": list(value.shape), + "dtype": str(value.dtype), + "sha256": hashlib.sha256(value.reshape(-1).view(torch.uint8).numpy().tobytes()).hexdigest(), + } + + +def tree_record(value): + import torch + + if torch.is_tensor(value): + return tensor_record(value) + if isinstance(value, dict): + return {str(k): tree_record(v) for k, v in value.items()} + if isinstance(value, (tuple, list)): + return [tree_record(v) for v in value] + if value is None or isinstance(value, (str, bool, int, float)): + return value + raise TypeError(f"Unsupported state: {type(value)}") + + +def emit(kind, **data): + path = Path(os.environ["VIME_RESUME_AUDIT_DIR"]) / f"resume-{os.getpid()}.jsonl" + with path.open("a") as stream: + stream.write(json.dumps({"kind": kind, "pid": os.getpid(), "time": time.time(), **data}) + "\n") + + +def rng_record(): + import numpy as np + import torch + from megatron.core.tensor_parallel.random import get_cuda_rng_tracker + + return { + "python": hashlib.sha256(pickle.dumps(random.getstate())).hexdigest(), + "numpy": hashlib.sha256(pickle.dumps(np.random.get_state())).hexdigest(), + "torch_cpu": tensor_record(torch.get_rng_state()), + "torch_cuda": tensor_record(torch.cuda.get_rng_state()), + "megatron": tree_record(get_cuda_rng_tracker().get_states()), + } + + +def state_record(model, optimizer, scheduler): + optimizers = getattr(optimizer, "chained_optimizers", [optimizer]) + states = [] + for opt in optimizers: + inner = opt.optimizer + states.append( + { + "class": type(opt).__name__, + "inner_class": type(inner).__name__, + "state": tree_record(inner.state_dict()), + "master_parameters": tree_record([p for group in inner.param_groups for p in group["params"]]), + } + ) + return { + "model": {f"{i}/{name}": tensor_record(p) for i, m in enumerate(model) for name, p in m.named_parameters()}, + "optimizer": states, + "scheduler": tree_record(scheduler.state_dict()), + "rng": rng_record(), + } + + +def before_train(args, rollout_id, step_id, model, optimizer, scheduler): + import torch + + emit( + "before_update", + rollout_id=rollout_id, + step_id=step_id, + state=state_record(model, optimizer, scheduler), + load_flags={ + k: getattr(args, k, None) for k in ["load", "no_load_optim", "no_load_rng", "finetune", "offload_train"] + }, + ) + optimizer._resume_audit_id = (rollout_id, step_id) + if getattr(optimizer, "_resume_audit_installed", False): + return + optimizer._resume_audit_installed = True + original_step = optimizer.step + original_scheduler_step = scheduler.step + + def audited_step(*a, **kw): + rid, sid = optimizer._resume_audit_id + gradients = {} + for i, m in enumerate(model): + for name, p in m.named_parameters(): + grad = getattr(p, "main_grad", None) + if grad is None: + grad = p.grad + if grad is not None: + gradients[f"{i}/{name}"] = tensor_record(grad) + emit("gradients", rollout_id=rid, step_id=sid, tensors=gradients) + result = original_step(*a, **kw) + emit("optimizer_result", rollout_id=rid, step_id=sid, result=tree_record(result)) + return result + + def audited_scheduler_step(*a, **kw): + result = original_scheduler_step(*a, **kw) + rid, sid = optimizer._resume_audit_id + torch.cuda.synchronize() + emit("after_update", rollout_id=rid, step_id=sid, state=state_record(model, optimizer, scheduler)) + return result + + optimizer.step = audited_step + scheduler.step = audited_scheduler_step + + +def require_equal(left, right, path="root"): + if type(left) is not type(right): + raise AssertionError(f"{path}: types differ: {type(left).__name__} / {type(right).__name__}") + if isinstance(left, dict): + if left.keys() != right.keys(): + raise AssertionError(f"{path}: keys differ: {left.keys() ^ right.keys()}") + for key in left: + require_equal(left[key], right[key], f"{path}.{key}") + elif isinstance(left, list): + if len(left) != len(right): + raise AssertionError(f"{path}: lengths differ") + for index, (a, b) in enumerate(zip(left, right, strict=True)): + require_equal(a, b, f"{path}[{index}]") + elif left != right: + raise AssertionError(f"{path}: values differ: {str(left)[:90]} / {str(right)[:90]}") + + +def read_run(root, expected_ids): + records = [json.loads(line) for p in root.glob("resume-*.jsonl") for line in p.read_text().splitlines()] + by_kind = {} + pids = set() + for kind in ["before_update", "gradients", "optimizer_result", "after_update"]: + rows = [row for row in records if row["kind"] == kind] + ids = [row["rollout_id"] for row in rows] + require_equal(sorted(ids), list(expected_ids), f"{root.name}.{kind}.rollout_ids") + for row in rows: + assert row["step_id"] == 0, "This fixture expects one optimizer update per rollout" + pids.add(row["pid"]) + by_kind[kind] = {row["rollout_id"]: row for row in rows} + assert len(pids) == 1, f"Expected a single trainer process, got {pids}" + assert json.loads((root / "train-returned.json").read_text())["completed"] + metrics = {} + for line in (root / "train.log").read_text().splitlines(): + match = re.search(r"step (\d+): (\{'train/loss'.*)", line) + if match: + row = ast.literal_eval(match.group(2)) + metrics[int(match.group(1))] = row + require_equal(sorted(metrics), list(expected_ids), f"{root.name}.metric_steps") + for rid in expected_ids: + row = metrics[rid] + assert all(math.isfinite(v) for v in row.values() if isinstance(v, (float, int))) + assert row["train/grad_norm"] > 0 + assert by_kind["gradients"][rid]["tensors"], "Missing gradient evidence" + assert by_kind["before_update"][rid]["state"]["model"], "Missing model evidence" + assert by_kind["after_update"][rid]["state"]["optimizer"], "Missing optimizer evidence" + assert by_kind["optimizer_result"][rid]["result"][0] is True + assert (root / f"dumps/train_data/{rid}.pt").is_file(), "Missing fixed-batch evidence" + assert ( + by_kind["after_update"][rid]["state"]["model"] != by_kind["before_update"][rid]["state"]["model"] + ), "No parameter update" + return by_kind, metrics, pids.pop() + + +def verify(continuous, first, resumed, split=4, steps=8): + import torch + + assert 0 < split < steps + baseline, base_metrics, a_pid = read_run(continuous, range(steps)) + before, before_metrics, b_pid = read_run(first, range(split)) + after, after_metrics, c_pid = read_run(resumed, range(split, steps)) + assert len({a_pid, b_pid, c_pid}) == 3, "Runs must use distinct trainer processes" + for rid in range(steps): + branch, metrics, root = (before, before_metrics, first) if rid < split else (after, after_metrics, resumed) + for kind, value_key in [ + ("before_update", "state"), + ("gradients", "tensors"), + ("optimizer_result", "result"), + ("after_update", "state"), + ]: + require_equal(baseline[kind][rid][value_key], branch[kind][rid][value_key], f"rollout{rid}.{kind}") + require_equal(base_metrics[rid], metrics[rid], f"rollout{rid}.metrics") + a = torch.load(continuous / f"dumps/train_data/{rid}.pt", map_location="cpu", weights_only=False) + b = torch.load(root / f"dumps/train_data/{rid}.pt", map_location="cpu", weights_only=False) + require_equal(tree_record(a), tree_record(b), f"rollout{rid}.training_data") + flags = after["before_update"][split]["load_flags"] + for name in ["no_load_optim", "no_load_rng", "finetune"]: + assert flags[name] is False, f"Restore bypassed {name}: {flags}" + require_equal( + before["after_update"][split - 1]["state"], after["before_update"][split]["state"], "checkpoint_boundary" + ) + return { + "status": "FIXED_BATCH_FRESH_PROCESS_RESUME_EXACT", + "steps": steps, + "split": split, + "trainer_pids": [a_pid, b_pid, c_pid], + "model_parameters": len(baseline["before_update"][0]["state"]["model"]), + "compared": [ + "model", + "FP32 master parameters", + "Adam moments and step", + "scheduler", + "RNG", + "gradients", + "optimizer results", + "loss metrics", + "training tokens masks logprobs advantages and schedule", + ], + "offload": flags["offload_train"], + "tolerance": "bitwise state and tensor hashes; exact scalar equality", + "scope": "TP=PP=CP=1, dense fixed-batch training. Live serving synchronization is a separate integration check.", + } diff --git a/tests/checkpoint_resume_parity.md b/tests/checkpoint_resume_parity.md new file mode 100644 index 00000000..2476054e --- /dev/null +++ b/tests/checkpoint_resume_parity.md @@ -0,0 +1,48 @@ +# Native Adam checkpoint resume and dense parity + +The Megatron patch handles a fresh native Torch Adam/AdamW optimizer when the checkpoint loader requests a state template, and restores an independent step scalar for each parameter. A shared scalar is incremented once per parameter by Torch Adam, corrupting the next update. The patch also removes the checkpoint-only group step field after rebuilding per-parameter steps and keeps the restored step when loading parameter shards, which do not contain it. + +`tests/utils/test_native_adam_checkpoint.py` exercises the real MCore methods using two CPU AdamW parameters: the original code fails the empty-state and step-independence checks; the patched code preserves the fifth update exactly and still rejects inconsistent saved steps or missing Adam moments. This test requires the Megatron version and patch installed by the container. + +Existing checkpoint smoke tests demonstrate that loading runs. The integration test compares +the state and the next updates against uninterrupted training with identical saved +batches. It uses the real training loop, Megatron checkpoint loader, optimizer, +and optional `--offload-train` cycles. The audit hook only hashes detached tensors; +it does not replace training or restore any state itself. + +Prepare a JSON array containing the arguments for a working single-GPU dense +Megatron training run, including `--num-rollout 8`, a fixed initial `--load`, and +`--load-debug-rollout-data /absolute/dumps/{rollout_id}.pt`. Keep the same eight +saved rollout files in all runs. They must retain tokens, masks, sampled logprobs, +rewards and grouping metadata. Use nonzero-gradient batches, zero dropout, and a +fixed runtime. Omit vLLM-only options because debug replay does not start serving. +This is a deterministic correctness test, not a timing benchmark. + +Run from the repository root in three separate processes. The runner captures +stdout/stderr in `train.log` in each output directory. Each output must be new. + +```bash +python -m tests.test_checkpoint_resume_parity record --train-args initial.json --output /tmp/parity-A +python -m tests.test_checkpoint_resume_parity record --train-args initial.json --output /tmp/parity-B0 --stop-after 4 +# In resumed.json change only --load to /tmp/parity-B0/checkpoints. +# Keep --num-rollout 8 and omit --start-rollout-id, --finetune, +# --no-load-optim and --no-load-rng so the loader controls restoration. +python -m tests.test_checkpoint_resume_parity record --train-args resumed.json --output /tmp/parity-B1 +python -m tests.test_checkpoint_resume_parity compare --continuous /tmp/parity-A --first /tmp/parity-B0 --resumed /tmp/parity-B1 --output /tmp/parity-result.json +``` + +The first branch runs all eight updates. The split prefix stops after update four +without changing the scheduler horizon. The resumed process must load the prefix +checkpoint and run updates five through eight. Verification requires separate +trainer PIDs, complete per-step evidence, nonzero updates, identical model and FP32 +master parameter hashes, Adam moments/step, scheduler, RNG, gradients, loss metrics, +and training dump values (including selected-token forward logprobs and advantages). +The restored state before update five must also exactly match the prefix's final +state. Missing, duplicate or skipped evidence fails verification. + +The current fixture is TP=PP=CP=1 and one update per rollout. It intentionally does +not claim topology resharding, full-vocabulary logits, stochastic dropout, dataset +cursor recovery from online generation, or live serving synchronization. Test the +first post-resume weight transfer and subsequent live generation separately. GPU +allocation and process cleanup belong to the calling CI/supervisor; never use +host-wide `pkill` or `ray stop` to run this test on shared machines. diff --git a/tests/test_checkpoint_resume_parity.py b/tests/test_checkpoint_resume_parity.py new file mode 100644 index 00000000..f0454130 --- /dev/null +++ b/tests/test_checkpoint_resume_parity.py @@ -0,0 +1,107 @@ +"""Fixed-batch checkpoint resume runner; see tests/checkpoint_resume_parity.md.""" + +import argparse +import importlib.util +import json +import os +import sys +import tempfile +from pathlib import Path + + +def record(args): + import ray + + checkout = Path(__file__).resolve().parents[1] + output = args.output.resolve() + output.mkdir(parents=True, exist_ok=False) + tokens = json.loads(args.train_args.read_text()) + assert isinstance(tokens, list) and all(isinstance(t, str) for t in tokens) + assert "--load-debug-rollout-data" in tokens, "Use saved rollout batches for resume parity" + assert "--num-rollout" in tokens and "--load" in tokens + + def setarg(flag, value): + if flag in tokens: + tokens[tokens.index(flag) + 1] = str(value) + else: + tokens.extend([flag, str(value)]) + + setarg("--custom-megatron-before-train-step-hook-path", "tests.checkpoint_resume_audit.before_train") + setarg("--dump-details", output / "dumps") + setarg("--save", output / "checkpoints") + setarg("--save-interval", args.split) + setarg("--num-gpus-per-node", 1) + setarg("--actor-num-gpus-per-node", 1) + for flag in ["--tensor-model-parallel-size", "--pipeline-model-parallel-size", "--context-parallel-size"]: + setarg(flag, 1) + os.environ["VIME_RESUME_AUDIT_DIR"] = str(output) + (output / "train-args.json").write_text(json.dumps(tokens, indent=2) + "\n") + sys.argv = [str(checkout / "train.py")] + tokens + print(f"Recording training and audit logs in {output}", flush=True) + saved_fds = [os.dup(1), os.dup(2)] + log = (output / "train.log").open("w") + os.dup2(log.fileno(), 1) + os.dup2(log.fileno(), 2) + # Each invocation owns a fresh Ray head. Shared-machine supervisors should + # bind an available GPU before launching and clean up only this invocation. + try: + ray_temp = tempfile.mkdtemp(prefix="vime-resume-ray-") + (output / "ray-temp-dir.txt").write_text(ray_temp + "\n") + ray.init( + address="local", + num_cpus=12, + num_gpus=1, + include_dashboard=False, + object_store_memory=2 * 1024**3, + namespace=output.name, + _temp_dir=ray_temp, + _node_ip_address="127.0.0.1", + runtime_env={"env_vars": {"VIME_RESUME_AUDIT_DIR": str(output)}}, + ) + spec = importlib.util.spec_from_file_location("resume_parity_train", checkout / "train.py") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + train_args = module.parse_args() + if args.stop_after is not None: + assert 0 < args.stop_after < train_args.num_rollout + # Do not change num_rollout: the scheduler horizon must be identical + # in the continuous run, split prefix, and resumed suffix. + module.range = lambda start, stop: range(start, min(stop, args.stop_after)) + module.train(train_args) + (output / "train-returned.json").write_text(json.dumps({"completed": True})) + finally: + ray.shutdown() + sys.stdout.flush() + sys.stderr.flush() + for fd, original in zip([1, 2], saved_fds, strict=True): + os.dup2(original, fd) + os.close(original) + log.close() + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + subparsers = parser.add_subparsers(dest="action", required=True) + run = subparsers.add_parser("record") + run.add_argument("--train-args", type=Path, required=True) + run.add_argument("--output", type=Path, required=True) + run.add_argument("--split", type=int, default=4) + run.add_argument("--stop-after", type=int) + compare = subparsers.add_parser("compare") + for name in ["continuous", "first", "resumed", "output"]: + compare.add_argument("--" + name, type=Path, required=True) + compare.add_argument("--split", type=int, default=4) + compare.add_argument("--steps", type=int, default=8) + args = parser.parse_args() + if args.action == "record": + record(args) + else: + from tests.checkpoint_resume_audit import verify + + result = verify(args.continuous, args.first, args.resumed, args.split, args.steps) + args.output.write_text(json.dumps(result, indent=2) + "\n") + print(json.dumps(result, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/tests/utils/test_checkpoint_resume_audit.py b/tests/utils/test_checkpoint_resume_audit.py new file mode 100644 index 00000000..af35afd0 --- /dev/null +++ b/tests/utils/test_checkpoint_resume_audit.py @@ -0,0 +1,132 @@ +"""Negative controls for the resume evidence verifier; CPU only.""" + +import json + +import pytest +import torch + +from tests.checkpoint_resume_audit import tensor_record, verify + + +@pytest.fixture +def evidence(tmp_path): + roots = [tmp_path / name for name in ["continuous", "first", "resumed"]] + for root, ids, pid in zip(roots, [range(2), range(1), range(1, 2)], [101, 102, 103], strict=True): + (root / "dumps/train_data").mkdir(parents=True) + (root / "train-returned.json").write_text('{"completed": true}') + records, logs = [], [] + + def state(step): + return { + "model": {"weight": tensor_record(torch.tensor([step], dtype=torch.bfloat16))}, + "optimizer": [ + { + "class": "DistributedOptimizer", + "inner_class": "AdamW", + "master_parameters": [tensor_record(torch.tensor([step], dtype=torch.float32))], + "state": {"step": step, "exp_avg": step / 10, "exp_avg_sq": step / 100}, + } + ], + "scheduler": {"num_steps": step * 8, "max_lr": 0.001}, + "rng": {"torch_cpu": "a", "torch_cuda": "b", "python": "c", "numpy": "d"}, + } + + for rid in ids: + common = {"pid": pid, "rollout_id": rid, "step_id": 0} + records.extend( + [ + { + **common, + "kind": "before_update", + "state": state(rid), + "load_flags": { + "no_load_optim": False, + "no_load_rng": False, + "finetune": False, + "offload_train": True, + }, + }, + {**common, "kind": "gradients", "tensors": {"weight": tensor_record(torch.tensor([rid + 0.5]))}}, + {**common, "kind": "optimizer_result", "result": [True, 0.5, 0]}, + {**common, "kind": "after_update", "state": state(rid + 1)}, + ] + ) + logs.append(f"step {rid}: {{'train/loss': 0.25, 'train/grad_norm': 0.5, 'train/step': {rid}}}") + torch.save( + { + "tokens": torch.tensor([1, 2]), + "loss_masks": torch.tensor([0, 1]), + "log_probs": torch.tensor([-0.5]), + "advantages": torch.tensor([0.3]), + "rollout_id": rid, + }, + root / f"dumps/train_data/{rid}.pt", + ) + (root / "resume-audit.jsonl").write_text("\n".join(json.dumps(x) for x in records) + "\n") + (root / "train.log").write_text("\n".join(logs) + "\n") + return roots + + +def mutate_record(root, kind, mutate): + path = root / "resume-audit.jsonl" + records = [json.loads(line) for line in path.read_text().splitlines()] + row = next(row for row in records if row["kind"] == kind) + mutate(row) + path.write_text("\n".join(json.dumps(row) for row in records) + "\n") + + +def test_accepts_complete_exact_resume(evidence): + assert verify(*evidence, split=1, steps=2)["status"] == "FIXED_BATCH_FRESH_PROCESS_RESUME_EXACT" + + +@pytest.mark.parametrize("component", ["model", "optimizer", "scheduler", "rng"]) +@pytest.mark.parametrize("kind", ["before_update", "after_update"]) +def test_rejects_state_changes(evidence, component, kind): + mutate_record(evidence[2], kind, lambda row: row["state"].__setitem__(component, {"corrupted": True})) + with pytest.raises(AssertionError, match=component): + verify(*evidence, split=1, steps=2) + + +@pytest.mark.parametrize("name", ["no_load_optim", "no_load_rng", "finetune"]) +def test_rejects_resume_bypass(evidence, name): + mutate_record(evidence[2], "before_update", lambda row: row["load_flags"].__setitem__(name, True)) + with pytest.raises(AssertionError, match="Restore bypassed"): + verify(*evidence, split=1, steps=2) + + +def test_rejects_missing_gradient_evidence(evidence): + mutate_record(evidence[2], "gradients", lambda row: row.__setitem__("tensors", {})) + with pytest.raises(AssertionError, match="Missing gradient"): + verify(*evidence, split=1, steps=2) + + +def test_rejects_duplicate_update(evidence): + path = evidence[2] / "resume-audit.jsonl" + lines = path.read_text().splitlines() + path.write_text("\n".join(lines + [lines[0]]) + "\n") + with pytest.raises(AssertionError, match="rollout_ids"): + verify(*evidence, split=1, steps=2) + + +@pytest.mark.parametrize("key", ["tokens", "loss_masks", "log_probs", "advantages"]) +def test_rejects_changed_training_batch(evidence, key): + path = evidence[2] / "dumps/train_data/1.pt" + batch = torch.load(path, weights_only=False) + batch[key] = batch[key] + 1 + torch.save(batch, path) + with pytest.raises(AssertionError, match="training_data"): + verify(*evidence, split=1, steps=2) + + +def test_rejects_changed_loss(evidence): + path = evidence[2] / "train.log" + path.write_text(path.read_text().replace("0.25", "0.26")) + with pytest.raises(AssertionError, match="metrics"): + verify(*evidence, split=1, steps=2) + + +def test_scalar_and_bfloat16_hashing(): + assert tensor_record(torch.tensor(1.0))["shape"] == [] + assert tensor_record(torch.tensor([1.0], dtype=torch.bfloat16)) != tensor_record( + torch.tensor([2.0], dtype=torch.bfloat16) + ) diff --git a/tests/utils/test_native_adam_checkpoint.py b/tests/utils/test_native_adam_checkpoint.py new file mode 100644 index 00000000..f126b29b --- /dev/null +++ b/tests/utils/test_native_adam_checkpoint.py @@ -0,0 +1,92 @@ +"""Exercise MCore's native Torch Adam checkpoint path with real CPU AdamW state.""" + +import copy +import types +from types import SimpleNamespace + +import pytest +import torch + +distributed_optimizer = pytest.importorskip("megatron.core.optimizer.distrib_optimizer") + + +def make_optimizer(): + parameters = [torch.nn.Parameter(torch.tensor([1.0, 2.0])), torch.nn.Parameter(torch.tensor([3.0, 4.0]))] + group = {key: False if key.startswith("is_") else 1.0 for key in distributed_optimizer.param_group_identifier_keys} + optimizer = torch.optim.AdamW([{"params": parameters, **group}], lr=0.01) + wrapper = SimpleNamespace( + optimizer=optimizer, + grad_scaler=None, + ddp_config=SimpleNamespace(use_megatron_fsdp=False), + config=SimpleNamespace(fp16=False, use_precision_aware_optimizer_no_fp8_or_ds_fp8=False), + model_param_group_index_map={parameter: (0, index) for index, parameter in enumerate(parameters)}, + ) + wrapper.state_dict = types.MethodType(distributed_optimizer.DistributedOptimizer.state_dict, wrapper) + wrapper.load_state_dict = types.MethodType(distributed_optimizer.DistributedOptimizer.load_state_dict, wrapper) + wrapper._set_main_param_and_optimizer_states = types.MethodType( + distributed_optimizer.DistributedOptimizer._set_main_param_and_optimizer_states, wrapper + ) + return parameters, optimizer, wrapper + + +def update(parameters, optimizer): + for index, parameter in enumerate(parameters): + parameter.grad = torch.full_like(parameter, index + 0.5) + optimizer.step() + + +def test_fresh_native_optimizer_has_valid_checkpoint_template(monkeypatch): + monkeypatch.setattr(distributed_optimizer, "HAVE_APEX_OR_TE", False) + _, optimizer, wrapper = make_optimizer() + state = wrapper.state_dict() + assert not optimizer.state # Template inspection must not run an optimizer update. + assert all(group["step"] == 0 for group in state["optimizer"]["param_groups"]) + + +def test_restored_native_adam_steps_are_independent_and_update_once(monkeypatch): + monkeypatch.setattr(distributed_optimizer, "HAVE_APEX_OR_TE", False) + parameters, optimizer, wrapper = make_optimizer() + for _ in range(4): + update(parameters, optimizer) + saved_common = copy.deepcopy(wrapper.state_dict()) + saved_tensors = copy.deepcopy(optimizer.state_dict()) + restored_parameters, restored, restored_wrapper = make_optimizer() + # Preallocate CPU state. Distributed checkpoint loading separately allocates + # GPU shards; this unit test targets the scalar/non-parameter load method. + update(restored_parameters, restored) + restored_wrapper.load_state_dict(saved_common) + with torch.no_grad(): + for index, (old, new) in enumerate(zip(parameters, restored_parameters, strict=True)): + # DP-reshardable parameter shards exclude step: common state above + # already restored it. Exercise the real shard setter as well. + restored_wrapper._set_main_param_and_optimizer_states( + new, {"param": old, **{k: saved_tensors["state"][index][k] for k in ["exp_avg", "exp_avg_sq"]}} + ) + steps = [restored.state[p]["step"] for p in restored_parameters] + assert len({step.data_ptr() for step in steps}) == len(steps) + assert [step.item() for step in steps] == [4.0, 4.0] + assert all("step" not in group for group in restored.param_groups) + update(parameters, optimizer) + update(restored_parameters, restored) + for old, new in zip(parameters, restored_parameters, strict=True): + assert torch.equal(old, new) + for field in ["step", "exp_avg", "exp_avg_sq"]: + assert torch.equal(optimizer.state[old][field], restored.state[new][field]) + assert [step.item() for step in steps] == [5.0, 5.0] + + +def test_inconsistent_native_steps_still_fail(monkeypatch): + monkeypatch.setattr(distributed_optimizer, "HAVE_APEX_OR_TE", False) + parameters, optimizer, wrapper = make_optimizer() + update(parameters, optimizer) + optimizer.state[parameters[0]]["step"].fill_(2) + with pytest.raises(AssertionError): + wrapper.state_dict() + + +def test_missing_parameter_moment_still_fails(monkeypatch): + monkeypatch.setattr(distributed_optimizer, "HAVE_APEX_OR_TE", False) + parameters, optimizer, wrapper = make_optimizer() + update(parameters, optimizer) + with torch.no_grad(), pytest.raises(KeyError, match="exp_avg"): + wrapper._set_main_param_and_optimizer_states(parameters[0], {"param": parameters[0]})