# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from __future__ import annotations

import asyncio
import concurrent.futures
import enum
import gc
import logging
import math
import os
import threading
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any, Literal, TypeVar

import cloudpickle
import torch
import torch.distributed as dist
import torchstore as ts
from torch.distributed._state_dict_utils import _create_cpu_state_dict
from torchstore import RankRole
from vllm import EngineArgs, LLMEngine, SamplingParams
from vllm.config import AttentionConfig, CompilationConfig
from vllm.config.compilation import CompilationMode, CUDAGraphMode, PassConfig
from vllm.outputs import RequestOutput
from vllm.sampling_params import RequestOutputKind
from vllm.v1.attention.backends.registry import AttentionBackendEnum

from torchtitan.components.checkpointer import CheckpointManager
from torchtitan.config import Configurable, DebugConfig, OverrideConfig
from torchtitan.distributed import maybe_apply_numa_binding
from torchtitan.distributed.batch_invariant import set_batch_invariance
from torchtitan.distributed.spmd_types import (
    dtensor_to_plain_tensor_state_dict,
    plain_tensor_to_dtensor_state_dict,
)
from torchtitan.models.common.attention import FlexInnerAttention, VarlenInnerAttention
from torchtitan.models.common.decoder import Decoder
from torchtitan.observability import structured_logger as sl
from torchtitan.observability.logging import init_logger
from torchtitan.rl.distributed.parallelism import InferenceParallelismConfig
from torchtitan.rl.distributed.routing.intra_generator import IntraGeneratorRouter
from torchtitan.rl.model.batch_invariance import force_logprobs_fn_for_batch_invariance
from torchtitan.rl.model.vllm_registry import (
    register_to_vllm,
    TORCHTITAN_CONFIG_FORMAT,
    TORCHTITAN_WORKER_CLS,
)
from torchtitan.rl.observability import metrics as m
from torchtitan.rl.observability.vllm import StatLoggerContext, VllmOtelStatLogger
from torchtitan.rl.types import Completion
from torchtitan.tools.utils import has_cuda_capability

logger = logging.getLogger(__name__)

_T = TypeVar("_T")

# TODO(async-rl): this file is large. Split a backend-agnostic BaseGenerator.


@dataclass(kw_only=True, slots=True)
class _RequestMetricsInputs:
    """Raw inputs needed to build a request's vLLM metrics. Used to pass
    metric related information when fan-in from DPs to rank 0.
    """

    num_cached_tokens: int | None
    has_stats: bool
    queued_ts: float = 0.0
    scheduled_ts: float = 0.0
    first_token_ts: float = 0.0
    last_token_ts: float = 0.0
    first_token_latency: float = 0.0
    num_generation_tokens: int = 0


def _extract_request_metrics_inputs(
    request_output: RequestOutput,
) -> _RequestMetricsInputs:
    """Pull the raw metric inputs off a finished ``RequestOutput``."""
    stats = request_output.metrics
    if stats is None:
        return _RequestMetricsInputs(
            num_cached_tokens=request_output.num_cached_tokens, has_stats=False
        )
    return _RequestMetricsInputs(
        num_cached_tokens=request_output.num_cached_tokens,
        has_stats=True,
        queued_ts=stats.queued_ts,
        scheduled_ts=stats.scheduled_ts,
        first_token_ts=stats.first_token_ts,
        last_token_ts=stats.last_token_ts,
        first_token_latency=stats.first_token_latency,
        num_generation_tokens=stats.num_generation_tokens,
    )


def _prepare_generation_request_metrics(
    inputs: _RequestMetricsInputs, *, prefix: str
) -> list[m.Metric]:
    """Prepare vLLM per-request metrics from the raw inputs.

    For `add_request` call, vLLM returns a RequestOutput carrying
    a single `RequestStateStats` (captured into `_RequestMetricsInputs`).

    Caveat under `SamplingParams.n > 1`: vLLM stores one `RequestStateStats`
    per child request; the parent output exposes the **last-finishing**
    child's timeline. `arrival_time` is shared across siblings, but
    [`queued_ts`, `scheduled_ts`, `first_token_ts`, `last_token_ts`,
    `num_generation_tokens`] describe one specific child - not an aggregate,
    not the first sibling's. The other `n-1` siblings' stats are dropped by
    vLLM at ``output_processor._finish_request``.
    """

    # TODO: Per-request fields here come from RequestOutput.metrics
    # (RequestStateStats). Engine-aggregate stats, such as KV-cache usage,
    # prefix-cache hit rate, preemptions, and batch occupancy, live in
    # SchedulerStats / IterationStats and require registering a
    # vllm.v1.metrics.loggers.StatLoggerBase via
    # LLMEngine.from_engine_args(..., stat_loggers=[...]).

    metric_values: dict[str, float] = {}
    if inputs.num_cached_tokens is not None:
        metric_values[f"{prefix}/num_cached_tokens"] = inputs.num_cached_tokens

    if inputs.has_stats:
        metric_values[f"{prefix}/queue_time_ms"] = (
            inputs.scheduled_ts - inputs.queued_ts
        ) * 1000

        if inputs.num_generation_tokens > 0:
            metric_values[f"{prefix}/time_to_first_token_ms"] = (
                inputs.first_token_latency * 1000
            )
            metric_values[f"{prefix}/prefill_time_ms"] = (
                inputs.first_token_ts - inputs.scheduled_ts
            ) * 1000

        if inputs.num_generation_tokens > 1:
            first_to_last_token_ms = (
                inputs.last_token_ts - inputs.first_token_ts
            ) * 1000
            metric_values[f"{prefix}/decode_time_ms"] = first_to_last_token_ms
            metric_values[
                f"{prefix}/inter_token_latency_ms"
            ] = first_to_last_token_ms / (inputs.num_generation_tokens - 1)

    # Emit each value with both Mean and Max aggregators.
    return [
        metric
        for key, value in metric_values.items()
        for metric in (m.Metric(key, m.Mean(value)), m.Metric(key, m.Max(value)))
    ]


# vLLM's default max_num_batched_tokens (vllm's per-step budget:
# prefill + decode tokens summed over the batch). Used as the CUDA graph capture
# cap for "FULL" (which graphs prefill / mixed batches, so capture sizes must
# reach the per-step budget or those batches fall back to eager) when
# ``Config.max_num_batched_tokens`` is unset; when that field is set, its value
# is used instead (and also drives the vLLM engine).
_DEFAULT_MAX_NUM_BATCHED_TOKENS = 2048


