matilda-jev-fp4 / kev /train_fsdp.py
yue-maincode's picture
Upload validated MATILDA JEV FP4 model and Decision Index scores
c69aaec verified
Raw History Blame Contribute Delete
26.3 kB
"""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 <run>/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()