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,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
|