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
|
@@ -0,0 +1,475 @@
|
|
|
1
|
+
"""Teacher-forced scoring against a self-hosted vLLM `/v1/completions` endpoint.
|
|
2
|
+
|
|
3
|
+
This is the cross-tokenizer teacher's ONLY network surface. `PromptLogprobClient.score`
|
|
4
|
+
submits an exact token sequence we already own and returns one logprob per position, so
|
|
5
|
+
the returned row indexes one for one into the teacher token ids the chunk aligner
|
|
6
|
+
produced. Nothing here tokenizes, renders, or samples.
|
|
7
|
+
|
|
8
|
+
Five wire facts are load-bearing (each was a real bug or a live probe finding):
|
|
9
|
+
|
|
10
|
+
1. The prompt goes on the wire as `list[int]`, NEVER as text. vLLM re-tokenizes a text
|
|
11
|
+
prompt server-side with `add_special_tokens` defaulting to True, which prepends
|
|
12
|
+
GLM's prefix/BOS and shifts EVERY position by one against our local offsets. There
|
|
13
|
+
is no error: the response looks perfectly well formed and every span sum is wrong.
|
|
14
|
+
2. `/v1/chat/completions` supports neither `echo` nor `prompt_logprobs`, so it cannot
|
|
15
|
+
score a prompt at all. The chat template is applied client-side (see
|
|
16
|
+
`wmo.distill.rendering`) and the rendered ids come here.
|
|
17
|
+
3. Position convention: `prompt_logprobs[p]` is the distribution FOR token p (entry 0
|
|
18
|
+
is null because token 0 has no context). That is exactly the Tinker
|
|
19
|
+
`compute_logprobs` convention `wmo.distill.teacher` already uses, so `score`
|
|
20
|
+
returns `len(token_ids)` entries with entry 0 = None and no shifting anywhere.
|
|
21
|
+
A short row is rejected rather than returned: it would silently corrupt every
|
|
22
|
+
downstream chunk sum.
|
|
23
|
+
4. The response shape varies by vLLM version: `prompt_logprobs` sits either at the top
|
|
24
|
+
level or under `choices[0]`. Both are read. Each position is a dict keyed by token
|
|
25
|
+
id (a JSON string) whose value carries a `logprob`, and the REALIZED token's entry
|
|
26
|
+
is the one we want, never the argmax.
|
|
27
|
+
5. Auth is a Bearer token when `api_key` is set. The repo convention for self-hosted
|
|
28
|
+
endpoints is the `WMO_ENDPOINT_API_KEY` env var (see `wmo.providers.openai`), but
|
|
29
|
+
this module never reads the environment: the caller passes the key in.
|
|
30
|
+
|
|
31
|
+
Deadlines are owned here rather than through `wmo.distill.deadlines`. That module
|
|
32
|
+
bounds Tinker SDK calls (futures and blocking calls with no timeout parameter) and its
|
|
33
|
+
knobs are Tinker-named env vars; httpx already bounds an HTTP request natively, so
|
|
34
|
+
wrapping it in a watchdog thread would add a second, weaker timer and a misleading
|
|
35
|
+
`TinkerDeadlineError`. The default `timeout_s` is 1200s (20 minutes), sized for the
|
|
36
|
+
real workload rather than for latency headroom: merged datum length is median ~14.5k
|
|
37
|
+
and up to 65.5k teacher tokens, and a distillation step fires many of these
|
|
38
|
+
concurrently, so a request waits in the server's queue behind other prefills before
|
|
39
|
+
its own runs. Every request is bounded on connect, read, write, and pool, so a wedged
|
|
40
|
+
connection raises `PromptLogprobTimeoutError` instead of hanging a run forever.
|
|
41
|
+
"""
|
|
42
|
+
|
|
43
|
+
from __future__ import annotations
|
|
44
|
+
|
|
45
|
+
import logging
|
|
46
|
+
import math
|
|
47
|
+
import threading
|
|
48
|
+
import time
|
|
49
|
+
from collections.abc import Callable
|
|
50
|
+
|
|
51
|
+
import httpx
|
|
52
|
+
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
|
53
|
+
|
|
54
|
+
from wmo.providers.base import ProviderKind, VerifyResult
|
|
55
|
+
|
|
56
|
+
logger = logging.getLogger(__name__)
|
|
57
|
+
|
|
58
|
+
DEFAULT_TIMEOUT_S = 1200.0
|
|
59
|
+
"""Per-request wall-clock bound; see the module docstring for the sizing."""
|
|
60
|
+
|
|
61
|
+
DEFAULT_MAX_ATTEMPTS = 3
|
|
62
|
+
"""Attempts per `score` call, counting the first; only transient failures retry."""
|
|
63
|
+
|
|
64
|
+
_RETRY_BASE_DELAY_S = 2.0
|
|
65
|
+
_RETRY_MAX_DELAY_S = 30.0
|
|
66
|
+
|
|
67
|
+
_RETRYABLE_STATUS = frozenset({408, 409, 425, 429, 500, 502, 503, 504})
|
|
68
|
+
"""Statuses worth another attempt: server-side capacity, restarts, cold boots."""
|
|
69
|
+
|
|
70
|
+
_COMPLETIONS_PATH = "/v1/completions"
|
|
71
|
+
|
|
72
|
+
_VERIFY_PROBE_TOKEN_IDS: tuple[int, ...] = (1, 2, 3, 4)
|
|
73
|
+
"""A tiny fixed sequence for verify(); small ids are valid in any real vocab."""
|
|
74
|
+
|
|
75
|
+
_MAX_CANDIDATES_IN_ERROR = 8
|
|
76
|
+
"""How many returned candidate ids an error message quotes before eliding."""
|
|
77
|
+
|
|
78
|
+
_MAX_BODY_CHARS_IN_ERROR = 300
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
class PromptLogprobError(RuntimeError):
|
|
82
|
+
"""A teacher scoring request failed, or its response was unusable.
|
|
83
|
+
|
|
84
|
+
Raised for HTTP failures, unparsable bodies, and shape violations (a row of
|
|
85
|
+
the wrong length, a missing realized token). The message names the endpoint
|
|
86
|
+
and the remedy, because a wrong shape here is silent data corruption
|
|
87
|
+
downstream rather than a crash.
|
|
88
|
+
"""
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
class PromptLogprobTimeoutError(PromptLogprobError, TimeoutError):
|
|
92
|
+
"""A teacher scoring request blew its wall-clock deadline.
|
|
93
|
+
|
|
94
|
+
Subclasses `TimeoutError` and keeps "timed out" in the message so retry
|
|
95
|
+
layers classify it as transient capacity, matching
|
|
96
|
+
`wmo.distill.deadlines.TinkerDeadlineError`'s contract.
|
|
97
|
+
"""
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
class _LogprobEntry(BaseModel):
|
|
101
|
+
"""One candidate in a position's `prompt_logprobs` dict (rank/decoded ignored)."""
|
|
102
|
+
|
|
103
|
+
model_config = ConfigDict(extra="ignore")
|
|
104
|
+
|
|
105
|
+
logprob: float
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
class _Choice(BaseModel):
|
|
109
|
+
"""The one completion choice, which carries `prompt_logprobs` on some versions."""
|
|
110
|
+
|
|
111
|
+
model_config = ConfigDict(extra="ignore")
|
|
112
|
+
|
|
113
|
+
prompt_logprobs: list[dict[str, _LogprobEntry] | None] | None = None
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
class _CompletionsResponse(BaseModel):
|
|
117
|
+
"""A `/v1/completions` response, tolerant of where `prompt_logprobs` lands."""
|
|
118
|
+
|
|
119
|
+
model_config = ConfigDict(extra="ignore")
|
|
120
|
+
|
|
121
|
+
prompt_logprobs: list[dict[str, _LogprobEntry] | None] | None = None
|
|
122
|
+
choices: list[_Choice] = Field(default_factory=list)
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
class _CompletionsRequest(BaseModel):
|
|
126
|
+
"""The scoring request body: a token-id prompt, one throwaway sampled token."""
|
|
127
|
+
|
|
128
|
+
model_config = ConfigDict(extra="forbid")
|
|
129
|
+
|
|
130
|
+
model: str
|
|
131
|
+
prompt: list[int]
|
|
132
|
+
max_tokens: int = 1
|
|
133
|
+
prompt_logprobs: int = 0
|
|
134
|
+
temperature: float = 0.0
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
def _completions_url(endpoint: str) -> str:
|
|
138
|
+
"""The `/v1/completions` URL for a configured endpoint.
|
|
139
|
+
|
|
140
|
+
Accepts either a bare server root (`https://host`) or a root that already
|
|
141
|
+
carries the OpenAI `/v1` prefix (the form `wmo providers` stores for
|
|
142
|
+
OpenAI-compatible servers), so a caller cannot accidentally produce
|
|
143
|
+
`/v1/v1/completions`.
|
|
144
|
+
|
|
145
|
+
Args:
|
|
146
|
+
endpoint: The teacher server base URL.
|
|
147
|
+
|
|
148
|
+
Returns:
|
|
149
|
+
The absolute URL to POST scoring requests to.
|
|
150
|
+
|
|
151
|
+
Raises:
|
|
152
|
+
ValueError: If `endpoint` is blank.
|
|
153
|
+
"""
|
|
154
|
+
base = endpoint.strip().rstrip("/")
|
|
155
|
+
if not base:
|
|
156
|
+
raise ValueError(
|
|
157
|
+
"PromptLogprobClient needs a teacher endpoint URL, for example "
|
|
158
|
+
"'https://my-vllm-host' or 'https://my-vllm-host/v1'; got an empty string"
|
|
159
|
+
)
|
|
160
|
+
if base.endswith("/v1"):
|
|
161
|
+
return base + "/completions"
|
|
162
|
+
return base + _COMPLETIONS_PATH
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
class PromptLogprobClient:
|
|
166
|
+
"""Scores exact teacher token ids on a vLLM `/v1/completions` endpoint.
|
|
167
|
+
|
|
168
|
+
One client is safe to share across threads: httpx connection pooling and the
|
|
169
|
+
usage counter are both synchronized, so a scoring pool can fan out over
|
|
170
|
+
datums against a single client (that is how the teacher is driven in a step).
|
|
171
|
+
|
|
172
|
+
Example:
|
|
173
|
+
>>> client = PromptLogprobClient("https://vllm-host", "zai-org/GLM-5.2")
|
|
174
|
+
>>> row = client.score(teacher_token_ids) # doctest: +SKIP
|
|
175
|
+
>>> row[0] is None # position 0 has no context # doctest: +SKIP
|
|
176
|
+
True
|
|
177
|
+
"""
|
|
178
|
+
|
|
179
|
+
def __init__(
|
|
180
|
+
self,
|
|
181
|
+
endpoint: str,
|
|
182
|
+
model: str,
|
|
183
|
+
*,
|
|
184
|
+
api_key: str | None = None,
|
|
185
|
+
timeout_s: float = DEFAULT_TIMEOUT_S,
|
|
186
|
+
transport: httpx.BaseTransport | None = None,
|
|
187
|
+
max_attempts: int = DEFAULT_MAX_ATTEMPTS,
|
|
188
|
+
sleep: Callable[[float], None] = time.sleep,
|
|
189
|
+
) -> None:
|
|
190
|
+
"""Build a client bound to one endpoint and one served model.
|
|
191
|
+
|
|
192
|
+
Args:
|
|
193
|
+
endpoint: Teacher server base URL, with or without a `/v1` suffix.
|
|
194
|
+
model: The model id the server serves, sent as the request's `model`.
|
|
195
|
+
api_key: Bearer token for the endpoint. The repo convention is to pass
|
|
196
|
+
`WMO_ENDPOINT_API_KEY`; this module never reads the environment
|
|
197
|
+
itself, so the real provider keys cannot leak to an arbitrary host.
|
|
198
|
+
None (the norm for a private vLLM host) sends no auth header.
|
|
199
|
+
timeout_s: Per-request wall-clock bound in seconds, applied to connect,
|
|
200
|
+
read, write, and pool waits. Defaults to `DEFAULT_TIMEOUT_S`
|
|
201
|
+
(1200s), sized for a 65k-token prefill queued behind other
|
|
202
|
+
requests. Retries multiply the worst case by `max_attempts`.
|
|
203
|
+
transport: httpx transport override. Tests pass an
|
|
204
|
+
`httpx.MockTransport` so no request ever leaves the process.
|
|
205
|
+
max_attempts: Total attempts per `score` call. Only transient failures
|
|
206
|
+
(timeouts, transport errors, retryable statuses) consume attempts.
|
|
207
|
+
sleep: Backoff sleeper, injectable so tests do not wait.
|
|
208
|
+
|
|
209
|
+
Raises:
|
|
210
|
+
ValueError: If the endpoint is blank, `timeout_s` is not a positive
|
|
211
|
+
finite number, or `max_attempts` is below 1.
|
|
212
|
+
"""
|
|
213
|
+
if not math.isfinite(timeout_s) or timeout_s <= 0:
|
|
214
|
+
raise ValueError(
|
|
215
|
+
f"timeout_s must be a positive finite number of seconds, got {timeout_s!r}; "
|
|
216
|
+
f"use the default ({DEFAULT_TIMEOUT_S:g}) unless the endpoint is known to be fast"
|
|
217
|
+
)
|
|
218
|
+
if max_attempts < 1:
|
|
219
|
+
raise ValueError(
|
|
220
|
+
f"max_attempts must be at least 1, got {max_attempts}; pass 1 to disable retries"
|
|
221
|
+
)
|
|
222
|
+
self._url = _completions_url(endpoint)
|
|
223
|
+
self._model = model
|
|
224
|
+
self._timeout_s = timeout_s
|
|
225
|
+
self._max_attempts = max_attempts
|
|
226
|
+
self._sleep = sleep
|
|
227
|
+
self._usage_tokens = 0
|
|
228
|
+
self._usage_lock = threading.Lock()
|
|
229
|
+
headers = {"Content-Type": "application/json"}
|
|
230
|
+
if api_key:
|
|
231
|
+
headers["Authorization"] = f"Bearer {api_key}"
|
|
232
|
+
self._client = httpx.Client(
|
|
233
|
+
timeout=httpx.Timeout(timeout_s),
|
|
234
|
+
transport=transport,
|
|
235
|
+
headers=headers,
|
|
236
|
+
)
|
|
237
|
+
|
|
238
|
+
@property
|
|
239
|
+
def url(self) -> str:
|
|
240
|
+
"""The absolute scoring URL this client posts to."""
|
|
241
|
+
return self._url
|
|
242
|
+
|
|
243
|
+
@property
|
|
244
|
+
def model(self) -> str:
|
|
245
|
+
"""The served model id sent with every request."""
|
|
246
|
+
return self._model
|
|
247
|
+
|
|
248
|
+
def score(self, token_ids: list[int]) -> list[float | None]:
|
|
249
|
+
"""Teacher logprobs for an exact token sequence, one entry per position.
|
|
250
|
+
|
|
251
|
+
Args:
|
|
252
|
+
token_ids: The teacher's own token ids, in the teacher's vocabulary.
|
|
253
|
+
They are sent verbatim as integers, so the returned row aligns
|
|
254
|
+
index for index with this list.
|
|
255
|
+
|
|
256
|
+
Returns:
|
|
257
|
+
A list of `len(token_ids)` entries. Entry 0 is always None (token 0
|
|
258
|
+
has no context); entry p is the teacher's logprob of `token_ids[p]`
|
|
259
|
+
given `token_ids[:p]`.
|
|
260
|
+
|
|
261
|
+
Raises:
|
|
262
|
+
ValueError: If `token_ids` is empty.
|
|
263
|
+
PromptLogprobTimeoutError: If every attempt blew the deadline.
|
|
264
|
+
PromptLogprobError: On a non-retryable HTTP status, an exhausted
|
|
265
|
+
retry budget, an unparsable body, a row length that does not
|
|
266
|
+
match the prompt, or a position whose dict lacks the realized
|
|
267
|
+
token.
|
|
268
|
+
"""
|
|
269
|
+
if not token_ids:
|
|
270
|
+
raise ValueError(
|
|
271
|
+
"score() needs at least one token id; an empty prompt has nothing to score "
|
|
272
|
+
"(filter empty spans out before scoring)"
|
|
273
|
+
)
|
|
274
|
+
return self._score(token_ids, count_usage=True)
|
|
275
|
+
|
|
276
|
+
def verify(self) -> VerifyResult:
|
|
277
|
+
"""One tiny scoring probe, reporting failure as `ok=False` instead of raising.
|
|
278
|
+
|
|
279
|
+
Mirrors `wmo.providers.base.verify_via_ping` and `TinkerTeacher.verify`, so
|
|
280
|
+
preflight can report every misconfigured backend at once. The probe's tokens
|
|
281
|
+
are excluded from `usage()`.
|
|
282
|
+
|
|
283
|
+
Returns:
|
|
284
|
+
`ok=True` when the endpoint answered with a well-formed row, otherwise
|
|
285
|
+
`ok=False` with the failure text in `detail`.
|
|
286
|
+
"""
|
|
287
|
+
try:
|
|
288
|
+
self._score(list(_VERIFY_PROBE_TOKEN_IDS), count_usage=False)
|
|
289
|
+
except Exception as exc: # noqa: BLE001 - verify reports failure, never raises
|
|
290
|
+
return VerifyResult(
|
|
291
|
+
ok=False, kind=ProviderKind.OPENAI, model=self._model, detail=str(exc)
|
|
292
|
+
)
|
|
293
|
+
return VerifyResult(ok=True, kind=ProviderKind.OPENAI, model=self._model, detail=self._url)
|
|
294
|
+
|
|
295
|
+
def usage(self) -> int:
|
|
296
|
+
"""Cumulative teacher tokens submitted for scoring (verify probes excluded).
|
|
297
|
+
|
|
298
|
+
Counts every dispatched attempt, not only successful ones: a request that
|
|
299
|
+
timed out or died mid-response has usually already run its prefill on the
|
|
300
|
+
server, so the work is real. This matches `TinkerTeacher.usage`'s
|
|
301
|
+
"submitted, not billed-on-success" contract and feeds the same
|
|
302
|
+
teacher_prefill meter.
|
|
303
|
+
"""
|
|
304
|
+
with self._usage_lock:
|
|
305
|
+
return self._usage_tokens
|
|
306
|
+
|
|
307
|
+
def close(self) -> None:
|
|
308
|
+
"""Close the underlying connection pool."""
|
|
309
|
+
self._client.close()
|
|
310
|
+
|
|
311
|
+
def __enter__(self) -> PromptLogprobClient:
|
|
312
|
+
return self
|
|
313
|
+
|
|
314
|
+
def __exit__(self, *exc_info: object) -> None:
|
|
315
|
+
self.close()
|
|
316
|
+
|
|
317
|
+
def _score(self, token_ids: list[int], *, count_usage: bool) -> list[float | None]:
|
|
318
|
+
body = _CompletionsRequest(model=self._model, prompt=token_ids).model_dump()
|
|
319
|
+
response = self._post(body, token_count=len(token_ids) if count_usage else 0)
|
|
320
|
+
rows = self._prompt_logprob_rows(response, expected=len(token_ids))
|
|
321
|
+
return self._realized_row(rows, token_ids)
|
|
322
|
+
|
|
323
|
+
def _post(self, body: dict[str, object], *, token_count: int) -> _CompletionsResponse:
|
|
324
|
+
"""POST one scoring request, retrying only transient failures."""
|
|
325
|
+
last_error: PromptLogprobError | None = None
|
|
326
|
+
for attempt in range(1, self._max_attempts + 1):
|
|
327
|
+
if token_count:
|
|
328
|
+
with self._usage_lock:
|
|
329
|
+
self._usage_tokens += token_count
|
|
330
|
+
try:
|
|
331
|
+
response = self._client.post(self._url, json=body)
|
|
332
|
+
except httpx.TimeoutException as exc:
|
|
333
|
+
last_error = PromptLogprobTimeoutError(
|
|
334
|
+
f"teacher scoring timed out after {self._timeout_s:g}s against "
|
|
335
|
+
f"{self._url} (attempt {attempt}/{self._max_attempts}): {exc}. Raise "
|
|
336
|
+
"timeout_s, lower the number of concurrent scoring calls, or check the "
|
|
337
|
+
"endpoint is up and not stuck in a cold boot"
|
|
338
|
+
)
|
|
339
|
+
except httpx.TransportError as exc:
|
|
340
|
+
last_error = PromptLogprobError(
|
|
341
|
+
f"teacher scoring could not reach {self._url} "
|
|
342
|
+
f"(attempt {attempt}/{self._max_attempts}): {exc!r}. Check the endpoint URL "
|
|
343
|
+
"and that the vLLM server is running and reachable from this host"
|
|
344
|
+
)
|
|
345
|
+
else:
|
|
346
|
+
if response.status_code < 400:
|
|
347
|
+
return self._parse(response)
|
|
348
|
+
error = self._status_error(response, attempt)
|
|
349
|
+
if response.status_code not in _RETRYABLE_STATUS:
|
|
350
|
+
raise error
|
|
351
|
+
last_error = error
|
|
352
|
+
if attempt < self._max_attempts:
|
|
353
|
+
delay = min(_RETRY_BASE_DELAY_S * 2 ** (attempt - 1), _RETRY_MAX_DELAY_S)
|
|
354
|
+
logger.warning(
|
|
355
|
+
"teacher scoring attempt %d/%d failed (%s); retrying in %.0fs",
|
|
356
|
+
attempt,
|
|
357
|
+
self._max_attempts,
|
|
358
|
+
last_error,
|
|
359
|
+
delay,
|
|
360
|
+
)
|
|
361
|
+
self._sleep(delay)
|
|
362
|
+
assert last_error is not None # noqa: S101 - the loop runs at least once
|
|
363
|
+
raise last_error
|
|
364
|
+
|
|
365
|
+
def _status_error(self, response: httpx.Response, attempt: int) -> PromptLogprobError:
|
|
366
|
+
"""A typed error for a failing HTTP status, with a status-specific remedy."""
|
|
367
|
+
body = response.text[:_MAX_BODY_CHARS_IN_ERROR]
|
|
368
|
+
if response.status_code in (401, 403):
|
|
369
|
+
remedy = (
|
|
370
|
+
"the endpoint rejected the credentials: pass the api_key this server expects "
|
|
371
|
+
"(the repo convention is the WMO_ENDPOINT_API_KEY env var, read by the caller)"
|
|
372
|
+
)
|
|
373
|
+
elif response.status_code == 404:
|
|
374
|
+
remedy = (
|
|
375
|
+
f"no such route or model: check the base URL and that the server serves model "
|
|
376
|
+
f"{self._model!r} (GET /v1/models lists it)"
|
|
377
|
+
)
|
|
378
|
+
elif response.status_code in _RETRYABLE_STATUS:
|
|
379
|
+
remedy = (
|
|
380
|
+
"the server failed transiently (capacity, restart, or cold boot); retries are "
|
|
381
|
+
"exhausted, so lower scoring concurrency or check the server logs"
|
|
382
|
+
)
|
|
383
|
+
else:
|
|
384
|
+
remedy = (
|
|
385
|
+
"the server rejected the request: confirm this vLLM build supports "
|
|
386
|
+
"prompt_logprobs on /v1/completions and that no token id exceeds its vocab"
|
|
387
|
+
)
|
|
388
|
+
return PromptLogprobError(
|
|
389
|
+
f"teacher scoring got HTTP {response.status_code} from {self._url} "
|
|
390
|
+
f"(attempt {attempt}/{self._max_attempts}): {body!r}. {remedy}"
|
|
391
|
+
)
|
|
392
|
+
|
|
393
|
+
def _parse(self, response: httpx.Response) -> _CompletionsResponse:
|
|
394
|
+
"""Parse a 2xx body, turning malformed JSON into a typed error."""
|
|
395
|
+
try:
|
|
396
|
+
payload = response.json()
|
|
397
|
+
except ValueError as exc:
|
|
398
|
+
raise PromptLogprobError(
|
|
399
|
+
f"teacher scoring got a non-JSON response from {self._url}: "
|
|
400
|
+
f"{response.text[:_MAX_BODY_CHARS_IN_ERROR]!r}. Check the URL points at a vLLM "
|
|
401
|
+
"OpenAI server and not at a proxy or web page"
|
|
402
|
+
) from exc
|
|
403
|
+
try:
|
|
404
|
+
return _CompletionsResponse.model_validate(payload)
|
|
405
|
+
except ValidationError as exc:
|
|
406
|
+
raise PromptLogprobError(
|
|
407
|
+
f"teacher scoring could not read the response from {self._url}: {exc}. Each "
|
|
408
|
+
"prompt_logprobs position must be null or a dict of token id to an object with "
|
|
409
|
+
"a 'logprob'; check the vLLM version's response format"
|
|
410
|
+
) from exc
|
|
411
|
+
|
|
412
|
+
def _prompt_logprob_rows(
|
|
413
|
+
self, response: _CompletionsResponse, *, expected: int
|
|
414
|
+
) -> list[dict[str, _LogprobEntry] | None]:
|
|
415
|
+
"""The per-position candidate dicts, from either shape, length-validated."""
|
|
416
|
+
rows = response.prompt_logprobs
|
|
417
|
+
if rows is None and response.choices:
|
|
418
|
+
rows = response.choices[0].prompt_logprobs
|
|
419
|
+
if rows is None:
|
|
420
|
+
raise PromptLogprobError(
|
|
421
|
+
f"teacher scoring response from {self._url} carried no prompt_logprobs (neither "
|
|
422
|
+
"at the top level nor under choices[0]). The request must go to /v1/completions "
|
|
423
|
+
"with prompt_logprobs set: /v1/chat/completions supports neither prompt_logprobs "
|
|
424
|
+
"nor echo, and older servers may not support it at all"
|
|
425
|
+
)
|
|
426
|
+
if len(rows) != expected:
|
|
427
|
+
raise PromptLogprobError(
|
|
428
|
+
f"teacher scoring returned {len(rows)} prompt_logprobs entries for a "
|
|
429
|
+
f"{expected}-token prompt at {self._url} (model {self._model}). Every position "
|
|
430
|
+
"must be scored, since a short or long row silently corrupts every downstream "
|
|
431
|
+
"span sum. Send the prompt as a list[int] (a text prompt is re-tokenized "
|
|
432
|
+
"server-side, which shifts every position), and confirm the server's "
|
|
433
|
+
f"max_model_len covers {expected} tokens"
|
|
434
|
+
)
|
|
435
|
+
return rows
|
|
436
|
+
|
|
437
|
+
def _realized_row(
|
|
438
|
+
self,
|
|
439
|
+
rows: list[dict[str, _LogprobEntry] | None],
|
|
440
|
+
token_ids: list[int],
|
|
441
|
+
) -> list[float | None]:
|
|
442
|
+
"""Pull each position's REALIZED token logprob, never the argmax."""
|
|
443
|
+
row: list[float | None] = [None]
|
|
444
|
+
for index in range(1, len(token_ids)):
|
|
445
|
+
candidates = rows[index]
|
|
446
|
+
token_id = token_ids[index]
|
|
447
|
+
entry = candidates.get(str(token_id)) if candidates else None
|
|
448
|
+
if entry is None:
|
|
449
|
+
raise PromptLogprobError(
|
|
450
|
+
f"teacher scoring returned no logprob for the realized token {token_id} at "
|
|
451
|
+
f"position {index} of {len(token_ids)} from {self._url} "
|
|
452
|
+
f"(candidates: {_describe_candidates(candidates)}). The realized token is "
|
|
453
|
+
"always included when the prompt is sent as token ids with prompt_logprobs "
|
|
454
|
+
f"set, so this means the ids are not in {self._model}'s vocabulary or the "
|
|
455
|
+
"prompt was re-tokenized server-side"
|
|
456
|
+
)
|
|
457
|
+
row.append(entry.logprob)
|
|
458
|
+
logger.debug(
|
|
459
|
+
"teacher scored %d position(s) at %s, %d tokens submitted so far",
|
|
460
|
+
len(token_ids),
|
|
461
|
+
self._url,
|
|
462
|
+
self.usage(),
|
|
463
|
+
)
|
|
464
|
+
return row
|
|
465
|
+
|
|
466
|
+
|
|
467
|
+
def _describe_candidates(candidates: dict[str, _LogprobEntry] | None) -> str:
|
|
468
|
+
"""A short, deterministic rendering of a position's returned candidate ids."""
|
|
469
|
+
if not candidates:
|
|
470
|
+
return "none (the position was null or empty)"
|
|
471
|
+
keys = sorted(candidates)
|
|
472
|
+
shown = ", ".join(keys[:_MAX_CANDIDATES_IN_ERROR])
|
|
473
|
+
if len(keys) > _MAX_CANDIDATES_IN_ERROR:
|
|
474
|
+
return f"{len(keys)} returned, first ids {shown}, ..."
|
|
475
|
+
return shown
|