#!/usr/bin/env python3 """One optimizer and one LR schedule across the complete S1->S10 ladder.""" from __future__ import annotations import argparse from concurrent.futures import ThreadPoolExecutor from contextlib import nullcontext from datetime import datetime, timedelta import hashlib import json import math import os from pathlib import Path import random import statistics import sys import time SC = Path("/e/scratch/reformo/schuhmann1_moss") HERE = Path(__file__).resolve().parent V3 = Path("/e/home/jusers/schuhmann1/jupiter/m2_600m_tts_training/cascade_v3_timed") M2 = SC / "code/m2_20k_score" SMALL = SC / "code/small_tts" for candidate in (str(HERE), str(V3), str(M2), str(SMALL)): while candidate in sys.path: sys.path.remove(candidate) sys.path[:0] = [str(HERE), str(V3), str(M2), str(SMALL)] def sha(path: str | Path) -> str: digest = hashlib.sha256() with Path(path).open("rb") as stream: for block in iter(lambda: stream.read(8 << 20), b""): digest.update(block) return digest.hexdigest() def validate(plan: dict, plan_path: Path) -> list[dict]: assert plan["kind"] == "m2_continuous_large_talker_v1" assert plan["architecture"] == "qwen3_0.6b_plus_fresh_sft3_width_talker" assert plan["scheduler"] == "single_linear_warmup_then_cosine_no_stage_restarts" assert plan["nodes"] >= 1 and plan["ranks_per_node"] in (1, 4) assert plan["world_size"] == plan["nodes"] * plan["ranks_per_node"] assert plan["global_batch"] == plan["world_size"] * plan["samples_per_gpu"] assert plan["lr_backbone_peak"] == 8e-5 assert plan["lr_talker_peak"] == 2.4e-4 assert plan["warmup_steps"] >= 1 and 0 <= plan["lr_floor_factor"] <= 1 assert plan["initialization"]["semantic"] == "original_pretrained_Qwen3-0.6B" assert plan["initialization"]["talker"] == "random_fresh" assert plan["checkpoint_policy"] == "all_stage_boundaries_plus_interval" expected_names = [f"S{i}" for i in range(1, 11)] assert [stage["name"] for stage in plan["stages"]] == expected_names specs = [] total_updates = 0 for stage in plan["stages"]: spec = json.loads(Path(stage["manifest"]).read_text()) assert spec["status"] == "complete" and spec["presentations"] > 0 assert sha(stage["manifest"]) == stage["manifest_sha256"] natural = math.ceil(spec["presentations"] / plan["global_batch"]) updates = int(stage.get("max_updates", natural)) assert 1 <= updates <= natural assert updates == stage["updates"] # A deliberately truncated smoke never reaches the natural final # partial batch. Validate that tail only when this plan consumes the # complete stage; every earlier update is a full global batch. if updates == natural: tail = spec["presentations"] - (natural - 1) * plan["global_batch"] assert tail >= plan["world_size"] total_updates += updates specs.append(spec) assert total_updates == plan["total_updates"] assert plan["warmup_steps"] < total_updates for name, expected in plan["code_sha256"].items(): source = HERE / name if name in ("continuous_train.py", "large_talker.py", "caption_curriculum_dataset.py") else V3 / name assert sha(source) == expected, (name, source) return specs def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("plan") args = parser.parse_args() plan_path = Path(args.plan).resolve() plan = json.loads(plan_path.read_text()) import numpy as np import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from transformers import AutoTokenizer from packing import ScorePacker from plan import step_indices, microbatches, dynamic_weights from continuation_control import latest_checkpoint, complete_checkpoint from cascade_dataset import CascadeRecords from caption_curriculum_dataset import CaptionCurriculumRecords import moss_small import state_io import va_loss from large_talker import build_fresh, parameter_counts specs = validate(plan, plan_path) rank, world, local = (int(os.environ[key]) for key in ("SLURM_PROCID", "SLURM_NTASKS", "SLURM_LOCALID")) assert world == int(plan["world_size"]) device_index = 0 if torch.cuda.device_count() == 1 else local torch.cuda.set_device(device_index) device = torch.device("cuda", device_index) os.environ.update(RANK=str(rank), WORLD_SIZE=str(world), LOCAL_RANK=str(device_index)) dist.init_process_group("nccl", timeout=timedelta(minutes=30), device_id=device) torch.set_num_threads(4) if plan.get("deterministic", False): torch.use_deterministic_algorithms(True) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False def log(message: str) -> None: if rank == 0: print(datetime.now().astimezone().isoformat(timespec="seconds"), message, flush=True) random.seed(plan["seed"]) np.random.seed(plan["seed"]) torch.manual_seed(plan["seed"]) torch.cuda.manual_seed_all(plan["seed"]) invocation = os.environ["SLURM_JOB_ID"] + os.environ.get("M2_INVOCATION_SUFFIX", "") invocation_dir = Path(plan["output"]) / "invocations" / invocation invocation_dir.mkdir(parents=True, exist_ok=True) if rank == 0: assert not (invocation_dir / "steps.jsonl").exists() schema_path = SC / "out/m2_20k_score/score_schema.json" schema = json.loads(schema_path.read_text()) model, config = build_fresh(schema, log=log) counts = parameter_counts(model) log("Fresh large-Talker model: " + json.dumps(counts, sort_keys=True)) from functools import partial from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( apply_activation_checkpointing, checkpoint_wrapper, CheckpointImpl) for module in model.modules(): if hasattr(module, "gradient_checkpointing"): module.gradient_checkpointing = False apply_activation_checkpointing( model, check_fn=lambda module: type(module).__name__ == "MossQwen3DecoderLayer", checkpoint_wrapper_fn=partial(checkpoint_wrapper, checkpoint_impl=CheckpointImpl.NO_REENTRANT), ) class Step(torch.nn.Module): def __init__(self, base): super().__init__() self.base = base def forward(self, input_ids, attention_mask, labels, score_conditioning, weights, scale): hidden = self.base(input_ids=input_ids, attention_mask=attention_mask, score_conditioning=score_conditioning, use_cache=False).last_hidden_state loss, per = va_loss.compute_supervised_loss_from_hidden( self.base, global_hidden_states=hidden, labels=labels, channelwise_loss_weight=weights, return_per_channel=True) return loss * scale, per step_model = DDP(Step(model.to(device)).train(), device_ids=[device_index], broadcast_buffers=False, find_unused_parameters=False, gradient_as_bucket_view=True) if plan["gradient_communication"] == "bfloat16": from torch.distributed.algorithms.ddp_comm_hooks.default_hooks import bf16_compress_hook step_model.register_comm_hook(None, bf16_compress_hook) else: assert plan["gradient_communication"] == "float32" groups = {"backbone": [], "talker": []} for name, parameter in model.named_parameters(): if not parameter.requires_grad: continue key = "backbone" if name.startswith(("transformer.", "text_lm_head.")) else "talker" groups[key].append(parameter) optimizer = torch.optim.AdamW([ {"params": groups["backbone"], "lr": plan["lr_backbone_peak"], "name": "backbone"}, {"params": groups["talker"], "lr": plan["lr_talker_peak"], "name": "talker"}, ], betas=(0.9, 0.95), eps=1e-8, weight_decay=0.1, foreach=False) total_updates = int(plan["total_updates"]) warmup = int(plan["warmup_steps"]) floor = float(plan["lr_floor_factor"]) def lr_factor(step: int) -> float: if step < warmup: return (step + 1) / warmup progress = min(1.0, (step - warmup) / max(1, total_updates - warmup)) return floor + (1.0 - floor) * 0.5 * (1.0 + math.cos(math.pi * progress)) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_factor) _, _, Processor = moss_small.export_classes() processor = Processor(tokenizer=AutoTokenizer.from_pretrained( moss_small.SFT3, trust_remote_code=True, local_files_only=True), audio_tokenizer=None, model_config=config) packer = ScorePacker(processor, config, schema) datasets = [] for stage, spec in zip(plan["stages"], specs): cls = CaptionCurriculumRecords if spec.get("format") == "m2-caption-curriculum-v1" else CascadeRecords datasets.append(cls(stage["manifest"])) assert all(len(dataset) == spec["presentations"] for dataset, spec in zip(datasets, specs)) manifest_digest = hashlib.sha256("".join( stage["manifest_sha256"] for stage in plan["stages"]).encode()).hexdigest() contract = { "plan_sha256": sha(plan_path), "schema_sha256": sha(schema_path), "manifest_sha256": manifest_digest, "world": world, "objective": "global_token_mean_continuous_ladder_v1", "seed": plan["seed"], } start_step = 0 selection = [None] if rank == 0: try: directory = latest_checkpoint(Path(plan["output"]) / "checkpoints", contract, world) if directory: complete_checkpoint(directory, contract, world, content=True) selection[0] = {"path": str(directory) if directory else None} except Exception as error: selection[0] = {"error": repr(error)} dist.broadcast_object_list(selection, src=0) assert "error" not in selection[0], selection[0] if selection[0]["path"]: start_step = state_io.load(selection[0]["path"], step_model, optimizer, scheduler, contract) log(f"Continuous resume point={start_step}/{total_updates}; one scheduler, no stage restarts") if start_step == total_updates: log(f"ALREADY_COMPLETE at {start_step}/{total_updates}") dist.barrier(); dist.destroy_process_group(); return assert start_step < total_updates def prepare(dataset, phase, phase_step): began = time.monotonic() indices = step_indices(phase, phase_step, rank, world, len(dataset)) examples = [dataset.example(index, packer) for index in indices] assert examples assert max(len(example["input_ids"]) for example in examples) <= int( getattr(config, "max_position_embeddings", 32768)) batches = microbatches(examples, plan["max_padded_tokens"], plan["max_examples"]) prepared = [] for examples_batch in batches: batch = packer.collate(examples_batch) frames = sum(example["accounting"]["frames"] for example in examples_batch) assert int(batch["labels"][:, :, 1].ge(0).sum()) == frames prepared.append((batch, frames, len(examples_batch))) stats = [len(examples), sum(x["accounting"]["frames"] for x in examples), sum(x["accounting"]["target_audio_hours"] for x in examples), sum(x["accounting"]["reference_frames"] for x in examples), sum(len(x["input_ids"]) for x in examples), sum(batch["attention_mask"].numel() for batch, _, _ in prepared)] stats += [sum(x["accounting"]["mode"] == mode for x in examples) for mode in ("instruction", "reference")] stats += [sum(x["accounting"]["form"] == form for x in examples) for form in ("A", "B")] allowed_prompt_ids = set(plan.get("allowed_prompt_format_ids", [plan["prompt_contract"]])) assert all(x["accounting"]["prompt_format_id"] in allowed_prompt_ids for x in examples) stats += [sum(x["accounting"]["prompt_timed"] for x in examples), sum(x["accounting"]["prompt_eligible"] for x in examples), sum(x["accounting"]["duration_tags"] for x in examples)] return prepared, stats, time.monotonic() - began def update(prepared, local_stats): stats = torch.tensor(local_stats, dtype=torch.float64, device=device) dist.all_reduce(stats) values = stats.tolist() total_samples, total_frames = int(values[0]), int(values[1]) optimizer.zero_grad(set_to_none=True) local_loss = torch.zeros((), device=device) per_sums = torch.zeros(13, device=device) for index, (batch, frames, samples) in enumerate(prepared): batch = {name: tuple(value.to(device, non_blocking=True) for value in content) if name == "score_conditioning" else content.to(device, non_blocking=True) for name, content in batch.items() if name not in ("modes", "meta")} weights, scale = dynamic_weights(frames, samples, total_frames, total_samples, world) context = step_model.no_sync() if index + 1 < len(prepared) else nullcontext() with context: with torch.autocast("cuda", dtype=torch.bfloat16): loss, per = step_model(**batch, weights=weights, scale=scale) assert torch.isfinite(loss) loss.backward() local_loss += loss.detach() / world per_sums += per * torch.tensor([frames + samples] + [frames] * 12, device=device) gradient_norm = torch.nn.utils.clip_grad_norm_(step_model.parameters(), 1.0) assert torch.isfinite(gradient_norm) optimizer.step() scheduler.step() dist.all_reduce(local_loss) dist.all_reduce(per_sums) per_sums /= torch.tensor([total_frames + total_samples] + [total_frames] * 12, device=device) torch.cuda.synchronize() return { "loss": float(local_loss), "per_channel": per_sums.tolist(), "global_samples": total_samples, "target_frames": total_frames, "target_audio_hours": values[2], "reference_frames": int(values[3]), "input_rows": int(values[4]), "padded_rows": int(values[5]), "mode_counts": dict(zip(("instruction", "reference"), map(int, values[6:8]))), "form_counts": dict(zip(("A", "B"), map(int, values[8:10]))), "prompt_timed": int(values[10]), "prompt_eligible": int(values[11]), "duration_tags": int(values[12]), "prompt_format_id": plan["prompt_contract"], "gradient_norm": float(gradient_norm), "lr_backbone": float(optimizer.param_groups[0]["lr"]), "lr_talker": float(optimizer.param_groups[1]["lr"]), } metrics = (invocation_dir / "steps.jsonl").open("a", buffering=1) if rank == 0 else None stage_offsets = [] cursor = 0 for stage in plan["stages"]: stage_offsets.append(cursor) cursor += stage["updates"] assert cursor == total_updates global_quarters = {math.ceil(total_updates * q / 4) for q in (1, 2, 3, 4)} stage_ends = {offset + stage["updates"] for offset, stage in zip(stage_offsets, plan["stages"])} if plan.get("smoke_mode"): # A full optimizer state is several GiB. The smoke proves both an early # boundary checkpoint and the final checkpoint without writing ten # redundant copies during its deliberately tiny stage transitions. global_quarters = set() stage_ends = set(map(int, plan["smoke_checkpoint_steps"])) last_save = time.monotonic() invocation_begin = time.monotonic() final_step = start_step best_loss = None best_step = None stopped = False observed_times = [] with ThreadPoolExecutor(max_workers=1) as pool: for stage_index, (stage, dataset, offset) in enumerate( zip(plan["stages"], datasets, stage_offsets)): phase = {"name": stage["name"], "offset": 0, "global_batch": plan["global_batch"], "updates": stage["updates"]} if start_step >= offset + stage["updates"]: continue pending = None local_start = max(0, start_step - offset) log(f"Entering {stage['name']} at local step {local_start}/{stage['updates']} " f"without optimizer/scheduler reset") for phase_step in range(local_start, stage["updates"]): if pending is None: pending = pool.submit(prepare, dataset, phase, phase_step) began = time.monotonic() ready, local_stats, preparation_seconds = pending.result() data_wait_seconds = time.monotonic() - began pending = None if phase_step + 1 < stage["updates"]: pending = pool.submit(prepare, dataset, phase, phase_step + 1) result = update(ready, local_stats) wall_seconds = time.monotonic() - began timing = torch.tensor([wall_seconds, data_wait_seconds, preparation_seconds], dtype=torch.float64, device=device) dist.all_reduce(timing, op=dist.ReduceOp.MAX) final_step = offset + phase_step + 1 observed_times.append(float(timing[0])) result.update( stage=stage["name"], stage_index=stage_index + 1, stage_step=phase_step + 1, stage_updates=stage["updates"], step=final_step, total_updates=total_updates, steady=final_step > warmup, wall_seconds=float(timing[0]), data_wait_seconds=float(timing[1]), preparation_seconds=float(timing[2]), job=os.environ["SLURM_JOB_ID"], invocation=invocation, ) if rank == 0: metrics.write(json.dumps(result) + "\n") log(f"{stage['name']} {phase_step+1}/{stage['updates']} global=" f"{final_step}/{total_updates} loss={result['loss']:.5f} " f"wall={result['wall_seconds']:.3f}s") if best_loss is None or result["loss"] < best_loss: best_loss, best_step = result["loss"], final_step interval_due = time.monotonic() - last_save >= plan["checkpoint_interval_seconds"] deadline_due = time.monotonic() - invocation_begin >= plan["max_invocation_seconds"] control = torch.tensor([int(interval_due), int(deadline_due)] if rank == 0 else [0, 0], device=device) dist.broadcast(control, src=0) stopped = bool(control[1]) and final_step < total_updates export = final_step in stage_ends or final_step in global_quarters if export or bool(control[0]) or stopped or final_step == total_updates: saved = state_io.save(Path(plan["output"]) / "checkpoints", step_model, optimizer, scheduler, final_step, contract, export=export or final_step == total_updates) last_save = time.monotonic() if rank == 0: log("Committed retained checkpoint " + str(saved)) dist.barrier() if stopped: log(f"STOPPED_RESUMABLE at global step {final_step}") break if stopped: break peak = torch.tensor(torch.cuda.max_memory_allocated() / 2**30, device=device) dist.all_reduce(peak, op=dist.ReduceOp.MAX) median_step = statistics.median(observed_times) threshold = plan.get("max_acceptable_median_step_seconds") throughput_failed = bool(threshold is not None and median_step > float(threshold)) if rank == 0: status = ("FAIL_THROUGHPUT" if throughput_failed else ("PASS" if final_step == total_updates else "STOPPED_RESUMABLE")) summary = { "status": status, "job": os.environ["SLURM_JOB_ID"], "invocation": invocation, "final_step": final_step, "total_updates": total_updates, "one_continuous_schedule": True, "optimizer_or_scheduler_restarts_at_stage_boundaries": 0, "best_observed_loss_this_invocation": best_loss, "best_observed_step_this_invocation": best_step, "median_step_seconds_this_invocation": median_step, "p90_step_seconds_this_invocation": float(np.quantile(observed_times, 0.9)), "peak_allocated_gib_max_rank": float(peak), **counts, "plan_sha256": sha(plan_path), "contract": contract, } (invocation_dir / "summary.json").write_text(json.dumps(summary, indent=2) + "\n") log(status + " " + str(invocation_dir / "summary.json")) dist.barrier() dist.destroy_process_group() if throughput_failed: raise SystemExit(42) if __name__ == "__main__": main()