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

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 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 (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):

  • 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) 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).

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):

  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:

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 constructs jobs automatically:


# 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:

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 (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.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →