Download training/code/train_decisions.py from Endikavi/Ines-1: direct link, hf CLI and curl.
- Browser
- Download file 9.97 kB
-
https://huggingface.co/Endikavi/Ines-1/resolve/main/training/code/train_decisions.py
- Command line
-
hf download hf://Endikavi/Ines-1/training/code/train_decisions.py
-
curl -L -o train_decisions.py https://huggingface.co/Endikavi/Ines-1/resolve/main/training/code/train_decisions.py
9.97 kB
| """Fine-tune mini-v41 to answer typed decisions (see decisions.py) and measure it. | |
| CUDA_VISIBLE_DEVICES=3 .venv/bin/python scripts/decisions/train_decisions.py \ | |
| --checkpoint runs/1p6b-pretrain-v3/inference_final \ | |
| --train a.jsonl --train b.jsonl --val va.jsonl --val vb.jsonl --test t.jsonl \ | |
| --out $DATA_ROOT/decisions/models/decisions-v1 [--epochs 3] [--eval-base] | |
| - Loss: cross-entropy against the target distribution over the presented options (a soft | |
| teacher distribution when the case carries `suave`, else one-hot), plus the MoE aux loss. | |
| Full fine-tune, fp32 master weights under bf16 autocast, one example per forward, | |
| gradient accumulation, warmup + linear decay. | |
| - choice options are re-shuffled every epoch (position bias); score keeps its order. | |
| - Epoch selection: mean accuracy over the --val FILES, each file weighing the same (so a big | |
| file does not drown a small one). The best epoch's weights are kept. | |
| - --test files are evaluated ONCE, at the end, with the selected weights; also with Engram | |
| switched off (engram_disabled) to see whether the n-gram tables matter, and a latency check | |
| of the scoring path (full forward vs prefill(num_logits=1)). | |
| - Saved with checkpoint.save_inference_checkpoint -> InferenceModel.from_checkpoint loads it. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import math | |
| import random | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import torch | |
| import torch.nn.functional as F | |
| sys.path.insert(0, str(Path(__file__).resolve().parent)) | |
| from decisions import Reader, accuracy, evaluate, examples, read_jsonl # noqa: E402 | |
| def evaluate_files(reader, files, details_dir=None, tag=""): | |
| res = {} | |
| for f in files: | |
| det = [] if details_dir else None | |
| r = evaluate(reader, read_jsonl(f), det) | |
| res[Path(f).stem] = r | |
| if details_dir: | |
| (details_dir / ("details_%s%s.jsonl" % (tag, Path(f).stem))).write_text( | |
| "".join(json.dumps(d, ensure_ascii=False) + "\n" for d in det)) | |
| return res | |
| def overall(r): | |
| tot = [v for k, v in r.items() if "/" not in k] | |
| a = sum(int(v.split("/")[0]) for v in tot) | |
| b = sum(int(v.split("/")[1].split(" ")[0]) for v in tot) | |
| return a / b if b else 0.0 | |
| def latency(reader, cases, n=50): | |
| rows = [] | |
| for c in cases[:n]: | |
| for qid, q in c["questions"].items(): | |
| rows.append(reader.prompt(c["state"], q, c.get("lang", "es"))) | |
| break | |
| model = reader.im.model | |
| model.eval() | |
| out = {} | |
| for name, fn in (("forward", lambda ids, k: model(torch.tensor([ids], device=reader.im.device))), | |
| ("prefill", lambda ids, k: reader.logits_eval(ids, k))): | |
| with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16): | |
| fn(*[rows[0][0], len(rows[0][1])]) | |
| torch.cuda.synchronize() | |
| t0 = time.perf_counter() | |
| for ids, keys in rows: | |
| fn(ids, len(keys)) | |
| torch.cuda.synchronize() | |
| out[name + "_ms"] = round(1000 * (time.perf_counter() - t0) / len(rows), 1) | |
| out["mean_prompt_tokens"] = round(sum(len(r[0]) for r in rows) / len(rows)) | |
| return out | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--checkpoint", type=Path, required=True) | |
| ap.add_argument("--train", type=Path, action="append", default=[]) | |
| ap.add_argument("--val", type=Path, action="append", default=[]) | |
| ap.add_argument("--test", type=Path, action="append", default=[]) | |
| ap.add_argument("--out", type=Path, required=True) | |
| ap.add_argument("--epochs", type=int, default=3) | |
| ap.add_argument("--lr", type=float, default=1e-5) | |
| ap.add_argument("--accum", type=int, default=16) | |
| ap.add_argument("--aux", type=float, default=0.01) | |
| ap.add_argument("--seed", type=int, default=7) | |
| ap.add_argument("--eval-base", action="store_true", help="also evaluate the untouched checkpoint on --test") | |
| ap.add_argument("--stage", default="decisions", help="training_stage written in the manifest identity") | |
| ap.add_argument("--save-epochs", action="store_true", | |
| help="also save every epoch's weights (fp32 state dict) to OUT/epochs/ for audit; changes nothing else") | |
| a = ap.parse_args() | |
| a.out.mkdir(parents=True, exist_ok=False) | |
| rng = random.Random(a.seed) | |
| torch.manual_seed(a.seed) | |
| from mini_v41.checkpoint import save_inference_checkpoint | |
| from mini_v41.engram import engram_disabled | |
| from mini_v41.inference import InferenceModel | |
| im = InferenceModel.from_checkpoint(a.checkpoint, device="cuda:0", dtype="fp32") | |
| if hasattr(im.model, "set_telemetry"): | |
| im.model.set_telemetry(True) # the fast-decode tree turns it off on load; aux_loss would read 0 | |
| reader = Reader(im) | |
| res = {"checkpoint": str(a.checkpoint), "train": [str(p) for p in a.train], "val": [str(p) for p in a.val], | |
| "test": [str(p) for p in a.test], "epochs": a.epochs, "lr": a.lr, "accum": a.accum, "aux": a.aux, | |
| "seed": a.seed} | |
| if a.eval_base: | |
| t0 = time.time() | |
| res["base_test"] = evaluate_files(reader, a.test) | |
| print("base (%.0f s) %s" % (time.time() - t0, json.dumps(res["base_test"], ensure_ascii=False)), flush=True) | |
| train_cases = [c for p in a.train for c in read_jsonl(p)] | |
| val_files = {str(p): read_jsonl(p) for p in a.val} | |
| model = im.model | |
| params = [p for p in model.parameters() if p.requires_grad] | |
| opt = torch.optim.AdamW(params, lr=a.lr, betas=(0.9, 0.95), weight_decay=0.0) | |
| n_q = sum(len(c["questions"]) for c in train_cases) | |
| steps = math.ceil(n_q * a.epochs / a.accum) | |
| warm = max(1, steps // 20) | |
| sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / warm) * max(0.0, 1 - s / steps)) | |
| print("train questions %d, optimizer steps %d" % (n_q, steps), flush=True) | |
| res["curve"], best, best_state, step = [], -1.0, None, 0 | |
| t0 = time.time() | |
| for ep in range(a.epochs): | |
| model.train() | |
| exs = examples(reader, train_cases, rng) | |
| tot = ok = n = 0 | |
| for i, (ids, k, target) in enumerate(exs): | |
| z, out = reader.logits_train(ids, k) | |
| tgt = torch.tensor(target, device=z.device) | |
| loss = -(tgt * F.log_softmax(z, -1)).sum() | |
| aux = getattr(out, "aux_loss", None) | |
| ((loss + (a.aux * aux if aux is not None else 0)) / a.accum).backward() | |
| tot += loss.item() | |
| ok += int(z.argmax()) == max(range(k), key=lambda j: target[j]) | |
| n += 1 | |
| if (i + 1) % a.accum == 0 or i + 1 == len(exs): | |
| torch.nn.utils.clip_grad_norm_(params, 1.0) | |
| opt.step() | |
| sched.step() | |
| opt.zero_grad(set_to_none=True) | |
| step += 1 | |
| if step % 100 == 0: | |
| print(" epoch %d step %d/%d loss %.3f acc %.3f %.0fs" % (ep + 1, step, steps, tot / n, ok / n, | |
| time.time() - t0), flush=True) | |
| vr = {Path(p).stem: evaluate(reader, cs) for p, cs in val_files.items()} | |
| score = sum(overall(r) for r in vr.values()) / max(1, len(vr)) | |
| res["curve"].append({"epoch": ep + 1, "loss": round(tot / n, 4), "train_acc": round(ok / n, 4), | |
| "val_score": round(score, 4), "val": vr, "s": round(time.time() - t0)}) | |
| print("epoch %d %s" % (ep + 1, json.dumps(res["curve"][-1], ensure_ascii=False)), flush=True) | |
| if a.save_epochs: | |
| (a.out / "epochs").mkdir(parents=True, exist_ok=True) | |
| torch.save({k: v.detach().to("cpu") for k, v in model.state_dict().items()}, a.out / "epochs" / ("epoch_%d.pt" % (ep + 1))) | |
| if score > best: | |
| best, res["selected_epoch"] = score, ep + 1 | |
| best_state = {k: v.detach().to("cpu", copy=True) for k, v in model.state_dict().items()} | |
| model.load_state_dict(best_state) | |
| model.eval() | |
| print("selected epoch %d (val %.4f)" % (res["selected_epoch"], best), flush=True) | |
| res["test"] = evaluate_files(reader, a.test, a.out, "") | |
| with engram_disabled(model) as n_off: | |
| res["test_engram_off"] = evaluate_files(reader, a.test) | |
| res["engram_modules_off"] = n_off | |
| if val_files: | |
| res["latency"] = latency(reader, next(iter(val_files.values()))) | |
| print("test %s" % json.dumps(res["test"], ensure_ascii=False), flush=True) | |
| print("test engram off %s" % json.dumps(res["test_engram_off"], ensure_ascii=False), flush=True) | |
| print("latency %s" % json.dumps(res.get("latency")), flush=True) | |
| # Same identity machinery as train_sft.py: stage, parent weights sha256, recipe hash, template. | |
| from mini_v41.chat import TEMPLATE_VERSION | |
| from mini_v41.run_identity import RunIdentity, file_sha256, recipe_hash | |
| parent_weights = next((p for p in (a.checkpoint / "model" / "model.pt", a.checkpoint / "model.pt") if p.exists()), None) | |
| recipe = {k: res[k] for k in ("train", "val", "epochs", "lr", "accum", "aux", "seed")} | |
| identity = RunIdentity.build( | |
| run_id=a.out.name, config=im.config, model=model, tokenizer=im.tokenizer, training_stage=a.stage, | |
| experiment_id="decisions", branch_id=a.out.name, parent_run_id=a.checkpoint.parent.name, | |
| parent_checkpoint_sha256=file_sha256(parent_weights) if parent_weights else "", | |
| training_recipe_hash=recipe_hash(recipe), template_version=TEMPLATE_VERSION, | |
| dataset_subset_id="+".join(Path(p).stem for p in a.train)) | |
| identity.validate() | |
| save_inference_checkpoint(a.out / "checkpoint", model=model, config=im.config, tokenizer=im.tokenizer, | |
| global_step=step, identity=identity, | |
| extra={"selected_epoch": res["selected_epoch"], "val_score": best}) | |
| (a.out / "results.json").write_text(json.dumps(res, indent=2, ensure_ascii=False) + "\n") | |
| if __name__ == "__main__": | |
| main() | |