world-model-optimizer 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (308) hide show
  1. llm_waterfall/LICENSE +21 -0
  2. llm_waterfall/__init__.py +53 -0
  3. llm_waterfall/adapters/__init__.py +36 -0
  4. llm_waterfall/adapters/anthropic.py +105 -0
  5. llm_waterfall/adapters/aws_mantle.py +47 -0
  6. llm_waterfall/adapters/azure_openai.py +71 -0
  7. llm_waterfall/adapters/base.py +51 -0
  8. llm_waterfall/adapters/bedrock.py +309 -0
  9. llm_waterfall/adapters/openai.py +130 -0
  10. llm_waterfall/classify.py +184 -0
  11. llm_waterfall/pricing.py +110 -0
  12. llm_waterfall/py.typed +0 -0
  13. llm_waterfall/types.py +295 -0
  14. llm_waterfall/waterfall.py +255 -0
  15. wmo/__init__.py +38 -0
  16. wmo/agents/__init__.py +7 -0
  17. wmo/agents/default.py +29 -0
  18. wmo/agents/meta.py +55 -0
  19. wmo/agents/optimizer.py +55 -0
  20. wmo/agents/project.py +928 -0
  21. wmo/cli/__init__.py +5 -0
  22. wmo/cli/agent_session.py +1123 -0
  23. wmo/cli/app.py +2489 -0
  24. wmo/cli/e2b_cmds.py +212 -0
  25. wmo/cli/eval_closed_loop.py +207 -0
  26. wmo/cli/harness_app.py +1147 -0
  27. wmo/cli/harness_distill.py +659 -0
  28. wmo/cli/hosted_session.py +880 -0
  29. wmo/cli/ingest_cmd.py +165 -0
  30. wmo/cli/model_roles.py +82 -0
  31. wmo/cli/platform_cmds.py +372 -0
  32. wmo/cli/route_app.py +274 -0
  33. wmo/cli/session_state.py +243 -0
  34. wmo/cli/ui.py +1107 -0
  35. wmo/cli/workspace_sync.py +504 -0
  36. wmo/config/__init__.py +60 -0
  37. wmo/config/card.py +129 -0
  38. wmo/config/config.py +367 -0
  39. wmo/config/dotenv.py +67 -0
  40. wmo/config/settings.py +128 -0
  41. wmo/config/store.py +177 -0
  42. wmo/conftest.py +19 -0
  43. wmo/connect/__init__.py +88 -0
  44. wmo/connect/apps.py +78 -0
  45. wmo/connect/brave.py +284 -0
  46. wmo/connect/connector.py +79 -0
  47. wmo/connect/credentials.py +164 -0
  48. wmo/connect/github.py +321 -0
  49. wmo/connect/google.py +627 -0
  50. wmo/connect/notion.py +790 -0
  51. wmo/connect/oauth.py +461 -0
  52. wmo/connect/slack.py +555 -0
  53. wmo/connect/store.py +199 -0
  54. wmo/connect/types.py +156 -0
  55. wmo/core/__init__.py +21 -0
  56. wmo/core/parsing.py +281 -0
  57. wmo/core/render.py +271 -0
  58. wmo/core/text.py +40 -0
  59. wmo/core/types.py +116 -0
  60. wmo/distill/__init__.py +14 -0
  61. wmo/distill/agents.py +140 -0
  62. wmo/distill/config.py +1006 -0
  63. wmo/distill/cost.py +437 -0
  64. wmo/distill/data.py +921 -0
  65. wmo/distill/deadlines.py +254 -0
  66. wmo/distill/fake_tinker.py +734 -0
  67. wmo/distill/gate.py +122 -0
  68. wmo/distill/loop.py +3499 -0
  69. wmo/distill/renderers.py +399 -0
  70. wmo/distill/rendering.py +620 -0
  71. wmo/distill/rollouts.py +726 -0
  72. wmo/distill/samples.py +195 -0
  73. wmo/distill/store.py +829 -0
  74. wmo/distill/teacher.py +714 -0
  75. wmo/distill/tokens.py +535 -0
  76. wmo/distill/tracking.py +552 -0
  77. wmo/distill/tripwire.py +411 -0
  78. wmo/distill/xtoken/byte_offsets.py +152 -0
  79. wmo/distill/xtoken/chunks.py +457 -0
  80. wmo/distill/xtoken/prompt_logprobs.py +475 -0
  81. wmo/distill/xtoken/teacher_render.py +346 -0
  82. wmo/engine/__init__.py +28 -0
  83. wmo/engine/autoconfig.py +367 -0
  84. wmo/engine/build.py +346 -0
  85. wmo/engine/demo.py +77 -0
  86. wmo/engine/eval_suites.py +245 -0
  87. wmo/engine/grounding.py +491 -0
  88. wmo/engine/knowledge.py +291 -0
  89. wmo/engine/loader.py +36 -0
  90. wmo/engine/play.py +92 -0
  91. wmo/engine/prompts.py +99 -0
  92. wmo/engine/replay.py +443 -0
  93. wmo/engine/reporting.py +58 -0
  94. wmo/engine/workspace.py +468 -0
  95. wmo/engine/world_model.py +568 -0
  96. wmo/env/__init__.py +22 -0
  97. wmo/env/base.py +121 -0
  98. wmo/env/closed_loop.py +229 -0
  99. wmo/env/episode.py +107 -0
  100. wmo/env/llm_agent.py +93 -0
  101. wmo/env/scenarios.py +73 -0
  102. wmo/evals/__init__.py +52 -0
  103. wmo/evals/agreement.py +110 -0
  104. wmo/evals/base.py +45 -0
  105. wmo/evals/closed_loop.py +480 -0
  106. wmo/evals/failover.py +96 -0
  107. wmo/evals/gold.py +127 -0
  108. wmo/evals/grid.py +394 -0
  109. wmo/evals/grid_plot.py +205 -0
  110. wmo/evals/harbor/__init__.py +27 -0
  111. wmo/evals/harbor/agent.py +573 -0
  112. wmo/evals/harbor/ctrf.py +171 -0
  113. wmo/evals/harbor/e2b_environment.py +587 -0
  114. wmo/evals/harbor/e2b_template_policy.py +144 -0
  115. wmo/evals/harbor/scorer.py +875 -0
  116. wmo/evals/harbor/tasks.py +140 -0
  117. wmo/evals/open_loop.py +194 -0
  118. wmo/evals/tasks.py +53 -0
  119. wmo/harness/__init__.py +51 -0
  120. wmo/harness/code_runtime.py +288 -0
  121. wmo/harness/create.py +1191 -0
  122. wmo/harness/delta.py +220 -0
  123. wmo/harness/doc.py +556 -0
  124. wmo/harness/e2b_ledger.py +342 -0
  125. wmo/harness/e2b_reap.py +476 -0
  126. wmo/harness/e2b_sandbox.py +350 -0
  127. wmo/harness/environment.py +35 -0
  128. wmo/harness/live_session.py +543 -0
  129. wmo/harness/mutate.py +343 -0
  130. wmo/harness/pi_e2b.py +1710 -0
  131. wmo/harness/pi_entry/entry.ts +268 -0
  132. wmo/harness/pi_entry/runner_frames.ts +92 -0
  133. wmo/harness/pi_entry/runner_live.ts +587 -0
  134. wmo/harness/pi_entry/runner_service.ts +270 -0
  135. wmo/harness/pi_entry/runner_stdio.ts +374 -0
  136. wmo/harness/pi_entry/runner_termination.ts +142 -0
  137. wmo/harness/pi_local.py +262 -0
  138. wmo/harness/pi_runtime.py +495 -0
  139. wmo/harness/pi_vendor.py +65 -0
  140. wmo/harness/population.py +509 -0
  141. wmo/harness/project_proposer.py +569 -0
  142. wmo/harness/proposer.py +977 -0
  143. wmo/harness/runner_link.py +619 -0
  144. wmo/harness/runtime.py +389 -0
  145. wmo/harness/scoring.py +247 -0
  146. wmo/harness/skills.py +116 -0
  147. wmo/harness/source_tree.py +319 -0
  148. wmo/harness/store.py +176 -0
  149. wmo/harness/tools.py +105 -0
  150. wmo/harness/vendor/manifest.sha256 +58 -0
  151. wmo/harness/vendor/pi-agent/CHANGELOG.md +556 -0
  152. wmo/harness/vendor/pi-agent/LICENSE +21 -0
  153. wmo/harness/vendor/pi-agent/README.md +488 -0
  154. wmo/harness/vendor/pi-agent/VENDOR.md +39 -0
  155. wmo/harness/vendor/pi-agent/docs/agent-harness.md +486 -0
  156. wmo/harness/vendor/pi-agent/docs/durable-harness.md +212 -0
  157. wmo/harness/vendor/pi-agent/docs/hooks.md +445 -0
  158. wmo/harness/vendor/pi-agent/docs/models.md +966 -0
  159. wmo/harness/vendor/pi-agent/docs/observability.md +376 -0
  160. wmo/harness/vendor/pi-agent/package.json +60 -0
  161. wmo/harness/vendor/pi-agent/src/agent-loop.ts +748 -0
  162. wmo/harness/vendor/pi-agent/src/agent.ts +575 -0
  163. wmo/harness/vendor/pi-agent/src/harness/agent-harness.ts +1029 -0
  164. wmo/harness/vendor/pi-agent/src/harness/compaction/branch-summarization.ts +261 -0
  165. wmo/harness/vendor/pi-agent/src/harness/compaction/compaction.ts +747 -0
  166. wmo/harness/vendor/pi-agent/src/harness/compaction/utils.ts +144 -0
  167. wmo/harness/vendor/pi-agent/src/harness/env/nodejs.ts +550 -0
  168. wmo/harness/vendor/pi-agent/src/harness/messages.ts +164 -0
  169. wmo/harness/vendor/pi-agent/src/harness/prompt-templates.ts +267 -0
  170. wmo/harness/vendor/pi-agent/src/harness/session/jsonl-repo.ts +177 -0
  171. wmo/harness/vendor/pi-agent/src/harness/session/jsonl-storage.ts +293 -0
  172. wmo/harness/vendor/pi-agent/src/harness/session/memory-repo.ts +50 -0
  173. wmo/harness/vendor/pi-agent/src/harness/session/memory-storage.ts +131 -0
  174. wmo/harness/vendor/pi-agent/src/harness/session/repo-utils.ts +51 -0
  175. wmo/harness/vendor/pi-agent/src/harness/session/session.ts +267 -0
  176. wmo/harness/vendor/pi-agent/src/harness/session/uuid.ts +54 -0
  177. wmo/harness/vendor/pi-agent/src/harness/skills.ts +375 -0
  178. wmo/harness/vendor/pi-agent/src/harness/system-prompt.ts +34 -0
  179. wmo/harness/vendor/pi-agent/src/harness/types.ts +836 -0
  180. wmo/harness/vendor/pi-agent/src/harness/utils/shell-output.ts +135 -0
  181. wmo/harness/vendor/pi-agent/src/harness/utils/truncate.ts +344 -0
  182. wmo/harness/vendor/pi-agent/src/index.ts +44 -0
  183. wmo/harness/vendor/pi-agent/src/node.ts +2 -0
  184. wmo/harness/vendor/pi-agent/src/proxy.ts +367 -0
  185. wmo/harness/vendor/pi-agent/src/types.ts +428 -0
  186. wmo/harness/vendor/pi-agent/test/agent-loop.test.ts +1351 -0
  187. wmo/harness/vendor/pi-agent/test/agent.test.ts +699 -0
  188. wmo/harness/vendor/pi-agent/test/e2e.test.ts +404 -0
  189. wmo/harness/vendor/pi-agent/test/harness/agent-harness-stream.test.ts +213 -0
  190. wmo/harness/vendor/pi-agent/test/harness/agent-harness.test.ts +608 -0
  191. wmo/harness/vendor/pi-agent/test/harness/compaction.test.ts +655 -0
  192. wmo/harness/vendor/pi-agent/test/harness/nodejs-env.test.ts +321 -0
  193. wmo/harness/vendor/pi-agent/test/harness/prompt-templates.test.ts +90 -0
  194. wmo/harness/vendor/pi-agent/test/harness/repo.test.ts +68 -0
  195. wmo/harness/vendor/pi-agent/test/harness/resource-formatting.test.ts +24 -0
  196. wmo/harness/vendor/pi-agent/test/harness/session-test-utils.ts +55 -0
  197. wmo/harness/vendor/pi-agent/test/harness/session-uuid.test.ts +50 -0
  198. wmo/harness/vendor/pi-agent/test/harness/session.test.ts +156 -0
  199. wmo/harness/vendor/pi-agent/test/harness/skills.test.ts +116 -0
  200. wmo/harness/vendor/pi-agent/test/harness/storage.test.ts +299 -0
  201. wmo/harness/vendor/pi-agent/test/harness/system-prompt.test.ts +66 -0
  202. wmo/harness/vendor/pi-agent/test/harness/truncate.test.ts +169 -0
  203. wmo/harness/vendor/pi-agent/test/scratch/simple.ts +72 -0
  204. wmo/harness/vendor/pi-agent/test/utils/calculate.ts +32 -0
  205. wmo/harness/vendor/pi-agent/test/utils/get-current-time.ts +46 -0
  206. wmo/harness/vendor/pi-agent/tsconfig.build.json +13 -0
  207. wmo/harness/vendor/pi-agent/vitest.config.ts +19 -0
  208. wmo/harness/vendor/pi-agent/vitest.harness.config.ts +28 -0
  209. wmo/harness/vendor/vendor_pi.sh +59 -0
  210. wmo/harness/workspace_patch.py +270 -0
  211. wmo/ingest/__init__.py +47 -0
  212. wmo/ingest/adapter.py +72 -0
  213. wmo/ingest/base.py +114 -0
  214. wmo/ingest/braintrust.py +339 -0
  215. wmo/ingest/detect.py +126 -0
  216. wmo/ingest/langfuse.py +291 -0
  217. wmo/ingest/langsmith.py +444 -0
  218. wmo/ingest/mastra.py +330 -0
  219. wmo/ingest/messages.py +170 -0
  220. wmo/ingest/normalize.py +679 -0
  221. wmo/ingest/otel_genai.py +69 -0
  222. wmo/ingest/otel_writer.py +100 -0
  223. wmo/ingest/phoenix.py +150 -0
  224. wmo/ingest/postgres.py +246 -0
  225. wmo/ingest/posthog.py +320 -0
  226. wmo/ingest/quality.py +28 -0
  227. wmo/ingest/stream.py +209 -0
  228. wmo/ingest/testdata/sample_otlp.json +60 -0
  229. wmo/ingest/testdata/sample_spans.jsonl +3 -0
  230. wmo/optimize/__init__.py +25 -0
  231. wmo/optimize/base.py +143 -0
  232. wmo/optimize/gepa.py +806 -0
  233. wmo/optimize/judge.py +262 -0
  234. wmo/optimize/judge_quality.py +359 -0
  235. wmo/optimize/knn.py +468 -0
  236. wmo/optimize/numeric.py +152 -0
  237. wmo/optimize/outcomes.py +103 -0
  238. wmo/optimize/policy.py +669 -0
  239. wmo/optimize/report.py +231 -0
  240. wmo/optimize/reward.py +129 -0
  241. wmo/optimize/routing.py +373 -0
  242. wmo/platform/__init__.py +6 -0
  243. wmo/platform/auth.py +115 -0
  244. wmo/platform/client.py +551 -0
  245. wmo/platform/credentials.py +126 -0
  246. wmo/platform/transfer.py +158 -0
  247. wmo/providers/__init__.py +40 -0
  248. wmo/providers/_bedrock_chat.py +155 -0
  249. wmo/providers/_openai_common.py +182 -0
  250. wmo/providers/_responses_common.py +472 -0
  251. wmo/providers/anthropic.py +134 -0
  252. wmo/providers/azure_openai.py +296 -0
  253. wmo/providers/base.py +300 -0
  254. wmo/providers/bedrock.py +312 -0
  255. wmo/providers/models.py +205 -0
  256. wmo/providers/openai.py +143 -0
  257. wmo/providers/openai_responses.py +240 -0
  258. wmo/providers/pool.py +170 -0
  259. wmo/providers/registry.py +73 -0
  260. wmo/providers/retry.py +151 -0
  261. wmo/providers/tinker.py +936 -0
  262. wmo/providers/waterfall.py +336 -0
  263. wmo/research/__init__.py +81 -0
  264. wmo/research/ablation.py +133 -0
  265. wmo/research/concurrency_plot.py +523 -0
  266. wmo/research/concurrency_run.py +240 -0
  267. wmo/research/concurrency_scaling.py +270 -0
  268. wmo/research/gepa_scaling.py +274 -0
  269. wmo/research/pipeline.py +198 -0
  270. wmo/research/scaling_split.py +82 -0
  271. wmo/research/scenario_fidelity.py +198 -0
  272. wmo/research/scenario_recovery.py +92 -0
  273. wmo/research/seed_stability.py +90 -0
  274. wmo/research/trace_scaling.py +348 -0
  275. wmo/retrieval/__init__.py +6 -0
  276. wmo/retrieval/embedders.py +105 -0
  277. wmo/retrieval/leakfree.py +52 -0
  278. wmo/retrieval/retriever.py +173 -0
  279. wmo/scenarios/__init__.py +58 -0
  280. wmo/scenarios/builder.py +152 -0
  281. wmo/scenarios/mining/__init__.py +27 -0
  282. wmo/scenarios/mining/clustering.py +171 -0
  283. wmo/scenarios/mining/facets.py +226 -0
  284. wmo/scenarios/mining/selection.py +220 -0
  285. wmo/scenarios/synthesis/__init__.py +6 -0
  286. wmo/scenarios/synthesis/scenario_set.py +63 -0
  287. wmo/scenarios/synthesis/synthesizer.py +85 -0
  288. wmo/scenarios/verification/__init__.py +17 -0
  289. wmo/scenarios/verification/judge.py +97 -0
  290. wmo/scenarios/verification/verify.py +135 -0
  291. wmo/serving/__init__.py +5 -0
  292. wmo/serving/builds.py +451 -0
  293. wmo/serving/chat.py +878 -0
  294. wmo/serving/endpoint_config.py +64 -0
  295. wmo/serving/savings.py +250 -0
  296. wmo/serving/server.py +553 -0
  297. wmo/serving/traces_source.py +206 -0
  298. wmo/telemetry.py +213 -0
  299. wmo/tracking/__init__.py +36 -0
  300. wmo/tracking/clock.py +24 -0
  301. wmo/tracking/metered.py +125 -0
  302. wmo/tracking/pricing.py +99 -0
  303. wmo/tracking/store.py +31 -0
  304. wmo/tracking/tracker.py +149 -0
  305. world_model_optimizer-0.2.0.dist-info/METADATA +203 -0
  306. world_model_optimizer-0.2.0.dist-info/RECORD +308 -0
  307. world_model_optimizer-0.2.0.dist-info/WHEEL +4 -0
  308. world_model_optimizer-0.2.0.dist-info/entry_points.txt +2 -0
