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,173 @@
|
|
|
1
|
+
"""Retrieval over the trace replay buffer (DreamGym Eq. 4).
|
|
2
|
+
|
|
3
|
+
At each step the world model retrieves the top-k past steps whose (state, action) is most similar to
|
|
4
|
+
the current one, by cosine similarity of an embedding `phi`:
|
|
5
|
+
|
|
6
|
+
{d_j} = Topk( cos( phi(s_t, a_t), phi(s_i, a_i) ) )
|
|
7
|
+
|
|
8
|
+
The buffer is initialized offline from ingested traces (`index`) and enriched online as the agent
|
|
9
|
+
steps (`add`).
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
import json
|
|
15
|
+
from pathlib import Path
|
|
16
|
+
from typing import Literal, Protocol, runtime_checkable
|
|
17
|
+
|
|
18
|
+
import numpy as np
|
|
19
|
+
from numpy.typing import NDArray
|
|
20
|
+
|
|
21
|
+
from wmo.core.render import encode_action, encode_state_action
|
|
22
|
+
from wmo.core.types import Action, EnvState, Observation, Step, Trace
|
|
23
|
+
from wmo.providers.base import Embedder
|
|
24
|
+
|
|
25
|
+
# What text phi embeds per step: the full (state, action) summary, or the command-only action.
|
|
26
|
+
RetrievalKey = Literal["state_action", "action"]
|
|
27
|
+
|
|
28
|
+
# A placeholder observation for query-only encoding: topk embeds (state, action), never the result.
|
|
29
|
+
_EMPTY_OBS = Observation(content="")
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@runtime_checkable
|
|
33
|
+
class Retriever(Protocol):
|
|
34
|
+
def index(self, traces: list[Trace]) -> None:
|
|
35
|
+
"""Build phase: embed every step's (state, action) and store it in the buffer."""
|
|
36
|
+
...
|
|
37
|
+
|
|
38
|
+
def topk(self, state: EnvState, action: Action, k: int) -> list[Step]:
|
|
39
|
+
"""Runtime: return the k most similar prior steps to (state, action)."""
|
|
40
|
+
...
|
|
41
|
+
|
|
42
|
+
def add(self, step: Step) -> None:
|
|
43
|
+
"""Online enrichment: add a freshly generated step to the buffer."""
|
|
44
|
+
...
|
|
45
|
+
|
|
46
|
+
def sample(self, n: int) -> list[Step]:
|
|
47
|
+
"""Return up to `n` steps from the buffer (e.g. to seed `wmo play` action suggestions)."""
|
|
48
|
+
...
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class EmbeddingRetriever:
|
|
52
|
+
"""Default Retriever: dense cosine similarity using a provider's embedding model.
|
|
53
|
+
|
|
54
|
+
The replay buffer is an in-memory embedding matrix (rows = steps) kept parallel to a
|
|
55
|
+
``list[Step]``. ``index`` embeds the whole corpus in one batched ``provider.embed`` call;
|
|
56
|
+
``add`` embeds a single step for online enrichment. ``topk`` ranks by cosine similarity,
|
|
57
|
+
matching DreamGym Eq. 4's ``Topk(cos(phi(s_t,a_t), phi(s_i,a_i)))``.
|
|
58
|
+
"""
|
|
59
|
+
|
|
60
|
+
def __init__(self, provider: Embedder, *, key_mode: RetrievalKey = "state_action") -> None:
|
|
61
|
+
self._provider = provider
|
|
62
|
+
# What text phi embeds per step: "state_action" (the full (state, action) summary, default)
|
|
63
|
+
# or "action" (command-only — no STATE/ACTION scaffolding, concentrating the signal for
|
|
64
|
+
# stateless traces). Index and query use the SAME mode, so the buffer stays self-consistent.
|
|
65
|
+
if key_mode not in ("state_action", "action"):
|
|
66
|
+
raise ValueError(f"key_mode must be 'state_action' or 'action', got {key_mode!r}")
|
|
67
|
+
self._key_mode = key_mode
|
|
68
|
+
# Parallel structures: row i of `_matrix` is the embedding of `_steps[i]`.
|
|
69
|
+
self._steps: list[Step] = []
|
|
70
|
+
self._matrix: NDArray[np.float64] | None = None
|
|
71
|
+
|
|
72
|
+
def _key_text(self, step: Step) -> str:
|
|
73
|
+
if self._key_mode == "action":
|
|
74
|
+
return encode_action(step.action)
|
|
75
|
+
return encode_state_action(step.state_before, step.action)
|
|
76
|
+
|
|
77
|
+
def _embed_steps(self, steps: list[Step]) -> NDArray[np.float64]:
|
|
78
|
+
# phi embeds the canonical step text from wmo.core.render — the same text the engine and
|
|
79
|
+
# GEPA render, so an embedded step and a shown demo match. `key_mode` selects which text.
|
|
80
|
+
texts = [self._key_text(s) for s in steps]
|
|
81
|
+
vectors = self._provider.embed(texts)
|
|
82
|
+
return np.asarray(vectors, dtype=np.float64)
|
|
83
|
+
|
|
84
|
+
def index(self, traces: list[Trace]) -> None:
|
|
85
|
+
"""Embed every step of every trace and (re)build the buffer from scratch."""
|
|
86
|
+
steps = [step for trace in traces for step in trace.steps]
|
|
87
|
+
self._steps = steps
|
|
88
|
+
if not steps:
|
|
89
|
+
self._matrix = None
|
|
90
|
+
return
|
|
91
|
+
self._matrix = self._embed_steps(steps)
|
|
92
|
+
|
|
93
|
+
def topk(self, state: EnvState, action: Action, k: int) -> list[Step]:
|
|
94
|
+
"""Return the up-to-k most similar prior steps by cosine similarity."""
|
|
95
|
+
if k <= 0 or self._matrix is None or not self._steps:
|
|
96
|
+
return []
|
|
97
|
+
query = self._embed_steps(
|
|
98
|
+
[Step(action=action, observation=_EMPTY_OBS, state_before=state)]
|
|
99
|
+
)[0]
|
|
100
|
+
if query.shape[0] != self._matrix.shape[1]:
|
|
101
|
+
raise ValueError(
|
|
102
|
+
f"embedder produces dim {query.shape[0]} but the indexed buffer has dim "
|
|
103
|
+
f"{self._matrix.shape[1]}; load the same embedder (embed_dim) used at build time"
|
|
104
|
+
)
|
|
105
|
+
scores = _cosine(query, self._matrix)
|
|
106
|
+
# argsort ascending, take the tail, reverse for descending-similarity order.
|
|
107
|
+
count = min(k, len(self._steps))
|
|
108
|
+
top = np.argsort(scores)[-count:][::-1]
|
|
109
|
+
return [self._steps[int(i)] for i in top]
|
|
110
|
+
|
|
111
|
+
def add(self, step: Step) -> None:
|
|
112
|
+
"""Append a freshly generated step to the buffer for online enrichment."""
|
|
113
|
+
vector = self._embed_steps([step])
|
|
114
|
+
self._steps.append(step)
|
|
115
|
+
if self._matrix is None:
|
|
116
|
+
self._matrix = vector
|
|
117
|
+
else:
|
|
118
|
+
self._matrix = np.vstack([self._matrix, vector])
|
|
119
|
+
|
|
120
|
+
def sample(self, n: int) -> list[Step]:
|
|
121
|
+
"""Return the first up-to-`n` steps from the buffer (deterministic; no RNG needed)."""
|
|
122
|
+
return self._steps[: max(0, n)]
|
|
123
|
+
|
|
124
|
+
def save(self, index_dir: str | Path) -> None:
|
|
125
|
+
"""Persist the buffer (embedding matrix + parallel steps) under `index_dir`.
|
|
126
|
+
|
|
127
|
+
`wmo build` writes this; `wmo serve` / `WorldModel.load` reloads it without re-embedding.
|
|
128
|
+
"""
|
|
129
|
+
path = Path(index_dir)
|
|
130
|
+
path.mkdir(parents=True, exist_ok=True)
|
|
131
|
+
matrix = self._matrix if self._matrix is not None else np.empty((0, 0), dtype=np.float64)
|
|
132
|
+
np.save(path / _MATRIX_FILE, matrix)
|
|
133
|
+
with (path / _STEPS_FILE).open("w", encoding="utf-8") as fh:
|
|
134
|
+
for step in self._steps:
|
|
135
|
+
fh.write(step.model_dump_json() + "\n")
|
|
136
|
+
# Persist key_mode: the matrix was embedded from this mode's key text, so a reload MUST
|
|
137
|
+
# query in the same mode or it cosine-compares mismatched embedding spaces (no dim error,
|
|
138
|
+
# just near-random neighbours). Without this, a reloaded index reverts to state_action.
|
|
139
|
+
(path / _META_FILE).write_text(json.dumps({"key_mode": self._key_mode}), encoding="utf-8")
|
|
140
|
+
|
|
141
|
+
def load(self, index_dir: str | Path) -> None:
|
|
142
|
+
"""Reload a buffer previously written by `save`, replacing any current contents."""
|
|
143
|
+
path = Path(index_dir)
|
|
144
|
+
matrix = np.load(path / _MATRIX_FILE)
|
|
145
|
+
steps = [
|
|
146
|
+
Step.model_validate_json(line)
|
|
147
|
+
for line in (path / _STEPS_FILE).read_text(encoding="utf-8").splitlines()
|
|
148
|
+
if line.strip()
|
|
149
|
+
]
|
|
150
|
+
self._steps = steps
|
|
151
|
+
self._matrix = matrix if matrix.size and steps else None
|
|
152
|
+
# Restore the mode the matrix was built with (older indexes predate meta.json -> default).
|
|
153
|
+
meta_path = path / _META_FILE
|
|
154
|
+
if meta_path.exists():
|
|
155
|
+
mode = json.loads(meta_path.read_text(encoding="utf-8")).get("key_mode", "state_action")
|
|
156
|
+
if mode not in ("state_action", "action"):
|
|
157
|
+
raise ValueError(f"index meta has invalid key_mode {mode!r}")
|
|
158
|
+
self._key_mode = mode
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
_MATRIX_FILE = "embeddings.npy"
|
|
162
|
+
_STEPS_FILE = "steps.jsonl"
|
|
163
|
+
_META_FILE = "meta.json"
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
def _cosine(query: NDArray[np.float64], matrix: NDArray[np.float64]) -> NDArray[np.float64]:
|
|
167
|
+
"""Cosine similarity of `query` against each row of `matrix`. Zero vectors score 0."""
|
|
168
|
+
query_norm = float(np.linalg.norm(query))
|
|
169
|
+
row_norms = np.linalg.norm(matrix, axis=1)
|
|
170
|
+
denom = row_norms * query_norm
|
|
171
|
+
dots = matrix @ query
|
|
172
|
+
# Avoid divide-by-zero: where either vector is zero, similarity is 0.
|
|
173
|
+
return np.divide(dots, denom, out=np.zeros_like(dots), where=denom > 0)
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
"""Scenario-set construction: distill a trace corpus into a representative eval scenario set.
|
|
2
|
+
|
|
3
|
+
The pipeline (Clio-style facets -> embed -> cluster -> select -> synthesize -> verify), organized
|
|
4
|
+
as one subpackage per stage — `mining/`, `synthesis/`, `verification/` — with `builder` on top:
|
|
5
|
+
|
|
6
|
+
facets = FacetExtractor(provider).extract_all(traces)
|
|
7
|
+
scenario_set = build_scenario_set(traces, facets, provider, embedder, config)
|
|
8
|
+
verdicts = verify_scenarios(scenario_set, traces, world_model, agent, judge_provider)
|
|
9
|
+
|
|
10
|
+
Exposed via `wmo scenarios build` / `wmo scenarios verify` on the CLI.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from wmo.scenarios.builder import ScenarioBuildConfig, build_scenario_set
|
|
14
|
+
from wmo.scenarios.mining import (
|
|
15
|
+
FacetExtractor,
|
|
16
|
+
Outcome,
|
|
17
|
+
SelectedTrace,
|
|
18
|
+
TraceCluster,
|
|
19
|
+
TraceFacet,
|
|
20
|
+
cluster_facets,
|
|
21
|
+
hybrid_select,
|
|
22
|
+
name_clusters,
|
|
23
|
+
semdedup_keep,
|
|
24
|
+
tool_signature,
|
|
25
|
+
trace_digest,
|
|
26
|
+
)
|
|
27
|
+
from wmo.scenarios.synthesis import EvalScenario, ScenarioSet, ScenarioSynthesizer
|
|
28
|
+
from wmo.scenarios.verification import (
|
|
29
|
+
ChecklistJudge,
|
|
30
|
+
ChecklistResult,
|
|
31
|
+
ScenarioVerdict,
|
|
32
|
+
VerificationReport,
|
|
33
|
+
verify_scenarios,
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
__all__ = [
|
|
37
|
+
"ChecklistJudge",
|
|
38
|
+
"ChecklistResult",
|
|
39
|
+
"EvalScenario",
|
|
40
|
+
"FacetExtractor",
|
|
41
|
+
"Outcome",
|
|
42
|
+
"ScenarioBuildConfig",
|
|
43
|
+
"ScenarioSet",
|
|
44
|
+
"ScenarioSynthesizer",
|
|
45
|
+
"ScenarioVerdict",
|
|
46
|
+
"SelectedTrace",
|
|
47
|
+
"TraceCluster",
|
|
48
|
+
"TraceFacet",
|
|
49
|
+
"VerificationReport",
|
|
50
|
+
"build_scenario_set",
|
|
51
|
+
"cluster_facets",
|
|
52
|
+
"hybrid_select",
|
|
53
|
+
"name_clusters",
|
|
54
|
+
"semdedup_keep",
|
|
55
|
+
"tool_signature",
|
|
56
|
+
"trace_digest",
|
|
57
|
+
"verify_scenarios",
|
|
58
|
+
]
|
wmo/scenarios/builder.py
ADDED
|
@@ -0,0 +1,152 @@
|
|
|
1
|
+
"""The scenario-set build pipeline behind `wmo scenarios build`.
|
|
2
|
+
|
|
3
|
+
facets -> embed -> cluster -> name -> select -> synthesize -> coverage. One entry point,
|
|
4
|
+
`build_scenario_set`, that takes already-extracted facets so callers (research runs, tests) can
|
|
5
|
+
cache or substitute them; `wmo scenarios build` extracts them fresh.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
11
|
+
|
|
12
|
+
import numpy as np
|
|
13
|
+
from pydantic import BaseModel
|
|
14
|
+
|
|
15
|
+
from wmo.core.types import Trace
|
|
16
|
+
from wmo.providers.base import Embedder, Provider
|
|
17
|
+
from wmo.scenarios.mining.clustering import cluster_facets, name_clusters, normalize_rows
|
|
18
|
+
from wmo.scenarios.mining.facets import TraceFacet
|
|
19
|
+
from wmo.scenarios.mining.selection import (
|
|
20
|
+
DEDUP_THRESHOLD,
|
|
21
|
+
PROPORTIONAL_FRACTION,
|
|
22
|
+
SelectedTrace,
|
|
23
|
+
hybrid_select,
|
|
24
|
+
)
|
|
25
|
+
from wmo.scenarios.synthesis import EvalScenario, ScenarioSet, ScenarioSynthesizer
|
|
26
|
+
from wmo.scenarios.verification import ChecklistJudge
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class ScenarioBuildConfig(BaseModel):
|
|
30
|
+
"""Knobs for one scenario-set build."""
|
|
31
|
+
|
|
32
|
+
budget: int = 20 # scenarios to construct
|
|
33
|
+
k: int | None = None # cluster count; default sqrt(n)
|
|
34
|
+
seed: int = 0
|
|
35
|
+
validate_checklists: bool = True # back-agreement gate inside the build (drop on repeat fail)
|
|
36
|
+
dedup_threshold: float = DEDUP_THRESHOLD
|
|
37
|
+
proportional_fraction: float = PROPORTIONAL_FRACTION
|
|
38
|
+
coverage_tau: float = 0.7 # facet counts as covered when cosine-within-tau of a selection
|
|
39
|
+
# LLM calls in the build (facet extraction, synthesis + back-agreement) are independent
|
|
40
|
+
# per trace/selection: run them on a small thread pool, order-preserving. 1 = sequential.
|
|
41
|
+
concurrency: int = 8
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def build_scenario_set(
|
|
45
|
+
traces: list[Trace],
|
|
46
|
+
facets: list[TraceFacet],
|
|
47
|
+
provider: Provider,
|
|
48
|
+
embedder: Embedder,
|
|
49
|
+
config: ScenarioBuildConfig,
|
|
50
|
+
*,
|
|
51
|
+
judge_provider: Provider | None = None,
|
|
52
|
+
) -> ScenarioSet:
|
|
53
|
+
"""Construct a representative scenario set from a facet-annotated trace corpus.
|
|
54
|
+
|
|
55
|
+
`provider` drives cluster naming and scenario synthesis; `embedder` embeds facet summaries.
|
|
56
|
+
`judge_provider` backs the inline checklist validation (defaults to `provider`) — pass a
|
|
57
|
+
different model to keep synthesis and validation families separate. Raises when traces/facets
|
|
58
|
+
are empty or misaligned.
|
|
59
|
+
"""
|
|
60
|
+
if not traces or not facets:
|
|
61
|
+
raise ValueError("need a non-empty trace corpus and facets to build a scenario set")
|
|
62
|
+
if len(traces) != len(facets):
|
|
63
|
+
raise ValueError(f"{len(traces)} traces but {len(facets)} facets")
|
|
64
|
+
|
|
65
|
+
embeddings = np.asarray(embedder.embed([facet.embed_text() for facet in facets]))
|
|
66
|
+
labels, clusters = cluster_facets(facets, embeddings, k=config.k, seed=config.seed)
|
|
67
|
+
name_clusters(provider, clusters, facets)
|
|
68
|
+
|
|
69
|
+
selections = hybrid_select(
|
|
70
|
+
facets,
|
|
71
|
+
embeddings,
|
|
72
|
+
labels,
|
|
73
|
+
config.budget,
|
|
74
|
+
proportional_fraction=config.proportional_fraction,
|
|
75
|
+
dedup_threshold=config.dedup_threshold,
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
traces_by_id = {trace.trace_id: trace for trace in traces}
|
|
79
|
+
facets_by_id = {facet.trace_id: facet for facet in facets}
|
|
80
|
+
cluster_names = {cluster.cluster_id: cluster.name for cluster in clusters}
|
|
81
|
+
synthesizer = ScenarioSynthesizer(provider)
|
|
82
|
+
judge = ChecklistJudge(judge_provider or provider) if config.validate_checklists else None
|
|
83
|
+
|
|
84
|
+
def _synthesize_one(selection: SelectedTrace) -> EvalScenario | None:
|
|
85
|
+
source = traces_by_id[selection.trace_id]
|
|
86
|
+
scenario = synthesizer.synthesize(source, facets_by_id[selection.trace_id])
|
|
87
|
+
if judge is not None:
|
|
88
|
+
# A generated checklist must correctly grade the very episode it was distilled
|
|
89
|
+
# from; one that misgrades its own source can't be trusted on new trajectories.
|
|
90
|
+
# One regeneration retry, then drop — an invalid scenario never leaves the build.
|
|
91
|
+
if not _checklist_agrees(judge, scenario, source):
|
|
92
|
+
scenario = synthesizer.synthesize(source, facets_by_id[selection.trace_id])
|
|
93
|
+
if not _checklist_agrees(judge, scenario, source):
|
|
94
|
+
return None
|
|
95
|
+
scenario.cluster_name = cluster_names.get(selection.cluster_id, "")
|
|
96
|
+
scenario.weight = selection.weight
|
|
97
|
+
if selection.pinned_failure is not None:
|
|
98
|
+
scenario.failure_category = selection.pinned_failure
|
|
99
|
+
return scenario
|
|
100
|
+
|
|
101
|
+
# Selections are independent (synthesis + back-agreement are per-trace LLM round trips), so
|
|
102
|
+
# they run on a small thread pool; `pool.map` preserves selection order, keeping the built
|
|
103
|
+
# set (and its weight renormalization) deterministic. concurrency=1 is the sequential loop.
|
|
104
|
+
if config.concurrency > 1 and len(selections) > 1:
|
|
105
|
+
with ThreadPoolExecutor(max_workers=min(config.concurrency, len(selections))) as pool:
|
|
106
|
+
maybe_scenarios = list(pool.map(_synthesize_one, selections))
|
|
107
|
+
else:
|
|
108
|
+
maybe_scenarios = [_synthesize_one(selection) for selection in selections]
|
|
109
|
+
scenarios = [scenario for scenario in maybe_scenarios if scenario is not None]
|
|
110
|
+
dropped = sum(1 for scenario in maybe_scenarios if scenario is None)
|
|
111
|
+
|
|
112
|
+
selected_ids = {scenario.provenance[0] for scenario in scenarios}
|
|
113
|
+
coverage = _corpus_coverage(facets, embeddings, selected_ids, tau=config.coverage_tau)
|
|
114
|
+
total_weight = sum(scenario.weight for scenario in scenarios)
|
|
115
|
+
if dropped and total_weight > 0: # dropped scenarios must not leave weights summing < 1
|
|
116
|
+
for scenario in scenarios:
|
|
117
|
+
scenario.weight /= total_weight
|
|
118
|
+
return ScenarioSet(
|
|
119
|
+
scenarios=scenarios,
|
|
120
|
+
clusters=clusters,
|
|
121
|
+
corpus_traces=len(traces),
|
|
122
|
+
corpus_coverage=coverage,
|
|
123
|
+
coverage_tau=config.coverage_tau,
|
|
124
|
+
)
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def _checklist_agrees(judge: ChecklistJudge, scenario: EvalScenario, source: Trace) -> bool:
|
|
128
|
+
"""Back-agreement: the judge's verdict on the SOURCE trajectory must match its recorded
|
|
129
|
+
outcome. Traces without a recorded outcome can't disagree, so they pass."""
|
|
130
|
+
if not scenario.checklist:
|
|
131
|
+
return False
|
|
132
|
+
reward = source.metadata.get("reward")
|
|
133
|
+
if not isinstance(reward, int | float):
|
|
134
|
+
return True
|
|
135
|
+
verdict = judge.score(scenario.task, scenario.checklist, source.steps)
|
|
136
|
+
return verdict.success == (float(reward) >= 1.0)
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
def _corpus_coverage(
|
|
140
|
+
facets: list[TraceFacet],
|
|
141
|
+
embeddings: np.ndarray,
|
|
142
|
+
selected_ids: set[str],
|
|
143
|
+
*,
|
|
144
|
+
tau: float,
|
|
145
|
+
) -> float:
|
|
146
|
+
"""Fraction of corpus facets within cosine `tau` of at least one selected facet."""
|
|
147
|
+
selected_rows = [i for i, facet in enumerate(facets) if facet.trace_id in selected_ids]
|
|
148
|
+
if not selected_rows:
|
|
149
|
+
return 0.0
|
|
150
|
+
matrix = normalize_rows(embeddings)
|
|
151
|
+
similarities = matrix @ matrix[np.asarray(selected_rows)].T
|
|
152
|
+
return float((similarities.max(axis=1) >= tau).mean())
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
"""Mining: reduce raw traces to facets, cluster them, and select representative source traces."""
|
|
2
|
+
|
|
3
|
+
from wmo.scenarios.mining.clustering import TraceCluster, cluster_facets, name_clusters
|
|
4
|
+
from wmo.scenarios.mining.facets import (
|
|
5
|
+
FacetExtractor,
|
|
6
|
+
Outcome,
|
|
7
|
+
TraceFacet,
|
|
8
|
+
tool_signature,
|
|
9
|
+
trace_digest,
|
|
10
|
+
trace_domain,
|
|
11
|
+
)
|
|
12
|
+
from wmo.scenarios.mining.selection import SelectedTrace, hybrid_select, semdedup_keep
|
|
13
|
+
|
|
14
|
+
__all__ = [
|
|
15
|
+
"FacetExtractor",
|
|
16
|
+
"Outcome",
|
|
17
|
+
"SelectedTrace",
|
|
18
|
+
"TraceCluster",
|
|
19
|
+
"TraceFacet",
|
|
20
|
+
"cluster_facets",
|
|
21
|
+
"hybrid_select",
|
|
22
|
+
"name_clusters",
|
|
23
|
+
"semdedup_keep",
|
|
24
|
+
"tool_signature",
|
|
25
|
+
"trace_digest",
|
|
26
|
+
"trace_domain",
|
|
27
|
+
]
|
|
@@ -0,0 +1,171 @@
|
|
|
1
|
+
"""Clustering of facet embeddings: numpy k-means (cosine) + LLM cluster naming.
|
|
2
|
+
|
|
3
|
+
k-means over L2-normalized facet embeddings (so squared-euclidean ranks like cosine), kmeans++
|
|
4
|
+
init, deterministic under a seed. Cluster naming is the Clio step: an LLM reads a sample of each
|
|
5
|
+
cluster's task summaries and writes a short name + description, which is what makes the resulting
|
|
6
|
+
scenario set auditable ("8 scenarios about baggage claims" instead of "cluster 3").
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import numpy as np
|
|
12
|
+
from pydantic import BaseModel, ValidationError
|
|
13
|
+
|
|
14
|
+
from wmo.core.parsing import extract_json_object
|
|
15
|
+
from wmo.providers.base import Message, Provider
|
|
16
|
+
from wmo.scenarios.mining.facets import TraceFacet
|
|
17
|
+
|
|
18
|
+
_KMEANS_ITERS = 50
|
|
19
|
+
_NAME_SAMPLE = 10
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class TraceCluster(BaseModel):
|
|
23
|
+
"""One discovered intent cluster over the facet corpus."""
|
|
24
|
+
|
|
25
|
+
cluster_id: int
|
|
26
|
+
name: str = ""
|
|
27
|
+
description: str = ""
|
|
28
|
+
member_trace_ids: list[str]
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def default_k(n: int) -> int:
|
|
32
|
+
"""Heuristic base-layer cluster count: sqrt(n), clamped to [2, n]."""
|
|
33
|
+
if n <= 2:
|
|
34
|
+
return max(1, n)
|
|
35
|
+
return min(n, max(2, round(float(np.sqrt(n)))))
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def normalize_rows(embeddings: np.ndarray) -> np.ndarray:
|
|
39
|
+
"""L2-normalize rows (zero rows stay zero) so euclidean k-means ranks like cosine."""
|
|
40
|
+
matrix = np.asarray(embeddings, dtype=np.float64)
|
|
41
|
+
norms = np.linalg.norm(matrix, axis=1, keepdims=True)
|
|
42
|
+
norms[norms == 0.0] = 1.0
|
|
43
|
+
return matrix / norms
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def kmeans_labels(embeddings: np.ndarray, k: int, *, seed: int = 0) -> np.ndarray:
|
|
47
|
+
"""Deterministic k-means (kmeans++ init, Lloyd iterations) over unit-normalized rows.
|
|
48
|
+
|
|
49
|
+
Returns an int label per row. Empty clusters are re-seeded on the farthest point from its
|
|
50
|
+
centroid so exactly `k` non-empty clusters come back whenever `k <= n_distinct_rows`.
|
|
51
|
+
"""
|
|
52
|
+
matrix = normalize_rows(embeddings)
|
|
53
|
+
n = matrix.shape[0]
|
|
54
|
+
if k < 1:
|
|
55
|
+
raise ValueError(f"k must be >= 1, got {k}")
|
|
56
|
+
if k >= n:
|
|
57
|
+
return np.arange(n, dtype=np.int64)
|
|
58
|
+
rng = np.random.default_rng(seed)
|
|
59
|
+
centroids = _kmeans_pp_init(matrix, k, rng)
|
|
60
|
+
labels = np.full(n, -1, dtype=np.int64) # impossible sentinel: never false-converges on iter 1
|
|
61
|
+
for _ in range(_KMEANS_ITERS):
|
|
62
|
+
distances = _sq_distances(matrix, centroids)
|
|
63
|
+
new_labels = distances.argmin(axis=1)
|
|
64
|
+
for cluster in range(k):
|
|
65
|
+
members = matrix[new_labels == cluster]
|
|
66
|
+
if len(members) > 0:
|
|
67
|
+
centroids[cluster] = members.mean(axis=0)
|
|
68
|
+
else:
|
|
69
|
+
# Re-seed an empty cluster on the point farthest from its current centroid.
|
|
70
|
+
farthest = int(np.argmax(distances.min(axis=1)))
|
|
71
|
+
centroids[cluster] = matrix[farthest]
|
|
72
|
+
new_labels[farthest] = cluster
|
|
73
|
+
if np.array_equal(new_labels, labels):
|
|
74
|
+
break
|
|
75
|
+
labels = new_labels
|
|
76
|
+
return labels
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def _kmeans_pp_init(matrix: np.ndarray, k: int, rng: np.random.Generator) -> np.ndarray:
|
|
80
|
+
"""kmeans++ seeding: spread initial centroids proportionally to squared distance."""
|
|
81
|
+
n = matrix.shape[0]
|
|
82
|
+
centroids = np.empty((k, matrix.shape[1]), dtype=np.float64)
|
|
83
|
+
centroids[0] = matrix[rng.integers(n)]
|
|
84
|
+
closest = _sq_distances(matrix, centroids[:1]).min(axis=1)
|
|
85
|
+
for i in range(1, k):
|
|
86
|
+
total = float(closest.sum())
|
|
87
|
+
if total <= 0.0: # all remaining points coincide with a centroid
|
|
88
|
+
centroids[i:] = centroids[0]
|
|
89
|
+
break
|
|
90
|
+
probabilities = closest / total
|
|
91
|
+
centroids[i] = matrix[rng.choice(n, p=probabilities)]
|
|
92
|
+
closest = np.minimum(closest, _sq_distances(matrix, centroids[i : i + 1]).min(axis=1))
|
|
93
|
+
return centroids
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def _sq_distances(matrix: np.ndarray, centroids: np.ndarray) -> np.ndarray:
|
|
97
|
+
"""Squared euclidean distance from every row to every centroid, shape (n, k)."""
|
|
98
|
+
diff = matrix[:, None, :] - centroids[None, :, :]
|
|
99
|
+
return np.einsum("nkd,nkd->nk", diff, diff)
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def cluster_facets(
|
|
103
|
+
facets: list[TraceFacet],
|
|
104
|
+
embeddings: np.ndarray,
|
|
105
|
+
*,
|
|
106
|
+
k: int | None = None,
|
|
107
|
+
seed: int = 0,
|
|
108
|
+
) -> tuple[np.ndarray, list[TraceCluster]]:
|
|
109
|
+
"""Cluster the facet corpus; returns (labels per facet, clusters ordered by descending size)."""
|
|
110
|
+
if len(facets) != len(embeddings):
|
|
111
|
+
raise ValueError(f"{len(facets)} facets but {len(embeddings)} embeddings")
|
|
112
|
+
if not facets:
|
|
113
|
+
return np.empty(0, dtype=np.int64), []
|
|
114
|
+
chosen_k = k if k is not None else default_k(len(facets))
|
|
115
|
+
labels = kmeans_labels(embeddings, chosen_k, seed=seed)
|
|
116
|
+
clusters: list[TraceCluster] = []
|
|
117
|
+
for cluster_id in sorted(set(labels.tolist())):
|
|
118
|
+
member_ids = [facets[i].trace_id for i in np.flatnonzero(labels == cluster_id)]
|
|
119
|
+
clusters.append(TraceCluster(cluster_id=int(cluster_id), member_trace_ids=member_ids))
|
|
120
|
+
clusters.sort(key=lambda c: len(c.member_trace_ids), reverse=True)
|
|
121
|
+
return labels, clusters
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
NAMING_SYSTEM = """You name one cluster of related AI-agent tasks. You see a sample of short task
|
|
125
|
+
summaries that all landed in the same cluster.
|
|
126
|
+
|
|
127
|
+
Respond with ONLY a JSON object, no prose around it:
|
|
128
|
+
{"name": "<2-5 word noun phrase naming the shared task intent>",
|
|
129
|
+
"description": "<one sentence describing what these tasks have in common>"}"""
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
class _RawName(BaseModel):
|
|
133
|
+
name: str
|
|
134
|
+
description: str = ""
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
def name_clusters(
|
|
138
|
+
provider: Provider,
|
|
139
|
+
clusters: list[TraceCluster],
|
|
140
|
+
facets: list[TraceFacet],
|
|
141
|
+
*,
|
|
142
|
+
sample_size: int = _NAME_SAMPLE,
|
|
143
|
+
) -> None:
|
|
144
|
+
"""Fill in `name`/`description` on every cluster via one LLM call each (mutates in place)."""
|
|
145
|
+
by_id = {facet.trace_id: facet for facet in facets}
|
|
146
|
+
for cluster in clusters:
|
|
147
|
+
summaries = [
|
|
148
|
+
by_id[trace_id].task_summary
|
|
149
|
+
for trace_id in cluster.member_trace_ids[:sample_size]
|
|
150
|
+
if trace_id in by_id
|
|
151
|
+
]
|
|
152
|
+
prompt = "TASK SUMMARIES:\n" + "\n".join(f"- {s}" for s in summaries)
|
|
153
|
+
completion = provider.complete(
|
|
154
|
+
NAMING_SYSTEM,
|
|
155
|
+
[Message(role="user", content=prompt)],
|
|
156
|
+
temperature=0.0,
|
|
157
|
+
max_tokens=256,
|
|
158
|
+
)
|
|
159
|
+
raw = extract_json_object(completion.text)
|
|
160
|
+
parsed: _RawName | None = None
|
|
161
|
+
if raw is not None:
|
|
162
|
+
try:
|
|
163
|
+
parsed = _RawName.model_validate_json(raw)
|
|
164
|
+
except ValidationError:
|
|
165
|
+
parsed = None
|
|
166
|
+
if parsed is not None and parsed.name.strip():
|
|
167
|
+
cluster.name = parsed.name.strip()
|
|
168
|
+
cluster.description = parsed.description.strip()
|
|
169
|
+
else:
|
|
170
|
+
cluster.name = f"cluster {cluster.cluster_id}"
|
|
171
|
+
cluster.description = summaries[0] if summaries else ""
|