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

"""
RL training loop using Monarch Actors.

This demonstrates:
1. Distributed actor architecture with VLLMGenerator (vLLM) and Trainer (TorchTitan)
   running on separate GPU meshes
2. Weight synchronization across meshes via TorchStore: the trainer publishes its
   model state dict and the generator pulls it into its own parallelism layout,
   with direct GPU-to-GPU RDMA transfer when available
3. Envs driven rollouts; reward and advantage computation live inline
   in the controller.

Command to run:
python3 -m torchtitan.rl.train \
    --module my_rl_configs --config rl_grpo_qwen3_0_6b_varlen

Set ``hf_assets_path`` in ``my_rl_configs.py`` to the model checkpoint path.
"""

import asyncio
import logging
from collections.abc import Callable
from dataclasses import dataclass

from monarch.actor import default_bootstrap_cmd, HostMesh, ProcMesh, this_host

from torchtitan.config import ConfigLoader
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.observability import structured_logger as sl
from torchtitan.observability.logging import init_logger
from torchtitan.rl.controller import Controller
from torchtitan.rl.distributed.parallelism import InferenceParallelismConfig


logger = logging.getLogger(__name__)


def _preimport_torch() -> None:
    """``bootstrap`` setup callable: pre-import torch on the spawned proc."""
    # TODO: Remove once Monarch/PyTorch fixes concurrent import during unpickling.
    import torch  # noqa: F401


def _bootstrap_generator() -> None:
    """``bootstrap`` setup callable for VLLMGenerator."""
    import os

    # TODO: this can be removed if it is addressed upstream in: https://github.com/tile-ai/tvm/issues/55
    # VLLMGenerator encounters the following issue:
    # 1. vLLM's jit_monitor transitively imports readline through the following chain:
    #      vLLM -> tilelang  -> tvm  -> tvm's base.py -> import readline
    #      https://github.com/vllm-project/vllm/blob/b23bd73f540175f9e117eaee5029cd7d8df63964/vllm/utils/jit_monitor.py#L451
    #      https://github.com/tile-ai/tvm/blob/28a0d34420d2fa9bc71fc891445a3f1396fca759/python/tvm/base.py#L71
    #    readline does tcsetattr on its import, which touches the controlling terminal.
    # 2. Monarch spawns procs in a background process group.
    #
    # As a result, when spawning the VLLMGenerator mesh, readline-import touches the
    # controlling terminal in background processes, and kernel sends SIGTTOU, which stops
    # the process.
    #
    # Strictly speaking this should be a tvm bug, since it should not assume the import
    # happens in a foreground process. To workaround it, we do a proactive readline-import
    # warmup here with SIGTTOU being temporarily blocked. Then the subsequent readline-import
    # in tvm will not do tcsetattr again since it is already done here.
    if os.isatty(0):
        import signal

        _prev_mask = signal.pthread_sigmask(signal.SIG_BLOCK, {signal.SIGTTOU})
        try:
            import readline  # noqa: F401
        finally:
            signal.pthread_sigmask(signal.SIG_SETMASK, _prev_mask)

    _preimport_torch()


class PerHostProvisioner:
    """Allocates non-overlapping GPU ranges within a single host.

    On the same host, the trainer and generator run on separate GPU
    meshes (e.g. GPUs 0-3 for training, GPUs 4-7 for generation). Each
    call to `allocate(n)` reserves the next *n* GPUs and returns the launch
    env (`CUDA_VISIBLE_DEVICES` plus any `extra_env`) for those GPUs. The
    returned env vars are applied to the spawned process via ``bootstrap_command``.
    """

    def __init__(self, total_gpus: int = 8):
        self.total_gpus = total_gpus
        self.next_gpu = 0

    @property
    def available(self) -> int:
        return self.total_gpus - self.next_gpu

    def allocate(
        self, num_gpus: int, *, extra_env: dict[str, str] | None = None
    ) -> dict[str, str]:
        if num_gpus > self.available:
            raise RuntimeError(
                f"Requested {num_gpus} GPUs but only {self.available} "
                f"available (total={self.total_gpus}, allocated={self.next_gpu})"
            )
        gpu_ids = list(range(self.next_gpu, self.next_gpu + num_gpus))
        self.next_gpu += num_gpus

        env = {"CUDA_VISIBLE_DEVICES": ",".join(str(g) for g in gpu_ids)}
        if extra_env:
            env.update(extra_env)
        return env


@dataclass
class HostMeshes:
    trainer: HostMesh
    generators: list[HostMesh]
    gpus_per_node: int


def _compute_trainer_world_size(p: ParallelismConfig) -> int:
    """Compute world size from all parallel dimensions."""
    dp_shard = max(p.data_parallel_shard_degree, 1)
    return (
        p.data_parallel_replicate_degree
        * dp_shard
        * p.tensor_parallel_degree
        * p.pipeline_parallel_degree
        * p.context_parallel_degree
    )


def _compute_generator_world_size(p: InferenceParallelismConfig) -> int:
    """Number of GPU processes for one generator (vLLM) instance."""
    return p.data_parallel_degree * p.tensor_parallel_degree


