Download nexora/training.py from devildasdf/NEXORA: direct link, hf CLI and curl.
- Browser
- Download file 7.88 kB
-
https://huggingface.co/devildasdf/NEXORA/resolve/main/nexora/training.py
- Command line
-
hf download hf://devildasdf/NEXORA/nexora/training.py
-
curl -L -o training.py https://huggingface.co/devildasdf/NEXORA/resolve/main/nexora/training.py
7.88 kB
| """Deterministic reference training with transactional, checksummed checkpoints.""" | |
| from dataclasses import asdict | |
| from hashlib import sha256 | |
| from pathlib import Path | |
| import json | |
| import math | |
| import os | |
| import random | |
| import time | |
| import uuid | |
| import numpy as np | |
| import torch | |
| from safetensors.torch import save_file | |
| from .model import ModelConfig, NexoraLM | |
| def seed_all(seed): | |
| random.seed(seed) | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| if torch.cuda.is_available(): | |
| torch.cuda.manual_seed_all(seed) | |
| def save_checkpoint(root, model, optimizer, step, generator, metadata): | |
| root = Path(root) | |
| root.mkdir(parents=True, exist_ok=True) | |
| name = f"step-{step:06d}-{uuid.uuid4().hex[:8]}.pt" | |
| target = root / name | |
| temp = root / (name + ".tmp") | |
| state = {"model": model.state_dict(), "optimizer": optimizer.state_dict(), "step": step, | |
| "config": asdict(model.config), "torch_rng": torch.get_rng_state(), | |
| "cuda_rng": torch.cuda.get_rng_state_all() if torch.cuda.is_available() else [], | |
| "batch_rng": generator.get_state(), "python_rng": random.getstate(), | |
| "numpy_rng": (np.random.get_state()[0], np.random.get_state()[1].tolist(), *np.random.get_state()[2:]), | |
| "metadata": metadata} | |
| torch.save(state, temp) | |
| os.replace(temp, target) | |
| receipt = {"file": name, "sha256": sha256(target.read_bytes()).hexdigest(), "step": step} | |
| receipt_tmp = root / "latest.json.tmp" | |
| receipt_tmp.write_text(json.dumps(receipt), encoding="utf-8") | |
| os.replace(receipt_tmp, root / "latest.json") | |
| return receipt | |
| def load_checkpoint(root, model, optimizer, generator): | |
| root = Path(root) | |
| receipt = json.loads((root / "latest.json").read_text(encoding="utf-8")) | |
| path = (root / receipt["file"]).resolve() | |
| if path.parent != root.resolve() or sha256(path.read_bytes()).hexdigest() != receipt["sha256"]: | |
| raise ValueError("Checkpoint path or checksum mismatch") | |
| state = torch.load(path, map_location="cpu", weights_only=True) | |
| if state["config"] != asdict(model.config): | |
| raise ValueError("Checkpoint architecture mismatch") | |
| model.load_state_dict(state["model"]) | |
| optimizer.load_state_dict(state["optimizer"]) | |
| torch.set_rng_state(state["torch_rng"]) | |
| if state["cuda_rng"] and torch.cuda.is_available(): | |
| torch.cuda.set_rng_state_all(state["cuda_rng"]) | |
| generator.set_state(state["batch_rng"]) | |
| random.setstate(state["python_rng"]) | |
| n = state["numpy_rng"] | |
| np.random.set_state((n[0], np.asarray(n[1], dtype=np.uint32), *n[2:])) | |
| return state | |
| def batch(tokens, size, length, generator, device): | |
| if len(tokens) <= length: | |
| raise ValueError("Shard too short for sequence length") | |
| starts = torch.randint(len(tokens) - length, (size,), generator=generator) | |
| x = torch.stack([tokens[i:i + length] for i in starts]).to(device) | |
| y = torch.stack([tokens[i + 1:i + length + 1] for i in starts]).to(device) | |
| return x, y | |
| def train(config_path, data_dir, output, *, resume=False, stop_after=None): | |
| config = json.loads(Path(config_path).read_text(encoding="utf-8")) | |
| c, t = ModelConfig(**config["model"]), config["training"] | |
| if t["steps"] < 1 or t["batch_size"] < 1 or not 0 < t["sequence_length"] <= c.max_context: | |
| raise ValueError("Invalid training configuration") | |
| torch.set_num_threads(t.get("threads", 4)) | |
| seed_all(t["seed"]) | |
| device = t.get("device", "auto") | |
| device = ("cuda" if torch.cuda.is_available() else "cpu") if device == "auto" else device | |
| model = NexoraLM(c).to(device) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=t["learning_rate"], weight_decay=0.1) | |
| generator = torch.Generator().manual_seed(t["seed"] + 1) | |
| out, data = Path(output), Path(data_dir) | |
| out.mkdir(parents=True, exist_ok=True) | |
| manifest_path = data / "manifest.json" | |
| manifest = json.loads(manifest_path.read_text(encoding="utf-8")) | |
| if manifest["tokenizer"] != "utf8-byte-v1" or c.vocab_size != 259: | |
| raise ValueError("Tokenizer/model incompatibility") | |
| for split in ("train", "validation"): | |
| info = manifest["shards"][split] | |
| if sha256((data / info["file"]).read_bytes()).hexdigest() != info["sha256"]: | |
| raise ValueError("Dataset checksum mismatch") | |
| train_ids = torch.from_numpy(np.load(data / "train.npy", allow_pickle=False).astype(np.int64)) | |
| val_ids = torch.from_numpy(np.load(data / "validation.npy", allow_pickle=False).astype(np.int64)) | |
| metadata = {"data_manifest_sha256": sha256(manifest_path.read_bytes()).hexdigest(), | |
| "training": t, "experiment_id": uuid.uuid4().hex} | |
| start = 0 | |
| if resume: | |
| state = load_checkpoint(out / "checkpoints", model, optimizer, generator) | |
| if state["metadata"]["data_manifest_sha256"] != metadata["data_manifest_sha256"] or state["metadata"]["training"] != t: | |
| raise ValueError("Resume requires identical data and training configuration") | |
| metadata, start = state["metadata"], state["step"] | |
| metrics = [] | |
| start_time = time.perf_counter() | |
| if device.startswith("cuda"): | |
| torch.cuda.reset_peak_memory_stats() | |
| final_step = min(t["steps"], stop_after) if stop_after is not None else t["steps"] | |
| if final_step < start: | |
| raise ValueError("Stop step precedes checkpoint") | |
| for step in range(start, final_step): | |
| model.train() | |
| # Schedule depends on planned total steps, so interrupted runs resume identically. | |
| lr = t["learning_rate"] * (0.1 + 0.9 * (1 + math.cos(math.pi * step / t["steps"])) / 2) | |
| for group in optimizer.param_groups: | |
| group["lr"] = lr | |
| x, y = batch(train_ids, t["batch_size"], t["sequence_length"], generator, device) | |
| optimizer.zero_grad(set_to_none=True) | |
| _, loss = model(x, y) | |
| if not torch.isfinite(loss): | |
| raise FloatingPointError("Non-finite training loss") | |
| loss.backward() | |
| grad = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0, error_if_nonfinite=True) | |
| optimizer.step() | |
| if step == start or (step + 1) % t["eval_every"] == 0 or step + 1 == final_step: | |
| model.eval() | |
| vg = torch.Generator().manual_seed(917) | |
| with torch.no_grad(): | |
| vx, vy = batch(val_ids, t["batch_size"], t["sequence_length"], vg, device) | |
| _, vl = model(vx, vy) | |
| row = {"step": step + 1, "train_loss": loss.item(), "validation_loss": vl.item(), "grad_norm": float(grad), "lr": lr} | |
| metrics.append(row) | |
| print(json.dumps(row), flush=True) | |
| if (step + 1) % t["checkpoint_every"] == 0 or step + 1 == final_step: | |
| save_checkpoint(out / "checkpoints", model, optimizer, step + 1, generator, metadata) | |
| if device.startswith("cuda"): | |
| torch.cuda.synchronize() | |
| elapsed = time.perf_counter() - start_time | |
| save_file({k: v.detach().cpu().contiguous() for k, v in model.state_dict().items()}, str(out / "model.safetensors")) | |
| (out / "config.json").write_text(json.dumps(asdict(c), indent=2), encoding="utf-8") | |
| report = {"status": "VALIDATED_SMALL_TRAINING_ONLY", "parameters": model.parameter_count(), "device": device, | |
| "elapsed_seconds": elapsed, "tokens_per_second_including_eval_and_checkpoints": (final_step-start)*t["batch_size"]*t["sequence_length"]/max(elapsed, 1e-9), | |
| "peak_vram_bytes": torch.cuda.max_memory_allocated() if device.startswith("cuda") else None, | |
| "steps_completed": final_step, "metadata": metadata, "metrics": metrics, | |
| "limitations": "Tiny educational corpus; no general assistant, reasoning or coding capability claim"} | |
| (out / ("resume-report.json" if resume else "training-report.json")).write_text(json.dumps(report, indent=2), encoding="utf-8") | |
| return report | |