@dataclass(kw_only=True, slots=True)
class VLLMCudaGraphConfig:
    """CUDA graph capture settings for the vLLM inference engine.

    Local compile regions come from the model config's ``local_compile_regions``.
    Only CUDA graph capture, which is vLLM-specific, is controlled here.

    ``mode`` selects which vLLM CUDA graph mode to capture; see that field and
    ``get_vllm_compilation_config`` for the per-mode trade-offs. The default,
    ``FULL``, graphs the whole forward (prefill included).
    """

    mode: Literal["NONE", "FULL_DECODE_ONLY", "FULL"] = "FULL"
    """Which vLLM CUDA graph mode to capture:

    - ``"NONE"``: disable compilation and CUDA graph capture.
    - ``"FULL_DECODE_ONLY"``: graph pure-decode batches; prefill / mixed
      batches run eager. Cheap (no inductor compile).
    - ``"FULL"`` (default): graph the whole forward, prefill included, attention
      captured too.
    """

    capture_sizes: list[int] | None = None
    """Explicit CUDA graph capture batch sizes. When ``None`` (default), sizes are
    auto-derived: powers of 2 up to the cap, plus ``max_num_seqs`` and the cap as
    exact sizes. When set, these sizes are deduped and sorted. When expert
    sequence parallelism is enabled, capture sizes that are not multiples of
    its degree are removed."""

    # TODO: Validate CUDA graph capture with MoE / Expert Parallelism.
    # MoE routing produces dynamic shapes that may conflict with full
    # CUDA graph capture despite being torch.compile-compatible
    # post https://github.com/pytorch/torchtitan/pull/3142

    # TODO: Explore applying CUDA graph capture on the torchtitan trainer
    # side as well (not just the vLLM generator).
    # https://github.com/pytorch/torchtitan/issues/3175

    def get_vllm_compilation_config(
        self,
        *,
        max_num_seqs: int,
        expert_sequence_parallel_size: int,
        enable_sequence_parallel: bool,
        max_num_batched_tokens: int | None = None,
    ) -> CompilationConfig:
        """Build a vLLM ``CompilationConfig`` for the generator.

        When ``capture_sizes`` is set, those exact sizes are captured. Otherwise
        sizes are auto-derived: powers of 2 up to the cap, plus ``max_num_seqs`` and
        the cap itself as exact sizes so the largest capture size is always the cap
        (even when it is not a power of 2). The cap is ``max_num_seqs`` for
        ``FULL_DECODE_ONLY`` (decode batch == num_seqs). ``FULL`` also graphs
        prefill, whose per-step token count is bounded by
        ``max_num_batched_tokens`` (the configured value, else
        ``_DEFAULT_MAX_NUM_BATCHED_TOKENS``), so the cap extends to it -- otherwise
        prefill chunks larger than the cap fall back to eager.

        ``expert_sequence_parallel_size`` is the TP-axis shard count used by the
        internally sequence-sharded MoE path. A value greater than one removes
        capture sizes that are not multiples of this count, preventing CUDA graph
        padding from producing a token count that cannot be evenly TP-sharded.

        ``enable_sequence_parallel`` is forwarded to vLLM's sequence parallelism
        pass. vLLM filters dense-SP CUDA graph sizes using its own TP size.

        All modes capture with ``mode=CompilationMode.NONE`` to avoid nesting
        vLLM's Inductor compile with TorchTitan local compile.
        """
        if self.mode == "NONE":
            return CompilationConfig(
                cudagraph_mode=CUDAGraphMode.NONE,
                mode=CompilationMode.NONE,
                pass_config=PassConfig(
                    enable_sp=enable_sequence_parallel,
                    sp_min_token_num=1 if enable_sequence_parallel else None,
                ),
            )
        if max_num_seqs <= 0:
            raise ValueError(f"max_num_seqs must be positive, got {max_num_seqs}")
        if expert_sequence_parallel_size <= 0:
            raise ValueError(
                "expert_sequence_parallel_size must be positive, got "
                f"{expert_sequence_parallel_size}"
            )
        if max_num_batched_tokens is not None:
            _max_cuda_graph_capture_size = max_num_batched_tokens
        else:
            _max_cuda_graph_capture_size = _DEFAULT_MAX_NUM_BATCHED_TOKENS
        cap = max_num_seqs
        if self.mode == "FULL":
            cap = max(cap, _max_cuda_graph_capture_size)
        if self.capture_sizes is not None:
            if not self.capture_sizes or any(s <= 0 for s in self.capture_sizes):
                raise ValueError(
                    "cuda_graph.capture_sizes must be a non-empty list of positive "
                    f"ints, got {self.capture_sizes}"
                )
            sizes = sorted(set(self.capture_sizes))
        else:
            sizes = [1 << i for i in range(int(math.log2(cap)) + 1)]
            # Always include max_num_seqs (decode batch) and the cap (largest
            # prefill chunk) as exact sizes so the largest capture size is the cap
            # even when it is not a power of 2
            if max_num_seqs not in sizes:
                sizes.append(max_num_seqs)
            if cap not in sizes:
                sizes.append(cap)
            sizes = sorted(sizes)

        if expert_sequence_parallel_size > 1:
            removed_sizes = [
                size for size in sizes if size % expert_sequence_parallel_size != 0
            ]
            if removed_sizes:
                logger.warning(
                    "CUDA graph capture sizes %s are removed because they are not "
                    "multiples of expert_sequence_parallel_size %d",
                    removed_sizes,
                    expert_sequence_parallel_size,
                )
            sizes = [
                size for size in sizes if size % expert_sequence_parallel_size == 0
            ]
            if not sizes:
                raise ValueError(
                    "No CUDA graph capture sizes are divisible by "
                    f"expert_sequence_parallel_size {expert_sequence_parallel_size}"
                )

        return CompilationConfig(
            cudagraph_mode=self.mode,
            mode=CompilationMode.NONE,
            cudagraph_capture_sizes=sizes,
            pass_config=PassConfig(
                enable_sp=enable_sequence_parallel,
                sp_min_token_num=1 if enable_sequence_parallel else None,
            ),
        )


@dataclass(kw_only=True, slots=True)
class SamplingConfig:
    """Sampling parameters passed to vLLM's SamplingParams."""

    temperature: float = 0.8
    """Sampling temperature. 0.0 = greedy, higher = more random."""

    top_p: float = 1.0
    """Nucleus sampling threshold. Must be 1.0: the trainer scores tokens over the full vocabulary,
    so sampling from a truncated nucleus would bias the gradient."""

    max_tokens: int = 100
    """Maximum number of tokens to generate per completion."""

    seed: int | None = None
    """Per-request RNG seed. The rollouter offsets this per sample so a group's
    n=1 requests stay diverse while remaining reproducible (None = nondeterministic)."""

    stop_token_ids: list[int] | None = None
    """Renderer role-boundary stop tokens; filled by the controller. Required at
    generation time: these are the only ids that end a request (vLLM's EOS stops
    are off)."""

    def __post_init__(self) -> None:
        # TODO(mask-replay): to allow top_p < 1, turn on vLLM's `return_sampling_mask` (needs a vLLM
        # upgrade, the V2 model runner, top_k > 0 and processed logprobs), carry each token's kept
        # ids next to generator_logprobs, and take the trainer's logsumexp of logits / T over them.
        if self.top_p < 1.0:
            raise ValueError(
                f"top_p must be 1.0, got {self.top_p}: the trainer computes logprobs "
                "over the full vocabulary."
            )