def _spawn_proc_mesh(
    host_mesh: HostMesh,
    role_world_size: int,
    gpus_per_node: int,
    *,
    bootstrap: Callable[[], None],
    role: str,
    extra_env: dict[str, str] | None = None,
) -> ProcMesh:
    """Spawn one role's proc mesh on ``host_mesh``, splitting ``role_world_size``
    evenly across the mesh's hosts. ``extra_env`` is applied in each proc's bootstrap.
    """
    nodes = len(host_mesh)
    assert role_world_size % nodes == 0, (
        f"{role} world size ({role_world_size}) must be evenly divisible by its "
        f"host count ({nodes})"
    )
    role_gpus_per_node = role_world_size // nodes
    provisioner = PerHostProvisioner(total_gpus=gpus_per_node)
    env = provisioner.allocate(role_gpus_per_node, extra_env=extra_env)
    return host_mesh.spawn_procs(
        per_host={"gpus": role_gpus_per_node},
        bootstrap=bootstrap,
        bootstrap_command=default_bootstrap_cmd().with_env(env),
    )


def spawn_proc_mesh(
    trainer_world_size: int,
    per_generator_world_size: int,
    host_meshes: HostMeshes | None = None,
    *,
    num_generators: int = 1,
    generator_env: dict[str, str] | None = None,
) -> tuple[ProcMesh, list[ProcMesh]]:
    """Spawn the trainer and generator proc meshes.

    Args:
        trainer_world_size: Number of GPU procs to spawn for the trainer.
        per_generator_world_size: Number of GPU procs to spawn for each generator.
        host_meshes: Caller-provided trainer/generator host meshes. When
            provided, each role is spawned on its provided host mesh. None means
            both roles are spawned on ``this_host()`` by using non-overlapping
            GPU ranges.
        num_generators: Number of generator proc meshes to spawn.

    Returns:
        The ``(trainer_mesh, generator_meshes)`` proc meshes.
    """
    total_generator_gpus = num_generators * per_generator_world_size
    total_gpus = trainer_world_size + total_generator_gpus
    logger.info(
        f"{num_generators} generator(s) * {per_generator_world_size} GPUs + "
        f"{trainer_world_size} trainer GPUs = {total_gpus} total"
    )

    if host_meshes is not None:
        trainer_host_mesh = host_meshes.trainer
        generator_host_meshes = host_meshes.generators
        gpus_per_node = host_meshes.gpus_per_node

        assert len(generator_host_meshes) == num_generators, (
            f"expected {num_generators} generator host mesh(es), "
            f"got {len(generator_host_meshes)}"
        )

        trainer_mesh = _spawn_proc_mesh(
            trainer_host_mesh,
            trainer_world_size,
            gpus_per_node,
            bootstrap=_preimport_torch,
            role="trainer",
        )
        generator_meshes = [
            _spawn_proc_mesh(
                gen_host_mesh,
                per_generator_world_size,
                gpus_per_node,
                bootstrap=_bootstrap_generator,
                role="generator",
                extra_env=generator_env,
            )
            for gen_host_mesh in generator_host_meshes
        ]
    else:
        # Single-node mode: partition GPUs on this_host() via
        # CUDA_VISIBLE_DEVICES
        host_mesh = this_host()
        provisioner = PerHostProvisioner(total_gpus=total_gpus)
        trainer_mesh = host_mesh.spawn_procs(
            per_host={"gpus": trainer_world_size},
            bootstrap=_preimport_torch,
            bootstrap_command=default_bootstrap_cmd().with_env(
                provisioner.allocate(trainer_world_size)
            ),
        )
        generator_meshes = [
            host_mesh.spawn_procs(
                per_host={"gpus": per_generator_world_size},
                bootstrap=_bootstrap_generator,
                bootstrap_command=default_bootstrap_cmd().with_env(
                    provisioner.allocate(
                        per_generator_world_size, extra_env=generator_env
                    )
                ),
            )
            for _ in range(num_generators)
        ]

    return trainer_mesh, generator_meshes


async def main():
    init_logger()
    config = ConfigLoader().load()
    assert isinstance(config, Controller.Config)
    sl.init_structured_logger(
        source="rl_controller",
        output_dir=config.dump_folder,
        rank=0,
        enable=config.trainer.debug.enable_structured_logging,
    )
    sl.log_trace_instant("structured_logger_started")

    rl_trainer: Controller = config.build()
    try:
        trainer_world_size = _compute_trainer_world_size(config.trainer.parallelism)
        per_generator_world_size = _compute_generator_world_size(
            config.generator.parallelism
        )
        trainer_mesh, generator_meshes = spawn_proc_mesh(
            trainer_world_size,
            per_generator_world_size,
            host_meshes=None,
            num_generators=config.num_generators,
        )
        await rl_trainer.setup_async(
            trainer_mesh=trainer_mesh,
            generator_meshes=generator_meshes,
        )
        await rl_trainer.run()
    except (KeyboardInterrupt, asyncio.CancelledError):
        logger.info("Interrupted; attempting graceful shutdown...")
        raise
    finally:
        await rl_trainer.close()


if __name__ == "__main__":
    asyncio.run(main())
