#!/usr/bin/env python3 """Freeze immutable smoke and production plans after virtual manifests pass.""" from __future__ import annotations from datetime import datetime import hashlib import json import math from pathlib import Path HERE = Path(__file__).resolve().parent ROOT = Path("/e/scratch/reformo/schuhmann1_moss/out/m2_600m_ladder/caption_curriculum_v1") V3 = Path("/e/home/jusers/schuhmann1/jupiter/m2_600m_tts_training/cascade_v3_timed") def sha(path): h = hashlib.sha256() with Path(path).open("rb") as stream: for block in iter(lambda: stream.read(8 << 20), b""): h.update(block) return h.hexdigest() def stages(global_batch, max_updates=None): result = [] for number in range(1, 11): manifest = ROOT / "manifests" / f"S{number}" / "manifest.json" spec = json.loads(manifest.read_text()) assert spec["status"] == "complete" and spec["format"] == "m2-caption-curriculum-v1" natural = math.ceil(spec["presentations"] / global_batch) updates = natural if max_updates is None else min(max_updates, natural) result.append({"name": f"S{number}", "tier": spec["tier"], "manifest": str(manifest), "manifest_sha256": sha(manifest), "presentations": spec["presentations"], "measured_hours": spec["hours"], "updates": updates, **({"max_updates": updates} if max_updates is not None else {})}) return result def base(nodes, samples_per_gpu): world = nodes * 4 code = { "continuous_train.py": sha(HERE / "continuous_train.py"), "large_talker.py": sha(HERE / "large_talker.py"), "caption_curriculum_dataset.py": sha(HERE / "caption_curriculum_dataset.py"), "cascade_dataset.py": sha(V3 / "cascade_dataset.py"), "timed_prompt.py": sha(V3 / "timed_prompt.py"), "anno_timed.py": sha(V3 / "anno_timed.py"), } return { "kind": "m2_continuous_large_talker_v1", "architecture": "qwen3_0.6b_plus_fresh_sft3_width_talker", "initialization": {"semantic": "original_pretrained_Qwen3-0.6B", "talker": "random_fresh"}, "nodes": nodes, "ranks_per_node": 4, "world_size": world, "samples_per_gpu": samples_per_gpu, "global_batch": world * samples_per_gpu, "seed": 20260923, "max_padded_tokens": 12000, "max_examples": 16, "lr_backbone_peak": 8e-5, "lr_talker_peak": 2.4e-4, "lr_floor_factor": 0.05, "scheduler": "single_linear_warmup_then_cosine_no_stage_restarts", "gradient_communication": "bfloat16", "deterministic": False, "checkpoint_policy": "all_stage_boundaries_plus_interval", "checkpoint_interval_seconds": 3600, "max_invocation_seconds": 40000, "resume_policy": "latest_complete_same_plan_world", "prompt_contract": "m2-multiformat-caption-v1-20260923", "allowed_prompt_format_ids": ["m2-cascade-timed-repair-v1-20260920", "m2-multiformat-caption-v1-20260923"], "codec_contract": "MOSS-Audio-Tokenizer-v2; 12 codebooks/frame; 12.5 frames/s; 150 codebook tokens/s", "code_sha256": code, "created_at": datetime.now().astimezone().isoformat(), } def write_once(path, payload): assert not path.exists(), f"Refusing to overwrite immutable plan {path}" path.write_text(json.dumps(payload, indent=2) + "\n") def main(): one = base(1, 2); one["stages"] = stages(one["global_batch"], 1) one["total_updates"] = sum(x["updates"] for x in one["stages"]); one["warmup_steps"] = 2 one["checkpoint_interval_seconds"] = 600; one["max_invocation_seconds"] = 2400 one["output"] = str(ROOT / "smoke_4gpu"); one["smoke_mode"] = True one["smoke_checkpoint_steps"] = [1, one["total_updates"]] one["purpose"] = "real-container schema, prompt, backward and resume smoke" write_once(HERE / "plan_caption_smoke_4gpu.json", one) wide = base(32, 32); wide["stages"] = stages(wide["global_batch"], 3) wide["total_updates"] = sum(x["updates"] for x in wide["stages"]); wide["warmup_steps"] = 3 wide["checkpoint_interval_seconds"] = 600; wide["max_invocation_seconds"] = 3600 wide["output"] = str(ROOT / "smoke_32nodes"); wide["smoke_mode"] = True wide["smoke_checkpoint_steps"] = [1, wide["total_updates"]] wide["max_acceptable_median_step_seconds"] = 5.0 wide["purpose"] = "32-node throughput acceptance gate for expanded caption mix" write_once(HERE / "plan_caption_smoke_32nodes.json", wide) production = base(32, 32); production["stages"] = stages(production["global_batch"]) production["total_updates"] = sum(x["updates"] for x in production["stages"]) production["warmup_steps"] = round(0.05 * production["total_updates"]) production["output"] = str(ROOT / "training") production["purpose"] = "authorized continuous S1-to-S10 multi-caption production" write_once(HERE / "plan_caption_production_32nodes.json", production) print(json.dumps({"smoke_4gpu": one["total_updates"], "smoke_32nodes": wide["total_updates"], "production": production["total_updates"], "warmup": production["warmup_steps"]}, indent=2)) if __name__ == "__main__": main()