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,457 @@
1
+ """Chunk plans and cross-tokenizer chunk advantages.
2
+
3
+ A `ChunkPlan` says, for one `TrainDatum`, which student token ranges are
4
+ scoreable and which teacher token range covers the same bytes. Chunks are the
5
+ unit of comparison in the cross-tokenizer loss: the teacher cannot score the
6
+ student's token ids (different vocabulary), so it scores its own tokenization
7
+ of the same text and the two are compared span by span.
8
+
9
+ `attach_chunk_advantages` turns those spans plus the teacher's per-position
10
+ logprobs into the per-token advantage array the existing `importance_sampling`
11
+ wire format already carries. Three properties of this module are load-bearing
12
+ and easy to get wrong:
13
+
14
+ 1. A chunk's influence on the gradient is its reverse-KL gap, not its length.
15
+ The loss sums `advantage * grad log pi` over positions, so broadcasting
16
+ `(teacher_sum - student_sum) / student_len` to a chunk's student tokens
17
+ makes the chunk contribute exactly `teacher_sum - student_sum` no matter
18
+ how many student tokens it spans. Dividing by the TEACHER span length
19
+ instead would scale every chunk by the tokenizers' verbosity ratio.
20
+
21
+ 2. Centering is over CHUNK TOTALS, never over tokens. Subtracting a constant
22
+ from every token shifts each chunk's total by that constant times the
23
+ chunk's length, which is length-dependent and can inuert a long chunk: two
24
+ chunks with totals +1.0 at lengths 10 and 1000 come out at +0.98 and -0.98
25
+ under token centering, so the long one trains in the wrong direction (its
26
+ total inverts). Subtracting a constant from each chunk's TOTAL preserves
27
+ ordering.
28
+
29
+ 3. Positions no chunk covers keep advantage 0.0 and are never touched by
30
+ centering. Advantage 0.0 IS the mask on the wire (`to_tinker_datums` has no
31
+ mask key), so a token that centering nudged off zero would train on noise.
32
+ The student's own structural tokens (end-of-turn framing, tool-call
33
+ wrappers) have no byte-identical counterpart under the teacher's chat
34
+ template and land here, so this is the common case, not an edge case.
35
+ """
36
+
37
+ from __future__ import annotations
38
+
39
+ import logging
40
+ from collections.abc import Sequence
41
+
42
+ from pydantic import BaseModel, ConfigDict, Field, model_validator
43
+
44
+ from wmo.distill.config import DistillConfig
45
+ from wmo.distill.data import TrainDatum
46
+
47
+ logger = logging.getLogger(__name__)
48
+
49
+
50
+ class ChunkSpan(BaseModel):
51
+ """One aligned chunk: a student token range and the teacher range for the same bytes.
52
+
53
+ Both ranges are half-open (`start` inclusive, `end` exclusive) and index
54
+ into their own side's token sequence.
55
+ """
56
+
57
+ model_config = ConfigDict(frozen=True, extra="forbid")
58
+
59
+ student_start: int = Field(ge=1)
60
+ """Student position 0 is excluded: `to_tinker_datums` ships
61
+ `advantages[1:]` for the next-token shift, so an advantage written at
62
+ position 0 is silently discarded. A chunk starting there would lose part
63
+ of its influence with no error, so the plan builder must start at 1."""
64
+
65
+ student_end: int = Field(gt=1)
66
+ teacher_start: int = Field(ge=1)
67
+ """Teacher position 0 has no context and can never carry a logprob, so a
68
+ scoreable chunk never starts there."""
69
+
70
+ teacher_end: int = Field(gt=1)
71
+ exact: bool = True
72
+ """Whether the two ranges' canonicalized text matched exactly (as opposed
73
+ to being paired by the aligner across a mismatch)."""
74
+
75
+ @model_validator(mode="after")
76
+ def _check_ranges(self) -> ChunkSpan:
77
+ """Reject empty or inverted ranges on either side."""
78
+ if self.student_end <= self.student_start:
79
+ raise ValueError(
80
+ f"student range [{self.student_start}, {self.student_end}) is empty or "
81
+ "inverted; a chunk must cover at least one student token"
82
+ )
83
+ if self.teacher_end <= self.teacher_start:
84
+ raise ValueError(
85
+ f"teacher range [{self.teacher_start}, {self.teacher_end}) is empty or "
86
+ "inverted; a chunk must cover at least one teacher token"
87
+ )
88
+ return self
89
+
90
+ @property
91
+ def student_len(self) -> int:
92
+ """How many student tokens this chunk covers."""
93
+ return self.student_end - self.student_start
94
+
95
+ @property
96
+ def teacher_len(self) -> int:
97
+ """How many teacher tokens this chunk covers."""
98
+ return self.teacher_end - self.teacher_start
99
+
100
+
101
+ class ChunkPlan(BaseModel):
102
+ """The chunk alignment for one datum, plus the teacher sequence it scores against."""
103
+
104
+ model_config = ConfigDict(frozen=True, extra="forbid")
105
+
106
+ trial_name: str = Field(min_length=1)
107
+ fragment_index: int = Field(ge=0)
108
+ chunks: list[ChunkSpan] = Field(default_factory=list)
109
+ teacher_token_count: int = Field(ge=0)
110
+ """Length of the teacher token sequence the chunks index into."""
111
+
112
+ @model_validator(mode="after")
113
+ def _check_monotonic(self) -> ChunkPlan:
114
+ """Require chunks to be sorted and non-overlapping on both sides.
115
+
116
+ Overlap would double-count a token's logprob into two chunks, and out
117
+ of order chunks would mean the aligner produced a crossing alignment,
118
+ which the DP forbids by construction.
119
+ """
120
+ previous_student = 1
121
+ previous_teacher = 1
122
+ for index, chunk in enumerate(self.chunks):
123
+ if chunk.student_start < previous_student:
124
+ raise ValueError(
125
+ f"chunk {index} starts at student position {chunk.student_start}, "
126
+ f"before the previous chunk ended ({previous_student}); chunks must "
127
+ "be sorted and non-overlapping"
128
+ )
129
+ if chunk.teacher_start < previous_teacher:
130
+ raise ValueError(
131
+ f"chunk {index} starts at teacher position {chunk.teacher_start}, "
132
+ f"before the previous chunk ended ({previous_teacher}); chunks must "
133
+ "be sorted and non-overlapping"
134
+ )
135
+ if chunk.teacher_end > self.teacher_token_count:
136
+ raise ValueError(
137
+ f"chunk {index} ends at teacher position {chunk.teacher_end}, past "
138
+ f"the teacher sequence length {self.teacher_token_count}"
139
+ )
140
+ previous_student = chunk.student_end
141
+ previous_teacher = chunk.teacher_end
142
+ return self
143
+
144
+ @property
145
+ def scored_student_tokens(self) -> int:
146
+ """How many student tokens are covered by some chunk."""
147
+ return sum(chunk.student_len for chunk in self.chunks)
148
+
149
+ def validate_against(self, datum: TrainDatum) -> None:
150
+ """Check this plan against the datum it will score.
151
+
152
+ Args:
153
+ datum: The datum whose `model_input_tokens` the student ranges
154
+ index into.
155
+
156
+ Raises:
157
+ ValueError: If the plan names a different datum, a chunk runs past
158
+ the token sequence, or a chunk covers a non-loss position.
159
+ That last one is the subtle case: `_merge_trial_spans` fills
160
+ `sampled_logprobs` with 0.0 at context positions as PADDING,
161
+ not as a real logprob, so a chunk straddling a loss-mask
162
+ transition would silently fold zeros into the student sum.
163
+ """
164
+ if datum.trial_name != self.trial_name or datum.fragment_index != self.fragment_index:
165
+ raise ValueError(
166
+ f"chunk plan is for trial {self.trial_name!r} fragment "
167
+ f"{self.fragment_index}, but the datum is trial {datum.trial_name!r} "
168
+ f"fragment {datum.fragment_index}; plans must be paired with their datum"
169
+ )
170
+ length = len(datum.model_input_tokens)
171
+ for index, chunk in enumerate(self.chunks):
172
+ if chunk.student_end > length:
173
+ raise ValueError(
174
+ f"chunk {index} ends at student position {chunk.student_end}, past "
175
+ f"the datum's {length} token(s)"
176
+ )
177
+ for position in range(chunk.student_start, chunk.student_end):
178
+ if datum.loss_mask[position] != 1.0:
179
+ raise ValueError(
180
+ f"chunk {index} covers student position {position}, which is a "
181
+ "context position (loss mask 0.0). Its sampled_logprobs entry is "
182
+ "0.0 filler rather than a real logprob, so scoring it would "
183
+ "corrupt the chunk's student sum; split chunks at every loss "
184
+ "mask transition"
185
+ )
186
+
187
+
188
+ class ChunkAdvantageStats(BaseModel):
189
+ """Accounting for one `attach_chunk_advantages` call.
190
+
191
+ Counters cover only the ATTACHED datums; a dropped datum is never trained
192
+ on, so its tokens are not signal.
193
+ """
194
+
195
+ model_config = ConfigDict(frozen=True, extra="forbid")
196
+
197
+ datums: int = Field(ge=0)
198
+ mismatch_drops: int = Field(ge=0)
199
+ """Datums dropped because their teacher row or plan did not line up."""
200
+
201
+ empty_coverage_drops: int = Field(ge=0)
202
+ """Datums dropped because no chunk covered any loss token, so the datum
203
+ carries no signal at all (an all-zero advantage array would be a wasted
204
+ forward pass, not a neutral one)."""
205
+
206
+ chunks: int = Field(ge=0)
207
+ scored_loss_tokens: int = Field(ge=0)
208
+ """Loss tokens covered by some chunk (the ones that carry gradient)."""
209
+
210
+ unscored_loss_tokens: int = Field(ge=0)
211
+ """Loss tokens no chunk covered; these keep advantage 0.0."""
212
+
213
+ clipped_chunks: int = Field(ge=0)
214
+ """Chunks whose per-token advantage hit the clip bound before centering;
215
+ always 0 when `train.advantage_clip` is None (clipping off)."""
216
+
217
+ chunk_reverse_kl: float | None
218
+ """`mean(student_lp - teacher_lp)` over scored loss tokens, the
219
+ cross-tokenizer analogue of the same-tokenizer reverse-KL metric; None
220
+ when nothing was scored."""
221
+
222
+ advantage_mean: float | None
223
+ """Mean advantage over scored loss tokens exactly as trained (after any
224
+ clipping and any centering). With both off (the defaults) it is the mean
225
+ chunk gap, so it reads the objective; under `train.center_advantages` it
226
+ is ~0.0 by construction."""
227
+
228
+ advantage_std: float | None
229
+ """Population standard deviation over the same tokens."""
230
+
231
+ @property
232
+ def coverage_rate(self) -> float:
233
+ """Fraction of the attached datums' loss tokens that a chunk covered."""
234
+ total = self.scored_loss_tokens + self.unscored_loss_tokens
235
+ return self.scored_loss_tokens / total if total else 0.0
236
+
237
+
238
+ def _chunk_totals(
239
+ datum: TrainDatum,
240
+ plan: ChunkPlan,
241
+ teacher_logprobs: Sequence[float | None],
242
+ clip: float | None,
243
+ ) -> tuple[list[float], int] | None:
244
+ """Per-chunk totals for one datum, or None when the teacher row fails.
245
+
246
+ Returns `(totals, clipped)` where `totals[i]` is chunk i's contribution
247
+ after per-token clipping, and `clipped` counts chunks that hit the bound
248
+ (`clip=None` clips nothing, so `clipped` is 0).
249
+ """
250
+ totals: list[float] = []
251
+ clipped_count = 0
252
+ for index, chunk in enumerate(plan.chunks):
253
+ teacher_sum = 0.0
254
+ for position in range(chunk.teacher_start, chunk.teacher_end):
255
+ value = teacher_logprobs[position]
256
+ if value is None:
257
+ logger.warning(
258
+ "dropping datum (trial %s, fragment %d): teacher logprob at "
259
+ "position %d is None but chunk %d needs it; the teacher must score "
260
+ "every position its chunks cover",
261
+ datum.trial_name,
262
+ datum.fragment_index,
263
+ position,
264
+ index,
265
+ )
266
+ return None
267
+ teacher_sum += value
268
+ student_sum = sum(
269
+ datum.sampled_logprobs[position]
270
+ for position in range(chunk.student_start, chunk.student_end)
271
+ )
272
+ # Divide by the STUDENT length so the chunk's total influence is
273
+ # exactly its reverse-KL gap (see the module docstring).
274
+ per_token = (teacher_sum - student_sum) / chunk.student_len
275
+ bounded = per_token if clip is None else min(max(per_token, -clip), clip)
276
+ if bounded != per_token:
277
+ clipped_count += 1
278
+ totals.append(bounded * chunk.student_len)
279
+ return totals, clipped_count
280
+
281
+
282
+ def attach_chunk_advantages(
283
+ datums: Sequence[TrainDatum],
284
+ plans: Sequence[ChunkPlan],
285
+ teacher_logprobs: Sequence[Sequence[float | None]],
286
+ cfg: DistillConfig,
287
+ ) -> tuple[list[TrainDatum], ChunkAdvantageStats]:
288
+ """Fill per-token advantages from chunk-aligned teacher logprobs.
289
+
290
+ Each chunk gets `(teacher_sum - student_sum) / student_len` (bounded to
291
+ `+-train.advantage_clip` when that bound is set; None, the default, clips
292
+ nothing) broadcast to its student tokens, so the chunk contributes its
293
+ reverse-KL gap regardless of length. Under `train.center_advantages` the mean over
294
+ CHUNK TOTALS is then subtracted from every chunk's total (see the module
295
+ docstring for why token-level centering would invert long chunks).
296
+ Positions no chunk covers stay at 0.0 and are never centered.
297
+
298
+ Args:
299
+ datums: Datums from `build_datums` (advantages not yet attached).
300
+ plans: One chunk plan per datum, in datum order.
301
+ teacher_logprobs: One per-position teacher logprob row per datum, in
302
+ the teacher's OWN tokenization (length `teacher_token_count`);
303
+ entry p is the logprob of teacher token p given tokens before it.
304
+ cfg: The run config; reads `train.advantage_clip` (None = no
305
+ clipping) and `train.center_advantages`.
306
+
307
+ Returns:
308
+ New datums with advantages attached (drops removed, order preserved)
309
+ and the stats, including chunk coverage and the chunk reverse KL.
310
+
311
+ Raises:
312
+ ValueError: If `plans` or `teacher_logprobs` do not have exactly one
313
+ entry per datum; that is a caller bug, not per-datum evidence.
314
+ """
315
+ if len(plans) != len(datums):
316
+ raise ValueError(
317
+ f"got {len(plans)} chunk plan(s) for {len(datums)} datum(s); pass exactly "
318
+ "one plan per datum, in datum order"
319
+ )
320
+ if len(teacher_logprobs) != len(datums):
321
+ raise ValueError(
322
+ f"got {len(teacher_logprobs)} teacher logprob row(s) for {len(datums)} "
323
+ "datum(s); pass exactly one row per datum, in datum order"
324
+ )
325
+ clip = cfg.train.advantage_clip
326
+ kept: list[TrainDatum] = []
327
+ kept_plans: list[ChunkPlan] = []
328
+ kept_totals: list[list[float]] = []
329
+ kept_rows: list[Sequence[float | None]] = []
330
+ mismatch_drops = 0
331
+ empty_coverage_drops = 0
332
+ clipped_chunks = 0
333
+ unscored = 0
334
+ for datum, plan, row in zip(datums, plans, teacher_logprobs, strict=True):
335
+ if len(row) != plan.teacher_token_count:
336
+ mismatch_drops += 1
337
+ logger.warning(
338
+ "dropping datum (trial %s, fragment %d) from training: teacher "
339
+ "returned %d logprob(s) for a %d-token teacher sequence; the row must "
340
+ "cover the exact sequence the chunk plan was built against",
341
+ datum.trial_name,
342
+ datum.fragment_index,
343
+ len(row),
344
+ plan.teacher_token_count,
345
+ )
346
+ continue
347
+ try:
348
+ plan.validate_against(datum)
349
+ except ValueError as exc:
350
+ mismatch_drops += 1
351
+ logger.warning(
352
+ "dropping datum (trial %s, fragment %d) from training: %s",
353
+ datum.trial_name,
354
+ datum.fragment_index,
355
+ exc,
356
+ )
357
+ continue
358
+ scored = plan.scored_student_tokens
359
+ if not scored:
360
+ empty_coverage_drops += 1
361
+ logger.warning(
362
+ "dropping datum (trial %s, fragment %d) from training: no chunk covered "
363
+ "any loss token, so the datum carries no gradient. Check the teacher "
364
+ "render and the aligner's fallback rate",
365
+ datum.trial_name,
366
+ datum.fragment_index,
367
+ )
368
+ continue
369
+ computed = _chunk_totals(datum, plan, row, clip)
370
+ if computed is None:
371
+ mismatch_drops += 1
372
+ continue
373
+ totals, clipped = computed
374
+ kept.append(datum)
375
+ kept_plans.append(plan)
376
+ kept_totals.append(totals)
377
+ kept_rows.append(row)
378
+ clipped_chunks += clipped
379
+ unscored += datum.loss_token_count - scored
380
+
381
+ # Centering over chunk totals: subtract one constant from each chunk's
382
+ # TOTAL so every chunk keeps its relative weight (module docstring, point 2).
383
+ if cfg.train.center_advantages:
384
+ chunk_count = sum(len(totals) for totals in kept_totals)
385
+ if chunk_count:
386
+ mean_total = sum(sum(totals) for totals in kept_totals) / chunk_count
387
+ kept_totals = [[total - mean_total for total in totals] for totals in kept_totals]
388
+
389
+ attached: list[TrainDatum] = []
390
+ scored_values: list[float] = []
391
+ for datum, plan, totals in zip(kept, kept_plans, kept_totals, strict=True):
392
+ advantages = [0.0] * len(datum.model_input_tokens)
393
+ for chunk, total in zip(plan.chunks, totals, strict=True):
394
+ per_token = total / chunk.student_len
395
+ for position in range(chunk.student_start, chunk.student_end):
396
+ advantages[position] = per_token
397
+ scored_values.append(per_token)
398
+ attached.append(
399
+ TrainDatum(
400
+ trial_name=datum.trial_name,
401
+ fragment_index=datum.fragment_index,
402
+ model_input_tokens=datum.model_input_tokens,
403
+ loss_mask=datum.loss_mask,
404
+ sampled_logprobs=datum.sampled_logprobs,
405
+ advantages=advantages,
406
+ )
407
+ )
408
+
409
+ # The chunk reverse KL is computed from the PRE-clip, PRE-centering gaps so
410
+ # it stays a comparable measurement of teacher-student divergence rather
411
+ # than a readout of the training transform.
412
+ kl_gap = 0.0
413
+ kl_tokens = 0
414
+ for datum, plan, row in zip(kept, kept_plans, kept_rows, strict=True):
415
+ for chunk in plan.chunks:
416
+ teacher_sum = sum(
417
+ value
418
+ for position in range(chunk.teacher_start, chunk.teacher_end)
419
+ if (value := row[position]) is not None
420
+ )
421
+ student_sum = sum(
422
+ datum.sampled_logprobs[position]
423
+ for position in range(chunk.student_start, chunk.student_end)
424
+ )
425
+ kl_gap += student_sum - teacher_sum
426
+ kl_tokens += chunk.student_len
427
+
428
+ advantage_mean: float | None = None
429
+ advantage_std: float | None = None
430
+ if scored_values:
431
+ advantage_mean = sum(scored_values) / len(scored_values)
432
+ variance = sum((value - advantage_mean) ** 2 for value in scored_values) / len(
433
+ scored_values
434
+ )
435
+ advantage_std = variance**0.5
436
+ stats = ChunkAdvantageStats(
437
+ datums=len(attached),
438
+ mismatch_drops=mismatch_drops,
439
+ empty_coverage_drops=empty_coverage_drops,
440
+ chunks=sum(len(plan.chunks) for plan in kept_plans),
441
+ scored_loss_tokens=len(scored_values),
442
+ unscored_loss_tokens=unscored,
443
+ clipped_chunks=clipped_chunks,
444
+ chunk_reverse_kl=kl_gap / kl_tokens if kl_tokens else None,
445
+ advantage_mean=advantage_mean,
446
+ advantage_std=advantage_std,
447
+ )
448
+ if stats.datums and stats.coverage_rate < 0.95:
449
+ logger.warning(
450
+ "chunk coverage is %.1f%% of loss tokens (%d scored, %d unscored); the "
451
+ "cross-tokenizer path expects >95%%, so check the teacher render's message "
452
+ "content islands and the aligner fallback rate",
453
+ stats.coverage_rate * 100.0,
454
+ stats.scored_loss_tokens,
455
+ stats.unscored_loss_tokens,
456
+ )
457
+ return attached, stats