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.
- llm_waterfall/LICENSE +21 -0
- llm_waterfall/__init__.py +53 -0
- llm_waterfall/adapters/__init__.py +36 -0
- llm_waterfall/adapters/anthropic.py +105 -0
- llm_waterfall/adapters/aws_mantle.py +47 -0
- llm_waterfall/adapters/azure_openai.py +71 -0
- llm_waterfall/adapters/base.py +51 -0
- llm_waterfall/adapters/bedrock.py +309 -0
- llm_waterfall/adapters/openai.py +130 -0
- llm_waterfall/classify.py +184 -0
- llm_waterfall/pricing.py +110 -0
- llm_waterfall/py.typed +0 -0
- llm_waterfall/types.py +295 -0
- llm_waterfall/waterfall.py +255 -0
- wmo/__init__.py +38 -0
- wmo/agents/__init__.py +7 -0
- wmo/agents/default.py +29 -0
- wmo/agents/meta.py +55 -0
- wmo/agents/optimizer.py +55 -0
- wmo/agents/project.py +928 -0
- wmo/cli/__init__.py +5 -0
- wmo/cli/agent_session.py +1123 -0
- wmo/cli/app.py +2489 -0
- wmo/cli/e2b_cmds.py +212 -0
- wmo/cli/eval_closed_loop.py +207 -0
- wmo/cli/harness_app.py +1147 -0
- wmo/cli/harness_distill.py +659 -0
- wmo/cli/hosted_session.py +880 -0
- wmo/cli/ingest_cmd.py +165 -0
- wmo/cli/model_roles.py +82 -0
- wmo/cli/platform_cmds.py +372 -0
- wmo/cli/route_app.py +274 -0
- wmo/cli/session_state.py +243 -0
- wmo/cli/ui.py +1107 -0
- wmo/cli/workspace_sync.py +504 -0
- wmo/config/__init__.py +60 -0
- wmo/config/card.py +129 -0
- wmo/config/config.py +367 -0
- wmo/config/dotenv.py +67 -0
- wmo/config/settings.py +128 -0
- wmo/config/store.py +177 -0
- wmo/conftest.py +19 -0
- wmo/connect/__init__.py +88 -0
- wmo/connect/apps.py +78 -0
- wmo/connect/brave.py +284 -0
- wmo/connect/connector.py +79 -0
- wmo/connect/credentials.py +164 -0
- wmo/connect/github.py +321 -0
- wmo/connect/google.py +627 -0
- wmo/connect/notion.py +790 -0
- wmo/connect/oauth.py +461 -0
- wmo/connect/slack.py +555 -0
- wmo/connect/store.py +199 -0
- wmo/connect/types.py +156 -0
- wmo/core/__init__.py +21 -0
- wmo/core/parsing.py +281 -0
- wmo/core/render.py +271 -0
- wmo/core/text.py +40 -0
- wmo/core/types.py +116 -0
- wmo/distill/__init__.py +14 -0
- wmo/distill/agents.py +140 -0
- wmo/distill/config.py +1006 -0
- wmo/distill/cost.py +437 -0
- wmo/distill/data.py +921 -0
- wmo/distill/deadlines.py +254 -0
- wmo/distill/fake_tinker.py +734 -0
- wmo/distill/gate.py +122 -0
- wmo/distill/loop.py +3499 -0
- wmo/distill/renderers.py +399 -0
- wmo/distill/rendering.py +620 -0
- wmo/distill/rollouts.py +726 -0
- wmo/distill/samples.py +195 -0
- wmo/distill/store.py +829 -0
- wmo/distill/teacher.py +714 -0
- wmo/distill/tokens.py +535 -0
- wmo/distill/tracking.py +552 -0
- wmo/distill/tripwire.py +411 -0
- wmo/distill/xtoken/byte_offsets.py +152 -0
- wmo/distill/xtoken/chunks.py +457 -0
- wmo/distill/xtoken/prompt_logprobs.py +475 -0
- wmo/distill/xtoken/teacher_render.py +346 -0
- wmo/engine/__init__.py +28 -0
- wmo/engine/autoconfig.py +367 -0
- wmo/engine/build.py +346 -0
- wmo/engine/demo.py +77 -0
- wmo/engine/eval_suites.py +245 -0
- wmo/engine/grounding.py +491 -0
- wmo/engine/knowledge.py +291 -0
- wmo/engine/loader.py +36 -0
- wmo/engine/play.py +92 -0
- wmo/engine/prompts.py +99 -0
- wmo/engine/replay.py +443 -0
- wmo/engine/reporting.py +58 -0
- wmo/engine/workspace.py +468 -0
- wmo/engine/world_model.py +568 -0
- wmo/env/__init__.py +22 -0
- wmo/env/base.py +121 -0
- wmo/env/closed_loop.py +229 -0
- wmo/env/episode.py +107 -0
- wmo/env/llm_agent.py +93 -0
- wmo/env/scenarios.py +73 -0
- wmo/evals/__init__.py +52 -0
- wmo/evals/agreement.py +110 -0
- wmo/evals/base.py +45 -0
- wmo/evals/closed_loop.py +480 -0
- wmo/evals/failover.py +96 -0
- wmo/evals/gold.py +127 -0
- wmo/evals/grid.py +394 -0
- wmo/evals/grid_plot.py +205 -0
- wmo/evals/harbor/__init__.py +27 -0
- wmo/evals/harbor/agent.py +573 -0
- wmo/evals/harbor/ctrf.py +171 -0
- wmo/evals/harbor/e2b_environment.py +587 -0
- wmo/evals/harbor/e2b_template_policy.py +144 -0
- wmo/evals/harbor/scorer.py +875 -0
- wmo/evals/harbor/tasks.py +140 -0
- wmo/evals/open_loop.py +194 -0
- wmo/evals/tasks.py +53 -0
- wmo/harness/__init__.py +51 -0
- wmo/harness/code_runtime.py +288 -0
- wmo/harness/create.py +1191 -0
- wmo/harness/delta.py +220 -0
- wmo/harness/doc.py +556 -0
- wmo/harness/e2b_ledger.py +342 -0
- wmo/harness/e2b_reap.py +476 -0
- wmo/harness/e2b_sandbox.py +350 -0
- wmo/harness/environment.py +35 -0
- wmo/harness/live_session.py +543 -0
- wmo/harness/mutate.py +343 -0
- wmo/harness/pi_e2b.py +1710 -0
- wmo/harness/pi_entry/entry.ts +268 -0
- wmo/harness/pi_entry/runner_frames.ts +92 -0
- wmo/harness/pi_entry/runner_live.ts +587 -0
- wmo/harness/pi_entry/runner_service.ts +270 -0
- wmo/harness/pi_entry/runner_stdio.ts +374 -0
- wmo/harness/pi_entry/runner_termination.ts +142 -0
- wmo/harness/pi_local.py +262 -0
- wmo/harness/pi_runtime.py +495 -0
- wmo/harness/pi_vendor.py +65 -0
- wmo/harness/population.py +509 -0
- wmo/harness/project_proposer.py +569 -0
- wmo/harness/proposer.py +977 -0
- wmo/harness/runner_link.py +619 -0
- wmo/harness/runtime.py +389 -0
- wmo/harness/scoring.py +247 -0
- wmo/harness/skills.py +116 -0
- wmo/harness/source_tree.py +319 -0
- wmo/harness/store.py +176 -0
- wmo/harness/tools.py +105 -0
- wmo/harness/vendor/manifest.sha256 +58 -0
- wmo/harness/vendor/pi-agent/CHANGELOG.md +556 -0
- wmo/harness/vendor/pi-agent/LICENSE +21 -0
- wmo/harness/vendor/pi-agent/README.md +488 -0
- wmo/harness/vendor/pi-agent/VENDOR.md +39 -0
- wmo/harness/vendor/pi-agent/docs/agent-harness.md +486 -0
- wmo/harness/vendor/pi-agent/docs/durable-harness.md +212 -0
- wmo/harness/vendor/pi-agent/docs/hooks.md +445 -0
- wmo/harness/vendor/pi-agent/docs/models.md +966 -0
- wmo/harness/vendor/pi-agent/docs/observability.md +376 -0
- wmo/harness/vendor/pi-agent/package.json +60 -0
- wmo/harness/vendor/pi-agent/src/agent-loop.ts +748 -0
- wmo/harness/vendor/pi-agent/src/agent.ts +575 -0
- wmo/harness/vendor/pi-agent/src/harness/agent-harness.ts +1029 -0
- wmo/harness/vendor/pi-agent/src/harness/compaction/branch-summarization.ts +261 -0
- wmo/harness/vendor/pi-agent/src/harness/compaction/compaction.ts +747 -0
- wmo/harness/vendor/pi-agent/src/harness/compaction/utils.ts +144 -0
- wmo/harness/vendor/pi-agent/src/harness/env/nodejs.ts +550 -0
- wmo/harness/vendor/pi-agent/src/harness/messages.ts +164 -0
- wmo/harness/vendor/pi-agent/src/harness/prompt-templates.ts +267 -0
- wmo/harness/vendor/pi-agent/src/harness/session/jsonl-repo.ts +177 -0
- wmo/harness/vendor/pi-agent/src/harness/session/jsonl-storage.ts +293 -0
- wmo/harness/vendor/pi-agent/src/harness/session/memory-repo.ts +50 -0
- wmo/harness/vendor/pi-agent/src/harness/session/memory-storage.ts +131 -0
- wmo/harness/vendor/pi-agent/src/harness/session/repo-utils.ts +51 -0
- wmo/harness/vendor/pi-agent/src/harness/session/session.ts +267 -0
- wmo/harness/vendor/pi-agent/src/harness/session/uuid.ts +54 -0
- wmo/harness/vendor/pi-agent/src/harness/skills.ts +375 -0
- wmo/harness/vendor/pi-agent/src/harness/system-prompt.ts +34 -0
- wmo/harness/vendor/pi-agent/src/harness/types.ts +836 -0
- wmo/harness/vendor/pi-agent/src/harness/utils/shell-output.ts +135 -0
- wmo/harness/vendor/pi-agent/src/harness/utils/truncate.ts +344 -0
- wmo/harness/vendor/pi-agent/src/index.ts +44 -0
- wmo/harness/vendor/pi-agent/src/node.ts +2 -0
- wmo/harness/vendor/pi-agent/src/proxy.ts +367 -0
- wmo/harness/vendor/pi-agent/src/types.ts +428 -0
- wmo/harness/vendor/pi-agent/test/agent-loop.test.ts +1351 -0
- wmo/harness/vendor/pi-agent/test/agent.test.ts +699 -0
- wmo/harness/vendor/pi-agent/test/e2e.test.ts +404 -0
- wmo/harness/vendor/pi-agent/test/harness/agent-harness-stream.test.ts +213 -0
- wmo/harness/vendor/pi-agent/test/harness/agent-harness.test.ts +608 -0
- wmo/harness/vendor/pi-agent/test/harness/compaction.test.ts +655 -0
- wmo/harness/vendor/pi-agent/test/harness/nodejs-env.test.ts +321 -0
- wmo/harness/vendor/pi-agent/test/harness/prompt-templates.test.ts +90 -0
- wmo/harness/vendor/pi-agent/test/harness/repo.test.ts +68 -0
- wmo/harness/vendor/pi-agent/test/harness/resource-formatting.test.ts +24 -0
- wmo/harness/vendor/pi-agent/test/harness/session-test-utils.ts +55 -0
- wmo/harness/vendor/pi-agent/test/harness/session-uuid.test.ts +50 -0
- wmo/harness/vendor/pi-agent/test/harness/session.test.ts +156 -0
- wmo/harness/vendor/pi-agent/test/harness/skills.test.ts +116 -0
- wmo/harness/vendor/pi-agent/test/harness/storage.test.ts +299 -0
- wmo/harness/vendor/pi-agent/test/harness/system-prompt.test.ts +66 -0
- wmo/harness/vendor/pi-agent/test/harness/truncate.test.ts +169 -0
- wmo/harness/vendor/pi-agent/test/scratch/simple.ts +72 -0
- wmo/harness/vendor/pi-agent/test/utils/calculate.ts +32 -0
- wmo/harness/vendor/pi-agent/test/utils/get-current-time.ts +46 -0
- wmo/harness/vendor/pi-agent/tsconfig.build.json +13 -0
- wmo/harness/vendor/pi-agent/vitest.config.ts +19 -0
- wmo/harness/vendor/pi-agent/vitest.harness.config.ts +28 -0
- wmo/harness/vendor/vendor_pi.sh +59 -0
- wmo/harness/workspace_patch.py +270 -0
- wmo/ingest/__init__.py +47 -0
- wmo/ingest/adapter.py +72 -0
- wmo/ingest/base.py +114 -0
- wmo/ingest/braintrust.py +339 -0
- wmo/ingest/detect.py +126 -0
- wmo/ingest/langfuse.py +291 -0
- wmo/ingest/langsmith.py +444 -0
- wmo/ingest/mastra.py +330 -0
- wmo/ingest/messages.py +170 -0
- wmo/ingest/normalize.py +679 -0
- wmo/ingest/otel_genai.py +69 -0
- wmo/ingest/otel_writer.py +100 -0
- wmo/ingest/phoenix.py +150 -0
- wmo/ingest/postgres.py +246 -0
- wmo/ingest/posthog.py +320 -0
- wmo/ingest/quality.py +28 -0
- wmo/ingest/stream.py +209 -0
- wmo/ingest/testdata/sample_otlp.json +60 -0
- wmo/ingest/testdata/sample_spans.jsonl +3 -0
- wmo/optimize/__init__.py +25 -0
- wmo/optimize/base.py +143 -0
- wmo/optimize/gepa.py +806 -0
- wmo/optimize/judge.py +262 -0
- wmo/optimize/judge_quality.py +359 -0
- wmo/optimize/knn.py +468 -0
- wmo/optimize/numeric.py +152 -0
- wmo/optimize/outcomes.py +103 -0
- wmo/optimize/policy.py +669 -0
- wmo/optimize/report.py +231 -0
- wmo/optimize/reward.py +129 -0
- wmo/optimize/routing.py +373 -0
- wmo/platform/__init__.py +6 -0
- wmo/platform/auth.py +115 -0
- wmo/platform/client.py +551 -0
- wmo/platform/credentials.py +126 -0
- wmo/platform/transfer.py +158 -0
- wmo/providers/__init__.py +40 -0
- wmo/providers/_bedrock_chat.py +155 -0
- wmo/providers/_openai_common.py +182 -0
- wmo/providers/_responses_common.py +472 -0
- wmo/providers/anthropic.py +134 -0
- wmo/providers/azure_openai.py +296 -0
- wmo/providers/base.py +300 -0
- wmo/providers/bedrock.py +312 -0
- wmo/providers/models.py +205 -0
- wmo/providers/openai.py +143 -0
- wmo/providers/openai_responses.py +240 -0
- wmo/providers/pool.py +170 -0
- wmo/providers/registry.py +73 -0
- wmo/providers/retry.py +151 -0
- wmo/providers/tinker.py +936 -0
- wmo/providers/waterfall.py +336 -0
- wmo/research/__init__.py +81 -0
- wmo/research/ablation.py +133 -0
- wmo/research/concurrency_plot.py +523 -0
- wmo/research/concurrency_run.py +240 -0
- wmo/research/concurrency_scaling.py +270 -0
- wmo/research/gepa_scaling.py +274 -0
- wmo/research/pipeline.py +198 -0
- wmo/research/scaling_split.py +82 -0
- wmo/research/scenario_fidelity.py +198 -0
- wmo/research/scenario_recovery.py +92 -0
- wmo/research/seed_stability.py +90 -0
- wmo/research/trace_scaling.py +348 -0
- wmo/retrieval/__init__.py +6 -0
- wmo/retrieval/embedders.py +105 -0
- wmo/retrieval/leakfree.py +52 -0
- wmo/retrieval/retriever.py +173 -0
- wmo/scenarios/__init__.py +58 -0
- wmo/scenarios/builder.py +152 -0
- wmo/scenarios/mining/__init__.py +27 -0
- wmo/scenarios/mining/clustering.py +171 -0
- wmo/scenarios/mining/facets.py +226 -0
- wmo/scenarios/mining/selection.py +220 -0
- wmo/scenarios/synthesis/__init__.py +6 -0
- wmo/scenarios/synthesis/scenario_set.py +63 -0
- wmo/scenarios/synthesis/synthesizer.py +85 -0
- wmo/scenarios/verification/__init__.py +17 -0
- wmo/scenarios/verification/judge.py +97 -0
- wmo/scenarios/verification/verify.py +135 -0
- wmo/serving/__init__.py +5 -0
- wmo/serving/builds.py +451 -0
- wmo/serving/chat.py +878 -0
- wmo/serving/endpoint_config.py +64 -0
- wmo/serving/savings.py +250 -0
- wmo/serving/server.py +553 -0
- wmo/serving/traces_source.py +206 -0
- wmo/telemetry.py +213 -0
- wmo/tracking/__init__.py +36 -0
- wmo/tracking/clock.py +24 -0
- wmo/tracking/metered.py +125 -0
- wmo/tracking/pricing.py +99 -0
- wmo/tracking/store.py +31 -0
- wmo/tracking/tracker.py +149 -0
- world_model_optimizer-0.2.0.dist-info/METADATA +203 -0
- world_model_optimizer-0.2.0.dist-info/RECORD +308 -0
- world_model_optimizer-0.2.0.dist-info/WHEEL +4 -0
- 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]
|