#!/usr/bin/env python3 """Generate immutable production and two-stage transition-smoke plans.""" 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") 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() STAGES = [ ("S1", "x100", ROOT / "manifests/S1/manifest.json"), ("S2", "x50", ROOT / "manifests_v2/S2/manifest.json"), ("S3", "tier7", ROOT / "manifests/S2/manifest.json"), ("S4", "tier6", ROOT / "manifests/S3/manifest.json"), ("S5", "tier5", ROOT / "manifests/S4/manifest.json"), ("S6", "tier4", ROOT / "manifests/S5/manifest.json"), ("S7", "tier3", ROOT / "manifests/S5b/manifest.json"), ("S8", "tier2", ROOT / "manifests/S6/manifest.json"), ("S9", "tier1", ROOT / "manifests/S5c/manifest.json"), ("S10", "tier0", ROOT / "manifests/S7/manifest.json"), ] def stage(name, tier, manifest, global_batch, max_updates=None): spec = json.loads(manifest.read_text()) updates = math.ceil(spec["presentations"] / global_batch) if max_updates is not None: updates = min(updates, max_updates) return { "name": name, "tier": tier, "manifest": str(manifest), "manifest_sha256": sha(manifest), "presentations": spec["presentations"], "measured_hours": spec.get("hours"), "updates": updates, **({"max_updates": updates} if max_updates is not None else {}), } 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"), "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, "warmup_steps": 600, "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-cascade-timed-repair-v1-20260920", "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): if path.exists(): existing = json.loads(path.read_text()) comparable_existing = {key: value for key, value in existing.items() if key != "created_at"} comparable_payload = {key: value for key, value in payload.items() if key != "created_at"} assert comparable_existing == comparable_payload, f"Refusing to overwrite immutable plan: {path}" return path.write_text(json.dumps(payload, indent=2) + "\n") def main(): production = base(32, 32) production["stages"] = [stage(name, tier, manifest, production["global_batch"]) for name, tier, manifest in STAGES] production["total_updates"] = sum(item["updates"] for item in production["stages"]) production["output"] = str(ROOT / "training_continuous_large_talker_v1") production["purpose"] = "production_waiting_for_explicit_user_go" write_once(HERE / "plan_production_32nodes.json", production) smoke = base(1, 2) smoke["warmup_steps"] = 2 smoke["checkpoint_interval_seconds"] = 600 smoke["max_invocation_seconds"] = 2400 manifest = ROOT / "manifests/S1_io_smoke/manifest.json" smoke["stages"] = [stage(f"S{i}", f"smoke-transition-{i}", manifest, smoke["global_batch"], max_updates=2) for i in range(1, 11)] smoke["total_updates"] = sum(item["updates"] for item in smoke["stages"]) smoke["output"] = str(ROOT / "training_continuous_large_talker_v1_smoke") smoke["purpose"] = "4-GPU architecture, backward, checkpoint, and all-stage-transition smoke" smoke["smoke_mode"] = True smoke["smoke_checkpoint_steps"] = [2, smoke["total_updates"]] smoke["max_examples"] = 2 write_once(HERE / "plan_smoke_4gpu.json", smoke) # Nine samples/rank makes every real manifest's natural final batch contain # at least one sample on each of the four ranks; this is also validated even # though the canary intentionally executes only the first update per stage. real = base(1, 9) real["warmup_steps"] = 2 real["checkpoint_interval_seconds"] = 600 real["max_invocation_seconds"] = 2400 real["stages"] = [stage(name, tier, manifest, real["global_batch"], max_updates=1) for name, tier, manifest in STAGES] real["total_updates"] = sum(item["updates"] for item in real["stages"]) real["output"] = str(ROOT / "training_continuous_large_talker_v1_real_ladder_smoke_v2") real["purpose"] = "4-GPU exact-production-manifest canary across all ten stages" real["smoke_mode"] = True real["smoke_checkpoint_steps"] = [1, real["total_updates"]] real["max_examples"] = 3 write_once(HERE / "plan_real_ladder_smoke_4gpu_v2.json", real) print(json.dumps({"production_updates": production["total_updates"], "smoke_updates": smoke["total_updates"], "real_ladder_smoke_updates": real["total_updates"]}, indent=2)) if __name__ == "__main__": main()