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

"""
TLDR:
_data_input_loop -> _rollout_loop -> _batcher_loop -> training_batch_queue -> _trainer_loop
           |               ^                    ^
           v               |                    |
           +-------RolloutGroupWorkBuffer-------+

Detailed diagram:

_data_input_loop                                      _rollout_loop[N] (group workers)
+--------------------------------------------------+  +--------------------------------------------------+
| group_buffer.wait_for_slot()                     |  | work = group_buffer.claim_next()                  |
| sample = rollouter.get_training_sample()         |  | group = rollouter.run_group_rollouts(work.sample) |
| work = RolloutGroupWork(group_id, sample)        |  | group_buffer.finalize_work(group)                 |
| group_buffer.add_work(work)                      |  +-----------------------+--------------------------+
+-----------------------+--------------------------+                          ^ |
                        |                                                     | |
                        | adds work entry                                     | | updates same entry
                        v                                                     | v
RolloutGroupWorkBuffer
+---------------------------------------------------------------------------------------------------------------------+
| active slots = (target_offpolicy_steps + 1) * num_prompts_per_train_step                                                |
|                                                                                                                     |
| caller            group_buffer call                                            state / active slot                  |
| _data_input_loop  add_work(RolloutGroupWork)                                   WAITING; slot acquired               |
| _rollout_loop[N]  claim_next()                                                 WAITING -> INFLIGHT                  |
| _rollout_loop[N]  finalize_work(RolloutGroup)                                  INFLIGHT -> FINALIZED                |
| _batcher_loop     RolloutGroup = take_finalized()                              FINALIZED -> taken (slot still held) |
| _batcher_loop     release_active_groups(1, "untrainable_group")                slot released                        |
| _trainer_loop     release_active_groups(num_prompts_per_train_step, "trained")  slots released after weight pull     |
+---------------------------------------------------------------------------------------------------------------------+
                                                  |
                                                  | group = group_buffer.take_finalized()
                                                  v
_batcher_loop
+----------------------------------------------------------------------------------------+
| training_sample_group = training_sample_builder.build_from_group(rollout_group=group)  |
| if no trainable samples: group_buffer.release_active_groups(1, "untrainable_group")    |
| batch, trainable = batcher.add_training_samples(training_sample_group)                  |
| training_batch_queue.put(TrainerStepBatch)                                             |
+-----------------------+------------------------------+---------------------------------+
                        |  ^                           |
          add group     |  | maybe_training_batch      | put batch
                        v  |                           v
Batcher                                              training_batch_queue
+-------------------------------------------------+   +----------------------------------------------+
| accumulated TrainingSampleGroups                |   | size 1; holds TrainerStepBatch | None        |
| pack at num_prompts_per_train_step               |   +---------------------+------------------------+
+-------------------------------------------------+                         |
                                                                         | packed = training_batch_queue.get()
                                                                         v
_trainer_loop
+----------------------------------------------------------------------------------------------------------+
| train batch -> optim -> push/pull weights -> buffer.release_active_groups(num_prompts_per_train_step)     |
+----------------------------------------------------------------------------------------------------------+

Backpressure (each loop: what it consumes/produces, and what gates each side):
_data_input_loop
  produces: RolloutGroupWork into group_buffer
    waits for:    a free active slot (group_buffer.wait_for_slot)
    unblocked by: _trainer_loop release_active_groups(num_prompts_per_train_step, "trained") after the pull
                  (and _batcher_loop release_active_groups(1,"untrainable_group"))
_rollout_loop[N]
  consumes: a WAITING RolloutGroupWork (group_buffer.claim_next)
    waits for:    a claimable WAITING entry
    unblocked by: _data_input_loop group_buffer.add_work()
  produces: RolloutGroup (group_buffer.finalize_work)
    waits for:    nothing (admits its own claimed slot)
    unblocked by: n/a

_batcher_loop
  consumes: the oldest FINALIZED group inside the window (group_buffer.take_finalized)
    waits for:    a group inside the window becoming FINALIZED (any group when windowed_fifo_batches is None)
    unblocked by: _rollout_loop[N] group_buffer.finalize_work()
  produces: TrainerStepBatch (training_batch_queue.put)
    waits for:    a free training_batch_queue slot (maxsize=1)
    unblocked by: _trainer_loop training_batch_queue.get()
_trainer_loop
  consumes: a TrainerStepBatch (training_batch_queue.get)
    waits for:    a TrainerStepBatch in the queue
    unblocked by: _batcher_loop training_batch_queue.put()
"""

import asyncio
import json
import logging
import math
import os
import time
import warnings
from dataclasses import dataclass, field, replace

# PYTORCH_CUDA_ALLOC_CONF is set in torchtitan/rl/__init__.py (before torch is imported)
# and in train.py; see the note there.
import torch  # noqa: F401
import torchstore as ts

from monarch.actor import ProcMesh, this_host

from torchtitan.components.renderer import RendererConfig
from torchtitan.components.tokenizer import HuggingFaceTokenizer
from torchtitan.config import Configurable
from torchtitan.models.common.decoder import Decoder
from torchtitan.observability import structured_logger as sl
from torchtitan.rl.components.batcher import Batcher
from torchtitan.rl.components.training_sample_builder import TrainingSampleBuilder
from torchtitan.rl.components.work_buffer import (
    RolloutGroupWork,
    RolloutGroupWorkBuffer,
)
from torchtitan.rl.distributed.actors.generator import VLLMGeneratorActor
from torchtitan.rl.distributed.actors.trainer import TrainerActor
from torchtitan.rl.distributed.routing.inter_generator import InterGeneratorRouter
from torchtitan.rl.distributed.torch_elastic import setup_torch_elastic_env
from torchtitan.rl.distributed.weight_sync import WeightSyncManager
from torchtitan.rl.generator import SamplingConfig, VLLMGenerator
from torchtitan.rl.observability import metrics as m
from torchtitan.rl.observability.controller import (
    compute_perf_ratio_metrics,
    compute_policy_age_metrics,
    compute_rollout_metrics,
    MetricsTimer,
)
from torchtitan.rl.observability.rollout_recorder import RolloutSampleRecorder
from torchtitan.rl.rollout import RolloutGroup
from torchtitan.rl.rollout.rollouter import Rollouter
from torchtitan.rl.rollout.types import GenerateFn
from torchtitan.rl.trainer import Trainer
from torchtitan.rl.types import Completion, TrainerStepBatch

