# How Batched MTP Decode Works in MTPLX's `mtp_batch` Lane

> Learn how batched MTP decode works in MTPLX's mtp_batch lane. Explore its three-stage pipeline: job creation, cohort formation, and decode window measurement for efficient processing.

- Repository: [Youssof Altoukhi/MTPLX](https://github.com/youssofal/MTPLX)
- Tags: internals
- Published: 2026-09-13

---

**Batched MTP decode in MTPLX processes multiple generation requests as a fixed-width cohort through a three-stage pipeline: job creation with `MTPBatchJob`, cohort formation via `MTPBatchGenerationService`, and decode window measurement using `mark_decode_started` and `terminal_perf_s` timestamps.**

The MTPLX inference engine accelerates large language model serving through Multi-Token-Parallel (MTP) decoding. Its `mtp_batch` lane groups individual requests into sealed cohorts that execute together in a single forward pass, enabling efficient batched MTP decode while preserving per-request latency metrics and throughput statistics.

## Stage 1: Job Creation and Enqueue

When an OpenAI-compatible request routes to the `mtp_batch` lane, the system instantiates an **`MTPBatchJob`** object (defined in [`mtplx/server/mtp_batch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/mtp_batch.py) at lines 84‑102). This job encapsulates the prompt token IDs, sampling configuration, token callbacks, and a **compatibility key** that determines which fixed-width lane the cohort will use.

The job creation process captures critical metadata for the batched decode:

- **`compatibility_key`** – A tuple containing model identifiers and sampling configurations that dictates cohort eligibility and lane selection.
- **`session_bank` hooks** – Optional bypass flags for the session cache, observable in the request handling logic at [`mtplx/server/openai.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/openai.py) (lines 23104‑23107).
- **Cancellation and error handling** – Each job carries a `cancel_event` and `cancel_error` callback to handle preemption during the batch window.

## Stage 2: Cohort Formation and Execution

The **`MTPBatchGenerationService`** accumulates jobs until sufficient requests arrive to form a sealed cohort. The service implements the `_lane_for_real_width` method to select the narrowest installed native graph capable of accommodating the cohort width.

For each job entering the cohort, the service constructs an **`A3BMTPBatchRequest`** that registers two critical callbacks (lines 86‑89 in [`mtplx/server/mtp_batch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/mtp_batch.py)):

- **`on_token=job.emit_token`** – Streams generated tokens back to the client.
- **`on_decode_start=job.mark_decode_started`** – Triggers timing measurement when the first token enters the decode phase.

The low-level driver **`generate_a3b_mtp_batch`** (located in [`mtplx/a3b_mtp_batch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/a3b_mtp_batch.py)) executes the fixed-width lane once, returning a result stream for each constituent request.

## Stage 3: Decode Window and Statistics

While the driver executes, the batched MTP decode system measures per-request latency with microsecond precision. When the first decoded token reaches the model-owner thread, **`job.mark_decode_started`** records the wall-clock time via `time.perf_counter()` (lines 151‑154 in [`mtplx/server/mtp_batch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/mtp_batch.py)).

Upon stream completion, the **`_complete_cohort_job`** method computes the decode window through the following logic (lines 638‑677 in [`mtplx/server/mtp_batch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/mtp_batch.py)):

1. **Determine decode end time** – Uses `terminal_perf_s` from the driver if available and valid; otherwise falls back to `completed_s` (current time).
2. **Calculate elapsed duration** – `decode_elapsed_s = decode_ended_s - decode_started_s`.
3. **Derive throughput metrics** – `decode_tok_s = completion_tokens / decode_elapsed_s`, alongside `end_to_end_tok_s` using the full request-level elapsed time.

The system then assembles a statistics dictionary containing **`decode_elapsed_s`**, **`decode_tok_s`**, **`scheduler_policy`**, **`mtp_batch_fixed_width`**, and **`mtp_batch_route_id`**. These values merge into `job.request_observability` (lines 666‑679), ensuring downstream envelope stages receive the sealed-width truth rather than submit-time estimates.

## Working with the Batched MTP API

### Creating a Batch Job

Instantiate `MTPBatchJob` directly when building custom routing logic:

```python
from mtplx.server.mtp_batch import MTPBatchJob
from mtplx.sampling import SamplerConfig

job = MTPBatchJob(
    request_id="mtpbatch-1234",
    prompt_ids=[101, 102, 103],
    max_tokens=256,
    sampler=SamplerConfig(temperature=0.7),
    draft_sampler=SamplerConfig(),
    seed=42,
    stop_token_ids={0},
    token_callback=lambda toks: print("token:", toks),
    compatibility_key=("my-model", False, sampler, draft_sampler),
    generation_limits={},
    solo_runner=None,
    cancel_error=lambda j: RuntimeError("cancelled"),
)

```

### Enqueueing from the OpenAI Server Path

The standard entry point in [`mtplx/server/openai.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/openai.py) constructs jobs automatically:

```python

# Excerpt from mtplx/server/openai.py

job = MTPBatchJob(
    request_id=response_id or f"mtpbatch-{uuid.uuid4().hex}",
    prompt_ids=prompt_ids,
    max_tokens=response_max,
    sampler=sampler,
    draft_sampler=draft_sampler,
    seed=generation_seed,
    stop_token_ids=_default_stop_tokens(state.runtime.tokenizer),
    token_callback=kwargs.get("token_callback"),
    prefill_callback=kwargs.get("prefill_callback"),
    compatibility_key=_mtp_batch_compatibility_key(lane, omit_bonus, sampler, draft_sampler),
    generation_limits=generation_limits,
    solo_runner=lambda _job: _run_generation(state, prompt_ids, **solo_kwargs),
    cancel_error=lambda item: _StreamCancelled(f"request {item.request_id} cancelled"),
    cancel_event=cancel_event,
    request_observability=request_observability,
    omit_speculative_bonus=omit_bonus,
    session_id=kwargs.get("session_id"),
    session_restore=session_restore_hook,
    session_commit=session_commit_hook,
)

```

### Inspecting Decode Statistics

After the batch completes, access detailed timing metrics from the result future:

```python
result = future.result()
stats = result["stats"]
print(f"Decode time: {stats['decode_elapsed_s']:.3f}s")
print(f"Tokens per second: {stats['decode_tok_s']:.1f}")
print(f"Fixed width used: {stats['mtp_batch_fixed_width']}")

```

## Summary

- **Batched MTP decode** groups requests into fixed-width cohorts processed by `generate_a3b_mtp_batch` in a single model pass.
- **`MTPBatchJob`** encapsulates request state and carries the **compatibility key** that determines lane eligibility.
- **`mark_decode_start`** and **`terminal_perf_s`** provide microsecond-accurate decode window measurement despite concurrent execution.
- Statistics including **`decode_tok_s`** and **`mtp_batch_fixed_width`** propagate through `request_observability` for downstream telemetry.

## Frequently Asked Questions

### What determines which fixed-width lane a request uses in MTPLX?

The **`compatibility_key`** tuple—comprising model identifiers, speculative bonus flags, and sampler configurations—determines lane eligibility. The `MTPBatchGenerationService` selects the narrowest installed native graph capable of accommodating the cohort width via `_lane_for_real_width`.

### How does MTPLX measure decode time for individual requests in a batch?

Each `MTPBatchJob` records the start time via **`mark_decode_started`** (using `time.perf_counter()`) when the first token reaches the decode callback. After the driver returns, **`_complete_cohort_job`** calculates `decode_elapsed_s` by subtracting this start time from `terminal_perf_s` (or the completion timestamp if the driver metric is unavailable).

### Can session caching be disabled for specific batched requests?

Yes. The **`session_bank`** hook in [`mtplx/server/openai.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/openai.py) (lines 23104‑23107) allows requests to bypass the session cache. When constructing the `MTPBatchJob`, pass `session_restore` and `session_commit` hooks to override default caching behavior.

### What statistics does the batched MTP lane expose after completion?

The completion payload includes **`decode_elapsed_s`** (wall-clock decode duration), **`decode_tok_s`** (tokens per second during decode), **`mtp_batch_fixed_width`** (the sealed cohort width), **`mtp_batch_route_id`**, and **`scheduler_policy`**. These values update the job's `request_observability` dictionary for downstream consumption.