vishwr commited on
Commit
2b18580
·
verified ·
1 Parent(s): 77cc9ef

Refactor: replace budget/cost tooling with cost-free planning module

Browse files
Makefile CHANGED
@@ -1,4 +1,4 @@
1
- .PHONY: help setup test validate preflight data sft dpo rl cost graphs all clean
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 incl. the Tinker API (costs a few cents)
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 --budget 120
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 what the pipeline costs.
3
 
4
  Three callers need the same arithmetic and used to disagree about it:
5
 
6
- scripts/estimate_cost.py prices the run before you start
7
- training/train_*.py size the eval cadence so the run fits the budget
8
- claim_drafter/progress.py meters spend while the run is in flight
9
 
10
- The eval numbers are the ones that used to be wrong. The old estimator assumed
11
- 5 rounds of 150 generations; the trainers were configured for 96 rounds of 440,
12
- a 56x gap that put the real total at roughly twice the budget. `stage_plan`
13
- below derives eval volume from the *same* cadence constants the trainers pass to
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
- # $ per 1M tokens, from https://tinker-docs.thinkingmachines.ai/tinker/models/
23
- # (checked 2026-07-20).
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 and under-counts eval
103
- cost -- verified against a real run: 20 steps at every=10 produced rounds at
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 cost for every stage, using the shipped cadences.
118
 
119
- Returns {stage: {...}} plus a "total" row. Everything downstream -- the
120
- estimator's table, the trainers' ETA banner, the budget check -- reads this.
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
- rl_prefill = rollouts * sft["avg_prompt"]
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 running cost 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,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, throughput and cumulative cost. metrics.jsonl has
15
- the learning curves but no timing or cost, 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,24 +57,15 @@ class StageProgress:
57
  tests can exercise the formatting.
58
  """
59
 
60
- def __init__(self, stage, total_steps, log_dir, price_train,
61
  tokens_per_step=0, eval_every=0, save_every=0,
62
- eval_cost_per_round=0.0, eval_seconds_per_round=None,
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, price_train, eval_rounds,
279
- eval_gens_per_round, eval_cost, stream=None, train_cost=None):
280
- """Print the pre-flight estimate for a stage before the first step runs.
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, price_train, tokens_per_step=0,
309
- eval_every=0, save_every=0, eval_cost_per_round=0.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
- price_train=price_train, tokens_per_step=tokens_per_step,
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] [--budget 120]
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, throughput and spend (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,14 +195,10 @@ def plot_slice_breakdown(plt, data, out):
195
  return save(fig, out, "03_slice_breakdown.png")
196
 
197
 
198
- def plot_timing_and_cost(plt, data, out, budget_usd):
199
- fig, axes = plt.subplots(1, 3, figsize=(15, 4))
200
- ax_rate, ax_cost, ax_wall = axes
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, spend and wall clock")
237
- return save(fig, out, "04_timing_and_cost.png")
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 what the rest cost.
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
- cost_x = [r.get("step", 0) for r in prog if r.get("cost_total_usd") is not None]
298
- cost_y = [r["cost_total_usd"] for r in prog if r.get("cost_total_usd") is not None]
299
- total_cost = cost_y[-1] if cost_y else 0.0
300
- cost_at_conv = 0.0
301
- for sx, sy in zip(cost_x, cost_y):
302
  if sx <= conv_step:
303
- cost_at_conv = sy
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 cost_y:
323
- ax.plot(cost_x, cost_y, color=C_SERIES_1, linewidth=2)
324
  ax.axvline(conv_step, color=C_MID, linewidth=2, linestyle="--")
325
- wasted = max(0.0, total_cost - cost_at_conv)
326
- ax.fill_between([x for x in cost_x if x >= conv_step],
327
- [cost_at_conv] * sum(1 for x in cost_x if x >= conv_step),
328
- [y for x, y in zip(cost_x, cost_y) if x >= conv_step],
329
  color=C_WARN, alpha=0.18)
330
- ax.annotate("$%.2f spent after convergence\n(of $%.2f total)" % (wasted, total_cost),
331
  xy=(0.42, 0.18), xycoords="axes fraction",
332
  fontsize=10, color=C_WARN)
333
- ax.set_title("Cumulative spend", color=C_INK)
334
  ax.set_xlabel("training step", color=C_INK_2)
335
- ax.set_ylabel("USD", color=C_INK_2)
336
  ax.grid(alpha=0.25)
337
 
338
  fig.suptitle("Did the run earn its length?", color=C_INK)
339
- return save(fig, out, "06_convergence_and_spend.png")
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 for less." % (xs[conv_i], xs[-1], frac * 100))
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, budget_usd):
672
  lines = ["# Run summary", ""]
673
- total_cost = total_time = 0.0
674
- lines.append("| stage | steps | wall clock | est. cost | eval rounds |")
675
- lines.append("|---|---|---|---|---|")
676
  for stage in STAGES:
677
  if stage not in data:
678
- lines.append("| %s | _not run_ | | | |" % stage.upper())
679
  continue
680
  p = data[stage]["progress"]
681
  if not p:
682
- lines.append("| %s | ? | ? | ? | ? |" % stage.upper())
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
- total_cost += cost
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), cost,
692
  sum(1 for r in p if r.get("is_eval"))))
693
- lines += ["", "**Total** %dh%02dm, $%.2f estimated"
694
- % (int(total_time // 3600), int(total_time % 3600 // 60), total_cost)]
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
- plot_timing_and_cost(plt, data, args.out, args.budget),
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, args.budget))
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, args.budget))
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 budget, progress
50
  from claim_drafter.evaluator import ClaimEvaluatorBuilder
51
 
52
- plan = budget.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,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"], budget.prices_for(args.model)["train"],
69
- plan["eval_rounds"], budget.N_SLICES * budget.GENS_PER_SLICE, plan["eval_cost"],
70
- train_cost=plan["train_cost"])
71
  progress.install(
72
- "dpo", total_steps, args.log_path, budget.prices_for(args.model)["train"],
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 budget, progress
53
  from claim_drafter.evaluator import ClaimEvaluatorBuilder
54
  from claim_drafter.rl_env import ClaimDatasetBuilder
55
 
56
- plan = budget.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"], budget.prices_for(args.model)["train"],
62
- plan["eval_rounds"], budget.N_SLICES * budget.GENS_PER_SLICE, plan["eval_cost"],
63
- train_cost=plan["train_cost"])
64
  progress.install(
65
- "rl", plan["steps"], args.log_path, budget.prices_for(args.model)["train"],
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. See docs/experiment.md for why each value is set.
7
 
8
- NOTE: this has not been executed against the live API (no key at authoring time).
9
- Run `python3 scripts/preflight.py` first — it exercises the same code path for a
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 budget, progress
44
  from claim_drafter.evaluator import ClaimEvaluatorBuilder
45
 
46
- plan = budget.stage_plan(model=args.model, sft_epochs=args.epochs,
47
- sft_batch=args.batch_size)["sft"]
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 cost three orders of
64
- # magnitude apart. The NLL holdout is a single forward pass -- cheap enough
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 expensive one, so it runs rarely. Running the slice evaluator on
68
- # the NLL cadence is what put the pipeline at ~2x budget.
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"], budget.prices_for(args.model)["train"],
76
- plan["eval_rounds"], budget.N_SLICES * budget.GENS_PER_SLICE, plan["eval_cost"],
77
- train_cost=plan["train_cost"])
78
  progress.install(
79
- "sft", total_steps, args.log_path, budget.prices_for(args.model)["train"],
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,