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

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (308) hide show
  1. llm_waterfall/LICENSE +21 -0
  2. llm_waterfall/__init__.py +53 -0
  3. llm_waterfall/adapters/__init__.py +36 -0
  4. llm_waterfall/adapters/anthropic.py +105 -0
  5. llm_waterfall/adapters/aws_mantle.py +47 -0
  6. llm_waterfall/adapters/azure_openai.py +71 -0
  7. llm_waterfall/adapters/base.py +51 -0
  8. llm_waterfall/adapters/bedrock.py +309 -0
  9. llm_waterfall/adapters/openai.py +130 -0
  10. llm_waterfall/classify.py +184 -0
  11. llm_waterfall/pricing.py +110 -0
  12. llm_waterfall/py.typed +0 -0
  13. llm_waterfall/types.py +295 -0
  14. llm_waterfall/waterfall.py +255 -0
  15. wmo/__init__.py +38 -0
  16. wmo/agents/__init__.py +7 -0
  17. wmo/agents/default.py +29 -0
  18. wmo/agents/meta.py +55 -0
  19. wmo/agents/optimizer.py +55 -0
  20. wmo/agents/project.py +928 -0
  21. wmo/cli/__init__.py +5 -0
  22. wmo/cli/agent_session.py +1123 -0
  23. wmo/cli/app.py +2489 -0
  24. wmo/cli/e2b_cmds.py +212 -0
  25. wmo/cli/eval_closed_loop.py +207 -0
  26. wmo/cli/harness_app.py +1147 -0
  27. wmo/cli/harness_distill.py +659 -0
  28. wmo/cli/hosted_session.py +880 -0
  29. wmo/cli/ingest_cmd.py +165 -0
  30. wmo/cli/model_roles.py +82 -0
  31. wmo/cli/platform_cmds.py +372 -0
  32. wmo/cli/route_app.py +274 -0
  33. wmo/cli/session_state.py +243 -0
  34. wmo/cli/ui.py +1107 -0
  35. wmo/cli/workspace_sync.py +504 -0
  36. wmo/config/__init__.py +60 -0
  37. wmo/config/card.py +129 -0
  38. wmo/config/config.py +367 -0
  39. wmo/config/dotenv.py +67 -0
  40. wmo/config/settings.py +128 -0
  41. wmo/config/store.py +177 -0
  42. wmo/conftest.py +19 -0
  43. wmo/connect/__init__.py +88 -0
  44. wmo/connect/apps.py +78 -0
  45. wmo/connect/brave.py +284 -0
  46. wmo/connect/connector.py +79 -0
  47. wmo/connect/credentials.py +164 -0
  48. wmo/connect/github.py +321 -0
  49. wmo/connect/google.py +627 -0
  50. wmo/connect/notion.py +790 -0
  51. wmo/connect/oauth.py +461 -0
  52. wmo/connect/slack.py +555 -0
  53. wmo/connect/store.py +199 -0
  54. wmo/connect/types.py +156 -0
  55. wmo/core/__init__.py +21 -0
  56. wmo/core/parsing.py +281 -0
  57. wmo/core/render.py +271 -0
  58. wmo/core/text.py +40 -0
  59. wmo/core/types.py +116 -0
  60. wmo/distill/__init__.py +14 -0
  61. wmo/distill/agents.py +140 -0
  62. wmo/distill/config.py +1006 -0
  63. wmo/distill/cost.py +437 -0
  64. wmo/distill/data.py +921 -0
  65. wmo/distill/deadlines.py +254 -0
  66. wmo/distill/fake_tinker.py +734 -0
  67. wmo/distill/gate.py +122 -0
  68. wmo/distill/loop.py +3499 -0
  69. wmo/distill/renderers.py +399 -0
  70. wmo/distill/rendering.py +620 -0
  71. wmo/distill/rollouts.py +726 -0
  72. wmo/distill/samples.py +195 -0
  73. wmo/distill/store.py +829 -0
  74. wmo/distill/teacher.py +714 -0
  75. wmo/distill/tokens.py +535 -0
  76. wmo/distill/tracking.py +552 -0
  77. wmo/distill/tripwire.py +411 -0
  78. wmo/distill/xtoken/byte_offsets.py +152 -0
  79. wmo/distill/xtoken/chunks.py +457 -0
  80. wmo/distill/xtoken/prompt_logprobs.py +475 -0
  81. wmo/distill/xtoken/teacher_render.py +346 -0
  82. wmo/engine/__init__.py +28 -0
  83. wmo/engine/autoconfig.py +367 -0
  84. wmo/engine/build.py +346 -0
  85. wmo/engine/demo.py +77 -0
  86. wmo/engine/eval_suites.py +245 -0
  87. wmo/engine/grounding.py +491 -0
  88. wmo/engine/knowledge.py +291 -0
  89. wmo/engine/loader.py +36 -0
  90. wmo/engine/play.py +92 -0
  91. wmo/engine/prompts.py +99 -0
  92. wmo/engine/replay.py +443 -0
  93. wmo/engine/reporting.py +58 -0
  94. wmo/engine/workspace.py +468 -0
  95. wmo/engine/world_model.py +568 -0
  96. wmo/env/__init__.py +22 -0
  97. wmo/env/base.py +121 -0
  98. wmo/env/closed_loop.py +229 -0
  99. wmo/env/episode.py +107 -0
  100. wmo/env/llm_agent.py +93 -0
  101. wmo/env/scenarios.py +73 -0
  102. wmo/evals/__init__.py +52 -0
  103. wmo/evals/agreement.py +110 -0
  104. wmo/evals/base.py +45 -0
  105. wmo/evals/closed_loop.py +480 -0
  106. wmo/evals/failover.py +96 -0
  107. wmo/evals/gold.py +127 -0
  108. wmo/evals/grid.py +394 -0
  109. wmo/evals/grid_plot.py +205 -0
  110. wmo/evals/harbor/__init__.py +27 -0
  111. wmo/evals/harbor/agent.py +573 -0
  112. wmo/evals/harbor/ctrf.py +171 -0
  113. wmo/evals/harbor/e2b_environment.py +587 -0
  114. wmo/evals/harbor/e2b_template_policy.py +144 -0
  115. wmo/evals/harbor/scorer.py +875 -0
  116. wmo/evals/harbor/tasks.py +140 -0
  117. wmo/evals/open_loop.py +194 -0
  118. wmo/evals/tasks.py +53 -0
  119. wmo/harness/__init__.py +51 -0
  120. wmo/harness/code_runtime.py +288 -0
  121. wmo/harness/create.py +1191 -0
  122. wmo/harness/delta.py +220 -0
  123. wmo/harness/doc.py +556 -0
  124. wmo/harness/e2b_ledger.py +342 -0
  125. wmo/harness/e2b_reap.py +476 -0
  126. wmo/harness/e2b_sandbox.py +350 -0
  127. wmo/harness/environment.py +35 -0
  128. wmo/harness/live_session.py +543 -0
  129. wmo/harness/mutate.py +343 -0
  130. wmo/harness/pi_e2b.py +1710 -0
  131. wmo/harness/pi_entry/entry.ts +268 -0
  132. wmo/harness/pi_entry/runner_frames.ts +92 -0
  133. wmo/harness/pi_entry/runner_live.ts +587 -0
  134. wmo/harness/pi_entry/runner_service.ts +270 -0
  135. wmo/harness/pi_entry/runner_stdio.ts +374 -0
  136. wmo/harness/pi_entry/runner_termination.ts +142 -0
  137. wmo/harness/pi_local.py +262 -0
  138. wmo/harness/pi_runtime.py +495 -0
  139. wmo/harness/pi_vendor.py +65 -0
  140. wmo/harness/population.py +509 -0
  141. wmo/harness/project_proposer.py +569 -0
  142. wmo/harness/proposer.py +977 -0
  143. wmo/harness/runner_link.py +619 -0
  144. wmo/harness/runtime.py +389 -0
  145. wmo/harness/scoring.py +247 -0
  146. wmo/harness/skills.py +116 -0
  147. wmo/harness/source_tree.py +319 -0
  148. wmo/harness/store.py +176 -0
  149. wmo/harness/tools.py +105 -0
  150. wmo/harness/vendor/manifest.sha256 +58 -0
  151. wmo/harness/vendor/pi-agent/CHANGELOG.md +556 -0
  152. wmo/harness/vendor/pi-agent/LICENSE +21 -0
  153. wmo/harness/vendor/pi-agent/README.md +488 -0
  154. wmo/harness/vendor/pi-agent/VENDOR.md +39 -0
  155. wmo/harness/vendor/pi-agent/docs/agent-harness.md +486 -0
  156. wmo/harness/vendor/pi-agent/docs/durable-harness.md +212 -0
  157. wmo/harness/vendor/pi-agent/docs/hooks.md +445 -0
  158. wmo/harness/vendor/pi-agent/docs/models.md +966 -0
  159. wmo/harness/vendor/pi-agent/docs/observability.md +376 -0
  160. wmo/harness/vendor/pi-agent/package.json +60 -0
  161. wmo/harness/vendor/pi-agent/src/agent-loop.ts +748 -0
  162. wmo/harness/vendor/pi-agent/src/agent.ts +575 -0
  163. wmo/harness/vendor/pi-agent/src/harness/agent-harness.ts +1029 -0
  164. wmo/harness/vendor/pi-agent/src/harness/compaction/branch-summarization.ts +261 -0
  165. wmo/harness/vendor/pi-agent/src/harness/compaction/compaction.ts +747 -0
  166. wmo/harness/vendor/pi-agent/src/harness/compaction/utils.ts +144 -0
  167. wmo/harness/vendor/pi-agent/src/harness/env/nodejs.ts +550 -0
  168. wmo/harness/vendor/pi-agent/src/harness/messages.ts +164 -0
  169. wmo/harness/vendor/pi-agent/src/harness/prompt-templates.ts +267 -0
  170. wmo/harness/vendor/pi-agent/src/harness/session/jsonl-repo.ts +177 -0
  171. wmo/harness/vendor/pi-agent/src/harness/session/jsonl-storage.ts +293 -0
  172. wmo/harness/vendor/pi-agent/src/harness/session/memory-repo.ts +50 -0
  173. wmo/harness/vendor/pi-agent/src/harness/session/memory-storage.ts +131 -0
  174. wmo/harness/vendor/pi-agent/src/harness/session/repo-utils.ts +51 -0
  175. wmo/harness/vendor/pi-agent/src/harness/session/session.ts +267 -0
  176. wmo/harness/vendor/pi-agent/src/harness/session/uuid.ts +54 -0
  177. wmo/harness/vendor/pi-agent/src/harness/skills.ts +375 -0
  178. wmo/harness/vendor/pi-agent/src/harness/system-prompt.ts +34 -0
  179. wmo/harness/vendor/pi-agent/src/harness/types.ts +836 -0
  180. wmo/harness/vendor/pi-agent/src/harness/utils/shell-output.ts +135 -0
  181. wmo/harness/vendor/pi-agent/src/harness/utils/truncate.ts +344 -0
  182. wmo/harness/vendor/pi-agent/src/index.ts +44 -0
  183. wmo/harness/vendor/pi-agent/src/node.ts +2 -0
  184. wmo/harness/vendor/pi-agent/src/proxy.ts +367 -0
  185. wmo/harness/vendor/pi-agent/src/types.ts +428 -0
  186. wmo/harness/vendor/pi-agent/test/agent-loop.test.ts +1351 -0
  187. wmo/harness/vendor/pi-agent/test/agent.test.ts +699 -0
  188. wmo/harness/vendor/pi-agent/test/e2e.test.ts +404 -0
  189. wmo/harness/vendor/pi-agent/test/harness/agent-harness-stream.test.ts +213 -0
  190. wmo/harness/vendor/pi-agent/test/harness/agent-harness.test.ts +608 -0
  191. wmo/harness/vendor/pi-agent/test/harness/compaction.test.ts +655 -0
  192. wmo/harness/vendor/pi-agent/test/harness/nodejs-env.test.ts +321 -0
  193. wmo/harness/vendor/pi-agent/test/harness/prompt-templates.test.ts +90 -0
  194. wmo/harness/vendor/pi-agent/test/harness/repo.test.ts +68 -0
  195. wmo/harness/vendor/pi-agent/test/harness/resource-formatting.test.ts +24 -0
  196. wmo/harness/vendor/pi-agent/test/harness/session-test-utils.ts +55 -0
  197. wmo/harness/vendor/pi-agent/test/harness/session-uuid.test.ts +50 -0
  198. wmo/harness/vendor/pi-agent/test/harness/session.test.ts +156 -0
  199. wmo/harness/vendor/pi-agent/test/harness/skills.test.ts +116 -0
  200. wmo/harness/vendor/pi-agent/test/harness/storage.test.ts +299 -0
  201. wmo/harness/vendor/pi-agent/test/harness/system-prompt.test.ts +66 -0
  202. wmo/harness/vendor/pi-agent/test/harness/truncate.test.ts +169 -0
  203. wmo/harness/vendor/pi-agent/test/scratch/simple.ts +72 -0
  204. wmo/harness/vendor/pi-agent/test/utils/calculate.ts +32 -0
  205. wmo/harness/vendor/pi-agent/test/utils/get-current-time.ts +46 -0
  206. wmo/harness/vendor/pi-agent/tsconfig.build.json +13 -0
  207. wmo/harness/vendor/pi-agent/vitest.config.ts +19 -0
  208. wmo/harness/vendor/pi-agent/vitest.harness.config.ts +28 -0
  209. wmo/harness/vendor/vendor_pi.sh +59 -0
  210. wmo/harness/workspace_patch.py +270 -0
  211. wmo/ingest/__init__.py +47 -0
  212. wmo/ingest/adapter.py +72 -0
  213. wmo/ingest/base.py +114 -0
  214. wmo/ingest/braintrust.py +339 -0
  215. wmo/ingest/detect.py +126 -0
  216. wmo/ingest/langfuse.py +291 -0
  217. wmo/ingest/langsmith.py +444 -0
  218. wmo/ingest/mastra.py +330 -0
  219. wmo/ingest/messages.py +170 -0
  220. wmo/ingest/normalize.py +679 -0
  221. wmo/ingest/otel_genai.py +69 -0
  222. wmo/ingest/otel_writer.py +100 -0
  223. wmo/ingest/phoenix.py +150 -0
  224. wmo/ingest/postgres.py +246 -0
  225. wmo/ingest/posthog.py +320 -0
  226. wmo/ingest/quality.py +28 -0
  227. wmo/ingest/stream.py +209 -0
  228. wmo/ingest/testdata/sample_otlp.json +60 -0
  229. wmo/ingest/testdata/sample_spans.jsonl +3 -0
  230. wmo/optimize/__init__.py +25 -0
  231. wmo/optimize/base.py +143 -0
  232. wmo/optimize/gepa.py +806 -0
  233. wmo/optimize/judge.py +262 -0
  234. wmo/optimize/judge_quality.py +359 -0
  235. wmo/optimize/knn.py +468 -0
  236. wmo/optimize/numeric.py +152 -0
  237. wmo/optimize/outcomes.py +103 -0
  238. wmo/optimize/policy.py +669 -0
  239. wmo/optimize/report.py +231 -0
  240. wmo/optimize/reward.py +129 -0
  241. wmo/optimize/routing.py +373 -0
  242. wmo/platform/__init__.py +6 -0
  243. wmo/platform/auth.py +115 -0
  244. wmo/platform/client.py +551 -0
  245. wmo/platform/credentials.py +126 -0
  246. wmo/platform/transfer.py +158 -0
  247. wmo/providers/__init__.py +40 -0
  248. wmo/providers/_bedrock_chat.py +155 -0
  249. wmo/providers/_openai_common.py +182 -0
  250. wmo/providers/_responses_common.py +472 -0
  251. wmo/providers/anthropic.py +134 -0
  252. wmo/providers/azure_openai.py +296 -0
  253. wmo/providers/base.py +300 -0
  254. wmo/providers/bedrock.py +312 -0
  255. wmo/providers/models.py +205 -0
  256. wmo/providers/openai.py +143 -0
  257. wmo/providers/openai_responses.py +240 -0
  258. wmo/providers/pool.py +170 -0
  259. wmo/providers/registry.py +73 -0
  260. wmo/providers/retry.py +151 -0
  261. wmo/providers/tinker.py +936 -0
  262. wmo/providers/waterfall.py +336 -0
  263. wmo/research/__init__.py +81 -0
  264. wmo/research/ablation.py +133 -0
  265. wmo/research/concurrency_plot.py +523 -0
  266. wmo/research/concurrency_run.py +240 -0
  267. wmo/research/concurrency_scaling.py +270 -0
  268. wmo/research/gepa_scaling.py +274 -0
  269. wmo/research/pipeline.py +198 -0
  270. wmo/research/scaling_split.py +82 -0
  271. wmo/research/scenario_fidelity.py +198 -0
  272. wmo/research/scenario_recovery.py +92 -0
  273. wmo/research/seed_stability.py +90 -0
  274. wmo/research/trace_scaling.py +348 -0
  275. wmo/retrieval/__init__.py +6 -0
  276. wmo/retrieval/embedders.py +105 -0
  277. wmo/retrieval/leakfree.py +52 -0
  278. wmo/retrieval/retriever.py +173 -0
  279. wmo/scenarios/__init__.py +58 -0
  280. wmo/scenarios/builder.py +152 -0
  281. wmo/scenarios/mining/__init__.py +27 -0
  282. wmo/scenarios/mining/clustering.py +171 -0
  283. wmo/scenarios/mining/facets.py +226 -0
  284. wmo/scenarios/mining/selection.py +220 -0
  285. wmo/scenarios/synthesis/__init__.py +6 -0
  286. wmo/scenarios/synthesis/scenario_set.py +63 -0
  287. wmo/scenarios/synthesis/synthesizer.py +85 -0
  288. wmo/scenarios/verification/__init__.py +17 -0
  289. wmo/scenarios/verification/judge.py +97 -0
  290. wmo/scenarios/verification/verify.py +135 -0
  291. wmo/serving/__init__.py +5 -0
  292. wmo/serving/builds.py +451 -0
  293. wmo/serving/chat.py +878 -0
  294. wmo/serving/endpoint_config.py +64 -0
  295. wmo/serving/savings.py +250 -0
  296. wmo/serving/server.py +553 -0
  297. wmo/serving/traces_source.py +206 -0
  298. wmo/telemetry.py +213 -0
  299. wmo/tracking/__init__.py +36 -0
  300. wmo/tracking/clock.py +24 -0
  301. wmo/tracking/metered.py +125 -0
  302. wmo/tracking/pricing.py +99 -0
  303. wmo/tracking/store.py +31 -0
  304. wmo/tracking/tracker.py +149 -0
  305. world_model_optimizer-0.2.0.dist-info/METADATA +203 -0
  306. world_model_optimizer-0.2.0.dist-info/RECORD +308 -0
  307. world_model_optimizer-0.2.0.dist-info/WHEEL +4 -0
  308. world_model_optimizer-0.2.0.dist-info/entry_points.txt +2 -0
@@ -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
+ ]
@@ -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 ""