class RequestDispatcher:
    """Handles the generator's DP/TP request dispatch.

    Every rank holds one dispatcher; methods act according to the rank's role:
    - Rank 0 is the coordinator (and DP0's tp_rank=0): it holds the outstanding
      generations, routes requests, opens the fan-in port, and resolves every
      completion -- its own replica's locally, peers' via the drain task. State
      and methods only ever used on rank 0 are prefixed ``rank0_``.
    - Other DP's tp_rank=0 build finished completions and fan them in
      to rank 0 over the port.
    - tp_rank!=0 hold no outputs; their dispatcher only carries the layout.

    Supported rank layout:

        global_rank = dp_rank * tp_degree + tp_rank

    EP reuses the same global-rank -> (dp, tp) mapping, so it needs no special
    handling here.

    Data flow:

    Take DP=2, TP=2 for example. Only rank 0 holds the replies and talks to the
    controller, so completions produced by any other DP replica must be sent
    back ("fanned in") to rank 0:

        controller --generate--> rank 0  (registers the generation)
                                   |
            rank0_route(): pick a DP rank for the queued requests
                                   |   (broadcast in the LoopDecision, elsewhere)
              +--------------------+--------------------+
              v                                         v
        DP0's tp_rank=0 (i.e. rank 0)               DP1's tp_rank=0 (i.e. rank 2)
          engine.step()                             engine.step()
          build (completion, metrics_inputs)        build (completion, metrics_inputs)
          resolve own replies locally   <--port--   send completions to rank 0
                                                    (tp_rank!=0: hold no outputs)
        rank 0 background drain task: recv from port -> resolve those replies
    """

    def __init__(
        self,
        *,
        rank: int,
        dp_rank: int,
        tp_rank: int,
        dp_degree: int,
        broadcast_group: dist.ProcessGroup,
        intra_generator_router: IntraGeneratorRouter.Config,
        open_result_channel: Callable[[], tuple[Any, Any]] | None,
    ):
        self._rank = rank
        self._dp_degree = dp_degree
        self._dp_rank = dp_rank
        self._tp_rank = tp_rank
        # Reused for the one-time result-port broadcast (see ``setup``).
        self._broadcast_group = broadcast_group

        # RANK 0: generations taken off the queue and not yet answered, keyed by
        # request id. An entry is removed once its reply is resolved, with the
        # completion or an error. Only rank 0 ever populates this.
        self._rank0_outstanding_generations: dict[str, OutstandingGeneration] = {}

        # --- Result fan-in (only when DP>1) ---
        # rank 0 opens the channel, keeps the receiving end, and its drain task
        # resolves whatever peer tp_rank=0 send. ``_result_port`` is the sending
        # end, held only by peer tp_rank=0; it stays None on rank 0 (resolves
        # locally) and on tp_rank!=0 (no outputs). All None for DP=1 since there
        # is no peer DP to send from.
        self._open_result_channel = open_result_channel
        self._result_port: Any | None = None
        self._rank0_result_receiver: Any | None = None
        self._rank0_drain_task: asyncio.Task | None = None

        # --- DP routing ---
        # RANK-0 DP routing: pick a DP rank per request, reserving its load until
        # the completion resolves.
        self._rank0_dp_router: IntraGeneratorRouter | None = (
            intra_generator_router.build(dp_degree=self._dp_degree)
            if self._rank == 0 and self._dp_degree > 1
            else None
        )

    def rank0_register_generation(
        self,
        request_id: str,
        metrics_prefix: str,
        reply: concurrent.futures.Future[Completion],
    ) -> None:
        """RANK 0: register a generation taken off the queue, whose ``reply`` is resolved with
        ``request_id``'s completion."""
        if request_id in self._rank0_outstanding_generations:
            raise AssertionError(f"request_id {request_id!r} is already in flight")
        self._rank0_outstanding_generations[request_id] = OutstandingGeneration(
            reply=reply, metrics_prefix=metrics_prefix
        )

    def rank0_has_outstanding_generations(self) -> bool:
        """RANK 0: whether any generation is outstanding (reply unresolved).

        A generation stays registered until its completion comes back, so this stays
        True while any peer DP rank is still running it.
        """
        return bool(self._rank0_outstanding_generations)

    def rank0_route(self, requests: list[EngineRequest]) -> list[list[EngineRequest]]:
        """RANK 0: pick which DP rank serves each queued request.

        Returns a fixed-length (``dp_degree``) list; index == DP rank. Each rank
        later admits only its own slice. With a single DP rank, everything goes
        to DP rank 0; otherwise ``IntraGeneratorRouter`` reserves a DP rank per
        request and that reservation is released when the request's completion
        resolves.
        """
        requests_per_dp_rank: list[list[EngineRequest]] = [
            [] for _ in range(self._dp_degree)
        ]
        for request in requests:
            if self._rank0_dp_router is None:
                dp_rank = 0
            else:
                # Pick a DP rank for this request, and increment this DP rank's
                # load by 1 (i.e. measured by request count).
                dp_rank = self._rank0_dp_router.reserve(
                    request.request_id,
                    routing_session_id=request.routing_session_id,
                )
            requests_per_dp_rank[dp_rank].append(request)

        return requests_per_dp_rank

    def rank0_stamp_min_policy_version(
        self,
        requests_per_dp_rank: list[list[EngineRequest]],
    ) -> None:
        """RANK 0: stamp each request's min policy version on its outstanding generation, for every
        request in this `LoopAction.STEP` decision, across all DP ranks. Without a KV reset the request
        may reuse KV cached under that version, so it is the oldest policy the completion can depend
        on. Rank 0 holds the outstanding generations regardless of which DP rank serves the request,
        so it stamps them all here."""
        for dp_requests in requests_per_dp_rank:
            for request in dp_requests:
                self._rank0_outstanding_generations[
                    request.request_id
                ].min_policy_version = request.min_policy_version

    def setup(self) -> None:
        """One-time setup before the engine loop starts (DP>1): distribute rank 0's
        result-fan-in port and start rank 0's drain task.

        All ranks call this so the broadcast over the all-ranks ``_broadcast_group``
        completes; only TP rank 0 keep the port.
        """
        if self._dp_degree == 1:
            return

        # rank 0 opens the channel and keeps the receiving end; it broadcasts only
        # the sending port to the peers.
        if self._rank == 0:
            if self._open_result_channel is None:
                raise RuntimeError(
                    "data-parallel generation requires a result-channel factory"
                )
            port, self._rank0_result_receiver = self._open_result_channel()
            # Monarch Port objects need cloudpickle, so we cloudpickle it into bytes
            # first. Otherwise, broadcast_object_list will attempt to pickle the
            # port object with stdlib pickle and result in error.
            container = [cloudpickle.dumps(port)]
        else:
            container = [None]
        dist.broadcast_object_list(
            container, src=0, group=self._broadcast_group, device=torch.device("cpu")
        )
        assert container[0] is not None
        # Only peer TP rank 0s send completions to global rank 0, so only they
        # keep the port. Global rank 0 resolves locally and tp_rank!=0 produce
        # no outputs.
        if self._rank != 0 and self._tp_rank == 0:
            self._result_port = cloudpickle.loads(container[0])
        if self._rank == 0:
            self._rank0_drain_task = asyncio.create_task(self._rank0_drain_results())

    def process_finished_requests(
        self, request_outputs: list[RequestOutput], policy_version: int
    ) -> None:
        """TP rank 0s send finished completions to global rank 0 after each
        ``engine.step()``:
          - Global Rank 0 resolves its own DP replica's completions locally;
          - every other TP rank 0 sends them over the port to global rank 0's drain
            task.
          - Other ranks hold no finished outputs and do nothing.
        """
        if self._tp_rank != 0:
            return

        completions = self._build_completions(request_outputs, policy_version)
        if self._rank == 0:
            self._rank0_resolve_generations(completions)
        elif completions:
            self._result_port.send(completions)

    def _build_completions(
        self, request_outputs: list[RequestOutput], policy_version: int
    ) -> list[tuple[str, Completion, _RequestMetricsInputs]]:
        """Turn finished ``RequestOutput``s into ``(request_id, Completion, metrics_inputs)``."""
        completions: list[tuple[str, Completion, _RequestMetricsInputs]] = []
        for request_output in request_outputs:
            # We enforce n=1 in sampling params -> exactly one CompletionOutput per finished request
            # Here we just sanity check it (a single engine.step may still finish several requests).
            if len(request_output.outputs) != 1:
                raise ValueError(
                    f"expected n=1 (one sample per request), got "
                    f"{len(request_output.outputs)} for {request_output.request_id}"
                )

            # flat_logprobs=True: vLLM returns logprobs as plain lists instead of one dict per token.
            # logprobs=0 keeps only the sampled token, so `.logprobs` has exactly one float per generated token.
            completion_output = request_output.outputs[0]
            flat_logprobs = completion_output.logprobs
            token_logprobs = list(flat_logprobs.logprobs)

            completions.append(
                (
                    request_output.request_id,
                    Completion(
                        # NOTE: min_policy_version is a PLACEHOLDER here, set equal to max (the finish
                        # version). The serving rank has no access to the outstanding generation that
                        # holds the true admitted version, so rank 0 REPLACES this with the real value in
                        # _rank0_resolve_generations. min == max here ONLY until that replacement.
                        min_policy_version=policy_version,
                        max_policy_version=policy_version,
                        request_id=request_output.request_id,
                        token_ids=list(completion_output.token_ids),
                        token_logprobs=token_logprobs,
                        finish_reason=completion_output.finish_reason,
                    ),
                    _extract_request_metrics_inputs(request_output),
                )
            )
        return completions

    def _rank0_resolve_generations(
        self, completions: list[tuple[str, Completion, _RequestMetricsInputs]]
    ) -> None:
        """RANK 0: build each completion's metrics (the only place that knows the
        request's ``metrics_prefix``), then resolve its reply.

        TODO: metrics are built in two phases -- a DP-leader produces the raw
        ``_RequestMetricsInputs`` alongside the ``Completion``, and rank 0
        finalizes ``completion.metrics`` in place here, where it has the
        ``inflight_requests_at_completion`` count. Consider unifying into a
        single build step once that count can travel with (or be derived
        without) the rank-0 outstanding-generation bookkeeping.
        """
        for request_id, completion, metrics_inputs in completions:
            # in flight when this one finished (includes itself; counted before the pop)
            inflight_requests_at_completion = float(
                len(self._rank0_outstanding_generations)
            )
            generation = self._rank0_outstanding_generations.pop(request_id)

            # Replace the placeholder min (the builder set min == max) with the true admitted
            # version stamped on the outstanding generation at admission.
            completion.min_policy_version = generation.min_policy_version
            metrics_prefix = generation.metrics_prefix

            metrics = _prepare_generation_request_metrics(
                metrics_inputs, prefix=metrics_prefix
            )
            for metric_type in [m.Max, m.Mean]:
                metrics.append(
                    m.Metric(
                        f"{metrics_prefix}/inflight_requests_at_completion",
                        metric_type(inflight_requests_at_completion),
                    )
                )
            completion.metrics = metrics

            generation.reply.set_result(completion)
            # Free the request's reserved load on its DP rank so load-aware
            # routing sees the accurate loads on DPs.
            if self._rank0_dp_router is not None:
                self._rank0_dp_router.release(request_id)

    async def _rank0_drain_results(self) -> None:
        """RANK 0 background task which receives and resolves completions pushed
        by peer TP rank 0s.
        """
        while True:
            completions = await self._rank0_result_receiver.recv()
            self._rank0_resolve_generations(completions)

    def fail_outstanding_generations(self, exc: BaseException) -> None:
        """RANK 0: fail the reply of every outstanding generation after an exception or
        teardown (no-op elsewhere, where the map is empty)."""
        for generation in self._rank0_outstanding_generations.values():
            if not generation.reply.done():
                generation.reply.set_exception(exc)
        self._rank0_outstanding_generations.clear()

    async def shutdown(self) -> None:
        """Stop rank 0's drain task, if any (no-op elsewhere)."""
        if self._rank0_drain_task is not None:
            self._rank0_drain_task.cancel()
            try:
                await self._rank0_drain_task
            except (asyncio.CancelledError, Exception):
                pass
            self._rank0_drain_task = None


