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,226 @@
|
|
|
1
|
+
"""Facet extraction: one compact, embeddable summary per trace (the Clio pattern).
|
|
2
|
+
|
|
3
|
+
Raw traces are dominated by boilerplate (tool schemas, retrieved content), so embedding them
|
|
4
|
+
directly washes out task intent — two traces with identical scaffolding but different tasks land
|
|
5
|
+
nearly on top of each other. Instead, a cheap LLM reads a compact digest of each trace and emits a
|
|
6
|
+
`TraceFacet`: a short task summary (what the user was trying to get done), the outcome, and a
|
|
7
|
+
failure category when the episode failed. The deterministic tool-call signature and the corpus
|
|
8
|
+
domain are computed in code, not by the LLM, and join the summary in the embedded text so
|
|
9
|
+
clustering groups by capability rather than phrasing. Downstream clustering/selection operates on
|
|
10
|
+
facet embeddings only.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
16
|
+
from enum import StrEnum
|
|
17
|
+
|
|
18
|
+
from pydantic import BaseModel, ValidationError
|
|
19
|
+
|
|
20
|
+
from wmo.core.parsing import extract_json_object
|
|
21
|
+
from wmo.core.types import ActionKind, Trace
|
|
22
|
+
from wmo.providers.base import Message, Provider
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class Outcome(StrEnum):
|
|
26
|
+
SUCCESS = "success"
|
|
27
|
+
FAILURE = "failure"
|
|
28
|
+
UNKNOWN = "unknown"
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class TraceFacet(BaseModel):
|
|
32
|
+
"""The embeddable summary of one trace; the unit clustering and selection operate on."""
|
|
33
|
+
|
|
34
|
+
trace_id: str
|
|
35
|
+
task_summary: str # <= ~30 words: what the user was trying to get done
|
|
36
|
+
tool_signature: str # deterministic "tool_a>tool_b>..." with consecutive repeats collapsed
|
|
37
|
+
domain: str | None = None # from trace metadata when the corpus records one
|
|
38
|
+
outcome: Outcome = Outcome.UNKNOWN
|
|
39
|
+
failure_category: str | None = None # short label, only when outcome == FAILURE
|
|
40
|
+
|
|
41
|
+
def embed_text(self) -> str:
|
|
42
|
+
"""The text clustering embeds: domain + task intent + capabilities exercised.
|
|
43
|
+
|
|
44
|
+
Embedding the summary alone clusters by phrasing, which splits one capability into
|
|
45
|
+
several clusters ("MMS troubleshooting" vs "International MMS troubleshooting") and lets
|
|
46
|
+
cluster-level allocation double-count it. Domain and the tool signature pull traces that
|
|
47
|
+
exercise the same capability together regardless of how the request was worded.
|
|
48
|
+
"""
|
|
49
|
+
parts = []
|
|
50
|
+
if self.domain:
|
|
51
|
+
parts.append(f"[{self.domain}]")
|
|
52
|
+
parts.append(self.task_summary)
|
|
53
|
+
text = " ".join(parts)
|
|
54
|
+
if self.tool_signature:
|
|
55
|
+
text = f"{text} | tools: {self.tool_signature}"
|
|
56
|
+
return text
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def trace_domain(trace: Trace) -> str | None:
|
|
60
|
+
"""The trace's domain from corpus metadata, when recorded (e.g. tau2's telecom/retail)."""
|
|
61
|
+
value = trace.metadata.get("domain")
|
|
62
|
+
if isinstance(value, str) and value.strip():
|
|
63
|
+
return value.strip()
|
|
64
|
+
return None
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def tool_signature(trace: Trace) -> str:
|
|
68
|
+
"""Deterministic tool-call sequence signature with consecutive repeats collapsed.
|
|
69
|
+
|
|
70
|
+
`search>search>book` becomes `search>book`: the signature captures *which* capabilities the
|
|
71
|
+
episode exercised in what order, not how many retries each took.
|
|
72
|
+
"""
|
|
73
|
+
names: list[str] = []
|
|
74
|
+
for step in trace.steps:
|
|
75
|
+
action = step.action
|
|
76
|
+
if action.kind is not ActionKind.TOOL_CALL or not action.name:
|
|
77
|
+
continue
|
|
78
|
+
if not names or names[-1] != action.name:
|
|
79
|
+
names.append(action.name)
|
|
80
|
+
return ">".join(names)
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
_MAX_DIGEST_STEPS = 30
|
|
84
|
+
_MAX_FIELD_CHARS = 300
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def trace_digest(trace: Trace, *, max_steps: int = _MAX_DIGEST_STEPS) -> str:
|
|
88
|
+
"""A compact plain-text rendering of a trace for the facet-extraction LLM.
|
|
89
|
+
|
|
90
|
+
Includes the task, then one line per step (tool name + truncated arguments + truncated
|
|
91
|
+
observation, error-flagged). Long traces keep the first and last steps and elide the middle:
|
|
92
|
+
intent lives at the start, resolution at the end.
|
|
93
|
+
"""
|
|
94
|
+
lines: list[str] = []
|
|
95
|
+
task = _trace_task(trace)
|
|
96
|
+
if task:
|
|
97
|
+
lines.append(f"TASK: {_truncate(task)}")
|
|
98
|
+
steps = trace.steps
|
|
99
|
+
if len(steps) > max_steps:
|
|
100
|
+
head = max_steps // 2
|
|
101
|
+
tail = max_steps - head
|
|
102
|
+
shown = list(enumerate(steps))[:head] + list(enumerate(steps))[-tail:]
|
|
103
|
+
elided = len(steps) - max_steps
|
|
104
|
+
else:
|
|
105
|
+
shown = list(enumerate(steps))
|
|
106
|
+
elided = 0
|
|
107
|
+
previous_index = -1
|
|
108
|
+
for index, step in shown:
|
|
109
|
+
if index > previous_index + 1:
|
|
110
|
+
lines.append(f"... ({elided} steps elided) ...")
|
|
111
|
+
previous_index = index
|
|
112
|
+
action = step.action
|
|
113
|
+
if action.kind is ActionKind.TOOL_CALL:
|
|
114
|
+
args = _truncate(str(action.arguments)) if action.arguments else ""
|
|
115
|
+
head_line = f"{index}. CALL {action.name}({args})"
|
|
116
|
+
else:
|
|
117
|
+
head_line = f"{index}. MSG {_truncate(action.content or '')}"
|
|
118
|
+
observation = step.observation
|
|
119
|
+
error_mark = " [ERROR]" if observation.is_error else ""
|
|
120
|
+
lines.append(f"{head_line} -> {_truncate(observation.content)}{error_mark}")
|
|
121
|
+
return "\n".join(lines)
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
def _truncate(text: str, limit: int = _MAX_FIELD_CHARS) -> str:
|
|
125
|
+
text = " ".join(text.split())
|
|
126
|
+
return text if len(text) <= limit else text[: limit - 1] + "…"
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
def _trace_task(trace: Trace) -> str | None:
|
|
130
|
+
for step in trace.steps:
|
|
131
|
+
if step.task and step.task.strip():
|
|
132
|
+
return step.task.strip()
|
|
133
|
+
return None
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
FACET_SYSTEM = """You summarize one AI-agent episode (a trace of tool calls and messages) into a
|
|
137
|
+
compact facet used to organize a large trace corpus.
|
|
138
|
+
|
|
139
|
+
Respond with ONLY a JSON object, no prose around it:
|
|
140
|
+
{"task_summary": "<what the USER was trying to get done, <=30 words, self-contained, no ids>",
|
|
141
|
+
"outcome": "success" | "failure" | "unknown",
|
|
142
|
+
"failure_category": "<short snake_case label, e.g. wrong_tool_arguments; null unless failure>"}
|
|
143
|
+
|
|
144
|
+
Rules:
|
|
145
|
+
- task_summary states the user's goal, not the agent's mechanics ("cancel a flight booking and get
|
|
146
|
+
a refund", NOT "called cancel_reservation").
|
|
147
|
+
- outcome is "success" only if the episode visibly achieved the goal; "failure" if it visibly did
|
|
148
|
+
not (errors, refusals, wrong result); otherwise "unknown".
|
|
149
|
+
- failure_category names the dominant failure mode in 1-3 words; null when outcome != "failure"."""
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
class _RawFacet(BaseModel):
|
|
153
|
+
"""Lenient view of the extractor's JSON before normalization."""
|
|
154
|
+
|
|
155
|
+
task_summary: str
|
|
156
|
+
outcome: str = "unknown"
|
|
157
|
+
failure_category: str | None = None
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
class FacetExtractor:
|
|
161
|
+
"""LLM facet extraction over a trace corpus (one completion per trace)."""
|
|
162
|
+
|
|
163
|
+
def __init__(self, provider: Provider) -> None:
|
|
164
|
+
self._provider = provider
|
|
165
|
+
|
|
166
|
+
def extract(self, trace: Trace) -> TraceFacet:
|
|
167
|
+
"""Extract the facet for one trace; falls back to the raw task on an unparseable reply."""
|
|
168
|
+
completion = self._provider.complete(
|
|
169
|
+
FACET_SYSTEM,
|
|
170
|
+
[Message(role="user", content=trace_digest(trace))],
|
|
171
|
+
temperature=0.0,
|
|
172
|
+
max_tokens=512,
|
|
173
|
+
)
|
|
174
|
+
signature = tool_signature(trace)
|
|
175
|
+
domain = trace_domain(trace)
|
|
176
|
+
raw = extract_json_object(completion.text)
|
|
177
|
+
if raw is not None:
|
|
178
|
+
try:
|
|
179
|
+
parsed = _RawFacet.model_validate_json(raw)
|
|
180
|
+
except ValidationError:
|
|
181
|
+
parsed = None
|
|
182
|
+
if parsed is not None and parsed.task_summary.strip():
|
|
183
|
+
outcome = _parse_outcome(parsed.outcome)
|
|
184
|
+
category = parsed.failure_category if outcome is Outcome.FAILURE else None
|
|
185
|
+
return TraceFacet(
|
|
186
|
+
trace_id=trace.trace_id,
|
|
187
|
+
task_summary=parsed.task_summary.strip(),
|
|
188
|
+
tool_signature=signature,
|
|
189
|
+
domain=domain,
|
|
190
|
+
outcome=outcome,
|
|
191
|
+
failure_category=_normalize_category(category),
|
|
192
|
+
)
|
|
193
|
+
# Fallback: the recorded task prompt is still a usable intent summary; flag as UNKNOWN.
|
|
194
|
+
return TraceFacet(
|
|
195
|
+
trace_id=trace.trace_id,
|
|
196
|
+
task_summary=_truncate(_trace_task(trace) or "(no task recorded)", 200),
|
|
197
|
+
tool_signature=signature,
|
|
198
|
+
domain=domain,
|
|
199
|
+
outcome=Outcome.UNKNOWN,
|
|
200
|
+
)
|
|
201
|
+
|
|
202
|
+
def extract_all(self, traces: list[Trace], *, concurrency: int = 8) -> list[TraceFacet]:
|
|
203
|
+
"""Extract facets for every trace, in order.
|
|
204
|
+
|
|
205
|
+
Each facet is one independent LLM call, so they run on a small thread pool
|
|
206
|
+
(`pool.map` preserves input order and propagates exceptions — the replay.py
|
|
207
|
+
precedent); `concurrency=1` keeps the sequential loop.
|
|
208
|
+
"""
|
|
209
|
+
if concurrency > 1 and len(traces) > 1:
|
|
210
|
+
with ThreadPoolExecutor(max_workers=min(concurrency, len(traces))) as pool:
|
|
211
|
+
return list(pool.map(self.extract, traces))
|
|
212
|
+
return [self.extract(trace) for trace in traces]
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
def _parse_outcome(raw: str) -> Outcome:
|
|
216
|
+
try:
|
|
217
|
+
return Outcome(raw.strip().lower())
|
|
218
|
+
except ValueError:
|
|
219
|
+
return Outcome.UNKNOWN
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def _normalize_category(category: str | None) -> str | None:
|
|
223
|
+
if category is None:
|
|
224
|
+
return None
|
|
225
|
+
normalized = "_".join(category.strip().lower().split())
|
|
226
|
+
return normalized or None
|
|
@@ -0,0 +1,220 @@
|
|
|
1
|
+
"""Representative selection: SemDeDup + hybrid-allocation medoid picking with failure pinning.
|
|
2
|
+
|
|
3
|
+
Given clustered facet embeddings and a scenario budget K, pick which real traces become scenarios:
|
|
4
|
+
|
|
5
|
+
1. SemDeDup (arXiv 2303.09540): drop near-duplicate facets within a cluster before selection, so
|
|
6
|
+
thirty rewordings of the same request can't claim thirty slots.
|
|
7
|
+
2. Hybrid allocation: ~70% of the budget goes to clusters proportionally to their corpus mass (the
|
|
8
|
+
eval mirrors traffic), the rest round-robin across clusters (the long tail keeps coverage).
|
|
9
|
+
3. Within a cluster the first pick is the medoid (the real trace nearest everything else); extra
|
|
10
|
+
slots go farthest-first for intra-cluster diversity.
|
|
11
|
+
4. Failure pinning: every failure category present in the corpus keeps at least one exemplar,
|
|
12
|
+
regardless of frequency — rare-but-critical traces are exactly the ones proportional sampling
|
|
13
|
+
silently drops.
|
|
14
|
+
|
|
15
|
+
Each selection carries `weight`: the fraction of the (deduped) corpus it stands for, so downstream
|
|
16
|
+
scoring can report a traffic-weighted number. Weights sum to 1 over the selection.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import numpy as np
|
|
22
|
+
from pydantic import BaseModel
|
|
23
|
+
|
|
24
|
+
from wmo.scenarios.mining.clustering import normalize_rows
|
|
25
|
+
from wmo.scenarios.mining.facets import Outcome, TraceFacet
|
|
26
|
+
|
|
27
|
+
DEDUP_THRESHOLD = 0.95
|
|
28
|
+
PROPORTIONAL_FRACTION = 0.7
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class SelectedTrace(BaseModel):
|
|
32
|
+
"""One trace chosen to become a scenario, with the corpus mass it represents."""
|
|
33
|
+
|
|
34
|
+
trace_id: str
|
|
35
|
+
cluster_id: int
|
|
36
|
+
weight: float # fraction of the deduped corpus this selection stands for
|
|
37
|
+
pinned_failure: str | None = None # failure category this pick was retained for, if any
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def semdedup_keep(
|
|
41
|
+
embeddings: np.ndarray, labels: np.ndarray, *, threshold: float = DEDUP_THRESHOLD
|
|
42
|
+
) -> list[int]:
|
|
43
|
+
"""Indices that survive within-cluster near-duplicate removal (first occurrence wins).
|
|
44
|
+
|
|
45
|
+
Compares cosine similarity only within a cluster (the SemDeDup trick: k-means already grouped
|
|
46
|
+
near-duplicates, so the quadratic pass stays per-cluster).
|
|
47
|
+
"""
|
|
48
|
+
matrix = normalize_rows(embeddings)
|
|
49
|
+
kept: list[int] = []
|
|
50
|
+
for cluster_id in sorted(set(labels.tolist())):
|
|
51
|
+
member_indices = np.flatnonzero(labels == cluster_id)
|
|
52
|
+
cluster_kept: list[int] = []
|
|
53
|
+
for index in member_indices.tolist():
|
|
54
|
+
duplicate = any(
|
|
55
|
+
float(matrix[index] @ matrix[other]) > threshold for other in cluster_kept
|
|
56
|
+
)
|
|
57
|
+
if not duplicate:
|
|
58
|
+
cluster_kept.append(index)
|
|
59
|
+
kept.extend(cluster_kept)
|
|
60
|
+
return sorted(kept)
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def hybrid_select(
|
|
64
|
+
facets: list[TraceFacet],
|
|
65
|
+
embeddings: np.ndarray,
|
|
66
|
+
labels: np.ndarray,
|
|
67
|
+
budget: int,
|
|
68
|
+
*,
|
|
69
|
+
proportional_fraction: float = PROPORTIONAL_FRACTION,
|
|
70
|
+
dedup_threshold: float = DEDUP_THRESHOLD,
|
|
71
|
+
) -> list[SelectedTrace]:
|
|
72
|
+
"""Pick `budget` representative traces from a clustered facet corpus.
|
|
73
|
+
|
|
74
|
+
See the module docstring for the algorithm. Cluster mass (and thus weights) is measured on the
|
|
75
|
+
deduped corpus. Raises when the budget is not positive; a budget larger than the deduped corpus
|
|
76
|
+
returns everything.
|
|
77
|
+
"""
|
|
78
|
+
if budget < 1:
|
|
79
|
+
raise ValueError(f"budget must be >= 1, got {budget}")
|
|
80
|
+
if len(facets) != len(embeddings) or len(facets) != len(labels):
|
|
81
|
+
raise ValueError("facets, embeddings, and labels must be parallel")
|
|
82
|
+
if not facets:
|
|
83
|
+
return []
|
|
84
|
+
|
|
85
|
+
matrix = normalize_rows(embeddings)
|
|
86
|
+
kept = semdedup_keep(embeddings, labels, threshold=dedup_threshold)
|
|
87
|
+
by_cluster: dict[int, list[int]] = {}
|
|
88
|
+
for index in kept:
|
|
89
|
+
by_cluster.setdefault(int(labels[index]), []).append(index)
|
|
90
|
+
total_kept = len(kept)
|
|
91
|
+
if budget >= total_kept:
|
|
92
|
+
selections = [_selection(facets[i], int(labels[i]), 1.0 / total_kept) for i in sorted(kept)]
|
|
93
|
+
return _pin_failures(selections, facets, matrix, labels, by_cluster)
|
|
94
|
+
|
|
95
|
+
slots = _allocate_slots(by_cluster, budget, proportional_fraction)
|
|
96
|
+
selections: list[SelectedTrace] = []
|
|
97
|
+
for cluster_id, cluster_slots in slots.items():
|
|
98
|
+
member_indices = by_cluster[cluster_id]
|
|
99
|
+
chosen = _pick_representatives(matrix, member_indices, cluster_slots)
|
|
100
|
+
cluster_weight = len(member_indices) / total_kept
|
|
101
|
+
for index in chosen:
|
|
102
|
+
selections.append(_selection(facets[index], cluster_id, cluster_weight / len(chosen)))
|
|
103
|
+
selections = _pin_failures(selections, facets, matrix, labels, by_cluster)
|
|
104
|
+
# Clusters allocated zero slots contribute no selection, so their mass would silently vanish
|
|
105
|
+
# from the weights; renormalize so weights always sum to 1 over the returned selection.
|
|
106
|
+
total_weight = sum(s.weight for s in selections)
|
|
107
|
+
return [s.model_copy(update={"weight": s.weight / total_weight}) for s in selections]
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def _selection(facet: TraceFacet, cluster_id: int, weight: float) -> SelectedTrace:
|
|
111
|
+
return SelectedTrace(trace_id=facet.trace_id, cluster_id=cluster_id, weight=weight)
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def _allocate_slots(
|
|
115
|
+
by_cluster: dict[int, list[int]], budget: int, proportional_fraction: float
|
|
116
|
+
) -> dict[int, int]:
|
|
117
|
+
"""Split `budget` slots across clusters: proportional share + round-robin coverage share.
|
|
118
|
+
|
|
119
|
+
Proportional slots follow cluster mass (largest-remainder rounding); the remaining coverage
|
|
120
|
+
slots go one per cluster in descending-mass order, cycling. No cluster is allocated more slots
|
|
121
|
+
than it has (deduped) members; leftover slots spill to the largest clusters with capacity.
|
|
122
|
+
"""
|
|
123
|
+
cluster_ids = sorted(by_cluster, key=lambda c: len(by_cluster[c]), reverse=True)
|
|
124
|
+
capacity = {c: len(by_cluster[c]) for c in cluster_ids}
|
|
125
|
+
total = sum(capacity.values())
|
|
126
|
+
proportional_budget = min(budget, round(budget * proportional_fraction))
|
|
127
|
+
|
|
128
|
+
# Largest-remainder proportional allocation, capped by capacity.
|
|
129
|
+
quotas = {c: proportional_budget * capacity[c] / total for c in cluster_ids}
|
|
130
|
+
slots = {c: min(int(quotas[c]), capacity[c]) for c in cluster_ids}
|
|
131
|
+
remainders = sorted(cluster_ids, key=lambda c: quotas[c] - int(quotas[c]), reverse=True)
|
|
132
|
+
leftover = proportional_budget - sum(slots.values())
|
|
133
|
+
for cluster_id in remainders:
|
|
134
|
+
if leftover <= 0:
|
|
135
|
+
break
|
|
136
|
+
if slots[cluster_id] < capacity[cluster_id]:
|
|
137
|
+
slots[cluster_id] += 1
|
|
138
|
+
leftover -= 1
|
|
139
|
+
|
|
140
|
+
# Coverage slots: uncovered clusters first (the long tail is the whole point of this share),
|
|
141
|
+
# then cycle clusters by descending mass, one slot each, skipping full clusters.
|
|
142
|
+
remaining = budget - sum(slots.values())
|
|
143
|
+
for cluster_id in cluster_ids:
|
|
144
|
+
if remaining <= 0:
|
|
145
|
+
break
|
|
146
|
+
if slots[cluster_id] == 0 and capacity[cluster_id] > 0:
|
|
147
|
+
slots[cluster_id] = 1
|
|
148
|
+
remaining -= 1
|
|
149
|
+
while remaining > 0:
|
|
150
|
+
progressed = False
|
|
151
|
+
for cluster_id in cluster_ids:
|
|
152
|
+
if remaining <= 0:
|
|
153
|
+
break
|
|
154
|
+
if slots[cluster_id] < capacity[cluster_id]:
|
|
155
|
+
slots[cluster_id] += 1
|
|
156
|
+
remaining -= 1
|
|
157
|
+
progressed = True
|
|
158
|
+
if not progressed: # every cluster saturated; budget > corpus, handled by caller
|
|
159
|
+
break
|
|
160
|
+
return {c: s for c, s in slots.items() if s > 0}
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
def _pick_representatives(matrix: np.ndarray, member_indices: list[int], slots: int) -> list[int]:
|
|
164
|
+
"""Medoid first, then farthest-first: real, central exemplars with intra-cluster diversity."""
|
|
165
|
+
members = np.asarray(member_indices)
|
|
166
|
+
if slots >= len(members):
|
|
167
|
+
return members.tolist()
|
|
168
|
+
similarities = matrix[members] @ matrix[members].T
|
|
169
|
+
chosen: list[int] = [int(members[similarities.mean(axis=1).argmax()])] # the medoid
|
|
170
|
+
while len(chosen) < slots:
|
|
171
|
+
chosen_rows = matrix[np.asarray(chosen)]
|
|
172
|
+
best_similarity = (matrix[members] @ chosen_rows.T).max(axis=1)
|
|
173
|
+
best_similarity[np.isin(members, chosen)] = np.inf # never re-pick
|
|
174
|
+
chosen.append(int(members[best_similarity.argmin()]))
|
|
175
|
+
return chosen
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
def _pin_failures(
|
|
179
|
+
selections: list[SelectedTrace],
|
|
180
|
+
facets: list[TraceFacet],
|
|
181
|
+
matrix: np.ndarray,
|
|
182
|
+
labels: np.ndarray,
|
|
183
|
+
by_cluster: dict[int, list[int]],
|
|
184
|
+
) -> list[SelectedTrace]:
|
|
185
|
+
"""Ensure every failure category in the (deduped) corpus keeps at least one exemplar.
|
|
186
|
+
|
|
187
|
+
A missing category's medoid replaces the currently lowest-weight unpinned selection, so the
|
|
188
|
+
budget holds. The replaced selection's weight transfers, keeping weights summing to 1.
|
|
189
|
+
"""
|
|
190
|
+
facet_by_id = {facet.trace_id: facet for facet in facets}
|
|
191
|
+
kept_indices = [i for members in by_cluster.values() for i in members]
|
|
192
|
+
categories: dict[str, list[int]] = {}
|
|
193
|
+
for index in kept_indices:
|
|
194
|
+
facet = facets[index]
|
|
195
|
+
if facet.outcome is Outcome.FAILURE and facet.failure_category:
|
|
196
|
+
categories.setdefault(facet.failure_category, []).append(index)
|
|
197
|
+
|
|
198
|
+
covered = {
|
|
199
|
+
facet_by_id[s.trace_id].failure_category
|
|
200
|
+
for s in selections
|
|
201
|
+
if facet_by_id[s.trace_id].outcome is Outcome.FAILURE
|
|
202
|
+
}
|
|
203
|
+
for category, member_indices in sorted(categories.items()):
|
|
204
|
+
if category in covered:
|
|
205
|
+
continue
|
|
206
|
+
members = np.asarray(member_indices)
|
|
207
|
+
similarities = matrix[members] @ matrix[members].T
|
|
208
|
+
exemplar = int(members[similarities.mean(axis=1).argmax()])
|
|
209
|
+
replaceable = [s for s in selections if s.pinned_failure is None]
|
|
210
|
+
if not replaceable:
|
|
211
|
+
break
|
|
212
|
+
victim = min(replaceable, key=lambda s: s.weight)
|
|
213
|
+
selections[selections.index(victim)] = SelectedTrace(
|
|
214
|
+
trace_id=facets[exemplar].trace_id,
|
|
215
|
+
cluster_id=int(labels[exemplar]),
|
|
216
|
+
weight=victim.weight,
|
|
217
|
+
pinned_failure=category,
|
|
218
|
+
)
|
|
219
|
+
covered.add(category)
|
|
220
|
+
return selections
|
|
@@ -0,0 +1,6 @@
|
|
|
1
|
+
"""Synthesis: write self-contained, judgeable scenarios from selected traces."""
|
|
2
|
+
|
|
3
|
+
from wmo.scenarios.synthesis.scenario_set import EvalScenario, ScenarioSet
|
|
4
|
+
from wmo.scenarios.synthesis.synthesizer import ScenarioSynthesizer
|
|
5
|
+
|
|
6
|
+
__all__ = ["EvalScenario", "ScenarioSet", "ScenarioSynthesizer"]
|
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
"""The scenario data model: `EvalScenario` records and the versioned `ScenarioSet` artifact."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
|
|
7
|
+
from pydantic import BaseModel, Field
|
|
8
|
+
|
|
9
|
+
from wmo.core.types import EnvState
|
|
10
|
+
from wmo.env.scenarios import Scenario
|
|
11
|
+
from wmo.scenarios.mining.clustering import TraceCluster
|
|
12
|
+
from wmo.scenarios.mining.facets import Outcome
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class EvalScenario(BaseModel):
|
|
16
|
+
"""One reusable eval scenario distilled from a real trace."""
|
|
17
|
+
|
|
18
|
+
scenario_id: str
|
|
19
|
+
task: str # self-contained task statement handed to the agent
|
|
20
|
+
seed_state: EnvState = Field(default_factory=EnvState) # initial env state for the world model
|
|
21
|
+
checklist: list[str] = Field(default_factory=list) # judgeable success criteria
|
|
22
|
+
provenance: list[str] = Field(default_factory=list) # source trace_ids
|
|
23
|
+
cluster_name: str = ""
|
|
24
|
+
weight: float = 0.0 # fraction of the corpus this scenario represents
|
|
25
|
+
source_outcome: Outcome = Outcome.UNKNOWN
|
|
26
|
+
failure_category: str | None = None
|
|
27
|
+
|
|
28
|
+
def to_scenario(self) -> Scenario:
|
|
29
|
+
"""The minimal `Scenario` view consumed by existing rollout code."""
|
|
30
|
+
return Scenario(task=self.task, provenance=list(self.provenance))
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class ScenarioSet(BaseModel):
|
|
34
|
+
"""The constructed scenario set plus the corpus statistics that justify it."""
|
|
35
|
+
|
|
36
|
+
scenarios: list[EvalScenario]
|
|
37
|
+
clusters: list[TraceCluster] = Field(default_factory=list)
|
|
38
|
+
corpus_traces: int = 0
|
|
39
|
+
corpus_coverage: float = 0.0 # fraction of corpus facets within tau of a selected facet
|
|
40
|
+
coverage_tau: float = 0.0
|
|
41
|
+
|
|
42
|
+
def retain(self, scenario_ids: set[str]) -> None:
|
|
43
|
+
"""Keep only `scenario_ids`, renormalizing weights and invalidating coverage.
|
|
44
|
+
|
|
45
|
+
Dropping scenarios (e.g. `wmo scenarios verify --drop`) breaks two invariants the artifact
|
|
46
|
+
promises: weights sum to 1 over the set, and `corpus_coverage` describes the current
|
|
47
|
+
scenarios. Weights are renormalized over the survivors; coverage needs the facet
|
|
48
|
+
embeddings (gone by verify time), so it is zeroed rather than left stale.
|
|
49
|
+
"""
|
|
50
|
+
self.scenarios = [s for s in self.scenarios if s.scenario_id in scenario_ids]
|
|
51
|
+
total_weight = sum(s.weight for s in self.scenarios)
|
|
52
|
+
if total_weight > 0:
|
|
53
|
+
for scenario in self.scenarios:
|
|
54
|
+
scenario.weight /= total_weight
|
|
55
|
+
self.corpus_coverage = 0.0
|
|
56
|
+
self.coverage_tau = 0.0
|
|
57
|
+
|
|
58
|
+
def save(self, path: str | Path) -> None:
|
|
59
|
+
Path(path).write_text(self.model_dump_json(indent=2), encoding="utf-8")
|
|
60
|
+
|
|
61
|
+
@classmethod
|
|
62
|
+
def load(cls, path: str | Path) -> ScenarioSet:
|
|
63
|
+
return cls.model_validate_json(Path(path).read_text(encoding="utf-8"))
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
"""Scenario synthesis: turn a selected trace into a self-contained, judgeable eval scenario.
|
|
2
|
+
|
|
3
|
+
The WildBench pattern: an LLM reads the source trace and writes (1) a self-contained task
|
|
4
|
+
statement (the user's goal plus constraints revealed mid-episode), (2) the minimal initial
|
|
5
|
+
environment state the episode needs (seeds the world model's scratchpad), and (3) a short
|
|
6
|
+
checklist of success criteria a judge can grade a new trajectory against. Every scenario keeps
|
|
7
|
+
provenance to its source trace so it stays auditable.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
from pydantic import BaseModel, Field, ValidationError
|
|
13
|
+
|
|
14
|
+
from wmo.core.parsing import extract_json_object
|
|
15
|
+
from wmo.core.types import EnvState, Trace
|
|
16
|
+
from wmo.providers.base import Message, Provider
|
|
17
|
+
from wmo.scenarios.mining.facets import TraceFacet, trace_digest
|
|
18
|
+
from wmo.scenarios.synthesis.scenario_set import EvalScenario
|
|
19
|
+
|
|
20
|
+
SYNTHESIS_SYSTEM = """You convert one recorded AI-agent episode into a reusable evaluation
|
|
21
|
+
scenario. You see a digest of the episode (task, tool calls, observations).
|
|
22
|
+
|
|
23
|
+
Respond with ONLY a JSON object, no prose around it:
|
|
24
|
+
{"task": "<self-contained task statement for a fresh agent: the user's goal plus any constraints
|
|
25
|
+
revealed during the episode; no references to 'the trace' or 'above'>",
|
|
26
|
+
"initial_state": "<2-6 sentences of environment facts the episode started from (accounts,
|
|
27
|
+
records, files, balances) that a simulator needs to answer the agent consistently>",
|
|
28
|
+
"checklist": ["<3-6 concrete, independently checkable success criteria for a NEW attempt>"]}
|
|
29
|
+
|
|
30
|
+
Rules:
|
|
31
|
+
- The task must be attemptable without seeing the original episode.
|
|
32
|
+
- initial_state states facts about the world, not about the agent's behavior.
|
|
33
|
+
- Checklist items grade the OUTCOME (what ended up true / communicated), not the exact tool
|
|
34
|
+
sequence — a different valid strategy must be able to pass."""
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class _RawSynthesis(BaseModel):
|
|
38
|
+
task: str
|
|
39
|
+
initial_state: str = ""
|
|
40
|
+
checklist: list[str] = Field(default_factory=list)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
class ScenarioSynthesizer:
|
|
44
|
+
"""LLM synthesis of one `EvalScenario` per selected trace."""
|
|
45
|
+
|
|
46
|
+
def __init__(self, provider: Provider) -> None:
|
|
47
|
+
self._provider = provider
|
|
48
|
+
|
|
49
|
+
def synthesize(self, trace: Trace, facet: TraceFacet) -> EvalScenario:
|
|
50
|
+
"""Synthesize the scenario for one selected trace.
|
|
51
|
+
|
|
52
|
+
On an unparseable reply, falls back to the facet's task summary with an empty checklist —
|
|
53
|
+
the scenario stays usable for rollouts, and verification will flag it (no checklist means
|
|
54
|
+
nothing to grade against).
|
|
55
|
+
"""
|
|
56
|
+
completion = self._provider.complete(
|
|
57
|
+
SYNTHESIS_SYSTEM,
|
|
58
|
+
[Message(role="user", content=trace_digest(trace))],
|
|
59
|
+
temperature=0.0,
|
|
60
|
+
max_tokens=1024,
|
|
61
|
+
)
|
|
62
|
+
raw = extract_json_object(completion.text)
|
|
63
|
+
parsed: _RawSynthesis | None = None
|
|
64
|
+
if raw is not None:
|
|
65
|
+
try:
|
|
66
|
+
parsed = _RawSynthesis.model_validate_json(raw)
|
|
67
|
+
except ValidationError:
|
|
68
|
+
parsed = None
|
|
69
|
+
if parsed is not None and parsed.task.strip():
|
|
70
|
+
task = parsed.task.strip()
|
|
71
|
+
seed_state = EnvState(scratchpad=parsed.initial_state.strip())
|
|
72
|
+
checklist = [item.strip() for item in parsed.checklist if item.strip()]
|
|
73
|
+
else:
|
|
74
|
+
task = facet.task_summary
|
|
75
|
+
seed_state = EnvState()
|
|
76
|
+
checklist = []
|
|
77
|
+
return EvalScenario(
|
|
78
|
+
scenario_id=f"scenario-{trace.trace_id}",
|
|
79
|
+
task=task,
|
|
80
|
+
seed_state=seed_state,
|
|
81
|
+
checklist=checklist,
|
|
82
|
+
provenance=[trace.trace_id],
|
|
83
|
+
source_outcome=facet.outcome,
|
|
84
|
+
failure_category=facet.failure_category,
|
|
85
|
+
)
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""Verification: back-agreement + solvability gates and the checklist judge that powers them."""
|
|
2
|
+
|
|
3
|
+
from wmo.scenarios.verification.judge import CHECKLIST_SYSTEM, ChecklistJudge, ChecklistResult
|
|
4
|
+
from wmo.scenarios.verification.verify import (
|
|
5
|
+
ScenarioVerdict,
|
|
6
|
+
VerificationReport,
|
|
7
|
+
verify_scenarios,
|
|
8
|
+
)
|
|
9
|
+
|
|
10
|
+
__all__ = [
|
|
11
|
+
"CHECKLIST_SYSTEM",
|
|
12
|
+
"ChecklistJudge",
|
|
13
|
+
"ChecklistResult",
|
|
14
|
+
"ScenarioVerdict",
|
|
15
|
+
"VerificationReport",
|
|
16
|
+
"verify_scenarios",
|
|
17
|
+
]
|