Feature Extraction
Transformers
Safetensors
qwen3_5
matilda
jev
fp4
quantized
maincode
8-bit precision
Instructions to use Maincode/matilda-jev-fp4 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Maincode/matilda-jev-fp4 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="Maincode/matilda-jev-fp4")# pip install -U transformers accelerate # Load model directly from transformers import AutoProcessor, AutoModel processor = AutoProcessor.from_pretrained("Maincode/matilda-jev-fp4") model = AutoModel.from_pretrained("Maincode/matilda-jev-fp4", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download kev/train_fsdp.py from Maincode/matilda-jev-fp4: direct link, hf CLI and curl.
- Browser
- Download file 26.3 kB
-
https://huggingface.co/Maincode/matilda-jev-fp4/resolve/main/kev/train_fsdp.py
- Command line
-
hf download hf://Maincode/matilda-jev-fp4/kev/train_fsdp.py
-
curl -L -o train_fsdp.py https://huggingface.co/Maincode/matilda-jev-fp4/resolve/main/kev/train_fsdp.py
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 | |
| # 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() | |