Download code/continuous_train.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 21.7 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/continuous_train.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/continuous_train.py
-
curl -L -o continuous_train.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/continuous_train.py
21.7 kB
| #!/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() | |