class VLLMGenerator(Configurable):
    """vLLM engine to drive concurrent `generate` calls through one SPMD engine loop.

    The controller fires independent calls (`generate`, `pull_model_state_dict`, `close`).
    With CPU weight prefetch enabled, the router also calls
    `prefetch_model_state_dict` before `pull_model_state_dict`.
    Rank 0 puts each call on a thread-safe queue and awaits its reply. One background `_engine_loop` per rank
    executes the `LoopDecision` rank 0 makes from the queue. Rank 0 resolves each reply when its request finishes and
    return the result back to the controller.

    Notice that vLLM `engine.step`, which is a TP collective, and the request-intake are decoupled, so a new request
    can join mid-flight, instead of waiting for the current batch to drain.

    One loop iteration, RequestDispatcher is used to dispatch requests to different ranks, and collect results.
    Take DP=1, TP=2 for example, after the controller fired generate(prompt_0) and generate(prompt_1):

        # request intake in `generate` takes a prompt, puts in a queue, releases control back to the controller
        generate(prompt_0): enqueue prompt_0, await reply_0   ┐ rank 0 owns the queue + replies
        generate(prompt_1): enqueue prompt_1, await reply_1   ┘ (other ranks are no-op)

        # meanwhile, the engine-loop, which is its own coroutine, is continuously running.
        rank 0   _decide_next_action -> LoopDecision(STEP, [prompt_0, prompt_1])─┐  broadcast_object_list (gloo)
        rank 1   (blocked inside the broadcast) ─────────────────────────────────┘  decision from rank0 is broadcast

        # The loop now executes the decision. In this example: STEP.
        ALL      add_request(prompt_0), add_request(prompt_1)
        ALL      engine.step() * max_engine_steps_between_decisions # N step burst before a new decision

        # resolve the reply, waking up `generate` so it returns the result to the controller.
        # Note that prompt_1 can be done before prompt_0. The result is per request, not per batch.
        rank 0   request_dispatcher.process_finished_requests -> prompt_1 done? reply_1.set_result(Completion)
        rank 1   request_dispatcher.process_finished_requests -> no-op (tp_rank != 0, holds no replies)

    For DP>1, the requests will be routed among DPs first. See RequestDispatcher's docstring for more details.

    Threading: `engine.step()` blocks, so the engine and everything the engine loop touches (queue, replies,
    dispatcher) live on a dedicated engine thread, which runs its own event loop. The endpoints run on the
    actor's event loop and reach the engine loop only through the thread-safe queue and the
    `concurrent.futures.Future`s it resolves, so the actor's loop stays free to take calls while the engine steps.
    Two endpoints are exceptions. `release_groups` pops cache-salt pins directly: each pop, like the engine loop's
    `setdefault`, is a single dict operation, atomic under the GIL. `prefetch_model_state_dict` writes the staging
    buffers the engine loop reads on a pull: the router awaits it before `pull_model_state_dict`, and the
    controller awaits that pull before the next weight sync, so the two never overlap.

    A weight sync rides the same loop: `pull_model_state_dict` puts a `ModelStateDictPullMessage` on the queue, which
    rank 0 turns into a `LoopDecision(LoopAction.PULL_MODEL_STATE_DICT)` applied between step bursts. The engine does
    NOT drain in-flight requests first ("hotswap"). This behavior can be changed in the inter-generator router, by
    blocking new requests until the engine is drained.

    With CPU weight prefetch enabled, the network transfer into pinned CPU memory
    happens before this loop action, which then performs the local CPU-to-GPU copy.

    Args:
        config: Generator-specific configuration.
        model_config: TorchTitan model configuration.
        model_path: Path to the HF model checkpoint.
        max_num_seqs: vLLM's upper bound on concurrently scheduled sequences (vLLM admits fewer if KV
            is tight); also sets the CUDA-graph capture sizes.
        output_dir: Structured-logger output directory.
        open_result_channel: Opens the asynchronous channel used to fan in
            completions from nonzero data-parallel replicas to rank 0. It must
            return the sending port and receiving endpoint. The Monarch actor
            adapter supplies ``Channel.open``; standalone runtimes can supply
            an equivalent transport. It is unused when data parallelism is
            disabled.
    """

    @dataclass(kw_only=True, slots=True)
    class Config(Configurable.Config):
        """Generator actor configuration.
        TODO: Expose a EngineConfig field to passing config to vLLM Engine"""

        parallelism: InferenceParallelismConfig = field(
            default_factory=InferenceParallelismConfig
        )
        """Parallelism configuration for the vLLM engine."""

        intra_generator_router: IntraGeneratorRouter.Config = field(
            default_factory=IntraGeneratorRouter.Config
        )
        """In-mesh DP routing config: how rank 0 partitions requests across the
        engine's data-parallel ranks (no effect when data_parallel_degree == 1,
        where there is a single DP rank)."""

        sampling: SamplingConfig = field(default_factory=SamplingConfig)
        """Default sampling parameters for generation."""

        override: OverrideConfig = field(default_factory=OverrideConfig)
        """Config overrides (e.g. ``torchtitan_recipes.overrides.fused_swiglu.fused_swiglu``)
        applied to this generator's model spec before model finalization and build.
        Separate from the trainer's override so the two can differ."""

        model_dtype: str = "bfloat16"
        """Data type for model weights, passed directly to vLLM (auto, float16, bfloat16, float32)."""

        gpu_memory_limit: float = 0.9
        """Fraction of GPU memory to use for the vLLM engine (0.0 to 1.0)."""

        max_num_batched_tokens: int | None = None
        """vLLM chunked-prefill chunk size: max tokens scheduled per engine step
        (prefill + decode, summed over the batch). ``None`` (default) leaves
        vLLM's own engine default in place."""

        cuda_graph: VLLMCudaGraphConfig = field(default_factory=VLLMCudaGraphConfig)
        """CUDA graph capture settings for the vLLM engine."""

        checkpointer: CheckpointManager.Config | None = None
        """Optional initial-weight loader for the vLLM wrapper.

        In the RL loop this stays ``None`` because weights arrive from
        TorchStore. Standalone inference supplies a config that loads the
        initial Hugging Face weights.
        """

        debug: DebugConfig = field(default_factory=DebugConfig)
        """Debug and determinism settings."""

        max_engine_steps_between_decisions: int = 16
        """Controls how many `engine.step()` calls the `engine_loop` performs before processing a new decision.
        Every generation call is queued for execution by the `engine_loop`. A higher value enables buffering
        of more requests to avoid a prefill between every engine decode step, which is inefficient."""

        # TODO: check if we should put these under WeightSyncConfig
        enable_cpu_weight_prefetch: bool = True
        """Prefetch model weights into pinned CPU memory before applying them on GPU.

        When ``enable_cpu_weight_prefetch=False``:

        Use vLLM's CuMem pool for model weights transferred directly over RDMA.

        TorchTitan enables PyTorch's expandable-segments allocator to reduce
        fragmentation. It can change the physical GPU memory behind an address,
        invalidating NIXL's RDMA registration for that memory.

        vLLM's CuMem pool disables expandable segments for its allocations,
        keeping their memory mappings stable. This option puts model weights in
        that pool.

        It is not needed when ``enable_cpu_weight_prefetch=True``
        because RDMA targets the persistent CPU buffers instead.
        """

        reset_kv_cache_on_weight_sync: bool = False
        """Reset cached and running-request KV after each weight sync.

        The default preserves in-flight requests and their KV: a rollout group keeps the
        cache salt pinned when this generator first admitted it, so its rollouts and
        later turns reuse its KV across weight syncs, while new groups use the current
        version.
        Enable this to clear prefix-cache entries and preempt running requests; vLLM
        then recomputes their KV under the new weights when they resume."""

        vllm_stat_logger: VllmOtelStatLogger.Config | None = None
        """Optional logger instantiated on TP rank 0 to export vLLM metrics."""

        def __post_init__(self):
            # The generator runs vLLM full expert parallelism: vLLM forms the EP
            # group from all DP*TP ranks, so expert_parallel_degree must equal
            # data_parallel_degree * tensor_parallel_degree (or 1 to disable EP).
            p = self.parallelism
            full_ep = p.data_parallel_degree * p.tensor_parallel_degree
            if p.data_parallel_degree > 1 and p.expert_parallel_degree == 1:
                raise ValueError(
                    "generator data_parallel_degree may be greater than 1 only "
                    "to supply ranks for expert parallelism. For independent "
                    "generator replicas, set data_parallel_degree=1 and increase "
                    "Controller.Config.num_generators instead."
                )
            if p.expert_parallel_degree not in (1, full_ep):
                raise ValueError(
                    f"expert_parallel_degree ({p.expert_parallel_degree}) must be 1 "
                    f"(no expert parallelism) or equal data_parallel_degree * "
                    f"tensor_parallel_degree ({full_ep}) in the generator."
                )

            if self.debug.batch_invariant and not self.reset_kv_cache_on_weight_sync:
                raise ValueError(
                    "batch_invariant requires reset_kv_cache_on_weight_sync=True so "
                    "cached KV cannot cross a policy update"
                )

    def __init__(
        self,
        config: Config,
        *,
        model_config: Decoder.Config,
        model_path: str,
        max_num_seqs: int,
        output_dir: str,
        rank: int | None = None,
        generator_name: str = "generator",
        open_result_channel: Callable[[], tuple[Any, Any]] | None = None,
    ):
        init_logger()
        # TODO: Quiet torchstore's per-op transport-resolve INFO spam (very noisy in CI).
        logging.getLogger("torchstore.transport").setLevel(logging.WARNING)
        sl.init_structured_logger(
            source="rl_generator",
            output_dir=output_dir,
            rank=rank if rank is not None else int(os.environ.get("RANK", "0")),
            enable=config.debug.enable_structured_logging,
        )
        sl.log_trace_instant("structured_logger_started")

        self.config = config
        self.model_config = model_config

        self._max_num_seqs = max_num_seqs

        self._rank = rank if rank is not None else int(os.environ.get("RANK", "0"))
        self._dp_degree = config.parallelism.data_parallel_degree
        tp_degree = config.parallelism.tensor_parallel_degree
        # TODO: revisit if PP/CP are added.
        self._dp_rank = self._rank // tp_degree
        self._tp_rank = self._rank % tp_degree

        # Register TorchTitan model + parser with vLLM
        register_to_vllm(
            model_config,
            parallelism=config.parallelism,
            checkpointer_config=config.checkpointer,
            override=config.override,
        )

        # Set vLLM environment variables from config before any vLLM initialization
        attention_backend = model_config.first_base_attention_backend
        assert isinstance(
            attention_backend,
            (VarlenInnerAttention.Config, FlexInnerAttention.Config),
        ), "Only varlen and flex attention backends are allowed."

        os.environ["VLLM_USE_V2_MODEL_RUNNER"] = "0"
        set_batch_invariance(config.debug.batch_invariant)
        if config.debug.batch_invariant:
            # The vLLM v2 logprob Triton kernel bypasses the aten overrides above;
            # route it through trainer's function to match the trainer exactly.
            force_logprobs_fn_for_batch_invariance()

        self._set_determinism(config.debug)

        self.model_path = model_path

        # Build vLLM engine
        enable_ep = config.parallelism.expert_parallel_degree > 1
        engine_kwargs = dict(
            # ``model`` is the path to the HF checkpoint directory. The
            # config is sourced from TorchTitan's model config via
            # ``config_format=TORCHTITAN_CONFIG_FORMAT`` (no config.json
            # read), but vLLM still uses this path to locate the
            # tokenizer assets and the safetensors weight shards.
            model=model_path,
            trust_remote_code=True,
            # Use the torchtitan custom config parser (registered by
            # register_to_vllm above). It builds PretrainedConfig from
            # model config instead of reading config.json from disk.
            config_format=TORCHTITAN_CONFIG_FORMAT,
            dtype=config.model_dtype,
            tensor_parallel_size=config.parallelism.tensor_parallel_degree,
            data_parallel_size=config.parallelism.data_parallel_degree,
            # NOTE: Monarch launches the generator workers and sets the torch
            # elastic distributed env; with external_launcher, vLLM uses that
            # world to build its process groups. vLLM does not take an
            # explicit EP degree: when this boolean is set, it converts all
            # DP * TP ranks into the expert-parallel group for MoE layers.
            enable_expert_parallel=enable_ep,
            worker_cls=TORCHTITAN_WORKER_CLS,
            # Monarch already spawned TP workers via proc mesh. "external_launcher"
            # tells vLLM to run one worker per process (no subprocess spawning)
            distributed_executor_backend="external_launcher",
            gpu_memory_utilization=config.gpu_memory_limit,
            enforce_eager=config.cuda_graph.mode == "NONE",
            attention_config=AttentionConfig(
                backend=(
                    AttentionBackendEnum.FLEX_ATTENTION
                    if isinstance(attention_backend, FlexInnerAttention.Config)
                    else AttentionBackendEnum.CUSTOM
                ),
            ),
            # Enables RequestOutput.metrics, so generator metrics can be returned
            disable_log_stats=False,
            enable_cumem_allocator=not config.enable_cpu_weight_prefetch,
            # Token-in-token-out: prompts and outputs are token ids, so vLLM
            # needs no tokenizer. This also drops the tokenizer's eos_token_id
            # as a stop; the generation config's eos ids are dropped by
            # ignore_eos in _build_sampling_params.
            skip_tokenizer_init=True,
        )
        engine_kwargs["max_model_len"] = model_config.max_context_length
        # Return logprobs of the distribution vLLM samples from (after temperature). vLLM's default
        # returns the raw model's logprobs, before temperature.
        engine_kwargs["logprobs_mode"] = "processed_logprobs"
        engine_kwargs["max_num_seqs"] = self._max_num_seqs
        if config.max_num_batched_tokens is not None:
            engine_kwargs["max_num_batched_tokens"] = config.max_num_batched_tokens
        # Continuous batching requires FCFS scheduling: admission order must equal the
        # broadcast order on every rank
        engine_kwargs["scheduling_policy"] = "fcfs"
        # Which sliding-window / Mamba (GDN) state blocks are hashed into the prefix
        # cache; full-attention groups hash every block regardless.
        # - 0 (vLLM default since v0.29): only the replay boundary (the last block
        #   boundary before the prompt's final token) and shared-prefix junctions.
        #   Other prompt blocks and every block crossed during decode are freed
        #   unhashed, so a multi-turn rollout's next turn, whose prompt contains the
        #   previous completion, cannot reuse anything past the previous prompt.
        # - None: every full block boundary, as before v0.29. Blocks are still freed
        #   when the request moves past them; hashed ones stay reusable until evicted.
        engine_kwargs["prefix_cache_retention_interval"] = None
        # FA2 requires block_size to be a multiple of 256
        if not has_cuda_capability(9, 0):
            engine_kwargs["block_size"] = 256
        expert_sequence_parallel_size = config.parallelism.expert_sequence_parallel_size
        vllm_compilation_config = config.cuda_graph.get_vllm_compilation_config(
            max_num_seqs=self._max_num_seqs,
            max_num_batched_tokens=config.max_num_batched_tokens,
            expert_sequence_parallel_size=expert_sequence_parallel_size,
            enable_sequence_parallel=config.parallelism.enable_sequence_parallel,
        )
        if vllm_compilation_config is not None:
            engine_kwargs["compilation_config"] = vllm_compilation_config
        if config.debug.seed is not None:
            engine_kwargs["seed"] = config.debug.seed
        engine_args = EngineArgs(**engine_kwargs)

        with sl.log_trace_span("vllm_init"):
            logger.info("Initializing LLMEngine from EngineArgs...")
            stat_loggers = None
            if self._tp_rank == 0:
                if config.vllm_stat_logger is None:
                    logger.info(
                        "VllmOtelStatLogger inactive because "
                        "vllm_stat_logger=None. To record vLLM metrics, set it "
                        "to VllmOtelStatLogger.Config() and set "
                        "OTEL_METRICS_EXPORTER=jsonl or otlp"
                    )
                else:
                    vllm_stat_logger_config = config.vllm_stat_logger
                    logger_context = StatLoggerContext(
                        rank=self._rank,
                        tp_rank=self._tp_rank,
                        dp_rank=self._dp_rank,
                        generator_name=generator_name,
                        output_dir=output_dir,
                    )

                    def build_stat_logger(vllm_config, engine_index):
                        return vllm_stat_logger_config.build(
                            vllm_config=vllm_config,
                            engine_index=engine_index,
                            context=logger_context,
                        )

                    stat_loggers = [build_stat_logger]

            # Start the thread that runs vllm engine.
            self._engine_event_loop = asyncio.new_event_loop()
            threading.Thread(
                target=self._engine_event_loop.run_forever,
                name="vllm-engine",
                daemon=True,
            ).start()
            self._engine: LLMEngine | None = self._call_on_engine_thread(
                lambda: LLMEngine.from_engine_args(
                    engine_args, stat_loggers=stat_loggers
                )
            )
            if config.enable_cpu_weight_prefetch:
                # The prefetch buffers below and each prefetch's TorchStore read pin memory through
                # the CUDA runtime, on the calling thread's device. Both run on this thread, which
                # would otherwise use device 0, rank 0's GPU.
                torch.cuda.set_device(
                    self._call_on_engine_thread(torch.cuda.current_device)
                )
            logger.info("vLLM rollout engine initialized")

        # The default PG was initialized during engine build. Confirm the configured
        # rank matches the torch-distributed global rank so the two views cannot
        # silently diverge.
        torch_distributed_rank = dist.get_rank()
        if self._rank != torch_distributed_rank:
            raise RuntimeError(
                f"rank mismatch: configured rank ({self._rank}) != "
                f"torch dist.get_rank() ({torch_distributed_rank})"
            )
        # Confirm the DP layout we computed above matches what vLLM derived
        # independently during engine build, so the two views can't silently diverge.
        vllm_parallel_config = self._engine.vllm_config.parallel_config
        if vllm_parallel_config.data_parallel_size != self._dp_degree:
            raise RuntimeError(
                f"DP layout mismatch on rank {self._rank}: our dp_size "
                f"({self._dp_degree}) != vLLM data_parallel_size "
                f"({vllm_parallel_config.data_parallel_size})"
            )
        if vllm_parallel_config.data_parallel_rank != self._dp_rank:
            raise RuntimeError(
                f"DP layout mismatch on rank {self._rank}: our dp_rank "
                f"({self._dp_rank}) != vLLM data_parallel_rank "
                f"({vllm_parallel_config.data_parallel_rank})"
            )

        self.policy_version = 0
        # RANK 0: group id -> min policy version the group is pinned to, set at the
        # group's first admission and used as its prefix cache salt. Unused with
        # reset_kv_cache_on_weight_sync. All rollouts of a group share the pin, so a
        # rollout first admitted after a pull still reuses its group's prompt KV, at the
        # cost of depending on the group's older version. Only the controller knows when
        # a group makes no more generation calls, so entries live until it calls
        # `release_groups`.
        self._group_min_policy_versions: dict[int, int] = {}
        self._prefetched_model_state_dict = (
            self._setup_prefetch_staging_state_dict()
            if config.enable_cpu_weight_prefetch
            else None
        )

        # --- Continuous-batching state (see the class docstring) ---
        self._broadcast_group = dist.new_group(backend="gloo")  # for LoopDecisions

        # --- Request dispatch ---
        # The dispatcher owns the DP/TP rank layout and the request dispatch /
        # completion fan-in (see its docstring).
        self._request_dispatcher = RequestDispatcher(
            rank=self._rank,
            dp_rank=self._dp_rank,
            tp_rank=self._tp_rank,
            dp_degree=self._dp_degree,
            broadcast_group=self._broadcast_group,
            intra_generator_router=config.intra_generator_router,
            open_result_channel=open_result_channel,
        )

        # Engine-loop queue (rank 0): messages the endpoints put; the loop takes them off to decide.
        self._engine_loop_queue = EngineLoopQueue(self._engine_event_loop)

        # `_engine_loop` running on the engine thread's event loop, as a future any thread can await;
        # None until start_engine_loop starts it.
        self._engine_loop_future: concurrent.futures.Future[None] | None = None

        logger.info("Generator initialized with vLLM engine")

    def _call_on_engine_thread(self, fn: Callable[[], _T]) -> _T:
        """Call `fn` on the engine thread, in the caller's contextvars, and block until it returns.
        For setup in `__init__`, before the engine loop starts."""
        result: concurrent.futures.Future[_T] = concurrent.futures.Future()

        def call() -> None:
            # A plain callback rather than a task, which would re-raise a `SystemExit` out of
            # `run_forever` before resolving its future, leaving the caller waiting forever.
            try:
                result.set_result(fn())
            except BaseException as exc:
                result.set_exception(exc)

        # `call_soon_threadsafe` runs `call` in a copy of the caller's contextvars.
        self._engine_event_loop.call_soon_threadsafe(call)
        return result.result()

    def _setup_prefetch_staging_state_dict(self) -> dict[str, Any]:
        """Allocate persistent pinned CPU buffers local to this rank's GPU."""
        # Bind before allocation so first-touch places the pinned buffers on
        # the NUMA node local to this rank's GPU.
        maybe_apply_numa_binding(torch.cuda.current_device(), "cuda")
        model = self._get_model()
        model_sd = plain_tensor_to_dtensor_state_dict(
            model.model.state_dict(),
            state_dict_layouts=model.get_state_dict_layouts(),
            parallelism_context=model.parallelism_context,
        )
        # Preserve the DTensor layouts while replacing their local storage
        # with persistent pinned CPU buffers.
        return _create_cpu_state_dict(model_sd, pin_memory=True)

    @staticmethod
    def _set_determinism(debug: DebugConfig) -> None:
        """Apply deterministic flags for the generator.

        The generator doesn't use torchtitan's ParallelismContext, so we apply
        the deterministic flags directly instead of using set_determinism().
        """
        if debug.deterministic:
            torch.use_deterministic_algorithms(
                True, warn_only=debug.deterministic_warn_only
            )
            torch.backends.cudnn.deterministic = True
            torch.backends.cudnn.benchmark = False
            os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"

        if debug.seed is not None:
            torch.manual_seed(debug.seed)

    def _get_model(self):
        """Access the model from the vLLM engine.
        Returns a VLLMModelWrapper instance.
        """
        return self._engine.model_executor.driver_worker.get_model()

    async def sync_log_step(self, step: int, relative_step: int | None = None) -> None:
        """Sync the structured-logger step counter from the controller."""
        sl.set_step(step, relative_step=relative_step)

    async def start_engine_loop(self) -> None:
        """Start the background engine loop on every rank (one-time, idempotent)."""
        if self._engine_loop_future is None:
            self._engine_loop_future = asyncio.run_coroutine_threadsafe(
                self._engine_loop(), self._engine_event_loop
            )

    def _rank0_check_engine_loop_running(self, endpoint_name: str) -> None:
        """Guard for the rank-0-only endpoints"""
        assert self._rank == 0, f"{endpoint_name} must be routed to rank 0 only"
        if self._engine_loop_queue.closed:
            raise RuntimeError(f"generator is closed; cannot call {endpoint_name}")
        if self._engine_loop_future is None:
            raise RuntimeError(
                "engine loop not started; call start_engine_loop on all ranks "
                f"before {endpoint_name}"
            )

    @sl.log_trace_span("generate")
    async def generate(
        self,
        prompt_token_ids: list[int],
        *,
        request_id: str,
        group_id: int,
        routing_session_id: str,
        sampling_config: SamplingConfig | None = None,
        metrics_prefix: str = "generator",
    ) -> Completion:
        """Generates one completion for one prompt.

        Can be accepted by rank 0 only (rank 0 owns the queue + replies and
        drives the followers through the engine loop). Returns the `Completion`,
        which carries its own per-generation metrics (`Completion.metrics`) that
        the controller attaches to the rollout turn.

        Args:
            prompt_token_ids: One tokenized prompt `[token_ids]`.
            request_id: Unique id for this request, echoed on the `Completion`.
            group_id: Rollout group id. Requests of one group share a prefix cache
                salt; `release_groups` drops it.
            routing_session_id: Stable session key for in-mesh DP routing.
            sampling_config: Optional per-call override for the generator's
                default SamplingConfig.
            metrics_prefix: Namespace prepended to every metric key on the returned
                `Completion` (default ``"generator"``). Callers that need to keep streams
                separate, e.g. ``"validation/generator"``, can override it.

        Example:

            completion = await generator.slice(hosts=0, gpus=0).generate.call_one(
                [1, 2, 3],
                request_id="step=3/group=0/sample=0/turn=0",
                group_id=0,
                routing_session_id="group=0/rollout=0",
            )
        """
        self._rank0_check_engine_loop_running("generate")

        sampling = (
            sampling_config if sampling_config is not None else self.config.sampling
        )
        assert (
            sampling.stop_token_ids is not None
        ), f"{request_id}: stop_token_ids must be set from the renderer"

        # Put the call on the queue; the engine loop will admit + process it, then resolve `reply`.
        reply: concurrent.futures.Future[Completion] = concurrent.futures.Future()
        self._engine_loop_queue.put(
            GenerationMessage(
                engine_request=EngineRequest(
                    request_id=request_id,
                    prompt_token_ids=prompt_token_ids,
                    sampling=sampling,
                    group_id=group_id,
                    routing_session_id=routing_session_id,
                ),
                metrics_prefix=metrics_prefix,
                reply=reply,
            )
        )
        return await asyncio.wrap_future(reply)

    @sl.log_trace_span("engine_loop")
    async def _engine_loop(self) -> None:
        """Non-stop loop running on all ranks to produce new tokens.

        Rank 0 decides a `LoopDecision` and broadcasts it; ALL ranks apply it in
        lockstep until `LoopAction.CLOSE`. On crash, fail every outstanding reply so callers don't hang. On exit,
        release the dispatcher and the vLLM engine.

        `_decide_next_action` is consulted once every `max_engine_steps_between_decisions` steps (a burst),
        so new requests buffer and prefill together instead of on every step.

        Example:
            max_engine_steps_between_decisions = 16
            check `_decide_next_action` --> "STEP"         --> run engine.step 16 times
            check `_decide_next_action` --> "PULL_MODEL_STATE_DICT" --> run `_pull_model_state_dict`
            check `_decide_next_action` --> "STEP"         --> run engine.step 16 times
            check `_decide_next_action` --> "CLOSE"        --> stop
        """
        # Engine-loop state (rank 0): engine requests taken off the queue but not yet put in a `LoopAction.STEP`.
        # We need to cache such requests in this field, so they can be carried across `_decide_next_action` calls.
        # Note these requests' parent `GenerationMessage`s have been registered as `OutstandingGeneration` in dispatcher.
        pending_engine_requests: list[EngineRequest] = []
        # Engine-loop state (rank 0): weight pulls taken off the queue but not yet applied.
        pending_pull_messages: list[ModelStateDictPullMessage] = []
        try:
            # One-time dispatcher setup before the loop starts.
            self._request_dispatcher.setup()
            while True:
                # Rank 0 decides next decision; followers pass None and learn from the broadcast.
                decision = (
                    await self._decide_next_action(
                        pending_engine_requests, pending_pull_messages
                    )
                    if self._rank == 0
                    else None
                )

                # Barrier(gloo, CPU): Ship rank 0's decision (incl. prompts) to every TP rank via gloo/CPU, off the
                # NCCL stream. broadcast_object_list mutates a list in place; `to_thread` runs the
                # blocking call in a worker so the event loop can still serve generate/pull/close.
                # TODO(perf): overlap this broadcast with the step burst (pipeline the next decision) so
                # gloo transfer hides behind GPU compute.
                # TODO(perf): swap broadcast_object_list (double serialization of pickle+broadcast) for a byte
                # broadcast, to cut pickle overhead. i.e. serialize the decision to bytes on rank 0, then broadcast.
                # TODO: Revisit when enabling DP>1
                decision_broadcast_container = [decision] if self._rank == 0 else [None]
                await asyncio.to_thread(
                    dist.broadcast_object_list,
                    decision_broadcast_container,
                    src=0,
                    group=self._broadcast_group,
                    device=torch.device("cpu"),
                )
                # get rank 0's broadcasted decision
                decision = decision_broadcast_container[0]  # [num_ranks]

                if decision.action is LoopAction.CLOSE:
                    if self._rank == 0:
                        _fail_pulls(
                            pending_pull_messages,
                            RuntimeError(
                                "generator closed before the pull was applied"
                            ),
                        )
                    return

                if decision.action is LoopAction.PULL_MODEL_STATE_DICT:
                    await self._pull_model_state_dict(decision.pull_version)
                    # One pull applied every pull this decision coalesced (only rank 0 holds any).
                    for pull in pending_pull_messages:
                        pull.reply.set_result(None)
                    pending_pull_messages.clear()
                    continue  # back to the start for the next decision

                if decision.action is LoopAction.STEP:
                    # Rank 0 holds every outstanding generation, so it stamps the admitted (min) version for the
                    # whole decision.
                    # TODO: move under the engine_step call (register at generation_start, not admission).
                    # The way to do it is probably to change to RequestOutputKind.CUMULATIVE and mark per token.
                    if self._rank == 0:
                        self._request_dispatcher.rank0_stamp_min_policy_version(
                            decision.requests_per_dp_rank
                        )
                    # Admit only this rank's DP replica slice. TP ranks in the same
                    # replica compute the same _dp_rank, so they add the identical
                    # set in the same FCFS order.
                    local_requests = decision.requests_per_dp_rank[
                        self._request_dispatcher._dp_rank
                    ]
                    if local_requests:
                        # render_cmpl is vLLM's input pipeline (tokenize is a no-op for tokenized prompts);
                        # the high-level entry stays resilient to vLLM internals vs vllm.inputs.tokens_input.
                        prompts = []
                        for request in local_requests:
                            prompt = {"prompt_token_ids": request.prompt_token_ids}
                            if not self.config.reset_kv_cache_on_weight_sync:
                                # Salt by the pinned version so a request only reuses KV
                                # computed under that version.
                                prompt["cache_salt"] = str(request.min_policy_version)
                            prompts.append(prompt)
                        engine_inputs = self._engine.renderer.render_cmpl(prompts)
                        for request, engine_input in zip(
                            local_requests, engine_inputs, strict=True
                        ):
                            self._engine.add_request(
                                request_id=request.request_id,
                                prompt=engine_input,
                                params=self._build_sampling_params(request.sampling),
                            )

                # Barrier (NCCL): engine.step() runs SPMD in lockstep.
                # The step burst `max_engine_steps_between_decisions` gives the generator time to buffer
                # new requests and avoid a prefill in between every engine decode step, which is inefficient.
                with sl.log_trace_span("vllm_engine_step_burst"):
                    for _ in range(self.config.max_engine_steps_between_decisions):
                        if not self._engine.has_unfinished_requests():
                            break
                        with torch.no_grad():
                            request_outputs = self._engine.step()
                        self._request_dispatcher.process_finished_requests(
                            request_outputs, self.policy_version
                        )
                        await asyncio.sleep(0)  # let pending generate() calls enqueue

        except Exception as exc:
            logger.exception("engine loop crashed; failing all outstanding replies")
            if self._rank == 0:
                self._request_dispatcher.fail_outstanding_generations(exc)
                _fail_pulls(pending_pull_messages, exc)
                await self._rank0_close_and_fail_queue(exc)
            raise
        finally:
            await self._release_loop_resources()

    async def _rank0_close_and_fail_queue(self, exc: Exception) -> None:
        """RANK 0: after the engine loop crashed, close the queue so later calls raise, and fail the
        replies of the calls still on it."""
        self._engine_loop_queue.close(CloseMessage(), reason="engine loop crashed")
        # Puts that beat the close have scheduled their enqueue ahead of this task's next step.
        await asyncio.sleep(0)
        while not self._engine_loop_queue.empty():
            message = self._engine_loop_queue.get_nowait()
            if isinstance(message, CloseMessage):
                continue
            if message.reply.set_running_or_notify_cancel():
                message.reply.set_exception(exc)

    async def _decide_next_action(
        self,
        pending_engine_requests: list[EngineRequest],
        pending_pull_messages: list[ModelStateDictPullMessage],
    ) -> LoopDecision:
        """RANK 0: takes everything off the queue and picks the next action. Sleeps until there is
        something to do.
        """
        if pending_pull_messages:
            raise AssertionError(
                "the engine loop applies or fails every pull before deciding again"
            )

        # * If there are outstanding generations, we drain the queue and make the next decision right away,
        #   so outstanding generations can keep stepping.
        # * If there are none, there is nothing to step, so we wait on the queue for the next decision.
        messages: list[EngineLoopMessage]
        if self._request_dispatcher.rank0_has_outstanding_generations():
            messages = []
        else:
            if pending_engine_requests:
                raise AssertionError(
                    f"{len(pending_engine_requests)} pending engine requests but no outstanding "
                    "generation"
                )
            messages = [await self._engine_loop_queue.get()]

        # Take everything currently in the queue if there is any, so they all can be part of this decision.
        while not self._engine_loop_queue.empty():
            messages.append(self._engine_loop_queue.get_nowait())

        for message in messages:
            if isinstance(message, CloseMessage):
                # Drops nothing: the queue rejects puts once closed, so the `CloseMessage` is the last message.
                return LoopDecision(action=LoopAction.CLOSE, requests_per_dp_rank=[])
            # Skip a call its caller already cancelled; once running, the reply ignores cancellation, so only
            # the engine loop resolves it.
            if not message.reply.set_running_or_notify_cancel():
                continue
            if isinstance(message, ModelStateDictPullMessage):
                pending_pull_messages.append(message)
            else:
                self._request_dispatcher.rank0_register_generation(
                    message.engine_request.request_id,
                    message.metrics_prefix,
                    message.reply,
                )
                pending_engine_requests.append(message.engine_request)

        # A weight pull takes priority over admitting new requests. Pulls taken off the queue
        # together are coalesced into one, at the highest version: every pull reads the latest push from
        # one TorchStore key, so the weights read are at least as new as any version requested.
        if pending_pull_messages:
            return LoopDecision(
                action=LoopAction.PULL_MODEL_STATE_DICT,
                requests_per_dp_rank=[],
                pull_version=max(pull.version for pull in pending_pull_messages),
            )

        # `LoopAction.STEP`: admit whatever is pending (may be empty -> just keep stepping in-flight work).
        for request in pending_engine_requests:
            if self.config.reset_kv_cache_on_weight_sync:
                # Each pull resets all KV, so requests need no pin or salt.
                request.min_policy_version = self.policy_version
            else:
                request.min_policy_version = self._group_min_policy_versions.setdefault(
                    request.group_id, self.policy_version
                )
        requests_per_dp_rank = self._request_dispatcher.rank0_route(
            pending_engine_requests
        )
        pending_engine_requests.clear()
        return LoopDecision(
            action=LoopAction.STEP,
            requests_per_dp_rank=requests_per_dp_rank,
        )

    def _build_sampling_params(self, sampling: SamplingConfig) -> SamplingParams:
        """Translate a `SamplingConfig` into vLLM `SamplingParams` (n=1).

        ``seed`` and ``stop_token_ids`` are carried on the ``SamplingConfig``
        (the controller fills ``stop_token_ids`` and the rollouter offsets
        ``seed`` per sample), so each sample in a group is a distinct ``n=1``
        request that stays diverse and bitwise-reproducible.

        The engine loads no tokenizer, so its ``eos_token_id`` is not a stop,
        but vLLM still adds the generation config's ``eos_token_id`` (checkpoint
        ``generation_config.json`` or the ``vllm_registry`` HF config).
        ``ignore_eos`` turns that off, so the renderer's ``stop_token_ids``
        (which include EOS) are the only stops.
        """
        return SamplingParams(
            temperature=sampling.temperature,
            top_p=sampling.top_p,
            max_tokens=sampling.max_tokens,
            n=1,  # always expects a single sample per request. Caller can call N times.
            stop_token_ids=sampling.stop_token_ids,
            # Drops the generation config's eos ids, which vLLM merges into
            # stop_token_ids even with skip_tokenizer_init.
            ignore_eos=True,
            seed=sampling.seed,
            logprobs=0,  # return only the sampled token's logprob (for the GRPO ratio)
            # Token ids in, token ids and logprob floats out: stops are token ids and nothing reads
            # text, so skip vLLM's per-token detokenization and per-token logprob dicts.
            detokenize=False,
            flat_logprobs=True,
            # Return each request's result once, when it is fully done, instead of streaming partial
            # outputs as tokens arrive.
            # TODO(async-rl): use RequestOutputKind.CUMULATIVE for exact per-token
            #   (start_token, version) boundaries; today we keep only the per-turn min/max.
            output_kind=RequestOutputKind.FINAL_ONLY,
        )

    async def release_groups(self, group_ids: list[int]) -> None:
        """Drop the pinned cache salts of finished rollout groups.

        Args:
            group_ids: Groups with no more generation calls.
        """
        # A pop can land partway through `_decide_next_action`'s stamping loop, so requests of a
        # released group in one batch can get different versions. If every batch must see one
        # consistent snapshot of the pins, send releases through the queue instead.
        for group_id in group_ids:
            self._group_min_policy_versions.pop(group_id, None)

    async def initialize_torchstore_client(self, requester_index: int) -> None:
        """Initialize this process as a TorchStore routing requester.

        Args:
            requester_index: Index used to namespace this generator mesh.
        """
        await ts.client(role=RankRole.REQUESTER, group=requester_index)

    @sl.log_trace_span("pull_model_state_dict")
    async def pull_model_state_dict(self, version: int) -> None:
        """Queues a weight pull for `version` and blocks until the engine loop has finished pulling.
        Pulls queued together are applied once, at the highest version.

        With CPU weight prefetch enabled, the network transfer has already
        completed and this pull applies the prefetched weights to the GPU.

        NOTE: In-flight requests are NOT drained here — the endpoint never drains; a caller that wants
        an idle engine holds off new `generate` calls until the queue drains, then calls this.

        Args:
            version: Policy version to pull
        """
        self._rank0_check_engine_loop_running("pull_model_state_dict")

        reply: concurrent.futures.Future[None] = concurrent.futures.Future()
        self._engine_loop_queue.put(
            ModelStateDictPullMessage(version=version, reply=reply)
        )
        await asyncio.wrap_future(reply)

    @sl.log_trace_span("prefetch_model_state_dict")
    async def prefetch_model_state_dict(self) -> None:
        """Fetch weights into pinned CPU memory without interrupting generation."""
        assert self.config.enable_cpu_weight_prefetch
        assert self._prefetched_model_state_dict is not None

        await ts.get_state_dict(
            "model_state_dict",
            user_state_dict=self._prefetched_model_state_dict,
            strict=False,
            direct_rdma=False,
        )

    @sl.log_trace_span("pull_model_state_dict_copy")
    async def _pull_model_state_dict(self, version: int) -> None:
        """ALL RANKS: collectively copy the latest weights from TorchStore, optionally drop the
        prefix cache when configured, and bump the policy version.

        With CPU weight prefetch enabled, copy the already-fetched weights from
        pinned CPU memory instead of fetching them from TorchStore here.
        """
        # Async RL uses a StorageVolume snapshot so generators do not read
        # live trainer GPU tensors while optimizer steps may be mutating them.
        model = self._get_model()
        model_sd = model.model.state_dict()
        await self._get_spmd_state_dict(model_sd, model=model)
        # With CPU prefetch, model_sd instead contains the prefetched CPU tensors,
        # and this load performs the local CPU-to-GPU copy.
        model.model.load_state_dict(model_sd, strict=True)
        self.policy_version = version
        if self.config.reset_kv_cache_on_weight_sync:
            # Always reset running requests too: the only reason to reset is a strict
            # recompute under the new weights. Keeping running requests' KV while hiding
            # old KV from new requests is already what the default (no reset) does via the salt.
            self._engine.reset_prefix_cache(
                reset_running_requests=True,
            )
        gc.collect()

    async def _get_spmd_state_dict(self, model_sd: dict, *, model) -> None:
        """Fetch trainer-pushed weights into a spmd_types generator state dict.

        spmd_types generators hold plain local tensors, but TorchStore already
        knows how to fill DTensor state-dict entries. Wrap each local tensor as
        a DTensor using its declared SPMD layout, fetch through the normal
        state-dict path, then put the local tensors back before load_state_dict.

        With CPU weight prefetch enabled, use the previously fetched DTensor
        state dict instead.
        """
        if self.config.enable_cpu_weight_prefetch:
            assert self._prefetched_model_state_dict is not None
            dtensor_model_sd = self._prefetched_model_state_dict
        else:
            dtensor_model_sd = plain_tensor_to_dtensor_state_dict(
                model_sd,
                state_dict_layouts=model.get_state_dict_layouts(),
                parallelism_context=model.parallelism_context,
            )

            await ts.get_state_dict(
                "model_state_dict",
                user_state_dict=dtensor_model_sd,
                strict=False,
                direct_rdma=False,
            )

        model_sd.update(dtensor_to_plain_tensor_state_dict(dtensor_model_sd))

    async def close(self) -> None:
        """Stop the engine loop, which releases the dispatcher and the vLLM engine on exit.

        Rank 0 closes the queue: calls already queued stay ahead of the `CloseMessage` that makes
        the engine loop quit the while-loop, and later calls raise.
        """
        if self._rank == 0:
            self._engine_loop_queue.close(CloseMessage(), reason="generator is closed")

        if self._engine_loop_future is None:
            # The loop never started (or an earlier `close` already awaited it), so release here, on
            # the engine thread, which makes every engine call.
            await asyncio.wrap_future(
                asyncio.run_coroutine_threadsafe(
                    self._release_loop_resources(), self._engine_event_loop
                )
            )
            return
        try:
            await asyncio.wrap_future(self._engine_loop_future)
        except Exception:
            logger.exception("engine loop raised during shutdown")
        self._engine_loop_future = None

    async def _release_loop_resources(self) -> None:
        """Stop the dispatcher, fail the replies it left unresolved, and drop the vLLM engine. No-op
        once released.

        Engine teardown: with `external_launcher`, vLLM reuses the process group and actor
        lifetime that Monarch owns. Calling vLLM's internal `engine_core.shutdown()` can block
        while Monarch is also trying to stop the same actor mesh, so this only closes
        renderer-local resources and leaves process teardown to `ProcMesh.stop()`.
        """
        # Stop the result-drain task on rank 0.
        await self._request_dispatcher.shutdown()

        # The loop has stopped; fail any replies it left unresolved so awaiting callers get an
        # exception instead of hanging.
        self._request_dispatcher.fail_outstanding_generations(
            RuntimeError("generator closed before the request finished")
        )

        # Shut down engine parts
        if self._engine is not None:
            renderer = getattr(self._engine, "renderer", None)
            if renderer is not None:
                logger.info("Shutting down vLLM renderer")
                renderer.shutdown()
            self._engine = None