logger = logging.getLogger(__name__)


@dataclass(kw_only=True, slots=True)
class ValidationConfig:
    """Held-out validation that runs at the start and end of training"""

    # TODO: enable periodic validation with proper overlapping

    num_samples: int = 20
    """Held-out prompts scored greedily (temp=0, n=1) per validation pass. 0 skips validation."""


@dataclass(kw_only=True, slots=True)
class AsyncLoopConfig(Configurable.Config):
    num_training_steps: int = 10
    """Optimizer steps to run."""

    num_prompts_per_train_step: int = 8
    """Global number of prompt groups, across all DPs, whose surviving rollouts compose
    one train step (the global_batch_size, in groups)."""

    num_samples_per_prompt: int = 8
    """Sibling rollouts sampled per prompt (the GRPO group)."""

    target_offpolicy_steps: int = 3
    """Target steady-state offpolicy steps used to set the active buffer size to
    `(S + 1) * P`. Observed offpolicy steps are not guaranteed to equal this
    target: when rollout generation is the bottleneck, the buffer may not fill
    and observed offpolicy steps will be lower. A finite `windowed_fifo_batches` bounds
    how far a slow group may exceed this target; None leaves it unbounded. See
    ``torchtitan/rl/docs/windowed_fifo.md`` for details."""

    windowed_fifo_batches: int | None = None
    """FIFO look-ahead window in train batches.

    None (the default) is greedy: the batcher takes the oldest finished group
    anywhere in the buffer. An unfinished older group does not block a younger
    finished group. Maximum policy age is unbounded.

    Set to `n >= 1` to limit consumption to `n * P` group ids from the oldest
    group in the buffer. Maximum policy age is bounded by
    `target_offpolicy_steps + n`. A value of 1 is FIFO by batch. See
    ``torchtitan/rl/docs/windowed_fifo.md``."""

    group_buffer: RolloutGroupWorkBuffer.Config = field(
        default_factory=RolloutGroupWorkBuffer.Config
    )
    training_sample_builder: TrainingSampleBuilder.Config = field(
        default_factory=TrainingSampleBuilder.Config
    )
    batcher: Batcher.Config = field(default_factory=Batcher.Config)
    validation: ValidationConfig = field(default_factory=ValidationConfig)

    def __post_init__(self) -> None:
        if self.num_prompts_per_train_step < 1:
            raise ValueError(
                "num_prompts_per_train_step must be >= 1, got "
                f"{self.num_prompts_per_train_step}"
            )
        if self.target_offpolicy_steps < 0:
            raise ValueError(
                f"target_offpolicy_steps must be >= 0, got {self.target_offpolicy_steps}"
            )
        if self.windowed_fifo_batches is not None and self.windowed_fifo_batches < 1:
            raise ValueError(
                f"windowed_fifo_batches must be None or >= 1, got {self.windowed_fifo_batches}"
            )

    @property
    def max_active_rollout_groups(self) -> int:
        return (self.target_offpolicy_steps + 1) * self.num_prompts_per_train_step

    @property
    def window_size(self) -> int | None:
        """FIFO look-ahead window in group ids, `windowed_fifo_batches * P`; None means no window."""
        if self.windowed_fifo_batches is None:
            return None
        return self.windowed_fifo_batches * self.num_prompts_per_train_step

    @property
    def max_offpolicy_steps(self) -> int | None:
        """Return the worst case consume-time offpolicy bound, or None without a window.

        For active buffer size `B`, window size `W`, and prompts per train step
        `P`, the bound is `(B + W - 2) // P`, which is `S + windowed_fifo_batches`.
        """
        if self.window_size is None:
            return None
        return (
            self.max_active_rollout_groups + self.window_size - 2
        ) // self.num_prompts_per_train_step


