Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
USER_GUIDE.md704 linesDownload Raw Back to docs
1# Composer Replication Framework — User Guide2 3A zero-to-training walkthrough for the open replication of Cursor Composer 2.5.4Pace: an ML engineer who knows GRPO/DPO at a textbook level but has never5opened this repo. Every step references real code, and every kwarg name6listed below has been imported and verified against7`composer_replication/` source.8 9---10 11## 1. What is this framework?12 13A pure-PyTorch replication of the **3-channel composer loss** that powers14agentic-coding model training. One model, one optimizer, three additive15loss terms — composed every step:16 17```18                      ┌────────────────────────────────────────────┐19                      │            compose_loss(model, batch)       │20                      └────────────────────────────────────────────┘21                                          │22        ┌─────────────────────────────────┼─────────────────────────────────┐23        ▼                                 ▼                                 ▼24┌───────────────────┐          ┌──────────────────────┐         ┌──────────────────────┐25│  Channel 1 (RL)   │          │  Channel 2 (SDPO)    │         │  Channel 3 (replay)  │26│  GRPO            │          │  hint-distillation   │         │  multi-teacher DPO   │27│  → lm_ce stub in │          │  generalized JSD     │         │  on (chosen,         │28│  verification    │          │  student vs teacher  │         │  rejected) pairs     │29│  harness         │          │  (hint-conditioned)  │         │  from N teachers     │30└─────────┬─────────┘          └──────────┬───────────┘         └──────────┬───────────┘31          │  weight = 1 (always on)       │ alpha_sdpo            beta_replay │32          └────────────────┬──────────────┴─────────────────┬──────────────────┘33                           ▼                                ▼34                   total = lm_ce + α·sdpo_jsd + β·trace_replay_dpo35                   (channel auto-disables if its weight=0 OR its inputs are missing)36```37 38Two API surfaces, on purpose:39 40- **Verification harness** — `compose_loss(model, batch, ...)` is a free41  function (channel 1 = LM cross-entropy, the GRPO limit under deterministic42  rewards). Use it for CPU smokes, unit tests, and gradient-flow debugging.43- **Production trainer** — `ComposerReplicationTrainer` is a `trl.GRPOTrainer`44  subclass that overrides `_compute_loss` with the same 3 channels on top of45  TRL's real reward + advantage machinery.46 47The verification harness is what you'll use for sections 2–6; the production48trainer (and its alternates VeRL/PRIME-RL/Monarch) is section 8.49 50Source of truth: `composer_replication/loss.py` for `compose_loss`,51`composer_replication/trainer/composer_trainer.py` for the trainer subclass.52 53---54 55## 2. Install — which extras to pick56 57Always start with the core install:58 59```bash60git clone https://huggingface.co/Codeseys/composer-replication-framework61cd composer-replication-framework62pip install -e .63```64 65> **Branch note (resolved 2026-06-09).** `main` is the canonical branch and is kept in66> sync with `master` (`main == master`). A fresh clone of `main` has the complete tree67> (incl. `make_dr_grpo_config` / `make_po_config`), so no branch switch is needed.68> Historically `main` lagged `master` — if you ever see an `ImportError` on those symbols,69> the clone is stale; `git fetch && git checkout main` (or pin a current SHA) fixes it.70> See [`docs/HF_REPO_LAYOUT.md`](HF_REPO_LAYOUT.md).71 72That gets you `torch>=2.0` + `transformers>=4.46` and is enough for the73verification harness on CPU (sections 3, 5, 6).74 75The seven optional extras are declared in `pyproject.toml` `[project.optional-dependencies]`:76 77```78                              Do you need …79                                    │80        ┌──────────────────────────┼──────────────────────────┐81        ▼                          ▼                          ▼82   real teacher calls          DiLoCo on                  production83   over OpenRouter?            >1 GPU?                    GRPO training?84        │                          │                          │85        │ yes                      │ yes                      │ yes86        ▼                          ▼                          ▼87   pip install -e ".[replay]"  pip install -e ".[diloco]"  pip install -e ".[train]"88   (httpx)                     (torchft-nightly)           (trl, peft, accelerate, datasets)89        │                          │                          │90        │ + want CPU-side          │ + scaling beyond a        │ + want PRIME-RL91        │ DPO normalization?       │ single host?              │ (Recipe C)?92        ▼                          ▼                          ▼93   pip install -e \".[replaysim]\"  pip install -e \".[serverless]\"  pip install -e \".[prime-rl]\"94   (data-juicer; depends         (fsspec, huggingface_hub)    (prime-rl>=0.5)95    on [replay])96                                                              │ + Monarch actor mesh?97                                                              ▼98                                                          pip install -e \".[monarch]\"99                                                          (monarch>=0.4.1)100```101 102Quick decision table:103 104| Goal                                                  | Install                                  |105|-------------------------------------------------------|------------------------------------------|106| CPU smoke / verification (sections 3, 5, 6)           | `pip install -e .`                       |107| Section 4 (replaysim DJNormalizer)                    | `pip install -e ".[replaysim]"`          |108| Section 7 dev loop (LocalProcessExecutor + file://)   | `pip install -e ".[serverless]"`         |109| Real DiLoCo outer-loop                                | `pip install -e ".[diloco,serverless]"`  |110| Section 8 Recipe A (TRL GRPO)                         | `pip install -e ".[train]"`              |111| Section 8 Recipe C (PRIME-RL)                         | `pip install -e ".[prime-rl]"`           |112| Section 8 Recipe C+D (PRIME-RL + Monarch)             | `pip install -e ".[prime-rl,monarch]"`   |113| Everything for development                            | `pip install -e ".[dev]"`                |114 115---116 117## 3. Quickstart: `examples/qwen_05b_quickstart` end-to-end on CPU118 119The fastest way to convince yourself the framework works on a real HF model.120~3–5 min wall-clock on CPU, ~1 GB disk for Qwen2.5-0.5B weights.121 122```bash123pip install -e .124python examples/qwen_05b_quickstart/run.py125```126 127What the script does (read the source at128`examples/qwen_05b_quickstart/run.py`):129 1301. Pin RNG (`random.seed(42)`, `torch.manual_seed(42)`) so the per-step131   numbers below are reproducible.1322. Load `Qwen/Qwen2.5-0.5B-Instruct` on CPU in fp32, set `model.train()`.1333. `batch = build_batch(tokenizer, device="cpu")` — a real chat-template-formatted134   batch with all keys the 3-channel composer might consume.1354. Five backward steps with `compose_loss(model, batch, alpha_sdpo=0.1,136   beta_replay=0.05)`; an `AdamW(lr=1e-5)` optimizer; finite-grad check137   after each step.138 139Expected output (transcribed from `examples/qwen_05b_quickstart/run.log`):140 141```142[quickstart] loading Qwen/Qwen2.5-0.5B-Instruct (CPU, fp32) ...143[quickstart] loaded — 0.494B params144[quickstart] building real chat-template batch ...145[quickstart] running 5 backward steps ...146  step 0: total=0.7390  lm_ce=0.7358  sdpo=0.0000  dpo=0.0639  finite=True147  step 1: total=0.0379  lm_ce=0.0351  sdpo=0.0000  dpo=0.0563  finite=True148  step 2: total=0.0122  lm_ce=0.0110  sdpo=0.0000  dpo=0.0240  finite=True149  step 3: total=0.0060  lm_ce=0.0055  sdpo=0.0000  dpo=0.0098  finite=True150  step 4: total=0.0031  lm_ce=0.0029  sdpo=0.0000  dpo=0.0044  finite=True151========================================================152  Initial loss: 0.7390  →  Final loss: 0.0031  →  Reduction: 99.6%153  Verdict: PASS154========================================================155```156 157How to read this:158 159- **`total` collapses by ~99%.** The model successfully memorizes the160  single batch — exactly what you expect from an SGD pass on a 0.5B model161  with one fixed input. This is a wiring check, not a generalization claim.162- **`lm_ce` carries almost all the magnitude.** Channel 1 (the GRPO stub)163  is doing the work — the response tokens are short and have low entropy164  under the trained model.165- **`sdpo=0.0000` on every step.** Channel 2 has auto-disabled because the166  default `build_batch` does not include `ctx_teacher_input_ids`. Compare167  the conditional in `compose_loss`:168  ```python169  if (alpha_sdpo > 0.0170      and "ctx_teacher_input_ids" in inputs171      and inputs["ctx_teacher_input_ids"].numel() > 0):172  ```173  — channel auto-off if either the weight or the inputs are missing.174- **`dpo > 0` and trending down.** The batch *does* include175  `dpo_chosen_input_ids`, `dpo_chosen_response_mask`,176  `dpo_chosen_ref_logprobs` (and the rejected counterparts), so channel 3177  is live.178- **`finite=True`** — every step's `p.grad` was finite for every parameter.179  This is the wiring contract; if it ever flips to `False` the smoke fails.180 181If you see `Verdict: PASS`, the framework is correctly installed and182gradients flow through all live channels. You are ready for section 4.183 184---185 186## 4. Adding the trace-replay channel187 188The quickstart batch *had* DPO inputs, but they were synthetic — the189`build_batch` helper bakes them in. To get **real** DPO pairs from190multi-teacher disagreement, use the replaysim package.191 192### 4a. Spin up `replay_trace`193 194```python195import asyncio196from composer_replication import (197    DEFAULT_TEACHERS, replay_trace, extract_dpo_pairs,198)199 200# Trace must be a list[TraceState]; see composer_replication/teacher_replay.py201# for the exact TypedDict shape. Each state holds a chat-messages prefix +202# the student's actual action at that step.203states = [...]   # your frozen agentic trace; see spike 001 for a 50-step example204 205teacher_actions = asyncio.run(206    replay_trace(207        states=states,208        teachers=DEFAULT_TEACHERS,    # claude-opus-4.7 + gpt-5 + deepseek-v4-pro209        max_total_usd=10.0,           # hard ceiling (spike 001 measured $0.98/trace mean)210    )211)212```213 214The 3 teachers are queried in parallel via OpenRouter215(`OPENROUTER_API_KEY` in env or `~/.hermes/.env`), latencies recorded,216costs tracked.217 218### 4b. Get `DPOPair`s from disagreement219 220```python221pairs = extract_dpo_pairs(222    states=states,223    teacher_actions=teacher_actions,224    agreement_threshold=2,    # at least 2/3 teachers must agree on the chosen action225)226```227 228Each pair is a `DPOPair` TypedDict with the exact shape the229`DJNormalizer` and downstream training expects:230 231```python232class DPOPair(TypedDict):233    state_id:           str234    state_messages:     list[dict]    # conversation context235    chosen:             str           # teacher-consensus action236    rejected:           str           # student action237    n_teachers_agreeing: int238```239 240(verified in `composer_replication/teacher_replay.py:99–105`).241 242### 4c. Run `DJNormalizer` with `default.yaml`243 244```python245from composer_replication.replaysim import DJNormalizer246 247normalizer = DJNormalizer()        # uses recipes/replaysim/default.yaml248normalized = normalizer.normalize(pairs)249# → list[NormalizedDPOPair]250```251 252`DJNormalizer` shells out to data-juicer's `DefaultExecutor` under the hood253(file-in / file-out contract). The default recipe at254`composer_replication/recipes/replaysim/default.yaml` runs four CPU-only ops255in order:256 2571. `text_length_filter` (8 ≤ chars ≤ 32000) on `chosen` and `rejected`2582. `words_num_filter` (2 ≤ words ≤ 4096) on both2593. `special_characters_filter` (≤50% non-alpha) on both2604. `document_deduplicator` (per-batch hashing, lowercase, ignore non-character) on `chosen`261 262Records carry **two parallel shapes** for `chosen`/`rejected`:263- flat strings (`chosen`, `rejected`) → consumed by data-juicer's text_key-based filters264- chat-messages lists (`chosen_messages`, `rejected_messages`) → preserved for chat-aware ops + round-trip265 266This dual-shape design (verified in the test267`test_dpo_pair_to_dj_record_shape`,268`composer_replication/replaysim/tests/test_replaysim.py:44`) is what269unblocked the data-juicer integration in Wave 14.270 271### 4d. The 3-record fixture from spike 001272 273The fixture lives at274`spikes/001-teacher-replay-cost/states.jsonl` (50 states) and275`spikes/001-teacher-replay-cost/results.jsonl` (the teacher responses, all276priced and timed). The first 3 states are:277 278```jsonl279{"id": "state-000", "task": "Fix the failing test in tests/test_auth.py::test_login_with_email", ...}280{"id": "state-001", "task": "Add rate-limiting middleware to the Flask app", ...}281{"id": "state-002", "task": "Refactor the parse_config function — it's 200 lines and has 3 responsibilities", ...}282```283 284For each, all 3 teachers answered (claude-opus-4.7, gpt-5, deepseek-v4-pro);285agreement on the `(c)` choice for state-000 and state-001 (read more286files / check schema first) drives a clean DPO pair where the student's287action becomes the rejected. For state-002, all 3 agreed on `(c)` (write288characterization tests first) → another clean pair. These three records289pass through the `DJNormalizer` default recipe unchanged (length, words,290special-char ratios all in bounds; no duplicates).291 292The full 50-state trace cost **$0.98 mean** end-to-end across all three293teachers (spike 001 verdict). The framework's cost ceiling294(`max_total_usd`) and VOI gating drop this to ~$0.30/trace projected.295 296### 4e. End-to-end one-liner297 298```python299from composer_replication.replaysim import replay_and_normalize_trace300 301teacher_actions, normalized_pairs = await replay_and_normalize_trace(302    states=states,303    teachers=DEFAULT_TEACHERS,304    agreement_threshold=2,305    max_total_usd=10.0,306)307```308 309(`async def`; for sync callers use the sibling `replay_and_normalize_trace_sync`310in `composer_replication.replaysim.normalize`.)311 312---313 314## 5. Switching DPO → SimPO: one kwarg315 316```python317components = compose_loss(318    model, batch,319    alpha_sdpo=0.1,320    beta_replay=0.05,321    dpo_variant="simpo",      # ← the only line that changes322    simpo_beta=2.0,           # paper default323    simpo_gamma=1.0,          # paper default324)325```326 327The kwarg is verified in `compose_loss`'s signature328(`composer_replication/loss.py:81`):329 330```python331dpo_variant: Literal["dpo", "simpo"] = "dpo",332```333 334### What changes in the loss curve335 336- **Channel 3 input requirements drop.** `compose_loss` no longer reads337  `dpo_chosen_ref_logprobs` / `dpo_rejected_ref_logprobs`. Reference-model338  VRAM cost goes to zero. (Source: `composer_replication/loss.py:111–113`339  and `composer_replication/distillation/simpo.py:23–27`.)340- **Loss scale shifts.** Standard DPO is341  `-logsigmoid(β·[(logπ(c) - logπ_ref(c)) - (logπ(r) - logπ_ref(r))])`.342  SimPO is `-logsigmoid(β·[avg_logπ(c) - avg_logπ(r)] - γ)` — average343  per-token log-prob (length-normalized) and a constant target margin γ.344- **Loss is ≤ DPO loss when chosen/rejected separation is large.** The345  unit test `test_simpo_loss_lower_for_better_separation`346  (`composer_replication/distillation/tests/test_distillation_losses.py:35`)347  verifies that a wider chosen-vs-rejected gap drives lower SimPO loss —348  meaning, in practice, SimPO curves are *steeper* than DPO when the349  preference signal is strong, and *flatter* when it's weak.350- **No KL-against-reference regularization.** This is both the upside (no351  ref-model serving) and the risk (more tendency to drift). Watch for352  reward-hacking-style degeneracies if your preference data has noise.353 354### When to use SimPO355 356- **GPU-poor.** You can't afford to keep a frozen reference policy resident357  alongside the trainer.358- **Cold-start preference data.** Length-normalization (avg_logπ vs sum)359  helps when chosen/rejected lengths are wildly imbalanced — common in360  agentic traces where the student's failed attempt is short and the361  teacher's correction is long.362- **You don't have ref logprobs precomputed.** SimPO needs nothing from363  the reference policy.364 365When to **stay on DPO**: when you need the explicit KL anchor against366a known-good reference (e.g., when training over a long horizon and you367want to bound the drift), or when your preference data is very noisy and368the reference acts as a regularizer.369 370---371 372## 6. Adding TAID / Entropy-Aware OPD wrappers373 374Channel 2 (SDPO/OPSD) can be replaced by **TAID** (Sakana AI,375arXiv:2501.16937) for capacity-gap distillation, or by376**Entropy-Aware OPD** (ICLR 2026 Spotlight) for per-token forward/reverse-KL377gating. Both are wired through `compose_loss`:378 379```python380sdpo_wrapper: Literal["none", "taid", "entropy_opd"] = "none",381taid_t: float | None = None,         # current TAID interpolation coeff382entropy_opd_h_max: float | None = None,383```384 385(verified at `composer_replication/loss.py:82–93`.)386 387### TAID (upstream-faithful port)388 389> **Wave 15 rewrite, breaking change.** The previous in-tree TAID was390> algorithmically different from the paper (it mixed in probability space391> against a frozen step-0 student snapshot and wrapped a symmetric JSD392> criterion). It has been replaced with an upstream-faithful port:393> logit-space mix, current-student-detached anchor, forward-KL criterion.394> Old kwargs `taid_schedule_step`, `taid_total_steps`, `taid_schedule`,395> `taid_alpha_min`, `taid_alpha_max`, plus `inputs["student_init_logits"]` /396> `inputs["student_init_input_ids"]` are **gone**. They have no upstream397> analogue. Use `taid_t` (and optionally `TAIDScheduler`) instead.398 399The TAID criterion is forward-KL against a logit-space-interpolated target:400 401```402p_t = softmax( (1 - t) · stop_grad(student_logits) + t · teacher_logits )403L   = - mean_token  Σ_v  p_t(v) · log_softmax(student_logits)(v)404```405 406where `t ∈ [0, 1]` is the interpolation coefficient. At `t=0` the target407is the (detached) student itself — the loss is the entropy of that408distribution and contributes no gradient to the student. At `t=1` it409reduces to standard forward-KL distillation against the teacher.410 411The schedule that produces `t` is the **trainer's** responsibility. The412package ships an optional `TAIDScheduler` that mirrors the paper's413adaptive momentum scheme:414 415```python416from composer_replication.distillation import TAIDScheduler417 418sched = TAIDScheduler(num_train_steps=10_000)   # paper defaults419for step in range(num_train_steps):420    components = compose_loss(421        model, batch,422        sdpo_wrapper="taid",423        taid_t=sched.t,424    )425    components.total.backward(); optimizer.step()426    sched.update_t(components.sdpo_jsd.detach(), global_step=step)427```428 429`TAIDScheduler` defaults match upstream: `t_start=0.4`, `t_end=1.0`,430`alpha=5e-4`, `beta=0.99`. Pass `disable_adaptive=True` to fall back to431the deterministic linear schedule432`t = t_start + progress · (t_end - t_start)`.433 434If you want a simple fixed schedule (no scheduler), just compute `t`435yourself and pass it in — `compose_loss` validates `taid_t ∈ [0, 1]`.436 437### Upstream-parity test438 439`composer_replication/distillation/tests/test_taid_parity.py` skip-imports440the upstream reference at `/tmp/taid-clone` (clone with441`git clone --depth 1 https://github.com/SakanaAI/TAID /tmp/taid-clone`)442and asserts our `taid_loss(student, teacher, mask, t)` matches upstream443`TAID.compute_loss(...)` within `atol=rtol=1e-5` across `t ∈ {0.0, 0.1, 0.4,4440.5, 0.9, 1.0}`. This is the load-bearing parity guarantee.445 446### Entropy-Aware OPD447 448Drop-in for channel 2 — gates between forward KL (mode-covering) and449reverse KL (mode-seeking) per token, weighted by the teacher's entropy:450 451```452L = Σ_t  w(t) · KL_fwd_t  +  (1 - w(t)) · KL_rev_t453w(t) = clamp(H_teacher(t) / h_max, 0, 1)454```455 456`entropy_opd_h_max=None` (the default) auto-sets `h_max = log(V)` (the457maximum-entropy bound for a vocab-V softmax).458 459### Boundary-condition unit test (proof of correctness)460 461The test `test_taid_loss_t_zero_target_matches_detached_student`462(`composer_replication/distillation/tests/test_distillation_losses.py`)463pins TAID's `t=0` invariant — the teacher is *completely* hidden from the464gradient because the target collapses to `softmax(student.detach())`:465 466```python467def test_taid_loss_t_zero_target_matches_detached_student():468    s1 = torch.randn(1, 2, 4, requires_grad=True)469    teacher_a = torch.zeros(1, 2, 4); teacher_a[..., 0] = 10.0470    teacher_b = torch.zeros(1, 2, 4); teacher_b[..., 3] = 10.0471    mask = torch.ones(1, 2)472    loss_a = taid_loss(s1, teacher_a, mask, t=0.0)473    loss_b = taid_loss(s1, teacher_b, mask, t=0.0)474    # Two completely different teachers must give the same loss at t=0.475    assert abs(float(loss_a) - float(loss_b)) < 1e-6476```477 478This is the load-bearing test for TAID: if the `t=0` endpoint ever leaks479teacher signal into the gradient, this test fires and the contract is480broken. The companion test `test_taid_loss_t_one_is_pure_forward_kl`481pins the `t=1` endpoint by hand-computing `-Σ p_teacher · log_q` and482asserting equality.483---484 485## 7. Going multi-replica with serverless DiLoCo486 487DiLoCo is the outer-loop optimizer that lets you run N replicas in488parallel, sync them every H inner steps, and tolerate slow links — see489`docs/adrs/ADR-005-serverless-diloco.md` for the design. The framework490gives you three increasingly-distant deployments:491 492### Step 1 — `LocalProcessExecutor` for development493 494```python495from composer_replication.diloco.serverless import (496    LocalProcessExecutor, ObjectStoreAllReduce,497)498import tempfile499 500with tempfile.TemporaryDirectory() as td:501    rendezvous = ObjectStoreAllReduce(td, rank=0, world_size=4)502    executor = LocalProcessExecutor()503    handles = executor.launch_replicas(504        n_replicas=4,505        entrypoint="composer_replication.diloco.serverless.replica_entrypoint",506        entrypoint_args={"rendezvous_uri": td, "rank_env": "REPLICA_RANK"},507    )508    results = executor.collect(handles, timeout=600)509```510 511`LocalProcessExecutor` (`composer_replication/diloco/serverless/executor.py:160`)512spawns N child processes via `multiprocessing.get_context("spawn")` and513sets `REPLICA_RANK={0..N-1}` in each child's env. It satisfies the514`ServerlessExecutor` Protocol (line 35) — the same Protocol the cloud515adapters implement. So the dev-loop code is byte-identical to the cloud516deploy: only the executor instance changes.517 518### Step 2 — `ObjectStoreAllReduce` as the rendezvous519 520```python521# Local file:// for tests522rendezvous = ObjectStoreAllReduce("/tmp/diloco-runs/run42/", rank=0, world_size=4)523 524# Real S3 (after `pip install -e .[serverless]`)525rendezvous = ObjectStoreAllReduce(526    "s3://my-bucket/diloco-runs/run42/",527    rank=0, world_size=4,528    timeout_s=1800.0,529)530```531 532The communication pattern is `S3 PutObject + N GetObjects` once per533inner H steps (matches DiLoCo's actual sync cadence,534arXiv:2311.08105 §3.2). For 1B-param bf16, that's ~2 GB / 30 minutes535per replica — well within S3 free-tier. On the inner side the framework536exposes a `MockManager` that drops into the `torchft.Manager` slot, so537you can validate the rendezvous logic before plugging in real torchft538(verified by `test_serverless_diloco_integration.py`).539 540### Step 3 — point at `ModalExecutor` / `HFJobsExecutor`541 542```python543# Modal (skeleton at composer_replication/diloco/serverless/modal.py)544from composer_replication.diloco.serverless.modal import ModalExecutor545executor = ModalExecutor(image="modal:python3.11", gpu="A100")546 547# HuggingFace Jobs (skeleton at composer_replication/diloco/serverless/hf_jobs.py)548from composer_replication.diloco.serverless.hf_jobs import HFJobsExecutor549executor = HFJobsExecutor(hardware="a10g-large")550 551# Same Protocol — same launch_replicas / poll / collect calls as Local552handles = executor.launch_replicas(n_replicas=4, ...)553```554 555Both adapters check their cloud SDK at `__init__` time (not at module556import) so they don't break the package if you don't have `modal` or557`huggingface_hub` installed. Production maturity: dev-ready for cloud558trial; per ADR-005, full HA-cluster fan-out lives in v0.2+.559 560---561 562## 8. Picking an RL backend563 564Four canonical recipes, each tied to an upstream framework. Source:565`docs/INTEGRATION_ARCHITECTURE.md` Recipes A–D plus566`docs/adrs/ADR-006-rl-frameworks.md`.567 568### Recipe A — TRL `GRPOTrainer` subclass569 570`ComposerReplicationTrainer` is a `trl.GRPOTrainer` subclass that571overrides `_compute_loss(model, inputs)` to compose the same 3 channels572on top of TRL's real reward + advantage machinery. Install:573`pip install -e ".[train]"`. Then:574 575```python576from composer_replication import ComposerReplicationTrainer577trainer = ComposerReplicationTrainer(model=..., reward_funcs=[...], ...)578trainer.train()579```580 581**When to use it:** This is the v0.0/v0.1 recommended path. You want582real GRPO with rewards, you have HF model + dataset + (mostly) standard583GRPO infrastructure, and you don't need >100B-param scale. TRL is584mature, the trainer is a small subclass, and the same `compose_loss`585math runs in both the verification harness and in production with no586re-coding.587 588→ See `docs/INTEGRATION_ARCHITECTURE.md` § "Recipe A: TRL `GRPOTrainer`589subclass" (line 205).590 591### Recipe B — VeRL custom `adv_estimator` + DataProto extension592 593VeRL replaces TRL's reward+advantage machinery with a Ray-driven actor594graph that's specifically optimized for distributed RL training of595large LMs. Composition with the framework: extend `DataProto` with the596hint-conditioned columns + DPO pair fields, register a custom597`adv_estimator` that calls the same `compose_loss` body.598 599**When to use it:** You're past 7B-param, you have multi-host setup600(Ray cluster), and TRL's single-process trainer is the bottleneck. VeRL601is the recommended scale path for v0.2+. Trade-off: the extension surface602is larger and the docs are sparser than TRL's.603 604→ See `docs/INTEGRATION_ARCHITECTURE.md` § "Recipe B: VeRL custom605`adv_estimator`" (line 289).606 607### Recipe C — PRIME-RL with DPPO-clip details608 609`composer_replication/recipes/prime_rl/composer_loss.py` ships a610`loss_fn(inputs, *, alpha_sdpo=0.0, beta_dpo=0.0, dppo_mask_high=0.2,611dppo_mask_low=0.2, adv_tau=1.0, kl_tau=1e-3)` adapter that maps612PRIME-RL's `LossInputs` struct (1-D per-sample tensors:613`trainer_logprobs`, `inference_logprobs`, `teacher_logprobs`,614`advantages`, `loss_mask`) into our 3-channel composition.615 616The DPPO+KL bit is what makes PRIME-RL distinctive — and we mirror617PRIME-RL's upstream `default_loss_fn` exactly (verified against618`prime_rl/trainer/rl/loss.py` lines 116-165):619 620```python621log_ir       = trainer_logprobs - inference_logprobs622ir           = exp(log_ir)                                  # importance ratio623probs_diff   = exp(trainer_logprobs) - exp(inference_logprobs)624invalid_high = probs_diff >  dppo_mask_high                 # for positive-advantage tokens625invalid_low  = probs_diff < -dppo_mask_low                  # for negative-advantage tokens626invalid      = where(advantages > 0, invalid_high, invalid_low)627keep         = loss_mask & ~invalid628pg_loss      = keep      * (adv_tau * advantages) * ir629kl_loss      = loss_mask * log_ir**2630loss         = (-pg_loss + kl_tau * kl_loss).sum()631```632 633Three things to remember: (1) the mask gate is on **probability-space**634`exp(trainer_lp) - exp(inference_lp)`, not on the log-ratio; (2) the635policy-gradient term is multiplied by the importance ratio636`exp(trainer_lp - inference_lp)`, not by `trainer_lp` directly (proper637IS-corrected gradient, not REINFORCE); (3) the mask is **conditioned on638the sign of the advantage** — positive-advantage tokens are dropped on639the upper bound, negative-advantage tokens on the lower. Defaults640`dppo_mask_high=dppo_mask_low=0.2` and `adv_tau=1.0, kl_tau=1e-3`641match PRIME-RL's `DefaultLossConfig` (all fields `Field(..., ge=0)`).642SDPO (channel 2) is gated `NotImplementedError` in v0 because PRIME-RL643exposes log-probs, not full vocab logits; channel 3 (trace-replay DPO)644emits a warning if `beta_dpo != 0`.645 646**When to use it:** You're already in the PRIME-Intellect / decentralized647training universe, you want INTELLECT-style scaling on a long-horizon648training run, and DPPO masking is part of your existing reward+advantage649recipe. Install: `pip install -e ".[prime-rl]"`.650 651→ See `composer_replication/recipes/prime_rl/prime_rl_recipe.md` and652`docs/INTEGRATION_ARCHITECTURE.md` § "Recipe C: TorchForge + Monarch"653(line 356).654 655### Recipe C+D — Monarch as actor mesh656 657Monarch (the actor framework underpinning TorchForge) hosts the658trainer/generator/manager actors in a topology-aware mesh. The framework659ships *skeleton* actor definitions at660`composer_replication/recipes/monarch/actors.py` (TrainerActor,661GeneratorActor) and a layout doc at `monarch_actor_layout.md`. v0662intentionally *fails fast* if you try to instantiate them663(`raise NotImplementedError("v0 skeleton; deferred to v0.2 per ADR-006")`)664because the upstream Monarch API is still moving.665 666**When to use it:** Reference-pattern reading only in v0. Decision point667is v0.2 once the upstream actor API stabilizes. Treat the skeleton as668shape-of-the-final-answer documentation, not as a production target.669Install: `pip install -e ".[prime-rl,monarch]"` for the full surface.670 671→ See `composer_replication/recipes/monarch/monarch_actor_layout.md`672and `docs/adrs/ADR-006-rl-frameworks.md`.673 674---675 676## Common pitfalls + what tests catch them677 678The framework's test suite (266 passing / 62 skipped, canonical count in docs/V1_V8_COVERAGE.md) is structured so each pitfall has a679specific test-file home. If you hit one of these in production, the680corresponding test is your fastest reproducer.681 682| Pitfall                                                                                       | Symptom                                              | Test file (catches it)                                                                                              |683|-----------------------------------------------------------------------------------------------|------------------------------------------------------|----------------------------------------------------------------------------------------------------------------------|684| Forgetting `taid_schedule_step` when `sdpo_wrapper="taid"`                                    | `ValueError` at first step                           | `composer_replication/tests/test_compose_loss_integration.py` (kwarg validation)                                    |685| TAID α=0 endpoint leaks teacher signal                                                        | Teacher swap changes the loss when α should be 0     | `test_taid_loss_alpha_zero_ignores_teacher` in `composer_replication/distillation/tests/test_distillation_losses.py:153` |686| TAID α=1 endpoint differs from plain SDPO                                                     | Bit-difference vs reference SDPO at the schedule end | `test_taid_blended_logits_endpoints` in `composer_replication/distillation/tests/test_distillation_losses.py:115`   |687| SimPO loss not differentiable through the loss-of-sigmoid path                                | `chosen.grad is None` after backward                 | `test_simpo_loss_differentiable` in `composer_replication/distillation/tests/test_distillation_losses.py:50`        |688| SimPO shape-mismatch slips through silently                                                   | Broadcasting bug, NaN downstream                     | `test_simpo_loss_shape_mismatch_raises` in `composer_replication/distillation/tests/test_distillation_losses.py:61` |689| Entropy-OPD failing to zero out when distributions match                                      | Loss > 0 when student≡teacher                        | `test_entropy_aware_opd_zero_when_distributions_match` in `composer_replication/distillation/tests/test_distillation_losses.py:217` |690| Entropy of one-hot ≠ 0 / entropy of uniform ≠ log(V)                                          | Wrong gating weights w(t)                            | `test_teacher_entropy_one_hot_is_zero` and `test_teacher_entropy_uniform_is_log_v` in `composer_replication/distillation/tests/test_distillation_losses.py:175,183` |691| `DJNormalizer` records missing the chat-messages shape                                        | Filters silently no-op or crash                      | `test_dpo_pair_to_dj_record_shape` in `composer_replication/replaysim/tests/test_replaysim.py:44`                   |692| `DJNormalizer` round-trip drops `state_messages` / metadata                                   | Lost provenance                                      | `test_dj_record_to_normalized_roundtrip` and `test_dj_record_to_normalized_preserves_state_messages` in `composer_replication/replaysim/tests/test_replaysim.py` |693| `ObjectStoreAllReduce` accepts an out-of-bounds rank                                          | Silent corruption of the all-reduce average          | `test_object_store_allreduce_init_validates_rank` in `composer_replication/diloco/serverless/tests/test_serverless_local.py:31` |694| `ObjectStoreAllReduce(world_size=1)` doesn't passthrough cleanly                              | False all-reduce on single replica                   | `test_object_store_allreduce_world_size_1_passthrough` in `composer_replication/diloco/serverless/tests/test_serverless_local.py:46` |695| `LocalProcessExecutor` doesn't propagate child failures to `collect()`                        | Silent test pass when a replica crashed              | `test_serverless_diloco_integration.py` in `composer_replication/diloco/serverless/tests/`                          |696| PRIME-RL adapter accidentally uses `(B, T)` shape instead of per-sample `(seq,)`              | Shape mismatch / wrong reduction                     | `composer_replication/recipes/prime_rl/tests/test_composer_loss.py` (10 tests covering shape and DPPO mask edges)   |697| Channel 2/3 fails to auto-disable when its inputs are absent                                  | Crash on missing key, not graceful skip              | `composer_replication/tests/test_compose_loss_integration.py` (`(a) defaults reproduce existing compose_loss output bit-exact`) |698 699Run the full suite with `pytest` from the repo root.700 701---702 703**File path:** `docs/USER_GUIDE.md` (repo-relative)704