# ===================== helpers =====================


# ---- Engine-loop queue: rank 0's queued calls; a LoopDecision broadcasts only their EngineRequests. ----


@dataclass(kw_only=True, slots=True)
class EngineRequest:
    """One `generate` call's request for the engine, awaiting admission; what a `LoopDecision` broadcasts."""

    request_id: str
    prompt_token_ids: list[int]  # [prompt_tokens]
    sampling: SamplingConfig
    group_id: int
    routing_session_id: str
    min_policy_version: int = field(init=False)
    """Oldest policy version this request's KV can come from; rank 0 sets it at admission.
    Without a KV reset on weight sync it is the group's pinned version and salts the
    prefix cache."""


@dataclass(kw_only=True, slots=True)
class GenerationMessage:
    """A queued `generate` call: its `engine_request`, and what only rank 0 needs to answer the caller.
    Never broadcast; a `LoopDecision` carries just the `engine_request`."""

    engine_request: EngineRequest
    metrics_prefix: str
    reply: concurrent.futures.Future[Completion]


@dataclass(kw_only=True, slots=True)
class ModelStateDictPullMessage:
    """A queued weight pull: the policy `version` to copy from TorchStore, and the `reply` the engine
    loop resolves once the pull has been applied."""

    version: int
    reply: concurrent.futures.Future[None]


def _fail_pulls(
    pending_pull_messages: list[ModelStateDictPullMessage], exc: BaseException
) -> None:
    """Fail the replies of `pending_pull_messages`, which the engine loop took off the queue but never
    applied."""
    for pull in pending_pull_messages:
        if not pull.reply.done():
            pull.reply.set_exception(exc)


@dataclass(kw_only=True, slots=True)
class CloseMessage:
    """The shutdown signal (no payload) `VLLMGenerator.close` closes the queue with; the engine loop
    returns `LoopAction.CLOSE` when it sees one."""


EngineLoopMessage = GenerationMessage | ModelStateDictPullMessage | CloseMessage


class EngineLoopQueue:
    """Rank 0's queue of `EngineLoopMessage`s for the engine loop.

    `put`, `close` and `closed` may run on any thread. The other methods must run on `event_loop`,
    which owns the queue.

    Args:
        event_loop: The engine thread's event loop.
    """

    def __init__(self, event_loop: asyncio.AbstractEventLoop) -> None:
        self._event_loop = event_loop
        self._queue: asyncio.Queue[EngineLoopMessage] = asyncio.Queue()
        self._closed_reason: str | None = None
        # Orders each put against `close`, so that no message can land behind the last one.
        self._lock = threading.Lock()

    @property
    def closed(self) -> bool:
        """Whether `close` has been called."""
        return self._closed_reason is not None

    def put(self, message: EngineLoopMessage) -> None:
        """Enqueue `message`; it lands on the next iteration of `event_loop`. Raises once closed."""
        with self._lock:
            if self._closed_reason is not None:
                raise RuntimeError(self._closed_reason)
            # asyncio.Queue is not thread-safe, so hand the put to the loop that owns it.
            self._event_loop.call_soon_threadsafe(self._queue.put_nowait, message)

    def close(self, last_message: EngineLoopMessage, reason: str) -> None:
        """Queue `last_message` behind every message already put; later puts raise `RuntimeError(reason)`.
        No-op once closed."""
        with self._lock:
            if self._closed_reason is not None:
                return
            self._closed_reason = reason
            self._event_loop.call_soon_threadsafe(self._queue.put_nowait, last_message)

    async def get(self) -> EngineLoopMessage:
        """Wait for the next message."""
        return await self._queue.get()

    def get_nowait(self) -> EngineLoopMessage:
        """Take the next message; raises `asyncio.QueueEmpty` if there is none."""
        return self._queue.get_nowait()

    def empty(self) -> bool:
        """Whether no message is waiting."""
        return self._queue.empty()


# ---- Rank 0's record of each outstanding generation. ----


@dataclass(kw_only=True, slots=True)
class OutstandingGeneration:
    """A generation rank 0 took off the queue and has not answered yet: the `reply` it resolves with
    the `Completion`, and what it needs to finish building that `Completion`."""

    reply: concurrent.futures.Future[Completion]
    metrics_prefix: str
    """Namespaces this generation's metrics (e.g. `generator` vs `validation_generator`)."""
    min_policy_version: int = field(init=False)
    """Policy version the request was admitted (sampled) under; the max is read at finish (see `Completion`)."""


class LoopAction(enum.Enum):
    """What the engine loop does each loop iteration (rank 0 decides; the choice is broadcast to all)."""

    STEP = "step"
    # run a burst of engine.step(); first admit any newly-queued requests.

    PULL_MODEL_STATE_DICT = "pull_model_state_dict"
    # pull the latest weights between bursts

    CLOSE = "close"
    # stop the loop


@dataclass(kw_only=True, slots=True)
class LoopDecision:
    """The per-iteration decision rank 0 broadcasts so every rank acts identically.
    This is broadcasted (pickled) in the engine_loop"""

    action: LoopAction

    requests_per_dp_rank: list[list[EngineRequest]] | None = None
    # Per-DP-rank requests to admit before a STEP burst; index == DP rank, fixed
    # length data_parallel_degree. Each rank admits only its own DP-rank slice.
    # (empty unless any queued)

    pull_version: int | None = None
    # set iff action is PULL_MODEL_STATE_DICT
