"""Full-weight SFT on one node with FSDP2: `torchrun --nproc_per_node=8 -m kev.train_fsdp ...`. Same objective, schedule, augmentation, evaluation and checkpoint selection as kev.train, with parameters, gradients and AdamW state sharded across ranks (FP32 master weights, bf16 compute, FP32 gradient reduction) instead of a CPU-offloaded optimizer. Every FSDP forward/backward is a collective, so every rank must run the same number of them: each global batch is packed into length-sorted microbatches, dealt round-robin, and padded with zero-weight repeats. Evaluation shards rows the same way. """ import gc import json import math import os import random import shutil import subprocess import time from collections.abc import Sequence from datetime import timedelta from pathlib import Path from typing import Any, cast import torch import torch.distributed as dist import torch.distributed.checkpoint as dcp import torch.nn.functional as F from torch.distributed.checkpoint.state_dict import ( StateDictOptions, get_model_state_dict, get_optimizer_state_dict, set_model_state_dict, set_optimizer_state_dict, ) from torch.distributed.device_mesh import init_device_mesh from torch.distributed.fsdp import MixedPrecisionPolicy, fully_shard from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5Model from kev.evaluate import ( Metrics, by_panel, calibration_ok, evaluate_logits, fit_temperature, hard_label, label_index, metrics, options, read_predictions, read_rows, selection_key, validate_coverage, write_json, ) from kev.events import record from kev.continuation import validate_extension from kev.model import DecisionModel, save_artifact from kev.tracking import Tracker from kev.train import ( RESUME_INVARIANT, Arguments, append, augment, check_partitions, digest, learning_rate_factor, length_estimate, microbatches, parse_arguments, prune_checkpoints, save_predictions, select_checkpoint, sync_directory, targets, truncate_logs, ) from kev.types import Example, JSONValue type Plan = list[tuple[list[Example], bool]] def shard_group(group: Sequence[Example], world: int, batch_size: int, token_budget: int) -> list[Plan]: """Split a length-sorted group into contiguous, token-balanced segments, one per rank. Rows of similar length share a rank, so left-padding waste stays small, and every rank carries about the same number of tokens. Each rank packs its segment into microbatches; shorter plans are padded with a one-row, zero-weight microbatch (cheap) so every rank runs the same number of FSDP forward/backward passes. """ if not group: raise ValueError("An optimizer step cannot be empty") if len(group) < world: # Keep tiny final groups: unused ranks run zero-weight collective padding. return distribute([[row] for row in group], world) lengths = [length_estimate(row) for row in group] total, cumulative, cuts = sum(lengths), 0, [0] for index, length in enumerate(lengths): cumulative += length rank = len(cuts) if rank < world and cumulative >= total * rank / world: cuts.append(min(max(index + 1, cuts[-1] + 1), len(group) - (world - rank))) while len(cuts) < world: cuts.append(cuts[-1] + 1) cuts.append(len(group)) plans = [microbatches(group[cuts[rank]:cuts[rank + 1]], batch_size, token_budget) for rank in range(world)] depth = max(len(plan) for plan in plans) return [[(batch, True) for batch in plan] + [([plan[0][0]], False)] * (depth - len(plan)) for plan in plans] def distribute(batches: Sequence[list[Example]], world: int) -> list[Plan]: """Deal microbatches round-robin; pad with zero-weight repeats so every rank runs the same count.""" padding = -len(batches) % world padded = list(batches) + [batches[-1]] * padding real = [True] * len(batches) + [False] * padding return [[(padded[index], real[index]) for index in range(rank, len(padded), world)] for rank in range(world)] class Trainer: def __init__(self, args: Arguments) -> None: self.args = args self.run = Path(args.run) self.output = self.run / "checkpoints" use_gpu = torch.cuda.is_available() and args.device != "cpu" local_rank = int(os.environ.get("LOCAL_RANK", "0")) self.device = torch.device("cuda", local_rank) if use_gpu else torch.device("cpu") if use_gpu: torch.cuda.set_device(local_rank) dist.init_process_group("nccl" if use_gpu else "gloo", timeout=timedelta(minutes=60), device_id=self.device if use_gpu else None) self.rank, self.world = dist.get_rank(), dist.get_world_size() self.mesh = init_device_mesh(self.device.type, (self.world,)) self.main = self.rank == 0 def log(self, kind: str, **fields: object) -> dict[str, object]: return record(kind, **fields) if self.main else {} def barrier(self) -> None: dist.barrier() def load_model(self, train_vision: bool) -> None: """Every rank loads FP32 weights on CPU; sharding moves each rank's shard to its GPU.""" started = time.monotonic() model = DecisionModel(train=True, base_model=self.args.base_model, device="cpu", dtype=torch.float32, gradient_checkpointing=True, cpu_threads=max(1, self.args.cpu_threads // self.world)) if not train_vision: # Text-only data never reaches the vision tower: freeze it (no gradients, no optimizer state). model.backbone.visual.requires_grad_(False) self.backbone_config = model.backbone.config policy = MixedPrecisionPolicy(param_dtype=torch.bfloat16, reduce_dtype=torch.float32) for layer in model.backbone.language_model.layers: fully_shard(layer, mesh=self.mesh, mp_policy=policy) fully_shard(model.backbone.visual, mesh=self.mesh, mp_policy=policy) fully_shard(model, mesh=self.mesh, mp_policy=policy) model.device_name = str(self.device) gc.collect() self.model = model self.parameters = [parameter for parameter in model.parameters() if parameter.requires_grad] self.optimizer = torch.optim.AdamW(self.parameters, lr=self.args.lr, weight_decay=self.args.weight_decay) self.log("model_sharded", seconds=time.monotonic() - started, world=self.world, device=str(self.device), parameters=sum(p.numel() for p in self.parameters)) # ------------------------------------------------------------------ evaluation @torch.no_grad() # not inference_mode: FSDP reuses all-gather buffers later in training def infer(self, rows: Sequence[Example]) -> list[list[float]]: """Sharded inference; rank 0 receives logits for every row in input order (others get []).""" was_training = self.model.training self.model.eval() order = sorted(range(len(rows)), key=lambda index: length_estimate(rows[index])) position = {id(rows[index]): index for index in order} batches = microbatches([rows[index] for index in order], self.args.eval_batch_size, self.args.eval_token_budget) local: list[tuple[int, list[float]]] = [] for batch, real in distribute(batches, self.world)[self.rank]: logits = self.model(self.model.prepare(batch, max_length=self.args.max_length)).float().cpu() if real: for row, values in zip(batch, logits, strict=True): local.append((position[id(row)], cast(list[float], values[:len(options(row["question"]))].tolist()))) gathered: list[list[tuple[int, list[float]]] | None] = [None] * self.world dist.gather_object(local, gathered if self.main else None, dst=0) self.model.train(was_training) if not self.main: return [] result: list[list[float]] = [[] for _ in rows] for part in gathered: for index, values in cast(list[tuple[int, list[float]]], part): result[index] = values if any(not values for values in result): raise RuntimeError("Sharded inference missed rows") return result # ------------------------------------------------------------------ checkpoints def save_selected(self, step: int, temperature: float, provenance: dict[str, str], protected: int | None) -> None: """Gather the full model (collective); rank 0 writes a standard decision checkpoint in bf16.""" full = get_model_state_dict(self.model, options=StateDictOptions(full_state_dict=True, cpu_offload=True)) if self.main: began = time.monotonic() destination = self.output / f"step-{step:05d}" if destination.exists(): shutil.rmtree(destination) tensors = cast(dict[str, torch.Tensor], full) backbone = {name.removeprefix("backbone."): tensor.to(torch.bfloat16) for name, tensor in tensors.items() if name.startswith("backbone.")} with torch.device("meta"): shell = Qwen3_5Model._from_config(self.backbone_config) model = self.model save_artifact(destination, lambda path: shell.save_pretrained(str(path), state_dict=backbone, max_shard_size="5GB"), tensors["readout.weight"].to(torch.bfloat16), model.processor, base_model=model.base_model, revision=model.revision, codes=model.codes, token_ids=model.token_ids, temperature=temperature, metadata={"step": step, "provenance": cast(JSONValue, provenance), "trainer": "fsdp", "world_size": self.world}) for file in destination.rglob("*"): if file.is_file(): with file.open("rb") as stream: os.fsync(stream.fileno()) sync_directory(destination) select_checkpoint(self.output, step) prune_checkpoints(self.output, self.args.keep_checkpoints, protected) self.log("checkpoint_saved", step=step, path=str(destination), seconds=time.monotonic() - began) del full self.barrier() def save_resume(self, meta: dict[str, object]) -> None: """Sharded model + optimizer state via torch.distributed.checkpoint; rank 0 adds the cursor/RNG.""" began = time.monotonic() pending, final = self.run / "resume.pending", self.run / "resume" if self.main and pending.exists(): shutil.rmtree(pending) self.barrier() state = {"model": get_model_state_dict(self.model), "optimizer": get_optimizer_state_dict(self.model, self.optimizer)} dcp.save(state, checkpoint_id=str(pending)) # type: ignore[attr-defined] if self.main: torch.save(meta, pending / "trainer.pt") sync_directory(pending) previous = self.run / "resume.previous" if final.exists(): final.rename(previous) pending.rename(final) sync_directory(self.run) shutil.rmtree(previous, ignore_errors=True) self.log("resume_saved", step=meta["step"], seconds=time.monotonic() - began, bytes=sum(path.stat().st_size for path in final.rglob("*") if path.is_file())) self.barrier() def load_resume(self, source: Path | None = None) -> dict[str, object]: final = source if source is not None else self.run / "resume" state = {"model": get_model_state_dict(self.model), "optimizer": get_optimizer_state_dict(self.model, self.optimizer)} dcp.load(state, checkpoint_id=str(final)) # type: ignore[attr-defined] set_model_state_dict(self.model, cast(dict[str, Any], state["model"])) set_optimizer_state_dict(self.model, self.optimizer, cast(dict[str, Any], state["optimizer"])) return cast(dict[str, object], torch.load(final / "trainer.pt", map_location="cpu", weights_only=False)) def main(argv: Sequence[str] | None = None) -> None: args = parse_arguments(argv) trainer = Trainer(args) run, output = trainer.run, trainer.output if trainer.main: if (run / "config.json").exists() and not args.resume: raise ValueError("Run already exists; pass --resume or choose a new run directory") if args.resume and not (run / "resume" / "trainer.pt").exists(): raise ValueError("--resume needs an existing /resume/") if args.extend_from and not args.resume: if Path(args.extend_from).resolve() == run.resolve(): raise ValueError("An extension must use a separate run directory") if not (Path(args.extend_from) / "resume/trainer.pt").is_file(): raise ValueError("The source run has no optimizer resume state") output.mkdir(parents=True, exist_ok=True) trainer.barrier() os.environ.setdefault("KEV_EVENTS", str(run / "events.jsonl")) train = read_rows(Path(args.train)) development, temperature_rows = read_rows(Path(args.development)), read_rows(Path(args.temperature)) public_rows = read_rows(Path(args.public)) if args.public else [] check_partitions({"train": train, "development": development, "temperature": temperature_rows, "public": public_rows}) reference: Metrics | None = None if args.reference: reference_predictions = read_predictions(Path(args.reference)) validate_coverage(development, reference_predictions) reference = metrics(reference_predictions) inputs = {"train": args.train, "development": args.development, "temperature": args.temperature, "reference": args.reference, "public": args.public} hashes = {name: digest(path) for name, path in inputs.items() if path} package = Path(__file__).resolve().parent code_hashes = {name: digest(package / name) for name in ("train_fsdp.py", "train.py", "continuation.py", "model.py", "evaluate.py", "types.py")} lock = package.parents[1] / "uv.lock" if lock.exists(): code_hashes["uv.lock"] = digest(lock) try: git_commit: str | None = subprocess.check_output(["git", "rev-parse", "HEAD"], cwd=package, text=True, stderr=subprocess.DEVNULL).strip() except (subprocess.CalledProcessError, FileNotFoundError): git_commit = None config = {**vars(args), "trainer": "fsdp", "world_size": trainer.world, "data_sha256": hashes, "code_sha256": code_hashes, "git_commit": git_commit, "train_rows": len(train), "development_rows": len(development), "temperature_rows": len(temperature_rows), "public_rows": len(public_rows)} if trainer.main and not args.resume: write_json(run / "config.json", config) tracker = Tracker(run, args.wandb_project, args.wandb_mode, config, enabled=trainer.main) random.seed(args.seed) torch.manual_seed(args.seed) trainer.load_model(train_vision=any(row.get("images") for row in train)) model, optimizer = trainer.model, trainer.optimizer groups: list[list[Example]] = [] for epoch in range(args.epochs): ordered = list(train) random.Random(args.seed + epoch).shuffle(ordered) for offset in range(0, len(ordered), args.effective_batch_size): groups.append(sorted(ordered[offset:offset + args.effective_batch_size], key=length_estimate)) total_steps = len(groups) step, examples_seen = 0, 0 best: Metrics | None = None best_step: int | None = None resumable_best_step: int | None = None selected_temperature: float | None = None schedule_offset_step = 0 started = time.monotonic() if args.resume or args.extend_from: extension = bool(args.extend_from and not args.resume) source = Path(args.extend_from) / "resume" if extension else None meta = trainer.load_resume(source) saved_config = cast(dict[str, object], meta["config"]) if extension: schedule_offset_step = validate_extension(meta, config, total_steps, hashes, Path(args.extend_from)) else: if meta["data_sha256"] != hashes or meta["total_steps"] != total_steps: raise ValueError("Resume data differs from the saved run") if saved_config["code_sha256"] != code_hashes: raise ValueError("Training implementation differs from the saved run") for key in RESUME_INVARIANT: if saved_config[key] != vars(args)[key]: raise ValueError(f"Resume configuration differs: {key}") schedule_offset_step = int(meta.get("schedule_offset_step", 0)) step, examples_seen = cast(int, meta["step"]), cast(int, meta["examples_seen"]) if args.stop_after is not None and args.stop_after <= step: raise ValueError("The requested stopping step must follow the saved step") best, best_step = cast(Metrics | None, meta["best"]), cast(int | None, meta["best_step"]) selected_temperature = cast(float | None, meta["selected_temperature"]) resumable_best_step = best_step if trainer.main: truncate_logs(run, step) for path in output.glob("step-*"): if path.is_dir() and int(path.name.rsplit("-", 1)[1]) > step: shutil.rmtree(path) if best_step is not None: select_checkpoint(output, best_step) trainer.barrier() trainer.log("training_started", run=run.name, git_commit=git_commit, data_sha256=hashes, world=trainer.world, initialization="exact_resume" if args.resume else "completed_run_extension" if args.extend_from else "base_fresh_optimizer", step=step, total_steps=total_steps, schedule_offset_step=schedule_offset_step, optimizer_state_restored=bool(args.resume or args.extend_from)) def evaluate(include_public: bool = False) -> None: nonlocal best, best_step, selected_temperature began = time.monotonic() temperature_logits = trainer.infer(temperature_rows) logits = trainer.infer(development) public_logits = trainer.infer(public_rows) if include_public and public_rows else [] decision: list[object] = [None] if trainer.main: fitted = fit_temperature(temperature_logits, [label_index(options(row["question"]), hard_label(row)) for row in temperature_rows]) raw_predictions, fitted_predictions = evaluate_logits(development, logits), evaluate_logits(development, logits, fitted) raw, calibrated = metrics(raw_predictions), metrics(fitted_predictions) eligible = reference is None or calibration_ok(calibrated, reference) improved = eligible and (best is None or selection_key(calibrated, step) < selection_key(best, cast(int, best_step))) save_predictions(run / f"development-{step:05d}.jsonl", fitted_predictions) save_predictions(run / f"temperature-{step:05d}.jsonl", evaluate_logits(temperature_rows, temperature_logits, fitted)) panels = by_panel(development, fitted_predictions) tracker.log_evaluation(step, raw, calibrated, panels, fitted, time.monotonic() - began) value = record("evaluation", run=run.name, step=step, examples_seen=examples_seen, raw=raw, fitted=calibrated, panels=panels, temperature=fitted, reference=reference, eligible=eligible, improved=improved, elapsed_seconds=time.monotonic() - started, evaluation_seconds=time.monotonic() - began) append(run / "evaluations.jsonl", value) print(json.dumps({"step": step, "temperature": round(fitted, 4), "accuracy": round(calibrated["accuracy"], 4), "ece": round(calibrated["ece"], 4), "brier": round(calibrated["brier"], 4), "improved": improved, "panels": {name: round(m["accuracy"], 4) for name, m in panels.items()}}), flush=True) if public_logits: public_predictions = evaluate_logits(public_rows, public_logits, fitted) save_predictions(run / f"public-{step:05d}.jsonl", public_predictions) append(run / "public-evaluations.jsonl", record("public_evaluation", run=run.name, step=step, temperature=fitted, metrics=metrics(public_predictions))) decision = [(improved, fitted, calibrated)] dist.broadcast_object_list(decision, src=0) improved, fitted, calibrated = cast(tuple[bool, float, Metrics], decision[0]) if improved: trainer.save_selected(step, fitted, {"run": run.name, "git_commit": git_commit or "", **hashes, **code_hashes}, resumable_best_step) best, best_step, selected_temperature = calibrated, step, fitted def save_resume() -> None: nonlocal resumable_best_step trainer.save_resume({"step": step, "examples_seen": examples_seen, "total_steps": total_steps, "data_sha256": hashes, "config": config, "best": best, "best_step": best_step, "selected_temperature": selected_temperature, "python_rng": random.getstate(), "schedule_offset_step": schedule_offset_step}) resumable_best_step = best_step if trainer.main: prune_checkpoints(output, args.keep_checkpoints, resumable_best_step) trainer.barrier() if step == 0: evaluate(include_public=bool(public_rows)) stop = min(total_steps, args.stop_after) if args.stop_after is not None else total_steps while step < stop: began = time.monotonic() model.train() group = augment(groups[step], random.Random(args.seed + 100003 * (step + 1))) plan = shard_group(group, trainer.world, args.batch_size, args.token_budget)[trainer.rank] optimizer.zero_grad(set_to_none=True) local = torch.zeros(3, dtype=torch.float64, device=trainer.device) # loss sum, rows, input tokens for batch, real in plan: prepared = model.prepare(batch, max_length=args.max_length) logits = model(prepared) losses = -(targets(batch, logits.device) * F.log_softmax(logits, dim=-1)).sum(-1) if not torch.isfinite(losses).all(): raise RuntimeError(f"Nonfinite loss at step {step + 1}") # FSDP averages gradients over ranks, so scale by world size to get the mean over the whole group. weight = trainer.world / len(group) if real else 0.0 (losses.sum() * weight).backward() if real: local += torch.tensor([float(losses.detach().sum()), len(batch), prepared.input_tokens], dtype=torch.float64, device=trainer.device) del logits, losses, prepared dist.all_reduce(local) gradient_norm = torch.nn.utils.clip_grad_norm_(trainer.parameters, 1.0) norm = float(gradient_norm.full_tensor() if hasattr(gradient_norm, "full_tensor") else gradient_norm) if not math.isfinite(norm): raise RuntimeError(f"Nonfinite gradient at step {step + 1}") step += 1 factor = learning_rate_factor(step - schedule_offset_step, total_steps - schedule_offset_step, args.warmup_fraction, args.min_lr_ratio) for parameter_group in optimizer.param_groups: parameter_group["lr"] = args.lr * factor optimizer.step() examples_seen += len(group) peak = torch.tensor([torch.cuda.max_memory_allocated() / 1e9 if trainer.device.type == "cuda" else 0.0], device=trainer.device) dist.all_reduce(peak, op=dist.ReduceOp.MAX) if trainer.main: seconds = time.monotonic() - began value = record("training_step", run=run.name, step=step, loss=float(local[0] / local[1]), learning_rate=args.lr * factor, examples_seen=examples_seen, group_examples=len(group), input_tokens=int(local[2]), tokens_per_second=float(local[2]) / seconds, microbatches_per_rank=len(plan), gradient_norm=norm, step_seconds=seconds, elapsed_seconds=time.monotonic() - started, gpu_peak_gb=float(peak[0])) append(run / "training.jsonl", value) tracker.log_training(value) print(json.dumps(value), flush=True) if step % args.eval_every == 0 or step == stop: evaluate(include_public=step % args.public_eval_every == 0 or step == total_steps) if step % args.resume_every == 0 or step == stop: save_resume() if trainer.main: summary = {"run": run.name, "steps": step, "planned_steps": total_steps, "complete": step == total_steps, "examples_seen": examples_seen, "best_step": best_step, "selected_temperature": selected_temperature, "selected_metrics": best, "reference": reference, "data_sha256": hashes, "git_commit": git_commit, "world_size": trainer.world, "elapsed_seconds": time.monotonic() - started, "epochs": args.epochs, "extended_from": args.extend_from, "schedule_offset_step": schedule_offset_step, "checkpoint": str(output / "selected") if best else None} write_json(run / "summary.json", summary) record("training_finished" if step == total_steps else "training_paused", **summary) tracker.summary({"best_step": best_step, "selected_temperature": selected_temperature, "steps": step, **({f"best/{key}": best[key] for key in ("accuracy", "ece", "brier", "nll")} if best else {})}) print(json.dumps({key: summary[key] for key in ("steps", "planned_steps", "best_step", "checkpoint")}), flush=True) tracker.finish() trainer.barrier() dist.destroy_process_group() if __name__ == "__main__": main()