class Controller(Configurable):
    """Top-level RL async training orchestrator.

    Owns a `Trainer` actor (gradient updates), a `VLLMGenerator` actor
    (sampling), and a `Rollouter` (datasets + rubric + env construction).

    Check the docstring at the top of the file for more details.

    Example:

        config = recipes.rl_grpo_qwen3_0_6b_varlen()
        controller = config.build()
        trainer_mesh = ...        # provisioned by the caller (see train.py)
        generator_meshes = ...
        await controller.setup_async(
            trainer_mesh=trainer_mesh, generator_meshes=generator_meshes
        )
        await controller.run()
    """

    @dataclass(kw_only=True, slots=True)
    class Config(Configurable.Config):
        """Top-level config for RL training."""

        model: Decoder.Config | None = None
        """Model config shared by the trainer and generator."""

        hf_assets_path: str = "./tests/assets/tokenizer"
        """Path to HF assets folder (model weights, tokenizer, config files)."""

        dump_folder: str = "outputs/rl"
        """Root output folder for RL artifacts (temp weights, logs, etc.)."""

        async_loop: AsyncLoopConfig = field(default_factory=AsyncLoopConfig)
        """How the data->rollout->batch->train loop is sized and coordinated."""

        rollouter: Rollouter.Config
        """The rollouter: its datasets, envs, and rubric."""
        # TODO: support multiple rollouters for data mixing.

        tokenizer: HuggingFaceTokenizer.Config = field(
            default_factory=HuggingFaceTokenizer.Config
        )
        """Tokenizer loaded from `hf_assets_path`."""

        renderer: RendererConfig
        """The model's chat template; renders messages to token ids and parses completions
        back. E.g. `from_renderers(Qwen3RendererConfig(enable_thinking=False))`."""

        rollout_recorder: RolloutSampleRecorder.Config = field(
            default_factory=RolloutSampleRecorder.Config
        )
        """JSONL recorder to save sampled rollouts to disk for further inspection and debugging."""

        trainer: Trainer.Config
        """Trainer config. Controls optimizer, training, parallelism."""

        # TODO: put generator, num generators and generator router in a separate config
        generator: VLLMGenerator.Config = field(default_factory=VLLMGenerator.Config)
        """VLLMGenerator actor configuration (vLLM engine, sampling)."""

        num_generators: int = 1
        """Number of generator replicas to spawn as separate proc meshes.

        This is distinct from intra-generator parallelism controlled by
        ``generator.parallelism``. Total generator GPU/process usage is
        ``num_generators * generator_world_size``.
        """

        generator_router: InterGeneratorRouter.Config = field(
            default_factory=InterGeneratorRouter.Config
        )
        """Generator routing strategy configuration."""

        # TODO: rename it to metrics_processor
        metrics: m.MetricsProcessor.Config = field(
            default_factory=m.MetricsProcessor.Config
        )

        def maybe_log(self) -> None:
            debug = self.trainer.debug
            config_dict = self.to_dict()
            if debug.print_config:
                logger.info(
                    f"Running with configs: {json.dumps(config_dict, indent=2, ensure_ascii=False)}"
                )

            if debug.save_config_file is not None:
                config_file = os.path.join(self.dump_folder, debug.save_config_file)
                os.makedirs(os.path.dirname(config_file), exist_ok=True)
                with open(config_file, "w") as file:
                    json.dump(config_dict, file, indent=2)
                logger.info(f"Saved job configs to {config_file}")

        def __post_init__(self):
            if self.num_generators < 1:
                raise ValueError(
                    f"num_generators must be at least 1, got {self.num_generators}"
                )
            if self.generator.checkpointer is not None:
                raise ValueError(
                    "Generator checkpoint must be disabled in the RL loop "
                    "(weights are synced from the trainer via TorchStore). "
                    "Set generator.checkpointer=None."
                )
            if self.trainer.training.num_tokens_per_train_step != -1:
                warnings.warn(
                    "trainer.training.num_tokens_per_train_step is ignored by "
                    "the RL loop; configure async_loop.num_prompts_per_train_step "
                    "to control optimizer-step boundaries.",
                    stacklevel=2,
                )
            if self.trainer.parallelism.enable_sequence_parallel:
                sp_degree = self.trainer.parallelism.tensor_parallel_degree
                max_context_length = self.trainer.training.max_context_length
                if sp_degree > 1 and max_context_length % sp_degree != 0:
                    raise ValueError(
                        "training.max_context_length "
                        f"({max_context_length}) must be divisible "
                        f"by sequence parallel degree ({sp_degree})."
                    )

            # TODO: add a check so that all seq_len related variables make sense
            # e.g. rollout max length cannot be larger than the model max_seq_len
            # or the packing len, etc.

            if self.trainer.debug.batch_invariant:
                if torch.version.hip is not None:
                    raise ValueError(
                        "batch_invariant mode is not supported on ROCm: the varlen "
                        "attention path cannot force num_splits=1 (rejected by ROCm), "
                        "so split-k reductions are non-deterministic."
                    )
                if not self.trainer.debug.deterministic:
                    raise ValueError("batch_invariant requires deterministic=True")
                # The trainer forward must compute in bf16 to match the bf16
                # generator, via FSDP mixed precision. The trainer always wraps
                # the model in FSDP (even at data_parallel_shard_degree=1, where
                # FSDP acts purely as a mixed-precision boundary), so
                # mixed_precision_param == "bfloat16" casts the fp32 master
                # weights to bf16 for the forward before any matmul.
                if self.trainer.training.mixed_precision_param != "bfloat16":
                    raise ValueError(
                        "batch_invariant requires the trainer forward to compute "
                        "in bfloat16 to match the generator. Set "
                        "training.mixed_precision_param='bfloat16' (fp32 master "
                        "weights, bf16-cast forward via FSDP mixed precision). "
                        "Got mixed_precision_param="
                        f"{self.trainer.training.mixed_precision_param!r}."
                    )
                if self.generator.model_dtype != "bfloat16":
                    raise ValueError(
                        f"batch_invariant requires bfloat16 generator dtype, "
                        f"got {self.generator.model_dtype!r}"
                    )
                if self.trainer.parallelism.enable_sequence_parallel:
                    raise ValueError(
                        "batch_invariant mode doesn't support SP now. "
                        "SP uses reduce-scatter which only supports Ring in NCCL "
                        "and has not been validated for determinism."
                    )

    def __init__(self, config: Config):
        self.config = config
        config.maybe_log()
        self.trainer: Trainer | None = None
        self.generator_router: InterGeneratorRouter | None = None
        # Resume step (0 = fresh); set in setup_async from the loaded checkpoint.
        self.start_step = 0
        self._proc_meshes = []
        self.metrics_processor: m.MetricsProcessor = config.metrics.build(
            log_dir=config.dump_folder,
            job_config=config.to_dict(),
        )
        self.tokenizer = config.tokenizer.build(tokenizer_path=config.hf_assets_path)
        self.renderer = config.renderer.build(tokenizer=self.tokenizer)

        # Carry the base seed and renderer stop tokens on the sampling config so
        # the generator reads them off each request; the rollouter offsets the
        # seed per sample. Avoids the generator depending on request_id format.
        self._sampling = replace(
            config.generator.sampling,
            seed=config.generator.debug.seed,
            stop_token_ids=list(self.renderer.get_stop_token_ids()),
        )
        self._rollouter: Rollouter = config.rollouter.build()
        self.rollout_recorder = config.rollout_recorder.build(
            dump_dir=config.dump_folder
        )

    async def close(self):
        """Best-effort: tear down actors, close metric backends, then stop proc meshes."""
        logger.info("Closing: tearing down actors and process meshes.")

        if self.trainer is not None:
            try:
                await self.trainer.close.call()
            except Exception:
                logger.exception("trainer.close failed")

        try:
            await self._rollouter.close()
        except Exception:
            logger.exception("rollouter.close failed")

        if self.generator_router is not None:
            try:
                close_results = await self.generator_router.close_generators.call_one()
                for idx, result in enumerate(close_results):
                    if isinstance(result, BaseException):
                        actor_name = (
                            "generator"
                            if len(close_results) == 1
                            else f"generator[{idx}]"
                        )
                        logger.error(
                            "%s.close failed",
                            actor_name,
                            exc_info=(type(result), result, result.__traceback__),
                        )
            except Exception:
                logger.exception("generator_router.close_generators failed")

        try:
            self.metrics_processor.close()
        except Exception:
            logger.exception("metrics_processor close failed")

        for i, mesh in enumerate(self._proc_meshes):
            try:
                await mesh.stop()
            except Exception:
                logger.exception("mesh.stop[%d] failed", i)
        self._proc_meshes = []

    def _get_rank_0_value(self, result):
        """Extract rank 0 result from a Monarch ValueMesh.

        Monarch actor endpoints return results from all ranks in the mesh.
        This method picks out rank 0's result. This should be used in cases
        where all ranks return the same result.
        """
        return result.get(0)

    def _make_generate_fn(self, metrics_prefix: str) -> GenerateFn:
        """Build the rollouter's `GenerateFn`: route a completion via the generator router, namespacing
        generation metrics with `metrics_prefix` and pinning sticky routing on `routing_session_id` (a sample's
        turns reuse one generator's prefix KV)."""
        # TODO: make this a pluggable config (a GenerateFn factory) so non-router generate backends can be swapped in.
        # Bind the router handle to a local so the closure captures it instead of
        # `self`. A GenerateFn may be shipped to another process, where an actor
        # handle serializes cheaply and the whole controller does not.
        generator_router = self.generator_router

        @sl.log_trace_span("generate")
        async def generate(
            prompt_token_ids: list[int],
            *,
            request_id: str,
            group_id: int,
            routing_session_id: str | None = None,
            sampling_config: SamplingConfig | None = None,
        ) -> Completion | None:
            return await generator_router.generate.call_one(
                prompt_token_ids,
                request_id=request_id,
                group_id=group_id,
                routing_session_id=routing_session_id,
                sampling_config=sampling_config,
                metrics_prefix=metrics_prefix,
            )

        return generate

    @sl.log_trace_span("setup_async")
    async def setup_async(
        self,
        *,
        trainer_mesh: ProcMesh,
        generator_meshes: list[ProcMesh],
    ):
        """Spawn Monarch actors on separate meshes and initialize weights.

        Kept separate from ``__init__`` because actor spawning, torch
        elastic env setup, TorchStore initialization, and the initial
        weight push/pull are all ``await``-based runtime side effects
        that cannot run in a synchronous constructor.

        The trainer and generator meshes are provisioned by the caller (see
        ``spawn_proc_mesh``). The router and rollout worker meshes are created
        on the controller host. This method spawns the actors and synchronizes
        initial weights from trainer to generator. Must be called before
        :meth:`run`.

        Args:
            trainer_mesh: ProcMesh the trainer actor is spawned on.
            generator_meshes: ProcMesh objects the generator actors are spawned on.
        """
        # Peak concurrent rollout sequences (groups * num_samples_per_prompt, or the validation pass); sizes max_num_seqs below.
        async_loop = self.config.async_loop
        max_active_rollout_groups = async_loop.max_active_rollout_groups
        rollout_concurrency = max(
            max_active_rollout_groups * async_loop.num_samples_per_prompt,
            async_loop.validation.num_samples,
        )
        config = self.config
        if not generator_meshes:
            raise ValueError("setup_async requires at least one generator mesh")

        trainer_parallelism = config.trainer.parallelism
        dp_shard = max(trainer_parallelism.data_parallel_shard_degree, 1)
        self.trainer_dp_degree = (
            trainer_parallelism.data_parallel_replicate_degree * dp_shard
        )

        generator_dp_degree = max(config.generator.parallelism.data_parallel_degree, 1)
        num_generator_dp_shards = len(generator_meshes) * generator_dp_degree

        # Ceiling (not target) for the generator's max_num_seqs: the per-generator
        # upper bound on concurrently scheduled sequences. vLLM may admit fewer if KV
        # is tight; this also sets CUDA-graph capture sizes.
        max_num_seqs = min(
            math.ceil(rollout_concurrency / num_generator_dp_shards), 512
        )

        logger.info(
            "max_num_seqs=%d per generator (rollout_concurrency=%d / generator_dp_shards=%d)",
            max_num_seqs,
            rollout_concurrency,
            num_generator_dp_shards,
        )

        # TODO(observability): the mesh_spawn span wraps ~80 LoC of branching
        # provisioner logic. Pull a PerHostProvisioner.spawn_meshes(...) helper and
        # shrink this span to a single call.
        with sl.log_trace_span("mesh_spawn"):
            # One process, so the router is a singleton and every caller reaches
            # it with `call_one`. It gets its own mesh rather than sharing the
            # controller's process so routing does not contend with the training
            # loop for the controller's GIL.
            router_mesh = this_host().spawn_procs(per_host={"cpus": 1})
            # Store proc meshes for cleanup
            self._proc_meshes = [router_mesh, trainer_mesh, *generator_meshes]

            await setup_torch_elastic_env(trainer_mesh)
            for generator_mesh in generator_meshes:
                await setup_torch_elastic_env(generator_mesh)

            # Spawn actors on their respective meshes
            self.trainer = trainer_mesh.spawn(
                "trainer",
                TrainerActor,
                config.trainer,
                model_config=config.model,
                hf_assets_path=config.hf_assets_path,
                generator_dtype=config.generator.model_dtype,
                max_num_documents=config.async_loop.batcher.max_num_documents,
                output_dir=config.dump_folder,
            )

            # TODO: torch.compile with aot_eager backend (inductor crashes the vLLM engine on the shared model path).
            generators = []
            for idx, generator_mesh in enumerate(generator_meshes):
                actor_name = (
                    "generator" if len(generator_meshes) == 1 else f"generator_{idx}"
                )
                generator = generator_mesh.spawn(
                    actor_name,
                    VLLMGeneratorActor,
                    config.generator,
                    model_config=config.model,
                    model_path=config.hf_assets_path,
                    max_num_seqs=max_num_seqs,
                    output_dir=config.dump_folder,
                )
                generators.append(generator)
            self.generator_router = router_mesh.spawn(
                "generator_router",
                InterGeneratorRouter,
                config.generator_router,
                generators=generators,
                enable_cpu_weight_prefetch=config.generator.enable_cpu_weight_prefetch,
            )

            await self._rollouter.setup_async(
                tokenizer_config=config.tokenizer,
                renderer_config=config.renderer,
                hf_assets_path=config.hf_assets_path,
            )

        # Initialize TorchStore for weight sync between trainer and generator.
        # StorageVolumes are spawned on the trainer mesh so they are colocated
        # with the weight source for faster data access in the non-RDMA path.
        # https://github.com/meta-pytorch/torchstore
        with sl.log_trace_span("torchstore_init"):
            await ts.initialize(
                mesh=trainer_mesh,
                strategy=ts.TorchStoreStrategy(),
                client_type=ts.ClientType.ROUTING,
            )
            # TorchStore clients are cached per process. Initialize every actor
            # rank before the role-less state-dict APIs use that cache.
            await asyncio.gather(
                self.trainer.initialize_torchstore_client.call(),
                *(
                    generator.initialize_torchstore_client.call(
                        requester_index=requester_index
                    )
                    for requester_index, generator in enumerate(generators)
                ),
            )

        # Resume: __init__ ran CheckpointManager.load(); read back the restored policy_version
        # (0 if fresh) so the loop resumes at the right step and generators pull at that version.
        # TODO(resume): only model/optimizer/policy_version are restored. The active-slot rollout
        #   buffer (in-flight rollouts) and the dataset stream position are NOT restored -- a resumed
        #   run refills the buffer and re-reads data from the start. Need to recycle prompts.
        self.start_step = self._get_rank_0_value(
            await self.trainer.get_policy_version.call()
        )
        if self.start_step > 0:
            logger.info(f"Resuming RL training from step {self.start_step}")

        # Start each generator's engine loop on all ranks once, before any
        # rank-0-only generate / pull (rank 0 drives the followers through this
        # loop, so every rank must be running it first).
        with sl.log_trace_span("generator_start_engine_loop"):
            await self.generator_router.start_engine_loop.call_one()

        # Initial weight sync: only the trainer loads weights; generators pull at start_step.
        with sl.log_trace_span("trainer_push_model_state_dict"):
            await self.trainer.push_model_state_dict.call()
        with sl.log_trace_span("generator_pull_model_state_dict"):
            await self.generator_router.pull_model_state_dict.call_one(self.start_step)

    # TODO: fold validation into a Validator(Configurable) the controller attaches, instead of 4 methods.
    @sl.log_trace_span("_collect_validation_rollouts")
    async def _collect_validation_rollouts(
        self, *, num_groups: int, sampling: SamplingConfig, step: int
    ) -> tuple[list[RolloutGroup], list[m.Metric]]:
        """Sample held-out prompts, run each greedily (n=1) concurrently, and emit validation metrics."""
        # TODO: group_size=1 (best-of-1) only. Support best-of-N.
        generate = self._make_generate_fn(metrics_prefix="validation_generator")
        # TODO(naming): reserve "sample" for TrainingSample; rename the rollouter's raw-prompt "sample" -> "prompt"/"data_input".
        samples = [self._rollouter.get_validation_sample() for _ in range(num_groups)]
        group_results = await asyncio.gather(
            *(
                self._rollouter.run_group_rollouts(
                    generate_fn=generate,
                    sample=sample,
                    # Negative ids keep validation disjoint from training group ids, so their
                    # request_ids can't collide in the shared engine (e.g. post-validation).
                    group_id=-(i + 1),
                    group_size=1,
                    sampling=sampling,
                )
                for i, sample in enumerate(samples)
            ),
            return_exceptions=True,
        )
        # Validation group ids are reused every validation, so their cache salts must not outlive it.
        await self.generator_router.release_groups.call_one(
            [-(i + 1) for i in range(num_groups)]
        )

        # Keep the groups that succeeded; log + count the ones that raised.
        rollout_groups: list[RolloutGroup] = []
        num_failed_groups = 0
        for i, result in enumerate(group_results):
            if isinstance(result, BaseException):
                logger.error(
                    f"validation group {-(i + 1)} (step={step}) failed; dropping",
                    exc_info=(type(result), result, result.__traceback__),
                )
                num_failed_groups += 1
                continue
            rollout_groups.append(result)

        metrics = compute_rollout_metrics(
            prefix="validation",
            rollouts=[
                rollout for group in rollout_groups for rollout in group.rollouts
            ],
        )
        metrics.append(
            m.Metric("validation/group_failures", m.Sum(float(num_failed_groups)))
        )
        return rollout_groups, metrics

    # TODO: we currently determine validation.num_samples
    # but what if i want to run the entire dataset?
    @sl.log_trace_span("validate")
    async def validate(self, *, step: int) -> list[m.Metric]:
        """Run greedy validation on held-out prompts.

        Args:
            step: Training step this validation pass belongs to (0 for the
                pre-training pass); tagged into logged rollout samples.

        Returns:
            Validation rollout metrics, generation metrics, and validation
            timing.
        """
        # TODO: investigate using pass@k for validation.
        t_validate_start = time.perf_counter()
        num_samples = self.config.async_loop.validation.num_samples
        if num_samples == 0:  # skip validation (e.g. loss guard CI)
            return []
        greedy = replace(self._sampling, temperature=0.0, top_p=1.0)

        rollout_groups, validation_metrics = await self._collect_validation_rollouts(
            num_groups=num_samples, sampling=greedy, step=step
        )

        self.rollout_recorder.record(is_validation=True, rollout_groups=rollout_groups)

        t_validate_s = time.perf_counter() - t_validate_start
        validation_metrics.append(m.Metric("timing/validate", m.NoReduce(t_validate_s)))
        return validation_metrics

    async def run(self) -> None:
        """Start every async loop and run until training completes or a stage crashes.

        Producers (_data_input_loop, _rollout_loop[N], _batcher_loop) loop forever; _trainer_loop is the
        only finite loop -- it runs num_training_steps, then returns, which drives shutdown.

        Shutdown (healthy):  _trainer_loop finishes N steps -> run() finally ->
          group_buffer.close()  (wakes _data_input_loop / _rollout_loop / _batcher_loop blocked on the buffer)
          -> task.cancel()      (wakes anything blocked on training_batch_queue.put/get; close does NOT wake these)
          -> gather(..., return_exceptions=True)
        Shutdown (crash):    any loop raises -> appears in `done` -> run() re-raises -> same finally.
        """
        async_loop = self.config.async_loop
        num_training_steps = async_loop.num_training_steps
        logger.info(
            f"Running pre-training validation; then {num_training_steps} steps of async RL training"
        )

        sl.log_trace_instant("validation_start")
        pre_validation = await self._validate_and_log(step=self.start_step)
        sl.log_trace_instant("training_start")

        # Trainer policy version, seeded from the resumed step; advances at each optimizer step.
        self._trainer_policy_version = self.start_step

        # Depth (S + 1) * P targets the mean policy age; the window caps the max age.
        max_active_rollout_groups = async_loop.max_active_rollout_groups
        window_size = async_loop.window_size
        logger.info(
            f"max_active_rollout_groups={max_active_rollout_groups}, "
            f"target_offpolicy_steps={async_loop.target_offpolicy_steps}, "
            f"windowed_fifo_batches={async_loop.windowed_fifo_batches}, "
            f"max_offpolicy_steps={async_loop.max_offpolicy_steps}"
        )

        self._group_buffer = async_loop.group_buffer.build(
            max_active_rollout_groups=max_active_rollout_groups,
            window_size=window_size,
        )

        # Overlaps each step's weight handoff (push -> pull -> buffer-slot release) with the next step's fwd/bwd
        self._weight_sync = WeightSyncManager(
            trainer=self.trainer,
            generator_router=self.generator_router,
            group_buffer=self._group_buffer,
            num_prompts_per_train_step=async_loop.num_prompts_per_train_step,
        )

        # training_sample_builder
        training_sample_builder = async_loop.training_sample_builder.build()

        # batcher
        batcher = async_loop.batcher.build(
            num_tokens_per_microbatch_per_dp_rank=(
                self.config.trainer.training.num_tokens_per_microbatch_per_dp_rank
            ),
            max_context_length=self.config.trainer.training.max_context_length,
            num_prompts_per_train_step=async_loop.num_prompts_per_train_step,
            dp_degree=self.trainer_dp_degree,
            pad_id=self.tokenizer.eos_id,
            temperature=self._sampling.temperature,
        )

        # training_batch_queue
        training_batch_queue: asyncio.Queue[TrainerStepBatch | None] = asyncio.Queue(
            maxsize=1
        )

        # rollout_loop
        generate_fn = self._make_generate_fn(metrics_prefix="generator")

        # One rollout worker per active buffer slot: lets generation fill every active slot,
        # including the cold start (step 0 fills every active slot, not just num_prompts_per_train_step per wave).
        # TODO: support warm start
        rollout_tasks = [
            asyncio.create_task(
                self._rollout_loop(
                    group_buffer=self._group_buffer,
                    generate_fn=generate_fn,
                ),
                name=f"rollout_worker_{group_worker_id}",
            )
            for group_worker_id in range(max_active_rollout_groups)
        ]

        # data_input_loop
        data_input_task = asyncio.create_task(
            self._data_input_loop(self._group_buffer), name="data_input"
        )

        # training_sample_batcher_loop
        batcher_task = asyncio.create_task(
            self._batcher_loop(
                group_buffer=self._group_buffer,
                training_sample_builder=training_sample_builder,
                batcher=batcher,
                training_batch_queue=training_batch_queue,
            ),
            name="batcher",
        )

        # trainer_loop
        trainer_task = asyncio.create_task(
            self._trainer_loop(
                training_batch_queue, num_training_steps=num_training_steps
            ),
            name="trainer",
        )

        # run everything until trainer finishes its number of steps
        # or some other loop breaks
        background_tasks = [
            data_input_task,
            *rollout_tasks,
            batcher_task,
        ]
        try:
            done, _ = await asyncio.wait(
                [trainer_task, *background_tasks], return_when=asyncio.FIRST_COMPLETED
            )
            # The trainer is the finite clock: it runs num_training_steps then returns -> training is done.
            # Producers loop forever, so a producer in `done` means it crashed (await re-raises) or wrongly
            # returned cleanly (the RuntimeError). Check producers even when the trainer also finished this
            # wakeup, so a simultaneous producer crash isn't hidden behind the finished trainer.
            for task in done:
                if task is trainer_task:
                    continue
                await task  # raises if task crashed; returns if task ended cleanly
                raise RuntimeError(f"{task.get_name()} exited unexpectedly")
            if trainer_task in done:
                await trainer_task
        finally:
            # Graceful first: buffer.close() (awaited) wakes loops blocked on the buffer so they return.
            # Then cancel covers anything blocked on the queue (which close does not wake); gather awaits all.
            await self._group_buffer.close()
            for task in (*background_tasks, trainer_task):
                task.cancel()
            await asyncio.gather(
                *background_tasks, trainer_task, return_exceptions=True
            )

        # Post-training validation (held-out eval after the final step).
        post_validation = await self._validate_and_log(step=num_training_steps)
        self._log_reward_delta(pre_validation, post_validation)

    async def _validate_and_log(self, *, step: int) -> dict[str, float]:
        """Run one validation pass, log it, and return its aggregated values for the pre/post delta."""
        metrics = await self.validate(step=step)
        self.metrics_processor.log(step=step, metrics=metrics, is_validation=True)
        return m.MetricsProcessor._aggregate_metrics(metrics)

    def _log_reward_delta(self, pre: dict[str, float], post: dict[str, float]) -> None:
        """Console pre/post reward summary, visible without scrolling back through the loop."""
        reward_keys = sorted(key for key in set(pre) | set(post) if "reward" in key)
        logger.info("=" * 60)
        logger.info("Validation reward (pre / post):")
        for key in reward_keys:
            logger.info(
                f"  {key}:  {pre.get(key, float('nan')):+.3f}  /  {post.get(key, float('nan')):+.3f}"
            )
        logger.info("=" * 60)

    async def _data_input_loop(self, group_buffer: RolloutGroupWorkBuffer) -> None:
        """produces a RolloutGroupWork into group_buffer.
        waits for:    a free active slot (group_buffer.wait_for_slot)
        unblocked by: _trainer_loop release_active_groups(num_prompts_per_train_step, "trained")
            after the pull (and _batcher_loop release_active_groups(1,"untrainable_group"))

        Separate from `_rollout_loop`, so slow data prep (e.g. on-the-fly question generation) overlaps
        generation instead of serializing in front of it.
        """
        # TODO(resume): persist dataset position so a restarted job continues the data stream, not from scratch.
        group_index = 0

        # TODO(perf): Slots are current released in batches, while this loop is a single producer.
        # we could a) increase the number of threads; b) revisit how we release slots and see if
        # we can release them on the batcher while still preserving max offpolicy steps.
        # finally, c) we need to check how will this data input loop truly overlaps with the rollout loop.
        while await group_buffer.wait_for_slot():
            with sl.log_trace_span("get_training_sample"):
                # to_thread: Dont block on dataset reads
                sample = await asyncio.to_thread(self._rollouter.get_training_sample)
            await group_buffer.add_work(
                RolloutGroupWork(
                    group_id=group_index,
                    sample=sample,
                )
            )
            group_index += 1
        logger.info("Buffer closed; data input loop stopping")

    async def _rollout_loop(
        self,
        *,
        group_buffer: RolloutGroupWorkBuffer,
        generate_fn: GenerateFn,
    ) -> None:
        """Generate + score one group at a time; a failed group becomes an empty group + a failure metric.

        Staleness is bounded by the buffer's active-slot budget. Raw rollouts are recorded before any drop,
        so dropped groups stay inspectable on disk.

        consumes: a WAITING RolloutGroupWork (group_buffer.claim_next)
            waits for:    a claimable WAITING entry
            unblocked by: _data_input_loop group_buffer.add_work()
        produces: RolloutGroup (group_buffer.finalize_work)
            waits for:    nothing (admits its own claimed slot)
            unblocked by: n/a
        """
        while True:
            work = await group_buffer.claim_next()
            if work is None:  # group_buffer closed/shutdown signal
                logger.info("Buffer closed; rollout worker stopping")
                return
            try:
                with sl.log_trace_span("rollout_group"):
                    group = await self._rollouter.run_group_rollouts(
                        generate_fn=generate_fn,
                        sample=work.sample,
                        group_id=work.group_id,
                        group_size=self.config.async_loop.num_samples_per_prompt,
                        sampling=self._sampling,
                    )
                group.metrics = compute_rollout_metrics(
                    prefix="rollout", rollouts=group.rollouts
                )

                # save rollout for inspection
                self.rollout_recorder.record(
                    is_validation=False,
                    rollout_groups=[group],
                )
            except Exception:
                logger.exception(f"rollout group {work.group_id} failed; dropping")
                group = RolloutGroup(
                    group_id=work.group_id,
                    rollouts=[],
                    metrics=[m.Metric("rollout/group_failures", m.Sum(1.0))],
                )
            # The group makes no more generation calls, so its cache salts can go.
            await self.generator_router.release_groups.call_one([work.group_id])
            await group_buffer.finalize_work(group)

    async def _batcher_loop(
        self,
        *,
        group_buffer: RolloutGroupWorkBuffer,
        training_sample_builder: TrainingSampleBuilder,
        batcher: Batcher,
        training_batch_queue: "asyncio.Queue[TrainerStepBatch | None]",
    ) -> None:
        """Take finalized groups, build training_samples, accumulate them, and queue each ready training batch.

        On a clean close/shutdown the group_buffer drains and returns None; we forward a `None` sentinel
        so the trainer stops.

        consumes: the oldest FINALIZED group inside the window (group_buffer.take_finalized)
            waits for:    a group inside the window becoming FINALIZED (any group when windowed_fifo_batches is None)
            unblocked by: _rollout_loop[N] group_buffer.finalize_work()
        produces: TrainerStepBatch (training_batch_queue.put)
            waits for:    a free training_batch_queue slot (maxsize=1)
            unblocked by: _trainer_loop training_batch_queue.get()
        """
        while True:
            rollout_group = await group_buffer.take_finalized()
            if rollout_group is None:  # closed and drained
                logger.info("Buffer drained; batcher loop stopping")
                break
            with sl.log_trace_span("training_sample_builder"):
                training_sample_group = training_sample_builder.build_from_group(
                    rollout_group=rollout_group
                )

            # We put a group in. We may get a batch back
            # if there are enough accumulated trainable groups to return one.
            with sl.log_trace_span("batcher_pack"):
                maybe_training_batch, group_is_trainable = await asyncio.to_thread(
                    batcher.add_training_samples,
                    training_sample_group=training_sample_group,
                )
            if not group_is_trainable:
                await group_buffer.release_active_groups(1, reason="untrainable_group")
            if maybe_training_batch is not None:
                await training_batch_queue.put(maybe_training_batch)
        await training_batch_queue.put(None)
        # TODO(async-rl): if finite datasets are supported, drain a final partial batch here.

    async def _trainer_loop(
        self,
        training_batch_queue: "asyncio.Queue[TrainerStepBatch | None]",
        *,
        num_training_steps: int,
    ) -> None:
        """Run num_training_steps optimizer steps: train one packed batch, publish trainer weights,
        then pull them into generators, log metrics.

        NOTE: Weight sync is overlapped with the training step.
        Trainer push:
            - Called after optimizer.step()
            - Awaited before next optimizer.step (weights changes then)
        Generator pull:
            - Called after push completes.
            - Awaited before next push (weights changes then)

        Impact on off-policiness: The buffer guarantees that no sample will be born stale,
        as long as we call `self._group_buffer.release_active_groups` after the pull.

        consumes: a TrainerStepBatch (training_batch_queue.get)
            waits for:    a TrainerStepBatch in the queue
            unblocked by: _batcher_loop training_batch_queue.put()
        """
        for step in range(self.start_step + 1, num_training_steps + 1):
            sl.set_step(step)  # propagate the step counter to the actors
            with sl.log_trace_span("sync_log_step"):
                await self.trainer.sync_log_step.call(step)
                await self.generator_router.sync_log_step.call_one(step)
                await self._rollouter.sync_log_step(step)
            step_timer = MetricsTimer()

            with (
                sl.log_trace_span("train_step"),
                step_timer.record("timing/step/total"),
            ):
                # Waits for a TrainerStepBatch to be ready (or None on shutdown).
                with (
                    sl.log_trace_span("wait_for_training_batch"),
                    step_timer.record("timing/step/wait_for_training_batch"),
                ):
                    packed = await training_batch_queue.get()

                if packed is None:
                    logger.info("Batcher closed and drained; stopping training")
                    break

                # Policy age is computed HERE, at consumption time, against the live trainer version, so it is
                # faithful to what this step trains on -- not the version when the batch was packed.
                policy_age_panel = compute_policy_age_metrics(
                    trainer_policy_version=self._trainer_policy_version,
                    min_policy_versions=packed.min_policy_versions,
                    target_offpolicy_steps=(
                        self.config.async_loop.target_offpolicy_steps
                    ),
                    max_offpolicy_steps=self.config.async_loop.max_offpolicy_steps,
                )

                # TODO(async): can't stream microbatches (interleave pack->train) -- the loss is normalized by
                #   global counts over ALL microbatches, needed before any fwd/bwd. To
                #   support streaming, accumulate raw loss/token counts across microbatches and scale before optimizer.
                with (
                    sl.log_trace_span("forward_backward_steps"),
                    step_timer.record("timing/step/forward_backward"),
                ):
                    fwd_bwd_metrics = self._get_rank_0_value(
                        await self.trainer.forward_backward_steps.call(
                            packed.microbatches,
                            packed.global_loss_token_counts,
                            packed.global_routing_token_counts,
                        )
                    )

                    if not math.isfinite(fwd_bwd_metrics["loss/mean"]):
                        logger.error("Loss is NaN/Inf; training diverged")
                        break

                # Await trainer weight push before the optimizer mutates the weights.
                with (
                    sl.log_trace_span("blocking_trainer_push_model_state_dict"),
                    step_timer.record(
                        "timing/step/blocking_trainer_push_model_state_dict"
                    ),
                ):
                    push_metrics = await self._weight_sync.wait_prev_push()

                with (
                    sl.log_trace_span("optimizer_step"),
                    step_timer.record("timing/step/optimizer"),
                ):
                    optimizer_result = self._get_rank_0_value(
                        await self.trainer.optimizer_step.call(
                            last_step=(step == num_training_steps)
                        )
                    )
                self._trainer_policy_version = optimizer_result.policy_version

                # Await generator weight pull to finish before the trainer's next push.
                with (
                    sl.log_trace_span("blocking_generator_pull_model_state_dict"),
                    step_timer.record(
                        "timing/step/blocking_generator_pull_model_state_dict"
                    ),
                ):
                    pull_metrics = await self._weight_sync.wait_prev_pull()

                # Overlap this step's push -> pull -> buffer-slot release with the next step's fwd/bwd.
                self._weight_sync.start_async_push_pull(
                    version=optimizer_result.policy_version
                )

            # TODO(metrics): See if metrics are being computed at the right place. E.g. should we put all
            # rollout related metrics here, or move all of them to the rollouter.
            time_metrics = step_timer.flush()
            with sl.log_trace_span("metrics_log"):
                self.metrics_processor.log(
                    step=step,
                    is_validation=False,
                    metrics=[
                        *packed.metrics,
                        *[
                            m.Metric(key, m.NoReduce(value))
                            for key, value in fwd_bwd_metrics.items()
                        ],
                        *[
                            m.Metric(key, m.NoReduce(value))
                            for key, value in optimizer_result.metrics.items()
                        ],
                        *self._group_buffer.metrics(),
                        *time_metrics,
                        *policy_age_panel,
                        # Background push/pull work time; the trainer's wait for it is timing/step/blocking_*.
                        *push_metrics,
                        *pull_metrics,
                        *compute_perf_ratio_metrics(
                            num_global_valid_tokens=int(
                                packed.global_loss_token_counts[0]
                            ),
                            time_metrics=time_metrics,
                        ),
                    ],
                )

        # Finish the last in-flight sync so generators hold the final weights for post-validation.
        await self._weight_sync.wait_inflight_push_pull()
