Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
API_REFERENCE.md1880 linesDownload Raw Back to docs
1# API Reference — composer-replication-framework2 3Complete reference for every public symbol in `composer_replication`. Source-of-truth is the `.py` files in `composer_replication/`; docstrings have been pulled verbatim where they exist and supplemented where missing.4 5**Legend**6 7- ⚠️ **UNTESTED-CONTRACT** — symbol exists and is callable, but its behaviour is not pinned by an automated test in `composer_replication/**/tests/` or `spikes/**/tests/`.8- 🟡 **SKELETON** — class/method body raises `NotImplementedError`; ships as design-of-record per ADR-005 / ADR-006.9 10**Module groups (in this document)**11 121. `composer_replication` (top-level re-exports)132. `composer_replication.loss`143. `composer_replication.batch`154. `composer_replication.opsd`165. `composer_replication.distillation`176. `composer_replication.teacher_replay`187. `composer_replication.replaysim`198. `composer_replication.ingestion` (+ `.claude_code`)209. `composer_replication.hint_generator`2110. `composer_replication.trainer` (+ `.composer_trainer`, `.data_collator`)2211. `composer_replication.diloco`2312. `composer_replication.diloco.serverless` (+ `.executor`, `.allreduce`, `.modal`, `.hf_jobs`, `.replica_entrypoint`)2413. `composer_replication.recipes.prime_rl.composer_loss`2514. `composer_replication.recipes.monarch.actors`2615. `composer_replication.diloco.serverless` — cloud executors (`.eks`, `.sagemaker`)2716. `composer_replication.datagen.docker_sandbox`2817. `composer_replication.safety` (+ `.kill_switch`)29 30---31 32## 1. `composer_replication` — top-level package33 34The package re-exports the most common entry points from sub-modules. `__all__` is the canonical list of public top-level names.35 36### `composer_replication.__version__: str`37 38Package version string. Currently `"0.1.0"`.39 40```python41import composer_replication42print(composer_replication.__version__)  # "0.1.0"43```44 45### `composer_replication._DILOCO_AVAILABLE: bool`46 47`True` iff `torchft` is importable in the running Python environment (gates `make_diloco_outer_loop`). Set to `False` and `make_diloco_outer_loop` is set to `None` when `torchft` is missing.48 49```python50from composer_replication import _DILOCO_AVAILABLE51if _DILOCO_AVAILABLE:52    from composer_replication import make_diloco_outer_loop53```54 55### Re-exports56 57| Name | Source module |58|---|---|59| `compose_loss` | `composer_replication.loss` |60| `LossComponents` | `composer_replication.loss` |61| `build_batch` | `composer_replication.batch` |62| `generalized_jsd_loss` | `composer_replication.opsd` |63| `ClaudeCodeIngester` | `composer_replication.ingestion.claude_code` |64| `IngestionStats` | `composer_replication.ingestion.claude_code` |65| `SYSTEM_PROMPT` | `composer_replication.ingestion.claude_code` |66| `DEFAULT_TEACHERS` | `composer_replication.teacher_replay` |67| `DPOPair` | `composer_replication.teacher_replay` |68| `TeacherCallResult` | `composer_replication.teacher_replay` |69| `TeacherSpec` | `composer_replication.teacher_replay` |70| `TraceState` | `composer_replication.teacher_replay` |71| `extract_dpo_pairs` | `composer_replication.teacher_replay` |72| `replay_trace` | `composer_replication.teacher_replay` |73| `ComposerReplicationTrainer` | `composer_replication.trainer` |74| `make_diloco_outer_loop` | `composer_replication.diloco` (or `None` if `torchft` missing) |75 76See each source module below for full signatures.77 78---79 80## 2. `composer_replication.loss`81 82Verification-harness 3-channel loss. Free function, does not depend on `trl`.83 84### `class LossComponents`85 86```python87@dataclass88class LossComponents:89    lm_ce: torch.Tensor90    sdpo_jsd: torch.Tensor91    trace_replay_dpo: torch.Tensor92    total: torch.Tensor93 94    def detached(self) -> dict[str, float]: ...95```96 97Per-channel breakdown of the total loss for logging and ablation. All four fields are scalar `torch.Tensor`s (`shape=()`); `total = lm_ce + alpha_sdpo * sdpo_jsd + beta_replay * trace_replay_dpo`.98 99**`detached() -> dict[str, float]`** — returns Python-float copies of all four fields with no grad. Useful for W&B logging.100 101```python102from composer_replication import compose_loss, build_batch103components = compose_loss(model, build_batch(tokenizer))104print(components.detached())  # {'lm_ce': 2.34, 'sdpo_jsd': 0.12, ...}105components.total.backward()106```107 108### `compose_loss(model, inputs, *, ...) -> LossComponents`109<a id="compose_loss"></a>110 111```python112def compose_loss(113    model: torch.nn.Module,114    inputs: dict[str, torch.Tensor],115    *,116    alpha_sdpo: float = 0.1,117    beta_replay: float = 0.05,118    sdpo_jsd_beta: float = 0.5,119    sdpo_temperature: float = 1.0,120    sdpo_token_clip: float | None = None,121    replay_dpo_beta: float = 0.1,122    lm_ce_label_smoothing: float = 0.0,123    dpo_variant: Literal["dpo", "simpo"] = "dpo",124    sdpo_wrapper: Literal["none", "taid", "entropy_opd"] = "none",125    taid_t: float | None = None,126    simpo_beta: float = 2.0,127    simpo_gamma: float = 1.0,128    entropy_opd_h_max: float | None = None,129) -> LossComponents130```131 132Compute `total = lm_ce + alpha_sdpo * sdpo_jsd + beta_replay * trace_replay_dpo`.133 134**Required keys in `inputs`**135 136- `input_ids`: `(B, T_s)` student rollout token ids.137- `response_mask`: `(B, T_s)` 1 on assistant-response tokens, 0 elsewhere.138 139**Optional keys** (channel auto-disables if missing OR if its weight = 0):140 141- SDPO: `ctx_teacher_input_ids` `(B, T_t)`, `sdpo_loss_mask` `(B, T_t)`.142- DPO (`dpo_variant="dpo"`): `dpo_chosen_input_ids`, `dpo_chosen_response_mask`, `dpo_rejected_input_ids`, `dpo_rejected_response_mask`, `dpo_chosen_ref_logprobs`, `dpo_rejected_ref_logprobs` (precomputed).143- SimPO (`dpo_variant="simpo"`): same DPO ids/masks; reference logprobs are silently ignored.144- TAID (`sdpo_wrapper="taid"`): no extra `inputs` keys needed; the optional `sdpo_loss_mask` is reused as the per-token TAID mask. Pass `taid_t` directly (or drive it from `TAIDScheduler`).145 146**Parameters**147 148| Name | Type | Default | Meaning |149|---|---|---|---|150| `model` | `torch.nn.Module` | — | HF causal-LM. Must accept `input_ids=` and return an object with `.logits`. |151| `inputs` | `dict[str, torch.Tensor]` | — | Batch dict (see required/optional keys above). |152| `alpha_sdpo` | `float` | `0.1` | Weight on SDPO/JSD channel. `0.0` disables. |153| `beta_replay` | `float` | `0.05` | Weight on trace-replay DPO channel. `0.0` disables. |154| `sdpo_jsd_beta` | `float` | `0.5` | β param for `generalized_jsd_loss` (0=fwd KL, 0.5=JSD, 1=rev KL). Unused when `sdpo_wrapper="taid"`. |155| `sdpo_temperature` | `float` | `1.0` | Softmax temperature in SDPO. Unused when `sdpo_wrapper="taid"`. |156| `sdpo_token_clip` | `float \| None` | `None` | Per-token JSD clamp. |157| `replay_dpo_beta` | `float` | `0.1` | β in standard DPO logit. |158| `lm_ce_label_smoothing` | `float` | `0.0` | `F.cross_entropy(label_smoothing=)`. |159| `dpo_variant` | `Literal["dpo","simpo"]` | `"dpo"` | Channel-3 algorithm. |160| `sdpo_wrapper` | `Literal["none","taid","entropy_opd"]` | `"none"` | Channel-2 wrapper. |161| `taid_t` | `float \| None` | `None` | Current TAID interpolation coefficient in `[0, 1]`. Required when `sdpo_wrapper="taid"`. Drive from `TAIDScheduler` or pass a fixed value. |162| `simpo_beta` | `float` | `2.0` | SimPO β (paper default). |163| `simpo_gamma` | `float` | `1.0` | SimPO target margin γ (paper default). |164| `entropy_opd_h_max` | `float \| None` | `None` | Max-entropy normalizer; `None` ⇒ `log(V)`. |165 166**Returns** `LossComponents` (see above).167 168**Raises** `ValueError` if `dpo_variant` or `sdpo_wrapper` is unknown, if `sdpo_wrapper="taid"` is requested without `taid_t`, or if `taid_t` is outside `[0, 1]`.169 170```python171from composer_replication import compose_loss, build_batch172batch = build_batch(tokenizer)173out = compose_loss(model, batch, alpha_sdpo=0.1, beta_replay=0.05)174out.total.backward()175print(out.detached())176```177 178---179 180## 3. `composer_replication.batch`181 182Verification-harness batch builder.183 184### `build_batch(tokenizer, *, ...) -> dict[str, torch.Tensor]`185 186```python187def build_batch(188    tokenizer: Any,189    *,190    device: torch.device | str = "cpu",191    seed: int = 42,192    variant: str = "factorial",193    align_sdpo_shapes: bool = False,194) -> dict[str, torch.Tensor]195```196 197Construct a full 3-channel batch from a real HF tokenizer. The DPO ref-logprobs are dummy tensors (the smoke verifies loss composition wires together, not the reference-policy precompute).198 199**Returned keys**: `input_ids`, `response_mask`, `ctx_teacher_input_ids`, `sdpo_loss_mask`, `dpo_chosen_input_ids`, `dpo_chosen_response_mask`, `dpo_rejected_input_ids`, `dpo_rejected_response_mask`, `dpo_chosen_ref_logprobs`, `dpo_rejected_ref_logprobs`.200 201**Parameters**202 203| Name | Type | Default | Meaning |204|---|---|---|---|205| `tokenizer` | HF `AutoTokenizer` (duck-typed) | — | Must support `apply_chat_template` and `__call__`. |206| `device` | `torch.device \| str` | `"cpu"` | Target device for all returned tensors. |207| `seed` | `int` | `42` | Fixes `torch.manual_seed`. |208| `variant` | `str` | `"factorial"` | One of `"factorial"`, `"binary_search"`. |209| `align_sdpo_shapes` | `bool` | `False` | If True, truncate/pad `ctx_teacher_input_ids` to `input_ids` length so the SDPO channel actually fires. |210 211**Raises** `ValueError` if `variant` is unknown.212 213```python214from transformers import AutoTokenizer215from composer_replication import build_batch216tok = AutoTokenizer.from_pretrained("Qwen/Qwen2.5-0.5B-Instruct")217batch = build_batch(tok, variant="factorial", align_sdpo_shapes=True)218print({k: v.shape for k, v in batch.items()})219```220 221---222 223## 4. `composer_replication.opsd`224 225Self-distillation generalized-JSD loss, lifted verbatim from `siyan-zhao/OPSD` (MIT) per ADR-006.226 227### `generalized_jsd_loss(student_logits, teacher_logits, labels=None, beta=0.5, ...) -> torch.Tensor`228 229```python230def generalized_jsd_loss(231    student_logits: torch.Tensor,232    teacher_logits: torch.Tensor,233    labels: torch.Tensor | None = None,234    beta: float = 0.5,235    temperature: float = 1.0,236    reduction: str = "batchmean",237    logits_are_probs: bool = False,238    top_k: int | None = None,239    token_clip: float | None = None,240) -> torch.Tensor241```242 243Generalized JSD between student and teacher distributions. Same model on different contexts in the SDPO recipe; student and teacher params come from the SAME model.244 245**Parameters**246 247| Name | Type | Default | Meaning |248|---|---|---|---|249| `student_logits` | `Tensor (B, T, V)` | — | Student logits with grad. |250| `teacher_logits` | `Tensor (B, T, V)` | — | Teacher logits (no grad in SDPO). |251| `labels` | `Tensor (B, T) \| None` | `None` | Per-token mask. `-100` positions are ignored (HF convention). |252| `beta` | `float` in [0, 1] | `0.5` | 0=fwd KL, 1=rev KL, 0.5=symmetric JSD. |253| `temperature` | `float` | `1.0` | Softmax temperature. |254| `reduction` | `str` | `"batchmean"` | `"batchmean"`, `"sum"`, `"mean"`, `"none"`. |255| `logits_are_probs` | `bool` | `False` | Skip softmax if inputs are already probabilities. |256| `top_k` | `int \| None` | `None` | Restrict KL to teacher's top-k tokens. |257| `token_clip` | `float \| None` | `None` | Clip per-token JSD for stability. |258 259**Returns** scalar tensor (or `(B, T)` if `reduction="none"`).260 261**Raises** `ValueError` for unknown `reduction`.262 263```python264import torch265from composer_replication.opsd import generalized_jsd_loss266s = torch.randn(2, 8, 32, requires_grad=True)267t = torch.randn(2, 8, 32)268loss = generalized_jsd_loss(s, t, beta=0.5, reduction="batchmean")269loss.backward()270```271 272---273 274## 5. `composer_replication.distillation`275 276Pluggable self-distillation losses (ADR-007). All pure PyTorch.277 278### `simpo_loss(chosen_avg_logprobs, rejected_avg_logprobs, *, beta=2.0, gamma=1.0) -> torch.Tensor`279 280```python281def simpo_loss(282    chosen_avg_logprobs: torch.Tensor,283    rejected_avg_logprobs: torch.Tensor,284    *,285    beta: float = 2.0,286    gamma: float = 1.0,287) -> torch.Tensor288```289 290Reference-free DPO with target margin γ (Meng et al., NeurIPS 2024). `L = -log σ(β · (avg_logπ(c) − avg_logπ(r)) − γ)`.291 292**Parameters**293 294| Name | Type | Default | Meaning |295|---|---|---|---|296| `chosen_avg_logprobs` | `Tensor (B,)` | — | Per-sequence avg logprob over chosen response tokens. |297| `rejected_avg_logprobs` | `Tensor (B,)` | — | Same for rejected. |298| `beta` | `float` | `2.0` | Scaling factor (paper default). |299| `gamma` | `float` | `1.0` | Target margin (paper default). |300 301**Returns** scalar; **Raises** `ValueError` if shapes mismatch.302 303```python304import torch305from composer_replication.distillation import simpo_loss306loss = simpo_loss(torch.tensor([-2.1, -1.8]), torch.tensor([-3.0, -2.5]),307                  beta=2.0, gamma=1.0)308```309 310### `avg_sequence_logprob(model_logprobs, response_mask) -> torch.Tensor`311 312⚠️ UNTESTED-CONTRACT (helper exported from `simpo.py` but not asserted by a test).313 314```python315def avg_sequence_logprob(316    model_logprobs: torch.Tensor,317    response_mask: torch.Tensor,318) -> torch.Tensor319```320 321Convert `(B, T)` per-token logprobs + `(B, T)` response mask into `(B,)` per-sequence average over response tokens.322 323```python324from composer_replication.distillation.simpo import avg_sequence_logprob325import torch326lp = torch.randn(2, 8); m = torch.tensor([[0,0,1,1,1,0,0,0],[0,1,1,1,1,1,0,0]])327out = avg_sequence_logprob(lp, m)  # shape (2,)328```329 330### `taid_loss(student_logits, teacher_logits, mask=None, *, t) -> torch.Tensor`331 332```python333def taid_loss(334    student_logits: torch.Tensor,335    teacher_logits: torch.Tensor,336    mask: torch.Tensor | None = None,337    *,338    t: float | torch.Tensor,339) -> torch.Tensor340```341 342Faithful port of `SakanaAI/TAID` (arXiv:2501.16937). Forward-KL distillation against a logit-space-interpolated target whose anchor is the **current student detached**:343 344```345p_t = softmax( (1 - t) · stop_grad(student_logits) + t · teacher_logits )346L   = - mean_token  Σ_v  p_t(v) · log_softmax(student_logits)(v)347```348 349At `t=0` the target collapses to the detached student (no teacher signal in the gradient). At `t=1` it reduces to standard forward-KL distillation against the teacher.350 351**Wave 15 breaking change.** The previous signature `taid_loss(student, teacher, student_init, *, schedule_step, total_steps, schedule, alpha_min, alpha_max, jsd_beta, temperature, reduction)` was algorithmically wrong (probability-space mix, frozen step-0 anchor, JSD criterion). All those kwargs are removed; the schedule is now the caller's responsibility (see `TAIDScheduler` below for the upstream adaptive scheme).352 353**Parameters**354 355| Name | Type | Default | Meaning |356|---|---|---|---|357| `student_logits` | `Tensor (B, T, V)` | — | Current student (with grad). |358| `teacher_logits` | `Tensor (B, T, V)` | — | Teacher logits. |359| `mask` | `Tensor (B, T) \| None` | `None` | Token mask. `None` ⇒ all-ones. |360| `t` | `float \| Tensor` | — | Interpolation coefficient in `[0, 1]`. |361 362**Raises** `ValueError` for shape mismatch.363 364```python365from composer_replication.distillation import taid_loss366loss = taid_loss(s_logits, t_logits, mask, t=0.4)367```368 369### `TAIDScheduler(num_train_steps, *, t_start=0.4, t_end=1.0, alpha=5e-4, beta=0.99, disable_adaptive=False)`370 371Stateful schedule that mirrors upstream `TAID.update_t`. Monotone non-decreasing, bumped above the linear floor by an EMA on the relative loss change. Use as:372 373```python374from composer_replication.distillation import TAIDScheduler375 376sched = TAIDScheduler(num_train_steps=10_000)   # paper defaults377for step in range(num_train_steps):378    loss = taid_loss(s, t, mask, t=sched.t)379    loss.backward(); optimizer.step()380    sched.update_t(loss.detach(), global_step=step)381```382 383**Parameters**384 385| Name | Type | Default | Meaning |386|---|---|---|---|387| `num_train_steps` | `int` | — | Total planned training steps; sets the linear floor. |388| `t_start` | `float` | `0.4` | Initial `t` (paper default). |389| `t_end` | `float` | `1.0` | Terminal `t`; hard ceiling at every step. |390| `alpha` | `float` | `5e-4` | Adaptive bump magnitude. |391| `beta` | `float` | `0.99` | EMA decay on relative-loss-change momentum. |392| `disable_adaptive` | `bool` | `False` | If True, fall back to deterministic linear schedule. |393| `device` | `torch.device \| str` | `"cpu"` | Where to allocate state buffers. |394 395**Properties / methods**396 397- `sched.t -> float` — current `t` as a Python float (zero-arg property).398- `sched.update_t(loss, global_step) -> Tensor | None` — update internal state. First finite-loss call only seeds `prev_loss` and returns `None`; subsequent calls return the (positive) `delta_t` added on top of the linear floor.399 400### `entropy_aware_opd_loss(student_logits, teacher_logits, *, labels=None, h_max=None, temperature=1.0, reduction="batchmean") -> torch.Tensor`401 402```python403def entropy_aware_opd_loss(404    student_logits: torch.Tensor,405    teacher_logits: torch.Tensor,406    *,407    labels: torch.Tensor | None = None,408    h_max: float | None = None,409    temperature: float = 1.0,410    reduction: str = "batchmean",411) -> torch.Tensor412```413 414Per-token mixture of forward and reverse KL gated by teacher entropy: `w(t) = clamp(H_teacher(t)/h_max, 0, 1)`. High-entropy tokens use forward KL (mode-covering), low-entropy tokens use reverse KL (mode-seeking).415 416**Parameters**417 418| Name | Type | Default | Meaning |419|---|---|---|---|420| `student_logits` | `Tensor (B,T,V)` | — | Student logits (grad). |421| `teacher_logits` | `Tensor (B,T,V)` | — | Teacher logits (no grad). |422| `labels` | `Tensor (B,T) \| None` | `None` | 0/1 mask, applied multiplicatively after the per-token mix. |423| `h_max` | `float \| None` | `None` ⇒ `log(V)` | Max-entropy normalizer. |424| `temperature` | `float` | `1.0` | Softmax temperature on both. |425| `reduction` | `str` | `"batchmean"` | `"batchmean"`, `"sum"`, `"mean"`, `"none"`. |426 427**Raises** `ValueError` on shape mismatch (student vs teacher; labels vs per-token loss) or unknown `reduction`.428 429```python430from composer_replication.distillation import entropy_aware_opd_loss431loss = entropy_aware_opd_loss(s_logits, t_logits, temperature=1.0)432loss.backward()433```434 435### `teacher_entropy(teacher_logits) -> torch.Tensor`436 437⚠️ UNTESTED-CONTRACT (helper exposed from `entropy_aware_opd.py`'s `__all__` but not directly asserted).438 439Per-token entropy in nats. Input `(B,T,V)`, output `(B,T)`.440 441```python442from composer_replication.distillation.entropy_aware_opd import teacher_entropy443H = teacher_entropy(teacher_logits)  # (B, T)444```445 446---447 448## 6. `composer_replication.teacher_replay`449 450N-teacher OpenRouter parallel client + DPO-pair extractor. `httpx` is lazy-imported inside `replay_trace`; the deterministic local logic is testable without it.451 452### `DEFAULT_TEACHERS: list[TeacherSpec]`453 454Three-teacher default set: `anthropic/claude-opus-4.7`, `openai/gpt-5`, `deepseek/deepseek-v4-pro` with paper-baseline OpenRouter pricing.455 456```python457from composer_replication.teacher_replay import DEFAULT_TEACHERS458print([t["slug"] for t in DEFAULT_TEACHERS])459```460 461### `class TeacherSpec(TypedDict)`462 463```python464class TeacherSpec(TypedDict):465    slug: str466    input_per_mtok: float467    output_per_mtok: float468```469 470OpenRouter model slug + per-million-token pricing.471 472```python473spec: TeacherSpec = {"slug": "openai/gpt-5",474                     "input_per_mtok": 1.25, "output_per_mtok": 10.0}475```476 477### `class TraceState(TypedDict)`478 479```python480class TraceState(TypedDict):481    state_id: str          # unique within the trace482    messages: list[dict]   # OpenAI-style chat history up to (and incl.) this user prompt483    student_action: str    # what the student actually did at this step484```485 486One step of a frozen agentic trace. `student_action` is the raw text emitted by the student; teachers are queried with `messages` and asked to predict the assistant's next action.487 488```python489state: TraceState = {"state_id": "ex001::0042",490                     "messages": [{"role": "user", "content": "..."}],491                     "student_action": "[TOOL_USE] name=Read input={...}"}492```493 494### `class TeacherCallResult(TypedDict)`495 496```python497class TeacherCallResult(TypedDict):498    state_id: str499    teacher_slug: str500    response_text: str | None    # None on error501    latency_s: float502    prompt_tokens: int503    completion_tokens: int504    cost_usd: float505    error: str | None            # None on success506```507 508One row of N×T results from `replay_trace`.509 510```python511r: TeacherCallResult = {"state_id": "x", "teacher_slug": "openai/gpt-5",512    "response_text": "ok", "latency_s": 1.2, "prompt_tokens": 100,513    "completion_tokens": 5, "cost_usd": 0.001, "error": None}514```515 516### `class DPOPair(TypedDict)`517 518```python519class DPOPair(TypedDict):520    state_id: str521    state_messages: list[dict]522    chosen: str          # teacher-consensus action523    rejected: str        # student action524    n_teachers_agreeing: int525```526 527One preference pair extracted from teacher-vs-student disagreement.528 529```python530p: DPOPair = {"state_id": "x", "state_messages": [...], "chosen": "...",531              "rejected": "...", "n_teachers_agreeing": 2}532```533 534### `async replay_trace(states, teachers=DEFAULT_TEACHERS, max_total_usd=5.0, api_key=None) -> list[TeacherCallResult]`535 536```python537async def replay_trace(538    states: Sequence[TraceState],539    teachers: Sequence[TeacherSpec] = tuple(DEFAULT_TEACHERS),540    max_total_usd: float = 5.0,541    api_key: str | None = None,542) -> list[TeacherCallResult]543```544 545For each state, fan-out one parallel call per teacher via OpenRouter. Hard-caps cumulative spend at `max_total_usd` (stops after the offending state completes).546 547**Parameters**548 549| Name | Type | Default | Meaning |550|---|---|---|---|551| `states` | `Sequence[TraceState]` | — | Frozen trace, one entry per assistant turn. |552| `teachers` | `Sequence[TeacherSpec]` | `DEFAULT_TEACHERS` | Models to query in parallel. |553| `max_total_usd` | `float` | `5.0` | Cumulative spend cap. |554| `api_key` | `str \| None` | `None` | OpenRouter key; defaults to `OPENROUTER_API_KEY` env or `~/.hermes/.env`. |555 556**Returns** flat list of `TeacherCallResult`s (length `len(states) * len(teachers)` modulo budget cutoff).557 558**Raises** `RuntimeError` if `OPENROUTER_API_KEY` is not findable; `ImportError` if `httpx` is missing at call time.559 560```python561import asyncio562from composer_replication import replay_trace563results = asyncio.run(replay_trace(states=my_trace, max_total_usd=1.0))564```565 566### `extract_dpo_pairs(states, teacher_actions, agreement_threshold=2) -> list[DPOPair]`567 568```python569def extract_dpo_pairs(570    states: Sequence[TraceState],571    teacher_actions: Sequence[TeacherCallResult],572    agreement_threshold: int = 2,573) -> list[DPOPair]574```575 576Group teacher_actions by `state_id`, normalize whitespace, and emit one `DPOPair` per state where ≥`agreement_threshold` teachers agreed on an action that differs from the student's. `chosen` is the original (un-normalized) teacher response text.577 578**Parameters**579 580| Name | Type | Default | Meaning |581|---|---|---|---|582| `states` | `Sequence[TraceState]` | — | Same as passed to `replay_trace`. |583| `teacher_actions` | `Sequence[TeacherCallResult]` | — | Output of `replay_trace`. |584| `agreement_threshold` | `int` | `2` | Min teachers that must agree for a pair to fire. |585 586**Returns** list of `DPOPair`. At most one pair per state (the most-agreed-upon action wins).587 588```python589from composer_replication import extract_dpo_pairs590pairs = extract_dpo_pairs(my_states, results, agreement_threshold=2)591```592 593### `save_pairs(pairs, path) -> None`594 595⚠️ UNTESTED-CONTRACT.596 597```python598def save_pairs(pairs: Sequence[DPOPair], path: str | Path) -> None599```600 601Write pairs to JSONL (one dict per line). Creates parent dirs.602 603```python604from composer_replication.teacher_replay import save_pairs605save_pairs(pairs, "/tmp/dpo_pairs.jsonl")606```607 608---609 610## 7. `composer_replication.replaysim`611 612ADR-004 normalization layer over `teacher_replay`. Re-exports `DPOPair`, `TeacherCallResult`, `extract_dpo_pairs`, `replay_trace` from `teacher_replay`.613 614### `class NormalizedDPOPair`615 616```python617@dataclass618class NormalizedDPOPair:619    state_id: str620    state_messages: list[dict[str, Any]]621    chosen_messages: list[dict[str, Any]]622    rejected_messages: list[dict[str, Any]]623    n_teachers_agreeing: int624    metadata: dict[str, Any]625```626 627Post-normalization shape. `chosen_messages`/`rejected_messages` are chat-format (`[{"role": "assistant", "content": ...}]`). `metadata` carries op-graph provenance, including `{"skipped": True}` when the normalizer was bypassed (`skip_dj=True`).628 629```python630from composer_replication.replaysim import NormalizedDPOPair631n = NormalizedDPOPair(state_id="x", state_messages=[],632    chosen_messages=[{"role": "assistant", "content": "ok"}],633    rejected_messages=[{"role": "assistant", "content": "no"}],634    n_teachers_agreeing=2, metadata={})635```636 637### `class DJNormalizer`638 639```python640class DJNormalizer:641    DEFAULT_RECIPE: ClassVar[Path]  # composer_replication/recipes/replaysim/default.yaml642 643    def __init__(644        self,645        recipe_path: str | os.PathLike[str] | None = None,646        *,647        skip_dj: bool = False,648    ) -> None: ...649 650    def normalize(651        self,652        pairs: Iterable[DPOPair | dict[str, Any]],653    ) -> list[NormalizedDPOPair]: ...654```655 656`data-juicer`-backed normalizer. Pipeline: each `DPOPair` → JSONL record → `data_juicer.core.DefaultExecutor.run()` against the recipe → JSONL → `NormalizedDPOPair`.657 658**Constructor parameters**659 660| Name | Type | Default | Meaning |661|---|---|---|---|662| `recipe_path` | `str \| PathLike \| None` | `None` ⇒ default recipe | data-juicer YAML recipe path. |663| `skip_dj` | `bool` (kw-only) | `False` | If True: passthrough; records get `metadata={"skipped": True}` and no ops run. |664 665**`normalize(pairs) -> list[NormalizedDPOPair]`** runs the op-graph. Output may be shorter than input if filter ops drop records.666 667**Raises** `RuntimeError` at construction time if `skip_dj=False` and `data_juicer` is not importable. `FileNotFoundError` if `recipe_path` (default or explicit) is missing and `skip_dj=False`.668 669```python670from composer_replication.replaysim import DJNormalizer671norm = DJNormalizer(skip_dj=True)672out = norm.normalize(my_pairs)673```674 675### `async replay_and_normalize_trace(*, states, teachers=None, agreement_threshold=2, max_total_usd=5.0, normalizer=None, **replay_kwargs) -> tuple[list[TeacherCallResult], list[NormalizedDPOPair]]`676 677```python678async def replay_and_normalize_trace(679    *,680    states: Any,681    teachers: Any = None,682    agreement_threshold: int = 2,683    max_total_usd: float = 5.0,684    normalizer: DJNormalizer | None = None,685    **replay_kwargs: Any,686) -> tuple[list[TeacherCallResult], list[NormalizedDPOPair]]687```688 689End-to-end async: replay → extract pairs → normalize.690 691**Parameters**692 693| Name | Type | Default | Meaning |694|---|---|---|---|695| `states` | `Sequence[TraceState]` | — | Frozen trace. |696| `teachers` | `Sequence[TeacherSpec] \| None` | `None` ⇒ defaults | Forwarded to `replay_trace`. |697| `agreement_threshold` | `int` | `2` | Forwarded to `extract_dpo_pairs`. |698| `max_total_usd` | `float` | `5.0` | Spend cap. |699| `normalizer` | `DJNormalizer \| None` | `None` ⇒ `DJNormalizer()` | Pass `DJNormalizer(skip_dj=True)` to bypass. |700| `**replay_kwargs` | `Any` | — | Forwarded to `replay_trace` (e.g. `api_key`). |701 702**Returns** `(raw_teacher_actions, normalized_pairs)`.703 704```python705import asyncio706from composer_replication.replaysim import replay_and_normalize_trace, DJNormalizer707raw, norm = asyncio.run(replay_and_normalize_trace(708    states=my_states, normalizer=DJNormalizer(skip_dj=True)))709```710 711### `replay_and_normalize_trace_sync(*args, **kwargs) -> tuple[list[TeacherCallResult], list[NormalizedDPOPair]]`712 713⚠️ UNTESTED-CONTRACT (sync wrapper around the async function; tests call the async form via `asyncio.run`).714 715```python716def replay_and_normalize_trace_sync(*args, **kwargs) -> ...717```718 719Sync convenience wrapping `asyncio.run(replay_and_normalize_trace(...))`.720 721```python722from composer_replication.replaysim.normalize import replay_and_normalize_trace_sync723raw, norm = replay_and_normalize_trace_sync(states=my_states)724```725 726---727 728## 8. `composer_replication.ingestion` & `composer_replication.ingestion.claude_code`729 730Trace-source adapters (ADR-002). v0.1 supports Claude Code session JSONL.731 732### `SYSTEM_PROMPT: str`733 734Default synthetic system prompt injected at `messages[0]` for ingested traces (most Claude Code sessions don't write one). Truncated head: `"You are a senior software engineer working as a coding agent in a terminal environment..."`.735 736```python737from composer_replication import SYSTEM_PROMPT738print(SYSTEM_PROMPT[:60])739```740 741### `class IngestionStats`742 743```python744@dataclass745class IngestionStats:746    n_records_total: int = 0747    n_records_skipped: int = 0748    n_states_emitted: int = 0749    n_assistant_turns: int = 0750    n_tool_use_blocks: int = 0751    n_text_blocks: int = 0752    skipped_subagent: int = 0753    skipped_summary: int = 0754    skipped_truncated_lines: int = 0755    version_warnings: list[str] | None = None  # initialized to [] in __post_init__756```757 758Counters populated by `ClaudeCodeIngester.ingest()` and exposed as `ingester.last_stats`.759 760```python761from composer_replication import IngestionStats762s = IngestionStats(n_records_total=5)763print(s.version_warnings)  # []764```765 766### `class ClaudeCodeIngester`767 768```python769class ClaudeCodeIngester:770    def __init__(771        self,772        *,773        system_prompt: str = SYSTEM_PROMPT,774        skip_sidechain: bool = True,775        strip_thinking: bool = True,776        max_history_tokens: int | None = None,777    ) -> None: ...778 779    def ingest(self, path: Path) -> Iterator[TraceState]: ...780```781 782Convert a Claude Code session JSONL to a stream of `TraceState`s — one per assistant TURN (not per `tool_use` block).783 784**Constructor parameters**785 786| Name | Type | Default | Meaning |787|---|---|---|---|788| `system_prompt` | `str` | `SYSTEM_PROMPT` | Synthetic system message injected at history[0]. |789| `skip_sidechain` | `bool` | `True` | Skip subagent files (`agent-*.jsonl`) and records with `isSidechain=True`. |790| `strip_thinking` | `bool` | `True` | Remove `[THINKING]` blocks from history handed to teachers (kept inside `student_action`). |791| `max_history_tokens` | `int \| None` | `None` | ⚠️ UNTESTED-CONTRACT — accepted but currently not used to truncate. |792 793**`ingest(path) -> Iterator[TraceState]`**: generator over `TraceState` objects. Each turn's `state_id` is `f"{path.stem}::{idx:04d}"`. Side effect: replaces `self.last_stats` with a fresh `IngestionStats` and updates it as records stream.794 795```python796from pathlib import Path797from composer_replication import ClaudeCodeIngester798ing = ClaudeCodeIngester()799for state in ing.ingest(Path("session.jsonl")):800    print(state["state_id"])801print(ing.last_stats.n_states_emitted)802```803 804---805 806## 9. `composer_replication.hint_generator`807 808⚠️ UNTESTED-CONTRACT (entire module — used by the data collator config but not pinned by a test).809 810Template-based hint registry for SDPO error-site injection.811 812### `class HintContext(TypedDict, total=False)`813 814```python815class HintContext(TypedDict, total=False):816    error_kind: str817    error_message: str818    available_tools: list[str]819    tool_name: str820    tool_schema: dict821    intent: str822```823 824Per-error context dict consumed by hint templates.825 826### `HINT_TEMPLATES: dict[str, Callable[[HintContext], str]]`827 828Default registry keys: `"tool_not_found"`, `"json_decode"`, `"type_error"`, `"runtime_error"`, `"repeated_failure"`.829 830### `dispatch(error_kind, ctx=None) -> str | None`831 832```python833def dispatch(error_kind: str, ctx: HintContext | None = None) -> str | None834```835 836Look up `error_kind` in `HINT_TEMPLATES`. Returns the template's hint text, or `None` if the kind is unknown.837 838```python839from composer_replication.hint_generator import dispatch840hint = dispatch("json_decode")  # "Reminder: tool arguments must be valid JSON. ..."841```842 843### `register(error_kind, fn) -> None`844 845```python846def register(error_kind: str, fn: Callable[[HintContext], str]) -> None847```848 849Add or override a custom hint template.850 851```python852from composer_replication.hint_generator import register853register("my_error", lambda ctx: "Reminder: try X.")854```855 856### Individual template functions857 858⚠️ UNTESTED-CONTRACT — exported only via `HINT_TEMPLATES`, useful as building blocks:859 860- `hint_tool_not_found(ctx) -> str`861- `hint_json_decode(ctx) -> str`862- `hint_type_error(ctx) -> str`863- `hint_runtime_error(ctx) -> str`864- `hint_repeated_failure(ctx) -> str`865 866Each accepts a `HintContext` and returns hint text. Signatures are uniform: `Callable[[HintContext], str]`.867 868```python869from composer_replication.hint_generator import hint_tool_not_found870text = hint_tool_not_found({"available_tools": ["Read", "Write"]})871```872 873---874 875## 10. `composer_replication.trainer` & sub-modules876 877Production trainer (TRL `GRPOTrainer` subclass) plus data collator.878 879### `class ComposerReplicationTrainer`880 881```python882class ComposerReplicationTrainer(GRPOTrainer):883    def __init__(884        self,885        *args: Any,886        alpha_sdpo: float = 0.1,887        beta_replay: float = 0.05,888        sdpo_jsd_beta: float = 0.5,889        sdpo_temperature: float = 1.0,890        sdpo_token_clip: float | None = None,891        replay_dpo_beta: float = 0.1,892        **kwargs: Any,893    ) -> None: ...894 895    def _compute_loss(896        self,897        model: torch.nn.Module,898        inputs: dict[str, torch.Tensor],899    ) -> torch.Tensor: ...900```901 902`trl.GRPOTrainer` subclass that overrides `_compute_loss(model, inputs)` to compose `total = grpo + α·sdpo + β·trace_replay_dpo`. When `trl` is not installed, the parent class falls back to `object` so the module imports — but instantiation will fail because the parent's GRPO machinery is missing.903 904**Constructor (kw-only beyond GRPOTrainer's own `*args, **kwargs`)**905 906| Name | Type | Default | Meaning |907|---|---|---|---|908| `alpha_sdpo` | `float` | `0.1` | Channel-2 weight. |909| `beta_replay` | `float` | `0.05` | Channel-3 weight. |910| `sdpo_jsd_beta` | `float` | `0.5` | β for `generalized_jsd_loss`. |911| `sdpo_temperature` | `float` | `1.0` | SDPO softmax temperature. |912| `sdpo_token_clip` | `float \| None` | `None` | Per-token JSD clip. |913| `replay_dpo_beta` | `float` | `0.1` | DPO β. |914 915**`_compute_loss(model, inputs) -> torch.Tensor`** — overrides `GRPOTrainer._compute_loss`. Calls `super()._compute_loss` for channel 1, then `_compute_sdpo_loss` and `_compute_trace_replay_loss`, then composes. Logs per-channel components every `args.logging_steps` (default 50). **Raises** whatever `super()` raises (TRL-shaped errors).916 917**Internal methods (publicly accessible, exercised by spike tests)**918 919- ⚠️ UNTESTED-CONTRACT `_compute_sdpo_loss(model, inputs) -> torch.Tensor` — generalized-JSD between student forward and `ctx_teacher_input_ids` forward. Returns `0.0` (with grad) when `alpha_sdpo == 0`, the key is missing, or shapes mismatch. Logs a warning on shape mismatch.920- ⚠️ UNTESTED-CONTRACT `_compute_trace_replay_loss(model, inputs) -> torch.Tensor` — standard DPO over `dpo_chosen_*` and `dpo_rejected_*`, using precomputed `dpo_chosen_ref_logprobs` / `dpo_rejected_ref_logprobs`.921- ⚠️ UNTESTED-CONTRACT `@staticmethod _sequence_logprobs(model, input_ids, response_mask) -> torch.Tensor` — sum logprobs over response tokens; standard DPO accounting.922 923```python924from composer_replication import ComposerReplicationTrainer925trainer = ComposerReplicationTrainer(926    model=my_model, args=my_grpo_args, train_dataset=ds,927    data_collator=my_collator, alpha_sdpo=0.1, beta_replay=0.05,928)929# trainer.train()  # uses overridden _compute_loss930```931 932### `make_dr_grpo_config(**overrides) -> trl.GRPOConfig`933 934Builds a `trl.GRPOConfig` configured to the **Dr. GRPO** recipe (Composer 2.5's935base objective per the Composer 2 tech report, arXiv:2603.24477; Dr.GRPO =936Liu et al. arXiv:2503.20783). Forces three knobs unless explicitly overridden,937with drift-guard assertions:938 939- `loss_type="dr_grpo"` — removes GRPO's length-standardization length bias.940- `scale_rewards="none"` — NO std-dev advantage normalization (Dr.GRPO requirement).941- `num_iterations=1` — single-epoch / strict on-policy.942 943Any field is overridable via kwargs (`learning_rate=`, `output_dir=`, `beta=`, …).944**Honest KL-estimator delta** (ADR-012 #1): TRL 1.5.0's `GRPOTrainer._compute_loss`945uses the **k3** estimator `exp(ref_logp−logp)−(ref_logp−logp)−1`, NOT the k1946estimator `−log r` the Dr.GRPO/Composer report frames; the delta is small for r≈1947and TRL is not monkeypatched — the delta is documented, not hidden. Exported from948both `composer_replication` and `composer_replication.trainer`.949 950```python951from composer_replication import make_dr_grpo_config952args = make_dr_grpo_config(output_dir="runs/x", learning_rate=1e-6)953```954 955### `make_po_config(objective="dr_grpo", **overrides) -> trl.GRPOConfig`956 957Builds a `trl.GRPOConfig` for a **named policy-optimization objective** from the958`PO_OBJECTIVES` menu (ADR-014). All presets are PURE CONFIG over trl 1.5.0's959`GRPOTrainer` (verified by introspection) — no custom `_compute_loss` needed.960`**overrides` set/override any `GRPOConfig` field on top.961 962- Raises `ValueError` on an unknown objective (lists the valid menu).963- Raises `AssertionError` if a requested knob silently failed to apply (drift guard;964  e.g. GSPO guards `importance_sampling_level=="sequence"`).965 966```python967from composer_replication import make_po_config, PO_OBJECTIVES968args = make_po_config("dapo", output_dir="runs/dapo", learning_rate=2e-6)969```970 971### `PO_OBJECTIVES: dict[str, dict]`972 973The selectable base policy-optimization objectives (named presets over real trl9741.5.0 `GRPOConfig` knobs). Keys and what each sets:975 976| Objective | `loss_type` | `scale_rewards` | Distinguishing knob | Paper |977|---|---|---|---|---|978| `grpo` | `grpo` | `group` (std-norm) | IS=`token` | DeepSeekMath 2402.03300 |979| `dr_grpo` *(default)* | `dr_grpo` | `none` | length-bias removed | 2503.20783 |980| `bnpo` | `bnpo` | `batch` | batch-normalized | trl |981| `dapo` | `dapo` | `none` | `epsilon_high=0.28` (decoupled clip-higher), `mask_truncated_completions`, `beta=0` | 2503.14476 |982| `gspo` | `grpo` | `group` | `importance_sampling_level="sequence"` | Qwen 2507.18071 |983| `cispo` | `cispo` | `none` | `epsilon_high=5.0` (detached IS coef) | MiniMax-M1 2506.13585 |984 985> **Diagnostic gotcha:** for any PO-objective ablation, log the *distinguishing*986> diagnostic (`clip_ratio/high_mean` for DAPO, the sequence-level ratio for GSPO).987> A `0` means the knob never engaged — NOT that the objectives are equal. (This is988> exactly the inert-knob artifact the A1 DAPO-vs-Dr.GRPO washout hit at lr=1e-6.)989 990### `class TraceTurn(TypedDict, total=False)` — `trainer.data_collator`991 992```python993class TraceTurn(TypedDict, total=False):994    role: str                # "user" | "assistant" | "tool"995    content: str996    tool_call: dict | None997    tool_error: str | None998    error_meta: dict999```1000 1001One turn of an agentic trace as consumed by `ComposerDataCollator`.1002 1003### `class TraceExample(TypedDict, total=False)` — `trainer.data_collator`1004 1005```python1006class TraceExample(TypedDict, total=False):1007    trace_id: str1008    turns: list[TraceTurn]1009    final_reward: float1010    dpo_pairs: list[dict] | None1011```1012 1013One training example: `(turns, optional dpo_pairs)`. `dpo_pairs` shape matches `DPOPair`.1014 1015### `class TokenizerLike` — `trainer.data_collator`1016 1017⚠️ UNTESTED-CONTRACT (duck-typed protocol; used as a type hint).1018 1019```python1020class TokenizerLike:1021    pad_token_id: int1022    def __call__(self, text: str | list[str], **kwargs: Any) -> dict[str, list]: ...1023    def apply_chat_template(self, messages: list[dict], **kwargs: Any) -> str | list[int]: ...1024```1025 1026Minimal protocol the collator needs. Compatible with HF `AutoTokenizer`.1027 1028### `class CollatorConfig` — `trainer.data_collator`1029 1030```python1031@dataclass1032class CollatorConfig:1033    max_seq_len: int = 40961034    max_dpo_seq_len: int = 20481035    pad_token_id: int = 01036    ignore_index: int = -1001037    enable_sdpo: bool = True1038    hint_generator: Callable[[str, dict], str | None] | None = None1039    enable_replay_dpo: bool = True1040    rlvr_reward_key: str = "final_reward"1041```1042 1043Tunables for `ComposerDataCollator`.1044 1045| Field | Default | Meaning |1046|---|---|---|1047| `max_seq_len` | `4096` | Truncation cap for student/teacher sequences. |1048| `max_dpo_seq_len` | `2048` | Truncation cap for DPO chosen/rejected sequences. |1049| `pad_token_id` | `0` | Padding token id. |1050| `ignore_index` | `-100` | HF "ignore in loss" sentinel for SDPO mask. |1051| `enable_sdpo` | `True` | Toggle channel-2 fields. |1052| `hint_generator` | `Callable[[str, dict], str \| None] \| None` (`None`) | `(error_kind, error_meta) -> hint_text`. SDPO is no-op without this. |1053| `enable_replay_dpo` | `True` | Toggle channel-3 fields. |1054| `rlvr_reward_key` | `"final_reward"` | Key in `TraceExample` to read scalar reward. |1055 1056```python1057from composer_replication.trainer.data_collator import CollatorConfig1058cfg = CollatorConfig(max_seq_len=2048, hint_generator=my_dispatch)1059```1060 1061### `class ComposerDataCollator` — `trainer.data_collator`1062 1063```python1064@dataclass1065class ComposerDataCollator:1066    tokenizer: TokenizerLike1067    config: CollatorConfig = field(default_factory=CollatorConfig)1068 1069    def __call__(1070        self, batch: Sequence[TraceExample]1071    ) -> dict[str, torch.Tensor]: ...1072```1073 1074Build trainer-ready batches from raw traces + optional DPO pairs.1075 1076**Output dict keys** (tested in `spikes/005-integrated-trainer-skeleton/tests/test_data_collator.py`):1077 1078- Channel 1 (always): `input_ids`, `attention_mask`, `response_mask`, `rewards`.1079- Channel 2 (when `enable_sdpo=True` AND batch has at least one error site AND `hint_generator` is set): `ctx_teacher_input_ids`, `sdpo_loss_mask`.1080- Channel 3 (when `enable_replay_dpo=True` AND batch has at least one `dpo_pair`): `dpo_chosen_input_ids`, `dpo_chosen_response_mask`, `dpo_rejected_input_ids`, `dpo_rejected_response_mask`. (Reference logprobs are NOT computed here — the trainer does that pass.)1081 1082```python1083from composer_replication.trainer.data_collator import (1084    ComposerDataCollator, CollatorConfig)1085collator = ComposerDataCollator(tokenizer=tok, config=CollatorConfig())1086batch = collator([{"trace_id": "x", "turns": [...], "final_reward": 1.0}])1087```1088 1089---1090 1091## 11. `composer_replication.diloco`1092 1093DiLoCo outer-loop wrapper around `torchft.local_sgd.DiLoCo`. Optional dep — when `torchft` is missing the package re-export `composer_replication.make_diloco_outer_loop` is `None`.1094 1095### Module-level attributes1096 1097- `DiLoCo: Any` — `torchft.local_sgd.DiLoCo` if importable else `None`.1098- `Manager: Any` — `torchft.manager.Manager` if importable else `None`.1099- `_DummyWork: Any` — `torchft.work._DummyWork` if importable else `None`.1100- `_TORCHFT_AVAILABLE: bool` — whether the imports succeeded.1101 1102```python1103from composer_replication.diloco import _TORCHFT_AVAILABLE, DiLoCo1104```1105 1106### `make_diloco_outer_loop(manager, model_fragments, inner_optimizer, *, ...) -> torchft.local_sgd.DiLoCo`1107 1108```python1109def make_diloco_outer_loop(1110    manager: Any,1111    model_fragments: list[torch.nn.Module],1112    inner_optimizer: torch.optim.Optimizer,1113    *,1114    outer_lr: float = 0.7,1115    outer_momentum: float = 0.9,1116    nesterov: bool = True,1117    sync_every: int = 100,1118    fragment_sync_delay: int = 0,1119    fragment_update_alpha: float = 0.0,1120) -> Any1121```1122 1123Construct a `torchft.DiLoCo` configured with framework-default hyperparams (DiLoCo paper §3.2: `lr=0.7, momentum=0.9, Nesterov`).1124 1125**Parameters**1126 1127| Name | Type | Default | Meaning |1128|---|---|---|---|1129| `manager` | `torchft.Manager` (or duck-typed `MockManager`) | — | Provides `allreduce`, `should_commit`, `current_step`, `start_quorum`, etc. |1130| `model_fragments` | `list[torch.nn.Module]` | — | One module for vanilla DiLoCo; N modules for Streaming DiLoCo. |1131| `inner_optimizer` | `torch.optim.Optimizer` | — | Inner-step optimizer (steps every batch). |1132| `outer_lr` | `float` | `0.7` | Outer SGD lr. |1133| `outer_momentum` | `float` | `0.9` | Outer SGD momentum. |1134| `nesterov` | `bool` | `True` | Nesterov momentum on outer SGD. |1135| `sync_every` | `int` | `100` | Inner steps per outer round. |1136| `fragment_sync_delay` | `int` | `0` | 0 = vanilla; >0 = Streaming DiLoCo (requires CUDA streams). |1137| `fragment_update_alpha` | `float` | `0.0` | 0 = full replacement on sync; >0 = exponential mix. |1138 1139**Returns** a `torchft.local_sgd.DiLoCo` instance — usable as a context manager.1140 1141**Raises** `RuntimeError` if `torchft` is not installed.1142 1143```python1144import torch1145from composer_replication.diloco import make_diloco_outer_loop1146opt = torch.optim.AdamW(model.parameters(), lr=1e-5)1147outer = make_diloco_outer_loop(manager=mgr, model_fragments=[model],1148                               inner_optimizer=opt, sync_every=100)1149with outer:1150    for _ in range(N):1151        opt.zero_grad(); loss.backward(); opt.step()1152```1153 1154---1155 1156## 12. `composer_replication.diloco.serverless`1157 1158ADR-005 serverless DiLoCo executors + object-store all-reduce.1159 1160### `class ReplicaHandle` — `serverless.executor`1161 1162```python1163@dataclass1164class ReplicaHandle:1165    rank: int1166    backend_name: str1167    metadata: dict[str, Any] = field(default_factory=dict)1168```1169 1170Opaque handle returned by `ServerlessExecutor.launch_replicas`. `metadata` is backend-specific.1171 1172```python1173from composer_replication.diloco.serverless import ReplicaHandle1174h = ReplicaHandle(rank=0, backend_name="local_process",1175                  metadata={"pid": 12345})1176```1177 1178### `class ServerlessExecutor` (Protocol) — `serverless.executor`1179 1180```python1181@runtime_checkable1182class ServerlessExecutor(Protocol):1183    backend_name: str1184    supports_inter_replica_network: bool1185 1186    def launch_replicas(1187        self,1188        n_replicas: int,1189        entrypoint: str | Callable[..., Any],1190        entrypoint_args: Mapping[str, Any],1191        *,1192        gpu: str | None = None,1193        timeout: int = 3600,1194    ) -> list[ReplicaHandle]: ...1195 1196    def poll(self, handle: ReplicaHandle) -> str: ...1197    def stream_logs(self, handle: ReplicaHandle, *, n_lines: int = 200) -> str: ...1198    def cancel(self, handle: ReplicaHandle) -> None: ...1199    def collect(1200        self, handles: list[ReplicaHandle], *, timeout: int | None = None,

Showing the first 1,200 of 1880 lines. Download the file for the rest.