wmo/distill/cost.py ADDED
@@ -0,0 +1,437 @@
1
+ """Cost projection and budget metering for one distillation run.
2
+
3
+ `estimate_run_cost` turns the run config plus the task split sizes into
4
+ per-meter token projections priced from the `[pricing]` section, for the CLI's
5
+ cost-confirm prompt. `BudgetMeter` then accumulates the ACTUAL token counts as
6
+ the run spends, and `check()` enforces the optional `[budget] max_usd` hard
7
+ cap by raising `BudgetExhausted` (the loop saves state and prints the resume
8
+ command on that error).
9
+
10
+ Metering follows Tinker's PER-REQUEST billing, not unique context tokens:
11
+ every sampling request bills its whole prompt, so each agent turn re-bills the
12
+ episode's full context, with the verbatim repeated prefix billed at the
13
+ discounted cached rate (`episode_billing` documents the exact split). Rollout
14
+ episodes therefore charge three meters each (full prefill, cached prefill,
15
+ sample), and teacher-in-harness episodes bill their sampled tokens at the
16
+ teacher's SAMPLING rate. Ignoring the per-request term once under-reported a
17
+ console-reconciled run by ~6x (306M billed tokens vs ~50M unique).
18
+
19
+ The projection is a deliberately simple, documented heuristic: episode counts
20
+ come exactly from the config (steps x tasks x group size, plus warmup teacher
21
+ episodes, interim evals, and the gate/baseline episodes), and per-episode
22
+ tokens come from a turns x tokens-per-turn model capped by the rollout context
23
+ budget. Meters mirror the `[pricing]` fields (cached rates fall back to the
24
+ documented 20% derivation); a meter without a price surfaces as a None-usd
25
+ line so the CLI can warn instead of silently under-reporting.
26
+ """
27
+
28
+ from __future__ import annotations
29
+
30
+ import logging
31
+ import math
32
+ from collections.abc import Sequence
33
+ from typing import Literal
34
+
35
+ from pydantic import BaseModel, ConfigDict, Field
36
+
37
+ from wmo.distill.config import DistillConfig, PricingConfig
38
+ from wmo.distill.tokens import TrialRecord
39
+ from wmo.providers.tinker import TokenSpan
40
+
41
+ logger = logging.getLogger(__name__)
42
+
43
+ MeterName = Literal[
44
+ "student_prefill",
45
+ "student_cached_prefill",
46
+ "student_sample",
47
+ "student_train",
48
+ "teacher_prefill",
49
+ "teacher_cached_prefill",
50
+ "teacher_sample",
51
+ ]
52
+
53
+ METER_NAMES: tuple[MeterName, ...] = (
54
+ "student_prefill",
55
+ "student_cached_prefill",
56
+ "student_sample",
57
+ "student_train",
58
+ "teacher_prefill",
59
+ "teacher_cached_prefill",
60
+ "teacher_sample",
61
+ )
62
+
63
+ _TOKENS_PER_USD_UNIT = 1_000_000
64
+ """Prices in `[pricing]` are USD per million tokens."""
65
+
66
+ _AVG_TURN_FRACTION = 0.5
67
+ """Episodes are assumed to use half the configured turn cap on average."""
68
+
69
+ _SAMPLED_TOKENS_PER_TURN = 512
70
+ """Assumed sampled (assistant/tool-call) tokens per agent turn."""
71
+
72
+ _OBSERVATION_TOKENS_PER_TURN = 1024
73
+ """Assumed prompt growth per turn (tool results and scaffolding)."""
74
+
75
+ _BASE_PROMPT_TOKENS = 2048
76
+ """Assumed initial prompt (system prompt, task instruction, tool schemas)."""
77
+
78
+
79
+ class BudgetExhausted(RuntimeError):
80
+ """Raised by `BudgetMeter.check` when actual spend exceeds the hard cap."""
81
+
82
+ def __init__(self, spent_usd: float, max_usd: float) -> None:
83
+ self.spent_usd = spent_usd
84
+ self.max_usd = max_usd
85
+ super().__init__(
86
+ f"budget exhausted: ${spent_usd:.2f} spent against the ${max_usd:.2f} "
87
+ "cap (budget.max_usd); the run saves its training state on this error, "
88
+ "so raise budget.max_usd in the run config and resume the run to continue"
89
+ )
90
+
91
+
92
+ class CostLine(BaseModel):
93
+ """One meter's token projection (or actuals) with its optional price."""
94
+
95
+ model_config = ConfigDict(frozen=True, extra="forbid")
96
+
97
+ meter: MeterName
98
+ tokens: int = Field(ge=0)
99
+ price_per_mtok: float | None
100
+ """USD per million tokens from `[pricing]`; None means unpriced."""
101
+
102
+ usd: float | None
103
+ """tokens x price; None when the meter is unpriced (CLI warns on these)."""
104
+
105
+
106
+ class CostEstimate(BaseModel):
107
+ """Per-meter projections for one run, plus the episode counts behind them."""
108
+
109
+ model_config = ConfigDict(frozen=True, extra="forbid")
110
+
111
+ lines: list[CostLine]
112
+ train_episodes: int = Field(ge=0)
113
+ eval_episodes: int = Field(ge=0)
114
+ baseline_episodes: int = Field(ge=0)
115
+ """Gate/baseline episodes: student before + student after + teacher-in-harness."""
116
+
117
+ warmup_episodes: int = Field(ge=0)
118
+ """Warmup teacher episodes: train tasks x warmup.rollouts_per_task (0 when off)."""
119
+
120
+ @property
121
+ def priced_usd(self) -> float:
122
+ """Total USD over the priced lines only."""
123
+ return sum(line.usd for line in self.lines if line.usd is not None)
124
+
125
+ @property
126
+ def unpriced_meters(self) -> list[MeterName]:
127
+ """Meters with no `[pricing]` entry, for the CLI's warning."""
128
+ return [line.meter for line in self.lines if line.usd is None]
129
+
130
+ def is_fully_priced(self) -> bool:
131
+ """Whether every meter carries a price, so `priced_usd` is the whole run."""
132
+ return not self.unpriced_meters
133
+
134
+
135
+ def _meter_price(pricing: PricingConfig, meter: MeterName) -> float | None:
136
+ if meter == "student_prefill":
137
+ return pricing.student_prefill
138
+ if meter == "student_cached_prefill":
139
+ return pricing.effective_student_cached_prefill
140
+ if meter == "student_sample":
141
+ return pricing.student_sample
142
+ if meter == "student_train":
143
+ return pricing.student_train
144
+ if meter == "teacher_prefill":
145
+ return pricing.teacher_prefill
146
+ if meter == "teacher_cached_prefill":
147
+ return pricing.effective_teacher_cached_prefill
148
+ return pricing.teacher_sample
149
+
150
+
151
+ def _line(pricing: PricingConfig, meter: MeterName, tokens: int) -> CostLine:
152
+ price = _meter_price(pricing, meter)
153
+ usd = tokens / _TOKENS_PER_USD_UNIT * price if price is not None else None
154
+ return CostLine(meter=meter, tokens=tokens, price_per_mtok=price, usd=usd)
155
+
156
+
157
+ class SpanBilling(BaseModel):
158
+ """Per-request billing volumes measured from recorded rollout spans.
159
+
160
+ The three volumes map onto three meters: `unique_tokens` at the full
161
+ prefill rate, `cached_tokens` at the cached-prefill rate, and
162
+ `sampled_tokens` at the sampling rate.
163
+ """
164
+
165
+ model_config = ConfigDict(frozen=True, extra="forbid")
166
+
167
+ unique_tokens: int = Field(ge=0)
168
+ """Distinct episode tokens, billed once at the full prefill rate."""
169
+
170
+ cached_tokens: int = Field(ge=0)
171
+ """Repeated per-request prompt volume, billed at the cached-prefill rate."""
172
+
173
+ sampled_tokens: int = Field(ge=0)
174
+ """Sampled completion tokens, billed at the sampling rate."""
175
+
176
+ def __add__(self, other: SpanBilling) -> SpanBilling:
177
+ return SpanBilling(
178
+ unique_tokens=self.unique_tokens + other.unique_tokens,
179
+ cached_tokens=self.cached_tokens + other.cached_tokens,
180
+ sampled_tokens=self.sampled_tokens + other.sampled_tokens,
181
+ )
182
+
183
+
184
+ def episode_billing(spans: Sequence[TokenSpan]) -> SpanBilling:
185
+ """Per-request billing volumes for ONE episode's recorded spans.
186
+
187
+ Tinker bills prefill PER REQUEST over each call's full prompt: every
188
+ agent turn re-bills the episode's whole context, with the verbatim
189
+ repeated prefix billed at the cached rate. The model here:
190
+
191
+ - per-request volume = sum over spans of `len(prompt_token_ids)`.
192
+ - unique = tokens the episode put through the model for the first time.
193
+ For a prefix-clean episode (every prompt extends the previous prompt
194
+ plus its sampled tokens verbatim, the same test `build_datums` merges
195
+ on) this is exactly the final span's prompt plus sampled length. A
196
+ prefix break restarts the accumulation, so a fragmented episode's
197
+ unique volume is the sum over fragments: re-prefilled context counts
198
+ as unique again, matching what the service re-bills at the full rate.
199
+ - cached = the per-request volume beyond unique, clamped at zero (a
200
+ single-call episode repeats nothing). Under the prefix property every
201
+ repeat is verbatim, so the cached rate applies to all of it.
202
+
203
+ Args:
204
+ spans: One episode's recorded spans, in any order (sorted by
205
+ call_index here).
206
+
207
+ Returns:
208
+ The episode's billing volumes.
209
+ """
210
+ unique = 0
211
+ per_request = 0
212
+ sampled = 0
213
+ accumulated: list[int] = []
214
+ for span in sorted(spans, key=lambda item: item.call_index):
215
+ prompt = span.prompt_token_ids
216
+ per_request += len(prompt)
217
+ if accumulated and prompt[: len(accumulated)] == accumulated:
218
+ unique += len(prompt) - len(accumulated)
219
+ else:
220
+ unique += len(prompt)
221
+ unique += len(span.sampled_token_ids)
222
+ sampled += len(span.sampled_token_ids)
223
+ accumulated = list(prompt) + list(span.sampled_token_ids)
224
+ return SpanBilling(
225
+ unique_tokens=unique,
226
+ cached_tokens=max(per_request - unique, 0),
227
+ sampled_tokens=sampled,
228
+ )
229
+
230
+
231
+ def batch_billing(records: Sequence[TrialRecord]) -> SpanBilling:
232
+ """Summed `episode_billing` over one rollout batch's trial records.
233
+
234
+ Summed per episode (not over a flattened span list) so each episode's
235
+ cached volume clamps independently and one trial's prefix break never
236
+ bleeds into another's accounting.
237
+
238
+ Args:
239
+ records: The batch's trial records (span-less trials contribute 0).
240
+
241
+ Returns:
242
+ The batch's total billing volumes.
243
+ """
244
+ total = SpanBilling(unique_tokens=0, cached_tokens=0, sampled_tokens=0)
245
+ for record in records:
246
+ total = total + episode_billing(record.spans)
247
+ return total
248
+
249
+
250
+ def estimate_run_cost(cfg: DistillConfig, n_train_tasks: int, n_holdout_tasks: int) -> CostEstimate:
251
+ """Project the run's per-meter token volumes and price them.
252
+
253
+ Episode counts (exact, from the config):
254
+
255
+ - train: `steps x min(tasks_per_batch, n_train_tasks) x group_size`
256
+ - warmup: `n_train_tasks x warmup.rollouts_per_task` teacher episodes when
257
+ `warmup.steps > 0`, else 0
258
+ - interim evals: `steps // eval.every` evals (0 when eval.every is 0) of
259
+ `min(eval.tasks, n_train_tasks) x eval.k` student episodes each
260
+ - gate/baseline: student-before and student-after at
261
+ `n_holdout x gate.k` each, plus one teacher-in-harness baseline at
262
+ `n_holdout x gate.k`
263
+
264
+ Per-episode tokens (heuristic; module constants document the assumptions):
265
+
266
+ - `avg_turns = max(1, ceil(rollout.max_turns x 0.5))`
267
+ - `sampled = avg_turns x min(sampling.max_tokens, 512)`
268
+ - `episode_tokens = min(rollout.context_budget_tokens,
269
+ 2048 + avg_turns x (1024 + sampled_per_turn))`: the episode's final
270
+ unique sequence length under the prefix property (the estimate assumes
271
+ prefix-clean episodes, so this is also the unique billing volume)
272
+ - per-request prefill: turn k's prompt is
273
+ `min(2048 + k x 1024 + (k - 1) x sampled_per_turn,
274
+ context_budget_tokens)` and every turn re-bills it whole, so the
275
+ per-request volume is the sum over turns; the part beyond
276
+ `episode_tokens` is the verbatim repeat billed at the cached rate
277
+ (see `episode_billing` for the same split on actual spans)
278
+
279
+ Meter mapping: every student episode charges `episode_tokens` (its unique
280
+ volume) to student_prefill, the repeated per-request volume to
281
+ student_cached_prefill, and `sampled` to student_sample; every train
282
+ episode additionally charges `episode_tokens` to student_train
283
+ (forward_backward over the full datum; x `train.topk` under the
284
+ `topk_ce` loss, whose k rank replicas each carry the full sequence) and
285
+ `episode_tokens` to teacher_prefill (the teacher scores each episode's
286
+ full sequence once, one full-price request with no repeats to cache;
287
+ the topk_ce prefill-only request bills the same volume). Teacher-in-harness
288
+ episodes (the gate baseline and the warmup collection) charge `sampled`
289
+ to teacher_sample (they bill the teacher's SAMPLING rate on what they
290
+ generate) plus per-request prefill exactly like a student episode, onto
291
+ teacher_prefill and teacher_cached_prefill. Warmup SFT training tokens
292
+ are NOT projected: they depend on how many teacher trials pass
293
+ (unknowable up front) and are bounded by warmup.steps full-batch passes
294
+ over at most the warmup episodes' tokens.
295
+
296
+ Args:
297
+ cfg: The validated run config.
298
+ n_train_tasks: Size of the train task split (must be >= 1).
299
+ n_holdout_tasks: Size of the holdout task split (>= 0; 0 skips the
300
+ gate/baseline episodes entirely).
301
+
302
+ Returns:
303
+ The estimate, one line per meter in `METER_NAMES` order.
304
+
305
+ Raises:
306
+ ValueError: If the split sizes are out of range.
307
+ """
308
+ if n_train_tasks < 1:
309
+ raise ValueError(
310
+ f"n_train_tasks must be >= 1, got {n_train_tasks}; a distillation run "
311
+ "needs a non-empty train task split"
312
+ )
313
+ if n_holdout_tasks < 0:
314
+ raise ValueError(f"n_holdout_tasks must be >= 0, got {n_holdout_tasks}")
315
+
316
+ avg_turns = max(1, math.ceil(cfg.rollout.max_turns * _AVG_TURN_FRACTION))
317
+ sampled_per_turn = min(cfg.sampling.max_tokens, _SAMPLED_TOKENS_PER_TURN)
318
+ context_budget = cfg.rollout.context_budget_tokens
319
+ episode_tokens = min(
320
+ context_budget,
321
+ _BASE_PROMPT_TOKENS + avg_turns * (_OBSERVATION_TOKENS_PER_TURN + sampled_per_turn),
322
+ )
323
+ sampled_tokens = min(avg_turns * sampled_per_turn, episode_tokens)
324
+ # Per-request accounting: every turn re-bills its whole prompt, so the
325
+ # per-request volume sums the per-turn prompts; the episode's distinct
326
+ # tokens bill once at the full rate and the rest is the cached repeat.
327
+ per_request_tokens = sum(
328
+ min(
329
+ _BASE_PROMPT_TOKENS
330
+ + turn * _OBSERVATION_TOKENS_PER_TURN
331
+ + (turn - 1) * sampled_per_turn,
332
+ context_budget,
333
+ )
334
+ for turn in range(1, avg_turns + 1)
335
+ )
336
+ cached_tokens = max(per_request_tokens - episode_tokens, 0)
337
+
338
+ tasks_per_step = min(cfg.train.tasks_per_batch, n_train_tasks)
339
+ train_episodes = cfg.train.steps * tasks_per_step * cfg.train.group_size
340
+ warmup_episodes = n_train_tasks * cfg.warmup.rollouts_per_task if cfg.warmup.steps > 0 else 0
341
+ interim_evals = cfg.train.steps // cfg.eval.every if cfg.eval.every > 0 else 0
342
+ eval_episodes = interim_evals * min(cfg.eval.tasks, n_train_tasks) * cfg.eval.k
343
+ gate_attempts = n_holdout_tasks * cfg.gate.k
344
+ student_baseline_episodes = 2 * gate_attempts # student-before + student-after
345
+ teacher_baseline_episodes = gate_attempts
346
+
347
+ student_episodes = train_episodes + eval_episodes + student_baseline_episodes
348
+ teacher_harness_episodes = teacher_baseline_episodes + warmup_episodes
349
+ # The topk_ce loss trains k rank-aligned cross_entropy replicas per datum,
350
+ # each carrying the full sequence, so its train volume is k x the default
351
+ # loss's (the loop meters actuals the same way; see build_topk_ce_datums).
352
+ train_replication = cfg.train.topk if cfg.train.loss == "topk_ce" else 1
353
+ projections: dict[MeterName, int] = {
354
+ "student_prefill": student_episodes * episode_tokens,
355
+ "student_cached_prefill": student_episodes * cached_tokens,
356
+ "student_sample": student_episodes * sampled_tokens,
357
+ "student_train": train_episodes * episode_tokens * train_replication,
358
+ "teacher_prefill": (
359
+ train_episodes * episode_tokens + teacher_harness_episodes * episode_tokens
360
+ ),
361
+ "teacher_cached_prefill": teacher_harness_episodes * cached_tokens,
362
+ "teacher_sample": teacher_harness_episodes * sampled_tokens,
363
+ }
364
+ estimate = CostEstimate(
365
+ lines=[_line(cfg.pricing, meter, projections[meter]) for meter in METER_NAMES],
366
+ train_episodes=train_episodes,
367
+ eval_episodes=eval_episodes,
368
+ baseline_episodes=student_baseline_episodes + teacher_baseline_episodes,
369
+ warmup_episodes=warmup_episodes,
370
+ )
371
+ logger.debug(
372
+ "cost estimate: %d train + %d warmup + %d eval + %d baseline episode(s), priced $%.2f%s",
373
+ estimate.train_episodes,
374
+ estimate.warmup_episodes,
375
+ estimate.eval_episodes,
376
+ estimate.baseline_episodes,
377
+ estimate.priced_usd,
378
+ f" (unpriced meters: {', '.join(estimate.unpriced_meters)})"
379
+ if estimate.unpriced_meters
380
+ else "",
381
+ )
382
+ return estimate
383
+
384
+
385
+ class BudgetMeter:
386
+ """Accumulates actual metered tokens and enforces the hard USD cap.
387
+
388
+ Args:
389
+ pricing: The `[pricing]` section; unpriced meters accumulate tokens
390
+ but contribute no USD (mirroring the estimate's None lines).
391
+ max_usd: The `[budget] max_usd` hard cap; None disables enforcement.
392
+ """
393
+
394
+ def __init__(self, pricing: PricingConfig, max_usd: float | None = None) -> None:
395
+ self._pricing = pricing
396
+ self._max_usd = max_usd
397
+ self._tokens: dict[MeterName, int] = {meter: 0 for meter in METER_NAMES}
398
+ self._spent_usd = 0.0
399
+
400
+ def charge(self, meter: MeterName, tokens: int) -> None:
401
+ """Record actual token usage against one meter.
402
+
403
+ Args:
404
+ meter: Which meter the tokens belong to.
405
+ tokens: The token count to add (>= 0).
406
+
407
+ Raises:
408
+ ValueError: If `tokens` is negative.
409
+ """
410
+ if tokens < 0:
411
+ raise ValueError(f"cannot charge a negative token count ({tokens}) to {meter}")
412
+ self._tokens[meter] += tokens
413
+ price = _meter_price(self._pricing, meter)
414
+ if price is not None:
415
+ self._spent_usd += tokens / _TOKENS_PER_USD_UNIT * price
416
+
417
+ def check(self) -> None:
418
+ """Enforce the hard cap against the priced spend so far.
419
+
420
+ Raises:
421
+ BudgetExhausted: When the cap is set and the spend exceeds it.
422
+ """
423
+ if self._max_usd is not None and self._spent_usd > self._max_usd:
424
+ raise BudgetExhausted(self._spent_usd, self._max_usd)
425
+
426
+ def tokens(self, meter: MeterName) -> int:
427
+ """Actual tokens charged to one meter so far."""
428
+ return self._tokens[meter]
429
+
430
+ @property
431
+ def spent_usd(self) -> float:
432
+ """Priced USD spend so far (unpriced meters contribute nothing)."""
433
+ return self._spent_usd
434
+
435
+ def lines(self) -> list[CostLine]:
436
+ """The actuals in the same line shape the estimate uses, for reporting."""
437
+ return [_line(self._pricing, meter, self._tokens[meter]) for meter in METER_NAMES]