File size: 6,384 Bytes
5de8603 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | #!/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()
|