Instructions to use vishwr/claim_drafter with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use vishwr/claim_drafter with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3.5-9B") model = PeftModel.from_pretrained(base_model, "vishwr/claim_drafter") - Notebooks
- Google Colab
- Kaggle
Refactor: replace budget/cost tooling with cost-free planning module
Browse files- Makefile +3 -9
- claim_drafter/{budget.py → planning.py} +15 -66
- claim_drafter/progress.py +12 -44
- scripts/estimate_cost.py +0 -79
- scripts/plot_runs.py +40 -60
- tests/test_budget.py +0 -134
- training/train_dpo.py +7 -10
- training/train_rl.py +6 -9
- training/train_sft.py +15 -19
Makefile
CHANGED
|
@@ -1,4 +1,4 @@
|
|
| 1 |
-
.PHONY: help setup test validate preflight data sft dpo rl
|
| 2 |
.DEFAULT_GOAL := help
|
| 3 |
|
| 4 |
PY ?= .venv/bin/python
|
|
@@ -18,17 +18,13 @@ setup: ## Create .venv (Python 3.12) and install dependencies
|
|
| 18 |
test: ## Run the offline test suite (no API key, no data rebuild needed)
|
| 19 |
$(PY) tests/test_rewards.py
|
| 20 |
$(PY) tests/test_dataset.py
|
| 21 |
-
$(PY) tests/test_budget.py
|
| 22 |
|
| 23 |
validate: ## Structural checks on the built dataset
|
| 24 |
$(PY) scripts/validate_dataset.py data/sft
|
| 25 |
|
| 26 |
-
preflight: ## Verify everything
|
| 27 |
$(PY) scripts/preflight.py
|
| 28 |
|
| 29 |
-
cost: ## Estimate the full SFT->DPO->RL pipeline against the budget
|
| 30 |
-
$(PY) scripts/estimate_cost.py --budget 120
|
| 31 |
-
|
| 32 |
calibrate: ## Re-check the reward function against real granted claims
|
| 33 |
$(PY) claim_drafter/rewards.py data/sft/train.jsonl 2000
|
| 34 |
|
|
@@ -79,11 +75,9 @@ rl: ## Stage 3: GRPO. CKPT defaults to the DPO run's final checkpoint
|
|
| 79 |
|
| 80 |
# ---------------------------------------------------------------- reporting
|
| 81 |
graphs: ## Build charts + summary from every stage's logs (run after `all`)
|
| 82 |
-
$(PY) scripts/plot_runs.py --runs runs --out runs/graphs
|
| 83 |
|
| 84 |
all: ## Run the whole pipeline end to end, then draw the graphs
|
| 85 |
-
@$(PY) scripts/estimate_cost.py --budget 120 \
|
| 86 |
-
|| (echo "Refusing to start: the plan is over budget." && exit 1)
|
| 87 |
$(MAKE) sft
|
| 88 |
$(MAKE) dpo
|
| 89 |
$(MAKE) rl
|
|
|
|
| 1 |
+
.PHONY: help setup test validate preflight data sft dpo rl graphs all clean
|
| 2 |
.DEFAULT_GOAL := help
|
| 3 |
|
| 4 |
PY ?= .venv/bin/python
|
|
|
|
| 18 |
test: ## Run the offline test suite (no API key, no data rebuild needed)
|
| 19 |
$(PY) tests/test_rewards.py
|
| 20 |
$(PY) tests/test_dataset.py
|
|
|
|
| 21 |
|
| 22 |
validate: ## Structural checks on the built dataset
|
| 23 |
$(PY) scripts/validate_dataset.py data/sft
|
| 24 |
|
| 25 |
+
preflight: ## Verify everything end to end, including the Tinker API
|
| 26 |
$(PY) scripts/preflight.py
|
| 27 |
|
|
|
|
|
|
|
|
|
|
| 28 |
calibrate: ## Re-check the reward function against real granted claims
|
| 29 |
$(PY) claim_drafter/rewards.py data/sft/train.jsonl 2000
|
| 30 |
|
|
|
|
| 75 |
|
| 76 |
# ---------------------------------------------------------------- reporting
|
| 77 |
graphs: ## Build charts + summary from every stage's logs (run after `all`)
|
| 78 |
+
$(PY) scripts/plot_runs.py --runs runs --out runs/graphs
|
| 79 |
|
| 80 |
all: ## Run the whole pipeline end to end, then draw the graphs
|
|
|
|
|
|
|
| 81 |
$(MAKE) sft
|
| 82 |
$(MAKE) dpo
|
| 83 |
$(MAKE) rl
|
claim_drafter/{budget.py → planning.py}
RENAMED
|
@@ -1,17 +1,15 @@
|
|
| 1 |
#!/usr/bin/env python3
|
| 2 |
-
"""One place that knows
|
| 3 |
|
| 4 |
Three callers need the same arithmetic and used to disagree about it:
|
| 5 |
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
claim_drafter/progress.py meters spend while the run is in flight
|
| 9 |
|
| 10 |
-
The eval
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
the cookbook, so the estimate and the run cannot drift apart again.
|
| 15 |
"""
|
| 16 |
|
| 17 |
import json
|
|
@@ -19,17 +17,8 @@ import os
|
|
| 19 |
|
| 20 |
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
| 21 |
|
| 22 |
-
#
|
| 23 |
-
#
|
| 24 |
-
PRICES = {
|
| 25 |
-
"Qwen/Qwen3.5-9B": {"train": 1.463, "prefill": 0.66, "sample": 1.995},
|
| 26 |
-
"Qwen/Qwen3.6-35B-A3B": {"train": 1.177, "prefill": 0.54, "sample": 1.335},
|
| 27 |
-
"Qwen/Qwen3.5-4B": {"train": 0.737, "prefill": 0.33, "sample": 1.005},
|
| 28 |
-
"Qwen/Qwen3-8B": {"train": 0.44, "prefill": 0.195, "sample": 0.60},
|
| 29 |
-
}
|
| 30 |
-
|
| 31 |
-
# Eval cadences, shared by the trainers and the estimator. Raising any of these
|
| 32 |
-
# is the cheapest way to buy budget back; see docs/pipeline.md.
|
| 33 |
CADENCE = {
|
| 34 |
"sft": {"eval_every": 150, "infrequent_eval_every": 600, "save_every": 300,
|
| 35 |
"nll_holdout": 64},
|
|
@@ -44,11 +33,6 @@ N_SLICES = 11
|
|
| 44 |
GENS_PER_SLICE = 8
|
| 45 |
|
| 46 |
|
| 47 |
-
def prices_for(model):
|
| 48 |
-
"""Price row for `model`, falling back to the 9B row for unknown models."""
|
| 49 |
-
return PRICES.get(model, PRICES["Qwen/Qwen3.5-9B"])
|
| 50 |
-
|
| 51 |
-
|
| 52 |
def sft_stats(manifest=None):
|
| 53 |
"""Token totals for the SFT training split."""
|
| 54 |
path = manifest or os.path.join(HERE, "data", "sft", "manifest.jsonl")
|
|
@@ -88,20 +72,13 @@ def dpo_stats(manifest=None):
|
|
| 88 |
return {"n": n, "total_tokens": total}
|
| 89 |
|
| 90 |
|
| 91 |
-
def eval_round_cost(n_gens, avg_prompt, avg_completion, prices):
|
| 92 |
-
"""Cost of one slice-evaluator round: prefill the prompts, sample the claims."""
|
| 93 |
-
return (n_gens * avg_prompt / 1e6 * prices["prefill"]
|
| 94 |
-
+ n_gens * avg_completion / 1e6 * prices["sample"])
|
| 95 |
-
|
| 96 |
-
|
| 97 |
def _rounds(total_steps, every):
|
| 98 |
"""Evaluator rounds over `total_steps`, counting the one at step 0.
|
| 99 |
|
| 100 |
The cookbook fires evaluators when `step % every == 0` with a zero-based
|
| 101 |
step, so step 0 always evaluates, and the trainers force a final round on
|
| 102 |
-
the last step. `total_steps // every` misses both
|
| 103 |
-
|
| 104 |
-
steps 0, 10 and 19, not the 2 this used to predict.
|
| 105 |
"""
|
| 106 |
if not every or total_steps <= 0:
|
| 107 |
return 0
|
|
@@ -114,16 +91,14 @@ def _rounds(total_steps, every):
|
|
| 114 |
def stage_plan(model="Qwen/Qwen3.5-9B", sft_epochs=2, sft_batch=8,
|
| 115 |
dpo_epochs=1, dpo_batch=4, rl_prompts=300, rl_group=8,
|
| 116 |
rl_batch=16, sft_manifest=None, dpo_manifest=None):
|
| 117 |
-
"""Steps, tokens and
|
| 118 |
|
| 119 |
-
Returns {stage: {...}}
|
| 120 |
-
|
| 121 |
"""
|
| 122 |
-
prices = prices_for(model)
|
| 123 |
sft = sft_stats(sft_manifest)
|
| 124 |
dpo = dpo_stats(dpo_manifest)
|
| 125 |
gens = N_SLICES * GENS_PER_SLICE
|
| 126 |
-
round_cost = eval_round_cost(gens, sft["avg_prompt"], sft["avg_completion"], prices)
|
| 127 |
|
| 128 |
plan = {}
|
| 129 |
|
|
@@ -135,18 +110,13 @@ def stage_plan(model="Qwen/Qwen3.5-9B", sft_epochs=2, sft_batch=8,
|
|
| 135 |
# ceiled count makes the progress bar unreachable and overstates tokens.
|
| 136 |
sft_steps = (sft_train_n // sft_batch) * sft_epochs
|
| 137 |
sft_tokens = int(sft["total_tokens"] * sft_epochs * sft_train_n / sft["n"])
|
| 138 |
-
nll_rounds = _rounds(sft_steps, c["eval_every"])
|
| 139 |
-
nll_tokens = nll_rounds * c["nll_holdout"] * (sft["avg_prompt"] + sft["avg_completion"])
|
| 140 |
slice_rounds = _rounds(sft_steps, c["infrequent_eval_every"])
|
| 141 |
plan["sft"] = {
|
| 142 |
"steps": sft_steps,
|
| 143 |
"train_tokens": sft_tokens,
|
| 144 |
-
"train_cost": sft_tokens / 1e6 * prices["train"],
|
| 145 |
"eval_rounds": slice_rounds,
|
| 146 |
"eval_gens": slice_rounds * gens,
|
| 147 |
-
"eval_cost": slice_rounds * round_cost + nll_tokens / 1e6 * prices["train"],
|
| 148 |
"tokens_per_step": sft_tokens // max(sft_steps, 1),
|
| 149 |
-
"eval_cost_per_round": round_cost,
|
| 150 |
**c,
|
| 151 |
}
|
| 152 |
|
|
@@ -154,23 +124,14 @@ def stage_plan(model="Qwen/Qwen3.5-9B", sft_epochs=2, sft_batch=8,
|
|
| 154 |
c = CADENCE["dpo"]
|
| 155 |
dpo_steps = (dpo["n"] // dpo_batch) * dpo_epochs if dpo["n"] else 0
|
| 156 |
dpo_tokens = dpo["total_tokens"] * dpo_epochs
|
| 157 |
-
# DPO scores each pair under BOTH the policy and the frozen reference, so the
|
| 158 |
-
# same tokens go through a second forward pass. Omitting it understated
|
| 159 |
-
# stage 2 by roughly 40%.
|
| 160 |
-
dpo_reference_tokens = dpo_tokens
|
| 161 |
slice_rounds = _rounds(dpo_steps, c["infrequent_eval_every"])
|
| 162 |
plan["dpo"] = {
|
| 163 |
"steps": dpo_steps,
|
| 164 |
"pairs": dpo["n"],
|
| 165 |
"train_tokens": dpo_tokens,
|
| 166 |
-
"reference_tokens": dpo_reference_tokens,
|
| 167 |
-
"train_cost": (dpo_tokens / 1e6 * prices["train"]
|
| 168 |
-
+ dpo_reference_tokens / 1e6 * prices["prefill"]),
|
| 169 |
"eval_rounds": slice_rounds,
|
| 170 |
"eval_gens": slice_rounds * gens,
|
| 171 |
-
"eval_cost": slice_rounds * round_cost,
|
| 172 |
"tokens_per_step": dpo_tokens // max(dpo_steps, 1),
|
| 173 |
-
"eval_cost_per_round": round_cost,
|
| 174 |
**c,
|
| 175 |
}
|
| 176 |
|
|
@@ -179,34 +140,22 @@ def stage_plan(model="Qwen/Qwen3.5-9B", sft_epochs=2, sft_batch=8,
|
|
| 179 |
c = CADENCE["rl"]
|
| 180 |
rl_steps = rl_prompts // rl_batch
|
| 181 |
rollouts = rl_prompts * rl_group
|
| 182 |
-
|
| 183 |
-
rl_sample = rollouts * sft["avg_completion"]
|
| 184 |
rl_train = rollouts * (sft["avg_prompt"] + sft["avg_completion"])
|
| 185 |
slice_rounds = _rounds(rl_steps, c["eval_every"]) or 1
|
| 186 |
plan["rl"] = {
|
| 187 |
"steps": rl_steps,
|
| 188 |
"rollouts": rollouts,
|
| 189 |
-
# rl_train already covers prompt+completion for every rollout;
|
| 190 |
-
# adding rl_sample counted the completions twice.
|
| 191 |
"train_tokens": rl_train,
|
| 192 |
-
"train_cost": (rl_prefill / 1e6 * prices["prefill"]
|
| 193 |
-
+ rl_sample / 1e6 * prices["sample"]
|
| 194 |
-
+ rl_train / 1e6 * prices["train"]),
|
| 195 |
"eval_rounds": slice_rounds,
|
| 196 |
"eval_gens": slice_rounds * gens,
|
| 197 |
-
"eval_cost": slice_rounds * round_cost,
|
| 198 |
"tokens_per_step": rl_train // max(rl_steps, 1),
|
| 199 |
-
"eval_cost_per_round": round_cost,
|
| 200 |
**c,
|
| 201 |
}
|
| 202 |
|
| 203 |
plan["total"] = {
|
| 204 |
-
"train_cost": sum(plan[s]["train_cost"] for s in ("sft", "dpo", "rl")),
|
| 205 |
-
"eval_cost": sum(plan[s]["eval_cost"] for s in ("sft", "dpo", "rl")),
|
| 206 |
"steps": sum(plan[s]["steps"] for s in ("sft", "dpo", "rl")),
|
| 207 |
"eval_gens": sum(plan[s]["eval_gens"] for s in ("sft", "dpo", "rl")),
|
| 208 |
}
|
| 209 |
-
plan["total"]["cost"] = plan["total"]["train_cost"] + plan["total"]["eval_cost"]
|
| 210 |
-
plan["prices"] = prices
|
| 211 |
plan["model"] = model
|
| 212 |
return plan
|
|
|
|
| 1 |
#!/usr/bin/env python3
|
| 2 |
+
"""One place that knows the shape of each training stage: steps, tokens, cadence.
|
| 3 |
|
| 4 |
Three callers need the same arithmetic and used to disagree about it:
|
| 5 |
|
| 6 |
+
training/train_*.py size the eval cadence and progress bar for the run
|
| 7 |
+
claim_drafter/progress.py meters tokens/throughput while the run is in flight
|
|
|
|
| 8 |
|
| 9 |
+
The eval volume is the number that used to be wrong. An early configuration ran
|
| 10 |
+
the expensive slice evaluator on every step; `stage_plan` derives eval volume
|
| 11 |
+
from the *same* cadence constants the trainers pass to the cookbook, so the
|
| 12 |
+
progress display and the run cannot drift apart.
|
|
|
|
| 13 |
"""
|
| 14 |
|
| 15 |
import json
|
|
|
|
| 17 |
|
| 18 |
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
| 19 |
|
| 20 |
+
# Eval cadences, shared by the trainers and the progress meter. Raising any of
|
| 21 |
+
# these is the cheapest way to shorten a run; see the pipeline notes.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 22 |
CADENCE = {
|
| 23 |
"sft": {"eval_every": 150, "infrequent_eval_every": 600, "save_every": 300,
|
| 24 |
"nll_holdout": 64},
|
|
|
|
| 33 |
GENS_PER_SLICE = 8
|
| 34 |
|
| 35 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
def sft_stats(manifest=None):
|
| 37 |
"""Token totals for the SFT training split."""
|
| 38 |
path = manifest or os.path.join(HERE, "data", "sft", "manifest.jsonl")
|
|
|
|
| 72 |
return {"n": n, "total_tokens": total}
|
| 73 |
|
| 74 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 75 |
def _rounds(total_steps, every):
|
| 76 |
"""Evaluator rounds over `total_steps`, counting the one at step 0.
|
| 77 |
|
| 78 |
The cookbook fires evaluators when `step % every == 0` with a zero-based
|
| 79 |
step, so step 0 always evaluates, and the trainers force a final round on
|
| 80 |
+
the last step. `total_steps // every` misses both -- verified against a real
|
| 81 |
+
run: 20 steps at every=10 produced rounds at steps 0, 10 and 19.
|
|
|
|
| 82 |
"""
|
| 83 |
if not every or total_steps <= 0:
|
| 84 |
return 0
|
|
|
|
| 91 |
def stage_plan(model="Qwen/Qwen3.5-9B", sft_epochs=2, sft_batch=8,
|
| 92 |
dpo_epochs=1, dpo_batch=4, rl_prompts=300, rl_group=8,
|
| 93 |
rl_batch=16, sft_manifest=None, dpo_manifest=None):
|
| 94 |
+
"""Steps, tokens and eval volume for every stage, using the shipped cadences.
|
| 95 |
|
| 96 |
+
Returns {stage: {...}}. Everything downstream -- the trainers' ETA banner and
|
| 97 |
+
the progress meter -- reads this, so the plan and the run stay in sync.
|
| 98 |
"""
|
|
|
|
| 99 |
sft = sft_stats(sft_manifest)
|
| 100 |
dpo = dpo_stats(dpo_manifest)
|
| 101 |
gens = N_SLICES * GENS_PER_SLICE
|
|
|
|
| 102 |
|
| 103 |
plan = {}
|
| 104 |
|
|
|
|
| 110 |
# ceiled count makes the progress bar unreachable and overstates tokens.
|
| 111 |
sft_steps = (sft_train_n // sft_batch) * sft_epochs
|
| 112 |
sft_tokens = int(sft["total_tokens"] * sft_epochs * sft_train_n / sft["n"])
|
|
|
|
|
|
|
| 113 |
slice_rounds = _rounds(sft_steps, c["infrequent_eval_every"])
|
| 114 |
plan["sft"] = {
|
| 115 |
"steps": sft_steps,
|
| 116 |
"train_tokens": sft_tokens,
|
|
|
|
| 117 |
"eval_rounds": slice_rounds,
|
| 118 |
"eval_gens": slice_rounds * gens,
|
|
|
|
| 119 |
"tokens_per_step": sft_tokens // max(sft_steps, 1),
|
|
|
|
| 120 |
**c,
|
| 121 |
}
|
| 122 |
|
|
|
|
| 124 |
c = CADENCE["dpo"]
|
| 125 |
dpo_steps = (dpo["n"] // dpo_batch) * dpo_epochs if dpo["n"] else 0
|
| 126 |
dpo_tokens = dpo["total_tokens"] * dpo_epochs
|
|
|
|
|
|
|
|
|
|
|
|
|
| 127 |
slice_rounds = _rounds(dpo_steps, c["infrequent_eval_every"])
|
| 128 |
plan["dpo"] = {
|
| 129 |
"steps": dpo_steps,
|
| 130 |
"pairs": dpo["n"],
|
| 131 |
"train_tokens": dpo_tokens,
|
|
|
|
|
|
|
|
|
|
| 132 |
"eval_rounds": slice_rounds,
|
| 133 |
"eval_gens": slice_rounds * gens,
|
|
|
|
| 134 |
"tokens_per_step": dpo_tokens // max(dpo_steps, 1),
|
|
|
|
| 135 |
**c,
|
| 136 |
}
|
| 137 |
|
|
|
|
| 140 |
c = CADENCE["rl"]
|
| 141 |
rl_steps = rl_prompts // rl_batch
|
| 142 |
rollouts = rl_prompts * rl_group
|
| 143 |
+
# One rollout covers prompt+completion; that is the sequence trained on.
|
|
|
|
| 144 |
rl_train = rollouts * (sft["avg_prompt"] + sft["avg_completion"])
|
| 145 |
slice_rounds = _rounds(rl_steps, c["eval_every"]) or 1
|
| 146 |
plan["rl"] = {
|
| 147 |
"steps": rl_steps,
|
| 148 |
"rollouts": rollouts,
|
|
|
|
|
|
|
| 149 |
"train_tokens": rl_train,
|
|
|
|
|
|
|
|
|
|
| 150 |
"eval_rounds": slice_rounds,
|
| 151 |
"eval_gens": slice_rounds * gens,
|
|
|
|
| 152 |
"tokens_per_step": rl_train // max(rl_steps, 1),
|
|
|
|
| 153 |
**c,
|
| 154 |
}
|
| 155 |
|
| 156 |
plan["total"] = {
|
|
|
|
|
|
|
| 157 |
"steps": sum(plan[s]["steps"] for s in ("sft", "dpo", "rl")),
|
| 158 |
"eval_gens": sum(plan[s]["eval_gens"] for s in ("sft", "dpo", "rl")),
|
| 159 |
}
|
|
|
|
|
|
|
| 160 |
plan["model"] = model
|
| 161 |
return plan
|
claim_drafter/progress.py
CHANGED
|
@@ -1,5 +1,5 @@
|
|
| 1 |
#!/usr/bin/env python3
|
| 2 |
-
"""Live console progress, ETA and
|
| 3 |
|
| 4 |
The cookbook's own console output is a Rich table printed on *every* step --
|
| 5 |
2,416 tables for stage 1 alone, which scrolls the useful information away.
|
|
@@ -11,8 +11,8 @@ to build its MultiplexLogger. We drop PrettyPrintLogger from the returned
|
|
| 11 |
multiplex and append a StageProgress in its place.
|
| 12 |
|
| 13 |
Side effect worth knowing: this writes `runs/<stage>/progress.jsonl`, one record
|
| 14 |
-
per step, carrying wall-clock
|
| 15 |
-
|
| 16 |
|
| 17 |
No new dependencies -- the bar is ~40 lines of ANSI rather than tqdm/rich.
|
| 18 |
"""
|
|
@@ -57,24 +57,15 @@ class StageProgress:
|
|
| 57 |
tests can exercise the formatting.
|
| 58 |
"""
|
| 59 |
|
| 60 |
-
def __init__(self, stage, total_steps, log_dir,
|
| 61 |
tokens_per_step=0, eval_every=0, save_every=0,
|
| 62 |
-
|
| 63 |
-
stream=None, cost_per_step=None):
|
| 64 |
self.stage = stage
|
| 65 |
self.total_steps = max(int(total_steps or 0), 0)
|
| 66 |
self.log_dir = log_dir
|
| 67 |
-
self.price_train = price_train
|
| 68 |
self.tokens_per_step = tokens_per_step
|
| 69 |
-
# DPO bills each pair twice -- once through the policy at the train
|
| 70 |
-
# rate, once through the frozen reference at the prefill rate -- so
|
| 71 |
-
# tokens x train-price under-reports it. The plan knows the real
|
| 72 |
-
# number; use it when given.
|
| 73 |
-
self.cost_per_step = (cost_per_step if cost_per_step is not None
|
| 74 |
-
else tokens_per_step / 1e6 * price_train)
|
| 75 |
self.eval_every = eval_every
|
| 76 |
self.save_every = save_every
|
| 77 |
-
self.eval_cost_per_round = eval_cost_per_round
|
| 78 |
self.eval_seconds_per_round = eval_seconds_per_round
|
| 79 |
self.stream = stream or sys.stderr
|
| 80 |
self.tty = hasattr(self.stream, "isatty") and self.stream.isatty()
|
|
@@ -84,8 +75,6 @@ class StageProgress:
|
|
| 84 |
self.sec_per_step = None # EMA, seconds
|
| 85 |
self.steps_seen = 0
|
| 86 |
self.evals_seen = 0
|
| 87 |
-
self.train_cost = 0.0
|
| 88 |
-
self.eval_cost = 0.0
|
| 89 |
self._fh = None
|
| 90 |
if log_dir:
|
| 91 |
os.makedirs(log_dir, exist_ok=True)
|
|
@@ -139,10 +128,8 @@ class StageProgress:
|
|
| 139 |
self.last_step_at = now
|
| 140 |
self.steps_seen += 1
|
| 141 |
|
| 142 |
-
self.train_cost += self.cost_per_step
|
| 143 |
if is_eval:
|
| 144 |
self.evals_seen += 1
|
| 145 |
-
self.eval_cost += self.eval_cost_per_round
|
| 146 |
|
| 147 |
elapsed = now - self.start
|
| 148 |
remaining = self._eta(step)
|
|
@@ -163,8 +150,6 @@ class StageProgress:
|
|
| 163 |
"steps %d" % self.steps_seen,
|
| 164 |
"wall clock %s" % _fmt_duration(total),
|
| 165 |
"eval rounds %d" % self.evals_seen,
|
| 166 |
-
"estimated cost $%.2f (train $%.2f + eval $%.2f)"
|
| 167 |
-
% (self.train_cost + self.eval_cost, self.train_cost, self.eval_cost),
|
| 168 |
"logs %s" % self.log_dir,
|
| 169 |
])
|
| 170 |
if self._fh:
|
|
@@ -207,7 +192,6 @@ class StageProgress:
|
|
| 207 |
if self.sec_per_step:
|
| 208 |
parts.append("%.1fs/it" % self.sec_per_step)
|
| 209 |
parts.append("%s<%s" % (_fmt_duration(elapsed), _fmt_duration(remaining)))
|
| 210 |
-
parts.append("$%.2f" % (self.train_cost + self.eval_cost))
|
| 211 |
|
| 212 |
label, away = self._next_event(step)
|
| 213 |
if label and self.sec_per_step:
|
|
@@ -250,9 +234,6 @@ class StageProgress:
|
|
| 250 |
"eta_s": round(remaining, 1) if remaining is not None else None,
|
| 251 |
"sec_per_step": round(self.sec_per_step, 3) if self.sec_per_step else None,
|
| 252 |
"tokens_cum": self.steps_seen * self.tokens_per_step,
|
| 253 |
-
"cost_train_usd": round(self.train_cost, 4),
|
| 254 |
-
"cost_eval_usd": round(self.eval_cost, 4),
|
| 255 |
-
"cost_total_usd": round(self.train_cost + self.eval_cost, 4),
|
| 256 |
"is_eval": is_eval,
|
| 257 |
}
|
| 258 |
if loss is not None:
|
|
@@ -275,25 +256,15 @@ class StageProgress:
|
|
| 275 |
self.stream.flush()
|
| 276 |
|
| 277 |
|
| 278 |
-
def announce(stage, total_steps, tokens_total,
|
| 279 |
-
eval_gens_per_round,
|
| 280 |
-
"""Print the pre-flight
|
| 281 |
-
|
| 282 |
-
`train_cost` comes from the plan when the caller has it. Recomputing it from
|
| 283 |
-
tokens is only right when every token is billed at the train rate, which is
|
| 284 |
-
not true for DPO: each pair also goes through the frozen reference model at
|
| 285 |
-
the prefill rate, so recomputing understated stage 2 by about 45%.
|
| 286 |
-
"""
|
| 287 |
stream = stream or sys.stderr
|
| 288 |
-
if train_cost is None:
|
| 289 |
-
train_cost = tokens_total / 1e6 * price_train
|
| 290 |
lines = [
|
| 291 |
"STAGE %s -- starting" % stage.upper(),
|
| 292 |
"steps %s" % "{:,}".format(total_steps),
|
| 293 |
"train tokens %s" % "{:,}".format(int(tokens_total)),
|
| 294 |
"eval %d rounds x %d generations" % (eval_rounds, eval_gens_per_round),
|
| 295 |
-
"estimated cost $%.2f (train $%.2f + eval $%.2f)"
|
| 296 |
-
% (train_cost + eval_cost, train_cost, eval_cost),
|
| 297 |
"duration unknown until the first steps land; the bar below",
|
| 298 |
" shows a live ETA once it has measured s/it",
|
| 299 |
]
|
|
@@ -305,10 +276,9 @@ def announce(stage, total_steps, tokens_total, price_train, eval_rounds,
|
|
| 305 |
stream.flush()
|
| 306 |
|
| 307 |
|
| 308 |
-
def install(stage, total_steps, log_dir,
|
| 309 |
-
eval_every=0, save_every=0,
|
| 310 |
-
eval_seconds_per_round=None, replace_pretty=True
|
| 311 |
-
cost_per_step=None):
|
| 312 |
"""Attach a StageProgress to whatever logger the cookbook builds next.
|
| 313 |
|
| 314 |
Must be called before the trainer's `main()`. Returns the StageProgress so
|
|
@@ -319,11 +289,9 @@ def install(stage, total_steps, log_dir, price_train, tokens_per_step=0,
|
|
| 319 |
|
| 320 |
progress = StageProgress(
|
| 321 |
stage=stage, total_steps=total_steps, log_dir=log_dir,
|
| 322 |
-
|
| 323 |
eval_every=eval_every, save_every=save_every,
|
| 324 |
-
eval_cost_per_round=eval_cost_per_round,
|
| 325 |
eval_seconds_per_round=eval_seconds_per_round,
|
| 326 |
-
cost_per_step=cost_per_step,
|
| 327 |
)
|
| 328 |
|
| 329 |
original = ml_log.setup_logging
|
|
|
|
| 1 |
#!/usr/bin/env python3
|
| 2 |
+
"""Live console progress, ETA and throughput for the three training stages.
|
| 3 |
|
| 4 |
The cookbook's own console output is a Rich table printed on *every* step --
|
| 5 |
2,416 tables for stage 1 alone, which scrolls the useful information away.
|
|
|
|
| 11 |
multiplex and append a StageProgress in its place.
|
| 12 |
|
| 13 |
Side effect worth knowing: this writes `runs/<stage>/progress.jsonl`, one record
|
| 14 |
+
per step, carrying wall-clock and throughput. metrics.jsonl has the learning
|
| 15 |
+
curves but no timing, so scripts/plot_runs.py reads both.
|
| 16 |
|
| 17 |
No new dependencies -- the bar is ~40 lines of ANSI rather than tqdm/rich.
|
| 18 |
"""
|
|
|
|
| 57 |
tests can exercise the formatting.
|
| 58 |
"""
|
| 59 |
|
| 60 |
+
def __init__(self, stage, total_steps, log_dir,
|
| 61 |
tokens_per_step=0, eval_every=0, save_every=0,
|
| 62 |
+
eval_seconds_per_round=None, stream=None):
|
|
|
|
| 63 |
self.stage = stage
|
| 64 |
self.total_steps = max(int(total_steps or 0), 0)
|
| 65 |
self.log_dir = log_dir
|
|
|
|
| 66 |
self.tokens_per_step = tokens_per_step
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 67 |
self.eval_every = eval_every
|
| 68 |
self.save_every = save_every
|
|
|
|
| 69 |
self.eval_seconds_per_round = eval_seconds_per_round
|
| 70 |
self.stream = stream or sys.stderr
|
| 71 |
self.tty = hasattr(self.stream, "isatty") and self.stream.isatty()
|
|
|
|
| 75 |
self.sec_per_step = None # EMA, seconds
|
| 76 |
self.steps_seen = 0
|
| 77 |
self.evals_seen = 0
|
|
|
|
|
|
|
| 78 |
self._fh = None
|
| 79 |
if log_dir:
|
| 80 |
os.makedirs(log_dir, exist_ok=True)
|
|
|
|
| 128 |
self.last_step_at = now
|
| 129 |
self.steps_seen += 1
|
| 130 |
|
|
|
|
| 131 |
if is_eval:
|
| 132 |
self.evals_seen += 1
|
|
|
|
| 133 |
|
| 134 |
elapsed = now - self.start
|
| 135 |
remaining = self._eta(step)
|
|
|
|
| 150 |
"steps %d" % self.steps_seen,
|
| 151 |
"wall clock %s" % _fmt_duration(total),
|
| 152 |
"eval rounds %d" % self.evals_seen,
|
|
|
|
|
|
|
| 153 |
"logs %s" % self.log_dir,
|
| 154 |
])
|
| 155 |
if self._fh:
|
|
|
|
| 192 |
if self.sec_per_step:
|
| 193 |
parts.append("%.1fs/it" % self.sec_per_step)
|
| 194 |
parts.append("%s<%s" % (_fmt_duration(elapsed), _fmt_duration(remaining)))
|
|
|
|
| 195 |
|
| 196 |
label, away = self._next_event(step)
|
| 197 |
if label and self.sec_per_step:
|
|
|
|
| 234 |
"eta_s": round(remaining, 1) if remaining is not None else None,
|
| 235 |
"sec_per_step": round(self.sec_per_step, 3) if self.sec_per_step else None,
|
| 236 |
"tokens_cum": self.steps_seen * self.tokens_per_step,
|
|
|
|
|
|
|
|
|
|
| 237 |
"is_eval": is_eval,
|
| 238 |
}
|
| 239 |
if loss is not None:
|
|
|
|
| 256 |
self.stream.flush()
|
| 257 |
|
| 258 |
|
| 259 |
+
def announce(stage, total_steps, tokens_total, eval_rounds,
|
| 260 |
+
eval_gens_per_round, stream=None):
|
| 261 |
+
"""Print the pre-flight summary for a stage before the first step runs."""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 262 |
stream = stream or sys.stderr
|
|
|
|
|
|
|
| 263 |
lines = [
|
| 264 |
"STAGE %s -- starting" % stage.upper(),
|
| 265 |
"steps %s" % "{:,}".format(total_steps),
|
| 266 |
"train tokens %s" % "{:,}".format(int(tokens_total)),
|
| 267 |
"eval %d rounds x %d generations" % (eval_rounds, eval_gens_per_round),
|
|
|
|
|
|
|
| 268 |
"duration unknown until the first steps land; the bar below",
|
| 269 |
" shows a live ETA once it has measured s/it",
|
| 270 |
]
|
|
|
|
| 276 |
stream.flush()
|
| 277 |
|
| 278 |
|
| 279 |
+
def install(stage, total_steps, log_dir, tokens_per_step=0,
|
| 280 |
+
eval_every=0, save_every=0,
|
| 281 |
+
eval_seconds_per_round=None, replace_pretty=True):
|
|
|
|
| 282 |
"""Attach a StageProgress to whatever logger the cookbook builds next.
|
| 283 |
|
| 284 |
Must be called before the trainer's `main()`. Returns the StageProgress so
|
|
|
|
| 289 |
|
| 290 |
progress = StageProgress(
|
| 291 |
stage=stage, total_steps=total_steps, log_dir=log_dir,
|
| 292 |
+
tokens_per_step=tokens_per_step,
|
| 293 |
eval_every=eval_every, save_every=save_every,
|
|
|
|
| 294 |
eval_seconds_per_round=eval_seconds_per_round,
|
|
|
|
| 295 |
)
|
| 296 |
|
| 297 |
original = ml_log.setup_logging
|
scripts/estimate_cost.py
DELETED
|
@@ -1,79 +0,0 @@
|
|
| 1 |
-
#!/usr/bin/env python3
|
| 2 |
-
"""Cost the full SFT -> DPO -> RL pipeline against a fixed budget.
|
| 3 |
-
|
| 4 |
-
python3 scripts/estimate_cost.py [--budget 120] [--model Qwen/Qwen3.5-9B]
|
| 5 |
-
|
| 6 |
-
All arithmetic lives in claim_drafter/budget.py, which the trainers also import
|
| 7 |
-
for their eval cadences -- so what this prints is what the run will actually do.
|
| 8 |
-
Before that refactor this script hard-coded an eval assumption (5 rounds of 150
|
| 9 |
-
generations) that the trainers did not share (96 rounds of 440), and understated
|
| 10 |
-
the total by about half the budget.
|
| 11 |
-
|
| 12 |
-
Caveat carried from the original: Tinker documents the "train" meter as "forward
|
| 13 |
-
and backward pass for gradient computation" but does not state whether it bills
|
| 14 |
-
every token or only weight=1 tokens. Everything below assumes the full sequence,
|
| 15 |
-
which is the conservative reading.
|
| 16 |
-
"""
|
| 17 |
-
|
| 18 |
-
import argparse
|
| 19 |
-
import os
|
| 20 |
-
import sys
|
| 21 |
-
|
| 22 |
-
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
| 23 |
-
from claim_drafter import budget as B
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
def main():
|
| 27 |
-
ap = argparse.ArgumentParser()
|
| 28 |
-
ap.add_argument("--budget", type=float, default=120.0)
|
| 29 |
-
ap.add_argument("--model", default="Qwen/Qwen3.5-9B")
|
| 30 |
-
ap.add_argument("--sft-epochs", type=int, default=2)
|
| 31 |
-
ap.add_argument("--dpo-epochs", type=int, default=1)
|
| 32 |
-
ap.add_argument("--rl-prompts", type=int, default=300)
|
| 33 |
-
ap.add_argument("--rl-group", type=int, default=8, help="GRPO samples per prompt")
|
| 34 |
-
args = ap.parse_args()
|
| 35 |
-
|
| 36 |
-
plan = B.stage_plan(model=args.model, sft_epochs=args.sft_epochs,
|
| 37 |
-
dpo_epochs=args.dpo_epochs, rl_prompts=args.rl_prompts,
|
| 38 |
-
rl_group=args.rl_group)
|
| 39 |
-
p = plan["prices"]
|
| 40 |
-
|
| 41 |
-
print("Model: %s (train $%.3f / prefill $%.3f / sample $%.3f per 1M)"
|
| 42 |
-
% (args.model, p["train"], p["prefill"], p["sample"]))
|
| 43 |
-
print("\n%-34s %8s %13s %8s %8s" % ("stage", "steps", "train tokens", "train $", "eval $"))
|
| 44 |
-
print("-" * 76)
|
| 45 |
-
|
| 46 |
-
labels = {
|
| 47 |
-
"sft": "1. SFT (%d epochs)" % args.sft_epochs,
|
| 48 |
-
"dpo": "2. DPO (%d epoch, %s pairs)" % (args.dpo_epochs,
|
| 49 |
-
"{:,}".format(plan["dpo"]["pairs"])),
|
| 50 |
-
"rl": "3. RL (%d prompts x %d samples)" % (args.rl_prompts, args.rl_group),
|
| 51 |
-
}
|
| 52 |
-
for stage in ("sft", "dpo", "rl"):
|
| 53 |
-
r = plan[stage]
|
| 54 |
-
print("%-34s %8s %13s %8.2f %8.2f"
|
| 55 |
-
% (labels[stage], "{:,}".format(r["steps"]),
|
| 56 |
-
"{:,}".format(r["train_tokens"]), r["train_cost"], r["eval_cost"]))
|
| 57 |
-
|
| 58 |
-
t = plan["total"]
|
| 59 |
-
print("-" * 76)
|
| 60 |
-
print("%-34s %8s %13s %8.2f %8.2f"
|
| 61 |
-
% ("TOTAL", "{:,}".format(t["steps"]), "", t["train_cost"], t["eval_cost"]))
|
| 62 |
-
print("%-34s %39s %8.2f" % ("", "combined", t["cost"]))
|
| 63 |
-
|
| 64 |
-
print("\nEval volume: %s generations across %d rounds"
|
| 65 |
-
% ("{:,}".format(t["eval_gens"]),
|
| 66 |
-
sum(plan[s]["eval_rounds"] for s in ("sft", "dpo", "rl"))))
|
| 67 |
-
print("Budget $%.2f -> %s"
|
| 68 |
-
% (args.budget,
|
| 69 |
-
("fits, $%.2f left" % (args.budget - t["cost"])) if t["cost"] <= args.budget
|
| 70 |
-
else ("OVER by $%.2f" % (t["cost"] - args.budget))))
|
| 71 |
-
print("Stage share: SFT %.0f%% DPO %.0f%% RL %.0f%%"
|
| 72 |
-
% tuple(100.0 * (plan[s]["train_cost"] + plan[s]["eval_cost"]) / t["cost"]
|
| 73 |
-
for s in ("sft", "dpo", "rl")))
|
| 74 |
-
|
| 75 |
-
return 0 if t["cost"] <= args.budget else 1
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
if __name__ == "__main__":
|
| 79 |
-
sys.exit(main())
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
scripts/plot_runs.py
CHANGED
|
@@ -1,12 +1,12 @@
|
|
| 1 |
#!/usr/bin/env python3
|
| 2 |
"""Turn the logs from every stage into charts, once all three have finished.
|
| 3 |
|
| 4 |
-
python3 scripts/plot_runs.py [--runs runs] [--out runs/graphs]
|
| 5 |
|
| 6 |
Reads two files per stage:
|
| 7 |
|
| 8 |
runs/<stage>/metrics.jsonl the cookbook's learning curves
|
| 9 |
-
runs/<stage>/progress.jsonl wall-clock
|
| 10 |
|
| 11 |
Stages that have not run yet are skipped with a note rather than an error, so
|
| 12 |
this is safe to run mid-pipeline -- you just get fewer panels.
|
|
@@ -195,14 +195,10 @@ def plot_slice_breakdown(plt, data, out):
|
|
| 195 |
return save(fig, out, "03_slice_breakdown.png")
|
| 196 |
|
| 197 |
|
| 198 |
-
def
|
| 199 |
-
fig, axes = plt.subplots(1,
|
| 200 |
-
ax_rate,
|
| 201 |
|
| 202 |
-
# Spend is cumulative across the whole pipeline, so each stage continues
|
| 203 |
-
# where the previous one stopped -- on both axes. Restarting x at 0 per
|
| 204 |
-
# stage made a single long run look like three short ones.
|
| 205 |
-
cost_base = x_base = 0.0
|
| 206 |
for stage in STAGES:
|
| 207 |
if stage not in data:
|
| 208 |
continue
|
|
@@ -210,31 +206,20 @@ def plot_timing_and_cost(plt, data, out, budget_usd):
|
|
| 210 |
xs, ys = series(p, "sec_per_step")
|
| 211 |
if ys:
|
| 212 |
ax_rate.plot(xs, ys, lw=1, label=stage)
|
| 213 |
-
cxs, cys = series(p, "cost_total_usd")
|
| 214 |
-
if cys:
|
| 215 |
-
ax_cost.plot([x_base + i for i in range(len(cys))],
|
| 216 |
-
[cost_base + y for y in cys], lw=1.5, label=stage)
|
| 217 |
-
cost_base += cys[-1]
|
| 218 |
-
x_base += len(cys)
|
| 219 |
exs, eys = series(p, "elapsed_s")
|
| 220 |
if eys:
|
| 221 |
ax_wall.plot(exs, [y / 60.0 for y in eys], lw=1.2, label=stage)
|
| 222 |
-
if budget_usd:
|
| 223 |
-
ax_cost.axhline(budget_usd, ls="--", c="crimson", lw=1,
|
| 224 |
-
label="budget $%.0f" % budget_usd)
|
| 225 |
|
| 226 |
ax_rate.set_title("Seconds per step")
|
| 227 |
ax_rate.set_xlabel("step")
|
| 228 |
-
ax_cost.set_title("Cumulative estimated spend ($)")
|
| 229 |
-
ax_cost.set_xlabel("logged step (all stages)")
|
| 230 |
ax_wall.set_title("Wall clock (minutes)")
|
| 231 |
ax_wall.set_xlabel("step")
|
| 232 |
for ax in axes:
|
| 233 |
ax.grid(alpha=0.3)
|
| 234 |
if ax.get_legend_handles_labels()[0]:
|
| 235 |
ax.legend(fontsize=8)
|
| 236 |
-
fig.suptitle("Throughput
|
| 237 |
-
return save(fig, out, "
|
| 238 |
|
| 239 |
|
| 240 |
def plot_slice_trajectories(plt, data, out, metric="parse_rate"):
|
|
@@ -274,10 +259,11 @@ def plot_slice_trajectories(plt, data, out, metric="parse_rate"):
|
|
| 274 |
|
| 275 |
|
| 276 |
def plot_convergence(plt, data, out):
|
| 277 |
-
"""Where the held-out loss stopped improving, and
|
| 278 |
|
| 279 |
This is the panel that answers 'should I have trained this long', which the
|
| 280 |
-
aggregate learning curve does not.
|
|
|
|
| 281 |
"""
|
| 282 |
if "sft" not in data:
|
| 283 |
return None
|
|
@@ -294,13 +280,13 @@ def plot_convergence(plt, data, out):
|
|
| 294 |
conv_step = xs[conv_i]
|
| 295 |
|
| 296 |
prog = data["sft"]["progress"]
|
| 297 |
-
|
| 298 |
-
|
| 299 |
-
|
| 300 |
-
|
| 301 |
-
for sx, sy in zip(
|
| 302 |
if sx <= conv_step:
|
| 303 |
-
|
| 304 |
|
| 305 |
fig, axes = plt.subplots(1, 2, figsize=(12.5, 4.2))
|
| 306 |
|
|
@@ -319,24 +305,24 @@ def plot_convergence(plt, data, out):
|
|
| 319 |
ax.grid(alpha=0.25)
|
| 320 |
|
| 321 |
ax = axes[1]
|
| 322 |
-
if
|
| 323 |
-
ax.plot(
|
| 324 |
ax.axvline(conv_step, color=C_MID, linewidth=2, linestyle="--")
|
| 325 |
-
wasted = max(0.0,
|
| 326 |
-
ax.fill_between([x for x in
|
| 327 |
-
[
|
| 328 |
-
[y for x, y in zip(
|
| 329 |
color=C_WARN, alpha=0.18)
|
| 330 |
-
ax.annotate("
|
| 331 |
xy=(0.42, 0.18), xycoords="axes fraction",
|
| 332 |
fontsize=10, color=C_WARN)
|
| 333 |
-
ax.set_title("Cumulative
|
| 334 |
ax.set_xlabel("training step", color=C_INK_2)
|
| 335 |
-
ax.set_ylabel("
|
| 336 |
ax.grid(alpha=0.25)
|
| 337 |
|
| 338 |
fig.suptitle("Did the run earn its length?", color=C_INK)
|
| 339 |
-
return save(fig, out, "
|
| 340 |
|
| 341 |
|
| 342 |
def plot_slice_deltas(plt, data, out, metric="reward"):
|
|
@@ -574,7 +560,7 @@ def verdict(data):
|
|
| 574 |
notes.append(
|
| 575 |
"- **Held-out loss converged at step %d of %d** (%.0f%% of the run). "
|
| 576 |
"The remaining steps did not improve it, so a shorter run would "
|
| 577 |
-
"have reached the same place
|
| 578 |
|
| 579 |
# 2. Did sampled-generation quality move?
|
| 580 |
for key, label in (("parse_rate", "parse rate"), ("reward", "reward")):
|
|
@@ -668,33 +654,28 @@ def verdict(data):
|
|
| 668 |
return lines + (notes or ["Nothing anomalous in the logged metrics."])
|
| 669 |
|
| 670 |
|
| 671 |
-
def summarise(data
|
| 672 |
lines = ["# Run summary", ""]
|
| 673 |
-
|
| 674 |
-
lines.append("| stage | steps | wall clock |
|
| 675 |
-
lines.append("|---|---|---|---|
|
| 676 |
for stage in STAGES:
|
| 677 |
if stage not in data:
|
| 678 |
-
lines.append("| %s | _not run_ | | |
|
| 679 |
continue
|
| 680 |
p = data[stage]["progress"]
|
| 681 |
if not p:
|
| 682 |
-
lines.append("| %s | ? | ? | ? |
|
| 683 |
continue
|
| 684 |
last = p[-1]
|
| 685 |
secs = last.get("elapsed_s") or 0.0
|
| 686 |
-
cost = last.get("cost_total_usd") or 0.0
|
| 687 |
total_time += secs
|
| 688 |
-
|
| 689 |
-
lines.append("| %s | %d | %dh%02dm | $%.2f | %d |"
|
| 690 |
% (stage.upper(), last.get("step", 0) + 1,
|
| 691 |
-
int(secs // 3600), int(secs % 3600 // 60),
|
| 692 |
sum(1 for r in p if r.get("is_eval"))))
|
| 693 |
-
lines += ["", "**Total** %dh%02dm
|
| 694 |
-
% (int(total_time // 3600), int(total_time % 3600 // 60)
|
| 695 |
-
if budget_usd:
|
| 696 |
-
lines[-1] += " against a $%.0f budget (%s)" % (
|
| 697 |
-
budget_usd, "within" if total_cost <= budget_usd else "OVER")
|
| 698 |
lines += verdict(data)
|
| 699 |
return "\n".join(lines) + "\n"
|
| 700 |
|
|
@@ -703,7 +684,6 @@ def main():
|
|
| 703 |
ap = argparse.ArgumentParser()
|
| 704 |
ap.add_argument("--runs", default="runs")
|
| 705 |
ap.add_argument("--out", default="runs/graphs")
|
| 706 |
-
ap.add_argument("--budget", type=float, default=120.0)
|
| 707 |
args = ap.parse_args()
|
| 708 |
|
| 709 |
try:
|
|
@@ -726,7 +706,7 @@ def main():
|
|
| 726 |
plot_learning_curves(plt, data, args.out),
|
| 727 |
plot_eval_quality(plt, data, args.out),
|
| 728 |
plot_slice_breakdown(plt, data, args.out),
|
| 729 |
-
|
| 730 |
plot_slice_trajectories(plt, data, args.out),
|
| 731 |
plot_convergence(plt, data, args.out),
|
| 732 |
plot_slice_deltas(plt, data, args.out),
|
|
@@ -737,7 +717,7 @@ def main():
|
|
| 737 |
|
| 738 |
summary_path = os.path.join(args.out, "summary.md")
|
| 739 |
with open(summary_path, "w") as f:
|
| 740 |
-
f.write(summarise(data
|
| 741 |
|
| 742 |
print("\nWrote %d charts to %s/" % (sum(1 for w in written if w), args.out))
|
| 743 |
for w in written:
|
|
@@ -745,7 +725,7 @@ def main():
|
|
| 745 |
print(" " + w)
|
| 746 |
print(" " + summary_path)
|
| 747 |
print()
|
| 748 |
-
print(summarise(data
|
| 749 |
|
| 750 |
|
| 751 |
if __name__ == "__main__":
|
|
|
|
| 1 |
#!/usr/bin/env python3
|
| 2 |
"""Turn the logs from every stage into charts, once all three have finished.
|
| 3 |
|
| 4 |
+
python3 scripts/plot_runs.py [--runs runs] [--out runs/graphs]
|
| 5 |
|
| 6 |
Reads two files per stage:
|
| 7 |
|
| 8 |
runs/<stage>/metrics.jsonl the cookbook's learning curves
|
| 9 |
+
runs/<stage>/progress.jsonl wall-clock and throughput (progress.py)
|
| 10 |
|
| 11 |
Stages that have not run yet are skipped with a note rather than an error, so
|
| 12 |
this is safe to run mid-pipeline -- you just get fewer panels.
|
|
|
|
| 195 |
return save(fig, out, "03_slice_breakdown.png")
|
| 196 |
|
| 197 |
|
| 198 |
+
def plot_timing(plt, data, out):
|
| 199 |
+
fig, axes = plt.subplots(1, 2, figsize=(11, 4))
|
| 200 |
+
ax_rate, ax_wall = axes
|
| 201 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 202 |
for stage in STAGES:
|
| 203 |
if stage not in data:
|
| 204 |
continue
|
|
|
|
| 206 |
xs, ys = series(p, "sec_per_step")
|
| 207 |
if ys:
|
| 208 |
ax_rate.plot(xs, ys, lw=1, label=stage)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 209 |
exs, eys = series(p, "elapsed_s")
|
| 210 |
if eys:
|
| 211 |
ax_wall.plot(exs, [y / 60.0 for y in eys], lw=1.2, label=stage)
|
|
|
|
|
|
|
|
|
|
| 212 |
|
| 213 |
ax_rate.set_title("Seconds per step")
|
| 214 |
ax_rate.set_xlabel("step")
|
|
|
|
|
|
|
| 215 |
ax_wall.set_title("Wall clock (minutes)")
|
| 216 |
ax_wall.set_xlabel("step")
|
| 217 |
for ax in axes:
|
| 218 |
ax.grid(alpha=0.3)
|
| 219 |
if ax.get_legend_handles_labels()[0]:
|
| 220 |
ax.legend(fontsize=8)
|
| 221 |
+
fig.suptitle("Throughput and wall clock")
|
| 222 |
+
return save(fig, out, "04_timing.png")
|
| 223 |
|
| 224 |
|
| 225 |
def plot_slice_trajectories(plt, data, out, metric="parse_rate"):
|
|
|
|
| 259 |
|
| 260 |
|
| 261 |
def plot_convergence(plt, data, out):
|
| 262 |
+
"""Where the held-out loss stopped improving, and how much run remained after.
|
| 263 |
|
| 264 |
This is the panel that answers 'should I have trained this long', which the
|
| 265 |
+
aggregate learning curve does not. The second panel measures the run after
|
| 266 |
+
convergence in wall-clock time.
|
| 267 |
"""
|
| 268 |
if "sft" not in data:
|
| 269 |
return None
|
|
|
|
| 280 |
conv_step = xs[conv_i]
|
| 281 |
|
| 282 |
prog = data["sft"]["progress"]
|
| 283 |
+
t_x = [r.get("step", 0) for r in prog if r.get("elapsed_s") is not None]
|
| 284 |
+
t_y = [r["elapsed_s"] / 60.0 for r in prog if r.get("elapsed_s") is not None]
|
| 285 |
+
total_min = t_y[-1] if t_y else 0.0
|
| 286 |
+
min_at_conv = 0.0
|
| 287 |
+
for sx, sy in zip(t_x, t_y):
|
| 288 |
if sx <= conv_step:
|
| 289 |
+
min_at_conv = sy
|
| 290 |
|
| 291 |
fig, axes = plt.subplots(1, 2, figsize=(12.5, 4.2))
|
| 292 |
|
|
|
|
| 305 |
ax.grid(alpha=0.25)
|
| 306 |
|
| 307 |
ax = axes[1]
|
| 308 |
+
if t_y:
|
| 309 |
+
ax.plot(t_x, t_y, color=C_SERIES_1, linewidth=2)
|
| 310 |
ax.axvline(conv_step, color=C_MID, linewidth=2, linestyle="--")
|
| 311 |
+
wasted = max(0.0, total_min - min_at_conv)
|
| 312 |
+
ax.fill_between([x for x in t_x if x >= conv_step],
|
| 313 |
+
[min_at_conv] * sum(1 for x in t_x if x >= conv_step),
|
| 314 |
+
[y for x, y in zip(t_x, t_y) if x >= conv_step],
|
| 315 |
color=C_WARN, alpha=0.18)
|
| 316 |
+
ax.annotate("%.0f min after convergence\n(of %.0f min total)" % (wasted, total_min),
|
| 317 |
xy=(0.42, 0.18), xycoords="axes fraction",
|
| 318 |
fontsize=10, color=C_WARN)
|
| 319 |
+
ax.set_title("Cumulative wall clock", color=C_INK)
|
| 320 |
ax.set_xlabel("training step", color=C_INK_2)
|
| 321 |
+
ax.set_ylabel("minutes", color=C_INK_2)
|
| 322 |
ax.grid(alpha=0.25)
|
| 323 |
|
| 324 |
fig.suptitle("Did the run earn its length?", color=C_INK)
|
| 325 |
+
return save(fig, out, "06_convergence.png")
|
| 326 |
|
| 327 |
|
| 328 |
def plot_slice_deltas(plt, data, out, metric="reward"):
|
|
|
|
| 560 |
notes.append(
|
| 561 |
"- **Held-out loss converged at step %d of %d** (%.0f%% of the run). "
|
| 562 |
"The remaining steps did not improve it, so a shorter run would "
|
| 563 |
+
"have reached the same place." % (xs[conv_i], xs[-1], frac * 100))
|
| 564 |
|
| 565 |
# 2. Did sampled-generation quality move?
|
| 566 |
for key, label in (("parse_rate", "parse rate"), ("reward", "reward")):
|
|
|
|
| 654 |
return lines + (notes or ["Nothing anomalous in the logged metrics."])
|
| 655 |
|
| 656 |
|
| 657 |
+
def summarise(data):
|
| 658 |
lines = ["# Run summary", ""]
|
| 659 |
+
total_time = 0.0
|
| 660 |
+
lines.append("| stage | steps | wall clock | eval rounds |")
|
| 661 |
+
lines.append("|---|---|---|---|")
|
| 662 |
for stage in STAGES:
|
| 663 |
if stage not in data:
|
| 664 |
+
lines.append("| %s | _not run_ | | |" % stage.upper())
|
| 665 |
continue
|
| 666 |
p = data[stage]["progress"]
|
| 667 |
if not p:
|
| 668 |
+
lines.append("| %s | ? | ? | ? |" % stage.upper())
|
| 669 |
continue
|
| 670 |
last = p[-1]
|
| 671 |
secs = last.get("elapsed_s") or 0.0
|
|
|
|
| 672 |
total_time += secs
|
| 673 |
+
lines.append("| %s | %d | %dh%02dm | %d |"
|
|
|
|
| 674 |
% (stage.upper(), last.get("step", 0) + 1,
|
| 675 |
+
int(secs // 3600), int(secs % 3600 // 60),
|
| 676 |
sum(1 for r in p if r.get("is_eval"))))
|
| 677 |
+
lines += ["", "**Total wall clock** %dh%02dm"
|
| 678 |
+
% (int(total_time // 3600), int(total_time % 3600 // 60))]
|
|
|
|
|
|
|
|
|
|
| 679 |
lines += verdict(data)
|
| 680 |
return "\n".join(lines) + "\n"
|
| 681 |
|
|
|
|
| 684 |
ap = argparse.ArgumentParser()
|
| 685 |
ap.add_argument("--runs", default="runs")
|
| 686 |
ap.add_argument("--out", default="runs/graphs")
|
|
|
|
| 687 |
args = ap.parse_args()
|
| 688 |
|
| 689 |
try:
|
|
|
|
| 706 |
plot_learning_curves(plt, data, args.out),
|
| 707 |
plot_eval_quality(plt, data, args.out),
|
| 708 |
plot_slice_breakdown(plt, data, args.out),
|
| 709 |
+
plot_timing(plt, data, args.out),
|
| 710 |
plot_slice_trajectories(plt, data, args.out),
|
| 711 |
plot_convergence(plt, data, args.out),
|
| 712 |
plot_slice_deltas(plt, data, args.out),
|
|
|
|
| 717 |
|
| 718 |
summary_path = os.path.join(args.out, "summary.md")
|
| 719 |
with open(summary_path, "w") as f:
|
| 720 |
+
f.write(summarise(data))
|
| 721 |
|
| 722 |
print("\nWrote %d charts to %s/" % (sum(1 for w in written if w), args.out))
|
| 723 |
for w in written:
|
|
|
|
| 725 |
print(" " + w)
|
| 726 |
print(" " + summary_path)
|
| 727 |
print()
|
| 728 |
+
print(summarise(data))
|
| 729 |
|
| 730 |
|
| 731 |
if __name__ == "__main__":
|
tests/test_budget.py
DELETED
|
@@ -1,134 +0,0 @@
|
|
| 1 |
-
"""Tests for the cost plan and the progress meter.
|
| 2 |
-
|
| 3 |
-
The load-bearing test is `test_plan_fits_the_budget`. The pipeline previously
|
| 4 |
-
shipped with an eval cadence that cost more than the training it was measuring
|
| 5 |
-
-- 96 rounds of 440 generations, far more than the training it measured. Nothing caught it because the estimator and the trainers held
|
| 6 |
-
separate copies of the eval assumption. Now both read claim_drafter/budget.py
|
| 7 |
-
and this test asserts the total, so the two cannot drift apart again.
|
| 8 |
-
|
| 9 |
-
pytest tests/ -q # or: python3 tests/test_budget.py
|
| 10 |
-
"""
|
| 11 |
-
|
| 12 |
-
import json
|
| 13 |
-
import os
|
| 14 |
-
import sys
|
| 15 |
-
import tempfile
|
| 16 |
-
|
| 17 |
-
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
| 18 |
-
|
| 19 |
-
from claim_drafter import budget
|
| 20 |
-
from claim_drafter.progress import StageProgress, _fmt_duration
|
| 21 |
-
|
| 22 |
-
BUDGET_USD = 120.0
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
class _Sink:
|
| 26 |
-
"""Stand-in for a non-TTY stream, so the bar renders without ANSI."""
|
| 27 |
-
|
| 28 |
-
def __init__(self):
|
| 29 |
-
self.text = ""
|
| 30 |
-
|
| 31 |
-
def write(self, s):
|
| 32 |
-
self.text += s
|
| 33 |
-
|
| 34 |
-
def flush(self):
|
| 35 |
-
pass
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
def test_plan_fits_the_budget():
|
| 39 |
-
total = budget.stage_plan()["total"]["cost"]
|
| 40 |
-
assert total <= BUDGET_USD, "plan costs $%.2f, over the $%.0f budget" % (total, BUDGET_USD)
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
def test_eval_is_a_minority_of_spend():
|
| 44 |
-
# Eval exists to measure training, so it should not rival it. This is the
|
| 45 |
-
# specific failure the old cadence had: eval was ~55% of the total.
|
| 46 |
-
t = budget.stage_plan()["total"]
|
| 47 |
-
assert t["eval_cost"] < 0.15 * t["train_cost"], \
|
| 48 |
-
"eval is $%.2f against $%.2f of training" % (t["eval_cost"], t["train_cost"])
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
def test_eval_volume_stays_bounded():
|
| 52 |
-
t = budget.stage_plan()["total"]
|
| 53 |
-
assert t["eval_gens"] < 2000, "%d eval generations is a wall-clock problem" % t["eval_gens"]
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
def test_every_stage_evaluates_at_least_once():
|
| 57 |
-
plan = budget.stage_plan()
|
| 58 |
-
for stage in ("sft", "dpo", "rl"):
|
| 59 |
-
assert plan[stage]["eval_rounds"] >= 1, "%s never runs the evaluator" % stage
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
def test_cheaper_model_costs_less():
|
| 63 |
-
big = budget.stage_plan(model="Qwen/Qwen3.5-9B")["total"]["cost"]
|
| 64 |
-
small = budget.stage_plan(model="Qwen/Qwen3.5-4B")["total"]["cost"]
|
| 65 |
-
assert small < big
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
def test_unknown_model_falls_back_rather_than_crashing():
|
| 69 |
-
assert budget.prices_for("Nonexistent/Model-1B") == budget.PRICES["Qwen/Qwen3.5-9B"]
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
def test_tokens_per_step_is_positive_everywhere():
|
| 73 |
-
plan = budget.stage_plan()
|
| 74 |
-
for stage in ("sft", "dpo", "rl"):
|
| 75 |
-
assert plan[stage]["tokens_per_step"] > 0, stage
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
def test_progress_writes_one_record_per_step():
|
| 79 |
-
d = tempfile.mkdtemp()
|
| 80 |
-
p = StageProgress("sft", 10, d, 1.463, tokens_per_step=1000, stream=_Sink())
|
| 81 |
-
for step in range(5):
|
| 82 |
-
p.log_metrics({"train_mean_nll": 2.0 - step * 0.1}, step=step)
|
| 83 |
-
p.close()
|
| 84 |
-
with open(os.path.join(d, "progress.jsonl")) as f:
|
| 85 |
-
records = [json.loads(l) for l in f if l.strip()]
|
| 86 |
-
assert len(records) == 5
|
| 87 |
-
assert records[-1]["step"] == 4
|
| 88 |
-
assert records[-1]["cost_total_usd"] > records[0]["cost_total_usd"]
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
def test_progress_flags_eval_rounds():
|
| 92 |
-
d = tempfile.mkdtemp()
|
| 93 |
-
p = StageProgress("sft", 4, d, 1.463, eval_cost_per_round=0.97, stream=_Sink())
|
| 94 |
-
p.log_metrics({"train_mean_nll": 1.0}, step=0)
|
| 95 |
-
p.log_metrics({"train_mean_nll": 1.0, "overall/reward": 0.6}, step=1)
|
| 96 |
-
p.close()
|
| 97 |
-
with open(os.path.join(d, "progress.jsonl")) as f:
|
| 98 |
-
records = [json.loads(l) for l in f if l.strip()]
|
| 99 |
-
assert records[0]["is_eval"] is False
|
| 100 |
-
assert records[1]["is_eval"] is True
|
| 101 |
-
assert records[1]["cost_eval_usd"] > 0
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
def test_eval_steps_do_not_poison_the_eta():
|
| 105 |
-
# An eval step is 10-100x slower than a training step. Folding it into the
|
| 106 |
-
# per-step average would inflate the ETA for every remaining step.
|
| 107 |
-
d = tempfile.mkdtemp()
|
| 108 |
-
p = StageProgress("sft", 100, d, 1.463, stream=_Sink())
|
| 109 |
-
p.log_metrics({"train_mean_nll": 1.0}, step=0)
|
| 110 |
-
p.log_metrics({"train_mean_nll": 1.0}, step=1)
|
| 111 |
-
baseline = p.sec_per_step
|
| 112 |
-
p.log_metrics({"train_mean_nll": 1.0, "overall/reward": 0.5}, step=2)
|
| 113 |
-
assert p.sec_per_step == baseline, "eval round changed the per-step average"
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
def test_duration_formatting():
|
| 117 |
-
assert _fmt_duration(45) == "45s"
|
| 118 |
-
assert _fmt_duration(750) == "12m30s"
|
| 119 |
-
assert _fmt_duration(11220) == "3h07m"
|
| 120 |
-
assert _fmt_duration(None) == "--"
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
if __name__ == "__main__":
|
| 124 |
-
failures = 0
|
| 125 |
-
for name, fn in sorted(globals().items()):
|
| 126 |
-
if name.startswith("test_") and callable(fn):
|
| 127 |
-
try:
|
| 128 |
-
fn()
|
| 129 |
-
print(" PASS %s" % name)
|
| 130 |
-
except AssertionError as e:
|
| 131 |
-
failures += 1
|
| 132 |
-
print(" FAIL %s: %s" % (name, e))
|
| 133 |
-
print("\n%d failure(s)" % failures)
|
| 134 |
-
sys.exit(1 if failures else 0)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
training/train_dpo.py
CHANGED
|
@@ -46,11 +46,11 @@ def main():
|
|
| 46 |
from tinker_cookbook.preference.preference_datasets import ComparisonBuilderFromJsonl
|
| 47 |
from tinker_cookbook.supervised.types import ChatDatasetBuilderCommonConfig
|
| 48 |
from tinker_cookbook.renderers import TrainOnWhat
|
| 49 |
-
from claim_drafter import
|
| 50 |
from claim_drafter.evaluator import ClaimEvaluatorBuilder
|
| 51 |
|
| 52 |
-
plan =
|
| 53 |
-
|
| 54 |
|
| 55 |
common = ChatDatasetBuilderCommonConfig(
|
| 56 |
model_name_for_tokenizer=args.model,
|
|
@@ -65,15 +65,12 @@ def main():
|
|
| 65 |
|
| 66 |
total_steps = args.max_steps or plan["steps"]
|
| 67 |
progress.announce(
|
| 68 |
-
"dpo", total_steps, plan["train_tokens"],
|
| 69 |
-
plan["eval_rounds"],
|
| 70 |
-
train_cost=plan["train_cost"])
|
| 71 |
progress.install(
|
| 72 |
-
"dpo", total_steps, args.log_path,
|
| 73 |
tokens_per_step=plan["tokens_per_step"],
|
| 74 |
-
eval_every=plan["infrequent_eval_every"], save_every=plan["save_every"]
|
| 75 |
-
eval_cost_per_round=plan["eval_cost_per_round"],
|
| 76 |
-
cost_per_step=plan["train_cost"] / max(plan["steps"], 1))
|
| 77 |
|
| 78 |
config = train_dpo.Config(
|
| 79 |
log_path=args.log_path,
|
|
|
|
| 46 |
from tinker_cookbook.preference.preference_datasets import ComparisonBuilderFromJsonl
|
| 47 |
from tinker_cookbook.supervised.types import ChatDatasetBuilderCommonConfig
|
| 48 |
from tinker_cookbook.renderers import TrainOnWhat
|
| 49 |
+
from claim_drafter import planning, progress
|
| 50 |
from claim_drafter.evaluator import ClaimEvaluatorBuilder
|
| 51 |
|
| 52 |
+
plan = planning.stage_plan(model=args.model, dpo_epochs=args.epochs,
|
| 53 |
+
dpo_batch=args.batch_size)["dpo"]
|
| 54 |
|
| 55 |
common = ChatDatasetBuilderCommonConfig(
|
| 56 |
model_name_for_tokenizer=args.model,
|
|
|
|
| 65 |
|
| 66 |
total_steps = args.max_steps or plan["steps"]
|
| 67 |
progress.announce(
|
| 68 |
+
"dpo", total_steps, plan["train_tokens"],
|
| 69 |
+
plan["eval_rounds"], planning.N_SLICES * planning.GENS_PER_SLICE)
|
|
|
|
| 70 |
progress.install(
|
| 71 |
+
"dpo", total_steps, args.log_path,
|
| 72 |
tokens_per_step=plan["tokens_per_step"],
|
| 73 |
+
eval_every=plan["infrequent_eval_every"], save_every=plan["save_every"])
|
|
|
|
|
|
|
| 74 |
|
| 75 |
config = train_dpo.Config(
|
| 76 |
log_path=args.log_path,
|
training/train_rl.py
CHANGED
|
@@ -49,24 +49,21 @@ def main():
|
|
| 49 |
|
| 50 |
from tinker_cookbook.rl import train as rl_train
|
| 51 |
from tinker_cookbook.rl.train import KLReferenceConfig
|
| 52 |
-
from claim_drafter import
|
| 53 |
from claim_drafter.evaluator import ClaimEvaluatorBuilder
|
| 54 |
from claim_drafter.rl_env import ClaimDatasetBuilder
|
| 55 |
|
| 56 |
-
plan =
|
| 57 |
rl_group=args.group_size,
|
| 58 |
rl_batch=args.batch_size)["rl"]
|
| 59 |
|
| 60 |
progress.announce(
|
| 61 |
-
"rl", plan["steps"], plan["train_tokens"],
|
| 62 |
-
plan["eval_rounds"],
|
| 63 |
-
train_cost=plan["train_cost"])
|
| 64 |
progress.install(
|
| 65 |
-
"rl", plan["steps"], args.log_path,
|
| 66 |
tokens_per_step=plan["tokens_per_step"],
|
| 67 |
-
eval_every=plan["eval_every"], save_every=plan["save_every"]
|
| 68 |
-
eval_cost_per_round=plan["eval_cost_per_round"],
|
| 69 |
-
cost_per_step=plan["train_cost"] / max(plan["steps"], 1))
|
| 70 |
|
| 71 |
config = rl_train.Config(
|
| 72 |
log_path=args.log_path,
|
|
|
|
| 49 |
|
| 50 |
from tinker_cookbook.rl import train as rl_train
|
| 51 |
from tinker_cookbook.rl.train import KLReferenceConfig
|
| 52 |
+
from claim_drafter import planning, progress
|
| 53 |
from claim_drafter.evaluator import ClaimEvaluatorBuilder
|
| 54 |
from claim_drafter.rl_env import ClaimDatasetBuilder
|
| 55 |
|
| 56 |
+
plan = planning.stage_plan(model=args.model, rl_prompts=args.prompts,
|
| 57 |
rl_group=args.group_size,
|
| 58 |
rl_batch=args.batch_size)["rl"]
|
| 59 |
|
| 60 |
progress.announce(
|
| 61 |
+
"rl", plan["steps"], plan["train_tokens"],
|
| 62 |
+
plan["eval_rounds"], planning.N_SLICES * planning.GENS_PER_SLICE)
|
|
|
|
| 63 |
progress.install(
|
| 64 |
+
"rl", plan["steps"], args.log_path,
|
| 65 |
tokens_per_step=plan["tokens_per_step"],
|
| 66 |
+
eval_every=plan["eval_every"], save_every=plan["save_every"])
|
|
|
|
|
|
|
| 67 |
|
| 68 |
config = rl_train.Config(
|
| 69 |
log_path=args.log_path,
|
training/train_sft.py
CHANGED
|
@@ -3,11 +3,10 @@
|
|
| 3 |
|
| 4 |
python3 training/train_sft.py --epochs 2
|
| 5 |
|
| 6 |
-
Reads TINKER_API_KEY from .env.
|
| 7 |
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
few cents and will surface any signature drift before you spend the budget.
|
| 11 |
"""
|
| 12 |
|
| 13 |
import argparse
|
|
@@ -40,11 +39,11 @@ def main():
|
|
| 40 |
from tinker_cookbook.supervised.data import FromConversationFileBuilder
|
| 41 |
from tinker_cookbook.supervised.types import ChatDatasetBuilderCommonConfig
|
| 42 |
from tinker_cookbook.renderers import TrainOnWhat
|
| 43 |
-
from claim_drafter import
|
| 44 |
from claim_drafter.evaluator import ClaimEvaluatorBuilder
|
| 45 |
|
| 46 |
-
plan =
|
| 47 |
-
|
| 48 |
|
| 49 |
common = ChatDatasetBuilderCommonConfig(
|
| 50 |
model_name_for_tokenizer=args.model,
|
|
@@ -60,27 +59,24 @@ def main():
|
|
| 60 |
train_on_what=TrainOnWhat.LAST_ASSISTANT_MESSAGE,
|
| 61 |
)
|
| 62 |
|
| 63 |
-
# Two evaluators on two cadences, because they
|
| 64 |
-
# magnitude
|
| 65 |
-
# to run often, and the only thing that shows overfitting, which
|
| 66 |
# train_mean_nll cannot. The slice evaluator samples 88 full claim sets and
|
| 67 |
-
# is the
|
| 68 |
-
#
|
| 69 |
slice_evals = [] if args.no_eval else [
|
| 70 |
ClaimEvaluatorBuilder(model_name_for_tokenizer=args.model)
|
| 71 |
]
|
| 72 |
|
| 73 |
total_steps = args.max_steps or plan["steps"]
|
| 74 |
progress.announce(
|
| 75 |
-
"sft", total_steps, plan["train_tokens"],
|
| 76 |
-
plan["eval_rounds"],
|
| 77 |
-
train_cost=plan["train_cost"])
|
| 78 |
progress.install(
|
| 79 |
-
"sft", total_steps, args.log_path,
|
| 80 |
tokens_per_step=plan["tokens_per_step"],
|
| 81 |
-
eval_every=plan["infrequent_eval_every"], save_every=plan["save_every"]
|
| 82 |
-
eval_cost_per_round=plan["eval_cost_per_round"],
|
| 83 |
-
cost_per_step=plan["train_cost"] / max(plan["steps"], 1))
|
| 84 |
|
| 85 |
config = train.Config(
|
| 86 |
log_path=args.log_path,
|
|
|
|
| 3 |
|
| 4 |
python3 training/train_sft.py --epochs 2
|
| 5 |
|
| 6 |
+
Reads TINKER_API_KEY from .env.
|
| 7 |
|
| 8 |
+
Run `python3 scripts/preflight.py` first — it exercises the same code path on a
|
| 9 |
+
couple of examples and will surface any signature drift before a full run.
|
|
|
|
| 10 |
"""
|
| 11 |
|
| 12 |
import argparse
|
|
|
|
| 39 |
from tinker_cookbook.supervised.data import FromConversationFileBuilder
|
| 40 |
from tinker_cookbook.supervised.types import ChatDatasetBuilderCommonConfig
|
| 41 |
from tinker_cookbook.renderers import TrainOnWhat
|
| 42 |
+
from claim_drafter import planning, progress
|
| 43 |
from claim_drafter.evaluator import ClaimEvaluatorBuilder
|
| 44 |
|
| 45 |
+
plan = planning.stage_plan(model=args.model, sft_epochs=args.epochs,
|
| 46 |
+
sft_batch=args.batch_size)["sft"]
|
| 47 |
|
| 48 |
common = ChatDatasetBuilderCommonConfig(
|
| 49 |
model_name_for_tokenizer=args.model,
|
|
|
|
| 59 |
train_on_what=TrainOnWhat.LAST_ASSISTANT_MESSAGE,
|
| 60 |
)
|
| 61 |
|
| 62 |
+
# Two evaluators on two cadences, because they differ by three orders of
|
| 63 |
+
# magnitude in runtime. The NLL holdout is a single forward pass -- fast
|
| 64 |
+
# enough to run often, and the only thing that shows overfitting, which
|
| 65 |
# train_mean_nll cannot. The slice evaluator samples 88 full claim sets and
|
| 66 |
+
# is the slow one, so it runs rarely. Running the slice evaluator on the NLL
|
| 67 |
+
# cadence is what made an early run take roughly twice as long.
|
| 68 |
slice_evals = [] if args.no_eval else [
|
| 69 |
ClaimEvaluatorBuilder(model_name_for_tokenizer=args.model)
|
| 70 |
]
|
| 71 |
|
| 72 |
total_steps = args.max_steps or plan["steps"]
|
| 73 |
progress.announce(
|
| 74 |
+
"sft", total_steps, plan["train_tokens"],
|
| 75 |
+
plan["eval_rounds"], planning.N_SLICES * planning.GENS_PER_SLICE)
|
|
|
|
| 76 |
progress.install(
|
| 77 |
+
"sft", total_steps, args.log_path,
|
| 78 |
tokens_per_step=plan["tokens_per_step"],
|
| 79 |
+
eval_every=plan["infrequent_eval_every"], save_every=plan["save_every"])
|
|
|
|
|
|
|
| 80 |
|
| 81 |
config = train.Config(
|
| 82 |
log_path=args.log_path,
|