archit11/verl-code-corpus-track-a-file-split
archit11/verl-code-corpus-track-a-file-split Repository-specific code corpus extracted from the verl project and split by file for training/evaluation. What is in this dataset Source corpus: data/code_corpus_verl Total files: 214 Train files: 172 Validation files: 21 Test files: 21 File type filter: .py Split mode: file (file-level holdout) Each row has: file_name: flattened source file name text: full file contents Training context This dataset… See the full description on the dataset page: https://huggingface.co/datasets/archit11/verl-code-corpus-track-a-file-split.
047
1{"file_name": "verl__base_config.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport collections\nfrom dataclasses import FrozenInstanceError, dataclass, fields\nfrom typing import Any\n\n\n# BaseConfig class inherits from collections.abc.Mapping, which means it can act like a dictionary\n@dataclass\nclass BaseConfig(collections.abc.Mapping):\n \"\"\"The BaseConfig provides dict-like interface for a dataclass config.\n\n By default all fields in the config is not mutable, unless specified in\n \"_mutable_fields\". The BaseConfig class implements the Mapping Abstract Base Class.\n This allows instances of this class to be used like dictionaries.\n \"\"\"\n\n _mutable_fields = set()\n _target_: str = \"\"\n\n def __setattr__(self, name: str, value):\n \"\"\"Set the value of an attribute. Check if the attr is mutable before setting the value.\"\"\"\n # If the field already exists, it's considered frozen unless it's in _mutable_fields\n if name in self.__dict__ and name not in getattr(self, \"_mutable_fields\", set()):\n raise FrozenInstanceError(f\"Field '{name}' is frozen and cannot be modified\")\n super().__setattr__(name, value)\n\n def get(self, key: str, default: Any = None) -> Any:\n \"\"\"Get the value associated with the given key. If the key does not exist, return the default value.\n\n Args:\n key (str): The attribute name to retrieve.\n default (Any, optional): The value to return if the attribute does not exist. Defaults to None.\n\n Returns:\n Any: The value of the attribute or the default value.\n \"\"\"\n try:\n return getattr(self, key)\n except AttributeError:\n return default\n\n def __getitem__(self, key: str):\n \"\"\"Implement the [] operator for the class. Allows accessing attributes like dictionary items.\n\n Args:\n key (str): The attribute name to retrieve.\n\n Returns:\n Any: The value of the attribute.\n\n Raises:\n AttributeError: If the attribute does not exist.\n TypeError: If the key type is not string\n \"\"\"\n return getattr(self, key)\n\n def __iter__(self):\n \"\"\"Implement the iterator protocol. Allows iterating over the attribute names of the instance.\n\n Yields:\n str: The name of each field in the dataclass.\n \"\"\"\n for f in fields(self):\n yield f.name\n\n def __len__(self):\n \"\"\"\n Return the number of fields in the dataclass.\n\n Returns:\n int: The number of fields in the dataclass.\n \"\"\"\n return len(fields(self))\n"}2{"file_name": "verl__checkpoint_engine__base.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport asyncio\nfrom abc import ABC, abstractmethod\nfrom typing import Any, Generator, TypedDict\n\nimport ray\nimport torch\n\nfrom verl.single_controller.base import Worker\nfrom verl.single_controller.base.decorator import Dispatch, register\nfrom verl.single_controller.ray import RayClassWithInitArgs, RayWorkerGroup\nfrom verl.utils.distributed import initialize_global_process_group_ray\nfrom verl.utils.ray_utils import auto_await\nfrom verl.workers.config import HFModelConfig, RolloutConfig\nfrom verl.workers.rollout import BaseRollout, RolloutReplica, get_rollout_class\n\n\nclass TensorMeta(TypedDict):\n name: str\n shape: torch.Size\n dtype: torch.dtype\n offset: int\n\n\nclass CheckpointEngineRegistry:\n \"\"\"Checkpoint engine registry.\"\"\"\n\n _registry: dict[str, type[\"CheckpointEngine\"]] = {}\n\n def register(backend: str):\n \"\"\"Register a checkpoint engine.\n\n Args:\n backend: The backend of the checkpoint engine.\n \"\"\"\n\n def wrapper(cls: type[\"CheckpointEngine\"]):\n CheckpointEngineRegistry._registry[backend] = cls\n return cls\n\n return wrapper\n\n @classmethod\n def get(cls, backend: str) -> type[\"CheckpointEngine\"]:\n \"\"\"Get the checkpoint engine class.\n\n Args:\n backend: The backend of the checkpoint engine.\n\n Returns:\n The checkpoint engine class.\n \"\"\"\n return cls._registry[backend]\n\n @classmethod\n def new(cls, backend: str, *args, **kwargs) -> \"CheckpointEngine\":\n \"\"\"Create a new checkpoint engine instance.\n\n Args:\n backend: The backend of the checkpoint engine.\n *args: Variable length argument pass to the checkpoint engine constructor.\n **kwargs: Arbitrary keyword arguments pass to the checkpoint engine constructor.\n\n Returns:\n A new checkpoint engine instance.\n \"\"\"\n if backend not in cls._registry:\n raise ValueError(f\"Checkpoint engine {backend} not registered\")\n return cls._registry[backend](*args, **kwargs)\n\n\nclass CheckpointEngine(ABC):\n \"\"\"CheckpointEngine is an abstraction to transfer weights from trainer to rollout.\n\n In trainer process:\n >>> trainer = EngineRegistry.new(...) # FSDP, Megatron, VeOmini, TorchTitan, ...\n >>> engine = CheckpointEngine.new(...) # NCCLCheckpointEngine, NIXLCheckpointEngine, ...\n >>> await engine.send_weights(trainer.get_per_tensor_param())\n\n In rollout process:\n >>> engine = CheckpointEngine.new(...)\n >>> server_adapter = ServerAdapter()\n >>> await server_adapter.update_weights(engine.get_weights()) # update weights via cuda ipc\n \"\"\"\n\n @abstractmethod\n def prepare(self) -> dict[str, Any]:\n \"\"\"Prepare checkpoint engine before each step send_weights/receive_weights.\n\n 1. Allocate weight bucket.\n 2. [Optional] Register weight bucket for RDMA.\n 3. Return metadata to build communication topology: master ip:port, register RDMA description, etc.\n\n Args:\n worker_group: The worker group that the checkpoint engine will be used.\n\n Returns:\n A dictionary that contains the metadata of the worker group.\n \"\"\"\n raise NotImplementedError\n\n @classmethod\n @abstractmethod\n def build_topology(\n cls, trainer_world_size: int, rollout_world_size: int, metadata: list[dict]\n ) -> tuple[dict[str, list[Any]], dict[str, list[Any]]]:\n \"\"\"Build communication topology between all workers.\n\n Args:\n trainer_world_size: The world size of the trainer worker group.\n rollout_world_size: The world size of the rollout replica.\n metadata: A list of metadata `prepare` from all workers.\n\n Returns:\n A tuple of two dictionaries that contains the communication topology for trainer and rollout worker group.\n Each dict value should be a list argument equal to the world size of the worker group to dispatch to\n `init_process_group`.\n\n ```\n world_size = rollout.world_size + trainer.world_size\n kwargs = {\n \"rank\": list(range(world_size)),\n \"world_size\": [world_size] * world_size,\n \"master_metadata\": [metadata[0]] * world_size,\n }\n ```\n \"\"\"\n raise NotImplementedError\n\n @abstractmethod\n def init_process_group(self, **kwargs):\n \"\"\"Init process group for checkpoint engine.\n\n Args:\n **kwargs: Keyword arguments from `build_topology`.\n \"\"\"\n raise NotImplementedError\n\n @abstractmethod\n def finalize(self):\n \"\"\"Finalize checkpoint engine after each step send_weights/receive_weights.\n\n 1. Free weight bucket.\n 1. [Optional] Deregister weight bucket for RDMA.\n 2. [Optional] Destroy process group.\n \"\"\"\n raise NotImplementedError\n\n @abstractmethod\n async def send_weights(self, weights: Generator[tuple[str, torch.Tensor], None, None]):\n \"\"\"Send the weights of the model.\n\n Args:\n weights: A generator that yields the name of the weight tensor and the tensor itself.\n \"\"\"\n raise NotImplementedError\n\n @abstractmethod\n async def receive_weights(self) -> Generator[tuple[str, torch.Tensor], None, None]:\n \"\"\"Receive the weights of the model.\n\n Yields:\n A tuple of the name of the weight tensor and the tensor itself.\n \"\"\"\n raise NotImplementedError\n\n\nclass CheckpointEngineWithCache(CheckpointEngine):\n \"\"\"Checkpoint engine with local cache: shm, disk, etc. This allow to synchronize weights without interrupting\n rollout ongoing requests (partial rollout). After requests exhausted, rollout can get weights from local cache.\n\n Laminar: https://arxiv.org/abs/2510.12633\n \"\"\"\n\n @abstractmethod\n async def get_weights(self) -> Generator[tuple[str, torch.Tensor], None, None]:\n \"\"\"Get the weights of the model from local cache.\n\n Yields:\n A tuple of the name of the weight tensor and the tensor itself.\n \"\"\"\n raise NotImplementedError\n\n\n@CheckpointEngineRegistry.register(\"naive\")\nclass ColocatedCheckpointEngine(CheckpointEngine):\n \"\"\"Checkpoint engine for trainer and rollout colocated on same GPU.\n\n In trainer process:\n >>> engine = ColocatedCheckpointEngine()\n >>> trainer = Trainer()\n >>> server_adapter = ServerAdapter()\n >>> engine.send_weights(trainer.get_per_tensor_param())\n >>> server_adapter.update_weights(engine.receive_weights())\n \"\"\"\n\n def __init__(self, bucket_size: int, is_master: bool = False) -> None:\n self.bucket_size = bucket_size\n self.is_master = is_master\n\n def prepare(self):\n raise NotImplementedError\n\n def init_process_group(self, **kwargs):\n raise NotImplementedError\n\n def finalize(self):\n raise NotImplementedError\n\n @classmethod\n def build_topology(cls, *args, **kwargs):\n raise NotImplementedError\n\n def send_weights(self, weights: Generator[tuple[str, torch.Tensor], None, None]):\n \"\"\"Send the weights of the model.\n\n Args:\n weights: A generator that yields the name of the weight tensor and the tensor itself.\n \"\"\"\n self.weights = weights\n\n def receive_weights(self) -> Generator[tuple[str, torch.Tensor], None, None]:\n \"\"\"Receive the weights of the model.\n\n Yields:\n A tuple of the name of the weight tensor and the tensor itself.\n \"\"\"\n yield from self.weights\n self.weights = None\n\n\nclass CheckpointEngineWorker(Worker):\n \"\"\"CheckpointEngineWorker colocated with inference engine's WorkerProc on same GPU.\n\n Args:\n rollout_config: The rollout configuration.\n model_config: The model configuration.\n server_adapter: The server adapter to update weights.\n \"\"\"\n\n def __init__(\n self,\n rollout_config: RolloutConfig,\n model_config: HFModelConfig,\n server_adapter: BaseRollout = None,\n ) -> None:\n self.rollout_config = rollout_config\n self.model_config = model_config\n\n # sglang and trt-llm need device_mesh for internal communication\n initialize_global_process_group_ray(timeout_second=None, backend=\"cpu:gloo\")\n self.server_adapter: BaseRollout = server_adapter or get_rollout_class(\n rollout_config.name, rollout_config.mode\n )(config=rollout_config, model_config=model_config, device_mesh=None)\n\n backend = rollout_config.checkpoint_engine.backend\n bucket_size = rollout_config.checkpoint_engine.update_weights_bucket_megabytes << 20\n engine_kwargs = rollout_config.checkpoint_engine.engine_kwargs.get(backend, {})\n self.checkpoint_engine = CheckpointEngineRegistry.new(backend, bucket_size=bucket_size, **engine_kwargs)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=False)\n async def update_weights(self):\n weights = self.checkpoint_engine.receive_weights()\n await self.server_adapter.update_weights(weights)\n\n @register(dispatch_mode=Dispatch.DP_COMPUTE, blocking=False)\n def execute_checkpoint_engine(self, method: str, *args, **kwargs):\n return getattr(self.checkpoint_engine, method)(*args, **kwargs)\n\n\n_worker_cls = ray.remote(CheckpointEngineWorker)\n\n\nclass CheckpointEngineManager:\n \"\"\"Checkpoint engine manager to coordinate weight synchronization between trainer and rollout replicas.\n\n - ME: model engine, FSDP, MCore, VeOmni, export full tensor generator `get_per_tensor_param`\n - CE: checkpoint engine, NCCL, NIXL, etc\n\n In trainer, model engine and checkpoint engine are in same process.\n In rollout, checkpoint engine and rollout worker are in separate process, update weights via cuda ipc.\n\n ```\n ┌────────┬────────┬─────┬────────┐ ┌───────────────────┬───────────────────┐\n │ ┌────┐ │ ┌────┐ │ │ ┌────┐ │ │ Replica 0 │ Replica 1 │\n │ │ ME0│ │ │ ME1│ │ │ │ MEn│ │ ├────┬────┬────┬────┼────┬────┬────┬────┤\n │ └──┬─┘ │ └────┘ │ ... │ └────┘ │ │ 0 │ 1 │ 2 │ 3 │ 0 │ 1 │ 2 │ 3 │\n │ v | | | | └──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┘\n | ┌──┴─┐ │ ┌────┐ │ │ ┌────┐ │ ^ ^ ^ cuda ipc ^ ^ ^\n │ │ CE │ │ │ CE │ │ │ │ CE │ │ ┌──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┐\n │ └──┬─┘ │ └────┘ │ │ └────┘ │ │ CE │ CE │ CE │ CE │ CE │ CE │ CE │ CE |\n └────┼───┴────────┴─────┴────────┘ └──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┘\n v | | | | | | | |\n └─────────────(nccl/nixl/..)─────────────┴────┴────┴────┴────┴────┴────┴────┘\n ```\n\n Args:\n backend: The checkpoint engine backend.\n trainer: The trainer worker group.\n replicas: The list of rollout replicas.\n \"\"\"\n\n def __init__(\n self,\n backend: str,\n trainer: RayWorkerGroup,\n replicas: list[RolloutReplica],\n ) -> None:\n self.backend = backend\n self.backend_cls = CheckpointEngineRegistry.get(backend)\n self.trainer = trainer\n self.replicas = replicas\n\n def build_process_group(self, rollout: RayWorkerGroup):\n \"\"\"Build process group for trainer and rollout replicas.\"\"\"\n trainer = self.trainer\n\n # 1. prepare all workers\n metadata = ray.get(\n trainer.execute_checkpoint_engine([\"prepare\"] * trainer.world_size)\n + rollout.execute_checkpoint_engine([\"prepare\"] * rollout.world_size)\n )\n\n # 2. build communication topology between all workers\n trainer_kwargs, rollout_kwargs = self.backend_cls.build_topology(\n trainer.world_size, rollout.world_size, metadata\n )\n for k, v in trainer_kwargs.items():\n assert len(v) == trainer.world_size, f\"trainer_kwargs[{k}] must have length of {trainer.world_size}\"\n for k, v in rollout_kwargs.items():\n assert len(v) == rollout.world_size, f\"rollout_kwargs[{k}] must have length of {rollout.world_size}\"\n\n trainer_kwargs[\"method\"] = [\"init_process_group\"] * trainer.world_size\n rollout_kwargs[\"method\"] = [\"init_process_group\"] * rollout.world_size\n\n # 3. init process group between all workers\n ray.get(\n trainer.execute_checkpoint_engine(**trainer_kwargs) + rollout.execute_checkpoint_engine(**rollout_kwargs)\n )\n\n def add_replicas(self, replicas: list[RolloutReplica]):\n \"\"\"Add rollout replicas to the manager for elastic scale up, will rebuild process group.\n\n Args:\n replicas: The list of rollout replicas to add.\n \"\"\"\n self.replicas.extend(replicas)\n\n def remove_replicas(self, replicas: list[RolloutReplica]):\n \"\"\"Remove rollout replicas from the manager for elastic scale down, will rebuild process group.\n\n Args:\n replicas: The list of rollout replicas to remove.\n \"\"\"\n replicas_set = set(replicas)\n self.replicas = [r for r in self.replicas if r not in replicas_set]\n\n @auto_await\n async def sleep_replicas(self):\n \"\"\"Sleep all rollout replicas: free weight and kv_cache device memory.\"\"\"\n # skip sleep replicas for disaggregated rollout\n if self.backend != \"naive\":\n return\n await asyncio.gather(*[r.sleep() for r in self.replicas])\n\n @auto_await\n async def update_weights(self):\n \"\"\"Update weights from trainer to rollout replicas.\"\"\"\n\n # 0. update weights for sync training with colocated trainer and rollout\n if self.backend == \"naive\":\n ray.get(self.trainer.update_weights())\n return\n\n # 1. abort and save all unfinished requests for partial rollout\n await asyncio.gather(*[r.abort_all_requests() for r in self.replicas])\n\n # 2. create a temporay worker group for all replicas\n workers = []\n for replica in self.replicas:\n workers.extend(replica.workers)\n rollout = RayWorkerGroup(worker_handles=workers, ray_cls_with_init=RayClassWithInitArgs(cls=_worker_cls))\n trainer = self.trainer\n\n # 3. build process group\n self.build_process_group(rollout)\n\n # 4. update weights of all workers\n ray.get(trainer.update_weights() + rollout.update_weights())\n\n # 5. finalize all workers\n ray.get(\n trainer.execute_checkpoint_engine([\"finalize\"] * trainer.world_size)\n + rollout.execute_checkpoint_engine([\"finalize\"] * rollout.world_size)\n )\n\n # 6. resume all unfinished requests for partial rollout\n await asyncio.gather(*[r.resume_all_requests() for r in self.replicas])\n"}3{"file_name": "verl__checkpoint_engine__nixl_checkpoint_engine.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport asyncio\nimport logging\nimport os\nimport time\nimport uuid\nfrom collections import defaultdict, deque\nfrom dataclasses import dataclass\nfrom typing import AsyncGenerator, Generator\nfrom unittest.mock import patch\n\nwith patch(\"importlib.metadata.distributions\", return_value=[]):\n import cupy as cp\n\nimport nixl._api as nixl_api\nimport nixl._bindings as nixl_bindings\nimport ray\nimport torch\nimport zmq\nimport zmq.asyncio\n\nfrom verl.checkpoint_engine.base import CheckpointEngine, CheckpointEngineRegistry, TensorMeta\nfrom verl.utils.net_utils import get_free_port, is_valid_ipv6_address\n\nlogger = logging.getLogger(__name__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\n@dataclass\nclass NixlAgentMetadata:\n agent_name: str\n agent_metadata: bytes\n zmq_ip: str\n zmq_port: int\n\n\nclass NixlAgent:\n \"\"\"This is a wrapper class for nixl_agent, the main purpose is to use ZeroMQ instead of\n `nixl_agent.send_notif` to send bucket tensor metadata.\n \"\"\"\n\n def __init__(self):\n self.agent_name = str(uuid.uuid4())\n self.agent = nixl_api.nixl_agent(self.agent_name)\n self.notifications: dict[str, deque[bytes]] = defaultdict(deque)\n\n self.start_zmq_server()\n self.zmq_clients: dict[str, zmq.Socket] = {}\n self.messages: dict[str, deque[bytes]] = defaultdict(deque)\n\n def __getattr__(self, name):\n attr = getattr(self.agent, name)\n\n if callable(attr):\n\n def wrapper(*args, **kwargs):\n return attr(*args, **kwargs)\n\n return wrapper\n else:\n return attr\n\n def get_agent_metadata(self) -> NixlAgentMetadata:\n return NixlAgentMetadata(\n agent_name=self.agent_name,\n agent_metadata=self.agent.get_agent_metadata(),\n zmq_ip=self.ip,\n zmq_port=self.listen_port,\n )\n\n def start_zmq_server(self):\n self.ip = ray.util.get_node_ip_address().strip(\"[]\")\n self.listen_port, self.listen_sock = get_free_port(self.ip)\n\n context = zmq.asyncio.Context()\n self.socket = context.socket(zmq.PULL)\n if is_valid_ipv6_address(self.ip):\n address = f\"tcp://[{self.ip}]:{self.listen_port}\"\n self.socket.setsockopt(zmq.IPV6, 1)\n else:\n address = f\"tcp://{self.ip}:{self.listen_port}\"\n\n self.socket.bind(address)\n\n def add_remote_agent(self, metadata: NixlAgentMetadata) -> str:\n agent_name = self.agent.add_remote_agent(metadata.agent_metadata).decode(\"utf-8\")\n assert agent_name == metadata.agent_name, f\"Agent name {agent_name} not equal to {metadata.agent_name}\"\n\n context = zmq.Context()\n socket = context.socket(zmq.PUSH)\n if is_valid_ipv6_address(metadata.zmq_ip):\n address = f\"tcp://[{metadata.zmq_ip}]:{metadata.zmq_port}\"\n socket.setsockopt(zmq.IPV6, 1)\n else:\n address = f\"tcp://{metadata.zmq_ip}:{metadata.zmq_port}\"\n\n socket.connect(address)\n self.zmq_clients[agent_name] = socket\n return agent_name\n\n def remove_remote_agent(self, agent_name: str):\n self.agent.remove_remote_agent(agent_name)\n socket = self.zmq_clients.pop(agent_name)\n socket.close()\n\n def send_message(self, agent_name, message: dict):\n socket = self.zmq_clients[agent_name]\n socket.send_pyobj((self.agent_name, message), zmq.DONTWAIT)\n\n async def read_message(self, agent_name: str) -> dict:\n while len(self.messages[agent_name]) == 0:\n recv_agent_name, message = await self.socket.recv_pyobj()\n self.messages[recv_agent_name].append(message)\n return self.messages[agent_name].popleft()\n\n async def get_notification(self, remote_name: str) -> bytes:\n while len(self.notifications[remote_name]) == 0:\n notifs = self.agent.get_new_notifs()\n for remote_name, notif in notifs.items():\n self.notifications[remote_name].extend(notif)\n await asyncio.sleep(0)\n return self.notifications[remote_name].popleft()\n\n\nclass ReadableOperation:\n \"\"\"Encapsulates a readable operation to remote agent.\n 1. send metadata to remote agent\n 2. wait until remote agent read complete.\n\n Args:\n agent (NixlAgent): The Nixl agent.\n remote_agent (str): The name of the remote agent.\n local_descs (nixl_bindings.nixlXferDList): The local transfer descriptors.\n metadata (dict): Metadata for the read operation.\n bucket_size (int): The size of the bucket in bytes.\n \"\"\"\n\n def __init__(\n self,\n agent: NixlAgent,\n remote_agent: str,\n local_descs: nixl_bindings.nixlXferDList,\n metadata: dict,\n ):\n self.agent = agent\n self.remote_agent = remote_agent\n self.local_descs = local_descs\n self.notify_key = uuid.uuid4().bytes\n message = {\"notify_key\": self.notify_key, \"remote_descs\": self.local_descs, **metadata}\n self.agent.send_message(self.remote_agent, message)\n\n async def wait_for_complete(self):\n \"\"\"Block until remote agent read complete.\"\"\"\n notification = await self.agent.get_notification(self.remote_agent)\n assert self.notify_key == notification, f\"Notify key {self.notify_key} not equal to {notification}\"\n logger.debug(f\"ReadableOperation to {self.remote_agent} complete\")\n\n\nclass ReadOperation:\n \"\"\"Encapsulates a read operation from remote agent.\n 1. read medata from remote agent\n 2. start read transfer operation\n 3. wait until read complete\n\n Args:\n agent (NixlAgent): The Nixl agent.\n remote_agent (str): The name of the remote agent.\n local_descs (nixl_bindings.nixlXferDList): The local transfer descriptors.\n bucket_size (int): The size of the bucket in bytes.\n \"\"\"\n\n def __init__(self, agent: NixlAgent, remote_agent: str, local_descs: nixl_bindings.nixlXferDList, bucket_size: int):\n self.agent = agent\n self.remote_agent = remote_agent\n self.local_descs = local_descs\n self.remote_descs = None\n self.xfer_handle = None\n self.notify_key = None\n self.bucket_size = bucket_size\n self.start_time = None\n\n async def read_metadata(self) -> dict:\n \"\"\"Block until the remote agent sends the metadata.\n\n Returns:\n dict: Metadata from the remote agent.\n \"\"\"\n metadata = await self.agent.read_message(self.remote_agent)\n self.remote_descs = metadata.pop(\"remote_descs\")\n self.notify_key = metadata.pop(\"notify_key\")\n return metadata\n\n def begin_read(self):\n \"\"\"Start the read operation.\"\"\"\n assert self.remote_descs is not None and self.notify_key is not None\n self.xfer_handle = self.agent.initialize_xfer(\n \"READ\", self.local_descs, self.remote_descs, self.remote_agent, self.notify_key\n )\n state = self.agent.transfer(self.xfer_handle)\n assert state != \"ERR\", f\"Read from {self.remote_agent} got to {state} state.\"\n self.start_time = time.time()\n\n async def wait_for_complete(self):\n \"\"\"Block until the read operation complete.\"\"\"\n while True:\n state = self.agent.check_xfer_state(self.xfer_handle)\n if state == \"ERR\":\n logger.error(f\"Read from {self.remote_agent} got to {state} state.\")\n exit(-1)\n elif state == \"DONE\":\n break\n else:\n await asyncio.sleep(0)\n self.agent.release_xfer_handle(self.xfer_handle)\n end_time = time.time()\n bandwidth = self.bucket_size / (end_time - self.start_time) / (1024 * 1024 * 1024)\n logger.debug(f\"ReadOperation read data from {self.remote_agent} complete, bandwidth: {bandwidth:.2f} GB/s\")\n\n\n@CheckpointEngineRegistry.register(\"nixl\")\nclass NIXLCheckpointEngine(CheckpointEngine):\n \"\"\"NIXL checkpoint engine with p2p communication, support various backends: ucx, uccl, mooncacke, etc.\n\n For UCX backend, some environment variables need to be set: UCX_TLS, UCX_IB_GID_INDEX, UCX_IB_DEVICES, etc.\n Please refer to: https://openucx.readthedocs.io/en/master/faq.html\n\n Args:\n bucket_size (int): Bucket size in bytes to transfer multiple weights at one time. Note that we use\n two buffer to send and recv weights at same time, so the device memory overhead is 2 * bucket_size.\n device (str): The device to use for the checkpoint engine, \"cpu\" or \"cuda\".\n rollout_dtype (torch.dtype): The dtype of the weights received from rollout workers. Defaults to torch.bfloat16.\n \"\"\"\n\n def __init__(\n self,\n bucket_size: int,\n device: str = \"cuda\",\n rollout_dtype: torch.dtype = torch.bfloat16,\n is_master: bool = False,\n ):\n self.bucket_size = bucket_size\n self.device = device\n self.rollout_dtype = rollout_dtype\n self.agent = NixlAgent()\n self.is_master = is_master\n\n def prepare(self) -> NixlAgentMetadata:\n \"\"\"Prepare send and recv bucket.\n\n Returns:\n NixlAgentMetadata: The metadata of the current nixl agent.\n \"\"\"\n # For master process, use cupy instead of torch to avoid memory register error\n # when `PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True`.\n if self.device == \"cuda\":\n send_buf = cp.zeros(self.bucket_size, dtype=cp.uint8)\n recv_buf = cp.zeros(self.bucket_size, dtype=cp.uint8)\n self.send_buf = torch.as_tensor(send_buf, dtype=torch.uint8)\n self.recv_buf = torch.as_tensor(recv_buf, dtype=torch.uint8)\n else:\n self.send_buf = torch.zeros(self.bucket_size, dtype=torch.uint8, device=self.device, pin_memory=True)\n self.recv_buf = torch.zeros(self.bucket_size, dtype=torch.uint8, device=self.device, pin_memory=True)\n self.send_reg_descs = self.agent.register_memory(self.send_buf)\n self.recv_reg_descs = self.agent.register_memory(self.recv_buf)\n self.send_descs = self.agent.get_xfer_descs(self.send_buf)\n self.recv_descs = self.agent.get_xfer_descs(self.recv_buf)\n\n return self.agent.get_agent_metadata()\n\n @classmethod\n def build_topology(cls, trainer_world_size: int, rollout_world_size: int, metadata: list[dict]):\n trainer_kwargs = {\n \"method\": [\"init_process_group\"] * trainer_world_size,\n \"rank\": [0] + [-1] * (trainer_world_size - 1),\n \"world_size\": [rollout_world_size + 1] * trainer_world_size,\n \"prev_agent_metadata\": [None] * trainer_world_size,\n \"next_agent_metadata\": [metadata[-rollout_world_size]] + [None] * (trainer_world_size - 1),\n }\n\n rollout_kwargs = {\n \"method\": [\"init_process_group\"] * rollout_world_size,\n \"rank\": list(range(1, rollout_world_size + 1)),\n \"world_size\": [rollout_world_size + 1] * rollout_world_size,\n \"prev_agent_metadata\": [metadata[0]] + metadata[-rollout_world_size:-1],\n \"next_agent_metadata\": metadata[-rollout_world_size + 1 :] + [None],\n }\n return trainer_kwargs, rollout_kwargs\n\n def init_process_group(\n self, rank: int, world_size: int, prev_agent_metadata: NixlAgentMetadata, next_agent_metadata: NixlAgentMetadata\n ):\n \"\"\"Setup the communication with the previous and next agent.\n\n Args:\n rank (int): The rank of the current process.\n world_size (int): The total number of processes.\n prev_agent_metadata (NixlAgentMetadata): The metadata of the previous nixl agent.\n next_agent_metadata (NixlAgentMetadata): The metadata of the next nixl agent.\n \"\"\"\n if rank < 0:\n assert not prev_agent_metadata and not next_agent_metadata, (\n f\"rank {rank} should not have prev_agent_metadata or next_agent_metadata\"\n )\n elif rank == 0:\n assert not prev_agent_metadata and next_agent_metadata, f\"rank {rank} should have next_agent_metadata\"\n elif 0 < rank < world_size - 1:\n assert prev_agent_metadata and next_agent_metadata, (\n f\"rank {rank} should have prev_agent_metadata and next_agent_metadata\"\n )\n elif rank == world_size - 1:\n assert prev_agent_metadata and not next_agent_metadata, (\n f\"rank {rank} should have prev_agent_metadata and not next_agent_metadata\"\n )\n\n self.rank = rank\n self.world_size = world_size\n self.prev_agent = None\n self.next_agent = None\n\n if prev_agent_metadata is not None:\n self.prev_agent = self.agent.add_remote_agent(prev_agent_metadata)\n\n if next_agent_metadata is not None:\n self.next_agent = self.agent.add_remote_agent(next_agent_metadata)\n\n logger.info(\n f\"init_process_group rank: {self.rank}, world_size: {self.world_size}, \"\n f\"prev_agent: {self.prev_agent}, next_agent: {self.next_agent}\"\n )\n\n def finalize(self):\n \"\"\"Cleanup communication with the previous and next agent, and deregister the memory.\"\"\"\n if self.prev_agent:\n self.agent.remove_remote_agent(self.prev_agent)\n if self.next_agent:\n self.agent.remove_remote_agent(self.next_agent)\n\n self.agent.deregister_memory(self.send_reg_descs)\n self.agent.deregister_memory(self.recv_reg_descs)\n self.send_buf = None\n self.recv_buf = None\n self.send_reg_descs = None\n self.recv_reg_descs = None\n self.send_descs = None\n self.recv_descs = None\n\n self.rank = None\n self.world_size = None\n self.prev_agent = None\n self.next_agent = None\n\n @torch.no_grad()\n async def send_weights(self, weights: Generator[tuple[str, torch.Tensor], None, None]):\n \"\"\"Send the weights of the model.\n\n Args:\n weights: A generator that yields the name of the weight tensor and the tensor itself.\n \"\"\"\n assert self.rank <= 0, \"Trainer workers other than rank 0 should not send weights.\"\n\n # For trainer workers other than rank 0, just consume weights and do nothing.\n if self.rank < 0:\n for name, weight in weights:\n pass\n return\n\n assert self.next_agent is not None, \"Next agent is not set.\"\n send_buf, recv_buf = self.send_buf, self.recv_buf\n send_descs, recv_descs = self.send_descs, self.recv_descs\n readable_op = None\n\n start_time = time.time()\n bucket_meta: dict[str, TensorMeta] = {}\n offset = 0\n for name, weight in weights:\n # fill the tensor bucket\n if offset + weight.nbytes > self.bucket_size:\n torch.cuda.synchronize()\n\n # wait previous bucket to be received\n if readable_op is not None:\n await readable_op.wait_for_complete()\n\n # send bucket meta to next agent\n readable_op = ReadableOperation(\n self.agent,\n self.next_agent,\n send_descs,\n {\"bucket_meta\": bucket_meta, \"is_last\": False},\n )\n\n # swap send and recv buf\n send_buf, recv_buf = recv_buf, send_buf\n send_descs, recv_descs = recv_descs, send_descs\n bucket_meta = {}\n offset = 0\n\n assert offset + weight.nbytes <= self.bucket_size, (\n f\"Weight {name}({weight.shape}, {weight.dtype}) is too large to fit in the bucket.\"\n )\n\n bucket_meta[name] = {\n \"name\": name,\n \"shape\": weight.shape,\n \"dtype\": weight.dtype,\n \"offset\": offset,\n }\n send_buf[offset : offset + weight.nbytes].copy_(weight.view(-1).view(torch.uint8), non_blocking=True)\n offset += weight.nbytes\n\n # send last bucket meta to next agent\n torch.cuda.synchronize()\n if readable_op is not None:\n await readable_op.wait_for_complete()\n\n readable_op = ReadableOperation(\n self.agent, self.next_agent, send_descs, {\"bucket_meta\": bucket_meta, \"is_last\": True}\n )\n await readable_op.wait_for_complete()\n logger.info(f\"Rank {self.rank} send weights done, time cost: {time.time() - start_time:.2f}s\")\n\n @torch.no_grad()\n async def receive_weights(self) -> AsyncGenerator[tuple[str, torch.Tensor], None]:\n \"\"\"Receive the weights of the model.\n\n Yields:\n A tuple of the name of the weight tensor and the tensor itself.\n \"\"\"\n assert self.prev_agent is not None, \"Previous agent is not set.\"\n send_buf, recv_buf = self.send_buf, self.recv_buf\n send_descs, recv_descs = self.send_descs, self.recv_descs\n total_bytes, total_params = 0, 0\n\n # receive first bucket from previous agent\n start_time = time.time()\n read_op = ReadOperation(self.agent, self.prev_agent, recv_descs, self.bucket_size)\n metadata = await read_op.read_metadata()\n read_op.begin_read()\n await read_op.wait_for_complete()\n total_bytes += self.bucket_size\n total_params += len(metadata[\"bucket_meta\"])\n\n # swap send and recv buf\n send_buf, recv_buf = recv_buf, send_buf\n send_descs, recv_descs = recv_descs, send_descs\n while not metadata[\"is_last\"]:\n # 1. send bucket to next agent\n readable_op = None\n if self.next_agent is not None:\n readable_op = ReadableOperation(\n self.agent,\n self.next_agent,\n send_descs,\n metadata,\n )\n\n # 2. receive bucket from previous agent\n read_op = ReadOperation(self.agent, self.prev_agent, recv_descs, self.bucket_size)\n next_metadata = await read_op.read_metadata()\n read_op.begin_read()\n\n # 3. yield tensor from send_buf\n for name, meta in metadata[\"bucket_meta\"].items():\n dtype, shape = meta[\"dtype\"], meta[\"shape\"]\n size = dtype.itemsize * shape.numel()\n tensor = send_buf[meta[\"offset\"] : meta[\"offset\"] + size].view(dtype=dtype).view(shape)\n yield name, tensor\n\n # 4. wait for next agent read complete and read from previous agent complete\n if readable_op is not None:\n await readable_op.wait_for_complete()\n await read_op.wait_for_complete()\n total_bytes += self.bucket_size\n total_params += len(next_metadata[\"bucket_meta\"])\n\n # 5. swap send and recv buf\n torch.cuda.synchronize() # sync non-blocking copy\n metadata = next_metadata\n send_buf, recv_buf = recv_buf, send_buf\n send_descs, recv_descs = recv_descs, send_descs\n\n # send last bucket to next agent\n readable_op = None\n if self.next_agent is not None:\n readable_op = ReadableOperation(\n self.agent,\n self.next_agent,\n send_descs,\n metadata,\n )\n\n # yield tensor from send_buf\n for name, meta in metadata[\"bucket_meta\"].items():\n dtype, shape = meta[\"dtype\"], meta[\"shape\"]\n size = dtype.itemsize * shape.numel()\n tensor = send_buf[meta[\"offset\"] : meta[\"offset\"] + size].view(dtype=dtype).view(shape)\n yield name, tensor\n\n # wait for next agent read complete\n if readable_op is not None:\n await readable_op.wait_for_complete()\n time_cost = time.time() - start_time\n bandwidth = total_bytes / time_cost / (1024 * 1024 * 1024)\n logger.info(\n f\"Rank {self.rank} receive weights done, total_params: {total_params}, \"\n f\"time cost: {time_cost:.2f}s, bandwidth: {bandwidth:.2f} GB/s\"\n )\n"}4{"file_name": "verl__interactions__gsm8k_interaction.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nimport os\nfrom typing import Any, Optional\nfrom uuid import uuid4\n\nfrom verl.utils.reward_score import gsm8k\n\nfrom .base import BaseInteraction\n\nlogger = logging.getLogger(__name__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\nclass Gsm8kInteraction(BaseInteraction):\n \"\"\"A demo interaction for calculating the reward of gsm8k.\n\n - `start_interaction`: start a interaction instance for a trajectory.\n - `generate_response`: generate the response of the assistant.\n - `calculate_score`: calculate the score of the interaction.\n - `finalize_interaction`: finalize the interaction instance.\n \"\"\"\n\n def __init__(self, config: dict):\n super().__init__(config)\n self._instance_dict = {}\n\n async def start_interaction(\n self, instance_id: Optional[str] = None, ground_truth: Optional[str] = None, **kwargs\n ) -> str:\n if instance_id is None:\n instance_id = str(uuid4())\n self._instance_dict[instance_id] = {\n \"response\": \"\",\n \"ground_truth\": ground_truth,\n \"reward\": 0.0,\n }\n return instance_id\n\n async def generate_response(\n self, instance_id: str, messages: list[dict[str, Any]], **kwargs\n ) -> tuple[bool, str, float, dict]:\n content = \"\"\n for i in range(len(messages) - 1, -1, -1):\n item = messages[i]\n if item.get(\"role\") == \"assistant\":\n content = item.get(\"content\")\n break\n\n self._instance_dict[instance_id][\"response\"] = content\n\n reward = await self.calculate_score(instance_id)\n if reward == 1.0:\n response = \"Your response is correct!\"\n should_terminate_sequence = True\n else:\n response = \"Your response is incorrect! You need to reflect on your answer and try again.\"\n should_terminate_sequence = False\n\n return should_terminate_sequence, response, reward, {}\n\n async def calculate_score(self, instance_id: str, **kwargs) -> float:\n return gsm8k.compute_score(\n self._instance_dict[instance_id][\"response\"],\n self._instance_dict[instance_id][\"ground_truth\"],\n method=\"strict\",\n format_score=0.0,\n score=1.0,\n )\n\n async def finalize_interaction(self, instance_id: str, **kwargs) -> None:\n del self._instance_dict[instance_id]\n"}5{"file_name": "verl__interactions__utils__interaction_registry.py", "text": "# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport importlib.util\nimport logging\nimport os\nimport sys\n\nfrom omegaconf import OmegaConf\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef get_interaction_class(cls_name):\n \"\"\"Dynamically import and return the interaction class.\"\"\"\n module_name, class_name = cls_name.rsplit(\".\", 1)\n if module_name not in sys.modules:\n spec = importlib.util.find_spec(module_name)\n module = importlib.util.module_from_spec(spec)\n sys.modules[module_name] = module\n spec.loader.exec_module(module)\n else:\n module = sys.modules[module_name]\n\n interaction_cls = getattr(module, class_name)\n return interaction_cls\n\n\ndef initialize_interactions_from_config(interaction_config_file):\n \"\"\"Initialize interactions from configuration file.\n\n Args:\n interaction_config_file: Path to the interaction configuration file.\n\n Returns:\n dict: A dictionary mapping interaction names to BaseInteraction instances.\n \"\"\"\n interaction_config = OmegaConf.load(interaction_config_file)\n interaction_map = {}\n\n for interaction_item in interaction_config.interaction:\n cls_name = interaction_item.class_name\n interaction_cls = get_interaction_class(cls_name)\n\n # Extract config and name\n config = OmegaConf.to_container(interaction_item.config, resolve=True)\n\n # Get the interaction name - either from config or derive from class name\n name = interaction_item.get(\"name\", None)\n if name is None:\n # If no name is specified, use the class name as default\n class_simple_name = cls_name.split(\".\")[-1]\n # Remove \"Interaction\" suffix if present, otherwise use full class name\n if class_simple_name.endswith(\"Interaction\"):\n name = class_simple_name[:-11].lower() # Remove \"Interaction\" (11 chars)\n else:\n name = class_simple_name.lower()\n\n # Check for duplicate names\n if name in interaction_map:\n raise ValueError(f\"Duplicate interaction name '{name}' found. Each interaction must have a unique name.\")\n\n # Inject the name into the config\n config[\"name\"] = name\n\n # Create the interaction instance\n interaction = interaction_cls(config=config)\n interaction_map[name] = interaction\n\n logger.info(f\"Initialized interaction '{name}' with class '{cls_name}'\")\n\n return interaction_map\n"}6{"file_name": "verl__interactions__weather_interaction.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nimport os\nfrom typing import Any, Optional\nfrom uuid import uuid4\n\nfrom .base import BaseInteraction\n\nlogger = logging.getLogger(__name__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\nclass WeatherInteraction(BaseInteraction):\n \"\"\"A demo interaction for handling weather-related queries.\n\n - `start_interaction`: start a interaction instance for a trajectory.\n - `generate_response`: generate the response of the assistant.\n - `calculate_score`: calculate the score of the interaction.\n - `finalize_interaction`: finalize the interaction instance.\n \"\"\"\n\n def __init__(self, config: dict):\n super().__init__(config)\n self._instance_dict = {}\n\n async def start_interaction(\n self, instance_id: Optional[str] = None, ground_truth: Optional[str] = None, **kwargs\n ) -> str:\n if instance_id is None:\n instance_id = str(uuid4())\n self._instance_dict[instance_id] = {\n \"response\": \"\",\n \"ground_truth\": ground_truth,\n \"reward\": 0.0,\n }\n return instance_id\n\n async def generate_response(\n self, instance_id: str, messages: list[dict[str, Any]], **kwargs\n ) -> tuple[bool, str, float, dict]:\n content = \"no tool call\"\n for i in range(len(messages) - 1, -1, -1):\n item = messages[i]\n if item.get(\"role\") == \"tool\":\n content = item.get(\"content\")\n break\n self._instance_dict[instance_id][\"response\"] = content\n\n reward = await self.calculate_score(instance_id)\n if reward == 1.0:\n response = \"Thank you for your weather query!\"\n should_terminate_sequence = True\n else:\n response = \"Please use the weather tool to get the weather information.\"\n should_terminate_sequence = True\n return should_terminate_sequence, response, reward, {}\n\n async def calculate_score(self, instance_id: str, **kwargs) -> float:\n # For weather interaction, we can implement a more complex scoring logic\n # For now, we'll just return a default score of 1.0\n if self._instance_dict[instance_id][\"response\"] == \"no tool call\":\n return 0.0\n return 1.0\n\n async def finalize_interaction(self, instance_id: str, **kwargs) -> None:\n del self._instance_dict[instance_id]\n"}7{"file_name": "verl__model_merger__base_model_merger.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport argparse\nimport os\nfrom abc import ABC, abstractmethod\nfrom dataclasses import dataclass, field\nfrom typing import Optional\n\nimport torch\nfrom accelerate import init_empty_weights\nfrom transformers import (\n AutoConfig,\n AutoModelForCausalLM,\n AutoModelForTokenClassification,\n GenerationConfig,\n)\n\nfrom verl.utils import hf_processor, hf_tokenizer\n\n\ndef parse_args():\n parser = argparse.ArgumentParser(description=\"verl model merger\")\n subparsers = parser.add_subparsers(dest=\"operation\", required=True, help=\"Specify 'merge' or 'test' operation.\")\n\n base_op_parser = argparse.ArgumentParser(add_help=False)\n base_op_parser.add_argument(\n \"--backend\", type=str, required=True, choices=[\"fsdp\", \"megatron\"], help=\"The backend of the model\"\n )\n base_op_parser.add_argument(\"--local_dir\", type=str, default=None, help=\"Path to the saved model checkpoints.\")\n base_op_parser.add_argument(\n \"--tie-word-embedding\",\n action=\"store_true\",\n help=\"Whether to tie word embedding weights (currently only Megatron supported)\",\n )\n base_op_parser.add_argument(\"--trust-remote-code\", action=\"store_true\", help=\"Whether to trust remote code\")\n base_op_parser.add_argument(\n \"--is-value-model\",\n action=\"store_true\",\n help=\"Whether the model is a value model (currently only Megatron supported)\",\n )\n base_op_parser.add_argument(\n \"--use_cpu_initialization\",\n action=\"store_true\",\n help=\"Whether to use CPU initialization for the model. This is useful for large models that cannot \"\n \"fit into GPU memory during initialization.\",\n )\n\n merge_parser = subparsers.add_parser(\"merge\", parents=[base_op_parser], help=\"Merge model checkpoints and save.\")\n merge_parser.add_argument(\n \"--target_dir\", default=\"tmp\", type=str, help=\"Directory to save the merged huggingface model\"\n )\n merge_parser.add_argument(\n \"--hf_upload_path\", default=None, type=str, help=\"Hugging Face repository ID to upload the model\"\n )\n merge_parser.add_argument(\n \"--private\", action=\"store_true\", help=\"Whether to upload the model to a private Hugging Face repository\"\n )\n\n test_parser = subparsers.add_parser(\n \"test\", parents=[base_op_parser], help=\"Test merged model against a reference Hugging Face model\"\n )\n test_parser.add_argument(\n \"--test_hf_dir\", type=str, required=True, help=\"Path to the reference Hugging Face model directory for testing\"\n )\n\n args = parser.parse_args()\n return args\n\n\n@dataclass\nclass ModelMergerConfig:\n \"\"\"Configuration for model merger operations.\n\n Args:\n operation (str): Operation type - 'merge' or 'test'.\n backend (str): Backend type for the model ('fsdp' or 'megatron').\n target_dir (Optional[str]): Directory to save the merged huggingface model. Defaults to \"tmp\".\n hf_upload_path (Optional[str]): Hugging Face repository ID to upload the model. Defaults to None.\n private (bool): Whether to upload the model to a private Hugging Face repository. Defaults to False.\n test_hf_dir (Optional[str]): Path to the reference Hugging Face model directory for testing. Defaults to None.\n tie_word_embedding (bool): Whether to tie word embedding weights (currently only Megatron\n supported). Defaults to False.\n trust_remote_code (bool): Whether to trust remote code. Defaults to False.\n is_value_model (bool): Whether the model is a value model (currently only Megatron\n supported). Defaults to False.\n local_dir (Optional[str]): Path to the saved model checkpoints. Defaults to None.\n hf_model_config_path (Optional[str]): Path to HuggingFace model configuration files. Defaults to None.\n hf_upload (bool): Whether to upload to HuggingFace (computed automatically). Not for initialization.\n use_cpu_initialization (bool): Whether to use CPU initialization for large models. Defaults to False.\n \"\"\"\n\n operation: str # 'merge' or 'test'\n backend: str\n target_dir: Optional[str] = \"tmp\"\n hf_upload_path: Optional[str] = None\n private: bool = False\n test_hf_dir: Optional[str] = None\n tie_word_embedding: bool = False\n trust_remote_code: bool = False\n is_value_model: bool = False\n local_dir: Optional[str] = None\n hf_model_config_path: Optional[str] = None\n hf_upload: bool = field(init=False)\n use_cpu_initialization: bool = False\n\n def __post_init__(self):\n self.hf_upload = self.operation == \"merge\" and bool(self.hf_upload_path)\n if self.operation == \"test\":\n self.target_dir = None\n self.hf_upload_path = None\n self.private = False\n\n\ndef generate_config_from_args(args: argparse.Namespace) -> ModelMergerConfig:\n common_config_args = {\n \"operation\": args.operation,\n \"backend\": args.backend,\n \"tie_word_embedding\": args.tie_word_embedding,\n \"trust_remote_code\": args.trust_remote_code,\n \"is_value_model\": args.is_value_model,\n \"local_dir\": args.local_dir,\n \"hf_model_config_path\": os.path.join(args.local_dir, \"huggingface\"),\n \"use_cpu_initialization\": args.use_cpu_initialization,\n }\n\n if args.operation == \"merge\":\n config = ModelMergerConfig(\n **common_config_args,\n target_dir=args.target_dir,\n hf_upload_path=args.hf_upload_path,\n private=args.private,\n test_hf_dir=None,\n )\n os.makedirs(config.target_dir, exist_ok=True)\n elif args.operation == \"test\":\n config = ModelMergerConfig(\n **common_config_args,\n test_hf_dir=args.test_hf_dir,\n # the following args are not used by test operation\n target_dir=None,\n hf_upload_path=None,\n private=False,\n )\n else:\n raise NotImplementedError(f\"Unknown operation: {args.operation}\")\n return config\n\n\nclass BaseModelMerger(ABC):\n \"\"\"\n Abstract base class for merging distributed model checkpoints into HuggingFace format.\n\n This class provides common functionality for converting model checkpoints from different\n distributed training backends (FSDP, Megatron) into standard HuggingFace format that\n can be easily loaded and used for inference or further training.\n\n The merger supports two main operations:\n - merge: Convert and save checkpoints to HuggingFace format\n - test: Validate merged checkpoints against a reference model\n\n Args:\n config (ModelMergerConfig): Configuration object containing paths, backend type,\n and operation parameters.\n\n Attributes:\n config (ModelMergerConfig): The configuration object passed during initialization.\n hf_model_config_path (str): Path to the HuggingFace model configuration files.\n model_config (PretrainedConfig): Loaded HuggingFace model configuration.\n \"\"\"\n\n def __init__(self, config: ModelMergerConfig):\n self.config = config\n self.hf_model_config_path = config.hf_model_config_path\n self.model_config = AutoConfig.from_pretrained(\n self.hf_model_config_path, trust_remote_code=self.config.trust_remote_code\n )\n\n def get_transformers_auto_model_class(self):\n has_remote_code = hasattr(self.model_config, \"auto_map\") and any(\n self.model_config.architectures[0] in val for val in self.model_config.auto_map.values()\n )\n if has_remote_code:\n auto_class = next(\n k for k, v in self.model_config.auto_map.items() if self.model_config.architectures[0] in v\n )\n match auto_class:\n case \"AutoModelForCausalLM\":\n return AutoModelForCausalLM\n case \"AutoModelForTokenClassification\":\n return AutoModelForTokenClassification\n case \"AutoModelForVision2Seq\":\n # Handle different transformers versions for Vision2Seq models\n import transformers\n from packaging import version\n\n if version.parse(transformers.__version__) >= version.parse(\"4.54.0\"):\n # transformers >= 4.54.0 uses AutoModelForImageTextToText\n from transformers import AutoModelForImageTextToText\n\n return AutoModelForImageTextToText\n else:\n # transformers < 4.54.0 uses AutoModelForVision2Seq\n from transformers import AutoModelForVision2Seq\n\n return AutoModelForVision2Seq\n case _:\n raise NotImplementedError(f\"Unknown auto class {auto_class}\")\n else:\n if \"ForTokenClassification\" in self.model_config.architectures[0]:\n return AutoModelForTokenClassification\n elif \"ForCausalLM\" in self.model_config.architectures[0]:\n return AutoModelForCausalLM\n elif \"ForConditionalGeneration\" in self.model_config.architectures[0]:\n return AutoModelForVision2Seq\n\n raise NotImplementedError(f\"Unknown architecture {self.model_config.architectures}\")\n\n def patch_model_generation_config(self, model):\n \"\"\"\n The generation_config created from model config may be different to the pretrained model,\n this may lead to error when generating: https://github.com/volcengine/verl/issues/1246\n\n This function patch the generation_config created from model config to the pretrained model.\n \"\"\"\n if model.can_generate():\n try:\n model.generation_config = GenerationConfig.from_pretrained(self.hf_model_config_path)\n except OSError:\n print(\n f\"Warning: Generation config file not found in {self.hf_model_config_path}, using a \"\n f\"generation config created from the model config.\"\n )\n return model\n\n def save_lora_adapter(self, state_dict: dict[str, torch.Tensor]):\n \"\"\"\n Save lora adapter to safetensors.\n\n Returns:\n lora_path: str, the path to the lora adapter. None if no lora adapter found.\n\n Note:\n This function change the 'state_dict' in place.\n \"\"\"\n lora_params_names = [name for name in state_dict.keys() if \"lora_\" in name]\n\n if len(lora_params_names) == 0:\n return None\n\n import json\n from typing import OrderedDict\n\n import peft\n from safetensors.torch import save_file\n\n lora_params = OrderedDict()\n target_modules = set()\n lora_key = None\n\n for name in lora_params_names:\n lora_key = name.replace(\".default.weight\", \".weight\")\n target_modules.add(lora_key.split(\".\")[-3])\n lora_params[lora_key] = state_dict.pop(name)\n\n lora_rank = min(lora_params[lora_key].shape[0], lora_params[lora_key].shape[1])\n peft_dict = {\n \"r\": lora_rank,\n \"lora_alpha\": 0, # lora_alpha is not set. An error should be raised to inform the user to set it manually.\n \"target_modules\": list(target_modules),\n }\n peft_config = peft.LoraConfig(**peft_dict).to_dict()\n peft_config[\"task_type\"] = peft_config[\"task_type\"].value if peft_config[\"task_type\"] else None\n peft_config[\"peft_type\"] = peft_config[\"peft_type\"].value if peft_config[\"peft_type\"] else None\n peft_config[\"target_modules\"] = list(peft_config[\"target_modules\"])\n\n lora_path = os.path.join(self.config.target_dir, \"lora_adapter\")\n os.makedirs(lora_path, exist_ok=True)\n with open(os.path.join(lora_path, \"adapter_config.json\"), \"w\", encoding=\"utf-8\") as f:\n json.dump(peft_config, f, ensure_ascii=False, indent=4)\n save_file(lora_params, os.path.join(lora_path, \"adapter_model.safetensors\"))\n\n for name in list(state_dict.keys()):\n key = (\n name.replace(\"base_model.model.\", \"\")\n .replace(\".base_layer.weight\", \".weight\")\n .replace(\".base_layer.bias\", \".bias\")\n )\n state_dict[key] = state_dict.pop(name)\n\n return lora_path\n\n def save_hf_model_and_tokenizer(self, state_dict: dict[str, torch.Tensor]):\n auto_model_class = self.get_transformers_auto_model_class()\n with init_empty_weights():\n model = auto_model_class.from_config(\n self.model_config, torch_dtype=torch.bfloat16, trust_remote_code=self.config.trust_remote_code\n )\n model.to_empty(device=\"cpu\")\n model = self.patch_model_generation_config(model)\n\n lora_path = self.save_lora_adapter(state_dict)\n if lora_path:\n print(f\"Saving lora adapter to {lora_path}\")\n\n print(f\"Saving model to {self.config.target_dir}\")\n model.save_pretrained(self.config.target_dir, state_dict=state_dict)\n del state_dict\n del model\n\n processor = hf_processor(self.hf_model_config_path, trust_remote_code=self.config.trust_remote_code)\n tokenizer = hf_tokenizer(self.hf_model_config_path, trust_remote_code=self.config.trust_remote_code)\n if processor is not None:\n print(f\"Saving processor to {self.config.target_dir}\")\n processor.save_pretrained(self.config.target_dir)\n if tokenizer is not None:\n print(f\"Saving tokenizer to {self.config.target_dir}\")\n tokenizer.save_pretrained(self.config.target_dir)\n\n def upload_to_huggingface(self):\n import requests\n from huggingface_hub import HfApi\n from huggingface_hub.utils import HfHubHTTPError, RepositoryNotFoundError\n\n api = HfApi()\n try:\n # Attempt to create repository\n api.create_repo(repo_id=self.config.hf_upload_path, private=self.config.private, exist_ok=True)\n except HfHubHTTPError as e:\n # Handle authentication/API errors\n if e.response.status_code == 401:\n raise PermissionError(\n \"Hugging Face authentication failed. Verify your token is valid and has write permissions.\"\n ) from e\n elif e.response.status_code == 404:\n raise RepositoryNotFoundError(f\"Repository path not found: {self.config.hf_upload_path}\") from e\n else:\n raise ConnectionError(f\"Failed to create repository ({e.response.status_code}): {e}\") from e\n except requests.exceptions.ConnectionError as e:\n raise ConnectionError(\"Network connection failed. Check your internet connection.\") from e\n\n try:\n # Attempt folder upload\n api.upload_folder(folder_path=self.config.target_dir, repo_id=self.config.hf_upload_path, repo_type=\"model\")\n except HfHubHTTPError as e:\n if e.response.status_code == 401:\n raise PermissionError(\"Authentication failed during upload. Token may have expired.\") from e\n else:\n raise RuntimeError(f\"Upload failed ({e.response.status_code}): {e}\") from e\n except requests.exceptions.ConnectionError as e:\n raise ConnectionError(\"Network interruption during upload. Try again with stable connection.\") from e\n except OSError as e:\n raise FileNotFoundError(f\"Local folder error: {self.config.target_dir} - {str(e)}\") from e\n except Exception as e:\n raise RuntimeError(f\"Unexpected error during upload: {str(e)}\") from e\n\n @abstractmethod\n def merge_and_save(self):\n raise NotImplementedError(\"Subclasses should implement this method\")\n\n @abstractmethod\n def cleanup(self):\n raise NotImplementedError(\"Subclasses should implement this method to clean up resources if needed\")\n"}8{"file_name": "verl__model_merger__fsdp_model_merger.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport json\nimport os\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\n\nimport numpy as np\nimport torch\nfrom torch.distributed._tensor import Placement, Shard\n\ntry:\n # for torch 2.5+\n from torch.distributed.tensor import DTensor\nexcept ImportError:\n from torch.distributed._tensor import DTensor\n\nfrom tqdm import tqdm\n\nfrom .base_model_merger import BaseModelMerger\n\n\nclass FSDPModelMerger(BaseModelMerger):\n \"\"\"\n Model merger for FSDP (Fully Sharded Data Parallel) checkpoints.\n\n This class handles the conversion of FSDP distributed checkpoints into HuggingFace format.\n FSDP shards model parameters across multiple processes, and this merger reconstructs\n the full model by loading and concatenating the sharded parameters from all ranks.\n\n The merger supports various FSDP configurations including:\n - Pure FSDP (single dimension sharding)\n - FSDP + DDP (data parallel + fully sharded data parallel)\n - DTensor-based sharding with custom device meshes\n\n Key features:\n - Automatic detection of world size from checkpoint filenames\n - Support for DTensor and non-DTensor checkpoints\n - Parallel loading of checkpoint shards for efficiency\n - Validation against reference HuggingFace models\n\n Example:\n To merge FSDP checkpoints:\n ```python\n config = ModelMergerConfig(\n operation=\"merge\",\n backend=\"fsdp\",\n local_dir=\"path/to/fsdp/checkpoints\",\n target_dir=\"path/to/output\"\n )\n merger = FSDPModelMerger(config)\n merger.merge_and_save()\n ```\n \"\"\"\n\n def _get_world_size(self) -> int:\n \"\"\"_summary_\n From FSDP json config file, extract the world size.\n\n Returns:\n int: world size\n \"\"\"\n config_path = Path(self.config.local_dir) / \"fsdp_config.json\"\n if not config_path.exists():\n raise FileNotFoundError(f\"Config file {config_path} does not exist.\")\n\n with open(config_path) as f:\n config = json.load(f)\n\n # Extract world size from the config\n world_size = config.get(\"world_size\", None)\n if world_size is None:\n raise ValueError(\"World size not found in the config file.\")\n\n return world_size\n\n def _load_rank_zero_state_dict(self, world_size: int) -> dict:\n return torch.load(\n Path(self.config.local_dir) / f\"model_world_size_{world_size}_rank_0.pt\",\n map_location=\"cpu\",\n weights_only=False,\n )\n\n def _extract_device_mesh_info(self, state_dict: dict, world_size: int) -> tuple[np.ndarray, tuple[str, ...]]:\n \"\"\"\n Retrieves sharding information (device_mesh, mesh_dim_names) from a DTensor in the state_dict.\n If no DTensor is found, infers a simple FSDP mesh based on world_size.\n \"\"\"\n pivot_key = sorted(list(state_dict.keys()))[0]\n weight = state_dict[pivot_key]\n\n if isinstance(weight, DTensor):\n # get sharding info\n device_mesh = weight.device_mesh\n mesh = device_mesh.mesh\n mesh_dim_names = device_mesh.mesh_dim_names\n else:\n # for non-DTensor\n mesh = np.array([world_size], dtype=np.int64)\n mesh_dim_names = (\"fsdp\",)\n\n return mesh, mesh_dim_names\n\n def _calculate_shard_configuration(\n self, mesh: np.ndarray, mesh_dim_names: tuple[str, ...]\n ) -> tuple[int, tuple[int, ...]]:\n \"\"\"Calculates the total number of shards and the shape of the device mesh.\"\"\"\n assert mesh_dim_names in ((\"fsdp\",), (\"ddp\", \"fsdp\")), f\"Unsupported mesh_dim_names {mesh_dim_names}\"\n\n if \"tp\" in mesh_dim_names:\n # TODO: \"tp\" is not supported yet due to the above assert\n total_shards = mesh.shape[-1] * mesh.shape[-2]\n mesh_shape = (mesh.shape[-2], mesh.shape[-1])\n else:\n total_shards = mesh.shape[-1]\n mesh_shape = (mesh.shape[-1],)\n\n return total_shards, mesh_shape\n\n def _merge_by_placement(self, tensors: list[torch.Tensor], placement: Placement) -> torch.Tensor:\n \"\"\"Merges a list of tensors based on their DTensor placement\"\"\"\n if placement.is_replicate():\n return tensors[0]\n elif placement.is_partial():\n raise NotImplementedError(\"Partial placement is not supported yet\")\n elif placement.is_shard():\n return torch.cat(tensors, dim=placement.dim).contiguous()\n\n raise NotImplementedError(f\"Unsupported placement: {placement}\")\n\n def _load_and_merge_state_dicts(\n self, world_size: int, total_shards: int, mesh_shape: tuple[int, ...], mesh_dim_names: tuple[str, ...]\n ) -> dict[str, torch.Tensor]:\n model_state_dict_lst = [None] * total_shards\n\n def process_one_shard(rank: int, model_state_dict_lst: list):\n model_path = Path(self.config.local_dir) / f\"model_world_size_{world_size}_rank_{rank}.pt\"\n state_dict = torch.load(model_path, map_location=\"cpu\", weights_only=False)\n model_state_dict_lst[rank] = state_dict\n return state_dict\n\n with ThreadPoolExecutor(max_workers=min(32, os.cpu_count())) as executor:\n futures = [executor.submit(process_one_shard, rank, model_state_dict_lst) for rank in range(total_shards)]\n for future in tqdm(futures, desc=f\"Loading {total_shards} FSDP shards\", total=total_shards):\n future.result()\n\n # Merge state dicts from all shards\n state_dict = {}\n param_placements: dict[str, list] = {}\n\n for key in set(model_state_dict_lst[0].keys()):\n state_dict[key] = []\n for model_state_shard in model_state_dict_lst:\n # add tensor shard in order of rank to state_dict[key]\n tensor = model_state_shard.pop(key)\n if isinstance(tensor, DTensor):\n state_dict[key].append(tensor._local_tensor.bfloat16())\n\n placements = tuple(tensor.placements)\n # replicated placement at dp dimension can be discarded\n if mesh_dim_names[0] in (\"dp\", \"ddp\"):\n placements = placements[1:]\n\n if key not in param_placements:\n param_placements[key] = placements\n else:\n assert param_placements[key] == placements\n else:\n state_dict[key].append(tensor.bfloat16())\n\n del model_state_dict_lst\n\n # Merge tensors\n for key in sorted(state_dict):\n if not isinstance(state_dict[key], list):\n print(f\"No need to merge key {key}\")\n continue\n if key in param_placements:\n # merge shards\n placements: tuple[Shard] = param_placements[key]\n if len(mesh_shape) == 1:\n # 1-D list, FSDP without TP\n assert len(placements) == 1\n shards = state_dict[key]\n state_dict[key] = self._merge_by_placement(shards, placements[0])\n else:\n # 2-D list, FSDP + TP\n raise NotImplementedError(\"FSDP + TP is not supported yet\")\n else:\n state_dict[key] = torch.cat(state_dict[key], dim=0)\n\n return state_dict\n\n def merge_and_save(self):\n world_size = self._get_world_size()\n rank_zero_state_dict = self._load_rank_zero_state_dict(world_size)\n\n mesh, mesh_dim_names = self._extract_device_mesh_info(rank_zero_state_dict, world_size)\n print(f\"Got device mesh {mesh}, mesh_dim_names {mesh_dim_names}\")\n\n total_shards, mesh_shape = self._calculate_shard_configuration(mesh, mesh_dim_names)\n print(f\"Processing model shards with {total_shards} {mesh_shape} in total\")\n\n merged_state_dict = self._load_and_merge_state_dicts(world_size, total_shards, mesh_shape, mesh_dim_names)\n\n if self.config.operation == \"test\":\n if not self.config.test_hf_dir:\n raise ValueError(\"test_hf_dir must be provided for test operation\")\n self._validate_state_dict(merged_state_dict)\n elif self.config.operation == \"merge\":\n self.save_hf_model_and_tokenizer(merged_state_dict)\n if self.config.hf_upload:\n self.upload_to_huggingface()\n else:\n raise ValueError(f\"Unknown operation: {self.config.operation}\")\n\n def _validate_state_dict(self, state_dict: dict[str, torch.Tensor]):\n auto_model_class = self.get_transformers_auto_model_class()\n\n hf_model = auto_model_class.from_pretrained(self.config.test_hf_dir, torch_dtype=torch.bfloat16)\n hf_state_dict = hf_model.state_dict()\n del hf_model\n\n hf_model_keys = set(hf_state_dict.keys())\n collected_keys = set(state_dict.keys())\n\n missing_keys = hf_model_keys - collected_keys\n assert len(missing_keys) == 0, f\"Missing keys in collected state dict: {list(sorted(missing_keys))}\"\n\n extra_keys = collected_keys - hf_model_keys\n assert len(extra_keys) == 0, f\"Extra keys in collected state dict: {list(sorted(extra_keys))}\"\n\n for key in hf_model_keys:\n hf_shape = hf_state_dict[key].shape\n collected_shape = state_dict[key].shape\n assert hf_shape == collected_shape, (\n f\"Shape mismatch for key '{key}': original {hf_shape} vs collected {collected_shape}\"\n )\n\n hf_dtype = hf_state_dict[key].dtype\n collected_dtype = state_dict[key].dtype\n assert hf_dtype == collected_dtype, (\n f\"Dtype mismatch for key '{key}': original {hf_dtype} vs collected {collected_dtype}\"\n )\n\n torch.testing.assert_close(hf_state_dict[key], state_dict[key], atol=1e-6, rtol=1e-6)\n\n print(\"FSDP checks passed: The merged state_dict matches the hf model saved by FSDPCheckpointManager.\")\n\n def cleanup(self):\n \"\"\"Cleanup temporary files if needed.\"\"\"\n # FSDP merger does not create temporary files, so no cleanup is needed.\n pass\n"}9{"file_name": "verl__model_merger__megatron_model_merger.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport json\nimport os\nimport warnings\nfrom contextlib import contextmanager\nfrom pathlib import Path\nfrom typing import Any, Callable, ContextManager\n\nimport numpy as np\nimport torch\nimport torch.distributed as dist\n\ntry:\n # NPU patch\n import mindspeed.megatron_adaptor # noqa: F401\nexcept ImportError:\n pass\n\nfrom accelerate import init_empty_weights\nfrom megatron.core import mpu\nfrom megatron.core.models.gpt.gpt_model import ModelType\nfrom megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed\nfrom safetensors.torch import load_file\nfrom transformers import (\n AutoConfig,\n PretrainedConfig,\n)\n\nfrom verl.models.mcore import hf_to_mcore_config\nfrom verl.utils.device import get_device_name, get_nccl_backend, get_torch_device\nfrom verl.utils.distributed import set_numa_affinity\nfrom verl.utils.megatron.dist_checkpointing import load_dist_checkpointing\nfrom verl.utils.megatron_utils import get_model\nfrom verl.utils.tokenizer import hf_processor, hf_tokenizer\n\nfrom .base_model_merger import BaseModelMerger, ModelMergerConfig\n\n\n@contextmanager\ndef noop_context() -> Any:\n yield\n\n\ndef get_dynamic_pipeline_shards(layer_num: int, pp_size: int) -> list[int]:\n \"\"\"Calculate the pipeline sharding configuration for Megatron-LM.\n\n Args:\n layer_num: Total number of layers in the model.\n pp_size: Number of pipeline parallel ranks.\n\n Returns:\n layer number of each pp rank. Make the sharding of the pipeline as uniform as possible.\n \"\"\"\n if layer_num < pp_size:\n raise ValueError(f\"layer_num {layer_num} must be greater than pp_size {pp_size}.\")\n\n if pp_size < 1:\n raise ValueError(f\"pp_size must be at least 1, got {pp_size}.\")\n if pp_size == 1:\n return [layer_num]\n\n if pp_size == 2:\n return [\n layer_num // 2,\n layer_num - layer_num // 2,\n ]\n\n middle_size = pp_size - 2\n shards_strategy = []\n for middle_layer_num in range(layer_num):\n first_last_layer_num = layer_num - middle_layer_num * middle_size\n first_layer_num = first_last_layer_num // 2\n last_layer_num = first_last_layer_num - first_last_layer_num // 2\n if 0 < first_layer_num <= middle_layer_num and 0 < last_layer_num <= middle_layer_num:\n shards_strategy.append(\n (\n [first_layer_num] + [middle_layer_num] * middle_size + [last_layer_num],\n abs(first_layer_num - middle_layer_num),\n )\n )\n\n # sort by diff of layer_num, to make it as uniform as possible\n res = sorted(shards_strategy, key=lambda x: x[1])[0][0]\n assert sum(res) == layer_num, f\"sum(res)={sum(res)} != layer_num={layer_num}, pp_size={pp_size}\"\n return res\n\n\nclass MegatronModelMerger(BaseModelMerger):\n \"\"\"\n Model merger for Megatron-LM distributed checkpoints.\n\n This class handles the conversion of Megatron-LM distributed checkpoints into HuggingFace format.\n Megatron-LM uses tensor parallelism, pipeline parallelism, and data parallelism to distribute\n large language models across multiple GPUs. This merger reconstructs the full model by\n loading distributed checkpoints and applying the necessary transformations.\n\n Key features:\n - Support for tensor parallel, pipeline parallel, and data parallel configurations\n - Automatic parameter name mapping from Megatron to HuggingFace conventions\n - Handling of QKV and gate-up tensor splitting/merging\n - Support for tied word embeddings and value models\n - Integration with Megatron's distributed checkpointing system\n\n The merger handles various model architectures and configurations:\n - Standard transformer models (GPT-style)\n - Models with tied word embeddings\n - Value models for reinforcement learning\n - Multi-layer attention (MLA) architectures\n - Mixture of Experts (MoE) models\n\n Args:\n config (ModelMergerConfig): Configuration object with Megatron-specific settings\n including tie_word_embedding and is_value_model flags.\n\n Example:\n To merge Megatron checkpoints:\n ```python\n config = ModelMergerConfig(\n operation=\"merge\",\n backend=\"megatron\",\n local_dir=\"path/to/megatron/checkpoints\",\n target_dir=\"path/to/output\",\n tie_word_embedding=True\n )\n merger = MegatronModelMerger(config)\n merger.merge_and_save()\n ```\n \"\"\"\n\n def __init__(self, config: ModelMergerConfig):\n super().__init__(config)\n # Currently we use only 1 rank to merge the dist_ckpt, we will move to multi-process save shortly afterwards\n if \"WORLD_SIZE\" not in os.environ:\n os.environ[\"RANK\"] = \"0\"\n os.environ[\"LOCAL_RANK\"] = \"0\"\n os.environ[\"WORLD_SIZE\"] = \"1\"\n os.environ[\"MASTER_ADDR\"] = \"localhost\"\n os.environ[\"MASTER_PORT\"] = \"12355\"\n\n set_numa_affinity()\n torch.distributed.init_process_group(get_nccl_backend())\n\n self.rank = torch.distributed.get_rank()\n self.world_size = torch.distributed.get_world_size()\n local_rank = os.environ.get(\"LOCAL_RANK\", 0)\n get_torch_device().set_device(f\"{get_device_name()}:{local_rank}\")\n\n mpu.initialize_model_parallel(\n tensor_model_parallel_size=1,\n pipeline_model_parallel_size=self.world_size,\n virtual_pipeline_model_parallel_size=None,\n context_parallel_size=1,\n expert_model_parallel_size=1,\n )\n model_parallel_cuda_manual_seed(0)\n self.hf_config = AutoConfig.from_pretrained(\n self.config.hf_model_config_path, trust_remote_code=self.config.trust_remote_code\n )\n print(self.hf_config, flush=True)\n\n self.params_mapping = {\n # megatron core gpt model name, huggingface model name\n # NOTICE: It's a little bit tricky, when 2 keys have the same prefix, we need to make sure the\n # longer key within the containing relationship is processed first.\n \"embedding.word_embeddings\": \"model.embed_tokens\",\n # input layer norm for dpskv3\n \"input_layernorm.weight\": \"input_layernorm.weight\",\n \"input_layernorm.bias\": \"input_layernorm.bias\",\n # attn\n \"self_attention.linear_qkv.layer_norm_weight\": \"input_layernorm.weight\",\n \"self_attention.linear_qkv.layer_norm_bias\": \"input_layernorm.bias\",\n \"self_attention.linear_qkv\": \"self_attn.qkv_proj\",\n \"self_attention.q_layernorm\": \"self_attn.q_norm\",\n \"self_attention.k_layernorm\": \"self_attn.k_norm\",\n \"self_attention.linear_proj\": \"self_attn.o_proj\",\n # mla\n \"self_attention.linear_q_proj\": \"self_attn.q_proj\",\n \"self_attention.linear_q_down_proj\": \"self_attn.q_a_proj\",\n \"self_attention.linear_q_up_proj.layer_norm_weight\": \"self_attn.q_a_layernorm.weight\",\n \"self_attention.linear_q_up_proj\": \"self_attn.q_b_proj\",\n \"self_attention.linear_kv_down_proj\": \"self_attn.kv_a_proj_with_mqa\",\n \"self_attention.linear_kv_up_proj.layer_norm_weight\": \"self_attn.kv_a_layernorm.weight\",\n \"self_attention.linear_kv_up_proj\": \"self_attn.kv_b_proj\",\n # mlp\n \"pre_mlp_layernorm\": \"post_attention_layernorm\",\n \"mlp.linear_fc1.layer_norm_weight\": \"post_attention_layernorm.weight\",\n \"mlp.linear_fc1.layer_norm_bias\": \"post_attention_layernorm.bias\",\n \"mlp.linear_fc1\": \"mlp.gate_up_proj\",\n \"mlp.linear_fc2\": \"mlp.down_proj\",\n # moe\n \"mlp.router.expert_bias\": \"mlp.gate.e_score_correction_bias\",\n \"mlp.router\": \"mlp.gate\",\n \"mlp.shared_experts.linear_fc1\": \"mlp.shared_experts.gate_up_proj\",\n \"mlp.shared_experts.linear_fc2\": \"mlp.shared_experts.down_proj\",\n \"linear_fc1\": \"gate_up_proj\",\n \"linear_fc2\": \"down_proj\",\n # output\n \"final_layernorm\": \"norm\",\n \"output_layer\": \"lm_head\",\n }\n\n if \"Qwen2MoeForCausalLM\" in self.hf_config.architectures:\n self.params_mapping[\"mlp.shared_experts.linear_fc1\"] = \"mlp.shared_expert.gate_up_proj\"\n self.params_mapping[\"mlp.shared_experts.linear_fc2\"] = \"mlp.shared_expert.down_proj\"\n self.params_mapping[\"mlp.shared_experts.gate_weight\"] = \"mlp.shared_expert_gate.weight\"\n\n def _load_state_dicts(self, model_ckpt_path: str) -> dict[str, Any]:\n \"\"\"_summary_\n Use Megatron dist_checkpointing to load the model state dicts from the checkpoint directory.\n\n Args:\n model_ckpt_path (str): Path to the model checkpoint directory.\n\n Returns:\n State dict containing the model parameters.\n \"\"\"\n\n # init hf config\n self.pipeline_shards = get_dynamic_pipeline_shards(self.hf_config.num_hidden_layers, self.world_size)\n print(f\"Pipeline shards: {self.pipeline_shards}, total layers: {sum(self.pipeline_shards)}\")\n\n tf_config = hf_to_mcore_config(\n self.hf_config,\n torch.bfloat16,\n num_layers_in_first_pipeline_stage=self.pipeline_shards[0] if len(self.pipeline_shards) > 1 else None,\n num_layers_in_last_pipeline_stage=self.pipeline_shards[-1] if len(self.pipeline_shards) > 2 else None,\n )\n tf_config.use_cpu_initialization = self.config.use_cpu_initialization\n tie_word_embeddings = getattr(self.hf_config, \"tie_word_embeddings\", False)\n\n # init megatron model\n def megatron_model_provider(pre_process, post_process):\n from verl.models.mcore import init_mcore_model\n\n parallel_model = init_mcore_model(\n tf_config,\n self.hf_config,\n pre_process,\n post_process,\n share_embeddings_and_output_weights=tie_word_embeddings,\n value=False,\n )\n return parallel_model\n\n context: Callable[..., ContextManager] = (\n init_empty_weights if self.config.use_cpu_initialization else noop_context\n )\n with context():\n whole_model = get_model(\n model_provider_func=megatron_model_provider,\n model_type=ModelType.encoder_or_decoder,\n wrap_with_ddp=False,\n transformer_config=tf_config,\n )\n\n if self.config.use_cpu_initialization:\n # convert meta device to empty tensor so it can use `copy_` function\n whole_model[0].module = whole_model[0].module.to_empty(device=\"cpu\")\n\n # load state dicts\n sharded_state_dict = {}\n for vpp_rank, model in enumerate(whole_model):\n key = f\"model{vpp_rank}\" if len(whole_model) > 1 else \"model\"\n mpu.set_virtual_pipeline_model_parallel_rank(vpp_rank)\n sharded_state_dict[key] = model.sharded_state_dict()\n model_state_dict = load_dist_checkpointing(sharded_state_dict, model_ckpt_path)\n model_state_dict_list = []\n for vpp_rank, model in enumerate(whole_model):\n key = f\"model{vpp_rank}\" if len(whole_model) > 1 else \"model\"\n mpu.set_virtual_pipeline_model_parallel_rank(vpp_rank)\n model_state_dict_list.append(model_state_dict[key])\n\n return model_state_dict_list\n\n def _check_megatron_state_key(self, key: str) -> bool:\n \"\"\"\n Checks if the key is a valid Megatron state key.\n\n Now the model merger only supports keys that start with \"decoder/embedding/output_layer\" in TransformerLayer.\n Shall not use key starts with \"model.\"\n \"\"\"\n if key.startswith(\"model.\"):\n raise ValueError(\n f\"Invalid key {key} in Megatron state_dict. Expected keys to start with \"\n f\"'decoder/embedding/output_layer' in TransformerLayer.\"\n )\n\n skip_checking_keys = [\"embedding.word_embeddings\", \"output_layer\"]\n for skip_key in skip_checking_keys:\n if skip_key in key:\n print(f\"skip checking key {key}\")\n return\n\n # Exclude extra state keys\n if not key.startswith(\"decoder\"):\n raise ValueError(\n f\"Invalid key {key} in Megatron state_dict. Expected keys to start with 'decoder' in TransformerLayer.\"\n )\n\n def _split_tensors(\n self, key: str, tensor: torch.Tensor, config: PretrainedConfig, is_value_model: bool = False\n ) -> list[torch.Tensor]:\n \"\"\"\n Splits a tensor into multiple tensors based on the name.\n This is used to handle qkv and gate_up tensors.\n \"\"\"\n if \"linear_fc1.weight\" in key:\n # if the tensor is gate and proj\n gate_lst = []\n up_lst = []\n gate, up = tensor.chunk(2)\n gate_lst.append(gate)\n up_lst.append(up)\n gate = torch.cat(gate_lst, dim=0)\n up = torch.cat(up_lst, dim=0)\n return [gate, up]\n elif \"self_attention.linear_qkv.\" in key and \"layer_norm\" not in key:\n # if the tensor is qkv, for each param on tp, split into q, k, v\n # concat q, k, v separately.\n q_lst, k_lst, v_lst = [], [], []\n assert config.num_attention_heads % config.num_key_value_heads == 0\n num_q_per_kv = config.num_attention_heads // config.num_key_value_heads\n assert tensor.shape[0] % (num_q_per_kv + 2) == 0, (\n f\"Tensor shape {tensor.shape} is not divisible by {num_q_per_kv + 2}\"\n )\n kv_size = tensor.shape[0] // (num_q_per_kv + 2)\n split_size = [kv_size * num_q_per_kv, kv_size, kv_size]\n\n num_query_groups_per_partition = config.num_key_value_heads\n for chunk in tensor.chunk(num_query_groups_per_partition):\n split_size = [\n kv_size * num_q_per_kv // num_query_groups_per_partition,\n kv_size // num_query_groups_per_partition,\n kv_size // num_query_groups_per_partition,\n ]\n q, k, v = chunk.split(split_size)\n q_lst.append(q)\n k_lst.append(k)\n v_lst.append(v)\n\n return [torch.cat(q_lst, dim=0), torch.cat(k_lst, dim=0), torch.cat(v_lst, dim=0)]\n else:\n return [tensor]\n\n def _merge_state_dicts(self, model_state_dict_list: list[dict[str, Any]]) -> dict[str, torch.Tensor]:\n state_dict = {}\n layers_cum = 0\n if self.world_size > 1:\n pipeline_cumsum = np.cumsum(self.pipeline_shards)\n layers_cum = 0 if self.rank == 0 else pipeline_cumsum[self.rank - 1]\n\n print(f\"{layers_cum=}\")\n for model_state_dict in model_state_dict_list:\n layers_handled = 0\n keys = model_state_dict.keys()\n for key in keys:\n if \"extra_state\" in key:\n continue\n if self.config.tie_word_embedding and (\"output_layer\" in key):\n print(\"skip lm_head and reward_head loading because of tie_word_embeddings\")\n continue\n\n self._check_megatron_state_key(key)\n hf_name = self._replace_name(key, self.params_mapping)\n assert hf_name is not None, f\"Failed to convert layer name [{key}] from megatron to huggingface.\"\n if \"model.layers.\" in hf_name:\n local_layer_no = int(hf_name.split(\".\")[2])\n layers_handled = max(local_layer_no, layers_handled)\n global_layer_no = local_layer_no + layers_cum\n new_key_list = hf_name.split(\".\")\n new_key_list[2] = str(global_layer_no)\n hf_name = \".\".join(new_key_list)\n else:\n warnings.warn(f\"hf_name {hf_name} will not be fixed with layer number\", stacklevel=2)\n\n if \"mlp.experts.\" in hf_name and \".weight\" in hf_name:\n name_prefix, expert_id = hf_name.split(\".weight\")\n for proj in [\"gate_up\", \"down\"]:\n if f\"{proj}_proj\" in hf_name:\n hf_name = hf_name.replace(\n f\"mlp.experts.{proj}_proj.weight{expert_id}\",\n f\"mlp.experts.{expert_id}.{proj}_proj.weight\",\n )\n\n tensor = model_state_dict[key]\n split_tensor = self._split_tensors(\n key, tensor, self.hf_config, is_value_model=self.config.is_value_model\n )\n\n if len(split_tensor) == 1:\n state_dict[hf_name] = split_tensor[0]\n elif len(split_tensor) == 3:\n # split qkv\n for n, d in zip([\"q\", \"k\", \"v\"], split_tensor, strict=True):\n state_dict[hf_name.replace(\"qkv\", n)] = d\n elif len(split_tensor) == 2:\n # split gate up\n state_dict[hf_name.replace(\"gate_up\", \"gate\")] = split_tensor[0]\n state_dict[hf_name.replace(\"gate_up\", \"up\")] = split_tensor[1]\n shape_info = (\n split_tensor.shape if isinstance(split_tensor, torch.Tensor) else [t.shape for t in split_tensor]\n )\n print(f\"converted {key} to {hf_name} with shape {shape_info}\")\n\n layers_cum += layers_handled + 1 # zero based\n\n return state_dict\n\n def save_hf_model_and_tokenizer(self, merged_state_dict):\n if self.world_size == 1:\n return super().save_hf_model_and_tokenizer(merged_state_dict)\n\n from safetensors.torch import save_file\n\n layer_num = self.hf_config.num_hidden_layers\n\n # FIXME: make configurable\n saves_per_layer = 1 if layer_num < 30 else 2\n saves_total = saves_per_layer * layer_num\n saves_indexes = {}\n\n # calculate the layer start index and key chunks\n layer_this_rank = self.pipeline_shards[self.rank]\n pipeline_cumsum = np.cumsum(self.pipeline_shards)\n layer_start = 0 if self.rank == 0 else pipeline_cumsum[self.rank - 1]\n keys = list(merged_state_dict.keys())\n keys_chunk = np.array_split(np.array(keys), layer_this_rank * saves_per_layer)\n numel = 0\n\n assert len(keys_chunk) == layer_this_rank * saves_per_layer, (\n f\"Expected {len(keys_chunk)} chunks, but got {layer_this_rank * saves_per_layer} for rank {self.rank}.\"\n )\n\n # save to model shards manually\n target_dir = Path(self.config.target_dir)\n for i, keys in enumerate(keys_chunk):\n sd_to_save = {k: merged_state_dict[k] for k in keys}\n numel += sum([sd_to_save[i].numel() for i in sd_to_save])\n save_idx = layer_start * saves_per_layer + i\n save_path = target_dir / f\"model-{save_idx + 1:05d}-of-{saves_total:05d}.safetensors\"\n\n save_file(sd_to_save, save_path)\n for k in keys:\n saves_indexes[k] = str(save_path.name)\n\n tensor = torch.tensor([numel]).to(get_device_name())\n dist.all_reduce(tensor, op=dist.ReduceOp.SUM)\n numel = tensor.cpu().item()\n\n all_save_indexes = [{} for _ in range(self.world_size)]\n dist.all_gather_object(all_save_indexes, saves_indexes)\n saves_indexes = {k: v for i in all_save_indexes for k, v in i.items()}\n if self.rank == 0:\n with open(target_dir / \"model.safetensors.index.json\", \"w\") as f:\n json.dump(\n {\n \"metadata\": {\n \"total_size\": numel,\n },\n \"weight_map\": saves_indexes,\n },\n f,\n indent=4,\n )\n print(f\"model saved to {target_dir} with {numel=}\")\n\n self.model_config.save_pretrained(self.config.target_dir)\n\n processor = hf_processor(self.hf_model_config_path, trust_remote_code=self.config.trust_remote_code)\n tokenizer = hf_tokenizer(self.hf_model_config_path, trust_remote_code=self.config.trust_remote_code)\n if processor is not None:\n print(f\"Saving processor to {self.config.target_dir}\")\n processor.save_pretrained(self.config.target_dir)\n if tokenizer is not None:\n print(f\"Saving tokenizer to {self.config.target_dir}\")\n tokenizer.save_pretrained(self.config.target_dir)\n\n def merge_and_save(self):\n from verl.utils.megatron_utils import get_dist_checkpoint_path\n\n model_ckpt_path = get_dist_checkpoint_path(self.config.local_dir)\n\n model_state_dict = self._load_state_dicts(model_ckpt_path)\n merged_state_dict = self._merge_state_dicts(model_state_dict)\n del model_state_dict\n\n if self.config.operation == \"test\":\n if not self.config.test_hf_dir:\n raise ValueError(\"test_hf_dir must be provided for test operation\")\n self._validate_state_dict(merged_state_dict)\n elif self.config.operation == \"merge\":\n self.save_hf_model_and_tokenizer(merged_state_dict)\n if self.config.hf_upload:\n self.upload_to_huggingface()\n else:\n raise ValueError(f\"Unknown operation: {self.config.operation}\")\n\n def _validate_state_dict(self, state_dict: dict[str, torch.Tensor]):\n \"\"\"\n Compares the merged Megatron state_dict against a reference safetensors model.\n Applies necessary name mappings from Megatron to Hugging Face conventions using _replace_name.\n \"\"\"\n ref_state_dict = load_file(Path(self.config.test_hf_dir) / \"model.safetensors\")\n\n for name, loaded_weight in state_dict.items():\n # name = self._replace_name(original_name, self.params_mapping)\n if not name or name.endswith(\".bias\") and name not in ref_state_dict:\n continue\n if \"rotary_emb.inv_freq\" in name:\n continue\n if \"lm_head.weight\" in name:\n if self.config.is_value_model or self.config.tie_word_embedding:\n continue\n if name not in ref_state_dict:\n raise RuntimeError(f\"key: {name} not exist in state_dict\")\n param = ref_state_dict[name]\n assert loaded_weight.dtype == param.dtype\n torch.testing.assert_close(loaded_weight.to(\"cpu\"), param, atol=1e-2, rtol=5e-2)\n\n def _replace_name(self, megatron_name: str, name_mapping: dict[str, str]) -> str:\n for m_name, v_name in name_mapping.items():\n if m_name not in megatron_name:\n continue\n\n megatron_name = megatron_name.replace(\"decoder\", \"model\")\n param_name = megatron_name.replace(m_name, v_name)\n\n return param_name\n\n return None # Return None if no mapping found\n\n def cleanup(self):\n torch.distributed.destroy_process_group()\n"}10{"file_name": "verl__models__llama__megatron__checkpoint_utils__llama_loader.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport time\n\nimport torch\nimport torch.distributed as dist\n\nfrom verl.utils.device import get_device_id, get_torch_device\n\n\ndef _megatron_calc_layer_map(config):\n \"\"\"Calculate the mapping of global layer_idx to local layer_idx\n Returns:\n layer_map (Dict: int -> tuple(int, int, int)):\n mapping from the global layer index to\n a tuple of (pp_rank, virtual_pp_rank, layer_idx inside model)\n \"\"\"\n from megatron.core import mpu\n\n print(f\"get megatron data parallel size: {mpu.get_data_parallel_world_size()}\")\n\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n\n layer_map = dict()\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n for pp_rank_idx in range(pp_size):\n for virtual_pp_rank_idx in range(virtual_pp_size):\n layer_offset = (\n virtual_pp_rank_idx * (config.num_hidden_layers // virtual_pp_size) + pp_rank_idx * num_layers_per_model\n )\n for layer_idx in range(num_layers_per_model):\n layer_map[layer_offset + layer_idx] = (\n pp_rank_idx,\n virtual_pp_rank_idx,\n layer_idx,\n )\n return layer_map\n\n\ndef load_state_dict_to_megatron_llama(\n state_dict, wrapped_models, config, params_dtype, is_value_model=False, tie_word_embeddings=False\n):\n \"\"\"Load merged state_dict to sharded Megatron module in training.\"\"\"\n from megatron.core import DistributedDataParallel as LocalDDP\n from megatron.core import mpu\n from megatron.core.transformer.module import Float16Module\n from torch.nn.parallel import DistributedDataParallel as torchDDP\n\n from verl.utils.logger import print_rank_0\n from verl.utils.megatron_utils import unwrap_model\n\n start_time = time.time()\n\n def _get_gpt_model(model):\n return model\n\n def fetch_params(module):\n for param in module.parameters():\n torch.distributed.fetch(\n param.data, src=mpu.get_data_parallel_src_rank(), group=mpu.get_data_parallel_group()\n )\n\n dp_rank = mpu.get_data_parallel_rank()\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n mp_group = mpu.get_model_parallel_group()\n\n if torch.distributed.get_rank() == 0:\n assert mp_group.rank() == 0, f\"mp_rank:[{mp_group.rank}] != 0 on rank #0\"\n assert pp_rank == 0, f\"pp_rank:[{pp_rank}] != 0 on rank #0\"\n assert dp_rank == 0, f\"dp_rank:[{dp_rank}] != 0 on rank #0\"\n\n if not isinstance(wrapped_models, list | tuple):\n wrapped_models = list(wrapped_models)\n\n assert len(wrapped_models) == virtual_pp_size\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers, (\n f\"num_layers_per_model: {num_layers_per_model} * pp_size: {pp_size} * virtual_pp_size \"\n f\"{virtual_pp_size} != config.num_hidden_layers: {config.num_hidden_layers}\"\n )\n\n models = [None] * len(wrapped_models)\n\n for i, wrapped_model in enumerate(wrapped_models):\n models[i] = unwrap_model(wrapped_model, (torchDDP, LocalDDP, Float16Module))\n gpt_model_module = _get_gpt_model(models[i])\n assert len(gpt_model_module.model.layers) == num_layers_per_model\n\n def _fetch_tensor(tensor, name) -> torch.Tensor:\n \"\"\"fetch tensor\"\"\"\n nonlocal state_dict\n if tensor is not None:\n tensor.data.copy_(state_dict[name])\n\n def _fetch_tp_shard_tensor_vocab(tensor, name, chunk_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"fetch tensor in tp shards\"\"\"\n nonlocal state_dict\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n if name in state_dict:\n full_weight = state_dict[name]\n\n if mutate_func is not None:\n full_weight = mutate_func(full_weight)\n tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)\n if tensor is not None:\n tensor.data.copy_(tensor_chunk[tp_rank])\n else:\n print(f\"tp_shard tensor:[{name}] not in state_dict, skip loading\")\n\n def _fetch_tp_shard_tensor(tensor, name, chunk_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"fetch tensor in tp shards\"\"\"\n nonlocal state_dict\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n if name in state_dict:\n full_weight = state_dict[name]\n\n if mutate_func is not None:\n full_weight = mutate_func(full_weight)\n tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)\n if tensor is not None:\n tensor.data.copy_(tensor_chunk[tp_rank])\n else:\n print(f\"tp_shard tensor:[{name}] not in state_dict, skip loading\")\n\n def _fetch_tp_shard_tensor_gate_up(tensor, gate_name, up_name) -> torch.Tensor:\n \"\"\"fetch gate_up tensor in tp shards\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n if gate_name in state_dict and up_name in state_dict:\n gate_weight = state_dict[gate_name]\n up_weight = state_dict[up_name]\n new_gate_up_weight = torch.empty(\n config.intermediate_size * 2, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n for i in range(tp_size):\n intermediate_size_tp = config.intermediate_size // tp_size\n gate_weight_tp = gate_weight[i * intermediate_size_tp : (i + 1) * intermediate_size_tp]\n up_weight_tp = up_weight[i * intermediate_size_tp : (i + 1) * intermediate_size_tp]\n new_gate_up_weight[intermediate_size_tp * 2 * i : intermediate_size_tp * 2 * (i + 1)].copy_(\n torch.cat([gate_weight_tp, up_weight_tp], dim=0)\n )\n\n tensor_chunk = torch.chunk(new_gate_up_weight, tp_size, dim=0)\n if tensor is not None:\n tensor.data.copy_(tensor_chunk[tp_rank])\n else:\n print(f\"tp_shard tensor:[{gate_name}, {up_name}] not in state_dict, skip loading\")\n\n def _fetch_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name) -> torch.Tensor:\n \"\"\"fetch tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n assert q_name in state_dict and k_name in state_dict and v_name in state_dict\n full_weight_q = state_dict[q_name]\n full_weight_k = state_dict[k_name]\n full_weight_v = state_dict[v_name]\n\n hidden_size_per_head = config.hidden_size // config.num_attention_heads\n\n if config.num_key_value_heads >= tp_size:\n q_size_tp = config.hidden_size // tp_size\n kv_size_tp = hidden_size_per_head * config.num_key_value_heads // tp_size\n total_size = q_size_tp + 2 * kv_size_tp\n new_weight_qkv = torch.empty(\n total_size * tp_size, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n for i in range(tp_size):\n q_part = full_weight_q[i * q_size_tp : (i + 1) * q_size_tp]\n k_part = full_weight_k[i * kv_size_tp : (i + 1) * kv_size_tp]\n v_part = full_weight_v[i * kv_size_tp : (i + 1) * kv_size_tp]\n new_weight_qkv[i * total_size : (i + 1) * total_size].copy_(torch.cat([q_part, k_part, v_part], dim=0))\n\n else:\n q_size_tp = config.hidden_size // tp_size\n kv_size_tp = hidden_size_per_head\n total_size = q_size_tp + 2 * kv_size_tp\n new_weight_qkv = torch.empty(\n total_size * tp_size, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n for i in range(tp_size):\n q_part = full_weight_q[i * q_size_tp : (i + 1) * q_size_tp]\n start_idx = i * config.num_key_value_heads // tp_size * hidden_size_per_head\n end_idx = (i * config.num_key_value_heads // tp_size + 1) * hidden_size_per_head\n k_part = full_weight_k[start_idx:end_idx]\n v_part = full_weight_v[start_idx:end_idx]\n new_weight_qkv[i * total_size : (i + 1) * total_size].copy_(torch.cat([q_part, k_part, v_part], dim=0))\n\n tensor_chunk = torch.chunk(new_weight_qkv, tp_size, dim=0)\n if tensor is not None:\n tensor.data.copy_(tensor_chunk[tp_rank])\n\n # Embeddings\n # -------------------\n print_rank_0(\"loading embeddings...\")\n gpt_model_module = _get_gpt_model(models[0])\n embed_tokens_weight = None\n if pp_rank == 0:\n embed_tokens_weight = gpt_model_module.model.embed_tokens.weight\n _fetch_tp_shard_tensor_vocab(embed_tokens_weight, \"model.embed_tokens.weight\")\n\n # Transformer layers\n # -------------------\n layer_map = _megatron_calc_layer_map(config)\n\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n num_layer_per_pp = config.num_hidden_layers // pp_size\n vpp_size = mpu.get_virtual_pipeline_model_parallel_world_size()\n\n layer_list = []\n if vpp_size is not None:\n for vpp_rank in range(vpp_size):\n num_layer_vpp_chunk = num_layer_per_pp // vpp_size\n num_layer_this_model = num_layer_vpp_chunk\n offset = vpp_rank * (config.num_hidden_layers // mpu.get_virtual_pipeline_model_parallel_world_size()) + (\n mpu.get_pipeline_model_parallel_rank() * num_layer_vpp_chunk\n )\n layer_list.extend(list(range(offset, offset + num_layer_this_model)))\n else:\n num_layer_this_model = num_layer_per_pp\n offset = pp_rank * num_layer_per_pp\n layer_list.extend(list(range(offset, offset + num_layer_this_model)))\n\n for layer in layer_list:\n print_rank_0(f\"loading layer #{layer}...\")\n layer_name = f\"model.layers.{layer}\"\n dst_pp_rank, dst_virtual_pp_rank, dst_layer_idx = layer_map[layer]\n\n gpt_model_module = _get_gpt_model(models[dst_virtual_pp_rank])\n sync_layer = gpt_model_module.model.layers[dst_layer_idx]\n\n _fetch_tensor(\n sync_layer.input_layernorm.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.input_layernorm.weight\",\n )\n\n _fetch_tp_shard_tensor_qkv(\n sync_layer.self_attn.qkv_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.q_proj.weight\",\n f\"{layer_name}.self_attn.k_proj.weight\",\n f\"{layer_name}.self_attn.v_proj.weight\",\n )\n\n _fetch_tp_shard_tensor(\n sync_layer.self_attn.o_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.o_proj.weight\",\n chunk_dim=1,\n )\n\n _fetch_tensor(\n sync_layer.post_attention_layernorm.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.post_attention_layernorm.weight\",\n )\n\n _fetch_tp_shard_tensor_gate_up(\n sync_layer.mlp.gate_up_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.mlp.gate_proj.weight\",\n f\"{layer_name}.mlp.up_proj.weight\",\n )\n\n _fetch_tp_shard_tensor(\n sync_layer.mlp.down_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.mlp.down_proj.weight\",\n chunk_dim=1,\n )\n # Final Layernorm\n # -------------------\n print_rank_0(\"loading final layernorm...\")\n gpt_model_module = _get_gpt_model(models[-1])\n _fetch_tensor(\n getattr(gpt_model_module.model.norm, \"weight\", None),\n \"model.norm.weight\",\n )\n\n print_rank_0(\"loading lm_head...\")\n if pp_rank + 1 == pp_size:\n lm_head_weight = gpt_model_module.lm_head.weight\n\n if is_value_model:\n if \"lm_head.weight\" in state_dict and state_dict[\"lm_head.weight\"].shape[0] == 1:\n _fetch_tensor(lm_head_weight, \"lm_head.weight\")\n print_rank_0(\"load lm_head weight\")\n elif \"reward_head.weight\" in state_dict and state_dict[\"reward_head.weight\"].shape[0] == 1:\n _fetch_tensor(lm_head_weight, \"reward_head.weight\")\n print_rank_0(\"load lm_head from value_head weight\")\n else:\n _fetch_tensor(None, \"lm_head.weight\")\n print_rank_0(\"fail to match lm_head in value_model\")\n else:\n _fetch_tp_shard_tensor(lm_head_weight, \"lm_head.weight\")\n\n dist.barrier()\n get_torch_device().empty_cache()\n print_rank_0(f\"loading megatron ckpt done, time elapsed {time.time() - start_time}s\")\n"}11{"file_name": "verl__models__llama__megatron__checkpoint_utils__llama_saver.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport time\n\nimport torch\nimport torch.distributed as dist\nfrom megatron.core import mpu\nfrom megatron.core.distributed import DistributedDataParallel as LocalDDP\nfrom megatron.core.transformer.module import Float16Module\nfrom torch.nn.parallel import DistributedDataParallel as torchDDP\n\nfrom verl.utils.device import get_device_id, get_torch_device\nfrom verl.utils.logger import print_rank_0\nfrom verl.utils.megatron_utils import unwrap_model\n\n\ndef _megatron_calc_global_rank(tp_rank: int = 0, dp_rank: int = 0, pp_rank: int = 0):\n \"\"\"given TP,DP,PP rank to get the global rank.\"\"\"\n\n tp_size = mpu.get_tensor_model_parallel_world_size()\n dp_size = mpu.get_data_parallel_world_size()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n assert tp_size * dp_size * pp_size == torch.distributed.get_world_size(), (\n f\"{tp_size} x {dp_size} x {pp_size} != {torch.distributed.get_world_size()}\"\n )\n # We only support TP-DP-PP grouping, for correctness when resharding\n return (pp_rank * dp_size + dp_rank) * tp_size + tp_rank\n\n\ndef _megatron_calc_layer_map(config):\n \"\"\"Calculate the mapping of global layer_idx to local layer_idx\n Returns:\n layer_map (Dict: int -> tuple(int, int, int)):\n mapping from the global layer index to\n a tuple of (pp_rank, virtual_pp_rank, layer_idx inside model)\n \"\"\"\n from megatron.core import mpu\n\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n\n layer_map = dict()\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n for pp_rank_idx in range(pp_size):\n for virtual_pp_rank_idx in range(virtual_pp_size):\n layer_offset = (\n virtual_pp_rank_idx * (config.num_hidden_layers // virtual_pp_size) + pp_rank_idx * num_layers_per_model\n )\n for layer_idx in range(num_layers_per_model):\n layer_map[layer_offset + layer_idx] = (\n pp_rank_idx,\n virtual_pp_rank_idx,\n layer_idx,\n )\n return layer_map\n\n\ndef merge_megatron_ckpt_llama(wrapped_models, config, dtype, is_value_model=False, tie_word_embeddings=False):\n \"\"\"Merge sharded parameters of a Megatron module into a merged checkpoint.\n\n Args:\n wrapped_models (list of megatron.core.distributed.DistributedDataParallel):\n The local DDP wrapped megatron modules.\n config (str or None):\n HF config for model\n dtype: model params type\n is_value_model: if model is value model\n tie_word_embeddings: tie_word_embeddings, not used in llama, only to keep same interface with qwen2\n Returns:\n state_dict (dict):\n The merged state_dict in rank 0, and an empty dictionary in other ranks.\n \"\"\"\n start_time = time.time()\n\n def _get_gpt_model(model):\n return model\n\n dp_rank = mpu.get_data_parallel_rank()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n mp_group = mpu.get_model_parallel_group()\n\n if dist.get_rank() == 0:\n assert mp_group.rank() == 0, f\"mp_rank:[{mp_group.rank}] != 0 on rank #0\"\n assert pp_rank == 0, f\"pp_rank:[{pp_rank}] != 0 on rank #0\"\n assert dp_rank == 0, f\"dp_rank:[{dp_rank}] != 0 on rank #0\"\n\n if not isinstance(wrapped_models, list | tuple):\n wrapped_models = list(wrapped_models)\n\n assert len(wrapped_models) == virtual_pp_size\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n models = [None] * len(wrapped_models)\n\n for i, wrapped_model in enumerate(wrapped_models):\n models[i] = unwrap_model(wrapped_model, (torchDDP, LocalDDP, Float16Module))\n assert len(models[i].model.layers) == num_layers_per_model, (\n \"len model layers {} not equal to num_layers_per_model {}\".format(\n len(models[i].model.layers), num_layers_per_model\n )\n )\n\n state_dict = dict()\n\n def _get_cpu_tensor(tensor: torch.Tensor):\n if tensor is None:\n return None\n if tensor.device == torch.device(\"cpu\"):\n return tensor.detach().clone()\n return tensor.detach().cpu()\n\n def _broadcast_tensor(tensor, name, src_pp_rank) -> torch.Tensor:\n \"\"\"broadcast tensor across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=src_pp_rank)\n\n if torch.distributed.get_rank() == src_rank:\n if tensor is None:\n weight = None\n tensor_shape = None\n else:\n weight = tensor\n tensor_shape = weight.shape\n else:\n weight = None\n tensor_shape = None\n\n obj_list = [tensor_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n tensor_shape = obj_list[0]\n\n if tensor_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tensor:[{name}] not exist, skip collect\")\n return\n\n if weight is None:\n weight = torch.empty(\n tensor_shape,\n dtype=dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n\n dist.broadcast(weight, src=src_rank, group=mp_group)\n\n if torch.distributed.get_rank() == 0:\n state_dict[name] = _get_cpu_tensor(weight)\n\n def _broadcast_tp_shard_tensor(tensor, name, src_pp_rank, concat_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_size = mpu.get_tensor_model_parallel_world_size()\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=src_pp_rank)\n\n chunk_shape = tensor.shape if torch.distributed.get_rank() == src_rank else None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{name}] not exist, skip collecting\")\n return\n\n buffer_tensor = torch.empty(\n chunk_shape,\n dtype=dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n\n chunk_tensors = [None] * tp_size\n\n for i in range(tp_size):\n cur_src_rank = _megatron_calc_global_rank(tp_rank=i, dp_rank=0, pp_rank=src_pp_rank)\n sync_tensor = tensor if torch.distributed.get_rank() == cur_src_rank else buffer_tensor\n dist.broadcast(sync_tensor, src=cur_src_rank, group=mp_group)\n\n if torch.distributed.get_rank() == 0:\n chunk_tensors[i] = _get_cpu_tensor(sync_tensor)\n\n if torch.distributed.get_rank() == 0:\n full_tensor = torch.concat(chunk_tensors, dim=concat_dim)\n if mutate_func is not None:\n full_tensor = mutate_func(full_tensor)\n state_dict[name] = full_tensor\n\n def _broadcast_tp_shard_tensor_gate_up(tensor, gate_name, up_name, src_pp_rank) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_size = mpu.get_tensor_model_parallel_world_size()\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=src_pp_rank)\n\n chunk_shape = tensor.shape if torch.distributed.get_rank() == src_rank else None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{gate_name, up_name}] not exist, skip collecting\")\n return\n\n buffer_tensor = torch.empty(\n chunk_shape,\n dtype=dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n\n chunk_tensors = [None] * tp_size\n\n for i in range(tp_size):\n cur_src_rank = _megatron_calc_global_rank(tp_rank=i, dp_rank=0, pp_rank=src_pp_rank)\n sync_tensor = tensor if torch.distributed.get_rank() == cur_src_rank else buffer_tensor\n dist.broadcast(sync_tensor, src=cur_src_rank, group=mp_group)\n\n if torch.distributed.get_rank() == 0:\n chunk_tensors[i] = _get_cpu_tensor(sync_tensor)\n\n if torch.distributed.get_rank() == 0:\n full_tensor = torch.concat(chunk_tensors, dim=0)\n intermediate_size_tp = config.intermediate_size // tp_size\n gate_weight_list = []\n up_weight_list = []\n for i in range(tp_size):\n gate_up_weight_tp = full_tensor[intermediate_size_tp * 2 * i : intermediate_size_tp * 2 * (i + 1)]\n gate_weight_tp = gate_up_weight_tp[:intermediate_size_tp]\n up_weight_tp = gate_up_weight_tp[intermediate_size_tp:]\n gate_weight_list.append(gate_weight_tp)\n up_weight_list.append(up_weight_tp)\n\n state_dict[gate_name] = torch.cat(gate_weight_list, dim=0)\n state_dict[up_name] = torch.cat(up_weight_list, dim=0)\n\n def _broadcast_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name, src_pp_rank):\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_size = mpu.get_tensor_model_parallel_world_size()\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=src_pp_rank)\n\n chunk_shape = tensor.shape if torch.distributed.get_rank() == src_rank else None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{q_name}] not exist, skip collecting\")\n return\n\n buffer_tensor = torch.empty(\n chunk_shape,\n dtype=dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n\n chunk_tensors = [None] * tp_size\n\n for i in range(tp_size):\n cur_src_rank = _megatron_calc_global_rank(tp_rank=i, dp_rank=0, pp_rank=src_pp_rank)\n sync_tensor = tensor if torch.distributed.get_rank() == cur_src_rank else buffer_tensor\n dist.broadcast(sync_tensor, src=cur_src_rank, group=mp_group)\n\n if torch.distributed.get_rank() == 0:\n chunk_tensors[i] = _get_cpu_tensor(sync_tensor)\n\n if torch.distributed.get_rank() == 0:\n full_tensor = torch.concat(chunk_tensors, dim=0)\n q_weight_list = []\n k_weight_list = []\n v_weight_list = []\n hidden_size_per_head = config.hidden_size // config.num_attention_heads\n\n if config.num_key_value_heads >= tp_size:\n q_size_tp = config.hidden_size // tp_size\n kv_size_tp = hidden_size_per_head * config.num_key_value_heads // tp_size\n total_size = q_size_tp + 2 * kv_size_tp\n for i in range(tp_size):\n qkv_part = full_tensor[i * total_size : (i + 1) * total_size]\n q_part = qkv_part[:q_size_tp]\n k_part = qkv_part[q_size_tp : q_size_tp + kv_size_tp]\n v_part = qkv_part[q_size_tp + kv_size_tp : total_size]\n q_weight_list.append(q_part)\n k_weight_list.append(k_part)\n v_weight_list.append(v_part)\n else:\n q_size_tp = config.hidden_size // tp_size\n kv_size_tp = hidden_size_per_head\n total_size = q_size_tp + 2 * kv_size_tp\n for i in range(tp_size):\n qkv_part = full_tensor[i * total_size : (i + 1) * total_size]\n q_part = qkv_part[:q_size_tp]\n k_part = qkv_part[q_size_tp : q_size_tp + kv_size_tp]\n v_part = qkv_part[q_size_tp + kv_size_tp : total_size]\n q_weight_list.append(q_part)\n if i * config.num_key_value_heads % tp_size == 0:\n k_weight_list.append(k_part)\n v_weight_list.append(v_part)\n\n state_dict[q_name] = torch.cat(q_weight_list, dim=0)\n state_dict[k_name] = torch.cat(k_weight_list, dim=0)\n state_dict[v_name] = torch.cat(v_weight_list, dim=0)\n\n # empty cache before collecting weights\n get_torch_device().empty_cache()\n # Embeddings\n # -------------------\n if dp_rank == 0:\n # Embeddings\n # -------------------\n print_rank_0(\"collecting embeddings...\")\n gpt_model_module = _get_gpt_model(models[0])\n _broadcast_tp_shard_tensor(\n gpt_model_module.model.embed_tokens.weight if pp_rank == 0 else None,\n \"model.embed_tokens.weight\",\n src_pp_rank=0,\n )\n\n # Transformer layers\n # -------------------\n layer_map = _megatron_calc_layer_map(config)\n for layer in range(config.num_hidden_layers):\n print_rank_0(f\"collecting layer #{layer}...\")\n layer_name = f\"model.layers.{layer}\"\n src_pp_rank, src_virtual_pp_rank, src_layer_idx = layer_map[layer]\n\n gpt_model_module = _get_gpt_model(models[src_virtual_pp_rank])\n sync_layer = gpt_model_module.model.layers[src_layer_idx]\n\n _broadcast_tensor(\n sync_layer.input_layernorm.weight,\n f\"{layer_name}.input_layernorm.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tp_shard_tensor_qkv(\n sync_layer.self_attn.qkv_proj.weight,\n f\"{layer_name}.self_attn.q_proj.weight\",\n f\"{layer_name}.self_attn.k_proj.weight\",\n f\"{layer_name}.self_attn.v_proj.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tp_shard_tensor(\n sync_layer.self_attn.o_proj.weight,\n f\"{layer_name}.self_attn.o_proj.weight\",\n concat_dim=1,\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tensor(\n sync_layer.post_attention_layernorm.weight,\n f\"{layer_name}.post_attention_layernorm.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tp_shard_tensor_gate_up(\n sync_layer.mlp.gate_up_proj.weight,\n f\"{layer_name}.mlp.gate_proj.weight\",\n f\"{layer_name}.mlp.up_proj.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tp_shard_tensor(\n sync_layer.mlp.down_proj.weight,\n f\"{layer_name}.mlp.down_proj.weight\",\n concat_dim=1,\n src_pp_rank=src_pp_rank,\n )\n\n # Final Layernorm\n # -------------------\n print_rank_0(\"collecting final layernorm...\")\n gpt_model_module = _get_gpt_model(models[-1])\n _broadcast_tensor(\n getattr(gpt_model_module.model.norm, \"weight\", None),\n \"model.norm.weight\",\n src_pp_rank=pp_size - 1,\n )\n\n print_rank_0(\"collecting lm_head...\")\n\n if is_value_model:\n if pp_rank == pp_size - 1:\n print(f\"gpt_model_module.lm_head.weight: {gpt_model_module.lm_head.weight.shape}\")\n _broadcast_tensor(\n gpt_model_module.lm_head.weight if pp_rank == pp_size - 1 else None,\n \"lm_head.weight\",\n src_pp_rank=pp_size - 1,\n )\n _broadcast_tensor(\n gpt_model_module.reward_head.weight\n if pp_rank == pp_size - 1 and getattr(gpt_model_module, \"reward_weight\", None) is not None\n else None,\n \"reward_head.weight\",\n src_pp_rank=pp_size - 1,\n )\n\n else:\n _broadcast_tp_shard_tensor(\n getattr(gpt_model_module.lm_head, \"weight\", None) if pp_rank == pp_size - 1 else None,\n \"lm_head.weight\",\n src_pp_rank=pp_size - 1,\n )\n\n dist.barrier()\n\n get_torch_device().empty_cache()\n if torch.distributed.get_rank() == 0:\n if dtype not in [torch.float16, torch.bfloat16, torch.float32]:\n print(f'Unknown/unsupported dtype to save: {dtype}\"')\n exit(1)\n for k, v in state_dict.items():\n if dtype != v.dtype:\n state_dict[k] = v.to(dtype)\n\n print_rank_0(f\"merge megatron ckpt done, time elapsed {time.time() - start_time}s\")\n return state_dict\n"}12{"file_name": "verl__models__llama__megatron__layers__parallel_decoder.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.\n#\n# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX\n# and OPT implementations in this library. It has been modified from its\n# original forms to accommodate minor architectural differences compared\n# to GPT-NeoX and OPT used by the Meta AI team that trained the model.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Optional\n\nimport torch\nfrom megatron.core import ModelParallelConfig\nfrom torch import nn\nfrom transformers import LlamaConfig\n\nfrom verl.utils.megatron_utils import TransformerConfig, convert_config\n\nfrom .parallel_attention import ParallelLlamaAttention, ParallelLlamaAttentionRmPad\nfrom .parallel_mlp import ParallelLlamaMLP\nfrom .parallel_rmsnorm import ParallelLlamaRMSNorm\n\n\nclass ParallelLlamaDecoderLayer(nn.Module):\n def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig, layer_idx: int):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.layer_idx = layer_idx\n self.hidden_size = config.hidden_size\n self.self_attn = ParallelLlamaAttention(config=config, megatron_config=megatron_config)\n\n self.mlp = ParallelLlamaMLP(config, megatron_config=megatron_config)\n self.input_layernorm = ParallelLlamaRMSNorm(config, megatron_config)\n self.post_attention_layernorm = ParallelLlamaRMSNorm(config, megatron_config)\n\n def forward(\n self,\n hidden_states: torch.Tensor,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple[torch.FloatTensor, Optional[tuple[torch.FloatTensor, torch.FloatTensor]]]:\n \"\"\"\n Args:\n hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`\n attention_mask (`torch.FloatTensor`, *optional*): attention mask of size\n `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.\n output_attentions (`bool`, *optional*):\n Whether or not to return the attentions tensors of all attention layers. See `attentions` under\n returned tensors for more detail.\n use_cache (`bool`, *optional*):\n If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding\n (see `past_key_values`).\n past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states\n \"\"\"\n\n residual = hidden_states\n\n hidden_states = self.input_layernorm(hidden_states)\n\n # Note: sequence parallel is hidden inside ColumnParallelLinear\n # reduce scatter is hidden inside RowParallelLinear\n\n # Self Attention\n hidden_states = self.self_attn(\n hidden_states=hidden_states,\n attention_mask=attention_mask,\n position_ids=position_ids,\n )\n\n # TODO: add sequence parallel operator reduce_scatter here\n\n hidden_states = residual + hidden_states\n\n # Fully Connected\n residual = hidden_states\n hidden_states = self.post_attention_layernorm(hidden_states)\n\n # TODO: add sequence parallel operator all_gather here\n\n hidden_states = self.mlp(hidden_states)\n\n # TODO: add sequence parallel operator reduce_scatter here\n\n hidden_states = residual + hidden_states\n\n outputs = hidden_states\n\n return outputs\n\n\nclass ParallelLlamaDecoderLayerRmPad(nn.Module):\n def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig, layer_idx: int):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.layer_idx = layer_idx\n self.hidden_size = config.hidden_size\n self.self_attn = ParallelLlamaAttentionRmPad(config=config, megatron_config=megatron_config)\n\n self.mlp = ParallelLlamaMLP(config, megatron_config=megatron_config)\n self.input_layernorm = ParallelLlamaRMSNorm(config, megatron_config)\n self.post_attention_layernorm = ParallelLlamaRMSNorm(config, megatron_config)\n\n def forward(\n self,\n hidden_states: torch.Tensor,\n position_ids: Optional[torch.LongTensor] = None,\n sequence_length: int = None,\n indices: torch.Tensor = None,\n cu_seqlens: int = None,\n max_seqlen_in_batch: int = None,\n ) -> tuple[torch.FloatTensor, Optional[tuple[torch.FloatTensor, torch.FloatTensor]]]:\n residual = hidden_states # (total_nnz // sp, 1, hidden_size)\n\n hidden_states = self.input_layernorm(hidden_states)\n\n # Self Attention\n # (total_nnz // sp, 1, hidden_size) -> all-gather (total_nnz, 1, hidden_size)\n # -> col + row -> reduce-scatter -> (total_nnz // sp, 1, hidden_size)\n hidden_states = self.self_attn(\n hidden_states=hidden_states,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n hidden_states = residual + hidden_states\n\n # Fully Connected\n # shape changes same as attn\n residual = hidden_states\n hidden_states = self.post_attention_layernorm(hidden_states)\n hidden_states = self.mlp(hidden_states)\n hidden_states = residual + hidden_states\n\n outputs = hidden_states\n\n return outputs\n"}13{"file_name": "verl__models__llama__megatron__layers__parallel_linear.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2023 The vLLM team.\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n# Adapted from https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/linear.py\n\nimport torch\nfrom megatron.core import tensor_parallel\n\n\nclass QKVParallelLinear(tensor_parallel.ColumnParallelLinear):\n def __init__(\n self,\n input_size,\n num_heads,\n num_key_value_heads,\n head_dim,\n *,\n bias=True,\n gather_output=True,\n skip_bias_add=False,\n **kwargs,\n ):\n # Keep input parameters, and already restrict the head numbers\n self.input_size = input_size\n self.q_output_size = num_heads * head_dim\n self.kv_output_size = num_key_value_heads * head_dim\n self.head_dim = head_dim\n self.gather_output = gather_output\n self.skip_bias_add = skip_bias_add\n\n input_size = self.input_size\n output_size = (num_heads + 2 * num_key_value_heads) * self.head_dim\n\n super().__init__(\n input_size=input_size,\n output_size=output_size,\n bias=bias,\n gather_output=gather_output,\n skip_bias_add=skip_bias_add,\n **kwargs,\n )\n\n\nclass MergedColumnParallelLinear(tensor_parallel.ColumnParallelLinear):\n def __init__(\n self,\n input_size,\n gate_ouput_size,\n up_output_size,\n *,\n bias=True,\n gather_output=True,\n skip_bias_add=False,\n **kwargs,\n ):\n # Keep input parameters, and already restrict the head numbers\n self.input_size = input_size\n self.output_size = gate_ouput_size + up_output_size\n self.gather_output = gather_output\n self.skip_bias_add = skip_bias_add\n\n super().__init__(\n input_size=self.input_size,\n output_size=self.output_size,\n bias=bias,\n gather_output=gather_output,\n skip_bias_add=skip_bias_add,\n **kwargs,\n )\n\n\nclass LinearForLastLayer(torch.nn.Linear):\n def __init__(\n self,\n input_size,\n output_size,\n *,\n config,\n bias=True,\n ):\n super().__init__(in_features=input_size, out_features=output_size, bias=bias)\n self.sequence_parallel = config.sequence_parallel\n if self.sequence_parallel:\n self.weight.sequence_parallel = True\n\n def forward(\n self,\n input_,\n weight=None,\n runtime_gather_output=None,\n ):\n logits = super().forward(input_)\n logits = logits.float()\n if self.sequence_parallel:\n logits = tensor_parallel.gather_from_sequence_parallel_region(logits, tensor_parallel_output_grad=False)\n return logits, None\n"}14{"file_name": "verl__models__llama__megatron__layers__parallel_mlp.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.\n#\n# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX\n# and OPT implementations in this library. It has been modified from its\n# original forms to accommodate minor architectural differences compared\n# to GPT-NeoX and OPT used by the Meta AI team that trained the model.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom megatron.core import ModelParallelConfig, tensor_parallel\nfrom megatron.core import parallel_state as mpu\nfrom torch import nn\nfrom transformers.activations import ACT2FN\n\nfrom verl.models.llama.megatron.layers.parallel_linear import MergedColumnParallelLinear\nfrom verl.utils.megatron import tensor_parallel as tp_utils\n\n\nclass ParallelLlamaMLP(nn.Module):\n def __init__(self, config, megatron_config: ModelParallelConfig = None) -> None:\n super().__init__()\n self.config = config\n self.hidden_size = config.hidden_size\n self.intermediate_size = config.intermediate_size\n # The weight is only [hidden_size, intermediate_size // model_parallel_world_size]\n\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n row_kwargs = tp_utils.get_default_kwargs_for_row_parallel_linear()\n\n if megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n assert row_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(row_kwargs, megatron_config)\n tp_utils.update_kwargs_with_config(column_kwargs, megatron_config)\n\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n self.gate_up_proj = MergedColumnParallelLinear(\n input_size=self.hidden_size,\n gate_ouput_size=self.intermediate_size,\n up_output_size=self.intermediate_size,\n bias=False,\n gather_output=False,\n skip_bias_add=False,\n **column_kwargs,\n )\n self.gate_size = self.intermediate_size // tp_size\n\n self.down_proj = tensor_parallel.RowParallelLinear(\n input_size=self.intermediate_size,\n output_size=self.hidden_size,\n bias=False,\n input_is_parallel=True,\n skip_bias_add=False,\n **row_kwargs,\n )\n\n self.act_fn = ACT2FN[config.hidden_act]\n\n def forward(self, x):\n gate_up = self.gate_up_proj(x)[0]\n gate, up = gate_up.split(self.gate_size, dim=-1)\n return self.down_proj(self.act_fn(gate) * up)[0]\n"}15{"file_name": "verl__models__llama__megatron__layers__parallel_rmsnorm.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport numbers\n\nimport torch\nfrom megatron.core import ModelParallelConfig\nfrom torch import nn\nfrom transformers import LlamaConfig\n\nfrom verl.utils.megatron import sequence_parallel as sp_utils\n\n\nclass ParallelLlamaRMSNorm(nn.Module):\n def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig):\n \"\"\"\n LlamaRMSNorm is equivalent to T5LayerNorm\n \"\"\"\n super().__init__()\n if isinstance(config.hidden_size, numbers.Integral):\n normalized_shape = (config.hidden_size,)\n self.normalized_shape = torch.Size(normalized_shape)\n self.weight = nn.Parameter(torch.ones(self.normalized_shape))\n self.variance_epsilon = config.rms_norm_eps\n\n if megatron_config.sequence_parallel:\n sp_utils.mark_parameter_as_sequence_parallel(self.weight)\n\n def forward(self, hidden_states):\n from apex.normalization.fused_layer_norm import fused_rms_norm_affine\n\n return fused_rms_norm_affine(\n input=hidden_states,\n weight=self.weight,\n normalized_shape=self.normalized_shape,\n eps=self.variance_epsilon,\n memory_efficient=True,\n )\n"}16{"file_name": "verl__models__llama__megatron__modeling_llama_megatron.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.\n#\n# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX\n# and OPT implementations in this library. It has been modified from its\n# original forms to accommodate minor architectural differences compared\n# to GPT-NeoX and OPT used by the Meta AI team that trained the model.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"PyTorch LLaMA model with Megatron-style acceleration.\"\"\"\n\nfrom typing import Optional\n\nimport torch\nimport torch.utils.checkpoint\nfrom megatron.core import ModelParallelConfig, mpu, tensor_parallel\nfrom torch import nn\nfrom transformers.modeling_outputs import BaseModelOutputWithPast\nfrom transformers.models.llama.configuration_llama import LlamaConfig\nfrom transformers.models.llama.modeling_llama import CausalLMOutputWithPast\n\nfrom verl.utils.megatron import sequence_parallel as sp_utils\nfrom verl.utils.megatron import tensor_parallel as tp_utils\nfrom verl.utils.megatron_utils import TransformerConfig, convert_config\n\nfrom .layers import ParallelLlamaDecoderLayer, ParallelLlamaDecoderLayerRmPad, ParallelLlamaRMSNorm\n\n\"\"\"\nTODO: \n1. Add weight initialization. Here we need to be careful on TP weight init.\n2. Add sequence parallel\n3. Load checkpoint from meta LLama pretrained checkpoint\n\"\"\"\n\n\n# Copied from transformers.models.bart.modeling_bart._make_causal_mask\ndef _make_causal_mask(input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device):\n \"\"\"\n Make causal mask used for bi-directional self-attention.\n \"\"\"\n bsz, tgt_len = input_ids_shape\n mask = torch.full((tgt_len, tgt_len), torch.finfo(dtype).min, device=device)\n mask_cond = torch.arange(mask.size(-1), device=device)\n mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)\n mask = mask.to(dtype)\n return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len)\n\n\n# Copied from transformers.models.bart.modeling_bart._expand_mask\ndef _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):\n \"\"\"\n Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.\n \"\"\"\n bsz, src_len = mask.size()\n tgt_len = tgt_len if tgt_len is not None else src_len\n\n expanded_mask = mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)\n\n inverted_mask = 1.0 - expanded_mask\n\n return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)\n\n\nclass ParallelLlamaModel(nn.Module):\n \"\"\"\n Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`LlamaDecoderLayer`]\n\n Args:\n config: LlamaConfig\n \"\"\"\n\n def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.padding_idx = config.pad_token_id\n self.vocab_size = config.vocab_size\n embedding_kwargs = tp_utils.get_default_kwargs_for_parallel_embedding()\n if megatron_config is not None:\n assert embedding_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(embedding_kwargs, self.megatron_config)\n self.embed_tokens = tensor_parallel.VocabParallelEmbedding(\n num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, **embedding_kwargs\n )\n\n self.layers = nn.ModuleList(\n [ParallelLlamaDecoderLayer(config, megatron_config) for _ in range(config.num_hidden_layers)]\n )\n self.norm = ParallelLlamaRMSNorm(config, megatron_config)\n\n # Copied from transformers.models.bart.modeling_bart.BartDecoder._prepare_decoder_attention_mask\n def _prepare_decoder_attention_mask(self, attention_mask, input_shape, inputs_embeds):\n # create causal mask\n # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]\n combined_attention_mask = None\n if input_shape[-1] > 1:\n combined_attention_mask = _make_causal_mask(\n input_shape,\n inputs_embeds.dtype,\n device=inputs_embeds.device,\n )\n\n if attention_mask is not None:\n # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]\n expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to(\n inputs_embeds.device\n )\n combined_attention_mask = (\n expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask\n )\n\n return combined_attention_mask\n\n def forward(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | BaseModelOutputWithPast:\n \"\"\"\n\n Args:\n input_ids: input ids. shape (batch_size, seq_length)\n attention_mask: attention_mask. shape (batch_size, seq_length)\n position_ids: position ids. shape (batch_size, seq_length)\n\n Returns:\n\n \"\"\"\n batch_size, seq_length = input_ids.shape\n inputs_embeds = self.embed_tokens(input_ids)\n # embed positions\n\n attention_mask = self._prepare_decoder_attention_mask(attention_mask, (batch_size, seq_length), inputs_embeds)\n\n hidden_states = inputs_embeds\n\n for idx, decoder_layer in enumerate(self.layers):\n layer_outputs = decoder_layer(\n hidden_states,\n attention_mask=attention_mask,\n position_ids=position_ids,\n )\n\n hidden_states = layer_outputs\n\n hidden_states = self.norm(hidden_states)\n\n return hidden_states\n\n\nclass ParallelLlamaForCausalLM(nn.Module):\n def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.model = ParallelLlamaModel(config, megatron_config=megatron_config)\n self.vocab_size = config.vocab_size\n\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n\n self.lm_head = tensor_parallel.ColumnParallelLinear(\n input_size=config.hidden_size,\n output_size=config.vocab_size,\n bias=False,\n gather_output=False,\n skip_bias_add=False,\n **column_kwargs,\n )\n\n def forward(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n r\"\"\"\n Args:\n labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):\n Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,\n config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored\n (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.\n\n Returns:\n ```\"\"\"\n\n # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)\n outputs = self.model(\n input_ids=input_ids,\n attention_mask=attention_mask,\n position_ids=position_ids,\n )\n\n hidden_states = outputs\n logits = self.lm_head(hidden_states)[0]\n\n logits = tensor_parallel.gather_from_tensor_model_parallel_region(logits)\n\n logits = logits.float()\n return CausalLMOutputWithPast(\n loss=None,\n logits=logits,\n past_key_values=None,\n hidden_states=None,\n attentions=None,\n )\n\n\nfrom flash_attn.bert_padding import index_first_axis, pad_input, unpad_input # noqa: F401, E402\n\n\nclass ParallelLlamaModelRmPad(nn.Module):\n \"\"\"\n Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`LlamaDecoderLayer`]\n\n Args:\n config: LlamaConfig\n \"\"\"\n\n def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.padding_idx = config.pad_token_id\n self.vocab_size = config.vocab_size\n embedding_kwargs = tp_utils.get_default_kwargs_for_parallel_embedding()\n self.megatron_config = megatron_config\n if megatron_config is not None:\n assert embedding_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(embedding_kwargs, self.megatron_config)\n self.embed_tokens = tensor_parallel.VocabParallelEmbedding(\n num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, **embedding_kwargs\n )\n\n self.layers = nn.ModuleList(\n [ParallelLlamaDecoderLayerRmPad(config, megatron_config) for _ in range(config.num_hidden_layers)]\n )\n self.norm = ParallelLlamaRMSNorm(config, megatron_config)\n\n def forward(\n self,\n input_ids: torch.Tensor,\n position_ids: Optional[torch.LongTensor] = None,\n sequence_length: int = None,\n indices: torch.Tensor = None,\n cu_seqlens: int = None,\n max_seqlen_in_batch: int = None,\n ) -> tuple | BaseModelOutputWithPast:\n \"\"\"\n\n Args:\n input_ids: input ids. shape (1, totol_nnz)\n position_ids: position ids. shape (batch_size, seq_length)\n\n Returns:\n\n \"\"\"\n inputs_embeds = self.embed_tokens(input_ids) # (1, total_nnz) -> (1, total_nnz, hidden_size)\n\n # (1, total_nnz, hidden_size) -> (total_nnz, 1, hidden_size) -> (total_nnz // sp, 1, hidden_size)\n inputs_embeds = inputs_embeds.transpose(0, 1)\n if self.megatron_config.sequence_parallel:\n inputs_embeds = tensor_parallel.scatter_to_sequence_parallel_region(inputs_embeds)\n\n hidden_states = inputs_embeds\n for idx, decoder_layer in enumerate(self.layers):\n layer_outputs = decoder_layer(\n hidden_states,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n hidden_states = layer_outputs\n\n hidden_states = self.norm(hidden_states)\n\n return hidden_states\n\n\nclass ParallelLlamaForCausalLMRmPad(nn.Module):\n def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.megatron_config = megatron_config\n self.model = ParallelLlamaModelRmPad(config, megatron_config=megatron_config)\n self.vocab_size = config.vocab_size\n self._init_head(config)\n\n def _init_head(self, config):\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if self.megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n self.lm_head = tensor_parallel.ColumnParallelLinear(\n input_size=config.hidden_size,\n output_size=config.vocab_size,\n bias=False,\n gather_output=False,\n skip_bias_add=False,\n **column_kwargs,\n )\n\n def _forward_head(self, hidden_states):\n # all_gather from sequence parallel region is performed inside lm_head\n logits = self.lm_head(hidden_states)[0]\n logits = logits.float() # (total_nnz_padded, 1, vocab_size // tp)\n logits = tensor_parallel.gather_from_tensor_model_parallel_region(logits) # (total_nnz_padded, 1, vocab_size)\n return logits\n\n def forward(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n r\"\"\"\n Args:\n labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):\n Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,\n config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored\n (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.\n\n Returns:\n ```\"\"\"\n batch_size, sequence_length = input_ids.shape\n\n # remove padding here\n input_ids, indices, cu_seqlens, max_seqlen_in_batch, *_ = unpad_input(\n input_ids.unsqueeze(dim=-1), attention_mask\n ) # (total_nnz, 1)\n\n # pad input_ids to multiple of tp for all tp ranks\n # TODO: for better performance, the sp padding should be removed at each layer. Not sure the performance gap\n if self.megatron_config.sequence_parallel:\n input_ids = sp_utils.pad_to_sequence_parallel(input_ids)\n\n input_ids = input_ids.transpose(0, 1) # (1, total_nnz+pad)\n\n outputs = self.model(\n input_ids=input_ids,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n hidden_states = outputs\n\n logits = self._forward_head(hidden_states)\n\n # remove padding from sequence parallel\n if self.megatron_config.sequence_parallel:\n totol_nnz = cu_seqlens[-1]\n logits = logits[:totol_nnz] # (total_nnz_padded)\n\n logits = torch.squeeze(logits, dim=1) # remove the artificial batch dimension\n # add removed padding back\n logits = pad_input(\n logits, indices, batch_size, seqlen=sequence_length\n ) # (batch_size, sequence_length, vocab_size)\n\n return CausalLMOutputWithPast(\n loss=None,\n logits=logits,\n past_key_values=None,\n hidden_states=None,\n attentions=None,\n )\n\n\nclass ParallelLlamaForValueRmPad(ParallelLlamaForCausalLMRmPad):\n def _init_head(self, config):\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if self.megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n self.lm_head = nn.Linear(in_features=config.hidden_size, out_features=1, bias=False)\n # lm_head is effectively the same as sequence parallel\n sp_utils.mark_parameter_as_sequence_parallel(self.lm_head.weight)\n\n def _forward_head(self, hidden_states):\n logits = self.lm_head(hidden_states) # (total_nnz_padded // tp, 1, 1)\n logits = logits.float()\n if self.megatron_config.sequence_parallel:\n logits = tensor_parallel.gather_from_sequence_parallel_region(logits, tensor_parallel_output_grad=False)\n return logits\n\n def forward(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n output = super().forward(input_ids, attention_mask, position_ids)\n output.logits = torch.squeeze(output.logits, dim=-1)\n return output\n\n\n\"\"\"\nSupport pipeline parallelism\n\"\"\"\n\n\nclass ParallelLlamaModelRmPadPP(nn.Module):\n \"\"\"\n Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`LlamaDecoderLayer`]\n This model definition supports pipeline parallelism. To support pp and vpp,\n - This model only contains layer in this pp stage and vpp chunk\n - When calling get_model in Megatron, this rank will instantiate all the vpp chunks in this pp.\n Args:\n config: LlamaConfig\n \"\"\"\n\n def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig, pre_process, post_process):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.padding_idx = config.pad_token_id\n self.vocab_size = config.vocab_size\n self.pre_process = pre_process\n self.post_process = post_process\n self.megatron_config = megatron_config\n embedding_kwargs = tp_utils.get_default_kwargs_for_parallel_embedding()\n if megatron_config is not None:\n assert embedding_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(embedding_kwargs, self.megatron_config)\n if pre_process:\n self.embed_tokens = tensor_parallel.VocabParallelEmbedding(\n num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, **embedding_kwargs\n )\n else:\n self.embed_tokens = None\n\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n pp_size = megatron_config.pipeline_model_parallel_size\n self.num_layer_per_pp = config.num_hidden_layers // pp_size\n vpp_size = megatron_config.virtual_pipeline_model_parallel_size\n vpp_rank = mpu.get_virtual_pipeline_model_parallel_rank()\n\n if vpp_size is not None:\n self.layers = nn.ModuleList()\n self.num_layer_vpp_chunk = self.num_layer_per_pp // vpp_size\n self.num_layer_this_model = self.num_layer_vpp_chunk\n offset = vpp_rank * (config.num_hidden_layers // vpp_size) + (pp_rank * self.num_layer_vpp_chunk)\n else:\n self.num_layer_this_model = self.num_layer_per_pp\n offset = pp_rank * self.num_layer_per_pp\n\n self.layers = nn.ModuleList()\n for i in range(self.num_layer_this_model):\n layer = ParallelLlamaDecoderLayerRmPad(config, megatron_config, layer_idx=offset + i)\n self.layers.add_module(f\"{i}\", layer)\n\n if post_process:\n self.norm = ParallelLlamaRMSNorm(config, megatron_config)\n else:\n self.norm = None\n\n def set_input_tensor(self, input_tensor):\n \"\"\"Set input tensor to be used instead of forward()'s input.\n\n When doing pipeline parallelism the input from the previous\n stage comes from communication, not from the input, so the\n model's forward_step_func won't have it. This function is thus\n used by internal code to bypass the input provided by the\n forward_step_func\"\"\"\n self.input_tensor = input_tensor\n\n def forward(\n self,\n input_ids: torch.Tensor,\n position_ids: Optional[torch.LongTensor] = None,\n sequence_length: int = None,\n indices: torch.Tensor = None,\n cu_seqlens: int = None,\n max_seqlen_in_batch: int = None,\n ) -> tuple | BaseModelOutputWithPast:\n \"\"\"\n\n Args:\n input_ids: input ids. shape (1, totol_nnz)\n position_ids: position ids. shape (batch_size, seq_length)\n\n Returns:\n\n \"\"\"\n if self.pre_process:\n inputs_embeds = self.embed_tokens(input_ids) # (1, total_nnz) -> (1, total_nnz, hidden_size)\n\n # vocab parallel embedding will not do sequence parallel reduce-scatter in open source megatron\n # so need to deal with it by handle here:\n # (1, total_nnz, hidden_size) -> (total_nnz, 1, hidden_size) -> (total_nnz // sp, 1, hidden_size)\n inputs_embeds = inputs_embeds.transpose(0, 1)\n if self.megatron_config.sequence_parallel:\n inputs_embeds = tensor_parallel.scatter_to_sequence_parallel_region(inputs_embeds)\n\n hidden_states = inputs_embeds\n else:\n # self.hidden_states should be passed by Megatron\n hidden_states = self.input_tensor\n\n for idx, decoder_layer in enumerate(self.layers):\n layer_outputs = decoder_layer(\n hidden_states,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n hidden_states = layer_outputs\n\n if self.post_process:\n hidden_states = self.norm(hidden_states)\n\n return hidden_states\n\n\nclass ParallelLlamaForCausalLMRmPadPP(nn.Module):\n def __init__(\n self,\n config: LlamaConfig,\n megatron_config: ModelParallelConfig,\n pre_process,\n post_process,\n share_embeddings_and_output_weights=False,\n ):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.megatron_config = megatron_config\n self.model = ParallelLlamaModelRmPadPP(\n config, megatron_config=megatron_config, pre_process=pre_process, post_process=post_process\n )\n assert share_embeddings_and_output_weights is False, (\n \"Llama Model not supports sharing embedding and output weights\"\n )\n self.share_embeddings_and_output_weights = share_embeddings_and_output_weights\n self.vocab_size = config.vocab_size\n self.pre_process = pre_process\n self.post_process = post_process\n if post_process:\n self._init_head(config)\n\n def set_input_tensor(self, input_tensor):\n \"\"\"Set input tensor to be used instead of forward()'s input.\n\n When doing pipeline parallelism the input from the previous\n stage comes from communication, not from the input, so the\n model's forward_step_func won't have it. This function is thus\n used by internal code to bypass the input provided by the\n forward_step_func\"\"\"\n assert len(input_tensor) == 1\n self.model.set_input_tensor(input_tensor[0])\n\n def _init_head(self, config):\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if self.megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n self.lm_head = tensor_parallel.ColumnParallelLinear(\n input_size=config.hidden_size,\n output_size=config.vocab_size,\n bias=False,\n gather_output=False,\n skip_bias_add=False,\n **column_kwargs,\n )\n\n def _forward_head(self, hidden_states):\n # all_gather from sequence parallel region is performed inside lm_head\n # logits shape before forward_head hidden_states.shape: [4, 32, 4096]\n logits = self.lm_head(hidden_states)[0]\n # logits shape after forward_head logits.shape: [8, 32, 8]\n logits = logits.float() # (total_nnz_padded, 1, vocab_size // tp)\n return logits\n\n def forward(\n self,\n # original input\n *,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n r\"\"\"\n Args:\n labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):\n Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,\n config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored\n (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.\n\n Returns:\n ```\"\"\"\n\n # Note that input_ids, attention_mask and position_ids should be passed to every pp layer.\n # In the first pp, input_ids will be used, in other pp layers hidden_states will be used inside self.model\n batch_size, sequence_length = input_ids.shape\n # remove padding here\n input_ids_rmpad, indices, cu_seqlens, max_seqlen_in_batch, *_ = unpad_input(\n input_ids.unsqueeze(dim=-1), attention_mask\n ) # (total_nnz, 1)\n\n # pad input_ids to multiple of tp for all tp ranks\n # TODO: for better performance, the sp padding should be removed at each layer. Not sure the performance gap\n if self.megatron_config.sequence_parallel:\n input_ids_rmpad = sp_utils.pad_to_sequence_parallel(input_ids_rmpad)\n\n input_ids_rmpad = input_ids_rmpad.transpose(0, 1) # (1, total_nnz+pad)\n\n outputs = self.model(\n input_ids=input_ids_rmpad,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n if self.post_process:\n hidden_states = outputs\n # print(f'hidden_states.shape = {hidden_states.shape}') # torch.Size([4, 32, 4096])\n logits = self._forward_head(hidden_states)\n logits = torch.squeeze(logits, dim=1) # remove the artificial batch dimension # torch.Size([8, 32, 16])\n\n # remove padding from sequence parallel\n if self.megatron_config.sequence_parallel:\n totol_nnz = cu_seqlens[-1]\n logits = logits[:totol_nnz] # (total_nnz_padded)\n # add removed padding back. If input is already rmpad, we let the caller pad_input\n logits = pad_input(\n logits, indices, batch_size, seqlen=sequence_length\n ) # (batch_size, sequence_length, vocab_size)\n\n return CausalLMOutputWithPast(\n loss=None,\n logits=logits,\n past_key_values=None,\n hidden_states=None,\n attentions=None,\n )\n else:\n return outputs\n\n\nclass ParallelLlamaForValueRmPadPP(ParallelLlamaForCausalLMRmPadPP):\n def _init_head(self, config):\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if self.megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n self.lm_head = nn.Linear(in_features=config.hidden_size, out_features=1, bias=False)\n # lm_head is effectively the same as sequence parallel\n sp_utils.mark_parameter_as_sequence_parallel(self.lm_head.weight)\n\n def _forward_head(self, hidden_states):\n logits = self.lm_head(hidden_states) # (total_nnz_padded // tp, 1, 1)\n logits = logits.float()\n if self.megatron_config.sequence_parallel:\n logits = tensor_parallel.gather_from_sequence_parallel_region(logits, tensor_parallel_output_grad=False)\n return logits\n\n def forward(\n self,\n *,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n output = super().forward(input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids)\n if self.post_process:\n output.logits = torch.squeeze(output.logits, dim=-1)\n return output\n else:\n return output\n"}17{"file_name": "verl__models__mcore__config_converter.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n# convert huggingface config to mcore transformer config\n\n\nimport warnings\nfrom typing import TypeVar\n\nimport torch\nimport torch.nn.functional as F\nfrom megatron.core import parallel_state as mpu\nfrom megatron.core.transformer import MLATransformerConfig, TransformerConfig\nfrom transformers import PretrainedConfig\n\nT = TypeVar(\"T\", bound=TransformerConfig)\n\n\ndef _get_base_transformer_config(\n hf_config: PretrainedConfig, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> dict:\n \"\"\"\n Create a base TransformerConfig with common parameters across different model architectures.\n TODO: (ycl) use dataclass or converter config?\n\n Args:\n hf_config: HuggingFace model configuration\n dtype: Data type for the model\n override_transformer_config_kwargs: Additional parameters to override defaults\n\n Returns:\n TransformerConfig with common parameters\n \"\"\"\n\n # Common parallel state parameters\n overlap_p2p_comm = (\n mpu.get_virtual_pipeline_model_parallel_world_size() is not None\n and mpu.get_virtual_pipeline_model_parallel_world_size() > 1\n )\n batch_p2p_comm = False\n\n # Base configuration with common parameters\n base_config = {\n # Model architecture parameters\n \"num_layers\": hf_config.num_hidden_layers,\n \"hidden_size\": hf_config.hidden_size,\n \"num_attention_heads\": hf_config.num_attention_heads,\n \"num_query_groups\": hf_config.num_key_value_heads,\n \"ffn_hidden_size\": hf_config.intermediate_size,\n \"attention_dropout\": hf_config.attention_dropout,\n \"hidden_dropout\": getattr(hf_config, \"hidden_dropout\", 0.0),\n \"kv_channels\": getattr(hf_config, \"head_dim\", None),\n \"layernorm_epsilon\": hf_config.rms_norm_eps,\n \"add_bias_linear\": True,\n # Activation and normalization\n \"activation_func\": F.silu,\n \"normalization\": \"RMSNorm\",\n \"gated_linear_unit\": True,\n # Data types\n \"pipeline_dtype\": dtype,\n \"params_dtype\": dtype,\n \"bf16\": dtype is torch.bfloat16,\n # Parallel configuration\n \"tensor_model_parallel_size\": mpu.get_tensor_model_parallel_world_size(),\n \"pipeline_model_parallel_size\": mpu.get_pipeline_model_parallel_world_size(),\n \"expert_model_parallel_size\": mpu.get_expert_model_parallel_world_size(),\n \"expert_tensor_parallel_size\": mpu.get_expert_tensor_parallel_world_size(),\n \"virtual_pipeline_model_parallel_size\": mpu.get_virtual_pipeline_model_parallel_world_size(),\n \"context_parallel_size\": mpu.get_context_parallel_world_size(),\n \"overlap_p2p_comm\": overlap_p2p_comm,\n \"batch_p2p_comm\": batch_p2p_comm,\n \"sequence_parallel\": mpu.get_tensor_model_parallel_world_size() > 1,\n # Common settings\n \"variable_seq_lengths\": True,\n \"masked_softmax_fusion\": True,\n \"moe_token_dispatcher_type\": \"alltoall\",\n }\n\n # Update with any provided overrides\n # override_transformer_config_kwargs as kwargs shall never be none\n base_config.update(override_transformer_config_kwargs)\n\n return base_config\n\n\ndef _get_mla_transformer_config(\n hf_config: PretrainedConfig, mla_rope_config: dict, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> dict:\n \"\"\"\n Create a MLATransformerConfig with common parameters across different model architectures.\n This is specifically for MLA models like DeepseekV3.\n\n Args:\n hf_config: HuggingFace model configuration\n mla_rope_config: MLA specific RoPE configuration\n dtype: Data type for the model\n override_transformer_config_kwargs: Additional parameters to override defaults\n\n Returns:\n MLATransformerConfig with common parameters\n \"\"\"\n base_config = _get_base_transformer_config(hf_config=hf_config, dtype=dtype, **override_transformer_config_kwargs)\n mla_config = {\n # MLA specific parameters\n \"q_lora_rank\": hf_config.q_lora_rank,\n \"kv_lora_rank\": hf_config.kv_lora_rank,\n \"qk_head_dim\": hf_config.qk_nope_head_dim,\n \"qk_pos_emb_head_dim\": hf_config.qk_rope_head_dim,\n \"v_head_dim\": hf_config.v_head_dim,\n \"rotary_base\": hf_config.rope_theta,\n \"rotary_scaling_factor\": mla_rope_config[\"factor\"],\n \"rope_type\": mla_rope_config[\"type\"],\n \"max_position_embeddings\": mla_rope_config[\"original_max_position_embeddings\"],\n \"beta_fast\": mla_rope_config[\"beta_fast\"],\n \"beta_slow\": mla_rope_config[\"beta_slow\"],\n \"mscale\": mla_rope_config[\"mscale\"],\n \"mscale_all_dim\": mla_rope_config[\"mscale_all_dim\"],\n }\n\n base_config.update(mla_config)\n return base_config\n\n\ndef check_and_construct_configs(original_config: dict, cls: type[T]) -> T:\n \"\"\"\n Check and disable incompatible configurations for older Megatron version.\n\n Args:\n original_config (dict): The original model configuration.\n\n Returns:\n dict: The updated model configuration with incompatible settings disabled.\n \"\"\"\n removed_keys = []\n for key in original_config.keys():\n if not hasattr(cls, key):\n removed_keys.append(key)\n if removed_keys:\n warnings.warn(\n f\"The following keys are not supported in the current Megatron version and will be removed: {removed_keys}\",\n stacklevel=2,\n )\n for key in removed_keys:\n original_config.pop(key)\n\n original_config = mapping_string_to_attn_backend(original_config)\n if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0:\n print(f\"Overridden {cls.__name__} init config: {original_config}\")\n return cls(**original_config)\n\n\ndef hf_to_mcore_config_dense(\n hf_config: PretrainedConfig, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> TransformerConfig:\n # for LlamaForCausalLM or Qwen2ForCausalLM\n qkv_bias = True if \"Qwen2\" in hf_config.architectures[0] else getattr(hf_config, \"attention_bias\", False)\n qk_layernorm = True if \"Qwen3\" in hf_config.architectures[0] else False\n\n args: dict = _get_base_transformer_config(\n hf_config=hf_config,\n dtype=dtype,\n use_cpu_initialization=False,\n add_bias_linear=False,\n add_qkv_bias=qkv_bias,\n qk_layernorm=qk_layernorm,\n )\n # override_transformer_config_kwargs as kwargs shall never be none\n args.update(override_transformer_config_kwargs)\n return check_and_construct_configs(args, TransformerConfig)\n\n\ndef hf_to_mcore_config_qwen2moe(\n hf_config: PretrainedConfig, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> TransformerConfig:\n args: dict = _get_base_transformer_config(\n hf_config=hf_config,\n dtype=dtype,\n use_cpu_initialization=False,\n add_bias_linear=False,\n layernorm_epsilon=hf_config.rms_norm_eps,\n # MoE specific\n moe_ffn_hidden_size=hf_config.moe_intermediate_size,\n moe_router_bias_update_rate=0.001,\n moe_router_topk=hf_config.num_experts_per_tok,\n num_moe_experts=hf_config.num_experts,\n moe_shared_expert_intermediate_size=hf_config.shared_expert_intermediate_size,\n moe_aux_loss_coeff=hf_config.router_aux_loss_coef,\n # moe_aux_loss_coeff=0.0,\n moe_router_load_balancing_type=\"none\", # turn off aux_loss as it hurts perf in RL\n moe_shared_expert_overlap=True,\n moe_grouped_gemm=True,\n moe_router_score_function=\"softmax\",\n # Other optimizations\n persist_layer_norm=True,\n bias_activation_fusion=True,\n bias_dropout_fusion=True,\n # Qwen specific\n moe_router_pre_softmax=True,\n add_qkv_bias=True,\n )\n # override_transformer_config_kwargs as kwargs shall never be none\n args.update(override_transformer_config_kwargs)\n return check_and_construct_configs(args, TransformerConfig)\n\n\ndef hf_to_mcore_config_mixtral(\n hf_config: PretrainedConfig, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> TransformerConfig:\n args: dict = _get_base_transformer_config(\n hf_config=hf_config,\n dtype=dtype,\n use_cpu_initialization=False,\n add_bias_linear=False,\n layernorm_epsilon=hf_config.rms_norm_eps,\n # MoE specific\n num_moe_experts=hf_config.num_local_experts,\n moe_aux_loss_coeff=hf_config.router_aux_loss_coef,\n moe_router_topk=hf_config.num_experts_per_tok,\n moe_router_pre_softmax=True,\n moe_router_load_balancing_type=\"none\", # turn off aux_loss as it hurts perf in RL\n moe_router_score_function=\"softmax\",\n moe_shared_expert_intermediate_size=None, # mixtral has no shared expert\n moe_shared_expert_overlap=False, # mixtral has no shared expert\n moe_ffn_hidden_size=hf_config.intermediate_size,\n moe_router_bias_update_rate=0.001,\n # moe_permute_fusion=True, # need TE 2.1+\n moe_grouped_gemm=True,\n # Other optimizations\n persist_layer_norm=True,\n apply_rope_fusion=True,\n bias_activation_fusion=True,\n bias_dropout_fusion=True,\n )\n # override_transformer_config_kwargs as kwargs shall never be none\n args.update(override_transformer_config_kwargs)\n return check_and_construct_configs(args, TransformerConfig)\n\n\ndef hf_to_mcore_config_qwen3moe(\n hf_config: PretrainedConfig, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> TransformerConfig:\n args: dict = _get_base_transformer_config(\n hf_config=hf_config,\n dtype=dtype,\n use_cpu_initialization=False,\n add_bias_linear=False,\n layernorm_epsilon=hf_config.rms_norm_eps,\n # MoE specific\n moe_ffn_hidden_size=hf_config.moe_intermediate_size,\n moe_router_bias_update_rate=0.001,\n moe_router_topk=hf_config.num_experts_per_tok,\n num_moe_experts=hf_config.num_experts,\n moe_aux_loss_coeff=hf_config.router_aux_loss_coef,\n # moe_aux_loss_coeff=0.0,\n moe_router_load_balancing_type=\"none\", # turn off aux_loss as it hurts perf in RL\n moe_grouped_gemm=True,\n moe_router_score_function=\"softmax\",\n # Other optimizations\n persist_layer_norm=True,\n bias_activation_fusion=True,\n bias_dropout_fusion=True,\n # Qwen specific\n moe_router_pre_softmax=False,\n qk_layernorm=True,\n )\n # override_transformer_config_kwargs as kwargs shall never be none\n args.update(override_transformer_config_kwargs)\n return check_and_construct_configs(args, TransformerConfig)\n\n\ndef hf_to_mcore_config_dpskv3(\n hf_config: PretrainedConfig, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> MLATransformerConfig:\n # DeepseekV3ForCausalLM\n from megatron.core.config import set_experimental_flag\n from megatron.core.transformer.enums import AttnBackend\n\n set_experimental_flag(True)\n\n from .patch import apply_patch\n\n apply_patch()\n\n mla_rope_config = {\n \"beta_fast\": 32,\n \"beta_slow\": 1,\n \"factor\": 1,\n \"mscale\": 1.0,\n \"mscale_all_dim\": 1.0,\n \"original_max_position_embeddings\": 4096,\n \"type\": \"rope\",\n }\n if \"rope_scaling\" in hf_config and hf_config.rope_scaling is not None:\n mla_rope_config.update(hf_config.rope_scaling)\n moe_layer_freq = [1] * hf_config.num_hidden_layers\n for i in range(min(hf_config.first_k_dense_replace, hf_config.num_hidden_layers)):\n moe_layer_freq[i] = 0\n\n # disable MTP and quantization for now\n if \"num_nextn_predict_layers\" in hf_config:\n assert hf_config.num_nextn_predict_layers == 0, (\n \"MTP is not supported for now, please modify the config.json to set num_nextn_predict_layers to 0\"\n )\n assert \"quantization_config\" not in hf_config or not hf_config.quantization_config, (\n \"quantization is not supported for now, please modify the config.json to remove quantization_config\"\n )\n\n args: dict = _get_mla_transformer_config(\n hf_config=hf_config,\n mla_rope_config=mla_rope_config,\n dtype=dtype,\n # Additional parameters\n use_cpu_initialization=False,\n add_bias_linear=False,\n attention_backend=AttnBackend.fused,\n qk_layernorm=True,\n # Standard MoE parameters\n moe_ffn_hidden_size=hf_config.moe_intermediate_size,\n moe_token_dispatcher_type=\"alltoall\",\n moe_router_bias_update_rate=0.001,\n moe_router_enable_expert_bias=True,\n moe_router_topk=hf_config.num_experts_per_tok,\n num_moe_experts=hf_config.n_routed_experts,\n moe_shared_expert_intermediate_size=hf_config.moe_intermediate_size * hf_config.n_shared_experts,\n moe_aux_loss_coeff=getattr(hf_config, \"aux_loss_alpha\", 0.001),\n moe_router_load_balancing_type=\"seq_aux_loss\",\n moe_shared_expert_overlap=True,\n # moe_permute_fusion=True, # need TE 2.1+\n moe_grouped_gemm=True,\n moe_router_score_function=\"sigmoid\",\n moe_router_pre_softmax=True,\n moe_router_topk_scaling_factor=hf_config.routed_scaling_factor,\n moe_layer_freq=moe_layer_freq,\n # mcore 0.12 moe\n moe_router_dtype=\"fp64\",\n disable_bf16_reduced_precision_matmul=True,\n # Other optimizations\n # deallocate_pipeline_outputs=True,\n # gradient_accumulation_fusion=True,\n persist_layer_norm=True,\n bias_activation_fusion=True,\n bias_dropout_fusion=True,\n )\n # override_transformer_config_kwargs as kwargs shall never be none\n args.update(override_transformer_config_kwargs)\n transformer_config = check_and_construct_configs(args, MLATransformerConfig)\n # MTP\n if \"num_nextn_predict_layers\" in hf_config:\n transformer_config.mtp_num_layers = hf_config.num_nextn_predict_layers\n transformer_config.mtp_loss_scaling_factor = 0.1\n\n return transformer_config\n\n\ndef hf_to_mcore_config_qwen2_5_vl(\n hf_config: PretrainedConfig, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> TransformerConfig:\n # Qwen2_5_VLForConditionalGeneration\n\n args = _get_base_transformer_config(\n hf_config=hf_config,\n dtype=dtype,\n add_bias_linear=False,\n # qwen specific\n add_qkv_bias=True,\n mrope_section=hf_config.rope_scaling[\"mrope_section\"],\n )\n # override_transformer_config_kwargs as kwargs shall never be none\n args.update(override_transformer_config_kwargs)\n args = mapping_string_to_attn_backend(args)\n return TransformerConfig(**args)\n\n\ndef hf_to_mcore_config_llama4(\n hf_config: PretrainedConfig, dtype: torch.dtype, **override_transformer_config_kwargs\n) -> TransformerConfig:\n # Llama4ForConditionalGeneration\n raise NotImplementedError(\"Llama4ForConditionalGeneration is not supported yet\")\n\n\ndef mapping_string_to_attn_backend(args: dict) -> dict:\n if \"attention_backend\" in args and isinstance(args[\"attention_backend\"], str):\n from megatron.core.transformer.enums import AttnBackend\n\n args[\"attention_backend\"] = AttnBackend[args[\"attention_backend\"]]\n return args\n"}18{"file_name": "verl__models__mcore__loader.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport time\n\nimport torch\nimport torch.distributed as dist\n\nfrom verl.utils.device import get_device_id, get_torch_device\n\nfrom .saver import _megatron_calc_global_rank\n\n\ndef _megatron_calc_layer_map(config):\n \"\"\"Calculate the mapping of global layer_idx to local layer_idx\n Returns:\n layer_map (Dict: int -> tuple(int, int, int)):\n mapping from the global layer index to\n a tuple of (pp_rank, virtual_pp_rank, layer_idx inside model)\n \"\"\"\n from megatron.core import mpu\n\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n\n layer_map = dict()\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n for pp_rank_idx in range(pp_size):\n for virtual_pp_rank_idx in range(virtual_pp_size):\n layer_offset = (\n virtual_pp_rank_idx * (config.num_hidden_layers // virtual_pp_size) + pp_rank_idx * num_layers_per_model\n )\n for layer_idx in range(num_layers_per_model):\n layer_map[layer_offset + layer_idx] = (\n pp_rank_idx,\n virtual_pp_rank_idx,\n layer_idx,\n )\n return layer_map\n\n\ndef load_state_dict_to_megatron_gptmodel(state_dict, wrapped_models, config, params_dtype, is_value_model=False):\n \"\"\"Load merged state_dict to sharded Megatron module in training.\"\"\"\n from megatron.core import DistributedDataParallel as LocalDDP\n from megatron.core import mpu\n from megatron.core.transformer.module import Float16Module\n from torch.nn.parallel import DistributedDataParallel as torchDDP\n\n from verl.utils.logger import print_rank_0\n from verl.utils.megatron_utils import unwrap_model\n\n start_time = time.time()\n\n def _get_gpt_model(model):\n return model\n\n def broadcast_params(module):\n for param in module.parameters():\n torch.distributed.broadcast(\n param.data, src=mpu.get_data_parallel_src_rank(), group=mpu.get_data_parallel_group()\n )\n\n dp_rank = mpu.get_data_parallel_rank()\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n cp_rank = mpu.get_context_parallel_rank()\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=0, cp_rank=cp_rank)\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n mp_group = mpu.get_model_parallel_group()\n\n if torch.distributed.get_rank() == src_rank:\n assert mp_group.rank() == 0, f\"mp_rank:[{mp_group.rank}] != 0 on rank #0\"\n assert pp_rank == 0, f\"pp_rank:[{pp_rank}] != 0 on rank #0\"\n assert dp_rank == 0, f\"dp_rank:[{dp_rank}] != 0 on rank #0\"\n\n if not isinstance(wrapped_models, list | tuple):\n wrapped_models = list(wrapped_models)\n\n assert len(wrapped_models) == virtual_pp_size\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n models = [None] * len(wrapped_models)\n\n for i, wrapped_model in enumerate(wrapped_models):\n models[i] = unwrap_model(wrapped_model, (torchDDP, LocalDDP, Float16Module))\n gpt_model_module = _get_gpt_model(models[i])\n assert len(gpt_model_module.decoder.layers) == num_layers_per_model\n\n def _broadcast_tensor(tensor, name) -> torch.Tensor:\n \"\"\"broadcast tensor from rank0 across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n if torch.distributed.get_rank() == src_rank:\n if name in state_dict:\n weight = state_dict[name]\n tensor_shape = weight.shape\n else:\n tensor_shape = None\n else:\n weight = None\n tensor_shape = None\n\n obj_list = [tensor_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n tensor_shape = obj_list[0]\n\n if tensor_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tensor:[{name}] not in state_dict, skip load\")\n return\n\n if tensor is None:\n tensor = torch.empty(\n tensor_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n if torch.distributed.get_rank() == src_rank:\n tensor.data.copy_(weight)\n dist.broadcast(tensor, src=src_rank, group=mp_group)\n\n def _broadcast_tp_shard_tensor_vocab(tensor, name, chunk_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n if torch.distributed.get_rank() == src_rank:\n if name in state_dict:\n full_weight = state_dict[name]\n\n if mutate_func is not None:\n full_weight = mutate_func(full_weight)\n tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)\n chunk_shape = tensor_chunk[0].shape\n else:\n chunk_shape = None\n else:\n chunk_shape = None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{name}] not in state_dict, skip loading\")\n return\n\n if tensor is None:\n sync_tensor = torch.empty(\n chunk_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n else:\n assert tensor.shape == chunk_shape, (\n f\"rank #{torch.distributed.get_rank()} tensor {name} shape {tensor.shape} != {chunk_shape}\"\n )\n sync_tensor = torch.empty_like(tensor, device=get_device_id(), requires_grad=False)\n\n for i in range(tp_size):\n if torch.distributed.get_rank() == src_rank:\n sync_tensor.data.copy_(tensor_chunk[i])\n dist.broadcast(sync_tensor, src=src_rank, group=mp_group)\n if (i == tp_rank) and (tensor is not None):\n tensor.data.copy_(sync_tensor)\n\n def _broadcast_tp_shard_tensor(tensor, name, chunk_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n if torch.distributed.get_rank() == src_rank:\n if name in state_dict:\n full_weight = state_dict[name]\n if mutate_func is not None:\n full_weight = mutate_func(full_weight)\n tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)\n chunk_shape = tensor_chunk[0].shape\n else:\n chunk_shape = None\n else:\n chunk_shape = None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{name}] not in state_dict, skip loading\")\n return\n\n if tensor is None:\n sync_tensor = torch.empty(\n chunk_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n else:\n assert tensor.shape == chunk_shape, (\n f\"rank #{torch.distributed.get_rank()} tensor {name} shape {tensor.shape} != {chunk_shape}\"\n )\n sync_tensor = torch.empty_like(tensor, device=get_device_id(), requires_grad=False)\n\n for i in range(tp_size):\n if torch.distributed.get_rank() == src_rank:\n sync_tensor.data.copy_(tensor_chunk[i])\n dist.broadcast(sync_tensor, src=src_rank, group=mp_group)\n if (i == tp_rank) and (tensor is not None):\n tensor.data.copy_(sync_tensor)\n\n def _broadcast_tp_shard_tensor_gate_up(tensor, gate_name, up_name) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n if torch.distributed.get_rank() == src_rank:\n gate_weight = state_dict[gate_name]\n up_weight = state_dict[up_name]\n new_gate_up_weight = torch.empty(\n config.intermediate_size * 2, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n for i in range(tp_size):\n intermediate_size_tp = config.intermediate_size // tp_size\n gate_weight_tp = gate_weight[i * intermediate_size_tp : (i + 1) * intermediate_size_tp]\n up_weight_tp = up_weight[i * intermediate_size_tp : (i + 1) * intermediate_size_tp]\n new_gate_up_weight[intermediate_size_tp * 2 * i : intermediate_size_tp * 2 * (i + 1)].copy_(\n torch.cat([gate_weight_tp, up_weight_tp], dim=0)\n )\n\n tensor_chunk = torch.chunk(new_gate_up_weight, tp_size, dim=0)\n chunk_shape = tensor_chunk[0].shape\n else:\n chunk_shape = None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{gate_name, up_name}] not in state_dict, skip loading\")\n return\n\n if tensor is None:\n sync_tensor = torch.empty(\n chunk_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n else:\n assert tensor.shape == chunk_shape, (\n f\"rank #{torch.distributed.get_rank() == src_rank:} tensor {gate_name, up_name} shape \"\n f\"{tensor.shape} != {chunk_shape}\"\n )\n sync_tensor = torch.empty_like(tensor, device=get_device_id(), requires_grad=False)\n\n for i in range(tp_size):\n if torch.distributed.get_rank() == src_rank:\n sync_tensor.data.copy_(tensor_chunk[i])\n dist.broadcast(sync_tensor, src=src_rank, group=mp_group)\n if (i == tp_rank) and (tensor is not None):\n tensor.data.copy_(sync_tensor)\n\n def _broadcast_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name, bias=False) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n if torch.distributed.get_rank() == src_rank:\n assert q_name in state_dict and k_name in state_dict and v_name in state_dict\n full_weight_q = state_dict[q_name]\n full_weight_k = state_dict[k_name]\n full_weight_v = state_dict[v_name]\n\n hidden_size_per_head = getattr(config, \"head_dim\", config.hidden_size // config.num_attention_heads)\n\n if config.num_key_value_heads >= tp_size:\n q_size_tp = hidden_size_per_head * config.num_attention_heads // tp_size\n kv_size_tp = hidden_size_per_head * config.num_key_value_heads // tp_size\n total_size = q_size_tp + 2 * kv_size_tp\n sizes = [total_size * tp_size]\n if not bias:\n sizes.append(config.hidden_size)\n new_weight_qkv = torch.empty(*sizes, dtype=params_dtype, device=get_device_id())\n for i in range(tp_size):\n q_part = full_weight_q[i * q_size_tp : (i + 1) * q_size_tp]\n k_part = full_weight_k[i * kv_size_tp : (i + 1) * kv_size_tp]\n v_part = full_weight_v[i * kv_size_tp : (i + 1) * kv_size_tp]\n num_query_groups_per_partition = models[0].config.num_query_groups // tp_size\n new_weight_qkv_this_tp = new_weight_qkv[i * total_size : (i + 1) * total_size]\n q_part_per_head = torch.chunk(q_part, num_query_groups_per_partition, dim=0)\n k_part_per_head = torch.chunk(k_part, num_query_groups_per_partition, dim=0)\n v_part_per_head = torch.chunk(v_part, num_query_groups_per_partition, dim=0)\n total_size_per_head = total_size // num_query_groups_per_partition\n for j in range(num_query_groups_per_partition):\n new_weight_qkv_this_tp[j * total_size_per_head : (j + 1) * total_size_per_head].copy_(\n torch.cat([q_part_per_head[j], k_part_per_head[j], v_part_per_head[j]], dim=0)\n )\n\n else:\n q_size_tp = hidden_size_per_head * config.num_attention_heads // tp_size\n kv_size_tp = hidden_size_per_head\n total_size = q_size_tp + 2 * kv_size_tp\n sizes = [total_size * tp_size]\n if not bias:\n sizes.append(config.hidden_size)\n new_weight_qkv = torch.empty(*sizes, dtype=params_dtype, device=get_device_id())\n for i in range(tp_size):\n q_part = full_weight_q[i * q_size_tp : (i + 1) * q_size_tp]\n start_idx = i * config.num_key_value_heads // tp_size * hidden_size_per_head\n end_idx = (i * config.num_key_value_heads // tp_size + 1) * hidden_size_per_head\n k_part = full_weight_k[start_idx:end_idx]\n v_part = full_weight_v[start_idx:end_idx]\n new_weight_qkv_this_tp = new_weight_qkv[i * total_size : (i + 1) * total_size]\n q_part_per_head = torch.chunk(q_part, config.num_attention_heads, dim=0)\n k_part_per_head = torch.chunk(k_part, config.num_attention_heads, dim=0)\n v_part_per_head = torch.chunk(v_part, config.num_attention_heads, dim=0)\n total_size_per_head = total_size // config.num_attention_heads\n for j in range(config.num_attention_heads):\n new_weight_qkv_this_tp[j * total_size_per_head : (j + 1) * total_size_per_head].copy_(\n torch.cat([q_part_per_head[j], k_part_per_head[j], v_part_per_head[j]], dim=0)\n )\n\n tensor_chunk = torch.chunk(new_weight_qkv, tp_size, dim=0)\n chunk_shape = tensor_chunk[0].shape\n else:\n chunk_shape = None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{q_name, k_name, v_name}] not in state_dict, skip loading\")\n return\n\n if tensor is None:\n sync_tensor = torch.empty(\n chunk_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n else:\n assert tensor.shape == chunk_shape, (\n f\"rank #{torch.distributed.get_rank()} tensor {q_name} shape {tensor.shape} != {chunk_shape}\"\n )\n sync_tensor = torch.empty_like(tensor, device=get_device_id(), requires_grad=False)\n\n for i in range(tp_size):\n if torch.distributed.get_rank() == src_rank:\n sync_tensor.data.copy_(tensor_chunk[i])\n dist.broadcast(sync_tensor, src=src_rank, group=mp_group)\n if (i == tp_rank) and (tensor is not None):\n tensor.data.copy_(sync_tensor)\n\n if dp_rank == 0:\n # Embeddings\n # -------------------\n print_rank_0(\"loading embeddings...\")\n gpt_model_module = _get_gpt_model(models[0])\n embed_tokens_weight = None\n if pp_rank == 0:\n embed_tokens_weight = gpt_model_module.embedding.word_embeddings.weight\n _broadcast_tp_shard_tensor_vocab(embed_tokens_weight, \"model.embed_tokens.weight\")\n\n # Transformer layers\n # -------------------\n layer_map = _megatron_calc_layer_map(config)\n\n for layer in range(config.num_hidden_layers):\n layer_name = f\"model.layers.{layer}\"\n print_rank_0(f\"loading layer #{layer}, with layer_name model.layers.{layer}...\")\n dst_pp_rank, dst_virtual_pp_rank, dst_layer_idx = layer_map[layer]\n\n gpt_model_module = _get_gpt_model(models[dst_virtual_pp_rank])\n sync_layer = gpt_model_module.decoder.layers[dst_layer_idx]\n\n _broadcast_tensor(\n sync_layer.self_attention.linear_qkv.layer_norm_weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.input_layernorm.weight\",\n )\n\n if f\"{layer_name}.self_attn.q_norm.weight\" in state_dict:\n _broadcast_tensor(\n sync_layer.self_attention.q_layernorm.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.q_norm.weight\",\n )\n _broadcast_tensor(\n sync_layer.self_attention.k_layernorm.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.k_norm.weight\",\n )\n\n _broadcast_tp_shard_tensor_qkv(\n sync_layer.self_attention.linear_qkv.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.q_proj.weight\",\n f\"{layer_name}.self_attn.k_proj.weight\",\n f\"{layer_name}.self_attn.v_proj.weight\",\n )\n if f\"{layer_name}.self_attn.q_proj.bias\" in state_dict:\n _broadcast_tp_shard_tensor_qkv(\n sync_layer.self_attention.linear_qkv.bias if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.q_proj.bias\",\n f\"{layer_name}.self_attn.k_proj.bias\",\n f\"{layer_name}.self_attn.v_proj.bias\",\n bias=True,\n )\n\n _broadcast_tp_shard_tensor(\n sync_layer.self_attention.linear_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.o_proj.weight\",\n chunk_dim=1,\n )\n _broadcast_tensor(\n sync_layer.mlp.linear_fc1.layer_norm_weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.post_attention_layernorm.weight\",\n )\n\n _broadcast_tp_shard_tensor_gate_up(\n sync_layer.mlp.linear_fc1.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.mlp.gate_proj.weight\",\n f\"{layer_name}.mlp.up_proj.weight\",\n )\n\n _broadcast_tp_shard_tensor(\n sync_layer.mlp.linear_fc2.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.mlp.down_proj.weight\",\n chunk_dim=1,\n )\n # Final Layernorm\n # -------------------\n print_rank_0(\"loading final layernorm...\")\n gpt_model_module = _get_gpt_model(models[-1])\n _broadcast_tensor(\n getattr(gpt_model_module.decoder.final_layernorm, \"weight\", None),\n \"model.norm.weight\",\n )\n\n print_rank_0(\"loading lm_head...\")\n lm_head_weight = None\n if pp_rank + 1 == pp_size:\n lm_head_weight = gpt_model_module.output_layer.weight\n\n if is_value_model:\n # if torch.distributed.get_rank() == src_rank:\n if \"lm_head.weight\" in state_dict and state_dict[\"lm_head.weight\"].shape[0] == 1:\n _broadcast_tensor(lm_head_weight, \"lm_head.weight\")\n elif \"reward_head.weight\" in state_dict and state_dict[\"reward_head.weight\"].shape[0] == 1:\n _broadcast_tensor(lm_head_weight, \"reward_head.weight\")\n print_rank_0(\"load lm_head from value_head weight\")\n elif \"score.weight\" in state_dict and state_dict[\"score.weight\"].shape[0] == 1:\n _broadcast_tensor(lm_head_weight, \"score.weight\")\n print_rank_0(\"load lm_head from score weight\")\n else:\n _broadcast_tensor(None, \"lm_head.weight\")\n print_rank_0(\"fail to match lm_head in value_model\")\n # else:\n\n # _broadcast_tensor(lm_head_weight, \"lm_head.weight\")\n\n else:\n _broadcast_tp_shard_tensor(lm_head_weight, \"lm_head.weight\")\n dist.barrier()\n # Broadcast weights inside data parallel groups\n for wrapped_model in wrapped_models:\n broadcast_params(wrapped_model)\n pass\n get_torch_device().empty_cache()\n print_rank_0(f\"loading megatron ckpt done, time elapsed {time.time() - start_time}s\")\n"}19{"file_name": "verl__models__mcore__model_forward.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport torch\n\nfrom verl.utils.megatron_utils import unwrap_model\nfrom verl.workers.config import MtpConfig\n\nfrom .util import (\n postprocess_bshd,\n postprocess_bshd_no_padding,\n postprocess_packed_seqs,\n postprocess_thd_no_padding,\n preprocess_bshd,\n preprocess_bshd_no_padding,\n preprocess_packed_seqs,\n preprocess_thd_no_padding,\n)\n\n\ndef model_forward_gen(vision_model: bool = False):\n def model_forward(\n model,\n input_ids,\n attention_mask,\n position_ids,\n multi_modal_inputs: dict,\n logits_processor=None,\n logits_processor_args: dict = None,\n value_model=False,\n data_format: str = \"thd\",\n mtp_config: MtpConfig = None,\n ):\n \"\"\"Forward pass for models with sequence packing.\"\"\"\n assert data_format in [\"thd\", \"bshd\"], \"data_format must be 'thd' or 'bshd'\"\n pre_process = (\n unwrap_model(model).pre_process if not vision_model else False\n ) # vision model does not need pre_process, because we pack the input_ids to thd in the forward function\n post_process = unwrap_model(model).post_process\n sp = unwrap_model(model).config.sequence_parallel\n fp8 = unwrap_model(model).config.fp8\n use_fp8_padding = fp8 in [\"e4m3\", \"hybrid\"]\n\n model_kwargs = {}\n if \"pixel_values\" in multi_modal_inputs:\n model_kwargs[\"pixel_values\"] = multi_modal_inputs[\"pixel_values\"].to(input_ids.device)\n if \"image_grid_thw\" in multi_modal_inputs:\n model_kwargs[\"image_grid_thw\"] = multi_modal_inputs[\"image_grid_thw\"].to(input_ids.device)\n if \"pixel_values_videos\" in multi_modal_inputs:\n model_kwargs[\"pixel_values_videos\"] = multi_modal_inputs[\"pixel_values_videos\"].to(input_ids.device)\n if \"video_grid_thw\" in multi_modal_inputs:\n model_kwargs[\"video_grid_thw\"] = multi_modal_inputs[\"video_grid_thw\"].to(input_ids.device)\n\n batch_size, seq_len = attention_mask.shape[:2]\n if data_format == \"thd\":\n input_ids_rmpad, packed_seq_params = preprocess_packed_seqs(\n input_ids, attention_mask, pre_process=pre_process or post_process, use_fp8_padding=use_fp8_padding\n )\n input_ids_rmpad = input_ids_rmpad.contiguous()\n\n # when pp > 1 and processor is not None, we need to pass the labels and loss_mask to the model\n if mtp_config and mtp_config.enable_train and post_process:\n args = {\n k: preprocess_packed_seqs(v, attention_mask, pre_process=True, use_fp8_padding=use_fp8_padding)[0]\n for k, v in logits_processor_args.items()\n }\n model_kwargs[\"labels\"] = args[\"label\"].contiguous()\n model_kwargs[\"loss_mask\"] = args[\"label_mask\"].contiguous()\n\n input_args = dict(\n input_ids=input_ids_rmpad,\n attention_mask=None,\n position_ids=position_ids if not vision_model else None, # vision models will calculate position_ids\n packed_seq_params=packed_seq_params,\n **model_kwargs,\n )\n\n if vision_model:\n # workaround for supporting sequence packing with context parallelism\n # cp split with sequence packing will make model lose vision token information, so we need to keep\n # the original input_ids and pack them after vision embedding is calculated,\n # cooporate with mbridge\n input_args[\"input_ids\"] = input_ids\n input_args[\"attention_mask\"] = attention_mask\n\n output_orig = model(**input_args)\n\n if post_process and logits_processor is not None:\n args = {\n k: preprocess_packed_seqs(v, attention_mask, pre_process=True, use_fp8_padding=use_fp8_padding)[0]\n for k, v in logits_processor_args.items()\n }\n output_dict = logits_processor(output_orig, **args)\n output = {\n k: postprocess_packed_seqs(\n v, packed_seq_params, attention_mask, batch_size, seq_len, post_process=post_process\n )\n for k, v in output_dict.items()\n }\n else:\n output = postprocess_packed_seqs(\n output_orig, packed_seq_params, attention_mask, batch_size, seq_len, post_process=post_process\n )\n elif data_format == \"bshd\":\n \"\"\"\n data_format: \"thd\" or \"bshd\", default is \"thd\",\n why we need this?\n for some new models, GPT-OSS, the thd format is not supported, so we need to use the bshd format.\n When using the bshd format, we have to add paddings to the input_ids to meet the longest sequence length, \n so it is recommended to disable dynamic batch size and set batch size to 1\n \"\"\"\n assert not vision_model, \"vision model does not support bshd format\"\n assert fp8 is None, \"fp8 is not supported for bshd format yet\"\n\n batch_size, sequence_length = attention_mask.shape[:2]\n new_input_ids, new_attention_mask, new_position_ids = preprocess_bshd(\n input_ids, attention_mask, position_ids, sequence_parallel=sp, pre_process=pre_process\n )\n output_orig = model(\n input_ids=new_input_ids,\n position_ids=new_position_ids,\n attention_mask=new_attention_mask,\n **model_kwargs,\n )\n if post_process and logits_processor is not None:\n args = {\n k: preprocess_bshd(v, attention_mask, position_ids, sequence_parallel=sp, pre_process=True)[0]\n for k, v in logits_processor_args.items()\n }\n output_dict = logits_processor(output_orig, **args)\n output = {\n k: postprocess_bshd(\n v, new_attention_mask, attention_mask, sequence_length, post_process=post_process\n )\n for k, v in output_dict.items()\n }\n else:\n output = postprocess_bshd(\n output_orig, new_attention_mask, attention_mask, sequence_length, post_process=post_process\n )\n if value_model and post_process:\n output = output[..., 0]\n return output\n\n return model_forward\n\n\ndef gptmodel_forward_no_padding(\n model,\n input_ids,\n multi_modal_inputs: dict,\n logits_processor=None,\n logits_processor_args: dict = None,\n value_model=False,\n vision_model=False,\n pad_token_id=None,\n data_format: str = \"thd\",\n enable_mtp: bool = False,\n):\n \"\"\"Default forward pass for GPT models with optional sequence packing.\"\"\"\n\n assert data_format in [\"thd\", \"bshd\"], \"data_format must be 'thd' or 'bshd'\"\n pre_process = unwrap_model(model).pre_process\n post_process = unwrap_model(model).post_process\n\n model_kwargs = {}\n if \"pixel_values\" in multi_modal_inputs:\n model_kwargs[\"pixel_values\"] = multi_modal_inputs[\"pixel_values\"].to(input_ids.device)\n if \"image_grid_thw\" in multi_modal_inputs:\n model_kwargs[\"image_grid_thw\"] = multi_modal_inputs[\"image_grid_thw\"].to(input_ids.device)\n if \"pixel_values_videos\" in multi_modal_inputs:\n model_kwargs[\"pixel_values_videos\"] = multi_modal_inputs[\"pixel_values_videos\"].to(input_ids.device)\n if \"video_grid_thw\" in multi_modal_inputs:\n model_kwargs[\"video_grid_thw\"] = multi_modal_inputs[\"video_grid_thw\"].to(input_ids.device)\n\n batch_size = input_ids.shape[0]\n if data_format == \"thd\":\n input_ids_rmpad, packed_seq_params = preprocess_thd_no_padding(input_ids, pre_process=pre_process)\n input_ids_rmpad = input_ids_rmpad.contiguous()\n\n if enable_mtp and post_process:\n args = {\n k: preprocess_thd_no_padding(v, pre_process=True, need_roll=(k == \"label\" or k == \"loss_mask\"))[0]\n for k, v in logits_processor_args.items()\n }\n model_kwargs[\"labels\"] = args[\"label\"].contiguous()\n model_kwargs[\"loss_mask\"] = args[\"loss_mask\"].contiguous()\n if logits_processor_args and \"loss_mask\" in logits_processor_args:\n logits_processor_args.pop(\"loss_mask\")\n\n # For VLM model, need to pass bshd format `input_ids` and `attention_mask`.\n attention_mask = None\n if vision_model:\n input_ids_rmpad = input_ids.to_padded_tensor(pad_token_id)\n seqlens_in_batch = input_ids.offsets().diff()\n attention_mask = torch.zeros_like(input_ids_rmpad, dtype=torch.bool)\n for i, seqlen in enumerate(seqlens_in_batch):\n attention_mask[i, :seqlen] = True\n\n output_orig = model(\n input_ids=input_ids_rmpad,\n attention_mask=attention_mask,\n position_ids=None,\n packed_seq_params=packed_seq_params,\n **model_kwargs,\n )\n\n if post_process and logits_processor is not None:\n args = {\n k: preprocess_thd_no_padding(v, pre_process=True, need_roll=(k == \"label\"))[0]\n for k, v in logits_processor_args.items()\n }\n output_dict = logits_processor(output_orig, **args)\n output = {\n k: postprocess_thd_no_padding(v, packed_seq_params, input_ids, batch_size, post_process=post_process)\n for k, v in output_dict.items()\n }\n else:\n output = postprocess_thd_no_padding(\n output_orig, packed_seq_params, input_ids, batch_size, post_process=post_process\n )\n else:\n \"\"\"\n data_format: \"thd\" or \"bshd\", default is \"thd\",\n why we need this?\n for some new models, GPT-OSS, the thd format is not supported, so we need to use the bshd format.\n When using the bshd format, we have to add paddings to the input_ids to meet the longest sequence length, \n so it is recommended to disable dynamic batch size and set batch size to 1\n \"\"\"\n\n input_ids_bshd, attention_mask_bshd, position_ids_bshd = preprocess_bshd_no_padding(\n input_ids, pre_process=pre_process\n )\n\n if enable_mtp and post_process:\n args = {\n k: preprocess_bshd_no_padding(v, pre_process=True, need_roll=(k == \"label\" or k == \"loss_mask\"))[0]\n for k, v in logits_processor_args.items()\n }\n model_kwargs[\"labels\"] = args[\"label\"].contiguous()\n model_kwargs[\"loss_mask\"] = args[\"loss_mask\"].contiguous()\n if logits_processor_args and \"loss_mask\" in logits_processor_args:\n logits_processor_args.pop(\"loss_mask\")\n\n output_orig = model(\n input_ids=input_ids_bshd,\n attention_mask=attention_mask_bshd,\n position_ids=position_ids_bshd,\n **model_kwargs,\n )\n if post_process and logits_processor is not None:\n args = {\n k: preprocess_bshd_no_padding(v, pre_process=True, need_roll=(k == \"label\"))[0]\n for k, v in logits_processor_args.items()\n }\n output_dict = logits_processor(output_orig, **args)\n output = {\n k: postprocess_bshd_no_padding(v, attention_mask_bshd, post_process=post_process)\n for k, v in output_dict.items()\n }\n else:\n output = postprocess_bshd_no_padding(output_orig, attention_mask_bshd, post_process=post_process)\n\n if value_model and post_process:\n # output = output[..., 0]\n # while using nested tensor, the advanced indexing operation above will result in an error at backward, i.e.\n # ValueError: NestedTensor _nested_select_backward_default(grad_output: t, self: jt_all, dim: any, index: any)\n # so we use `squeeze` to remove the last dimension\n output = output.squeeze(-1)\n\n return output\n"}20{"file_name": "verl__models__mcore__model_forward_1f1b_overlap.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Callable, Optional\n\nimport torch\nfrom megatron.core.models.common.model_chunk_schedule_plan import TransformerModelChunkSchedulePlan\nfrom megatron.core.models.gpt.gpt_model import GPTModel\nfrom megatron.core.utils import make_viewless_tensor\nfrom torch import Tensor\n\nfrom verl.models.mcore.util import preprocess_packed_seqs\nfrom verl.utils.kernel.linear_cross_entropy import linear_cross_entropy\nfrom verl.utils.megatron_utils import unwrap_model\nfrom verl.utils.model import CausalLMOutputForPPO\n\nfrom .util import postprocess_packed_seqs, postprocess_packed_seqs_for_dict_output\n\n\ndef gptmodel_forward_1f1b_overlap(\n model: GPTModel,\n input_ids: Tensor,\n position_ids: Tensor,\n attention_mask: Tensor,\n labels: Tensor = None,\n labels_mask: Tensor = None,\n multi_modal_inputs: Optional[dict] = None,\n logits_processor: Optional[Callable] = None,\n logits_processor_args: Optional[dict] = None,\n temperature: float = 1.0,\n) -> TransformerModelChunkSchedulePlan:\n pre_process: bool = unwrap_model(model).pre_process\n post_process: bool = unwrap_model(model).post_process\n assert logits_processor is None, \"only support fused kernel\"\n batch_size, seq_len = attention_mask.shape[:2]\n input_ids_rmpad, packed_seq_params = preprocess_packed_seqs(input_ids, attention_mask, pre_process=pre_process)\n input_ids_rmpad = input_ids_rmpad.contiguous()\n\n schedule_plan = model.build_schedule_plan(\n input_ids=input_ids_rmpad,\n attention_mask=attention_mask,\n labels=labels,\n position_ids=position_ids,\n packed_seq_params=packed_seq_params,\n )\n if post_process:\n attention_mask_out = attention_mask\n\n def _postprocess(\n self,\n hidden_states,\n input_ids,\n position_ids,\n labels,\n rotary_pos_emb,\n rotary_pos_cos,\n rotary_pos_sin,\n mtp_in_postprocess=None,\n loss_mask=None,\n decoder_input=None,\n attention_mask=None,\n inference_params=None,\n packed_seq_params=None,\n sequence_len_offset=None,\n runtime_gather_output=None,\n extra_block_kwargs=None,\n inference_context=None,\n ):\n \"\"\"patched from https://github.com/NVIDIA/Megatron-LM/blob/core_r0.14.0/megatron/core/models/gpt/gpt_model.py#L412\"\"\"\n \"\"\"Postprocesses decoder hidden states to generate logits or compute loss.\n\n Applies Multi-Token Prediction if enabled, generates output logits through\n the output layer, and computes language model loss when labels are provided.\n \"\"\"\n from megatron.core import parallel_state\n from megatron.core.tensor_parallel import gather_from_sequence_parallel_region\n\n in_inference_mode = inference_context is not None and not self.training\n if in_inference_mode:\n assert runtime_gather_output, \"Inference must always gather TP logits\"\n\n # logits and loss\n output_weight = None\n if self.share_embeddings_and_output_weights:\n output_weight = self.shared_embedding_or_output_weight()\n\n if mtp_in_postprocess:\n hidden_states = self.mtp(\n input_ids=input_ids,\n position_ids=position_ids,\n hidden_states=hidden_states,\n attention_mask=attention_mask,\n inference_params=inference_params,\n rotary_pos_emb=rotary_pos_emb,\n rotary_pos_cos=rotary_pos_cos,\n rotary_pos_sin=rotary_pos_sin,\n packed_seq_params=packed_seq_params,\n sequence_len_offset=sequence_len_offset,\n embedding=self.embedding,\n **(extra_block_kwargs or {}),\n )\n\n if not self.post_process:\n return hidden_states\n\n if self.mtp_process:\n from megatron.core.transformer.multi_token_prediction import (\n MTPLossAutoScaler,\n MTPLossLoggingHelper,\n roll_tensor,\n )\n\n mtp_labels = labels.clone()\n hidden_states_list = torch.chunk(hidden_states, 1 + self.config.mtp_num_layers, dim=0)\n hidden_states = hidden_states_list[0]\n if loss_mask is None:\n # if loss_mask is not provided, use all ones as loss_mask\n loss_mask = torch.ones_like(mtp_labels)\n for mtp_layer_number in range(self.config.mtp_num_layers):\n # output\n mtp_logits, _ = self.output_layer(\n hidden_states_list[mtp_layer_number + 1],\n weight=output_weight,\n runtime_gather_output=runtime_gather_output,\n )\n # Calc loss for the current Multi-Token Prediction (MTP) layers.\n mtp_labels, _ = roll_tensor(mtp_labels, shifts=-1, dims=-1, cp_group=self.cp_group)\n loss_mask, num_tokens = roll_tensor(loss_mask, shifts=-1, dims=-1, cp_group=self.cp_group)\n mtp_loss = self.compute_language_model_loss(mtp_labels, mtp_logits)\n mtp_loss = loss_mask * mtp_loss\n if self.training:\n # TODO(shifangx): remove the use of parallel_state here\n # after moving loss logging to loss_func in pretrain_gpt.py\n MTPLossLoggingHelper.save_loss_to_tracker(\n torch.sum(mtp_loss) / num_tokens,\n mtp_layer_number,\n self.config.mtp_num_layers,\n avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True),\n )\n mtp_loss_scale = self.config.mtp_loss_scaling_factor / self.config.mtp_num_layers\n if self.config.calculate_per_token_loss:\n hidden_states = MTPLossAutoScaler.apply(hidden_states, mtp_loss_scale * mtp_loss)\n else:\n hidden_states = MTPLossAutoScaler.apply(hidden_states, mtp_loss_scale * mtp_loss / num_tokens)\n\n if logits_processor is not None:\n logits, _ = self.output_layer(\n hidden_states, weight=output_weight, runtime_gather_output=runtime_gather_output\n )\n output_orig = logits.transpose(0, 1).contiguous()\n args = {\n k: preprocess_packed_seqs(v, attention_mask_out, pre_process=True)[0]\n for k, v in logits_processor_args.items()\n }\n output_dict = logits_processor(output_orig, **args)\n output = {\n k: postprocess_packed_seqs(\n v, packed_seq_params, attention_mask_out, batch_size, seq_len, post_process=post_process\n )\n for k, v in output_dict.items()\n }\n else:\n # fused kernel\n\n labels_rmpad, _ = preprocess_packed_seqs(labels, attention_mask, pre_process=True)\n labels_mask_rmpad, _ = preprocess_packed_seqs(labels_mask, attention_mask, pre_process=True)\n labels_rmpad = labels_rmpad.contiguous()\n labels_mask_rmpad = labels_mask_rmpad.contiguous()\n\n output = CausalLMOutputForPPO(\n loss=None,\n logits=None,\n past_key_values=None,\n hidden_states=hidden_states,\n attentions=None,\n )\n if self.config.sequence_parallel:\n hidden_states = gather_from_sequence_parallel_region(hidden_states)\n logprobs, entropy = linear_cross_entropy(\n hidden_states,\n self.output_layer.weight,\n labels_rmpad,\n temperature,\n \"none\",\n parallel_state.get_tensor_model_parallel_group(),\n )\n output.entropy = entropy\n output.log_probs = logprobs\n\n output = postprocess_packed_seqs_for_dict_output(\n labels_mask_rmpad,\n output,\n packed_seq_params,\n attention_mask,\n batch_size,\n seq_len,\n post_process=post_process,\n )\n output_ = [output[\"log_probs\"]]\n # TODO NOW 1f1b overlap only support one tensor output\n # if \"entropy\" in output:\n # output_.append(output[\"entropy\"])\n output_ = tuple(output_)\n return output_\n\n def _custom_post_process_node_forward_impl(self, hidden_states):\n if self.gpt_model.decoder.final_layernorm and not self.gpt_model.mtp_process:\n hidden_states = self.gpt_model.decoder.final_layernorm(hidden_states)\n # TENorm produces a \"viewed\" tensor. This will result in schedule.py's\n # deallocate_output_tensor() throwing an error, so a viewless tensor is\n # created to prevent this.\n hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True)\n\n # Run GPTModel._postprocess\n output = self.gpt_model._postprocess(\n hidden_states=hidden_states,\n input_ids=self.chunk_state.input_ids,\n position_ids=self.chunk_state.position_ids,\n labels=self.chunk_state.labels,\n decoder_input=self.chunk_state.decoder_input,\n rotary_pos_emb=self.chunk_state.rotary_pos_emb,\n rotary_pos_cos=self.chunk_state.rotary_pos_cos,\n rotary_pos_sin=self.chunk_state.rotary_pos_sin,\n mtp_in_postprocess=False,\n loss_mask=self.chunk_state.loss_mask,\n attention_mask=self.chunk_state.attention_mask,\n packed_seq_params=self.chunk_state.packed_seq_params,\n sequence_len_offset=self.chunk_state.sequence_len_offset,\n runtime_gather_output=self.chunk_state.runtime_gather_output,\n extra_block_kwargs=self.chunk_state.extra_block_kwargs,\n )\n return output\n\n schedule_plan.post_process.forward_impl = _custom_post_process_node_forward_impl.__get__(\n schedule_plan.post_process, schedule_plan.post_process.__class__\n )\n unwrap_model(model)._postprocess = _postprocess.__get__(unwrap_model(model), unwrap_model(model).__class__)\n\n return schedule_plan\n"}21{"file_name": "verl__models__mcore__model_forward_fused.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom collections import OrderedDict\nfrom typing import Optional\n\nimport megatron.core as mcore\nimport torch\nfrom megatron.core import parallel_state\nfrom megatron.core.config_logger import has_config_logger_enabled, log_config_to_disk\nfrom megatron.core.inference.contexts import BaseInferenceContext\nfrom megatron.core.models.gpt.gpt_model import GPTModel\nfrom megatron.core.packed_seq_params import PackedSeqParams\nfrom megatron.core.tensor_parallel.mappings import gather_from_sequence_parallel_region\nfrom megatron.core.utils import deprecate_inference_params\nfrom packaging import version\nfrom torch import Tensor\n\nfrom verl.models.mcore.util import preprocess_packed_seqs, preprocess_thd_no_padding\nfrom verl.utils.kernel.linear_cross_entropy import linear_cross_entropy\nfrom verl.utils.megatron_utils import unwrap_model\nfrom verl.utils.model import CausalLMOutputForPPO\n\nfrom .util import postprocess_packed_seqs_for_dict_output, postprocess_thd_no_padding\n\n\ndef _get_patching_model(model: torch.nn.Module):\n model = unwrap_model(model)\n if isinstance(model, GPTModel):\n return model\n\n if not (hasattr(model, \"language_model\") and isinstance(model.language_model, GPTModel)):\n print(f\"Model {model.__class__.__name__} is not a supported for fused forward\")\n return None\n\n return model.language_model\n\n\ndef patch_fused_forward(model: torch.nn.Module):\n assert version.parse(mcore.__version__) >= version.parse(\"0.13.0\"), (\n \"Fused forward patching requires mecore >= 0.13.0\"\n )\n model = _get_patching_model(model)\n if model is not None:\n model.forward_backup = model.forward\n model.forward = _fused_GPTModel_forward.__get__(model, model.__class__)\n\n\ndef unpatch_fused_forward(model: torch.nn.Module):\n model = _get_patching_model(model)\n if model is not None:\n model.forward = model.forward_backup\n\n\ndef fused_forward_model_gen(vision_model: bool = False):\n def fused_forward_model(\n model,\n input_ids: Tensor,\n position_ids: Tensor,\n attention_mask: Tensor,\n labels: Tensor,\n labels_mask: Tensor,\n temperature: float,\n multi_modal_inputs: dict,\n ):\n pre_process: bool = (\n unwrap_model(model).pre_process if not vision_model else False\n ) # vision model does not need pre_process, because we pack the input_ids to thd in the forward function\n post_process: bool = unwrap_model(model).post_process\n\n model_kwargs = {}\n if \"pixel_values\" in multi_modal_inputs:\n model_kwargs[\"pixel_values\"] = multi_modal_inputs[\"pixel_values\"].to(input_ids.device)\n if \"image_grid_thw\" in multi_modal_inputs:\n model_kwargs[\"image_grid_thw\"] = multi_modal_inputs[\"image_grid_thw\"].to(input_ids.device)\n if \"pixel_values_videos\" in multi_modal_inputs:\n model_kwargs[\"pixel_values_videos\"] = multi_modal_inputs[\"pixel_values_videos\"].to(input_ids.device)\n if \"video_grid_thw\" in multi_modal_inputs:\n model_kwargs[\"video_grid_thw\"] = multi_modal_inputs[\"video_grid_thw\"].to(input_ids.device)\n\n batch_size, seq_len = attention_mask.shape[:2]\n input_ids_rmpad, packed_seq_params = preprocess_packed_seqs(input_ids, attention_mask, pre_process=pre_process)\n input_ids_rmpad = input_ids_rmpad.contiguous()\n labels_rmpad, _ = preprocess_packed_seqs(labels, attention_mask, pre_process=True)\n labels_mask_rmpad, _ = preprocess_packed_seqs(labels_mask, attention_mask, pre_process=True)\n labels_rmpad = labels_rmpad.contiguous()\n labels_mask_rmpad = labels_mask_rmpad.contiguous()\n\n input_args = dict(\n input_ids=input_ids_rmpad,\n attention_mask=None,\n position_ids=position_ids if not vision_model else None, # vision models will calculate position_ids\n packed_seq_params=packed_seq_params,\n labels=labels_rmpad,\n temperature=temperature,\n **model_kwargs,\n )\n\n if vision_model:\n # workaround for supporting sequence packing with context parallelism\n # cp split with sequence packing will make model lose vision token information, so we need to keep\n # the original input_ids and pack them after vision embedding is calculated,\n # cooporate with mbridge\n input_args[\"input_ids\"] = input_ids\n input_args[\"attention_mask\"] = attention_mask\n\n output_orig: CausalLMOutputForPPO = model(**input_args)\n\n if post_process:\n # output_orig is in type of CausalLMOutputForPPO\n output = postprocess_packed_seqs_for_dict_output(\n labels_mask_rmpad,\n output_orig,\n packed_seq_params,\n attention_mask,\n batch_size,\n seq_len,\n post_process=post_process,\n )\n else:\n output = output_orig\n return output\n\n return fused_forward_model\n\n\ndef fused_forward_no_padding_gen(vision_model: bool = False):\n def fused_forward_no_padding(\n model,\n input_ids: Tensor,\n labels: Tensor,\n multi_modal_inputs: dict,\n temperature: float,\n calculate_entropy: bool,\n pad_token_id: int,\n ):\n pre_process = unwrap_model(model).pre_process\n post_process = unwrap_model(model).post_process\n\n input_ids_rmpad, packed_seq_params = preprocess_thd_no_padding(input_ids, pre_process=pre_process)\n input_ids_rmpad = input_ids_rmpad.contiguous()\n\n model_kwargs = {}\n if \"pixel_values\" in multi_modal_inputs:\n model_kwargs[\"pixel_values\"] = multi_modal_inputs[\"pixel_values\"].to(input_ids.device)\n if \"image_grid_thw\" in multi_modal_inputs:\n model_kwargs[\"image_grid_thw\"] = multi_modal_inputs[\"image_grid_thw\"].to(input_ids.device)\n if \"pixel_values_videos\" in multi_modal_inputs:\n model_kwargs[\"pixel_values_videos\"] = multi_modal_inputs[\"pixel_values_videos\"].to(input_ids.device)\n if \"video_grid_thw\" in multi_modal_inputs:\n model_kwargs[\"video_grid_thw\"] = multi_modal_inputs[\"video_grid_thw\"].to(input_ids.device)\n\n attention_mask = None\n if vision_model:\n input_ids_rmpad = input_ids.to_padded_tensor(pad_token_id)\n seqlens_in_batch = input_ids.offsets().diff().to(input_ids.device)\n max_seq_len = input_ids_rmpad.shape[1]\n attention_mask = torch.arange(max_seq_len, device=input_ids.device).unsqueeze(\n 0\n ) < seqlens_in_batch.unsqueeze(1)\n\n labels_rmpad, _ = preprocess_thd_no_padding(labels, pre_process=True, need_roll=True)\n labels_rmpad = labels_rmpad.contiguous()\n output_orig: CausalLMOutputForPPO = model(\n input_ids=input_ids_rmpad,\n attention_mask=attention_mask,\n position_ids=None,\n packed_seq_params=packed_seq_params,\n labels=labels_rmpad,\n temperature=temperature,\n **model_kwargs,\n )\n\n if not post_process:\n return output_orig\n\n log_probs = output_orig.log_probs\n if log_probs.dim() == 1:\n log_probs = log_probs.unsqueeze(0)\n log_probs = postprocess_thd_no_padding(\n log_probs, packed_seq_params, input_ids, input_ids.shape[0], post_process=post_process\n )\n\n output = {\"log_probs\": log_probs}\n\n if calculate_entropy:\n entropy = output_orig.entropy\n if entropy.dim() == 1:\n entropy = entropy.unsqueeze(0)\n entropy = postprocess_thd_no_padding(\n entropy, packed_seq_params, input_ids, input_ids.shape[0], post_process=post_process\n )\n output[\"entropy\"] = entropy\n\n return output\n\n return fused_forward_no_padding\n\n\ndef _fused_GPTModel_forward(\n model,\n input_ids: Tensor,\n position_ids: Tensor,\n attention_mask: Tensor,\n decoder_input: Tensor = None,\n labels: Tensor = None,\n inference_context: BaseInferenceContext = None,\n packed_seq_params: PackedSeqParams = None,\n extra_block_kwargs: dict = None,\n runtime_gather_output: Optional[bool] = None,\n *,\n inference_params: Optional[BaseInferenceContext] = None,\n loss_mask: Optional[Tensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> CausalLMOutputForPPO:\n \"\"\"\n Patch self._postprocess in forward for GPT models to enable fused kernel support.\n https://github.com/NVIDIA/Megatron-LM/blob/core_v0.13.0/megatron/core/models/gpt/gpt_model.py\n\n TODO: Currently we still need to patch `forward` because we need to pass `temperature`\n explicitly to `self._postprocess` when calling, maybe there can be a better way to handle this?\n \"\"\"\n\n inference_context = deprecate_inference_params(inference_context, inference_params)\n\n preproc_output = model._preprocess(\n input_ids=input_ids,\n position_ids=position_ids,\n decoder_input=decoder_input,\n inference_context=inference_context,\n packed_seq_params=packed_seq_params,\n )\n\n (decoder_input, rotary_pos_emb, rotary_pos_cos, rotary_pos_sin, sequence_len_offset) = preproc_output[:5]\n\n # Run decoder.\n hidden_states = model.decoder(\n hidden_states=decoder_input,\n attention_mask=attention_mask,\n inference_context=inference_context,\n rotary_pos_emb=rotary_pos_emb,\n rotary_pos_cos=rotary_pos_cos,\n rotary_pos_sin=rotary_pos_sin,\n packed_seq_params=packed_seq_params,\n sequence_len_offset=sequence_len_offset,\n **(extra_block_kwargs or {}),\n **kwargs,\n )\n\n if not model.post_process:\n return hidden_states\n\n output = CausalLMOutputForPPO(\n loss=None,\n logits=None,\n past_key_values=None,\n hidden_states=hidden_states,\n attentions=None,\n )\n\n if model.config.sequence_parallel:\n hidden_states = gather_from_sequence_parallel_region(hidden_states)\n\n # Get the output weight - use embedding weight if output_layer is None or weight is shared\n if hasattr(model, \"output_layer\") and model.output_layer is not None and model.output_layer.weight is not None:\n output_weight = model.output_layer.weight\n else:\n # When embeddings are tied, use the embedding weight\n output_weight = model.embedding.word_embeddings.weight\n\n logprobs, entropy = linear_cross_entropy(\n hidden_states,\n output_weight,\n labels,\n temperature,\n \"none\",\n parallel_state.get_tensor_model_parallel_group(),\n )\n\n if has_config_logger_enabled(model.config):\n payload = OrderedDict(\n {\n \"input_ids\": input_ids,\n \"position_ids\": position_ids,\n \"attention_mask\": attention_mask,\n \"decoder_input\": decoder_input,\n \"logprobs\": logprobs,\n \"entropy\": entropy,\n }\n )\n log_config_to_disk(model.config, payload, prefix=\"input_and_logits\")\n\n output.entropy = entropy\n output.log_probs = logprobs\n\n return output\n"}22{"file_name": "verl__models__mcore__model_initializer.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n# use mcore transformer config to initialize the model\nimport inspect\nfrom abc import ABC, abstractmethod\n\nfrom megatron.core.models.gpt.gpt_layer_specs import get_gpt_decoder_block_spec, get_gpt_mtp_block_spec\nfrom megatron.core.models.gpt.gpt_model import GPTModel\n\nfrom .config_converter import PretrainedConfig, TransformerConfig\n\n\nclass BaseModelInitializer(ABC):\n \"\"\"Base class for model initializers.\"\"\"\n\n def __init__(self, tfconfig: TransformerConfig, hf_config: PretrainedConfig):\n self.tfconfig = tfconfig\n self.hf_config = hf_config\n self.has_vp_stage = inspect.signature(get_gpt_decoder_block_spec).parameters.get(\"vp_stage\", None) is not None\n\n @abstractmethod\n def get_transformer_layer_spec(self, vp_stage=None):\n \"\"\"Get the transformer layer specification.\n https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/core/models/gpt/gpt_layer_specs.py\"\"\"\n pass\n\n def get_rope_scaling_args(self) -> dict:\n \"\"\"Get rope scaling args.\"\"\"\n rope_scaling_args = {}\n if \"rope_scaling\" in self.hf_config:\n if self.hf_config.rope_scaling is not None:\n # assert self.hf_config.rope_scaling[\"type\"] == \"linear\", \"only linear scaling is supported for now\"\n rope_scaling_args[\"seq_len_interpolation_factor\"] = self.hf_config.rope_scaling[\"factor\"]\n return rope_scaling_args\n\n def initialize(\n self,\n pre_process: bool = True,\n post_process: bool = True,\n share_embeddings_and_output_weights: bool = False,\n value: bool = False,\n **extra_kwargs,\n ) -> GPTModel:\n \"\"\"Initialize a GPT model with the given configuration.\n https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/core/models/gpt/gpt_model.py\n\n Args:\n pre_process (bool): include embedding layer.\n post_process (bool): including an output layer.\n share_embeddings_and_output_weights (bool): input embeddings and output logit weights are shared.\n value (bool): add an extra linear layer for classification or regression.\n\n Returns:\n GPTModel: An initialized GPT model instance\n \"\"\"\n vp_stage = extra_kwargs.get(\"vp_stage\", None)\n transformer_layer_spec = self.get_transformer_layer_spec(vp_stage=vp_stage)\n rope_scaling_args = self.get_rope_scaling_args()\n mtp_block_spec = extra_kwargs.get(\"mtp_block_spec\", None)\n model = GPTModel(\n config=self.tfconfig,\n transformer_layer_spec=transformer_layer_spec,\n vocab_size=self.hf_config.vocab_size,\n max_sequence_length=self.hf_config.max_position_embeddings,\n pre_process=pre_process,\n post_process=post_process,\n share_embeddings_and_output_weights=share_embeddings_and_output_weights,\n position_embedding_type=\"rope\",\n rotary_base=self.hf_config.rope_theta,\n **rope_scaling_args,\n mtp_block_spec=mtp_block_spec,\n **({} if not self.has_vp_stage else {\"vp_stage\": vp_stage}),\n )\n\n if post_process and value:\n from verl.models.llama.megatron.layers.parallel_linear import LinearForLastLayer\n\n model.output_layer = LinearForLastLayer(\n input_size=self.tfconfig.hidden_size, output_size=1, config=self.tfconfig\n )\n\n return model\n\n\nclass DenseModel(BaseModelInitializer):\n \"\"\"Initializer for dense models like Llama and Qwen2.\"\"\"\n\n def get_transformer_layer_spec(self, vp_stage=None):\n assert self.tfconfig.normalization == \"RMSNorm\", \"only RMSNorm is supported for now\"\n extra_kwargs = {} if not self.has_vp_stage else {\"vp_stage\": vp_stage}\n return get_gpt_decoder_block_spec(self.tfconfig, use_transformer_engine=True, **extra_kwargs)\n\n\nclass Qwen2MoEModel(BaseModelInitializer):\n \"\"\"Initializer for Qwen2 MoE models.\"\"\"\n\n def get_transformer_layer_spec(self, vp_stage=None):\n assert self.tfconfig.normalization == \"RMSNorm\", \"only RMSNorm is supported for now\"\n extra_kwargs = {} if not self.has_vp_stage else {\"vp_stage\": vp_stage}\n transformer_layer_spec = get_gpt_decoder_block_spec(self.tfconfig, use_transformer_engine=True, **extra_kwargs)\n\n # Patch layer spec for shared experts\n for i in range(len(transformer_layer_spec.layer_specs)):\n transformer_layer_spec.layer_specs[i].submodules.mlp.submodules.shared_experts.params[\"gate\"] = True\n\n return transformer_layer_spec\n\n def initialize(self, **kwargs):\n # Qwen default freeze_moe_router: true\n model = super().initialize(**kwargs)\n freeze_moe_router = kwargs.get(\"freeze_moe_router\", True)\n if freeze_moe_router:\n for layer in model.decoder.layers:\n layer.mlp.router.weight.requires_grad = False\n return model\n\n\nclass MixtralModel(BaseModelInitializer):\n \"\"\"Initializer for Mixtral models.\"\"\"\n\n def get_transformer_layer_spec(self, vp_stage=None):\n assert self.tfconfig.normalization == \"RMSNorm\", \"only RMSNorm is supported for now\"\n extra_kwargs = {} if not self.has_vp_stage else {\"vp_stage\": vp_stage}\n transformer_layer_spec = get_gpt_decoder_block_spec(self.tfconfig, use_transformer_engine=True, **extra_kwargs)\n return transformer_layer_spec\n\n def initialize(self, **kwargs):\n model = super().initialize(**kwargs)\n freeze_moe_router = kwargs.get(\"freeze_moe_router\", False)\n if freeze_moe_router:\n for layer in model.decoder.layers:\n layer.mlp.router.weight.requires_grad = False\n return model\n\n\nclass Qwen3MoEModel(BaseModelInitializer):\n \"\"\"Initializer for Qwen3 MoE models.\"\"\"\n\n def get_transformer_layer_spec(self, vp_stage=None):\n assert self.tfconfig.normalization == \"RMSNorm\", \"only RMSNorm is supported for now\"\n extra_kwargs = {} if not self.has_vp_stage else {\"vp_stage\": vp_stage}\n transformer_layer_spec = get_gpt_decoder_block_spec(self.tfconfig, use_transformer_engine=True, **extra_kwargs)\n return transformer_layer_spec\n\n def initialize(self, **kwargs):\n # Qwen default freeze_moe_router: true\n model = super().initialize(**kwargs)\n freeze_moe_router = kwargs.get(\"freeze_moe_router\", True)\n if freeze_moe_router:\n for layer in model.decoder.layers:\n layer.mlp.router.weight.requires_grad = False\n return model\n\n\nclass DeepseekV3Model(BaseModelInitializer):\n \"\"\"Initializer for DeepseekV3 models.\"\"\"\n\n def get_transformer_layer_spec(self, vp_stage=None):\n extra_kwargs = {} if not self.has_vp_stage else {\"vp_stage\": vp_stage}\n transformer_layer_spec = get_gpt_decoder_block_spec(self.tfconfig, use_transformer_engine=True, **extra_kwargs)\n return transformer_layer_spec\n\n def get_rope_scaling_args(self) -> dict:\n \"\"\"Get rope scaling args.\"\"\"\n rope_scaling_args = {}\n return rope_scaling_args\n\n def initialize(\n self,\n **kwargs,\n ):\n vp_stage = kwargs.get(\"vp_stage\", None)\n freeze_moe_router = kwargs.get(\"freeze_moe_router\", True)\n if freeze_moe_router:\n self.tfconfig.moe_router_load_balancing_type = \"none\"\n # MTP\n if self.tfconfig.mtp_num_layers is not None and self.tfconfig.mtp_num_layers > 0:\n transformer_layer_spec = self.get_transformer_layer_spec(vp_stage=vp_stage)\n mtp_block_spec = get_gpt_mtp_block_spec(\n self.tfconfig, transformer_layer_spec, use_transformer_engine=True, vp_stage=vp_stage\n )\n kwargs[\"mtp_block_spec\"] = mtp_block_spec\n\n model = super().initialize(**kwargs)\n if freeze_moe_router:\n for layer in model.decoder.layers:\n if hasattr(layer.mlp, \"router\"):\n layer.mlp.router.weight.requires_grad = False\n return model\n\n\nclass Qwen25VLModel(BaseModelInitializer):\n \"\"\"Initializer for Qwen2.5 VL models.\"\"\"\n\n def get_transformer_layer_spec(self, vp_stage=None):\n extra_kwargs = {} if not self.has_vp_stage else {\"vp_stage\": vp_stage}\n transformer_layer_spec = get_gpt_decoder_block_spec(self.tfconfig, use_transformer_engine=True, **extra_kwargs)\n return transformer_layer_spec\n\n def initialize(\n self,\n pre_process=None,\n post_process=None,\n share_embeddings_and_output_weights=False,\n value=False,\n **extra_kwargs,\n ):\n tfconfig = self.tfconfig\n hf_config = self.hf_config\n # Qwen2_5_VLForConditionalGeneration\n from copy import deepcopy\n\n transformer_layer_spec = self.get_transformer_layer_spec()\n\n from megatron.core.extensions.transformer_engine import TEColumnParallelLinear, TERowParallelLinear\n from megatron.core.models.gpt.moe_module_specs import MLPSubmodules\n from megatron.core.models.vision.vit_layer_specs import get_vit_layer_with_transformer_engine_spec\n\n from .qwen2_5_vl import Qwen2_5VLModel, get_vision_model_config, get_vision_projection_config\n\n vision_transformer_config = get_vision_model_config(deepcopy(tfconfig))\n vision_transformer_config.pipeline_model_parallel_size = 1\n vision_transformer_config.first_pipeline_num_layers = None\n\n vision_projection_config = get_vision_projection_config(\n deepcopy(tfconfig),\n vision_transformer_config.hidden_size,\n spatial_merge_size=hf_config.vision_config.spatial_merge_size,\n )\n vision_projection_layer_spec = MLPSubmodules(\n linear_fc1=TEColumnParallelLinear,\n linear_fc2=TERowParallelLinear,\n )\n vision_transformer_layer_spec = get_vit_layer_with_transformer_engine_spec()\n\n qwen25_vl_model = Qwen2_5VLModel(\n language_transformer_config=tfconfig,\n language_transformer_layer_spec=transformer_layer_spec,\n language_vocab_size=hf_config.vocab_size,\n language_max_sequence_length=hf_config.max_position_embeddings,\n vision_transformer_config=vision_transformer_config,\n vision_transformer_layer_spec=vision_transformer_layer_spec,\n vision_projection_config=vision_projection_config,\n vision_projection_layer_spec=vision_projection_layer_spec,\n vision_projection_type=\"mlp\",\n language_rotary_base=hf_config.rope_theta,\n pre_process=pre_process,\n post_process=post_process,\n add_decoder=True,\n add_encoder=True,\n parallel_output=True,\n language_share_embeddings_and_output_weights=share_embeddings_and_output_weights,\n )\n\n if post_process and value:\n from verl.models.llama.megatron.layers.parallel_linear import LinearForLastLayer\n\n qwen25_vl_model.language_model.output_layer = LinearForLastLayer(\n input_size=tfconfig.hidden_size, output_size=1, config=tfconfig\n )\n\n return qwen25_vl_model\n"}23{"file_name": "verl__models__mcore__mtp_patch.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.\n# Copyright 2025 Meituan Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Callable\n\nimport torch\nfrom megatron.core import parallel_state\nfrom megatron.core.models.gpt.gpt_model import GPTModel\nfrom megatron.core.transformer.multi_token_prediction import (\n MTPLossAutoScaler,\n MTPLossLoggingHelper,\n roll_tensor,\n)\n\ntry:\n from megatron.core.utils import unwrap_model\nexcept ImportError:\n from verl.utils.megatron_utils import unwrap_model\n\n\ndef _get_patching_model(model: torch.nn.Module):\n model = unwrap_model(model)\n if isinstance(model, GPTModel):\n return model\n\n if not (hasattr(model, \"language_model\") and isinstance(model.language_model, GPTModel)):\n print(f\"Model {model.__class__.__name__} is not a supported for fused forward\")\n return None\n\n return model.language_model\n\n\ndef patch_postprocess(model: torch.nn.Module):\n model = _get_patching_model(model)\n if model is not None:\n model._postprocess_backup = model._postprocess\n model._postprocess = _megatron_gptmodel_postprocess.__get__(model, model.__class__)\n\n\ndef unpatch_postprocess(model: torch.nn.Module):\n model = _get_patching_model(model)\n if model is not None:\n model._postprocess = model._postprocess_backup\n\n\n# copy from https://github.com/NVIDIA/Megatron-LM/blob/23e092f41ec8bc659020e401ddac9576c1cfed7e/megatron/core/models/gpt/gpt_model.py\n# patch the postprocess method of GPTModel to support advanced features like MTP, 1f1b overlap, etc.\ndef _megatron_gptmodel_postprocess(\n self,\n hidden_states,\n input_ids,\n position_ids,\n labels,\n rotary_pos_emb,\n rotary_pos_cos,\n rotary_pos_sin,\n mtp_in_postprocess=None,\n loss_mask=None,\n decoder_input=None,\n attention_mask=None,\n inference_params=None,\n packed_seq_params=None,\n sequence_len_offset=None,\n runtime_gather_output=None,\n extra_block_kwargs=None,\n inference_context=None,\n):\n \"\"\"Postprocesses decoder hidden states to generate logits or compute loss.\n\n Applies Multi-Token Prediction if enabled, generates output logits through\n the output layer, and computes language model loss when labels are provided.\n \"\"\"\n\n # logits and loss\n output_weight = None\n if self.share_embeddings_and_output_weights:\n output_weight = self.shared_embedding_or_output_weight()\n\n if mtp_in_postprocess and labels is not None:\n hidden_states = self.mtp(\n input_ids=input_ids,\n position_ids=position_ids,\n hidden_states=hidden_states,\n attention_mask=attention_mask,\n inference_params=inference_params,\n rotary_pos_emb=rotary_pos_emb,\n rotary_pos_cos=rotary_pos_cos,\n rotary_pos_sin=rotary_pos_sin,\n packed_seq_params=packed_seq_params,\n sequence_len_offset=sequence_len_offset,\n embedding=self.embedding,\n **(extra_block_kwargs or {}),\n )\n\n if not self.post_process:\n return hidden_states\n\n # Skip when mtp_num_layers is None or 0\n if self.config.mtp_num_layers and labels is not None:\n mtp_labels = labels.clone()\n\n hidden_states_list = torch.chunk(hidden_states, 1 + self.config.mtp_num_layers, dim=0)\n hidden_states = hidden_states_list[0]\n if loss_mask is None:\n # if loss_mask is not provided, use all ones as loss_mask\n loss_mask = torch.ones_like(mtp_labels)\n for mtp_layer_number in range(self.config.mtp_num_layers):\n # Calc loss for the current Multi-Token Prediction (MTP) layers.\n mtp_labels, _ = roll_tensor(\n mtp_labels,\n shifts=-1,\n dims=-1,\n cp_group=self.cp_group,\n packed_seq_params=packed_seq_params,\n )\n loss_mask, num_tokens = roll_tensor(\n loss_mask,\n shifts=-1,\n dims=-1,\n cp_group=self.cp_group,\n packed_seq_params=packed_seq_params,\n )\n\n # Compute mtp loss without storing logits to save memory.\n mtp_loss = self.compute_output_layer_and_language_model_loss(\n hidden_states_list[mtp_layer_number + 1],\n labels=mtp_labels,\n weight=self.shared_embedding_or_output_weight(),\n sequence_parallel_enabled=self.output_layer.sequence_parallel,\n column_parallel_linear=self.output_layer,\n col_linear_kwargs={\n \"weight\": output_weight,\n \"runtime_gather_output\": runtime_gather_output,\n },\n )\n\n mtp_loss = loss_mask * mtp_loss\n if self.training:\n # TODO(shifangx): remove the use of parallel_state here\n # after moving loss logging to loss_func in pretrain_gpt.py\n MTPLossLoggingHelper.save_loss_to_tracker(\n torch.sum(mtp_loss) / num_tokens,\n mtp_layer_number,\n self.config.mtp_num_layers,\n avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True),\n )\n mtp_loss_scale = self.config.mtp_loss_scaling_factor / self.config.mtp_num_layers\n if self.config.calculate_per_token_loss:\n hidden_states = MTPLossAutoScaler.apply(hidden_states, mtp_loss_scale * mtp_loss)\n else:\n hidden_states = MTPLossAutoScaler.apply(hidden_states, mtp_loss_scale * mtp_loss / num_tokens)\n\n logits, _ = self.output_layer(hidden_states, weight=output_weight, runtime_gather_output=runtime_gather_output)\n # [s b h] => [b s h]\n return logits.transpose(0, 1).contiguous()\n\n\ndef patch_mtp_layer_get_embeddings(model: torch.nn.Module):\n \"\"\"Patch the _get_embeddings method of MultiTokenPredictionLayer\"\"\"\n from megatron.core.models.gpt.gpt_model import GPTModel\n from megatron.core.transformer.multi_token_prediction import MultiTokenPredictionLayer\n\n # Unwrap each model in the actor_module to get the actual GPTModel\n model = _get_patching_model(model)\n # Collect all MultiTokenPredictionLayer instances\n target_layers = []\n\n if isinstance(model, GPTModel):\n # Check if GPTModel has MTP and find the layers\n if hasattr(model, \"mtp\") and hasattr(model.mtp, \"layers\"):\n for layer in model.mtp.layers:\n if isinstance(layer, MultiTokenPredictionLayer):\n target_layers.append(layer)\n elif hasattr(model, \"layers\"):\n # Check if any layer in the model is MultiTokenPredictionLayer\n for layer in model.layers:\n if isinstance(layer, MultiTokenPredictionLayer):\n target_layers.append(layer)\n\n if target_layers:\n for layer in target_layers:\n layer._get_embeddings_backup = layer._get_embeddings\n layer._get_embeddings = _patched_get_embeddings_for_detach.__get__(layer, layer.__class__)\n print(f\"Found and patched {len(target_layers)} MTP layer(s) in any of the actor modules\")\n return True\n else:\n print(\"No MTP layers found to patch in any of the actor modules\")\n return False\n\n\ndef unpatch_mtp_layer_get_embeddings(model: torch.nn.Module):\n \"\"\"Unpatch the _get_embeddings method of MultiTokenPredictionLayer\"\"\"\n from megatron.core.models.gpt.gpt_model import GPTModel\n from megatron.core.transformer.multi_token_prediction import MultiTokenPredictionLayer\n\n # Unwrap each model in the actor_module to get the actual GPTModel\n model = _get_patching_model(model)\n\n # Collect all MultiTokenPredictionLayer instances\n target_layers = []\n\n if isinstance(model, GPTModel):\n # Check if GPTModel has MTP and find the layers\n if hasattr(model, \"mtp\") and hasattr(model.mtp, \"layers\"):\n for layer in model.mtp.layers:\n if isinstance(layer, MultiTokenPredictionLayer):\n target_layers.append(layer)\n elif hasattr(model, \"layers\"):\n # Check if any layer in the model is MultiTokenPredictionLayer\n for layer in model.layers:\n if isinstance(layer, MultiTokenPredictionLayer):\n target_layers.append(layer)\n\n unpatched_count = 0\n for layer in target_layers:\n if hasattr(layer, \"_get_embeddings_backup\"):\n layer._get_embeddings = layer._get_embeddings_backup\n delattr(layer, \"_get_embeddings_backup\")\n unpatched_count += 1\n\n if unpatched_count > 0:\n print(f\"Unpatched {unpatched_count} MTP layer(s)\")\n return True\n return False\n\n\ndef _patched_get_embeddings_for_detach(\n self,\n input_ids: torch.Tensor,\n position_ids: torch.Tensor,\n embedding: Callable,\n hidden_states: torch.Tensor,\n packed_seq_params=None,\n):\n \"\"\"\n Patched version of _get_embeddings method for MultiTokenPredictionLayer.\n\n This is a modified version that you can customize according to your needs.\n The original implementation is preserved below with modifications.\n \"\"\"\n\n # You can modify the logic here as needed\n # For example, you could:\n # - Change the shift amount in roll_tensor\n # - Apply custom transformations to input_ids or position_ids\n # - Add debugging information\n # - Modify the embedding computation\n\n # Original logic with custom modifications\n from megatron.core.transformer.multi_token_prediction import roll_tensor\n from megatron.core.utils import make_viewless_tensor\n\n # Calc logits for the current Multi-Token Prediction (MTP) layers.\n input_ids, _ = roll_tensor(\n input_ids,\n shifts=-1, # You can modify this shift value\n dims=-1,\n cp_group=self.cp_group,\n packed_seq_params=packed_seq_params,\n )\n position_ids, _ = roll_tensor(\n position_ids,\n shifts=-1, # You can modify this shift value\n dims=-1,\n cp_group=self.cp_group,\n packed_seq_params=packed_seq_params,\n )\n\n # embedding computation - you can modify this part\n decoder_input = embedding(input_ids=input_ids, position_ids=position_ids)\n\n # Apply custom transformations if needed\n # For example: decoder_input = some_custom_function(decoder_input)\n\n hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True)\n\n # detach decoder_input and hidden_states\n decoder_input = decoder_input.detach()\n hidden_states = hidden_states.detach()\n\n return input_ids, position_ids, decoder_input, hidden_states\n"}24{"file_name": "verl__models__mcore__patch.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n# there is some bug in mcore 0.12, so we need to patch it\n# 1. `get_query_key_value_tensors` in `multi_latent_attention.py` works wrong when packed_seq_params is not None\n\n\ndef apply_patch():\n import megatron.core\n import torch\n import torch.nn.functional as F\n from megatron.core import parallel_state, tensor_parallel\n from megatron.core.transformer.multi_latent_attention import (\n MLASelfAttention,\n MultiLatentAttention,\n apply_rotary_pos_emb,\n deprecate_inference_params,\n gather_from_sequence_parallel_region,\n gather_from_tensor_model_parallel_region,\n scatter_to_sequence_parallel_region,\n )\n from packaging import version\n\n mcore_ge_013 = version.parse(megatron.core.__version__) >= version.parse(\"0.13.0\")\n\n def patch_get_query_key_value_tensors(\n self,\n hidden_states,\n key_value_states=None,\n position_ids=None,\n packed_seq_params=None,\n inference_context=None,\n *,\n inference_params=None,\n ):\n \"\"\"\n Derives `query`, `key` and `value` tensors from `hidden_states`.\n \"\"\"\n # s = sequence length, b = batch size, h = hidden size, n = num attention heads\n # Attention heads [s, b, n*h]\n assert hidden_states.ndim == 3, f\"hidden_states should be 3D, [s, b, n*h], got {hidden_states.ndim}D\"\n\n inference_context = deprecate_inference_params(inference_context, inference_params)\n\n # =========================================\n # Prepare RoPE and seqlen related params\n # =========================================\n rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len(\n inference_context, None, hidden_states, self.config, packed_seq_params\n )\n\n # rotary_pos_emb:[s, b, 1, 64]\n mscale = 1.0\n if self.config.rope_type == \"rope\":\n packed_seq = packed_seq_params is not None and packed_seq_params.qkv_format == \"thd\"\n try:\n # In case of TypeError: RotaryEmbedding.forward() got an unexpected keyword argument 'packed_seq'\n rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len, packed_seq=packed_seq)\n except TypeError:\n rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len)\n else:\n rotary_pos_emb, mscale = self.rotary_pos_emb(rotary_seq_len)\n\n # =========================================\n # QKV down projection and layernorm\n # =========================================\n if self.config.q_lora_rank is not None:\n # if linear_q_down_proj is ColumnParallelLinear:\n # q_compressed: [s, b, q_lora_rank / TP]\n # elif linear_q_down_proj is Linear:\n # q_compressed: [s / TP, b, q_lora_rank]\n q_compressed, _ = self.linear_q_down_proj(hidden_states)\n\n # When output is sharded (ColumnParallelLinear), two things are needed to be\n # identical to a normal Linear.\n # 1. Manually gather output to restore output dim q_lora_rank;\n # 2. Scatter sequence back to s / TP if sequence-parallel since it was\n # gathered by ColumnParallelLinear.\n if q_compressed.size(-1) != self.config.q_lora_rank:\n q_compressed = gather_from_tensor_model_parallel_region(q_compressed)\n if self.config.sequence_parallel:\n q_compressed = scatter_to_sequence_parallel_region(q_compressed)\n\n q_compressed = self.q_layernorm(q_compressed)\n else:\n q_compressed = hidden_states\n\n # if linear_kv_down_proj is ColumnParallelLinear:\n # kv_combined: [s, b, (kv_lora_rank + qk_pos_emb_head_dim) / TP]\n # elif linear_kv_down_proj is Linear:\n # kv_combined: [s / TP, b, (kv_lora_rank + qk_pos_emb_head_dim)]\n kv_combined, _ = self.linear_kv_down_proj(hidden_states)\n if kv_combined.size(-1) != self.config.kv_lora_rank + self.config.qk_pos_emb_head_dim:\n # kv_combined: [s, b, (kv_lora_rank + qk_pos_emb_head_dim)]\n kv_combined = gather_from_tensor_model_parallel_region(kv_combined)\n # kv_compressed:[s, b, kv_lora_rank], k_pos_emb: [s, b, qk_pos_emb_head_dim]\n kv_compressed, k_pos_emb = torch.split(\n kv_combined, [self.config.kv_lora_rank, self.config.qk_pos_emb_head_dim], dim=-1\n )\n if self.config.sequence_parallel:\n # kv_compressed:[s / TP, b, kv_lora_rank]\n kv_compressed = scatter_to_sequence_parallel_region(kv_compressed)\n else:\n # kv_compressed:[s / TP, b, kv_lora_rank], k_pos_emb: [s / TP, b, qk_pos_emb_head_dim]\n kv_compressed, k_pos_emb = torch.split(\n kv_combined, [self.config.kv_lora_rank, self.config.qk_pos_emb_head_dim], dim=-1\n )\n if parallel_state.get_tensor_model_parallel_world_size() > 1:\n # k_pos_emb: [s, b, qk_pos_emb_head_dim]\n k_pos_emb = gather_from_sequence_parallel_region(k_pos_emb)\n\n kv_compressed = self.kv_layernorm(kv_compressed)\n\n # =========================================\n # QKV up projection and RoPE apply\n # =========================================\n def qkv_up_proj_and_rope_apply(q_compressed, kv_compressed, k_pos_emb, rotary_pos_emb):\n if self.config.q_lora_rank is not None:\n q, _ = self.linear_q_up_proj(q_compressed)\n else:\n # hidden_states:[s, b, 2048], q: [s, b, n * 192]\n q, _ = self.linear_q_proj(q_compressed)\n\n q_len, bsz, _ = q.size()\n\n # q: [s, b, n, 192]\n q = q.view(q_len, bsz, self.num_attention_heads_per_partition, self.q_head_dim)\n\n # kv: [s, b, 2048]\n kv, _ = self.linear_kv_up_proj(kv_compressed)\n\n # kv: [s, b, n, 256]\n kv = kv.view(\n q_len,\n bsz,\n self.num_attention_heads_per_partition,\n self.config.qk_head_dim + self.config.v_head_dim,\n )\n\n cp_size = parallel_state.get_context_parallel_world_size()\n if inference_context is not None:\n # add offset to the sequence start for inference\n sequence_start = inference_context.sequence_len_offset\n sequence_end = sequence_start + q_len\n rotary_pos_emb = rotary_pos_emb[sequence_start:sequence_end]\n elif packed_seq_params is None or cp_size == 1:\n # Shorten rotary_pos_emb to the sequence length when inference_params\n # is not provided. This makes sure we can run forward directly with\n # any sequence length. During training, the sequence length is always\n # the full rotary_pos_emb length, except for sequence packing + CP.\n # When sequence packing and context parallel are both enabled, the\n # position embedding will not split rotary_pos_emb, so it may exceed\n # the sequence length on this CP rank, but we need the full rotary_pos_emb\n # to cover the full sequence, so we do not shorten it here.\n rotary_pos_emb = rotary_pos_emb[0:q_len]\n\n # [s, b, 64] -> [s, b, 1, 64]\n k_pos_emb = torch.unsqueeze(k_pos_emb, 2)\n\n # q: [s, b, n, 128], q_pos_emb: [s, b, n, 64]\n q_no_pe, q_pos_emb = torch.split(q, [self.config.qk_head_dim, self.config.qk_pos_emb_head_dim], dim=-1)\n\n # k_no_pe: [s, b, n, 128], value: [s, b, n, 128]\n k_no_pe, value = torch.split(kv, [self.config.qk_head_dim, self.config.v_head_dim], dim=-1)\n\n if packed_seq_params is not None:\n cu_seqlens_q = packed_seq_params.cu_seqlens_q\n cu_seqlens_kv = packed_seq_params.cu_seqlens_kv\n q_pos_emb = q_pos_emb.squeeze(1)\n k_pos_emb = k_pos_emb.squeeze(1)\n q_no_pe = q_no_pe.squeeze(1)\n k_no_pe = k_no_pe.squeeze(1)\n value = value.squeeze(1)\n else:\n cu_seqlens_q = cu_seqlens_kv = None\n\n # q_pos_emb: [s, b, n, 64], k_pos_emb:[s, b, 1, 64]\n q_pos_emb = apply_rotary_pos_emb(\n q_pos_emb,\n rotary_pos_emb,\n config=self.config,\n cu_seqlens=cu_seqlens_q,\n mscale=mscale,\n )\n k_pos_emb = apply_rotary_pos_emb(\n k_pos_emb,\n rotary_pos_emb,\n config=self.config,\n cu_seqlens=cu_seqlens_kv,\n mscale=mscale,\n )\n\n # query: [s, b, n, 192]\n query = torch.cat([q_no_pe, q_pos_emb], dim=-1)\n if packed_seq_params is not None:\n k_pos_emb = k_pos_emb.expand(-1, self.num_attention_heads_per_partition, -1)\n key = torch.cat([k_no_pe, k_pos_emb], dim=-1)\n else:\n # key: [s, b, n, 192]\n k_pos_emb = k_pos_emb.expand(-1, -1, self.num_attention_heads_per_partition, -1)\n key = torch.cat([k_no_pe, k_pos_emb], dim=-1)\n\n query = query.contiguous()\n key = key.contiguous()\n value = value.contiguous()\n return query, key, value\n\n if self.recompute_up_proj:\n self.qkv_up_checkpoint = tensor_parallel.CheckpointWithoutOutput()\n query, key, value = self.qkv_up_checkpoint.checkpoint(\n qkv_up_proj_and_rope_apply, q_compressed, kv_compressed, k_pos_emb, rotary_pos_emb\n )\n else:\n query, key, value = qkv_up_proj_and_rope_apply(q_compressed, kv_compressed, k_pos_emb, rotary_pos_emb)\n\n return query, key, value\n\n def patch_forward(\n self,\n hidden_states,\n attention_mask,\n key_value_states=None,\n inference_context=None,\n rotary_pos_emb=None,\n rotary_pos_cos=None,\n rotary_pos_sin=None,\n attention_bias=None,\n packed_seq_params=None,\n position_ids=None,\n sequence_len_offset=None,\n *,\n inference_params=None,\n **kwargs,\n ):\n \"\"\"Forward pass for multi-latent attention\"\"\"\n assert attention_bias is None, \"Attention bias should not be passed into MLA.\"\n assert rotary_pos_cos is None and rotary_pos_sin is None, \"MLA does not support Flash Decoding\"\n\n # hidden_states: [sq, b, h]\n\n inference_context = deprecate_inference_params(inference_context, inference_params)\n\n # =====================\n # Query, Key, and Value\n # =====================\n # Get the query, key and value tensors based on the type of attention -\n # self or cross attn.\n # query: [96, 1, 16, 128], key:[96, 1, 16, 128], value:[96, 1, 16, 128]\n query, key, value = self.get_query_key_value_tensors(\n hidden_states,\n key_value_states,\n position_ids,\n packed_seq_params,\n inference_context=inference_context,\n )\n\n # ===================================================\n # Adjust key, value for inference\n # ===================================================\n # rotary_pos_emb = None\n if mcore_ge_013:\n query, key, value, _, attn_mask_type, _ = self._adjust_key_value_for_inference(\n inference_context, query, key, value, rotary_pos_emb=None\n )\n else:\n query, key, value, _, attn_mask_type = self._adjust_key_value_for_inference(\n inference_context, query, key, value, rotary_pos_emb=None\n )\n\n # TODO: Currently, TE can only accept contiguous tensors for MLA\n query = query.contiguous()\n key = key.contiguous()\n value = value.contiguous()\n\n # ==================================\n # core attention computation\n # ==================================\n # Need corresponding TE change\n thd_qkv_format = packed_seq_params and packed_seq_params.qkv_format == \"thd\"\n v_dim = value.shape[-1]\n if thd_qkv_format and query.shape[-1] != v_dim:\n value = F.pad(value, [0, query.shape[-1] - v_dim])\n self.core_attention.hidden_size_per_attention_head_v = value.shape[-1]\n if self.checkpoint_core_attention and self.training:\n core_attn_out = self._checkpointed_attention_forward(\n query, key, value, attention_mask, packed_seq_params=packed_seq_params\n )\n else:\n core_attn_out = self.core_attention(\n query,\n key,\n value,\n attention_mask,\n packed_seq_params=packed_seq_params,\n attn_mask_type=attn_mask_type,\n )\n if thd_qkv_format:\n if core_attn_out.ndim == 2:\n core_attn_out = core_attn_out.reshape(*core_attn_out.shape[:-1], -1, value.shape[-1])\n if query.shape[-1] != v_dim:\n core_attn_out = core_attn_out[..., :v_dim]\n # reshape to same output shape as unpacked case\n # (t, np, hn) -> (t, b=1, h=np*hn)\n # t is the pack size = sum (sq_i)\n # note that batch is a dummy dimension in the packed case\n core_attn_out = core_attn_out.reshape(core_attn_out.size(0), 1, -1)\n\n if self.recompute_up_proj:\n assert self.qkv_up_checkpoint is not None\n self.qkv_up_checkpoint.discard_output_and_register_recompute(core_attn_out)\n self.qkv_up_checkpoint = None\n\n # =================\n # Output. [sq, b, h]\n # =================\n output, bias = self.linear_proj(core_attn_out)\n\n return output, bias\n\n MLASelfAttention.get_query_key_value_tensors = patch_get_query_key_value_tensors\n\n MultiLatentAttention.forward = patch_forward\n\n\ndef apply_patch_mbridge():\n try:\n from megatron.core.utils import get_tensor_model_parallel_group_if_none\n except ImportError:\n import warnings\n\n import megatron.core.utils\n import torch\n from megatron.core import parallel_state\n\n def get_tensor_model_parallel_group_if_none(tp_group, is_expert=False, check_initialized=True):\n \"\"\"Issue a deprecation warning if tp_group is None and return the default tp group.\"\"\"\n if not torch.distributed.is_initialized():\n return None\n if tp_group is None:\n if torch.distributed.is_initialized() and torch.distributed.get_rank() == 0:\n warnings.warn(\n \"Warning: tp_group is None, using default tp group. Passing tp_group will be mandatory soon\",\n DeprecationWarning,\n stacklevel=2,\n )\n if is_expert:\n tp_group = parallel_state.get_expert_tensor_parallel_group(check_initialized=check_initialized)\n else:\n tp_group = parallel_state.get_tensor_model_parallel_group(check_initialized=check_initialized)\n return tp_group\n\n megatron.core.utils.get_tensor_model_parallel_group_if_none = get_tensor_model_parallel_group_if_none\n"}25{"file_name": "verl__models__mcore__qwen2_5_vl__model.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright (c) 2024 Alibaba PAI Team.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nimport logging\n\nimport torch\nfrom megatron.core import InferenceParams, mpu, tensor_parallel\nfrom megatron.core.models.gpt.gpt_model import GPTModel\n\n# from .transformer_config import Qwen2VLTransformerConfig\nfrom megatron.core.packed_seq_params import PackedSeqParams\nfrom megatron.core.transformer import MegatronModule\nfrom megatron.core.transformer.spec_utils import ModuleSpec\nfrom megatron.core.transformer.transformer_config import TransformerConfig\n\nfrom verl.models.mcore.util import preprocess_packed_seqs\n\nfrom .attention import Qwen2_5VLSelfAttention\nfrom .vision_model import Qwen2_5VisionModel\n\n\n# Note: This is under development and may be missing features.\nclass Qwen2_5VLModel(MegatronModule):\n \"\"\"Qwen2.5VL multi-modal model.\n\n Args:\n language_transformer_config (TransformerConfig): Transformer config for the language model.\n language_transformer_layer_spec (ModuleSpec): Specifies module to use for transformer layers of the\n language model.\n language_vocab_size (int): Language model vocabulary size.\n language_max_sequence_length (int): Language model maximum sequence length. This is used for\n positional embedding.\n vision_transformer_config (TransformerConfig): Transformer config for the vision model.\n vision_transformer_layer_spec (ModuleSpec): Specifies module to use for transformer layers of the\n vision model.\n vision_projection_config (TransformerConfig): Config for the projection from vision model outputs to\n language model inputs.\n vision_projection_layer_spec (ModuleSpec): Specifies the module to use for the vision\n projection.\n vision_projection_type (str): Type of the vision projection to use. Default is a 2-layer MLP.\n parallel_output (bool): Do not gather the outputs, keep them split across tensor parallel ranks. This\n is typically True for training and False for inference.\n language_rotary_percent (float): Percent of rotary dimension to use for rotary position embeddings\n in the language model. Defaults to 1.0.\n pre_process (bool): Include the embedding layer in the gpt decoder (used with pipeline parallelism).\n Defaults to True.\n post_process (bool): Include an output layer and a layernorm in the gpt decoder (used with pipeline\n parallelism). Defaults to True.\n add_encoder (bool): Construct the encoder module (used with pipeline parallelism). Defaults to True.\n When we use pipelining, the encoder\n will live on only a subset of the pipeline stages (specifically, only the first stage).\n add_decoder (bool): Construct the decoder module (used with pipeline parallelism). Defaults to True.\n When we use pipelining, the decoder\n will live on only a subset of the pipeline stages (specifically, every stage after the first one).\n img_h (int): The height of each image that the ViT will see.\n img_w (int): The width of each image that the ViT will see.\n patch_dim (int): The size of each patch side.\n img_embedding_idx (int): Index in the language_embeddings tensor where image_embeddings should be\n inserted. Defaults to 0.\n \"\"\"\n\n def __init__(\n self,\n language_transformer_config: TransformerConfig,\n language_transformer_layer_spec: ModuleSpec,\n language_vocab_size: int,\n language_max_sequence_length: int,\n vision_transformer_config: TransformerConfig,\n vision_transformer_layer_spec: ModuleSpec,\n vision_projection_config: TransformerConfig,\n vision_projection_layer_spec: ModuleSpec,\n vision_projection_type: str = \"mlp\",\n parallel_output: bool = True,\n language_rotary_percent: float = 1.0,\n pre_process: bool = True,\n post_process: bool = True,\n add_encoder: bool = True,\n add_decoder: bool = True,\n language_rotary_base: int = 10000,\n fp16_lm_cross_entropy: bool = False,\n language_share_embeddings_and_output_weights: bool = False,\n image_token_id: int = 151655,\n video_token_id: int = 151656,\n ) -> None:\n super().__init__(config=language_transformer_config)\n\n # patch self_attention to use qwen2_5_vl attention\n vision_transformer_layer_spec.submodules.self_attention.module = Qwen2_5VLSelfAttention\n for layer_spec in language_transformer_layer_spec.layer_specs:\n layer_spec.submodules.self_attention.module = Qwen2_5VLSelfAttention\n\n logging.getLogger(__name__).warning(\"Qwen2VL model is under development and may be missing features.\")\n\n self.pre_process = pre_process\n self.post_process = post_process\n self.add_encoder = add_encoder\n self.add_decoder = add_decoder\n\n self.encoder_hidden_state = None\n self.vision_model = None\n self.vision_projection = None\n self.language_model = None\n self.image_token_id = image_token_id\n self.video_token_id = video_token_id\n\n self.square_merge_size = vision_projection_config.ffn_hidden_size // vision_transformer_config.hidden_size\n\n # This attribute is needed to check if an all-reduce is required\n # on the word embeddings inside `finalize_model_grads._allreduce_word_embedding_grads`.\n self.share_embeddings_and_output_weights = False\n if self.pre_process:\n self.vision_model = Qwen2_5VisionModel(\n vision_transformer_config,\n vision_transformer_layer_spec,\n vision_projection_config,\n vision_projection_layer_spec,\n projection_type=vision_projection_type,\n pre_process=True,\n post_process=True,\n )\n\n self.language_model = GPTModel(\n config=language_transformer_config,\n transformer_layer_spec=language_transformer_layer_spec,\n vocab_size=language_vocab_size,\n max_sequence_length=language_max_sequence_length,\n parallel_output=parallel_output,\n position_embedding_type=\"mrope\",\n rotary_percent=language_rotary_percent,\n pre_process=self.pre_process,\n post_process=self.post_process,\n rotary_base=language_rotary_base,\n fp16_lm_cross_entropy=fp16_lm_cross_entropy,\n share_embeddings_and_output_weights=language_share_embeddings_and_output_weights,\n scatter_embedding_sequence_parallel=False,\n )\n assert mpu.get_context_parallel_world_size() <= 1, \"please use mbridge for qwen2_5_vl with context parallelism\"\n self.share_embeddings_and_output_weights = self.language_model.share_embeddings_and_output_weights\n\n def shared_embedding_or_output_weight(self):\n \"\"\"This is a convenience method to surface the language model's word embeddings, which is\n necessary for `finalize_model_grads._allreduce_word_embedding_grads`.\"\"\"\n if self.add_decoder:\n return self.language_model.shared_embedding_or_output_weight()\n return None\n\n def set_input_tensor(self, input_tensor) -> None:\n # This is usually handled in schedules.py but some inference code still\n # gives us non-lists or None\n if not isinstance(input_tensor, list):\n input_tensor = [input_tensor]\n assert len(input_tensor) == 1, \"input_tensor should only be length 1 for Qwen2VL\"\n\n if self.pre_process:\n self.encoder_hidden_state = input_tensor[0]\n else:\n self.language_model.set_input_tensor(input_tensor[0])\n\n def freeze(self, freeze_language_model: bool, freeze_vision_model: bool, freeze_vision_projection: bool):\n \"\"\"Freeze model modules.\n\n Make specific modules non-trainable by setting requires_grad to False for the module's parameters.\n\n Args:\n freeze_language_model (bool): Freeze the language model module.\n freeze_vision_model (bool): Freeze the vision model module.\n freeze_vision_projection (bool): Freeze the vision projection module.\n \"\"\"\n modules = []\n if freeze_language_model and self.language_model is not None:\n modules.append(self.language_model)\n if freeze_vision_model and self.vision_model is not None:\n modules.append(self.vision_model)\n if freeze_vision_projection and self.vision_projection is not None:\n modules.append(self.vision_projection)\n\n for module in modules:\n for param in module.parameters():\n param.requires_grad = False\n\n def forward(\n self,\n input_ids: torch.Tensor,\n position_ids: torch.Tensor,\n attention_mask: torch.Tensor = None,\n labels: torch.Tensor = None,\n inference_params: InferenceParams = None,\n packed_seq_params: PackedSeqParams = None,\n extra_block_kwargs: dict = None,\n pixel_values: torch.Tensor = None,\n pixel_values_videos: torch.Tensor = None,\n image_grid_thw: torch.Tensor = None,\n video_grid_thw: torch.Tensor = None,\n **kwargs,\n ) -> torch.Tensor:\n \"\"\"Forward function of the Qwen2VL model.\n ### there is a workaround for supporting sequence packing with context parallelism\n # cp split with sequence packing will make model lose vision token information, so we need to keep\n # the original input_ids and pack them after vision embedding is calculated,\n # cooporate with verl's models/mcore/model_forward.py\n # pack the combined_embeddings to thd here, we check if packed_seq_params is None to determine if\n # we need to pack the combined_embeddings to thd\n # this function needs the position_ids and attention_mask in BSHD format, no matter use packed_seq or not\n\n Args:\n image_data (torch.Tensor): input image of shape [total_thw_size, n_features].\n input_ids (torch.Tensor): input text ids [batch, text_seq_len].\n position_ids (torch.Tensor): input text position ids [batch, text_seq_len].\n attention_mask (torch.Tensor): attention mask for the language model [batch, 1, combined_seq_len,\n combined_seq_len].\n labels (torch.Tensor): Optional target text labels [batch, combined_seq_len].\n inference_params (InferenceParams): Inference-time parameters including KV cache.\n\n video_start_index:\n 0 -- all video\n len(video_seq) -- all image\n others -- mixture\n *_input_mask: should not be None in the first PP stage\n Returns:\n output (torch.Tensor): Loss of shape [b, s] if labels are provided, otherwise logits of shape\n [b, s, vocab_size].\n \"\"\"\n video_start_index = 0\n vision_grid_thw = None\n vision_data = None\n if image_grid_thw is not None:\n image_mask = input_ids == self.image_token_id\n vision_grid_thw = image_grid_thw\n vision_data = pixel_values\n video_start_index = image_mask.sum().item()\n if video_grid_thw is not None:\n video_mask = input_ids == self.video_token_id\n if vision_grid_thw is not None:\n vision_grid_thw = torch.cat([vision_grid_thw, video_grid_thw], dim=0)\n vision_data = torch.cat([vision_data, pixel_values_videos], dim=0)\n else:\n vision_grid_thw = video_grid_thw\n vision_data = pixel_values_videos\n use_inference_kv_cache = (\n inference_params is not None and \"image_tokens_count\" in inference_params.key_value_memory_dict\n )\n if use_inference_kv_cache:\n raise NotImplementedError()\n\n if self.pre_process:\n vision_embeds = None\n if vision_grid_thw is not None and vision_grid_thw.shape[0] > 0:\n vision_embeds = self.vision_model(\n vision_data=vision_data, # If None, vision model should use intermediate outputs (EPP > 1)\n grid_thw=vision_grid_thw, # should provided in each EPP stage\n )\n\n # If running inference, the language model KV cache will be updated for image token positions.\n # Here we store the image tokens sequence length, which can be used as an offset to the KV cache later.\n if inference_params is not None:\n raise NotImplementedError()\n # inference_params.key_value_memory_dict[\"image_tokens_count\"] = (\n # vision_embeddings.shape[0]\n # )\n\n # If running inference, we can skip image token computation if they were computed already earlier\n # for this sample.\n if use_inference_kv_cache:\n language_embeddings: torch.Tensor = self.language_model.embedding(\n input_ids=input_ids,\n position_ids=None, # NOTE: disable\n ) # [text_seq_len, b, h_language]\n # NOTE: why not cat here? is it the combined embeddings useless?\n combined_embeddings = language_embeddings\n elif vision_embeds is not None:\n if video_start_index == 0:\n image_embeds = None\n video_embeds = vision_embeds\n elif video_start_index == vision_embeds.shape[0]:\n image_embeds = vision_embeds\n video_embeds = None\n elif 0 < video_start_index < vision_embeds.shape[0]:\n image_embeds = vision_embeds[:video_start_index]\n video_embeds = vision_embeds[video_start_index:]\n else:\n raise ValueError(\n f\"Expect video token start index in range [0, {vision_embeds.shape[0]}], but got \"\n f\"{video_start_index}\"\n )\n\n combined_embeddings = self.language_model.embedding(\n input_ids=input_ids,\n position_ids=None, # NOTE: disable\n ) # [text_seq_len, b, h_language]\n\n if image_embeds is not None or video_embeds is not None:\n combined_embeddings = combined_embeddings.transpose(0, 1).contiguous()\n if image_embeds is not None:\n image_mask = (input_ids == self.image_token_id).contiguous()\n if image_mask.sum() > 0:\n combined_embeddings = combined_embeddings.clone()\n combined_embeddings[image_mask] = image_embeds.to(\n dtype=combined_embeddings.dtype, device=combined_embeddings.device\n )\n if video_embeds is not None:\n video_mask = (input_ids == self.video_token_id).contiguous()\n if video_mask.sum() > 0:\n combined_embeddings = combined_embeddings.clone()\n combined_embeddings[video_mask] = video_embeds.to(\n dtype=combined_embeddings.dtype, device=combined_embeddings.device\n )\n combined_embeddings = combined_embeddings.transpose(0, 1).contiguous()\n\n else:\n combined_embeddings = self.language_model.embedding(\n input_ids=input_ids,\n position_ids=None, # NOTE: disable\n ) # [text_seq_len, b, h_language]\n\n if packed_seq_params is not None:\n combined_embeddings = (\n preprocess_packed_seqs(\n combined_embeddings.transpose(0, 1).contiguous(), attention_mask, pre_process=True\n )[0]\n .transpose(0, 1)\n .contiguous()\n )\n if self.config.sequence_parallel:\n combined_embeddings = tensor_parallel.scatter_to_sequence_parallel_region(combined_embeddings)\n combined_embeddings = combined_embeddings.contiguous()\n else:\n combined_embeddings = None\n from .rope_utils import get_rope_index\n\n # BSHD\n position_ids, _ = get_rope_index(\n input_ids,\n image_grid_thw=image_grid_thw,\n video_grid_thw=video_grid_thw,\n attention_mask=attention_mask,\n )\n # THD\n if packed_seq_params is not None:\n position_ids = (\n preprocess_packed_seqs(position_ids.permute(1, 2, 0), attention_mask, pre_process=True)[0]\n .permute(2, 0, 1)\n .contiguous()\n )\n attention_mask = None\n\n output = self.language_model(\n input_ids=None,\n position_ids=position_ids, # None in encoder\n attention_mask=attention_mask, # None in encoder\n decoder_input=combined_embeddings, # only not None in the first decoder PP stage\n labels=labels, # only not None in the last decoder PP stage\n # inference_params=inference_params, # currently always None\n packed_seq_params=packed_seq_params, # currently always None\n **(extra_block_kwargs or {}),\n **kwargs,\n )\n\n return output\n"}26{"file_name": "verl__models__mcore__qwen2_5_vl__rope_utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright (c) 2024 Alibaba PAI Team.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nfrom __future__ import annotations\n\nimport logging\nfrom typing import Optional\n\nimport torch\nfrom megatron.core.models.common.embeddings.rope_utils import *\nfrom megatron.core.models.common.embeddings.rope_utils import _apply_rotary_pos_emb_bshd\nfrom torch import Tensor\n\nlogger = logging.getLogger(__name__)\n\n\n# Slightly modified from Qwen2VLForConditionalGeneration.get_rope_index\ndef get_rope_index(\n input_ids: Optional[torch.LongTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n second_per_grid_ts: Optional[torch.Tensor] = None,\n attention_mask: Optional[torch.Tensor] = None,\n):\n \"\"\"\n Calculate the 3D rope index based on image and video's temporal, height and width in LLM.\n\n Explanation:\n\n Each embedding sequence contains vision embedding and text embedding or just contains text embedding.\n\n For pure text embedding sequence, the rotary position embedding has no difference with modern LLMs.\n\n Examples:\n\n input_ids: [T T T T T], here T is for text.\n temporal position_ids: [0, 1, 2, 3, 4]\n height position_ids: [0, 1, 2, 3, 4]\n width position_ids: [0, 1, 2, 3, 4]\n\n For vision and text embedding sequence, we calculate 3D rotary position embedding for vision part\n and 1D rotary position embedding for text part.\n\n Examples:\n\n Temporal (Time): 3 patches, representing different segments of the video in time.\n Height: 2 patches, dividing each frame vertically.\n Width: 2 patches, dividing each frame horizontally.\n We also have some important parameters:\n fps (Frames Per Second): The video's frame rate, set to 1. This means one frame is processed each\n second.\n tokens_per_second: This is a crucial parameter. It dictates how many \"time-steps\" or \"temporal\n tokens\" are conceptually packed into a one-second interval of the video.\n In this case, we have 25 tokens per second. So each second of the video will be\n represented with 25 separate time points. It essentially defines the temporal\n granularity.\n temporal_patch_size: The number of frames that compose one temporal patch. Here, it's 2 frames.\n interval: The step size for the temporal position IDs, calculated as tokens_per_second *\n temporal_patch_size / fps. In this case, 25 * 2 / 1 = 50. This means that each temporal patch will be\n have a difference of 50 in the temporal position IDs.\n input_ids: [V V V V V V V V V V V V T T T T T], here V is for vision.\n vision temporal position_ids: [0, 0, 0, 0, 50, 50, 50, 50, 100, 100, 100, 100]\n vision height position_ids: [0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1]\n vision width position_ids: [0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1]\n text temporal position_ids: [101, 102, 103, 104, 105]\n text height position_ids: [101, 102, 103, 104, 105]\n text width position_ids: [101, 102, 103, 104, 105]\n Here we calculate the text start position_ids as the max vision position_ids plus 1.\n\n Args:\n input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):\n Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide\n it.\n image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):\n The temporal, height and width of feature shape of each image in LLM.\n video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):\n The temporal, height and width of feature shape of each video in LLM.\n second_per_grid_ts (`torch.Tensor` of shape `(num_videos)`, *optional*):\n The time interval (in seconds) for each grid along the temporal dimension in the 3D position IDs.\n attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):\n Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:\n\n - 1 for tokens that are **not masked**,\n - 0 for tokens that are **masked**.\n\n Returns:\n position_ids (`torch.LongTensor` of shape `(3, batch_size, sequence_length)`)\n mrope_position_deltas (`torch.Tensor` of shape `(batch_size)`)\n \"\"\"\n spatial_merge_size = 2\n tokens_per_second = 2\n image_token_id = 151655\n video_token_id = 151656\n vision_start_token_id = 151652\n mrope_position_deltas = []\n if input_ids is not None and (image_grid_thw is not None or video_grid_thw is not None):\n total_input_ids = input_ids\n if attention_mask is None:\n attention_mask = torch.ones_like(total_input_ids)\n position_ids = torch.ones(\n 3,\n input_ids.shape[0],\n input_ids.shape[1],\n dtype=input_ids.dtype,\n device=input_ids.device,\n )\n image_index, video_index = 0, 0\n attention_mask = attention_mask.to(total_input_ids.device)\n for i, input_ids in enumerate(total_input_ids):\n input_ids = input_ids[attention_mask[i] == 1]\n image_nums, video_nums = 0, 0\n vision_start_indices = torch.argwhere(input_ids == vision_start_token_id).squeeze(1)\n vision_tokens = input_ids[vision_start_indices + 1]\n image_nums = (vision_tokens == image_token_id).sum()\n video_nums = (vision_tokens == video_token_id).sum()\n input_tokens = input_ids.tolist()\n llm_pos_ids_list: list = []\n st = 0\n remain_images, remain_videos = image_nums, video_nums\n for _ in range(image_nums + video_nums):\n if image_token_id in input_tokens and remain_images > 0:\n ed_image = input_tokens.index(image_token_id, st)\n else:\n ed_image = len(input_tokens) + 1\n if video_token_id in input_tokens and remain_videos > 0:\n ed_video = input_tokens.index(video_token_id, st)\n else:\n ed_video = len(input_tokens) + 1\n if ed_image < ed_video:\n t, h, w = (\n image_grid_thw[image_index][0],\n image_grid_thw[image_index][1],\n image_grid_thw[image_index][2],\n )\n second_per_grid_t = 0\n image_index += 1\n remain_images -= 1\n ed = ed_image\n\n else:\n t, h, w = (\n video_grid_thw[video_index][0],\n video_grid_thw[video_index][1],\n video_grid_thw[video_index][2],\n )\n if second_per_grid_ts is not None:\n second_per_grid_t = second_per_grid_ts[video_index]\n else:\n second_per_grid_t = 1.0\n video_index += 1\n remain_videos -= 1\n ed = ed_video\n llm_grid_t, llm_grid_h, llm_grid_w = (\n t.item(),\n h.item() // spatial_merge_size,\n w.item() // spatial_merge_size,\n )\n text_len = ed - st\n\n st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0\n llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)\n\n range_tensor = torch.arange(llm_grid_t).view(-1, 1)\n expanded_range = range_tensor.expand(-1, llm_grid_h * llm_grid_w)\n\n time_tensor = expanded_range * second_per_grid_t * tokens_per_second\n\n time_tensor_long = time_tensor.long()\n t_index = time_tensor_long.flatten()\n\n h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten()\n w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten()\n llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + text_len + st_idx)\n st = ed + llm_grid_t * llm_grid_h * llm_grid_w\n\n if st < len(input_tokens):\n st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0\n text_len = len(input_tokens) - st\n llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)\n\n llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)\n position_ids[..., i, attention_mask[i] == 1] = llm_positions.to(position_ids.device)\n mrope_position_deltas.append(llm_positions.max() + 1 - len(total_input_ids[i]))\n mrope_position_deltas = torch.tensor(mrope_position_deltas, device=input_ids.device).unsqueeze(1)\n return position_ids, mrope_position_deltas\n else:\n if attention_mask is not None:\n position_ids = attention_mask.long().cumsum(-1) - 1\n position_ids.masked_fill_(attention_mask == 0, 1)\n position_ids = position_ids.unsqueeze(0).expand(3, -1, -1).to(attention_mask.device)\n max_position_ids = position_ids.max(0, keepdim=False)[0].max(-1, keepdim=True)[0]\n mrope_position_deltas = max_position_ids + 1 - attention_mask.shape[-1]\n else:\n position_ids = (\n torch.arange(input_ids.shape[1], device=input_ids.device)\n .view(1, 1, -1)\n .expand(3, input_ids.shape[0], -1)\n )\n mrope_position_deltas = torch.zeros(\n [input_ids.shape[0], 1],\n device=input_ids.device,\n dtype=input_ids.dtype,\n )\n\n return position_ids, mrope_position_deltas\n\n\ndef apply_rotary_pos_emb_thd_absolute(\n t: Tensor, cu_seqlens: Tensor, freqs: Tensor, rotary_interleaved: bool = False\n) -> Tensor:\n \"\"\"A baseline implementation of applying RoPE for `thd` format.\n\n Args:\n t (Tensor): Input tensor T is of shape [t, h, d]\n cu_seqlens(Tensor): Cumulative sum of sequence lengths in a batch for `t`,\n with shape [b + 1] and dtype torch.int32.\n freqs (Tensor): Rotary Positional embedding tensor freq is of shape [max_s, 1, 1, d]\n\n Returns:\n Tensor: Shape [t, h, d]. The input tensor after applying RoPE.\n \"\"\"\n return _apply_rotary_pos_emb_bshd(t[:, None], freqs, rotary_interleaved=rotary_interleaved).squeeze(1)\n\n\ndef apply_rotary_pos_emb_absolute(\n t: Tensor,\n freqs: Tensor,\n config: TransformerConfig,\n cu_seqlens: Optional[Tensor] = None,\n):\n \"\"\"\n Reroute to the appropriate apply_rotary_pos_emb function depending on\n bshd (conventional) / thd (packed seq) format\n\n In Qwen2-VL, the shape of freqs is (seq_length, bs, 1, 2 * dim) instead of [max_seqlen, 1, 1, 2 * dim]\n \"\"\"\n\n if config.apply_rope_fusion:\n if cu_seqlens is None:\n # NOTE: TE backends do not support mRoPE in bshd format when bs > 1\n if freqs.shape[1] > 1:\n return _apply_rotary_pos_emb_bshd(t, freqs, rotary_interleaved=config.rotary_interleaved)\n else:\n return fused_apply_rotary_pos_emb(t, freqs)\n else:\n # NOTE: as expected, thd format can use bshd\n return fused_apply_rotary_pos_emb(t[:, None], freqs).squeeze(1)\n else:\n if cu_seqlens is None:\n return _apply_rotary_pos_emb_bshd(t, freqs, rotary_interleaved=config.rotary_interleaved)\n else:\n return apply_rotary_pos_emb_thd_absolute(t, cu_seqlens, freqs, rotary_interleaved=config.rotary_interleaved)\n"}27{"file_name": "verl__models__mcore__qwen2_5_vl__vision_config.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright (c) 2024 Alibaba PAI Team.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport torch\nfrom megatron.core import parallel_state\nfrom megatron.core.transformer import TransformerConfig\n\n\ndef get_vision_model_config(config: TransformerConfig) -> TransformerConfig:\n # Given a Transformer Config from decoder, build vision encoder config\n # diff: out_hidden_size & intermediate_size\n\n # mlp: hidden_size -> intermediate_size -> embed_dim, silu\n # NOTE: here we provide a workaround to solve the wrong layer amount when VPP of decoder is on\n if config.num_layers in [28, 36]:\n config.ffn_hidden_size = 3420\n else:\n config.ffn_hidden_size = 3456\n\n if parallel_state.get_virtual_pipeline_model_parallel_world_size() is not None:\n config.num_layers = 32 * parallel_state.get_virtual_pipeline_model_parallel_world_size() # depth\n else:\n config.num_layers = 32 # depth\n config.num_attention_heads = 16 # num_heads\n config.add_bias_linear = True # all nn.Linear has bias (MLP, attn)\n config.add_qkv_bias = True # qkv_proj in attn has bias\n config.hidden_size = 1280 # hidden_size\n config.hidden_dropout = 0.0\n config.attention_dropout = 0.0\n\n # config.gated_linear_unit = False # no gated\n # config.activation_func = quick_gelu # hidden_act\n config.kv_channels = config.hidden_size // config.num_attention_heads\n config.num_query_groups = config.num_attention_heads # no GQA\n config.layernorm_zero_centered_gamma = False # False\n config.apply_query_key_layer_scaling = False # factor=math.sqrt(head_dim)\n config.bias_activation_fusion = False # no swiglu, set false\n config.bias_dropout_fusion = False # no dropout, set false\n config.attention_softmax_in_fp32 = True # use True\n # config.normalization = 'LayerNorm' # use RMSNorm\n config.seq_length = 1\n\n config.tp_comm_overlap = False\n config.sequence_parallel = False\n config.temporal_patch_size = 2\n config.patch_size = 14\n config.in_channels = 3\n config.spatial_merge_size = 2\n\n config.fullatt_block_indexes = [7, 15, 23, 31]\n config._qwen2_5_vl_window_size = 112\n return config\n\n\ndef get_vision_projection_config(\n config: TransformerConfig, embed_dim: int, spatial_merge_size: int\n) -> TransformerConfig:\n # merger:\n # context_dim = hidden_size * merge_size**2\n # out_hidden_size = hidden_size\n # context_dim -> context_dim -> out_hidden_size\n # MLP:\n # input_size -> ffn_hidden_size -> hidden_size\n # spec: LN -> Linear(bias=True) -> GELU -> Linear(bias=True)\n config.gated_linear_unit = False\n config.bias_activation_fusion = False\n config.add_bias_linear = True\n config.ffn_hidden_size = embed_dim * (spatial_merge_size**2)\n config.activation_func = torch.nn.functional.gelu\n config.tp_comm_overlap = False\n config.sequence_parallel = False\n return config\n"}28{"file_name": "verl__models__mcore__qwen2_5_vl__vision_model.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright (c) 2024 Alibaba PAI Team.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Optional\n\nimport torch\nfrom megatron.core import InferenceParams\nfrom megatron.core.models.common.vision_module.vision_module import VisionModule\nfrom megatron.core.models.vision.multimodal_projector import MultimodalProjector\nfrom megatron.core.packed_seq_params import PackedSeqParams\nfrom megatron.core.transformer.enums import ModelType\nfrom megatron.core.transformer.spec_utils import ModuleSpec\nfrom megatron.core.transformer.transformer_config import TransformerConfig\nfrom torch import nn\nfrom torch.nn import functional as F\n\nfrom .vision_transformer_block import Qwen2_5VisionTransformerBlock as TransformerBlock\n\n\n# copied from https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen2_vl/modeling_qwen2_vl.py\nclass PatchEmbed(nn.Module):\n def __init__(\n self,\n patch_size: int = 14,\n temporal_patch_size: int = 2,\n in_channels: int = 3,\n embed_dim: int = 1152,\n ) -> None:\n super().__init__()\n self.patch_size = patch_size\n self.temporal_patch_size = temporal_patch_size\n self.in_channels = in_channels\n self.embed_dim = embed_dim\n\n kernel_size = [temporal_patch_size, patch_size, patch_size]\n self.proj = nn.Conv3d(in_channels, embed_dim, kernel_size=kernel_size, stride=kernel_size, bias=False)\n\n def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:\n target_dtype = self.proj.weight.dtype\n hidden_states = hidden_states.view(\n -1, self.in_channels, self.temporal_patch_size, self.patch_size, self.patch_size\n )\n hidden_states = self.proj(hidden_states.to(dtype=target_dtype)).view(-1, self.embed_dim)\n return hidden_states\n\n\n# copied from https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen2_vl/modeling_qwen2_vl.py\nclass VisionRotaryEmbedding(nn.Module):\n def __init__(self, dim: int, theta: float = 10000.0) -> None:\n super().__init__()\n inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float) / dim))\n self.register_buffer(\"inv_freq\", inv_freq, persistent=False)\n\n def forward(self, seqlen: int) -> torch.Tensor:\n seq = torch.arange(seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype)\n freqs = torch.outer(seq, self.inv_freq)\n return freqs.float()\n\n\nclass Qwen2_5VisionModel(VisionModule):\n \"\"\"Qwen2.5 ViT vision model.\n\n Args:\n transformer_config (TransformerConfig): Transformer config.\n transformer_layer_spec (ModuleSpec): Specifies module to use for transformer layers.\n ln_pre_impl (ModuleSpec or type): Specifies the layer norm type to use for ln_pre.\n add_class_token (bool, optional): Include a class token. Defaults to True.\n class_token_len (int): Class token length. Defaults to 1 but 8 may be faster.\n patch_dim (int): Image patch size.\n img_h (int): Input image height.\n img_w (int): Input image width.\n \"\"\"\n\n def __init__(\n self,\n transformer_config: TransformerConfig,\n transformer_layer_spec: ModuleSpec,\n projection_config: TransformerConfig,\n projection_layer_spec: ModuleSpec,\n projection_type: str = \"mlp\",\n pre_process: bool = True,\n post_process: bool = False,\n ) -> None:\n super().__init__(config=transformer_config)\n\n self.spatial_merge_size = transformer_config.spatial_merge_size\n\n embed_dim = transformer_config.hidden_size\n num_heads = transformer_config.num_attention_heads\n temporal_patch_size = transformer_config.temporal_patch_size\n patch_size = transformer_config.patch_size\n in_channels = transformer_config.in_channels\n\n self.patch_size = transformer_config.patch_size\n self.fullatt_block_indexes = transformer_config.fullatt_block_indexes\n self.window_size = transformer_config._qwen2_5_vl_window_size\n self.spatial_merge_unit = self.spatial_merge_size * self.spatial_merge_size\n\n self.max_sequence_length = transformer_config.seq_length\n self.patch_embed = PatchEmbed(\n patch_size=patch_size,\n temporal_patch_size=temporal_patch_size,\n in_channels=in_channels,\n embed_dim=embed_dim,\n )\n\n head_dim = embed_dim // num_heads\n self.rotary_pos_emb = VisionRotaryEmbedding(head_dim // 2)\n\n self.model_type = ModelType.encoder_or_decoder\n self.pre_process = pre_process\n self.post_process = post_process\n\n # Transformer layers.\n # TODO: Follow-up changes will make pre and post_process configurable. They are needed for supporting\n # pipeline parallelism.\n # NOTE: a final layer norm and/or linear layer present in some implementations are omitted here.\n self.decoder = TransformerBlock(\n config=transformer_config,\n spec=transformer_layer_spec,\n pre_process=self.pre_process,\n post_process=self.post_process,\n post_layer_norm=True,\n )\n\n self.merge_hidden_size = projection_config.ffn_hidden_size\n self.square_merge_size = self.merge_hidden_size // embed_dim\n\n if self.post_process:\n self.projection = MultimodalProjector(\n projection_config, projection_layer_spec, projection_type, projection_config.ffn_hidden_size\n )\n else:\n self.projection = None\n\n self.input_tensor = None\n\n def set_input_tensor(self, input_tensor: torch.Tensor) -> None:\n \"\"\"Sets input tensor to the model.\n\n Args:\n input_tensor (Tensor): Sets the input tensor for the model.\n \"\"\"\n if self.pre_process: # always True\n self.input_tensor = input_tensor\n else:\n raise NotImplementedError()\n\n def rot_pos_emb(self, grid_thw):\n pos_ids = []\n for t, h, w in grid_thw:\n hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w)\n hpos_ids = hpos_ids.reshape(\n h // self.spatial_merge_size,\n self.spatial_merge_size,\n w // self.spatial_merge_size,\n self.spatial_merge_size,\n )\n hpos_ids = hpos_ids.permute(0, 2, 1, 3)\n hpos_ids = hpos_ids.flatten()\n\n wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1)\n wpos_ids = wpos_ids.reshape(\n h // self.spatial_merge_size,\n self.spatial_merge_size,\n w // self.spatial_merge_size,\n self.spatial_merge_size,\n )\n wpos_ids = wpos_ids.permute(0, 2, 1, 3)\n wpos_ids = wpos_ids.flatten()\n pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1))\n pos_ids = torch.cat(pos_ids, dim=0).to(grid_thw.device)\n max_grid_size = grid_thw[:, 1:].max()\n rotary_pos_emb_full = self.rotary_pos_emb(max_grid_size).to(grid_thw.device)\n rotary_pos_emb = rotary_pos_emb_full[pos_ids].flatten(1)\n return rotary_pos_emb\n\n def get_window_index(self, grid_thw):\n window_index: list = []\n cu_window_seqlens: list = [0]\n window_index_id = 0\n vit_merger_window_size = self.window_size // self.spatial_merge_size // self.patch_size\n\n for grid_t, grid_h, grid_w in grid_thw:\n llm_grid_h, llm_grid_w = (\n grid_h // self.spatial_merge_size,\n grid_w // self.spatial_merge_size,\n )\n index = torch.arange(grid_t * llm_grid_h * llm_grid_w).reshape(grid_t, llm_grid_h, llm_grid_w)\n pad_h = vit_merger_window_size - llm_grid_h % vit_merger_window_size\n pad_w = vit_merger_window_size - llm_grid_w % vit_merger_window_size\n num_windows_h = (llm_grid_h + pad_h) // vit_merger_window_size\n num_windows_w = (llm_grid_w + pad_w) // vit_merger_window_size\n index_padded = F.pad(index, (0, pad_w, 0, pad_h), \"constant\", -100)\n index_padded = index_padded.reshape(\n grid_t,\n num_windows_h,\n vit_merger_window_size,\n num_windows_w,\n vit_merger_window_size,\n )\n index_padded = index_padded.permute(0, 1, 3, 2, 4).reshape(\n grid_t,\n num_windows_h * num_windows_w,\n vit_merger_window_size,\n vit_merger_window_size,\n )\n seqlens = (index_padded != -100).sum([2, 3]).reshape(-1)\n index_padded = index_padded.reshape(-1)\n index_new = index_padded[index_padded != -100]\n window_index.append(index_new + window_index_id)\n cu_seqlens_tmp = seqlens.cumsum(0) * self.spatial_merge_unit + cu_window_seqlens[-1]\n cu_window_seqlens.extend(cu_seqlens_tmp.tolist())\n window_index_id += (grid_t * llm_grid_h * llm_grid_w).item()\n window_index = torch.cat(window_index, dim=0)\n\n return window_index, cu_window_seqlens\n\n def forward(\n self,\n vision_data: Optional[torch.Tensor],\n grid_thw: torch.Tensor,\n inference_params: Optional[InferenceParams] = None,\n extra_block_kwargs: dict = None,\n ) -> torch.Tensor:\n \"\"\"Forward function of the Qwen2 Vision Model. This function passes the input tensors\n through the embedding layer and then the transformer.\n\n Args:\n x (torch.Tensor): input image/video data of shape [n_tokens, n_dims]\n grid_thw (torch.Tensor): the size tensor indicates grid size of each image/frame\n packed_seq_params (PackedSeqParams): parameters to build attention mask in the backend\n\n Returns:\n x (torch.Tensor): output after final transformer block of shape [b, s, h].\n \"\"\"\n assert grid_thw is not None\n assert self.input_tensor is None\n assert inference_params is None\n\n # Rotary positional embeddings (embedding is None for PP intermediate devices)\n vision_data = self.patch_embed(vision_data)\n window_index, cu_window_seqlens = self.get_window_index(grid_thw)\n cu_window_seqlens = torch.tensor(\n cu_window_seqlens,\n device=vision_data.device,\n dtype=torch.int32,\n )\n cu_window_seqlens = torch.unique_consecutive(cu_window_seqlens)\n\n seq_len, _ = vision_data.size()\n vision_data = vision_data.reshape(seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1)\n vision_data = vision_data[window_index, :, :]\n vision_data = vision_data.reshape(seq_len, 1, -1)\n\n rotary_pos_emb = self.rot_pos_emb(grid_thw)\n rotary_pos_emb = rotary_pos_emb.reshape(seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1)\n rotary_pos_emb = rotary_pos_emb[window_index, :, :]\n rotary_pos_emb = rotary_pos_emb.reshape(seq_len, 1, 1, -1).repeat(1, 1, 1, 2)\n\n hidden_states = self.decoder(\n hidden_states=vision_data,\n attention_mask=None,\n inference_params=inference_params,\n rotary_pos_emb=rotary_pos_emb,\n packed_seq_params=self.build_packed_seq_params(None, cu_window_seqlens),\n packed_seq_params_full=self.build_packed_seq_params(grid_thw),\n fullatt_block_indexes=self.fullatt_block_indexes,\n **(extra_block_kwargs or {}),\n )\n\n hidden_states = self.projection(hidden_states.view(-1, self.merge_hidden_size))\n reverse_indices = torch.argsort(window_index)\n return hidden_states[reverse_indices, :]\n\n def build_packed_seq_params(\n self,\n grid_thw: Optional[torch.Tensor],\n cu_seqlens: Optional[torch.Tensor] = None,\n ) -> PackedSeqParams:\n # NOTE: each frame is a sequence (rather than each grid)\n if grid_thw is not None:\n seqlens = torch.repeat_interleave(grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0])\n cu_seqlens = seqlens.cumsum(dim=0)\n cu_seqlens = F.pad(cu_seqlens, (1, 0), value=0).int()\n else:\n seqlens = cu_seqlens[1:] - cu_seqlens[:-1]\n\n max_seqlen_q = seqlens.max()\n return PackedSeqParams(\n cu_seqlens_q=cu_seqlens,\n cu_seqlens_kv=cu_seqlens,\n qkv_format=\"thd\",\n max_seqlen_q=max_seqlen_q,\n max_seqlen_kv=max_seqlen_q,\n )\n"}29{"file_name": "verl__models__mcore__qwen2_5_vl__vision_transformer_block.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright (c) 2024 Alibaba PAI Team.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nfrom megatron.core.transformer.transformer_block import *\n\n\nclass Qwen2_5VisionTransformerBlock(TransformerBlock):\n def _checkpointed_forward(\n self,\n hidden_states: Tensor,\n attention_mask: Tensor,\n context: Tensor,\n context_mask: Tensor,\n rotary_pos_emb: Tensor,\n attention_bias: Tensor,\n packed_seq_params: PackedSeqParams,\n packed_seq_params_full: PackedSeqParams,\n fullatt_block_indexes,\n ):\n \"\"\"Forward method with activation checkpointing.\"\"\"\n\n def custom(start: int, end: int):\n def custom_forward(hidden_states, attention_mask, context, context_mask, rotary_pos_emb):\n for index in range(start, end):\n if index in fullatt_block_indexes:\n packed_seq_params_now = packed_seq_params_full\n else:\n packed_seq_params_now = packed_seq_params\n layer = self._get_layer(index)\n hidden_states, context = layer(\n hidden_states=hidden_states,\n attention_mask=attention_mask,\n context=context,\n context_mask=context_mask,\n rotary_pos_emb=rotary_pos_emb,\n attention_bias=attention_bias,\n inference_context=None,\n packed_seq_params=packed_seq_params_now,\n )\n return hidden_states, context\n\n return custom_forward\n\n def checkpoint_handler(forward_func):\n \"\"\"Determines whether to use the `te_checkpoint` or `tensor_parallel.checkpoint`\"\"\"\n if self.config.fp8:\n return te_checkpoint(\n forward_func,\n self.config.distribute_saved_activations,\n tensor_parallel.random.get_cuda_rng_tracker,\n parallel_state.get_tensor_model_parallel_group(),\n hidden_states,\n attention_mask,\n context,\n context_mask,\n rotary_pos_emb,\n )\n else:\n return tensor_parallel.checkpoint(\n forward_func,\n self.config.distribute_saved_activations,\n hidden_states,\n attention_mask,\n context,\n context_mask,\n rotary_pos_emb,\n )\n\n if self.config.recompute_method == \"uniform\":\n # Uniformly divide the total number of Transformer layers and checkpoint\n # the input activation of each divided chunk.\n # A method to further reduce memory usage reducing checkpoints.\n layer_idx = 0\n while layer_idx < self.num_layers_per_pipeline_rank:\n hidden_states, context = checkpoint_handler(\n custom(layer_idx, layer_idx + self.config.recompute_num_layers)\n )\n\n layer_idx += self.config.recompute_num_layers\n\n elif self.config.recompute_method == \"block\":\n # Checkpoint the input activation of only a set number of individual\n # Transformer layers and skip the rest.\n # A method fully use the device memory removing redundant re-computation.\n recompute_skip_num_layers = 0\n for layer_idx in range(self.num_layers_per_pipeline_rank):\n # Skip recomputation when input grad computation is not needed.\n # Need to have at least one input tensor with gradient computation\n # for re-enterant autograd engine.\n if self.config.fp8 and not hidden_states.requires_grad:\n recompute_skip_num_layers += 1\n if (\n layer_idx >= recompute_skip_num_layers\n and layer_idx < self.config.recompute_num_layers + recompute_skip_num_layers\n ):\n hidden_states, context = checkpoint_handler(custom(layer_idx, layer_idx + 1))\n else:\n hidden_states, context = custom(layer_idx, layer_idx + 1)(\n hidden_states, attention_mask, context, context_mask, rotary_pos_emb\n )\n else:\n raise ValueError(\"Invalid activation recompute method.\")\n\n return hidden_states\n\n def forward(\n self,\n hidden_states: Union[Tensor, WrappedTensor],\n attention_mask: Optional[Tensor],\n context: Optional[Tensor] = None,\n context_mask: Optional[Tensor] = None,\n rotary_pos_emb: Optional[Tensor] = None,\n rotary_pos_cos: Optional[Tensor] = None,\n rotary_pos_sin: Optional[Tensor] = None,\n attention_bias: Optional[Tensor] = None,\n inference_context: Optional[BaseInferenceContext] = None,\n packed_seq_params: Optional[PackedSeqParams] = None,\n sequence_len_offset: Optional[Tensor] = None,\n packed_seq_params_full: PackedSeqParams = None,\n fullatt_block_indexes=None,\n *,\n inference_params: Optional[BaseInferenceContext] = None,\n ):\n \"\"\"\n Perform the forward pass through the transformer block.\n\n This method handles the core computation of the transformer, including\n self-attention, optional cross-attention, and feed-forward operations.\n\n Args:\n hidden_states (Union[Tensor, WrappedTensor]): Input tensor of shape [s, b, h]\n where s is the sequence length, b is the batch size, and h is the hidden size.\n Can be passed as a WrappedTensor during inference to avoid an obsolete\n reference in the calling function.\n attention_mask (Tensor): Boolean tensor of shape [1, 1, s, s] for masking\n self-attention.\n context (Tensor, optional): Context tensor for cross-attention.\n context_mask (Tensor, optional): Mask for cross-attention context\n rotary_pos_emb (Tensor, optional): Rotary positional embeddings.\n attention_bias (Tensor): Bias tensor for Q * K.T of shape in shape broadcastable\n to [b, num_head, sq, skv], e.g. [1, 1, sq, skv].\n Used as an alternative to apply attention mask for TE cuDNN attention.\n inference_context (BaseInferenceContext, optional): Parameters for inference-time\n optimizations.\n packed_seq_params (PackedSeqParams, optional): Parameters for packed sequence\n processing.\n\n Returns:\n Union[Tensor, Tuple[Tensor, Tensor]]: The output hidden states tensor of shape\n [s, b, h], and optionally the updated context tensor if cross-attention is used.\n \"\"\"\n\n inference_context = deprecate_inference_params(inference_context, inference_params)\n\n # Delete the obsolete reference to the initial input tensor if necessary\n if isinstance(hidden_states, WrappedTensor):\n hidden_states = hidden_states.unwrap()\n\n if not self.pre_process:\n # See set_input_tensor()\n hidden_states = self.input_tensor\n\n # Update the inference parameters with the current batch size in case it is variable\n if inference_context and not self.training:\n inference_context.current_batch_size = hidden_states.size(1)\n\n # Viewless tensor.\n # - We only need to create a viewless tensor in the case of micro batch\n # size (mbs) == 1, since in this case, 'hidden_states.transpose()'\n # above creates a view tensor, and '.contiguous()' is a pass-through.\n # For mbs >= 2, '.contiguous()' creates a new tensor, eliminating\n # the need to make it viewless.\n #\n # However, we don't explicitly check mbs == 1 here because\n # make_viewless_tensor() has negligible overhead when its input\n # is already viewless.\n #\n # - For the 'else' case above, calling make_viewless_tensor() here is\n # likely redundant, since p2p_communication.py (likely originator)\n # already creates viewless tensors. That said, make_viewless_tensor()\n # is called here to be future-proof and corner-case-proof.\n hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True)\n\n if self.config.sequence_parallel:\n rng_context = tensor_parallel.get_cuda_rng_tracker().fork()\n else:\n rng_context = nullcontext()\n\n # If fp8_recipe is delayed, wrap the entire pass with get_fp8_context(),\n # otherwise do nothing extra at the outer level\n # if we are using other fp8 recipes, then the context manager enter&exit are free\n # we can wrap fp8_context within the for loop over layers, so that we can fine-grained\n # control which layer will be fp8 or bf16\n use_outer_fp8_context = self.config.fp8 and self.config.fp8_recipe == Fp8Recipe.delayed\n use_inner_fp8_context = self.config.fp8 and self.config.fp8_recipe != Fp8Recipe.delayed\n outer_fp8_context = get_fp8_context(self.config) if use_outer_fp8_context else nullcontext()\n\n with rng_context, outer_fp8_context:\n # Forward pass.\n if self.config.recompute_granularity == \"full\" and self.training:\n hidden_states = self._checkpointed_forward(\n hidden_states=hidden_states,\n attention_mask=attention_mask,\n context=context,\n context_mask=context_mask,\n rotary_pos_emb=rotary_pos_emb,\n attention_bias=attention_bias,\n packed_seq_params=packed_seq_params,\n packed_seq_params_full=packed_seq_params_full,\n fullatt_block_indexes=fullatt_block_indexes,\n )\n else:\n for l_no, layer in enumerate(self.layers):\n inner_fp8_context = (\n get_fp8_context(self.config, layer.layer_number - 1) if use_inner_fp8_context else nullcontext()\n )\n if l_no in fullatt_block_indexes:\n packed_seq_params_now = packed_seq_params_full\n else:\n packed_seq_params_now = packed_seq_params\n with self.offload_context, inner_fp8_context:\n hidden_states, context = layer(\n hidden_states=hidden_states,\n attention_mask=attention_mask,\n context=context,\n context_mask=context_mask,\n rotary_pos_emb=rotary_pos_emb,\n rotary_pos_cos=rotary_pos_cos,\n rotary_pos_sin=rotary_pos_sin,\n attention_bias=attention_bias,\n inference_context=inference_context,\n packed_seq_params=packed_seq_params_now,\n sequence_len_offset=sequence_len_offset,\n )\n\n if (\n torch.is_grad_enabled()\n and self.config.cpu_offloading\n and self.group_prefetch_offload_commit_async is not None\n ):\n hidden_states = self.group_prefetch_offload_commit_async(hidden_states)\n\n # Final layer norm.\n if self.final_layernorm is not None:\n hidden_states = self.final_layernorm(hidden_states)\n # TENorm produces a \"viewed\" tensor. This will result in schedule.py's\n # deallocate_output_tensor() throwing an error, so a viewless tensor is\n # created to prevent this.\n hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True)\n\n return hidden_states\n"}30{"file_name": "verl__models__mcore__saver.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport time\n\nimport torch\nimport torch.distributed as dist\nfrom megatron.core import mpu\nfrom megatron.core.distributed import DistributedDataParallel as LocalDDP\nfrom megatron.core.transformer.module import Float16Module\nfrom torch.nn.parallel import DistributedDataParallel as torchDDP\n\nfrom verl.utils.device import get_device_id, get_torch_device\nfrom verl.utils.logger import print_rank_0\nfrom verl.utils.megatron_utils import unwrap_model\n\n\ndef _megatron_calc_global_rank(\n tp_rank: int = 0, dp_rank: int = 0, pp_rank: int = 0, cp_rank: int = 0, ep_rank: int = 0\n):\n \"\"\"Calculate global rank with support for CP/EP parallelism\"\"\"\n\n # Get parallel sizes for each dimension\n tp_size = mpu.get_tensor_model_parallel_world_size()\n dp_size = mpu.get_data_parallel_world_size()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n cp_size = mpu.get_context_parallel_world_size()\n # ep_size = mpu.get_expert_model_parallel_world_size()\n\n # Verify total GPU count matches (must be consistent with parallel_state.py)\n total_size = tp_size * dp_size * pp_size * cp_size\n assert total_size == torch.distributed.get_world_size(), (\n f\"{tp_size}x{dp_size}x{pp_size}x{cp_size} != {torch.distributed.get_world_size()}\"\n )\n\n # Core calculation logic (corresponds to RankGenerator order parameter)\n # Assumes default order is \"tp-cp-ep-dp-pp\"\n return ((pp_rank * dp_size + dp_rank) * cp_size + cp_rank) * tp_size + tp_rank\n\n\ndef _megatron_calc_layer_map(config):\n \"\"\"Calculate the mapping of global layer_idx to local layer_idx\n Returns:\n layer_map (Dict: int -> tuple(int, int, int)):\n mapping from the global layer index to\n a tuple of (pp_rank, virtual_pp_rank, layer_idx inside model)\n \"\"\"\n from megatron.core import mpu\n\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n\n layer_map = dict()\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n for pp_rank_idx in range(pp_size):\n for virtual_pp_rank_idx in range(virtual_pp_size):\n layer_offset = (\n virtual_pp_rank_idx * (config.num_hidden_layers // virtual_pp_size) + pp_rank_idx * num_layers_per_model\n )\n for layer_idx in range(num_layers_per_model):\n layer_map[layer_offset + layer_idx] = (\n pp_rank_idx,\n virtual_pp_rank_idx,\n layer_idx,\n )\n return layer_map\n\n\ndef merge_megatron_ckpt_gptmodel(wrapped_models, config, dtype, is_value_model=False, tie_word_embeddings=False):\n \"\"\"Merge sharded parameters of a Megatron module into a merged checkpoint.\n\n Args:\n wrapped_models (list of megatron.core.distributed.DistributedDataParallel):\n The local DDP wrapped megatron modules.\n config (str or None):\n HF config for model\n dtype: model params type\n is_value_model: if model is value model\n tie_word_embeddings: tie_word_embeddings\n Returns:\n state_dict (dict):\n The merged state_dict in rank 0, and an empty dictionary in other ranks.\n \"\"\"\n start_time = time.time()\n\n def _get_gpt_model(model):\n return model\n\n dp_rank = mpu.get_data_parallel_rank()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n cp_rank = mpu.get_context_parallel_rank()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n mp_group = mpu.get_model_parallel_group()\n\n if dist.get_rank() == 0:\n assert mp_group.rank() == 0, f\"mp_rank:[{mp_group.rank}] != 0 on rank #0\"\n assert pp_rank == 0, f\"pp_rank:[{pp_rank}] != 0 on rank #0\"\n assert dp_rank == 0, f\"dp_rank:[{dp_rank}] != 0 on rank #0\"\n\n if not isinstance(wrapped_models, list | tuple):\n wrapped_models = list(wrapped_models)\n\n assert len(wrapped_models) == virtual_pp_size\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n models = [None] * len(wrapped_models)\n\n for i, wrapped_model in enumerate(wrapped_models):\n models[i] = unwrap_model(wrapped_model, (torchDDP, LocalDDP, Float16Module))\n assert len(models[i].decoder.layers) == num_layers_per_model, (\n \"len model layers {} not equal to num_layers_per_model {}\".format(\n len(models[i].decoder.layers), num_layers_per_model\n )\n )\n\n state_dict = dict()\n\n def _get_cpu_tensor(tensor: torch.Tensor):\n if tensor is None:\n return None\n if tensor.device == torch.device(\"cpu\"):\n return tensor.detach().clone()\n return tensor.detach().cpu()\n\n def _broadcast_tensor(tensor, name, src_pp_rank) -> torch.Tensor:\n \"\"\"broadcast tensor across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=src_pp_rank, cp_rank=cp_rank)\n\n if torch.distributed.get_rank() == src_rank:\n if tensor is None:\n weight = None\n tensor_shape = None\n else:\n weight = tensor\n tensor_shape = weight.shape\n else:\n weight = None\n tensor_shape = None\n\n obj_list = [tensor_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n tensor_shape = obj_list[0]\n\n if tensor_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tensor:[{name}] not exist, skip collect\")\n return\n\n if weight is None:\n weight = torch.empty(\n tensor_shape,\n dtype=dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n\n dist.broadcast(weight, src=src_rank, group=mp_group)\n\n if torch.distributed.get_rank() == 0:\n state_dict[name] = _get_cpu_tensor(weight)\n\n def _broadcast_tp_shard_tensor(tensor, name, src_pp_rank, concat_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n # tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=src_pp_rank, cp_rank=cp_rank)\n\n chunk_shape = tensor.shape if torch.distributed.get_rank() == src_rank else None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{name}] not exist, skip collecting\")\n return\n\n buffer_tensor = torch.empty(\n chunk_shape,\n dtype=dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n\n chunk_tensors = [None] * tp_size\n\n for i in range(tp_size):\n cur_src_rank = _megatron_calc_global_rank(tp_rank=i, dp_rank=0, pp_rank=src_pp_rank, cp_rank=cp_rank)\n sync_tensor = tensor if torch.distributed.get_rank() == cur_src_rank else buffer_tensor\n dist.broadcast(sync_tensor, src=cur_src_rank, group=mp_group)\n\n if torch.distributed.get_rank() == 0:\n chunk_tensors[i] = _get_cpu_tensor(sync_tensor)\n\n if torch.distributed.get_rank() == 0:\n full_tensor = torch.concat(chunk_tensors, dim=concat_dim)\n if mutate_func is not None:\n full_tensor = mutate_func(full_tensor)\n state_dict[name] = full_tensor\n\n def _broadcast_tp_shard_tensor_gate_up(tensor, gate_name, up_name, src_pp_rank) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n # tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=src_pp_rank, cp_rank=cp_rank)\n\n chunk_shape = tensor.shape if torch.distributed.get_rank() == src_rank else None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{gate_name, up_name}] not exist, skip collecting\")\n return\n\n buffer_tensor = torch.empty(\n chunk_shape,\n dtype=dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n\n chunk_tensors = [None] * tp_size\n\n for i in range(tp_size):\n cur_src_rank = _megatron_calc_global_rank(tp_rank=i, dp_rank=0, pp_rank=src_pp_rank, cp_rank=cp_rank)\n sync_tensor = tensor if torch.distributed.get_rank() == cur_src_rank else buffer_tensor\n dist.broadcast(sync_tensor, src=cur_src_rank, group=mp_group)\n\n if torch.distributed.get_rank() == 0:\n chunk_tensors[i] = _get_cpu_tensor(sync_tensor)\n\n if torch.distributed.get_rank() == 0:\n full_tensor = torch.concat(chunk_tensors, dim=0)\n intermediate_size_tp = config.intermediate_size // tp_size\n gate_weight_list = []\n up_weight_list = []\n for i in range(tp_size):\n gate_up_weight_tp = full_tensor[intermediate_size_tp * 2 * i : intermediate_size_tp * 2 * (i + 1)]\n gate_weight_tp = gate_up_weight_tp[:intermediate_size_tp]\n up_weight_tp = gate_up_weight_tp[intermediate_size_tp:]\n gate_weight_list.append(gate_weight_tp)\n up_weight_list.append(up_weight_tp)\n\n state_dict[gate_name] = torch.cat(gate_weight_list, dim=0)\n state_dict[up_name] = torch.cat(up_weight_list, dim=0)\n\n def _broadcast_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name, src_pp_rank):\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n # tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=src_pp_rank, cp_rank=cp_rank)\n\n chunk_shape = tensor.shape if torch.distributed.get_rank() == src_rank else None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=src_rank, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{q_name}] not exist, skip collecting\")\n return\n\n buffer_tensor = torch.empty(\n chunk_shape,\n dtype=dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n\n chunk_tensors = [None] * tp_size\n\n for i in range(tp_size):\n cur_src_rank = _megatron_calc_global_rank(tp_rank=i, dp_rank=0, pp_rank=src_pp_rank, cp_rank=cp_rank)\n sync_tensor = tensor if torch.distributed.get_rank() == cur_src_rank else buffer_tensor\n dist.broadcast(sync_tensor, src=cur_src_rank, group=mp_group)\n\n if torch.distributed.get_rank() == 0:\n chunk_tensors[i] = _get_cpu_tensor(sync_tensor)\n\n if torch.distributed.get_rank() == 0:\n full_tensor = torch.concat(chunk_tensors, dim=0)\n q_weight_list = []\n k_weight_list = []\n v_weight_list = []\n hidden_size_per_head = getattr(config, \"head_dim\", config.hidden_size // config.num_attention_heads)\n\n if config.num_key_value_heads >= tp_size:\n q_size_tp = hidden_size_per_head * config.num_attention_heads // tp_size\n kv_size_tp = hidden_size_per_head * config.num_key_value_heads // tp_size\n total_size = q_size_tp + 2 * kv_size_tp\n for i in range(tp_size):\n num_query_groups_per_partition = wrapped_models[0].config.num_query_groups // tp_size\n qkv_part = full_tensor[i * total_size : (i + 1) * total_size]\n q_size_chunk = q_size_tp // num_query_groups_per_partition\n kv_size_chunk = kv_size_tp // num_query_groups_per_partition\n for qkv_part_chunk in qkv_part.chunk(num_query_groups_per_partition):\n q_part = qkv_part_chunk[:q_size_chunk]\n k_part = qkv_part_chunk[q_size_chunk : q_size_chunk + kv_size_chunk]\n v_part = qkv_part_chunk[q_size_chunk + kv_size_chunk :]\n q_weight_list.append(q_part)\n k_weight_list.append(k_part)\n v_weight_list.append(v_part)\n else:\n q_size_tp = hidden_size_per_head * config.num_attention_heads // tp_size\n kv_size_tp = hidden_size_per_head\n total_size = q_size_tp + 2 * kv_size_tp\n for i in range(tp_size):\n num_query_groups_per_partition = wrapped_models[0].config.num_query_groups // tp_size\n qkv_part = full_tensor[i * total_size : (i + 1) * total_size]\n q_size_chunk = q_size_tp // num_query_groups_per_partition\n kv_size_chunk = kv_size_tp // num_query_groups_per_partition\n for qkv_part_chunk in qkv_part.chunk(num_query_groups_per_partition):\n q_part = qkv_part_chunk[:q_size_chunk]\n k_part = qkv_part_chunk[q_size_chunk : q_size_chunk + kv_size_chunk]\n v_part = qkv_part_chunk[q_size_chunk + kv_size_chunk :]\n q_weight_list.append(q_part)\n if i * config.num_key_value_heads % tp_size == 0:\n k_weight_list.append(k_part)\n v_weight_list.append(v_part)\n\n state_dict[q_name] = torch.cat(q_weight_list, dim=0)\n state_dict[k_name] = torch.cat(k_weight_list, dim=0)\n state_dict[v_name] = torch.cat(v_weight_list, dim=0)\n\n # empty cache before collecting weights\n get_torch_device().empty_cache()\n # Embeddings\n # -------------------\n if dp_rank == 0 and cp_rank == 0: # models are identical across cp ranks\n # Embeddings\n # -------------------\n print_rank_0(\"collecting embeddings...\")\n gpt_model_module = _get_gpt_model(models[0])\n _broadcast_tp_shard_tensor(\n gpt_model_module.embedding.word_embeddings.weight if pp_rank == 0 else None,\n \"model.embed_tokens.weight\",\n src_pp_rank=0,\n )\n\n # Transformer layers\n # -------------------\n layer_map = _megatron_calc_layer_map(config)\n for layer in range(config.num_hidden_layers):\n print_rank_0(f\"collecting layer #{layer}...\")\n layer_name = f\"model.layers.{layer}\"\n src_pp_rank, src_virtual_pp_rank, src_layer_idx = layer_map[layer]\n\n gpt_model_module = _get_gpt_model(models[src_virtual_pp_rank])\n sync_layer = gpt_model_module.decoder.layers[src_layer_idx]\n\n _broadcast_tensor(\n sync_layer.self_attention.linear_qkv.layer_norm_weight,\n f\"{layer_name}.input_layernorm.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n if gpt_model_module.config.qk_layernorm:\n _broadcast_tensor(\n sync_layer.self_attention.q_layernorm.weight,\n f\"{layer_name}.self_attn.q_norm.weight\",\n src_pp_rank=src_pp_rank,\n )\n _broadcast_tensor(\n sync_layer.self_attention.k_layernorm.weight,\n f\"{layer_name}.self_attn.k_norm.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tp_shard_tensor_qkv(\n sync_layer.self_attention.linear_qkv.weight,\n f\"{layer_name}.self_attn.q_proj.weight\",\n f\"{layer_name}.self_attn.k_proj.weight\",\n f\"{layer_name}.self_attn.v_proj.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n if gpt_model_module.config.add_qkv_bias:\n _broadcast_tp_shard_tensor_qkv(\n sync_layer.self_attention.linear_qkv.bias,\n f\"{layer_name}.self_attn.q_proj.bias\",\n f\"{layer_name}.self_attn.k_proj.bias\",\n f\"{layer_name}.self_attn.v_proj.bias\",\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tp_shard_tensor(\n sync_layer.self_attention.linear_proj.weight,\n f\"{layer_name}.self_attn.o_proj.weight\",\n concat_dim=1,\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tensor(\n sync_layer.mlp.linear_fc1.layer_norm_weight,\n f\"{layer_name}.post_attention_layernorm.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tp_shard_tensor_gate_up(\n sync_layer.mlp.linear_fc1.weight,\n f\"{layer_name}.mlp.gate_proj.weight\",\n f\"{layer_name}.mlp.up_proj.weight\",\n src_pp_rank=src_pp_rank,\n )\n\n _broadcast_tp_shard_tensor(\n sync_layer.mlp.linear_fc2.weight,\n f\"{layer_name}.mlp.down_proj.weight\",\n concat_dim=1,\n src_pp_rank=src_pp_rank,\n )\n\n # Final Layernorm\n # -------------------\n print_rank_0(\"collecting final layernorm...\")\n gpt_model_module = _get_gpt_model(models[-1])\n _broadcast_tensor(\n getattr(gpt_model_module.decoder.final_layernorm, \"weight\", None),\n \"model.norm.weight\",\n src_pp_rank=pp_size - 1,\n )\n\n if tie_word_embeddings:\n print_rank_0(\"tie word embedding skip load lm_head...\")\n else:\n print_rank_0(\"collecting lm_head...\")\n\n if is_value_model:\n lm_head_weight = None\n if pp_rank == pp_size - 1:\n lm_head_weight = getattr(gpt_model_module.output_layer, \"weight\", None)\n _broadcast_tensor(lm_head_weight, \"lm_head.weight\", src_pp_rank=pp_size - 1)\n\n else:\n _broadcast_tp_shard_tensor(\n getattr(gpt_model_module.output_layer, \"weight\", None) if pp_rank == pp_size - 1 else None,\n \"lm_head.weight\",\n src_pp_rank=pp_size - 1,\n )\n\n dist.barrier()\n get_torch_device().empty_cache()\n if torch.distributed.get_rank() == 0:\n for k, v in state_dict.items():\n if dtype != v.dtype:\n state_dict[k] = v.to(dtype)\n\n print_rank_0(f\"merge megatron ckpt done, time elapsed {time.time() - start_time}s\")\n return state_dict\n\n\ndef merge_megatron_ckpt_gptmodel_qwen_moe(\n wrapped_models, config, dtype, is_value_model=False, tie_word_embeddings=False\n):\n raise NotImplementedError(\"merge_megatron_ckpt_gptmodel_qwen_moe is not implemented\")\n\n\ndef merge_megatron_ckpt_gptmodel_qwen2_5_vl(\n wrapped_models, config, dtype, is_value_model=False, tie_word_embeddings=False\n):\n raise NotImplementedError(\"merge_megatron_ckpt_gptmodel_qwen2_5_vl is not implemented\")\n\n\ndef merge_megatron_ckpt_gptmodel_dpskv3(wrapped_models, config, dtype, is_value_model=False, tie_word_embeddings=False):\n raise NotImplementedError(\"merge_megatron_ckpt_gptmodel_dpskv3 is not implemented\")\n\n\ndef merge_megatron_ckpt_gptmodel_mixtral(\n wrapped_models, config, dtype, is_value_model=False, tie_word_embeddings=False\n):\n raise NotImplementedError(\"merge_megatron_ckpt_gptmodel_mixtral is not implemented\")\n"}31{"file_name": "verl__models__mcore__util.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport math\n\nimport torch\nfrom megatron.core import parallel_state as mpu\nfrom megatron.core.packed_seq_params import PackedSeqParams\n\nfrom verl.utils.model import CausalLMOutputForPPO\n\n\ndef preprocess_packed_seqs(\n input_ids: torch.Tensor, attention_mask: torch.Tensor, pre_process: bool = True, use_fp8_padding=False\n) -> tuple[torch.Tensor, PackedSeqParams]:\n \"\"\"\n Preprocess packed sequences\n CP splits sequence into CP*2 chunks, and each GPU gets 2 chunks (GPU0 gets first and last chunks, GPU1\n gets second and second last chunks, and so on), this is for load balancing with causal masking.\n See https://github.com/NVIDIA/TransformerEngine/issues/1368\n \"\"\"\n batch_size = input_ids.shape[0]\n\n seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)\n tp_size = mpu.get_tensor_model_parallel_world_size()\n cp_size = mpu.get_context_parallel_world_size()\n cp_rank = mpu.get_context_parallel_rank()\n align_size = tp_size * cp_size * 2 if cp_size > 1 else tp_size\n if use_fp8_padding:\n # if fp8 is enabled, ensure the sequence is padded to multiples of 16 for better performance\n original_align_size = align_size\n align_size = math.lcm(16, align_size)\n\n pad_size = (align_size - seqlens_in_batch % align_size) % align_size\n seqlens_in_batch_padded = seqlens_in_batch + pad_size\n\n cu_seqlens = torch.zeros(batch_size + 1, dtype=torch.int32, device=input_ids.device)\n cu_seqlens[1:] = torch.cumsum(seqlens_in_batch, dim=0)\n cu_seqlens_padded = torch.zeros(batch_size + 1, dtype=torch.int32, device=input_ids.device)\n cu_seqlens_padded[1:] = torch.cumsum(seqlens_in_batch_padded, dim=0)\n\n if use_fp8_padding:\n # make sure all the sequences are padded to multiples of 128 for TE compatibility\n align_size_last = original_align_size * 128\n pad_size_last = (align_size_last - cu_seqlens_padded[-1] % align_size_last) % align_size_last\n cu_seqlens_padded[-1] += pad_size_last\n seqlens_in_batch_padded[-1] += pad_size_last\n\n # ----------------------------------------------------------------------------\n # Move the index information needed in the subsequent loop to the CPU at once,\n # to avoid frequent .item() calls in the loop that cause D2H synchronization\n # ----------------------------------------------------------------------------\n seqlens_in_batch_cpu: list[int] = seqlens_in_batch.tolist() # original valid lengths\n seqlens_in_batch_padded_cpu: list[int] = seqlens_in_batch_padded.tolist() # lengths after padding\n cu_seqlens_padded_cpu: list[int] = cu_seqlens_padded.tolist() # start positions (after padding)\n\n # Pure Python int calculation to avoid further synchronization\n max_seqlen_in_batch = max(seqlens_in_batch_padded_cpu)\n\n shape = list(input_ids.shape[1:])\n shape[0] = sum(seqlens_in_batch_padded_cpu) // cp_size\n if pre_process:\n input_ids_rmpad = torch.zeros(shape, dtype=input_ids.dtype, device=input_ids.device)\n for i in range(batch_size):\n # Use Python int, so no GPU→CPU sync in the loop\n if cp_size <= 1:\n seqlen = seqlens_in_batch_cpu[i]\n start_idx = cu_seqlens_padded_cpu[i]\n input_ids_rmpad[start_idx : start_idx + seqlen] = input_ids[i, attention_mask[i]]\n continue\n\n seqlen_padded_i = seqlens_in_batch_padded_cpu[i]\n seqlen = seqlen_padded_i // cp_size\n half_seqlen = seqlen // 2\n start_idx = cu_seqlens_padded_cpu[i] // cp_size\n # split to 2 chunks\n d = input_ids[i, attention_mask[i]]\n input_ids_rmpad[start_idx : start_idx + half_seqlen] = d[\n half_seqlen * cp_rank : half_seqlen * (cp_rank + 1)\n ]\n\n remain_start = seqlen_padded_i - half_seqlen * (cp_rank + 1)\n remain_end = seqlen_padded_i - half_seqlen * cp_rank\n remain_end = min(remain_end, d.shape[0])\n remain_len = remain_end - remain_start\n if remain_len > 0:\n input_ids_rmpad[start_idx + half_seqlen : start_idx + half_seqlen + remain_len] = d[\n remain_start:remain_end\n ]\n\n packed_seq_params = PackedSeqParams(\n qkv_format=\"thd\",\n cu_seqlens_q=cu_seqlens_padded,\n max_seqlen_q=max_seqlen_in_batch,\n cu_seqlens_kv=cu_seqlens_padded,\n max_seqlen_kv=max_seqlen_in_batch,\n cu_seqlens_q_padded=cu_seqlens_padded,\n cu_seqlens_kv_padded=cu_seqlens_padded,\n )\n if pre_process:\n return input_ids_rmpad.unsqueeze(0), packed_seq_params\n else:\n return input_ids, packed_seq_params\n\n\ndef postprocess_packed_seqs(\n output: torch.Tensor,\n packed_seq_params: PackedSeqParams,\n attention_mask: torch.Tensor,\n batch_size: int,\n seq_len: int,\n post_process: bool = True,\n) -> torch.Tensor:\n \"\"\"\n Postprocess packed sequences\n \"\"\"\n if not post_process:\n return output\n\n # -------------------------------------------------------------------------\n # Move the lengths and offsets needed for subsequent Python-level indexing to the CPU in advance,\n # to avoid a large number of .item() calls in the loop\n # -------------------------------------------------------------------------\n cu_padded_cpu: list[int] = packed_seq_params.cu_seqlens_q_padded.tolist()\n seq_lens_cpu: list[int] = attention_mask.sum(dim=1, dtype=torch.int32).cpu().tolist()\n\n shape = [batch_size, seq_len] + list(output.shape[2:]) # 1,packed, dim -> batch_size, seq_len, dim\n output_new = torch.zeros(shape, dtype=output.dtype, device=output.device)\n\n cp_size = mpu.get_context_parallel_world_size()\n # all gather output across context parallel group\n if cp_size > 1:\n # output shape: [1, packed_len, hidden_dim]\n # need to gather across cp group and concatenate in sequence dimension\n output_list = [torch.empty_like(output, dtype=output.dtype) for _ in range(cp_size)]\n torch.distributed.all_gather(output_list, output.detach(), group=mpu.get_context_parallel_group())\n output_list[mpu.get_context_parallel_rank()] = output\n else:\n output_list = [output]\n for i in range(batch_size):\n if cp_size <= 1:\n s = seq_lens_cpu[i]\n start_idx = cu_padded_cpu[i]\n output_new[i, attention_mask[i]] = output[0][start_idx : start_idx + s]\n continue\n s_len_padded_chunk = (cu_padded_cpu[i + 1] - cu_padded_cpu[i]) // cp_size\n half_seqlen = s_len_padded_chunk // 2\n s_len = seq_lens_cpu[i]\n s_len_padded = s_len_padded_chunk * cp_size\n tmp = torch.empty(s_len_padded, *output.shape[2:], device=output.device, dtype=output.dtype)\n for j in range(cp_size):\n o = output_list[j][0]\n # split to 2 chunks\n packed_start_idx = cu_padded_cpu[i] // cp_size\n o0, o1 = (\n o[packed_start_idx : packed_start_idx + half_seqlen],\n o[packed_start_idx + half_seqlen : packed_start_idx + s_len_padded_chunk],\n )\n tmp[j * half_seqlen : (j + 1) * half_seqlen] = o0\n tmp[s_len_padded - (j + 1) * half_seqlen : s_len_padded - j * half_seqlen] = o1\n output_new[i, attention_mask[i]] = tmp[:s_len]\n\n return output_new\n\n\ndef preprocess_bshd(\n input_ids: torch.Tensor,\n attention_mask: torch.Tensor,\n position_ids: torch.Tensor,\n sequence_parallel: bool = False,\n pre_process: bool = True,\n):\n \"\"\"\n Remove left padding from input_ids, attention_mask and position_ids\n return new_input_ids, new_attention_mask, new_position_ids\n \"\"\"\n assert attention_mask.ndim == 2\n assert position_ids.ndim == 2\n cp_size = mpu.get_context_parallel_world_size()\n assert cp_size == 1, \"Context parallel size without seq_pack is not supported\"\n batch_size = input_ids.shape[0]\n shape = list(input_ids.shape) # batch_size, seq_len,...\n seq_lens = attention_mask.sum(dim=1)\n seq_len = seq_lens.max().item()\n if sequence_parallel:\n sp_world_size = mpu.get_tensor_model_parallel_world_size()\n pad_size = (sp_world_size - seq_len % sp_world_size) % sp_world_size\n seq_len = seq_len + pad_size\n shape[1] = seq_len\n if pre_process:\n new_input_ids = torch.zeros(dtype=input_ids.dtype, device=input_ids.device, size=shape)\n new_attention_mask = torch.zeros(\n dtype=attention_mask.dtype, device=attention_mask.device, size=(batch_size, seq_len)\n )\n new_position_ids = torch.zeros(dtype=position_ids.dtype, device=position_ids.device, size=(batch_size, seq_len))\n for i in range(batch_size):\n if pre_process:\n new_input_ids[i, : seq_lens[i]] = input_ids[i, attention_mask[i]]\n new_attention_mask[i, : seq_lens[i]] = attention_mask[i, attention_mask[i]]\n new_position_ids[i, : seq_lens[i]] = position_ids[i, attention_mask[i]]\n if pre_process:\n return new_input_ids, new_attention_mask, new_position_ids\n else:\n return input_ids, new_attention_mask, new_position_ids\n\n\ndef postprocess_bshd(\n result,\n attention_mask: torch.Tensor,\n original_attention_mask: torch.Tensor,\n origin_seqlen: int,\n post_process: bool = True,\n):\n \"\"\"\n Recover left padding from result\n return result\n \"\"\"\n if not post_process:\n return result\n shape = list(result.shape)\n batch_size = shape[0]\n shape[1] = origin_seqlen\n new_result = torch.zeros(dtype=result.dtype, device=result.device, size=shape)\n for i in range(batch_size):\n new_result[i, original_attention_mask[i]] = result[i, attention_mask[i]]\n return new_result\n\n\ndef postprocess_packed_seqs_for_dict_output(\n labels_mask: torch.Tensor,\n output: CausalLMOutputForPPO,\n packed_seq_params: PackedSeqParams,\n attention_mask: torch.Tensor,\n batch_size: int,\n seq_len: int,\n post_process: bool = True,\n) -> dict[str, torch.Tensor]:\n \"\"\"_summary_\n For fused kernels, the output is a dictionary with keys like 'log_probs', 'entropy', etc.\n This function post-processes each tensor in the output dictionary.\n Args:\n output (CausalLMOutputForPPO): _description_\n packed_seq_params (PackedSeqParams): _description_\n attention_mask (torch.Tensor): _description_\n batch_size (int): _description_\n seq_len (int): _description_\n post_process (bool, optional): _description_. Defaults to True.\n Returns:\n CausalLMOutputForPPO: _description_\n \"\"\"\n ret = {}\n output.entropy = output.entropy.view(1, -1)\n output.log_probs = output.log_probs.view(1, -1)\n output.log_probs = output.log_probs.masked_fill(~labels_mask, 0.0)\n ret[\"entropy\"] = postprocess_packed_seqs(\n output.entropy, packed_seq_params, attention_mask, batch_size, seq_len, post_process=post_process\n )\n ret[\"log_probs\"] = postprocess_packed_seqs(\n output.log_probs, packed_seq_params, attention_mask, batch_size, seq_len, post_process=post_process\n )\n return ret\n\n\n### No padding versions for model engine\n### inputs are nested tensors\n\n\ndef preprocess_thd_no_padding(\n input_ids: torch.Tensor, pre_process: bool = True, need_roll: bool = False\n) -> tuple[torch.Tensor, PackedSeqParams]:\n \"\"\"\n Preprocess packed sequences\n CP splits sequence into CP*2 chunks, and each GPU gets 2 chunks (GPU0 gets first and last chunks, GPU1\n gets second and second last chunks, and so on), this is for load balancing with causal masking.\n See https://github.com/NVIDIA/TransformerEngine/issues/1368\n \"\"\"\n batch_size = input_ids.shape[0]\n\n tp_size = mpu.get_tensor_model_parallel_world_size()\n cp_size = mpu.get_context_parallel_world_size()\n cp_rank = mpu.get_context_parallel_rank()\n align_size = tp_size * cp_size * 2 if cp_size > 1 else tp_size\n seqlens_in_batch = input_ids.offsets().diff()\n\n pad_size = (align_size - seqlens_in_batch % align_size) % align_size\n seqlens_in_batch_padded = seqlens_in_batch + pad_size\n\n cu_seqlens = torch.zeros(batch_size + 1, dtype=torch.int32, device=input_ids.device)\n cu_seqlens[1:] = torch.cumsum(seqlens_in_batch, dim=0)\n cu_seqlens_padded = torch.zeros(batch_size + 1, dtype=torch.int32, device=input_ids.device)\n cu_seqlens_padded[1:] = torch.cumsum(seqlens_in_batch_padded, dim=0)\n\n # ----------------------------------------------------------------------------\n # Move the index information needed in the subsequent loop to the CPU at once,\n # to avoid frequent .item() calls in the loop that cause D2H synchronization\n # ----------------------------------------------------------------------------\n seqlens_in_batch_cpu: list[int] = seqlens_in_batch.tolist() # original valid lengths\n seqlens_in_batch_padded_cpu: list[int] = seqlens_in_batch_padded.tolist() # lengths after padding\n cu_seqlens_padded_cpu: list[int] = cu_seqlens_padded.tolist() # start positions (after padding)\n\n # Pure Python int calculation to avoid further synchronization\n max_seqlen_in_batch = max(seqlens_in_batch_padded_cpu)\n\n shape = list(input_ids.shape[1:])\n shape[0] = sum(seqlens_in_batch_padded_cpu) // cp_size\n if pre_process:\n input_ids_rmpad = torch.zeros(shape, dtype=input_ids.dtype, device=input_ids.device)\n if need_roll:\n saved_roll_dict = {}\n for i in range(batch_size):\n # Use Python int, so no GPU→CPU sync in the loop\n if cp_size <= 1:\n seqlen = seqlens_in_batch_cpu[i]\n start_idx = cu_seqlens_padded_cpu[i]\n input_ids_rmpad[start_idx : start_idx + seqlen] = input_ids[i]\n continue\n\n seqlen_padded_i = seqlens_in_batch_padded_cpu[i]\n seqlen = seqlen_padded_i // cp_size\n half_seqlen = seqlen // 2\n start_idx = cu_seqlens_padded_cpu[i] // cp_size\n # split to 2 chunks\n d = input_ids[i]\n input_ids_rmpad[start_idx : start_idx + half_seqlen] = d[\n half_seqlen * cp_rank : half_seqlen * (cp_rank + 1)\n ]\n\n remain_start = seqlen_padded_i - half_seqlen * (cp_rank + 1)\n remain_end = seqlen_padded_i - half_seqlen * cp_rank\n remain_end = min(remain_end, d.shape[0])\n remain_len = remain_end - remain_start\n if remain_len > 0:\n input_ids_rmpad[start_idx + half_seqlen : start_idx + half_seqlen + remain_len] = d[\n remain_start:remain_end\n ]\n\n if need_roll:\n # Handle roll for cp_size > 1 case\n saved_roll_dict[start_idx + half_seqlen - 1] = d[(cp_rank + 1) * half_seqlen]\n if remain_len > 0:\n if remain_end == d.shape[0]:\n saved_roll_dict[start_idx + half_seqlen + remain_len - 1] = d[0]\n else:\n saved_roll_dict[start_idx + half_seqlen + remain_len - 1] = d[remain_end]\n\n if need_roll:\n input_ids_rmpad = torch.roll(input_ids_rmpad, shifts=-1, dims=0)\n if len(saved_roll_dict) > 0:\n for k, v in saved_roll_dict.items():\n input_ids_rmpad[k] = v\n\n packed_seq_params = PackedSeqParams(\n qkv_format=\"thd\",\n cu_seqlens_q=cu_seqlens_padded,\n max_seqlen_q=max_seqlen_in_batch,\n cu_seqlens_kv=cu_seqlens_padded,\n max_seqlen_kv=max_seqlen_in_batch,\n cu_seqlens_q_padded=cu_seqlens_padded,\n cu_seqlens_kv_padded=cu_seqlens_padded,\n )\n if pre_process:\n return input_ids_rmpad.unsqueeze(0), packed_seq_params\n else:\n return input_ids, packed_seq_params\n\n\ndef postprocess_thd_no_padding(\n output: torch.Tensor,\n packed_seq_params: PackedSeqParams,\n input_ids: torch.Tensor,\n batch_size: int,\n post_process: bool = True,\n) -> torch.Tensor:\n \"\"\"\n Postprocess packed sequences\n \"\"\"\n if not post_process:\n return output\n\n # -------------------------------------------------------------------------\n # Move the lengths and offsets needed for subsequent Python-level indexing to the CPU in advance,\n # to avoid a large number of .item() calls in the loop\n # -------------------------------------------------------------------------\n cu_padded_cpu: list[int] = packed_seq_params.cu_seqlens_q_padded.tolist()\n # The reason why we use input_ids.offsets() instead of packed_seq_params.cu_seqlens_q.diff()\n # is that the latter one is the padded length, while the former one is the original length.\n cu_seqlens = input_ids.offsets()\n seq_lens_cpu: list[int] = cu_seqlens.diff().tolist()\n\n output_new = []\n\n cp_size = mpu.get_context_parallel_world_size()\n # all gather output across context parallel group\n if cp_size > 1:\n # output shape: [1, packed_len, hidden_dim]\n # need to gather across cp group and concatenate in sequence dimension\n output_list = [torch.empty_like(output) for _ in range(cp_size)]\n torch.distributed.all_gather(output_list, output.detach(), group=mpu.get_context_parallel_group())\n output_list[mpu.get_context_parallel_rank()] = output\n else:\n output_list = [output]\n\n for i in range(batch_size):\n if cp_size <= 1:\n s = seq_lens_cpu[i]\n start_idx = cu_padded_cpu[i]\n output_new.append(output[0][start_idx : start_idx + s])\n continue\n s_len_padded_chunk = (cu_padded_cpu[i + 1] - cu_padded_cpu[i]) // cp_size\n half_seqlen = s_len_padded_chunk // 2\n s_len = seq_lens_cpu[i]\n s_len_padded = s_len_padded_chunk * cp_size\n tmp = torch.empty(s_len_padded, *output.shape[2:], device=output.device)\n for j in range(cp_size):\n o = output_list[j][0]\n # split to 2 chunks\n packed_start_idx = cu_padded_cpu[i] // cp_size\n o0, o1 = (\n o[packed_start_idx : packed_start_idx + half_seqlen],\n o[packed_start_idx + half_seqlen : packed_start_idx + s_len_padded_chunk],\n )\n tmp[j * half_seqlen : (j + 1) * half_seqlen] = o0\n tmp[s_len_padded - (j + 1) * half_seqlen : s_len_padded - j * half_seqlen] = o1\n output_new.append(tmp[:s_len])\n\n output_new_tensor = torch.nested.as_nested_tensor(output_new, layout=torch.jagged)\n\n return output_new_tensor\n\n\ndef preprocess_bshd_no_padding(input_ids: torch.Tensor, pre_process: bool = True, need_roll: bool = False):\n \"\"\"\n Preprocess bshd sequences\n return \"input_ids, attention_mask, position_ids\"\n \"\"\"\n cp_size = mpu.get_context_parallel_world_size()\n # TODO: support context parallel size > 1\n assert cp_size == 1, \"Context parallel size without bshd is not supported yet\"\n\n batch_size = input_ids.shape[0]\n seqlens_in_batch = input_ids.offsets().diff()\n max_seqlen = seqlens_in_batch.max().item()\n if mpu.get_tensor_model_parallel_world_size() > 1:\n sp_world_size = mpu.get_tensor_model_parallel_world_size()\n pad_size = (sp_world_size - max_seqlen % sp_world_size) % sp_world_size\n max_seqlen = max_seqlen + pad_size\n\n attention_mask = torch.zeros(batch_size, max_seqlen, dtype=torch.bool, device=input_ids.device)\n input_ids_bshd = torch.zeros(batch_size, max_seqlen, dtype=input_ids.dtype, device=input_ids.device)\n for i in range(batch_size):\n attention_mask[i, : seqlens_in_batch[i]] = True\n input_ids_bshd[i, : seqlens_in_batch[i]] = input_ids[i]\n position_ids = torch.arange(max_seqlen, dtype=torch.long, device=input_ids.device)\n position_ids = position_ids.unsqueeze(0).expand_as(input_ids_bshd)\n if need_roll:\n input_ids_bshd = torch.roll(input_ids_bshd, shifts=-1, dims=1)\n\n return input_ids_bshd, attention_mask, position_ids\n\n\ndef postprocess_bshd_no_padding(\n output: torch.Tensor,\n attention_mask: torch.Tensor,\n post_process: bool = True,\n) -> torch.Tensor:\n \"\"\"\n Postprocess bshd sequences\n \"\"\"\n if not post_process:\n return output\n\n batch_size = output.shape[0]\n output_new = []\n\n for i in range(batch_size):\n mask = attention_mask[i].bool()\n output_new.append(output[i][mask])\n\n output_new_tensor = torch.nested.as_nested_tensor(output_new, layout=torch.jagged)\n\n return output_new_tensor\n"}32{"file_name": "verl__models__mcore__weight_converter.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n# online convert mcore weight to pure huggingface weight, no any fusion\n# including format conversion and name mapping\n# not including resharding\nimport torch\nfrom megatron.core.transformer import TransformerConfig\nfrom transformers import PretrainedConfig\n\n\nclass McoreToHFWeightConverterBase:\n def __init__(self, hf_config: PretrainedConfig, mcore_config: TransformerConfig):\n self.hf_config = hf_config\n self.mcore_config = mcore_config\n\n def convert_param(self, name: str, params_one_group: list[torch.Tensor]) -> torch.Tensor:\n raise NotImplementedError\n\n\nclass McoreToHFWeightConverterDense(McoreToHFWeightConverterBase):\n def _convert_attention_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n # 'decoder.layers.0.self_attention.linear_proj.weight'\n # 'decoder.layers.0.self_attention.linear_qkv.layer_norm_weight'\n # 'decoder.layers.0.self_attention.linear_qkv.weight'\n # 'decoder.layers.0.self_attention.linear_qkv.bias'\n layer_number = name.split(\".\")[2]\n convert_names = []\n if \"self_attention.linear_qkv.bias\" in name or \"self_attention.linear_qkv.weight\" in name:\n param_type = name.split(\".\")[-1]\n assert param_type == \"bias\" or param_type == \"weight\"\n convert_names.append(f\"model.layers.{layer_number}.self_attn.q_proj.{param_type}\")\n convert_names.append(f\"model.layers.{layer_number}.self_attn.k_proj.{param_type}\")\n convert_names.append(f\"model.layers.{layer_number}.self_attn.v_proj.{param_type}\")\n assert len(params) == 3\n elif \"self_attention.linear_proj.weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.self_attn.o_proj.weight\")\n assert len(params) == 1\n elif \"self_attention.linear_qkv.layer_norm_weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.input_layernorm.weight\")\n assert len(params) == 1\n elif \"self_attention.q_layernorm.weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.self_attn.q_norm.weight\")\n assert len(params) == 1\n elif \"self_attention.k_layernorm.weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.self_attn.k_norm.weight\")\n assert len(params) == 1\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n return convert_names, params\n\n def _convert_mlp_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n # 'decoder.layers.0.mlp.linear_fc1.layer_norm_weight'\n # 'decoder.layers.0.mlp.linear_fc1.weight'\n # 'decoder.layers.0.mlp.linear_fc2.weight'\n layer_number = name.split(\".\")[2]\n convert_names = []\n if \"mlp.linear_fc1.weight\" in name:\n # split gate_proj and up_proj\n convert_names.append(f\"model.layers.{layer_number}.mlp.gate_proj.weight\")\n convert_names.append(f\"model.layers.{layer_number}.mlp.up_proj.weight\")\n assert len(params) == 2\n elif \"mlp.linear_fc1.layer_norm_weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.post_attention_layernorm.weight\")\n assert len(params) == 1\n elif \"mlp.linear_fc2.weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.mlp.down_proj.weight\")\n assert len(params) == 1\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n return convert_names, params\n\n def convert_param(self, name: str, params_one_group: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n direct_name_mapping = {\n \"embedding.word_embeddings.weight\": \"model.embed_tokens.weight\",\n \"decoder.final_layernorm.weight\": \"model.norm.weight\",\n \"output_layer.weight\": \"lm_head.weight\",\n }\n if name in direct_name_mapping:\n return [direct_name_mapping[name]], [params_one_group[0]]\n\n if \"self_attention\" in name:\n return self._convert_attention_param(name, params_one_group)\n elif \"mlp\" in name:\n return self._convert_mlp_param(name, params_one_group)\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n\n\nclass McoreToHFWeightConverterQwen2Moe(McoreToHFWeightConverterDense):\n def _convert_mlp_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n # 'decoder.layers.0.pre_mlp_layernorm.weight',\n # 'decoder.layers.0.mlp.router.weight',\n # 'decoder.layers.0.mlp.shared_experts.gate_weight',\n # 'decoder.layers.0.mlp.shared_experts.linear_fc1.weight',\n # 'decoder.layers.0.mlp.shared_experts.linear_fc2.weight'\n # moe1\n # 'decoder.layers.0.mlp.experts.linear_fc1.weight0',\n # 'decoder.layers.0.mlp.experts.linear_fc1.weight1',\n # 'decoder.layers.0.mlp.experts.linear_fc1.weight2',\n # 'decoder.layers.0.mlp.experts.linear_fc1.weight3',\n # moe2\n # 'decoder.layers.0.mlp.experts.linear_fc2.weight0',\n # 'decoder.layers.0.mlp.experts.linear_fc2.weight1',\n layer_number = name.split(\".\")[2]\n convert_names = []\n if \"pre_mlp_layernorm\" in name:\n convert_names.append(f\"model.layers.{layer_number}.post_attention_layernorm.weight\")\n assert len(params) == 1\n elif \"mlp.router.weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.mlp.gate.weight\")\n assert len(params) == 1\n elif \"shared_experts.gate_weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.mlp.shared_expert_gate.weight\")\n assert len(params) == 1\n elif \"shared_experts.linear_fc1.weight\" in name: # split gate_proj and up_proj\n convert_names.append(f\"model.layers.{layer_number}.mlp.shared_expert.gate_proj.weight\")\n convert_names.append(f\"model.layers.{layer_number}.mlp.shared_expert.up_proj.weight\")\n assert len(params) == 2\n elif \"shared_experts.linear_fc2.weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.mlp.shared_expert.down_proj.weight\")\n assert len(params) == 1\n elif \"mlp.experts.linear_fc1\" in name: # split gate_proj and up_proj\n expert_id = name.split(\"weight\")[-1]\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.gate_proj.weight\")\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.up_proj.weight\")\n assert len(params) == 2\n elif \"mlp.experts.linear_fc2\" in name:\n expert_id = name.split(\"weight\")[-1]\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.down_proj.weight\")\n assert len(params) == 1\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n return convert_names, params\n\n\nclass McoreToHFWeightConverterQwen2_5_VL(McoreToHFWeightConverterDense):\n def convert_param(self, name: str, params_one_group: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n direct_name_mapping = {\n \"language_model.embedding.word_embeddings.weight\": \"model.embed_tokens.weight\",\n \"language_model.decoder.final_layernorm.weight\": \"model.norm.weight\",\n \"language_model.output_layer.weight\": \"lm_head.weight\",\n \"vision_model.patch_embed.proj.weight\": \"visual.patch_embed.proj.weight\",\n \"vision_model.decoder.final_layernorm.weight\": \"visual.merger.ln_q.weight\",\n \"vision_model.projection.encoder.linear_fc1.weight\": \"visual.merger.mlp.0.weight\",\n \"vision_model.projection.encoder.linear_fc1.bias\": \"visual.merger.mlp.0.bias\",\n \"vision_model.projection.encoder.linear_fc2.weight\": \"visual.merger.mlp.2.weight\",\n \"vision_model.projection.encoder.linear_fc2.bias\": \"visual.merger.mlp.2.bias\",\n }\n if name in direct_name_mapping:\n return [direct_name_mapping[name]], [params_one_group[0]]\n\n if \"self_attention\" in name:\n return self._convert_attention_param(name, params_one_group)\n elif \"mlp\" in name:\n return self._convert_mlp_param(name, params_one_group)\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n\n def _convert_attention_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n model_type, _, _, layer_number = name.split(\".\")[:4]\n\n convert_names = []\n if model_type == \"language_model\":\n name_map_after_layer = {\n \"self_attention.linear_qkv.bias\": [\n \"self_attn.q_proj.bias\",\n \"self_attn.k_proj.bias\",\n \"self_attn.v_proj.bias\",\n ],\n \"self_attention.linear_qkv.weight\": [\n \"self_attn.q_proj.weight\",\n \"self_attn.k_proj.weight\",\n \"self_attn.v_proj.weight\",\n ],\n \"self_attention.linear_proj.weight\": \"self_attn.o_proj.weight\",\n \"self_attention.linear_qkv.layer_norm_weight\": \"input_layernorm.weight\",\n }\n name_after_layer = \".\".join(name.split(\".\")[-3:])\n mapped_name = name_map_after_layer.get(name_after_layer)\n if isinstance(mapped_name, list):\n assert len(params) == len(mapped_name)\n for one in mapped_name:\n convert_names.append(f\"model.layers.{layer_number}.{one}\")\n else:\n assert len(params) == 1\n convert_names.append(f\"model.layers.{layer_number}.{mapped_name}\")\n elif model_type == \"vision_model\":\n name_map_after_layer = {\n \"self_attention.linear_proj.weight\": \"attn.proj.weight\",\n \"self_attention.linear_proj.bias\": \"attn.proj.bias\",\n \"self_attention.linear_qkv.layer_norm_weight\": \"norm1.weight\",\n }\n name_after_layer = \".\".join(name.split(\".\")[-3:])\n mapped_name = name_map_after_layer.get(name_after_layer, None)\n if mapped_name is None:\n assert \"linear_qkv\" in name_after_layer\n assert len(params) == 3\n new_param = torch.cat(params, dim=0)\n params = [new_param]\n if \"bias\" in name_after_layer:\n convert_names.append(f\"visual.blocks.{layer_number}.attn.qkv.bias\")\n else:\n convert_names.append(f\"visual.blocks.{layer_number}.attn.qkv.weight\")\n else:\n assert len(params) == 1\n convert_names.append(f\"visual.blocks.{layer_number}.{mapped_name}\")\n else:\n raise NotImplementedError(f\"Unsupported model type: {model_type}\")\n return convert_names, params\n\n def _convert_mlp_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n model_type, _, _, layer_number = name.split(\".\")[:4]\n\n convert_names = []\n if model_type == \"language_model\":\n name_map_after_layer = {\n \"mlp.linear_fc1.weight\": [\"mlp.gate_proj.weight\", \"mlp.up_proj.weight\"],\n \"mlp.linear_fc1.bias\": [\"mlp.gate_proj.bias\", \"mlp.up_proj.bias\"],\n \"mlp.linear_fc2.weight\": \"mlp.down_proj.weight\",\n \"mlp.linear_fc2.bias\": \"mlp.down_proj.bias\",\n \"mlp.linear_fc1.layer_norm_weight\": \"post_attention_layernorm.weight\",\n }\n name_after_layer = \".\".join(name.split(\".\")[-3:])\n mapped_name = name_map_after_layer.get(name_after_layer)\n if isinstance(mapped_name, list):\n assert len(params) == len(mapped_name)\n for one in mapped_name:\n convert_names.append(f\"model.layers.{layer_number}.{one}\")\n else:\n assert len(params) == 1\n convert_names.append(f\"model.layers.{layer_number}.{mapped_name}\")\n\n elif model_type == \"vision_model\":\n name_map_after_layer = {\n \"mlp.linear_fc1.weight\": [\"mlp.gate_proj.weight\", \"mlp.up_proj.weight\"],\n \"mlp.linear_fc1.bias\": [\"mlp.gate_proj.bias\", \"mlp.up_proj.bias\"],\n \"mlp.linear_fc2.weight\": \"mlp.down_proj.weight\",\n \"mlp.linear_fc2.bias\": \"mlp.down_proj.bias\",\n \"mlp.linear_fc1.layer_norm_weight\": \"norm2.weight\",\n }\n name_after_layer = \".\".join(name.split(\".\")[-3:])\n mapped_name = name_map_after_layer.get(name_after_layer)\n if isinstance(mapped_name, list):\n assert len(params) == len(mapped_name)\n for one in mapped_name:\n convert_names.append(f\"visual.blocks.{layer_number}.{one}\")\n else:\n assert len(params) == 1\n convert_names.append(f\"visual.blocks.{layer_number}.{mapped_name}\")\n else:\n raise NotImplementedError(f\"Unsupported model type: {model_type}\")\n return convert_names, params\n\n\nclass McoreToHFWeightConverterDpskv3(McoreToHFWeightConverterBase):\n def _convert_attention_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n # mcore\n # 'decoder.layers.0.input_layernorm.weight'\n # 'decoder.layers.0.self_attention.linear_proj.weight'\n # 'decoder.layers.0.self_attention.linear_q_proj.weight'\n # 'decoder.layers.0.self_attention.linear_kv_down_proj.weight'\n # 'decoder.layers.0.self_attention.linear_kv_up_proj.layer_norm_weight'\n # 'decoder.layers.0.self_attention.linear_kv_up_proj.weight'\n # 'decoder.layers.0.self_attention.linear_q_down_proj.weight'\n # 'decoder.layers.0.self_attention.linear_q_up_proj.weight'\n # 'decoder.layers.0.self_attention.linear_q_up_proj.layer_norm_weight'\n # hf\n # 'model.layers.0.input_layernorm.weight'\n # 'model.layers.0.self_attn.o_proj.weight'\n # 'model.layers.0.self_attn.q_proj.weight'\n # 'model.layers.0.self_attn.kv_a_proj_with_mqa.weight'\n # 'model.layers.0.self_attn.kv_a_layernorm.weight'\n # 'model.layers.0.self_attn.kv_b_proj.weight'\n # 'model.layers.0.self_attn.q_a_proj.weight'\n # 'model.layers.0.self_attn.q_b_proj.weight'\n # 'model.layers.0.self_attn.q_a_layernorm.weight'\n name_map_after_layer = {\n \"input_layernorm.weight\": \"input_layernorm.weight\",\n \"self_attention.linear_proj.weight\": \"self_attn.o_proj.weight\",\n \"self_attention.linear_q_proj.weight\": \"self_attn.q_proj.weight\",\n \"self_attention.linear_kv_down_proj.weight\": \"self_attn.kv_a_proj_with_mqa.weight\",\n \"self_attention.linear_kv_up_proj.layer_norm_weight\": \"self_attn.kv_a_layernorm.weight\",\n \"self_attention.linear_kv_up_proj.weight\": \"self_attn.kv_b_proj.weight\",\n \"self_attention.linear_q_down_proj.weight\": \"self_attn.q_a_proj.weight\",\n \"self_attention.linear_q_up_proj.weight\": \"self_attn.q_b_proj.weight\",\n \"self_attention.linear_q_up_proj.layer_norm_weight\": \"self_attn.q_a_layernorm.weight\",\n }\n assert len(params) == 1\n convert_names = []\n layer_number = name.split(\".\")[2]\n name_after_layer = name.split(f\".{layer_number}.\")[1]\n convert_names.append(f\"model.layers.{layer_number}.{name_map_after_layer[name_after_layer]}\")\n return convert_names, params\n\n def _convert_mlp_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n # mcore dense\n # 'decoder.layers.0.mlp.linear_fc1.layer_norm_weight'\n # 'decoder.layers.0.mlp.linear_fc2.weight'\n # 'decoder.layers.0.mlp.linear_fc1.weight'\n # ---\n # 'decoder.layers.1.mlp.shared_experts.linear_fc1.weight'\n # ---\n # 'decoder.layers.1.mlp.shared_experts.linear_fc2.weight'\n # hf dense\n # 'model.layers.0.post_attention_layernorm.weight'\n # 'model.layers.0.mlp.down_proj.weight'\n # 'model.layers.0.mlp.gate_proj.weight'\n # 'model.layers.0.mlp.up_proj.weight'\n # 'model.layers.1.mlp.shared_experts.gate_proj.weight'\n # 'model.layers.1.mlp.shared_experts.up_proj.weight'\n # 'model.layers.1.mlp.shared_experts.down_proj.weight'\n\n # mcore moe\n # 'decoder.layers.1.pre_mlp_layernorm.weight'\n # 'decoder.layers.1.mlp.router.weight'\n # 'decoder.layers.1.mlp.router.expert_bias'\n # 'decoder.layers.1.mlp.experts.linear_fc1.weight0'\n # ---\n # 'decoder.layers.1.mlp.experts.linear_fc2.weight0'\n # hf moe\n # 'model.layers.1.post_attention_layernorm.weight'\n # 'model.layers.1.mlp.gate.weight'\n # 'model.layers.1.mlp.gate.e_score_correction_bias'\n # 'model.layers.1.mlp.experts.0.gate_proj.weight'\n # 'model.layers.1.mlp.experts.0.up_proj.weight'\n # 'model.layers.1.mlp.experts.0.down_proj.weight'\n\n name_map_after_layer = {\n \"mlp.linear_fc1.layer_norm_weight\": \"post_attention_layernorm.weight\",\n \"mlp.linear_fc2.weight\": \"mlp.down_proj.weight\",\n \"mlp.shared_experts.linear_fc2.weight\": \"mlp.shared_experts.down_proj.weight\",\n \"mlp.linear_fc1.weight\": [\"mlp.gate_proj.weight\", \"mlp.up_proj.weight\"],\n \"mlp.shared_experts.linear_fc1.weight\": [\n \"mlp.shared_experts.gate_proj.weight\",\n \"mlp.shared_experts.up_proj.weight\",\n ],\n \"pre_mlp_layernorm.weight\": \"post_attention_layernorm.weight\",\n \"mlp.router.weight\": \"mlp.gate.weight\",\n \"mlp.router.expert_bias\": \"mlp.gate.e_score_correction_bias\",\n }\n convert_names = []\n layer_number = name.split(\".\")[2]\n name_after_layer = name.split(f\".{layer_number}.\")[1]\n if name_after_layer in name_map_after_layer:\n mapped_name = name_map_after_layer[name_after_layer]\n if isinstance(mapped_name, list):\n assert len(params) == len(mapped_name)\n for one in mapped_name:\n convert_names.append(f\"model.layers.{layer_number}.{one}\")\n else:\n assert len(params) == 1\n convert_names.append(f\"model.layers.{layer_number}.{mapped_name}\")\n else:\n if \"mlp.experts.linear_fc1.weight\" in name:\n expert_id = name.split(\"weight\")[-1]\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.gate_proj.weight\")\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.up_proj.weight\")\n assert len(params) == 2\n elif \"mlp.experts.linear_fc2.weight\" in name:\n expert_id = name.split(\"weight\")[-1]\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.down_proj.weight\")\n assert len(params) == 1\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n\n return convert_names, params\n\n def _convert_mtp_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n assert self.mcore_config.mtp_num_layers == 1, \"only support one mtp layer for now\"\n assert self.mcore_config.num_layers == 61, \"only support 61 layers for now\"\n direct_name_mapping = {\n \"mtp.layers.0.enorm.weight\": \"model.layers.61.enorm.weight\",\n \"mtp.layers.0.hnorm.weight\": \"model.layers.61.hnorm.weight\",\n \"mtp.layers.0.eh_proj.weight\": \"model.layers.61.eh_proj.weight\",\n \"mtp.layers.0.final_layernorm.weight\": \"model.layers.61.shared_head.norm.weight\",\n }\n if name in direct_name_mapping:\n return [direct_name_mapping[name]], [params[0]]\n assert \"mtp.layers.0.transformer_layer\" in name, \"only support transformer layer for now\"\n # use proxy name to convert\n proxy_name = name.replace(\"mtp.layers.0.transformer_layer\", \"decoder.layers.61\")\n if \"self_attention\" in proxy_name or \"input_layernorm.weight\" in proxy_name:\n convert_names, params = self._convert_attention_param(proxy_name, params)\n elif \"mlp\" in proxy_name:\n convert_names, params = self._convert_mlp_param(proxy_name, params)\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n return convert_names, params\n\n def convert_param(self, name: str, params_one_group: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n direct_name_mapping = {\n \"embedding.word_embeddings.weight\": \"model.embed_tokens.weight\",\n \"decoder.final_layernorm.weight\": \"model.norm.weight\",\n \"output_layer.weight\": \"lm_head.weight\",\n }\n if name in direct_name_mapping:\n return [direct_name_mapping[name]], [params_one_group[0]]\n if \"mtp\" in name:\n return self._convert_mtp_param(name, params_one_group)\n elif \"self_attention\" in name or \"input_layernorm.weight\" in name:\n return self._convert_attention_param(name, params_one_group)\n elif \"mlp\" in name:\n return self._convert_mlp_param(name, params_one_group)\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n\n\nclass McoreToHFWeightConverterMixtral(McoreToHFWeightConverterDense):\n def _convert_mlp_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n # decoder.layers.0.mlp.router.weight\n # decoder.layers.0.mlp.experts.linear_fc1.weight0 - weight7\n # decoder.layers.0.mlp.experts.linear_fc2.weight0 - weight7\n\n layer_number = name.split(\".\")[2]\n convert_names = []\n if \"pre_mlp_layernorm\" in name:\n convert_names.append(f\"model.layers.{layer_number}.post_attention_layernorm.weight\")\n elif \"mlp.router.weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.block_sparse_moe.gate.weight\")\n elif \"mlp.experts.linear_fc1.weight\" in name:\n expert_id = name.split(\"weight\")[-1]\n convert_names.append(f\"model.layers.{layer_number}.block_sparse_moe.experts.{expert_id}.w1.weight\")\n convert_names.append(f\"model.layers.{layer_number}.block_sparse_moe.experts.{expert_id}.w3.weight\")\n elif \"mlp.experts.linear_fc2.weight\" in name:\n expert_id = name.split(\"weight\")[-1]\n convert_names.append(f\"model.layers.{layer_number}.block_sparse_moe.experts.{expert_id}.w2.weight\")\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n return convert_names, params\n\n\nclass McoreToHFWeightConverterQwen3Moe(McoreToHFWeightConverterDense):\n def _convert_mlp_param(self, name: str, params: list[torch.Tensor]) -> tuple[list[str], list[torch.Tensor]]:\n # qwen3 moe no share expert\n\n # 'decoder.layers.0.pre_mlp_layernorm.weight',\n # 'decoder.layers.0.mlp.router.weight',\n # moe1\n # 'decoder.layers.0.mlp.experts.linear_fc1.weight0',\n # 'decoder.layers.0.mlp.experts.linear_fc1.weight1',\n # 'decoder.layers.0.mlp.experts.linear_fc1.weight2',\n # 'decoder.layers.0.mlp.experts.linear_fc1.weight3',\n # moe2\n # 'decoder.layers.0.mlp.experts.linear_fc2.weight0',\n # 'decoder.layers.0.mlp.experts.linear_fc2.weight1',\n layer_number = name.split(\".\")[2]\n convert_names = []\n if \"pre_mlp_layernorm\" in name:\n convert_names.append(f\"model.layers.{layer_number}.post_attention_layernorm.weight\")\n assert len(params) == 1\n elif \"mlp.router.weight\" in name:\n convert_names.append(f\"model.layers.{layer_number}.mlp.gate.weight\")\n assert len(params) == 1\n elif \"mlp.experts.linear_fc1\" in name: # split gate_proj and up_proj\n expert_id = name.split(\"weight\")[-1]\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.gate_proj.weight\")\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.up_proj.weight\")\n assert len(params) == 2\n elif \"mlp.experts.linear_fc2\" in name:\n expert_id = name.split(\"weight\")[-1]\n convert_names.append(f\"model.layers.{layer_number}.mlp.experts.{expert_id}.down_proj.weight\")\n assert len(params) == 1\n else:\n raise NotImplementedError(f\"Unsupported parameter name: {name}\")\n return convert_names, params\n"}33{"file_name": "verl__models__qwen2__megatron__checkpoint_utils__qwen2_loader.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport time\n\nimport torch\nimport torch.distributed as dist\n\nfrom verl.utils.device import get_device_id, get_torch_device\n\n\ndef _megatron_calc_layer_map(config):\n \"\"\"Calculate the mapping of global layer_idx to local layer_idx\n Returns:\n layer_map (Dict: int -> tuple(int, int, int)):\n mapping from the global layer index to\n a tuple of (pp_rank, virtual_pp_rank, layer_idx inside model)\n \"\"\"\n from megatron.core import mpu\n\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n\n layer_map = dict()\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n for pp_rank_idx in range(pp_size):\n for virtual_pp_rank_idx in range(virtual_pp_size):\n layer_offset = (\n virtual_pp_rank_idx * (config.num_hidden_layers // virtual_pp_size) + pp_rank_idx * num_layers_per_model\n )\n for layer_idx in range(num_layers_per_model):\n layer_map[layer_offset + layer_idx] = (\n pp_rank_idx,\n virtual_pp_rank_idx,\n layer_idx,\n )\n return layer_map\n\n\ndef load_state_dict_to_megatron_qwen2(\n state_dict, wrapped_models, config, params_dtype, is_value_model=False, tie_word_embeddings=False\n):\n \"\"\"Load merged state_dict to sharded Megatron module in training.\"\"\"\n from megatron.core import DistributedDataParallel as LocalDDP\n from megatron.core import mpu\n from megatron.core.transformer.module import Float16Module\n from torch.nn.parallel import DistributedDataParallel as torchDDP\n\n from verl.utils.logger import print_rank_0\n from verl.utils.megatron_utils import unwrap_model\n\n start_time = time.time()\n\n def _get_gpt_model(model):\n return model\n\n def fetch_params(module):\n for param in module.parameters():\n torch.distributed.fetch(\n param.data, src=mpu.get_data_parallel_src_rank(), group=mpu.get_data_parallel_group()\n )\n\n dp_rank = mpu.get_data_parallel_rank()\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n mp_group = mpu.get_model_parallel_group()\n\n if torch.distributed.get_rank() == 0:\n assert mp_group.rank() == 0, f\"mp_rank:[{mp_group.rank}] != 0 on rank #0\"\n assert pp_rank == 0, f\"pp_rank:[{pp_rank}] != 0 on rank #0\"\n assert dp_rank == 0, f\"dp_rank:[{dp_rank}] != 0 on rank #0\"\n\n if not isinstance(wrapped_models, list | tuple):\n wrapped_models = list(wrapped_models)\n\n assert len(wrapped_models) == virtual_pp_size\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers, (\n f\"num_layers_per_model: {num_layers_per_model} * pp_size: {pp_size} * virtual_pp_size: \"\n f\"{virtual_pp_size} != config.num_hidden_layers: {config.num_hidden_layers}\"\n )\n\n models = [None] * len(wrapped_models)\n\n for i, wrapped_model in enumerate(wrapped_models):\n models[i] = unwrap_model(wrapped_model, (torchDDP, LocalDDP, Float16Module))\n gpt_model_module = _get_gpt_model(models[i])\n assert len(gpt_model_module.model.layers) == num_layers_per_model\n\n def _fetch_tensor(tensor, name) -> torch.Tensor:\n \"\"\"fetch tensor\"\"\"\n nonlocal state_dict\n if tensor is not None:\n tensor = tensor.data.copy_(state_dict[name], non_blocking=True)\n\n def _fetch_tp_shard_tensor_vocab(tensor, name, chunk_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"fetch tensor in tp shards\"\"\"\n nonlocal state_dict\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n if name in state_dict:\n full_weight = state_dict[name]\n\n if mutate_func is not None:\n full_weight = mutate_func(full_weight)\n tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)\n if tensor is not None:\n tensor = tensor.data.copy_(tensor_chunk[tp_rank], non_blocking=True)\n else:\n print(f\"tp_shard tensor:[{name}] not in state_dict, skip loading\")\n\n def _fetch_tp_shard_tensor(tensor, name, chunk_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"fetch tensor in tp shards\"\"\"\n nonlocal state_dict\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n if name in state_dict:\n full_weight = state_dict[name]\n\n if mutate_func is not None:\n full_weight = mutate_func(full_weight)\n tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)\n if tensor is not None:\n tensor = tensor.data.copy_(tensor_chunk[tp_rank], non_blocking=True)\n else:\n print(f\"tp_shard tensor:[{name}] not in state_dict, skip loading\")\n\n def _fetch_tp_shard_tensor_gate_up(tensor, gate_name, up_name) -> torch.Tensor:\n \"\"\"fetch gate_up tensor in tp shards\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n if gate_name in state_dict and up_name in state_dict:\n gate_weight = state_dict[gate_name]\n up_weight = state_dict[up_name]\n new_gate_up_weight = torch.empty(\n config.intermediate_size * 2, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n for i in range(tp_size):\n intermediate_size_tp = config.intermediate_size // tp_size\n gate_weight_tp = gate_weight[i * intermediate_size_tp : (i + 1) * intermediate_size_tp]\n up_weight_tp = up_weight[i * intermediate_size_tp : (i + 1) * intermediate_size_tp]\n new_gate_up_weight[intermediate_size_tp * 2 * i : intermediate_size_tp * 2 * (i + 1)].copy_(\n torch.cat([gate_weight_tp, up_weight_tp], dim=0)\n )\n\n tensor_chunk = torch.chunk(new_gate_up_weight, tp_size, dim=0)\n if tensor is not None:\n tensor = tensor.data.copy_(tensor_chunk[tp_rank], non_blocking=True)\n else:\n print(f\"tp_shard tensor:[{gate_name}, {up_name}] not in state_dict, skip loading\")\n\n def _fetch_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name, bias=False) -> torch.Tensor:\n \"\"\"fetch tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n assert q_name in state_dict and k_name in state_dict and v_name in state_dict\n full_weight_q = state_dict[q_name]\n full_weight_k = state_dict[k_name]\n full_weight_v = state_dict[v_name]\n\n hidden_size_per_head = config.hidden_size // config.num_attention_heads\n\n if config.num_key_value_heads >= tp_size:\n q_size_tp = config.hidden_size // tp_size\n kv_size_tp = hidden_size_per_head * config.num_key_value_heads // tp_size\n total_size = q_size_tp + 2 * kv_size_tp\n if not bias:\n new_weight_qkv = torch.empty(\n total_size * tp_size, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n else:\n new_weight_qkv = torch.empty(total_size * tp_size, dtype=params_dtype, device=get_device_id())\n for i in range(tp_size):\n q_part = full_weight_q[i * q_size_tp : (i + 1) * q_size_tp]\n k_part = full_weight_k[i * kv_size_tp : (i + 1) * kv_size_tp]\n v_part = full_weight_v[i * kv_size_tp : (i + 1) * kv_size_tp]\n new_weight_qkv[i * total_size : (i + 1) * total_size].copy_(torch.cat([q_part, k_part, v_part], dim=0))\n\n else:\n q_size_tp = config.hidden_size // tp_size\n kv_size_tp = hidden_size_per_head\n total_size = q_size_tp + 2 * kv_size_tp\n if not bias:\n new_weight_qkv = torch.empty(\n total_size * tp_size, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n else:\n new_weight_qkv = torch.empty(total_size * tp_size, dtype=params_dtype, device=get_device_id())\n for i in range(tp_size):\n q_part = full_weight_q[i * q_size_tp : (i + 1) * q_size_tp]\n start_idx = i * config.num_key_value_heads // tp_size * hidden_size_per_head\n end_idx = (i * config.num_key_value_heads // tp_size + 1) * hidden_size_per_head\n k_part = full_weight_k[start_idx:end_idx]\n v_part = full_weight_v[start_idx:end_idx]\n new_weight_qkv[i * total_size : (i + 1) * total_size].copy_(torch.cat([q_part, k_part, v_part], dim=0))\n\n tensor_chunk = torch.chunk(new_weight_qkv, tp_size, dim=0)\n if tensor is not None:\n tensor = tensor.data.copy_(tensor_chunk[tp_rank], non_blocking=True)\n\n # Embeddings\n # -------------------\n print_rank_0(\"loading embeddings...\")\n gpt_model_module = _get_gpt_model(models[0])\n if pp_rank == 0:\n embed_tokens_weight = gpt_model_module.model.embed_tokens.weight\n _fetch_tp_shard_tensor_vocab(embed_tokens_weight, \"model.embed_tokens.weight\")\n\n # Transformer layers\n # -------------------\n layer_map = _megatron_calc_layer_map(config)\n\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n num_layer_per_pp = config.num_hidden_layers // pp_size\n vpp_size = mpu.get_virtual_pipeline_model_parallel_world_size()\n\n layer_list = []\n if vpp_size is not None:\n for vpp_rank in range(vpp_size):\n num_layer_vpp_chunk = num_layer_per_pp // vpp_size\n num_layer_this_model = num_layer_vpp_chunk\n offset = vpp_rank * (config.num_hidden_layers // mpu.get_virtual_pipeline_model_parallel_world_size()) + (\n mpu.get_pipeline_model_parallel_rank() * num_layer_vpp_chunk\n )\n layer_list.extend(list(range(offset, offset + num_layer_this_model)))\n else:\n num_layer_this_model = num_layer_per_pp\n offset = pp_rank * num_layer_per_pp\n layer_list.extend(list(range(offset, offset + num_layer_this_model)))\n\n for layer in layer_list:\n print(f\"{torch.distributed.get_rank()} loading layer #{layer}...\")\n layer_name = f\"model.layers.{layer}\"\n dst_pp_rank, dst_virtual_pp_rank, dst_layer_idx = layer_map[layer]\n\n print(\n f\"{torch.distributed.get_rank()} offset: {offset}, num_layer_this_model: {num_layer_this_model}, \"\n f\"layer_name: {layer_name}, layer_map[layer]: {layer_map[layer]}\"\n )\n\n gpt_model_module = _get_gpt_model(models[dst_virtual_pp_rank])\n sync_layer = gpt_model_module.model.layers[dst_layer_idx]\n\n _fetch_tensor(\n sync_layer.input_layernorm.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.input_layernorm.weight\",\n )\n\n _fetch_tp_shard_tensor_qkv(\n sync_layer.self_attn.qkv_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.q_proj.weight\",\n f\"{layer_name}.self_attn.k_proj.weight\",\n f\"{layer_name}.self_attn.v_proj.weight\",\n )\n\n _fetch_tp_shard_tensor_qkv(\n sync_layer.self_attn.qkv_proj.bias if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.q_proj.bias\",\n f\"{layer_name}.self_attn.k_proj.bias\",\n f\"{layer_name}.self_attn.v_proj.bias\",\n bias=True,\n )\n\n _fetch_tp_shard_tensor(\n sync_layer.self_attn.o_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.o_proj.weight\",\n chunk_dim=1,\n )\n\n _fetch_tensor(\n sync_layer.post_attention_layernorm.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.post_attention_layernorm.weight\",\n )\n\n _fetch_tp_shard_tensor_gate_up(\n sync_layer.mlp.gate_up_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.mlp.gate_proj.weight\",\n f\"{layer_name}.mlp.up_proj.weight\",\n )\n\n _fetch_tp_shard_tensor(\n sync_layer.mlp.down_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.mlp.down_proj.weight\",\n chunk_dim=1,\n )\n # Final Layernorm\n # -------------------\n print_rank_0(\"loading final layernorm...\")\n gpt_model_module = _get_gpt_model(models[-1])\n _fetch_tensor(\n getattr(gpt_model_module.model.norm, \"weight\", None),\n \"model.norm.weight\",\n )\n\n if tie_word_embeddings:\n print_rank_0(\"tie_word_embeddings skip load lm_head\")\n else:\n print_rank_0(\"loading lm_head...\")\n if pp_rank + 1 == pp_size:\n lm_head_weight = gpt_model_module.lm_head.weight\n\n if is_value_model:\n if \"lm_head.weight\" in state_dict and state_dict[\"lm_head.weight\"].shape[0] == 1:\n _fetch_tensor(lm_head_weight, \"lm_head.weight\")\n print_rank_0(\"load lm_head from value_head weight\")\n elif \"reward_head.weight\" in state_dict and state_dict[\"reward_head.weight\"].shape[0] == 1:\n _fetch_tensor(lm_head_weight, \"reward_head.weight\")\n print_rank_0(\"load lm_head from value_head weight\")\n else:\n _fetch_tensor(None, \"lm_head.weight\")\n print_rank_0(\"fail to match lm_head in value_model\")\n\n else:\n _fetch_tp_shard_tensor(lm_head_weight, \"lm_head.weight\")\n\n dist.barrier()\n get_torch_device().empty_cache()\n print_rank_0(f\"loading megatron ckpt done, time elapsed {time.time() - start_time}s\")\n"}34{"file_name": "verl__models__qwen2__megatron__checkpoint_utils__qwen2_loader_depracated.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport time\n\nimport torch\nimport torch.distributed as dist\n\nfrom verl.utils.device import get_device_id, get_torch_device\n\n\ndef _megatron_calc_layer_map(config):\n \"\"\"Calculate the mapping of global layer_idx to local layer_idx\n Returns:\n layer_map (Dict: int -> tuple(int, int, int)):\n mapping from the global layer index to\n a tuple of (pp_rank, virtual_pp_rank, layer_idx inside model)\n \"\"\"\n from megatron.core import mpu\n\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n\n layer_map = dict()\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers\n\n for pp_rank_idx in range(pp_size):\n for virtual_pp_rank_idx in range(virtual_pp_size):\n layer_offset = (\n virtual_pp_rank_idx * (config.num_hidden_layers // virtual_pp_size) + pp_rank_idx * num_layers_per_model\n )\n for layer_idx in range(num_layers_per_model):\n layer_map[layer_offset + layer_idx] = (\n pp_rank_idx,\n virtual_pp_rank_idx,\n layer_idx,\n )\n return layer_map\n\n\ndef load_state_dict_to_megatron_qwen2(\n state_dict, wrapped_models, config, params_dtype, is_value_model=False, tie_word_embeddings=False\n):\n \"\"\"Load merged state_dict to sharded Megatron module in training.\"\"\"\n from megatron.core import DistributedDataParallel as LocalDDP\n from megatron.core import mpu\n from megatron.core.transformer.module import Float16Module\n from torch.nn.parallel import DistributedDataParallel as torchDDP\n\n from verl.utils.logger import print_rank_0\n from verl.utils.megatron_utils import unwrap_model\n\n start_time = time.time()\n\n def _get_gpt_model(model):\n return model\n\n def broadcast_params(module):\n for param in module.parameters():\n torch.distributed.broadcast(\n param.data, src=mpu.get_data_parallel_src_rank(), group=mpu.get_data_parallel_group()\n )\n\n dp_rank = mpu.get_data_parallel_rank()\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n virtual_pp_size = mpu.get_virtual_pipeline_model_parallel_world_size() or 1\n mp_group = mpu.get_model_parallel_group()\n\n if torch.distributed.get_rank() == 0:\n assert mp_group.rank() == 0, f\"mp_rank:[{mp_group.rank}] != 0 on rank #0\"\n assert pp_rank == 0, f\"pp_rank:[{pp_rank}] != 0 on rank #0\"\n assert dp_rank == 0, f\"dp_rank:[{dp_rank}] != 0 on rank #0\"\n\n if not isinstance(wrapped_models, list | tuple):\n wrapped_models = list(wrapped_models)\n\n assert len(wrapped_models) == virtual_pp_size\n num_layers_per_model = config.num_hidden_layers // pp_size // virtual_pp_size\n assert num_layers_per_model * pp_size * virtual_pp_size == config.num_hidden_layers, (\n f\"num_layers_per_model: {num_layers_per_model} * pp_size: {pp_size} * virtual_pp_size: \"\n f\"{virtual_pp_size} != config.num_hidden_layers: {config.num_hidden_layers}\"\n )\n\n models = [None] * len(wrapped_models)\n\n for i, wrapped_model in enumerate(wrapped_models):\n models[i] = unwrap_model(wrapped_model, (torchDDP, LocalDDP, Float16Module))\n gpt_model_module = _get_gpt_model(models[i])\n assert len(gpt_model_module.model.layers) == num_layers_per_model\n\n def _broadcast_tensor(tensor, name) -> torch.Tensor:\n \"\"\"broadcast tensor from rank0 across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n if torch.distributed.get_rank() == 0:\n if name in state_dict:\n weight = state_dict[name]\n tensor_shape = weight.shape\n else:\n tensor_shape = None\n else:\n weight = None\n tensor_shape = None\n\n obj_list = [tensor_shape]\n dist.broadcast_object_list(obj_list, src=0, group=mp_group)\n tensor_shape = obj_list[0]\n\n if tensor_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tensor:[{name}] not in state_dict, skip load\")\n return\n\n if tensor is None:\n tensor = torch.empty(\n tensor_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n if torch.distributed.get_rank() == 0:\n tensor.data.copy_(weight)\n dist.broadcast(tensor, src=0, group=mp_group)\n\n def _broadcast_tp_shard_tensor_vocab(tensor, name, chunk_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n if torch.distributed.get_rank() == 0:\n if name in state_dict:\n full_weight = state_dict[name]\n\n if mutate_func is not None:\n full_weight = mutate_func(full_weight)\n tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)\n chunk_shape = tensor_chunk[0].shape\n else:\n chunk_shape = None\n else:\n chunk_shape = None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=0, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{name}] not in state_dict, skip loading\")\n return\n\n if tensor is None:\n sync_tensor = torch.empty(\n chunk_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n else:\n assert tensor.shape == chunk_shape, (\n f\"rank #{torch.distributed.get_rank()} tensor {name} shape {tensor.shape} != {chunk_shape}\"\n )\n sync_tensor = torch.empty_like(tensor, device=get_device_id(), requires_grad=False)\n\n for i in range(tp_size):\n if torch.distributed.get_rank() == 0:\n sync_tensor.data.copy_(tensor_chunk[i])\n dist.broadcast(sync_tensor, src=0, group=mp_group)\n if (i == tp_rank) and (tensor is not None):\n tensor.data.copy_(sync_tensor)\n\n def _broadcast_tp_shard_tensor(tensor, name, chunk_dim=0, mutate_func=None) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n if torch.distributed.get_rank() == 0:\n if name in state_dict:\n full_weight = state_dict[name]\n if mutate_func is not None:\n full_weight = mutate_func(full_weight)\n tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)\n chunk_shape = tensor_chunk[0].shape\n else:\n chunk_shape = None\n else:\n chunk_shape = None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=0, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{name}] not in state_dict, skip loading\")\n return\n\n if tensor is None:\n sync_tensor = torch.empty(\n chunk_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n else:\n assert tensor.shape == chunk_shape, (\n f\"rank #{torch.distributed.get_rank()} tensor {name} shape {tensor.shape} != {chunk_shape}\"\n )\n sync_tensor = torch.empty_like(tensor, device=get_device_id(), requires_grad=False)\n\n for i in range(tp_size):\n if torch.distributed.get_rank() == 0:\n sync_tensor.data.copy_(tensor_chunk[i])\n dist.broadcast(sync_tensor, src=0, group=mp_group)\n if (i == tp_rank) and (tensor is not None):\n tensor.data.copy_(sync_tensor)\n\n def _broadcast_tp_shard_tensor_gate_up(tensor, gate_name, up_name) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n if torch.distributed.get_rank() == 0:\n gate_weight = state_dict[gate_name]\n up_weight = state_dict[up_name]\n new_gate_up_weight = torch.empty(\n config.intermediate_size * 2, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n for i in range(tp_size):\n intermediate_size_tp = config.intermediate_size // tp_size\n gate_weight_tp = gate_weight[i * intermediate_size_tp : (i + 1) * intermediate_size_tp]\n up_weight_tp = up_weight[i * intermediate_size_tp : (i + 1) * intermediate_size_tp]\n new_gate_up_weight[intermediate_size_tp * 2 * i : intermediate_size_tp * 2 * (i + 1)].copy_(\n torch.cat([gate_weight_tp, up_weight_tp], dim=0)\n )\n\n tensor_chunk = torch.chunk(new_gate_up_weight, tp_size, dim=0)\n chunk_shape = tensor_chunk[0].shape\n else:\n chunk_shape = None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=0, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{gate_name, up_name}] not in state_dict, skip loading\")\n return\n\n if tensor is None:\n sync_tensor = torch.empty(\n chunk_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n else:\n assert tensor.shape == chunk_shape, (\n f\"rank #{torch.distributed.get_rank() == 0:} tensor {gate_name, up_name} shape \"\n f\"{tensor.shape} != {chunk_shape}\"\n )\n sync_tensor = torch.empty_like(tensor, device=get_device_id(), requires_grad=False)\n\n for i in range(tp_size):\n if torch.distributed.get_rank() == 0:\n sync_tensor.data.copy_(tensor_chunk[i])\n dist.broadcast(sync_tensor, src=0, group=mp_group)\n if (i == tp_rank) and (tensor is not None):\n tensor.data.copy_(sync_tensor)\n\n def _broadcast_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name, bias=False) -> torch.Tensor:\n \"\"\"broadcast tensor in tp shards across mp_group\"\"\"\n nonlocal state_dict\n nonlocal mp_group\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n if torch.distributed.get_rank() == 0:\n assert q_name in state_dict and k_name in state_dict and v_name in state_dict\n full_weight_q = state_dict[q_name]\n full_weight_k = state_dict[k_name]\n full_weight_v = state_dict[v_name]\n\n hidden_size_per_head = config.hidden_size // config.num_attention_heads\n\n if config.num_key_value_heads >= tp_size:\n q_size_tp = config.hidden_size // tp_size\n kv_size_tp = hidden_size_per_head * config.num_key_value_heads // tp_size\n total_size = q_size_tp + 2 * kv_size_tp\n if not bias:\n new_weight_qkv = torch.empty(\n total_size * tp_size, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n else:\n new_weight_qkv = torch.empty(total_size * tp_size, dtype=params_dtype, device=get_device_id())\n for i in range(tp_size):\n q_part = full_weight_q[i * q_size_tp : (i + 1) * q_size_tp]\n k_part = full_weight_k[i * kv_size_tp : (i + 1) * kv_size_tp]\n v_part = full_weight_v[i * kv_size_tp : (i + 1) * kv_size_tp]\n new_weight_qkv[i * total_size : (i + 1) * total_size].copy_(\n torch.cat([q_part, k_part, v_part], dim=0)\n )\n\n else:\n q_size_tp = config.hidden_size // tp_size\n kv_size_tp = hidden_size_per_head\n total_size = q_size_tp + 2 * kv_size_tp\n if not bias:\n new_weight_qkv = torch.empty(\n total_size * tp_size, config.hidden_size, dtype=params_dtype, device=get_device_id()\n )\n else:\n new_weight_qkv = torch.empty(total_size * tp_size, dtype=params_dtype, device=get_device_id())\n for i in range(tp_size):\n q_part = full_weight_q[i * q_size_tp : (i + 1) * q_size_tp]\n start_idx = i * config.num_key_value_heads // tp_size * hidden_size_per_head\n end_idx = (i * config.num_key_value_heads // tp_size + 1) * hidden_size_per_head\n k_part = full_weight_k[start_idx:end_idx]\n v_part = full_weight_v[start_idx:end_idx]\n new_weight_qkv[i * total_size : (i + 1) * total_size].copy_(\n torch.cat([q_part, k_part, v_part], dim=0)\n )\n\n tensor_chunk = torch.chunk(new_weight_qkv, tp_size, dim=0)\n chunk_shape = tensor_chunk[0].shape\n else:\n chunk_shape = None\n\n obj_list = [chunk_shape]\n dist.broadcast_object_list(obj_list, src=0, group=mp_group)\n chunk_shape = obj_list[0]\n if chunk_shape is None:\n # all or none ranks in the mp_group should reach here\n print_rank_0(f\"tp_shard tensor:[{q_name, k_name, v_name}] not in state_dict, skip loading\")\n return\n\n if tensor is None:\n sync_tensor = torch.empty(\n chunk_shape,\n dtype=params_dtype,\n device=get_device_id(),\n requires_grad=False,\n )\n else:\n assert tensor.shape == chunk_shape, (\n f\"rank #{torch.distributed.get_rank()} tensor {q_name} shape {tensor.shape} != {chunk_shape}\"\n )\n sync_tensor = torch.empty_like(tensor, device=get_device_id(), requires_grad=False)\n\n for i in range(tp_size):\n if torch.distributed.get_rank() == 0:\n sync_tensor.data.copy_(tensor_chunk[i])\n dist.broadcast(sync_tensor, src=0, group=mp_group)\n if (i == tp_rank) and (tensor is not None):\n tensor.data.copy_(sync_tensor)\n\n if dp_rank == 0:\n # Embeddings\n # -------------------\n print_rank_0(\"loading embeddings...\")\n gpt_model_module = _get_gpt_model(models[0])\n embed_tokens_weight = None\n if pp_rank == 0:\n embed_tokens_weight = gpt_model_module.model.embed_tokens.weight\n _broadcast_tp_shard_tensor_vocab(embed_tokens_weight, \"model.embed_tokens.weight\")\n\n # Transformer layers\n # -------------------\n layer_map = _megatron_calc_layer_map(config)\n\n for layer in range(config.num_hidden_layers):\n print_rank_0(f\"loading layer #{layer}...\")\n layer_name = f\"model.layers.{layer}\"\n dst_pp_rank, dst_virtual_pp_rank, dst_layer_idx = layer_map[layer]\n\n gpt_model_module = _get_gpt_model(models[dst_virtual_pp_rank])\n sync_layer = gpt_model_module.model.layers[dst_layer_idx]\n\n _broadcast_tensor(\n sync_layer.input_layernorm.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.input_layernorm.weight\",\n )\n\n _broadcast_tp_shard_tensor_qkv(\n sync_layer.self_attn.qkv_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.q_proj.weight\",\n f\"{layer_name}.self_attn.k_proj.weight\",\n f\"{layer_name}.self_attn.v_proj.weight\",\n )\n\n _broadcast_tp_shard_tensor_qkv(\n sync_layer.self_attn.qkv_proj.bias if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.q_proj.bias\",\n f\"{layer_name}.self_attn.k_proj.bias\",\n f\"{layer_name}.self_attn.v_proj.bias\",\n bias=True,\n )\n\n _broadcast_tp_shard_tensor(\n sync_layer.self_attn.o_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.self_attn.o_proj.weight\",\n chunk_dim=1,\n )\n\n _broadcast_tensor(\n sync_layer.post_attention_layernorm.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.post_attention_layernorm.weight\",\n )\n\n _broadcast_tp_shard_tensor_gate_up(\n sync_layer.mlp.gate_up_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.mlp.gate_proj.weight\",\n f\"{layer_name}.mlp.up_proj.weight\",\n )\n\n _broadcast_tp_shard_tensor(\n sync_layer.mlp.down_proj.weight if dst_pp_rank == pp_rank else None,\n f\"{layer_name}.mlp.down_proj.weight\",\n chunk_dim=1,\n )\n # Final Layernorm\n # -------------------\n print_rank_0(\"loading final layernorm...\")\n gpt_model_module = _get_gpt_model(models[-1])\n _broadcast_tensor(\n getattr(gpt_model_module.model.norm, \"weight\", None),\n \"model.norm.weight\",\n )\n\n if tie_word_embeddings:\n print_rank_0(\"tie_word_embeddings skip load lm_head\")\n else:\n print_rank_0(\"loading lm_head...\")\n lm_head_weight = None\n if pp_rank + 1 == pp_size:\n lm_head_weight = gpt_model_module.lm_head.weight\n\n if is_value_model:\n if \"lm_head.weight\" in state_dict and state_dict[\"lm_head.weight\"].shape[0] == 1:\n _broadcast_tensor(lm_head_weight, \"lm_head.weight\")\n print_rank_0(\"load lm_head from value_head weight\")\n elif \"reward_head.weight\" in state_dict and state_dict[\"reward_head.weight\"].shape[0] == 1:\n _broadcast_tensor(lm_head_weight, \"reward_head.weight\")\n print_rank_0(\"load lm_head from value_head weight\")\n else:\n _broadcast_tensor(None, \"lm_head.weight\")\n print_rank_0(\"fail to match lm_head in value_model\")\n\n else:\n _broadcast_tp_shard_tensor(lm_head_weight, \"lm_head.weight\")\n\n dist.barrier()\n # Broadcast weights inside data parallel groups\n for wrapped_model in wrapped_models:\n broadcast_params(wrapped_model)\n\n get_torch_device().empty_cache()\n print_rank_0(f\"loading megatron ckpt done, time elapsed {time.time() - start_time}s\")\n"}35{"file_name": "verl__models__qwen2__megatron__layers__parallel_attention.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.\n#\n# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX\n# and OPT implementations in this library. It has been modified from its\n# original forms to accommodate minor architectural differences compared\n# to GPT-NeoX and OPT used by the Meta AI team that trained the model.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport math\nfrom typing import Optional\n\nimport torch.nn.functional as F\nfrom einops import rearrange\nfrom transformers.utils import is_flash_attn_2_available\n\nif is_flash_attn_2_available():\n from flash_attn import flash_attn_varlen_func\n from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input # noqa: F401\n\nimport torch\nfrom flash_attn.layers.rotary import apply_rotary_emb\nfrom megatron.core import ModelParallelConfig, tensor_parallel\nfrom megatron.core import parallel_state as mpu\nfrom torch import nn\nfrom transformers import Qwen2Config\n\nfrom verl.models.qwen2.megatron.layers.parallel_linear import QKVParallelLinear\nfrom verl.utils.megatron import tensor_parallel as tp_utils\n\n\nclass Qwen2RotaryEmbedding(nn.Module):\n def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):\n super().__init__()\n\n self.dim = dim\n self.max_position_embeddings = max_position_embeddings\n self.base = base\n inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))\n self.register_buffer(\"inv_freq\", inv_freq, persistent=False)\n\n # Build here to make `torch.jit.trace` work.\n self._set_cos_sin_cache(\n seq_len=max_position_embeddings, device=self.inv_freq.device, dtype=torch.get_default_dtype()\n )\n\n def _set_cos_sin_cache(self, seq_len, device, dtype):\n self.max_seq_len_cached = seq_len\n t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)\n\n freqs = torch.einsum(\"i,j->ij\", t, self.inv_freq)\n # Different from paper, but it uses a different permutation in order to obtain the same calculation\n emb = torch.cat((freqs, freqs), dim=-1)\n self.register_buffer(\"cos_cached\", emb.cos().to(dtype), persistent=False)\n self.register_buffer(\"sin_cached\", emb.sin().to(dtype), persistent=False)\n\n def forward(self, x, seq_len=None):\n # x: [bs, num_attention_heads, seq_len, head_size]\n if seq_len > self.max_seq_len_cached:\n self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=x.dtype)\n\n return (\n self.cos_cached[:seq_len].to(dtype=x.dtype),\n self.sin_cached[:seq_len].to(dtype=x.dtype),\n )\n\n\nclass Qwen2LinearScalingRotaryEmbedding(Qwen2RotaryEmbedding):\n \"\"\"Qwen2RotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev\"\"\"\n\n def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):\n self.scaling_factor = scaling_factor\n super().__init__(dim, max_position_embeddings, base, device)\n\n def _set_cos_sin_cache(self, seq_len, device, dtype):\n self.max_seq_len_cached = seq_len\n t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)\n t = t / self.scaling_factor\n\n freqs = torch.einsum(\"i,j->ij\", t, self.inv_freq)\n # Different from paper, but it uses a different permutation in order to obtain the same calculation\n emb = torch.cat((freqs, freqs), dim=-1)\n self.register_buffer(\"cos_cached\", emb.cos().to(dtype), persistent=False)\n self.register_buffer(\"sin_cached\", emb.sin().to(dtype), persistent=False)\n\n\nclass Qwen2DynamicNTKScalingRotaryEmbedding(Qwen2RotaryEmbedding):\n \"\"\"Qwen2RotaryEmbedding extended with Dynamic NTK scaling. Credits to the Reddit users /u/bloc97 and /u/emozilla\"\"\"\n\n def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):\n self.scaling_factor = scaling_factor\n super().__init__(dim, max_position_embeddings, base, device)\n\n def _set_cos_sin_cache(self, seq_len, device, dtype):\n self.max_seq_len_cached = seq_len\n\n if seq_len > self.max_position_embeddings:\n base = self.base * (\n (self.scaling_factor * seq_len / self.max_position_embeddings) - (self.scaling_factor - 1)\n ) ** (self.dim / (self.dim - 2))\n inv_freq = 1.0 / (base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))\n self.register_buffer(\"inv_freq\", inv_freq, persistent=False)\n\n t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)\n\n freqs = torch.einsum(\"i,j->ij\", t, self.inv_freq)\n # Different from paper, but it uses a different permutation in order to obtain the same calculation\n emb = torch.cat((freqs, freqs), dim=-1)\n self.register_buffer(\"cos_cached\", emb.cos().to(dtype), persistent=False)\n self.register_buffer(\"sin_cached\", emb.sin().to(dtype), persistent=False)\n\n\ndef rotate_half(x):\n \"\"\"Rotates half the hidden dims of the input.\"\"\"\n x1 = x[..., : x.shape[-1] // 2]\n x2 = x[..., x.shape[-1] // 2 :]\n return torch.cat((-x2, x1), dim=-1)\n\n\ndef apply_rotary_pos_emb(q, k, cos, sin, position_ids):\n cos = cos[position_ids].unsqueeze(1) # [bs, 1, seq_len, dim]\n sin = sin[position_ids].unsqueeze(1) # [bs, 1, seq_len, dim]\n q_embed = (q * cos) + (rotate_half(q) * sin)\n k_embed = (k * cos) + (rotate_half(k) * sin)\n return q_embed, k_embed\n\n\ndef repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:\n \"\"\"\n This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,\n num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)\n \"\"\"\n batch, num_key_value_heads, slen, head_dim = hidden_states.shape\n if n_rep == 1:\n return hidden_states\n hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)\n return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)\n\n\nclass ParallelQwen2Attention(nn.Module):\n \"\"\"Multi-headed attention from 'Attention Is All You Need' paper\"\"\"\n\n def __init__(self, config: Qwen2Config, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config = config\n self.megatron_config = megatron_config\n self.hidden_size = config.hidden_size\n self.num_heads = config.num_attention_heads\n self.head_dim = self.hidden_size // self.num_heads\n self.num_key_value_heads = config.num_key_value_heads\n self.num_key_value_groups = self.num_heads // self.num_key_value_heads\n self.max_position_embeddings = config.max_position_embeddings\n self.rope_theta = config.rope_theta\n\n # assign values after tp\n tp_size = mpu.get_tensor_model_parallel_world_size()\n assert self.num_heads % tp_size == 0, (\n f\"num_head must be divisible by tp_size. Got num_head={self.num_heads}, tp_size={tp_size}\"\n )\n assert self.num_key_value_heads % tp_size == 0, (\n f\"num_key_value_heads must be divisible by tp_size. Got num_key_value_heads=\"\n f\"{self.num_key_value_heads}, tp_size={tp_size}\"\n )\n\n self.num_heads_per_tp = self.num_heads // tp_size\n self.num_key_value_heads_per_tp = self.num_key_value_heads // tp_size\n self.hidden_size_per_tp = self.hidden_size // tp_size\n\n if (self.head_dim * self.num_heads) != self.hidden_size:\n raise ValueError(\n f\"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size} and \"\n f\"`num_heads`: {self.num_heads}).\"\n )\n\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n row_kwargs = tp_utils.get_default_kwargs_for_row_parallel_linear()\n\n if megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n assert row_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, megatron_config)\n tp_utils.update_kwargs_with_config(row_kwargs, megatron_config)\n\n # [self.q_size, self.k_size, self.v_size]\n self.qkv_proj = QKVParallelLinear(\n input_size=self.hidden_size,\n num_heads=self.num_heads,\n num_key_value_heads=self.num_key_value_heads,\n head_dim=self.head_dim,\n # bias=config.attention_bias,\n bias=True,\n gather_output=False,\n skip_bias_add=False,\n **column_kwargs,\n )\n\n self.q_size = self.num_heads_per_tp * self.head_dim\n self.k_size = self.num_key_value_heads_per_tp * self.head_dim\n self.v_size = self.num_key_value_heads_per_tp * self.head_dim\n\n self.o_proj = tensor_parallel.RowParallelLinear(\n input_size=self.num_heads * self.head_dim,\n output_size=self.hidden_size,\n # bias=config.attention_bias,\n bias=False,\n input_is_parallel=True,\n skip_bias_add=False,\n **row_kwargs,\n )\n\n self._init_rope()\n\n def _init_rope(self):\n self.rotary_emb = Qwen2RotaryEmbedding(\n self.head_dim,\n max_position_embeddings=self.max_position_embeddings,\n base=self.rope_theta,\n )\n\n def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):\n return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()\n\n def forward(\n self,\n hidden_states: torch.Tensor,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]:\n bsz, q_len, _ = hidden_states.size()\n qkv = self.qkv_proj(hidden_states)[0]\n query_states, key_states, value_states = qkv.split([self.q_size, self.k_size, self.v_size], dim=-1)\n\n query_states = query_states.view(bsz, q_len, self.num_heads_per_tp, self.head_dim).transpose(1, 2)\n key_states = key_states.view(bsz, q_len, self.num_key_value_heads_per_tp, self.head_dim).transpose(1, 2)\n value_states = value_states.view(bsz, q_len, self.num_key_value_heads_per_tp, self.head_dim).transpose(1, 2)\n\n kv_seq_len = key_states.shape[-2]\n cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)\n query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)\n\n key_states = repeat_kv(key_states, self.num_key_value_groups)\n value_states = repeat_kv(value_states, self.num_key_value_groups)\n\n attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)\n\n if attn_weights.size() != (bsz, self.num_heads_per_tp, q_len, kv_seq_len):\n raise ValueError(\n f\"Attention weights should be of size {(bsz, self.num_heads_per_tp, q_len, kv_seq_len)}, \"\n f\"but is {attn_weights.size()}\"\n )\n\n if attention_mask is not None:\n if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):\n raise ValueError(\n f\"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}\"\n )\n attn_weights = attn_weights + attention_mask\n\n # upcast attention to fp32\n attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)\n attn_output = torch.matmul(attn_weights, value_states)\n\n if attn_output.size() != (bsz, self.num_heads_per_tp, q_len, self.head_dim):\n raise ValueError(\n f\"`attn_output` should be of size {(bsz, self.num_heads_per_tp, q_len, self.head_dim)}, \"\n f\"but is {attn_output.size()}\"\n )\n\n attn_output = attn_output.transpose(1, 2).contiguous()\n attn_output = attn_output.reshape(bsz, q_len, self.hidden_size_per_tp)\n attn_output = self.o_proj(attn_output)[0]\n return attn_output\n\n\n\"\"\"\nRemove padding Attention\n- Using Flash-attn 2\n- Compatible with sequence parallel\n\"\"\"\n\n\ndef apply_rotary_pos_emb_rmpad(q, k, cos, sin, position_ids, indices, sequence_length):\n batch_size = position_ids.shape[0]\n\n q = pad_input(q, indices, batch_size, sequence_length) # (batch_size, seqlen, num_head, head_dim)\n k = pad_input(k, indices, batch_size, sequence_length)\n cos = cos[position_ids].unsqueeze(2) # [bs, seq_len, 1, dim]\n sin = sin[position_ids].unsqueeze(2) # [bs, seq_len, 1, dim]\n q_embed = (q * cos) + (rotate_half(q) * sin)\n k_embed = (k * cos) + (rotate_half(k) * sin)\n\n q_embed = index_first_axis(rearrange(q_embed, \"b s ... -> (b s) ...\"), indices)\n k_embed = index_first_axis(rearrange(k_embed, \"b s ... -> (b s) ...\"), indices)\n\n return q_embed, k_embed\n\n\n# use flash-attn rotary embeddings with rmpad\n# cos/sin shoudl be: (seq_length, rotary_dim / 2)\ndef apply_rotary_pos_emb_rmpad_flash(q, k, cos, sin, cu_seqlens, max_seqlen):\n q_embed = apply_rotary_emb(\n q, cos, sin, interleaved=False, inplace=False, cu_seqlens=cu_seqlens, max_seqlen=max_seqlen\n )\n k_embed = apply_rotary_emb(\n k, cos, sin, interleaved=False, inplace=False, cu_seqlens=cu_seqlens, max_seqlen=max_seqlen\n )\n return q_embed, k_embed\n\n\nclass ParallelQwen2AttentionRmPad(ParallelQwen2Attention):\n def forward(\n self,\n hidden_states: torch.Tensor,\n position_ids: Optional[torch.LongTensor] = None,\n sequence_length: int = None,\n indices: torch.Tensor = None,\n cu_seqlens: torch.Tensor = None,\n max_seqlen_in_batch: int = None,\n ):\n total_nnz, _, _ = hidden_states.size() # This is the total_nnz padded after sequence parallel\n\n if self.megatron_config.sequence_parallel:\n total_nnz = total_nnz * mpu.get_tensor_model_parallel_world_size()\n\n qkv = self.qkv_proj(hidden_states)[0]\n query_states, key_states, value_states = qkv.split(\n [self.q_size, self.k_size, self.v_size], dim=-1\n ) # (total_nnz, 1, hidden_size)\n\n if self.megatron_config.sequence_parallel:\n sequence_parallel_pad = total_nnz - cu_seqlens[-1]\n total_nnz = cu_seqlens[-1] # total_nnz before sp padding\n query_states = query_states[:total_nnz]\n key_states = key_states[:total_nnz]\n value_states = value_states[:total_nnz]\n\n # Flash attention requires the input to have the shape\n # batch_size x seq_length x head_dime x hidden_dim\n # therefore we just need to keep the original shape\n query_states = query_states.view(total_nnz, self.num_heads_per_tp, self.head_dim)\n key_states = key_states.view(total_nnz, self.num_key_value_heads_per_tp, self.head_dim)\n value_states = value_states.view(total_nnz, self.num_key_value_heads_per_tp, self.head_dim)\n\n cos, sin = self.rotary_emb(value_states, seq_len=sequence_length)\n cos, sin = cos[:, : cos.shape[1] // 2], sin[:, : sin.shape[1] // 2] # flash attn only needs half\n query_states, key_states = apply_rotary_pos_emb_rmpad_flash(\n query_states, key_states, cos, sin, cu_seqlens=cu_seqlens, max_seqlen=max_seqlen_in_batch\n )\n # query_states, key_states = apply_rotary_pos_emb_rmpad(query_states, key_states, cos, sin,\n # position_ids, indices,\n\n # It is recommended to use dropout with FA according to the docs\n # when training.\n dropout_rate = 0.0 # if not self.training else self.attn_dropout\n\n # In PEFT, usually we cast the layer norms in float32 for training stability reasons\n # therefore the input hidden states gets silently casted in float32. Hence, we need\n # cast them back in float16 just to be sure everything works as expected.\n # This might slowdown training & inference so it is recommended to not cast the LayerNorms\n # in fp32. (Qwen2RMSNorm handles it correctly)\n input_dtype = query_states.dtype\n if input_dtype == torch.float32:\n query_states = query_states.to(torch.float16)\n key_states = key_states.to(torch.float16)\n value_states = value_states.to(torch.float16)\n\n attn_output_unpad = flash_attn_varlen_func(\n query_states,\n key_states,\n value_states,\n cu_seqlens_q=cu_seqlens,\n cu_seqlens_k=cu_seqlens,\n max_seqlen_q=max_seqlen_in_batch,\n max_seqlen_k=max_seqlen_in_batch,\n dropout_p=dropout_rate,\n softmax_scale=None,\n causal=True,\n )\n\n attn_output_unpad = attn_output_unpad.to(input_dtype)\n attn_output_unpad = attn_output_unpad.reshape(total_nnz, 1, self.hidden_size_per_tp).contiguous()\n\n # sequence parallel reduce_scatter is performed inside RowColumnParallel if enabled\n # Here we need to repad\n if self.megatron_config.sequence_parallel:\n attn_output_unpad = F.pad(attn_output_unpad, pad=(0, 0, 0, 0, 0, sequence_parallel_pad))\n\n attn_output_unpad = self.o_proj(attn_output_unpad)[0]\n return attn_output_unpad\n"}36{"file_name": "verl__models__qwen2__megatron__layers__parallel_mlp.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.\n#\n# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX\n# and OPT implementations in this library. It has been modified from its\n# original forms to accommodate minor architectural differences compared\n# to GPT-NeoX and OPT used by the Meta AI team that trained the model.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom megatron.core import ModelParallelConfig, tensor_parallel\nfrom megatron.core import parallel_state as mpu\nfrom torch import nn\nfrom transformers.activations import ACT2FN\n\nfrom verl.models.qwen2.megatron.layers.parallel_linear import MergedColumnParallelLinear\nfrom verl.utils.megatron import tensor_parallel as tp_utils\n\n\nclass ParallelQwen2MLP(nn.Module):\n def __init__(self, config, megatron_config: ModelParallelConfig = None) -> None:\n super().__init__()\n self.config = config\n self.hidden_size = config.hidden_size\n self.intermediate_size = config.intermediate_size\n # The weight is only [hidden_size, intermediate_size // model_parallel_world_size]\n\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n row_kwargs = tp_utils.get_default_kwargs_for_row_parallel_linear()\n\n if megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n assert row_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(row_kwargs, megatron_config)\n tp_utils.update_kwargs_with_config(column_kwargs, megatron_config)\n\n tp_size = mpu.get_tensor_model_parallel_world_size()\n\n self.gate_up_proj = MergedColumnParallelLinear(\n input_size=self.hidden_size,\n gate_ouput_size=self.intermediate_size,\n up_output_size=self.intermediate_size,\n bias=False,\n gather_output=False,\n skip_bias_add=False,\n **column_kwargs,\n )\n self.gate_size = self.intermediate_size // tp_size\n\n self.down_proj = tensor_parallel.RowParallelLinear(\n input_size=self.intermediate_size,\n output_size=self.hidden_size,\n bias=False,\n input_is_parallel=True,\n skip_bias_add=False,\n **row_kwargs,\n )\n\n self.act_fn = ACT2FN[config.hidden_act]\n\n def forward(self, x):\n gate_up = self.gate_up_proj(x)[0]\n gate, up = gate_up.split(self.gate_size, dim=-1)\n return self.down_proj(self.act_fn(gate) * up)[0]\n"}37{"file_name": "verl__models__qwen2__megatron__modeling_qwen2_megatron.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.\n#\n# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX\n# and OPT implementations in this library. It has been modified from its\n# original forms to accommodate minor architectural differences compared\n# to GPT-NeoX and OPT used by the Meta AI team that trained the model.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"PyTorch Qwen2 model.\"\"\"\n\nfrom typing import Optional\n\nimport torch\nimport torch.utils.checkpoint\nfrom megatron.core import ModelParallelConfig, mpu, parallel_state, tensor_parallel\nfrom torch import nn\nfrom transformers.modeling_outputs import BaseModelOutputWithPast\nfrom transformers.models.qwen2.configuration_qwen2 import Qwen2Config\nfrom transformers.models.qwen2.modeling_qwen2 import CausalLMOutputWithPast\n\nfrom verl.utils.device import get_device_name\nfrom verl.utils.megatron import sequence_parallel as sp_utils\nfrom verl.utils.megatron import tensor_parallel as tp_utils\nfrom verl.utils.megatron_utils import TransformerConfig, convert_config\n\nfrom .layers import ParallelQwen2DecoderLayer, ParallelQwen2DecoderLayerRmPad, ParallelQwen2RMSNorm\n\n\"\"\"\nTODO: \n1. Add weight initialization. Here we need to be careful on TP weight init.\n2. Add sequence parallel\n3. Load checkpoint from Qwen2 pretrained checkpoint\n\"\"\"\n\n\n# Copied from transformers.models.bart.modeling_bart._make_causal_mask\ndef _make_causal_mask(input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device):\n \"\"\"\n Make causal mask used for bi-directional self-attention.\n \"\"\"\n bsz, tgt_len = input_ids_shape\n mask = torch.full((tgt_len, tgt_len), torch.finfo(dtype).min, device=device)\n mask_cond = torch.arange(mask.size(-1), device=device)\n mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)\n mask = mask.to(dtype)\n return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len)\n\n\n# Copied from transformers.models.bart.modeling_bart._expand_mask\ndef _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):\n \"\"\"\n Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.\n \"\"\"\n bsz, src_len = mask.size()\n tgt_len = tgt_len if tgt_len is not None else src_len\n\n expanded_mask = mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)\n\n inverted_mask = 1.0 - expanded_mask\n\n return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)\n\n\nclass ParallelQwen2Model(nn.Module):\n \"\"\"\n Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`Qwen2DecoderLayer`]\n\n Args:\n config: Qwen2Config\n \"\"\"\n\n def __init__(self, config: Qwen2Config, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.padding_idx = config.pad_token_id\n self.vocab_size = config.vocab_size\n embedding_kwargs = tp_utils.get_default_kwargs_for_parallel_embedding()\n if megatron_config is not None:\n assert embedding_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(embedding_kwargs, megatron_config)\n self.embed_tokens = tensor_parallel.VocabParallelEmbedding(\n num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, **embedding_kwargs\n )\n\n self.layers = nn.ModuleList(\n [ParallelQwen2DecoderLayer(config, megatron_config) for _ in range(config.num_hidden_layers)]\n )\n self.norm = ParallelQwen2RMSNorm(config, megatron_config)\n\n # Copied from transformers.models.bart.modeling_bart.BartDecoder._prepare_decoder_attention_mask\n def _prepare_decoder_attention_mask(self, attention_mask, input_shape, inputs_embeds):\n # create causal mask\n # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]\n combined_attention_mask = None\n if input_shape[-1] > 1:\n combined_attention_mask = _make_causal_mask(\n input_shape,\n inputs_embeds.dtype,\n device=inputs_embeds.device,\n )\n\n if attention_mask is not None:\n # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]\n expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to(\n inputs_embeds.device\n )\n combined_attention_mask = (\n expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask\n )\n\n return combined_attention_mask\n\n def forward(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | BaseModelOutputWithPast:\n \"\"\"\n\n Args:\n input_ids: input ids. shape (batch_size, seq_length)\n attention_mask: attention_mask. shape (batch_size, seq_length)\n position_ids: position ids. shape (batch_size, seq_length)\n\n Returns:\n\n \"\"\"\n batch_size, seq_length = input_ids.shape\n inputs_embeds = self.embed_tokens(input_ids)\n # embed positions\n\n attention_mask = self._prepare_decoder_attention_mask(attention_mask, (batch_size, seq_length), inputs_embeds)\n\n hidden_states = inputs_embeds\n\n for idx, decoder_layer in enumerate(self.layers):\n layer_outputs = decoder_layer(\n hidden_states,\n attention_mask=attention_mask,\n position_ids=position_ids,\n )\n\n hidden_states = layer_outputs\n\n hidden_states = self.norm(hidden_states)\n\n return hidden_states\n\n\nclass ParallelQwen2ForCausalLM(nn.Module):\n def __init__(self, config: Qwen2Config, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.model = ParallelQwen2Model(config, megatron_config=megatron_config)\n self.vocab_size = config.vocab_size\n\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n\n self.lm_head = tensor_parallel.ColumnParallelLinear(\n input_size=config.hidden_size,\n output_size=config.vocab_size,\n bias=False,\n gather_output=False,\n skip_bias_add=False,\n **column_kwargs,\n )\n\n def forward(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n r\"\"\"\n Args:\n labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):\n Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,\n config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored\n (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.\n\n Returns:\n ```\"\"\"\n\n # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)\n outputs = self.model(\n input_ids=input_ids,\n attention_mask=attention_mask,\n position_ids=position_ids,\n )\n\n hidden_states = outputs\n logits = self.lm_head(hidden_states)[0]\n\n logits = tensor_parallel.gather_from_tensor_model_parallel_region(logits)\n\n logits = logits.float()\n return CausalLMOutputWithPast(\n loss=None,\n logits=logits,\n past_key_values=None,\n hidden_states=None,\n attentions=None,\n )\n\n\nfrom flash_attn.bert_padding import index_first_axis, pad_input, unpad_input # noqa: F401, E402\n\n\nclass ParallelQwen2ModelRmPad(nn.Module):\n \"\"\"\n Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`Qwen2DecoderLayer`]\n\n Args:\n config: Qwen2Config\n \"\"\"\n\n def __init__(self, config: Qwen2Config, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.padding_idx = config.pad_token_id\n self.vocab_size = config.vocab_size\n embedding_kwargs = tp_utils.get_default_kwargs_for_parallel_embedding()\n self.megatron_config = megatron_config\n if megatron_config is not None:\n assert embedding_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(embedding_kwargs, self.megatron_config)\n self.embed_tokens = tensor_parallel.VocabParallelEmbedding(\n num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, **embedding_kwargs\n )\n\n self.layers = nn.ModuleList(\n [ParallelQwen2DecoderLayerRmPad(config, megatron_config) for _ in range(config.num_hidden_layers)]\n )\n self.norm = ParallelQwen2RMSNorm(config, megatron_config)\n\n def forward(\n self,\n input_ids: torch.Tensor,\n position_ids: Optional[torch.LongTensor] = None,\n sequence_length: int = None,\n indices: torch.Tensor = None,\n cu_seqlens: int = None,\n max_seqlen_in_batch: int = None,\n ) -> tuple | BaseModelOutputWithPast:\n \"\"\"\n\n Args:\n input_ids: input ids. shape (1, totol_nnz)\n position_ids: position ids. shape (batch_size, seq_length)\n\n Returns:\n\n \"\"\"\n inputs_embeds = self.embed_tokens(input_ids) # (1, total_nnz) -> (1, total_nnz, hidden_size)\n\n # (1, total_nnz, hidden_size) -> (total_nnz, 1, hidden_size) -> (total_nnz // sp, 1, hidden_size)\n inputs_embeds = inputs_embeds.transpose(0, 1)\n if self.megatron_config.sequence_parallel:\n inputs_embeds = tensor_parallel.scatter_to_sequence_parallel_region(inputs_embeds)\n\n hidden_states = inputs_embeds\n for idx, decoder_layer in enumerate(self.layers):\n layer_outputs = decoder_layer(\n hidden_states,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n hidden_states = layer_outputs\n\n hidden_states = self.norm(hidden_states)\n\n return hidden_states\n\n\nclass ParallelQwen2ForCausalLMRmPad(nn.Module):\n def __init__(self, config: Qwen2Config, megatron_config: ModelParallelConfig):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.megatron_config = megatron_config\n self.model = ParallelQwen2ModelRmPad(config, megatron_config=megatron_config)\n self.vocab_size = config.vocab_size\n self._init_head(config)\n\n def _init_head(self, config: Qwen2Config):\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if self.megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n self.lm_head = tensor_parallel.ColumnParallelLinear(\n input_size=config.hidden_size,\n output_size=config.vocab_size,\n bias=False,\n gather_output=False,\n skip_bias_add=False,\n **column_kwargs,\n )\n\n def _forward_head(self, hidden_states):\n # all_gather from sequence parallel region is performed inside lm_head\n logits = self.lm_head(hidden_states)[0]\n logits = logits.float() # (total_nnz_padded, 1, vocab_size // tp)\n logits = tensor_parallel.gather_from_tensor_model_parallel_region(logits) # (total_nnz_padded, 1, vocab_size)\n return logits\n\n def forward(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n r\"\"\"\n Args:\n labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):\n Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,\n config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored\n (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.\n\n Returns:\n ```\"\"\"\n batch_size, sequence_length = input_ids.shape\n\n # remove padding here\n input_ids, indices, cu_seqlens, max_seqlen_in_batch, *_ = unpad_input(\n input_ids.unsqueeze(dim=-1), attention_mask\n ) # (total_nnz, 1)\n\n # pad input_ids to multiple of tp for all tp ranks\n # TODO: for better performance, the sp padding should be removed at each layer. Not sure the performance gap\n if self.megatron_config.sequence_parallel:\n input_ids = sp_utils.pad_to_sequence_parallel(input_ids)\n\n input_ids = input_ids.transpose(0, 1) # (1, total_nnz+pad)\n\n outputs = self.model(\n input_ids=input_ids,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n hidden_states = outputs\n\n logits = self._forward_head(hidden_states)\n\n # remove padding from sequence parallel\n if self.megatron_config.sequence_parallel:\n totol_nnz = cu_seqlens[-1]\n logits = logits[:totol_nnz] # (total_nnz_padded)\n\n logits = torch.squeeze(logits, dim=1) # remove the artificial batch dimension\n # add removed padding back\n logits = pad_input(\n logits, indices, batch_size, seqlen=sequence_length\n ) # (batch_size, sequence_length, vocab_size)\n\n return CausalLMOutputWithPast(\n loss=None,\n logits=logits,\n past_key_values=None,\n hidden_states=None,\n attentions=None,\n )\n\n\nclass ParallelQwen2ForValueRmPad(ParallelQwen2ForCausalLMRmPad):\n def _init_head(self, config):\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if self.megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n self.lm_head = nn.Linear(in_features=config.hidden_size, out_features=1, bias=False)\n # lm_head is effectively the same as sequence parallel\n sp_utils.mark_parameter_as_sequence_parallel(self.lm_head.weight)\n\n def _forward_head(self, hidden_states):\n logits = self.lm_head(hidden_states) # (total_nnz_padded // tp, 1, 1)\n logits = logits.float()\n if self.megatron_config.sequence_parallel:\n logits = tensor_parallel.gather_from_sequence_parallel_region(logits, tensor_parallel_output_grad=False)\n return logits\n\n def forward(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n output = super().forward(input_ids, attention_mask, position_ids)\n output.logits = torch.squeeze(output.logits, dim=-1)\n return output\n\n\n\"\"\"\nSupport pipeline parallelism\n\"\"\"\n\n\nclass ParallelQwen2ModelRmPadPP(nn.Module):\n \"\"\"\n Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`Qwen2DecoderLayer`]\n This model definition supports pipeline parallelism. To support pp and vpp,\n - This model only contains layer in this pp stage and vpp chunk\n - When calling get_model in Megatron, this rank will instantiate all the vpp chunks in this pp.\n Args:\n config: Qwen2Config\n \"\"\"\n\n def __init__(self, config: Qwen2Config, megatron_config: ModelParallelConfig, pre_process, post_process):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.padding_idx = config.pad_token_id\n self.vocab_size = config.vocab_size\n self.pre_process = pre_process\n self.post_process = post_process\n self.megatron_config = megatron_config\n embedding_kwargs = tp_utils.get_default_kwargs_for_parallel_embedding()\n if megatron_config is not None:\n assert embedding_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(embedding_kwargs, self.megatron_config)\n if pre_process:\n self.embed_tokens = tensor_parallel.VocabParallelEmbedding(\n num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, **embedding_kwargs\n )\n else:\n self.embed_tokens = None\n\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n pp_size = megatron_config.pipeline_model_parallel_size\n self.num_layer_per_pp = config.num_hidden_layers // pp_size\n vpp_size = megatron_config.virtual_pipeline_model_parallel_size\n vpp_rank = mpu.get_virtual_pipeline_model_parallel_rank()\n\n if vpp_size is not None:\n self.num_layer_vpp_chunk = self.num_layer_per_pp // vpp_size\n self.num_layer_this_model = self.num_layer_vpp_chunk\n offset = vpp_rank * (config.num_hidden_layers // vpp_size) + (pp_rank * self.num_layer_vpp_chunk)\n else:\n self.num_layer_this_model = self.num_layer_per_pp\n offset = pp_rank * self.num_layer_per_pp\n\n self.layers = nn.ModuleList()\n for i in range(self.num_layer_this_model):\n layer = ParallelQwen2DecoderLayerRmPad(config, megatron_config, layer_idx=i + offset)\n self.layers.add_module(f\"{i}\", layer)\n\n if post_process:\n self.norm = ParallelQwen2RMSNorm(config, megatron_config)\n else:\n self.norm = None\n\n def set_input_tensor(self, input_tensor):\n \"\"\"Set input tensor to be used instead of forward()'s input.\n\n When doing pipeline parallelism the input from the previous\n stage comes from communication, not from the input, so the\n model's forward_step_func won't have it. This function is thus\n used by internal code to bypass the input provided by the\n forward_step_func\"\"\"\n self.input_tensor = input_tensor\n\n def forward(\n self,\n input_ids: torch.Tensor,\n position_ids: Optional[torch.LongTensor] = None,\n sequence_length: int = None,\n indices: torch.Tensor = None,\n cu_seqlens: int = None,\n max_seqlen_in_batch: int = None,\n ) -> tuple | BaseModelOutputWithPast:\n \"\"\"\n\n Args:\n input_ids: input ids. shape (1, totol_nnz)\n position_ids: position ids. shape (batch_size, seq_length)\n\n Returns:\n\n \"\"\"\n if self.pre_process:\n inputs_embeds = self.embed_tokens(input_ids) # (1, total_nnz) -> (1, total_nnz, hidden_size)\n\n # vocab parallel embedding will not do sequence parallel reduce-scatter in open source megatron\n # so need to deal with it by handle here:\n # (1, total_nnz, hidden_size) -> (total_nnz, 1, hidden_size) -> (total_nnz // sp, 1, hidden_size)\n inputs_embeds = inputs_embeds.transpose(0, 1)\n if self.megatron_config.sequence_parallel:\n inputs_embeds = tensor_parallel.scatter_to_sequence_parallel_region(inputs_embeds)\n\n hidden_states = inputs_embeds\n else:\n # self.hidden_states should be passed by Megatron\n hidden_states = self.input_tensor\n\n for idx, decoder_layer in enumerate(self.layers):\n layer_outputs = decoder_layer(\n hidden_states,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n hidden_states = layer_outputs\n\n if self.post_process:\n hidden_states = self.norm(hidden_states)\n\n return hidden_states\n\n\nclass ParallelQwen2ForCausalLMRmPadPP(nn.Module):\n def __init__(\n self,\n config: Qwen2Config,\n megatron_config: ModelParallelConfig,\n pre_process,\n post_process,\n share_embeddings_and_output_weights,\n ):\n super().__init__()\n self.config: TransformerConfig = convert_config(config, megatron_config)\n self.megatron_config = megatron_config\n self.model = ParallelQwen2ModelRmPadPP(\n config, megatron_config=megatron_config, pre_process=pre_process, post_process=post_process\n )\n self.share_embeddings_and_output_weights = share_embeddings_and_output_weights\n self.vocab_size = config.vocab_size\n self.pre_process = pre_process\n self.post_process = post_process\n if post_process:\n self._init_head(config)\n if pre_process or post_process:\n self.setup_embeddings_and_output_layer()\n\n def set_input_tensor(self, input_tensor):\n \"\"\"Set input tensor to be used instead of forward()'s input.\n\n When doing pipeline parallelism the input from the previous\n stage comes from communication, not from the input, so the\n model's forward_step_func won't have it. This function is thus\n used by internal code to bypass the input provided by the\n forward_step_func\"\"\"\n assert len(input_tensor) == 1\n self.model.set_input_tensor(input_tensor[0])\n\n def _init_head(self, config):\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if self.megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n self.lm_head = tensor_parallel.ColumnParallelLinear(\n input_size=config.hidden_size,\n output_size=config.vocab_size,\n bias=False,\n gather_output=False,\n skip_bias_add=False,\n skip_weight_param_allocation=self.pre_process and self.share_embeddings_and_output_weights,\n **column_kwargs,\n )\n\n def setup_embeddings_and_output_layer(self) -> None:\n \"\"\"Sets up embedding layer in first stage and output layer in last stage.\n\n This function initializes word embeddings in the final stage when we are\n using pipeline parallelism and sharing word embeddings, and sets up param\n attributes on the embedding and output layers.\n \"\"\"\n # Set `is_embedding_or_output_parameter` attribute.\n if self.pre_process:\n self.model.embed_tokens.weight.is_embedding_or_output_parameter = True\n if self.post_process and self.lm_head.weight is not None:\n self.lm_head.weight.is_embedding_or_output_parameter = True\n\n if not self.share_embeddings_and_output_weights:\n return\n\n if parallel_state.get_pipeline_model_parallel_world_size() == 1:\n # Zero out wgrad if sharing embeddings between two layers on same\n # pipeline stage to make sure grad accumulation into main_grad is\n # correct and does not include garbage values (e.g., from torch.empty).\n self.shared_embedding_or_output_weight().zero_out_wgrad = True\n return\n\n if parallel_state.is_pipeline_first_stage() and self.pre_process and not self.post_process:\n self.shared_embedding_or_output_weight().shared_embedding = True\n\n if self.post_process and not self.pre_process:\n assert not parallel_state.is_pipeline_first_stage()\n # set word_embeddings weights to 0 here, then copy first\n # stage's weights using all_reduce below.\n self.lm_head.weight.data.fill_(0)\n self.lm_head.weight.shared = True\n self.lm_head.weight.shared_embedding = True\n\n if torch.distributed.is_initialized() and parallel_state.is_rank_in_embedding_group():\n weight = self.shared_embedding_or_output_weight()\n weight.data = weight.data.to(get_device_name())\n torch.distributed.all_reduce(weight.data, group=parallel_state.get_embedding_group())\n\n def shared_embedding_or_output_weight(self) -> torch.Tensor:\n if self.pre_process:\n return self.model.embed_tokens.weight\n elif self.post_process:\n return self.lm_head.weight\n return None\n\n def _forward_head(self, hidden_states):\n # all_gather from sequence parallel region is performed inside lm_head\n # print(f'logits shape before forward_head: {hidden_states.shape}, vocab_size = '\n # f'{self.config.vocab_size}') # [4, 32, 4096]\n output_weight = None\n if self.share_embeddings_and_output_weights:\n output_weight = self.shared_embedding_or_output_weight()\n logits = self.lm_head(hidden_states, weight=output_weight)[0]\n # print(f'logits shape after forward_head: {logits.shape}') # [8, 32, 8]\n logits = logits.float() # (total_nnz_padded, 1, vocab_size // tp)\n return logits\n\n def forward(\n self,\n # original input\n *,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n r\"\"\"\n Args:\n labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):\n Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,\n config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored\n (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.\n\n Returns:\n ```\"\"\"\n\n # Note that input_ids, attention_mask and position_ids should be passed to every pp layer.\n # In the first pp, input_ids will be used, in other pp layers hidden_states will be used inside self.model\n batch_size, sequence_length = input_ids.shape\n # remove padding here\n input_ids_rmpad, indices, cu_seqlens, max_seqlen_in_batch, *_ = unpad_input(\n input_ids.unsqueeze(dim=-1), attention_mask\n ) # (total_nnz, 1)\n\n # pad input_ids to multiple of tp for all tp ranks\n # TODO: for better performance, the sp padding should be removed at each layer. Not sure the performance gap\n if self.megatron_config.sequence_parallel:\n input_ids_rmpad = sp_utils.pad_to_sequence_parallel(input_ids_rmpad)\n\n input_ids_rmpad = input_ids_rmpad.transpose(0, 1) # (1, total_nnz+pad)\n\n outputs = self.model(\n input_ids=input_ids_rmpad,\n position_ids=position_ids,\n sequence_length=sequence_length,\n indices=indices,\n cu_seqlens=cu_seqlens,\n max_seqlen_in_batch=max_seqlen_in_batch,\n )\n\n if self.post_process:\n hidden_states = outputs\n logits = self._forward_head(hidden_states)\n logits = torch.squeeze(logits, dim=1) # remove the artificial batch dimension # torch.Size([8, 32, 16])\n\n # remove padding from sequence parallel\n if self.megatron_config.sequence_parallel:\n totol_nnz = cu_seqlens[-1]\n logits = logits[:totol_nnz] # (total_nnz_padded)\n # add removed padding back. If input is already rmpad, we let the caller pad_input\n logits = pad_input(\n logits, indices, batch_size, seqlen=sequence_length\n ) # (batch_size, sequence_length, vocab_size)\n\n return CausalLMOutputWithPast(\n loss=None,\n logits=logits,\n past_key_values=None,\n hidden_states=None,\n attentions=None,\n )\n else:\n return outputs\n\n\nclass ParallelQwen2ForValueRmPadPP(ParallelQwen2ForCausalLMRmPadPP):\n def _init_head(self, config):\n column_kwargs = tp_utils.get_default_kwargs_for_column_parallel_linear()\n if self.megatron_config is not None:\n assert column_kwargs.get(\"config\", False), \"must have ModelParallelConfig\"\n tp_utils.update_kwargs_with_config(column_kwargs, self.megatron_config)\n self.lm_head = nn.Linear(in_features=config.hidden_size, out_features=1, bias=False)\n # lm_head is effectively the same as sequence parallel\n sp_utils.mark_parameter_as_sequence_parallel(self.lm_head.weight)\n\n def _forward_head(self, hidden_states):\n logits = self.lm_head(hidden_states) # (total_nnz_padded // tp, 1, 1)\n logits = logits.float()\n if self.megatron_config.sequence_parallel:\n logits = tensor_parallel.gather_from_sequence_parallel_region(logits, tensor_parallel_output_grad=False)\n return logits\n\n def forward(\n self,\n *,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n ) -> tuple | CausalLMOutputWithPast:\n output = super().forward(input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids)\n if self.post_process:\n output.logits = torch.squeeze(output.logits, dim=-1)\n return output\n else:\n return output\n"}38{"file_name": "verl__models__registry.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport importlib\nfrom typing import Optional\n\nimport torch.nn as nn\n\n# Supported models in Megatron-LM\n# Architecture -> (module, class).\n_MODELS = {\n \"LlamaForCausalLM\": (\n \"llama\",\n (\"ParallelLlamaForCausalLMRmPadPP\", \"ParallelLlamaForValueRmPadPP\", \"ParallelLlamaForCausalLMRmPad\"),\n ),\n \"Qwen2ForCausalLM\": (\n \"qwen2\",\n (\"ParallelQwen2ForCausalLMRmPadPP\", \"ParallelQwen2ForValueRmPadPP\", \"ParallelQwen2ForCausalLMRmPad\"),\n ),\n \"MistralForCausalLM\": (\n \"mistral\",\n (\"ParallelMistralForCausalLMRmPadPP\", \"ParallelMistralForValueRmPadPP\", \"ParallelMistralForCausalLMRmPad\"),\n ),\n \"ApertusForCausalLM\": (\n \"apertus\",\n (\"ParallelApertusForCausalLMRmPadPP\", \"ParallelApertusForValueRmPadPP\", \"ParallelApertusForCausalLMRmPad\"),\n ),\n}\n\n\n# return model class\nclass ModelRegistry:\n @staticmethod\n def load_model_cls(model_arch: str, value=False) -> Optional[type[nn.Module]]:\n if model_arch not in _MODELS:\n return None\n\n megatron = \"megatron\"\n\n module_name, model_cls_name = _MODELS[model_arch]\n if not value: # actor/ref\n model_cls_name = model_cls_name[0]\n elif value: # critic/rm\n model_cls_name = model_cls_name[1]\n\n module = importlib.import_module(f\"verl.models.{module_name}.{megatron}.modeling_{module_name}_megatron\")\n return getattr(module, model_cls_name, None)\n\n @staticmethod\n def get_supported_archs() -> list[str]:\n return list(_MODELS.keys())\n"}39{"file_name": "verl__models__transformers__dense_common.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom dataclasses import dataclass\nfrom typing import Optional, Union\n\nimport torch\nfrom transformers.cache_utils import Cache\nfrom transformers.modeling_outputs import CausalLMOutputWithPast\n\n\n@dataclass\nclass CausalLMOutputForPPO(CausalLMOutputWithPast):\n log_probs: Optional[torch.FloatTensor] = None\n entropy: Optional[torch.FloatTensor] = None\n\n\ndef forward_base_model(\n self,\n input_ids: Optional[torch.LongTensor] = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n past_key_values: Optional[Cache] = None,\n inputs_embeds: Optional[torch.FloatTensor] = None,\n use_cache: Optional[bool] = None,\n output_attentions: Optional[bool] = None,\n output_hidden_states: Optional[bool] = None,\n return_dict: Optional[bool] = None,\n cache_position: Optional[torch.LongTensor] = None,\n) -> CausalLMOutputWithPast:\n r\"\"\"\n Copy paste LLaMa's forward\n https://github.com/linkedin/Liger-Kernel/blob/main/src/liger_kernel/transformers/model/llama.py\n\n This function should be generic enough for all pure text models.\n ```\"\"\"\n\n output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions\n output_hidden_states = (\n output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states\n )\n\n # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)\n outputs = self.model(\n input_ids=input_ids,\n attention_mask=attention_mask,\n position_ids=position_ids,\n past_key_values=past_key_values,\n inputs_embeds=inputs_embeds,\n use_cache=use_cache,\n output_attentions=output_attentions,\n output_hidden_states=output_hidden_states,\n return_dict=return_dict,\n cache_position=cache_position,\n )\n\n return outputs\n\n\ndef forward_with_torch_backend(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n past_key_values: Optional[Union[\"Cache\", list[torch.FloatTensor]]] = None,\n inputs_embeds: Optional[torch.FloatTensor] = None,\n labels: Optional[torch.LongTensor] = None,\n use_cache: Optional[bool] = None,\n output_attentions: Optional[bool] = None,\n output_hidden_states: Optional[bool] = None,\n return_dict: Optional[bool] = None,\n cache_position: Optional[torch.LongTensor] = None,\n logits_to_keep: int | torch.Tensor = 0,\n temperature: float = 1.0,\n **loss_kwargs,\n) -> tuple | CausalLMOutputForPPO:\n from verl.utils.experimental.torch_functional import FusedLinearForPPO\n\n outputs = forward_base_model(\n self,\n input_ids=input_ids,\n attention_mask=attention_mask,\n position_ids=position_ids,\n past_key_values=past_key_values,\n inputs_embeds=inputs_embeds,\n use_cache=use_cache,\n output_attentions=output_attentions,\n output_hidden_states=output_hidden_states,\n cache_position=cache_position,\n )\n\n hidden_states = outputs[0]\n\n if not return_dict:\n raise NotImplementedError(\"forward_with_torch_backend has to return_dict\")\n\n # Loss calculations\n if labels is not None:\n rolled_labels = torch.roll(labels, shifts=-1, dims=-1)\n elif input_ids is not None:\n rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)\n else:\n raise RuntimeError(\"To use forward_with_torch_backend, either labels or input_ids must be provided.\")\n\n fused_linear_for_ppo = FusedLinearForPPO()\n log_probs, entropy = fused_linear_for_ppo.forward(\n hidden_states=hidden_states,\n vocab_weights=self.lm_head.weight,\n input_ids=rolled_labels,\n temperature=temperature,\n )\n\n return CausalLMOutputForPPO(\n log_probs=log_probs,\n entropy=entropy,\n past_key_values=outputs.past_key_values,\n hidden_states=outputs.hidden_states,\n attentions=outputs.attentions,\n )\n\n\ndef forward_with_triton_backend(\n self,\n input_ids: torch.LongTensor = None,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n past_key_values: Optional[Union[\"Cache\", list[torch.FloatTensor]]] = None,\n inputs_embeds: Optional[torch.FloatTensor] = None,\n labels: Optional[torch.LongTensor] = None,\n use_cache: Optional[bool] = None,\n output_attentions: Optional[bool] = None,\n output_hidden_states: Optional[bool] = None,\n return_dict: Optional[bool] = None,\n cache_position: Optional[torch.LongTensor] = None,\n logits_to_keep: int | torch.Tensor = 0,\n temperature: float = 1.0,\n **loss_kwargs,\n) -> tuple | CausalLMOutputForPPO:\n from verl.utils.kernel.linear_cross_entropy import linear_cross_entropy\n\n outputs = forward_base_model(\n self,\n input_ids=input_ids,\n attention_mask=attention_mask,\n position_ids=position_ids,\n past_key_values=past_key_values,\n inputs_embeds=inputs_embeds,\n use_cache=use_cache,\n output_attentions=output_attentions,\n output_hidden_states=output_hidden_states,\n return_dict=return_dict,\n cache_position=cache_position,\n )\n\n hidden_states = outputs[0]\n\n if not return_dict:\n raise NotImplementedError(\"forward_with_triton_backend has to return_dict\")\n\n # Loss calculations\n if labels is not None:\n rolled_labels = torch.roll(labels, shifts=-1, dims=-1)\n elif input_ids is not None:\n rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)\n else:\n raise RuntimeError(\"To use forward_with_triton_backend, either labels or input_ids must be provided.\")\n\n log_probs, entropy = linear_cross_entropy(\n hidden_states,\n self.lm_head.weight,\n rolled_labels,\n temperature,\n \"none\",\n )\n\n return CausalLMOutputForPPO(\n log_probs=log_probs,\n entropy=entropy,\n past_key_values=outputs.past_key_values,\n hidden_states=outputs.hidden_states,\n attentions=outputs.attentions,\n )\n"}40{"file_name": "verl__models__transformers__glm4v.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport inspect\nimport itertools\nimport logging\nimport os\nfrom dataclasses import dataclass\nfrom typing import Optional\n\nimport torch\nimport torch.distributed as dist\nfrom transformers.modeling_flash_attention_utils import _flash_attention_forward, fa_peft_integration_check\nfrom transformers.models.glm4v.modeling_glm4v import (\n Glm4vCausalLMOutputWithPast,\n Glm4vForConditionalGeneration,\n Glm4vTextAttention,\n)\nfrom transformers.utils import is_flash_attn_2_available, is_flash_attn_greater_or_equal_2_10\n\nfrom verl.utils.device import is_npu_available\nfrom verl.utils.ulysses import (\n gather_heads_scatter_seq,\n gather_seq_scatter_heads,\n get_ulysses_sequence_parallel_group,\n get_ulysses_sequence_parallel_world_size,\n validate_ulysses_config,\n)\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\nif is_flash_attn_2_available():\n from flash_attn import flash_attn_func, flash_attn_varlen_func\n\n _flash_supports_window_size = \"window_size\" in inspect.signature(flash_attn_func).parameters\n _flash_supports_deterministic = \"deterministic\" in inspect.signature(flash_attn_func).parameters\n _flash_use_top_left_mask = not is_flash_attn_greater_or_equal_2_10()\n\nif is_npu_available:\n from transformers.integrations.npu_flash_attention import npu_flash_attn_func as flash_attn_func\n from transformers.integrations.npu_flash_attention import npu_flash_attn_varlen_func as flash_attn_varlen_func\n from transformers.modeling_flash_attention_utils import flash_attn_supports_top_left_mask\n\n _flash_supports_window_size = \"window_size\" in inspect.signature(flash_attn_func).parameters\n _flash_supports_deterministic = \"deterministic\" in inspect.signature(flash_attn_func).parameters\n _flash_use_top_left_mask = flash_attn_supports_top_left_mask()\n\n_flash_deterministic_enabled = os.getenv(\"FLASH_ATTENTION_DETERMINISTIC\", \"0\") == \"1\"\n\n\ndef get_rope_index(\n processor,\n input_ids: torch.Tensor,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n attention_mask: Optional[torch.Tensor] = None,\n) -> torch.Tensor:\n \"\"\"\n Gets the position ids for GLM4V in padding-free format.\n The batch dim has been removed and the input_ids should be a 1D tensor representing a single example.\n \"\"\"\n spatial_merge_size = processor.image_processor.merge_size\n image_token_id = processor.tokenizer.convert_tokens_to_ids(\"<|image|>\")\n video_start_token_id = processor.tokenizer.convert_tokens_to_ids(\"<|begin_of_video|>\")\n video_end_token_id = processor.tokenizer.convert_tokens_to_ids(\"<|end_of_video|>\")\n\n if input_ids is not None and (image_grid_thw is not None or video_grid_thw is not None):\n if attention_mask is None:\n attention_mask = torch.ones_like(input_ids)\n\n position_ids = torch.ones(3, input_ids.size(0), dtype=input_ids.dtype, device=input_ids.device) # (3, seqlen)\n image_index, video_index = 0, 0\n video_group_index = 0\n\n input_ids_filtered = input_ids[attention_mask == 1]\n input_tokens = input_ids_filtered.tolist()\n\n input_token_type = []\n video_check_flg = False\n for token in input_tokens:\n if token == video_start_token_id:\n video_check_flg = True\n elif token == video_end_token_id:\n video_check_flg = False\n\n if token == image_token_id and not video_check_flg:\n input_token_type.append(\"image\")\n elif token == image_token_id and video_check_flg:\n input_token_type.append(\"video\")\n else:\n input_token_type.append(\"text\")\n\n input_type_group = []\n for key, group in itertools.groupby(enumerate(input_token_type), lambda x: x[1]):\n group = list(group)\n start_index = group[0][0]\n end_index = group[-1][0] + 1\n input_type_group.append((key, start_index, end_index))\n\n llm_pos_ids_list = []\n video_frame_num = 1\n\n for modality_type, start_idx, end_idx in input_type_group:\n st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0\n\n if modality_type == \"image\":\n t, h, w = (\n image_grid_thw[image_index][0],\n image_grid_thw[image_index][1],\n image_grid_thw[image_index][2],\n )\n llm_grid_t, llm_grid_h, llm_grid_w = (\n t.item(),\n h.item() // spatial_merge_size,\n w.item() // spatial_merge_size,\n )\n\n t_index = torch.arange(llm_grid_t).view(-1, 1).expand(-1, llm_grid_h * llm_grid_w).flatten()\n h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten()\n w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten()\n llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + st_idx)\n\n image_index += 1\n video_frame_num = 1\n\n elif modality_type == \"video\":\n t, h, w = (\n video_frame_num,\n video_grid_thw[video_index][1],\n video_grid_thw[video_index][2],\n )\n\n llm_grid_t, llm_grid_h, llm_grid_w = (\n t,\n h.item() // spatial_merge_size,\n w.item() // spatial_merge_size,\n )\n\n for t_idx in range(llm_grid_t):\n t_index = torch.tensor(t_idx).view(-1, 1).expand(-1, llm_grid_h * llm_grid_w).flatten()\n h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(1, -1, llm_grid_w).flatten()\n w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(1, llm_grid_h, -1).flatten()\n llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + st_idx)\n\n video_group_index += 1\n\n if video_group_index >= video_grid_thw[video_index][0]:\n video_index += 1\n video_group_index = 0\n\n video_frame_num += 1\n\n else:\n text_len = end_idx - start_idx\n llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)\n video_frame_num = 1\n\n llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)\n position_ids[..., attention_mask == 1] = llm_positions.to(position_ids.device)\n else:\n if attention_mask is not None:\n position_ids = attention_mask.long().cumsum(-1) - 1\n position_ids.masked_fill_(attention_mask == 0, 1)\n position_ids = position_ids.unsqueeze(0).expand(3, -1).to(input_ids.device)\n else:\n position_ids = torch.arange(input_ids.shape[0], device=input_ids.device).view(1, -1).expand(3, -1)\n\n return position_ids\n\n\ndef prepare_fa2_from_position_ids(\n query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, position_ids: torch.Tensor\n):\n assert position_ids.ndim == 2 # (batch_size, seq_length)\n query = query.contiguous().view(-1, query.size(-2), query.size(-1))\n key = key.contiguous().view(-1, key.size(-2), key.size(-1))\n value = value.contiguous().view(-1, value.size(-2), value.size(-1))\n position_ids = position_ids.view(-1)\n cu_seqlens = torch.cat(\n (\n (position_ids == 0).nonzero().view(-1).to(torch.int32),\n torch.tensor(position_ids.size(), device=position_ids.device, dtype=torch.int32),\n )\n )\n max_length = cu_seqlens.diff().max() # use cu_seqlens to infer max_length for qwen2vl mrope\n return (query, key, value, (cu_seqlens, cu_seqlens), (max_length, max_length))\n\n\ndef _custom_flash_attention_forward(\n query_states: torch.Tensor,\n key_states: torch.Tensor,\n value_states: torch.Tensor,\n attention_mask: Optional[torch.Tensor],\n query_length: int,\n is_causal: bool = True,\n position_ids: Optional[torch.Tensor] = None,\n use_top_left_mask: bool = False,\n deterministic: Optional[bool] = None,\n **kwargs,\n):\n \"\"\"\n Patches flash attention forward to handle 3D position ids in mrope. (3, batch_size, seq_length)\n \"\"\"\n # Assuming 4D tensors, key_states.shape[1] is the key/value sequence length (source length).\n flash_kwargs = {}\n\n if _flash_supports_deterministic:\n flash_kwargs[\"deterministic\"] = deterministic if deterministic is not None else _flash_deterministic_enabled\n\n if kwargs.get(\"softcap\") is not None:\n flash_kwargs[\"softcap\"] = kwargs.pop(\"softcap\")\n\n query_states, key_states, value_states = fa_peft_integration_check(\n query_states, key_states, value_states, target_dtype=torch.bfloat16\n )\n\n if position_ids is not None:\n assert position_ids.ndim == 2 # (batch_size, seq_length / sp_size)\n\n sp_size = get_ulysses_sequence_parallel_world_size()\n if sp_size > 1:\n # qkv: (batch_size, seq_length / sp_size, num_head, head_size)\n validate_ulysses_config(query_states.size(2), sp_size)\n query_states = gather_seq_scatter_heads(query_states, seq_dim=1, head_dim=2)\n key_states = gather_seq_scatter_heads(key_states, seq_dim=1, head_dim=2)\n value_states = gather_seq_scatter_heads(value_states, seq_dim=1, head_dim=2)\n position_ids_lst = [torch.empty_like(position_ids) for _ in range(sp_size)]\n position_ids = dist.all_gather(position_ids_lst, position_ids, group=get_ulysses_sequence_parallel_group())\n position_ids = torch.cat(position_ids_lst, dim=-1) # (batch_size, seq_length)\n\n if position_ids is not None and query_length != 1 and not (torch.diff(position_ids, dim=-1) >= 0).all():\n batch_size = query_states.size(0)\n q, k, v, (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = prepare_fa2_from_position_ids(\n query_states, key_states, value_states, position_ids\n )\n attn_output = flash_attn_varlen_func(\n q=q,\n k=k,\n v=v,\n cu_seqlens_q=cu_seqlens_q,\n cu_seqlens_k=cu_seqlens_k,\n max_seqlen_q=max_seqlen_q,\n max_seqlen_k=max_seqlen_k,\n dropout_p=kwargs.pop(\"dropout\", 0.0),\n softmax_scale=kwargs.pop(\"softmax_scale\", None),\n causal=is_causal,\n **flash_kwargs,\n )\n attn_output = attn_output.view(batch_size, -1, attn_output.size(-2), attn_output.size(-1))\n else:\n attn_output = _flash_attention_forward(\n query_states,\n key_states,\n value_states,\n attention_mask,\n query_length,\n is_causal=is_causal,\n use_top_left_mask=use_top_left_mask,\n deterministic=deterministic,\n **kwargs,\n ) # do not pass position_ids to old flash_attention_forward\n\n if sp_size > 1:\n # (batch_size, seq_length, num_head, head_size)\n attn_output = gather_heads_scatter_seq(attn_output, head_dim=2, seq_dim=1)\n\n return attn_output\n\n\ndef glm4v_attn_forward(\n self: \"Glm4vTextAttention\",\n hidden_states: torch.Tensor,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46\n **kwargs,\n) -> tuple[torch.Tensor, None, None]:\n from transformers.models.glm4v.modeling_glm4v import apply_multimodal_rotary_pos_emb, repeat_kv\n\n bsz, q_len, _ = hidden_states.size() # q_len = seq_length / sp_size\n query_states = self.q_proj(hidden_states) # (batch_size, seq_length / sp_size, num_heads * head_size)\n key_states = self.k_proj(hidden_states)\n value_states = self.v_proj(hidden_states)\n\n query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)\n key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)\n value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)\n\n # Because the input can be padded, the absolute sequence length depends on the max position id.\n cos, sin = position_embeddings\n query_states, key_states = apply_multimodal_rotary_pos_emb(\n query_states, key_states, cos, sin, self.rope_scaling[\"mrope_section\"]\n )\n key_states = repeat_kv(key_states, self.num_key_value_groups)\n value_states = repeat_kv(value_states, self.num_key_value_groups)\n dropout_rate = 0.0 if not self.training else self.attention_dropout\n\n # This is before the transpose\n q_len = query_states.shape[2]\n\n # FA2 uses non-transposed inputs\n query_states = query_states.transpose(1, 2)\n key_states = key_states.transpose(1, 2)\n value_states = value_states.transpose(1, 2)\n\n attn_output = _custom_flash_attention_forward(\n query_states,\n key_states,\n value_states,\n attention_mask,\n query_length=q_len,\n is_causal=getattr(self, \"is_causal\", True),\n dropout=dropout_rate,\n use_top_left_mask=_flash_use_top_left_mask,\n position_ids=position_ids, # important: pass position ids\n ) # (batch_size, seq_length / sp_size, num_head, head_size)\n attn_output = attn_output.reshape(bsz, q_len, self.hidden_size).contiguous()\n attn_output = self.o_proj(attn_output)\n return attn_output, None\n\n\ndef _get_input_embeds(\n model: \"Glm4vForConditionalGeneration\",\n input_ids: torch.LongTensor,\n attention_mask: Optional[torch.Tensor] = None,\n pixel_values: Optional[torch.FloatTensor] = None,\n pixel_values_videos: Optional[torch.FloatTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n):\n inputs_embeds = model.get_input_embeddings()(input_ids)\n if pixel_values is not None:\n pixel_values = pixel_values.type(model.visual.dtype)\n image_embeds = model.visual(pixel_values, grid_thw=image_grid_thw)\n n_image_tokens = (input_ids == model.config.image_token_id).sum().item()\n n_image_features = image_embeds.shape[0]\n if n_image_tokens != n_image_features:\n raise ValueError(\n f\"Image features and image tokens do not match: tokens: {n_image_tokens}, features {n_image_features}\"\n )\n\n mask = input_ids == model.config.image_token_id\n mask_unsqueezed = mask.unsqueeze(-1)\n mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)\n image_mask = mask_expanded.to(inputs_embeds.device)\n\n image_embeds = image_embeds.to(inputs_embeds.device, inputs_embeds.dtype)\n inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)\n\n if pixel_values_videos is not None:\n pixel_values_videos = pixel_values_videos.type(model.visual.dtype)\n video_embeds = model.visual(pixel_values_videos, grid_thw=video_grid_thw)\n n_video_tokens = (input_ids == model.config.video_token_id).sum().item()\n n_video_features = video_embeds.shape[0]\n if n_video_tokens != n_video_features:\n raise ValueError(\n f\"Video features and video tokens do not match: tokens: {n_video_tokens}, features {n_video_features}\"\n )\n\n mask = input_ids == model.config.video_token_id\n mask_unsqueezed = mask.unsqueeze(-1)\n mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)\n video_mask = mask_expanded.to(inputs_embeds.device)\n\n video_embeds = video_embeds.to(inputs_embeds.device, inputs_embeds.dtype)\n inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)\n\n if pixel_values is None and pixel_values_videos is None: # handle mixed text-image data\n pixel_values = torch.zeros((16, 1176), dtype=inputs_embeds.dtype, device=inputs_embeds.device)\n image_grid_thw = torch.tensor([[1, 4, 4]], dtype=torch.long, device=inputs_embeds.device)\n image_embeds = model.visual(pixel_values, grid_thw=image_grid_thw)\n inputs_embeds += 0.0 * image_embeds.mean()\n\n if attention_mask is not None:\n attention_mask = attention_mask.to(inputs_embeds.device)\n\n return inputs_embeds, attention_mask\n\n\ndef process_position_ids(position_ids: torch.Tensor) -> torch.Tensor:\n if position_ids.ndim != 3 or position_ids.size(0) != 4:\n # we concat the text position ids with the 3D vision position ids by default\n # see https://github.com/huggingface/transformers/pull/39447\n raise ValueError(\"position_ids should be a 3D tensor of shape (4, batch_size, seq_length).\")\n\n return position_ids\n\n\n@dataclass\nclass Glm4vCausalLMOutputForPPO(Glm4vCausalLMOutputWithPast):\n log_probs: Optional[torch.FloatTensor] = None\n entropy: Optional[torch.FloatTensor] = None\n\n\ndef glm4v_base_forward(\n self: \"Glm4vForConditionalGeneration\",\n input_ids: torch.LongTensor,\n attention_mask: Optional[torch.Tensor] = None,\n labels: Optional[torch.LongTensor] = None,\n pixel_values: Optional[torch.FloatTensor] = None,\n pixel_values_videos: Optional[torch.FloatTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n **kwargs,\n):\n kwargs[\"inputs_embeds\"], kwargs[\"attention_mask\"] = _get_input_embeds(\n self, input_ids, attention_mask, pixel_values, pixel_values_videos, image_grid_thw, video_grid_thw\n ) # avoid lora module having multiple keyword arguments\n return self.language_model(\n input_ids=None,\n **kwargs,\n )\n\n\ndef glm4v_forward(\n self: \"Glm4vForConditionalGeneration\",\n input_ids: torch.LongTensor,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n pixel_values: Optional[torch.FloatTensor] = None,\n pixel_values_videos: Optional[torch.FloatTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n **kwargs,\n):\n return self.model(\n input_ids=input_ids,\n attention_mask=attention_mask,\n position_ids=process_position_ids(position_ids),\n pixel_values=pixel_values,\n pixel_values_videos=pixel_values_videos,\n image_grid_thw=image_grid_thw,\n video_grid_thw=video_grid_thw,\n **kwargs,\n )\n\n\ndef forward_with_normal_backend(\n self: Glm4vForConditionalGeneration,\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> \"Glm4vCausalLMOutputWithPast\":\n outputs = glm4v_forward(self, input_ids, **kwargs)\n hidden_states = outputs[0]\n logits = self.lm_head(hidden_states)\n\n return Glm4vCausalLMOutputWithPast(\n logits=logits,\n hidden_states=outputs.hidden_states,\n )\n\n\ndef forward_with_torch_backend(\n self: Glm4vForConditionalGeneration,\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> tuple | Glm4vCausalLMOutputForPPO:\n from verl.utils.experimental.torch_functional import FusedLinearForPPO\n\n outputs = glm4v_forward(self, input_ids, **kwargs)\n hidden_states = outputs[0]\n\n # Loss calculations\n if labels is not None:\n rolled_labels = torch.roll(labels, shifts=-1, dims=-1)\n elif input_ids is not None:\n rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)\n else:\n raise RuntimeError(\"To use forward_with_torch_backend, either labels or input_ids must be provided.\")\n\n fused_linear_for_ppo = FusedLinearForPPO()\n log_probs, entropy = fused_linear_for_ppo.forward(\n hidden_states=hidden_states,\n vocab_weights=self.lm_head.weight,\n input_ids=rolled_labels,\n temperature=temperature,\n )\n return Glm4vCausalLMOutputForPPO(\n log_probs=log_probs,\n entropy=entropy,\n hidden_states=outputs.hidden_states,\n )\n\n\ndef forward_with_triton_backend(\n self: Glm4vForConditionalGeneration,\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> tuple | Glm4vCausalLMOutputForPPO:\n from verl.utils.kernel.linear_cross_entropy import linear_cross_entropy\n\n outputs = glm4v_forward(self, input_ids, **kwargs)\n hidden_states = outputs[0]\n\n # Loss calculations\n if labels is not None:\n rolled_labels = torch.roll(labels, shifts=-1, dims=-1)\n elif input_ids is not None:\n rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)\n else:\n raise RuntimeError(\"To use forward_with_triton_backend, either labels or input_ids must be provided.\")\n\n log_probs, entropy = linear_cross_entropy(\n hidden_states,\n self.lm_head.weight,\n rolled_labels,\n temperature,\n \"none\",\n )\n return Glm4vCausalLMOutputForPPO(\n log_probs=log_probs,\n entropy=entropy,\n hidden_states=outputs.hidden_states,\n )\n"}41{"file_name": "verl__models__transformers__kimi_vl.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Optional\n\nimport torch\nimport torch.nn.functional as F\nfrom transformers.cache_utils import Cache\nfrom transformers.modeling_flash_attention_utils import _flash_attention_forward\n\nfrom verl.models.transformers.monkey_patch import is_transformers_version_in_range\n\n# Import compatibility wrapper for flash_attn_supports_top_left_mask\nfrom verl.utils.transformers_compat import flash_attn_supports_top_left_mask\nfrom verl.utils.ulysses import (\n gather_heads_scatter_seq,\n gather_seq_scatter_heads,\n get_ulysses_sequence_parallel_world_size,\n validate_ulysses_config,\n)\n\n\n# Copied from transformers.models.llama.modeling_llama.rotate_half\ndef rotate_half(x):\n \"\"\"Rotates half the hidden dims of the input.\"\"\"\n x1 = x[..., : x.shape[-1] // 2]\n x2 = x[..., x.shape[-1] // 2 :]\n return torch.cat((-x2, x1), dim=-1)\n\n\n# Copied from transformers.models.llama.modeling_llama.apply_rotary_pos_emb\ndef apply_rotary_pos_emb(q, k, cos, sin, position_ids, unsqueeze_dim=1):\n \"\"\"Applies Rotary Position Embedding to the query and key tensors.\n\n Args:\n q (`torch.Tensor`): The query tensor.\n k (`torch.Tensor`): The key tensor.\n cos (`torch.Tensor`): The cosine part of the rotary embedding.\n sin (`torch.Tensor`): The sine part of the rotary embedding.\n position_ids (`torch.Tensor`):\n The position indices of the tokens corresponding to the query and key tensors. For example, this can be\n used to pass offsetted position ids when working with a KV-cache.\n unsqueeze_dim (`int`, *optional*, defaults to 1):\n The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and\n sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note\n that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and\n k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes\n cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have\n the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.\n Returns:\n `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.\n \"\"\"\n cos = cos[position_ids].unsqueeze(unsqueeze_dim)\n sin = sin[position_ids].unsqueeze(unsqueeze_dim)\n\n b, h, s, d = q.shape\n q = q.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d)\n\n b, h, s, d = k.shape\n k = k.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d)\n\n q_embed = (q * cos) + (rotate_half(q) * sin)\n k_embed = (k * cos) + (rotate_half(k) * sin)\n return q_embed, k_embed\n\n\n# Copied from transformers.models.llama.modeling_llama.repeat_kv\ndef repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:\n \"\"\"\n This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,\n num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)\n \"\"\"\n batch, num_key_value_heads, slen, head_dim = hidden_states.shape\n if n_rep == 1:\n return hidden_states\n hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)\n return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)\n\n\ndef _ulysses_flash_attn_forward(\n self,\n hidden_states: torch.Tensor,\n attention_mask: Optional[torch.LongTensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n past_key_value: Optional[Cache] = None,\n output_attentions: bool = False,\n use_cache: bool = False,\n **kwargs,\n) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]:\n bsz, q_len, _ = hidden_states.size()\n\n if self.q_lora_rank is None:\n q = self.q_proj(hidden_states)\n else:\n q = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states)))\n q = q.view(bsz, q_len, self.num_heads, self.q_head_dim).transpose(1, 2)\n\n # Flash attention requires the input to have the shape\n # batch_size x seq_length x head_dim x hidden_dim\n # therefore we just need to keep the original shape\n compressed_kv = self.kv_a_proj_with_mqa(hidden_states)\n compressed_kv, k_pe = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)\n k_pe = k_pe.view(bsz, q_len, 1, self.qk_rope_head_dim).transpose(1, 2)\n kv = (\n self.kv_b_proj(self.kv_a_layernorm(compressed_kv))\n .view(bsz, q_len, self.num_heads, self.qk_nope_head_dim + self.v_head_dim)\n .transpose(1, 2)\n )\n\n k_nope, value_states = torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)\n\n # patch\n ulysses_sp_size = get_ulysses_sequence_parallel_world_size()\n if ulysses_sp_size > 1:\n validate_ulysses_config(self.num_heads, ulysses_sp_size)\n\n num_key_value_groups = self.config.num_attention_heads // self.config.num_key_value_heads\n k_pe = repeat_kv(k_pe, ulysses_sp_size) # to keep heads=1 after a2a\n k_nope = repeat_kv(k_nope, num_key_value_groups)\n value_states = repeat_kv(value_states, num_key_value_groups)\n q = gather_seq_scatter_heads(q, seq_dim=2, head_dim=1)\n k_pe = gather_seq_scatter_heads(k_pe, seq_dim=2, head_dim=1)\n k_nope = gather_seq_scatter_heads(k_nope, seq_dim=2, head_dim=1)\n value_states = gather_seq_scatter_heads(value_states, seq_dim=2, head_dim=1)\n # (batch_size, num_head / sp_size, seq_length, head_size)\n full_q_len = q.size(2) # full_q_len = seq_length\n\n else:\n full_q_len = q_len\n\n q_nope, q_pe = torch.split(q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)\n cos, sin = self.rotary_emb(value_states, seq_len=full_q_len)\n q_pe, k_pe = apply_rotary_pos_emb(q_pe, k_pe, cos, sin, position_ids)\n\n query_states = k_pe.new_empty(bsz, self.num_heads // ulysses_sp_size, full_q_len, self.q_head_dim)\n query_states[:, :, :, : self.qk_nope_head_dim] = q_nope\n query_states[:, :, :, self.qk_nope_head_dim :] = q_pe\n\n key_states = k_pe.new_empty(bsz, self.num_heads // ulysses_sp_size, full_q_len, self.q_head_dim)\n key_states[:, :, :, : self.qk_nope_head_dim] = k_nope\n key_states[:, :, :, self.qk_nope_head_dim :] = k_pe\n\n if self.q_head_dim != self.v_head_dim:\n value_states = F.pad(value_states, [0, self.q_head_dim - self.v_head_dim])\n\n # TODO: These transpose are quite inefficient but Flash Attention requires the layout\n # [batch_size, sequence_length, num_heads, head_dim]. We would need to refactor the KV cache\n # to be able to avoid many of these transpose/reshape/view.\n query_states = query_states.transpose(1, 2)\n key_states = key_states.transpose(1, 2)\n value_states = value_states.transpose(1, 2)\n\n dropout_rate = self.attention_dropout if self.training else 0.0\n\n attn_output = _flash_attention_forward(\n query_states,\n key_states,\n value_states,\n attention_mask,\n full_q_len,\n dropout=dropout_rate,\n sliding_window=None,\n is_causal=self.is_causal,\n use_top_left_mask=flash_attn_supports_top_left_mask(),\n position_ids=position_ids, # important: pass position ids\n softmax_scale=self.softmax_scale,\n )\n\n if ulysses_sp_size > 1:\n attn_output = gather_heads_scatter_seq(attn_output, head_dim=2, seq_dim=1)\n\n if self.q_head_dim != self.v_head_dim:\n attn_output = attn_output[:, :, :, : self.v_head_dim]\n\n attn_output = attn_output.reshape(bsz, q_len, self.num_heads * self.v_head_dim).contiguous()\n attn_output = self.o_proj(attn_output)\n\n if is_transformers_version_in_range(min_version=\"4.53.0\"):\n return attn_output, None\n else:\n return attn_output, None, None\n"}42{"file_name": "verl__models__transformers__llama.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport sys\nfrom typing import Callable, Optional\n\nimport torch\n\nif sys.version_info >= (3, 11):\n pass\nelse:\n pass\n\nfrom transformers.cache_utils import Cache\nfrom transformers.modeling_flash_attention_utils import _flash_attention_forward\nfrom transformers.models.llama.modeling_llama import apply_rotary_pos_emb\nfrom transformers.utils import logging\n\n# Import compatibility wrapper for flash_attn_supports_top_left_mask\nfrom verl.utils.transformers_compat import flash_attn_supports_top_left_mask\nfrom verl.utils.ulysses import (\n gather_heads_scatter_seq,\n gather_seq_scatter_heads,\n get_ulysses_sequence_parallel_world_size,\n validate_ulysses_config,\n)\n\nlogger = logging.get_logger(__name__)\n\n\ndef llama_flash_attn_forward(\n self,\n hidden_states: torch.Tensor,\n attention_mask: Optional[torch.LongTensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n past_key_value: Optional[Cache] = None,\n output_attentions: bool = False,\n use_cache: bool = False,\n cache_position: Optional[torch.LongTensor] = None,\n position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46\n **kwargs,\n) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]:\n \"\"\"\n Adapted from transformers 4.47.1 to support Ulysses sequence parallelism.\n\n NOTE: This function is used for transformers versions in the range [4.45.0, 4.47.1].\n \"\"\"\n output_attentions = False\n\n bsz, q_len, _ = hidden_states.size()\n\n query_states = self.q_proj(hidden_states)\n key_states = self.k_proj(hidden_states)\n value_states = self.v_proj(hidden_states)\n\n # Flash attention requires the input to have the shape\n # batch_size x seq_length x head_dim x hidden_dim\n # therefore we just need to keep the original shape\n query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)\n key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)\n value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)\n\n # trade off: repeat first and then all to all\n # key_states = repeat_kv(key_states, self.num_key_value_groups)\n # value_states = repeat_kv(value_states, self.num_key_value_groups)\n\n ########## AlltoAll for Ulysses ##########\n ulysses_sp_size = get_ulysses_sequence_parallel_world_size()\n\n if ulysses_sp_size > 1:\n validate_ulysses_config(self.num_heads, ulysses_sp_size)\n\n # (bsz, n_head, seq_len/n, head_dim) -> (bsz, n_head/n, seq_len, head_dim)\n query_states = gather_seq_scatter_heads(query_states, seq_dim=2, head_dim=1)\n key_states = gather_seq_scatter_heads(key_states, seq_dim=2, head_dim=1)\n value_states = gather_seq_scatter_heads(value_states, seq_dim=2, head_dim=1)\n\n full_q_len = query_states.size(2) # full seq length\n\n if position_embeddings is None:\n logger.warning_once(\n \"The attention layers in this model are transitioning from computing the RoPE embeddings internally \"\n \"through `position_ids` (2D tensor with the indexes of the tokens), to using externally computed \"\n \"`position_embeddings` (Tuple of tensors, containing cos and sin). In v4.46 `position_ids` will be \"\n \"removed and `position_embeddings` will be mandatory.\"\n )\n cos, sin = self.rotary_emb(value_states, position_ids)\n else:\n cos, sin = position_embeddings\n query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)\n\n if past_key_value is not None:\n # sin and cos are specific to RoPE models; cache_position needed for the static cache\n cache_kwargs = {\"sin\": sin, \"cos\": cos, \"cache_position\": cache_position}\n key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)\n\n # TODO: These transpose are quite inefficient but Flash Attention requires the layout\n # [batch_size, sequence_length, num_heads, head_dim]. We would need to refactor the KV cache\n # to be able to avoid many of these transpose/reshape/view.\n query_states = query_states.transpose(1, 2)\n key_states = key_states.transpose(1, 2)\n value_states = value_states.transpose(1, 2)\n\n dropout_rate = self.attention_dropout if self.training else 0.0\n\n # In PEFT, usually we cast the layer norms in float32 for training stability reasons\n # therefore the input hidden states gets silently casted in float32. Hence, we need\n # cast them back in the correct dtype just to be sure everything works as expected.\n # This might slowdown training & inference so it is recommended to not cast the LayerNorms\n # in fp32. (LlamaRMSNorm handles it correctly)\n\n input_dtype = query_states.dtype\n if input_dtype == torch.float32:\n if torch.is_autocast_enabled():\n target_dtype = torch.get_autocast_gpu_dtype()\n # Handle the case where the model is quantized\n elif hasattr(self.config, \"_pre_quantization_dtype\"):\n target_dtype = self.config._pre_quantization_dtype\n else:\n target_dtype = self.q_proj.weight.dtype\n\n logger.warning_once(\n f\"The input hidden states seems to be silently casted in float32, this might be related to \"\n f\"the fact you have upcasted embedding or layer norm layers in float32. We will cast back the \"\n f\"input in {target_dtype}.\"\n )\n\n query_states = query_states.to(target_dtype)\n key_states = key_states.to(target_dtype)\n value_states = value_states.to(target_dtype)\n\n attn_output = _flash_attention_forward(\n query_states,\n key_states,\n value_states,\n attention_mask,\n full_q_len,\n position_ids=position_ids,\n dropout=dropout_rate,\n sliding_window=getattr(self, \"sliding_window\", None),\n use_top_left_mask=flash_attn_supports_top_left_mask(),\n is_causal=self.is_causal,\n **kwargs,\n )\n\n attn_output = attn_output.reshape(bsz, full_q_len, -1, self.head_dim).contiguous()\n ########## AlltoAll for Ulysses ##########\n if ulysses_sp_size > 1:\n attn_output = gather_heads_scatter_seq(attn_output, seq_dim=1, head_dim=2)\n attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()\n attn_output = self.o_proj(attn_output)\n\n if not output_attentions:\n attn_weights = None\n\n return attn_output, attn_weights, past_key_value\n\n\ndef llama_attn_forward(\n self,\n hidden_states: torch.Tensor,\n position_embeddings: tuple[torch.Tensor, torch.Tensor],\n attention_mask: Optional[torch.Tensor],\n past_key_value: Optional[Cache] = None,\n cache_position: Optional[torch.LongTensor] = None,\n **kwargs,\n) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]:\n \"\"\"\n Adapted from transformers 4.49.0 to support Ulysses sequence parallelism for transformers >= 4.48.0.\n\n NOTE: This function has been tested only on transformers versions between 4.48.0 and 4.50.0.\n \"\"\"\n from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS\n from transformers.models.llama.modeling_llama import eager_attention_forward\n\n bsz, q_len, _ = hidden_states.shape\n\n query_states = self.q_proj(hidden_states).view(bsz, q_len, -1, self.head_dim).transpose(1, 2)\n key_states = self.k_proj(hidden_states).view(bsz, q_len, -1, self.head_dim).transpose(1, 2)\n value_states = self.v_proj(hidden_states).view(bsz, q_len, -1, self.head_dim).transpose(1, 2)\n\n ########## AlltoAll for Ulysses ##########\n ulysses_sp_size = get_ulysses_sequence_parallel_world_size()\n\n if ulysses_sp_size > 1:\n validate_ulysses_config(self.config.num_attention_heads, ulysses_sp_size)\n\n query_states = gather_seq_scatter_heads(query_states, seq_dim=2, head_dim=1)\n key_states = gather_seq_scatter_heads(key_states, seq_dim=2, head_dim=1)\n value_states = gather_seq_scatter_heads(value_states, seq_dim=2, head_dim=1)\n\n full_q_len = query_states.size(2)\n\n cos, sin = position_embeddings\n query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)\n\n if past_key_value is not None:\n # sin and cos are specific to RoPE models; cache_position needed for the static cache\n cache_kwargs = {\"sin\": sin, \"cos\": cos, \"cache_position\": cache_position}\n key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)\n\n attention_interface: Callable = eager_attention_forward\n if self.config._attn_implementation != \"eager\":\n if self.config._attn_implementation == \"sdpa\" and kwargs.get(\"output_attentions\", False):\n logger.warning_once(\n \"`torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. \"\n \"Falling back to eager attention. This warning can be removed using the argument \"\n '`attn_implementation=\"eager\"` when loading the model.'\n )\n else:\n attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]\n\n attn_output, attn_weights = attention_interface(\n self,\n query_states,\n key_states,\n value_states,\n attention_mask,\n dropout=0.0 if not self.training else self.attention_dropout,\n scaling=self.scaling,\n **kwargs,\n )\n\n attn_output = attn_output.reshape(bsz, full_q_len, -1, self.head_dim).contiguous()\n ########## AlltoAll for Ulysses ##########\n if ulysses_sp_size > 1:\n attn_output = gather_heads_scatter_seq(attn_output, seq_dim=1, head_dim=2)\n attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()\n attn_output = self.o_proj(attn_output)\n return attn_output, attn_weights\n"}43{"file_name": "verl__models__transformers__monkey_patch.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nApply monkey-patch function to models\n\"\"\"\n\nimport sys\nfrom types import SimpleNamespace\nfrom typing import Optional\n\nimport torch\nfrom transformers.modeling_flash_attention_utils import _flash_attention_forward\nfrom transformers.modeling_utils import PreTrainedModel\n\nfrom verl.utils.import_utils import is_trl_available\nfrom verl.utils.transformers_compat import is_transformers_version_in_range\nfrom verl.utils.ulysses import (\n gather_heads_scatter_seq,\n gather_seq_scatter_heads,\n get_ulysses_sequence_parallel_group,\n get_ulysses_sequence_parallel_world_size,\n slice_input_tensor,\n)\n\n_PREFIX_GROUPER_PATCHED = False\n_PREFIX_GROUPER_SUPPORTED_ATTENTIONS = {\"flash_attention_2\", \"flash_attention_3\", \"sdpa\", \"flex_attention\", \"eager\"}\n\n\ndef _create_prefix_grouper_wrapper(original_fn):\n \"\"\"Wrap attention function to support prefix_grouper in kwargs.\"\"\"\n\n def wrapped(module, query, key, value, attention_mask, *args, **kwargs):\n prefix_grouper = kwargs.pop(\"prefix_grouper\", None)\n if prefix_grouper is None:\n return original_fn(module, query, key, value, attention_mask, *args, **kwargs)\n\n def attn_func(q, k, v, attn_mask, *inner_args, **inner_kwargs):\n out, _ = original_fn(module, q, k, v, attn_mask, *inner_args, **inner_kwargs)\n return out\n\n return prefix_grouper.forward(attn_func, query, key, value, *args, **kwargs), None\n\n return wrapped\n\n\ndef apply_prefix_grouper_patch():\n \"\"\"Patch ALL_ATTENTION_FUNCTIONS to support prefix_grouper parameter.\"\"\"\n global _PREFIX_GROUPER_PATCHED\n if _PREFIX_GROUPER_PATCHED:\n return\n\n from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS\n\n patched = []\n for name in list(ALL_ATTENTION_FUNCTIONS.keys()):\n if name in _PREFIX_GROUPER_SUPPORTED_ATTENTIONS:\n ALL_ATTENTION_FUNCTIONS[name] = _create_prefix_grouper_wrapper(ALL_ATTENTION_FUNCTIONS[name])\n patched.append(name)\n\n _PREFIX_GROUPER_PATCHED = True\n print(f\"[PrefixGrouper] Patched: {patched}\")\n\n\ndef repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:\n \"\"\"\n This is the equivalent of torch.repeat_interleave(x, dim=2, repeats=n_rep). The hidden states go from (batch,\n seqlen, num_key_value_heads, head_dim) to (batch, seqlen, num_attention_heads, head_dim)\n \"\"\"\n batch, slen, num_key_value_heads, head_dim = hidden_states.shape\n if n_rep == 1:\n return hidden_states\n hidden_states = hidden_states[:, :, :, None, :].expand(batch, slen, num_key_value_heads, n_rep, head_dim)\n return hidden_states.reshape(batch, slen, num_key_value_heads * n_rep, head_dim)\n\n\ndef _ulysses_flash_attention_forward(\n query_states: torch.Tensor,\n key_states: torch.Tensor,\n value_states: torch.Tensor,\n attention_mask: Optional[torch.Tensor],\n query_length: int,\n *args,\n position_ids: Optional[torch.Tensor] = None,\n **kwargs,\n):\n \"\"\"Insert all-to-all before and after flash attention.\n DeepSpeed-Ulysses: https://arxiv.org/pdf/2309.14509\n\n For transformers>=4.55, the flash attention api has changed,\n we need to pass the query_length after doing ulysses all2all.\n See https://github.com/huggingface/transformers/issues/40399\n\n Args:\n query_states (torch.Tensor): (batch_size, seqlen/sp_size, nheads, head_dim)\n key_states (torch.Tensor): (batch_size, seqlen/sp_size, nheads_k, head_dim)\n value_states (torch.Tensor): (batch_size, seqlen/sp_size, nheads_k, head_dim)\n position_ids (torch.Tensor, optional): (batch_size, seqlen/sp_size)\n\n Returns:\n torch.Tensor: (batch_size, seqlen/sp_size, nheads, head_dim)\n\n \"\"\"\n ulysses_sp_size = get_ulysses_sequence_parallel_world_size()\n\n ########## AlltoAll for Ulysses ##########\n # TODO: Disable sp for ViT, there's no elegent way to determine whether it's ViT or not.\n # Use `position_ids` as condition since ViT doesn't pass it to flash attention.\n if ulysses_sp_size > 1 and position_ids is not None:\n # NOTE: repeat kv heads to be divided by sequence parallel. Instead of repeating nheads_q//nheads_k,\n # we choose to repeat sp_size//nheads_k, since flash_attention supports MQA/GQA.\n # For example:\n # - nheads_k=4, sp=8, repeats=2\n # - nheads_k=8, sp=8, repeats=1\n # - nheads_k=16, sp=8, repeats=1\n repeats = max(ulysses_sp_size // key_states.size(2), 1)\n key_states = repeat_kv(key_states, repeats)\n value_states = repeat_kv(value_states, repeats)\n\n # (bsz, seq_len/n, n_head, head_dim) -> (bsz, seq_len, n_head/n, head_dim)\n query_states = gather_seq_scatter_heads(query_states, seq_dim=1, head_dim=2)\n key_states = gather_seq_scatter_heads(key_states, seq_dim=1, head_dim=2)\n value_states = gather_seq_scatter_heads(value_states, seq_dim=1, head_dim=2)\n\n # TODO: all_gather position_ids because `prepare_fa2_from_position_ids` needs it, we can eliminate\n # this all_gather by passing cu_seq_lens_q, cu_seq_lens_k, max_length_k, max_length_q explicitly.\n # https://github.com/huggingface/transformers/pull/33932\n\n # (bsz, seq_len/n) -> (bsz, seq_len)\n position_ids_list = [torch.empty_like(position_ids) for _ in range(ulysses_sp_size)]\n torch.distributed.all_gather(position_ids_list, position_ids, group=get_ulysses_sequence_parallel_group())\n position_ids = torch.concat(position_ids_list, dim=-1)\n\n # (bsz, seq_len, n_head/n, head_dim)\n query_length = query_states.size(1)\n attn_output = _flash_attention_forward(\n query_states, key_states, value_states, attention_mask, query_length, *args, position_ids=position_ids, **kwargs\n )\n\n ########## AlltoAll for Ulysses ##########\n if ulysses_sp_size > 1 and position_ids is not None:\n # (bsz, seq_len, n_head/n, head_dim) -> (bsz, seq_len/n, n_head, head_dim)\n attn_output = gather_heads_scatter_seq(attn_output, seq_dim=1, head_dim=2)\n\n return attn_output\n\n\ndef patch_vlm_for_ulysses_input_slicing(model_class: type):\n \"\"\"\n Applies a monkey patch to the forward method of a given model class\n to enable Ulysses sequence parallelism input slicing.\n \"\"\"\n\n def _create_ulysses_wrapped_decoder_forward(original_forward):\n def ulysses_wrapped_decoder_forward(self, *args, **kwargs):\n inputs_embeds = kwargs.get(\"inputs_embeds\")\n position_ids = kwargs.get(\"position_ids\")\n visual_pos_masks = kwargs.get(\"visual_pos_masks\")\n deepstack_visual_embeds = kwargs.get(\"deepstack_visual_embeds\")\n call_kwargs = kwargs.copy()\n\n current_ulysses_sp_size = get_ulysses_sequence_parallel_world_size()\n\n slice_now = (\n inputs_embeds is not None\n and current_ulysses_sp_size > 1\n and getattr(self, \"_needs_initial_slice\", True)\n )\n if slice_now:\n call_kwargs[\"inputs_embeds\"] = slice_input_tensor(inputs_embeds, dim=1, padding=False)\n call_kwargs[\"position_ids\"] = slice_input_tensor(position_ids, dim=-1, padding=False)\n # Also slice visual_pos_masks and deepstack_visual_embeds for Qwen3 VL models\n if visual_pos_masks is not None:\n original_visual_mask = visual_pos_masks\n sliced_visual_mask = slice_input_tensor(visual_pos_masks, dim=1, padding=False)\n call_kwargs[\"visual_pos_masks\"] = sliced_visual_mask\n\n if deepstack_visual_embeds is not None:\n sliced_embeds = []\n\n num_visual_before = original_visual_mask.sum().item()\n num_visual_in_shard = sliced_visual_mask.sum().item()\n\n if num_visual_in_shard > 0 and num_visual_before > 0:\n # Calculate which visual embeddings belong to this shard\n # We need to find the offset of visual tokens in this shard\n from verl.utils.ulysses import get_ulysses_sequence_parallel_rank\n\n rank = get_ulysses_sequence_parallel_rank()\n seq_len = original_visual_mask.shape[1]\n local_seq_len = seq_len // current_ulysses_sp_size\n start_idx = rank * local_seq_len\n end_idx = start_idx + local_seq_len\n\n # Get total visual tokens before and up to the end of the shard's sequence slice\n # This correctly handles batches by summing across all samples\n visual_start = original_visual_mask[:, :start_idx].sum().item() if start_idx > 0 else 0\n visual_end = original_visual_mask[:, :end_idx].sum().item()\n\n # Slice each tensor in deepstack_visual_embeds\n for embed in deepstack_visual_embeds:\n sliced_embeds.append(embed[visual_start:visual_end])\n else:\n # No visual tokens in this shard, create empty tensors to maintain gradient flow\n for embed in deepstack_visual_embeds:\n sliced_embeds.append(embed[:0])\n call_kwargs[\"deepstack_visual_embeds\"] = sliced_embeds\n\n self._needs_initial_slice = False\n try:\n return original_forward(self, *args, **call_kwargs)\n finally:\n if slice_now:\n self._needs_initial_slice = True\n\n return ulysses_wrapped_decoder_forward\n\n original_forward = model_class.forward\n wrapped_forward = _create_ulysses_wrapped_decoder_forward(original_forward)\n model_class.forward = wrapped_forward\n print(f\"Monkey patch {model_class.__name__}.forward for Ulysses SP input slicing.\")\n\n\ndef patch_forward_with_backends(\n model: PreTrainedModel,\n use_fused_kernels: bool = False,\n fused_kernels_backend: str = None,\n):\n \"\"\"\n Choose the forward function based on the model and backend.\n Args:\n model (PreTrainedModel): The model to apply the monkey patch.\n use_fused_kernels (bool): Whether to use fused kernels.\n fused_kernels_backend (str): The backend to use for fused kernels.\n \"\"\"\n if not use_fused_kernels or fused_kernels_backend not in [\"triton\", \"torch\"]:\n print(\n f\"Skipping monkey patch for {model.__class__.__name__} as use_fused_kernels is \"\n f\"{use_fused_kernels} or fused_kernels_backend is {fused_kernels_backend}\"\n )\n return\n\n forward_with_torch_backend_function = model.__class__.forward\n forward_with_triton_backend_function = model.__class__.forward\n if model.config.model_type in [\"qwen2_5_vl\", \"qwen2_vl\"]:\n from verl.models.transformers.qwen2_vl import forward_with_torch_backend, forward_with_triton_backend\n\n forward_with_torch_backend_function = forward_with_torch_backend\n forward_with_triton_backend_function = forward_with_triton_backend\n elif model.config.model_type in [\"qwen3_vl\", \"qwen3_vl_moe\"]:\n from verl.models.transformers.qwen3_vl import forward_with_torch_backend, forward_with_triton_backend\n\n forward_with_torch_backend_function = forward_with_torch_backend\n forward_with_triton_backend_function = forward_with_triton_backend\n elif model.config.model_type == \"glm4v\":\n from verl.models.transformers.glm4v import forward_with_torch_backend, forward_with_triton_backend\n\n forward_with_torch_backend_function = forward_with_torch_backend\n forward_with_triton_backend_function = forward_with_triton_backend\n else:\n from verl.models.transformers.dense_common import forward_with_torch_backend, forward_with_triton_backend\n\n forward_with_torch_backend_function = forward_with_torch_backend\n forward_with_triton_backend_function = forward_with_triton_backend\n\n if fused_kernels_backend == \"triton\":\n model.__class__.forward = forward_with_triton_backend_function\n print(f\"Using Triton backend for fused kernels in {model.__class__.__name__}\")\n elif fused_kernels_backend == \"torch\":\n model.__class__.forward = forward_with_torch_backend_function\n print(f\"Using Torch backend for fused kernels in {model.__class__.__name__}\")\n else:\n raise ValueError(f\"Unsupported fused_kernels_backend: {fused_kernels_backend}. Choose 'triton' or 'torch'.\")\n\n\ndef apply_monkey_patch(\n model: PreTrainedModel,\n ulysses_sp_size: int = 1,\n use_remove_padding: bool = True,\n use_fused_kernels: bool = False,\n fused_kernels_backend: str = None,\n use_prefix_grouper: bool = False,\n use_tiled_mlp: bool = False,\n tiled_mlp_shards: int = 4,\n):\n \"\"\"\n Apply monkey patch to the models for ulysses sequence parallel, fused kernel, tiled MLP and prefix grouper.\n\n In the end of this function forward function of the model is patched for fused kernel.\n If the model is not supported with fused kernel, please return after patch.\n\n Args:\n model: The model to apply the monkey patch.\n ulysses_sp_size: The size of ulysses sequence parallel.\n use_remove_padding: Whether to use remove padding.\n use_fused_kernels: Whether to use fused kernels.\n fused_kernels_backend: The backend to use for fused kernels.\n use_tiled_mlp: Whether to use TiledMLP for memory-efficient MLP computation.\n tiled_mlp_shards: Number of shards for TiledMLP (higher = lower memory, slightly slower).\n \"\"\"\n\n # Apply TiledMLP monkey patch for memory-efficient MLP computation\n if use_tiled_mlp:\n from verl.models.transformers.tiled_mlp import apply_tiled_mlp_monkey_patch\n\n model_type = getattr(model.config, \"model_type\", None)\n apply_tiled_mlp_monkey_patch(num_shards=tiled_mlp_shards, model_type=model_type)\n # Apply PrefixGrouper patch if enabled\n if use_prefix_grouper:\n apply_prefix_grouper_patch()\n\n \"\"\"Replace _flash_attention_forward to _ulysses_flash_attention_forward\"\"\"\n module = sys.modules[model.__module__]\n\n try:\n num_attention_heads, num_key_value_heads = model.config.num_attention_heads, model.config.num_key_value_heads\n except AttributeError:\n num_attention_heads, num_key_value_heads = (\n model.config.text_config.num_attention_heads,\n model.config.text_config.num_key_value_heads,\n )\n\n assert num_attention_heads % ulysses_sp_size == 0, (\n f\"num_attention_heads {num_attention_heads} must be divisible by ulysses_sp_size {ulysses_sp_size}\"\n )\n assert num_key_value_heads % ulysses_sp_size == 0 or ulysses_sp_size % num_key_value_heads == 0, (\n f\"num_key_value_heads {num_key_value_heads} must be divisible by ulysses_sp_size \"\n f\"{ulysses_sp_size}or vise versa. Upon ulysses_sp_size % num_key_value_heads == 0,\"\n f\"kv heads are repeated to ensure correctness.\"\n )\n\n if is_trl_available():\n from trl import AutoModelForCausalLMWithValueHead # type: ignore\n\n def state_dict(self, *args, **kwargs):\n return torch.nn.Module.state_dict(self, *args, **kwargs)\n\n AutoModelForCausalLMWithValueHead.state_dict = state_dict\n print(\"Monkey patch state_dict in AutoModelForCausalLMWithValueHead. \")\n\n # TODO: VLM models only, unify monkey patch to LLM models.\n if model.config.model_type in [\"qwen2_5_vl\", \"qwen2_vl\"]:\n # Step 1: patch model to support image-text mixed data\n if is_transformers_version_in_range(min_version=\"4.52.0\"):\n from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (\n Qwen2_5_VLForConditionalGeneration,\n Qwen2_5_VLModel,\n Qwen2_5_VLTextModel,\n )\n from transformers.models.qwen2_vl.modeling_qwen2_vl import (\n Qwen2VLForConditionalGeneration,\n Qwen2VLModel,\n Qwen2VLTextModel,\n )\n else:\n from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VLForConditionalGeneration\n from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VLModel as Qwen2_5_VLTextModel\n from transformers.models.qwen2_vl.modeling_qwen2_vl import Qwen2VLForConditionalGeneration\n from transformers.models.qwen2_vl.modeling_qwen2_vl import Qwen2VLModel as Qwen2VLTextModel\n\n Qwen2_5_VLModel = SimpleNamespace(forward=None)\n Qwen2VLModel = SimpleNamespace(forward=None)\n\n from verl.models.transformers.qwen2_vl import forward_with_normal_backend, qwen2_vl_base_forward\n\n Qwen2_5_VLModel.forward = qwen2_vl_base_forward\n Qwen2VLModel.forward = qwen2_vl_base_forward\n Qwen2_5_VLForConditionalGeneration.forward = forward_with_normal_backend\n Qwen2VLForConditionalGeneration.forward = forward_with_normal_backend\n print(f\"Monkey patch {model.__class__.__name__} model forward\")\n\n # Step 2: patch attention to support ulysses parallelism\n if is_transformers_version_in_range(min_version=\"4.54.0\"):\n from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VLAttention\n from transformers.models.qwen2_vl.modeling_qwen2_vl import Qwen2VLAttention\n elif is_transformers_version_in_range(min_version=\"4.53.0\"):\n raise RuntimeError(\"Transformers 4.53.* is bugged. Use transformers 4.54.0 or later.\")\n else:\n from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (\n Qwen2_5_VLFlashAttention2 as Qwen2_5_VLAttention,\n )\n from transformers.models.qwen2_vl.modeling_qwen2_vl import Qwen2VLFlashAttention2 as Qwen2VLAttention\n\n if use_remove_padding or ulysses_sp_size > 1:\n from verl.models.transformers.qwen2_vl import qwen2_vl_attn_forward\n\n Qwen2_5_VLAttention.forward = qwen2_vl_attn_forward\n Qwen2VLAttention.forward = qwen2_vl_attn_forward\n print(f\"Monkey patch {model.__class__.__name__} attention layer\")\n\n # Step 3: patch input for multimodal sequence parallelism\n if ulysses_sp_size > 1:\n patch_vlm_for_ulysses_input_slicing(Qwen2_5_VLTextModel)\n patch_vlm_for_ulysses_input_slicing(Qwen2VLTextModel)\n\n elif model.config.model_type in [\"qwen3_vl\", \"qwen3_vl_moe\"]:\n # Step 1: patch model to support image-text mixed data\n from transformers.models.qwen3_vl.modeling_qwen3_vl import (\n Qwen3VLForConditionalGeneration,\n Qwen3VLModel,\n Qwen3VLTextModel,\n )\n from transformers.models.qwen3_vl_moe.modeling_qwen3_vl_moe import (\n Qwen3VLMoeForConditionalGeneration,\n Qwen3VLMoeModel,\n Qwen3VLMoeTextModel,\n )\n\n from verl.models.transformers.qwen3_vl import (\n forward_with_normal_backend,\n patch_qwen3_vl_moe_sparse_moe_block_forward,\n qwen3_vl_base_forward,\n )\n\n Qwen3VLModel.forward = qwen3_vl_base_forward\n Qwen3VLMoeModel.forward = qwen3_vl_base_forward\n Qwen3VLForConditionalGeneration.forward = forward_with_normal_backend\n Qwen3VLMoeForConditionalGeneration.forward = forward_with_normal_backend\n print(f\"Monkey patch {model.__class__.__name__} model forward\")\n\n # Step 1.5: patch Qwen3VLMoeTextSparseMoeBlock to fix transformers 4.57.3 bug\n if model.config.model_type == \"qwen3_vl_moe\" and is_transformers_version_in_range(max_version=\"4.57.3\"):\n patch_qwen3_vl_moe_sparse_moe_block_forward()\n\n # Step 2: patch input for multimodal sequence parallelism\n if ulysses_sp_size > 1:\n patch_vlm_for_ulysses_input_slicing(Qwen3VLTextModel)\n patch_vlm_for_ulysses_input_slicing(Qwen3VLMoeTextModel)\n\n elif model.config.model_type == \"glm4v\":\n # Step 1: patch model to support image-text mixed data\n\n from transformers.models.glm4v.modeling_glm4v import (\n Glm4vForConditionalGeneration,\n Glm4vModel,\n Glm4vTextAttention,\n Glm4vTextModel,\n )\n\n from verl.models.transformers.glm4v import forward_with_normal_backend, glm4v_base_forward\n\n Glm4vModel.forward = glm4v_base_forward\n Glm4vForConditionalGeneration.forward = forward_with_normal_backend\n print(f\"Monkey patch {model.__class__.__name__} model forward\")\n\n # Step 2: patch attention to support ulysses parallelism\n if use_remove_padding or ulysses_sp_size > 1:\n from verl.models.transformers.glm4v import glm4v_attn_forward\n\n Glm4vTextAttention.forward = glm4v_attn_forward\n print(f\"Monkey patch {model.__class__.__name__} attention layer\")\n\n # Step 3: patch input for multimodal sequence parallelism\n if ulysses_sp_size > 1:\n patch_vlm_for_ulysses_input_slicing(Glm4vTextModel)\n\n elif model.config.model_type == \"kimi_vl\":\n if use_remove_padding or ulysses_sp_size > 1:\n # TODO: Changes need to be made when transformers are adapted.\n from verl.models.transformers.kimi_vl import _ulysses_flash_attn_forward\n\n module.DeepseekV3FlashAttention2.forward = _ulysses_flash_attn_forward\n print(\"Monkey patch FlashAttention2.forward in KimiVL\")\n\n if ulysses_sp_size > 1:\n patch_vlm_for_ulysses_input_slicing(module.DeepseekV3ForCausalLM)\n\n if use_fused_kernels:\n print(\"Not support fused kernels for KimiVL\")\n\n return\n\n if use_remove_padding or ulysses_sp_size > 1:\n if hasattr(module, \"_flash_attention_forward\"): # transformers <= 4.47.1 or legacy models\n module._flash_attention_forward = _ulysses_flash_attention_forward\n print(f\"Monkey patch _flash_attention_forward in {model.__module__}\")\n else:\n from transformers.integrations import flash_attention\n\n flash_attention._flash_attention_forward = _ulysses_flash_attention_forward\n print(f\"Monkey patch _flash_attention_forward in {flash_attention.__name__}\")\n\n patch_forward_with_backends(model, use_fused_kernels=use_fused_kernels, fused_kernels_backend=fused_kernels_backend)\n"}44{"file_name": "verl__models__transformers__npu_patch.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Copyright 2025 The Qwen Team and The HuggingFace Inc. team\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nimport torch\nimport torch.nn.functional as F\nimport torch_npu\nfrom torch import nn\nfrom transformers.activations import ACT2FN\nfrom transformers.models.qwen2 import modeling_qwen2\nfrom transformers.models.qwen2_5_vl import modeling_qwen2_5_vl\nfrom transformers.models.qwen3 import modeling_qwen3\nfrom transformers.models.qwen3_moe import modeling_qwen3_moe\nfrom transformers.models.qwen3_next import modeling_qwen3_next\nfrom transformers.models.qwen3_vl import modeling_qwen3_vl\nfrom transformers.models.qwen3_vl_moe import modeling_qwen3_vl_moe\nfrom transformers.utils import logging\n\nlogger = logging.get_logger(__name__)\n\n\ndef rms_norm_forward_npu(self, x):\n \"\"\"NPU optimized implementation for RMSNorm.\"\"\"\n if x.dtype != self.weight.dtype:\n x = x.to(self.weight.dtype)\n return torch_npu.npu_rms_norm(x, self.weight, epsilon=self.variance_epsilon)[0]\n\n\ndef silu_forward_npu(self, hidden_state):\n \"\"\"NPU optimized implementation for SiLU in `forward` func in MLP layer.\"\"\"\n gate_up = torch.cat((self.gate_proj(hidden_state), self.up_proj(hidden_state)), dim=-1)\n return self.down_proj(torch_npu.npu_swiglu(gate_up, dim=-1))\n\n\ndef apply_rotary_pos_emb_npu(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):\n \"\"\"NPU optimized implementation for RoPE.\"\"\"\n cos = cos.unsqueeze(unsqueeze_dim)\n sin = sin.unsqueeze(unsqueeze_dim)\n q_embed = torch_npu.npu_rotary_mul(q, cos, sin)\n k_embed = torch_npu.npu_rotary_mul(k, cos, sin)\n return q_embed.to(q.dtype), k_embed.to(k.dtype)\n\n\ndef qwen3_next_rms_norm_forward_npu(self, x):\n return torch_npu.npu_rms_norm(x.float(), 1.0 + self.weight.float(), epsilon=self.eps)[0].type_as(x)\n\n\ndef qwen3_next_rms_norm_forward_gated_npu(self, hidden_states, gate=None):\n input_dtype = hidden_states.dtype\n hidden_states = hidden_states.to(torch.float32)\n hidden_states = torch_npu.npu_rms_norm(hidden_states, self.weight.float(), epsilon=self.variance_epsilon)[0]\n hidden_states = hidden_states * F.silu(gate.to(torch.float32))\n return hidden_states.to(input_dtype)\n\n\ndef qwen3_next_apply_rotary_pos_emb_npu(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):\n cos = cos.unsqueeze(unsqueeze_dim)\n sin = sin.unsqueeze(unsqueeze_dim)\n\n # Keep half or full tensor for later concatenation\n rotary_dim = cos.shape[-1]\n q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]\n k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]\n\n q_embed = torch_npu.npu_rotary_mul(q_rot, cos, sin).to(q.dtype)\n k_embed = torch_npu.npu_rotary_mul(k_rot, cos, sin).to(k.dtype)\n q_embed = torch.cat([q_embed, q_pass], dim=-1)\n k_embed = torch.cat([k_embed, k_pass], dim=-1)\n return q_embed, k_embed\n\n\nclass NPUGmmFunction(torch.autograd.Function):\n @staticmethod\n def forward(ctx, x, weight, group_list, group_list_type=1):\n \"\"\"\n Grouped Matmul(GMM) for Ascend NPU.\n\n Args:\n x (torch.Tensor): Input tensor, shape (tokens_num * top_k, hidden_size)\n weight (torch.Tensor): Expert weights, shape (n_experts, hidden_size, intermediate_size)\n group_list (torch.Tensor): Expert token counts, shape (n_experts,)\n - type 0: cumsum of tokens per expert\n - type 1: direct tokens per expert (default)\n \"\"\"\n ctx.save_for_backward(x, weight)\n ctx.group_list = group_list\n ctx.group_list_type = group_list_type\n\n output = torch_npu.npu_grouped_matmul(\n [x], [weight], bias=None, group_list=group_list, split_item=2, group_type=0, group_list_type=group_list_type\n )[0]\n\n return output\n\n @staticmethod\n def backward(ctx, grad_output):\n x, weight = ctx.saved_tensors\n group_list = ctx.group_list\n group_list_type = ctx.group_list_type\n\n dx = torch_npu.npu_grouped_matmul(\n [grad_output],\n [weight.transpose(1, 2)],\n bias=None,\n group_list=group_list,\n split_item=2,\n group_type=0,\n group_list_type=group_list_type,\n )[0]\n\n dw = torch_npu.npu_grouped_matmul(\n [x.transpose(0, 1)],\n [grad_output],\n bias=None,\n group_list=group_list,\n split_item=3,\n group_type=2,\n group_list_type=group_list_type,\n )[0]\n\n return dx, dw, None, None\n\n\ndef _qwen3_sparse_moe_routed_forward_npu(self, hidden_states: torch.Tensor):\n \"\"\"\n Shared NPU routed-expert path for Qwen3Moe/Qwen3Next sparse MoE blocks.\n\n Returns:\n tuple: (flattened_input, routed_hidden_states, router_logits)\n \"\"\"\n hidden_dim = hidden_states.shape[-1]\n hidden_states = hidden_states.view(-1, hidden_dim)\n # router_logits: (batch * sequence_length, n_experts)\n router_logits = self.gate(hidden_states)\n\n routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float)\n routing_weights, selected_experts = torch.topk(routing_weights, self.top_k, dim=-1)\n if self.norm_topk_prob: # only diff with mixtral sparse moe block!\n routing_weights /= routing_weights.sum(dim=-1, keepdim=True)\n # we cast back to the input dtype\n routing_weights = routing_weights.to(hidden_states.dtype)\n\n # Loop over all available experts in the model and perform the computation on each expert\n # Concat all weights\n input_dtype = hidden_states.dtype\n up_weight_list = [e.up_proj.weight for e in self.experts]\n gate_weight_list = [e.gate_proj.weight for e in self.experts]\n down_weight_list = [e.down_proj.weight for e in self.experts]\n w1 = torch.stack(up_weight_list).transpose(1, 2).to(input_dtype)\n w2 = torch.stack(gate_weight_list).transpose(1, 2).to(input_dtype)\n w3 = torch.stack(down_weight_list).transpose(1, 2).to(input_dtype)\n\n permuted_tokens, row_ids_map = torch_npu.npu_moe_token_permute(hidden_states, selected_experts.to(torch.int32))\n tokens_per_expert = torch.histc(selected_experts, bins=self.num_experts, min=0, max=self.num_experts)\n\n up_res = NPUGmmFunction.apply(permuted_tokens, w1, tokens_per_expert)\n gate_res = NPUGmmFunction.apply(permuted_tokens, w2, tokens_per_expert)\n act_res = torch_npu.npu_swiglu(torch.cat([gate_res, up_res], dim=-1))\n down_res = NPUGmmFunction.apply(act_res, w3, tokens_per_expert)\n\n routed_hidden_states = torch_npu.npu_moe_token_unpermute(down_res, row_ids_map, probs=routing_weights)\n\n return hidden_states, routed_hidden_states, router_logits\n\n\ndef qwen3_moe_sparse_moe_block_forward_npu(self, hidden_states: torch.Tensor) -> torch.Tensor:\n \"\"\"NPU optimized implementation for `forward` in Qwen3MoeSparseMoeBlock.\"\"\"\n output_shape = hidden_states.shape\n _, routed_hidden_states, router_logits = _qwen3_sparse_moe_routed_forward_npu(self, hidden_states)\n final_hidden_states = routed_hidden_states.reshape(output_shape)\n return final_hidden_states, router_logits\n\n\ndef qwen3_next_sparse_moe_block_forward_npu(self, hidden_states: torch.Tensor) -> torch.Tensor:\n \"\"\"NPU optimized implementation for `forward` in Qwen3NextSparseMoeBlock.\"\"\"\n output_shape = hidden_states.shape\n hidden_states, routed_hidden_states, router_logits = _qwen3_sparse_moe_routed_forward_npu(self, hidden_states)\n\n shared_expert_output = self.shared_expert(hidden_states)\n shared_expert_output = torch.sigmoid(self.shared_expert_gate(hidden_states)) * shared_expert_output\n\n final_hidden_states = (routed_hidden_states + shared_expert_output).reshape(output_shape)\n return final_hidden_states, router_logits\n\n\nclass NPUQwen3VLMoeTextExperts(nn.Module):\n \"\"\"NPU optimized implementation for Qwen3VLMoeTextExperts.\"\"\"\n\n def __init__(self, config):\n super().__init__()\n self.num_experts = config.num_experts\n self.intermediate_size = config.moe_intermediate_size\n self.hidden_size = config.hidden_size\n self.expert_dim = self.intermediate_size\n self.gate_up_proj = nn.Parameter(torch.empty(self.num_experts, self.hidden_size, 2 * self.expert_dim))\n self.down_proj = nn.Parameter(torch.empty((self.num_experts, self.expert_dim, self.hidden_size)))\n self.act_fn = ACT2FN[config.hidden_act]\n\n def forward(\n self, hidden_states: torch.Tensor, routing_weights: torch.Tensor, router_indices: torch.Tensor\n ) -> torch.Tensor:\n \"\"\"\n When training it is more efficient to just loop over the experts and compute the output for each expert\n as otherwise the memory would explode.\n\n For inference we can sacrifice some memory and compute the output for all experts at once.\n By repeating the inputs.\n\n Args:\n hidden_states (torch.Tensor): (batch_size * token_num, hidden_size)\n routing_weights (torch.Tensor): (batch_size * token_num, num_experts)\n router_indices (torch.Tensor): (batch_size * token_num, top_k)\n Returns:\n torch.Tensor\n \"\"\"\n batch_size = hidden_states.shape[0]\n hidden_states = hidden_states.reshape(-1, self.hidden_size) # (num_tokens, hidden_size)\n if self.training:\n permuted_hidden_states, row_ids_map = torch_npu.npu_moe_token_permute(\n hidden_states, router_indices.to(torch.int32)\n )\n tokens_per_expert = torch.histc(router_indices, bins=self.num_experts, min=0, max=self.num_experts)\n intermediate_hidden_states = NPUGmmFunction.apply(\n permuted_hidden_states, self.gate_up_proj, tokens_per_expert\n )\n intermediate_activations = torch_npu.npu_swiglu(intermediate_hidden_states, dim=-1)\n output = NPUGmmFunction.apply(intermediate_activations, self.down_proj, tokens_per_expert)\n num_tokens = hidden_states.shape[0]\n top_k = router_indices.shape[1]\n batch_idx = torch.arange(num_tokens, device=routing_weights.device)\n batch_idx = batch_idx.unsqueeze(1).expand(-1, top_k)\n selected_probs = routing_weights[batch_idx, router_indices]\n next_states = torch_npu.npu_moe_token_unpermute(output, row_ids_map, probs=selected_probs)\n next_states = next_states.view(batch_size, -1, self.hidden_size)\n else:\n hidden_states = hidden_states.repeat(self.num_experts, 1)\n hidden_states = hidden_states.view(self.num_experts, -1, self.hidden_size)\n gate_up = torch.bmm(hidden_states, self.gate_up_proj)\n gate, up = gate_up.chunk(2, dim=-1) # not supported for DTensors\n next_states = torch.bmm((up * self.act_fn(gate)), self.down_proj)\n next_states = next_states.reshape(self.num_experts, batch_size, -1, self.hidden_size)\n next_states = (\n next_states * routing_weights.transpose(0, 1).view(self.num_experts, batch_size, -1)[..., None]\n )\n next_states = next_states.sum(dim=0)\n return next_states\n\n\nclass NPUQwen3VLMoeTextSparseMoeBlock(nn.Module):\n \"\"\"NPU optimized implementation for Qwen3VLMoeTextSparseMoeBlock.\"\"\"\n\n def __init__(self, config):\n super().__init__()\n self.hidden_size = config.hidden_size\n self.num_experts = config.num_experts\n self.top_k = config.num_experts_per_tok\n self.gate = nn.Linear(config.hidden_size, config.num_experts, bias=False)\n self.experts = NPUQwen3VLMoeTextExperts(config)\n\n def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:\n batch_size = hidden_states.shape[0]\n hidden_states = hidden_states.reshape(-1, self.hidden_size)\n router_logits = self.gate(hidden_states)\n routing_weights = torch.nn.functional.softmax(router_logits, dim=-1, dtype=torch.float)\n routing_weights, router_indices = torch.topk(routing_weights, self.top_k, dim=-1)\n routing_weights = routing_weights / routing_weights.sum(dim=-1, keepdim=True)\n routing_weights = routing_weights.to(router_logits.dtype)\n hidden_states = hidden_states.reshape(batch_size, -1, self.hidden_size)\n if not self.training:\n routing_weights = torch.zeros_like(router_logits).scatter_(1, router_indices, routing_weights)\n routed_out = self.experts(hidden_states, routing_weights, router_indices)\n return routed_out\n\n\n# Patches for Qwen2 Model\nmodeling_qwen2.Qwen2RMSNorm.forward = rms_norm_forward_npu\nmodeling_qwen2.Qwen2MLP.forward = silu_forward_npu\nmodeling_qwen2.apply_rotary_pos_emb = apply_rotary_pos_emb_npu\n\n# Patches for Qwen2.5-VL Model\nmodeling_qwen2_5_vl.Qwen2RMSNorm.forward = rms_norm_forward_npu\nmodeling_qwen2_5_vl.Qwen2_5_VLMLP.forward = silu_forward_npu\n\n# Patches for Qwen3 Model\nmodeling_qwen3.Qwen3RMSNorm.forward = rms_norm_forward_npu\nmodeling_qwen3.Qwen3MLP.forward = silu_forward_npu\nmodeling_qwen3.apply_rotary_pos_emb = apply_rotary_pos_emb_npu\n\n# Patches for Qwen3 MoE Model\nmodeling_qwen3_moe.Qwen3MoeRMSNorm.forward = rms_norm_forward_npu\nmodeling_qwen3_moe.Qwen3MoeSparseMoeBlock.forward = qwen3_moe_sparse_moe_block_forward_npu\nmodeling_qwen3_moe.apply_rotary_pos_emb = apply_rotary_pos_emb_npu\n\n# Patches for Qwen3 VL Model\nmodeling_qwen3_vl.Qwen3VLTextRMSNorm.forward = rms_norm_forward_npu\nmodeling_qwen3_vl.Qwen3VLTextMLP.forward = silu_forward_npu\n\n# Patches for Qwen3-VL MoE Model\nmodeling_qwen3_vl_moe.Qwen3VLMoeTextSparseMoeBlock = NPUQwen3VLMoeTextSparseMoeBlock\nmodeling_qwen3_vl_moe.Qwen3VLMoeTextRMSNorm.forward = rms_norm_forward_npu\nmodeling_qwen3_vl_moe.apply_rotary_pos_emb = apply_rotary_pos_emb_npu\n\n# Patches for Qwen3 Next Model\nmodeling_qwen3_next.Qwen3NextSparseMoeBlock.forward = qwen3_next_sparse_moe_block_forward_npu\nmodeling_qwen3_next.Qwen3NextRMSNormGated.forward = qwen3_next_rms_norm_forward_gated_npu\nmodeling_qwen3_next.Qwen3NextRMSNorm.forward = qwen3_next_rms_norm_forward_npu\nmodeling_qwen3_next.apply_rotary_pos_emb = qwen3_next_apply_rotary_pos_emb_npu\n"}45{"file_name": "verl__models__transformers__qwen2.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Callable, Optional\n\nimport torch\nfrom transformers.cache_utils import Cache\nfrom transformers.modeling_flash_attention_utils import _flash_attention_forward\nfrom transformers.models.llama.modeling_llama import apply_rotary_pos_emb, repeat_kv\nfrom transformers.utils import logging\n\n# Import compatibility wrapper for flash_attn_supports_top_left_mask\nfrom verl.utils.transformers_compat import flash_attn_supports_top_left_mask\nfrom verl.utils.ulysses import (\n gather_heads_scatter_seq,\n gather_seq_scatter_heads,\n get_ulysses_sequence_parallel_world_size,\n validate_ulysses_config,\n)\n\nlogger = logging.get_logger(__name__)\n\n\ndef qwen2_flash_attn_forward(\n self,\n hidden_states: torch.Tensor,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n past_key_value: Optional[Cache] = None,\n output_attentions: bool = False,\n use_cache: bool = False,\n cache_position: Optional[torch.LongTensor] = None,\n position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46\n):\n \"\"\"\n Adapted from transformers 4.47.1 to support Ulysses sequence parallelism.\n\n NOTE: This function is only tested on transformers versions between 4.45.0 and 4.47.1.\n \"\"\"\n bsz, q_len, _ = hidden_states.size()\n\n query_states = self.q_proj(hidden_states)\n key_states = self.k_proj(hidden_states)\n value_states = self.v_proj(hidden_states)\n\n query_states = query_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)\n key_states = key_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)\n value_states = value_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)\n\n ########## AlltoAll for Ulysses ##########\n ulysses_sp_size = get_ulysses_sequence_parallel_world_size()\n\n if ulysses_sp_size > 1:\n validate_ulysses_config(self.num_heads, ulysses_sp_size)\n\n # (bsz, n_head, seq_len/n, head_dim) -> (bsz, n_head/n, seq_len, head_dim)\n query_states = gather_seq_scatter_heads(query_states, seq_dim=2, head_dim=1)\n key_states = gather_seq_scatter_heads(key_states, seq_dim=2, head_dim=1)\n value_states = gather_seq_scatter_heads(value_states, seq_dim=2, head_dim=1)\n\n full_q_len = query_states.size(2) # full seq length\n\n if position_embeddings is None:\n logger.warning_once(\n \"The attention layers in this model are transitioning from computing the RoPE embeddings internally \"\n \"through `position_ids` (2D tensor with the indexes of the tokens), to using externally computed \"\n \"`position_embeddings` (Tuple of tensors, containing cos and sin). In v4.46 `position_ids` will be \"\n \"removed and `position_embeddings` will be mandatory.\"\n )\n cos, sin = self.rotary_emb(value_states, position_ids)\n else:\n cos, sin = position_embeddings\n query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)\n\n if past_key_value is not None:\n cache_kwargs = {\"sin\": sin, \"cos\": cos, \"cache_position\": cache_position} # Specific to RoPE models\n key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)\n\n # repeat k/v heads if n_kv_heads < n_heads\n key_states = repeat_kv(key_states, self.num_key_value_groups)\n value_states = repeat_kv(value_states, self.num_key_value_groups)\n dropout_rate = 0.0 if not self.training else self.attention_dropout\n\n # In PEFT, usually we cast the layer norms in float32 for training stability reasons\n # therefore the input hidden states gets silently casted in float32. Hence, we need\n # cast them back in float16 just to be sure everything works as expected.\n input_dtype = query_states.dtype\n if input_dtype == torch.float32:\n if torch.is_autocast_enabled():\n target_dtype = torch.get_autocast_gpu_dtype()\n # Handle the case where the model is quantized\n elif hasattr(self.config, \"_pre_quantization_dtype\"):\n target_dtype = self.config._pre_quantization_dtype\n else:\n target_dtype = self.q_proj.weight.dtype\n\n logger.warning_once(\n f\"The input hidden states seems to be silently casted in float32, this might be related to \"\n f\"the fact you have upcasted embedding or layer norm layers in float32. We will cast back the \"\n f\"input in {target_dtype}.\"\n )\n\n query_states = query_states.to(target_dtype)\n key_states = key_states.to(target_dtype)\n value_states = value_states.to(target_dtype)\n\n # Reashape to the expected shape for Flash Attention\n query_states = query_states.transpose(1, 2)\n key_states = key_states.transpose(1, 2)\n value_states = value_states.transpose(1, 2)\n\n if (\n self.config.use_sliding_window\n and getattr(self.config, \"sliding_window\", None) is not None\n and self.layer_idx >= self.config.max_window_layers\n ):\n sliding_window = self.config.sliding_window\n else:\n sliding_window = None\n\n attn_output = _flash_attention_forward(\n query_states,\n key_states,\n value_states,\n attention_mask,\n full_q_len,\n position_ids=position_ids,\n dropout=dropout_rate,\n sliding_window=sliding_window,\n is_causal=self.is_causal,\n use_top_left_mask=flash_attn_supports_top_left_mask(),\n )\n\n # use full_q_len to reshape\n attn_output = attn_output.reshape(bsz, full_q_len, -1, self.head_dim).contiguous()\n ########## AlltoAll for Ulysses ##########\n if ulysses_sp_size > 1:\n attn_output = gather_heads_scatter_seq(attn_output, seq_dim=1, head_dim=2)\n attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()\n attn_output = self.o_proj(attn_output)\n\n if not output_attentions:\n attn_weights = None\n\n return attn_output, attn_weights, past_key_value\n\n\ndef qwen2_attn_forward(\n self,\n hidden_states: torch.Tensor,\n position_embeddings: tuple[torch.Tensor, torch.Tensor],\n attention_mask: Optional[torch.Tensor],\n past_key_value: Optional[Cache] = None,\n cache_position: Optional[torch.LongTensor] = None,\n **kwargs,\n) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]:\n \"\"\"\n Adapted from transformers 4.49.0 to support Ulysses sequence parallelism for transformers >= 4.48.0.\n\n NOTE: This function has been tested only on transformers versions between 4.48.0 and 4.50.0.\n \"\"\"\n from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS\n\n bsz, q_len, _ = hidden_states.shape\n hidden_shape = (bsz, q_len, -1, self.head_dim)\n\n query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)\n key_states = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)\n value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)\n\n ########## AlltoAll for Ulysses ##########\n ulysses_sp_size = get_ulysses_sequence_parallel_world_size()\n\n if ulysses_sp_size > 1:\n validate_ulysses_config(self.config.num_attention_heads, ulysses_sp_size)\n\n # (bsz, n_head, seq_len/n, head_dim) -> (bsz, n_head/n, seq_len, head_dim)\n query_states = gather_seq_scatter_heads(query_states, seq_dim=2, head_dim=1)\n key_states = gather_seq_scatter_heads(key_states, seq_dim=2, head_dim=1)\n value_states = gather_seq_scatter_heads(value_states, seq_dim=2, head_dim=1)\n\n full_q_len = query_states.size(2)\n\n cos, sin = position_embeddings\n query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)\n\n if past_key_value is not None:\n # sin and cos are specific to RoPE models; cache_position needed for the static cache\n cache_kwargs = {\"sin\": sin, \"cos\": cos, \"cache_position\": cache_position}\n key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)\n\n sliding_window = None\n if (\n self.config.use_sliding_window\n and getattr(self.config, \"sliding_window\", None) is not None\n and self.layer_idx >= self.config.max_window_layers\n ):\n sliding_window = self.config.sliding_window\n\n from transformers.models.qwen2.modeling_qwen2 import eager_attention_forward\n\n attention_interface: Callable = eager_attention_forward\n if self.config._attn_implementation != \"eager\":\n if self.config._attn_implementation == \"sdpa\" and kwargs.get(\"output_attentions\", False):\n logger.warning_once(\n \"`torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. \"\n \"Falling back to eager attention. This warning can be removed using the argument \"\n '`attn_implementation=\"eager\"` when loading the model.'\n )\n else:\n attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]\n\n attn_output, attn_weights = attention_interface(\n self,\n query_states,\n key_states,\n value_states,\n attention_mask,\n dropout=0.0 if not self.training else self.attention_dropout,\n scaling=self.scaling,\n sliding_window=sliding_window, # main diff with Llama\n **kwargs,\n )\n\n attn_output = attn_output.reshape(bsz, full_q_len, -1, self.head_dim).contiguous()\n ########## AlltoAll for Ulysses ##########\n if ulysses_sp_size > 1:\n # (bsz, seq_len, n_head/n, head_dim) -> (bsz, seq_len/n, n_head, head_dim)\n attn_output = gather_heads_scatter_seq(attn_output, seq_dim=1, head_dim=2)\n attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()\n attn_output = self.o_proj(attn_output)\n return attn_output, attn_weights\n"}46{"file_name": "verl__models__transformers__qwen2_vl.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport inspect\nimport logging\nimport os\nfrom dataclasses import dataclass\nfrom typing import Optional\n\nimport torch\nimport torch.distributed as dist\nfrom transformers.modeling_flash_attention_utils import _flash_attention_forward, fa_peft_integration_check\nfrom transformers.models.qwen2_vl.modeling_qwen2_vl import (\n Qwen2VLAttention,\n Qwen2VLCausalLMOutputWithPast,\n Qwen2VLForConditionalGeneration,\n)\nfrom transformers.utils import is_flash_attn_2_available, is_flash_attn_greater_or_equal_2_10\n\nfrom verl.utils.device import is_npu_available\nfrom verl.utils.transformers_compat import is_transformers_version_in_range\nfrom verl.utils.ulysses import (\n gather_heads_scatter_seq,\n gather_seq_scatter_heads,\n get_ulysses_sequence_parallel_group,\n get_ulysses_sequence_parallel_world_size,\n validate_ulysses_config,\n)\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\nif is_flash_attn_2_available():\n from flash_attn import flash_attn_func, flash_attn_varlen_func\n\n _flash_supports_window_size = \"window_size\" in inspect.signature(flash_attn_func).parameters\n _flash_supports_deterministic = \"deterministic\" in inspect.signature(flash_attn_func).parameters\n _flash_use_top_left_mask = not is_flash_attn_greater_or_equal_2_10()\n\nif is_npu_available:\n from transformers.integrations.npu_flash_attention import npu_flash_attn_func as flash_attn_func\n from transformers.integrations.npu_flash_attention import npu_flash_attn_varlen_func as flash_attn_varlen_func\n from transformers.modeling_flash_attention_utils import flash_attn_supports_top_left_mask\n\n _flash_supports_window_size = \"window_size\" in inspect.signature(flash_attn_func).parameters\n _flash_supports_deterministic = \"deterministic\" in inspect.signature(flash_attn_func).parameters\n _flash_use_top_left_mask = flash_attn_supports_top_left_mask()\n\n_flash_deterministic_enabled = os.getenv(\"FLASH_ATTENTION_DETERMINISTIC\", \"0\") == \"1\"\n\n\ndef get_rope_index(\n processor,\n input_ids: torch.Tensor,\n image_grid_thw: Optional[torch.Tensor] = None,\n video_grid_thw: Optional[torch.Tensor] = None,\n second_per_grid_ts: Optional[torch.Tensor] = None,\n attention_mask: Optional[torch.Tensor] = None,\n) -> torch.Tensor:\n \"\"\"\n Gets the position ids for Qwen2-VL, it should be generated before sharding the sequence.\n The batch dim has been removed and the input_ids should be a 1D tensor representing a single example.\n https://github.com/huggingface/transformers/blob/v4.52.4/src/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py#L1405\n \"\"\"\n spatial_merge_size = processor.image_processor.merge_size\n tokens_per_second = 2\n image_token_id = processor.tokenizer.convert_tokens_to_ids(\"<|image_pad|>\")\n video_token_id = processor.tokenizer.convert_tokens_to_ids(\"<|video_pad|>\")\n vision_start_token_id = processor.tokenizer.convert_tokens_to_ids(\"<|vision_start|>\")\n if input_ids is not None and (image_grid_thw is not None or video_grid_thw is not None):\n if attention_mask is None:\n attention_mask = torch.ones_like(input_ids)\n\n position_ids = torch.ones(3, input_ids.size(0), dtype=input_ids.dtype, device=input_ids.device) # (3, seqlen)\n image_index, video_index = 0, 0\n input_ids = input_ids[attention_mask == 1]\n image_nums, video_nums = 0, 0\n vision_start_indices = torch.argwhere(input_ids == vision_start_token_id)\n vision_tokens = input_ids[vision_start_indices + 1]\n image_nums = (vision_tokens == image_token_id).sum()\n video_nums = (vision_tokens == video_token_id).sum()\n input_tokens = input_ids.tolist()\n llm_pos_ids_list: list = []\n st = 0\n remain_images, remain_videos = image_nums, video_nums\n for _ in range(image_nums + video_nums):\n if image_token_id in input_tokens and remain_images > 0:\n ed_image = input_tokens.index(image_token_id, st)\n else:\n ed_image = len(input_tokens) + 1\n if video_token_id in input_tokens and remain_videos > 0:\n ed_video = input_tokens.index(video_token_id, st)\n else:\n ed_video = len(input_tokens) + 1\n if ed_image < ed_video:\n t, h, w = (\n image_grid_thw[image_index][0],\n image_grid_thw[image_index][1],\n image_grid_thw[image_index][2],\n )\n second_per_grid_t = 0\n image_index += 1\n remain_images -= 1\n ed = ed_image\n else:\n t, h, w = (\n video_grid_thw[video_index][0],\n video_grid_thw[video_index][1],\n video_grid_thw[video_index][2],\n )\n second_per_grid_t = second_per_grid_ts[video_index] if second_per_grid_ts is not None else 1.0\n\n video_index += 1\n remain_videos -= 1\n ed = ed_video\n\n llm_grid_t, llm_grid_h, llm_grid_w = (\n t.item(),\n h.item() // spatial_merge_size,\n w.item() // spatial_merge_size,\n )\n text_len = ed - st\n\n st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0\n llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)\n\n t_index = torch.arange(llm_grid_t).view(-1, 1).expand(-1, llm_grid_h * llm_grid_w)\n t_index = (t_index * second_per_grid_t * tokens_per_second).long().flatten()\n h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten()\n w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten()\n llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + text_len + st_idx)\n st = ed + llm_grid_t * llm_grid_h * llm_grid_w\n\n if st < len(input_tokens):\n st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0\n text_len = len(input_tokens) - st\n llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)\n\n llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)\n position_ids[..., attention_mask == 1] = llm_positions.to(position_ids.device)\n else:\n if attention_mask is not None:\n position_ids = attention_mask.long().cumsum(-1) - 1\n position_ids.masked_fill_(attention_mask == 0, 1)\n position_ids = position_ids.unsqueeze(0).expand(3, -1).to(input_ids.device)\n else:\n position_ids = torch.arange(input_ids.shape[1], device=input_ids.device).view(1, -1).expand(3, -1)\n\n return position_ids\n\n\ndef prepare_fa2_from_position_ids(\n query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, position_ids: torch.Tensor\n):\n assert position_ids.ndim == 2 # (batch_size, seq_length)\n query = query.contiguous().view(-1, query.size(-2), query.size(-1))\n key = key.contiguous().view(-1, key.size(-2), key.size(-1))\n value = value.contiguous().view(-1, value.size(-2), value.size(-1))\n position_ids = position_ids.view(-1)\n cu_seqlens = torch.cat(\n (\n (position_ids == 0).nonzero().view(-1).to(torch.int32),\n torch.tensor(position_ids.size(), device=position_ids.device, dtype=torch.int32),\n )\n )\n max_length = cu_seqlens.diff().max() # use cu_seqlens to infer max_length for qwen2vl mrope\n return (query, key, value, (cu_seqlens, cu_seqlens), (max_length, max_length))\n\n\ndef _custom_flash_attention_forward(\n query_states: torch.Tensor,\n key_states: torch.Tensor,\n value_states: torch.Tensor,\n attention_mask: Optional[torch.Tensor],\n query_length: int,\n is_causal: bool = True,\n position_ids: Optional[torch.Tensor] = None,\n sliding_window: Optional[int] = None,\n use_top_left_mask: bool = False,\n deterministic: Optional[bool] = None,\n **kwargs,\n):\n \"\"\"\n Patches flash attention forward to handle 3D position ids in mrope. (3, batch_size, seq_length)\n \"\"\"\n # Assuming 4D tensors, key_states.shape[1] is the key/value sequence length (source length).\n use_sliding_windows = (\n _flash_supports_window_size and sliding_window is not None and key_states.shape[1] > sliding_window\n )\n flash_kwargs = {\"window_size\": (sliding_window, sliding_window)} if use_sliding_windows else {}\n\n if _flash_supports_deterministic:\n flash_kwargs[\"deterministic\"] = deterministic if deterministic is not None else _flash_deterministic_enabled\n\n if kwargs.get(\"softcap\") is not None:\n flash_kwargs[\"softcap\"] = kwargs.pop(\"softcap\")\n\n query_states, key_states, value_states = fa_peft_integration_check(\n query_states, key_states, value_states, target_dtype=torch.bfloat16\n )\n\n if position_ids is not None:\n assert position_ids.ndim == 2 # (batch_size, seq_length / sp_size)\n\n sp_size = get_ulysses_sequence_parallel_world_size()\n if sp_size > 1:\n # qkv: (batch_size, seq_length / sp_size, num_head, head_size)\n validate_ulysses_config(query_states.size(2), sp_size)\n query_states = gather_seq_scatter_heads(query_states, seq_dim=1, head_dim=2)\n key_states = gather_seq_scatter_heads(key_states, seq_dim=1, head_dim=2)\n value_states = gather_seq_scatter_heads(value_states, seq_dim=1, head_dim=2)\n position_ids_lst = [torch.empty_like(position_ids) for _ in range(sp_size)]\n position_ids = dist.all_gather(position_ids_lst, position_ids, group=get_ulysses_sequence_parallel_group())\n position_ids = torch.cat(position_ids_lst, dim=-1) # (batch_size, seq_length)\n\n if position_ids is not None and query_length != 1 and not (torch.diff(position_ids, dim=-1) >= 0).all():\n batch_size = query_states.size(0)\n q, k, v, (cu_seqlens_q, cu_seqlens_k), (max_seqlen_q, max_seqlen_k) = prepare_fa2_from_position_ids(\n query_states, key_states, value_states, position_ids\n )\n attn_output = flash_attn_varlen_func(\n q=q,\n k=k,\n v=v,\n cu_seqlens_q=cu_seqlens_q,\n cu_seqlens_k=cu_seqlens_k,\n max_seqlen_q=max_seqlen_q,\n max_seqlen_k=max_seqlen_k,\n dropout_p=kwargs.pop(\"dropout\", 0.0),\n softmax_scale=kwargs.pop(\"softmax_scale\", None),\n causal=is_causal,\n **flash_kwargs,\n )\n attn_output = attn_output.view(batch_size, -1, attn_output.size(-2), attn_output.size(-1))\n else:\n attn_output = _flash_attention_forward(\n query_states,\n key_states,\n value_states,\n attention_mask,\n query_length,\n is_causal=is_causal,\n sliding_window=sliding_window,\n use_top_left_mask=use_top_left_mask,\n deterministic=deterministic,\n **kwargs,\n ) # do not pass position_ids to old flash_attention_forward\n\n if sp_size > 1:\n # (batch_size, seq_length, num_head, head_size)\n attn_output = gather_heads_scatter_seq(attn_output, head_dim=2, seq_dim=1)\n\n return attn_output\n\n\ndef qwen2_vl_attn_forward(\n self: \"Qwen2VLAttention\",\n hidden_states: torch.Tensor,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46\n **kwargs,\n) -> tuple[torch.Tensor, None, None]:\n from transformers.models.qwen2_vl.modeling_qwen2_vl import apply_multimodal_rotary_pos_emb, repeat_kv\n\n bsz, q_len, _ = hidden_states.size() # q_len = seq_length / sp_size\n query_states = self.q_proj(hidden_states) # (batch_size, seq_length / sp_size, num_heads * head_size)\n key_states = self.k_proj(hidden_states)\n value_states = self.v_proj(hidden_states)\n\n query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)\n key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)\n value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)\n\n # Because the input can be padded, the absolute sequence length depends on the max position id.\n cos, sin = position_embeddings\n query_states, key_states = apply_multimodal_rotary_pos_emb(\n query_states, key_states, cos, sin, self.rope_scaling[\"mrope_section\"]\n )\n key_states = repeat_kv(key_states, self.num_key_value_groups)\n value_states = repeat_kv(value_states, self.num_key_value_groups)\n dropout_rate = 0.0 if not self.training else self.attention_dropout\n\n sliding_window = None\n if (\n self.config.use_sliding_window\n and getattr(self.config, \"sliding_window\", None) is not None\n and self.layer_idx >= self.config.max_window_layers\n ):\n sliding_window = self.config.sliding_window\n\n # This is before the transpose\n q_len = query_states.shape[2]\n\n # FA2 uses non-transposed inputs\n query_states = query_states.transpose(1, 2)\n key_states = key_states.transpose(1, 2)\n value_states = value_states.transpose(1, 2)\n\n if position_ids.ndim == 3:\n position_ids = position_ids[0]\n\n attn_output = _custom_flash_attention_forward(\n query_states,\n key_states,\n value_states,\n attention_mask,\n query_length=q_len,\n is_causal=getattr(self, \"is_causal\", True),\n dropout=dropout_rate,\n sliding_window=sliding_window,\n use_top_left_mask=_flash_use_top_left_mask,\n position_ids=position_ids, # important: pass position ids\n ) # (batch_size, seq_length / sp_size, num_head, head_size)\n attn_output = attn_output.reshape(bsz, q_len, self.hidden_size).contiguous()\n attn_output = self.o_proj(attn_output)\n if is_transformers_version_in_range(min_version=\"4.54.0\"):\n return attn_output, None\n else:\n return attn_output, None, None\n\n\ndef _get_input_embeds(\n model: \"Qwen2VLForConditionalGeneration\",\n input_ids: torch.LongTensor,\n attention_mask: Optional[torch.Tensor] = None,\n pixel_values: Optional[torch.FloatTensor] = None,\n pixel_values_videos: Optional[torch.FloatTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n):\n inputs_embeds = model.get_input_embeddings()(input_ids)\n if pixel_values is not None:\n pixel_values = pixel_values.type(model.visual.dtype)\n image_embeds = model.visual(pixel_values, grid_thw=image_grid_thw)\n n_image_tokens = (input_ids == model.config.image_token_id).sum().item()\n n_image_features = image_embeds.shape[0]\n if n_image_tokens != n_image_features:\n raise ValueError(\n f\"Image features and image tokens do not match: tokens: {n_image_tokens}, features {n_image_features}\"\n )\n\n mask = input_ids == model.config.image_token_id\n mask_unsqueezed = mask.unsqueeze(-1)\n mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)\n image_mask = mask_expanded.to(inputs_embeds.device)\n\n image_embeds = image_embeds.to(inputs_embeds.device, inputs_embeds.dtype)\n inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)\n\n if pixel_values_videos is not None:\n pixel_values_videos = pixel_values_videos.type(model.visual.dtype)\n video_embeds = model.visual(pixel_values_videos, grid_thw=video_grid_thw)\n n_video_tokens = (input_ids == model.config.video_token_id).sum().item()\n n_video_features = video_embeds.shape[0]\n if n_video_tokens != n_video_features:\n raise ValueError(\n f\"Video features and video tokens do not match: tokens: {n_video_tokens}, features {n_video_features}\"\n )\n\n mask = input_ids == model.config.video_token_id\n mask_unsqueezed = mask.unsqueeze(-1)\n mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)\n video_mask = mask_expanded.to(inputs_embeds.device)\n\n video_embeds = video_embeds.to(inputs_embeds.device, inputs_embeds.dtype)\n inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)\n\n if pixel_values is None and pixel_values_videos is None: # handle mixed text-image data\n config = model.config.vision_config\n patch_dim = config.in_channels * config.temporal_patch_size * config.patch_size**2\n pixel_values = torch.zeros((16, patch_dim), dtype=inputs_embeds.dtype, device=inputs_embeds.device)\n image_grid_thw = torch.tensor([[1, 4, 4]], dtype=torch.long, device=inputs_embeds.device)\n image_embeds = model.visual(pixel_values, grid_thw=image_grid_thw)\n inputs_embeds += 0.0 * image_embeds.mean()\n\n if attention_mask is not None:\n attention_mask = attention_mask.to(inputs_embeds.device)\n\n return inputs_embeds, attention_mask\n\n\ndef process_position_ids(position_ids: torch.Tensor) -> torch.Tensor:\n if position_ids.ndim != 3 or position_ids.size(0) != 4:\n # we concat the text position ids with the 3D vision position ids by default\n # see https://github.com/huggingface/transformers/pull/39447\n raise ValueError(\"position_ids should be a 3D tensor of shape (4, batch_size, seq_length).\")\n\n if is_transformers_version_in_range(max_version=\"4.53.3\"):\n # transformers < 4.54.0 only accepts vision position ids, so we discard the text position ids here\n position_ids = position_ids[1:]\n\n return position_ids\n\n\n@dataclass\nclass Qwen2VLCausalLMOutputForPPO(Qwen2VLCausalLMOutputWithPast):\n log_probs: Optional[torch.FloatTensor] = None\n entropy: Optional[torch.FloatTensor] = None\n\n\ndef qwen2_vl_base_forward(\n self: \"Qwen2VLForConditionalGeneration\",\n input_ids: torch.LongTensor,\n attention_mask: Optional[torch.Tensor] = None,\n labels: Optional[torch.LongTensor] = None,\n pixel_values: Optional[torch.FloatTensor] = None,\n pixel_values_videos: Optional[torch.FloatTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n **kwargs,\n):\n kwargs[\"inputs_embeds\"], kwargs[\"attention_mask\"] = _get_input_embeds(\n self, input_ids, attention_mask, pixel_values, pixel_values_videos, image_grid_thw, video_grid_thw\n ) # avoid lora module having multiple keyword arguments\n return self.language_model(input_ids=None, **kwargs)\n\n\ndef qwen2_vl_forward(\n self: \"Qwen2VLForConditionalGeneration\",\n input_ids: torch.LongTensor,\n attention_mask: Optional[torch.Tensor] = None,\n position_ids: Optional[torch.LongTensor] = None,\n pixel_values: Optional[torch.FloatTensor] = None,\n pixel_values_videos: Optional[torch.FloatTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n **kwargs,\n):\n if is_transformers_version_in_range(min_version=\"4.52.0\"):\n return self.model(\n input_ids=input_ids,\n attention_mask=attention_mask,\n position_ids=process_position_ids(position_ids),\n pixel_values=pixel_values,\n pixel_values_videos=pixel_values_videos,\n image_grid_thw=image_grid_thw,\n video_grid_thw=video_grid_thw,\n **kwargs,\n )\n else:\n inputs_embeds, attention_mask = _get_input_embeds(\n self, input_ids, attention_mask, pixel_values, pixel_values_videos, image_grid_thw, video_grid_thw\n )\n return self.model(\n input_ids=None,\n attention_mask=attention_mask,\n position_ids=process_position_ids(position_ids),\n inputs_embeds=inputs_embeds,\n **kwargs,\n )\n\n\ndef forward_with_normal_backend(\n self: Qwen2VLForConditionalGeneration,\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> \"Qwen2VLCausalLMOutputWithPast\":\n outputs = qwen2_vl_forward(self, input_ids, **kwargs)\n hidden_states = outputs[0]\n logits = self.lm_head(hidden_states)\n\n return Qwen2VLCausalLMOutputWithPast(\n logits=logits,\n hidden_states=outputs.hidden_states,\n )\n\n\ndef forward_with_torch_backend(\n self: Qwen2VLForConditionalGeneration,\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> tuple | Qwen2VLCausalLMOutputForPPO:\n from verl.utils.experimental.torch_functional import FusedLinearForPPO\n\n outputs = qwen2_vl_forward(self, input_ids, **kwargs)\n hidden_states = outputs[0]\n\n # Loss calculations\n if labels is not None:\n rolled_labels = torch.roll(labels, shifts=-1, dims=-1)\n elif input_ids is not None:\n rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)\n else:\n raise RuntimeError(\"To use forward_with_torch_backend, either labels or input_ids must be provided.\")\n\n fused_linear_for_ppo = FusedLinearForPPO()\n log_probs, entropy = fused_linear_for_ppo.forward(\n hidden_states=hidden_states,\n vocab_weights=self.lm_head.weight,\n input_ids=rolled_labels,\n temperature=temperature,\n )\n return Qwen2VLCausalLMOutputForPPO(\n log_probs=log_probs,\n entropy=entropy,\n hidden_states=outputs.hidden_states,\n )\n\n\ndef forward_with_triton_backend(\n self: Qwen2VLForConditionalGeneration,\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> tuple | Qwen2VLCausalLMOutputForPPO:\n from verl.utils.kernel.linear_cross_entropy import linear_cross_entropy\n\n outputs = qwen2_vl_forward(self, input_ids, **kwargs)\n hidden_states = outputs[0]\n\n # Loss calculations\n if labels is not None:\n rolled_labels = torch.roll(labels, shifts=-1, dims=-1)\n elif input_ids is not None:\n rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)\n else:\n raise RuntimeError(\"To use forward_with_triton_backend, either labels or input_ids must be provided.\")\n\n log_probs, entropy = linear_cross_entropy(\n hidden_states,\n self.lm_head.weight,\n rolled_labels,\n temperature,\n \"none\",\n )\n return Qwen2VLCausalLMOutputForPPO(\n log_probs=log_probs,\n entropy=entropy,\n hidden_states=outputs.hidden_states,\n )\n"}47{"file_name": "verl__models__transformers__qwen3_vl.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport functools\nimport logging\nimport os\nfrom dataclasses import dataclass\nfrom typing import Optional\n\nimport torch\nfrom transformers.models.qwen3_vl.modeling_qwen3_vl import (\n Qwen3VLCausalLMOutputWithPast,\n Qwen3VLForConditionalGeneration,\n)\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef get_rope_index(\n processor,\n input_ids: torch.Tensor,\n image_grid_thw: Optional[torch.Tensor] = None,\n video_grid_thw: Optional[torch.Tensor] = None,\n attention_mask: Optional[torch.Tensor] = None,\n **kwargs,\n) -> torch.Tensor:\n \"\"\"\n Gets the position ids for Qwen3-VL, it should be generated before sharding the sequence.\n The batch dim has been removed and the input_ids should be a 1D tensor representing a single example.\n https://github.com/huggingface/transformers/blob/v4.57.0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L916\n \"\"\"\n spatial_merge_size = processor.image_processor.merge_size\n image_token_id = processor.image_token_id\n video_token_id = processor.video_token_id\n vision_start_token_id = processor.vision_start_token_id\n\n # Since we use timestamps to separate videos,\n # like <t1> <vision_start> <frame1> <vision_end> <t2> <vision_start> <frame2> <vision_end>,\n # the video_grid_thw should also be split\n if video_grid_thw is not None:\n video_grid_thw = torch.repeat_interleave(video_grid_thw, video_grid_thw[:, 0], dim=0)\n video_grid_thw[:, 0] = 1\n\n if input_ids is not None and (image_grid_thw is not None or video_grid_thw is not None):\n if attention_mask is None:\n attention_mask = torch.ones_like(input_ids)\n\n position_ids = torch.ones(3, input_ids.shape[0], dtype=input_ids.dtype, device=input_ids.device)\n image_index, video_index = 0, 0\n attention_mask = attention_mask.to(input_ids.device)\n input_ids = input_ids[attention_mask == 1]\n image_nums, video_nums = 0, 0\n vision_start_indices = torch.argwhere(input_ids == vision_start_token_id)\n vision_tokens = input_ids[vision_start_indices + 1]\n image_nums = (vision_tokens == image_token_id).sum()\n video_nums = (vision_tokens == video_token_id).sum()\n input_tokens = input_ids.tolist()\n llm_pos_ids_list: list = []\n st = 0\n remain_images, remain_videos = image_nums, video_nums\n for _ in range(image_nums + video_nums):\n if image_token_id in input_tokens and remain_images > 0:\n ed_image = input_tokens.index(image_token_id, st)\n else:\n ed_image = len(input_tokens) + 1\n if video_token_id in input_tokens and remain_videos > 0:\n ed_video = input_tokens.index(video_token_id, st)\n else:\n ed_video = len(input_tokens) + 1\n if ed_image < ed_video:\n t, h, w = (\n image_grid_thw[image_index][0],\n image_grid_thw[image_index][1],\n image_grid_thw[image_index][2],\n )\n image_index += 1\n remain_images -= 1\n ed = ed_image\n else:\n t, h, w = (\n video_grid_thw[video_index][0],\n video_grid_thw[video_index][1],\n video_grid_thw[video_index][2],\n )\n video_index += 1\n remain_videos -= 1\n ed = ed_video\n\n llm_grid_t, llm_grid_h, llm_grid_w = (\n t.item(),\n h.item() // spatial_merge_size,\n w.item() // spatial_merge_size,\n )\n text_len = ed - st\n\n st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0\n llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)\n\n # t_index is always 0 because llm_grid_t is always 1\n # (we use timestamps to encode the temporal information for videos)\n t_index = torch.arange(llm_grid_t).view(-1, 1).expand(-1, llm_grid_h * llm_grid_w).flatten()\n h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten()\n w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten()\n llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + text_len + st_idx)\n st = ed + llm_grid_t * llm_grid_h * llm_grid_w\n\n if st < len(input_tokens):\n st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0\n text_len = len(input_tokens) - st\n llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)\n\n llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)\n position_ids[..., attention_mask == 1] = llm_positions.to(position_ids.device)\n else:\n if attention_mask is not None:\n position_ids = attention_mask.long().cumsum(-1) - 1\n position_ids.masked_fill_(attention_mask == 0, 1)\n position_ids = position_ids.unsqueeze(0).expand(3, -1).to(attention_mask.device)\n else:\n position_ids = torch.arange(input_ids.shape[1], device=input_ids.device).view(1, -1).expand(3, -1)\n\n return position_ids\n\n\ndef _get_input_embeds(\n model: \"Qwen3VLForConditionalGeneration\",\n input_ids: torch.LongTensor,\n attention_mask: Optional[torch.Tensor] = None,\n pixel_values: Optional[torch.FloatTensor] = None,\n pixel_values_videos: Optional[torch.FloatTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n):\n inputs_embeds = model.get_input_embeddings()(input_ids)\n image_mask, video_mask = None, None\n if pixel_values is not None:\n pixel_values = pixel_values.type(model.visual.dtype)\n image_embeds, deepstack_image_embeds = model.visual(pixel_values, grid_thw=image_grid_thw)\n n_image_tokens = (input_ids == model.config.image_token_id).sum().item()\n n_image_features = image_embeds.shape[0]\n if n_image_tokens != n_image_features:\n raise ValueError(\n f\"Image features and image tokens do not match: tokens: {n_image_tokens}, features {n_image_features}\"\n )\n\n mask = input_ids == model.config.image_token_id\n mask_unsqueezed = mask.unsqueeze(-1)\n mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)\n image_mask = mask_expanded.to(inputs_embeds.device)\n\n image_embeds = image_embeds.to(inputs_embeds.device, inputs_embeds.dtype)\n inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)\n\n if pixel_values_videos is not None:\n pixel_values_videos = pixel_values_videos.type(model.visual.dtype)\n video_embeds, deepstack_video_embeds = model.visual(pixel_values_videos, grid_thw=video_grid_thw)\n n_video_tokens = (input_ids == model.config.video_token_id).sum().item()\n n_video_features = video_embeds.shape[0]\n if n_video_tokens != n_video_features:\n raise ValueError(\n f\"Video features and video tokens do not match: tokens: {n_video_tokens}, features {n_video_features}\"\n )\n\n mask = input_ids == model.config.video_token_id\n mask_unsqueezed = mask.unsqueeze(-1)\n mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)\n video_mask = mask_expanded.to(inputs_embeds.device)\n\n video_embeds = video_embeds.to(inputs_embeds.device, inputs_embeds.dtype)\n inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)\n\n visual_pos_masks = None\n deepstack_visual_embeds = None\n if image_mask is not None and video_mask is not None:\n # aggregate visual_pos_masks and deepstack_visual_embeds\n image_mask = image_mask[..., 0]\n video_mask = video_mask[..., 0]\n visual_pos_masks = image_mask | video_mask\n deepstack_visual_embeds = []\n image_mask_joint = image_mask[visual_pos_masks]\n video_mask_joint = video_mask[visual_pos_masks]\n for img_embed, vid_embed in zip(deepstack_image_embeds, deepstack_video_embeds, strict=False):\n embed_joint = img_embed.new_zeros(visual_pos_masks.sum(), img_embed.shape[-1]).to(img_embed.device)\n embed_joint[image_mask_joint, :] = img_embed\n embed_joint[video_mask_joint, :] = vid_embed\n deepstack_visual_embeds.append(embed_joint)\n elif image_mask is not None:\n image_mask = image_mask[..., 0]\n visual_pos_masks = image_mask\n deepstack_visual_embeds = deepstack_image_embeds\n elif video_mask is not None:\n video_mask = video_mask[..., 0]\n visual_pos_masks = video_mask\n deepstack_visual_embeds = deepstack_video_embeds\n\n if pixel_values is None and pixel_values_videos is None:\n config = model.config.vision_config\n patch_dim = config.in_channels * config.temporal_patch_size * config.patch_size**2\n pixel_values = torch.zeros((16, patch_dim), dtype=inputs_embeds.dtype, device=inputs_embeds.device)\n image_grid_thw = torch.tensor([[1, 4, 4]], dtype=torch.long, device=inputs_embeds.device)\n image_embeds, dummy_deepstack_image_embeds = model.visual(pixel_values, grid_thw=image_grid_thw)\n inputs_embeds += 0.0 * image_embeds.mean()\n for emb in dummy_deepstack_image_embeds or []:\n inputs_embeds += 0.0 * emb.mean()\n\n if attention_mask is not None:\n attention_mask = attention_mask.to(inputs_embeds.device)\n\n return {\n \"inputs_embeds\": inputs_embeds,\n \"attention_mask\": attention_mask,\n \"visual_pos_masks\": visual_pos_masks,\n \"deepstack_visual_embeds\": deepstack_visual_embeds,\n }\n\n\n@dataclass\nclass Qwen3VLCausalLMOutputForPPO(Qwen3VLCausalLMOutputWithPast):\n log_probs: Optional[torch.FloatTensor] = None\n entropy: Optional[torch.FloatTensor] = None\n\n\ndef qwen3_vl_base_forward(\n self: \"Qwen3VLForConditionalGeneration\",\n input_ids: torch.LongTensor,\n attention_mask: Optional[torch.Tensor] = None,\n pixel_values: Optional[torch.FloatTensor] = None,\n pixel_values_videos: Optional[torch.FloatTensor] = None,\n image_grid_thw: Optional[torch.LongTensor] = None,\n video_grid_thw: Optional[torch.LongTensor] = None,\n **kwargs,\n):\n input_kwargs = _get_input_embeds(\n self, input_ids, attention_mask, pixel_values, pixel_values_videos, image_grid_thw, video_grid_thw\n ) # avoid lora module having multiple keyword arguments\n kwargs.update(input_kwargs)\n return self.language_model(\n input_ids=None,\n **kwargs,\n )\n\n\ndef forward_with_normal_backend(\n self: \"Qwen3VLForConditionalGeneration\",\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> \"Qwen3VLCausalLMOutputForPPO\":\n outputs = self.model(input_ids, **kwargs)\n hidden_states = outputs[0]\n logits = self.lm_head(hidden_states)\n\n return Qwen3VLCausalLMOutputForPPO(\n logits=logits,\n hidden_states=outputs.hidden_states,\n )\n\n\ndef forward_with_torch_backend(\n self: \"Qwen3VLForConditionalGeneration\",\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> \"Qwen3VLCausalLMOutputForPPO\":\n from verl.utils.experimental.torch_functional import FusedLinearForPPO\n\n outputs = self.model(input_ids, **kwargs)\n hidden_states = outputs[0]\n\n # Loss calculations\n if labels is not None:\n rolled_labels = torch.roll(labels, shifts=-1, dims=-1)\n elif input_ids is not None:\n rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)\n else:\n raise RuntimeError(\"To use forward_with_torch_backend, either labels or input_ids must be provided.\")\n\n fused_linear_for_ppo = FusedLinearForPPO()\n log_probs, entropy = fused_linear_for_ppo.forward(\n hidden_states=hidden_states,\n vocab_weights=self.lm_head.weight,\n input_ids=rolled_labels,\n temperature=temperature,\n )\n return Qwen3VLCausalLMOutputForPPO(\n log_probs=log_probs,\n entropy=entropy,\n hidden_states=outputs.hidden_states,\n )\n\n\ndef forward_with_triton_backend(\n self: \"Qwen3VLForConditionalGeneration\",\n input_ids: torch.LongTensor = None,\n labels: Optional[torch.LongTensor] = None,\n temperature: float = 1.0,\n **kwargs,\n) -> \"Qwen3VLCausalLMOutputForPPO\":\n from verl.utils.kernel.linear_cross_entropy import linear_cross_entropy\n\n outputs = self.model(input_ids, **kwargs)\n hidden_states = outputs[0]\n\n # Loss calculations\n if labels is not None:\n rolled_labels = torch.roll(labels, shifts=-1, dims=-1)\n elif input_ids is not None:\n rolled_labels = torch.roll(input_ids, shifts=-1, dims=-1)\n else:\n raise RuntimeError(\"To use forward_with_triton_backend, either labels or input_ids must be provided.\")\n\n log_probs, entropy = linear_cross_entropy(\n hidden_states,\n self.lm_head.weight,\n rolled_labels,\n temperature,\n \"none\",\n )\n return Qwen3VLCausalLMOutputForPPO(\n log_probs=log_probs,\n entropy=entropy,\n hidden_states=outputs.hidden_states,\n )\n\n\ndef patch_qwen3_vl_moe_sparse_moe_block_forward():\n \"\"\"\n Monkey patch to fix a bug in transformers 4.57.3 where Qwen3VLMoeTextSparseMoeBlock.forward\n incorrectly uses torch.zeros_like(hidden_states) instead of torch.zeros_like(router_logits)\n when creating router_weights (line 148 in modeling_qwen3_vl_moe.py).\n\n This is a minimal fix that only changes the problematic line while keeping the rest of the\n original implementation intact.\n \"\"\"\n try:\n from transformers.models.qwen3_vl_moe.modeling_qwen3_vl_moe import Qwen3VLMoeTextSparseMoeBlock\n except ImportError:\n # Model not available, skip patching\n return\n\n # Store the original forward method for reference\n original_forward = Qwen3VLMoeTextSparseMoeBlock.forward\n\n @functools.wraps(original_forward)\n def patched_forward(self, hidden_states: torch.Tensor) -> torch.Tensor:\n batch_size = hidden_states.shape[0]\n hidden_states = hidden_states.reshape(-1, self.hidden_size)\n router_logits = self.gate(hidden_states)\n routing_weights = torch.nn.functional.softmax(router_logits, dim=-1, dtype=torch.float)\n routing_weights, router_indices = torch.topk(routing_weights, self.top_k, dim=-1)\n routing_weights = routing_weights / routing_weights.sum(dim=-1, keepdim=True)\n # BUG FIX: Original code incorrectly uses hidden_states here, should use router_logits\n routing_weights = routing_weights.to(router_logits.dtype)\n router_weights = torch.zeros_like(router_logits).scatter_(1, router_indices, routing_weights)\n hidden_states = hidden_states.reshape(batch_size, -1, self.hidden_size)\n routed_out = self.experts(hidden_states, router_weights, router_indices)\n return routed_out\n\n # Apply the patch\n Qwen3VLMoeTextSparseMoeBlock.forward = patched_forward\n logger.info(\"Monkey patched Qwen3VLMoeTextSparseMoeBlock.forward to fix router_weights bug\")\n"}48{"file_name": "verl__models__transformers__tiled_mlp.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nFSDP2-compatible TiledMLP implementation for memory-efficient MLP computation.\n\nThis module provides a tiled MLP implementation that reduces peak memory usage\nby processing the MLP forward/backward pass in chunks (tiles). This is particularly\nuseful for large models with FSDP2 training.\n\"\"\"\n\nimport threading\nfrom typing import Optional\n\nimport torch\nimport torch.nn as nn\n\n\nclass GradientAccumulator:\n \"\"\"Gradient accumulator for TiledMLP (FSDP compatible).\n\n This class manages gradient accumulation across multiple shards during\n the backward pass of TiledMLP. It ensures correct gradient computation\n when processing input in chunks.\n \"\"\"\n\n def __init__(self, params: list[torch.nn.Parameter], total_shards: int, dtype: torch.dtype = None):\n self.params = params\n self.total_shards = total_shards\n self.grad_accumulation_dtype = dtype or torch.float32\n self.accumulated_grads = {}\n self.hooks = []\n self.lock = threading.Lock()\n\n for param in self.params:\n if param.grad is not None:\n self.accumulated_grads[param] = param.grad.to(self.grad_accumulation_dtype)\n param.grad = None\n else:\n self.accumulated_grads[param] = torch.zeros_like(param, dtype=self.grad_accumulation_dtype)\n\n def install_hooks(self, is_last_shard: bool):\n \"\"\"Install gradient hooks for the current shard.\"\"\"\n self._remove_hooks()\n\n def create_hook(param):\n def hook(grad):\n with self.lock:\n grad_to_accum_dtype = grad.to(self.grad_accumulation_dtype)\n self.accumulated_grads[param] += grad_to_accum_dtype\n\n if is_last_shard:\n param.grad = None # Critical: prevent double accumulation\n final_grad = self.accumulated_grads[param].to(param.dtype)\n return final_grad\n return None\n\n return hook\n\n for param in self.params:\n if param.requires_grad:\n hook = param.register_hook(create_hook(param))\n self.hooks.append(hook)\n\n def _remove_hooks(self):\n \"\"\"Remove all registered hooks.\"\"\"\n for hook in self.hooks:\n hook.remove()\n self.hooks.clear()\n\n def cleanup(self):\n \"\"\"Cleanup hooks and resources.\"\"\"\n self._remove_hooks()\n\n\nclass TiledMLP(torch.autograd.Function):\n \"\"\"TiledMLP implementation for memory-efficient MLP computation.\n\n This autograd function processes MLP forward/backward in tiles (chunks)\n to reduce peak memory usage. Compatible with FSDP2.\n \"\"\"\n\n @staticmethod\n def forward(ctx, fn, module, x, shards, compute_params):\n ctx.fn = fn\n ctx.module = module\n ctx.shards = shards\n ctx.compute_params = [p for p in compute_params if p.requires_grad]\n ctx.save_for_backward(x)\n\n # Split on dim=-2 (seqlen dimension) following Liger Kernel style\n x_shards = list(torch.chunk(x, chunks=shards, dim=-2))\n with torch.no_grad():\n output_shards = [fn(module, x_shard) for x_shard in x_shards]\n output_unsharded = torch.cat(output_shards, dim=-2)\n return output_unsharded\n\n @staticmethod\n def backward(ctx, *grads):\n fn = ctx.fn\n (x,) = ctx.saved_tensors\n module = ctx.module\n shards = ctx.shards\n compute_params = ctx.compute_params\n\n x_requires_grad = x.requires_grad\n x = x.detach()\n x.requires_grad_(x_requires_grad)\n\n # Flatten to [bs*seqlen, hidden_size]\n hidden_size = x.shape[-1]\n x_shape_orig = x.shape\n x = x.view(-1, hidden_size)\n incoming_grad = grads[0].view(-1, hidden_size)\n\n # Pre-allocate input gradient\n x_grad = torch.zeros_like(x)\n\n # Split on dim=0\n x_shards = list(torch.chunk(x, chunks=shards, dim=0))\n\n grad_accumulator = GradientAccumulator(compute_params, shards, dtype=x.dtype)\n\n for i, x_shard in enumerate(x_shards):\n x_shard.requires_grad_(x_requires_grad)\n\n shard_step = x_shards[i].shape[0]\n shard_offset = i * x_shards[0].shape[0]\n\n # narrow(0, ...) creates a contiguous view that can receive gradients\n x_shard.grad = x_grad.narrow(0, shard_offset, shard_step)\n incoming_grad_shard = incoming_grad.narrow(0, shard_offset, shard_step)\n\n is_last_shard = i + 1 == shards\n grad_accumulator.install_hooks(is_last_shard)\n\n with torch.enable_grad():\n output = fn(module, x_shard)\n torch.autograd.backward(output, incoming_grad_shard)\n\n grad_accumulator.cleanup()\n del grad_accumulator\n\n # Restore original shape\n x_grad = x_grad.view(x_shape_orig) if x_requires_grad else None\n return (None, None, x_grad, None, None)\n\n\ndef _mlp_forward_fn(module, x):\n \"\"\"Forward function for LlamaMLP / Qwen2MLP / Qwen3MLP style.\"\"\"\n return module.down_proj(module.act_fn(module.gate_proj(x)) * module.up_proj(x))\n\n\n# ============================================================================\n# Monkey Patch Functions\n# ============================================================================\n\n# Model type to MLP class mapping\n_MODEL_TYPE_TO_MLP_CLASS = {\n \"llama\": (\"transformers.models.llama.modeling_llama\", \"LlamaMLP\"),\n \"qwen2\": (\"transformers.models.qwen2.modeling_qwen2\", \"Qwen2MLP\"),\n \"qwen2_5\": (\"transformers.models.qwen2.modeling_qwen2\", \"Qwen2MLP\"), # Qwen2.5 uses Qwen2 MLP\n \"qwen3\": (\"transformers.models.qwen3.modeling_qwen3\", \"Qwen3MLP\"),\n}\n\n\ndef apply_tiled_mlp_monkey_patch(\n num_shards: int = 4,\n model_type: Optional[str] = None,\n):\n \"\"\"Apply TiledMLP monkey patch based on model_type.\n\n This function MUST be called BEFORE model instantiation to take effect.\n It patches the MLP classes in transformers library to use TiledMLP for\n memory-efficient computation during training.\n\n Args:\n num_shards: Number of shards to split the input into. Higher values\n reduce peak memory but may slightly impact performance.\n model_type: The model type string (e.g., \"llama\", \"qwen2\", \"qwen3\").\n If None, patches all supported model types.\n\n Returns:\n List of patched class names.\n \"\"\"\n if model_type is None:\n types_to_patch = list(_MODEL_TYPE_TO_MLP_CLASS.keys())\n elif model_type in _MODEL_TYPE_TO_MLP_CLASS:\n types_to_patch = [model_type]\n else:\n raise ValueError(\n f\"TiledMLP does not support model_type='{model_type}'. \"\n f\"Supported types: {list(_MODEL_TYPE_TO_MLP_CLASS.keys())}. \"\n f\"For SwiGLU-style MLPs, you can add support by extending _MODEL_TYPE_TO_MLP_CLASS \"\n f\"in verl/models/transformers/tiled_mlp.py\"\n )\n\n patched_classes = []\n\n for mtype in types_to_patch:\n module_path, class_name = _MODEL_TYPE_TO_MLP_CLASS[mtype]\n try:\n import importlib\n\n module = importlib.import_module(module_path)\n mlp_class = getattr(module, class_name)\n _patch_mlp_class(mlp_class, _mlp_forward_fn, num_shards)\n if class_name not in patched_classes:\n patched_classes.append(class_name)\n except (ImportError, AttributeError) as e:\n print(f\"Warning: Could not patch {mtype} MLP: {e}\")\n\n if patched_classes:\n print(f\"TiledMLP monkey patch applied to: {', '.join(patched_classes)} (shards={num_shards})\")\n\n return patched_classes\n\n\ndef _patch_mlp_class(mlp_class: type[nn.Module], forward_fn, num_shards: int):\n \"\"\"Patch a single MLP class to use TiledMLP.\"\"\"\n\n def tiled_forward(self, x):\n compute_params = [p for p in self.parameters() if p.requires_grad]\n return TiledMLP.apply(forward_fn, self, x, num_shards, compute_params)\n\n mlp_class.forward = tiled_forward\n"}49{"file_name": "verl__single_controller__base__decorator.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport inspect\nfrom functools import partial, wraps\nfrom types import FunctionType\n\nfrom tensordict import TensorDict\n\nfrom verl.protocol import DataProtoFuture, _padding_size_key\nfrom verl.utils.py_functional import DynamicEnum\nfrom verl.utils.tensordict_utils import chunk_tensordict, concat_tensordict, contiguous\nfrom verl.utils.transferqueue_utils import BatchMeta\n\n# here we add a magic number of avoid user-defined function already have this attribute\nMAGIC_ATTR = \"attrs_3141562937\"\n\n\nclass Dispatch(DynamicEnum):\n \"\"\"Enum class defining different dispatch modes for distributed computation.\n\n Each mode represents a specific strategy for distributing data across\n different ranks in a distributed system. The modes are used to control\n how data is partitioned and processed across different worker groups.\n \"\"\"\n\n _registry = {}\n _next_value = 0\n\n\ndef init_predefined_dispatch_mode():\n Dispatch.register(\"RANK_ZERO\")\n Dispatch.register(\"ONE_TO_ALL\")\n Dispatch.register(\"ALL_TO_ALL\")\n Dispatch.register(\"DP_COMPUTE\")\n Dispatch.register(\"DP_COMPUTE_PROTO\")\n Dispatch.register(\"DP_COMPUTE_PROTO_WITH_FUNC\")\n Dispatch.register(\"DP_COMPUTE_METRIC\")\n # This is a special dispatch mode for vllm ExternalRayDistributedExecutor\n Dispatch.register(\"DIRECT_ROLLOUT_METHOD\")\n\n\nclass Execute(DynamicEnum):\n \"\"\"Enum class defining different execution modes for distributed computation.\n\n These modes control how a function should be executed across different ranks\n in a distributed system.\n \"\"\"\n\n _registry = {}\n _next_value = 0\n\n\ndef init_predefined_execute_mode():\n Execute.register(\"ALL\")\n Execute.register(\"RANK_ZERO\")\n\n\n# Initialize the two Dynamic Enum Classes\ninit_predefined_dispatch_mode()\ninit_predefined_execute_mode()\n\n\ndef _consolidate_tuple_td(chunked_arg):\n return tuple(contiguous(val).consolidate() for val in chunked_arg)\n\n\ndef _split_args_kwargs_data_proto(chunks, *args, **kwargs):\n from verl.protocol import DataProto, DataProtoFuture\n\n splitted_args = []\n for arg in args:\n assert isinstance(arg, DataProto | DataProtoFuture | BatchMeta | TensorDict)\n if isinstance(arg, TensorDict):\n chunked_arg = chunk_tensordict(arg, chunks)\n chunked_arg = _consolidate_tuple_td(chunked_arg)\n else:\n chunked_arg = arg.chunk(chunks=chunks)\n assert len(chunked_arg) == chunks\n splitted_args.append(chunked_arg)\n\n splitted_kwargs = {}\n for key, val in kwargs.items():\n assert isinstance(val, DataProto | DataProtoFuture | BatchMeta | TensorDict)\n if isinstance(val, TensorDict):\n chunked_kwarg = chunk_tensordict(val, chunks)\n chunked_kwarg = _consolidate_tuple_td(chunked_kwarg)\n else:\n chunked_kwarg = val.chunk(chunks=chunks)\n assert len(chunked_kwarg) == chunks\n splitted_kwargs[key] = chunked_kwarg\n\n return splitted_args, splitted_kwargs\n\n\ndef _split_args_kwargs_data_proto_with_auto_padding(chunks, *args, **kwargs):\n from verl.protocol import DataProto, DataProtoFuture\n\n data_proto_len = None\n padding_size = None\n\n def _padding_and_split_data(obj, chunks):\n nonlocal data_proto_len, padding_size\n assert isinstance(obj, DataProto | DataProtoFuture)\n if isinstance(obj, DataProto) and obj.is_padding_enabled():\n # for padding, we only support DataProto with same length\n if data_proto_len is None:\n data_proto_len = len(obj)\n padding_size = (chunks - (data_proto_len % chunks)) if (data_proto_len % chunks > 0) else 0\n else:\n assert data_proto_len == len(obj), (\n f\"expecting all arg share same length of {data_proto_len}, but got {len(obj)}\"\n )\n obj.padding(padding_size=padding_size)\n return obj.chunk(chunks=chunks)\n\n splitted_args = [_padding_and_split_data(arg, chunks) for arg in args]\n splitted_kwargs = {key: _padding_and_split_data(val, chunks) for key, val in kwargs.items()}\n if padding_size is not None:\n splitted_kwargs[_padding_size_key] = padding_size\n\n return splitted_args, splitted_kwargs\n\n\ndef dispatch_one_to_all(worker_group, *args, **kwargs):\n args = tuple([arg] * worker_group.world_size for arg in args)\n kwargs = {k: [v] * worker_group.world_size for k, v in kwargs.items()}\n return args, kwargs\n\n\ndef dummy_direct_rollout_call(worker_group, *args, **kwargs):\n raise NotImplementedError(\"Direct rollout call is forbidden.\")\n\n\ndef dispatch_all_to_all(worker_group, *args, **kwargs):\n return args, kwargs\n\n\ndef collect_all_to_all(worker_group, output):\n return output\n\n\ndef _concat_data_proto_or_future(output: list):\n import ray\n\n from verl.protocol import DataProto, DataProtoFuture\n\n # make sure all the elements in output has the same type\n for o in output:\n assert type(o) is type(output[0])\n\n o = output[0]\n\n if isinstance(o, DataProto):\n return DataProto.concat(output)\n elif isinstance(o, ray.ObjectRef):\n return DataProtoFuture.concat(output)\n elif isinstance(o, BatchMeta):\n return BatchMeta.concat(output)\n elif isinstance(o, TensorDict):\n return concat_tensordict(output)\n else:\n raise NotImplementedError\n\n\ndef dispatch_dp_compute(worker_group, *args, **kwargs):\n from verl.single_controller.base.worker_group import WorkerGroup\n\n assert isinstance(worker_group, WorkerGroup)\n for arg in args:\n assert isinstance(arg, tuple | list) and len(arg) == worker_group.world_size\n for k, v in kwargs.items():\n assert isinstance(v, tuple | list) and len(v) == worker_group.world_size\n return args, kwargs\n\n\ndef collect_dp_compute(worker_group, output):\n from verl.single_controller.base.worker_group import WorkerGroup\n\n assert isinstance(worker_group, WorkerGroup)\n assert len(output) == worker_group.world_size\n return output\n\n\ndef dispatch_dp_compute_data_proto(worker_group, *args, **kwargs):\n from verl.single_controller.base.worker_group import WorkerGroup\n\n assert isinstance(worker_group, WorkerGroup)\n # Note: enable auto padding for dp compute DatapProto\n splitted_args, splitted_kwargs = _split_args_kwargs_data_proto_with_auto_padding(\n worker_group.world_size,\n *args,\n **kwargs,\n )\n return splitted_args, splitted_kwargs\n\n\ndef dispatch_dp_compute_data_proto_with_func(worker_group, *args, **kwargs):\n from verl.single_controller.base.worker_group import WorkerGroup\n\n assert isinstance(worker_group, WorkerGroup)\n assert isinstance(args[0], FunctionType) # NOTE: The first one args is a function!\n\n splitted_args, splitted_kwargs = _split_args_kwargs_data_proto(worker_group.world_size, *args[1:], **kwargs)\n splitted_args_with_func = [[args[0]] * worker_group.world_size] + splitted_args\n return splitted_args_with_func, splitted_kwargs\n\n\ndef collect_dp_compute_data_proto(worker_group, output):\n import ray\n\n from verl.protocol import DataProto\n\n for o in output:\n assert isinstance(o, DataProto | ray.ObjectRef), f\"expecting {o} to be DataProto, but got {type(o)}\"\n\n output = collect_dp_compute(worker_group, output)\n return _concat_data_proto_or_future(output)\n\n\ndef dispatch_nd_compute(dp_rank_mapping: list[int], dp_size, worker_group, *args, **kwargs):\n import os\n\n from verl.single_controller.base.worker_group import WorkerGroup\n from verl.utils.ray_utils import parallel_put\n\n assert isinstance(worker_group, WorkerGroup)\n\n max_workers = max(1, min(len(args[0]), os.cpu_count()))\n\n args = [parallel_put(arg, max_workers=max_workers) for arg in args]\n kwargs = {k: parallel_put(v, max_workers=max_workers) for k, v in kwargs.items()}\n\n all_args = []\n for arg in args:\n assert isinstance(arg, tuple | list) and len(arg) == dp_size\n transformed_args = []\n for i in range(worker_group.world_size):\n local_dp_rank = dp_rank_mapping[i]\n transformed_args.append(arg[local_dp_rank])\n all_args.append(transformed_args)\n all_args = tuple(all_args)\n\n all_kwargs = {}\n for k, v in kwargs.items():\n assert isinstance(v, tuple | list) and len(v) == dp_size\n transformed_v = []\n for i in range(worker_group.world_size):\n local_dp_rank = dp_rank_mapping[i]\n transformed_v.append(v[local_dp_rank])\n all_kwargs[k] = transformed_v\n return all_args, all_kwargs\n\n\ndef collect_nd_compute(collect_mask: list[bool], worker_group, output):\n from verl.single_controller.base.worker_group import WorkerGroup\n\n assert isinstance(worker_group, WorkerGroup)\n assert len(output) == worker_group.world_size\n\n output_in_dp = []\n for global_rank in range(worker_group.world_size):\n collect_dp_rank = collect_mask[global_rank]\n if collect_dp_rank:\n output_in_dp.append(output[global_rank])\n return output_in_dp\n\n\ndef dispatch_nd_compute_dataproto(dp_rank_mapping: list[int], dp_size, worker_group, *args, **kwargs):\n splitted_args, splitted_kwargs = _split_args_kwargs_data_proto(dp_size, *args, **kwargs)\n return dispatch_nd_compute(dp_rank_mapping, dp_size, worker_group, *splitted_args, **splitted_kwargs)\n\n\ndef collect_nd_compute_dataproto(collect_mask: list[bool], worker_group, output):\n output = collect_nd_compute(collect_mask, worker_group, output)\n import ray\n\n from verl.protocol import DataProto\n\n for o in output:\n assert isinstance(o, DataProto | ray.ObjectRef | BatchMeta | TensorDict), (\n f\"expecting {o} to be DataProto | ray.ObjectRef | BatchMeta | TensorDict, but got {type(o)}\"\n )\n return _concat_data_proto_or_future(output)\n\n\ndef dispatch_lazy_compute_data_proto(mesh_name, worker_group, *args, **kwargs):\n from verl.single_controller.base.worker_group import WorkerGroup\n\n assert isinstance(worker_group, WorkerGroup)\n\n # query dispatch info of the worker group\n if mesh_name not in worker_group._dispatch_info:\n worker_group._dispatch_info[mesh_name] = worker_group._query_dispatch_info(mesh_name)\n assert len(worker_group._dispatch_info[mesh_name]) == worker_group.world_size\n\n dp_rank_mapping = worker_group._dispatch_info[mesh_name]\n # perform dispatch\n dp_size = max(dp_rank_mapping) + 1\n return dispatch_nd_compute_dataproto(dp_rank_mapping, dp_size, worker_group, *args, **kwargs)\n\n\ndef collect_lazy_compute_data_proto(mesh_name, worker_group, *args, **kwargs):\n from verl.single_controller.base.worker_group import WorkerGroup\n\n assert isinstance(worker_group, WorkerGroup)\n\n # the dispatch info is stored in the worker group\n assert mesh_name in worker_group._dispatch_info\n\n if mesh_name not in worker_group._collect_info:\n worker_group._collect_info[mesh_name] = worker_group._query_collect_info(mesh_name)\n assert len(worker_group._collect_info[mesh_name]) == worker_group.world_size\n\n # a boolean of whether the dp_rank is used for collect\n collect_mask = worker_group._collect_info[mesh_name]\n # perform dispatch\n return collect_nd_compute_dataproto(collect_mask, worker_group, *args, **kwargs)\n\n\ndef make_nd_compute_dataproto_dispatch_fn(mesh_name):\n return {\n \"dispatch_fn\": partial(dispatch_lazy_compute_data_proto, mesh_name),\n \"collect_fn\": partial(collect_lazy_compute_data_proto, mesh_name),\n }\n\n\n# Global registry for dispatch mode.\nDISPATCH_MODE_FN_REGISTRY = {\n Dispatch.ONE_TO_ALL: {\n \"dispatch_fn\": dispatch_one_to_all,\n \"collect_fn\": collect_all_to_all,\n },\n Dispatch.ALL_TO_ALL: {\n \"dispatch_fn\": dispatch_all_to_all,\n \"collect_fn\": collect_all_to_all,\n },\n Dispatch.DP_COMPUTE: {\"dispatch_fn\": dispatch_dp_compute, \"collect_fn\": collect_dp_compute},\n Dispatch.DP_COMPUTE_PROTO: {\n \"dispatch_fn\": dispatch_dp_compute_data_proto,\n \"collect_fn\": collect_dp_compute_data_proto,\n },\n Dispatch.DP_COMPUTE_PROTO_WITH_FUNC: {\n \"dispatch_fn\": dispatch_dp_compute_data_proto_with_func,\n \"collect_fn\": collect_dp_compute_data_proto,\n },\n Dispatch.DP_COMPUTE_METRIC: {\"dispatch_fn\": dispatch_dp_compute_data_proto, \"collect_fn\": collect_dp_compute},\n Dispatch.DIRECT_ROLLOUT_METHOD: {\n \"dispatch_fn\": dummy_direct_rollout_call,\n \"collect_fn\": dummy_direct_rollout_call,\n },\n}\n\n\ndef get_predefined_dispatch_fn(dispatch_mode):\n return DISPATCH_MODE_FN_REGISTRY[dispatch_mode]\n\n\ndef register_dispatch_mode(dispatch_mode_name, dispatch_fn, collect_fn):\n \"\"\"\n Register a new dispatch mode.\n \"\"\"\n dispatch_mode = Dispatch.register(dispatch_mode_name)\n _check_dispatch_mode(dispatch_mode)\n assert dispatch_mode not in DISPATCH_MODE_FN_REGISTRY, f\"dispatch_mode_name {dispatch_mode_name} already exists\"\n DISPATCH_MODE_FN_REGISTRY[dispatch_mode] = {\"dispatch_fn\": dispatch_fn, \"collect_fn\": collect_fn}\n\n\ndef update_dispatch_mode(dispatch_mode, dispatch_fn, collect_fn):\n \"\"\"\n Update the dispatch mode.\n \"\"\"\n _check_dispatch_mode(dispatch_mode)\n assert dispatch_mode in DISPATCH_MODE_FN_REGISTRY, f\"dispatch_mode {dispatch_mode} not found\"\n DISPATCH_MODE_FN_REGISTRY[dispatch_mode] = {\"dispatch_fn\": dispatch_fn, \"collect_fn\": collect_fn}\n\n\ndef get_predefined_execute_fn(execute_mode):\n \"\"\"\n Note that here we only asks execute_all and execute_rank_zero to be implemented\n Leave the choice of how these two functions handle argument 'blocking' to users\n \"\"\"\n predefined_execute_mode_fn = {\n Execute.ALL: {\"execute_fn_name\": \"execute_all\"},\n Execute.RANK_ZERO: {\"execute_fn_name\": \"execute_rank_zero\"},\n }\n return predefined_execute_mode_fn[execute_mode]\n\n\ndef _check_dispatch_mode(dispatch_mode):\n assert isinstance(dispatch_mode, Dispatch | dict), (\n f\"dispatch_mode must be a Dispatch or a Dict. Got {dispatch_mode}\"\n )\n if isinstance(dispatch_mode, dict):\n necessary_keys = [\"dispatch_fn\", \"collect_fn\"]\n for key in necessary_keys:\n assert key in dispatch_mode, f\"key {key} should be in dispatch_mode if it is a dictionary\"\n\n\ndef _check_execute_mode(execute_mode):\n assert isinstance(execute_mode, Execute), f\"execute_mode must be a Execute. Got {execute_mode}\"\n\n\ndef _materialize_futures(*args, **kwargs):\n new_args = []\n for arg in args:\n if isinstance(arg, DataProtoFuture):\n arg = arg.get()\n # add more type to materialize\n new_args.append(arg)\n for k, v in kwargs.items():\n if isinstance(v, DataProtoFuture):\n kwargs[k] = v.get()\n\n new_args = tuple(new_args)\n return new_args, kwargs\n\n\ndef register(dispatch_mode=Dispatch.ALL_TO_ALL, execute_mode=Execute.ALL, blocking=True, materialize_futures=True):\n \"\"\"Register a function with distributed execution configuration.\n\n This decorator registers a function with specific dispatch and execution modes\n for distributed computation. It handles both synchronous and asynchronous\n functions, and optionally materializes futures before execution.\n\n Args:\n dispatch_mode:\n Dispatch mode for computation distribution. Default: Dispatch.ALL_TO_ALL.\n execute_mode:\n Execute mode for computation distribution. Default: Execute.ALL.\n blocking:\n Whether the execution should be blocking. Defaults to True.\n materialize_futures:\n Whether to materialize the data before dispatching. Defaults to True.\n\n Returns:\n A decorator that wraps the original function with distributed execution\n configuration.\n \"\"\"\n from verl.utils.transferqueue_utils import tqbridge\n\n _check_dispatch_mode(dispatch_mode=dispatch_mode)\n _check_execute_mode(execute_mode=execute_mode)\n\n def decorator(func):\n func = tqbridge(dispatch_mode=dispatch_mode)(func)\n\n @wraps(func)\n def inner(*args, **kwargs):\n if materialize_futures:\n args, kwargs = _materialize_futures(*args, **kwargs)\n return func(*args, **kwargs)\n\n @wraps(func)\n async def async_inner(*args, **kwargs):\n if materialize_futures:\n args, kwargs = _materialize_futures(*args, **kwargs)\n return await func(*args, **kwargs)\n\n wrapper = async_inner if inspect.iscoroutinefunction(func) else inner\n attrs = {\"dispatch_mode\": dispatch_mode, \"execute_mode\": execute_mode, \"blocking\": blocking}\n setattr(wrapper, MAGIC_ATTR, attrs)\n return wrapper\n\n return decorator\n"}50{"file_name": "verl__single_controller__base__worker.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nthe class for Worker\n\"\"\"\n\nimport os\nimport socket\nimport warnings\nfrom dataclasses import dataclass\n\nimport ray\n\nfrom verl.utils.device import (\n get_torch_device,\n get_visible_devices_keyword,\n is_npu_available,\n)\n\nfrom .decorator import Dispatch, Execute, register\n\n\n@dataclass\nclass DistRankInfo:\n tp_rank: int\n dp_rank: int\n pp_rank: int\n cp_rank: int\n\n\n@dataclass\nclass DistGlobalInfo:\n tp_size: int\n dp_size: int\n pp_size: int\n cp_size: int\n\n\nclass WorkerHelper:\n @staticmethod\n def _get_node_ip():\n if os.getenv(\"WG_BACKEND\", None) == \"ray\":\n return ray.util.get_node_ip_address()\n else:\n raise NotImplementedError(\"WG_BACKEND now just support ray mode.\")\n\n @staticmethod\n def _get_free_port():\n with socket.socket() as sock:\n sock.bind((\"\", 0))\n return sock.getsockname()[1]\n\n def get_availale_master_addr_port(self):\n warnings.warn(\n \"This function is deprecated due to typo in name; Please use `get_available_master_addr_port` instead\",\n stacklevel=2,\n )\n return self.get_available_master_addr_port()\n\n def get_available_master_addr_port(self):\n return self._get_node_ip().strip(\"[]\"), str(self._get_free_port())\n\n\n# we assume that in each WorkerGroup, there is a Master Worker\nclass Worker(WorkerHelper):\n \"\"\"A distributed worker that handles initialization and configuration for distributed training.\n\n This class manages worker initialization, configuration, and provides methods for executing\n distributed operations. It handles communication settings, device configuration, and worker\n metadata management.\n \"\"\"\n\n fused_worker_attr_name = \"fused_worker_dict\"\n\n def _register_dispatch_collect_info(self, mesh_name: str, dp_rank: int, is_collect: bool):\n \"\"\"Register the dp_rank for a given mesh name. This function is meant to be called by the worker\n\n Args:\n mesh_name (str):\n Name of the mesh to register dp_rank for.\n dp_rank (int):\n dp_rank to register for the given mesh name.\n is_collect (bool):\n Whether the dp_rank is used for collect.\n \"\"\"\n if mesh_name in self.__dispatch_dp_rank or mesh_name in self.__collect_dp_rank:\n raise ValueError(f\"mesh_name {mesh_name} has been registered\")\n self.__dispatch_dp_rank[mesh_name] = dp_rank\n self.__collect_dp_rank[mesh_name] = is_collect\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def _query_dispatch_info(self, mesh_name: str):\n \"\"\"Query the dispatch info for a given mesh name.\n\n Args:\n mesh_name (str):\n Name of the mesh to query dispatch info for.\n\n Returns:\n int:\n The dp_rank for the given mesh name.\n \"\"\"\n assert mesh_name in self.__dispatch_dp_rank, f\"{mesh_name} is not registered in {self.__class__.__name__}\"\n # note that each rank store its own dp_rank\n return self.__dispatch_dp_rank[mesh_name]\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def _query_collect_info(self, mesh_name: str):\n return self.query_collect_info(mesh_name)\n\n def query_collect_info(self, mesh_name: str):\n \"\"\"Query the collect info for a given mesh name.\n\n Args:\n mesh_name (str):\n Name of the mesh to query collect info for.\n\n Returns:\n bool:\n Whether the dp_rank is used for collect.\n \"\"\"\n assert mesh_name in self.__collect_dp_rank, f\"{mesh_name} is not registered in {self.__class__.__name__}\"\n return self.__collect_dp_rank[mesh_name]\n\n def get_dispatch_collect(self):\n \"\"\"Get all registered dispatch and collect dp_ranks.\n\n Returns:\n dict[str, int]:\n A dictionary mapping mesh names to their dispatch dp_ranks.\n dict[str, bool]:\n A dictionary mapping mesh names to whether they are used for collect.\n \"\"\"\n return {\"dispatch_dp_rank\": self.__dispatch_dp_rank, \"collect_dp_rank\": self.__collect_dp_rank}\n\n def set_dispatch_collect(self, mesh_name: str, dispatch_dp_rank: dict[str, int], collect_dp_rank: dict[str, bool]):\n \"\"\"Set the dispatch and collect dp_ranks for all registered meshes.\n\n Args:\n mesh_name (str): Mesh name to set dispatch and collect dp_ranks for.\n dispatch_dp_rank (dict[str, int]):\n A dictionary mapping mesh names to their dispatch dp_ranks.\n collect_dp_rank (dict[str, bool]):\n A dictionary mapping mesh names to whether they are used for collect.\n \"\"\"\n assert mesh_name not in self.__dispatch_dp_rank, (\n f\"{mesh_name} is already registered, {self.__dispatch_dp_rank.keys()}\"\n )\n assert mesh_name not in self.__collect_dp_rank, (\n f\"{mesh_name} is already registered, {self.__collect_dp_rank.keys()}\"\n )\n for dp_rank in dispatch_dp_rank.values():\n self.__dispatch_dp_rank[mesh_name] = dp_rank\n for is_collect in collect_dp_rank.values():\n self.__collect_dp_rank[mesh_name] = is_collect\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=True)\n def create_transferqueue_client(self, config):\n from verl.utils.transferqueue_utils import create_transferqueue_client\n\n create_transferqueue_client(\n client_id=f\"worker_{self.rank}\",\n config=config.transfer_queue,\n )\n\n @classmethod\n def env_keys(cls):\n \"\"\"The keys of the environment variables that are used to configure the Worker.\"\"\"\n return [\n \"WORLD_SIZE\",\n \"RANK\",\n \"LOCAL_WORLD_SIZE\",\n \"LOCAL_RANK\",\n \"MASTER_ADDR\",\n \"MASTER_PORT\",\n get_visible_devices_keyword().upper(),\n ]\n\n def __init__(self, cuda_visible_devices=None) -> None:\n \"\"\"Initialize the worker with environment settings and device configuration.\n\n Args:\n cuda_visible_devices (str, optional):\n CUDA visible devices configuration. Defaults to None.\n \"\"\"\n # construct a meta from environment variable. Note that the import must be inside the class because\n # it is executed remotely\n import os\n\n self._setup_env_cuda_visible_devices()\n\n world_size = int(os.environ[\"WORLD_SIZE\"])\n rank = int(os.environ[\"RANK\"])\n self._rank = rank\n self._world_size = world_size\n\n master_addr = os.environ[\"MASTER_ADDR\"]\n master_port = os.environ[\"MASTER_PORT\"]\n\n local_world_size = int(os.getenv(\"LOCAL_WORLD_SIZE\", \"1\"))\n local_rank = int(os.getenv(\"LOCAL_RANK\", \"0\"))\n\n store = {\n \"_world_size\": world_size,\n \"_rank\": rank,\n \"_local_world_size\": local_world_size,\n \"_local_rank\": local_rank,\n \"_master_addr\": master_addr,\n \"_master_port\": master_port,\n }\n if cuda_visible_devices is not None:\n store[f\"_{get_visible_devices_keyword()}\".lower()] = cuda_visible_devices\n\n self._configure_with_store(store=store)\n\n self.fused_worker_dict = {}\n self.__dispatch_dp_rank = {}\n self.__collect_dp_rank = {}\n\n def get_fused_worker_by_name(self, worker_name: str):\n \"\"\"Get a fused worker by its name.\n\n Args:\n worker_name (str):\n Name of the worker to retrieve\n \"\"\"\n return self.fused_worker_dict.get(worker_name, None)\n\n def _setup_env_cuda_visible_devices(self):\n from verl.utils.ray_utils import ray_noset_visible_devices\n\n is_ray_noset_visible_devices = ray_noset_visible_devices()\n\n # Prevent use of clashing `{CUDA/HIP/ROCR}_VISIBLE_DEVICES``\n rocr_val = os.environ.get(\"ROCR_VISIBLE_DEVICES\", None)\n hip_val = os.environ.get(\"HIP_VISIBLE_DEVICES\", None)\n cuda_val = os.environ.get(\"CUDA_VISIBLE_DEVICES\", None)\n if hip_val:\n # Switch the use of HIP_VISIBLE_DEVICES to CUDA_VISIBLE_DEVICES for consistency.\n # Make sure that the HIP_VISIBLE_DEVICES is set to the same value as CUDA_VISIBLE_DEVICES\n # at this point.\n val = os.environ.pop(\"HIP_VISIBLE_DEVICES\")\n hip_val = None\n if cuda_val:\n assert val == cuda_val, (\n f\"Please use the same HIP_VISIBLE_DEVICES or CUDA_VISIBLE_DEVICES, inconsistant values \"\n f\"found: {val} and {cuda_val}.\"\n )\n else:\n cuda_val = val\n os.environ[\"CUDA_VISIBLE_DEVICES\"] = val\n # os.environ[\"HIP_VISIBLE_DEVICES\"] = val\n\n if rocr_val:\n # You must take care if both HIP/CUDA and ROCR env vars are set as they have\n # different meanings. Both env vars accept either a list of ints or a\n # list of UUIDs. The ROCR env var is processed first which then reduces\n # the number of GPUs that HIP can select from.\n # https://github.com/pytorch/pytorch/pull/144026\n # To avoid the complexity of this, we simply gives out error if both are set\n # (Also to keep consistency with ray's practice with 2.45.0).\n # Otherwise, we will set ROCR_VISIBLE_DEVICES to CUDA_VISIBLE_DEVICES\n # and remove ROCR_VISIBLE_DEVICES.\n if cuda_val:\n raise ValueError(\"Please don't set ROCR_VISIBLE_DEVICES when HIP/CUDA_VISIBLE_DEVICES is set.\")\n\n cuda_val = os.environ.pop(\"ROCR_VISIBLE_DEVICES\")\n os.environ[\"CUDA_VISIBLE_DEVICES\"] = cuda_val\n rocr_val = None\n\n if is_ray_noset_visible_devices:\n # NOTE: Ray will automatically set the *_VISIBLE_DEVICES\n # environment variable for each actor, unless\n # RAY_EXPERIMENTAL_NOSET_*_VISIBLE_DEVICES is set,\n # so we need to set local rank when the flag is set.\n device_name = \"NPU\" if is_npu_available else \"GPU\"\n local_rank = ray.get_runtime_context().get_accelerator_ids()[device_name][0]\n os.environ[\"LOCAL_RANK\"] = local_rank\n get_torch_device().set_device(int(local_rank))\n\n def _configure_with_store(self, store: dict):\n \"\"\"\n This function should only be called inside by WorkerGroup\n \"\"\"\n store_env_dict = {f\"_{key.lower()}\": store.get(f\"_{key.lower()}\", None) for key in type(self).env_keys()}\n self.__dict__.update(store_env_dict) # this is hacky\n # print(f\"__dict__: {self.__dict__}\")\n for key in type(self).env_keys():\n val = self.__dict__.get(f\"_{key.lower()}\", None)\n if val is not None:\n # print(f\"set {key} to {val}\")\n os.environ[key] = str(val)\n os.environ[\"REDIS_STORE_SERVER_HOST\"] = (\n str(self._master_addr).replace(\"[\", \"\").replace(\"]\", \"\") if self._master_addr else \"\"\n )\n\n def get_master_addr_port(self):\n \"\"\"Get the master address and port for distributed communication.\"\"\"\n return self._master_addr, self._master_port\n\n def get_cuda_visible_devices(self):\n \"\"\"Get the CUDA visible devices configuration.\"\"\"\n import os\n\n visible_devices = os.environ.get(get_visible_devices_keyword().upper(), \"not set\")\n return visible_devices\n\n @property\n def world_size(self):\n \"\"\"Get the total number of workers in the distributed setup.\"\"\"\n return self._world_size\n\n @property\n def rank(self):\n \"\"\"Get the rank of this worker in the distributed setup.\"\"\"\n return self._rank\n\n @register(dispatch_mode=Dispatch.DP_COMPUTE_PROTO_WITH_FUNC)\n def execute_with_func_generator(self, func, *args, **kwargs):\n \"\"\"Execute a function with function generator dispatch mode.\n\n Args:\n func:\n Function to execute\n *args:\n Positional arguments for the function\n **kwargs:\n Keyword arguments for the function\n \"\"\"\n ret_proto = func(self, *args, **kwargs)\n return ret_proto\n\n @register(dispatch_mode=Dispatch.ALL_TO_ALL, execute_mode=Execute.RANK_ZERO)\n def execute_func_rank_zero(self, func, *args, **kwargs):\n \"\"\"Execute a function in rank zero execution mode.\n\n Args:\n func:\n Function to execute\n *args:\n Positional arguments for the function\n **kwargs:\n Keyword arguments for the function\n \"\"\"\n result = func(*args, **kwargs)\n return result\n"}51{"file_name": "verl__single_controller__ray__base.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport inspect\nimport logging\nimport os\nimport socket\nfrom copy import deepcopy\nfrom dataclasses import dataclass, field\nfrom typing import Any, Optional\n\nimport numpy as np\nimport ray\nfrom ray.experimental.state.api import get_actor\nfrom ray.util.placement_group import PlacementGroup, placement_group\nfrom ray.util.scheduling_strategies import NodeAffinitySchedulingStrategy, PlacementGroupSchedulingStrategy\n\nfrom verl.protocol import DataProto, _padding_size_key\nfrom verl.single_controller.base import ClassWithInitArgs, ResourcePool, Worker, WorkerGroup\nfrom verl.single_controller.base.decorator import MAGIC_ATTR, Dispatch\nfrom verl.utils.device import get_device_name\nfrom verl.utils.py_functional import temp_env_var\n\n__all__ = [\"Worker\"]\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef get_random_string(length: int) -> str:\n import random\n import string\n\n letters_digits = string.ascii_letters + string.digits\n return \"\".join(random.choice(letters_digits) for _ in range(length))\n\n\ndef func_generator(self, method_name, dispatch_fn, collect_fn, execute_fn, blocking):\n class Functor:\n def __call__(this, *args, **kwargs):\n args, kwargs = dispatch_fn(self, *args, **kwargs)\n padding_count = kwargs.pop(_padding_size_key, 0)\n output = execute_fn(method_name, *args, **kwargs)\n if blocking:\n output = ray.get(output)\n output = collect_fn(self, output)\n if padding_count > 0:\n if isinstance(output, DataProto):\n indices = [i for i in range(len(output))][:-padding_count]\n output = output.select_idxs(indices)\n elif isinstance(output, list):\n output = output[:-padding_count]\n return output\n\n # use class type to pass the method_name to get a better observability\n return type(method_name, (Functor,), {})()\n\n\ndef sort_placement_group_by_node_ip(pgs: list[PlacementGroup]) -> list[PlacementGroup]:\n \"\"\"\n Sort the placement groups by node ip, all bundles in a single placement group should be on the same node.\n\n FSDPCheckpointManager saves sharded model states and optimizer states in local storage, which requires RANK\n to be consistent across nodes when resume from checkpoint.\n\n With this function, if there's only one resource pool and there's no node change, RANK should be consistent\n across nodes in multiple ray jobs, even if the whole ray cluster is restarted.\n \"\"\"\n node_ip = {node[\"NodeID\"]: node[\"NodeManagerAddress\"] for node in ray.nodes()}\n pg_ip = {}\n for pg in pgs:\n specs = ray._private.state.state.placement_group_table(pg.id)\n # all bunles should be on the same node\n node_id = specs[\"bundles_to_node_id\"][0]\n pg_ip[pg.id] = node_ip[node_id]\n return sorted(pgs, key=lambda pg: pg_ip[pg.id])\n\n\n@ray.remote\ndef get_master_addr_port(master_port_range: Optional[list[int]] = None) -> tuple[str, str]:\n addr = ray.util.get_node_ip_address().strip(\"[]\")\n\n if master_port_range is None:\n with socket.socket() as s:\n s.bind((\"\", 0))\n port = s.getsockname()[1]\n else:\n port = master_port_range[0]\n while port < master_port_range[1]:\n try:\n with socket.socket() as s:\n s.bind((\"\", port))\n break\n except OSError:\n port += 1 # Increment port number if already in use\n logger.info(\"Port %d is already in use, trying port %d\", port - 1, port)\n else:\n raise RuntimeError(f\"Could not find a free port in range {master_port_range}\")\n return addr, str(port)\n\n\nclass RayResourcePool(ResourcePool):\n def __init__(\n self,\n process_on_nodes: Optional[list[int]] = None,\n use_gpu: bool = True,\n name_prefix: str = None,\n max_colocate_count: int = 10,\n detached=False,\n accelerator_type: Optional[str] = None,\n ) -> None:\n super().__init__(process_on_nodes, max_colocate_count)\n self.use_gpu = use_gpu\n # print(f\"in RayProcessDispatchConfiguration: name_prefix = {name_prefix}\")\n self.name_prefix = get_random_string(length=6) if name_prefix is None else name_prefix\n self.pgs = None\n self.detached = detached\n self.accelerator_type = accelerator_type\n\n def get_placement_groups(self, strategy=\"STRICT_PACK\", name=None, device_name=\"cuda\"):\n if self.pgs is not None:\n return self.pgs\n\n pg_name_prefix = (\n name if name else f\"{self.name_prefix}verl_group_{'_'.join([str(count) for count in self._store])}:\"\n )\n # print(f\"pg_name_prefix = {pg_name_prefix}\")\n if device_name == \"npu\":\n device_name = \"NPU\"\n elif device_name == \"cuda\":\n device_name = \"GPU\"\n\n bundle = {\"CPU\": self.max_colocate_count}\n if self.use_gpu:\n bundle[device_name] = 1\n if self.accelerator_type is not None:\n bundle[self.accelerator_type] = 1e-4\n pg_scheme = [[bundle.copy() for _ in range(process_count)] for process_count in self._store]\n\n lifetime = \"detached\" if self.detached else None\n\n pgs = [\n placement_group(bundles=bundles, strategy=strategy, name=pg_name_prefix + str(idx), lifetime=lifetime)\n for idx, bundles in enumerate(pg_scheme)\n ]\n\n ray.get([pg.ready() for pg in pgs])\n\n self.pgs = sort_placement_group_by_node_ip(pgs)\n return pgs\n\n\nclass SubRayResourcePool(RayResourcePool):\n def __init__(\n self,\n placement_groups: list[PlacementGroup],\n start_bundle_index: int,\n subgroup_world_size: int,\n **kwargs,\n ) -> None:\n super().__init__(**kwargs)\n self.pgs = placement_groups\n self.start_bundle_index = start_bundle_index\n self.subgroup_world_size = subgroup_world_size\n\n @property\n def world_size(self):\n return self.subgroup_world_size\n\n\n@dataclass\nclass ResourcePoolManager:\n \"\"\"\n Define a resource pool specification. Resource pool will be initialized first.\n \"\"\"\n\n resource_pool_spec: dict[str, list[int]]\n mapping: dict[int, str]\n resource_pool_dict: dict[str, RayResourcePool] = field(default_factory=dict)\n\n def create_resource_pool(self):\n \"\"\"Create Ray resource pools for distributed training.\n\n Initializes resource pools based on the resource pool specification,\n with each pool managing GPU resources across multiple nodes.\n For FSDP backend, uses max_colocate_count=1 to merge WorkerGroups.\n For Megatron backend, uses max_colocate_count>1 for different models.\n \"\"\"\n for resource_pool_name, process_on_nodes in self.resource_pool_spec.items():\n # max_colocate_count means the number of WorkerGroups (i.e. processes) in each RayResourcePool\n # For FSDP backend, using max_colocate_count=3: actor_critic_ref, rollout, reward model (optional)\n # For Megatron backend, we recommend using max_colocate_count>1\n # that can utilize different WorkerGroup for differnt models\n resource_pool = RayResourcePool(\n process_on_nodes=process_on_nodes, use_gpu=True, max_colocate_count=3, name_prefix=resource_pool_name\n )\n self.resource_pool_dict[resource_pool_name] = resource_pool\n\n self._check_resource_available()\n\n def get_resource_pool(self, role) -> RayResourcePool:\n \"\"\"Get the resource pool of the worker_cls\"\"\"\n return self.resource_pool_dict[self.mapping[role]]\n\n def get_n_gpus(self) -> int:\n \"\"\"Get the number of gpus in this cluster.\"\"\"\n return sum([n_gpus for process_on_nodes in self.resource_pool_spec.values() for n_gpus in process_on_nodes])\n\n def _check_resource_available(self):\n \"\"\"Check if the resource pool can be satisfied in this ray cluster.\"\"\"\n node_available_resources = ray._private.state.available_resources_per_node()\n node_available_gpus = {\n node: node_info.get(\"GPU\", 0) if \"GPU\" in node_info else node_info.get(\"NPU\", 0)\n for node, node_info in node_available_resources.items()\n }\n\n # check total required gpus can be satisfied\n total_available_gpus = sum(node_available_gpus.values())\n total_required_gpus = sum(\n [n_gpus for process_on_nodes in self.resource_pool_spec.values() for n_gpus in process_on_nodes]\n )\n if total_available_gpus < total_required_gpus:\n raise ValueError(\n f\"Total available GPUs {total_available_gpus} is less than total desired GPUs {total_required_gpus}\"\n )\n\n\ndef extract_pg_from_exist(\n resource_pools: dict[str, RayResourcePool], src_role_names: list[str], resource_pool: RayResourcePool\n) -> list:\n src_pgs = [\n pg\n for role_name, resource_pool in resource_pools.items()\n for pg in resource_pool.get_placement_groups()\n if role_name in src_role_names\n ]\n\n sorted_src_pgs = sorted(src_pgs, key=lambda pg: pg.bundle_count, reverse=True)\n sorted_process_on_nodes = sorted([(val, idx) for idx, val in enumerate(resource_pool.store)], reverse=True)\n\n unsorted_pgs: list[tuple[int, PlacementGroup]] = []\n searching_idx = 0\n for request_process, original_idx in sorted_process_on_nodes:\n assert searching_idx < len(sorted_src_pgs), f\"no enough nodes for request: searching {searching_idx} th node\"\n assert request_process <= sorted_src_pgs[searching_idx].bundle_count, (\n f\"requesting {request_process} processes, bundle count cannot satisfy\"\n )\n unsorted_pgs.append((original_idx, sorted_src_pgs[searching_idx]))\n searching_idx += 1\n\n return [pg for _, pg in sorted(unsorted_pgs)]\n\n\n# split a RayResourcePool or SubRayResourcePool into multiple SubRayResourcePool\ndef split_resource_pool(\n resource_pool: RayResourcePool | SubRayResourcePool, split_size: int | list[int]\n) -> list[SubRayResourcePool]:\n \"\"\"\n Split a RayResourcePool into multiple SubRayResourcePool.\n resouce_pool can also be a SubRayResourcePool (have been splited) for multiple-time spliting.\n\n Args:\n resource_pool (RayResourcePool | SubRayResourcePool): The resource pool to split.\n split_size (int | list[int]): The size of each split. If int, all splits will have the same size.\n If list[int], each element in the list represents the size of a split.\n\n Returns:\n list[SubRayResourcePool]: A list of SubRayResourcePool after splitting.\n \"\"\"\n # convert split_size to list[int]\n if isinstance(split_size, int):\n assert resource_pool.world_size % split_size == 0, \"split_size must be a divisor of world_size\"\n num_replica = resource_pool.world_size // split_size\n split_size_list = [split_size] * num_replica\n else:\n split_size_list = split_size\n\n assert sum(split_size_list) == resource_pool.world_size, \"split_size must sum up to world_size\"\n\n # judge if this resource pool has been splited\n if isinstance(resource_pool, SubRayResourcePool):\n start_bundle_idx_list = np.cumsum([resource_pool.start_bundle_index] + split_size_list[:-1])\n else:\n start_bundle_idx_list = np.cumsum([0] + split_size_list[:-1])\n\n # ensure resource_pool.pgs has been initialized\n placement_groups = resource_pool.get_placement_groups()\n split_resource_pools = [\n SubRayResourcePool(\n process_on_nodes=resource_pool.store,\n use_gpu=resource_pool.use_gpu,\n name_prefix=f\"{resource_pool.name_prefix}_split_{split_idx}\",\n max_colocate_count=resource_pool.max_colocate_count,\n placement_groups=placement_groups,\n start_bundle_index=start_bundle_idx_list[split_idx],\n subgroup_world_size=split_size_list[split_idx],\n )\n for split_idx in range(len(split_size_list))\n ]\n return split_resource_pools\n\n\ndef merge_resource_pool(rp1: RayResourcePool, rp2: RayResourcePool) -> RayResourcePool:\n assert rp1.use_gpu == rp2.use_gpu, \"Both RayResourcePool must either use_gpu or not\"\n assert rp1.max_colocate_count == rp2.max_colocate_count, \"Both RayResourcePool must has the same max_colocate_count\"\n assert rp1.n_gpus_per_node == rp2.n_gpus_per_node, \"Both RayResourcePool must has the same n_gpus_per_node\"\n assert rp1.detached == rp2.detached, \"Detached ResourcePool cannot be merged with non-detached ResourcePool\"\n\n new_store = rp1.store + rp2.store\n\n merged = type(rp1)(\n new_store, rp1.use_gpu, f\"{rp1.name_prefix}_{rp2.name_prefix}\", rp1.max_colocate_count, rp1.detached\n )\n merged.pgs = rp1.get_placement_groups(device_name=get_device_name()) + rp2.get_placement_groups(\n device_name=get_device_name()\n )\n\n return merged\n\n\nclass RayClassWithInitArgs(ClassWithInitArgs):\n \"\"\"A wrapper class for Ray actors with initialization arguments.\n\n This class extends ClassWithInitArgs to provide additional functionality for\n configuring and creating Ray actors with specific resource requirements and\n scheduling strategies.\n \"\"\"\n\n def __init__(self, cls, *args, **kwargs) -> None:\n # self._options = kwargs.pop('options', dict())\n super().__init__(cls, *args, **kwargs)\n self._options = {}\n self._additional_resource = {}\n\n def set_additional_resource(self, additional_resource):\n \"\"\"Set additional resource requirements for the actor.\n\n Args:\n additional_resource: Dictionary specifying additional resource requirements\n \"\"\"\n self._additional_resource = additional_resource\n\n def update_options(self, options: dict):\n \"\"\"Update the Ray actor creation options.\n\n Args:\n options: Dictionary of options to update\n \"\"\"\n self._options.update(options)\n\n def __call__(\n self,\n placement_group,\n placement_group_bundle_idx,\n use_gpu: bool = True,\n num_gpus=1,\n sharing_with=None,\n device_name=\"cuda\",\n ) -> Any:\n \"\"\"Create and return a Ray actor with the configured options.\n\n Args:\n placement_group: Ray placement group for scheduling\n placement_group_bundle_idx: Index of the bundle in the placement group\n use_gpu: Whether to use GPU resources\n num_gpus: Number of GPUs to allocate\n sharing_with: Actor to share resources with\n device_name: Device for training\n\n Returns:\n A Ray actor handle with the configured options\n \"\"\"\n if sharing_with is not None:\n target_node_id = ray.get(sharing_with.get_node_id.remote())\n visible_devices = ray.get(sharing_with.get_cuda_visible_devices.remote())\n options = {\"scheduling_strategy\": NodeAffinitySchedulingStrategy(node_id=target_node_id, soft=False)}\n return self.cls.options(**options).remote(*self.args, cuda_visible_devices=visible_devices, **self.kwargs)\n\n options = {\n \"scheduling_strategy\": PlacementGroupSchedulingStrategy(\n placement_group=placement_group, placement_group_bundle_index=placement_group_bundle_idx\n )\n }\n options.update(self._options)\n\n if use_gpu and device_name == \"cuda\":\n options[\"num_gpus\"] = num_gpus\n if use_gpu and device_name == \"npu\":\n options[\"resources\"] = {\"NPU\": num_gpus}\n\n if len(self._additional_resource) > 1:\n for k, v in self._additional_resource.items():\n options[k] = v\n\n # print(\"cls:\", self.cls)\n # print(\"args: \", self.args)\n # print(\"kwargs: \", self.kwargs)\n return self.cls.options(**options).remote(*self.args, **self.kwargs)\n\n\nclass RayWorkerGroup(WorkerGroup):\n \"\"\"A group of Ray workers that can be managed collectively.\n\n This class extends WorkerGroup to provide Ray-specific functionality for\n creating and managing groups of Ray actors with specific resource requirements\n and scheduling strategies.\n \"\"\"\n\n def __init__(\n self,\n resource_pool: RayResourcePool = None,\n ray_cls_with_init: RayClassWithInitArgs = None,\n bin_pack: bool = True,\n name_prefix: str = None,\n detached=False,\n worker_names=None,\n worker_handles: list[ray.actor.ActorHandle] = None,\n ray_wait_register_center_timeout: int = 300,\n **kwargs,\n ) -> None:\n \"\"\"Initialize a RayWorkerGroup.\n\n Args:\n resource_pool: Resource pool for worker allocation\n ray_cls_with_init: Class with initialization arguments for workers\n bin_pack: Whether to use strict bin packing for resource allocation\n name_prefix: Prefix for worker names\n detached: Whether workers should be detached\n worker_names: Names of existing workers to attach to\n ray_wait_register_center_timeout: Timeout for waiting on register center\n **kwargs: Additional keyword arguments\n \"\"\"\n self._master_addr = kwargs.pop(\"master_addr\", None)\n self._master_port = kwargs.pop(\"master_port\", None)\n self.use_gpu = kwargs.pop(\"use_gpu\", resource_pool.use_gpu if resource_pool is not None else True)\n self._ray_master_port_range = kwargs.pop(\"master_port_range\", None)\n super().__init__(resource_pool=resource_pool, **kwargs)\n self.ray_cls_with_init = ray_cls_with_init\n self.name_prefix = get_random_string(length=6) if name_prefix is None else name_prefix\n self._ray_wait_register_center_timeout = ray_wait_register_center_timeout\n # Whether the WorkerGroup is a Colocate WorkerGroup created by FusedWorker.\n self.fused_worker_used = False if ray_cls_with_init is None else ray_cls_with_init.fused_worker_used\n # if a WorkerGroup is spawned from Colocate WorkerGroup, this indicates which sub-class is binded to\n # this WorkerGroup.\n self.sub_cls_name = \"\"\n self.device_name = kwargs.get(\"device_name\", \"cuda\")\n self.profile_steps = kwargs.get(\"profile_steps\", None)\n self.worker_nsight_options = kwargs.get(\"worker_nsight_options\", None)\n self.customized_worker_env = kwargs.get(\"worker_env\", {})\n if self.worker_nsight_options is not None and self.worker_nsight_options[\"capture-range-end\"] is None:\n self.worker_nsight_options[\"capture-range-end\"] = f\"repeat-shutdown:{6 * len(self.profile_steps)}\"\n\n if worker_names is not None and (not self.fused_worker_used):\n assert self._is_init_with_detached_workers\n self._worker_names = worker_names\n\n if self._is_init_with_detached_workers:\n self._init_with_detached_workers(worker_names=worker_names, worker_handles=worker_handles)\n elif isinstance(resource_pool, SubRayResourcePool):\n self._init_with_subresource_pool(\n resource_pool=resource_pool,\n ray_cls_with_init=ray_cls_with_init,\n bin_pack=bin_pack,\n detached=detached,\n worker_env=self.customized_worker_env,\n )\n else:\n self._init_with_resource_pool(\n resource_pool=resource_pool,\n ray_cls_with_init=ray_cls_with_init,\n bin_pack=bin_pack,\n detached=detached,\n worker_env=self.customized_worker_env,\n )\n\n if ray_cls_with_init is not None:\n self._bind_worker_method(self.ray_cls_with_init.cls, func_generator)\n\n self.wg_dict = None\n self.method_names = []\n\n def _is_worker_alive(self, worker: ray.actor.ActorHandle):\n \"\"\"Check if a worker actor is still alive.\n\n Args:\n worker: Ray actor handle to check\n\n Returns:\n bool: True if the worker is alive, False otherwise\n \"\"\"\n worker_state_dict = get_actor(worker._actor_id.hex())\n return worker_state_dict.get(\"state\", \"undefined\") == \"ALIVE\" if worker_state_dict is not None else False\n\n def _init_with_detached_workers(self, worker_names, worker_handles):\n # ray.get_actor holds a weak reference to the actor, which causes actors garbage collected unexpectedly\n # if we only hold spawn RayWorkerGroup. By passing actor handle explicitly, spawn RayWorkerGroup have\n # strong reference to these actors.\n # https://github.com/ray-project/ray/pull/45699\n workers = worker_handles if worker_handles else [ray.get_actor(name=name) for name in worker_names]\n self._workers = workers\n self._world_size = len(workers)\n\n def _get_master_addr_port(self, pg, bundle_index=0, master_port_range=None):\n \"\"\"Get master addr and port for this worker group\"\"\"\n if self._master_addr is None and self._master_port is None:\n self._master_addr, self._master_port = ray.get(\n get_master_addr_port.options(\n scheduling_strategy=PlacementGroupSchedulingStrategy(\n placement_group=pg, placement_group_bundle_index=bundle_index\n ),\n ).remote(master_port_range=master_port_range)\n )\n elif self._master_addr is not None and self._master_port is not None:\n logger.debug(f\"{self._master_addr=} {self._master_port=}\")\n else:\n raise ValueError(\n \"Both 'master_addr' and 'master_port' must be provided if you intend to manually specify them, \"\n \"or neither should be provided to use Ray's default assignment.\"\n )\n\n def _init_with_resource_pool(\n self,\n resource_pool,\n ray_cls_with_init,\n bin_pack,\n detached,\n worker_env=None,\n ):\n \"\"\"Initialize the worker group by creating new workers from a resource pool.\n\n Args:\n resource_pool: Resource pool for worker allocation\n ray_cls_with_init: Class with initialization arguments for workers\n bin_pack: Whether to use strict bin packing for resource allocation\n detached: Whether workers should be detached\n \"\"\"\n self.resource_pool = resource_pool\n strategy = \"PACK\"\n if bin_pack:\n strategy = \"STRICT_PACK\"\n pgs = resource_pool.get_placement_groups(strategy=strategy, device_name=self.device_name)\n world_size = resource_pool.world_size\n self._world_size = world_size\n # cia.add_kwarg(\"_world_size\", world_size)\n\n rank = -1\n local_world_size = resource_pool.store[0]\n for pg_idx, pg in enumerate(sort_placement_group_by_node_ip(pgs)):\n assert local_world_size <= pg.bundle_count, f\"when generating for {self.name_prefix}, for the \"\n if pg_idx == 0:\n self._get_master_addr_port(pg, bundle_index=0, master_port_range=self._ray_master_port_range)\n\n for local_rank in range(local_world_size):\n rank += 1\n self._create_worker(\n rank=rank,\n pg_idx=pg_idx,\n pg=pg,\n local_rank=local_rank,\n resource_pool=resource_pool,\n ray_cls_with_init=ray_cls_with_init,\n worker_env=worker_env,\n detached=detached,\n )\n\n def _init_with_subresource_pool(self, resource_pool, ray_cls_with_init, bin_pack, detached, worker_env=None):\n \"\"\"Initialize the worker group by creating new workers from a resource pool or sub resource pool.\n Args:\n resource_pool: Resource pool for worker allocation\n ray_cls_with_init: Class with initialization arguments for workers\n bin_pack: Whether to use strict bin packing for resource allocation\n detached: Whether workers should be detached\n \"\"\"\n strategy = \"PACK\"\n if bin_pack:\n strategy = \"STRICT_PACK\"\n pgs = resource_pool.get_placement_groups(strategy=strategy, device_name=self.device_name)\n world_size = resource_pool.world_size\n self._world_size = world_size\n\n rank = -1\n local_world_size = resource_pool.store[0]\n self._get_master_addr_port(\n pgs[resource_pool.start_bundle_index // local_world_size],\n bundle_index=resource_pool.start_bundle_index % local_world_size,\n master_port_range=self._ray_master_port_range,\n )\n for curr_rank in range(resource_pool.start_bundle_index, resource_pool.start_bundle_index + world_size):\n pg_idx = curr_rank // local_world_size\n pg = pgs[pg_idx]\n local_rank = curr_rank % local_world_size\n assert local_world_size <= pg.bundle_count, f\"when generating for {self.name_prefix}, for the \"\n\n rank += 1\n self._create_worker(\n rank=rank,\n pg_idx=pg_idx,\n pg=pg,\n local_rank=local_rank,\n resource_pool=resource_pool,\n ray_cls_with_init=ray_cls_with_init,\n worker_env=worker_env,\n detached=detached,\n )\n\n def _create_worker(self, rank, pg_idx, pg, local_rank, resource_pool, ray_cls_with_init, worker_env, detached):\n world_size = resource_pool.world_size\n use_gpu = resource_pool.use_gpu\n if self.use_gpu and not use_gpu:\n raise ValueError(\"use_gpu is True but resource_pool.use_gpu is False\")\n local_world_size = resource_pool.store[0]\n num_gpus = 1 / resource_pool.max_colocate_count\n\n # we pass in environment variable at option so that Worker can use environment variable to set\n env_vars = {\n \"WORLD_SIZE\": str(world_size),\n \"RANK\": str(rank),\n \"WG_PREFIX\": self.name_prefix,\n \"WG_BACKEND\": \"ray\",\n \"RAY_LOCAL_WORLD_SIZE\": str(local_world_size),\n \"MASTER_ADDR\": self._master_addr,\n \"MASTER_PORT\": self._master_port,\n }\n if worker_env is not None:\n logging.debug(f\"Appending ray class env, origin: {env_vars}, customized env: {worker_env}\")\n conflict_env_vars = set(env_vars.keys()) & set(worker_env.keys())\n if len(conflict_env_vars) > 0:\n logging.error(\n f\"User customized env vars conflict with system env: {conflict_env_vars} \"\n f\"Overriding may cause unexpected behavior.\"\n )\n raise ValueError(f\"Cannot override protected system env: {conflict_env_vars}\")\n env_vars.update(worker_env)\n import re\n\n cia_name = type(ray_cls_with_init.cls).__name__\n match = re.search(r\"ActorClass\\(([^)]+)\\)\", cia_name) # ray.remote(Obj) -> \"ActorClass(Obj)\"\n cia_name = match.group(1) if match else cia_name # \"ActorClass(Obj)\" -> \"Obj\"\n name = f\"{self.name_prefix}{cia_name}_{pg_idx}:{local_rank}\" # e.g. Worker_2:5\n\n if self.profile_steps and self.device_name == \"cuda\":\n ray_cls_with_init.update_options(\n {\n \"runtime_env\": {\n \"env_vars\": env_vars,\n \"nsight\": self.worker_nsight_options,\n },\n \"name\": name,\n }\n )\n else:\n ray_cls_with_init.update_options({\"runtime_env\": {\"env_vars\": env_vars}, \"name\": name})\n\n if detached:\n ray_cls_with_init.update_options({\"lifetime\": \"detached\"})\n\n # create a worker\n worker = ray_cls_with_init(\n placement_group=pg,\n placement_group_bundle_idx=local_rank,\n use_gpu=self.use_gpu,\n num_gpus=num_gpus,\n device_name=self.device_name,\n )\n self._workers.append(worker)\n self._worker_names.append(name)\n\n @property\n def worker_names(self):\n return self._worker_names\n\n @classmethod\n def from_detached(\n cls,\n name_prefix=None,\n worker_names=None,\n worker_handles=None,\n ray_cls_with_init=None,\n **kwargs,\n ):\n \"\"\"Create a worker group from existing detached workers.\n\n Args:\n name_prefix: Prefix for worker names\n worker_names: Names of existing workers to attach to\n ray_cls_with_init: Class with initialization arguments for workers\n\n Returns:\n A new RayWorkerGroup instance\n \"\"\"\n worker_group = cls(\n resource_pool=None,\n ray_cls_with_init=ray_cls_with_init,\n name_prefix=name_prefix,\n worker_names=worker_names,\n worker_handles=worker_handles,\n **kwargs,\n )\n return worker_group\n\n def spawn(self, prefix_set):\n \"\"\"Spawn to a dictionary of worker groups, each with a subset of method with prefix.\n\n Args:\n prefix_set: Set of prefixes to create worker groups for\n\n Returns:\n Dictionary of worker groups keyed by prefix\n \"\"\"\n if self.fused_worker_used:\n return self.spawn_fused(prefix_set)\n\n def _rebind_actor_methods(worker_group, actor_name):\n prefix: str = actor_name + \"_\"\n for method_name in dir(worker_group):\n if method_name.startswith(prefix):\n original_method_name = method_name.removeprefix(prefix)\n method = getattr(worker_group, method_name)\n setattr(worker_group, original_method_name, method)\n\n new_worker_group_dict = {}\n for prefix in prefix_set:\n new_worker_group = self.from_detached(\n name_prefix=self.name_prefix,\n worker_names=self._worker_names,\n worker_handles=self._workers,\n ray_cls_with_init=self.ray_cls_with_init,\n profile_steps=self.profile_steps,\n worker_nsight_options=self.worker_nsight_options,\n )\n\n _rebind_actor_methods(new_worker_group, prefix)\n new_worker_group_dict[prefix] = new_worker_group\n return new_worker_group_dict\n\n def spawn_fused(self, prefix_set):\n \"\"\"Create a dictionary of worker groups for fused workers.\n\n Args:\n prefix_set: Set of prefixes to create worker groups for\n\n Returns:\n Dictionary of worker groups keyed by prefix\n \"\"\"\n wg_dict = dict()\n for key in prefix_set:\n new_wg = deepcopy(self)\n new_wg._bind_worker_method(self.ray_cls_with_init.cls.raw_cls_dict[key], func_generator)\n new_wg.sub_cls_name = key\n wg_dict[key] = new_wg\n return wg_dict\n\n def fuse(self, prefix_set):\n \"\"\"Fuse multiple worker groups into the current worker group.\n\n Args:\n prefix_set: Set of prefixes to fuse into the worker group\n \"\"\"\n if self.wg_dict is None:\n self.wg_dict = self.spawn(prefix_set)\n for role_name, role_wg in self.wg_dict.items():\n setattr(self, role_name, role_wg)\n self.method_names = self._bind_worker_method(self.ray_cls_with_init.cls, func_generator)\n\n def _execute_remote_single_worker(self, worker, method_name: str, *args, **kwargs):\n \"\"\"Execute a method on a single worker remotely.\n\n Args:\n worker: The worker actor handle\n method_name: Name of the method to execute\n *args: Positional arguments for the method\n **kwargs: Keyword arguments for the method\n\n Returns:\n Remote object reference to the method execution\n \"\"\"\n if self.fused_worker_used and method_name not in self.method_names:\n remote_call = getattr(worker, self.fused_worker_execute_fn_name)\n return remote_call.remote(f\"{self.sub_cls_name}_fwmn_{method_name}\", *args, **kwargs)\n # fused worker not used\n remote_call = getattr(worker, method_name)\n return remote_call.remote(*args, **kwargs)\n\n def execute_rank_zero_sync(self, method_name: str, *args, **kwargs):\n \"\"\"Execute a method on rank zero worker synchronously.\n\n Args:\n method_name: Name of the method to execute\n *args: Positional arguments for the method\n **kwargs: Keyword arguments for the method\n\n Returns:\n Result of the method execution\n \"\"\"\n return ray.get(self.execute_rank_zero_async(method_name, *args, **kwargs))\n\n def execute_rank_zero_async(self, method_name: str, *args, **kwargs):\n \"\"\"Execute a method on rank zero worker asynchronously.\n\n Args:\n method_name: Name of the method to execute\n *args: Positional arguments for the method\n **kwargs: Keyword arguments for the method\n\n Returns:\n Remote object reference to the method execution\n \"\"\"\n return self._execute_remote_single_worker(self._workers[0], method_name, *args, **kwargs)\n\n def execute_rank_zero(self, method_name: str, *args, **kwargs):\n \"\"\"Alias for execute_rank_zero_async.\n\n Args:\n method_name: Name of the method to execute\n *args: Positional arguments for the method\n **kwargs: Keyword arguments for the method\n\n Returns:\n Remote object reference to the method execution\n \"\"\"\n return self.execute_rank_zero_async(method_name, *args, **kwargs)\n\n def execute_all(self, method_name: str, *args, **kwargs):\n \"\"\"Alias for execute_all_async.\n\n Args:\n method_name: Name of the method to execute\n *args: Positional arguments for the method\n **kwargs: Keyword arguments for the method\n\n Returns:\n List of remote object references to the method executions\n \"\"\"\n return self.execute_all_async(method_name, *args, **kwargs)\n\n def execute_all_sync(self, method_name: str, *args, **kwargs):\n \"\"\"Execute a method on all workers synchronously.\n\n Args:\n method_name: Name of the method to execute\n *args: Positional arguments for the method\n **kwargs: Keyword arguments for the method\n\n Returns:\n List of results from all workers\n \"\"\"\n return ray.get(self.execute_all_async(method_name, *args, **kwargs))\n\n def execute_all_async(self, method_name: str, *args, **kwargs):\n \"\"\"Execute a method on all workers asynchronously.\n\n Args:\n method_name: Name of the method to execute\n *args: Positional arguments for the method\n **kwargs: Keyword arguments for the method\n\n Returns:\n List of remote object references to the method executions\n \"\"\"\n # Here, we assume that if all arguments in args and kwargs are lists,\n # and their lengths match len(self._workers), we'll distribute each\n # element in these lists to the corresponding worker\n # print(f\"execute_all_async: method {method_name}({args}, {kwargs})\")\n length = len(self._workers)\n if all(isinstance(arg, list) for arg in args) and all(isinstance(kwarg, list) for kwarg in kwargs.values()):\n if all(len(arg) == length for arg in args) and all(len(kwarg) == length for kwarg in kwargs.values()):\n # print(f\"splitting args and kwargs into {length} shards\")\n result = []\n for i in range(length):\n sliced_args = tuple(arg[i] for arg in args)\n sliced_kwargs = {k: v[i] for k, v in kwargs.items()}\n result.append(\n self._execute_remote_single_worker(self._workers[i], method_name, *sliced_args, **sliced_kwargs)\n )\n return result\n\n return [self._execute_remote_single_worker(worker, method_name, *args, **kwargs) for worker in self._workers]\n\n @property\n def master_address(self):\n return self._master_addr\n\n @property\n def master_port(self):\n return self._master_port\n\n @property\n def workers(self):\n return self._workers\n\n @property\n def world_size(self):\n return self._world_size\n\n\n\"\"\"\nUtilities that enables creating workers inside the same ray.Actor,\nwith code written in separate ray.Actors.\n\"\"\"\n\n\n# deprecated, switching to FusedWorker\ndef _bind_workers_method_to_parent(cls, key, user_defined_cls):\n \"\"\"\n Binds the methods of each worker to the WorkerDict.\n Note that we only bind public methods that are decorated by register\n \"\"\"\n\n for method_name in dir(user_defined_cls):\n try:\n method = getattr(user_defined_cls, method_name)\n assert callable(method), f\"{method_name} in {user_defined_cls} is not callable\"\n except Exception:\n # if it is a property, it will fail because Class doesn't have instance property\n continue\n\n if hasattr(method, MAGIC_ATTR):\n\n def generate_function(name, key=key):\n def func(self, *args, **kwargs):\n # dispatch to the actual worker\n return getattr(self.worker_dict[key], name)(*args, **kwargs)\n\n async def async_func(self, *args, **kwargs):\n # dispatch to the actual worker\n return await getattr(self.worker_dict[key], name)(*args, **kwargs)\n\n wrapper = async_func if inspect.iscoroutinefunction(method) else func # noqa: B023\n\n return wrapper\n\n func = generate_function(method_name)\n # pass MAGIC_ATTR for outer worker group\n attrs = getattr(method, MAGIC_ATTR)\n setattr(func, MAGIC_ATTR, attrs)\n try:\n # bind direct rollout method to class without prefix\n if attrs[\"dispatch_mode\"] == Dispatch.DIRECT_ROLLOUT_METHOD and \"rollout\" in key:\n assert not hasattr(cls, method_name), (\n f\"conflict direct rollout method {method_name} with role {key}\"\n )\n setattr(cls, method_name, func)\n print(f\"bind role {key} method {method_name} to class {cls}\")\n else:\n method_name_with_prefix = key + \"_\" + method_name\n setattr(cls, method_name_with_prefix, func)\n except Exception as e:\n raise ValueError(f\"Fail to set method_name {method_name}\") from e\n\n\ndef _unwrap_ray_remote(cls):\n if hasattr(cls, \"__ray_actor_class__\"):\n cls = cls.__ray_actor_class__\n return cls\n\n\ndef _determine_fsdp_megatron_base_class(mros: list):\n \"\"\"\n - megatron: base class should be MegatronWorker\n - fsdp: base class should be Worker\n \"\"\"\n for cls in mros[0]:\n if cls.__name__ == \"MegatronWorker\":\n return cls\n if cls.__name__ == \"Worker\":\n return cls\n raise ValueError(f\"Cannot determine base class for {mros}\")\n\n\n# deprecated, switching to FusedWorker\ndef create_colocated_worker_cls(class_dict: dict[str, RayClassWithInitArgs]):\n \"\"\"\n This function should return a class instance that delegates the calls to every\n cls in cls_dict\n \"\"\"\n cls_dict = {}\n init_args_dict = {}\n worker_cls = _determine_fsdp_megatron_base_class(\n [cls.cls.__ray_actor_class__.__mro__ for cls in class_dict.values()]\n )\n assert issubclass(worker_cls, Worker), f\"worker_cls {worker_cls} should be a subclass of Worker\"\n print(f\"colocated worker base class {worker_cls}\")\n\n for key, cls in class_dict.items():\n cls_dict[key] = cls.cls\n init_args_dict[key] = {\"args\": cls.args, \"kwargs\": cls.kwargs}\n\n assert cls_dict.keys() == init_args_dict.keys()\n\n # TODO: create a class with customizable name\n class WorkerDict(worker_cls):\n def __init__(self):\n super().__init__()\n self.worker_dict = {}\n for key, user_defined_cls in cls_dict.items():\n user_defined_cls = _unwrap_ray_remote(user_defined_cls)\n # directly instantiate the class without remote\n # in worker class, e.g. <verl.single_controller.base.worker.Worker>\n # when DISABLE_WORKER_INIT == 1 it will return immediately\n with temp_env_var(\"DISABLE_WORKER_INIT\", \"1\"):\n self.worker_dict[key] = user_defined_cls(\n *init_args_dict[key].get(\"args\", ()), **init_args_dict[key].get(\"kwargs\", {})\n )\n\n # now monkey-patch the methods from inner class to WorkerDict\n for key, user_defined_cls in cls_dict.items():\n user_defined_cls = _unwrap_ray_remote(user_defined_cls)\n _bind_workers_method_to_parent(WorkerDict, key, user_defined_cls)\n\n remote_cls = ray.remote(WorkerDict)\n remote_cls = RayClassWithInitArgs(cls=remote_cls)\n return remote_cls\n\n\nFusedWorkerCLSName = \"FusedWorker\"\n\n\ndef create_colocated_worker_raw_cls(class_dict: dict[str, RayClassWithInitArgs]):\n \"\"\"\n This function returns a FusedWorker class.\n\n `FusedWorker.{class_name}` -> FusedClass\n Use `class_name` as a param to directly access the underlying class.\n\n `FusedWorker._fuw_execute(\"{class_name}_fwmn_{method_name}\", *args, **kwargs)`\n First param must be \"{class_name}_fwmn_{method_name}\" in order to access `method_name`\n of underlying class `{class_name}`.\n\n `FusedWorker.fused_worker_dict` -> {\"class_name\": FusedClass}\n Stores all underlying classes.\n\n `FusedClass.fused_worker_dict` -> {\"class_name\": FusedClass}\n The same as `FusedWorker.fused_worker_dict`, enables underlying class to access other\n underlying classes.\n \"\"\"\n raw_cls_dict = {cls_name: _unwrap_ray_remote(cia.cls) for cls_name, cia in class_dict.items()}\n init_args_dict = {cls_name: cia.args for cls_name, cia in class_dict.items()}\n init_kwargs_dict = {cls_name: cia.kwargs for cls_name, cia in class_dict.items()}\n cls_names = list(class_dict.keys())\n\n # FusedWorker_Actor_Critic\n class_name_renamed = \"_\".join([FusedWorkerCLSName] + cls_names)\n\n class FusedWorker(Worker):\n def __init__(self, *args, **kwargs):\n super().__init__(*args, **kwargs)\n self.cls_names = cls_names\n self.raw_cls_dict = raw_cls_dict\n self.init_args_dict = init_args_dict\n self.init_kwargs_dict = init_kwargs_dict\n\n for cls_name, udc, ud_args, ud_kwargs in zip(\n self.cls_names,\n self.raw_cls_dict.values(),\n self.init_args_dict.values(),\n self.init_kwargs_dict.values(),\n strict=True,\n ):\n with temp_env_var(\"DISABLE_WORKER_INIT\", \"1\"):\n udc._get_ray_actor_cls_name = lambda x, name_renamed=class_name_renamed: name_renamed\n udc._get_ray_method_prefix = lambda x, name_prefixed=cls_name: f\"{name_prefixed}_\"\n # cls_name = \"actor\", \"critic\", udc = ActorWorker, CriticWorker\n self.fused_worker_dict[cls_name] = udc(*ud_args, **ud_kwargs)\n setattr(self, cls_name, self.fused_worker_dict[cls_name])\n\n # injecting fused_worker to each sub worker so they can be aware of existence of each other\n for _, worker in self.fused_worker_dict.items():\n setattr(worker, Worker.fused_worker_attr_name, self.fused_worker_dict)\n\n def _fuw_execute(self, method_name: str, *args, **kwargs):\n # for fused_worker, method_name is in a form of \"{cls_name}_fwmn_{method_name}\"\n # where fwmn stands \"fused worker method name\"\n names = method_name.split(\"_fwmn_\")\n cls_name = names[0]\n method_name = names[1]\n\n assert cls_name in self.fused_worker_dict, (\n f\"calling {cls_name}'s {method_name}, but {cls_name} not in fused_worker_dict\"\n )\n udc_method = getattr(self.fused_worker_dict[cls_name], method_name)\n return udc_method(*args, **kwargs)\n\n renamed_fused_worker_cls = type(class_name_renamed, (FusedWorker,), {})\n renamed_fused_worker_cls.is_fused_worker = True\n renamed_fused_worker_cls.raw_cls_dict = raw_cls_dict\n\n return renamed_fused_worker_cls\n\n\ndef create_colocated_worker_cls_fused(class_dict: dict[str, RayClassWithInitArgs]):\n \"\"\"\n This function returns a RayClassWithInitArgs instance of FusedWorker, which is an replacement\n of `create_colocated_worker_cls`. WorkerGroup constructed using this class will be a colocated\n WorkerGroup, which will be referenced as `ColocateWorkerGroup` below.\n\n `ColocateWorkerGroup.spawn(prefix_set)`\n returns a dict of WorkerGroup {\"class_name\": WorkerGroup}, WorkerGroup in this dict will\n have methods of underlying class `class_name` attached.\n\n `ColocateWorkerGroup.fuse(prefix_set)`\n After executing this function, `ColocateWorkerGroup.{class_name}` will return WorkerGroup\n with methods of underlying class `class_name` attached.\n \"\"\"\n raw_colocated_worker_cls = create_colocated_worker_raw_cls(class_dict)\n\n remote_cls = ray.remote(raw_colocated_worker_cls)\n cia = RayClassWithInitArgs(cls=remote_cls)\n cia.fused_worker_used = True\n\n return cia\n"}52{"file_name": "verl__trainer__config__config.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom dataclasses import dataclass, field\nfrom typing import Any, Optional\n\nfrom verl.base_config import BaseConfig\n\n__all__ = [\"CheckpointConfig\", \"ProfileConfig\", \"BaseModelConfig\"]\n\n\n@dataclass\nclass CheckpointConfig(BaseConfig):\n \"\"\"Configuration for model checkpointing.\n\n The inheritance from BaseConfig provides omegaconf.DictConfig-like interface for a dataclass config.\n\n Args:\n save_contents (list[str]): What to include in saved checkpoints.\n Options: 'model', 'optimizer', 'extra', 'hf_model'.\n load_contents (list[str]): Contents to load from checkpoint. Defaults to same as save_contents.\n async_save (bool): Whether to save checkpoints asynchronously. Only implemented for Megatron as of now.\n \"\"\"\n\n save_contents: list[str] = field(default_factory=lambda: [\"model\", \"optimizer\", \"extra\"])\n load_contents: list[str] = field(default_factory=lambda: [\"model\", \"optimizer\", \"extra\"])\n async_save: bool = False\n mbridge_config: dict[str, Any] = field(default_factory=dict)\n\n\n@dataclass\nclass ProfileConfig(BaseConfig):\n \"\"\"Configuration for profiling.\n\n The inheritance from BaseConfig provides omegaconf.DictConfig-like interface for a dataclass config.\n\n Args:\n profile_ranks (Optional[list[int]]): List of ranks to profile. None means all ranks.\n step_start (int): Starting step for profiling.\n step_end (int): Ending step for profiling.\n save_path (Optional[str]): Path to save profiling results.\n \"\"\"\n\n profile_ranks: Optional[list[int]] = None\n step_start: int = -1\n step_end: int = -1\n save_path: Optional[str] = None\n\n\n@dataclass\nclass BaseModelConfig(BaseConfig):\n \"\"\"Base configuration for a model.\n Contains core settings for loading and initializing a pretrained model checkpoint.\n\n Args:\n path (str): Path to pretrained model weights.\n tokenizer_path (Optional[str]): Tokenizer path (defaults to actor's model path if not set).\n override_config (dict): Hugging Face config override.\n external_lib (Optional[str]): External model implementation (optional).\n trust_remote_code (bool): Whether to trust remote code from Hugging Face models.\n lora (dict[str, Any]): LoRA configuration dictionary.\n \"\"\"\n\n path: str = \"~/models/deepseek-llm-7b-chat\"\n tokenizer_path: Optional[str] = None\n override_config: dict[str, Any] = field(default_factory=dict)\n external_lib: Optional[str] = None\n trust_remote_code: bool = False\n lora: dict[str, Any] = field(default_factory=dict)\n\n\n@dataclass\nclass ModuleConfig(BaseConfig):\n \"\"\"Configuration for external Python module, which can be loaded, executed (and optionally, ``import``ed).\n\n Args:\n path (str, optional): Path to the module file to load and execute.\n name (str, optional): Name of the module to ``import``. Format: ``\"import.path.to.module\"``.\n If ``None``, the module will be loaded with a hased name and\n will not be added to ``sys.modules``, thus can not be ``import``ed as ``name``.\n \"\"\"\n\n path: Optional[str] = None\n name: Optional[str] = None\n"}53{"file_name": "verl__trainer__fsdp_sft_trainer.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nA lightweight one-file FSDP SFT Trainer\nTODO(zhangchi.usc1992)\n- Add calculation of mfu\n- Add validation\n\"\"\"\n\nimport os\n\nos.environ[\"NCCL_DEBUG\"] = \"WARN\"\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"true\"\n\nimport logging\nimport re\nimport time\nfrom contextlib import nullcontext\n\nimport hydra\nimport torch\nimport torch.distributed\nfrom omegaconf import DictConfig, OmegaConf\nfrom peft import LoraConfig, TaskType, get_peft_model\nfrom tensordict import TensorDict\nfrom torch import nn\nfrom torch.distributed.device_mesh import DeviceMesh, init_device_mesh\nfrom torch.distributed.fsdp import CPUOffload, MixedPrecision, ShardingStrategy\nfrom torch.distributed.fsdp import FullyShardedDataParallel as FSDP\nfrom torch.utils.data import Dataset, DistributedSampler\nfrom torchdata.stateful_dataloader import StatefulDataLoader\nfrom tqdm import tqdm\nfrom transformers import AutoConfig, AutoModelForCausalLM, PreTrainedModel\n\nimport verl.utils.hdfs_io as hdfs_io\nfrom verl.utils.attention_utils import index_first_axis, pad_input, rearrange, unpad_input\nfrom verl.utils.checkpoint.checkpoint_manager import find_latest_ckpt_path, get_checkpoint_tracker_filename\nfrom verl.utils.checkpoint.fsdp_checkpoint_manager import FSDPCheckpointManager\nfrom verl.utils.dataset import SFTDataset\nfrom verl.utils.dataset.multiturn_sft_dataset import MultiTurnSFTDataset\nfrom verl.utils.device import (\n auto_set_device,\n get_device_id,\n get_device_name,\n is_cuda_available,\n is_npu_available,\n)\nfrom verl.utils.distributed import destroy_global_process_group, initialize_global_process_group\nfrom verl.utils.fs import copy_to_local\nfrom verl.utils.fsdp_utils import (\n CPUOffloadPolicy,\n MixedPrecisionPolicy,\n apply_fsdp2,\n fsdp2_clip_grad_norm_,\n fsdp2_load_full_state_dict,\n get_fsdp_wrap_policy,\n get_init_weight_context_manager,\n init_fn,\n)\nfrom verl.utils.logger import log_with_rank\nfrom verl.utils.profiler import log_gpu_memory_usage\nfrom verl.utils.py_functional import convert_to_regular_types\nfrom verl.utils.torch_dtypes import PrecisionType\nfrom verl.utils.torch_functional import get_cosine_schedule_with_warmup, get_wsd_schedule_with_warmup\nfrom verl.utils.tracking import Tracking\nfrom verl.utils.ulysses import (\n gather_outputs_and_unpad,\n get_ulysses_sequence_parallel_world_size,\n ulysses_pad_and_slice_inputs,\n)\nfrom verl.workers.config.optimizer import build_optimizer\nfrom verl.workers.sharding_manager.fsdp_ulysses import FSDPUlyssesShardingManager\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_SFT_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef extract_step(path):\n match = re.search(r\"global_step_(\\d+)\", path)\n if match:\n return int(match.group(1))\n return None\n\n\nclass FSDPSFTTrainer:\n def __init__(\n self,\n config,\n device_mesh: DeviceMesh,\n ulysses_device_mesh: DeviceMesh,\n tokenizer,\n train_dataset: Dataset,\n val_dataset: Dataset,\n ):\n self.config = config\n self.device_mesh = device_mesh\n self.ulysses_device_mesh = ulysses_device_mesh\n self.sharding_manager = FSDPUlyssesShardingManager(self.ulysses_device_mesh)\n self.tokenizer = tokenizer\n if self.config.data.chat_template is not None:\n raise ValueError(\"Apply Chat template from config is not supported yet.\")\n\n # normalize dp size\n self._normalize_config_bsz()\n\n # Set sequence parallel size\n self.config.ulysses_sequence_parallel_size = getattr(self.config, \"ulysses_sequence_parallel_size\", 1)\n self.use_remove_padding = getattr(self.config, \"use_remove_padding\", False)\n if self.device_mesh.get_rank() == 0:\n print(f\"Using sequence parallel size: {self.config.ulysses_sequence_parallel_size}\")\n print(f\"Using remove padding: {self.use_remove_padding}\")\n\n self._build_dataloader(train_dataset, val_dataset)\n\n self.lora = self.config.model.get(\"lora_adapter_path\") is not None or self.config.model.lora_rank > 0\n\n # Initialize resume-related variables\n self.resume_global_step = 0\n\n # build model\n self._build_model_optimizer()\n\n # Initialize checkpoint manager\n self._init_checkpoint_manager()\n\n self.load_checkpoint()\n\n if self.device_mesh.get_rank() == 0:\n print(self.config)\n\n self.device_name = self.config.trainer.device\n\n def _normalize_config_bsz(self):\n dp_size = self.device_mesh.size(0) if not self.ulysses_device_mesh else self.ulysses_device_mesh.size(0)\n if self.device_mesh.get_rank() == 0:\n print(f\"Normalize batch size by dp {dp_size}\")\n\n assert self.config.data.train_batch_size % dp_size == 0, (\n f\"Global batch size {self.config.data.train_batch_size} is not divisible by dp size {dp_size}\"\n )\n\n self.config.data.train_batch_size //= dp_size\n\n assert self.config.data.train_batch_size % self.config.data.micro_batch_size_per_gpu == 0\n\n def _build_dataloader(self, train_dataset, val_dataset):\n # build dataset\n config = self.config\n self.train_dataset, self.val_dataset = train_dataset, val_dataset\n\n # build dataloader\n # Use data parallel rank and size instead of global rank and world size\n\n # If doing SP, we need to use the local rank and size\n if self.config.ulysses_sequence_parallel_size > 1:\n rank = self.ulysses_device_mesh.get_local_rank(\"dp\")\n world_size = self.ulysses_device_mesh.size(0)\n if self.ulysses_device_mesh.get_rank() == 0:\n print(f\"Using SP rank {rank} and size {world_size} for data distribution\")\n print(\"Each SP rank gets different data, but the same data WITHIN the same rank\")\n else:\n rank = self.device_mesh.get_rank()\n world_size = self.device_mesh.size()\n if self.device_mesh.get_rank() == 0:\n print(f\"Using FSDP rank {rank} and size {world_size} for data distribution\")\n\n # Set pin_memory_device when pin_memory is enabled.\n device_name = get_device_name()\n\n self.train_sampler = DistributedSampler(\n self.train_dataset, shuffle=True, num_replicas=world_size, rank=rank, drop_last=True\n )\n self.train_dataloader = StatefulDataLoader(\n dataset=self.train_dataset,\n batch_size=config.data.train_batch_size,\n sampler=self.train_sampler,\n num_workers=8,\n pin_memory=True,\n drop_last=True,\n pin_memory_device=device_name,\n )\n\n self.val_sampler = DistributedSampler(\n self.val_dataset, shuffle=False, num_replicas=world_size, rank=rank, drop_last=True\n )\n self.val_dataloader = StatefulDataLoader(\n dataset=self.val_dataset,\n batch_size=config.data.micro_batch_size_per_gpu,\n sampler=self.val_sampler,\n num_workers=8,\n pin_memory=True,\n drop_last=True,\n pin_memory_device=device_name,\n )\n\n def _build_model_optimizer(self):\n # TODO (zhangchi.usc1992):\n # 1. support pretrain from random weights\n # 2. support init directly from sharded weights\n local_model_path = copy_to_local(src=self.config.model.partial_pretrain, verbose=True)\n\n if self.config.model.get(\"external_lib\", None) is not None:\n # This is used to import external_lib into the huggingface systems\n import importlib\n\n importlib.import_module(self.config.model.external_lib)\n\n log_gpu_memory_usage(\"Before model allocation\", logger=logger)\n\n trust_remote_code = self.config.model.trust_remote_code\n torch_dtype = self.config.model.fsdp_config.get(\"model_dtype\", \"fp32\")\n torch_dtype = PrecisionType.to_dtype(torch_dtype)\n # load config first\n config = AutoConfig.from_pretrained(local_model_path, trust_remote_code=trust_remote_code)\n self.model_config = config\n if hasattr(self.model_config, \"max_position_embeddings\"):\n self.model_config.max_position_embeddings = max(\n self.model_config.max_position_embeddings, self.config.data.max_length\n )\n if self.config.ulysses_sequence_parallel_size > 1:\n assert self.use_remove_padding, \"Sequence parallel is only supported when remove_padding is enabled\"\n\n # This may be very large\n init_context = get_init_weight_context_manager(\n use_meta_tensor=not config.tie_word_embeddings, mesh=self.device_mesh\n )\n\n with init_context():\n self.model: PreTrainedModel = AutoModelForCausalLM.from_pretrained(\n local_model_path,\n config=config,\n torch_dtype=torch_dtype,\n attn_implementation=\"flash_attention_2\",\n trust_remote_code=trust_remote_code,\n )\n\n if self.use_remove_padding or self.config.ulysses_sequence_parallel_size > 1:\n from verl.models.transformers.monkey_patch import apply_monkey_patch\n\n apply_monkey_patch(model=self.model, ulysses_sp_size=self.config.ulysses_sequence_parallel_size)\n\n # Apply Liger kernel if use_liger is enabled\n if self.config.model.get(\"use_liger\", False):\n from liger_kernel.transformers.monkey_patch import _apply_liger_kernel_to_instance\n\n _apply_liger_kernel_to_instance(model=self.model)\n\n if self.lora:\n self.model.enable_input_require_grads()\n\n lora_adapter_path = self.config.model.get(\"lora_adapter_path\")\n if lora_adapter_path is not None:\n from peft import PeftModel\n\n print(f\"Loading pre-trained LoRA adapter for sft from: {lora_adapter_path}\")\n\n local_adapter_path = copy_to_local(lora_adapter_path, use_shm=self.config.model.use_shm)\n\n self.model = PeftModel.from_pretrained(self.model, local_adapter_path, is_trainable=True)\n peft_config = self.model.peft_config[\"default\"]\n # Ensure task_type is TaskType enum, not string\n if isinstance(peft_config.task_type, str):\n peft_config.task_type = TaskType.CAUSAL_LM\n else:\n # Convert config to regular Python types before creating PEFT model\n lora_config = {\n \"task_type\": TaskType.CAUSAL_LM,\n \"r\": self.config.model.lora_rank,\n \"lora_alpha\": self.config.model.lora_alpha,\n \"target_modules\": convert_to_regular_types(self.config.model.target_modules),\n \"bias\": \"none\",\n }\n self.model = get_peft_model(self.model, LoraConfig(**lora_config))\n self.model = self.model.to(torch_dtype)\n\n if self.config.model.enable_gradient_checkpointing:\n self.model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={\"use_reentrant\": False})\n\n log_gpu_memory_usage(\"After model allocation\", logger=logger)\n\n mixed_precision = MixedPrecision(\n param_dtype=torch.bfloat16, reduce_dtype=torch.float32, buffer_dtype=torch.float32\n )\n\n auto_wrap_policy = get_fsdp_wrap_policy(\n self.model,\n config=self.config.model.fsdp_config.wrap_policy,\n is_lora=self.lora,\n )\n\n if self.device_mesh.get_rank() == 0:\n print(auto_wrap_policy)\n\n if not self.config.model.fsdp_config.cpu_offload:\n cpu_offload = None\n else:\n cpu_offload = CPUOffload(offload_params=self.config.model.fsdp_config.offload_params)\n\n fsdp_strategy = self.config.model.strategy\n if fsdp_strategy == \"fsdp\":\n self.fsdp_model = FSDP(\n self.model,\n cpu_offload=cpu_offload,\n param_init_fn=init_fn,\n use_orig_params=False,\n auto_wrap_policy=auto_wrap_policy,\n device_id=get_device_id(),\n sharding_strategy=ShardingStrategy.FULL_SHARD,\n mixed_precision=mixed_precision,\n sync_module_states=True,\n device_mesh=self.device_mesh,\n forward_prefetch=False,\n )\n elif fsdp_strategy == \"fsdp2\":\n assert CPUOffloadPolicy is not None, \"PyTorch version >= 2.4 is required for using fully_shard API (FSDP2)\"\n mp_policy = MixedPrecisionPolicy(\n param_dtype=torch.bfloat16, reduce_dtype=torch.float32, cast_forward_inputs=True\n )\n\n fsdp_kwargs = {\n \"mesh\": self.device_mesh,\n \"mp_policy\": mp_policy,\n \"offload_policy\": cpu_offload,\n \"reshard_after_forward\": True,\n }\n full_state = self.model.state_dict()\n apply_fsdp2(self.model, fsdp_kwargs, self.config.model.fsdp_config)\n fsdp2_load_full_state_dict(self.model, full_state, self.device_mesh, cpu_offload)\n self.fsdp_model = self.model\n else:\n raise NotImplementedError(f\"not implement {fsdp_strategy}\")\n\n log_gpu_memory_usage(\"After FSDP wrapping\", logger=logger)\n\n self.optimizer = build_optimizer(self.fsdp_model.parameters(), self.config.optim)\n\n log_gpu_memory_usage(\"After initialize optimizer\", logger=logger)\n\n self.steps_per_epoch = len(self.train_dataloader)\n self.total_steps = self.steps_per_epoch * self.config.trainer.total_epochs\n\n if self.device_mesh.get_rank() == 0:\n print(\n f\"Number of steps/epoch {self.steps_per_epoch}, number of epochs \"\n f\"{self.config.trainer.total_epochs}, total number of steps {self.total_steps}\"\n )\n\n num_warmup_steps = int(self.total_steps * self.config.optim.lr_warmup_steps_ratio)\n\n if not hasattr(self.config.optim, \"lr_scheduler\") or self.config.optim.lr_scheduler == \"cosine\":\n self.lr_scheduler = get_cosine_schedule_with_warmup(\n optimizer=self.optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=self.total_steps\n )\n elif self.config.optim.lr_scheduler == \"wsd\":\n self.lr_scheduler = get_wsd_schedule_with_warmup(\n optimizer=self.optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=self.total_steps\n )\n else:\n raise ValueError(f\"Unknown lr scheduler: {self.config.optim.lr_scheduler}\")\n\n def _compute_loss_and_backward(self, batch, do_backward=True, n_micro_batches=1):\n \"\"\"Compute loss with optional sequence parallelism and remove padding features\"\"\"\n use_sp = self.use_remove_padding and self.config.ulysses_sequence_parallel_size > 1\n\n # Move inputs to GPU and prepare loss mask\n input_ids = batch[\"input_ids\"].to(self.device_name)\n attention_mask = batch[\"attention_mask\"].to(self.device_name)\n position_ids = batch[\"position_ids\"].to(self.device_name)\n loss_mask = batch.pop(\"loss_mask\")[:, 1:].reshape(-1).to(self.device_name)\n loss_fct = nn.CrossEntropyLoss(reduction=\"none\")\n\n # Context manager for sequence parallel if needed\n context = self.sharding_manager if use_sp else nullcontext()\n with context, torch.autocast(device_type=self.device_name, dtype=torch.bfloat16):\n if not use_sp:\n # Standard forward pass without sequence parallel\n labels = input_ids[:, 1:].contiguous()\n output = self.fsdp_model(\n input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, use_cache=False\n )\n logits = output.logits\n\n shift_logits = logits[..., :-1, :].contiguous()\n shift_labels = labels.contiguous()\n # Flatten the tokens\n shift_logits = shift_logits.view(-1, self.model.config.vocab_size)\n shift_labels = shift_labels.view(-1)\n # Enable model parallelism\n shift_labels = shift_labels.to(shift_logits.device)\n loss = loss_fct(shift_logits, shift_labels)\n loss = loss * loss_mask.to(loss.device)\n else:\n # IMPORTANT: We have a big assumption here, so we can shard the SAME sequence across SP ranks\n # i.e., each GPU has <1 sequence, and each SP group has 1 sequence\n # 1. All SP ranks will receive the *SAME* batch\n # 2. Different SP groups will receive *DIFFERENT* batches\n # This is implemented by the DistributedSampler\n\n batch_size, seqlen = input_ids.shape\n # Remove padding\n input_ids_rmpad, indices, *_ = unpad_input(\n input_ids.unsqueeze(-1), attention_mask\n ) # input_ids_rmpad (total_nnz, ...)\n input_ids_rmpad = input_ids_rmpad.transpose(0, 1) # (1, total_nnz)\n\n # Unpad position_ids to align rotary\n position_ids_rmpad = index_first_axis(\n rearrange(position_ids.unsqueeze(-1), \"b s ... -> (b s) ...\"), indices\n ).transpose(0, 1)\n\n # Pad and slice inputs for sequence parallelism\n input_ids_rmpad_sliced, position_ids_rmpad_padded, pad_size = ulysses_pad_and_slice_inputs(\n input_ids_rmpad, position_ids_rmpad, sp_size=get_ulysses_sequence_parallel_world_size()\n )\n # For computing loss\n input_ids_rmpad_rolled = torch.roll(input_ids_rmpad, shifts=-1, dims=1) # (1, total_nnz)\n input_ids_rmpad_rolled, _, _ = ulysses_pad_and_slice_inputs(\n input_ids_rmpad_rolled, None, get_ulysses_sequence_parallel_world_size()\n )\n input_ids_rmpad_rolled = input_ids_rmpad_rolled.squeeze(0) # ((total_nnz / sp) + pad)\n\n # Forward pass\n output = self.fsdp_model(\n input_ids=input_ids_rmpad_sliced,\n attention_mask=None, # Not needed with flash attention varlen\n position_ids=position_ids_rmpad_padded,\n use_cache=False,\n )\n\n # Compute loss locally then aggregate\n logits_rmpad = output.logits.squeeze(0)\n input_ids_rmpad_rolled = input_ids_rmpad_rolled.to(logits_rmpad.device)\n loss = loss_fct(logits_rmpad, input_ids_rmpad_rolled)\n # Gather and unpad for sequence parallelism\n loss = gather_outputs_and_unpad(loss, gather_dim=0, unpad_dim=0, padding_size=pad_size)\n\n # This is the loss collected from all ulysses ranks\n full_loss = pad_input(\n hidden_states=loss.unsqueeze(-1), indices=indices, batch=batch_size, seqlen=seqlen\n )\n full_loss = full_loss.squeeze(-1)[:, :-1] # Remove last token's loss\n full_loss = full_loss.reshape(-1)\n loss_mask = loss_mask.to(full_loss.device)\n loss = full_loss * loss_mask\n\n valid_token_this_rank = torch.sum(loss_mask)\n\n if self.config.data.balance_dp_token:\n torch.distributed.all_reduce(valid_token_this_rank)\n dp_size = self.ulysses_device_mesh.size(\"dp\") if use_sp else torch.distributed.get_world_size()\n else:\n dp_size = 1\n\n loss = torch.sum(loss) / (valid_token_this_rank + 1e-8) * dp_size\n\n loss = loss / n_micro_batches # normalize loss\n\n if do_backward:\n loss.backward()\n return loss\n\n def training_step(self, batch: TensorDict):\n start_time = time.time()\n\n self.fsdp_model.train()\n\n log_gpu_memory_usage(\"Before optimizer zero_grad\", logger=logger)\n\n self.optimizer.zero_grad()\n\n log_gpu_memory_usage(\"After optimizer zero_grad\", logger=logger)\n\n micro_batches = batch.split(self.config.data.micro_batch_size_per_gpu)\n n_micro_batches = len(micro_batches)\n step_loss = 0\n for micro_batch in micro_batches:\n loss = self._compute_loss_and_backward(batch=micro_batch, n_micro_batches=n_micro_batches)\n step_loss += loss.item()\n\n if self.config.model.strategy == \"fsdp\":\n grad_norm = self.fsdp_model.clip_grad_norm_(max_norm=self.config.optim.clip_grad)\n elif self.config.model.strategy == \"fsdp2\":\n grad_norm = fsdp2_clip_grad_norm_(self.fsdp_model.parameters(), max_norm=self.config.optim.clip_grad)\n else:\n raise NotImplementedError(f\"not implement {self.config.model.strategy}\")\n\n log_gpu_memory_usage(\"Before optimizer step\", logger=logger)\n\n # if grad_norm is not finite, skip the update\n if not torch.isfinite(grad_norm):\n print(f\"WARN: grad_norm is not finite: {grad_norm}\")\n self.optimizer.zero_grad()\n else:\n self.optimizer.step()\n\n log_gpu_memory_usage(\"After optimizer step\", logger=logger)\n\n self.lr_scheduler.step()\n\n # reduce loss across dp ranks\n lr = self.lr_scheduler.get_last_lr()[0]\n\n log_gpu_memory_usage(\"After offload weights\", logger=logger)\n\n step_loss = torch.tensor(step_loss).to(self.device_name)\n\n # compute time spent per step\n end_time = time.time()\n spend_time_per_step = end_time - start_time\n\n if is_cuda_available:\n torch.distributed.all_reduce(step_loss, op=torch.distributed.ReduceOp.AVG)\n elif is_npu_available:\n torch.distributed.all_reduce(step_loss)\n step_loss /= self.device_mesh.size(0)\n return {\n \"train/loss\": step_loss.detach().item(),\n \"train/lr(1e-3)\": lr * 1e3,\n \"train/time(s)\": spend_time_per_step,\n }\n\n def validation_step(self, batch: TensorDict):\n self.fsdp_model.eval()\n with torch.no_grad():\n loss = self._compute_loss_and_backward(batch, do_backward=False)\n if is_cuda_available:\n torch.distributed.all_reduce(loss, op=torch.distributed.ReduceOp.AVG)\n elif is_npu_available:\n torch.distributed.all_reduce(loss)\n loss /= self.device_mesh.size(0)\n return loss\n\n def save_checkpoint(self, step):\n \"\"\"Save checkpoint using FSDPCheckpointManager with improved tracking\"\"\"\n from verl.utils.fs import local_mkdir_safe\n\n # Determine checkpoint path\n local_global_step_folder = os.path.join(self.config.trainer.default_local_dir, f\"global_step_{step}\")\n\n if self.device_mesh.get_rank() == 0:\n print(f\"Saving checkpoint to: {local_global_step_folder}\")\n\n # Get max checkpoints to keep\n max_ckpt_to_keep = getattr(self.config.trainer, \"max_ckpt_to_keep\", None)\n\n # Use checkpoint manager to save\n self.checkpoint_manager.save_checkpoint(\n local_path=local_global_step_folder, global_step=step, max_ckpt_to_keep=max_ckpt_to_keep\n )\n\n # Save dataloader state\n if self.device_mesh.get_rank() == 0:\n local_mkdir_safe(local_global_step_folder)\n dataloader_local_path = os.path.join(local_global_step_folder, \"data.pt\")\n\n # Use StatefulDataLoader's built-in state dict functionality\n dataloader_state_dict = self.train_dataloader.state_dict()\n torch.save(dataloader_state_dict, dataloader_local_path)\n print(f\"Saved dataloader state to: {dataloader_local_path}\")\n\n # Update latest checkpoint tracker (atomic write)\n tracker_file = get_checkpoint_tracker_filename(self.config.trainer.default_local_dir)\n temp_tracker_file = tracker_file + \".tmp\"\n with open(temp_tracker_file, \"w\") as f:\n f.write(str(step))\n os.rename(temp_tracker_file, tracker_file)\n print(f\"Updated checkpoint tracker: {tracker_file}\")\n\n # Copy to HDFS if configured\n if self.device_mesh.get_rank() == 0 and getattr(self.config.trainer, \"default_hdfs_dir\", None):\n hdfs_io.makedirs(self.config.trainer.default_hdfs_dir, exist_ok=True)\n hdfs_io.copy(src=local_global_step_folder, dst=self.config.trainer.default_hdfs_dir, dirs_exist_ok=True)\n\n torch.distributed.barrier()\n\n def _init_checkpoint_manager(self):\n \"\"\"Initialize checkpoint manager with proper configuration\"\"\"\n # Get checkpoint configuration from config, with defaults\n checkpoint_config = getattr(self.config.trainer, \"checkpoint\", {})\n\n # Set default values if not specified\n save_contents = checkpoint_config.get(\"save_contents\", [\"model\", \"optimizer\", \"extra\"])\n load_contents = checkpoint_config.get(\"load_contents\", save_contents)\n\n # Create checkpoint config dict\n checkpoint_config_dict = {\n \"load_contents\": load_contents,\n \"save_contents\": save_contents,\n }\n\n # Convert to DictConfig for compatibility\n checkpoint_config_dict = DictConfig(checkpoint_config_dict)\n\n # Initialize checkpoint manager\n self.checkpoint_manager = FSDPCheckpointManager(\n model=self.fsdp_model,\n optimizer=self.optimizer,\n lr_scheduler=self.lr_scheduler,\n processing_class=self.tokenizer,\n checkpoint_config=checkpoint_config_dict,\n trust_remote_code=self.config.model.trust_remote_code,\n )\n\n def load_checkpoint(self):\n # Determine resume path based on configuration\n checkpoint_path = self._determine_resume_path()\n\n if checkpoint_path is None:\n return 0\n\n # extract resume step from checkpoint path\n resume_step = extract_step(checkpoint_path)\n if resume_step is None:\n log_with_rank(\n f\"Warning: Could not extract step number from {checkpoint_path}, starting from step 0\",\n logger=logger,\n rank=self.device_mesh.get_rank(),\n level=logging.WARNING,\n log_only_rank_0=True,\n )\n return 0\n self.resume_global_step = resume_step\n\n # Use checkpoint manager to load model state\n self.checkpoint_manager.load_checkpoint(checkpoint_path)\n log_with_rank(\n f\"Successfully loaded model checkpoint from {checkpoint_path} (step {resume_step})\",\n logger=logger,\n rank=self.device_mesh.get_rank(),\n log_only_rank_0=True,\n )\n\n # Always load dataloader state for StatefulDataLoader\n self._load_dataloader_state(checkpoint_path)\n\n return resume_step\n\n def _load_dataloader_state(self, checkpoint_path: str):\n \"\"\"Load dataloader state from checkpoint\"\"\"\n dataloader_path = os.path.join(checkpoint_path, \"data.pt\")\n\n if os.path.exists(dataloader_path):\n # Use StatefulDataLoader's built-in state dict functionality\n dataloader_state_dict = torch.load(dataloader_path, map_location=\"cpu\", weights_only=False)\n self.train_dataloader.load_state_dict(dataloader_state_dict)\n\n log_with_rank(\n f\"Successfully loaded dataloader state from {dataloader_path}\",\n logger=logger,\n rank=self.device_mesh.get_rank(),\n log_only_rank_0=True,\n )\n\n else:\n log_with_rank(\n f\"Warning: No dataloader state found at {dataloader_path}, will start from scratch\",\n logger=logger,\n rank=self.device_mesh.get_rank(),\n level=logging.WARNING,\n log_only_rank_0=True,\n )\n\n def _determine_resume_path(self):\n \"\"\"Determine the path to resume from based on resume_mode configuration\"\"\"\n resume_mode = getattr(self.config.trainer, \"resume_mode\", \"auto\")\n resume_from_path = getattr(self.config.trainer, \"resume_from_path\", None)\n\n if resume_mode == \"disable\":\n return None\n elif resume_mode == \"auto\":\n if resume_from_path is not None:\n assert os.path.exists(resume_from_path), (\n \"resume_from_path must be null or an existing path when resume_mode is 'auto'\"\n )\n assert \"global_step_\" in resume_from_path, \"resume_from_path must specify the global_steps\"\n return resume_from_path\n # Try to find the latest checkpoint in the default directory\n return self._find_latest_checkpoint()\n elif resume_mode == \"resume_path\":\n assert os.path.exists(resume_from_path), (\n \"resume_from_path must be an existing path when resume_mode is 'resume_path'\"\n )\n assert \"global_step_\" in resume_from_path, \"resume_from_path must specify the global_steps\"\n return resume_from_path\n else:\n raise ValueError(f\"Invalid resume_mode: {resume_mode}. Must be 'auto', 'disable', or 'resume_path'\")\n\n def _find_latest_checkpoint(self):\n \"\"\"Find the latest checkpoint in the default local directory\"\"\"\n checkpoint_dir = self.config.trainer.default_local_dir\n\n if not os.path.exists(checkpoint_dir):\n return None\n\n latest_checkpoint = find_latest_ckpt_path(checkpoint_dir)\n\n if latest_checkpoint and self.device_mesh.get_rank() == 0:\n step_num = extract_step(latest_checkpoint)\n print(f\"Found latest checkpoint: {latest_checkpoint} (step {step_num})\")\n\n return latest_checkpoint\n\n def fit(self):\n rank = self.device_mesh.get_rank()\n\n # TODO: add a unified tracking\n if rank == 0:\n tracking = Tracking(\n project_name=self.config.trainer.project_name,\n experiment_name=self.config.trainer.experiment_name,\n default_backend=self.config.trainer.logger,\n config=OmegaConf.to_container(self.config, resolve=True),\n )\n\n global_step = self.resume_global_step # Start from resumed step\n last_valid_metric = None\n # compute the total training steps.\n # the total training steps in SFT is mainly for early exit\n total_training_steps = len(self.train_dataloader) * self.config.trainer.total_epochs\n\n if self.config.trainer.total_training_steps is not None:\n total_training_steps = self.config.trainer.total_training_steps\n\n self.total_training_steps = total_training_steps\n log_with_rank(\n f\"Total training steps: {self.total_training_steps},\",\n logger=logger,\n rank=self.device_mesh.get_rank(),\n log_only_rank_0=True,\n )\n\n # With StatefulDataLoader, we don't need to manually calculate epochs and steps\n # The dataloader will automatically resume from where it left off\n if global_step > 0:\n log_with_rank(\n f\"StatefulDataLoader will automatically resume from global step: {global_step}\",\n logger=logger,\n rank=self.device_mesh.get_rank(),\n log_only_rank_0=True,\n )\n\n # Calculate which epoch we're starting from for sampler.set_epoch()\n start_epoch = global_step // self.steps_per_epoch\n\n train_time = 0\n for epoch in range(start_epoch, self.config.trainer.total_epochs):\n self.train_sampler.set_epoch(epoch=epoch)\n\n for step_in_epoch, data in enumerate(\n tqdm(\n self.train_dataloader,\n initial=global_step % self.steps_per_epoch if epoch == start_epoch else 0,\n total=self.steps_per_epoch,\n desc=f\"Epoch {epoch + 1}/{self.config.trainer.total_epochs}\",\n disable=rank != 0,\n )\n ):\n global_step += 1\n data = TensorDict(data, batch_size=self.config.data.train_batch_size).to(self.device_name)\n metric = self.training_step(data)\n train_time += metric[\"train/time(s)\"]\n if rank == 0:\n tracking.log(data=metric, step=global_step)\n\n is_last_step = global_step >= self.total_training_steps\n is_valid_step = global_step % self.config.trainer.test_freq == 0\n is_save_step = global_step % self.config.trainer.save_freq == 0\n\n # early exit or validation step\n if is_last_step or (self.config.trainer.test_freq > 0 and is_valid_step):\n # Perform validation\n val_losses = []\n for val_data in self.val_dataloader:\n val_data = TensorDict(val_data, batch_size=self.config.data.micro_batch_size_per_gpu).to(\n self.device_name\n )\n val_loss = self.validation_step(val_data)\n val_losses.append(val_loss)\n if rank == 0:\n val_loss = torch.mean(torch.stack(val_losses))\n metric = {\"val/loss\": val_loss.detach().item()}\n tracking.log(data=metric, step=global_step)\n last_valid_metric = metric\n torch.distributed.barrier()\n\n if is_last_step or (self.config.trainer.save_freq > 0 and is_save_step):\n self.save_checkpoint(step=global_step)\n\n if is_last_step:\n if rank == 0:\n print(f\"Total time for train steps: {train_time:.2f}s\")\n print(f\"Final validation metrics: {last_valid_metric}\")\n return\n\n\ndef run_sft(config):\n device_name = get_device_name()\n local_rank, rank, world_size = initialize_global_process_group()\n\n device_mesh = init_device_mesh(device_type=device_name, mesh_shape=(world_size,), mesh_dim_names=(\"fsdp\",))\n dp_size = world_size // config.ulysses_sequence_parallel_size\n ulysses_device_mesh = init_device_mesh(\n device_type=device_name,\n mesh_shape=(dp_size, config.ulysses_sequence_parallel_size),\n mesh_dim_names=(\"dp\", \"sp\"),\n )\n # build tokenizer and datasets first\n from verl.utils import hf_tokenizer\n\n local_model_path = copy_to_local(src=config.model.partial_pretrain, verbose=True)\n tokenizer = hf_tokenizer(local_model_path, trust_remote_code=config.model.trust_remote_code)\n train_dataset = create_sft_dataset(\n config.data.train_files, config.data, tokenizer, max_samples=config.data.get(\"train_max_samples\", -1)\n )\n val_dataset = create_sft_dataset(\n config.data.val_files, config.data, tokenizer, max_samples=config.data.get(\"val_max_samples\", -1)\n )\n\n trainer = FSDPSFTTrainer(\n config=config,\n device_mesh=device_mesh,\n ulysses_device_mesh=ulysses_device_mesh,\n tokenizer=tokenizer,\n train_dataset=train_dataset,\n val_dataset=val_dataset,\n )\n\n trainer.fit()\n\n destroy_global_process_group()\n\n\n@hydra.main(config_path=\"config\", config_name=\"sft_trainer\", version_base=None)\ndef main(config):\n # Automatically set `config.trainer.device = npu` when running on Ascend NPU.\n auto_set_device(config)\n\n run_sft(config)\n\n\ndef create_sft_dataset(data_paths, data_config, tokenizer, max_samples=-1):\n \"\"\"Create a dataset.\"\"\"\n # build dataset\n # First check if a custom dataset class is specified\n if data_config.custom_cls.get(\"path\", None):\n from verl.utils.import_utils import load_extern_object\n\n dataset_cls = load_extern_object(data_config.custom_cls.path, data_config.custom_cls.name)\n # Then check if multi-turn dataset should be used\n elif data_config.get(\"multiturn\", {}).get(\"enable\", False):\n dataset_cls = MultiTurnSFTDataset\n # Default to single-turn dataset\n else:\n dataset_cls = SFTDataset\n\n # Create datasets based on the selected class\n dataset = dataset_cls(parquet_files=data_paths, tokenizer=tokenizer, config=data_config, max_samples=max_samples)\n return dataset\n\n\nif __name__ == \"__main__\":\n main()\n"}54{"file_name": "verl__trainer__main_eval.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nOffline evaluate the performance of a generated file using reward model and ground truth verifier.\nThe input is a parquet file that contains N generated sequences and (optional) the ground truth.\n\n\"\"\"\n\nfrom collections import defaultdict\n\nimport hydra\nimport numpy as np\nimport pandas as pd\nimport ray\nfrom omegaconf import OmegaConf\nfrom tqdm import tqdm\n\nfrom verl.trainer.ppo.reward import get_custom_reward_fn\nfrom verl.utils.fs import copy_to_local\n\n\n@ray.remote\ndef process_item(config, data_source, response_lst, reward_data):\n reward_fn = get_custom_reward_fn(config)\n ground_truth = reward_data[\"ground_truth\"]\n score_lst = [reward_fn(data_source, r, ground_truth) for r in response_lst]\n return data_source, np.mean(score_lst)\n\n\n@hydra.main(config_path=\"config\", config_name=\"evaluation\", version_base=None)\ndef main(config):\n local_path = copy_to_local(config.data.path, use_shm=config.data.get(\"use_shm\", False))\n dataset = pd.read_parquet(local_path)\n responses = dataset[config.data.response_key]\n data_sources = dataset[config.data.data_source_key]\n reward_model_data = dataset[config.data.reward_model_key]\n\n total = len(dataset)\n\n # Initialize Ray\n if not ray.is_initialized():\n ray.init(**OmegaConf.to_container(config.ray_kwargs.get(\"ray_init\", {})))\n\n # evaluate test_score based on data source\n data_source_reward = defaultdict(list)\n # Create remote tasks\n remote_tasks = [\n process_item.remote(config, data_sources[i], responses[i], reward_model_data[i]) for i in range(total)\n ]\n\n # Process results as they come in\n with tqdm(total=total) as pbar:\n while len(remote_tasks) > 0:\n # Use ray.wait to get completed tasks\n done_ids, remote_tasks = ray.wait(remote_tasks)\n for result_id in done_ids:\n data_source, score = ray.get(result_id)\n data_source_reward[data_source].append(score)\n pbar.update(1)\n\n metric_dict = {}\n for data_source, rewards in data_source_reward.items():\n metric_dict[f\"test_score/{data_source}\"] = np.mean(rewards)\n\n print(metric_dict)\n\n\nif __name__ == \"__main__\":\n main()\n"}55{"file_name": "verl__trainer__main_generation.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nGenerate responses given a dataset of prompts\n\"\"\"\n\nimport os\n\nimport hydra\nimport numpy as np\nimport ray\n\nos.environ[\"NCCL_DEBUG\"] = \"WARN\"\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"true\"\n# os.environ['TORCH_COMPILE_DISABLE'] = '1'\n\nfrom pprint import pprint\n\nimport pandas as pd\nfrom omegaconf import OmegaConf\n\nfrom verl import DataProto\nfrom verl.protocol import pad_dataproto_to_divisor, unpad_dataproto\nfrom verl.single_controller.ray import RayClassWithInitArgs, RayResourcePool, RayWorkerGroup\nfrom verl.utils import hf_tokenizer\nfrom verl.utils.fs import copy_to_local\nfrom verl.utils.hdfs_io import makedirs\nfrom verl.utils.model import compute_position_id_with_mask\nfrom verl.workers.fsdp_workers import ActorRolloutRefWorker\n\n\n@hydra.main(config_path=\"config\", config_name=\"generation\", version_base=None)\ndef main(config):\n run_generation(config)\n\n\ndef run_generation(config) -> None:\n if not ray.is_initialized():\n # this is for local ray cluster\n default_runtime_env = {\"env_vars\": {\"TOKENIZERS_PARALLELISM\": \"true\", \"NCCL_DEBUG\": \"WARN\"}}\n ray_init_kwargs = config.ray_kwargs.get(\"ray_init\", {})\n runtime_env_kwargs = ray_init_kwargs.get(\"runtime_env\", {})\n runtime_env = OmegaConf.merge(default_runtime_env, runtime_env_kwargs)\n ray_init_kwargs = OmegaConf.create({**ray_init_kwargs, \"runtime_env\": runtime_env})\n print(f\"ray init kwargs: {ray_init_kwargs}\")\n ray.init(**OmegaConf.to_container(ray_init_kwargs))\n\n ray.get(main_task.remote(config))\n\n\n@ray.remote(num_cpus=1)\ndef main_task(config):\n pprint(OmegaConf.to_container(config, resolve=True)) # resolve=True will eval symbol values\n OmegaConf.resolve(config)\n\n local_path = copy_to_local(config.model.path)\n trust_remote_code = config.data.get(\"trust_remote_code\", False)\n tokenizer = hf_tokenizer(local_path, trust_remote_code=trust_remote_code)\n\n if config.rollout.temperature == 0.0:\n assert config.data.n_samples == 1, \"When temperature=0, n_samples must be 1.\"\n assert config.data.n_samples >= 1, \"n_samples should always >= 1\"\n\n # read dataset. Note that the dataset should directly contain chat template format (e.g., a list of dictionary)\n dataset = pd.read_parquet(config.data.path)\n chat_lst = dataset[config.data.prompt_key].tolist()\n\n chat_lst = [chat.tolist() for chat in chat_lst]\n\n tokenizer.padding_side = \"left\"\n if tokenizer.pad_token is None:\n tokenizer.pad_token = tokenizer.eos_token\n\n ray_cls_with_init = RayClassWithInitArgs(cls=ray.remote(ActorRolloutRefWorker), config=config, role=\"rollout\")\n resource_pool = RayResourcePool(process_on_nodes=[config.trainer.n_gpus_per_node] * config.trainer.nnodes)\n\n wg = RayWorkerGroup(\n resource_pool=resource_pool,\n ray_cls_with_init=ray_cls_with_init,\n device_name=config.trainer.device,\n )\n wg.init_model()\n\n total_samples = len(dataset)\n config_batch_size = config.data.batch_size\n apply_chat_template_kwargs = config.data.get(\"apply_chat_template_kwargs\", {})\n num_batch = -(-total_samples // config_batch_size)\n output_lst = [[] for _ in range(config.data.n_samples)]\n\n for batch_idx in range(num_batch):\n print(f\"[{batch_idx + 1}/{num_batch}] Start to process.\")\n batch_chat_lst = chat_lst[batch_idx * config_batch_size : (batch_idx + 1) * config_batch_size]\n inputs = tokenizer.apply_chat_template(\n batch_chat_lst,\n add_generation_prompt=True,\n padding=True,\n truncation=True,\n max_length=config.rollout.prompt_length,\n return_tensors=\"pt\",\n return_dict=True,\n tokenize=True,\n **apply_chat_template_kwargs,\n )\n input_ids = inputs[\"input_ids\"]\n attention_mask = inputs[\"attention_mask\"]\n position_ids = compute_position_id_with_mask(attention_mask)\n batch_dict = {\"input_ids\": input_ids, \"attention_mask\": attention_mask, \"position_ids\": position_ids}\n\n data = DataProto.from_dict(batch_dict)\n data_padded, pad_size = pad_dataproto_to_divisor(data, wg.world_size)\n\n # START TO GENERATE FOR n_samples TIMES\n print(f\"[{batch_idx + 1}/{num_batch}] Start to generate.\")\n for n_sample in range(config.data.n_samples):\n output_padded = wg.generate_sequences(data_padded)\n output = unpad_dataproto(output_padded, pad_size=pad_size)\n\n output_texts = []\n for i in range(len(output)):\n data_item = output[i]\n prompt_length = data_item.batch[\"prompts\"].shape[-1]\n valid_response_length = data_item.batch[\"attention_mask\"][prompt_length:].sum()\n valid_response_ids = data_item.batch[\"responses\"][:valid_response_length]\n response_str = tokenizer.decode(valid_response_ids, skip_special_tokens=True)\n output_texts.append(response_str)\n\n output_lst[n_sample].extend(output_texts)\n\n # convert output_lst from (n_samples, n_data) to (n_data, n_sampels)\n output_lst = np.array(output_lst, dtype=object)\n output_lst = np.transpose(output_lst, axes=(1, 0)).tolist()\n\n # add to the data frame\n dataset[\"responses\"] = output_lst\n\n # write to a new parquet\n output_dir = os.path.dirname(config.data.output_path)\n makedirs(output_dir, exist_ok=True)\n dataset.to_parquet(config.data.output_path)\n\n\nif __name__ == \"__main__\":\n main()\n"}56{"file_name": "verl__trainer__main_generation_server.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nGenerate responses given a dataset of prompts\n\"\"\"\n\nimport os\n\nimport aiohttp\nimport hydra\nimport numpy as np\nimport ray\n\nos.environ[\"NCCL_DEBUG\"] = \"WARN\"\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"true\"\n# os.environ['TORCH_COMPILE_DISABLE'] = '1'\n\nimport asyncio\nfrom pprint import pprint\n\nimport pandas as pd\nfrom omegaconf import OmegaConf\nfrom openai.types.chat import ChatCompletion\n\nfrom verl.utils.hdfs_io import makedirs\nfrom verl.workers.rollout.replica import get_rollout_replica_class\n\n\nasync def start_server(config):\n tp_size = config.actor_rollout_ref.rollout.tensor_model_parallel_size\n num_replicas = (config.trainer.n_gpus_per_node * config.trainer.nnodes) // tp_size\n rollout_config = config.actor_rollout_ref.rollout\n model_config = config.actor_rollout_ref.model\n # create standalone rollout server\n rollout_server_class = get_rollout_replica_class(config.actor_rollout_ref.rollout.name)\n rollout_servers = [\n rollout_server_class(\n replica_rank=replica_rank,\n config=rollout_config,\n model_config=model_config,\n gpus_per_node=config.trainer.n_gpus_per_node,\n )\n for replica_rank in range(num_replicas)\n ]\n await asyncio.gather(*[server.init_standalone() for server in rollout_servers])\n\n server_handles = [server._server_handle for server in rollout_servers]\n server_addresses = [server._server_address for server in rollout_servers]\n assert len(server_handles) == num_replicas\n assert len(server_addresses) == num_replicas\n\n return server_handles, server_addresses\n\n\nasync def submit_request(server_address, **chat_complete_request):\n try:\n extra_headers = chat_complete_request.pop(\"extra_headers\", {})\n timeout = aiohttp.ClientTimeout(total=None)\n session = aiohttp.ClientSession(timeout=timeout)\n async with session.post(\n url=f\"http://{server_address}/v1/chat/completions\",\n headers={\"Authorization\": \"Bearer token-abc123\", **extra_headers},\n json=chat_complete_request,\n ) as resp:\n data = await resp.json()\n return ChatCompletion(**data)\n finally:\n await session.close()\n\n\nasync def generate_per_replica(server_address, model_path: str, n_samples: int, sampling_params: dict, chat_lst: list):\n # here we should sample n_samples for each chat_lst.\n # we use aiohttp to avoid hang in AsyncOpenAI when the number of requests is large.\n\n # client = AsyncOpenAI(\n # api_key=\"123-abc\",\n # base_url=f\"http://{server_address}/v1\",\n # )\n\n chat_complete_request = [\n {\n \"model\": model_path,\n \"messages\": messages,\n **sampling_params,\n }\n for messages in chat_lst\n for _ in range(n_samples)\n ]\n\n tasks = [submit_request(server_address, **req) for req in chat_complete_request]\n results = await asyncio.gather(*tasks)\n return results\n\n\nasync def generate(\n server_addresses: list, model_path: str, n_samples: int, sampling_params: dict, chat_numpy: np.ndarray\n):\n num_replicas = len(server_addresses)\n chat_sub_array = np.array_split(chat_numpy, num_replicas)\n chat_sub_array = [chat.tolist() for chat in chat_sub_array]\n assert len(server_addresses) == len(chat_sub_array)\n results = await asyncio.gather(\n *[\n generate_per_replica(server_addresses[i], model_path, n_samples, sampling_params, chat_sub_array[i])\n for i in range(num_replicas)\n ]\n )\n return results\n\n\n@hydra.main(config_path=\"config\", config_name=\"ppo_trainer\", version_base=None)\ndef main(config):\n ray.init(runtime_env={\"env_vars\": {\"TOKENIZERS_PARALLELISM\": \"true\", \"NCCL_DEBUG\": \"WARN\", \"VLLM_USE_V1\": \"1\"}})\n\n pprint(OmegaConf.to_container(config, resolve=True)) # resolve=True will eval symbol values\n OmegaConf.resolve(config)\n\n n_samples = config.actor_rollout_ref.rollout.n\n\n if config.actor_rollout_ref.rollout.temperature == 0.0:\n assert n_samples == 1, \"When temperature=0, n_samples must be 1.\"\n assert n_samples >= 1, \"n_samples should always >= 1\"\n\n sampling_params = {\n \"temperature\": config.actor_rollout_ref.rollout.temperature,\n \"top_p\": config.actor_rollout_ref.rollout.top_p,\n # \"top_k\": config.actor_rollout_ref.rollout.top_k,\n \"max_tokens\": config.actor_rollout_ref.rollout.response_length,\n }\n\n from omegaconf import ListConfig\n\n train_files = config.data.train_files\n if not isinstance(train_files, list | ListConfig):\n train_files = [train_files]\n\n # read dataset. Note that the dataset should directly contain chat template format (e.g., a list of dictionary)\n\n datasets = []\n for train_file in train_files:\n dataset = pd.read_parquet(train_file)\n datasets.append(dataset)\n\n # concat dataset\n dataset = pd.concat(datasets, axis=0, ignore_index=True)\n chat_lst = dataset[config.data.prompt_key].tolist()\n chat_lst = [chat.tolist() for chat in chat_lst]\n chat_numpy = np.array(chat_lst)\n\n # start native server\n server_handles, server_addresses = asyncio.run(start_server(config))\n\n # run generate\n gen_results = asyncio.run(\n generate(server_addresses, config.actor_rollout_ref.model.path, n_samples, sampling_params, chat_numpy)\n )\n\n # reshape results into a numpy array\n import itertools\n\n results = list(itertools.chain.from_iterable(gen_results))\n\n # extract content from results\n results = np.array([result.choices[0].message.content for result in results])\n results = np.reshape(results, (-1, n_samples))\n\n assert results.shape == (len(chat_lst), n_samples)\n\n results = results.tolist()\n\n # add to the data frame\n dataset[\"responses\"] = results\n\n # write to a new parquet\n output_dir = os.path.dirname(config.data.output_path)\n makedirs(output_dir, exist_ok=True)\n print(f\"Saving results to {config.data.output_path}\")\n dataset.to_parquet(config.data.output_path)\n\n\nif __name__ == \"__main__\":\n main()\n"}57{"file_name": "verl__trainer__main_ppo.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nNote that we don't combine the main with ray_trainer as ray_trainer is used by other mpain.\n\"\"\"\n\nimport os\nimport socket\n\nimport hydra\nimport ray\nfrom omegaconf import OmegaConf\n\nfrom verl.experimental.dataset.sampler import AbstractSampler\nfrom verl.experimental.reward_loop import migrate_legacy_reward_impl\nfrom verl.trainer.constants_ppo import get_ppo_ray_runtime_env\nfrom verl.trainer.ppo.ray_trainer import RayPPOTrainer\nfrom verl.trainer.ppo.utils import need_critic, need_reference_policy\nfrom verl.utils.config import validate_config\nfrom verl.utils.device import auto_set_device, is_cuda_available\nfrom verl.utils.import_utils import load_extern_object\n\n\n@hydra.main(config_path=\"config\", config_name=\"ppo_trainer\", version_base=None)\ndef main(config):\n \"\"\"Main entry point for PPO training with Hydra configuration management.\n\n Args:\n config: Hydra configuration dictionary containing training parameters.\n \"\"\"\n # Automatically set `config.trainer.device = npu` when running on Ascend NPU.\n auto_set_device(config)\n config = migrate_legacy_reward_impl(config)\n run_ppo(config)\n\n\n# Define a function to run the PPO-like training process\ndef run_ppo(config, task_runner_class=None) -> None:\n \"\"\"Initialize Ray cluster and run distributed PPO training process.\n\n Args:\n config: Training configuration object containing all necessary parameters\n for distributed PPO training including Ray initialization settings,\n model paths, and training hyperparameters.\n task_runner_class: For recipe to change TaskRunner.\n \"\"\"\n # Check if Ray is not initialized\n if not ray.is_initialized():\n # Initialize Ray with a local cluster configuration\n # Set environment variables in the runtime environment to control tokenizer parallelism,\n # NCCL debug level, VLLM logging level, and allow runtime LoRA updating\n # `num_cpus` specifies the number of CPU cores Ray can use, obtained from the configuration\n default_runtime_env = get_ppo_ray_runtime_env()\n ray_init_kwargs = config.ray_kwargs.get(\"ray_init\", {})\n runtime_env_kwargs = ray_init_kwargs.get(\"runtime_env\", {})\n\n if config.transfer_queue.enable:\n # Add runtime environment variables for transfer queue\n runtime_env_vars = runtime_env_kwargs.get(\"env_vars\", {})\n runtime_env_vars[\"TRANSFER_QUEUE_ENABLE\"] = \"1\"\n runtime_env_kwargs[\"env_vars\"] = runtime_env_vars\n\n runtime_env = OmegaConf.merge(default_runtime_env, runtime_env_kwargs)\n ray_init_kwargs = OmegaConf.create({**ray_init_kwargs, \"runtime_env\": runtime_env})\n print(f\"ray init kwargs: {ray_init_kwargs}\")\n ray.init(**OmegaConf.to_container(ray_init_kwargs))\n\n if task_runner_class is None:\n task_runner_class = ray.remote(num_cpus=1)(TaskRunner) # please make sure main_task is not scheduled on head\n\n # Create a remote instance of the TaskRunner class, and\n # Execute the `run` method of the TaskRunner instance remotely and wait for it to complete\n if (\n is_cuda_available\n and config.global_profiler.tool == \"nsys\"\n and config.global_profiler.get(\"steps\") is not None\n and len(config.global_profiler.get(\"steps\", [])) > 0\n ):\n from verl.utils.import_utils import is_nvtx_available\n\n assert is_nvtx_available(), \"nvtx is not available in CUDA platform. Please 'pip3 install nvtx'\"\n nsight_options = OmegaConf.to_container(\n config.global_profiler.global_tool_config.nsys.controller_nsight_options\n )\n runner = task_runner_class.options(runtime_env={\"nsight\": nsight_options}).remote()\n else:\n runner = task_runner_class.remote()\n ray.get(runner.run.remote(config))\n\n # [Optional] get the path of the timeline trace file from the configuration, default to None\n # This file is used for performance analysis\n timeline_json_file = config.ray_kwargs.get(\"timeline_json_file\", None)\n if timeline_json_file:\n ray.timeline(filename=timeline_json_file)\n\n\nclass TaskRunner:\n \"\"\"Ray remote class for executing distributed PPO training tasks.\n\n This class encapsulates the main training logic and runs as a Ray remote actor\n to enable distributed execution across multiple nodes and GPUs.\n\n Attributes:\n role_worker_mapping: Dictionary mapping Role enums to Ray remote worker classes\n mapping: Dictionary mapping Role enums to resource pool IDs for GPU allocation\n \"\"\"\n\n def __init__(self):\n self.role_worker_mapping = {}\n self.mapping = {}\n\n def add_actor_rollout_worker(self, config):\n \"\"\"Add actor rollout worker based on the actor strategy.\"\"\"\n from verl.single_controller.ray import RayWorkerGroup\n from verl.trainer.ppo.ray_trainer import Role\n\n use_legacy_worker_impl = config.trainer.get(\"use_legacy_worker_impl\", \"auto\")\n\n # use new model engine implementation\n if use_legacy_worker_impl == \"disable\":\n from verl.workers.engine_workers import ActorRolloutRefWorker\n\n actor_rollout_cls = ActorRolloutRefWorker\n ray_worker_group_cls = RayWorkerGroup\n\n lora_rank = config.actor_rollout_ref.model.get(\"lora\", {}).get(\"rank\", 0)\n if lora_rank <= 0:\n lora_rank = config.actor_rollout_ref.model.get(\"lora_rank\", 0)\n ref_in_actor = lora_rank > 0 or config.actor_rollout_ref.model.get(\"lora_adapter_path\") is not None\n # NOTE: In new model engine, ref policy and actor rollout are in same ActorRolloutRefWorker,\n # while in legacy model engine, ref policy is in a separate ActorRolloutRefWorker.\n if need_reference_policy(config) and not ref_in_actor:\n role = Role.ActorRolloutRef\n else:\n role = Role.ActorRollout\n self.role_worker_mapping[role] = ray.remote(actor_rollout_cls)\n self.mapping[role] = \"global_pool\"\n return actor_rollout_cls, ray_worker_group_cls\n\n # Note: sync mode validation is now handled in RolloutConfig.__post_init__\n # Always use async worker since sync mode is deprecated and rejected\n if config.actor_rollout_ref.actor.strategy in {\"fsdp\", \"fsdp2\"}:\n from verl.workers.fsdp_workers import AsyncActorRolloutRefWorker\n\n actor_rollout_cls = AsyncActorRolloutRefWorker\n ray_worker_group_cls = RayWorkerGroup\n\n elif config.actor_rollout_ref.actor.strategy == \"megatron\":\n from verl.workers.megatron_workers import AsyncActorRolloutRefWorker\n\n actor_rollout_cls = AsyncActorRolloutRefWorker\n ray_worker_group_cls = RayWorkerGroup\n\n elif config.actor_rollout_ref.actor.strategy == \"veomni\":\n raise NotImplementedError(\"VeOmni does not support legacy worker implementation\")\n\n else:\n raise NotImplementedError\n\n self.role_worker_mapping[Role.ActorRollout] = ray.remote(actor_rollout_cls)\n self.mapping[Role.ActorRollout] = \"global_pool\"\n return actor_rollout_cls, ray_worker_group_cls\n\n def add_critic_worker(self, config):\n \"\"\"Add critic worker to role mapping.\"\"\"\n use_legacy_worker_impl = config.trainer.get(\"use_legacy_worker_impl\", \"auto\")\n if config.critic.strategy in {\"fsdp\", \"fsdp2\"}:\n if use_legacy_worker_impl in [\"auto\", \"enable\"]:\n from verl.workers.fsdp_workers import CriticWorker\n elif use_legacy_worker_impl == \"disable\":\n # we don't need to specialize critic worker. Just use TrainingWorker\n from verl.workers.engine_workers import TrainingWorker\n\n CriticWorker = TrainingWorker\n print(\"Using new worker implementation\")\n else:\n raise ValueError(f\"Invalid use_legacy_worker_impl: {use_legacy_worker_impl}\")\n\n elif config.critic.strategy == \"megatron\":\n # TODO: switch this to TrainingWorker as well\n from verl.workers.megatron_workers import CriticWorker\n\n elif config.critic.strategy == \"veomni\":\n if use_legacy_worker_impl == \"disable\":\n from verl.workers.engine_workers import TrainingWorker\n\n CriticWorker = TrainingWorker\n print(\"Using new worker implementation\")\n else:\n raise ValueError(f\"Invalid use_legacy_worker_impl: {use_legacy_worker_impl}\")\n\n else:\n raise NotImplementedError\n\n from verl.trainer.ppo.ray_trainer import Role\n\n self.role_worker_mapping[Role.Critic] = ray.remote(CriticWorker)\n self.mapping[Role.Critic] = \"global_pool\"\n\n def init_resource_pool_mgr(self, config):\n \"\"\"Initialize resource pool manager.\"\"\"\n\n global_pool_id = \"global_pool\"\n resource_pool_spec = {\n global_pool_id: [config.trainer.n_gpus_per_node] * config.trainer.nnodes,\n }\n\n if config.reward.reward_model.enable_resource_pool:\n if config.reward.reward_model.n_gpus_per_node <= 0:\n raise ValueError(\"config.reward.reward_model.n_gpus_per_node must be greater than 0\")\n if config.reward.reward_model.nnodes <= 0:\n raise ValueError(\"config.reward.reward_model.nnodes must be greater than 0\")\n\n reward_pool = [config.reward.reward_model.n_gpus_per_node] * config.reward.reward_model.nnodes\n resource_pool_spec[\"reward_pool\"] = reward_pool\n else:\n config.reward.reward_model.nnodes = config.trainer.nnodes\n config.reward.reward_model.n_gpus_per_node = config.trainer.n_gpus_per_node\n\n from verl.trainer.ppo.ray_trainer import ResourcePoolManager\n\n resource_pool_manager = ResourcePoolManager(resource_pool_spec=resource_pool_spec, mapping=self.mapping)\n return resource_pool_manager\n\n def add_reward_model_resource_pool(self, config):\n \"\"\"Add reward model worker if enabled.\"\"\"\n from verl.trainer.ppo.ray_trainer import Role\n\n if config.reward.reward_model.enable:\n # we do not use reward model workers, so we only register reward model in resource pool\n # without continue to register reward model worker in role mapping\n if config.reward.reward_model.enable_resource_pool:\n self.mapping[Role.RewardModel] = \"reward_pool\"\n else:\n self.mapping[Role.RewardModel] = \"global_pool\"\n\n def add_ref_policy_worker(self, config, ref_policy_cls):\n \"\"\"Add reference policy worker if KL loss or KL reward is used.\"\"\"\n from verl.trainer.ppo.ray_trainer import Role\n\n # Ref policy has been fused into ActorRolloutRefWorker in new model engine,\n # we don't need to add a separate ref policy worker group.\n use_legacy_worker_impl = config.trainer.get(\"use_legacy_worker_impl\", \"auto\")\n if use_legacy_worker_impl == \"disable\":\n return\n\n if need_reference_policy(config):\n self.role_worker_mapping[Role.RefPolicy] = ray.remote(ref_policy_cls)\n self.mapping[Role.RefPolicy] = \"global_pool\"\n\n def run(self, config):\n \"\"\"Execute the main PPO training workflow.\n\n This method sets up the distributed training environment, initializes\n workers, datasets, and reward functions, then starts the training process.\n\n Args:\n config: Training configuration object containing all parameters needed\n for setting up and running the PPO training process.\n \"\"\"\n # Print the initial configuration. `resolve=True` will evaluate symbolic values.\n from pprint import pprint\n\n from omegaconf import OmegaConf\n\n from verl.utils.fs import copy_to_local\n\n print(f\"TaskRunner hostname: {socket.gethostname()}, PID: {os.getpid()}\")\n pprint(OmegaConf.to_container(config, resolve=True))\n OmegaConf.resolve(config)\n\n actor_rollout_cls, ray_worker_group_cls = self.add_actor_rollout_worker(config)\n self.add_critic_worker(config)\n\n self.add_reward_model_resource_pool(config)\n\n # Add a reference policy worker if KL loss or KL reward is used.\n self.add_ref_policy_worker(config, actor_rollout_cls)\n\n # validate config\n validate_config(\n config=config,\n use_reference_policy=need_reference_policy(config),\n use_critic=need_critic(config),\n )\n\n # Download the checkpoint from HDFS to the local machine.\n # `use_shm` determines whether to use shared memory, which could lead to faster model loading if turned on\n local_path = copy_to_local(\n config.actor_rollout_ref.model.path, use_shm=config.actor_rollout_ref.model.get(\"use_shm\", False)\n )\n\n # Instantiate the tokenizer and processor.\n from verl.utils import hf_processor, hf_tokenizer\n\n trust_remote_code = config.data.get(\"trust_remote_code\", False)\n tokenizer = hf_tokenizer(local_path, trust_remote_code=trust_remote_code)\n # Used for multimodal LLM, could be None\n processor = hf_processor(local_path, trust_remote_code=trust_remote_code, use_fast=True)\n\n resource_pool_manager = self.init_resource_pool_mgr(config)\n\n from verl.utils.dataset.rl_dataset import collate_fn\n\n # Create training and validation datasets.\n train_dataset = create_rl_dataset(\n config.data.train_files,\n config.data,\n tokenizer,\n processor,\n is_train=True,\n max_samples=config.data.get(\"train_max_samples\", -1),\n )\n val_dataset = create_rl_dataset(\n config.data.val_files,\n config.data,\n tokenizer,\n processor,\n is_train=False,\n max_samples=config.data.get(\"val_max_samples\", -1),\n )\n train_sampler = create_rl_sampler(config.data, train_dataset)\n\n # Initialize the PPO trainer.\n trainer = RayPPOTrainer(\n config=config,\n tokenizer=tokenizer,\n processor=processor,\n role_worker_mapping=self.role_worker_mapping,\n resource_pool_manager=resource_pool_manager,\n ray_worker_group_cls=ray_worker_group_cls,\n train_dataset=train_dataset,\n val_dataset=val_dataset,\n collate_fn=collate_fn,\n train_sampler=train_sampler,\n )\n # Initialize the workers of the trainer.\n trainer.init_workers()\n\n # Start the training process.\n trainer.fit()\n\n\ndef create_rl_dataset(data_paths, data_config, tokenizer, processor, is_train=True, max_samples: int = -1):\n \"\"\"Create a dataset.\n\n Arguments:\n data_paths: List of paths to data files.\n data_config: The data config.\n tokenizer (Tokenizer): The tokenizer.\n processor (Processor): The processor.\n\n Returns:\n dataset (Dataset): The dataset.\n \"\"\"\n\n from verl.utils.dataset.rl_dataset import get_dataset_class\n\n # Get the dataset class\n dataset_cls = get_dataset_class(data_config)\n\n # Instantiate the dataset using the determined dataset class\n dataset = dataset_cls(\n data_files=data_paths,\n tokenizer=tokenizer,\n processor=processor,\n config=data_config,\n max_samples=max_samples,\n )\n\n return dataset\n\n\ndef create_rl_sampler(data_config, dataset):\n \"\"\"Create a sampler for the dataset.\n\n Arguments:\n data_config: The data config.\n dataset (Dataset): The dataset.\n\n Returns:\n sampler (Sampler): The sampler.\n \"\"\"\n import torch\n from torch.utils.data import SequentialSampler\n\n # torch.utils.data.RandomSampler could not recover properly\n from torchdata.stateful_dataloader.sampler import RandomSampler\n\n if data_config.sampler is not None and data_config.sampler.get(\"class_path\", None) is not None:\n curriculum_class = load_extern_object(\n data_config.sampler.class_path,\n data_config.sampler.class_name,\n )\n sampler = curriculum_class(\n data_source=dataset,\n data_config=data_config,\n )\n assert isinstance(sampler, AbstractSampler)\n assert data_config.get(\"dataloader_num_workers\", 8) == 0, (\n \"If using curriculum, num_workers must be 0 to prevent data caching. \"\n \"If the dataloader caches data before the batch is done the \"\n \"curriculum sampler won't have the opportunity to reorder it. \"\n )\n\n # Use a sampler to facilitate checkpoint resumption.\n # If shuffling is enabled in the data configuration, create a random sampler.\n elif data_config.shuffle:\n train_dataloader_generator = torch.Generator()\n seed = data_config.get(\"seed\")\n if seed is not None:\n train_dataloader_generator.manual_seed(seed)\n sampler = RandomSampler(data_source=dataset, generator=train_dataloader_generator)\n else:\n # If shuffling is disabled, use a sequential sampler to iterate through the dataset in order.\n sampler = SequentialSampler(data_source=dataset)\n\n return sampler\n\n\nif __name__ == \"__main__\":\n main()\n"}58{"file_name": "verl__trainer__ppo__metric_utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nMetrics related to the PPO trainer.\n\"\"\"\n\nfrom collections import defaultdict\nfrom functools import partial\nfrom typing import Any, Callable\n\nimport numpy as np\nimport torch\n\nimport verl.utils.torch_functional as verl_F\nfrom verl import DataProto\nfrom verl.utils.import_utils import deprecated\n\n\n@deprecated(\"verl.utils.metric.reduce_metrics\")\ndef reduce_metrics(metrics: dict[str, list[Any]]) -> dict[str, Any]:\n \"\"\"\n Reduces a dictionary of metric lists by computing the mean of each list.\n\n Args:\n metrics: A dictionary mapping metric names to lists of metric values.\n\n Returns:\n A dictionary with the same keys but with each list replaced by its mean value.\n\n Example:\n >>> metrics = {\"loss\": [1.0, 2.0, 3.0], \"accuracy\": [0.8, 0.9, 0.7]}\n >>> reduce_metrics(metrics)\n {\"loss\": 2.0, \"accuracy\": 0.8}\n \"\"\"\n from verl.utils.metric import reduce_metrics\n\n return reduce_metrics(metrics)\n\n\ndef _compute_response_info(batch: DataProto) -> dict[str, Any]:\n \"\"\"\n Computes information about prompts and responses from a batch.\n\n This is an internal helper function that extracts masks and lengths for prompts and responses.\n\n Args:\n batch: A DataProto object containing batch data with responses and attention masks.\n\n Returns:\n A dictionary containing:\n - response_mask: Attention mask for the response tokens\n - prompt_length: Tensor of prompt lengths for each item in the batch\n - response_length: Tensor of response lengths for each item in the batch\n \"\"\"\n response_length = batch.batch[\"responses\"].shape[-1]\n\n prompt_mask = batch.batch[\"attention_mask\"][:, :-response_length]\n response_mask = batch.batch[\"attention_mask\"][:, -response_length:]\n\n prompt_length = prompt_mask.sum(-1).float()\n response_length = response_mask.sum(-1).float() # (batch_size,)\n\n return dict(\n response_mask=response_mask,\n prompt_length=prompt_length,\n response_length=response_length,\n )\n\n\ndef compute_data_metrics(batch: DataProto, use_critic: bool = True) -> dict[str, Any]:\n \"\"\"\n Computes various metrics from a batch of data for PPO training.\n\n This function calculates metrics related to scores, rewards, advantages, returns, values,\n and sequence lengths from a batch of data. It provides statistical information (mean, max, min)\n for each metric category.\n\n Args:\n batch: A DataProto object containing batch data with token-level scores, rewards, advantages, etc.\n use_critic: Whether to include critic-specific metrics. Defaults to True.\n\n Returns:\n A dictionary of metrics including:\n - critic/score/mean, max, min: Statistics about sequence scores\n - critic/rewards/mean, max, min: Statistics about sequence rewards\n - critic/advantages/mean, max, min: Statistics about advantages\n - critic/returns/mean, max, min: Statistics about returns\n - critic/values/mean, max, min: Statistics about critic values (if use_critic=True)\n - critic/vf_explained_var: Explained variance of the value function (if use_critic=True)\n - response_length/mean, max, min, clip_ratio: Statistics about response lengths\n - prompt_length/mean, max, min, clip_ratio: Statistics about prompt lengths\n - num_turns/mean, max, min: Statistics about the number of multi-turn conversations\n \"\"\"\n sequence_score = batch.batch[\"token_level_scores\"].sum(-1)\n sequence_reward = batch.batch[\"token_level_rewards\"].sum(-1)\n\n advantages = batch.batch[\"advantages\"]\n returns = batch.batch[\"returns\"]\n\n max_response_length = batch.batch[\"responses\"].shape[-1]\n\n prompt_mask = batch.batch[\"attention_mask\"][:, :-max_response_length].bool()\n response_mask = batch.batch[\"response_mask\"].bool()\n\n max_prompt_length = prompt_mask.size(-1)\n\n response_info = _compute_response_info(batch)\n prompt_length = response_info[\"prompt_length\"]\n response_length = response_info[\"response_length\"]\n\n aborted_mask = (response_length == 0).bool()\n non_aborted_mask = ~aborted_mask\n\n non_aborted_sequence_score = sequence_score[non_aborted_mask]\n non_aborted_sequence_reward = sequence_reward[non_aborted_mask]\n\n score_mean = torch.mean(non_aborted_sequence_score).detach().item()\n score_max = torch.max(non_aborted_sequence_score).detach().item()\n score_min = torch.min(non_aborted_sequence_score).detach().item()\n\n reward_mean = torch.mean(non_aborted_sequence_reward).detach().item()\n reward_max = torch.max(non_aborted_sequence_reward).detach().item()\n reward_min = torch.min(non_aborted_sequence_reward).detach().item()\n\n valid_adv = torch.masked_select(advantages, response_mask)\n valid_returns = torch.masked_select(returns, response_mask)\n\n if use_critic:\n values = batch.batch[\"values\"]\n valid_values = torch.masked_select(values, response_mask)\n return_diff_var = torch.var(valid_returns - valid_values)\n return_var = torch.var(valid_returns)\n\n # Aborted samples and non-aborted response length statistics\n # response_length_non_aborted/*: statistics computed on non-aborted samples only\n aborted_ratio = torch.mean(aborted_mask.float()).detach().item()\n\n non_aborted_response_length = response_length[non_aborted_mask]\n if non_aborted_response_length.numel() > 0:\n non_aborted_response_length_mean = torch.mean(non_aborted_response_length).detach().item()\n non_aborted_response_length_max = torch.max(non_aborted_response_length).detach().item()\n non_aborted_response_length_min = torch.min(non_aborted_response_length).detach().item()\n non_aborted_response_length_clip_ratio = (\n torch.mean(torch.eq(non_aborted_response_length, max_response_length).float()).detach().item()\n )\n else:\n raise ValueError(\"All samples are aborted, this should not happen.\")\n\n metrics = {\n # score\n \"critic/score/mean\": score_mean,\n \"critic/score/max\": score_max,\n \"critic/score/min\": score_min,\n # reward\n \"critic/rewards/mean\": reward_mean,\n \"critic/rewards/max\": reward_max,\n \"critic/rewards/min\": reward_min,\n # adv\n \"critic/advantages/mean\": torch.mean(valid_adv).detach().item(),\n \"critic/advantages/max\": torch.max(valid_adv).detach().item(),\n \"critic/advantages/min\": torch.min(valid_adv).detach().item(),\n # returns\n \"critic/returns/mean\": torch.mean(valid_returns).detach().item(),\n \"critic/returns/max\": torch.max(valid_returns).detach().item(),\n \"critic/returns/min\": torch.min(valid_returns).detach().item(),\n **(\n {\n # values\n \"critic/values/mean\": torch.mean(valid_values).detach().item(),\n \"critic/values/max\": torch.max(valid_values).detach().item(),\n \"critic/values/min\": torch.min(valid_values).detach().item(),\n # vf explained var\n \"critic/vf_explained_var\": (1.0 - return_diff_var / (return_var + 1e-5)).detach().item(),\n }\n if use_critic\n else {}\n ),\n # response length\n \"response_length/mean\": torch.mean(response_length).detach().item(),\n \"response_length/max\": torch.max(response_length).detach().item(),\n \"response_length/min\": torch.min(response_length).detach().item(),\n \"response_length/clip_ratio\": torch.mean(torch.eq(response_length, max_response_length).float())\n .detach()\n .item(),\n # response length (non-aborted only)\n # These statistics exclude aborted samples to avoid skew from zeros\n \"response_length_non_aborted/mean\": non_aborted_response_length_mean,\n \"response_length_non_aborted/max\": non_aborted_response_length_max,\n \"response_length_non_aborted/min\": non_aborted_response_length_min,\n \"response_length_non_aborted/clip_ratio\": non_aborted_response_length_clip_ratio,\n # aborted ratio\n # Fraction of samples whose response length is zero\n \"response/aborted_ratio\": aborted_ratio,\n # prompt length\n \"prompt_length/mean\": torch.mean(prompt_length).detach().item(),\n \"prompt_length/max\": torch.max(prompt_length).detach().item(),\n \"prompt_length/min\": torch.min(prompt_length).detach().item(),\n \"prompt_length/clip_ratio\": torch.mean(torch.eq(prompt_length, max_prompt_length).float()).detach().item(),\n }\n\n # multi-turn conversation\n if \"__num_turns__\" in batch.non_tensor_batch:\n num_turns = batch.non_tensor_batch[\"__num_turns__\"]\n metrics[\"num_turns/min\"] = num_turns.min()\n metrics[\"num_turns/max\"] = num_turns.max()\n metrics[\"num_turns/mean\"] = num_turns.mean()\n\n if \"tool_call_counts\" in batch.non_tensor_batch:\n tool_call_counts = batch.non_tensor_batch[\"tool_call_counts\"]\n metrics[\"tool_call_counts/min\"] = tool_call_counts.min()\n metrics[\"tool_call_counts/max\"] = tool_call_counts.max()\n metrics[\"tool_call_counts/mean\"] = tool_call_counts.mean()\n\n return metrics\n\n\ndef compute_timing_metrics(batch: DataProto, timing_raw: dict[str, float]) -> dict[str, Any]:\n \"\"\"\n Computes timing metrics for different processing stages in PPO training.\n\n This function calculates both raw timing metrics (in seconds) and per-token timing metrics\n (in milliseconds) for various processing stages like generation, reference computation,\n value computation, advantage computation, and model updates.\n\n Args:\n batch: A DataProto object containing batch data with responses and attention masks.\n timing_raw: A dictionary mapping stage names to their execution times in seconds.\n\n Returns:\n A dictionary containing:\n - timing_s/{name}: Raw timing in seconds for each stage\n - timing_per_token_ms/{name}: Per-token timing in milliseconds for each stage\n\n Note:\n Different stages use different token counts for normalization:\n - \"gen\" uses only response tokens\n - Other stages (\"ref\", \"values\", \"adv\", \"update_critic\", \"update_actor\") use all tokens\n (prompt + response)\n \"\"\"\n response_info = _compute_response_info(batch)\n num_prompt_tokens = torch.sum(response_info[\"prompt_length\"]).item()\n num_response_tokens = torch.sum(response_info[\"response_length\"]).item()\n num_overall_tokens = num_prompt_tokens + num_response_tokens\n\n num_tokens_of_section = {\n \"gen\": num_response_tokens,\n **{name: num_overall_tokens for name in [\"ref\", \"values\", \"adv\", \"update_critic\", \"update_actor\"]},\n }\n\n return {\n **{f\"timing_s/{name}\": value for name, value in timing_raw.items()},\n **{\n f\"timing_per_token_ms/{name}\": timing_raw[name] * 1000 / num_tokens_of_section[name]\n for name in set(num_tokens_of_section.keys()) & set(timing_raw.keys())\n },\n }\n\n\ndef compute_throughout_metrics(batch: DataProto, timing_raw: dict[str, float], n_gpus: int) -> dict[str, Any]:\n \"\"\"\n Computes throughput metrics for PPO training.\n\n This function calculates performance metrics related to token processing speed,\n including the total number of tokens processed, time per step, and throughput\n (tokens per second per GPU).\n\n Args:\n batch: A DataProto object containing batch data with meta information about token counts.\n timing_raw: A dictionary mapping stage names to their execution times in seconds.\n Must contain a \"step\" key with the total step time.\n n_gpus: Number of GPUs used for training.\n\n Returns:\n A dictionary containing:\n - perf/total_num_tokens: Total number of tokens processed in the batch\n - perf/time_per_step: Time taken for the step in seconds\n - perf/throughput: Tokens processed per second per GPU\n\n Note:\n The throughput is calculated as total_tokens / (time * n_gpus) to normalize\n across different GPU counts.\n \"\"\"\n total_num_tokens = sum(batch.meta_info[\"global_token_num\"])\n time = timing_raw[\"step\"]\n # estimated_flops, promised_flops = flops_function.estimate_flops(num_tokens, time)\n # f'Actual TFLOPs/s/GPU': estimated_flops/(n_gpus),\n # f'Theoretical TFLOPs/s/GPU': promised_flops,\n return {\n \"perf/total_num_tokens\": total_num_tokens,\n \"perf/time_per_step\": time,\n \"perf/throughput\": total_num_tokens / (time * n_gpus),\n }\n\n\ndef compute_variance_proxy_metrics(batch: DataProto, gradient_norm: float = None) -> dict[str, float]:\n \"\"\"\n Compute variance proxy metrics using the simplified expected squared norm approach.\n\n This metric provides a computationally efficient way to monitor gradient variance\n during training. It works for any advantage estimator as long as sum_pi_squared\n is available from the actor.\n\n Theory:\n - Full variance: Var(g̃) = E[||g̃||²] - ||g_true||²\n - Simplified proxy (when ||g_true||² ≈ 0): Var(g̃) ≈ E[||g̃||²]\n - Using W-score approximation: E[||g̃||²] ≈ E[A² × W(τ)]\n\n Where W(τ) = Σ_t[1 - 2π_t(y_t) + Σπ²] is the score-norm proxy.\n \"\"\"\n metrics = {}\n\n # Check if we have the necessary data (sum_pi_squared is required for W-score)\n if \"sum_pi_squared\" not in batch.batch or \"old_log_probs\" not in batch.batch or \"advantages\" not in batch.batch:\n return metrics\n\n # Compute W(τ) = Σ_t[1 - 2π_t(y_t) + Σπ²]\n pi_t = torch.exp(batch.batch[\"old_log_probs\"])\n w_per_timestep = 1 - 2 * pi_t + batch.batch[\"sum_pi_squared\"]\n\n # Get response mask to only consider valid tokens\n response_mask = batch.batch[\"response_mask\"]\n\n # Use pre-computed rollout IS weights from batch (for variance proxy consistency with training loss)\n # IS weights are computed centrally in ray_trainer.py to avoid duplication\n rollout_is_weights = None\n if \"rollout_is_weights\" in batch.batch:\n # Extract pre-computed IS weights from batch (already computed in trainer)\n rollout_is_weights = batch.batch[\"rollout_is_weights\"]\n\n # Scale W by (rollout IS weight)² for optimal baseline under biased estimation\n w_per_timestep = w_per_timestep * (rollout_is_weights**2).detach()\n\n # Note: IS weight statistics and mismatch metrics are logged in ray_trainer.py\n\n # Get scalar advantages (mean over timesteps)\n advantages = batch.batch[\"advantages\"]\n # Compute mean advantage per trajectory using masked_mean\n advantages_scalar = verl_F.masked_mean(advantages, response_mask, axis=-1)\n\n # Compute W values (sum over timesteps)\n w_values = verl_F.masked_sum(w_per_timestep, response_mask, axis=-1)\n\n # ====== COMPUTE VARIANCE PROXIES ======\n # Variance proxy should match the actual gradient computation:\n # - If IS weights were computed/applied: use them in variance proxy calculation\n # - Otherwise: compute on-policy variance proxy\n\n # ====== PROXY 1: Signal Strength ||ḡ||² ======\n # The squared norm of the mean gradient (provided from training loop)\n proxy1_signal_strength = gradient_norm**2 if gradient_norm is not None else None\n\n # ====== PROXY 2: Total Power E[||ĝ_τ||²] ======\n # Measures the average of squared gradient norms (Signal + Noise)\n if rollout_is_weights is not None:\n # Off-policy with IS correction applied: use clamped weights consistently with actual gradient computation\n rollout_is_weights_scalar = verl_F.masked_mean(rollout_is_weights, response_mask, axis=-1)\n # Recover original W (before IS correction was applied in line 657)\n # Clamp to avoid division by zero when IS weights are zero\n w_original = verl_F.masked_sum(\n w_per_timestep / torch.clamp((rollout_is_weights**2).detach(), min=1e-10), response_mask, axis=-1\n )\n # Clamp W to avoid negative values (which would cause NaN in sqrt)\n w_original = torch.clamp(w_original, min=0.0)\n # Proxy 2 for off-policy: E[ρ̄² × A² × W]\n proxy2_total_power = ((rollout_is_weights_scalar**2) * (advantages_scalar**2) * w_original).mean()\n\n else:\n # On-policy Proxy 2: E[A² × W]\n # Clamp W to avoid negative values (which would cause NaN in sqrt)\n w_values_clamped = torch.clamp(w_values, min=0.0)\n proxy2_total_power = (advantages_scalar**2 * w_values_clamped).mean()\n\n # ====== PROXY 3: Pure Noise - Variance of Mean Vector ======\n # Requires ||ḡ||² from actual batch gradient\n # Formula: (1/(N-1)) × (Proxy2 - Proxy1)\n proxy3_pure_noise = None\n if proxy1_signal_strength is not None:\n batch_size = advantages_scalar.shape[0]\n if batch_size > 1:\n proxy3_pure_noise = (1.0 / (batch_size - 1)) * (proxy2_total_power - proxy1_signal_strength)\n # Ensure non-negative (can be negative due to numerical errors)\n proxy3_pure_noise = max(\n 0.0, proxy3_pure_noise.item() if torch.is_tensor(proxy3_pure_noise) else proxy3_pure_noise\n )\n\n # Decompose into components for analysis\n expected_a_squared = (advantages_scalar**2).mean()\n expected_w = w_values.mean()\n\n metrics.update(\n {\n # Proxy 1: Signal Strength ||ḡ||²\n \"variance_proxy/proxy1_signal_strength\": (\n proxy1_signal_strength if proxy1_signal_strength is not None else 0.0\n ),\n # Proxy 2: Total Power E[||ĝ_τ||²]\n \"variance_proxy/proxy2_total_power\": proxy2_total_power.detach().item(),\n # Proxy 3: Pure Noise - Variance of Mean Vector\n \"variance_proxy/proxy3_pure_noise\": proxy3_pure_noise if proxy3_pure_noise is not None else 0.0,\n # Component metrics for debugging\n \"variance_proxy/expected_a_squared\": expected_a_squared.detach().item(),\n \"variance_proxy/expected_w\": expected_w.detach().item(),\n }\n )\n\n return metrics\n\n\ndef bootstrap_metric(\n data: list[Any],\n subset_size: int,\n reduce_fns: list[Callable[[np.ndarray], float]],\n n_bootstrap: int = 1000,\n seed: int = 42,\n) -> list[tuple[float, float]]:\n \"\"\"\n Performs bootstrap resampling to estimate statistics of metrics.\n\n This function uses bootstrap resampling to estimate the mean and standard deviation\n of metrics computed by the provided reduction functions on random subsets of the data.\n\n Args:\n data: List of data points to bootstrap from.\n subset_size: Size of each bootstrap sample.\n reduce_fns: List of functions that compute a metric from a subset of data.\n n_bootstrap: Number of bootstrap iterations. Defaults to 1000.\n seed: Random seed for reproducibility. Defaults to 42.\n\n Returns:\n A list of tuples, where each tuple contains (mean, std) for a metric\n corresponding to each reduction function in reduce_fns.\n\n Example:\n >>> data = [1, 2, 3, 4, 5]\n >>> reduce_fns = [np.mean, np.max]\n >>> bootstrap_metric(data, 3, reduce_fns)\n [(3.0, 0.5), (4.5, 0.3)] # Example values\n \"\"\"\n np.random.seed(seed)\n data_np = np.array(data, dtype=object)\n n_data = len(data_np)\n\n # generate bootstrap indices, shape: (n_bootstrap, subset_size)\n bootstrap_idxs = np.random.choice(n_data, size=(n_bootstrap, subset_size), replace=True)\n\n # pre-allocate result array, shape: (n_fns, n_bootstrap)\n n_fns = len(reduce_fns)\n metric_results = np.empty((n_fns, n_bootstrap), dtype=np.float64)\n\n # compute metric results for each bootstrap sample\n for fn_idx, reduce_fn in enumerate(reduce_fns):\n # bootstrap sample and compute metric\n for boot_idx in range(n_bootstrap):\n sample = data_np[bootstrap_idxs[boot_idx]]\n metric_results[fn_idx, boot_idx] = reduce_fn(sample)\n\n # compute mean and std for each metric function\n result = [\n (float(np.mean(metric_results[fn_idx])), float(np.std(metric_results[fn_idx]))) for fn_idx in range(n_fns)\n ]\n return result\n\n\ndef calc_maj_val(data: list[dict[str, Any]], vote_key: str, val_key: str) -> float:\n \"\"\"\n Calculate a value based on majority voting.\n\n This function identifies the most common value for a specified vote key\n in the data, then returns the corresponding value for that majority vote.\n\n Args:\n data: List of dictionaries, where each dictionary contains both vote_key and val_key.\n vote_key: The key in each dictionary used for voting/counting.\n val_key: The key in each dictionary whose value will be returned for the majority vote.\n\n Returns:\n The value associated with the most common vote.\n\n Example:\n >>> data = [\n ... {\"pred\": \"A\", \"val\": 0.9},\n ... {\"pred\": \"B\", \"val\": 0.8},\n ... {\"pred\": \"A\", \"val\": 0.7}\n ... ]\n >>> calc_maj_val(data, vote_key=\"pred\", val_key=\"val\")\n 0.9 # Returns the first \"val\" for the majority vote \"A\"\n \"\"\"\n vote2vals = defaultdict(list)\n for d in data:\n vote2vals[d[vote_key]].append(d[val_key])\n\n vote2cnt = {k: len(v) for k, v in vote2vals.items()}\n maj_vote = max(vote2cnt, key=vote2cnt.get)\n\n maj_val = vote2vals[maj_vote][0]\n\n return maj_val\n\n\ndef process_validation_metrics(\n data_sources: list[str], sample_uids: list[str], infos_dict: dict[str, list[Any]], seed: int = 42\n) -> dict[str, dict[str, dict[str, float]]]:\n \"\"\"\n Process validation metrics into a structured format with statistical analysis.\n\n This function organizes validation metrics by data source and prompt, then computes\n various statistical measures including means, standard deviations, best/worst values,\n and majority voting results. It also performs bootstrap sampling to estimate statistics\n for different sample sizes.\n\n Args:\n data_sources: List of data source identifiers for each sample.\n sample_uids: List of sample uids corresponding to each sample.\n infos_dict: Dictionary mapping variable names to lists of values for each sample.\n seed: Random seed for bootstrap sampling. Defaults to 42.\n\n Returns:\n A nested dictionary with the structure:\n {\n data_source: {\n variable_name: {\n metric_name: value\n }\n }\n }\n\n Where metric_name includes:\n - \"mean@N\": Mean value across N samples\n - \"std@N\": Standard deviation across N samples\n - \"best@N/mean\": Mean of the best values in bootstrap samples of size N\n - \"best@N/std\": Standard deviation of the best values in bootstrap samples\n - \"worst@N/mean\": Mean of the worst values in bootstrap samples\n - \"worst@N/std\": Standard deviation of the worst values in bootstrap samples\n - \"maj@N/mean\": Mean of majority voting results in bootstrap samples (if \"pred\" exists)\n - \"maj@N/std\": Standard deviation of majority voting results (if \"pred\" exists)\n\n Example:\n >>> data_sources = [\"source1\", \"source1\", \"source2\"]\n >>> sample_uids = [\"uid1\", \"uid1\", \"uid2\"]\n >>> infos_dict = {\"score\": [0.8, 0.9, 0.7], \"pred\": [\"A\", \"A\", \"B\"]}\n >>> result = process_validation_metrics(data_sources, sample_uids, infos_dict)\n >>> # result will contain statistics for each data source and variable\n \"\"\"\n # Group metrics by data source, prompt and variable\n data_src2uid2var2vals = defaultdict(lambda: defaultdict(lambda: defaultdict(list)))\n for sample_idx, data_source in enumerate(data_sources):\n uid = sample_uids[sample_idx]\n var2vals = data_src2uid2var2vals[data_source][uid]\n for var_name, var_vals in infos_dict.items():\n var2vals[var_name].append(var_vals[sample_idx])\n\n np_mean = np.mean\n np_std = np.std\n reduce_fns_best_worst = [np.max, np.min]\n n_bootstrap = 1000\n\n # 2. cache ns list\n def gen_ns(n_resps: int) -> list[int]:\n if n_resps <= 1:\n return []\n ns = []\n n = 2\n while n < n_resps:\n ns.append(n)\n n *= 2\n ns.append(n_resps)\n return ns\n\n ns_cache = {}\n\n # 3. cache metric results\n data_src2uid2var2metric = {}\n\n # 4. flatten loop\n for data_source, uid2var2vals in data_src2uid2var2vals.items():\n # create uid dict\n uid_dict = data_src2uid2var2metric.setdefault(data_source, {})\n\n for uid, var2vals in uid2var2vals.items():\n pred_vals = var2vals.get(\"pred\")\n has_pred = pred_vals is not None\n var_dict = uid_dict.setdefault(uid, {})\n\n for var_name, var_vals in var2vals.items():\n # skip empty or string values\n if not var_vals or isinstance(var_vals[0], str):\n continue\n\n # compute mean and std\n n_resps = len(var_vals)\n metric = {f\"mean@{n_resps}\": float(np_mean(var_vals))}\n\n if n_resps > 1:\n metric[f\"std@{n_resps}\"] = float(np_std(var_vals))\n\n # cache ns list\n if n_resps not in ns_cache:\n ns_cache[n_resps] = gen_ns(n_resps)\n ns = ns_cache[n_resps]\n\n # compute best/worst metrics\n for n in ns:\n # compute best/worst metrics\n (bon_mean, bon_std), (won_mean, won_std) = bootstrap_metric(\n data=var_vals,\n subset_size=n,\n reduce_fns=reduce_fns_best_worst,\n n_bootstrap=n_bootstrap,\n seed=seed,\n )\n metric[f\"best@{n}/mean\"] = bon_mean\n metric[f\"best@{n}/std\"] = bon_std\n metric[f\"worst@{n}/mean\"] = won_mean\n metric[f\"worst@{n}/std\"] = won_std\n\n # compute maj metrics\n if has_pred:\n # create vote_data\n vote_data = [\n {\"val\": val, \"pred\": pred} for val, pred in zip(var_vals, pred_vals, strict=True)\n ]\n # compute maj metrics\n [(maj_n_mean, maj_n_std)] = bootstrap_metric(\n data=vote_data,\n subset_size=n,\n reduce_fns=[partial(calc_maj_val, vote_key=\"pred\", val_key=\"val\")],\n n_bootstrap=n_bootstrap,\n seed=seed,\n )\n metric[f\"maj@{n}/mean\"] = maj_n_mean\n metric[f\"maj@{n}/std\"] = maj_n_std\n\n var_dict[var_name] = metric\n\n # Aggregate metrics across uids\n data_src2var2metric2uid_vals = defaultdict(lambda: defaultdict(lambda: defaultdict(list)))\n for data_source, uid2var2metric in data_src2uid2var2metric.items():\n for uid, var2metric in uid2var2metric.items():\n for var_name, metric in var2metric.items():\n for metric_name, metric_val in metric.items():\n data_src2var2metric2uid_vals[data_source][var_name][metric_name].append(metric_val)\n\n data_src2var2metric2val = defaultdict(lambda: defaultdict(lambda: defaultdict(float)))\n for data_source, var2metric2uid_vals in data_src2var2metric2uid_vals.items():\n for var_name, metric2uid_vals in var2metric2uid_vals.items():\n for metric_name, uid_vals in metric2uid_vals.items():\n data_src2var2metric2val[data_source][var_name][metric_name] = np.mean(uid_vals)\n return data_src2var2metric2val\n"}59{"file_name": "verl__trainer__ppo__prefix_grouper_utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom __future__ import annotations\n\nimport torch\nfrom prefix_grouper import PrefixGrouper\n\nfrom verl.utils.torch_functional import logprobs_from_logits\n\n\ndef build_position_ids_for_prefix_grouper(prefix_grouper: PrefixGrouper) -> torch.Tensor:\n \"\"\"Build position_ids for PrefixGrouper where each response restarts from prefix_len.\"\"\"\n num_samples = len(prefix_grouper.group_info)\n max_len = prefix_grouper.padding_mask.size(1)\n device = prefix_grouper.padding_mask.device\n\n position_ids = torch.zeros(num_samples, max_len, dtype=torch.long, device=device)\n\n for i, group in enumerate(prefix_grouper.group_info):\n prefix_len = group.prefix_len\n\n position_ids[i, :prefix_len] = torch.arange(prefix_len, device=device)\n cur_pos = prefix_len\n for suffix_len in group.suffix_lens:\n if suffix_len > 0:\n position_ids[i, cur_pos : cur_pos + suffix_len] = torch.arange(\n prefix_len, prefix_len + suffix_len, device=device\n )\n cur_pos += suffix_len\n\n return position_ids\n\n\ndef build_pg_from_micro_batch(\n micro_batch: dict,\n pad_token_id: int,\n padding_mode: str = \"right\",\n):\n \"\"\"Build PrefixGrouper from micro_batch dict containing prompts, responses, response_mask, uid.\"\"\"\n prompts = micro_batch[\"prompts\"]\n responses = micro_batch[\"responses\"]\n response_mask = micro_batch[\"response_mask\"]\n uids = micro_batch[\"uid\"]\n\n bs = responses.size(0)\n\n group_sizes = []\n cur = 1\n for i in range(1, bs):\n if uids[i] == uids[i - 1]:\n cur += 1\n else:\n group_sizes.append(cur)\n cur = 1\n group_sizes.append(cur)\n\n prefix_indices = []\n cursor = 0\n for gs in group_sizes:\n prefix_indices.append(cursor)\n cursor += gs\n prefix_indices = torch.tensor(prefix_indices, device=prompts.device)\n\n prefix_ids = prompts.index_select(0, prefix_indices)\n prefix_mask = prefix_ids.ne(pad_token_id)\n\n prefix_grouper = PrefixGrouper.from_ungrouped_masks(\n prefix_mask=prefix_mask,\n suffix_mask=response_mask,\n group_sizes=group_sizes,\n padding_mode=padding_mode,\n device=prompts.device,\n )\n\n concat_input_ids = prefix_grouper.concat_input(prefix_ids, prefix_mask, responses, response_mask)\n\n attention_mask = prefix_grouper.padding_mask\n\n position_ids = build_position_ids_for_prefix_grouper(prefix_grouper)\n\n return (\n prefix_grouper,\n concat_input_ids,\n attention_mask,\n position_ids,\n responses,\n response_mask,\n )\n\n\ndef pg_forward(\n model,\n prefix_grouper,\n concat_input_ids,\n attention_mask,\n position_ids,\n completion_ids,\n completion_mask,\n *,\n temperature=1.0,\n padding_mode=\"right\",\n include_prefix_last=1,\n calculate_entropy=False,\n entropy_fn=None,\n):\n logits = model(\n input_ids=concat_input_ids,\n attention_mask=attention_mask,\n position_ids=position_ids,\n use_cache=False,\n prefix_grouper=prefix_grouper,\n ).logits\n\n prefix_out, prefix_mask, suffix_out_raw, suffix_mask_raw = prefix_grouper.split_output(\n logits, include_prefix_last=include_prefix_last\n )\n\n completion_ids_right = prefix_grouper.convert_padding(\n completion_ids,\n completion_mask,\n padding_mode=padding_mode,\n )\n\n suffix_out = suffix_out_raw[:, :-1].float()\n suffix_mask = suffix_mask_raw[:, 1:]\n\n suffix_out /= temperature\n\n log_probs = logprobs_from_logits(suffix_out, completion_ids_right)\n\n entropy = None\n if calculate_entropy and entropy_fn is not None:\n entropy = entropy_fn(suffix_out)\n\n return log_probs, entropy, suffix_mask\n\n\ndef forward_micro_batch_with_prefix_grouper(\n micro_batch: dict,\n model,\n temperature: float,\n calculate_entropy: bool,\n device_name: str,\n param_dtype,\n use_chunking_entropy: bool = False,\n):\n \"\"\"\n Forward pass using PrefixGrouper for shared-prefix optimization.\n\n Args:\n micro_batch: Dict containing prompts, responses, response_mask, uid, etc.\n model: The actor module.\n temperature: Temperature for logits scaling.\n calculate_entropy: Whether to compute entropy.\n device_name: Device name for autocast.\n param_dtype: Parameter dtype for autocast.\n use_chunking_entropy: Whether to use chunking entropy function.\n\n Returns:\n tuple: (entropy, log_probs) where entropy may be None if not calculated.\n \"\"\"\n import verl.utils.torch_functional as verl_F\n\n entropy_fn = None\n if calculate_entropy:\n if use_chunking_entropy:\n entropy_fn = verl_F.entropy_from_logits_with_chunking\n else:\n entropy_fn = verl_F.entropy_from_logits\n\n pad_token_id = micro_batch.get(\"pad_token_id\", 0)\n\n (\n prefix_grouper,\n concat_input_ids,\n attention_mask,\n position_ids,\n responses,\n response_mask,\n ) = build_pg_from_micro_batch(\n micro_batch,\n pad_token_id=pad_token_id,\n padding_mode=\"right\",\n )\n\n with torch.autocast(device_type=device_name, dtype=param_dtype):\n log_probs, entropy, suffix_mask_from_pg = pg_forward(\n model=model,\n prefix_grouper=prefix_grouper,\n concat_input_ids=concat_input_ids,\n attention_mask=attention_mask,\n position_ids=position_ids,\n completion_ids=responses,\n completion_mask=response_mask,\n temperature=temperature,\n padding_mode=\"right\",\n include_prefix_last=1,\n calculate_entropy=calculate_entropy,\n entropy_fn=entropy_fn,\n )\n\n # Zero out padding positions\n padding_mask = suffix_mask_from_pg == 0\n log_probs = log_probs.masked_fill(padding_mask, 0.0)\n if entropy is not None:\n entropy = entropy.masked_fill(padding_mask, 0.0)\n\n # Pad to target response length if needed\n target_response_length = responses.size(1)\n if log_probs.size(1) != target_response_length:\n batch_size = log_probs.size(0)\n current_len = log_probs.size(1)\n\n full_log_probs = log_probs.new_zeros(batch_size, target_response_length)\n full_log_probs[:, :current_len] = log_probs\n log_probs = full_log_probs\n\n if entropy is not None:\n full_entropy = entropy.new_zeros(batch_size, target_response_length)\n full_entropy[:, :current_len] = entropy\n entropy = full_entropy\n\n return entropy, log_probs\n"}60{"file_name": "verl__trainer__ppo__ray_trainer.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nPPO Trainer with Ray-based single controller.\nThis trainer supports model-agonistic model initialization with huggingface\n\"\"\"\n\nimport json\nimport os\nimport uuid\nfrom collections import defaultdict\nfrom copy import deepcopy\nfrom pprint import pprint\nfrom typing import Any, Optional\n\nimport numpy as np\nimport torch\nfrom omegaconf import OmegaConf, open_dict\nfrom torch.utils.data import Dataset, Sampler\nfrom torchdata.stateful_dataloader import StatefulDataLoader\nfrom tqdm import tqdm\n\nfrom verl import DataProto\nfrom verl.checkpoint_engine import CheckpointEngineManager\nfrom verl.experimental.dataset.sampler import AbstractCurriculumSampler\nfrom verl.protocol import pad_dataproto_to_divisor, unpad_dataproto\nfrom verl.single_controller.ray import RayClassWithInitArgs, RayWorkerGroup, ResourcePoolManager\nfrom verl.single_controller.ray.base import create_colocated_worker_cls\nfrom verl.trainer.config import AlgoConfig\nfrom verl.trainer.ppo import core_algos\nfrom verl.trainer.ppo.core_algos import AdvantageEstimator, agg_loss\nfrom verl.trainer.ppo.metric_utils import (\n compute_data_metrics,\n compute_throughout_metrics,\n compute_timing_metrics,\n compute_variance_proxy_metrics,\n process_validation_metrics,\n)\nfrom verl.trainer.ppo.reward import extract_reward\nfrom verl.trainer.ppo.utils import Role, WorkerType, need_critic, need_reference_policy, need_reward_model\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.checkpoint.checkpoint_manager import find_latest_ckpt_path, should_save_ckpt_esi\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.utils.debug import marked_timer\nfrom verl.utils.import_utils import load_class_from_fqn\nfrom verl.utils.metric import reduce_metrics\nfrom verl.utils.py_functional import rename_dict\nfrom verl.utils.rollout_skip import RolloutSkip\nfrom verl.utils.seqlen_balancing import calculate_workload, get_seqlen_balanced_partitions, log_seqlen_unbalance\nfrom verl.utils.torch_functional import masked_mean\nfrom verl.utils.tracking import ValidationGenerationsLogger\nfrom verl.workers.config import FSDPEngineConfig\nfrom verl.workers.utils.padding import left_right_2_no_padding, no_padding_2_padding\n\n\ndef apply_kl_penalty(data: DataProto, kl_ctrl: core_algos.AdaptiveKLController, kl_penalty=\"kl\"):\n \"\"\"Apply KL penalty to the token-level rewards.\n\n This function computes the KL divergence between the reference policy and current policy,\n then applies a penalty to the token-level rewards based on this divergence.\n\n Args:\n data (DataProto): The data containing batched model outputs and inputs.\n kl_ctrl (core_algos.AdaptiveKLController): Controller for adaptive KL penalty.\n kl_penalty (str, optional): Type of KL penalty to apply. Defaults to \"kl\".\n\n Returns:\n tuple: A tuple containing:\n - The updated data with token-level rewards adjusted by KL penalty\n - A dictionary of metrics related to the KL penalty\n \"\"\"\n response_mask = data.batch[\"response_mask\"]\n token_level_scores = data.batch[\"token_level_scores\"]\n batch_size = data.batch.batch_size[0]\n\n # compute kl between ref_policy and current policy\n # When apply_kl_penalty, algorithm.use_kl_in_reward=True, so the reference model has been enabled.\n kld = core_algos.kl_penalty(\n data.batch[\"old_log_probs\"], data.batch[\"ref_log_prob\"], kl_penalty=kl_penalty\n ) # (batch_size, response_length)\n kld = kld * response_mask\n beta = kl_ctrl.value\n\n token_level_rewards = token_level_scores - beta * kld\n\n current_kl = masked_mean(kld, mask=response_mask, axis=-1) # average over sequence\n current_kl = torch.mean(current_kl, dim=0).item()\n\n # according to https://github.com/huggingface/trl/blob/951ca1841f29114b969b57b26c7d3e80a39f75a0/trl/trainer/ppo_trainer.py#L837\n kl_ctrl.update(current_kl=current_kl, n_steps=batch_size)\n data.batch[\"token_level_rewards\"] = token_level_rewards\n\n metrics = {\"actor/reward_kl_penalty\": current_kl, \"actor/reward_kl_penalty_coeff\": beta}\n\n return data, metrics\n\n\ndef compute_response_mask(data: DataProto):\n \"\"\"Compute the attention mask for the response part of the sequence.\n\n This function extracts the portion of the attention mask that corresponds to the model's response,\n which is used for masking computations that should only apply to response tokens.\n\n Args:\n data (DataProto): The data containing batched model outputs and inputs.\n\n Returns:\n torch.Tensor: The attention mask for the response tokens.\n \"\"\"\n responses = data.batch[\"responses\"]\n response_length = responses.size(1)\n attention_mask = data.batch[\"attention_mask\"]\n return attention_mask[:, -response_length:]\n\n\ndef compute_advantage(\n data: DataProto,\n adv_estimator: AdvantageEstimator,\n gamma: float = 1.0,\n lam: float = 1.0,\n num_repeat: int = 1,\n norm_adv_by_std_in_grpo: bool = True,\n config: Optional[AlgoConfig] = None,\n) -> DataProto:\n \"\"\"Compute advantage estimates for policy optimization.\n\n This function computes advantage estimates using various estimators like GAE, GRPO, REINFORCE++, etc.\n The advantage estimates are used to guide policy optimization in RL algorithms.\n\n Args:\n data (DataProto): The data containing batched model outputs and inputs.\n adv_estimator (AdvantageEstimator): The advantage estimator to use (e.g., GAE, GRPO, REINFORCE++).\n gamma (float, optional): Discount factor for future rewards. Defaults to 1.0.\n lam (float, optional): Lambda parameter for GAE. Defaults to 1.0.\n num_repeat (int, optional): Number of times to repeat the computation. Defaults to 1.\n norm_adv_by_std_in_grpo (bool, optional): Whether to normalize advantages by standard deviation in\n GRPO. Defaults to True.\n config (dict, optional): Configuration dictionary for algorithm settings. Defaults to None.\n\n Returns:\n DataProto: The updated data with computed advantages and returns.\n \"\"\"\n # Back-compatible with trainers that do not compute response mask in fit\n if \"response_mask\" not in data.batch.keys():\n data.batch[\"response_mask\"] = compute_response_mask(data)\n # prepare response group\n if adv_estimator == AdvantageEstimator.GAE:\n # Compute advantages and returns using Generalized Advantage Estimation (GAE)\n advantages, returns = core_algos.compute_gae_advantage_return(\n token_level_rewards=data.batch[\"token_level_rewards\"],\n values=data.batch[\"values\"],\n response_mask=data.batch[\"response_mask\"],\n gamma=gamma,\n lam=lam,\n )\n data.batch[\"advantages\"] = advantages\n data.batch[\"returns\"] = returns\n if config.get(\"use_pf_ppo\", False):\n data = core_algos.compute_pf_ppo_reweight_data(\n data,\n config.pf_ppo.get(\"reweight_method\"),\n config.pf_ppo.get(\"weight_pow\"),\n )\n elif adv_estimator == AdvantageEstimator.GRPO:\n # Initialize the mask for GRPO calculation\n grpo_calculation_mask = data.batch[\"response_mask\"]\n\n # Call compute_grpo_outcome_advantage with parameters matching its definition\n advantages, returns = core_algos.compute_grpo_outcome_advantage(\n token_level_rewards=data.batch[\"token_level_rewards\"],\n response_mask=grpo_calculation_mask,\n index=data.non_tensor_batch[\"uid\"],\n norm_adv_by_std_in_grpo=norm_adv_by_std_in_grpo,\n )\n data.batch[\"advantages\"] = advantages\n data.batch[\"returns\"] = returns\n else:\n # handle all other adv estimator type other than GAE and GRPO\n adv_estimator_fn = core_algos.get_adv_estimator_fn(adv_estimator)\n adv_kwargs = {\n \"token_level_rewards\": data.batch[\"token_level_rewards\"],\n \"response_mask\": data.batch[\"response_mask\"],\n \"config\": config,\n }\n if \"uid\" in data.non_tensor_batch: # optional\n adv_kwargs[\"index\"] = data.non_tensor_batch[\"uid\"]\n if \"reward_baselines\" in data.batch: # optional\n adv_kwargs[\"reward_baselines\"] = data.batch[\"reward_baselines\"]\n # Add sum_pi_squared for Optimal Token Baseline\n if adv_estimator in (AdvantageEstimator.OPTIMAL_TOKEN_BASELINE, AdvantageEstimator.TIR_OPTIMAL_TOKEN_BASELINE):\n # Check if sum_pi_squared is available\n assert \"sum_pi_squared\" in data.batch, (\n \"Step-dependent optimal baseline requires sum_pi_squared from actor. \"\n \"Please set actor.calculate_sum_pi_squared=True in config.\"\n )\n adv_kwargs[\"sum_pi_squared\"] = data.batch[\"sum_pi_squared\"]\n # Get pre-computed rollout IS weights if available\n rollout_is_weights = data.batch.get(\"rollout_is_weights\", None)\n adv_kwargs[\"rollout_is_weights\"] = rollout_is_weights\n\n # calculate advantage estimator\n advantages, returns = adv_estimator_fn(**adv_kwargs)\n data.batch[\"advantages\"] = advantages\n data.batch[\"returns\"] = returns\n return data\n\n\nclass RayPPOTrainer:\n \"\"\"Distributed PPO trainer using Ray for scalable reinforcement learning.\n\n This trainer orchestrates distributed PPO training across multiple nodes and GPUs,\n managing actor rollouts, critic training, and reward computation with Ray backend.\n Supports various model architectures including FSDP, Megatron, vLLM, and SGLang integration.\n \"\"\"\n\n # TODO: support each role have individual ray_worker_group_cls,\n # i.e., support different backend of different role\n def __init__(\n self,\n config,\n tokenizer,\n role_worker_mapping: dict[Role, WorkerType],\n resource_pool_manager: ResourcePoolManager,\n ray_worker_group_cls: type[RayWorkerGroup] = RayWorkerGroup,\n processor=None,\n train_dataset: Optional[Dataset] = None,\n val_dataset: Optional[Dataset] = None,\n collate_fn=None,\n train_sampler: Optional[Sampler] = None,\n device_name=None,\n ):\n \"\"\"\n Initialize distributed PPO trainer with Ray backend.\n Note that this trainer runs on the driver process on a single CPU/GPU node.\n\n Args:\n config: Configuration object containing training parameters.\n tokenizer: Tokenizer used for encoding and decoding text.\n role_worker_mapping (dict[Role, WorkerType]): Mapping from roles to worker classes.\n resource_pool_manager (ResourcePoolManager): Manager for Ray resource pools.\n ray_worker_group_cls (RayWorkerGroup, optional): Class for Ray worker groups. Defaults to RayWorkerGroup.\n processor: Optional data processor, used for multimodal data\n train_dataset (Optional[Dataset], optional): Training dataset. Defaults to None.\n val_dataset (Optional[Dataset], optional): Validation dataset. Defaults to None.\n collate_fn: Function to collate data samples into batches.\n train_sampler (Optional[Sampler], optional): Sampler for the training dataset. Defaults to None.\n device_name (str, optional): Device name for training (e.g., \"cuda\", \"cpu\"). Defaults to None.\n \"\"\"\n\n # Store the tokenizer for text processing\n self.tokenizer = tokenizer\n self.processor = processor\n self.config = config\n\n self.hybrid_engine = config.actor_rollout_ref.hybrid_engine\n assert self.hybrid_engine, \"Currently, only support hybrid engine\"\n\n if self.hybrid_engine:\n assert Role.ActorRollout in role_worker_mapping or Role.ActorRolloutRef in role_worker_mapping, (\n f\"{role_worker_mapping.keys()=}\"\n )\n\n self.role_worker_mapping = role_worker_mapping\n self.resource_pool_manager = resource_pool_manager\n self.use_reference_policy = need_reference_policy(self.config)\n\n self.use_rm = need_reward_model(self.config)\n\n self.use_critic = need_critic(self.config)\n self.ray_worker_group_cls = ray_worker_group_cls\n self.device_name = device_name if device_name else self.config.trainer.device\n self.validation_generations_logger = ValidationGenerationsLogger(\n project_name=self.config.trainer.project_name,\n experiment_name=self.config.trainer.experiment_name,\n )\n\n # if ref_in_actor is True, the reference policy will be actor without lora applied\n lora_rank = config.actor_rollout_ref.model.get(\"lora\", {}).get(\"rank\", 0)\n if lora_rank <= 0:\n lora_rank = config.actor_rollout_ref.model.get(\"lora_rank\", 0)\n self.ref_in_actor = lora_rank > 0 or config.actor_rollout_ref.model.get(\"lora_adapter_path\") is not None\n\n # define in-reward KL control\n # kl loss control currently not suppoorted\n if self.config.algorithm.use_kl_in_reward:\n self.kl_ctrl_in_reward = core_algos.get_kl_controller(self.config.algorithm.kl_ctrl)\n\n self.use_prefix_grouper = self.config.actor_rollout_ref.actor.get(\"use_prefix_grouper\", False)\n self.use_legacy_worker_impl = config.trainer.get(\"use_legacy_worker_impl\", \"auto\")\n\n self._create_dataloader(train_dataset, val_dataset, collate_fn, train_sampler)\n\n def _create_dataloader(self, train_dataset, val_dataset, collate_fn, train_sampler: Optional[Sampler]):\n \"\"\"\n Creates the train and validation dataloaders.\n \"\"\"\n # TODO: we have to make sure the batch size is divisible by the dp size\n from verl.trainer.main_ppo import create_rl_dataset, create_rl_sampler\n\n if train_dataset is None:\n train_dataset = create_rl_dataset(\n self.config.data.train_files,\n self.config.data,\n self.tokenizer,\n self.processor,\n max_samples=self.config.data.get(\"train_max_samples\", -1),\n )\n if val_dataset is None:\n val_dataset = create_rl_dataset(\n self.config.data.val_files,\n self.config.data,\n self.tokenizer,\n self.processor,\n max_samples=self.config.data.get(\"val_max_samples\", -1),\n )\n self.train_dataset, self.val_dataset = train_dataset, val_dataset\n\n if train_sampler is None:\n train_sampler = create_rl_sampler(self.config.data, self.train_dataset)\n if collate_fn is None:\n from verl.utils.dataset.rl_dataset import collate_fn as default_collate_fn\n\n collate_fn = default_collate_fn\n\n num_workers = self.config.data[\"dataloader_num_workers\"]\n\n self.train_dataloader = StatefulDataLoader(\n dataset=self.train_dataset,\n batch_size=self.config.data.get(\"gen_batch_size\", self.config.data.train_batch_size),\n num_workers=num_workers,\n drop_last=True,\n collate_fn=collate_fn,\n sampler=train_sampler,\n )\n\n val_batch_size = self.config.data.val_batch_size # Prefer config value if set\n if val_batch_size is None:\n val_batch_size = len(self.val_dataset)\n\n self.val_dataloader = StatefulDataLoader(\n dataset=self.val_dataset,\n batch_size=val_batch_size,\n num_workers=num_workers,\n shuffle=self.config.data.get(\"validation_shuffle\", True),\n drop_last=False,\n collate_fn=collate_fn,\n )\n\n assert len(self.train_dataloader) >= 1, \"Train dataloader is empty!\"\n assert len(self.val_dataloader) >= 1, \"Validation dataloader is empty!\"\n\n print(\n f\"Size of train dataloader: {len(self.train_dataloader)}, Size of val dataloader: \"\n f\"{len(self.val_dataloader)}\"\n )\n\n total_training_steps = len(self.train_dataloader) * self.config.trainer.total_epochs\n\n if self.config.trainer.total_training_steps is not None:\n total_training_steps = self.config.trainer.total_training_steps\n\n self.total_training_steps = total_training_steps\n print(f\"Total training steps: {self.total_training_steps}\")\n\n try:\n OmegaConf.set_struct(self.config, True)\n with open_dict(self.config):\n if OmegaConf.select(self.config, \"actor_rollout_ref.actor.optim\"):\n self.config.actor_rollout_ref.actor.optim.total_training_steps = total_training_steps\n if OmegaConf.select(self.config, \"critic.optim\"):\n self.config.critic.optim.total_training_steps = total_training_steps\n except Exception as e:\n print(f\"Warning: Could not set total_training_steps in config. Structure missing? Error: {e}\")\n\n def _dump_generations(self, inputs, outputs, gts, scores, reward_extra_infos_dict, dump_path):\n \"\"\"Dump rollout/validation samples as JSONL.\"\"\"\n os.makedirs(dump_path, exist_ok=True)\n filename = os.path.join(dump_path, f\"{self.global_steps}.jsonl\")\n\n n = len(inputs)\n base_data = {\n \"input\": inputs,\n \"output\": outputs,\n \"gts\": gts,\n \"score\": scores,\n \"step\": [self.global_steps] * n,\n }\n\n for k, v in reward_extra_infos_dict.items():\n if len(v) == n:\n base_data[k] = v\n\n lines = []\n for i in range(n):\n entry = {k: v[i] for k, v in base_data.items()}\n lines.append(json.dumps(entry, ensure_ascii=False))\n\n with open(filename, \"w\") as f:\n f.write(\"\\n\".join(lines) + \"\\n\")\n\n print(f\"Dumped generations to {filename}\")\n\n def _log_rollout_data(\n self, batch: DataProto, reward_extra_infos_dict: dict, timing_raw: dict, rollout_data_dir: str\n ):\n \"\"\"Log rollout data to disk.\n Args:\n batch (DataProto): The batch containing rollout data\n reward_extra_infos_dict (dict): Additional reward information to log\n timing_raw (dict): Timing information for profiling\n rollout_data_dir (str): Directory path to save the rollout data\n \"\"\"\n with marked_timer(\"dump_rollout_generations\", timing_raw, color=\"green\"):\n inputs = self.tokenizer.batch_decode(batch.batch[\"prompts\"], skip_special_tokens=True)\n outputs = self.tokenizer.batch_decode(batch.batch[\"responses\"], skip_special_tokens=True)\n scores = batch.batch[\"token_level_scores\"].sum(-1).cpu().tolist()\n sample_gts = [item.non_tensor_batch.get(\"reward_model\", {}).get(\"ground_truth\", None) for item in batch]\n\n reward_extra_infos_to_dump = reward_extra_infos_dict.copy()\n if \"request_id\" in batch.non_tensor_batch:\n reward_extra_infos_dict.setdefault(\n \"request_id\",\n batch.non_tensor_batch[\"request_id\"].tolist(),\n )\n\n self._dump_generations(\n inputs=inputs,\n outputs=outputs,\n gts=sample_gts,\n scores=scores,\n reward_extra_infos_dict=reward_extra_infos_to_dump,\n dump_path=rollout_data_dir,\n )\n\n def _maybe_log_val_generations(self, inputs, outputs, scores):\n \"\"\"Log a table of validation samples to the configured logger (wandb or swanlab)\"\"\"\n\n generations_to_log = self.config.trainer.log_val_generations\n\n if generations_to_log == 0:\n return\n\n import numpy as np\n\n # Create tuples of (input, output, score) and sort by input text\n samples = list(zip(inputs, outputs, scores, strict=True))\n samples.sort(key=lambda x: x[0]) # Sort by input text\n\n # Use fixed random seed for deterministic shuffling\n rng = np.random.RandomState(42)\n rng.shuffle(samples)\n\n # Take first N samples after shuffling\n samples = samples[:generations_to_log]\n\n # Log to each configured logger\n self.validation_generations_logger.log(self.config.trainer.logger, samples, self.global_steps)\n\n def _get_gen_batch(self, batch: DataProto) -> DataProto:\n reward_keys = set({\"data_source\", \"reward_model\", \"extra_info\", \"uid\"}) & batch.non_tensor_batch.keys()\n\n # pop those keys for generation\n batch_keys_to_pop = []\n non_tensor_batch_keys_to_pop = set(batch.non_tensor_batch.keys()) - reward_keys\n gen_batch = batch.pop(\n batch_keys=batch_keys_to_pop,\n non_tensor_batch_keys=list(non_tensor_batch_keys_to_pop),\n )\n\n # For agent loop, we need reward model keys to compute score.\n gen_batch.non_tensor_batch.update(batch.non_tensor_batch)\n\n return gen_batch\n\n def _compute_reward_colocate(self, batch: DataProto) -> tuple[torch.Tensor, dict[str, Any]] | torch.Tensor:\n \"\"\"\n compute reward use colocate reward model\n \"\"\"\n assert self.reward_loop_manager is not None, \"RewardLoopManager is None\"\n batch_reward = self.reward_loop_manager.compute_rm_score(batch)\n return batch_reward\n\n def _validate(self, merged: bool = False):\n data_source_lst = []\n reward_extra_infos_dict: dict[str, list] = defaultdict(list)\n\n # Lists to collect samples for the table\n sample_inputs = []\n sample_outputs = []\n sample_gts = []\n sample_scores = []\n sample_turns = []\n sample_uids = []\n\n for test_data in self.val_dataloader:\n test_batch = DataProto.from_single_dict(test_data)\n\n if \"uid\" not in test_batch.non_tensor_batch:\n test_batch.non_tensor_batch[\"uid\"] = np.array(\n [str(uuid.uuid4()) for _ in range(len(test_batch.batch))], dtype=object\n )\n\n # repeat test batch\n test_batch = test_batch.repeat(\n repeat_times=self.config.actor_rollout_ref.rollout.val_kwargs.n, interleave=True\n )\n\n ground_truths = [\n item.non_tensor_batch.get(\"reward_model\", {}).get(\"ground_truth\", None) for item in test_batch\n ]\n sample_gts.extend(ground_truths)\n\n test_gen_batch = self._get_gen_batch(test_batch)\n test_gen_batch.meta_info = {\n \"eos_token_id\": self.tokenizer.eos_token_id,\n \"pad_token_id\": self.tokenizer.pad_token_id,\n \"recompute_log_prob\": False,\n \"do_sample\": self.config.actor_rollout_ref.rollout.val_kwargs.do_sample,\n \"validate\": True,\n \"global_steps\": self.global_steps,\n }\n print(f\"test_gen_batch meta info: {test_gen_batch.meta_info}\")\n\n # pad to be divisible by dp_size\n size_divisor = self.config.actor_rollout_ref.rollout.agent.num_workers\n test_gen_batch_padded, pad_size = pad_dataproto_to_divisor(test_gen_batch, size_divisor)\n test_output_gen_batch_padded = self.async_rollout_manager.generate_sequences(test_gen_batch_padded)\n\n if self.use_rm and \"rm_scores\" not in test_output_gen_batch_padded.batch.keys():\n # for colocate reward models, we need to sleep rollout model\n # to spare GPU memory for reward model\n self.checkpoint_manager.sleep_replicas()\n batch_reward = self._compute_reward_colocate(test_output_gen_batch_padded)\n test_output_gen_batch_padded = test_output_gen_batch_padded.union(batch_reward)\n # wake up rollout model\n # replace with wake_up method once supported\n self.checkpoint_manager.update_weights()\n\n # unpad\n test_output_gen_batch = unpad_dataproto(test_output_gen_batch_padded, pad_size=pad_size)\n\n print(\"validation generation end\")\n\n # Store generated outputs\n output_ids = test_output_gen_batch.batch[\"responses\"]\n output_texts = [self.tokenizer.decode(ids, skip_special_tokens=True) for ids in output_ids]\n sample_outputs.extend(output_texts)\n\n test_batch = test_batch.union(test_output_gen_batch)\n test_batch.meta_info[\"validate\"] = True\n\n # Store original inputs\n input_ids = test_batch.batch[\"prompts\"]\n # TODO: Can we keep special tokens except for padding tokens?\n input_texts = [self.tokenizer.decode(ids, skip_special_tokens=True) for ids in input_ids]\n sample_inputs.extend(input_texts)\n sample_uids.extend(test_batch.non_tensor_batch[\"uid\"])\n\n # evaluate using reward_function\n reward_tensor, reward_extra_info = extract_reward(test_batch)\n\n scores = reward_tensor.sum(-1).cpu().tolist()\n sample_scores.extend(scores)\n\n reward_extra_infos_dict[\"reward\"].extend(scores)\n for key, values in reward_extra_info.items():\n if key not in reward_extra_infos_dict:\n reward_extra_infos_dict[key] = []\n if isinstance(values, np.ndarray):\n reward_extra_infos_dict[key].extend(values.tolist())\n else:\n reward_extra_infos_dict[key].extend(values if isinstance(values, list) else [values])\n\n # collect num_turns of each prompt\n if \"__num_turns__\" in test_batch.non_tensor_batch:\n sample_turns.append(test_batch.non_tensor_batch[\"__num_turns__\"])\n\n data_source_lst.append(test_batch.non_tensor_batch.get(\"data_source\", [\"unknown\"] * reward_tensor.shape[0]))\n\n self._maybe_log_val_generations(inputs=sample_inputs, outputs=sample_outputs, scores=sample_scores)\n\n # dump generations\n val_data_dir = self.config.trainer.get(\"validation_data_dir\", None)\n if val_data_dir:\n self._dump_generations(\n inputs=sample_inputs,\n outputs=sample_outputs,\n gts=sample_gts,\n scores=sample_scores,\n reward_extra_infos_dict=reward_extra_infos_dict,\n dump_path=val_data_dir,\n )\n\n for key_info, lst in reward_extra_infos_dict.items():\n assert len(lst) == 0 or len(lst) == len(sample_scores), f\"{key_info}: {len(lst)=}, {len(sample_scores)=}\"\n\n if merged:\n print(\"_merge_validation_results validate result will be merged\")\n return {\n \"data_sources\": data_source_lst,\n \"sample_uids\": sample_uids,\n \"sample_turns\": sample_turns,\n \"reward_extra_infos_dict\": reward_extra_infos_dict,\n }\n data_sources = np.concatenate(data_source_lst, axis=0)\n return self._val_metrics_update(data_sources, sample_uids, reward_extra_infos_dict, sample_turns)\n\n def _val_metrics_update(self, data_sources, sample_uids, reward_extra_infos_dict, sample_turns):\n data_src2var2metric2val = process_validation_metrics(data_sources, sample_uids, reward_extra_infos_dict)\n metric_dict = {}\n for data_source, var2metric2val in data_src2var2metric2val.items():\n core_var = \"acc\" if \"acc\" in var2metric2val else \"reward\"\n for var_name, metric2val in var2metric2val.items():\n n_max = max([int(name.split(\"@\")[-1].split(\"/\")[0]) for name in metric2val.keys()])\n for metric_name, metric_val in metric2val.items():\n if (\n (var_name == core_var)\n and any(metric_name.startswith(pfx) for pfx in [\"mean\", \"maj\", \"best\"])\n and (f\"@{n_max}\" in metric_name)\n ):\n metric_sec = \"val-core\"\n else:\n metric_sec = \"val-aux\"\n pfx = f\"{metric_sec}/{data_source}/{var_name}/{metric_name}\"\n metric_dict[pfx] = metric_val\n\n if len(sample_turns) > 0:\n sample_turns = np.concatenate(sample_turns)\n metric_dict[\"val-aux/num_turns/min\"] = sample_turns.min()\n metric_dict[\"val-aux/num_turns/max\"] = sample_turns.max()\n metric_dict[\"val-aux/num_turns/mean\"] = sample_turns.mean()\n\n return metric_dict\n\n def _merge_validation_results(self, result_a, result_b):\n if result_a is None and result_b is None:\n return {}\n if result_a is None:\n result_a = {\"data_sources\": [], \"sample_uids\": [], \"sample_turns\": [], \"reward_extra_infos_dict\": {}}\n if result_b is None:\n result_b = {\"data_sources\": [], \"sample_uids\": [], \"sample_turns\": [], \"reward_extra_infos_dict\": {}}\n\n if not result_a.get(\"data_sources\") and not result_b.get(\"data_sources\"):\n return {}\n\n data_sources = np.concatenate(result_a[\"data_sources\"] + result_b[\"data_sources\"], axis=0)\n sample_uids = result_a[\"sample_uids\"] + result_b[\"sample_uids\"]\n sample_turns = result_a[\"sample_turns\"] + result_b[\"sample_turns\"]\n\n reward_extra_infos_dict = {}\n all_keys = set(result_a[\"reward_extra_infos_dict\"].keys()) | set(result_b[\"reward_extra_infos_dict\"].keys())\n for key in all_keys:\n list_a = result_a[\"reward_extra_infos_dict\"].get(key, [])\n list_b = result_b[\"reward_extra_infos_dict\"].get(key, [])\n reward_extra_infos_dict[key] = list_a + list_b\n\n return self._val_metrics_update(data_sources, sample_uids, reward_extra_infos_dict, sample_turns)\n\n def init_workers(self):\n \"\"\"Initialize distributed training workers using Ray backend.\n\n Creates:\n 1. Ray resource pools from configuration\n 2. Worker groups for each role (actor, critic, etc.)\n \"\"\"\n self.resource_pool_manager.create_resource_pool()\n\n self.resource_pool_to_cls = {pool: {} for pool in self.resource_pool_manager.resource_pool_dict.values()}\n\n # create actor and rollout\n actor_role = Role.ActorRolloutRef if Role.ActorRolloutRef in self.role_worker_mapping else Role.ActorRollout\n if self.hybrid_engine:\n actor_rollout_resource_pool = self.resource_pool_manager.get_resource_pool(actor_role)\n actor_rollout_cls = RayClassWithInitArgs(\n cls=self.role_worker_mapping[actor_role],\n config=self.config.actor_rollout_ref,\n role=str(actor_role),\n )\n self.resource_pool_to_cls[actor_rollout_resource_pool][str(actor_role)] = actor_rollout_cls\n else:\n raise NotImplementedError\n\n # create critic\n if self.use_critic:\n resource_pool = self.resource_pool_manager.get_resource_pool(Role.Critic)\n\n from verl.workers.config import CriticConfig\n\n critic_cfg: CriticConfig = omega_conf_to_dataclass(self.config.critic)\n\n if self.use_legacy_worker_impl == \"disable\":\n # convert critic_cfg into TrainingWorkerConfig\n from verl.workers.engine_workers import TrainingWorkerConfig\n\n orig_critic_cfg = critic_cfg\n if orig_critic_cfg.strategy == \"fsdp\":\n engine_config: FSDPEngineConfig = orig_critic_cfg.model.fsdp_config\n engine_config.infer_max_token_len_per_gpu = critic_cfg.ppo_infer_max_token_len_per_gpu\n engine_config.max_token_len_per_gpu = critic_cfg.ppo_max_token_len_per_gpu\n else:\n raise NotImplementedError(f\"Unknown strategy {orig_critic_cfg.strategy=}\")\n\n critic_cfg = TrainingWorkerConfig(\n model_type=\"value_model\",\n model_config=orig_critic_cfg.model_config,\n engine_config=engine_config,\n optimizer_config=orig_critic_cfg.optim,\n checkpoint_config=orig_critic_cfg.checkpoint,\n )\n\n critic_cls = RayClassWithInitArgs(cls=self.role_worker_mapping[Role.Critic], config=critic_cfg)\n self.resource_pool_to_cls[resource_pool][str(Role.Critic)] = critic_cls\n\n # create reference policy if needed\n if self.use_reference_policy and Role.RefPolicy in self.role_worker_mapping:\n resource_pool = self.resource_pool_manager.get_resource_pool(Role.RefPolicy)\n ref_policy_cls = RayClassWithInitArgs(\n self.role_worker_mapping[Role.RefPolicy],\n config=self.config.actor_rollout_ref,\n role=str(Role.RefPolicy),\n )\n self.resource_pool_to_cls[resource_pool][str(Role.RefPolicy)] = ref_policy_cls\n\n # initialize WorkerGroup\n # NOTE: if you want to use a different resource pool for each role, which can support different parallel size,\n # you should not use `create_colocated_worker_cls`.\n # Instead, directly pass different resource pool to different worker groups.\n # See https://github.com/volcengine/verl/blob/master/examples/ray/tutorial.ipynb for more information.\n all_wg = {}\n wg_kwargs = {} # Setting up kwargs for RayWorkerGroup\n if OmegaConf.select(self.config.trainer, \"ray_wait_register_center_timeout\") is not None:\n wg_kwargs[\"ray_wait_register_center_timeout\"] = self.config.trainer.ray_wait_register_center_timeout\n if OmegaConf.select(self.config.global_profiler, \"steps\") is not None:\n wg_kwargs[\"profile_steps\"] = OmegaConf.select(self.config.global_profiler, \"steps\")\n # Only require nsight worker options when tool is nsys\n if OmegaConf.select(self.config.global_profiler, \"tool\") == \"nsys\":\n assert (\n OmegaConf.select(self.config.global_profiler.global_tool_config.nsys, \"worker_nsight_options\")\n is not None\n ), \"worker_nsight_options must be set when using nsys with profile_steps\"\n wg_kwargs[\"worker_nsight_options\"] = OmegaConf.to_container(\n OmegaConf.select(self.config.global_profiler.global_tool_config.nsys, \"worker_nsight_options\")\n )\n wg_kwargs[\"device_name\"] = self.device_name\n\n for resource_pool, class_dict in self.resource_pool_to_cls.items():\n worker_dict_cls = create_colocated_worker_cls(class_dict=class_dict)\n wg_dict = self.ray_worker_group_cls(\n resource_pool=resource_pool,\n ray_cls_with_init=worker_dict_cls,\n **wg_kwargs,\n )\n spawn_wg = wg_dict.spawn(prefix_set=class_dict.keys())\n all_wg.update(spawn_wg)\n\n if self.use_critic:\n self.critic_wg = all_wg[str(Role.Critic)]\n if self.use_legacy_worker_impl == \"disable\":\n self.critic_wg.reset()\n # assign critic loss\n from functools import partial\n\n from verl.workers.utils.losses import value_loss\n\n value_loss_ = partial(value_loss, config=orig_critic_cfg)\n self.critic_wg.set_loss_fn(value_loss_)\n else:\n self.critic_wg.init_model()\n\n if self.use_reference_policy and not self.ref_in_actor:\n if str(Role.RefPolicy) in all_wg:\n self.ref_policy_wg = all_wg[str(Role.RefPolicy)]\n self.ref_policy_wg.init_model()\n else:\n # Model engine: ActorRolloutRefWorker\n assert str(Role.ActorRolloutRef) in all_wg, f\"{all_wg.keys()=}\"\n self.ref_policy_wg = all_wg[str(Role.ActorRolloutRef)]\n\n # we should create rollout at the end so that vllm can have a better estimation of kv cache memory\n self.actor_rollout_wg = all_wg[str(actor_role)]\n self.actor_rollout_wg.init_model()\n\n if self.ref_in_actor:\n self.ref_policy_wg = self.actor_rollout_wg\n\n # create reward loop manager\n from verl.experimental.reward_loop import RewardLoopManager\n\n # initalize reward loop manager\n # reward model (colocate or standalone): get resource_pool\n # no reward model: resource_pool = None\n resource_pool = self.resource_pool_manager.get_resource_pool(Role.RewardModel) if self.use_rm else None\n self.reward_loop_manager = RewardLoopManager(\n config=self.config,\n rm_resource_pool=resource_pool,\n )\n\n # create async rollout manager and request scheduler\n # Note: mode is always \"async\" since sync mode is deprecated\n self.async_rollout_mode = True\n\n # Support custom AgentLoopManager via config\n manager_class_fqn = self.config.actor_rollout_ref.rollout.get(\"agent\", {}).get(\"agent_loop_manager_class\")\n if manager_class_fqn:\n AgentLoopManager = load_class_from_fqn(manager_class_fqn, \"AgentLoopManager\")\n else:\n from verl.experimental.agent_loop import AgentLoopManager\n\n # infrastructure overview: https://verl.readthedocs.io/en/latest/advance/reward_loop.html#architecture-design\n # agent_reward_loop: streaming reward computation with actor rollout\n # two conditions satisfied: (1) no reward model, or (2) reward model with extra resource pool\n enable_agent_reward_loop = not self.use_rm or self.config.reward.reward_model.enable_resource_pool\n\n # if enable_agent_reward_loop, we directly pass reward_loop_workers to agent loop manager\n # to stream reward computation with actor rollout\n reward_loop_worker_handles = self.reward_loop_manager.reward_loop_workers if enable_agent_reward_loop else None\n self.async_rollout_manager = AgentLoopManager(\n config=self.config,\n worker_group=self.actor_rollout_wg,\n rollout_resource_pool=actor_rollout_resource_pool,\n reward_loop_worker_handles=reward_loop_worker_handles,\n )\n\n self.checkpoint_manager = CheckpointEngineManager(\n backend=self.config.actor_rollout_ref.rollout.checkpoint_engine.backend,\n trainer=self.actor_rollout_wg,\n replicas=self.async_rollout_manager.rollout_replicas,\n )\n\n # sleep all replicas to load checkpoint\n self.checkpoint_manager.sleep_replicas()\n\n def _save_checkpoint(self):\n from verl.utils.fs import local_mkdir_safe\n\n # path: given_path + `/global_step_{global_steps}` + `/actor`\n local_global_step_folder = os.path.join(\n self.config.trainer.default_local_dir, f\"global_step_{self.global_steps}\"\n )\n\n print(f\"local_global_step_folder: {local_global_step_folder}\")\n actor_local_path = os.path.join(local_global_step_folder, \"actor\")\n\n actor_remote_path = (\n None\n if self.config.trainer.default_hdfs_dir is None\n else os.path.join(self.config.trainer.default_hdfs_dir, f\"global_step_{self.global_steps}\", \"actor\")\n )\n\n remove_previous_ckpt_in_save = self.config.trainer.get(\"remove_previous_ckpt_in_save\", False)\n if remove_previous_ckpt_in_save:\n print(\n \"Warning: remove_previous_ckpt_in_save is deprecated,\"\n + \" set max_actor_ckpt_to_keep=1 and max_critic_ckpt_to_keep=1 instead\"\n )\n max_actor_ckpt_to_keep = (\n self.config.trainer.get(\"max_actor_ckpt_to_keep\", None) if not remove_previous_ckpt_in_save else 1\n )\n max_critic_ckpt_to_keep = (\n self.config.trainer.get(\"max_critic_ckpt_to_keep\", None) if not remove_previous_ckpt_in_save else 1\n )\n\n self.actor_rollout_wg.save_checkpoint(\n actor_local_path, actor_remote_path, self.global_steps, max_ckpt_to_keep=max_actor_ckpt_to_keep\n )\n\n if self.use_critic:\n critic_local_path = os.path.join(local_global_step_folder, str(Role.Critic))\n critic_remote_path = (\n None\n if self.config.trainer.default_hdfs_dir is None\n else os.path.join(\n self.config.trainer.default_hdfs_dir, f\"global_step_{self.global_steps}\", str(Role.Critic)\n )\n )\n self.critic_wg.save_checkpoint(\n critic_local_path, critic_remote_path, self.global_steps, max_ckpt_to_keep=max_critic_ckpt_to_keep\n )\n\n # save dataloader\n local_mkdir_safe(local_global_step_folder)\n dataloader_local_path = os.path.join(local_global_step_folder, \"data.pt\")\n dataloader_state_dict = self.train_dataloader.state_dict()\n torch.save(dataloader_state_dict, dataloader_local_path)\n\n # latest checkpointed iteration tracker (for atomic usage)\n if (\n hasattr(self.config.actor_rollout_ref.actor.checkpoint, \"async_save\")\n and self.config.actor_rollout_ref.actor.checkpoint.async_save\n ) or (\n \"async_save\" in self.config.actor_rollout_ref.actor.checkpoint\n and self.config.actor_rollout_ref.actor.checkpoint[\"async_save\"]\n ):\n print(\"skip write latest_checkpointed_iteration.txt when async_save is True\")\n return\n local_latest_checkpointed_iteration = os.path.join(\n self.config.trainer.default_local_dir, \"latest_checkpointed_iteration.txt\"\n )\n with open(local_latest_checkpointed_iteration, \"w\") as f:\n f.write(str(self.global_steps))\n\n def _load_checkpoint(self):\n if self.config.trainer.resume_mode == \"disable\":\n return 0\n\n # load from hdfs\n if self.config.trainer.default_hdfs_dir is not None:\n raise NotImplementedError(\"load from hdfs is not implemented yet\")\n else:\n checkpoint_folder = self.config.trainer.default_local_dir # TODO: check path\n if not os.path.isabs(checkpoint_folder):\n working_dir = os.getcwd()\n checkpoint_folder = os.path.join(working_dir, checkpoint_folder)\n global_step_folder = find_latest_ckpt_path(checkpoint_folder) # None if no latest\n\n # find global_step_folder\n if self.config.trainer.resume_mode == \"auto\":\n if global_step_folder is None:\n print(\"Training from scratch\")\n return 0\n else:\n if self.config.trainer.resume_mode == \"resume_path\":\n assert isinstance(self.config.trainer.resume_from_path, str), \"resume ckpt must be str type\"\n assert \"global_step_\" in self.config.trainer.resume_from_path, (\n \"resume ckpt must specify the global_steps\"\n )\n global_step_folder = self.config.trainer.resume_from_path\n if not os.path.isabs(global_step_folder):\n working_dir = os.getcwd()\n global_step_folder = os.path.join(working_dir, global_step_folder)\n print(f\"Load from checkpoint folder: {global_step_folder}\")\n # set global step\n self.global_steps = int(global_step_folder.split(\"global_step_\")[-1])\n\n print(f\"Setting global step to {self.global_steps}\")\n print(f\"Resuming from {global_step_folder}\")\n\n actor_path = os.path.join(global_step_folder, \"actor\")\n critic_path = os.path.join(global_step_folder, str(Role.Critic))\n # load actor\n self.actor_rollout_wg.load_checkpoint(\n actor_path, del_local_after_load=self.config.trainer.del_local_ckpt_after_load\n )\n # load critic\n if self.use_critic:\n self.critic_wg.load_checkpoint(\n critic_path, del_local_after_load=self.config.trainer.del_local_ckpt_after_load\n )\n\n # load dataloader,\n # TODO: from remote not implemented yet\n dataloader_local_path = os.path.join(global_step_folder, \"data.pt\")\n if os.path.exists(dataloader_local_path):\n dataloader_state_dict = torch.load(dataloader_local_path, weights_only=False)\n self.train_dataloader.load_state_dict(dataloader_state_dict)\n else:\n print(f\"Warning: No dataloader state found at {dataloader_local_path}, will start from scratch\")\n\n def _start_profiling(self, do_profile: bool) -> None:\n \"\"\"Start profiling for all worker groups if profiling is enabled.\"\"\"\n if do_profile:\n self.actor_rollout_wg.start_profile(role=\"e2e\", profile_step=self.global_steps)\n if self.use_reference_policy:\n self.ref_policy_wg.start_profile(profile_step=self.global_steps)\n if self.use_critic:\n self.critic_wg.start_profile(profile_step=self.global_steps)\n\n def _stop_profiling(self, do_profile: bool) -> None:\n \"\"\"Stop profiling for all worker groups if profiling is enabled.\"\"\"\n if do_profile:\n self.actor_rollout_wg.stop_profile()\n if self.use_reference_policy:\n self.ref_policy_wg.stop_profile()\n if self.use_critic:\n self.critic_wg.stop_profile()\n\n def _get_dp_size(self, worker_group, role: str) -> int:\n \"\"\"Get data parallel size from worker group dispatch info.\n\n This method retrieves the data parallel size by querying the dispatch info\n for the specified role. The dispatch info is cached for subsequent calls.\n\n Args:\n worker_group: The worker group to query dispatch info from.\n role: The role name (e.g., \"actor\", \"critic\") to get DP size for.\n\n Returns:\n The data parallel size (number of DP ranks).\n \"\"\"\n if role not in worker_group._dispatch_info:\n dp_rank_mapping = worker_group._query_dispatch_info(role)\n worker_group._dispatch_info[role] = dp_rank_mapping\n else:\n dp_rank_mapping = worker_group._dispatch_info[role]\n return max(dp_rank_mapping) + 1\n\n def _balance_batch(self, batch: DataProto, metrics, logging_prefix=\"global_seqlen\", keep_minibatch=False):\n \"\"\"Reorder the data on single controller such that each dp rank gets similar total tokens.\n\n When use_prefix_grouper is enabled, uses group-level balancing to keep samples with\n the same uid together on the same rank for prefix sharing optimization.\n \"\"\"\n attention_mask = batch.batch[\"attention_mask\"]\n batch_size = attention_mask.shape[0]\n global_seqlen_lst = batch.batch[\"attention_mask\"].view(batch_size, -1).sum(-1) # (train_batch_size,)\n workload_lst = calculate_workload(global_seqlen_lst)\n # Get dp_size from dispatch info to correctly balance across data parallel ranks\n # Note: world_size may include tensor/pipeline parallel dimensions, but we only want DP\n dp_size = self._get_dp_size(self.actor_rollout_wg, \"actor\")\n\n # Use group-level balancing for PrefixGrouper to keep same-uid samples together\n if getattr(self, \"use_prefix_grouper\", False) and \"uid\" in batch.non_tensor_batch:\n from verl.utils.seqlen_balancing import get_group_balanced_partitions\n\n uid_list = list(batch.non_tensor_batch[\"uid\"])\n seqlen_list = global_seqlen_lst.tolist()\n\n # Count number of uid groups\n num_groups = len(set(uid_list))\n\n if num_groups % dp_size != 0:\n raise ValueError(\n f\"PrefixGrouper with balance_batch requires num_uid_groups ({num_groups}) \"\n f\"% dp_size ({dp_size}) == 0. \"\n f\"This ensures each rank gets equal number of groups. \"\n f\"Current batch_size={batch_size}, adjust batch_size to be a multiple of \"\n f\"dp_size * rollout.n.\"\n )\n\n global_partition_lst = get_group_balanced_partitions(\n seqlen_list=seqlen_list,\n uid_list=uid_list,\n k_partitions=dp_size,\n )\n\n elif keep_minibatch:\n # Decouple the DP balancing and mini-batching.\n minibatch_size = self.config.actor_rollout_ref.actor.get(\"ppo_mini_batch_size\")\n minibatch_num = len(workload_lst) // minibatch_size\n global_partition_lst = [[] for _ in range(dp_size)]\n for i in range(minibatch_num):\n rearrange_minibatch_lst = get_seqlen_balanced_partitions(\n workload_lst[i * minibatch_size : (i + 1) * minibatch_size],\n k_partitions=dp_size,\n equal_size=True,\n )\n for j, part in enumerate(rearrange_minibatch_lst):\n global_partition_lst[j].extend([x + minibatch_size * i for x in part])\n else:\n global_partition_lst = get_seqlen_balanced_partitions(workload_lst, k_partitions=dp_size, equal_size=True)\n # Place smaller micro-batches at both ends to reduce the bubbles in pipeline parallel.\n # Skip reordering within partitions for PrefixGrouper to maintain uid grouping\n if not getattr(self, \"use_prefix_grouper\", False):\n for idx, partition in enumerate(global_partition_lst):\n partition.sort(key=lambda x: (workload_lst[x], x))\n ordered_partition = partition[::2] + partition[1::2][::-1]\n global_partition_lst[idx] = ordered_partition\n\n # reorder based on index. The data will be automatically equally partitioned by dispatch function\n global_idx = torch.tensor([j for partition in global_partition_lst for j in partition])\n batch.reorder(global_idx)\n global_balance_stats = log_seqlen_unbalance(\n seqlen_list=global_seqlen_lst.tolist(), partitions=global_partition_lst, prefix=logging_prefix\n )\n metrics.update(global_balance_stats)\n\n def _compute_values(self, batch: DataProto) -> DataProto:\n if self.use_legacy_worker_impl == \"disable\":\n batch_td = batch.to_tensordict()\n # step 2: convert from padding to nopadding\n batch_td = left_right_2_no_padding(batch_td)\n # step 3: add meta info\n tu.assign_non_tensor(batch_td, compute_loss=False)\n output = self.critic_wg.infer_batch(batch_td)\n output = output.get()\n values = tu.get(output, \"values\")\n values = no_padding_2_padding(values, batch_td)\n values = tu.get_tensordict({\"values\": values.float()})\n values = DataProto.from_tensordict(values)\n else:\n values = self.critic_wg.compute_values(batch)\n return values\n\n def _compute_ref_log_prob(self, batch: DataProto) -> DataProto:\n if self.use_legacy_worker_impl == \"disable\":\n # step 1: convert dataproto to tensordict.\n batch_td = batch.to_tensordict()\n # step 2: convert from padding to nopadding\n batch_td = left_right_2_no_padding(batch_td)\n # step 3: add meta info\n metadata = {\"calculate_entropy\": False, \"compute_loss\": False}\n if self.ref_in_actor:\n metadata[\"no_lora_adapter\"] = True\n tu.assign_non_tensor(batch_td, **metadata)\n if self.ref_in_actor:\n output = self.actor_rollout_wg.compute_log_prob(batch_td)\n else:\n output = self.ref_policy_wg.compute_ref_log_prob(batch_td)\n # gather output\n log_probs = tu.get(output, \"log_probs\")\n # step 4. No padding to padding\n log_probs = no_padding_2_padding(log_probs, batch_td)\n # step 5: rebuild a tensordict and convert to dataproto\n ref_log_prob = tu.get_tensordict({\"ref_log_prob\": log_probs.float()})\n ref_log_prob = DataProto.from_tensordict(ref_log_prob)\n else:\n ref_log_prob = self.ref_policy_wg.compute_ref_log_prob(batch)\n\n return ref_log_prob\n\n def _compute_old_log_prob(self, batch: DataProto):\n if self.use_legacy_worker_impl == \"disable\":\n # TODO: remove step 1, 2, 4 after we make the whole training tensordict and padding free\n # step 1: convert dataproto to tensordict.\n batch_td = batch.to_tensordict()\n # step 2: convert from padding to nopadding\n batch_td = left_right_2_no_padding(batch_td)\n # step 3: add meta info\n tu.assign_non_tensor(batch_td, calculate_entropy=True, compute_loss=False)\n output = self.actor_rollout_wg.compute_log_prob(batch_td)\n # gather output\n entropy = tu.get(output, \"entropy\")\n log_probs = tu.get(output, \"log_probs\")\n old_log_prob_mfu = tu.get(output, \"metrics\")[\"mfu\"]\n # step 4. No padding to padding\n entropy = no_padding_2_padding(entropy, batch_td)\n log_probs = no_padding_2_padding(log_probs, batch_td)\n # step 5: rebuild a tensordict and convert to dataproto\n old_log_prob = tu.get_tensordict({\"old_log_probs\": log_probs.float(), \"entropys\": entropy.float()})\n old_log_prob = DataProto.from_tensordict(old_log_prob)\n else:\n old_log_prob = self.actor_rollout_wg.compute_log_prob(batch)\n old_log_prob_mfu = 0\n return old_log_prob, old_log_prob_mfu\n\n def _update_actor(self, batch: DataProto) -> DataProto:\n rollout_config = self.config.actor_rollout_ref.rollout\n batch.meta_info[\"multi_turn\"] = rollout_config.multi_turn.enable\n # TODO: Make \"temperature\" single source of truth from generation.\n batch.meta_info[\"temperature\"] = rollout_config.temperature\n # update actor\n if self.use_legacy_worker_impl == \"disable\":\n batch_td = batch.to_tensordict()\n # step 2: convert from padding to no-padding\n batch_td = left_right_2_no_padding(batch_td)\n calculate_entropy = self.config.actor_rollout_ref.actor.entropy_coeff != 0.0\n ppo_mini_batch_size = self.config.actor_rollout_ref.actor.ppo_mini_batch_size\n ppo_mini_batch_size = ppo_mini_batch_size * self.config.actor_rollout_ref.rollout.n\n ppo_epochs = self.config.actor_rollout_ref.actor.ppo_epochs\n seed = self.config.actor_rollout_ref.actor.data_loader_seed\n shuffle = self.config.actor_rollout_ref.actor.shuffle\n tu.assign_non_tensor(\n batch_td,\n calculate_entropy=calculate_entropy,\n global_batch_size=ppo_mini_batch_size,\n mini_batch_size=ppo_mini_batch_size,\n epochs=ppo_epochs,\n seed=seed,\n dataloader_kwargs={\"shuffle\": shuffle},\n )\n\n actor_output = self.actor_rollout_wg.update_actor(batch_td)\n actor_output = tu.get(actor_output, \"metrics\")\n actor_output = rename_dict(actor_output, \"actor/\")\n # modify key name\n actor_output[\"perf/mfu/actor\"] = actor_output.pop(\"actor/mfu\")\n actor_output = DataProto.from_single_dict(data={}, meta_info={\"metrics\": actor_output})\n else:\n actor_output = self.actor_rollout_wg.update_actor(batch)\n\n return actor_output\n\n def _update_critic(self, batch: DataProto) -> DataProto:\n if self.use_legacy_worker_impl == \"disable\":\n batch_td = batch.to_tensordict()\n # step 2: convert from padding to no-padding\n batch_td = left_right_2_no_padding(batch_td)\n ppo_mini_batch_size = self.config.critic.ppo_mini_batch_size\n ppo_mini_batch_size = ppo_mini_batch_size * self.config.actor_rollout_ref.rollout.n\n ppo_epochs = self.config.critic.ppo_epochs\n seed = self.config.critic.data_loader_seed\n shuffle = self.config.critic.shuffle\n tu.assign_non_tensor(\n batch_td,\n global_batch_size=ppo_mini_batch_size,\n mini_batch_size=ppo_mini_batch_size,\n epochs=ppo_epochs,\n seed=seed,\n dataloader_kwargs={\"shuffle\": shuffle},\n )\n\n output = self.critic_wg.train_mini_batch(batch_td)\n output = output.get()\n output = tu.get(output, \"metrics\")\n output = rename_dict(output, \"critic/\")\n # modify key name\n output[\"perf/mfu/critic\"] = output.pop(\"critic/mfu\")\n critic_output = DataProto.from_single_dict(data={}, meta_info={\"metrics\": output})\n else:\n critic_output = self.critic_wg.update_critic(batch)\n return critic_output\n\n def fit(self):\n \"\"\"\n The training loop of PPO.\n The driver process only need to call the compute functions of the worker group through RPC\n to construct the PPO dataflow.\n The light-weight advantage computation is done on the driver process.\n \"\"\"\n from omegaconf import OmegaConf\n\n from verl.utils.tracking import Tracking\n\n logger = Tracking(\n project_name=self.config.trainer.project_name,\n experiment_name=self.config.trainer.experiment_name,\n default_backend=self.config.trainer.logger,\n config=OmegaConf.to_container(self.config, resolve=True),\n )\n\n self.global_steps = 0\n\n # load checkpoint and update weights before doing anything\n self._load_checkpoint()\n self.checkpoint_manager.update_weights()\n\n current_epoch = self.global_steps // len(self.train_dataloader)\n\n # perform validation before training\n # currently, we only support validation using the reward_function.\n if self.config.trainer.get(\"val_before_train\", True):\n val_metrics = self._validate()\n assert val_metrics, f\"{val_metrics=}\"\n pprint(f\"Initial validation metrics: {val_metrics}\")\n logger.log(data=val_metrics, step=self.global_steps)\n if self.config.trainer.get(\"val_only\", False):\n return\n\n if self.config.actor_rollout_ref.rollout.get(\"skip_rollout\", False):\n rollout_skip = RolloutSkip(self.config, self.async_rollout_manager)\n rollout_skip.wrap_generate_sequences()\n\n # add tqdm\n progress_bar = tqdm(total=self.total_training_steps, initial=self.global_steps, desc=\"Training Progress\")\n\n # we start from step 1\n self.global_steps += 1\n last_val_metrics = None\n self.max_steps_duration = 0\n\n prev_step_profile = False\n curr_step_profile = (\n self.global_steps in self.config.global_profiler.steps\n if self.config.global_profiler.steps is not None\n else False\n )\n next_step_profile = False\n\n for epoch in range(current_epoch, self.config.trainer.total_epochs):\n for batch_dict in self.train_dataloader:\n if hasattr(self.actor_rollout_wg, \"async_calls_finalize_fn_exec\"):\n self.actor_rollout_wg.async_calls_finalize_fn_exec(blocking=False)\n metrics = {}\n timing_raw = {}\n\n with marked_timer(\"start_profile\", timing_raw):\n self._start_profiling(\n not prev_step_profile and curr_step_profile\n if self.config.global_profiler.profile_continuous_steps\n else curr_step_profile\n )\n batch: DataProto = DataProto.from_single_dict(batch_dict)\n batch.meta_info[\"temperature\"] = self.config.actor_rollout_ref.rollout.temperature\n\n # add uid to batch\n batch.non_tensor_batch[\"uid\"] = np.array(\n [str(uuid.uuid4()) for _ in range(len(batch.batch))], dtype=object\n )\n\n gen_batch = self._get_gen_batch(batch)\n\n # pass global_steps to trace\n gen_batch.meta_info[\"global_steps\"] = self.global_steps\n gen_batch_output = gen_batch.repeat(\n repeat_times=self.config.actor_rollout_ref.rollout.n, interleave=True\n )\n\n is_last_step = self.global_steps >= self.total_training_steps\n with marked_timer(\"step\", timing_raw):\n # generate a batch\n with marked_timer(\"gen\", timing_raw, color=\"red\"):\n if curr_step_profile:\n self.async_rollout_manager.start_profile()\n gen_batch_output = self.async_rollout_manager.generate_sequences(gen_batch_output)\n self.checkpoint_manager.sleep_replicas()\n if curr_step_profile:\n self.async_rollout_manager.stop_profile()\n\n timing_raw.update(gen_batch_output.meta_info[\"timing\"])\n gen_batch_output.meta_info.pop(\"timing\", None)\n\n if self.config.algorithm.adv_estimator == AdvantageEstimator.REMAX:\n with marked_timer(\"gen_max\", timing_raw, color=\"purple\"):\n gen_baseline_batch = deepcopy(gen_batch)\n gen_baseline_batch.meta_info[\"do_sample\"] = False\n if curr_step_profile:\n self.async_rollout_manager.start_profile()\n gen_baseline_output = self.async_rollout_manager.generate_sequences(gen_baseline_batch)\n self.checkpoint_manager.sleep_replicas()\n if curr_step_profile:\n self.async_rollout_manager.stop_profile()\n batch = batch.union(gen_baseline_output)\n # compute reward model score on batch\n rm_scores = None\n if self.use_rm and \"rm_scores\" not in batch.batch.keys():\n batch_reward = self._compute_reward_colocate(batch)\n batch = batch.union(batch_reward)\n\n # Compute or extract reward for REMAX baseline\n reward_baseline_tensor = batch.batch[\"rm_scores\"].sum(dim=-1)\n\n keys_to_pop = set(gen_baseline_output.batch.keys())\n if rm_scores is not None:\n keys_to_pop.update(rm_scores.batch.keys())\n batch.pop(batch_keys=list(keys_to_pop))\n\n batch.batch[\"reward_baselines\"] = reward_baseline_tensor\n\n del rm_scores, gen_baseline_batch, gen_baseline_output\n # repeat to align with repeated responses in rollout\n batch = batch.repeat(repeat_times=self.config.actor_rollout_ref.rollout.n, interleave=True)\n batch = batch.union(gen_batch_output)\n\n if \"response_mask\" not in batch.batch.keys():\n batch.batch[\"response_mask\"] = compute_response_mask(batch)\n # Balance the number of valid tokens across DP ranks.\n # NOTE: This usually changes the order of data in the `batch`,\n # which won't affect the advantage calculation (since it's based on uid),\n # but might affect the loss calculation (due to the change of mini-batching).\n if self.config.trainer.balance_batch:\n self._balance_batch(batch, metrics=metrics)\n\n # compute global_valid tokens\n batch.meta_info[\"global_token_num\"] = torch.sum(batch.batch[\"attention_mask\"], dim=-1).tolist()\n # get images_seqlens\n images_seqlens_all = []\n for multi_modal_input in batch.non_tensor_batch[\"multi_modal_inputs\"]:\n if \"image_grid_thw\" not in multi_modal_input.keys():\n continue\n images_seqlens_all.extend(multi_modal_input[\"images_seqlens\"].tolist())\n batch.meta_info[\"images_seqlens\"] = images_seqlens_all\n with marked_timer(\"reward\", timing_raw, color=\"yellow\"):\n # compute reward model score\n if self.use_rm and \"rm_scores\" not in batch.batch.keys():\n batch_reward = self._compute_reward_colocate(batch)\n batch = batch.union(batch_reward)\n\n # extract reward_tensor and reward_extra_infos_dict for training\n reward_tensor, reward_extra_infos_dict = extract_reward(batch)\n\n # Operating Mode Selection:\n # - Bypass mode: Sets old_log_probs = rollout_log_probs (2 policies: π_rollout, π_θ)\n # - Decoupled mode: Recomputes old_log_probs as proximal anchor (3 policies: π_rollout, π_old, π_θ)\n # Note: π_old computed once per data batch, serves as stable reference during mini-batch updates\n rollout_corr_config = self.config.algorithm.get(\"rollout_correction\", None)\n bypass_recomputing_logprobs = rollout_corr_config and rollout_corr_config.get(\"bypass_mode\", False)\n if bypass_recomputing_logprobs: # Use `rollout_log_probs`\n from verl.trainer.ppo.rollout_corr_helper import apply_bypass_mode\n\n apply_bypass_mode(\n batch=batch,\n rollout_corr_config=rollout_corr_config,\n policy_loss_config=self.config.actor_rollout_ref.actor.policy_loss,\n )\n else: # Recompute old_log_probs\n with marked_timer(\"old_log_prob\", timing_raw, color=\"blue\"):\n old_log_prob, old_log_prob_mfu = self._compute_old_log_prob(batch)\n entropys = old_log_prob.batch[\"entropys\"]\n response_masks = batch.batch[\"response_mask\"]\n actor_config = self.config.actor_rollout_ref.actor\n entropy_agg = agg_loss(\n loss_mat=entropys,\n loss_mask=response_masks,\n loss_agg_mode=actor_config.loss_agg_mode,\n loss_scale_factor=actor_config.loss_scale_factor,\n )\n old_log_prob_metrics = {\n \"actor/entropy\": entropy_agg.detach().item(),\n \"perf/mfu/actor_infer\": old_log_prob_mfu,\n }\n metrics.update(old_log_prob_metrics)\n old_log_prob.batch.pop(\"entropys\")\n if \"routed_experts\" in batch.batch and \"routed_experts\" in old_log_prob.batch:\n router_mode = getattr(\n self.config.actor_rollout_ref.actor.router_replay, \"mode\", \"disabled\"\n )\n if router_mode == \"R2\":\n batch.batch.pop(\"routed_experts\")\n else:\n old_log_prob.batch.pop(\"routed_experts\")\n batch = batch.union(old_log_prob)\n if \"rollout_log_probs\" in batch.batch.keys():\n # TODO: we may want to add diff of probs too.\n from verl.utils.debug.metrics import calculate_debug_metrics\n\n metrics.update(calculate_debug_metrics(batch))\n\n assert \"old_log_probs\" in batch.batch, f'\"old_log_prob\" not in {batch.batch.keys()=}'\n\n if self.use_reference_policy:\n # compute reference log_prob\n with marked_timer(str(Role.RefPolicy), timing_raw, color=\"olive\"):\n ref_log_prob = self._compute_ref_log_prob(batch)\n batch = batch.union(ref_log_prob)\n\n # compute values\n if self.use_critic:\n with marked_timer(\"values\", timing_raw, color=\"cyan\"):\n values = self._compute_values(batch)\n batch = batch.union(values)\n\n with marked_timer(\"adv\", timing_raw, color=\"brown\"):\n # we combine with rule-based rm\n reward_extra_infos_dict: dict[str, list]\n batch.batch[\"token_level_scores\"] = reward_tensor\n\n if reward_extra_infos_dict:\n batch.non_tensor_batch.update({k: np.array(v) for k, v in reward_extra_infos_dict.items()})\n\n # compute rewards. apply_kl_penalty if available\n if self.config.algorithm.use_kl_in_reward:\n batch, kl_metrics = apply_kl_penalty(\n batch, kl_ctrl=self.kl_ctrl_in_reward, kl_penalty=self.config.algorithm.kl_penalty\n )\n metrics.update(kl_metrics)\n else:\n batch.batch[\"token_level_rewards\"] = batch.batch[\"token_level_scores\"]\n\n # Compute rollout correction: IS weights, rejection sampling, and metrics\n # Only runs in decoupled mode (computes once per batch using stable π_old)\n # In bypass mode, this is skipped - actor computes metrics from evolving π_θ vs π_rollout\n if (\n rollout_corr_config is not None\n and \"rollout_log_probs\" in batch.batch\n and not bypass_recomputing_logprobs # Only in decoupled mode\n ):\n from verl.trainer.ppo.rollout_corr_helper import compute_rollout_correction_and_add_to_batch\n\n # Compute IS weights, apply rejection sampling, compute metrics\n batch, is_metrics = compute_rollout_correction_and_add_to_batch(batch, rollout_corr_config)\n # IS and off-policy metrics already have rollout_corr/ prefix\n metrics.update(is_metrics)\n\n # compute advantages, executed on the driver process\n norm_adv_by_std_in_grpo = self.config.algorithm.get(\n \"norm_adv_by_std_in_grpo\", True\n ) # GRPO adv normalization factor\n\n batch = compute_advantage(\n batch,\n adv_estimator=self.config.algorithm.adv_estimator,\n gamma=self.config.algorithm.gamma,\n lam=self.config.algorithm.lam,\n num_repeat=self.config.actor_rollout_ref.rollout.n,\n norm_adv_by_std_in_grpo=norm_adv_by_std_in_grpo,\n config=self.config.algorithm,\n )\n\n # update critic\n if self.use_critic:\n with marked_timer(\"update_critic\", timing_raw, color=\"pink\"):\n critic_output = self._update_critic(batch)\n critic_output_metrics = reduce_metrics(critic_output.meta_info[\"metrics\"])\n metrics.update(critic_output_metrics)\n\n # implement critic warmup\n if self.config.trainer.critic_warmup <= self.global_steps:\n # update actor\n with marked_timer(\"update_actor\", timing_raw, color=\"red\"):\n actor_output = self._update_actor(batch)\n\n # Check if the ESI (Elastic Server Instance)/training plan is close to expiration.\n esi_close_to_expiration = should_save_ckpt_esi(\n max_steps_duration=self.max_steps_duration,\n redundant_time=self.config.trainer.esi_redundant_time,\n )\n # Check if the conditions for saving a checkpoint are met.\n # The conditions include a mandatory condition (1) and\n # one of the following optional conditions (2/3/4):\n # 1. The save frequency is set to a positive value.\n # 2. It's the last training step.\n # 3. The current step number is a multiple of the save frequency.\n # 4. The ESI(Elastic Server Instance)/training plan is close to expiration.\n if self.config.trainer.save_freq > 0 and (\n is_last_step\n or self.global_steps % self.config.trainer.save_freq == 0\n or esi_close_to_expiration\n ):\n if esi_close_to_expiration:\n print(\"Force saving checkpoint: ESI instance expiration approaching.\")\n with marked_timer(\"save_checkpoint\", timing_raw, color=\"green\"):\n self._save_checkpoint()\n\n # update weights from trainer to rollout\n with marked_timer(\"update_weights\", timing_raw, color=\"red\"):\n self.checkpoint_manager.update_weights()\n\n actor_output_metrics = reduce_metrics(actor_output.meta_info[\"metrics\"])\n metrics.update(actor_output_metrics)\n\n # Log rollout generations if enabled\n rollout_data_dir = self.config.trainer.get(\"rollout_data_dir\", None)\n if rollout_data_dir:\n self._log_rollout_data(batch, reward_extra_infos_dict, timing_raw, rollout_data_dir)\n\n # validate\n if self.config.trainer.test_freq > 0 and (\n is_last_step or self.global_steps % self.config.trainer.test_freq == 0\n ):\n with marked_timer(\"testing\", timing_raw, color=\"green\"):\n val_metrics: dict = self._validate()\n if is_last_step:\n last_val_metrics = val_metrics\n metrics.update(val_metrics)\n\n with marked_timer(\"stop_profile\", timing_raw):\n next_step_profile = (\n self.global_steps + 1 in self.config.global_profiler.steps\n if self.config.global_profiler.steps is not None\n else False\n )\n self._stop_profiling(\n curr_step_profile and not next_step_profile\n if self.config.global_profiler.profile_continuous_steps\n else curr_step_profile\n )\n prev_step_profile = curr_step_profile\n curr_step_profile = next_step_profile\n\n steps_duration = timing_raw[\"step\"]\n self.max_steps_duration = max(self.max_steps_duration, steps_duration)\n\n # training metrics\n metrics.update(\n {\n \"training/global_step\": self.global_steps,\n \"training/epoch\": epoch,\n }\n )\n # collect metrics\n metrics.update(compute_data_metrics(batch=batch, use_critic=self.use_critic))\n metrics.update(compute_timing_metrics(batch=batch, timing_raw=timing_raw))\n # TODO: implement actual tflpo and theoretical tflpo\n n_gpus = self.resource_pool_manager.get_n_gpus()\n metrics.update(compute_throughout_metrics(batch=batch, timing_raw=timing_raw, n_gpus=n_gpus))\n # compute variance proxy metrics\n gradient_norm = metrics.get(\"actor/grad_norm\", None)\n metrics.update(compute_variance_proxy_metrics(batch=batch, gradient_norm=gradient_norm))\n # Note: mismatch metrics (KL, PPL, etc.) are collected at line 1179 after advantage computation\n\n # this is experimental and may be changed/removed in the future in favor of a general-purpose one\n if isinstance(self.train_dataloader.sampler, AbstractCurriculumSampler):\n self.train_dataloader.sampler.update(batch=batch)\n\n # TODO: make a canonical logger that supports various backend\n logger.log(data=metrics, step=self.global_steps)\n\n progress_bar.update(1)\n self.global_steps += 1\n\n if (\n hasattr(self.config.actor_rollout_ref.actor, \"profiler\")\n and self.config.actor_rollout_ref.actor.profiler.tool == \"torch_memory\"\n ):\n self.actor_rollout_wg.dump_memory_snapshot(\n tag=f\"post_update_step{self.global_steps}\", sub_dir=f\"step{self.global_steps}\"\n )\n\n if is_last_step:\n if hasattr(self.actor_rollout_wg, \"async_calls_finalize_fn_exec\"):\n self.actor_rollout_wg.async_calls_finalize_fn_exec(blocking=True)\n pprint(f\"Final validation metrics: {last_val_metrics}\")\n progress_bar.close()\n return\n\n # this is experimental and may be changed/removed in the future\n # in favor of a general-purpose data buffer pool\n if hasattr(self.train_dataset, \"on_batch_end\"):\n # The dataset may be changed after each training batch\n self.train_dataset.on_batch_end(batch=batch)\n"}61{"file_name": "verl__trainer__sft_trainer.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nimport os\nfrom functools import partial\n\nfrom tensordict.tensorclass import NonTensorData\n\nos.environ[\"NCCL_DEBUG\"] = \"WARN\"\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"true\"\n\nimport logging\n\nimport hydra\nimport torch\nimport torch.distributed\nfrom omegaconf import OmegaConf\nfrom torch.utils.data import DistributedSampler\nfrom torchdata.stateful_dataloader import StatefulDataLoader\nfrom tqdm import tqdm\n\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.checkpoint import CheckpointHandler\nfrom verl.utils.dataset.dataset_utils import SFTTensorCollator\nfrom verl.utils.dataset.multiturn_sft_dataset import MultiTurnSFTDataset\nfrom verl.utils.device import auto_set_device, get_device_name\nfrom verl.utils.distributed import destroy_global_process_group\nfrom verl.utils.logger import log_with_rank\nfrom verl.utils.memory_utils import aggressive_empty_cache\nfrom verl.utils.profiler import log_gpu_memory_usage\nfrom verl.utils.tracking import Tracking\nfrom verl.workers.engine_workers import TrainingWorker\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_SFT_LOGGING_LEVEL\", \"WARN\"))\n\n\nclass SFTTrainer:\n def __init__(\n self,\n config,\n ):\n self.config = config\n\n log_gpu_memory_usage(f\"rank {torch.distributed.get_rank()}: Before SFTTrainer init\", logger=logger)\n\n self.rank = torch.distributed.get_rank()\n\n self._build_config()\n self._build_dataset()\n\n self._build_engine()\n\n self._build_dataloader()\n\n self._init_engine()\n\n self._build_ckpt_handler()\n\n # Initialize resume-related variables\n self.resume_global_step = self.ckpt_handler.load_checkpoint()\n\n self.device_name = self.config.trainer.device\n\n if self.rank == 0:\n print(self.config)\n\n log_gpu_memory_usage(f\"rank {self.rank}: After SFTTrainer init\", logger=logger)\n\n def _build_ckpt_handler(self):\n resume_mode = getattr(self.config.trainer, \"resume_mode\", \"auto\")\n resume_from_path = getattr(self.config.trainer, \"resume_from_path\", None)\n max_ckpt_to_keep = getattr(self.config.trainer, \"max_ckpt_to_keep\", None)\n default_hdfs_dir = getattr(self.config.trainer, \"default_hdfs_dir\", None)\n\n self.ckpt_handler = CheckpointHandler(\n engine=self.engine,\n train_dataloader=self.train_dataloader,\n default_local_dir=self.config.trainer.default_local_dir,\n max_ckpt_to_keep=max_ckpt_to_keep,\n default_hdfs_dir=default_hdfs_dir,\n resume_mode=resume_mode,\n resume_from_path=resume_from_path,\n )\n\n def _build_config(self):\n from verl.utils.config import omega_conf_to_dataclass\n\n self.model_config = omega_conf_to_dataclass(self.config.model)\n self.engine_config = omega_conf_to_dataclass(self.config.engine)\n self.optimizer_config = omega_conf_to_dataclass(self.config.optim)\n self.checkpoint_config = omega_conf_to_dataclass(self.config.checkpoint)\n self.profiler_config = omega_conf_to_dataclass(self.config.profiler)\n\n # check profile interval\n self.profiler_interval = self.config.trainer.profile_interval\n self._validate_profiler_interval()\n\n def _validate_profiler_interval(self):\n assert len(self.profiler_interval) == 2\n self.start_profile_step = self.profiler_interval[0]\n self.end_profile_step = self.profiler_interval[1]\n assert self.end_profile_step >= self.start_profile_step\n if self.start_profile_step < 0:\n assert self.end_profile_step < 0\n\n def _build_engine(self):\n from verl.workers.engine_workers import TrainingWorkerConfig\n from verl.workers.utils.losses import sft_loss\n\n self.loss_fn = partial(sft_loss, config=None)\n\n config = TrainingWorkerConfig(\n model_type=\"language_model\",\n model_config=self.model_config,\n engine_config=self.engine_config,\n optimizer_config=self.optimizer_config,\n checkpoint_config=self.checkpoint_config,\n profiler_config=self.profiler_config,\n )\n\n self.training_client = TrainingWorker(config=config)\n self.training_client.set_loss_fn(loss_fn=self.loss_fn)\n # Note that in SPMD world, this abstraction has to break\n self.engine = self.training_client.engine\n\n def _init_engine(self):\n # patch optimizer config\n if self.config.trainer.total_training_steps is not None:\n self.total_training_steps = self.config.trainer.total_training_steps\n else:\n self.total_training_steps = len(self.train_dataloader) * self.config.trainer.total_epochs\n self.optimizer_config.total_training_steps = self.total_training_steps\n\n self.steps_per_epoch = len(self.train_dataloader)\n\n # manage save and test frequency\n self.save_freq = self.config.trainer.save_freq\n if self.save_freq == \"after_each_epoch\":\n self.save_freq = self.steps_per_epoch\n\n self.test_freq = self.config.trainer.test_freq\n if self.test_freq == \"after_each_epoch\":\n self.test_freq = self.steps_per_epoch\n\n self.training_client.reset()\n\n def _build_dataset(self):\n config = self.config\n tokenizer = self.model_config.tokenizer\n processor = self.model_config.processor\n train_dataset = create_sft_dataset(\n config.data.train_files,\n config.data,\n tokenizer,\n processor,\n max_samples=config.data.get(\"train_max_samples\", -1),\n )\n if config.data.val_files:\n val_dataset = create_sft_dataset(\n config.data.val_files,\n config.data,\n tokenizer,\n processor,\n max_samples=config.data.get(\"val_max_samples\", -1),\n )\n else:\n val_dataset = None\n\n self.train_dataset, self.val_dataset = train_dataset, val_dataset\n\n def _build_dataloader(self):\n # build dataset\n config = self.config\n # build dataloader\n # Use data parallel rank and size instead of global rank and world size\n\n # Set pin_memory_device when pin_memory is enabled.\n device_name = get_device_name()\n\n dp_rank = self.engine.get_data_parallel_rank()\n dp_size = self.engine.get_data_parallel_size()\n\n self.train_sampler = DistributedSampler(\n self.train_dataset, shuffle=True, num_replicas=dp_size, rank=dp_rank, drop_last=True\n )\n\n self.global_batch_size = config.data.train_batch_size\n self.train_batch_size_per_dp = self.global_batch_size // dp_size\n self.collate_fn = SFTTensorCollator(config.data.pad_mode)\n\n self.train_dataloader = StatefulDataLoader(\n dataset=self.train_dataset,\n batch_size=self.train_batch_size_per_dp,\n sampler=self.train_sampler,\n collate_fn=self.collate_fn,\n num_workers=self.config.data.num_workers,\n pin_memory=False,\n drop_last=True,\n pin_memory_device=device_name,\n )\n\n if self.val_dataset:\n self.val_sampler = DistributedSampler(\n self.val_dataset, shuffle=False, num_replicas=dp_size, rank=dp_rank, drop_last=True\n )\n self.val_dataloader = StatefulDataLoader(\n dataset=self.val_dataset,\n batch_size=self.train_batch_size_per_dp,\n sampler=self.val_sampler,\n collate_fn=self.collate_fn,\n num_workers=self.config.data.num_workers,\n pin_memory=False,\n drop_last=True,\n pin_memory_device=device_name,\n )\n else:\n self.val_dataloader = None\n\n def _get_batch_seqlens(self, data):\n # mean over dp group\n is_nested = data[\"input_ids\"].is_nested\n if is_nested:\n batch_seqlens: torch.Tensor = data[\"input_ids\"].offsets().diff()\n else:\n batch_seqlens: torch.Tensor = data[\"attention_mask\"].sum(dim=-1)\n batch_seqlens = batch_seqlens.to(self.device_name) # (global_bsz // dp)\n\n output_tensor = torch.empty(\n (batch_seqlens.shape[0] * self.engine.get_data_parallel_size(),),\n dtype=batch_seqlens.dtype,\n device=self.device_name,\n ) # (global_bsz,)\n\n torch.distributed.all_gather_into_tensor(\n output_tensor=output_tensor,\n input_tensor=batch_seqlens,\n group=self.engine.get_data_parallel_group(),\n )\n\n batch_seqlens = output_tensor.tolist()\n return batch_seqlens\n\n def fit(self):\n is_logging = self.engine.is_mp_src_rank_with_outputs() and self.engine.get_data_parallel_rank() == 0\n\n # TODO: add a unified tracking\n if is_logging:\n tracking = Tracking(\n project_name=self.config.trainer.project_name,\n experiment_name=self.config.trainer.experiment_name,\n default_backend=self.config.trainer.logger,\n config=OmegaConf.to_container(self.config, resolve=True),\n )\n\n global_step = self.resume_global_step # Start from resumed step\n last_valid_metric = None\n\n log_with_rank(\n f\"Total training steps: {self.total_training_steps},\",\n logger=logger,\n rank=0,\n log_only_rank_0=True,\n )\n\n # With StatefulDataLoader, we don't need to manually calculate epochs and steps\n # The dataloader will automatically resume from where it left off\n if global_step > 0:\n log_with_rank(\n f\"StatefulDataLoader will automatically resume from global step: {global_step}\",\n logger=logger,\n rank=0,\n log_only_rank_0=True,\n )\n\n # Calculate which epoch we're starting from for sampler.set_epoch()\n start_epoch = global_step // self.steps_per_epoch\n\n meta_info = {\n \"use_remove_padding\": self.config.model.use_remove_padding,\n \"use_dynamic_bsz\": self.config.data.use_dynamic_bsz,\n \"max_token_len_per_gpu\": self.config.data.max_token_len_per_gpu,\n \"micro_batch_size_per_gpu\": self.config.data.micro_batch_size_per_gpu,\n \"temperature\": 1.0,\n \"global_batch_size\": self.global_batch_size,\n \"pad_mode\": self.config.data.pad_mode,\n \"pad_token_id\": self.model_config.tokenizer.pad_token_id,\n }\n\n train_time = 0\n total_tokens = 0\n for epoch in range(start_epoch, self.config.trainer.total_epochs):\n self.train_sampler.set_epoch(epoch=epoch)\n\n aggressive_empty_cache(force_sync=True)\n log_gpu_memory_usage(f\"rank {self.rank}: At start of epoch {epoch}\", logger=logger)\n\n for step_in_epoch, data in enumerate(\n tqdm(\n self.train_dataloader,\n initial=global_step % self.steps_per_epoch if epoch == start_epoch else 0,\n total=self.steps_per_epoch,\n desc=f\"Epoch {epoch + 1}/{self.config.trainer.total_epochs}\",\n disable=not is_logging,\n )\n ):\n global_step += 1\n\n # construct tensordict\n data = tu.get_tensordict(tensor_dict=data, non_tensor_dict=meta_info)\n batch_seqlens = self._get_batch_seqlens(data=data)\n # this is necessary. Otherwise, it is interpreted as NonTensorStack\n batch_seqlens_ntd = NonTensorData(batch_seqlens)\n\n tu.assign_non_tensor(data, update_lr_scheduler=True, global_token_num=batch_seqlens_ntd)\n\n # start profile in SPMD mode\n if global_step == self.start_profile_step:\n self.training_client.start_profile()\n # train for on batch\n output = self.training_client.train_batch(data=data)\n\n if global_step == self.end_profile_step:\n self.training_client.stop_profile()\n\n if self.engine.is_mp_src_rank_with_outputs():\n metrics = tu.get(output, \"metrics\")\n\n # TODO: we can actual accumulate metrics for N steps and perform aggregate metrics\n for k in [\"loss\", \"grad_norm\", \"lr\", \"mfu\"]:\n if k in metrics.keys():\n value = metrics.pop(k)\n metrics[f\"train/{k}\"] = value\n\n metrics[\"train/global_tokens\"] = torch.sum(\n torch.tensor(batch_seqlens, device=self.device_name)\n ).item()\n total_tokens += metrics[\"train/global_tokens\"]\n metrics[\"train/total_tokens(B)\"] = total_tokens / 1e9\n\n if self.engine.get_data_parallel_rank() == 0:\n tracking.log(data=metrics, step=global_step)\n\n is_last_step = global_step >= self.total_training_steps\n is_valid_step = global_step % self.test_freq == 0\n is_save_step = global_step % self.save_freq == 0\n\n # early exit or validation step\n if is_last_step and self.val_dataloader is not None or (self.test_freq > 0 and is_valid_step):\n # Perform validation\n val_losses = []\n for val_data in self.val_dataloader:\n val_data = tu.get_tensordict(tensor_dict=val_data, non_tensor_dict=meta_info)\n output = self.training_client.infer_batch(val_data)\n\n if self.engine.is_mp_src_rank_with_outputs():\n metrics = tu.get(output, \"metrics\")\n val_losses.append(metrics[\"loss\"])\n\n if self.engine.is_mp_src_rank_with_outputs():\n val_loss = torch.mean(torch.tensor(val_losses, device=self.device_name))\n # average over data parallel group\n torch.distributed.all_reduce(\n val_loss, op=torch.distributed.ReduceOp.AVG, group=self.engine.get_data_parallel_group()\n )\n\n if is_logging:\n metric = {\"val/loss\": val_loss.detach().item()}\n tracking.log(data=metric, step=global_step)\n last_valid_metric = metric\n torch.distributed.barrier()\n\n if is_last_step or (self.save_freq > 0 and is_save_step):\n aggressive_empty_cache(force_sync=True)\n self.ckpt_handler.save_checkpoint(step=global_step)\n\n if is_last_step:\n if is_logging:\n print(f\"Total time for train steps: {train_time:.2f}s\")\n print(f\"Final validation metrics: {last_valid_metric}\")\n return\n\n\ndef run_sft(config):\n from verl.utils.distributed import initialize_global_process_group\n\n initialize_global_process_group()\n trainer = SFTTrainer(config=config)\n trainer.fit()\n destroy_global_process_group()\n\n\n@hydra.main(config_path=\"config\", config_name=\"sft_trainer_engine\", version_base=None)\ndef main(config):\n # Automatically set `config.trainer.device = npu` when running on Ascend NPU.\n auto_set_device(config)\n run_sft(config)\n\n\ndef create_sft_dataset(data_paths, data_config, tokenizer, processor, max_samples=-1):\n \"\"\"Create a dataset.\"\"\"\n # build dataset\n # First check if a custom dataset class is specified\n if data_config.custom_cls.get(\"path\", None):\n from verl.utils.import_utils import load_extern_object\n\n dataset_cls = load_extern_object(data_config.custom_cls.path, data_config.custom_cls.name)\n else:\n # Default to multi-turn dataset\n dataset_cls = MultiTurnSFTDataset\n\n # Create datasets based on the selected class\n dataset = dataset_cls(\n parquet_files=data_paths, tokenizer=tokenizer, config=data_config, processor=processor, max_samples=max_samples\n )\n return dataset\n\n\nif __name__ == \"__main__\":\n main()\n"}62{"file_name": "verl__utils__activation_offload.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2022-2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\"\"\"Functionality for CPU offloading of tensors saved for backward pass.\"\"\"\n\nfrom __future__ import annotations\n\nimport functools\nimport logging\nimport os\nfrom typing import Any, Optional\n\nimport torch\nfrom torch.distributed.fsdp import FullyShardedDataParallel as FSDP\n\nfrom verl.utils.device import get_torch_device\nfrom verl.utils.fsdp_utils import FSDPModule as FSDP2\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef _get_unique_tensor_key(tensor):\n key = (tensor.untyped_storage().data_ptr() + tensor.storage_offset(), tensor.dtype)\n return key\n\n\nclass FSDPParameterFilter:\n def __init__(self):\n self.model_parameters_storage = set()\n\n def __call__(self, tensor):\n return tensor.untyped_storage().data_ptr() not in self.model_parameters_storage\n\n def update_model_parameters(self, model):\n new_storage = set()\n for p in model.parameters():\n new_storage.add(p.data.untyped_storage().data_ptr())\n self.model_parameters_storage = new_storage\n\n\nclass CpuOffloadHookWithOffloadHandler:\n \"\"\"Context-manager that offloads/recovers tensors through an offload hander.\n\n The hook just offloads/recovers the tensor object to the handler through `tensor_push`\n and `tensor_pop` interface. How the offload-handler manages the offloading, recovering\n or prefetching timing is transparent to this hook.\n \"\"\"\n\n def __init__(\n self,\n offload_handler: OffloadHandler,\n handler_extra_kwargs: Optional[dict[str, Any]] = None,\n ) -> None:\n if handler_extra_kwargs is None:\n handler_extra_kwargs = {}\n self.offload_handler: OffloadHandler = offload_handler\n self.handler_extra_kwargs: dict[str, Any] = handler_extra_kwargs\n self.inside_context = False\n\n def __enter__(self):\n self.inside_context = True\n torch._C._autograd._push_saved_tensors_default_hooks(self.on_save_for_backward, self.on_get_saved_tensor)\n\n def __exit__(self, *args: Any):\n self.inside_context = False\n torch._C._autograd._pop_saved_tensors_default_hooks()\n\n def on_save_for_backward(self, tensor: torch.Tensor) -> Any:\n retrieve_identifier = self.offload_handler.tensor_push(tensor, **self.handler_extra_kwargs)\n return retrieve_identifier\n\n def on_get_saved_tensor(self, saved_state: Any) -> torch.Tensor:\n tensor = self.offload_handler.tensor_pop(saved_state, **self.handler_extra_kwargs)\n return tensor\n\n\nclass OffloadHandler:\n \"\"\"A base class for CPU offload-handler.\"\"\"\n\n def __init__(self) -> None:\n pass\n\n def tensor_push(self, tensor: torch.Tensor, **kwargs) -> Any:\n \"\"\"Tensor push.\"\"\"\n raise NotImplementedError(\n \"`tensor_push is not implented in OffloadHandler class. Inherit this class and implement your \"\n \"custom tensor_push.\"\n )\n\n def tensor_pop(self, tensor_tag: Any, **kwargs):\n \"\"\"Tensor pop.\"\"\"\n raise NotImplementedError(\n \"`tensor_pop is not implented in OffloadHandler class. Inherit this class and implement your \"\n \"custom tensor_pop.\"\n )\n\n\nclass GroupCommitFunction(torch.autograd.Function):\n \"\"\"this is a dummy op with output identical to input.\n However, it is necessary for marking a timepoint for offload handler to\n accomplish all synchronizations. Implementing it as a function is necessary\n because we need to actions in both forward and backward.\n \"\"\"\n\n @staticmethod\n def forward(ctx, tensor, cpu_offload_handler):\n # pylint: disable=missing-function-docstring\n cpu_offload_handler.on_group_commit_forward()\n ctx.cpu_offload_handler = cpu_offload_handler\n # return the identical tensor\n return tensor\n\n @staticmethod\n def backward(ctx, grad_output):\n # pylint: disable=missing-function-docstring\n cpu_offload_handler = ctx.cpu_offload_handler\n cpu_offload_handler.on_group_commit_backward()\n return grad_output, None\n\n\ngroup_prefetch_offload_commit = GroupCommitFunction.apply\n\n\nclass SynchronizedGroupOffloadHandler(OffloadHandler):\n \"\"\"Offload Handler that offloads/reloads in a synchronized way.\n The device-to-host and host-to-device copying happen in the same stream\n as the computation kernels, thus the copying will block computation.\n \"\"\"\n\n def __init__(self, num_offload_group, tensor_need_offloading_checker=(lambda _: True)) -> None:\n super().__init__()\n\n self.num_offload_group = num_offload_group\n self.tensor_need_offloading_checker = tensor_need_offloading_checker\n\n self.groupid_reset()\n\n def groupid_reset(self):\n \"\"\"Groupid reset.\"\"\"\n # Data structures to label saved tensors and book-keep their cpu copies.\n # Currently, on push, create a new cpu tensor and copies; on pop, copies\n # the tensor back to gpu and deletes the cpu tensor.\n # These will increment whenever `group_commit()` is invoked\n self.current_group, self.tensor_count_current_group = (0, 0)\n self.torch_tensor_count = 0\n self.tensor_tag_to_state = {}\n\n def on_group_commit_forward(self):\n \"\"\"On group commit forward.\"\"\"\n # finishing up with updating current group and tensor count\n self.current_group += 1 # increment\n self.tensor_count_current_group = 0 # reset\n\n def on_group_commit_backward(self):\n \"\"\"On group commit backward.\"\"\"\n self.current_group -= 1\n assert self.current_group >= 0\n\n @staticmethod\n def offload(src_tensor, pin_memory=True):\n \"\"\"Offload.\"\"\"\n\n cpu_backup = torch.empty(\n src_tensor.size(),\n dtype=src_tensor.dtype,\n layout=src_tensor.layout,\n device=\"cpu\",\n pin_memory=pin_memory,\n )\n cpu_backup.copy_(src_tensor, non_blocking=True)\n state = (src_tensor.device, cpu_backup)\n return state\n\n @staticmethod\n def reload(state, non_blocking=None):\n \"\"\"Reload.\"\"\"\n dev, cpu_backup = state\n if non_blocking is None:\n non_blocking = cpu_backup.is_pinned()\n return cpu_backup.to(dev, non_blocking=non_blocking)\n\n def tensor_push(self, tensor: torch.Tensor, **kwargs):\n \"\"\"Tensor push.\"\"\"\n # obtain a unique tensor tag\n tensor_tag = (self.current_group, self.tensor_count_current_group)\n self.tensor_count_current_group += 1\n assert tensor_tag not in self.tensor_tag_to_state\n if self.current_group < self.num_offload_group and self.tensor_need_offloading_checker(tensor):\n state = SynchronizedGroupOffloadHandler.offload(tensor)\n self.tensor_tag_to_state[tensor_tag] = state\n else:\n # will be offloaded together after group commit\n self.tensor_tag_to_state[tensor_tag] = tensor\n\n return tensor_tag\n\n def tensor_pop(self, tensor_tag, **kwargs):\n \"\"\"Tensor pop.\"\"\"\n assert tensor_tag in self.tensor_tag_to_state\n state = self.tensor_tag_to_state.pop(tensor_tag)\n if isinstance(state, tuple):\n tensor = SynchronizedGroupOffloadHandler.reload(state)\n else:\n tensor = state\n return tensor\n\n\nclass AsyncDoubleBufferGroupOffloadHandler(SynchronizedGroupOffloadHandler):\n \"\"\"Compared to synchronize, this uses more memory because of the buffer but\n achieves better performance due to the overlapping. D2h and h2d copying are\n completely hidden behind computation if computation time of a layer is longer\n than host-device communication time. Bulk offloading with delay and bulk reloading\n with prefetch are implemented.\"\"\"\n\n def __init__(\n self,\n num_offload_group, # must be <= actual number of groups (number of commits)\n num_model_group,\n tensor_need_offloading_checker=(lambda t: True),\n ) -> None:\n super().__init__(\n num_offload_group=num_offload_group,\n tensor_need_offloading_checker=tensor_need_offloading_checker,\n )\n # Number of layers in the model\n self.num_layers = num_model_group\n # Data Structure to maintain reference to activation tensors\n self.tensor_tag_to_buf = {}\n # Tracking the number of layers offloaded\n self.offloaded_group_count = 0\n # Core data structure that decides the window for offloading\n self.layer_window_map = {}\n self.group_offload_mapping = {}\n\n # Logic to make offloading load balance across computation\n # for optimal CPU/GPU interconnect usage\n constant = 0\n for i in range(self.num_offload_group):\n self.layer_window_map[i] = ((self.num_layers // self.num_offload_group) * (i + 1)) - 1\n if i < (self.num_layers % self.num_offload_group):\n self.layer_window_map[i] += i + 1\n constant = i + 1\n else:\n self.layer_window_map[i] += constant\n\n # allocate streams and events for synchronization\n self.d2h_stream = get_torch_device().Stream()\n self.h2d_stream = get_torch_device().Stream()\n\n def tensor_push(self, tensor: torch.Tensor, **kwargs) -> Any:\n torch_stray_tensor = isinstance(\n tensor,\n torch._subclasses.fake_tensor.FakeTensor | torch._subclasses.functional_tensor.FunctionalTensor,\n )\n need_offload = not torch_stray_tensor\n need_offload = need_offload and self.tensor_need_offloading_checker(tensor)\n\n if need_offload:\n # obtain a unique tensor tag\n tensor_tag = (self.current_group, self.tensor_count_current_group)\n self.tensor_count_current_group += 1\n\n assert tensor_tag not in self.tensor_tag_to_state\n self.tensor_tag_to_state[tensor_tag] = tensor\n\n if self.current_group < self.num_offload_group:\n self.tensor_tag_to_buf[tensor_tag] = tensor\n else:\n tensor_tag = tensor\n return tensor_tag\n\n def tensor_pop(self, tensor_tag, **kwargs):\n \"\"\"Tensor pop.\"\"\"\n if isinstance(tensor_tag, torch.Tensor):\n return tensor_tag\n assert tensor_tag in self.tensor_tag_to_state\n tensor = self.tensor_tag_to_state.pop(tensor_tag)\n self.tensor_tag_to_buf.pop(tensor_tag, None)\n\n # the tensor should have been copied back in on_group_commit_backward()\n # which invokes bulk_reload_group.\n assert not isinstance(tensor, tuple)\n return tensor\n\n def bulk_offload_group(self, group_to_offload):\n \"\"\"Bulk offload group.\"\"\"\n offload_mapping = {}\n offload_size = 0\n with get_torch_device().stream(self.d2h_stream):\n for tensor_tag, state in self.tensor_tag_to_state.items():\n group_id, _ = tensor_tag\n if group_id == group_to_offload:\n assert not isinstance(state, tuple)\n key = _get_unique_tensor_key(state)\n if key not in offload_mapping:\n offload_mapping[key] = state\n # if offload, return the reference to cpu copy\n self.tensor_tag_to_state[tensor_tag] = (key, state.shape)\n for key, tensor in offload_mapping.items():\n state = SynchronizedGroupOffloadHandler.offload(tensor)\n offload_size += tensor.numel() * tensor.element_size()\n offload_mapping[key] = state\n\n self.group_offload_mapping[group_to_offload] = offload_mapping\n\n def synchronize_on_group_commit_forward(self, current_group):\n \"\"\"Synchronize on group commit forward.\"\"\"\n\n # For the first group, kickstart the offload after we have\n # the first compute completion\n if current_group == 0:\n self.d2h_stream.wait_stream(get_torch_device().current_stream())\n self.bulk_offload_group(current_group)\n\n # Window map data structure helps us synchronize based on number\n # of layers offloaded\n if self.layer_window_map[self.offloaded_group_count] == current_group:\n # Stream synchronization both ways\n self.d2h_stream.wait_stream(get_torch_device().current_stream())\n get_torch_device().current_stream().wait_stream(self.d2h_stream)\n\n # Time to free the activation memory after usage\n for tensor_tag, _ in self.tensor_tag_to_buf.items():\n if tensor_tag[0] == self.offloaded_group_count:\n self.tensor_tag_to_buf[tensor_tag] = None\n\n # Time to offload the next group\n if self.offloaded_group_count < (self.num_offload_group - 1):\n self.bulk_offload_group(self.offloaded_group_count + 1)\n\n # Increment the offload group count to keep track\n self.offloaded_group_count += 1\n\n def on_group_commit_forward(self):\n \"\"\"This function will cause host device synchronization\"\"\"\n # handle synchronization events\n self.synchronize_on_group_commit_forward(self.current_group)\n\n super().on_group_commit_forward()\n\n @torch.no_grad\n def bulk_reload_group(self, group_to_reload):\n \"\"\"Bulk reload group.\"\"\"\n assert group_to_reload < self.num_offload_group\n\n with get_torch_device().stream(self.h2d_stream):\n # move back tensors\n offload_mapping = self.group_offload_mapping.pop(group_to_reload)\n assert offload_mapping is not None\n for key, state in offload_mapping.items():\n offload_mapping[key] = SynchronizedGroupOffloadHandler.reload(state)\n for tensor_label, state in self.tensor_tag_to_state.items():\n group_id, _ = tensor_label\n if group_id == group_to_reload and not isinstance(state, torch.Tensor):\n assert isinstance(state, tuple), f\"{group_id} {state}\"\n key, shape = state\n recovered_tensor = offload_mapping[key].view(shape)\n self.tensor_tag_to_state[tensor_label] = recovered_tensor\n\n def on_group_commit_backward(self):\n # first decrement the current group.\n # after last commit in forward, the group will +1; in backward it -1.\n # Finally it should be decremented to 0.\n self.current_group -= 1\n assert self.current_group >= 0\n\n # Layer window data structure helps us to reload at right times\n if self.layer_window_map[self.offloaded_group_count - 1] == self.current_group:\n # Stream synchronization both ways\n self.h2d_stream.wait_stream(get_torch_device().current_stream())\n get_torch_device().current_stream().wait_stream(self.h2d_stream)\n\n # Time to reload the next group\n self.bulk_reload_group(self.offloaded_group_count - 1)\n\n # Decrease the offloading group counter\n self.offloaded_group_count -= 1 if self.offloaded_group_count > 1 else 0\n\n # Last group computation needs to wait till all the reloads complete\n if self.current_group == 0:\n get_torch_device().current_stream().wait_stream(self.h2d_stream)\n self.offloaded_group_count = 0\n\n\ndef get_activation_offload_context(\n num_layers: int = 1, model_layers: int = 1, tensor_need_offloading_checker=(lambda t: True)\n):\n cpu_offload_handler = AsyncDoubleBufferGroupOffloadHandler(\n num_offload_group=num_layers,\n num_model_group=model_layers,\n tensor_need_offloading_checker=tensor_need_offloading_checker,\n )\n\n def group_prefetch_offload_commit_async(tensor):\n return group_prefetch_offload_commit(tensor, cpu_offload_handler)\n\n return (\n CpuOffloadHookWithOffloadHandler(offload_handler=cpu_offload_handler),\n group_prefetch_offload_commit_async,\n )\n\n\nclass ActivationHandler:\n def __init__(self, offload_ctx, sync_func, tensor_filter, enable_ckpt):\n self._offload_ctx = offload_ctx\n self._sync_func = sync_func\n self._enable_ckpt = enable_ckpt\n self._tensor_filter = tensor_filter\n if enable_ckpt:\n self.checkpoint_fn = functools.partial(\n torch.utils.checkpoint.checkpoint,\n use_reentrant=True,\n )\n\n def pre_forward(self, module):\n if module.training:\n self._offload_ctx.__enter__()\n self._tensor_filter.update_model_parameters(module)\n\n def post_forward(self, module):\n if module.training:\n self._offload_ctx.__exit__(None, None, None)\n\n def _pack_kwargs(self, *args, **kwargs):\n kwarg_keys = []\n flat_args = list(args)\n for k, v in kwargs.items():\n kwarg_keys.append(k)\n flat_args.append(v)\n\n return tuple(flat_args), tuple(kwarg_keys)\n\n def _unpack_kwargs(self, flat_args, kwarg_keys):\n assert len(kwarg_keys) <= len(flat_args), f\"too many keys {len(kwarg_keys)} vs. {len(flat_args)}\"\n if len(kwarg_keys) == 0:\n return flat_args, {}\n args = flat_args[: -len(kwarg_keys)]\n kwargs = dict(zip(kwarg_keys, flat_args[-len(kwarg_keys) :], strict=True))\n return args, kwargs\n\n def _ckpt_forward(self, forward_method, *args, **kwargs):\n flat_args, kwarg_keys = self._pack_kwargs(*args, **kwargs)\n\n def my_function(*inputs):\n # unpack back into args and kwargs\n nonlocal forward_method, kwarg_keys\n unpacked_args, unpacked_kwargs = self._unpack_kwargs(inputs, kwarg_keys)\n # run original module\n return forward_method(*unpacked_args, **unpacked_kwargs)\n\n return self.checkpoint_fn(\n my_function,\n *flat_args,\n )\n\n def forward(self, module, forward_method, *args, **kwargs):\n if not module.training:\n return forward_method(*args, **kwargs)\n if not self._enable_ckpt:\n ret = forward_method(*args, **kwargs)\n else:\n ret = self._ckpt_forward(forward_method, *args, **kwargs)\n binded_tensor = ret\n if isinstance(ret, tuple):\n binded_tensor = ret[0]\n binded_tensor = self._sync_func(binded_tensor)\n final_ret = binded_tensor\n if isinstance(ret, tuple):\n final_ret = (final_ret,) + ret[1:]\n return final_ret\n\n def wrap_module_forward_method(self, module):\n orig_method = module.forward\n handler = self\n\n @functools.wraps(orig_method)\n def wrapped_method(model_self, *args, **kwargs):\n nonlocal handler\n handler.pre_forward(model_self)\n out = handler.forward(model_self, orig_method, *args, **kwargs)\n handler.post_forward(model_self)\n return out\n\n module.forward = wrapped_method.__get__(module, type(module))\n\n\ndef enable_activation_offloading(model, strategy, enable_ckpt=False):\n \"\"\"\n Enable activation offloading for the model. It groups activations by TransformerLayer and offloads activation\n groups asynchronously. This means that the offloading of the i-th activation group and the computation of the i+1-th\n activation group happen at the same time, and there are at most two activation groups in GPU memory.\n\n Args:\n model: the model to enable activation offloading\n strategy: the training strategy of the model, such as \"fsdp\"\n enable_ckpt: whether activation checkpointing(also called gradient checkpointing) has been enabled for the model\n\n Note:\n For best efficiency, activation offloading is usually combined with activation checkpointing. However, this\n implementation of activation offloading is conflicted with the implementation of activation checkpointing in\n some training strategies. This function resolves this conflict, and therefore requires the \"strategy\" and\n \"enable_ckpt\" arguments.\n\n Returns:\n\n \"\"\"\n\n assert strategy == \"fsdp\" or strategy == \"fsdp2\", \"activation offloading only supports fsdp strategy\"\n layers = []\n\n def get_layers(module):\n for name, child in module.named_children():\n if not isinstance(child, FSDP | FSDP2):\n get_layers(child)\n else:\n wrapped_module = child\n if isinstance(child, FSDP):\n wrapped_module = child._fsdp_wrapped_module\n # In some cases, torch.nn.Embedding is wrapped with FSDP alone. However, the activation\n # size of torch.nn.Embedding is small, so it's not necessary to offload it.\n if not isinstance(wrapped_module, torch.nn.Embedding):\n layers.append(child)\n\n get_layers(model)\n if len(layers) < 3:\n logger.warning(f\"Find only {len(layers)} fsdp layers, not necessary to enable async activation offloading\")\n return\n\n tensor_filter = FSDPParameterFilter()\n context, sync_func = get_activation_offload_context(len(layers) - 1, len(layers), tensor_filter)\n if enable_ckpt:\n # The implementation of activation checkpointing in transformers library is incompatible with\n # activation offloading,\n # so it will be disabled, but this implementation supports another version of activation checkpointing, so that\n # these two features can be enabled at the same time.\n for module in model.modules():\n if hasattr(module, \"gradient_checkpointing_disable\"):\n module.gradient_checkpointing_disable()\n\n handler = ActivationHandler(context, sync_func, tensor_filter, enable_ckpt)\n for layer in layers:\n module = layer\n if isinstance(layer, FSDP):\n module = module._fsdp_wrapped_module\n handler.wrap_module_forward_method(module)\n"}63{"file_name": "verl__utils__attention_utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Callable\n\n_index_first_axis, _pad_input, _rearrange, _unpad_input = None, None, None, None\n\n\ndef _get_attention_functions() -> tuple[Callable, Callable, Callable, Callable]:\n \"\"\"Dynamically import attention functions based on available hardware.\"\"\"\n\n from verl.utils.device import is_torch_npu_available\n\n global _index_first_axis, _pad_input, _rearrange, _unpad_input\n\n if is_torch_npu_available(check_device=False):\n from verl.utils.npu_flash_attn_utils import index_first_axis, pad_input, rearrange, unpad_input\n else:\n from flash_attn.bert_padding import index_first_axis, pad_input, rearrange, unpad_input\n\n _index_first_axis, _pad_input, _rearrange, _unpad_input = index_first_axis, pad_input, rearrange, unpad_input\n\n return _index_first_axis, _pad_input, _rearrange, _unpad_input\n\n\ndef index_first_axis(*args, **kwargs):\n \"\"\"\n Unified entry point for `index_first_axis` across CUDA and NPU backends.\n\n Dynamically dispatches to the appropriate device-specific implementation:\n - On CUDA: `flash_attn.bert_padding.index_first_axis`\n - On NPU: `transformers.integrations.npu_flash_attention.index_first_axis`\n (falls back to `transformers.modeling_flash_attention_utils._index_first_axis`\n in newer versions of transformers).\n\n Users can call this function directly without worrying about the underlying device.\n \"\"\"\n func, *_ = _get_attention_functions()\n return func(*args, **kwargs)\n\n\ndef pad_input(*args, **kwargs):\n \"\"\"\n Unified entry point for `pad_input` across CUDA and NPU backends.\n\n Dynamically dispatches to the appropriate device-specific implementation:\n - On CUDA: `flash_attn.bert_padding.pad_input`\n - On NPU: `transformers.integrations.npu_flash_attention.pad_input`\n (falls back to `transformers.modeling_flash_attention_utils._pad_input`\n in newer versions of transformers).\n\n Users can call this function directly without worrying about the underlying device.\n \"\"\"\n _, func, *_ = _get_attention_functions()\n return func(*args, **kwargs)\n\n\ndef rearrange(*args, **kwargs):\n \"\"\"\n Unified entry point for `rearrange` across CUDA and NPU backends.\n\n Dynamically dispatches to the appropriate device-specific implementation:\n - On CUDA: `flash_attn.bert_padding.rearrange`\n - On NPU: `transformers.integrations.npu_flash_attention.rearrange`\n (falls back to `einops.rearrange` if no dedicated NPU implementation exists).\n\n Users can call this function directly without worrying about the underlying device.\n \"\"\"\n *_, func, _ = _get_attention_functions()\n return func(*args, **kwargs)\n\n\ndef unpad_input(*args, **kwargs):\n \"\"\"\n Unified entry point for `unpad_input` across CUDA and NPU backends.\n\n Dynamically dispatches to the appropriate device-specific implementation:\n - On CUDA: `flash_attn.bert_padding.unpad_input`\n - On NPU: `transformers.integrations.npu_flash_attention.unpad_input`\n (falls back to `transformers.modeling_flash_attention_utils._unpad_input`\n in newer versions of transformers).\n\n Users can call this function directly without worrying about the underlying device.\n \"\"\"\n *_, func = _get_attention_functions()\n return func(*args, **kwargs)\n\n\n__all__ = [\"index_first_axis\", \"pad_input\", \"rearrange\", \"unpad_input\"]\n"}64{"file_name": "verl__utils__chat_template.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\nimport logging\nimport os\n\nlogger = logging.getLogger(__name__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef initialize_system_prompt(tokenizer, **apply_chat_template_kwargs) -> list[int]:\n \"\"\"\n Initialize system prompt tokens for chat templates that support them.\n\n Args:\n tokenizer: The tokenizer with a chat template\n **apply_chat_template_kwargs: Additional arguments for apply_chat_template\n\n Returns:\n List of token IDs for the system prompt, or empty list if not supported\n \"\"\"\n token1 = tokenizer.apply_chat_template(\n [{\"role\": \"user\", \"content\": \"\"}], add_generation_prompt=False, tokenize=True\n )\n token2 = tokenizer.apply_chat_template(\n [{\"role\": \"user\", \"content\": \"\"}] * 2, add_generation_prompt=False, tokenize=True\n )\n # get system prompt tokens\n system_prompt = token1[: -(len(token2) - len(token1))]\n return system_prompt\n\n\ndef extract_system_prompt_and_generation(tokenizer):\n token1 = tokenizer.apply_chat_template(\n [{\"role\": \"user\", \"content\": \"\"}], add_generation_prompt=False, tokenize=True\n )\n token2 = tokenizer.apply_chat_template(\n [{\"role\": \"user\", \"content\": \"\"}] * 2, add_generation_prompt=False, tokenize=True\n )\n # get system prompt tokens\n system_prompt = token1[: -(len(token2) - len(token1))]\n # get generate prompt tokens\n token3 = tokenizer.apply_chat_template([{\"role\": \"user\", \"content\": \"\"}], add_generation_prompt=True, tokenize=True)\n generate_prompt = token3[len(token1) :]\n\n return system_prompt, generate_prompt\n"}65{"file_name": "verl__utils__checkpoint__checkpoint_handler.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\n# TODO: add unit tests\n\nimport logging\nimport os\nimport re\nfrom enum import Enum\n\nimport torch\n\nimport verl.utils.hdfs_io as hdfs_io\nfrom verl.single_controller import WorkerGroup\nfrom verl.utils.checkpoint.checkpoint_manager import find_latest_ckpt_path, get_checkpoint_tracker_filename\nfrom verl.utils.logger import log_with_rank\nfrom verl.workers.engine import BaseEngine\n\n\ndef extract_step(path):\n match = re.search(r\"global_step_(\\d+)\", path)\n if match:\n return int(match.group(1))\n return None\n\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_SFT_LOGGING_LEVEL\", \"WARN\"))\n\n\nclass OrchestrationMode(Enum):\n SPMD = 0\n RAY = 1\n\n\nclass CheckpointHandler:\n \"\"\"\n Checkpoint handler handles the path, global_step of a checkpoint folder.\n Currently, it only works with a single model.\n We can expand it to support multiple models. It is expected to be used with SPMD style (e.g., torchrun)\n \"\"\"\n\n def __init__(\n self,\n engine: BaseEngine | WorkerGroup,\n train_dataloader,\n *,\n default_local_dir,\n max_ckpt_to_keep=None,\n default_hdfs_dir=None,\n resume_mode=\"auto\",\n resume_from_path=None,\n mode=OrchestrationMode.SPMD,\n ):\n self.default_local_dir = default_local_dir\n self.max_ckpt_to_keep = max_ckpt_to_keep\n self.default_hdfs_dir = default_hdfs_dir\n self.resume_mode = resume_mode\n self.resume_from_path = resume_from_path\n self.engine = engine\n self.train_dataloader = train_dataloader\n self.mode = mode\n\n if self.mode == OrchestrationMode.SPMD:\n self.rank = torch.distributed.get_rank()\n self.is_mp_src_rank_with_outputs = self.engine.is_mp_src_rank_with_outputs()\n self.dp_rank = self.engine.get_data_parallel_rank()\n elif self.mode == OrchestrationMode.RAY:\n self.rank = 0\n self.is_mp_src_rank_with_outputs = True\n self.dp_rank = 0\n else:\n raise ValueError(f\"Unknown {self.mode=}\")\n\n def save_checkpoint(self, step):\n \"\"\"Save checkpoint using FSDPCheckpointManager with improved tracking\"\"\"\n from verl.utils.fs import local_mkdir_safe\n\n # Determine checkpoint path\n local_global_step_folder = os.path.join(self.default_local_dir, f\"global_step_{step}\")\n if self.rank == 0:\n print(f\"Saving checkpoint to: {local_global_step_folder}\")\n\n # Get max checkpoints to keep\n max_ckpt_to_keep = self.max_ckpt_to_keep\n\n # Use checkpoint manager to save\n self.engine.save_checkpoint(\n local_path=local_global_step_folder, global_step=step, max_ckpt_to_keep=max_ckpt_to_keep\n )\n\n # Save dataloader state. Note that we only save the iterator in the train_dataloader.\n # So it's identical in each dp rank.\n if self.is_mp_src_rank_with_outputs:\n dp_rank = self.dp_rank\n local_mkdir_safe(local_global_step_folder)\n dataloader_local_path = os.path.join(local_global_step_folder, f\"data_{dp_rank}.pt\")\n\n # Use StatefulDataLoader's built-in state dict functionality\n dataloader_state_dict = self.train_dataloader.state_dict()\n torch.save(dataloader_state_dict, dataloader_local_path)\n print(f\"Saved dataloader state to: {dataloader_local_path}\")\n\n if self.rank == 0:\n # Update latest checkpoint tracker (atomic write)\n tracker_file = get_checkpoint_tracker_filename(self.default_local_dir)\n temp_tracker_file = tracker_file + \".tmp\"\n with open(temp_tracker_file, \"w\") as f:\n f.write(str(step))\n os.rename(temp_tracker_file, tracker_file)\n print(f\"Updated checkpoint tracker: {tracker_file}\")\n\n # Copy to HDFS if configured\n if self.rank == 0 and self.default_hdfs_dir:\n hdfs_io.makedirs(self.default_hdfs_dir, exist_ok=True)\n hdfs_io.copy(src=local_global_step_folder, dst=self.default_hdfs_dir, dirs_exist_ok=True)\n\n if self.mode == OrchestrationMode.SPMD:\n torch.distributed.barrier()\n\n def load_checkpoint(self):\n # Determine resume path based on configuration\n checkpoint_path = self._determine_resume_path()\n\n if checkpoint_path is None:\n return 0\n\n # extract resume step from checkpoint path\n resume_step = extract_step(checkpoint_path)\n if resume_step is None:\n log_with_rank(\n f\"Warning: Could not extract step number from {checkpoint_path}, starting from step 0\",\n logger=logger,\n rank=self.rank,\n level=logging.WARNING,\n log_only_rank_0=True,\n )\n return 0\n self.resume_global_step = resume_step\n\n # Use checkpoint manager to load model state\n self.engine.load_checkpoint(checkpoint_path)\n # Always load dataloader state for StatefulDataLoader\n self._load_dataloader_state(checkpoint_path)\n\n return resume_step\n\n def _load_dataloader_state(self, checkpoint_path: str):\n \"\"\"Load dataloader state from checkpoint\"\"\"\n dp_rank = self.dp_rank\n dataloader_path = os.path.join(checkpoint_path, f\"data_{dp_rank}.pt\")\n\n if os.path.exists(dataloader_path):\n # Use StatefulDataLoader's built-in state dict functionality\n dataloader_state_dict = torch.load(dataloader_path, map_location=\"cpu\", weights_only=False)\n self.train_dataloader.load_state_dict(dataloader_state_dict)\n\n log_with_rank(\n f\"Successfully loaded dataloader state from {dataloader_path}\",\n logger=logger,\n rank=self.rank,\n log_only_rank_0=True,\n )\n\n else:\n log_with_rank(\n f\"Warning: No dataloader state found at {dataloader_path}, will start from scratch\",\n logger=logger,\n rank=self.rank,\n level=logging.WARNING,\n log_only_rank_0=True,\n )\n\n def _determine_resume_path(self):\n \"\"\"Determine the path to resume from based on resume_mode configuration\"\"\"\n resume_mode = self.resume_mode\n resume_from_path = self.resume_from_path\n\n if resume_mode == \"disable\":\n return None\n elif resume_mode == \"auto\":\n if resume_from_path is not None:\n assert os.path.exists(resume_from_path), (\n \"resume_from_path must be null or an existing path when resume_mode is 'auto'\"\n )\n assert \"global_step_\" in resume_from_path, \"resume_from_path must specify the global_steps\"\n return resume_from_path\n # Try to find the latest checkpoint in the default directory\n return self._find_latest_checkpoint()\n elif resume_mode == \"resume_path\":\n assert os.path.exists(resume_from_path), (\n \"resume_from_path must be an existing path when resume_mode is 'resume_path'\"\n )\n assert \"global_step_\" in resume_from_path, \"resume_from_path must specify the global_steps\"\n return resume_from_path\n else:\n raise ValueError(f\"Invalid resume_mode: {resume_mode}. Must be 'auto', 'disable', or 'resume_path'\")\n\n def _find_latest_checkpoint(self):\n \"\"\"Find the latest checkpoint in the default local directory\"\"\"\n checkpoint_dir = self.default_local_dir\n\n if not os.path.exists(checkpoint_dir):\n return None\n\n latest_checkpoint = find_latest_ckpt_path(checkpoint_dir)\n\n if latest_checkpoint and self.rank == 0:\n step_num = extract_step(latest_checkpoint)\n print(f\"Found latest checkpoint: {latest_checkpoint} (step {step_num})\")\n\n return latest_checkpoint\n"}66{"file_name": "verl__utils__checkpoint__checkpoint_manager.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport os\nimport random\nimport shutil\n\nimport numpy as np\nimport torch\nimport torch.distributed\nfrom omegaconf import DictConfig\nfrom transformers import PreTrainedTokenizer, ProcessorMixin\n\nfrom verl.trainer.config import CheckpointConfig\nfrom verl.utils.device import get_device_name, get_torch_device\n\n\nclass BaseCheckpointManager:\n \"\"\"\n A checkpoint manager that saves and loads the following states in a SPMD way:\n - model\n - optimizer\n - lr_scheduler\n - extra_states\n\n We save\n - sharded model states and optimizer states\n - full lr_scheduler states\n - huggingface tokenizer and config for ckpt merge\n \"\"\"\n\n def __init__(\n self,\n model,\n optimizer: torch.optim.Optimizer,\n lr_scheduler: torch.optim.lr_scheduler.LRScheduler = None,\n processing_class: PreTrainedTokenizer | ProcessorMixin = None,\n checkpoint_config: DictConfig | CheckpointConfig = None,\n ):\n self.checkpoint_config = checkpoint_config\n checkpoint_load_contents = checkpoint_config.get(\"load_contents\", None) if checkpoint_config else None\n checkpoint_save_contents = checkpoint_config.get(\"save_contents\", None) if checkpoint_config else None\n if checkpoint_load_contents is None:\n checkpoint_load_contents = [\"model\", \"optimizer\", \"extra\"]\n if checkpoint_save_contents is None:\n checkpoint_save_contents = [\"model\", \"optimizer\", \"extra\"]\n self.previous_global_step = None\n self.previous_saved_paths = []\n\n self.model = model\n self.optimizer = optimizer\n self.lr_scheduler = lr_scheduler\n self.processing_class = processing_class\n self.checkpoint_load_contents = checkpoint_load_contents\n self.checkpoint_save_contents = checkpoint_save_contents\n\n self.rank = torch.distributed.get_rank()\n self.world_size = torch.distributed.get_world_size()\n\n @property\n def should_save_model(self) -> bool:\n \"\"\"\n Returns True if 'model' is in checkpoint_save_contents, indicating the model state should be saved.\n \"\"\"\n return \"model\" in self.checkpoint_save_contents\n\n @property\n def should_save_optimizer(self) -> bool:\n \"\"\"\n Returns True if 'optimizer' is in checkpoint_save_contents, indicating the optimizer state should be saved.\n \"\"\"\n return \"optimizer\" in self.checkpoint_save_contents\n\n @property\n def should_save_extra(self) -> bool:\n \"\"\"\n Returns True if 'extra' is in checkpoint_save_contents, indicating the extra state should be saved.\n \"\"\"\n return \"extra\" in self.checkpoint_save_contents\n\n @property\n def should_save_hf_model(self) -> bool:\n \"\"\"\n Returns True if 'hf_model' is in checkpoint_save_contents, indicating the model should be converted to hf\n model and saved.\n \"\"\"\n return \"hf_model\" in self.checkpoint_save_contents\n\n @property\n def should_load_model(self) -> bool:\n \"\"\"\n Returns True if 'model' is in checkpoint_load_contents, indicating the model state should be loaded.\n \"\"\"\n return \"model\" in self.checkpoint_load_contents\n\n @property\n def should_load_optimizer(self) -> bool:\n \"\"\"\n Returns True if 'optimizer' is in checkpoint_load_contents, indicating the optimizer state should be loaded.\n \"\"\"\n return \"optimizer\" in self.checkpoint_load_contents\n\n @property\n def should_load_extra(self) -> bool:\n \"\"\"\n Returns True if 'extra' is in checkpoint_load_contents, indicating the extra state should be loaded.\n \"\"\"\n return \"extra\" in self.checkpoint_load_contents\n\n def load_checkpoint(self, local_path: str, hdfs_path: str = None, del_local_after_load: bool = False):\n raise NotImplementedError\n\n def save_checkpoint(\n self, local_path: str, hdfs_path: str = None, global_step: int = 0, max_ckpt_to_keep: int = None\n ):\n raise NotImplementedError\n\n @staticmethod\n def checkpath(local_path: str, hdfs_path: str):\n assert local_path is not None or hdfs_path is not None, \"local_path and hdfs_path cannot be both None\"\n return local_path is not None, local_path if local_path is not None else hdfs_path\n\n def remove_previous_save_local_path(self, path):\n if isinstance(path, str):\n path = [path]\n for p in path:\n abs_path = os.path.abspath(p)\n print(f\"Checkpoint manager remove previous save local path: {abs_path}\")\n if not os.path.exists(abs_path):\n continue\n shutil.rmtree(abs_path, ignore_errors=True)\n\n def ensure_checkpoint_capacity(self, max_ckpt_to_keep: int):\n \"\"\"\n Remove old checkpoints to make room for a new one, keeping a safety buffer.\n\n With max_ckpt_to_keep=1, this does nothing - we keep the existing checkpoint\n until the new save completes successfully (handled by register_checkpoint).\n For max_ckpt_to_keep >= 2, we keep (max_ckpt_to_keep - 1) checkpoints before save.\n \"\"\"\n if not (max_ckpt_to_keep and isinstance(max_ckpt_to_keep, int) and max_ckpt_to_keep > 1):\n return\n if len(self.previous_saved_paths) >= max_ckpt_to_keep:\n keep_start = len(self.previous_saved_paths) - max_ckpt_to_keep + 1\n self.remove_previous_save_local_path(self.previous_saved_paths[:keep_start])\n self.previous_saved_paths = self.previous_saved_paths[keep_start:]\n\n def register_checkpoint(self, new_path: str, max_ckpt_to_keep: int):\n \"\"\"\n Register a successfully saved checkpoint and enforce retention limit.\n\n Adds the new checkpoint path to tracking and removes excess old\n checkpoints beyond max_ckpt_to_keep.\n \"\"\"\n self.previous_saved_paths.append(new_path)\n if not (max_ckpt_to_keep and isinstance(max_ckpt_to_keep, int) and max_ckpt_to_keep > 0):\n return\n if len(self.previous_saved_paths) > max_ckpt_to_keep:\n keep_start = len(self.previous_saved_paths) - max_ckpt_to_keep\n self.remove_previous_save_local_path(self.previous_saved_paths[:keep_start])\n self.previous_saved_paths = self.previous_saved_paths[keep_start:]\n\n @staticmethod\n def get_rng_state():\n rng_state = {\n \"cpu\": torch.get_rng_state(),\n \"numpy\": np.random.get_state(),\n \"random\": random.getstate(),\n }\n\n if get_device_name() != \"cpu\":\n rng_state[get_device_name()] = get_torch_device().get_rng_state()\n\n return rng_state\n\n @staticmethod\n def load_rng_state(rng_state):\n torch.set_rng_state(rng_state[\"cpu\"])\n np.random.set_state(rng_state[\"numpy\"])\n random.setstate(rng_state[\"random\"])\n\n if get_device_name() != \"cpu\":\n get_torch_device().set_rng_state(rng_state[get_device_name()])\n\n\ndef find_latest_ckpt_path(path, directory_format=\"global_step_{}\"):\n \"\"\"\n Return the most recent checkpoint directory based on a tracker file.\n\n Args:\n path (str): Base directory containing the checkpoint tracker.\n directory_format (str): Template for checkpoint subfolders with one\n placeholder for the iteration number (default \"global_step_{}\").\n\n Returns:\n str or None: Full path to the latest checkpoint directory, or\n None if the tracker or checkpoint folder is missing.\n \"\"\"\n if path is None:\n return None\n\n tracker_file = get_checkpoint_tracker_filename(path)\n if not os.path.exists(tracker_file):\n if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0:\n print(f\"Checkpoint tracker file does not exist: {tracker_file}\")\n return None\n\n with open(tracker_file, \"rb\") as f:\n iteration = int(f.read().decode())\n ckpt_path = os.path.join(path, directory_format.format(iteration))\n if not os.path.exists(ckpt_path):\n print(\"Checkpoint does not exist: %s\", ckpt_path)\n return None\n\n print(\"Found checkpoint: %s\", ckpt_path)\n return ckpt_path\n\n\ndef get_checkpoint_tracker_filename(root_path: str):\n \"\"\"\n Tracker file rescords the latest chckpoint during training to restart from.\n \"\"\"\n return os.path.join(root_path, \"latest_checkpointed_iteration.txt\")\n\n\ndef should_save_ckpt_esi(max_steps_duration: float, save_ckpt_duration: float = 60, redundant_time: float = 0) -> bool:\n \"\"\"\n Determine if checkpoint should be saved based on capacity esi expiration.\n\n Args:\n max_steps_duration: Max estimated time (seconds) required to complete one training step\n save_ckpt_duration: Estimated time (seconds) required to save checkpoint (default: 60)\n redundant_time: Additional buffer time (seconds) for unexpected delays (default: 0)\n \"\"\"\n exp_ts_mlp = os.getenv(\"MLP_CURRENT_CAPACITY_BLOCK_EXPIRATION_TIMESTAMP\") # vemlp\n exp_ts_aws = os.getenv(\"SAGEMAKER_CURRENT_CAPACITY_BLOCK_EXPIRATION_TIMESTAMP\") # aws\n if exp_ts_mlp:\n try:\n import time\n\n remaining = float(exp_ts_mlp) - time.time()\n except ValueError:\n return False\n return (\n remaining > 0\n and max_steps_duration > 0\n and remaining <= save_ckpt_duration + max_steps_duration + redundant_time\n )\n elif exp_ts_aws:\n from datetime import datetime, timedelta\n\n expiration_time = datetime.fromtimestamp(int(exp_ts_aws))\n time_difference = expiration_time - datetime.now()\n threshold_minutes = (save_ckpt_duration + max_steps_duration + redundant_time) / 60\n return time_difference < timedelta(minutes=threshold_minutes)\n else:\n return False\n"}67{"file_name": "verl__utils__checkpoint__fsdp_checkpoint_manager.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport json\nimport logging\nimport os\nimport warnings\nfrom dataclasses import asdict, dataclass\nfrom typing import Optional\n\nimport torch\nimport torch.distributed\nfrom accelerate import init_empty_weights\nfrom omegaconf import DictConfig\nfrom torch.distributed.fsdp import FullyShardedDataParallel as FSDP\nfrom torch.distributed.fsdp import ShardedOptimStateDictConfig, ShardedStateDictConfig, StateDictType\nfrom transformers import GenerationConfig, PreTrainedTokenizer, ProcessorMixin\nfrom transformers.dynamic_module_utils import custom_object_save\n\nfrom verl.utils.device import is_cuda_available\nfrom verl.utils.fs import copy_to_local, is_non_local, local_mkdir_safe\nfrom verl.utils.fsdp_utils import fsdp_version, get_fsdp_full_state_dict, get_fsdp_state_ctx\nfrom verl.utils.logger import log_with_rank\n\nfrom .checkpoint_manager import BaseCheckpointManager\n\n# Setup logging\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"INFO\"))\n\n\n@dataclass\nclass FSDPConfig:\n \"\"\"Configuration for FSDP checkpointing.\n\n Args:\n FSDP_version (int): Version of FSDP being used.\n world_size (int): Number of processes in the distributed training setup.\n \"\"\"\n\n FSDP_version: int\n world_size: int\n\n\nclass FSDPCheckpointManager(BaseCheckpointManager):\n \"\"\"\n Manage FSDP checkpointing in SPMD training.\n\n - Saves/loads per-rank sharded model & optimizer states\n - Persists full lr_scheduler and RNG state\n - Stores HF tokenizer/processor and model/config for unified restore\n\n Args:\n model (FSDP): Wrapped model instance.\n optimizer (Optimizer): Training optimizer.\n lr_scheduler (LRScheduler): Learning-rate scheduler.\n processing_class (PreTrainedTokenizer or ProcessorMixin, optional):\n Pre-/post-processing artifact handler.\n checkpoint_contents DictConfig: Configuration for checkpoint contents.\n - 'load': Components to load; must contain 'model'. Defaults to ['model', 'optimizer', 'extra'].\n - 'save': Components to save; must contain 'model'. Defaults to ['model', 'optimizer', 'extra'].\n trust_remote_code: Whether to trust_remote_code when loading the model configuration\n \"\"\"\n\n def __init__(\n self,\n model: FSDP,\n optimizer: Optional[torch.optim.Optimizer] = None,\n lr_scheduler: Optional[torch.optim.lr_scheduler.LRScheduler] = None,\n processing_class: PreTrainedTokenizer | ProcessorMixin = None,\n checkpoint_config: DictConfig = None,\n trust_remote_code: bool = False,\n **kwargs,\n ):\n if processing_class is None and \"tokenizer\" in kwargs:\n warnings.warn(\n \"`tokenizer` is deprecated. use `processing_class` instead.\", DeprecationWarning, stacklevel=2\n )\n processing_class = kwargs.pop(\"tokenizer\")\n\n super().__init__(\n model,\n optimizer,\n lr_scheduler=lr_scheduler,\n processing_class=processing_class,\n checkpoint_config=checkpoint_config,\n )\n self.trust_remote_code = trust_remote_code\n\n def load_checkpoint(self, local_path: str, hdfs_path: str = None, del_local_after_load=False):\n \"\"\"\n Load an FSDP checkpoint for this rank.\n\n Downloads and loads:\n - model and optimizer shards\n - extra state dict (scheduler + RNG)\n\n Args:\n local_path: Directory with per-rank checkpoint files.\n hdfs_path: Unused (for API compatibility).\n del_local_after_load: Remove local files after loading.\n \"\"\"\n if local_path is None:\n return\n\n # check if the checkpoint_load_contents is valid\n if self.should_load_model:\n assert self.model is not None, \"model must be provided when checkpoint_contents.load includes ['model']\"\n if self.should_load_optimizer:\n assert self.optimizer is not None, (\n \"optimizer must be provided when checkpoint_contents.load includes ['optimizer']\"\n )\n\n # every rank download its own checkpoint\n state_dict_cfg = (\n ShardedStateDictConfig(offload_to_cpu=True if is_cuda_available else False)\n if self.should_load_model\n else None\n )\n optim_cfg = (\n ShardedOptimStateDictConfig(offload_to_cpu=True if is_cuda_available else False)\n if self.should_load_optimizer\n else None\n )\n with get_fsdp_state_ctx(self.model, StateDictType.SHARDED_STATE_DICT, state_dict_cfg, optim_cfg):\n if self.should_load_model:\n remote_model_path = os.path.join(local_path, f\"model_world_size_{self.world_size}_rank_{self.rank}.pt\")\n local_model_path = copy_to_local(remote_model_path)\n model_state_dict = torch.load(local_model_path, weights_only=False)\n self.model.load_state_dict(model_state_dict)\n log_with_rank(f\"Loaded model from {remote_model_path}\", rank=self.rank, logger=logger)\n\n if self.should_load_optimizer:\n remote_optim_path = os.path.join(local_path, f\"optim_world_size_{self.world_size}_rank_{self.rank}.pt\")\n local_optim_path = copy_to_local(remote_optim_path)\n optimizer_state_dict = torch.load(local_optim_path, weights_only=False)\n self.optimizer.load_state_dict(optimizer_state_dict)\n log_with_rank(f\"Loaded optimizer from {remote_optim_path}\", rank=self.rank, logger=logger)\n\n if self.should_load_extra:\n remote_extra_state_path = os.path.join(\n local_path, f\"extra_state_world_size_{self.world_size}_rank_{self.rank}.pt\"\n )\n local_extra_state_path = copy_to_local(remote_extra_state_path)\n extra_state_dict = torch.load(local_extra_state_path, weights_only=False)\n # recover random state\n if \"rng\" in extra_state_dict:\n # 'rng' may not exist for backward compatibility\n self.load_rng_state(extra_state_dict[\"rng\"])\n log_with_rank(f\"Loaded rng from {remote_extra_state_path}\", rank=self.rank, logger=logger)\n\n lr_scheduler_state_dict = extra_state_dict[\"lr_scheduler\"]\n if lr_scheduler_state_dict is not None and self.lr_scheduler is not None:\n self.lr_scheduler.load_state_dict(lr_scheduler_state_dict)\n log_with_rank(f\"Loaded lr_scheduler from {remote_extra_state_path}\", rank=self.rank, logger=logger)\n\n if self.rank == 0 and del_local_after_load:\n try:\n os.remove(local_model_path) if is_non_local(local_model_path) else None\n os.remove(local_optim_path) if is_non_local(local_optim_path) else None\n os.remove(local_extra_state_path) if is_non_local(local_extra_state_path) else None\n except Exception as e:\n log_with_rank(\n f\"remove local resume ckpt file after loading failed, exception {e} will be ignored\",\n rank=self.rank,\n logger=logger,\n )\n\n # wait for everyone to load checkpoints\n torch.distributed.barrier()\n\n def save_checkpoint(self, local_path: str, hdfs_path: str = None, global_step: int = 0, max_ckpt_to_keep=None):\n \"\"\"\n Save an FSDP checkpoint for this rank.\n\n Writes:\n - model & optimizer shard files\n - extra state dict (scheduler + RNG)\n - HF tokenizer/processor and model/config on rank 0\n - optional full HF model under 'huggingface/' if requested\n\n Rotates old checkpoints, keeping at most `max_ckpt_to_keep`.\n\n Args:\n local_path: Target directory for checkpoint files.\n hdfs_path: Unused (for API compatibility).\n global_step: Current training step (used for bookkeeping).\n max_ckpt_to_keep: Number of recent checkpoints to retain.\n \"\"\"\n if local_path is None:\n return\n\n # record the previous global step\n self.previous_global_step = global_step\n\n if self.rank == 0:\n self.ensure_checkpoint_capacity(max_ckpt_to_keep)\n\n local_path = local_mkdir_safe(local_path)\n torch.distributed.barrier()\n\n # check if the checkpoint_save_contents is valid\n if self.should_save_model:\n assert self.model is not None, \"model must be provided when checkpoint_contents.save includes ['model']\"\n if self.should_save_optimizer:\n assert self.optimizer is not None, (\n \"optimizer must be provided when checkpoint_contents.save includes ['optimizer']\"\n )\n\n # every rank will save its own model and optim shard\n state_dict_cfg = ShardedStateDictConfig(offload_to_cpu=True if is_cuda_available else False)\n optim_cfg = ShardedOptimStateDictConfig(offload_to_cpu=True if is_cuda_available else False)\n with warnings.catch_warnings():\n warnings.simplefilter(\"ignore\")\n with get_fsdp_state_ctx(self.model, StateDictType.SHARDED_STATE_DICT, state_dict_cfg, optim_cfg):\n model_path = os.path.join(local_path, f\"model_world_size_{self.world_size}_rank_{self.rank}.pt\")\n optim_path = os.path.join(local_path, f\"optim_world_size_{self.world_size}_rank_{self.rank}.pt\")\n extra_path = os.path.join(local_path, f\"extra_state_world_size_{self.world_size}_rank_{self.rank}.pt\")\n\n if self.should_save_model:\n model_state_dict = self.model.state_dict()\n torch.save(model_state_dict, model_path)\n log_with_rank(f\"Saved model to {os.path.abspath(model_path)}\", rank=self.rank, logger=logger)\n\n if self.should_save_optimizer:\n optimizer_state_dict = self.optimizer.state_dict()\n torch.save(optimizer_state_dict, optim_path)\n log_with_rank(f\"Saved optim to {os.path.abspath(optim_path)}\", rank=self.rank, logger=logger)\n\n if self.should_save_extra:\n lr_scheduler_state_dict = self.lr_scheduler.state_dict() if self.lr_scheduler is not None else None\n extra_state_dict = {\n \"lr_scheduler\": lr_scheduler_state_dict,\n \"rng\": self.get_rng_state(),\n }\n torch.save(extra_state_dict, extra_path)\n log_with_rank(f\"Saved extra_state to {os.path.abspath(extra_path)}\", rank=self.rank, logger=logger)\n\n if self.rank == 0:\n # Save HF tokenizer/processor and model config on rank 0 to huggingface/ directory, no matter whether\n # huggingface model is requested to be saved or not.\n\n if fsdp_version(self.model) == 1:\n unwrap_model = self.model._fsdp_wrapped_module\n else:\n unwrap_model = self.model\n\n hf_config_tokenizer_path = os.path.join(local_path, \"huggingface\")\n local_mkdir_safe(hf_config_tokenizer_path)\n model_config = unwrap_model.config\n generation_config = None\n if unwrap_model.can_generate() and hasattr(model_config, \"name_or_path\") and model_config.name_or_path:\n try:\n # Some model's name_or_path is empty if not initialized from pretrained,\n # in this cases, we don't save generation config.\n generation_config = GenerationConfig.from_pretrained(model_config.name_or_path)\n generation_config.save_pretrained(hf_config_tokenizer_path)\n except Exception:\n # if the generation config isn't available, we don't save it\n pass\n\n if hasattr(model_config, \"auto_map\") and None in model_config.auto_map:\n model_config.auto_map = {k: v for k, v in model_config.auto_map.items() if k is not None}\n\n model_config.save_pretrained(hf_config_tokenizer_path)\n if self.processing_class is not None:\n self.processing_class.save_pretrained(hf_config_tokenizer_path)\n log_with_rank(\n f\"Saved model config and tokenizer class to {os.path.abspath(hf_config_tokenizer_path)}\",\n rank=self.rank,\n logger=logger,\n log_only_rank_0=True,\n )\n\n # If we have a custom model, we copy the file defining it in the folder and set the attributes so it can be\n # loaded from the Hub.\n if hasattr(model_config, \"auto_map\"):\n custom_object_save(unwrap_model, hf_config_tokenizer_path, config=model_config)\n\n # Also save runtime FSDP config\n fsdp_config_path = os.path.join(local_path, \"fsdp_config.json\")\n fsdp_config = FSDPConfig(\n FSDP_version=fsdp_version(self.model),\n world_size=self.world_size,\n )\n with open(fsdp_config_path, \"w\") as f:\n json.dump(asdict(fsdp_config), f, indent=4)\n\n # wait for everyone to dump to local\n torch.distributed.barrier()\n\n if self.should_save_hf_model:\n # Only rank 0 will save hf model and,\n # offload to cpu to save LLMs which may be too large to fit in one GPU\n state_dict = get_fsdp_full_state_dict(self.model, offload_to_cpu=True, rank0_only=True)\n\n if self.rank == 0:\n hf_local_path = os.path.join(local_path, \"huggingface\")\n os.makedirs(hf_local_path, exist_ok=True)\n\n if \"ForTokenClassification\" in model_config.architectures[0]:\n from transformers import AutoModelForTokenClassification\n\n auto_model_cls = AutoModelForTokenClassification\n elif \"ForCausalLM\" in model_config.architectures[0]:\n from transformers import AutoModelForCausalLM\n\n auto_model_cls = AutoModelForCausalLM\n elif \"ForConditionalGeneration\" in model_config.architectures[0]:\n # Handle different transformers versions for Vision2Seq models\n import transformers\n from packaging import version\n\n if version.parse(transformers.__version__) >= version.parse(\"4.54.0\"):\n # transformers >= 4.54.0 uses AutoModelForImageTextToText\n from transformers import AutoModelForImageTextToText\n\n auto_model_cls = AutoModelForImageTextToText\n else:\n # transformers < 4.54.0 uses AutoModelForVision2Seq\n from transformers import AutoModelForVision2Seq\n\n auto_model_cls = AutoModelForVision2Seq\n else:\n raise NotImplementedError(f\"Unknown architecture {model_config['architectures']}\")\n\n with init_empty_weights():\n save_model = auto_model_cls.from_config(\n model_config, torch_dtype=torch.bfloat16, trust_remote_code=self.trust_remote_code\n )\n\n save_model.to_empty(device=\"cpu\")\n\n if save_model.can_generate():\n if generation_config is not None:\n save_model.generation_config = generation_config\n else:\n print(\n f\"Warning: {self.__class__.__name__}.save_checkpoint: Generation config file not found \"\n f\"in, using a generation config created from the model config when saving hf_model.\"\n )\n\n save_model.save_pretrained(hf_local_path, state_dict=state_dict)\n log_with_rank(\n f\"Saved hf_model to {os.path.abspath(hf_local_path)}\",\n rank=self.rank,\n logger=logger,\n log_only_rank_0=True,\n )\n del state_dict\n del save_model\n\n # wait for rank0 to dump hf_model to local\n torch.distributed.barrier()\n\n if self.rank == 0:\n self.register_checkpoint(local_path, max_ckpt_to_keep)\n"}68{"file_name": "verl__utils__checkpoint__megatron_checkpoint_manager.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport inspect\nimport json\nimport logging\nimport os\nimport random\nfrom collections.abc import Callable\nfrom dataclasses import asdict\n\nimport megatron.core\nimport numpy as np\nimport torch\nimport torch.distributed\nfrom megatron.core import dist_checkpointing, mpu, tensor_parallel\nfrom megatron.core.dist_checkpointing.mapping import ShardedObject\nfrom megatron.core.transformer.enums import AttnBackend\nfrom packaging import version\nfrom transformers import GenerationConfig\n\nfrom verl.models.weight_loader_registry import get_weight_saver\nfrom verl.utils.device import get_device_name, get_torch_device\nfrom verl.utils.fs import is_non_local, local_mkdir_safe\nfrom verl.utils.logger import log_with_rank\nfrom verl.utils.megatron.dist_checkpointing import load_dist_checkpointing, save_dist_checkpointing\nfrom verl.utils.megatron_utils import (\n get_dist_checkpoint_path,\n get_hf_model_checkpoint_path,\n get_transformer_config_checkpoint_path,\n)\n\nfrom .checkpoint_manager import BaseCheckpointManager\n\n# Setup logging\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"INFO\"))\nmcore_ge_014 = version.parse(megatron.core.__version__) >= version.parse(\"0.14.0\")\nif not mcore_ge_014:\n logger.warning(\n \"Detected megatron.core %s, recommend upgrading to >= 0.14.0 for better checkpoint compatibility\",\n megatron.core.__version__,\n )\n\n\nclass MegatronCheckpointManager(BaseCheckpointManager):\n \"\"\"\n Checkpoint manager for Megatron-LM distributed training.\n\n This class manages the saving and loading of model checkpoints in a Megatron-LM\n distributed training environment. It handles various aspects of checkpointing\n including model states, optimizer states, learning rate schedulers, and random\n number generator states, ensuring compatibility with HuggingFace formats.\n\n Key features:\n - Distributed checkpoint saving and loading using Megatron's dist_checkpointing\n - Support for tensor parallel, pipeline parallel, and data parallel configurations\n - Automatic handling of model state dictionaries across multiple pipeline stages\n - Integration with HuggingFace model configurations and tokenizers\n - Random number generator state management for reproducibility\n - Support for both synchronous and asynchronous checkpoint operations\n\n The manager automatically handles:\n - Directory structure creation based on global steps and process ranks\n - Model configuration and tokenizer saving in HuggingFace format\n - Optimizer and scheduler state persistence\n - CUDA RNG state management for deterministic training\n - Checkpoint cleanup and retention policies\n\n Args:\n model: The Megatron model instance to checkpoint\n optimizer: The optimizer instance (optional)\n lr_scheduler: The learning rate scheduler instance (optional)\n\n Attributes:\n model: Reference to the Megatron model being checkpointed\n optimizer: Reference to the optimizer (if provided)\n lr_scheduler: Reference to the learning rate scheduler (if provided)\n rank: Current process rank in the distributed setup\n\n Example:\n ```python\n checkpoint_manager = MegatronCheckpointManager(\n model=megatron_model,\n optimizer=optimizer,\n lr_scheduler=scheduler\n )\n\n checkpoint_manager.save_checkpoint(\n local_path=\"checkpoints/step_1000\",\n global_step=1000\n )\n\n checkpoint_manager.load_checkpoint(\n local_path=\"checkpoints/step_1000\"\n )\n ```\n \"\"\"\n\n def __init__(\n self,\n config,\n checkpoint_config,\n model_config,\n transformer_config,\n role,\n model: torch.nn.ModuleList,\n arch: str,\n hf_config,\n param_dtype: torch.dtype,\n share_embeddings_and_output_weights: bool,\n processing_class,\n optimizer,\n optimizer_scheduler,\n use_distributed_optimizer: bool,\n use_checkpoint_opt_param_scheduler: bool = False,\n use_dist_checkpointing: bool = True,\n bridge=None,\n provider=None,\n peft_cls=None,\n **kwargs,\n ):\n super().__init__(\n model,\n optimizer=optimizer,\n lr_scheduler=optimizer_scheduler,\n processing_class=processing_class,\n checkpoint_config=checkpoint_config,\n )\n self.arch = arch\n self.config = config\n self.transformer_config = transformer_config\n self.role = role\n self.is_value_model = False\n if self.role in [\"reward\", \"critic\"]:\n self.is_value_model = True\n self.model_config = model_config\n self.hf_config = hf_config\n self.param_dtype = param_dtype\n self.share_embeddings_and_output_weights = share_embeddings_and_output_weights\n self.model_path = self.config.model.path\n self.use_distributed_optimizer = use_distributed_optimizer\n self.use_checkpoint_opt_param_scheduler = use_checkpoint_opt_param_scheduler\n self.bridge = bridge\n self.provider = provider\n self.vanilla_bridge = self.provider is None\n self.peft_cls = peft_cls\n self.rank = torch.distributed.get_rank()\n # Megatron-Bridge is Okay to load/save HF checkpoint for value model as well\n self.use_dist_checkpointing = (\n use_dist_checkpointing or not self.bridge or (self.vanilla_bridge and self.is_value_model)\n )\n self.use_hf_checkpoint = not self.use_dist_checkpointing\n\n self.weight_saver = None\n if self.bridge is None:\n self.weight_saver = get_weight_saver(self.arch)\n\n def get_rng_state(self, use_dist_ckpt: bool = True, data_parallel_random_init: bool = False):\n \"\"\"collect rng state across data parallel ranks\"\"\"\n rng_state = {\n \"random_rng_state\": random.getstate(),\n \"np_rng_state\": np.random.get_state(),\n \"torch_rng_state\": torch.get_rng_state(),\n \"rng_tracker_states\": tensor_parallel.get_cuda_rng_tracker().get_states(),\n }\n\n if get_device_name() != \"cpu\":\n rng_state[f\"{get_device_name()}_rng_state\"] = get_torch_device().get_rng_state()\n\n rng_state_list = None\n if torch.distributed.is_initialized() and mpu.get_data_parallel_world_size() > 1 and data_parallel_random_init:\n rng_state_list = [None for i in range(mpu.get_data_parallel_world_size())]\n torch.distributed.all_gather_object(rng_state_list, rng_state, group=mpu.get_data_parallel_group())\n else:\n rng_state_list = [rng_state]\n\n if use_dist_ckpt:\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n pp_size = mpu.get_pipeline_model_parallel_world_size()\n tp_rank = mpu.get_tensor_model_parallel_rank()\n tp_size = mpu.get_tensor_model_parallel_world_size()\n rng_state_list = ShardedObject(\n \"rng_state\",\n rng_state_list,\n (pp_size, tp_size),\n (pp_rank, tp_rank),\n replica_id=mpu.get_data_parallel_rank(with_context_parallel=True),\n )\n\n return rng_state_list\n\n def get_checkpoint_name(\n self,\n checkpoints_path,\n pipeline_parallel=None,\n tensor_rank=None,\n pipeline_rank=None,\n cp_rank=None,\n expert_parallel=None,\n expert_rank=None,\n return_base_dir=True,\n basename=\"model.pt\",\n ):\n \"\"\"Determine the directory name for this rank's checkpoint.\"\"\"\n # Use both the tensor and pipeline MP rank.\n if pipeline_parallel is None:\n pipeline_parallel = mpu.get_pipeline_model_parallel_world_size() > 1\n if tensor_rank is None:\n tensor_rank = mpu.get_tensor_model_parallel_rank()\n if pipeline_rank is None:\n pipeline_rank = mpu.get_pipeline_model_parallel_rank()\n if cp_rank is None:\n cp_rank = mpu.get_context_parallel_rank()\n if expert_parallel is None:\n expert_parallel = mpu.get_expert_model_parallel_world_size() > 1\n if expert_rank is None:\n expert_rank = mpu.get_expert_model_parallel_rank()\n\n # Use both the tensor and pipeline MP rank. If using the distributed\n # optimizer, then the optimizer's path must additionally include the\n # data parallel rank.\n\n # due to the fact that models are identical across cp ranks, cp rank is not used in the checkpoint path\n if not pipeline_parallel:\n common_path = os.path.join(checkpoints_path, f\"mp_rank_{tensor_rank:02d}\")\n else:\n common_path = os.path.join(checkpoints_path, f\"mp_rank_{tensor_rank:02d}_{pipeline_rank:03d}\")\n\n if expert_parallel:\n common_path = common_path + f\"_{expert_rank:03d}\"\n\n os.makedirs(common_path, exist_ok=True)\n\n if return_base_dir:\n return common_path\n return os.path.join(common_path, basename)\n\n def generate_state_dict(\n self,\n generate_model: bool = True,\n generate_optimizer: bool = True,\n generate_extra: bool = True,\n is_loading: bool = False,\n metadata: dict | None = None,\n ):\n # For save dist checkpointing\n state_dict = {}\n base_metadata = metadata or self._build_sharded_state_dict_metadata()\n\n # Should always generate model state dict\n # All ranks Save Model to reduce memory pressure\n # Get sharded state dict, notice that state_dict will collect among dp groups, causing memory pressure\n for vpp_rank, model in enumerate(self.model):\n if len(self.model) > 1:\n mpu.set_virtual_pipeline_model_parallel_rank(vpp_rank)\n key = f\"model{vpp_rank}\" if len(self.model) > 1 else \"model\"\n else:\n key = \"model\"\n if hasattr(model, \"module\"):\n model = model.module\n\n # GPTModel's sharded_state_dict function when having mtp requires metadata['dp_cp_group']\n model_metadata = dict(base_metadata)\n model_metadata[\"dp_cp_group\"] = mpu.get_data_parallel_group(with_context_parallel=True)\n kwargs = {\"metadata\": model_metadata}\n state_dict[key] = model.sharded_state_dict(**kwargs)\n\n # Optimizer State Dict\n if generate_optimizer:\n torch.distributed.barrier()\n sharded_state_dict_kwargs = {\"is_loading\": is_loading}\n if base_metadata is not None:\n # https://github.com/NVIDIA/Megatron-LM/blob/core_v0.14.0/megatron/core/optimizer/distrib_optimizer.py#L1109-L1123\n if mcore_ge_014:\n sharded_state_dict_kwargs[\"metadata\"] = base_metadata\n optimizer_sharded_states = self.optimizer.sharded_state_dict(state_dict, **sharded_state_dict_kwargs)\n state_dict[\"optimizer\"] = optimizer_sharded_states\n\n if self.lr_scheduler is not None:\n lr_state_dict = self.lr_scheduler.state_dict()\n state_dict[\"lr_scheduler\"] = lr_state_dict\n\n if not generate_model:\n state_dict.pop(\"model\", None)\n\n # RNG States State Dict\n if generate_extra:\n torch.distributed.barrier()\n rng_state = self.get_rng_state()\n state_dict[\"rng_state\"] = rng_state\n\n return state_dict\n\n def _build_sharded_state_dict_metadata(self) -> dict:\n \"\"\"Builds metadata used for sharded_state_dict versioning.\n\n\n The whole content metadata is passed to ``sharded_state_dict`` model and optimizer methods\n and therefore affects only the logic behind sharded_state_dict creation.\n The content metadata should be minimalistic, ideally flat (or with a single nesting level)\n and with semantically meaningful flag names (e.g. `distrib_optim_sharding_type`).\n In particular, a simple integer (or SemVer) versioning flag (e.g. `metadata['version'] = 3.4`)\n is discouraged, because the metadata serves for all models and optimizers and it's practically\n impossible to enforce a linearly increasing versioning for this whole space.\n \"\"\"\n metadata: dict = {}\n\n if not mcore_ge_014:\n # For backward compatibility with Megatron core < v0.14.0\n if self.use_distributed_optimizer:\n metadata[\"distrib_optim_sharding_type\"] = \"fully_sharded_model_space\"\n return metadata\n\n if self.use_distributed_optimizer:\n megatron_config = getattr(self.config, self.role, self.config).megatron\n dist_ckpt_optim_fully_reshardable = megatron_config.dist_ckpt_optim_fully_reshardable\n distrib_optim_fully_reshardable_mem_efficient = (\n megatron_config.distrib_optim_fully_reshardable_mem_efficient\n )\n if dist_ckpt_optim_fully_reshardable:\n metadata[\"distrib_optim_sharding_type\"] = \"fully_reshardable\"\n metadata[\"distrib_optim_fully_reshardable_mem_efficient\"] = (\n distrib_optim_fully_reshardable_mem_efficient\n )\n else:\n metadata[\"distrib_optim_sharding_type\"] = \"dp_reshardable\"\n\n metadata[\"singleton_local_shards\"] = False\n metadata[\"chained_optim_avoid_prefix\"] = True\n return metadata\n\n def load_rng_states(self, rng_states, data_parallel_random_init=False, use_dist_ckpt=True):\n # access rng_state for data parallel rank\n if data_parallel_random_init:\n rng_states = rng_states[mpu.get_data_parallel_rank()]\n else:\n rng_states = rng_states[0]\n random.setstate(rng_states[\"random_rng_state\"])\n np.random.set_state(rng_states[\"np_rng_state\"])\n torch.set_rng_state(rng_states[\"torch_rng_state\"])\n\n if get_device_name() != \"cpu\":\n get_torch_device().set_rng_state(rng_states[f\"{get_device_name()}_rng_state\"])\n\n # Check for empty states array\n if not rng_states[\"rng_tracker_states\"]:\n raise KeyError\n tensor_parallel.get_cuda_rng_tracker().set_states(rng_states[\"rng_tracker_states\"])\n\n def load_checkpoint(self, local_path: str, hdfs_path: str = None, del_local_after_load=False):\n if local_path is not None:\n assert os.path.exists(local_path), f\"Checkpoint path {local_path} does not exist.\"\n\n # For load optimizer dist_ckpt\n try:\n import transformer_engine\n\n torch.serialization.add_safe_globals([torch.optim.AdamW])\n torch.serialization.add_safe_globals([transformer_engine.pytorch.optimizers.fused_adam.FusedAdam])\n except Exception:\n pass\n\n dist_checkpoint_path = get_dist_checkpoint_path(local_path)\n\n load_content_metadata = getattr(dist_checkpointing, \"load_content_metadata\", None)\n if load_content_metadata is None:\n # For backward compatibility\n sharded_sd_metadata = None\n else:\n sharded_sd_metadata = load_content_metadata(checkpoint_dir=dist_checkpoint_path)\n if sharded_sd_metadata is None:\n if self.use_distributed_optimizer:\n # Backward-compatibility with old checkpoints which don't have content versioning\n # Can be removed after ending support for MLM optimizer checkpoints with MCore < v0.13\n # (for MCore v0.13+ checkpoints `sharded_sd_metadata is not None`)\n sharded_sd_metadata = {\n \"distrib_optim_sharding_type\": \"fully_sharded_model_space\",\n }\n else:\n sharded_sd_metadata = self._build_sharded_state_dict_metadata()\n\n # Get State Dict for loading\n sharded_state_dict = self.generate_state_dict(\n self.should_load_model and self.use_dist_checkpointing,\n self.should_load_optimizer,\n self.should_load_extra,\n is_loading=True,\n metadata=sharded_sd_metadata,\n )\n log_with_rank(f\"Generated state dict for loading: {sharded_state_dict.keys()}\", rank=self.rank, logger=logger)\n\n # Load Dist Checkpointing\n state_dict = load_dist_checkpointing(\n sharded_state_dict=sharded_state_dict,\n ckpt_dir=dist_checkpoint_path,\n )\n\n if self.should_load_model and self.use_dist_checkpointing:\n assert \"model\" in state_dict or any(\n f\"model{vpp_rank}\" in state_dict for vpp_rank in range(len(self.model))\n ), f\"Model state dict not found in {state_dict.keys()}. Please check the checkpoint file {local_path}.\"\n for vpp_rank, model in enumerate(self.model):\n if len(self.model) == 1:\n model_state_dict = state_dict[\"model\"]\n else:\n assert f\"model{vpp_rank}\" in state_dict, f\"model{vpp_rank} not found in state_dict\"\n model_state_dict = state_dict[f\"model{vpp_rank}\"]\n mpu.set_virtual_pipeline_model_parallel_rank(vpp_rank)\n self.model[vpp_rank].load_state_dict(model_state_dict)\n log_with_rank(f\"Loaded sharded model checkpoint from {local_path}\", rank=self.rank, logger=logger)\n\n # Skip HF checkpoint loading if PEFT is used\n elif self.should_load_model and self.use_hf_checkpoint and self.peft_cls is None:\n hf_model_path = get_hf_model_checkpoint_path(local_path)\n if self.vanilla_bridge:\n self.bridge.load_weights(self.model, hf_model_path)\n else:\n self.bridge.load_hf_weights(self.model, hf_model_path)\n log_with_rank(f\"Loaded HF model checkpoint from {hf_model_path} with bridge\", rank=self.rank, logger=logger)\n # Load PEFT adapter checkpoint if available\n if self.should_load_model and self.peft_cls is not None:\n adapter_ckpt_path = os.path.join(local_path, \"adapter_checkpoint\")\n if os.path.exists(adapter_ckpt_path):\n from verl.utils.megatron_peft_utils import load_adapter_checkpoint\n\n # TODO: a better format for adapter checkpoint, waiting megatron-bridge support\n\n load_adapter_checkpoint(\n self.model,\n adapter_ckpt_path,\n )\n log_with_rank(\n f\"Loaded adapter checkpoint from {adapter_ckpt_path}\",\n rank=self.rank,\n logger=logger,\n )\n else:\n log_with_rank(\n f\"PEFT config is set but no adapter checkpoint found at {adapter_ckpt_path}\",\n rank=self.rank,\n logger=logger,\n )\n\n if self.should_load_optimizer:\n assert \"optimizer\" in state_dict, (\n f\"Optimizer state dict not found in {state_dict.keys()}. Please check the checkpoint file {local_path}.\"\n )\n optimizer_state_dict = state_dict[\"optimizer\"]\n self.optimizer.load_state_dict(optimizer_state_dict)\n log_with_rank(f\"Loaded optimizer checkpoint from {local_path}\", rank=self.rank, logger=logger)\n if self.use_checkpoint_opt_param_scheduler:\n assert \"lr_scheduler\" in state_dict, (\n f\"LR scheduler state dict not found in {state_dict.keys()}. Please check the checkpoint file \"\n f\"{local_path}.\"\n )\n lr_scheduler_state_dict = state_dict[\"lr_scheduler\"]\n if self.lr_scheduler is not None:\n self.lr_scheduler.load_state_dict(lr_scheduler_state_dict)\n log_with_rank(f\"Loaded LR scheduler checkpoint from {local_path}\", rank=self.rank, logger=logger)\n\n if self.should_load_extra:\n assert \"rng_state\" in state_dict, (\n f\"RNG state dict not found in {state_dict.keys()}. Please check the checkpoint file {local_path}.\"\n )\n rng_state = state_dict[\"rng_state\"]\n self.load_rng_states(rng_state)\n log_with_rank(f\"Loaded RNG states from {local_path}\", rank=self.rank, logger=logger)\n\n if del_local_after_load:\n try:\n os.remove(local_path) if is_non_local(local_path) else None\n except Exception as e:\n log_with_rank(\n f\"remove local resume ckpt file after loading failed, exception {e} will be ignored\",\n rank=self.rank,\n logger=logger,\n )\n\n def save_checkpoint(self, local_path: str, hdfs_path: str = None, global_step: int = 0, max_ckpt_to_keep=None):\n # record the previous global step\n self.previous_global_step = global_step\n\n if not self.checkpoint_config.async_save:\n self.ensure_checkpoint_capacity(max_ckpt_to_keep)\n\n local_path = local_mkdir_safe(local_path)\n dist_checkpoint_path = get_dist_checkpoint_path(local_path)\n\n # Note that model weights, optimizer states, and extra states are generated\n # together in a state dict, we save them in one time\n if self.use_dist_checkpointing:\n # Generate state dict for saving\n sharded_sd_metadata = self._build_sharded_state_dict_metadata()\n state_dict = self.generate_state_dict(\n self.should_save_model,\n self.should_save_optimizer,\n self.should_save_extra,\n metadata=sharded_sd_metadata,\n )\n log_with_rank(f\"Generated state dict for saving: {state_dict.keys()}\", rank=self.rank, logger=logger)\n for vpp_rank, model in enumerate(self.model):\n if len(self.model) > 1:\n model_i_keys = state_dict[f\"model{vpp_rank}\"].keys()\n log_with_rank(f\"Generated state dict for saving: {model_i_keys}\", rank=self.rank, logger=logger)\n else:\n log_with_rank(\n f\"Generated state dict for saving: {state_dict['model'].keys()}\", rank=self.rank, logger=logger\n )\n # Start Async save if enabled\n async_save_request = save_dist_checkpointing(\n sharded_state_dict=state_dict,\n ckpt_path=dist_checkpoint_path,\n async_save=self.checkpoint_config.async_save,\n content_metadata=sharded_sd_metadata,\n )\n\n # Synchronize all async save requests\n if not self.checkpoint_config.async_save:\n assert async_save_request is None, \"Async save request should be None when not using async save.\"\n torch.distributed.barrier()\n else:\n assert self.use_hf_checkpoint, \"When not using distributed checkpointing, use_hf_checkpoint should be True.\"\n # Generate optimizer and exra state dicts\n sharded_sd_metadata = self._build_sharded_state_dict_metadata()\n state_dict = self.generate_state_dict(\n generate_model=False,\n generate_optimizer=self.should_save_optimizer,\n generate_extra=self.should_save_extra,\n metadata=sharded_sd_metadata,\n )\n # Save optimizer and extra states to local path\n # Start Async save if enabled\n async_save_request = save_dist_checkpointing(\n sharded_state_dict=state_dict,\n ckpt_path=dist_checkpoint_path,\n async_save=self.checkpoint_config.async_save,\n content_metadata=sharded_sd_metadata,\n )\n\n # Synchronize all async save requests\n if not self.checkpoint_config.async_save:\n assert async_save_request is None, \"Async save request should be None when not using async save.\"\n torch.distributed.barrier()\n\n if self.should_save_model:\n # Save adapter-only checkpoint if PEFT is enabled\n if self.peft_cls is not None:\n from verl.utils.megatron_peft_utils import save_adapter_checkpoint\n\n adapter_ckpt_path = os.path.join(local_path, \"adapter_checkpoint\")\n\n # Save adapter weights only (much smaller than full model)\n save_adapter_checkpoint(\n self.model,\n adapter_ckpt_path,\n self.rank,\n )\n\n log_with_rank(\n f\"Saved adapter-only checkpoint to {adapter_ckpt_path}\",\n rank=self.rank,\n logger=logger,\n log_only_rank_0=True,\n )\n elif self.use_hf_checkpoint:\n # Use mbridge to save HF model checkpoint\n log_with_rank(f\"Saving HF model checkpoint to {local_path} with bridge\", rank=self.rank, logger=logger)\n hf_ckpt_path = get_hf_model_checkpoint_path(local_path)\n if self.vanilla_bridge:\n extended_args = {}\n mbridge_config = getattr(self.checkpoint_config, \"mbridge_config\", None) or {}\n for sig in inspect.signature(self.bridge.save_weights).parameters:\n if sig == \"weights_path\" or sig == \"models\":\n continue\n if sig in mbridge_config:\n extended_args[sig] = mbridge_config[sig]\n self.bridge.save_weights(self.model, hf_ckpt_path, **extended_args)\n else:\n self.bridge.save_hf_weights(self.model, hf_ckpt_path)\n\n log_with_rank(f\"Saved bridge checkpoint to {hf_ckpt_path}\", rank=self.rank, logger=logger)\n\n # Only rank 0 saves the hf config and tokenizer to huggingface path\n # No matter whether we save hf model or not\n if self.rank == 0:\n # Save tokenizer\n hf_config_tokenizer_path = get_hf_model_checkpoint_path(local_path)\n if self.processing_class is not None:\n self.processing_class.save_pretrained(hf_config_tokenizer_path)\n # Save huggingface config\n self.hf_config.save_pretrained(hf_config_tokenizer_path)\n if hasattr(self.hf_config, \"name_or_path\") and self.hf_config.name_or_path:\n try:\n generation_config = GenerationConfig.from_pretrained(self.hf_config.name_or_path)\n generation_config.save_pretrained(hf_config_tokenizer_path)\n except Exception:\n # if the generation config isn't available, we don't save it\n pass\n log_with_rank(\n f\"Saved Huggingface config and tokenizer to {hf_config_tokenizer_path}\",\n rank=self.rank,\n logger=logger,\n log_only_rank_0=True,\n )\n\n if self.should_save_extra:\n if self.rank == 0:\n # Save transformer config\n print(self.transformer_config)\n bypass_keys = [\n \"finalize_model_grads_func\",\n \"grad_scale_func\",\n \"no_sync_func\",\n \"grad_sync_func\",\n \"param_sync_func\",\n \"generation_config\",\n \"_pg_collection\",\n ]\n backup = {}\n for k in bypass_keys:\n if hasattr(self.transformer_config, k):\n backup[k] = getattr(self.transformer_config, k, None)\n delattr(self.transformer_config, k)\n transformer_config_dict = asdict(self.transformer_config)\n for k in backup:\n setattr(self.transformer_config, k, backup[k])\n to_convert_types = {torch.dtype: str, AttnBackend: str}\n ignore_types = [Callable]\n pop_keys = []\n for key, value in transformer_config_dict.items():\n if type(value) in to_convert_types:\n transformer_config_dict[key] = to_convert_types[type(value)](value)\n if type(value) in ignore_types:\n pop_keys.append(key)\n if callable(value):\n pop_keys.append(key)\n for key in pop_keys:\n transformer_config_dict.pop(key)\n transformer_config_path = get_transformer_config_checkpoint_path(local_path)\n with open(transformer_config_path, \"w\") as f:\n json.dump(transformer_config_dict, f, indent=2)\n\n if self.should_save_hf_model and not self.use_hf_checkpoint:\n # wait for everyone to dump to local\n if self.bridge is not None:\n hf_model_ckpt_path = get_hf_model_checkpoint_path(local_path)\n if self.vanilla_bridge:\n extended_args = {}\n mbridge_config = getattr(self.checkpoint_config, \"mbridge_config\", None) or {}\n for sig in inspect.signature(self.bridge.save_weights).parameters:\n if sig == \"weights_path\" or sig == \"models\":\n continue\n if sig in mbridge_config:\n extended_args[sig] = mbridge_config[sig]\n self.bridge.save_weights(self.model, hf_model_ckpt_path, **extended_args)\n else:\n self.bridge.save_hf_weights(self.model, hf_model_ckpt_path)\n else:\n state_dict = self.weight_saver(\n self.model,\n self.hf_config,\n dtype=self.param_dtype,\n is_value_model=self.is_value_model,\n tie_word_embeddings=self.share_embeddings_and_output_weights,\n )\n\n torch.distributed.barrier()\n if self.rank == 0:\n hf_model_ckpt_path = get_hf_model_checkpoint_path(local_path)\n import warnings\n\n from accelerate import init_empty_weights\n\n with init_empty_weights(), warnings.catch_warnings():\n warnings.simplefilter(\"ignore\")\n if \"mistral7b-rm\" in self.config.model.path:\n from transformers import MistralForSequenceClassification\n\n model = MistralForSequenceClassification.from_pretrained(\n self.config.model.path\n ) # use score head instead of lm_head\n state_dict[\"score.weight\"] = state_dict[\"score.weight\"]\n else:\n from transformers import AutoModelForCausalLM\n\n model = AutoModelForCausalLM.from_pretrained(self.config.model.path, torch_dtype=\"auto\")\n model.save_pretrained(hf_model_ckpt_path, state_dict=state_dict)\n log_with_rank(\n f\"Saved Huggingface config and tokenizer to {hf_model_ckpt_path}\",\n rank=self.rank,\n logger=logger,\n log_only_rank_0=True,\n )\n\n if hdfs_path is not None:\n log_with_rank(\n f\"Uploading checkpoint to {hdfs_path}\", rank=self.rank, logger=logger, log_only_rank_0=True\n )\n from verl.utils import hdfs_io\n\n hdfs_io.makedirs(hdfs_path, exist_ok=True)\n hdfs_io.copy(src=hf_model_ckpt_path, dst=hdfs_path, dirs_exist_ok=True)\n log_with_rank(\n f\"HDFS checkpoint uploaded to {hdfs_path}\",\n rank=self.rank,\n logger=logger,\n log_only_rank_0=True,\n )\n\n def finalize_save_fn():\n # Rank 0 uploads checkpoint to HDFS if hdfs_path is provided\n log_with_rank(\n f\"Dist checkpointing save completed for {dist_checkpoint_path}\", rank=self.rank, logger=logger\n )\n if self.rank == 0:\n if hdfs_path is not None:\n log_with_rank(f\"Uploading checkpoint to {hdfs_path}\", rank=self.rank, logger=logger)\n from verl.utils import hdfs_io\n\n hdfs_io.makedirs(hdfs_path, exist_ok=True)\n hdfs_io.copy(src=dist_checkpoint_path, dst=hdfs_path, dirs_exist_ok=True)\n hdfs_io.copy(src=hf_config_tokenizer_path, dst=hdfs_path, dirs_exist_ok=True)\n\n # update latest_checkpointed_iteration.txt when async_save is True\n if self.checkpoint_config.async_save and self.rank == 0:\n log_with_rank(\n f\"Update latest_checkpointed_iteration.txt to step {global_step}\",\n rank=self.rank,\n logger=logger,\n )\n local_latest_checkpointed_iteration = os.path.join(\n os.path.dirname(os.path.dirname(local_path)), \"latest_checkpointed_iteration.txt\"\n )\n with open(local_latest_checkpointed_iteration, \"w\") as f:\n f.write(str(global_step))\n\n self.register_checkpoint(local_path, max_ckpt_to_keep)\n\n if self.checkpoint_config.async_save:\n assert async_save_request is not None, \"Async save request should not be None when using async save.\"\n async_save_request.add_finalize_fn(finalize_save_fn)\n from megatron.core.dist_checkpointing.strategies.base import async_calls\n\n async_calls.schedule_async_request(async_save_request)\n else:\n finalize_save_fn()\n"}69{"file_name": "verl__utils__dataset__dataset_utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n\n# http://www.apache.org/licenses/LICENSE-2.0\n\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nfrom enum import Enum\n\nimport torch\nfrom tensordict.tensorclass import NonTensorData\n\n\nclass DatasetPadMode(str, Enum):\n \"\"\"Padding mode for dataset\"\"\"\n\n RIGHT = \"right\"\n LEFT_RIGHT = \"left_right\"\n NO_PADDING = \"no_padding\"\n\n\nclass SFTTensorCollator:\n \"\"\"\n A custom collate_fn that handles batching of sequences.\n 1. for variable-length sequences, convert them into NestedTensors.\n 2. for fixed-length sequences, use default_collate.\n \"\"\"\n\n def __init__(self, pad_mode: DatasetPadMode = DatasetPadMode.LEFT_RIGHT):\n self.pad_mode = pad_mode\n\n def __call__(self, batch: list[dict[str, any]]) -> dict[str, any]:\n if self.pad_mode == DatasetPadMode.NO_PADDING:\n return self.collate_variable_batch(batch)\n elif self.pad_mode in [DatasetPadMode.RIGHT, DatasetPadMode.LEFT_RIGHT]:\n from torch.utils.data import default_collate\n\n return default_collate(batch)\n else:\n raise NotImplementedError(f\"pad_mode {self.pad_mode} not implemented\")\n\n def collate_variable_batch(self, batch: list[dict[str, any]]) -> dict[str, any]:\n \"\"\"\n Collates a list of samples into a single batch.\n\n Args:\n batch: A list of dictionary samples from the dataset.\n\n Returns:\n A dictionary representing the batched data, with variable-length\n sequences converted to NestedTensors.\n \"\"\"\n\n final_batch = {}\n\n tensor_keys = set().union(*(d.keys() for d in batch))\n\n # Handle tensor values by creating a NestedTensor.\n for key in tensor_keys:\n if isinstance(batch[0][key], torch.Tensor):\n tensors = [item[key] for item in batch]\n final_batch[key] = torch.nested.as_nested_tensor(tensors, layout=torch.jagged)\n else:\n tensors = [NonTensorData(item.get(key)) for item in batch]\n final_batch[key] = torch.stack(tensors, dim=0)\n\n return final_batch\n"}70{"file_name": "verl__utils__dataset__multiturn_sft_dataset.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n\n# http://www.apache.org/licenses/LICENSE-2.0\n\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nMulti-turn SFT dataset that supports training on conversation data with multiple turns\n\"\"\"\n\nimport logging\nimport os\nimport re\nfrom functools import wraps\nfrom typing import Any, Optional\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn.functional as F\nfrom omegaconf import DictConfig, ListConfig\nfrom torch.utils.data import Dataset\nfrom transformers import PreTrainedTokenizer, ProcessorMixin\n\nfrom verl.models.transformers.qwen2_vl import get_rope_index\nfrom verl.utils import hf_tokenizer\nfrom verl.utils.chat_template import extract_system_prompt_and_generation\nfrom verl.utils.dataset.dataset_utils import DatasetPadMode\nfrom verl.utils.dataset.vision_utils import process_image, process_video\nfrom verl.utils.fs import copy_local_path_from_hdfs\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef once(func):\n \"\"\"Decorator to ensure a function runs only once. Subsequent calls do nothing.\"\"\"\n\n @wraps(func)\n def wrapper(*args, **kwargs):\n if not hasattr(wrapper, \"called\"):\n wrapper.called = True\n return func(*args, **kwargs)\n\n return wrapper\n\n\n@once\ndef print_assembled_message(tokenizer, message_list, input_ids, loss_mask, attn_mask, tools):\n \"\"\"\n Print the message after applying the chat template\n \"\"\"\n\n tokenized = tokenizer.apply_chat_template(message_list, add_generation_prompt=False, tokenize=False, tools=tools)\n sep = \"\\n\\n\"\n str = f\"tokenized entire message:\\n{tokenized}\"\n str += sep\n str += f\"tokenized seperately :\\n{tokenizer.decode(input_ids)}\"\n\n logger.debug(str)\n\n\ndef convert_nested_value_to_list_recursive(data_item):\n if isinstance(data_item, dict):\n return {k: convert_nested_value_to_list_recursive(v) for k, v in data_item.items()}\n elif isinstance(data_item, list):\n return [convert_nested_value_to_list_recursive(elem) for elem in data_item]\n elif isinstance(data_item, np.ndarray):\n # Convert to list, then recursively process the elements of the new list\n return convert_nested_value_to_list_recursive(data_item.tolist())\n else:\n # Base case: item is already a primitive type (int, str, float, bool, etc.)\n return data_item\n\n\nclass MultiTurnSFTDataset(Dataset):\n \"\"\"\n Dataset for multi-turn conversations where each assistant response should be trained\n\n Args:\n data_files (str or list): Path(s) to Parquet file(s).\n tokenizer (PreTrainedTokenizer): For the tokenization of text to token IDs.\n config (DictConfig): Options like cache_dir, prompt_key, max_prompt_length, truncation, etc.\n processor (ProcessorMixin, optional): Multimodal preprocessor for images/videos.\n max_samples (int, optional): Limit the number of samples. Defaults to -1 (use all).\n \"\"\"\n\n def __init__(\n self,\n parquet_files: str | list[str],\n tokenizer: PreTrainedTokenizer,\n config: DictConfig,\n processor: Optional[ProcessorMixin] = None,\n max_samples: int = -1,\n ):\n # Set defaults and extract parameters from config if provided\n config = config or {}\n self.pad_mode = config.get(\"pad_mode\", \"right\")\n assert self.pad_mode in [\"right\", \"no_padding\"], (\n f\"Expect pad_mode to be 'right' or 'no_padding'. Got {self.pad_mode}\"\n )\n self.truncation = config.get(\"truncation\", \"error\")\n # for right padding\n self.max_length = config.get(\"max_length\", 1024)\n # Get messages_key from the new multiturn config structure\n self.messages_key = config.get(\"messages_key\", \"messages\")\n self.image_key = config.get(\"image_key\", \"images\")\n self.video_key = config.get(\"video_key\", \"videos\")\n self.image_patch_size = config.get(\n \"image_patch_size\", processor.image_processor.patch_size if processor else None\n )\n self.tools_key = config.get(\"tools_key\", \"tools\")\n self.enable_thinking_key = config.get(\"enable_thinking_key\", \"enable_thinking\")\n self.enable_thinking_default = config.get(\"enable_thinking_default\", None)\n self.apply_chat_template_kwargs = config.get(\"apply_chat_template_kwargs\", {})\n self.shuffle = config.get(\"shuffle\", False)\n self.seed = config.get(\"seed\")\n self.max_samples = max_samples\n self.ignore_input_ids_mismatch = config.get(\"ignore_input_ids_mismatch\", False)\n assert self.truncation in [\"error\", \"left\", \"right\"]\n\n if not isinstance(parquet_files, list | ListConfig):\n parquet_files = [parquet_files]\n\n self.parquet_files = parquet_files\n if isinstance(tokenizer, str):\n tokenizer = hf_tokenizer(tokenizer)\n self.tokenizer: PreTrainedTokenizer = tokenizer\n self.processor = processor\n\n self._download()\n self._read_files_and_process()\n\n def _download(self):\n for i, parquet_file in enumerate(self.parquet_files):\n self.parquet_files[i] = copy_local_path_from_hdfs(parquet_file, verbose=True)\n\n def _read_files_and_process(self):\n def series_to_item(ls):\n import numpy\n import pandas\n\n while isinstance(ls, pandas.core.series.Series | numpy.ndarray) and len(ls) == 1:\n ls = ls[0]\n return ls\n\n dataframes = []\n for parquet_file in self.parquet_files:\n # default loader loads some list as np.ndarray, which fails the tokenizer\n dataframe = pd.read_parquet(parquet_file, dtype_backend=\"pyarrow\")\n dataframes.append(dataframe)\n self.dataframe = pd.concat(dataframes)\n\n total = len(self.dataframe)\n print(f\"dataset len: {len(self.dataframe)}\")\n\n if self.max_samples > 0 and self.max_samples < total:\n if self.shuffle:\n rngs_args = (self.seed,) if self.seed is not None else ()\n rng = np.random.default_rng(*rngs_args)\n indices = rng.choice(total, size=self.max_samples, replace=False)\n else:\n indices = np.arange(self.max_samples)\n self.dataframe = self.dataframe.iloc[indices.tolist()]\n print(f\"selected {self.max_samples} random samples out of {total}\")\n\n # Extract messages list from dataframe\n self.messages = self.dataframe[self.messages_key].apply(convert_nested_value_to_list_recursive).tolist()\n\n # Extract tools list from dataframe\n if self.tools_key in self.dataframe.columns:\n self.tools = self.dataframe[self.tools_key].apply(convert_nested_value_to_list_recursive).tolist()\n else:\n self.tools = None\n # Extract enable_thinking list from dataframe\n if self.enable_thinking_key in self.dataframe.columns:\n self.enable_thinking = self.dataframe[self.enable_thinking_key].tolist()\n else:\n self.enable_thinking = None\n\n # system prompt: <|im_start|>system\\nYou are a helpful assistant.<|im_end|>\\n\n # generation prompt: <|im_start|>assistant\\n\n self.system_prompt, self.generation_prompt = extract_system_prompt_and_generation(self.tokenizer)\n\n def __len__(self):\n return len(self.messages)\n\n def _process_single_message(\n self,\n index: int,\n message: dict[str, Any],\n full_message: list,\n tools: Optional[list[dict[str, Any]]] = None,\n enable_thinking: Optional[bool] = None,\n ) -> tuple[list[int], list[int], list[int]]:\n \"\"\"\n Process a single message and return its tokenized representation.\n\n Args:\n index: turn index in the conversation\n message: A single message dictionary\n images: List of images to be used\n videos: List of videos to be used\n tools: List of tools to be used\n enable_thinking: Whether to enable thinking mode\n\n Returns:\n Tuple of (input_ids, loss_mask, attention_mask, dict[str, torch.Tensor])\n \"\"\"\n processor = self.processor if self.processor is not None else self.tokenizer\n apply_chat_template_kwargs = {**self.apply_chat_template_kwargs}\n if enable_thinking is not None:\n apply_chat_template_kwargs[\"enable_thinking\"] = enable_thinking\n\n inputs = processor.apply_chat_template(\n [message],\n tools=tools,\n add_generation_prompt=False,\n tokenize=True,\n return_dict=True,\n return_tensors=\"pt\",\n **apply_chat_template_kwargs,\n )\n\n inputs = dict(inputs)\n input_ids = inputs.pop(\"input_ids\")[0]\n attention_mask = inputs.pop(\"attention_mask\")[0]\n\n # remove system prompt if exists\n if index != 0 and message[\"role\"] != \"system\":\n input_ids = input_ids[len(self.system_prompt) :]\n attention_mask = attention_mask[len(self.system_prompt) :]\n\n if message[\"role\"] == \"assistant\":\n loss_mask = torch.ones_like(attention_mask)\n # mask out generation prompt if assistant message\n loss_mask[: len(self.generation_prompt)] = 0\n else:\n loss_mask = torch.zeros_like(attention_mask)\n\n return input_ids, loss_mask, attention_mask, inputs\n\n def _build_messages(self, example: dict):\n \"\"\"Replace <image> and <video> placeholder in messages with corresponding image and video\n which is required by processor.apply_chat_template.\n - <image>: {\"type\": \"image\", \"image\": image}\n - <video>: {\"type\": \"video\", \"video\": video}\n\n Args:\n example: Row dictionary from dataframe.\n\n Returns:\n messages: List of messages with replaced placeholder.\n \"\"\"\n messages: list = example[self.messages_key]\n images = example[self.image_key] if self.image_key in example else []\n videos = example[self.video_key] if self.video_key in example else []\n\n image_offset, video_offset = 0, 0\n for message in messages:\n if self.image_key not in example and self.video_key not in example:\n continue\n assert self.processor is not None, \"processor is needed to process image and video\"\n\n content = message[\"content\"]\n if not isinstance(content, str):\n continue\n\n content_list = []\n segments = re.split(\"(<image>|<video>)\", content)\n segments = [item for item in segments if item != \"\"]\n for segment in segments:\n if segment == \"<image>\":\n image = process_image(images[image_offset], image_patch_size=self.image_patch_size)\n content_list.append({\"type\": \"image\", \"image\": image})\n image_offset += 1\n elif segment == \"<video>\":\n video = process_video(videos[video_offset], image_patch_size=self.image_patch_size)\n content_list.append({\"type\": \"video\", \"video\": video})\n video_offset += 1\n else:\n content_list.append({\"type\": \"text\", \"text\": segment})\n message[\"content\"] = content_list\n\n assert image_offset == len(images), f\"image_offset {image_offset} != len(images) {len(images)}\"\n assert video_offset == len(videos), f\"video_offset {video_offset} != len(videos) {len(videos)}\"\n return messages\n\n def __getitem__(self, item):\n row_dict: dict = self.dataframe.iloc[item].to_dict()\n messages = self._build_messages(row_dict)\n tools = self.tools[item] if self.tools is not None else None\n enable_thinking = (\n self.enable_thinking[item] if self.enable_thinking is not None else self.enable_thinking_default\n )\n\n # 1. tokenize each message\n input_ids, loss_mask, attention_mask, multi_modal_inputs = [], [], [], {}\n for i, message in enumerate(messages):\n _input_ids, _loss_mask, _attention_mask, _inputs = self._process_single_message(\n index=i,\n message=message,\n full_message=messages,\n tools=tools if i == 0 else None,\n enable_thinking=enable_thinking,\n )\n input_ids.append(_input_ids)\n loss_mask.append(_loss_mask)\n attention_mask.append(_attention_mask)\n for k, v in _inputs.items():\n multi_modal_inputs.setdefault(k, []).append(v)\n\n input_ids = torch.cat(input_ids, dim=0)\n loss_mask = torch.cat(loss_mask, dim=0)\n attention_mask = torch.cat(attention_mask, dim=0)\n assert input_ids.shape == loss_mask.shape == attention_mask.shape, (\n f\"Shape mismatch: {input_ids.shape}, {loss_mask.shape}, {attention_mask.shape}\"\n )\n\n print_assembled_message(self.tokenizer, messages, input_ids, loss_mask, attention_mask, tools)\n self.sanity_check(input_ids, messages, tools, enable_thinking)\n\n # Since the tokenizer may return user-customized results, we need to filter out inconsistent tensor shapes\n keys_to_remove = []\n for k, v in multi_modal_inputs.items():\n if len(v) > 0 and v[0] is not None and isinstance(v[0], torch.Tensor):\n # Check if all tensors in the list have the same shape\n first_shape = v[0].shape[1:]\n if not all(tensor.shape[1:] == first_shape for tensor in v):\n keys_to_remove.append(k)\n\n for k in keys_to_remove:\n del multi_modal_inputs[k]\n\n for k, v in multi_modal_inputs.items():\n multi_modal_inputs[k] = torch.concat(v, dim=0)\n\n # 2. handle position_ids for Qwen-VL series models\n if self.processor is not None and \"Qwen2VLImageProcessor\" in self.processor.image_processor.__class__.__name__:\n image_grid_thw = multi_modal_inputs.get(\"image_grid_thw\", None)\n video_grid_thw = multi_modal_inputs.get(\"video_grid_thw\", None)\n second_per_grid_ts = multi_modal_inputs.get(\"second_per_grid_ts\", None)\n\n vision_position_ids = get_rope_index(\n self.processor,\n input_ids=input_ids,\n image_grid_thw=image_grid_thw,\n video_grid_thw=video_grid_thw,\n second_per_grid_ts=second_per_grid_ts,\n attention_mask=attention_mask,\n ) # (3, seq_len)\n text_position_ids = torch.arange(input_ids.shape[0], dtype=torch.long).unsqueeze(0) # (1, seq_len)\n position_ids = torch.cat((text_position_ids, vision_position_ids), dim=0) # (4, seq_length)\n else:\n position_ids = torch.arange(input_ids.shape[0], dtype=torch.long) # (seq_len,)\n\n # 3. handle padding\n sequence_length = input_ids.shape[0]\n # Handle sequence length\n if self.pad_mode == DatasetPadMode.RIGHT:\n if sequence_length < self.max_length:\n # Pad sequences\n pad_token_id = self.tokenizer.pad_token_id if self.tokenizer.pad_token_id is not None else 0\n padded_input_ids = torch.full((self.max_length - sequence_length,), pad_token_id, dtype=input_ids.dtype)\n padded_attention_mask = torch.zeros((self.max_length - sequence_length,), dtype=attention_mask.dtype)\n padded_loss_mask = torch.zeros((self.max_length - sequence_length,), dtype=loss_mask.dtype)\n\n input_ids = torch.cat((input_ids, padded_input_ids))\n attention_mask = torch.cat((attention_mask, padded_attention_mask))\n loss_mask = torch.cat((loss_mask, padded_loss_mask))\n position_ids = F.pad(position_ids, (0, self.max_length - sequence_length), value=0)\n elif sequence_length > self.max_length:\n if self.truncation == \"left\":\n input_ids = input_ids[-self.max_length :]\n attention_mask = attention_mask[-self.max_length :]\n loss_mask = loss_mask[-self.max_length :]\n position_ids = position_ids[..., -self.max_length :]\n elif self.truncation == \"right\":\n input_ids = input_ids[: self.max_length]\n attention_mask = attention_mask[: self.max_length]\n loss_mask = loss_mask[: self.max_length]\n position_ids = position_ids[..., : self.max_length]\n elif self.truncation == \"error\":\n raise ValueError(f\"{sequence_length=} is larger than {self.max_length=}\")\n else:\n raise ValueError(f\"Unknown truncation method {self.truncation}\")\n\n res = {\n \"input_ids\": input_ids,\n \"attention_mask\": attention_mask,\n \"position_ids\": position_ids,\n \"loss_mask\": loss_mask,\n }\n if len(multi_modal_inputs) > 0:\n res[\"multi_modal_inputs\"] = multi_modal_inputs\n return res\n elif self.pad_mode == DatasetPadMode.NO_PADDING:\n # truncate input_ids if it is longer than max_length\n if len(input_ids) > self.max_length:\n input_ids = input_ids[: self.max_length]\n loss_mask = loss_mask[: self.max_length]\n position_ids = position_ids[..., : self.max_length]\n\n # return nested tensor with out padding\n res = {\n \"input_ids\": input_ids,\n \"position_ids\": position_ids,\n \"loss_mask\": loss_mask,\n }\n if len(multi_modal_inputs) > 0:\n res[\"multi_modal_inputs\"] = multi_modal_inputs\n return res\n else:\n raise ValueError(f\"Unknown pad mode {self.pad_mode}\")\n\n def sanity_check(self, input_ids: torch.Tensor, messages: list[dict], tools: list[dict], enable_thinking: bool):\n \"\"\"Check concatenated input_ids of apply_chat_template to each turn equals\n apply_chat_template to whole messages.\n \"\"\"\n processor = self.processor if self.processor is not None else self.tokenizer\n apply_chat_template_kwargs = {**self.apply_chat_template_kwargs}\n if enable_thinking is not None:\n apply_chat_template_kwargs[\"enable_thinking\"] = enable_thinking\n inputs = processor.apply_chat_template(\n messages,\n tools=tools,\n add_generation_prompt=False,\n tokenize=True,\n return_dict=True,\n return_tensors=\"pt\",\n **apply_chat_template_kwargs,\n )\n\n error_message = (\n \"MultiTurnSFTDataset apply_chat_template to each turn separately and concat `input_ids` \"\n \"as a whole sequence, which may not equal to apply_chat_template to whole messages at once.\\n\"\n \"For example, Qwen Thinking series models add <think></think> tags to last turn, please check \"\n \"your tokenizer chat template settings.\\n\"\n \"Set `ignore_input_ids_mismatch=True` to ignore input_ids mismatch and use the concatenated \"\n \"input_ids as the final input_ids. \"\n )\n\n if not torch.equal(input_ids, inputs[\"input_ids\"].squeeze(0)):\n if self.ignore_input_ids_mismatch:\n logger.warning_once(error_message)\n else:\n raise AssertionError(error_message)\n"}71{"file_name": "verl__utils__dataset__rl_dataset.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport copy\nimport logging\nimport os\nimport re\nimport traceback\nfrom collections import defaultdict\nfrom io import BytesIO\nfrom typing import Optional\n\nimport datasets\nimport numpy as np\nimport torch\nfrom omegaconf import DictConfig, ListConfig\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nfrom transformers import PreTrainedTokenizer, ProcessorMixin\n\nfrom verl.utils.import_utils import load_extern_object\n\nlogger = logging.getLogger(__name__)\n\n\ndef collate_fn(data_list: list[dict]) -> dict:\n \"\"\"\n Collate a batch of sample dicts into batched tensors and arrays.\n\n Args:\n data_list: List of dicts mapping feature names to torch.Tensor or other values.\n\n Returns:\n Dict where tensor entries are stacked into a torch.Tensor of shape\n (batch_size, \\\\*dims) and non-tensor entries are converted to\n np.ndarray of dtype object with shape (batch_size,).\n \"\"\"\n tensors = defaultdict(list)\n non_tensors = defaultdict(list)\n\n for data in data_list:\n for key, val in data.items():\n if isinstance(val, torch.Tensor):\n tensors[key].append(val)\n else:\n non_tensors[key].append(val)\n\n for key, val in tensors.items():\n tensors[key] = torch.stack(val, dim=0)\n\n for key, val in non_tensors.items():\n non_tensors[key] = np.fromiter(val, dtype=object, count=len(val))\n\n return {**tensors, **non_tensors}\n\n\nclass RLHFDataset(Dataset):\n \"\"\"\n Load and preprocess RLHF data from Parquet files.\n\n - Caches files locally.\n - Reads into a HuggingFace Dataset and tokenizes prompts.\n - Optionally handles images/videos via a ProcessorMixin.\n - Filters prompts over a max length.\n - Supports resuming from checkpoints.\n\n Args:\n data_files (str or list): Path(s) to Parquet file(s).\n tokenizer (PreTrainedTokenizer): For the tokenization of text to token IDs.\n config (DictConfig): Options like cache_dir, prompt_key, max_prompt_length, truncation, etc.\n processor (ProcessorMixin, optional): Multimodal preprocessor for images/videos.\n \"\"\"\n\n def __init__(\n self,\n data_files: str | list[str],\n tokenizer: PreTrainedTokenizer,\n config: DictConfig,\n processor: Optional[ProcessorMixin] = None,\n max_samples: int = -1,\n ):\n if not isinstance(data_files, list | ListConfig):\n data_files = [data_files]\n\n self.data_files = copy.deepcopy(data_files)\n self.original_data_files = copy.deepcopy(data_files) # use for resume\n self.tokenizer = tokenizer\n self.processor = processor\n self.max_samples = max_samples\n self.config = config\n\n self.cache_dir = os.path.expanduser(config.get(\"cache_dir\", \"~/.cache/verl/rlhf\"))\n self.prompt_key = config.get(\"prompt_key\", \"prompt\")\n self.image_key = config.get(\"image_key\", \"images\")\n self.video_key = config.get(\"video_key\", \"videos\")\n self.image_patch_size = config.get(\"image_patch_size\", 14)\n self.max_prompt_length = config.get(\"max_prompt_length\", 1024)\n self.return_raw_chat = config.get(\"return_raw_chat\", False)\n self.return_full_prompt = config.get(\"return_full_prompt\", False)\n self.truncation = config.get(\"truncation\", \"error\")\n self.filter_overlong_prompts = config.get(\"filter_overlong_prompts\", True)\n self.apply_chat_template_kwargs = config.get(\"apply_chat_template_kwargs\", {})\n\n self.tool_config_path = config.get(\"tool_config_path\", None)\n self.tool_schemas = None\n if self.tool_config_path:\n try:\n from verl.tools.utils.tool_registry import initialize_tools_from_config\n\n tool_list = initialize_tools_from_config(self.tool_config_path)\n # match ToolAgentLoop behaviour: model_dump to plain dicts\n self.tool_schemas = [\n tool.tool_schema.model_dump(exclude_unset=True, exclude_none=True) for tool in tool_list\n ]\n except Exception as e:\n logger.warning(\"Failed to initialize tools from %s: %s\", self.tool_config_path, e)\n self.tool_schemas = None\n\n self.num_workers = config.get(\"filter_overlong_prompts_workers\", max(1, os.cpu_count() // 4))\n self.num_workers = min(self.num_workers, os.cpu_count()) if self.num_workers is not None else None\n self.use_shm = config.get(\"use_shm\", False)\n self.chat_template_func = config.get(\"chat_template_func\", None)\n self.need_tools_kwargs = config.get(\"need_tools_kwargs\", False)\n self.filter_prompts = config.get(\"filter_prompts\", True)\n self.serialize_dataset = False\n self.return_multi_modal_inputs = config.get(\"return_multi_modal_inputs\", True)\n self.shuffle = config.get(\"shuffle\", False)\n self.seed = config.get(\"seed\")\n\n self._download()\n self._read_files_and_tokenize()\n\n def _download(self, use_origin_parquet=False):\n from verl.utils.fs import copy_to_local\n\n data_files = self.data_files if not use_origin_parquet else self.original_data_files\n for i, parquet_file in enumerate(data_files):\n self.data_files[i] = copy_to_local(src=parquet_file, cache_dir=self.cache_dir, use_shm=self.use_shm)\n\n def _read_files_and_tokenize(self):\n dataframes = []\n for parquet_file in self.data_files:\n # read files and cache\n if parquet_file.endswith(\".parquet\"):\n dataframe = datasets.load_dataset(\"parquet\", data_files=parquet_file)[\"train\"]\n elif parquet_file.endswith(\".json\"):\n dataframe = datasets.load_dataset(\"json\", data_files=parquet_file)[\"train\"]\n else:\n raise ValueError(f\"Unsupported file format: {parquet_file}\")\n dataframes.append(dataframe)\n self.dataframe: datasets.Dataset = datasets.concatenate_datasets(dataframes)\n\n total = len(self.dataframe)\n print(f\"dataset len: {len(self.dataframe)}\")\n\n if self.max_samples > 0 and self.max_samples < total:\n if self.shuffle:\n rngs_args = (self.seed,) if self.seed is not None else ()\n rng = np.random.default_rng(*rngs_args)\n indices = rng.choice(total, size=self.max_samples, replace=False)\n else:\n indices = np.arange(self.max_samples)\n self.dataframe = self.dataframe.select(indices.tolist())\n print(f\"selected {self.max_samples} random samples out of {total}\")\n\n self.dataframe = self.maybe_filter_out_long_prompts(self.dataframe)\n\n def maybe_filter_out_long_prompts(self, dataframe: datasets.Dataset = None):\n # filter out too long prompts\n if self.filter_overlong_prompts:\n tokenizer = self.tokenizer\n processor = self.processor\n prompt_key = self.prompt_key\n image_key = self.image_key\n video_key = self.video_key\n\n if processor is not None:\n from verl.utils.dataset.vision_utils import process_image, process_video\n\n def doc2len(doc) -> int:\n try:\n messages = self._build_messages(doc)\n # pass tool schemas if available so the processor can format prompts\n apply_kwargs = dict(**self.apply_chat_template_kwargs)\n if self.tool_schemas is not None:\n apply_kwargs[\"tools\"] = self.tool_schemas\n\n raw_prompt = self.processor.apply_chat_template(\n messages, add_generation_prompt=True, tokenize=False, **apply_kwargs\n )\n if image_key in doc and doc[image_key]:\n images = [\n process_image(image, image_patch_size=self.image_patch_size) for image in doc[image_key]\n ]\n else:\n images = None\n\n if video_key in doc and doc[video_key]:\n videos, video_metadata = zip(\n *[\n process_video(\n video, image_patch_size=self.image_patch_size, return_video_metadata=True\n )\n for video in doc[video_key]\n ],\n strict=True,\n )\n videos = list(videos)\n video_metadata = list(video_metadata)\n videos_kwargs = {\"video_metadata\": video_metadata, \"do_sample_frames\": False}\n else:\n videos = None\n videos_kwargs = {}\n\n return len(\n processor(text=[raw_prompt], images=images, videos=videos, videos_kwargs=videos_kwargs)[\n \"input_ids\"\n ][0]\n )\n except Exception:\n print(\"Error processing one of the samples, skipping...\")\n traceback.print_exc()\n return self.max_prompt_length + 1\n\n else:\n\n def doc2len(doc) -> int:\n try:\n apply_kwargs = dict(**self.apply_chat_template_kwargs)\n if self.tool_schemas is not None:\n apply_kwargs[\"tools\"] = self.tool_schemas\n\n return len(\n tokenizer.apply_chat_template(doc[prompt_key], add_generation_prompt=True, **apply_kwargs)\n )\n except Exception:\n print(\"Error processing one of the samples, skipping...\")\n traceback.print_exc()\n return self.max_prompt_length + 1\n\n dataframe = dataframe.filter(\n lambda doc: doc2len(doc) <= self.max_prompt_length,\n num_proc=self.num_workers,\n desc=f\"Filtering prompts longer than {self.max_prompt_length} tokens\",\n )\n\n print(f\"filter dataset len: {len(dataframe)}\")\n return dataframe\n\n def resume_dataset_state(self):\n self.serialize_dataset = not hasattr(self, \"original_data_files\")\n # resume dataframe if not it's serialized in data.pt\n if not self.serialize_dataset:\n self._download(use_origin_parquet=True) # download and resume from original parquet files\n self._read_files_and_tokenize()\n else:\n print(r\"old dataloader ckpt file is used, please train from scratch for better ckpt performance\")\n\n def __getstate__(self):\n if not self.serialize_dataset:\n state = self.__dict__.copy()\n\n if \"dataframe\" in state:\n del state[\"dataframe\"]\n return state\n\n return self.__dict__.copy()\n\n def __len__(self):\n return len(self.dataframe)\n\n def _build_messages(self, example: dict):\n \"\"\"Replace <image> and <video> placeholder in messages with corresponding image and video\n which is required by processor.apply_chat_template.\n - <image>: {\"type\": \"image\", **image}\n - <video>: {\"type\": \"video\", **video}\n\n Args:\n example: Row dictionary from dataframe.\n\n Returns:\n messages: List of messages with replaced placeholder.\n \"\"\"\n messages: list = example[self.prompt_key]\n # When concatenating image and video datasets, pop will return None for image or video sample\n images = example.pop(self.image_key, None) or []\n videos = example.pop(self.video_key, None) or []\n\n image_offset, video_offset = 0, 0\n for message in messages:\n if not images and not videos:\n continue\n assert self.processor is not None, \"processor is needed to process image and video\"\n\n content = message[\"content\"]\n if not isinstance(content, str):\n continue\n\n content_list = []\n segments = re.split(\"(<image>|<video>)\", content)\n segments = [item for item in segments if item != \"\"]\n for segment in segments:\n if segment == \"<image>\":\n assert image_offset < len(images), f\"image_offset {image_offset} >= len(images) {len(images)}\"\n image = images[image_offset]\n if isinstance(image, Image.Image):\n image = image.convert(\"RGB\")\n content_list.append({\"type\": \"image\", \"image\": image})\n elif isinstance(image, dict):\n if \"bytes\" in image:\n image[\"image\"] = Image.open(BytesIO(image[\"bytes\"]))\n content_list.append({\"type\": \"image\", **image})\n else:\n raise TypeError(f\"image must be dict or PIL.Image, unsupported image type: {type(image)}\")\n image_offset += 1\n elif segment == \"<video>\":\n assert video_offset < len(videos), f\"video_offset {video_offset} >= len(videos) {len(videos)}\"\n content_list.append({\"type\": \"video\", **videos[video_offset]})\n video_offset += 1\n else:\n content_list.append({\"type\": \"text\", \"text\": segment})\n message[\"content\"] = content_list\n\n assert image_offset == len(images), f\"image_offset {image_offset} != len(images) {len(images)}\"\n assert video_offset == len(videos), f\"video_offset {video_offset} != len(videos) {len(videos)}\"\n return messages\n\n def __getitem__(self, item):\n \"\"\"For rollout, apply_chat_template has been moved to AgentLoop, so we only return raw_prompt here.\"\"\"\n row_dict: dict = self.dataframe[item]\n row_dict[\"raw_prompt\"] = self._build_messages(row_dict)\n\n # TODO(wuxibin): We still need a dummy tensor to make sure DataProto.batch is not empty.\n # Remove this after deprecate DataProto by TensorDict.\n row_dict[\"dummy_tensor\"] = torch.tensor([0], dtype=torch.uint8)\n\n # add index for each prompt\n if \"extra_info\" not in row_dict or row_dict[\"extra_info\"] is None:\n row_dict[\"extra_info\"] = dict()\n index = row_dict.get(\"extra_info\", {}).get(\"index\", 0)\n tools_kwargs = row_dict.get(\"extra_info\", {}).get(\"tools_kwargs\", {})\n interaction_kwargs = row_dict.get(\"extra_info\", {}).get(\"interaction_kwargs\", {})\n need_tools_kwargs = row_dict.get(\"extra_info\", {}).get(\"need_tools_kwargs\", self.need_tools_kwargs)\n if need_tools_kwargs and not tools_kwargs:\n logger.warning(\"tools_kwargs is empty for index {}, data source: {}\", index, row_dict[\"data_source\"])\n row_dict[\"index\"] = index\n row_dict[\"tools_kwargs\"] = tools_kwargs\n row_dict[\"interaction_kwargs\"] = interaction_kwargs\n return row_dict\n\n @classmethod\n async def process_vision_info(\n cls,\n messages: list[dict],\n image_patch_size,\n config: DictConfig,\n ) -> tuple[list[Image.Image], list[tuple[torch.Tensor, dict]]]:\n \"\"\"Extract images and videos from messages.\n\n This method is called by AgentLoop (e.g SingleTurnAgentLoop) before apply_chat_template to\n the `raw_prompt` from dataset. User may customize RLHFDataset and override this method to\n support custom vision extraction.\n\n >>> messages = kwargs[\"raw_prompt\"]\n >>> images, videos = RLHFDataset.process_vision_info(messages, image_patch_size)\n >>> videos, video_metadatas = zip(*videos)\n >>> raw_prompt = processor.apply_chat_template(messages, tokenize=False)\n >>> inputs = processor(text=[raw_prompt], images=images, videos=videos,\n ... video_metadata=video_metadatas, do_sample_frames=False)\n\n Args:\n messages: List of messages from dataset `raw_prompt`.\n image_patch_size: Image patch size for processor.\n config: Config for dataset.\n\n Returns:\n images: List of images.\n videos: List of videos, each video is a tuple of (video_tensor, video_metadata).\n \"\"\"\n from qwen_vl_utils import process_vision_info\n\n images, videos = process_vision_info(messages, image_patch_size=image_patch_size, return_video_metadata=True)\n return images, videos\n\n def split(self, num_splits: int):\n \"\"\"\n split the dataset into num_splits sub-datasets\n Args:\n num_splits: specified number of splits\n Returns:\n List[RLHFDataset]: list of RLHFDataset splits\n Raises:\n ValueError: if num_splits is not a positive integer\n \"\"\"\n if not isinstance(num_splits, int) or num_splits <= 0:\n raise ValueError(f\"num_splits must be a positive integer, got {num_splits}\")\n\n if not hasattr(self, \"dataframe\"):\n raise AttributeError(\n \"dataframe not found in RLHFDataset\\n\"\n \"reason: _read_files_and_tokenize() not called or Parquet file loading failed\"\n )\n if self.dataframe is None:\n raise ValueError(\"RLHFDataset dataframe 为 None!\")\n\n total_samples = len(self.dataframe)\n print(f\"total_samples: {total_samples}\")\n if total_samples == 0:\n raise ValueError(\"Cannot split an empty dataset\")\n if total_samples % num_splits != 0:\n raise ValueError(f\"Cannot split dataset size {total_samples} into {num_splits} splits\")\n split_size = total_samples // num_splits\n splits = []\n\n for i in range(num_splits):\n start_idx = i * split_size\n end_idx = (i + 1) * split_size if i < num_splits - 1 else total_samples\n\n split_dataframe = self.dataframe.select(range(start_idx, end_idx))\n\n split_dataset = RLHFDataset(\n data_files=self.data_files,\n tokenizer=self.tokenizer,\n config=self.config,\n processor=self.processor,\n max_samples=self.max_samples,\n )\n split_dataset.dataframe = split_dataframe\n split_dataset.serialize_dataset = self.serialize_dataset\n split_dataset.original_data_files = self.original_data_files\n\n splits.append(split_dataset)\n\n return splits\n\n\ndef get_dataset_class(data_config: DictConfig):\n \"\"\"Get RLHF dataset class.\n\n Args:\n data_config: The data config.\n\n Returns:\n dataset_cls: The dataset class.\n \"\"\"\n\n # Check if a custom dataset class is specified in the data configuration\n # and if the path to the custom class is provided\n if \"custom_cls\" in data_config and data_config.custom_cls.get(\"path\", None) is not None:\n # Dynamically load the custom dataset class\n dataset_cls = load_extern_object(data_config.custom_cls.path, data_config.custom_cls.name)\n # Verify that the custom dataset class inherits from torch.utils.data.Dataset\n if not issubclass(dataset_cls, Dataset):\n raise TypeError(\n f\"The custom dataset class '{data_config.custom_cls.name}' from \"\n f\"'{data_config.custom_cls.path}' must inherit from torch.utils.data.Dataset\"\n )\n else:\n # Use the default RLHFDataset class if no custom class is specified\n dataset_cls = RLHFDataset\n print(f\"Using dataset class: {dataset_cls.__name__}\")\n\n return dataset_cls\n"}72{"file_name": "verl__utils__dataset__rm_dataset.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport os\nfrom typing import Optional\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset\n\nfrom verl.utils import hf_tokenizer\n\n\ndef download_files_distributed(download_fn):\n import torch.distributed\n\n if torch.distributed.is_initialized():\n if torch.distributed.get_rank() == 0:\n # download files\n download_fn()\n\n torch.distributed.barrier()\n else:\n # download anyway\n download_fn()\n\n\nclass RMDataset(Dataset):\n def __init__(\n self,\n parquet_files: str | list[str],\n tokenizer,\n prompt_key=\"prompt\",\n chosen_key=\"chosen\",\n rejected_key=\"rejected\",\n max_length=1024,\n add_eos=True,\n cache_dir=\"~/.cache/verl/rm\",\n max_samples: int = -1,\n shuffle: bool = False,\n seed: Optional[int] = None,\n ):\n if not isinstance(parquet_files, list):\n parquet_files = [parquet_files]\n\n self.parquet_files = parquet_files\n self.max_samples = max_samples\n self.shuffle = shuffle\n self.seed = seed\n self.cache_dir = os.path.expanduser(cache_dir)\n if isinstance(tokenizer, str):\n tokenizer = hf_tokenizer(tokenizer)\n self.tokenizer = tokenizer\n\n self.prompt_key = prompt_key\n self.chosen_key = chosen_key\n self.rejected_key = rejected_key\n\n self.add_eos = add_eos\n self.max_length = max_length\n\n self._download()\n self._read_files_and_tokenize()\n\n def _download(self):\n def _download_files():\n from verl.utils.fs import copy, is_non_local\n\n os.makedirs(self.cache_dir, exist_ok=True)\n assert os.path.exists(self.cache_dir)\n for i, parquet_file in enumerate(self.parquet_files):\n if is_non_local(parquet_file):\n dst = os.path.join(self.cache_dir, os.path.basename(parquet_file))\n if not os.path.exists(dst):\n copy(src=parquet_file, dst=dst)\n self.parquet_files[i] = dst\n\n download_files_distributed(_download_files)\n\n def _read_files_and_tokenize(self):\n dataframes = []\n for parquet_file in self.parquet_files:\n # read parquet files and cache\n dataframe = pd.read_parquet(parquet_file)\n dataframes.append(dataframe)\n self.dataframe = pd.concat(dataframes)\n\n total = len(self.dataframe)\n print(f\"dataset len: {len(self.dataframe)}\")\n\n if self.max_samples > 0 and self.max_samples < total:\n if self.shuffle:\n rngs_args = (self.seed,) if self.seed is not None else ()\n rng = np.random.default_rng(*rngs_args)\n indices = rng.choice(total, size=self.max_samples, replace=False)\n else:\n indices = np.arange(self.max_samples)\n self.dataframe = self.dataframe.iloc[indices.tolist()]\n print(f\"selected {self.max_samples} random samples out of {total}\")\n\n self.prompts = self.dataframe[self.prompt_key].tolist()\n self.chosen_responses = self.dataframe[self.chosen_key].tolist()\n self.rejected_responses = self.dataframe[self.rejected_key].tolist()\n\n def __len__(self):\n return len(self.prompts)\n\n def _pad_to_length(self, input_ids, attention_mask):\n curr_length = input_ids.shape[-1]\n\n if curr_length < self.max_length:\n input_ids = torch.cat(\n (input_ids, torch.zeros(size=(self.max_length - curr_length,), dtype=input_ids.dtype)), dim=-1\n )\n attention_mask = torch.cat(\n (attention_mask, torch.zeros(size=(self.max_length - curr_length,), dtype=attention_mask.dtype)), dim=-1\n )\n elif curr_length > self.max_length:\n input_ids = input_ids[: self.max_length]\n attention_mask = attention_mask[: self.max_length]\n\n return input_ids, attention_mask\n\n def __getitem__(self, item):\n prompt = self.prompts[item]\n chosen_response = self.chosen_responses[item]\n rejected_response = self.rejected_responses[item]\n\n prompt_ids = self.tokenizer(prompt, return_tensors=\"pt\")[\"input_ids\"][0]\n chosen_response_ids = self.tokenizer(chosen_response, return_tensors=\"pt\")[\"input_ids\"][0]\n rejected_response_ids = self.tokenizer(rejected_response, return_tensors=\"pt\")[\"input_ids\"][0]\n\n if self.add_eos:\n chosen_response_ids = torch.cat((chosen_response_ids, torch.tensor([self.tokenizer.eos_token_id])), dim=-1)\n rejected_response_ids = torch.cat(\n (rejected_response_ids, torch.tensor([self.tokenizer.eos_token_id])), dim=-1\n )\n\n chosen_input_ids = torch.cat((prompt_ids, chosen_response_ids), dim=-1)\n chosen_attention_mask = torch.ones_like(chosen_input_ids)\n\n rejected_input_ids = torch.cat((prompt_ids, rejected_response_ids), dim=-1)\n rejected_attention_mask = torch.ones_like(rejected_input_ids)\n\n chosen_input_ids, chosen_attention_mask = self._pad_to_length(chosen_input_ids, chosen_attention_mask)\n rejected_input_ids, rejected_attention_mask = self._pad_to_length(rejected_input_ids, rejected_attention_mask)\n\n input_ids = torch.stack((chosen_input_ids, rejected_input_ids), dim=0)\n attention_mask = torch.stack((chosen_attention_mask, rejected_attention_mask), dim=0)\n\n return {\n \"input_ids\": input_ids,\n \"attention_mask\": attention_mask,\n }\n"}73{"file_name": "verl__utils__dataset__sft_dataset.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nSFT dataset\n- We assume user pass a single parquet file.\n- We load all the data into the memory.\nEach parquet file contains\n\"\"\"\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom omegaconf.listconfig import ListConfig\nfrom torch.utils.data import Dataset\nfrom transformers import PreTrainedTokenizer\n\nfrom verl.utils import hf_tokenizer\nfrom verl.utils.fs import copy_to_local\nfrom verl.utils.model import compute_position_id_with_mask\n\n\nclass SFTDataset(Dataset):\n \"\"\"\n This is an in-memory SFTDataset\n\n Arguments:\n config (OmegaConf): the data config\n \"\"\"\n\n def __init__(self, parquet_files: str | ListConfig, tokenizer, config, max_samples: int = -1):\n prompt_key = config.get(\"prompt_key\", \"prompt\")\n prompt_dict_keys = config.get(\"prompt_dict_keys\", None)\n response_key = config.get(\"response_key\", \"response\")\n response_dict_keys = config.get(\"response_dict_keys\", None)\n max_length = config.get(\"max_length\", 1024)\n truncation = config.get(\"truncation\", \"error\")\n use_shm = config.get(\"use_shm\", False)\n self.shuffle = config.get(\"shuffle\", False)\n self.seed = config.get(\"seed\")\n self.apply_chat_template_kwargs = config.get(\"apply_chat_template_kwargs\", {})\n\n assert truncation in [\"error\", \"left\", \"right\"]\n self.truncation = truncation\n self.use_shm = use_shm\n\n if not isinstance(parquet_files, ListConfig):\n parquet_files = [parquet_files]\n\n self.parquet_files = parquet_files\n self.max_samples = max_samples\n if isinstance(tokenizer, str):\n tokenizer = hf_tokenizer(tokenizer)\n self.tokenizer: PreTrainedTokenizer = tokenizer\n\n self.prompt_key = prompt_key if isinstance(prompt_key, tuple | list) else [prompt_key]\n self.response_key = response_key if isinstance(response_key, tuple | list) else [response_key]\n self.prompt_dict_keys = prompt_dict_keys if prompt_dict_keys else []\n self.response_dict_keys = response_dict_keys if response_dict_keys else []\n\n self.max_length = max_length\n\n self._download()\n self._read_files_and_tokenize()\n\n def _download(self):\n for i, parquet_file in enumerate(self.parquet_files):\n self.parquet_files[i] = copy_to_local(parquet_file, verbose=True, use_shm=self.use_shm)\n\n def _read_files_and_tokenize(self):\n def series_to_item(ls):\n import numpy\n import pandas\n\n while isinstance(ls, pandas.core.series.Series | numpy.ndarray) and len(ls) == 1:\n ls = ls[0]\n return ls\n\n dataframes = []\n for parquet_file in self.parquet_files:\n # read parquet files and cache\n dataframe = pd.read_parquet(parquet_file)\n dataframes.append(dataframe)\n self.dataframe = pd.concat(dataframes)\n\n total = len(self.dataframe)\n print(f\"dataset len: {len(self.dataframe)}\")\n\n if self.max_samples > 0 and self.max_samples < total:\n if self.shuffle:\n rngs_args = (self.seed,) if self.seed is not None else ()\n rng = np.random.default_rng(*rngs_args)\n indices = rng.choice(total, size=self.max_samples, replace=False)\n else:\n indices = np.arange(self.max_samples)\n self.dataframe = self.dataframe.iloc[indices.tolist()]\n print(f\"selected {self.max_samples} random samples out of {total}\")\n\n self.prompts = self.dataframe[self.prompt_key]\n for key in self.prompt_dict_keys:\n # type(x): pandas.core.series.Series\n # type(x[0]): numpy.ndarray\n # type(x[0][0]): dict\n try:\n self.prompts = self.prompts.apply(lambda x: series_to_item(x)[key], axis=1) # noqa: B023\n except Exception:\n print(f\"self.prompts={self.prompts}\")\n raise\n if isinstance(self.prompts, pd.DataFrame):\n self.prompts = self.prompts.squeeze()\n self.prompts = self.prompts.tolist()\n self.responses = self.dataframe[self.response_key]\n for key in self.response_dict_keys:\n try:\n self.responses = self.responses.apply(lambda x: series_to_item(x)[key], axis=1) # noqa: B023\n except Exception:\n print(f\"self.responses={self.responses}\")\n raise\n if isinstance(self.responses, pd.DataFrame):\n self.responses = self.responses.squeeze()\n self.responses = self.responses.tolist()\n\n def __len__(self):\n return len(self.prompts)\n\n def __getitem__(self, item):\n tokenizer = self.tokenizer\n\n prompt = self.prompts[item]\n response = self.responses[item]\n\n # apply chat template\n prompt_chat = [{\"role\": \"user\", \"content\": prompt}]\n\n # string\n prompt_chat_str = tokenizer.apply_chat_template(\n prompt_chat, add_generation_prompt=True, tokenize=False, **self.apply_chat_template_kwargs\n )\n response_chat_str = response + tokenizer.eos_token\n\n # tokenize\n prompt_ids_output = tokenizer(prompt_chat_str, return_tensors=\"pt\", add_special_tokens=False)\n prompt_ids = prompt_ids_output[\"input_ids\"][0]\n prompt_attention_mask = prompt_ids_output[\"attention_mask\"][0]\n\n response_ids_output = tokenizer(response_chat_str, return_tensors=\"pt\", add_special_tokens=False)\n response_ids = response_ids_output[\"input_ids\"][0]\n response_attention_mask = response_ids_output[\"attention_mask\"][0]\n\n prompt_length = prompt_ids.shape[0]\n response_length = response_ids.shape[0]\n\n input_ids = torch.cat((prompt_ids, response_ids), dim=-1)\n attention_mask = torch.cat((prompt_attention_mask, response_attention_mask), dim=-1)\n\n # padding to max length\n sequence_length = input_ids.shape[0]\n if sequence_length < self.max_length:\n padded_input_ids = (\n torch.ones(size=(self.max_length - sequence_length,), dtype=input_ids.dtype)\n * self.tokenizer.pad_token_id\n )\n padded_attention_mask = torch.zeros(size=(self.max_length - sequence_length,), dtype=attention_mask.dtype)\n\n input_ids = torch.cat((input_ids, padded_input_ids))\n attention_mask = torch.cat((attention_mask, padded_attention_mask))\n elif sequence_length > self.max_length:\n if self.truncation == \"left\":\n # actually, left truncation may not be reasonable\n input_ids = input_ids[-self.max_length :]\n attention_mask = attention_mask[-self.max_length :]\n elif self.truncation == \"right\":\n input_ids = input_ids[: self.max_length]\n attention_mask = attention_mask[: self.max_length]\n elif self.truncation == \"error\":\n raise NotImplementedError(f\"{sequence_length=} is larger than {self.max_length=}\")\n else:\n raise NotImplementedError(f\"Unknown truncation method {self.truncation}\")\n\n position_ids = compute_position_id_with_mask(attention_mask)\n\n loss_mask = attention_mask.clone()\n if prompt_length > 1:\n # mask out prompt for SFT.\n loss_mask[: min(prompt_length, loss_mask.size(0)) - 1] = 0\n # mask out the last token in response\n loss_mask[min(prompt_length + response_length, loss_mask.size(0)) - 1] = 0\n\n return {\n \"input_ids\": input_ids,\n \"attention_mask\": attention_mask,\n \"position_ids\": position_ids,\n \"loss_mask\": loss_mask,\n }\n"}74{"file_name": "verl__utils__dataset__vision_utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom io import BytesIO\nfrom typing import Optional\n\nimport torch\nfrom PIL import Image\n\n\ndef process_image(image: dict | Image.Image, image_patch_size: int = 14) -> Image.Image:\n from qwen_vl_utils import fetch_image\n\n if isinstance(image, Image.Image):\n return image.convert(\"RGB\")\n\n if \"bytes\" in image:\n assert \"image\" not in image, \"Cannot have both `bytes` and `image`\"\n image[\"image\"] = Image.open(BytesIO(image[\"bytes\"]))\n\n try:\n ans = fetch_image(image, image_patch_size=image_patch_size)\n except Exception:\n ans = fetch_image(image)\n return ans\n\n\nVIDEO_FORMAT_HELP = \"\"\"Currently, we only support the video formats introduced in qwen2-vl.\nRefer to https://github.com/QwenLM/Qwen2.5-VL?tab=readme-ov-file#using---transformers-to-chat.\n\neg.\n{\n \"type\": \"video\",\n \"video\": [\n \"file:///path/to/frame1.jpg\",\n \"file:///path/to/frame2.jpg\"\n ]\n}\n\n{\n \"type\": \"video\",\n \"video\": \"file:///path/to/video.mp4\"\n}\n# Defaults to fps=2, min_frames=4, max_frames=768\n\n{\n \"type\": \"video\",\n \"video\": \"file:///path/to/video.mp4\",\n \"fps\": 2,\n \"min_frames\": 1,\n \"max_frames\": 32\n}\n\"\"\"\n\n\ndef process_video(\n video: dict,\n image_patch_size: int = 14,\n nframes: Optional[int] = None,\n fps: Optional[float] = None,\n fps_min_frames: Optional[int] = None,\n fps_max_frames: Optional[int] = None,\n return_video_sample_fps: bool = False,\n return_video_metadata: bool = False,\n) -> torch.Tensor:\n \"\"\"Converts a video dict into a [n_frames, 3, H, W] tensor\n\n Add video sample FPS in a future MR\n \"\"\"\n from qwen_vl_utils import fetch_video\n\n if not isinstance(video, dict) or \"video\" not in video:\n raise NotImplementedError(VIDEO_FORMAT_HELP)\n assert nframes is None or fps is None, \"Can't use both `nframes` or `fps`\"\n\n # Shallow copy... since we might want to add some keys\n video = dict(video)\n\n contains_sampling_rules = \"nframes\" in video or \"fps\" in video\n if not contains_sampling_rules:\n if nframes is not None:\n video[\"nframes\"] = nframes\n elif fps is not None:\n video[\"fps\"] = fps\n if fps_min_frames is not None:\n video[\"min_frames\"] = fps_min_frames\n if fps_max_frames is not None:\n video[\"max_frames\"] = fps_max_frames\n\n return fetch_video(\n video,\n image_patch_size=image_patch_size,\n return_video_sample_fps=return_video_sample_fps,\n return_video_metadata=return_video_metadata,\n )\n\n\ndef process_multi_modal_inputs_for_minicpmo(input_ids, attention_mask, position_ids, cu_seqlens, multi_modal_inputs):\n # Adjust image bounds based on left padding and cumulative sequence lengths\n # This is necessary for MiniCPM-o's vision-language alignment\n left_padding_length = torch.argmax(attention_mask, dim=1)\n image_bounds = []\n for i in range(len(multi_modal_inputs[\"image_bound\"])):\n image_bound = (\n multi_modal_inputs[\"image_bound\"][i].to(left_padding_length.device) - left_padding_length[i] + cu_seqlens[i]\n )\n image_bounds.append(image_bound)\n\n # Flatten pixel values list for MiniCPM-o processing\n pixel_values = []\n for i in range(len(multi_modal_inputs[\"pixel_values\"])):\n pixel_values.extend([p for p in multi_modal_inputs[\"pixel_values\"][i]])\n\n multi_modal_inputs[\"pixel_values\"] = [pixel_values]\n multi_modal_inputs[\"image_bound\"] = [torch.vstack(image_bounds)]\n multi_modal_inputs[\"tgt_sizes\"] = [torch.vstack(multi_modal_inputs[\"tgt_sizes\"])]\n multi_modal_inputs[\"input_ids\"] = input_ids\n multi_modal_inputs[\"attention_mask\"] = attention_mask\n multi_modal_inputs[\"position_ids\"] = position_ids\n return {\"data\": multi_modal_inputs}\n"}75{"file_name": "verl__utils__debug__metrics.py", "text": "# Copyright 2025 Individual Contributor: TomQunChaoA\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\n\nimport torch\n\nfrom verl.protocol import DataProto\n\nlogger = logging.getLogger(__file__)\n\n\ndef calculate_token_list_diff(tensor1: torch.Tensor, tensor2: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:\n # verify inputs\n if tensor1.numel() == 0 or tensor2.numel() == 0:\n return torch.zeros(tensor1.shape[0], dtype=torch.long, device=tensor1.device)\n if tensor1.shape != tensor2.shape or mask.shape != tensor1.shape or mask.shape != tensor2.shape:\n print(\n f\"<WARN> dim of tensor1, tensor2, mask is not equal, {(tensor1.shape)=},{(tensor2.shape)=}, {(mask.shape)=}\"\n )\n return torch.ones_like(tensor1)\n # transfer to same device\n if tensor2.device != tensor1.device:\n tensor2 = tensor2.to(tensor1.device)\n if mask.device != tensor1.device:\n mask = mask.to(tensor1.device)\n\n # calculate diff\n diff_mask = tensor1 != tensor2\n\n valid_diff_mask = diff_mask & (mask == 1)\n\n diff_counts = valid_diff_mask.sum(dim=1)\n\n return diff_counts\n\n\ndef pearson_correlation_coefficient(tensor1: torch.Tensor, tensor2: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:\n # implemention of https://arxiv.org/pdf/2506.13585\n if tensor1.shape != tensor2.shape or mask.shape != tensor1.shape or mask.shape != tensor2.shape:\n return 0\n mt1 = torch.masked_select(tensor1, mask)\n mt2 = torch.masked_select(tensor2, mask)\n result = torch.corrcoef(torch.stack([mt1, mt2], dim=0))\n return result[0][1].detach().item()\n\n\ndef calculate_log_prob_diff(log_probs1: torch.Tensor, log_probs2: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:\n full_diff = torch.abs(log_probs1 - log_probs2)\n return torch.masked_select(full_diff, mask)\n\n\ndef calculate_debug_metrics(data: DataProto) -> dict:\n \"\"\"\n calculate rollout vs actor logprobs diff, for debugging purpose\n\n Args:\n data: DataProto\n the data batch to calculate\n rollout_log_probs: log_probs record when rollout forward tokens\n old_log_probs(actor log probs): log_probs record when actor forward tokens\n loss_mask or attention_mask: to mask unrelated token\n responses: the response tokens, for calculating size\n Returns:\n dict: metrics\n \"training/rollout_probs_diff_valid\": 1->input is valid, 0->input is invalid\n \"training/rollout_probs_diff_max\": max value of logprob diff of rollout vs. actor\n \"training/rollout_probs_diff_mean\": mean value of logprob diff of rollout vs. actor\n \"training/rollout_probs_diff_std\": std value of logprob diff of rollout vs. actor\n \"training/rollout_actor_probs_pearson_corr\": logprob's pearson corrcoef of rollout vs. actor, reference to https://arxiv.org/pdf/2506.13585\n \"\"\"\n\n rollout_old_log_probs = data.batch[\"rollout_log_probs\"]\n actor_old_log_probs = data.batch[\"old_log_probs\"]\n if \"response_mask\" in data.batch:\n logger.debug(\"response mask found, use it to mask log probs\")\n log_prob_mask = data.batch[\"response_mask\"]\n elif \"attention_mask\" in data.batch:\n log_prob_mask = data.batch[\"attention_mask\"]\n else:\n logger.warning(f\"no mask info found, use all log probs, {(data.batch.keys())=}\")\n log_prob_mask = torch.ones_like(rollout_old_log_probs)\n responses = data.batch[\"responses\"]\n response_length = responses.size(1)\n\n response_mask = log_prob_mask[:, -response_length:]\n # calculate pearson corrcoef\n actor_probs = torch.exp(actor_old_log_probs)\n rollout_probs = torch.exp(rollout_old_log_probs)\n response_mask_bool = response_mask.bool()\n pearson_corrcoef = pearson_correlation_coefficient(actor_probs, rollout_probs, response_mask_bool)\n rollout_probs_diff = calculate_log_prob_diff(actor_probs, rollout_probs, response_mask_bool)\n return {\n \"training/rollout_probs_diff_valid\": 1,\n \"training/rollout_probs_diff_max\": torch.max(rollout_probs_diff).detach().item(),\n \"training/rollout_probs_diff_mean\": torch.mean(rollout_probs_diff).detach().item(),\n \"training/rollout_probs_diff_std\": torch.std(rollout_probs_diff).detach().item(),\n \"training/rollout_actor_probs_pearson_corr\": pearson_corrcoef,\n }\n"}76{"file_name": "verl__utils__debug__trajectory_tracker.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nTrajectory tracker can be inserted into code to save the intermediate results.\nThe results will be dump to hdfs for offline comparison.\nEach process will have a client that first move all the tensors to CPU\n\"\"\"\n\nimport io\nimport os\nimport tempfile\nfrom collections import deque\n\nimport ray\nimport torch\n\nfrom verl.utils.hdfs_io import copy, makedirs\n\nremote_copy = ray.remote(copy)\n\n\n@ray.remote\ndef save_to_hdfs(data: io.BytesIO, name, hdfs_dir, verbose):\n filename = name + \".pth\"\n with tempfile.TemporaryDirectory() as tmpdirname:\n local_filepath = os.path.join(tmpdirname, filename)\n with open(local_filepath, \"wb\") as f:\n f.write(data.getbuffer())\n # upload to hdfs\n\n if verbose:\n print(f\"Saving {local_filepath} to {hdfs_dir}\")\n try:\n copy(local_filepath, hdfs_dir)\n except Exception as e:\n print(e)\n\n\n@ray.remote\nclass TrajectoryTracker:\n def __init__(self, hdfs_dir, verbose) -> None:\n self.hdfs_dir = hdfs_dir\n makedirs(hdfs_dir)\n self.verbose = verbose\n\n self.handle = deque()\n\n def dump(self, data: io.BytesIO, name):\n # get a temp file and write to it\n self.handle.append(save_to_hdfs.remote(data, name, self.hdfs_dir, self.verbose))\n\n def wait_for_hdfs(self):\n while len(self.handle) != 0:\n future = self.handle.popleft()\n ray.get(future)\n\n\ndef dump_data(data, name):\n enable = os.getenv(\"VERL_ENABLE_TRACKER\", \"0\") == \"1\"\n if not enable:\n return\n buffer = io.BytesIO()\n torch.save(data, buffer)\n tracker = get_trajectory_tracker()\n ray.get(tracker.dump.remote(buffer, name))\n\n\ndef get_trajectory_tracker():\n hdfs_dir = os.getenv(\"VERL_TRACKER_HDFS_DIR\", default=None)\n verbose = os.getenv(\"VERL_TRACKER_VERBOSE\", default=\"0\") == \"1\"\n assert hdfs_dir is not None\n tracker = TrajectoryTracker.options(name=\"global_tracker\", get_if_exists=True, lifetime=\"detached\").remote(\n hdfs_dir, verbose\n )\n return tracker\n\n\nif __name__ == \"__main__\":\n # testing\n os.environ[\"VERL_ENABLE_TRACKER\"] = \"1\"\n os.environ[\"VERL_TRACKER_HDFS_DIR\"] = \"~/debug/test\"\n\n @ray.remote\n def process(iter):\n data = {\"obs\": torch.randn(10, 20)}\n dump_data(data, f\"process_{iter}_obs\")\n\n ray.init()\n\n output_lst = []\n\n for i in range(10):\n output_lst.append(process.remote(i))\n\n out = ray.get(output_lst)\n\n tracker = get_trajectory_tracker()\n ray.get(tracker.wait_for_hdfs.remote())\n"}77{"file_name": "verl__utils__device.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# This code is inspired by the torchtune.\n# https://github.com/pytorch/torchtune/blob/main/torchtune/utils/_device.py\n#\n# Copyright (c) Meta Platforms, Inc. and affiliates.\n# All rights reserved.\n#\n# This source code is licensed under the BSD-style license in https://github.com/pytorch/torchtune/blob/main/LICENSE\n\nimport logging\nimport os\nimport platform\nimport subprocess\n\nimport torch\nfrom packaging import version\n\nlogger = logging.getLogger(__name__)\n\n\ndef is_torch_npu_available(check_device=True) -> bool:\n \"\"\"Check if Ascend NPU is available for PyTorch operations.\n\n Attempts to detect NPU availability by checking for the torch.npu module\n and its is_available() function.\n\n Args:\n check_device : only check torch_npu package or strictly check if NPU device is available\n\n Returns:\n bool: True if NPU is available, False otherwise.\n \"\"\"\n try:\n if not hasattr(torch, \"npu\"):\n return False\n\n if check_device:\n return torch.npu.is_available()\n else:\n return True\n except ImportError:\n return False\n\n\nis_cuda_available = torch.cuda.is_available()\nis_npu_available = is_torch_npu_available()\n\n\ndef get_resource_name() -> str:\n \"\"\"Function that return ray resource name based on the device type.\n Returns:\n ray resource name string, either \"GPU\" or \"NPU\".\n \"\"\"\n return \"GPU\" if is_cuda_available else \"NPU\"\n\n\ndef get_visible_devices_keyword() -> str:\n \"\"\"Get the environment variable name for visible device selection.\n\n Returns the appropriate environment variable name based on the available\n accelerator type (CUDA or Ascend NPU).\n\n Returns:\n str: 'CUDA_VISIBLE_DEVICES' if CUDA is available,\n 'ASCEND_RT_VISIBLE_DEVICES' otherwise.\n \"\"\"\n return \"CUDA_VISIBLE_DEVICES\" if not is_torch_npu_available(check_device=False) else \"ASCEND_RT_VISIBLE_DEVICES\"\n\n\ndef get_device_name() -> str:\n \"\"\"Get the device type string based on available accelerators.\n\n Detects the available accelerator and returns the corresponding PyTorch\n device type string. Currently supports CUDA, Ascend NPU, and CPU.\n\n Returns:\n str: Device type string ('cuda', 'npu', or 'cpu').\n \"\"\"\n if is_cuda_available:\n device = \"cuda\"\n elif is_npu_available:\n device = \"npu\"\n else:\n device = \"cpu\"\n return device\n\n\ndef get_torch_device():\n \"\"\"Get the PyTorch device module for the current accelerator.\n\n Returns the torch device namespace (e.g., torch.cuda, torch.npu) based on\n the detected accelerator type. Falls back to torch.cuda if the namespace\n is not found.\n\n Returns:\n module: The PyTorch device module (torch.cuda, torch.npu, etc.).\n \"\"\"\n device_name = get_device_name()\n try:\n return getattr(torch, device_name)\n except AttributeError:\n logger.warning(f\"Device namespace '{device_name}' not found in torch, try to load torch.cuda.\")\n return torch.cuda\n\n\ndef get_device_id() -> int:\n \"\"\"Get the index of the current accelerator device.\n\n Returns:\n int: The current device index (e.g., 0 for 'cuda:0').\n \"\"\"\n return get_torch_device().current_device()\n\n\ndef get_nccl_backend() -> str:\n \"\"\"Get the distributed communication backend based on device type.\n\n Returns the appropriate collective communication backend for the\n detected accelerator (HCCL for Ascend NPU, NCCL for CUDA).\n\n Returns:\n str: Backend name ('hccl' for NPU, 'nccl' for CUDA/default).\n \"\"\"\n if is_npu_available:\n return \"hccl\"\n else:\n # default to nccl\n return \"nccl\"\n\n\ndef set_expandable_segments(enable: bool) -> None:\n \"\"\"Configure CUDA memory allocator expandable segments setting.\n\n Expandable segments can help avoid out-of-memory (OOM) errors by allowing\n the memory allocator to expand existing memory segments rather than\n allocating new ones.\n\n Args:\n enable: If True, enable expandable segments. If False, disable them.\n\n Note:\n This function only has an effect when CUDA is available.\n \"\"\"\n if is_cuda_available:\n torch.cuda.memory._set_allocator_settings(f\"expandable_segments:{enable}\")\n\n\ndef auto_set_device(config) -> None:\n \"\"\"Automatically configure device name for different accelerators.\n\n For example, on Ascend NPU, this function defaults the trainer device to \"npu\"\n unless explicitly set to \"cpu\".\n\n Args:\n config: Configuration object with trainer.device attribute.\n \"\"\"\n if config and hasattr(config, \"trainer\") and hasattr(config.trainer, \"device\"):\n if is_torch_npu_available():\n if config.trainer.device not in [\"cpu\", \"npu\"]:\n logger.warning(\n f\"Detect setting config.trainer.device to {config.trainer.device} for Ascend NPU, maybe\"\n f\"from default value in config file, automatically set to `npu` instead.\"\n )\n\n config.trainer.device = \"npu\"\n # Other cases: set device to \"cuda\" via config file, no need to change.\n\n\ndef get_device_capability(device_id: int = 0) -> tuple[int | None, int | None]:\n \"\"\"Get the compute capability of a CUDA device.\n\n Args:\n device_id: The CUDA device index to query. Defaults to 0.\n\n Returns:\n tuple: A tuple of (major, minor) compute capability version,\n or (None, None) if CUDA is not available.\n \"\"\"\n major, minor = None, None\n if is_cuda_available:\n major, minor = torch.cuda.get_device_capability(device_id)\n\n return major, minor\n\n\ndef get_npu_versions() -> tuple[str, str]:\n \"\"\"Get the software version and CANN toolkit version for NPU devices.\n\n Returns:\n tuple[str, str]: A tuple of (software_version, cann_version)\n\n Raises:\n RuntimeError: If unable to retrieve version information\n \"\"\"\n # Check npu-smi software version\n result = subprocess.run([\"npu-smi\", \"info\", \"-t\", \"board\", \"-i\", \"1\"], capture_output=True, text=True, check=True)\n\n # Parse software version from output\n software_version = None\n for line in result.stdout.split(\"\\n\"):\n if \"Software Version\" in line:\n # Extract version from line like: \"Software Version : 25.3.rc1.2\"\n parts = line.split(\":\")\n if len(parts) > 1:\n software_version = parts[1].strip().lower()\n break\n\n if not software_version:\n raise RuntimeError(\"Could not find Software Version in npu-smi output\")\n\n # Check CANN toolkit version\n arch = platform.machine()\n if arch not in [\"arm64\", \"aarch64\", \"x86_64\"]:\n raise RuntimeError(f\"Unsupported architecture: {arch}\")\n\n ascend_home = os.environ.get(\"ASCEND_HOME_PATH\", \"/usr/local/Ascend/ascend-toolkit/latest\")\n cann_path = os.path.join(ascend_home, f\"{arch}-linux\")\n\n if not os.path.exists(cann_path):\n raise RuntimeError(f\"CANN toolkit path does not exist: {cann_path}\")\n\n info_file = os.path.join(cann_path, \"ascend_toolkit_install.info\")\n if not os.path.exists(info_file):\n raise RuntimeError(f\"CANN toolkit info file does not exist: {info_file}\")\n\n # Parse version from info file\n cann_version = None\n with open(info_file) as f:\n for line in f:\n if line.startswith(\"version=\"):\n cann_version = line.split(\"=\", 1)[1].strip().lower()\n break\n\n if not cann_version:\n raise RuntimeError(\"Could not find version in CANN toolkit info file\")\n\n return software_version, cann_version\n\n\ndef check_ipc_version_support(software_version: str, cann_version: str) -> bool:\n \"\"\"Check if the given software and CANN versions support IPC.\n\n Compares the software version and CANN toolkit version against minimum\n required versions for IPC support:\n - Software Version should be >= 25.3.rc1\n - CANN version should be >= 8.3.rc1\n\n Args:\n software_version: The software version string (e.g., \"25.5.0\", \"25.3.rc1.2\", \"25.5.t3.b001\")\n cann_version: The CANN toolkit version string (e.g., \"8.3.0\", \"8.3.rc1\")\n\n Returns:\n bool: True if IPC is supported, False otherwise.\n\n Raises:\n RuntimeError: If version format is invalid\n \"\"\"\n # For software_version like \"25.3.rc1.2\", \"25.5.0\", or \"25.5.t3.b001\",\n # we need to extract the base version\n # Use regex to extract version with the following rules:\n # - Standard version: 25.5.0 -> 25.5.0\n # - RC version: 25.3.rc1.2 -> 25.3.rc1\n # - t suffix version: 25.5.t3.b001 -> 25.5 (only first 2 parts if third part is lowercase t)\n # - RC version: 25.3.rc1 -> 25.3.rc1\n # For versions with more than 3 parts (e.g., 25.3.rc1.2), only match the first 3 parts\n import re\n\n # Match version with optional rc part or lowercase t suffix:\n # - If version has lowercase t (e.g., 25.5.t3.b001), only match first 2 parts\n # - Otherwise, match up to 3 parts (e.g., 25.5.0, 25.3.rc1.2)\n ascend_version_pattern = r\"(\\d+\\.\\d+(?=\\.t))|(\\d+\\.\\d+(?:\\.(?:rc\\d+|\\d+))?)\"\n software_match = re.match(ascend_version_pattern, software_version)\n if not software_match:\n raise RuntimeError(f\"Invalid software version format: {software_version}\")\n\n # Select the matched group (either first 2 parts or up to 3 parts)\n software_base = software_match.group(1) if software_match.group(1) else software_match.group(2)\n\n cann_match = re.match(ascend_version_pattern, cann_version)\n if not cann_match:\n raise RuntimeError(f\"Invalid CANN version format: {cann_version}\")\n else:\n # Select the matched group (either first 2 parts or up to 3 parts)\n cann_base = cann_match.group(1) if cann_match.group(1) else cann_match.group(2)\n\n if version.parse(software_base) >= version.parse(\"25.3.rc1\"):\n if version.parse(cann_base) >= version.parse(\"8.3.rc1\"):\n return True\n else:\n logger.info(f\"CANN version {cann_version} is below 8.3.RC1\")\n else:\n logger.info(f\"Software version {software_version} is below 25.3.rc1\")\n\n return False\n\n\ndef is_support_ipc() -> bool:\n \"\"\"Check if the device supports IPC (Inter-Process Communication).\n\n For GPU devices, always returns True.\n For NPU devices, checks the software version and CANN toolkit version\n to determine if IPC is supported.\n\n Returns:\n bool: True if IPC is supported, False otherwise.\n \"\"\"\n # If CUDA is available, it's a GPU device\n if is_cuda_available:\n return True\n\n # For NPU devices, check the software version and CANN toolkit version\n if is_npu_available:\n try:\n software_version, cann_version = get_npu_versions()\n return check_ipc_version_support(software_version, cann_version)\n\n except subprocess.CalledProcessError as e:\n raise RuntimeError(f\"Failed to execute npu-smi command: {e}\") from e\n except Exception as e:\n raise RuntimeError(f\"Error checking IPC support: {e}\") from e\n\n # For other devices (CPU), return False\n return False\n"}78{"file_name": "verl__utils__flops_counter.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport inspect\n\nimport torch\nfrom transformers import PretrainedConfig\n\nfrom verl.utils.device import get_torch_device\n\n_DEVICE_FLOPS = {\n \"CPU\": 448e9,\n \"GB200\": 2.5e15,\n \"B200\": 2.25e15,\n \"MI300X\": 1336e12,\n \"H100\": 989e12,\n \"H800\": 989e12,\n \"H200\": 989e12,\n \"A100\": 312e12,\n \"A800\": 312e12,\n \"L40S\": 362.05e12,\n \"L40\": 181.05e12,\n \"A40\": 149.7e12,\n \"L20\": 119.5e12,\n \"H20\": 148e12,\n \"910B\": 354e12,\n \"Ascend910\": 354e12,\n \"RTX 3070 Ti\": 21.75e12,\n}\n\n\ndef get_device_flops(unit=\"T\", device_name=None):\n \"\"\"Get the theoretical FLOPS (Floating Point Operations Per Second) capacity of the current device.\n\n Args:\n unit (str): The unit to return the FLOPS in. Supported values are:\n \"B\" - Billion (1e9)\n \"K\" - Thousand (1e3)\n \"M\" - Million (1e6)\n \"G\" - Giga (1e9)\n \"T\" - Tera (1e12, default)\n \"P\" - Peta (1e15)\n\n Returns:\n float: The theoretical FLOPS capacity of the current device in the specified unit.\n Returns float('inf') for unknown GPU types.\n \"\"\"\n\n def unit_convert(number, level):\n units = [\"B\", \"K\", \"M\", \"G\", \"T\", \"P\"]\n if number <= 0:\n return number\n ptr = 0\n while ptr < len(units) and units[ptr] != level:\n number /= 1000\n ptr += 1\n return number\n\n # pass device_name is for testing purpose only\n if device_name is None:\n device = get_torch_device()\n if device == torch.cpu:\n device_name = \"CPU\"\n else:\n device_name = get_torch_device().get_device_name()\n\n flops = float(\"inf\") # INF flops for unkown gpu type\n\n for key, value in sorted(_DEVICE_FLOPS.items(), reverse=True):\n if key in device_name:\n flops = value\n break\n flops_unit = unit_convert(flops, unit)\n return flops_unit\n\n\ndef _estimate_qwen2_flops(config, tokens_sum, batch_seqlens, delta_time):\n hidden_size = config.hidden_size\n vocab_size = config.vocab_size\n num_hidden_layers = config.num_hidden_layers\n num_key_value_heads = config.num_key_value_heads\n num_attention_heads = config.num_attention_heads\n intermediate_size = config.intermediate_size\n\n head_dim = getattr(config, \"head_dim\", config.hidden_size // config.num_attention_heads)\n q_size = num_attention_heads * head_dim\n k_size = num_key_value_heads * head_dim\n v_size = num_key_value_heads * head_dim\n\n # non-attn per layer parm\n # Qwen2/LLama use SwiGelu, gate, having up and down linear layer in mlp\n mlp_N = hidden_size * intermediate_size * 3\n attn_linear_N = hidden_size * (q_size + k_size + v_size + num_attention_heads * head_dim)\n emd_and_lm_head_N = vocab_size * hidden_size * 2\n # non-attn all_layer parm\n dense_N = (mlp_N + attn_linear_N) * num_hidden_layers + emd_and_lm_head_N\n # non-attn all_layer & all_token fwd & bwd flops\n dense_N_flops = 6 * dense_N * tokens_sum\n\n # attn all_layer & all_token fwd & bwd flops\n seqlen_square_sum = 0\n for seqlen in batch_seqlens:\n seqlen_square_sum += seqlen * seqlen\n attn_qkv_flops = 6 * seqlen_square_sum * head_dim * num_attention_heads * num_hidden_layers\n\n # all_layer & all_token fwd & bwd flops\n flops_all_token = dense_N_flops + attn_qkv_flops\n flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12\n return flops_achieved\n\n\ndef _estimate_qwen3_vl_flops(config, tokens_sum, batch_seqlens, delta_time, **kargs):\n # qwen3_vl uses text_config and vision_config to distinguish configs of different parts.\n hidden_size = config.text_config.hidden_size\n vocab_size = config.text_config.vocab_size\n num_hidden_layers = config.text_config.num_hidden_layers\n num_key_value_heads = config.text_config.num_key_value_heads\n num_attention_heads = config.text_config.num_attention_heads\n intermediate_size = config.text_config.intermediate_size\n\n head_dim = hidden_size // num_attention_heads\n q_size = num_attention_heads * head_dim\n k_size = num_key_value_heads * head_dim\n v_size = num_key_value_heads * head_dim\n\n # non-attn per layer parm\n mlp_N = hidden_size * intermediate_size * 3\n attn_linear_N = hidden_size * (q_size + k_size + v_size + num_attention_heads * head_dim)\n emd_and_lm_head_N = vocab_size * hidden_size * 2\n # non-attn all_layer parm\n dense_N = (mlp_N + attn_linear_N) * num_hidden_layers + emd_and_lm_head_N\n # non-attn all_layer & all_token fwd & bwd flops\n dense_N_flops = 6 * dense_N * tokens_sum\n\n # qwen3_vl uses deepstack to merge visual embeds and text embeds, but it has no tensor operation.\n\n # attn all_layer & all_token fwd & bwd flops\n seqlen_square_sum = 0\n for seqlen in batch_seqlens:\n seqlen_square_sum += seqlen * seqlen\n attn_qkv_flops = 6 * seqlen_square_sum * head_dim * num_attention_heads * num_hidden_layers\n\n # vit flops\n images_seqlens = kargs.get(\"images_seqlens\", None)\n if images_seqlens is not None:\n vit_flops = _estimate_qwen3_vit_flop(images_seqlens, config.vision_config)\n else:\n vit_flops = 0\n\n # all_layer & all_token fwd & bwd flops\n flops_all_token = dense_N_flops + attn_qkv_flops + vit_flops\n flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12\n return flops_achieved\n\n\ndef _estimate_qwen3_vl_moe_flops(config, tokens_sum, batch_seqlens, delta_time, **kargs):\n # qwen3_vl uses text_config and vision_config to distinguish configs of different parts.\n hidden_size = config.text_config.hidden_size\n vocab_size = config.text_config.vocab_size\n num_hidden_layers = config.text_config.num_hidden_layers\n num_key_value_heads = config.text_config.num_key_value_heads\n num_attention_heads = config.text_config.num_attention_heads\n moe_intermediate_size = config.text_config.moe_intermediate_size\n moe_num_expert = config.text_config.num_experts\n moe_topk = config.text_config.num_experts_per_tok\n\n head_dim = getattr(\n config.text_config, \"head_dim\", config.text_config.hidden_size // config.text_config.num_attention_heads\n )\n q_size = num_attention_heads * head_dim\n k_size = num_key_value_heads * head_dim\n v_size = num_key_value_heads * head_dim\n\n # non-attn per layer parm\n moe_gata_N = hidden_size * moe_num_expert\n # moe has gate_proj, up_proj and down_proj using SwiGLU in ExpertMlp layer & shared experts\n moe_expertmlp_N = hidden_size * moe_intermediate_size * (moe_topk) * 3\n attn_linear_N = hidden_size * (q_size + k_size + v_size + num_attention_heads * head_dim)\n emd_and_lm_head_N = vocab_size * hidden_size * 2\n # non-attn all_layer parm\n moe_N = (moe_gata_N + moe_expertmlp_N + attn_linear_N) * (num_hidden_layers) + emd_and_lm_head_N\n # non-attn all_layer & all_token fwd & bwd flops\n dense_N_flops = 6 * moe_N * tokens_sum\n\n # attn all_layer & all_token fwd & bwd flops\n seqlen_square_sum = 0\n for seqlen in batch_seqlens:\n seqlen_square_sum += seqlen * seqlen\n attn_qkv_flops = 6 * seqlen_square_sum * head_dim * num_attention_heads * num_hidden_layers\n\n # vit flops\n images_seqlens = kargs.get(\"images_seqlens\", None)\n if images_seqlens is not None:\n vit_flops = _estimate_qwen3_vit_flop(images_seqlens, config.vision_config)\n else:\n vit_flops = 0\n\n # all_layer & all_token fwd & bwd flops\n flops_all_token = dense_N_flops + attn_qkv_flops + vit_flops\n flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12\n return flops_achieved\n\n\ndef _estimate_qwen3_vit_flop(images_seqlens, config):\n \"\"\"\n Estimate the FLOPS of the vision encoder for Qwen3-VL\n \"\"\"\n\n if config is None:\n return 0\n tokens_sum = sum(images_seqlens)\n\n num_heads = config.num_heads\n depth = config.depth\n\n dim = config.hidden_size\n mlp_hidden_dim = config.intermediate_size\n out_hidden_size = config.out_hidden_size\n\n spatial_merge_size = config.spatial_merge_size\n\n head_dim = dim // num_heads\n\n # every vision token's patch_embed comes from a conv of (C, T, H, W) -> (dim,)\n patch_embed_N = dim * config.in_channels * config.temporal_patch_size * config.patch_size * config.patch_size\n # Qwen3 VL vision mlp does not use GLU, thus 2.\n mlp_N = dim * mlp_hidden_dim * 2\n attn_linear_N = dim * (4 * dim) # qkv and output proj\n merger_N = (out_hidden_size + (dim * (spatial_merge_size**2))) * (dim * (spatial_merge_size**2))\n\n # Qwen3 VL uses deep stack, one merger for every deepstack layer\n deepstack_merger_N = merger_N * len(config.deepstack_visual_indexes)\n # non-attn all_layer parm\n dense_N = patch_embed_N + (mlp_N + attn_linear_N) * depth + deepstack_merger_N + merger_N\n\n # non-attn all_layer & all_token fwd & bwd flops\n dense_N_flops = 6 * dense_N * tokens_sum\n\n # In Qwen3 VL, full attention is used in all vision layers.\n full_attn_layer_num = depth\n\n # full attn layer & all_token fwd & bwd flops\n seqlen_square_sum = 0\n for seqlen in images_seqlens:\n seqlen_square_sum += seqlen * seqlen\n attn_qkv_flops = 12 * seqlen_square_sum * head_dim * num_heads * full_attn_layer_num\n\n vit_flops = dense_N_flops + attn_qkv_flops\n\n return vit_flops\n\n\ndef _estimate_deepseek_v3_flops(config, tokens_sum, batch_seqlens, delta_time):\n hidden_size = config.hidden_size\n vocab_size = config.vocab_size\n moe_intermediate_size = config.moe_intermediate_size\n num_hidden_layers = config.num_hidden_layers\n first_k_dense_replace = config.first_k_dense_replace\n num_query_heads = config.num_attention_heads\n moe_num_expert = config.n_routed_experts\n\n moe_topk = config.num_experts_per_tok\n share_expert_num = config.n_shared_experts\n\n # non-attn per layer parm\n moe_gata_N = hidden_size * moe_num_expert\n # moe has fc1_1, fc1_2 and fc2 using SwiGLU in ExpertMlp layer & shared experts\n moe_expertmlp_N = hidden_size * moe_intermediate_size * (moe_topk + share_expert_num) * 3\n # MLA attn\n attn_linear_N = 0\n q_head_dim = config.qk_nope_head_dim + config.qk_rope_head_dim\n if config.q_lora_rank is None:\n attn_linear_N += hidden_size * num_query_heads * q_head_dim\n else:\n attn_linear_N += hidden_size * config.q_lora_rank\n attn_linear_N += num_query_heads * q_head_dim * config.q_lora_rank\n\n attn_linear_N += hidden_size * (config.kv_lora_rank + config.qk_rope_head_dim)\n attn_linear_N += num_query_heads * (q_head_dim - config.qk_rope_head_dim + config.v_head_dim) * config.kv_lora_rank\n attn_linear_N += num_query_heads * config.v_head_dim * hidden_size\n emd_and_lm_head_N = vocab_size * hidden_size * 2\n # non-attn all_layer parm\n moe_N = (\n (moe_gata_N + moe_expertmlp_N + attn_linear_N) * (num_hidden_layers - first_k_dense_replace)\n + (hidden_size * config.intermediate_size * 3 + attn_linear_N) * first_k_dense_replace\n + emd_and_lm_head_N\n )\n # non-attn all_layer & all_token fwd & bwd flops\n dense_N_flops = 6 * moe_N * tokens_sum\n\n # attn all_layer & all_token fwd & bwd flops\n seqlen_square_sum = 0\n for seqlen in batch_seqlens:\n seqlen_square_sum += seqlen * seqlen * num_hidden_layers\n\n # Core attention FLOPS for MLA with causal mask:\n # Q @ K^T: 3 * 2 * seq^2 * q_head_dim * num_heads / 2 (causal)\n # attn @ V: 3 * 2 * seq^2 * v_head_dim * num_heads / 2 (causal)\n attn_qkv_flops = 3 * seqlen_square_sum * (q_head_dim + config.v_head_dim) * num_query_heads\n # all_layer & all_token fwd & bwk flops\n flops_all_token = dense_N_flops + attn_qkv_flops\n flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12\n\n return flops_achieved\n\n\ndef _estimate_qwen2_moe_flops(config, tokens_sum, batch_seqlens, delta_time):\n hidden_size = config.hidden_size\n vocab_size = config.vocab_size\n num_hidden_layers = config.num_hidden_layers\n num_key_value_heads = config.num_key_value_heads\n num_attention_heads = config.num_attention_heads\n moe_intermediate_size = config.moe_intermediate_size\n moe_topk = config.num_experts_per_tok\n num_experts = config.num_experts\n\n head_dim = getattr(config, \"head_dim\", config.hidden_size // config.num_attention_heads)\n q_size = num_attention_heads * head_dim\n k_size = num_key_value_heads * head_dim\n v_size = num_key_value_heads * head_dim\n\n # non-attn per layer parm\n # gate + moe export\n moe_mlp_N = hidden_size * moe_topk * moe_intermediate_size * 3 + hidden_size * num_experts\n attn_linear_N = hidden_size * (q_size + k_size + v_size + num_attention_heads * head_dim)\n emd_and_lm_head_N = vocab_size * hidden_size * 2\n # non-attn all_layer parm\n dense_N = (moe_mlp_N + attn_linear_N) * num_hidden_layers + emd_and_lm_head_N\n # non-attn all_layer & all_token fwd & bwd flops\n dense_N_flops = 6 * dense_N * tokens_sum\n\n # attn all_layer & all_token fwd & bwd flops\n seqlen_square_sum = 0\n for seqlen in batch_seqlens:\n seqlen_square_sum += seqlen * seqlen\n attn_qkv_flops = 6 * seqlen_square_sum * head_dim * num_attention_heads * num_hidden_layers\n\n # all_layer & all_token fwd & bwd flops\n flops_all_token = dense_N_flops + attn_qkv_flops\n flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12\n return flops_achieved\n\n\ndef _estimate_gemma3_flops(config, tokens_sum, batch_seqlens, delta_time):\n hidden_size = config.hidden_size\n vocab_size = config.vocab_size\n num_hidden_layers = config.num_hidden_layers\n num_key_value_heads = config.num_key_value_heads\n num_attention_heads = config.num_attention_heads\n intermediate_size = config.intermediate_size\n\n head_dim = getattr(config, \"head_dim\", config.hidden_size // config.num_attention_heads)\n q_size = num_attention_heads * head_dim\n k_size = num_key_value_heads * head_dim\n v_size = num_key_value_heads * head_dim\n\n # non-attn per layer parm\n # Gemma3 uses GeGLU (gelu_pytorch_tanh), having 3 matrices in MLP (inherited from Gemma2MLP)\n mlp_N = hidden_size * intermediate_size * 3\n attn_linear_N = hidden_size * (q_size + k_size + v_size + num_attention_heads * head_dim)\n emd_and_lm_head_N = vocab_size * hidden_size * 2\n # non-attn all_layer parm\n dense_N = (mlp_N + attn_linear_N) * num_hidden_layers + emd_and_lm_head_N\n # non-attn all_layer & all_token fwd & bwd flops\n dense_N_flops = 6 * dense_N * tokens_sum\n\n # attn all_layer & all_token fwd & bwd flops\n # Gemma3 alternates between full and sliding window attention based on layer_types\n seqlen_square_sum = 0\n\n layer_types = getattr(config, \"layer_types\", None)\n sliding_window = getattr(config, \"sliding_window\", 1024) # default 1024\n # default pattern: every 6th layer is full\n sliding_window_pattern = getattr(config, \"sliding_window_pattern\", 6)\n\n # If layer_types is not provided, generate it based on sliding_window_pattern\n if layer_types is None and sliding_window is not None and sliding_window_pattern is not None:\n layer_types = [\n \"sliding_attention\" if bool((i + 1) % sliding_window_pattern) else \"full_attention\"\n for i in range(num_hidden_layers)\n ]\n\n if layer_types:\n # Calculate attention flops per layer based on attention type\n for layer_idx in range(num_hidden_layers):\n is_sliding = False\n if layer_types and layer_idx < len(layer_types):\n is_sliding = layer_types[layer_idx] == \"sliding_attention\"\n\n for seqlen in batch_seqlens:\n if is_sliding and sliding_window:\n # Sliding window limits each token to attend to at most window_size tokens\n effective_seqlen = min(seqlen, sliding_window)\n seqlen_square_sum += seqlen * effective_seqlen\n else:\n # Full attention\n seqlen_square_sum += seqlen * seqlen\n else:\n # If no layer_types config, assume all layers use full attention\n for seqlen in batch_seqlens:\n seqlen_square_sum += seqlen * seqlen\n seqlen_square_sum *= num_hidden_layers\n\n attn_qkv_flops = 6 * seqlen_square_sum * head_dim * num_attention_heads\n\n # all_layer & all_token fwd & bwd flops\n flops_all_token = dense_N_flops + attn_qkv_flops\n flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12\n return flops_achieved\n\n\ndef _estimate_apertus_flops(config, tokens_sum, batch_seqlens, delta_time):\n hidden_size = config.hidden_size\n vocab_size = config.vocab_size\n num_hidden_layers = config.num_hidden_layers\n num_key_value_heads = config.num_key_value_heads\n num_attention_heads = config.num_attention_heads\n intermediate_size = config.intermediate_size\n\n head_dim = getattr(config, \"head_dim\", config.hidden_size // config.num_attention_heads)\n q_size = num_attention_heads * head_dim\n k_size = num_key_value_heads * head_dim\n v_size = num_key_value_heads * head_dim\n\n # Apertus MLP with XIELU activation uses only 2 linear layers (up_proj, down_proj)\n # No gate_proj for XIELU, unlike SwiGLU which has 3 layers\n mlp_N = hidden_size * intermediate_size * 2\n attn_linear_N = hidden_size * (q_size + k_size + v_size + num_attention_heads * head_dim)\n\n # ApertusConfig has qk_norm defaulting to True.\n # This adds params for q_norm (on H) and k_norm (on num_kv_heads * head_dim)\n qk_norm_params_per_layer = hidden_size + num_key_value_heads * head_dim # q_norm + k_norm\n\n emd_and_lm_head_N = vocab_size * hidden_size * 2\n # non-attn all_layer params\n dense_N = (mlp_N + attn_linear_N + qk_norm_params_per_layer) * num_hidden_layers + emd_and_lm_head_N\n # non-attn all_layer & all_token fwd & bwd flops\n dense_N_flops = 6 * dense_N * tokens_sum\n\n # attn all_layer & all_token fwd & bwd flops\n seqlen_square_sum = 0\n for seqlen in batch_seqlens:\n seqlen_square_sum += seqlen * seqlen\n attn_qkv_flops = 6 * seqlen_square_sum * head_dim * num_attention_heads * num_hidden_layers\n\n # all_layer & all_token fwd & bwd flops\n flops_all_token = dense_N_flops + attn_qkv_flops\n flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12\n return flops_achieved\n\n\ndef _estimate_gpt_oss_flops(config, tokens_sum, batch_seqlens, delta_time):\n hidden_size = config.hidden_size\n vocab_size = config.vocab_size\n num_hidden_layers = config.num_hidden_layers\n num_key_value_heads = config.num_key_value_heads\n num_attention_heads = config.num_attention_heads\n\n # MoE params\n moe_intermediate_size = config.intermediate_size\n num_experts = config.num_local_experts\n num_experts_per_tok = config.num_experts_per_tok\n mlp_matrices = 3\n\n # Head dim\n head_dim = getattr(config, \"head_dim\", hidden_size // num_attention_heads)\n q_size = num_attention_heads * head_dim\n k_size = num_key_value_heads * head_dim\n v_size = num_key_value_heads * head_dim\n\n # 1. Attention Block (GQA)\n attn_linear_N = hidden_size * (q_size + k_size + v_size + num_attention_heads * head_dim)\n # 2. MLP / MoE Block\n # Gate network\n moe_gate_N = hidden_size * num_experts\n # Expert forward calculation, Active parameters: mlp_matrices * H * I * num_experts_per_tok\n moe_expert_N = hidden_size * moe_intermediate_size * mlp_matrices * num_experts_per_tok\n\n moe_mlp_N = moe_gate_N + moe_expert_N\n\n emd_and_lm_head_N = vocab_size * hidden_size * 2\n\n # Total non-attn params per layer * layers + embeddings\n # (moe_mlp_N + attn_linear_N) * layers\n dense_N = (moe_mlp_N + attn_linear_N) * num_hidden_layers + emd_and_lm_head_N\n\n # FLOPs for dense part (fwd + bwd = 6 * N)\n dense_N_flops = 6 * dense_N * tokens_sum\n\n # 3. Attention Matrix FLOPs\n seqlen_square_sum = 0\n\n # Handle sliding window attention\n layer_types = getattr(config, \"layer_types\", None)\n sliding_window = getattr(config, \"sliding_window\", 128)\n\n if layer_types:\n for layer_type in layer_types:\n is_sliding = layer_type == \"sliding_attention\"\n\n for seqlen in batch_seqlens:\n if is_sliding and sliding_window:\n # Sliding window limits each token to attend to at most window_size tokens\n effective_seqlen = min(seqlen, sliding_window)\n seqlen_square_sum += seqlen * effective_seqlen\n else:\n # Full attention\n seqlen_square_sum += seqlen * seqlen\n else:\n # Default to full attention for all layers\n for seqlen in batch_seqlens:\n seqlen_square_sum += seqlen * seqlen\n seqlen_square_sum *= num_hidden_layers\n\n attn_qkv_flops = 6 * seqlen_square_sum * head_dim * num_attention_heads\n\n # Total FLOPs\n flops_all_token = dense_N_flops + attn_qkv_flops\n flops_achieved = flops_all_token * (1.0 / delta_time) / 1e12\n return flops_achieved\n\n\ndef _estimate_unknown_flops(config, tokens_sum, batch_seqlens, delta_time):\n return 0\n\n\nESTIMATE_FUNC = {\n \"qwen2\": _estimate_qwen2_flops,\n \"llama\": _estimate_qwen2_flops,\n \"qwen2_moe\": _estimate_qwen2_moe_flops,\n \"qwen2_vl\": _estimate_qwen2_flops,\n \"qwen2_5_vl\": _estimate_qwen2_flops,\n \"qwen3\": _estimate_qwen2_flops,\n \"qwen3_moe\": _estimate_qwen2_moe_flops,\n \"qwen3_vl\": _estimate_qwen3_vl_flops,\n \"qwen3_vl_moe\": _estimate_qwen3_vl_moe_flops,\n \"deepseek_v3\": _estimate_deepseek_v3_flops,\n \"minicpmv\": _estimate_qwen2_flops,\n \"minicpmo\": _estimate_qwen2_flops,\n \"mistral\": _estimate_qwen2_flops,\n \"gemma3_text\": _estimate_gemma3_flops,\n \"seed_oss\": _estimate_qwen2_flops,\n \"apertus\": _estimate_apertus_flops,\n \"glm4v\": _estimate_qwen2_flops,\n \"gpt_oss\": _estimate_gpt_oss_flops,\n \"mimo\": _estimate_qwen2_flops,\n}\n\n\nclass FlopsCounter:\n \"\"\"\n Used to count mfu during training loop\n\n Example:\n flops_counter = FlopsCounter(config)\n flops_achieved, flops_promised = flops_counter.estimate_flops(tokens_list, delta_time)\n\n \"\"\"\n\n def __init__(self, config: PretrainedConfig):\n VALID_CONFIG_TYPE = ESTIMATE_FUNC.keys()\n if config.model_type not in VALID_CONFIG_TYPE:\n print(\n f\"Only support config type of {VALID_CONFIG_TYPE}, but got {config.model_type}. MFU will always be \"\n f\"zero.\"\n )\n\n self.config = config\n\n # TODO: actually we can make this a static method\n def estimate_flops(self, batch_seqlens, delta_time, **kargs):\n \"\"\"\n Estimate the FLOPS based on the number of valid tokens in the current batch and the time taken.\n\n Args:\n batch_seqlens (List[int]): A list where each element represents the number of valid tokens in the\n current batch.\n delta_time (float): The time taken to process the batch, in seconds.\n\n Returns:\n estimated_flops (float): The estimated FLOPS based on the input tokens and time.\n promised_flops (float): The expected FLOPS of the current device.\n \"\"\"\n tokens_sum = sum(batch_seqlens)\n func = ESTIMATE_FUNC.get(self.config.model_type, _estimate_unknown_flops)\n sig = inspect.signature(func)\n if any(p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()):\n estimated_flops = func(self.config, tokens_sum, batch_seqlens, delta_time, **kargs)\n else:\n estimated_flops = func(self.config, tokens_sum, batch_seqlens, delta_time)\n promised_flops = get_device_flops()\n return estimated_flops, promised_flops\n"}79{"file_name": "verl__utils__fsdp_utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport functools\nimport itertools\nimport json\nimport math\nimport os\nfrom abc import ABC\nfrom collections import OrderedDict\nfrom contextlib import contextmanager, nullcontext\nfrom typing import cast\n\nimport torch\nimport torch.distributed as dist\nimport torch.nn as nn\nfrom packaging import version\nfrom torch.distributed import DeviceMesh\nfrom torch.distributed.fsdp import FullyShardedDataParallel as FSDP\nfrom torch.distributed.fsdp._runtime_utils import _lazy_init\nfrom torch.distributed.fsdp.wrap import size_based_auto_wrap_policy, transformer_auto_wrap_policy\nfrom transformers.trainer_pt_utils import get_module_class_from_name\n\nfrom verl.utils.device import get_device_id, get_device_name, get_torch_device\nfrom verl.utils.model import check_exclude_modules, check_target_modules\n\nif version.parse(torch.__version__) >= version.parse(\"2.6\"):\n from torch.distributed.fsdp import CPUOffloadPolicy, FSDPModule, MixedPrecisionPolicy, fully_shard\n from torch.distributed.fsdp._fully_shard._fsdp_init import _get_post_forward_mesh_info\n from torch.distributed.tensor import Shard\n\n fully_shard_module = torch.distributed.fsdp._fully_shard._fully_shard\nelif version.parse(torch.__version__) >= version.parse(\"2.4\"):\n from torch.distributed._composable.fsdp import CPUOffloadPolicy, FSDPModule, MixedPrecisionPolicy, fully_shard\n\n fully_shard_module = torch.distributed._composable.fsdp\nelse:\n fully_shard, MixedPrecisionPolicy, FSDPModule, CPUOffloadPolicy, fully_shard_module = None, None, None, None, None\n\n\ndef init_fn(x: torch.nn.Module):\n if torch.distributed.get_rank() != 0:\n x = x.to_empty(device=get_device_id(), recurse=False)\n get_torch_device().empty_cache()\n return x\n\n\ndef get_init_weight_context_manager(use_meta_tensor=True, mesh: DeviceMesh = None):\n from accelerate import init_empty_weights\n\n cpu_init_weights = lambda: torch.device(\"cpu\")\n if use_meta_tensor:\n if mesh is None:\n init_context = init_empty_weights if torch.distributed.get_rank() != 0 else cpu_init_weights\n else:\n init_context = init_empty_weights if mesh.get_coordinate()[-1] != 0 else cpu_init_weights\n else:\n init_context = cpu_init_weights\n return init_context\n\n\n# Copyright 2020-present the HuggingFace Inc. team.\n# Adapted from https://github.com/huggingface/transformers/src/transformers/trainer.py\ndef get_fsdp_wrap_policy(module, config=None, is_lora=False):\n \"\"\"Get FSDP wrap policy for the module.\n\n Args:\n module: The module to get wrap policy for\n config: Configuration for wrap policy\n is_lora: Whether to enable lambda policy for LoRA modules\n \"\"\"\n if config is None:\n config = {}\n\n # NOTE: This is a temporary workaround to be compatible with the OmegaConf & dataclass. We will remove this\n # once we have make all config in verl from OmegaConf to data class.\n def _get_attr(attr_name, default_value=None):\n if hasattr(config, \"get\"):\n return config.get(attr_name, default_value)\n else:\n return config.__getattribute__(attr_name)\n\n if _get_attr(\"disable\", False):\n return None\n\n default_transformer_cls_names_to_wrap = getattr(module, \"_no_split_modules\", None)\n fsdp_transformer_layer_cls_to_wrap = _get_attr(\n \"transformer_layer_cls_to_wrap\", default_transformer_cls_names_to_wrap\n )\n min_num_params = _get_attr(\"min_num_params\", 0)\n auto_wrap_policy = None\n\n policies = []\n\n from torch.distributed.fsdp.wrap import _or_policy, lambda_auto_wrap_policy\n\n # Add lambda policy for LoRA modules if is_lora is True\n if is_lora:\n\n def lambda_policy_fn(module):\n return bool(\n len(list(module.named_children())) == 0\n and getattr(module, \"weight\", None) is not None\n and module.weight.requires_grad\n )\n\n lambda_policy = functools.partial(lambda_auto_wrap_policy, lambda_fn=lambda_policy_fn)\n policies.append(lambda_policy)\n\n if min_num_params > 0:\n size_policy = functools.partial(size_based_auto_wrap_policy, min_num_params=min_num_params)\n policies.append(size_policy)\n elif fsdp_transformer_layer_cls_to_wrap is not None:\n transformer_cls_to_wrap = set()\n for layer_class in fsdp_transformer_layer_cls_to_wrap:\n transformer_cls = get_module_class_from_name(module, layer_class)\n if transformer_cls is None:\n raise Exception(\"Could not find the transformer layer class to wrap in the model.\")\n else:\n transformer_cls_to_wrap.add(transformer_cls)\n\n transformer_policy = functools.partial(\n transformer_auto_wrap_policy,\n transformer_layer_cls=transformer_cls_to_wrap,\n )\n policies.append(transformer_policy)\n\n if len(policies) > 0:\n auto_wrap_policy = functools.partial(_or_policy, policies=policies)\n\n return auto_wrap_policy\n\n\n@torch.no_grad()\ndef offload_fsdp_model_to_cpu(model: FSDP, empty_cache: bool = True):\n if fsdp_version(model) == 2:\n offload_fsdp2_model_to_cpu(model, empty_cache)\n return\n\n assert isinstance(model, FSDP)\n # lazy init FSDP model\n _lazy_init(model, model)\n assert model._is_root, \"Only support root model offloading to CPU\"\n for handle in model._all_handles:\n if handle._offload_params:\n continue\n flat_param = handle.flat_param\n assert (\n flat_param.data.data_ptr() == flat_param._local_shard.data_ptr()\n and id(flat_param.data) != id(flat_param._local_shard)\n and flat_param.data.size() == flat_param._local_shard.size()\n )\n handle.flat_param_to(torch.device(\"cpu\"), non_blocking=True)\n # the following still keeps id(._local_shard) != id(.data)\n flat_param._local_shard = flat_param.data\n assert id(flat_param._local_shard) != id(flat_param.data)\n if empty_cache:\n get_torch_device().empty_cache()\n\n\n@torch.no_grad()\ndef offload_fsdp2_model_to_cpu(model, empty_cache: bool = True):\n model.cpu()\n if empty_cache:\n get_torch_device().empty_cache()\n\n\n@torch.no_grad()\ndef load_fsdp_model_to_gpu(model: FSDP):\n if fsdp_version(model) == 2:\n load_fsdp2_model_to_gpu(model)\n return\n\n assert isinstance(model, FSDP)\n # lazy init FSDP model\n _lazy_init(model, model)\n assert model._is_root, \"Only support root model loading to GPU\"\n device_id = get_device_id()\n for handle in model._all_handles:\n if handle._offload_params:\n continue\n flat_param = handle.flat_param\n handle.flat_param_to(torch.device(f\"{get_device_name()}:{device_id}\"), non_blocking=True)\n # the following still keeps id(._local_shard) != id(.data)\n flat_param._local_shard = flat_param.data\n\n\n@torch.no_grad()\ndef load_fsdp2_model_to_gpu(model):\n device = get_device_id()\n model.to(device)\n\n\n@torch.no_grad()\ndef offload_fsdp_optimizer(optimizer):\n if not optimizer.state:\n return\n for param_group in optimizer.param_groups:\n for param in param_group[\"params\"]:\n state = optimizer.state[param]\n for key, value in state.items():\n if isinstance(value, torch.Tensor):\n state[key] = value.to(\"cpu\", non_blocking=True)\n\n\n@torch.no_grad()\ndef load_fsdp_optimizer(optimizer, device_id):\n if not optimizer.state:\n return\n for param_group in optimizer.param_groups:\n for param in param_group[\"params\"]:\n state = optimizer.state[param]\n for key, value in state.items():\n if isinstance(value, torch.Tensor):\n state[key] = value.to(device_id, non_blocking=True)\n\n\n@contextmanager\ndef meta_device_init():\n \"\"\"\n Create model parameters with meta device.\n\n Note buffers in model will still be initialized in default device (e.g., CPU),\n since the buffers can be non-persistent and filled with expected values that can\n NOT be captured in meta device.\n \"\"\"\n device = torch.device(\"meta\")\n old_register_parameter = nn.Module.register_parameter\n registered = set()\n\n def register_empty_parameter(module, name, param):\n old_register_parameter(module, name, param)\n # we will skip register shared parameters as it\n # is already registered previously\n if param is not None and param not in registered:\n param_cls = type(module._parameters[name])\n kwargs = module._parameters[name].__dict__\n kwargs[\"requires_grad\"] = param.requires_grad\n module._parameters[name] = param_cls(module._parameters[name].to(device), **kwargs)\n registered.add(module._parameters[name])\n\n try:\n nn.Module.register_parameter = register_empty_parameter\n yield\n finally:\n registered.clear()\n nn.Module.register_parameter = old_register_parameter\n\n\ndef parallel_load_safetensors(filepath):\n \"\"\"\n Parallel load safetensors from huggingface checkpoint\n\n Huggingface checkpoint contains:\n\n - config.json: a json file for model configuration\n - model.safetensor.index.json: a json file for safetensors (parameters & buffers) index\n - model-000x-of-ooxx.safetensors: a binary file for safetensors (parameters & buffers) chunks\n\n Or (when model is small),\n\n - model.safetensors: a binary file for all parameters and buffers\n\n Each rank will own a part of model chunks and load them directly into GPU memory.\n \"\"\"\n from safetensors.torch import load_file\n\n safetensors2param = {}\n\n index_file = os.path.join(filepath, \"model.safetensors.index.json\")\n if os.path.exists(index_file):\n index = json.load(open(index_file, \"rb\"))\n for param_name, filename in index[\"weight_map\"].items():\n safetensors2param.setdefault(filename, []).append(param_name)\n else:\n # in this case, the model is small and we can load it all at once\n param_file = os.path.join(filepath, \"model.safetensors\")\n assert os.path.exists(param_file), f\"Cannot find {param_file}\"\n states = load_file(param_file)\n for param_name in states:\n safetensors2param.setdefault(\"model.safetensors\", []).append(param_name)\n del states\n\n total_files = len(safetensors2param)\n ckpt_chunks = sorted(safetensors2param.keys())\n world_size = dist.get_world_size()\n size = int(math.ceil(total_files / world_size))\n ckpt_chunks = [ckpt_chunks[rank * size : rank * size + size] for rank in range(world_size)]\n\n shard_states = {}\n device = get_device_id()\n for rank, files in enumerate(ckpt_chunks):\n if rank == dist.get_rank():\n for file in files:\n file = os.path.join(filepath, file)\n states = load_file(file, device=device)\n # print(f\"rank {rank} loading {file}...\")\n shard_states.update(states)\n else:\n for file in files:\n for param_name in safetensors2param[file]:\n shard_states[param_name] = rank\n return shard_states\n\n\ndef parallel_init_module_fn(module: torch.nn.Module, shard_states: dict[str, torch.nn.Parameter]):\n \"\"\"\n Generate a function to initialize sub-modules in the `module` with `shard_states`\n from huggingface checkpoint.\n\n Args:\n module (torch.nn.Module): the global module to be initialized\n shard_states (Dict[str, torch.nn.Parameter]): the shard states from huggingface checkpoint\n\n Returns:\n init_fn (Callable): a function to initialize sub-modules in the `module` with `shard_states`\n \"\"\"\n\n state2fqn = {}\n for name, state in itertools.chain(\n module.named_parameters(remove_duplicate=False), module.named_buffers(remove_duplicate=False)\n ):\n state2fqn.setdefault(state, []).append(name)\n # remove standalone parameters and buffers\n shared = {s for s, names in state2fqn.items() if len(names) > 1}\n materialized_states = {}\n\n @torch.no_grad()\n def create_and_sync_state(param_name, state, is_param):\n assert param_name in shard_states, f\"{param_name} not loaded\"\n device = get_device_id()\n if is_param:\n param = torch.nn.Parameter(torch.empty_like(state.data, device=device), requires_grad=state.requires_grad)\n else: # buffer\n param = torch.empty_like(state.data, device=device)\n loaded = shard_states[param_name]\n if isinstance(loaded, torch.nn.Parameter | torch.Tensor):\n # NOTE: loaded.dtype can be different with param.dtype\n param.data.copy_(loaded.data)\n dist.broadcast(param.data, src=dist.get_rank())\n else:\n assert isinstance(loaded, int) # the rank that holds the state\n dist.broadcast(param.data, src=loaded)\n shard_states.pop(param_name)\n del loaded\n return param\n\n def init_fn(sub_mod: torch.nn.Module, recurse: bool = True):\n param_and_buffers = tuple(sub_mod.named_parameters(recurse=False)) + tuple(sub_mod.named_buffers(recurse=False))\n # param_and_buffers = sorted(sub_mod.named_parameters(recurse=False), key=lambda x: x[0])\n for name, state in param_and_buffers:\n if not state.is_meta:\n continue\n is_param = name in sub_mod._parameters\n fqn = state2fqn[state].pop(0)\n # non-persistent buffers will not be saved in state dict, we can safely skip it\n if (not is_param) and fqn not in shard_states:\n if state.is_meta:\n raise RuntimeError(\n f\"find a non-persistent buffer ({fqn}) initiated with device meta. Such buffer is not saved \"\n f\"in checkpoint and user should guarantee to init in CPU / GPU device.\"\n )\n continue\n # for shared parameter, we get it from the first time it is created\n if state in shared:\n if state not in materialized_states:\n materialized_states[state] = create_and_sync_state(fqn, state, is_param)\n else:\n if fqn in shard_states:\n shard_states.pop(fqn)\n materialize_state = materialized_states[state]\n # for not shared parameter, we create it directly\n else:\n materialize_state = create_and_sync_state(fqn, state, is_param)\n if is_param:\n sub_mod._parameters[name] = materialize_state\n else:\n sub_mod._buffers[name] = materialize_state\n if recurse:\n for module in sub_mod.children():\n init_fn(module, recurse=True)\n\n # for debug\n # if len(shard_states) == 0: print(\"clear\")\n return sub_mod\n\n return init_fn\n\n\ndef fsdp_version(model):\n if isinstance(model, FSDP):\n return 1\n elif isinstance(model, FSDPModule):\n return 2\n else:\n return 0\n\n\ndef get_fsdp_state_ctx(model, state_type, state_cfg, optim_cfg):\n if fsdp_version(model) == 1:\n return FSDP.state_dict_type(model, state_type, state_cfg, optim_cfg)\n else:\n return nullcontext()\n\n\ndef get_fsdp_full_state_dict(model: torch.nn.Module, offload_to_cpu: bool = True, rank0_only: bool = True):\n \"\"\"\n Get the full state dict from an FSDP model.\n\n Args:\n model (torch.nn.Module): The FSDP model to get state dict from\n offload_to_cpu (bool, optional): Whether to offload the state dict to CPU. Defaults to True.\n rank0_only (bool, optional): Whether to only get state dict on rank 0. Defaults to True.\n\n Returns:\n dict: The full state dict of the model\n\n Raises:\n NotImplementedError: If the FSDP version is unknown\n \"\"\"\n if fsdp_version(model) == 1:\n from torch.distributed.fsdp import FullStateDictConfig, StateDictType\n\n state_dict_config = FullStateDictConfig(offload_to_cpu=offload_to_cpu, rank0_only=rank0_only)\n with get_fsdp_state_ctx(\n model, state_type=StateDictType.FULL_STATE_DICT, state_cfg=state_dict_config, optim_cfg=None\n ):\n state_dict = model.state_dict()\n return state_dict\n elif fsdp_version(model) == 2:\n from torch.distributed.checkpoint.state_dict import StateDictOptions, get_model_state_dict\n\n state_dict_config = StateDictOptions(\n full_state_dict=True, cpu_offload=offload_to_cpu, broadcast_from_rank0=not rank0_only\n )\n state_dict = get_model_state_dict(model, options=state_dict_config)\n return state_dict\n else:\n raise NotImplementedError(f\"Unknown FSDP version {fsdp_version}\")\n\n\ndef fsdp2_load_full_state_dict(model: torch.nn.Module, full_state: dict, device_mesh=None, cpu_offload=None):\n \"\"\"\n Loads the full state dict (could be only on rank 0) into the sharded model. This is done by broadcasting the\n parameters from rank 0 to all other ranks. This function modifies the model in-place.\n\n Args:\n model (`torch.nn.Module`): The model to load the state dict into\n full_state (`dict`): The full state dict to load, can only be on rank 0\n \"\"\"\n\n if version.parse(torch.__version__) >= version.parse(\"2.7.0\"):\n from torch.distributed.checkpoint.state_dict import StateDictOptions, set_model_state_dict\n else:\n # official torch 2.6.0 set_model_state_dict API leads to OOM\n # use torch 2.7.0 copy from verl/third_party/torch/distributed/checkpoint\n from verl.third_party.torch.distributed.checkpoint.state_dict import StateDictOptions, set_model_state_dict\n\n # To broadcast, it needs to be instantiated in the GPU.\n if dist.get_rank() == 0:\n model = model.to(device=get_device_id(), non_blocking=True)\n else:\n model = model.to_empty(device=get_device_id())\n\n cpu_offload = cpu_offload is not None\n options = StateDictOptions(full_state_dict=True, cpu_offload=cpu_offload, broadcast_from_rank0=True)\n set_model_state_dict(model, full_state, options=options)\n\n # rotary_emb is not in state_dict, so we need to broadcast it manually\n for name, buf in model.named_buffers():\n dist.broadcast(buf, src=0)\n\n if cpu_offload:\n model.to(\"cpu\", non_blocking=True)\n for buf in model.buffers():\n buf.data = buf.data.to(get_device_id())\n\n\n@contextmanager\ndef maybe_patch_fsdp_module(model):\n if fully_shard_module is None:\n yield\n return\n\n orig_fsdp_module = fully_shard_module.FSDPModule\n\n class FSDPModuleABC(ABC, orig_fsdp_module):\n pass\n\n try:\n if isinstance(model, ABC):\n fully_shard_module.FSDPModule = FSDPModuleABC\n yield\n finally:\n fully_shard_module.FSDPModule = orig_fsdp_module\n\n\ndef apply_fsdp2(model, fsdp_kwargs, config):\n \"\"\"model: AutoModelForCausalLM\"\"\"\n assert CPUOffloadPolicy is not None, \"PyTorch version >= 2.4 is required for using fully_shard API (FSDP2)\"\n\n default_transformer_cls_names_to_wrap = getattr(model, \"_no_split_modules\", None)\n fsdp_transformer_layer_cls_to_wrap = config.get(\"wrap_policy\", {}).get(\n \"transformer_layer_cls_to_wrap\", default_transformer_cls_names_to_wrap\n )\n\n if isinstance(fsdp_transformer_layer_cls_to_wrap, str):\n fsdp_transformer_layer_cls_to_wrap = [fsdp_transformer_layer_cls_to_wrap]\n\n assert len(fsdp_transformer_layer_cls_to_wrap) > 0 and fsdp_transformer_layer_cls_to_wrap[0] is not None\n\n modules = []\n for name, module in model.named_modules():\n if module.__class__.__name__ in fsdp_transformer_layer_cls_to_wrap or (\n isinstance(module, nn.Embedding) and not model.config.tie_word_embeddings\n ):\n modules.append(module)\n\n for idx, module in enumerate(modules):\n # if torch.distributed.is_initialized() and torch.distributed.get_rank() == 0:\n # print(f\"wrap module {module.__class__.__name__}\")\n with maybe_patch_fsdp_module(module):\n fully_shard(module, **fsdp_kwargs)\n\n # if torch.distributed.is_initialized() and torch.distributed.get_rank() == 0:\n # print(f\"wrap module {model.__class__.__name__}\")\n with maybe_patch_fsdp_module(model):\n fully_shard(model, **fsdp_kwargs) # fsdp2 will not reshard_after_forward for root module\n\n\ndef get_shard_placement_fn(fsdp_size):\n \"\"\"Choose the dimension that can divide fsdp_size to avoid padding\"\"\"\n\n def shard_placement_fn(param):\n shape = list(param.shape)\n for i in range(len(shape)):\n if shape[i] % fsdp_size == 0:\n return Shard(i)\n return Shard(0)\n\n return shard_placement_fn\n\n\ndef fsdp2_clip_grad_norm_(parameters, max_norm, norm_type=2.0, error_if_nonfinite=False, foreach=None):\n \"\"\"torch.nn.utils.clip_grad_norm_ cann't run on cpu parameter DTensor\"\"\"\n from torch.nn.utils.clip_grad import _clip_grads_with_norm_, _get_total_norm\n\n if isinstance(parameters, torch.Tensor):\n parameters = [parameters]\n else:\n # prevent generators from being exhausted\n parameters = list(parameters)\n grads = [p.grad for p in parameters if p.grad is not None]\n total_norm = _get_total_norm(grads, norm_type, error_if_nonfinite, foreach)\n total_norm = total_norm.to(get_device_id(), non_blocking=True)\n _clip_grads_with_norm_(parameters, max_norm, total_norm, foreach)\n return total_norm\n\n\ndef layered_summon_lora_params(fsdp_module) -> OrderedDict:\n from peft.utils.save_and_load import get_peft_model_state_dict\n\n def __prefix_submodules(module, prefix):\n for name, submodule in module.named_modules():\n if name.startswith(prefix) and \".\" not in name[len(prefix) :]:\n yield name, submodule\n\n lora_params = OrderedDict()\n prefix_list = [\n # fsdp\n \"_fsdp_wrapped_module.base_model.model.\",\n \"_fsdp_wrapped_module.base_model.model.model.\",\n \"_fsdp_wrapped_module.base_model.model.model.layers.\",\n \"_fsdp_wrapped_module.base_model.model.model.language_model.layers.\",\n # fsdp2\n \"base_model.model.\",\n \"base_model.model.model.\",\n \"base_model.model.model.layers.\",\n \"base_model.model.model.language_model.layers.\",\n ]\n peft_model = getattr(fsdp_module, \"_fsdp_wrapped_module\", fsdp_module)\n for prefix in prefix_list:\n for name, submodule in __prefix_submodules(fsdp_module, prefix):\n prefix = name.replace(\"_fsdp_wrapped_module.base_model.model.\", \"base_model.model.\")\n if name.endswith(\".model\") or name.endswith(\".layers\"):\n continue\n if fsdp_version(submodule) > 0:\n with FSDP.summon_full_params(submodule, writeback=False):\n sub_lora_params = get_peft_model_state_dict(peft_model, state_dict=submodule.state_dict())\n sub_lora_params = {\n f\"{prefix}.{name}\": param.full_tensor().detach().cpu()\n if hasattr(param, \"full_tensor\")\n else param.detach().cpu()\n for name, param in sub_lora_params.items()\n }\n lora_params.update(sub_lora_params)\n submodule._is_root = False\n get_torch_device().empty_cache()\n return lora_params\n\n\ndef collect_lora_params(module: FSDP, layered_summon: bool, base_sync_done: bool) -> OrderedDict:\n \"\"\"\n collect lora params or full params if base model is not ready in vllm\n work with if isinstance(self.module._fsdp_wrapped_module, PeftModel)\n \"\"\"\n from peft.utils.save_and_load import get_peft_model_state_dict\n\n lora_params = OrderedDict()\n peft_model = getattr(module, \"_fsdp_wrapped_module\", module)\n if fsdp_version(module) > 0:\n if layered_summon:\n if not base_sync_done:\n raise ValueError(\n \"To use layered_summon, you must make sure base-model is preloaded in vllm, e.g. let \"\n \"rollout.load_format=safetensors\"\n )\n lora_params = layered_summon_lora_params(module)\n else:\n with FSDP.summon_full_params(module, writeback=False):\n if base_sync_done:\n lora_params = get_peft_model_state_dict(peft_model)\n lora_params = {\n name: param.full_tensor().detach().cpu()\n if hasattr(param, \"full_tensor\")\n else param.detach().cpu()\n for name, param in lora_params.items()\n }\n else:\n model = peft_model.base_model.model\n orig_dev = \"cpu\" if \"cpu\" in str(next(model.parameters()).device) else get_device_name()\n model = model.to(\"cpu\")\n for name, param in model.state_dict().items():\n if any(x in name for x in [\"_flat_param\", \"lora_\"]):\n continue\n name = name.replace(\"_fsdp_wrapped_module.\", \"\").replace(\".base_layer\", \"\")\n lora_params[name] = (\n param.full_tensor().detach().cpu()\n if hasattr(param, \"full_tensor\")\n else param.detach().cpu()\n )\n model = model.to(orig_dev)\n get_torch_device().empty_cache()\n else:\n if base_sync_done:\n lora_params = get_peft_model_state_dict(peft_model)\n else:\n model = peft_model.base_model.model\n orig_dev = \"cpu\" if \"cpu\" in str(next(model.parameters()).device) else get_device_name()\n model = model.to(\"cpu\")\n for name, param in model.state_dict().items():\n if any(x in name for x in [\"_flat_param\", \"lora_\"]):\n continue\n name = name.replace(\"_fsdp_wrapped_module.\", \"\").replace(\".base_layer\", \"\")\n lora_params[name] = param.detach().cpu()\n model = model.to(orig_dev)\n return lora_params\n\n\ndef replace_lora_wrapper(k, peft_config):\n \"\"\"Replace LoRA parameter keys with base layer equivalents.\n\n Transforms LoRA parameter names to their corresponding base layer\n names for proper weight loading in vLLM when base model sync is not done.\n\n Args:\n k (str): Original parameter key name.\n\n Returns:\n str: Transformed parameter key for base layer.\n \"\"\"\n stacked_params = [\"q_proj\", \"k_proj\", \"v_proj\", \"o_proj\", \"gate_proj\", \"up_proj\", \"down_proj\"]\n if k.endswith(\".weight\"):\n module_k = k[: -len(\".weight\")]\n if check_exclude_modules(peft_config, module_k):\n return k\n elif any([module_k.endswith(s) for s in stacked_params]) or check_target_modules(peft_config, module_k):\n return f\"{module_k}.base_layer.weight\"\n if k.endswith(\".bias\"):\n module_k = k[: -len(\".bias\")]\n if check_exclude_modules(peft_config, module_k):\n return k\n elif any([module_k.endswith(s) for s in stacked_params]) or check_target_modules(peft_config, module_k):\n return f\"{module_k}.base_layer.bias\"\n return k\n\n\ndef set_reshard_after_forward(module: FSDPModule, reshard_after_forward: bool, recurse: bool = True) -> None:\n \"\"\"\n Sets if the module should reshard parameters after forward. This can be\n used to change the ``reshard_after_forward`` FSDP arg at runtime. For\n example, this can be used to set the FSDP root module's value to\n ``True`` (since it is otherwise specially set to ``False``), or it can\n set an FSDP module's value to ``False`` for running evals and set back\n to ``True`` for training.\n\n Args:\n reshard_after_forward (bool): Whether to reshard parameters after\n forward.\n recurse (bool): Whether to set for all FSDP submodules or just the\n passed-in module.\n\n ---\n Copied from https://github.com/pytorch/pytorch/blob/main/torch/distributed/fsdp/_fully_shard/_fully_shard.py to\n address the absence of the set_reshard_after_forward function in torch versions earlier than 2.8.0.\n \"\"\"\n\n if not isinstance(reshard_after_forward, bool):\n raise ValueError(f\"reshard_after_forward should be a bool, got {type(reshard_after_forward)}\")\n self_module = cast(nn.Module, module)\n modules = list(self_module.modules()) if recurse else [self_module]\n for module in modules:\n if isinstance(module, FSDPModule):\n state = module._get_fsdp_state()\n state._auto_reshard_after_forward = False\n if fsdp_param_group := state._fsdp_param_group:\n fsdp_param_group.post_forward_mesh_info = _get_post_forward_mesh_info(\n reshard_after_forward, fsdp_param_group.mesh_info\n )\n\n\ndef normalize_peft_param_name(params: dict) -> dict:\n \"\"\"\n Converts peft model parameter name to base parameter name\n For example,\n base_model.model.model.embed_tokens.weight -> model.embed_tokens.weight\n base_model.model.model.layers.0.self_attn.q_proj.base_layer.weight -> model.layers.0.self_attn.q_proj.weight\n and remove params such as base_model.model.model.layers.0.self_attn.q_proj.lora_A.default.weight,\n base_model.model.model.layers.0.self_attn.q_proj.lora_B.default.weight\n \"\"\"\n\n def _normalize_peft_name(name: str) -> str:\n return name.replace(\"base_model.model.\", \"\").replace(\"base_model.\", \"\").replace(\".base_layer\", \"\")\n\n def _is_lora_key(name: str) -> bool:\n # catch typical PEFT keys\n return (\"lora_\" in name) or (\".adapter_\" in name)\n\n params = [(_normalize_peft_name(k), v) for k, v in params.items()]\n # strip any residual LoRA tensors\n params = {k: v for k, v in params if not _is_lora_key(k)}\n return params\n\n\ndef _merge_or_unmerge_lora_(module, merge: bool):\n \"\"\"Merge or unmerge LoRA adapters in a module.\n\n Args:\n module: The module containing LoRA layers\n merge: If True, merge LoRA into base model; if False, unmerge LoRA\n \"\"\"\n from peft.tuners.lora import LoraLayer\n\n with torch.no_grad():\n for m in module.modules():\n if isinstance(m, LoraLayer):\n is_merged = getattr(m, \"merged\", False)\n if merge and not is_merged:\n m.merge()\n elif (not merge) and is_merged:\n m.unmerge()\n\n\n# merged_adapters\ndef _clean_merged_lora_(module):\n \"\"\"Cleans the merged lora adapters\"\"\"\n from peft.tuners.lora import LoraLayer\n\n with torch.no_grad():\n for m in module.modules():\n if isinstance(m, LoraLayer):\n merged_adapters = getattr(m, \"merged_adapters\", False)\n if merged_adapters:\n m.merged_adapters = []\n\n\ndef fsdp_merge_unmerge(module: nn.Module, do_merge: bool):\n \"\"\"Merge or unmerge LoRA adapters in FSDP module.\n\n For FSDP (v1), it gathers all model parameters to each device, which may cause OOM.\n For FSDP2, it gathers model parameters layer-by-layer to reduce memory footprint.\n\n Args:\n module: The FSDP module to merge/unmerge LoRA adapters\n do_merge: If True, merge LoRA into base model; if False, unmerge LoRA\n \"\"\"\n version = fsdp_version(module)\n assert version in [1, 2], f\"fsdp_merge_unmerge requires FSDP module, got version {version}\"\n\n if version == 1:\n # Unshard → merge → Reshard\n with FSDP.summon_full_params(module, writeback=True, with_grads=False):\n _merge_or_unmerge_lora_(module, merge=do_merge)\n else:\n # FSDP2: Unshard → merge → Reshard layer-by-layer\n for name, submodule in module.named_modules():\n if isinstance(submodule, FSDPModule) and name != \"\": # skip root model\n with FSDP.summon_full_params(submodule, writeback=True, with_grads=False):\n _merge_or_unmerge_lora_(submodule, merge=do_merge)\n\n\ndef backup_base_model_weights(module):\n \"\"\"Backup base model weights to CPU with LoRA temporarily disabled.\n\n This function temporarily disables LoRA adapters, backs up the clean base model weights\n to CPU, then re-enables the adapters.\n\n Args:\n module: The PEFT model with LoRA adapters\n\n Returns:\n dict: Dictionary mapping parameter name to CPU tensor backup of base model weights\n \"\"\"\n from peft import PeftModel\n\n backup = {}\n with torch.no_grad():\n # Check if module is a PEFT model\n if isinstance(module, PeftModel):\n # Temporarily disable adapters to get clean base model weights\n with module.disable_adapter():\n # Backup base model weights (excluding lora parameters)\n for name, param in module.named_parameters():\n if \"lora\" not in name.lower():\n backup[name] = param.data.clone().cpu()\n else:\n # For non-PEFT models, just backup all parameters\n for name, param in module.named_parameters():\n backup[name] = param.data.clone().cpu()\n return backup\n\n\ndef restore_base_model_weights(module, backup):\n \"\"\"Restore base model weights from CPU backup.\n\n This function restores the base model weights from the CPU backup, effectively\n undoing any LoRA merge operations.\n\n Args:\n module: The PEFT model with LoRA adapters\n backup: Dictionary mapping parameter name to CPU tensor backup of base model weights\n \"\"\"\n with torch.no_grad():\n for name, param in module.named_parameters():\n if name in backup:\n param.data.copy_(backup[name].to(param.device))\n\n\n@contextmanager\ndef merged_lora_context(actor, backup_adapters=False):\n \"\"\"Context manager to temporarily merge LoRA adapters.\n\n This context manager merges LoRA adapters into the base model weights,\n performs operations (like syncing weights to vLLM), then restores the base model\n weights from backup.\n\n Args:\n actor: The actor module with LoRA adapters to merge\n backup_adapters: If True, backup base model weights (with LoRA disabled) before\n merging and restore them after. This is more numerically stable than unmerging.\n\n Yields:\n None\n \"\"\"\n base_weights_backup = None\n if backup_adapters:\n # Backup base model weights with LoRA temporarily disabled\n base_weights_backup = backup_base_model_weights(actor)\n\n # Merge LoRA adapters into base model\n fsdp_merge_unmerge(actor, do_merge=True)\n try:\n # Do work while merged (sync_to_vllm / generate / etc.)\n yield\n finally:\n if backup_adapters and base_weights_backup is not None:\n # Restore base model weights from CPU backup (effectively undoing the merge)\n restore_base_model_weights(actor, base_weights_backup)\n _clean_merged_lora_(actor)\n else:\n # Fall back to unmerge if no backup was made\n fsdp_merge_unmerge(actor, do_merge=False)\n"}80{"file_name": "verl__utils__groupwise.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\"\"\"\nGroup-wise helpers for RL training utilities.\n\nPublic API:\n - as_torch_index(index, device=None) -> torch.LongTensor\n - group_mean_std(scores, gidx, eps=1e-6, device=None) -> (mean_g, std_g, count_g)\n\nDefault device policy:\n - If `device` is None:\n * In pytest (detected by env \"PYTEST_CURRENT_TEST\"): use CPU.\n * Else if CUDA is available: use CUDA.\n * Else: use CPU.\n - You can override via env \"VERL_FORCE_DEVICE\" (e.g., \"cuda:0\" / \"cpu\").\n\nNotes:\n- as_torch_index: canonicalizes arbitrary group labels to a contiguous 1-D torch.long\n tensor in range [0..G-1]. Robust to torch/numpy/list/tuple, ints/floats/bools,\n numeric strings, UUIDs, mixed object arrays. Near-integer floats (|x-round(x)|<=1e-6)\n are rounded; otherwise factorization is applied.\n- group_mean_std: pure-PyTorch per-group mean/std with Bessel correction for variance\n (denominator max(count-1, 1)). Singleton groups fallback to mean=0, std=1 for\n compatibility with common “native” conventions.\n\"\"\"\n\nfrom __future__ import annotations\n\nimport os\nfrom typing import Any, Optional\n\nimport numpy as np\nimport torch\n\nfrom verl.utils.device import get_device_name\n\n__all__ = [\"as_torch_index\", \"group_mean_std\"]\n\n\ndef _resolve_device(explicit: Optional[torch.device | str]) -> torch.device:\n \"\"\"\n Resolve device according to policy described in the module docstring.\n Priority:\n 1) explicit argument\n 2) VERL_FORCE_DEVICE env\n 3) pytest detection -> cpu\n 4) cuda if available, else cpu\n \"\"\"\n if explicit is not None:\n return torch.device(explicit)\n\n forced = os.getenv(\"VERL_FORCE_DEVICE\")\n if forced:\n return torch.device(forced)\n\n # Heuristic: pytest sets PYTEST_CURRENT_TEST\n if \"PYTEST_CURRENT_TEST\" in os.environ:\n return torch.device(\"cpu\")\n\n return torch.device(get_device_name())\n\n\ndef _to_1d_numpy_object_array(x: Any) -> np.ndarray:\n \"\"\"Best-effort: convert arbitrary input into a 1-D numpy array; fallback to object dtype.\"\"\"\n try:\n arr = np.asarray(x)\n except Exception:\n try:\n arr = np.array(list(x), dtype=object)\n except Exception:\n arr = np.array([x], dtype=object)\n if arr.ndim != 1:\n arr = arr.reshape(-1)\n return arr\n\n\ndef as_torch_index(index: Any, device: torch.device | str | None = None) -> torch.Tensor:\n \"\"\"\n Convert arbitrary group labels to a contiguous 1-D torch.long tensor (0..G-1).\n\n Args:\n index: Any iterable of labels or tensor/ndarray.\n device: Target device; if None, resolved via _resolve_device().\n\n Returns:\n torch.LongTensor with shape (N,)\n \"\"\"\n target = _resolve_device(device)\n\n # ---------- Fast path: torch.Tensor ----------\n if isinstance(index, torch.Tensor):\n t = index.reshape(-1)\n if t.dtype in (\n torch.int64,\n torch.int32,\n torch.int16,\n torch.int8,\n getattr(torch, \"uint8\", torch.uint8),\n torch.bool,\n ):\n return t.to(device=target, dtype=torch.long)\n\n if t.dtype in (torch.float16, torch.float32, torch.float64, torch.bfloat16):\n t64 = t.to(dtype=torch.float64)\n rounded = torch.round(t64)\n if torch.allclose(t64, rounded, rtol=0.0, atol=1e-6):\n return rounded.to(device=target, dtype=torch.long)\n arr = np.array([str(x.item()) for x in t], dtype=object)\n else:\n arr = np.array([str(x.item()) if hasattr(x, \"item\") else str(x) for x in t], dtype=object)\n\n else:\n # ---------- Non-torch: go through numpy ----------\n arr = _to_1d_numpy_object_array(index)\n\n # Pure integers (incl. bool)\n if arr.dtype != object and np.issubdtype(arr.dtype, np.integer):\n return torch.from_numpy(arr.astype(np.int64, copy=False)).to(device=target)\n\n # Floats nearly equal to integers\n if arr.dtype != object and np.issubdtype(arr.dtype, np.floating):\n arr64 = arr.astype(np.float64, copy=False)\n rounded = np.rint(arr64)\n if np.allclose(arr64, rounded, rtol=0.0, atol=1e-6):\n return torch.from_numpy(rounded.astype(np.int64)).to(device=target)\n # fall through\n\n # Try numeric string coercion\n try:\n coerced = arr.astype(np.int64)\n return torch.from_numpy(coerced).to(device=target)\n except Exception:\n pass\n\n if arr.dtype != object:\n arr = arr.astype(object)\n\n # ---------- Factorization (UUIDs / mixed types / arbitrary labels) ----------\n try:\n _, inv = np.unique(arr, return_inverse=True)\n except Exception:\n sarr = np.array([str(x) for x in arr], dtype=object)\n _, inv = np.unique(sarr, return_inverse=True)\n\n inv = inv.astype(np.int64, copy=False)\n return torch.from_numpy(inv).to(device=target)\n\n\n@torch.no_grad()\ndef group_mean_std(\n scores: torch.Tensor,\n gidx: torch.Tensor,\n eps: float = 1e-6,\n device: torch.device | str | None = None,\n) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n \"\"\"\n Compute per-group mean/std/count in pure PyTorch.\n\n mean_g = sum / count\n std_g = sqrt( max( (sum2 - sum^2/count) / max(count-1, 1), eps ) )\n\n Singleton groups fallback to mean=0, std=1.\n\n Args:\n scores: (N,) float tensor.\n gidx : (N,) long/int tensor with group indices (0..G-1).\n eps : Numerical floor for variance.\n device: Target device; if None, resolved via _resolve_device().\n\n Returns:\n mean_g: (G,) float32\n std_g : (G,) float32\n count : (G,) float32\n \"\"\"\n target = _resolve_device(device)\n\n scores = scores.reshape(-1).to(device=target, dtype=torch.float32)\n gidx = gidx.reshape(-1).to(device=target, dtype=torch.long)\n\n if scores.numel() != gidx.numel():\n raise ValueError(f\"scores and gidx length mismatch: {scores.numel()} vs {gidx.numel()}\")\n\n G = int(torch.max(gidx).item()) + 1 if gidx.numel() > 0 else 0\n if G == 0:\n # Return empty tensors on the selected device\n empty = torch.empty(0, device=target, dtype=torch.float32)\n return empty, empty, empty\n\n ones = torch.ones_like(scores, dtype=torch.float32)\n\n count = torch.zeros(G, device=target, dtype=torch.float32).index_add_(0, gidx, ones)\n s1 = torch.zeros(G, device=target, dtype=torch.float32).index_add_(0, gidx, scores)\n s2 = torch.zeros(G, device=target, dtype=torch.float32).index_add_(0, gidx, scores * scores)\n\n mean = s1 / count.clamp_min(1.0)\n var_num = s2 - (s1 * s1) / count.clamp_min(1.0)\n denom = (count - 1.0).clamp_min(1.0)\n var = var_num / denom\n std = torch.sqrt(torch.clamp(var, min=eps))\n\n # Singleton groups: mean=0, std=1\n single = count <= 1.0\n if torch.any(single):\n mean = mean.clone()\n std = std.clone()\n mean[single] = 0.0\n std[single] = 1.0\n\n return mean, std, count\n"}81{"file_name": "verl__utils__hdfs_io.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nimport os\nimport shutil\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_SFT_LOGGING_LEVEL\", \"WARN\"))\n\n_HDFS_PREFIX = \"hdfs://\"\n\n_HDFS_BIN_PATH = shutil.which(\"hdfs\")\n\n\ndef exists(path: str, **kwargs) -> bool:\n r\"\"\"Works like os.path.exists() but supports hdfs.\n\n Test whether a path exists. Returns False for broken symbolic links.\n\n Args:\n path (str): path to test\n\n Returns:\n bool: True if the path exists, False otherwise\n \"\"\"\n if _is_non_local(path):\n return _exists(path, **kwargs)\n return os.path.exists(path)\n\n\ndef _exists(file_path: str):\n \"\"\"hdfs capable to check whether a file_path is exists\"\"\"\n if file_path.startswith(\"hdfs\"):\n return _run_cmd(_hdfs_cmd(f\"-test -e {file_path}\")) == 0\n return os.path.exists(file_path)\n\n\ndef makedirs(name, mode=0o777, exist_ok=False, **kwargs) -> None:\n r\"\"\"Works like os.makedirs() but supports hdfs.\n\n Super-mkdir; create a leaf directory and all intermediate ones. Works like\n mkdir, except that any intermediate path segment (not just the rightmost)\n will be created if it does not exist. If the target directory already\n exists, raise an OSError if exist_ok is False. Otherwise no exception is\n raised. This is recursive.\n\n Args:\n name (str): directory to create\n mode (int): file mode bits\n exist_ok (bool): if True, do not raise an exception if the directory already exists\n kwargs: keyword arguments for hdfs\n\n \"\"\"\n if _is_non_local(name):\n # TODO(haibin.lin):\n # - handle OSError for hdfs(?)\n # - support exist_ok for hdfs(?)\n _mkdir(name, **kwargs)\n else:\n os.makedirs(name, mode=mode, exist_ok=exist_ok)\n\n\ndef _mkdir(file_path: str) -> bool:\n \"\"\"hdfs mkdir\"\"\"\n if file_path.startswith(\"hdfs\"):\n _run_cmd(_hdfs_cmd(f\"-mkdir -p {file_path}\"))\n else:\n os.makedirs(file_path, exist_ok=True)\n return True\n\n\ndef copy(src: str, dst: str, **kwargs) -> bool:\n r\"\"\"Works like shutil.copy() for file, and shutil.copytree for dir, and supports hdfs.\n\n Copy data and mode bits (\"cp src dst\"). Return the file's destination.\n The destination may be a directory.\n If source and destination are the same file, a SameFileError will be\n raised.\n\n Arg:\n src (str): source file path\n dst (str): destination file path\n kwargs: keyword arguments for hdfs copy\n\n Returns:\n str: destination file path\n\n \"\"\"\n if _is_non_local(src) or _is_non_local(dst):\n # TODO(haibin.lin):\n # - handle SameFileError for hdfs files(?)\n # - return file destination for hdfs files\n return _copy(src, dst)\n else:\n if os.path.isdir(src):\n return shutil.copytree(src, dst, **kwargs)\n else:\n return shutil.copy(src, dst, **kwargs)\n\n\ndef _copy(from_path: str, to_path: str, timeout: int = None) -> bool:\n if to_path.startswith(\"hdfs\"):\n if from_path.startswith(\"hdfs\"):\n returncode = _run_cmd(_hdfs_cmd(f\"-cp -f {from_path} {to_path}\"), timeout=timeout)\n else:\n returncode = _run_cmd(_hdfs_cmd(f\"-put -f {from_path} {to_path}\"), timeout=timeout)\n else:\n if from_path.startswith(\"hdfs\"):\n returncode = _run_cmd(\n _hdfs_cmd(\n f\"-get \\\n {from_path} {to_path}\"\n ),\n timeout=timeout,\n )\n else:\n try:\n shutil.copy(from_path, to_path)\n returncode = 0\n except shutil.SameFileError:\n returncode = 0\n except Exception as e:\n logger.warning(f\"copy {from_path} {to_path} failed: {e}\")\n returncode = -1\n return returncode == 0\n\n\ndef _run_cmd(cmd: str, timeout=None):\n return os.system(cmd)\n\n\ndef _hdfs_cmd(cmd: str) -> str:\n return f\"{_HDFS_BIN_PATH} dfs {cmd}\"\n\n\ndef _is_non_local(path: str):\n return path.startswith(_HDFS_PREFIX)\n"}82{"file_name": "verl__utils__kernel__fp8_kernel.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport logging\nimport os\n\nimport torch\n\nlogger = logging.getLogger(__name__)\n\n# Check if Triton is available\n_TRITON_AVAILABLE = False\ntry:\n import triton\n import triton.language as tl\n\n _TRITON_AVAILABLE = True\nexcept ImportError:\n logger.debug(\"Triton not available, FP8 Triton kernels will not be used\")\n\n# Environment variable to control Triton FP8 usage (set to \"1\" to disable)\n_DISABLE_TRITON_FP8 = os.environ.get(\"VERL_DISABLE_TRITON_FP8\", \"0\").lower() in (\"1\", \"true\", \"yes\")\n\n# FP8 constants\nFP8_DTYPE = torch.float8_e4m3fn\nFP8_MAX = torch.finfo(FP8_DTYPE).max\nFP8_MIN = -FP8_MAX\n\n\ndef ceil_div(x: int, y: int) -> int:\n \"\"\"Perform ceiling division of two integers.\"\"\"\n return (x + y - 1) // y\n\n\ndef is_triton_available() -> bool:\n \"\"\"Check if Triton is available for FP8 kernels.\"\"\"\n return _TRITON_AVAILABLE\n\n\nif _TRITON_AVAILABLE:\n\n @triton.jit\n def _blockwise_cast_to_fp8_kernel(\n X,\n Y,\n S,\n stride_xm,\n stride_xn,\n stride_ym,\n stride_yn,\n stride_sm,\n stride_sn,\n M,\n N,\n eps,\n fp8_min,\n fp8_max,\n BLOCK_M: tl.constexpr = 128,\n BLOCK_N: tl.constexpr = 128,\n ):\n \"\"\"Triton kernel for blockwise FP8 quantization.\n\n Each program instance handles one block of size (BLOCK_M, BLOCK_N).\n Computes per-block scale and quantizes to FP8 in a single pass.\n\n Refer to https://github.com/THUDM/slime/blob/main/slime/backends/megatron_utils/kernels/fp8_kernel.py\n \"\"\"\n pid_m = tl.cast(tl.program_id(axis=0), tl.int64)\n pid_n = tl.cast(tl.program_id(axis=1), tl.int64)\n\n # Compute block offsets\n off_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)\n off_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)\n\n # Create masks for boundary handling\n mask_m = off_m < M\n mask_n = off_n < N\n mask = mask_m[:, None] & mask_n[None, :]\n\n # Load input block and convert to float32 for precision\n x = tl.load(X + off_m[:, None] * stride_xm + off_n[None, :] * stride_xn, mask=mask, other=0.0).to(tl.float32)\n\n # Compute block-wise absolute maximum with epsilon for numerical stability\n _absmax = tl.maximum(tl.max(tl.abs(x)), eps)\n\n # Compute scale: scale = absmax / fp8_max\n x_s = _absmax / fp8_max\n\n # Compute inverse scale for quantization\n s_inv = 1.0 / x_s\n\n # Quantize: clamp(x * s_inv, fp8_min, fp8_max)\n y_q = tl.clamp(x * s_inv, fp8_min, fp8_max).to(Y.dtype.element_ty)\n\n # Store quantized values and scale\n tl.store(Y + off_m[:, None] * stride_ym + off_n[None, :] * stride_yn, y_q, mask=mask)\n tl.store(S + pid_m * stride_sm + pid_n * stride_sn, x_s)\n\n def blockwise_cast_to_fp8_triton(\n x: torch.Tensor,\n weight_block_size: list[int] | tuple[int, int] | None = None,\n ) -> tuple[torch.Tensor, torch.Tensor]:\n \"\"\"Quantize a 2D tensor to FP8 using blockwise quantization with Triton.\n\n This function provides high-performance FP8 quantization with minimal memory overhead.\n All computations (abs, max, scale, clamp) are performed in a single Triton kernel,\n eliminating intermediate tensor allocations.\n\n Args:\n x: Input tensor of shape (M, N), must be 2D.\n weight_block_size: Block size for quantization as [BLOCK_M, BLOCK_N].\n Defaults to [128, 128] if None.\n\n Returns:\n Tuple of (quantized_tensor, scale_tensor):\n - quantized_tensor: FP8 quantized tensor of shape (M, N)\n - scale_tensor: Per-block scale factors of shape (ceil(M/BLOCK_M), ceil(N/BLOCK_N))\n This is the inverse scale (multiply to dequantize).\n \"\"\"\n assert x.dim() == 2, f\"Expected 2D tensor, got {x.dim()}D\"\n\n # Default block size\n BLOCK_M, BLOCK_N = 128, 128\n if weight_block_size is not None:\n BLOCK_M, BLOCK_N = weight_block_size[0], weight_block_size[1]\n\n M, N = x.shape\n\n # Pre-allocate output tensors (only memory allocation in this function)\n y = torch.empty(M, N, device=x.device, dtype=FP8_DTYPE)\n s = torch.empty(ceil_div(M, BLOCK_M), ceil_div(N, BLOCK_N), dtype=torch.float32, device=x.device)\n\n # Grid: one program per block\n def grid(meta):\n return (triton.cdiv(M, meta[\"BLOCK_M\"]), triton.cdiv(N, meta[\"BLOCK_N\"]))\n\n # Tune kernel parameters based on memory layout\n if x.is_contiguous():\n kwargs = {\"BLOCK_M\": BLOCK_M, \"BLOCK_N\": BLOCK_N, \"num_warps\": 8, \"num_stages\": 2}\n else:\n kwargs = {\"BLOCK_M\": BLOCK_M, \"BLOCK_N\": BLOCK_N, \"num_warps\": 1, \"num_stages\": 4}\n\n # Launch kernel\n _blockwise_cast_to_fp8_kernel[grid](\n x,\n y,\n s,\n *x.stride(),\n *y.stride(),\n *s.stride(),\n M,\n N,\n 1e-10, # eps for numerical stability\n FP8_MIN,\n FP8_MAX,\n **kwargs,\n )\n\n return y, s\n\n\ndef scaled_fp8_blockwise_triton(\n data_hp: torch.Tensor,\n weight_block_size: list[int] | tuple[int, int],\n) -> tuple[torch.Tensor, torch.Tensor]:\n \"\"\"High-performance FP8 blockwise quantization using Triton kernel.\n\n This is the recommended function to use for FP8 quantization when Triton is available.\n It handles padding automatically and returns results in the expected format.\n\n Args:\n data_hp: Input high-precision tensor of shape (M, N).\n weight_block_size: Block size for quantization as [BLOCK_M, BLOCK_N].\n\n Returns:\n Tuple of (fp8_data, descale):\n - fp8_data: FP8 quantized tensor of original shape\n - descale: Per-block descale factors (inverse of scale, for dequantization)\n\n Raises:\n RuntimeError: If Triton is not available.\n \"\"\"\n if not _TRITON_AVAILABLE:\n raise RuntimeError(\"Triton is required for scaled_fp8_blockwise_triton but is not available\")\n\n block_size0 = weight_block_size[0]\n block_size1 = weight_block_size[1]\n\n # Save original shape for potential cropping\n original_shape = data_hp.shape\n\n # Pad dimensions to be multiples of block size if needed\n pad_dim0 = (block_size0 - data_hp.shape[0] % block_size0) % block_size0\n pad_dim1 = (block_size1 - data_hp.shape[1] % block_size1) % block_size1\n\n if pad_dim0 > 0 or pad_dim1 > 0:\n logger.debug(\n f\"Padding weight from {data_hp.shape} to \"\n f\"({data_hp.shape[0] + pad_dim0}, {data_hp.shape[1] + pad_dim1}) \"\n f\"for blockwise FP8 quantization\"\n )\n data_hp = torch.nn.functional.pad(data_hp, (0, pad_dim1, 0, pad_dim0), mode=\"constant\", value=0)\n\n # Call Triton kernel\n fp_data, scale = blockwise_cast_to_fp8_triton(data_hp, weight_block_size)\n\n # Remove padding to restore original shape\n if pad_dim0 > 0 or pad_dim1 > 0:\n fp_data = fp_data[: original_shape[0], : original_shape[1]].contiguous()\n\n # Return scale as descale (the Triton kernel returns scale, we need to return it as-is\n # since it's already the inverse scale format expected by vLLM/SGLang)\n return fp_data, scale\n\n\ndef _scaled_fp8_blockwise_pytorch(\n data_hp: torch.Tensor,\n weight_block_size: list[int] | tuple[int, int],\n) -> tuple[torch.Tensor, torch.Tensor]:\n \"\"\"PyTorch implementation of blockwise FP8 quantization.\n\n Memory-optimized implementation that:\n - Uses in-place operations where possible\n - Explicitly deletes intermediate tensors\n - Minimizes peak memory usage during quantization\n\n Args:\n data_hp: Input high-precision tensor of shape (M, N).\n weight_block_size: Block size for quantization as [BLOCK_M, BLOCK_N].\n\n Returns:\n Tuple of (fp8_data, descale):\n - fp8_data: FP8 quantized tensor\n - descale: Per-block descale factors for dequantization\n \"\"\"\n block_size0 = weight_block_size[0]\n block_size1 = weight_block_size[1]\n assert block_size0 == block_size1, \"Block sizes must be equal\"\n\n # Save unpadded shape for later cropping\n original_shape = data_hp.shape\n\n # Pad dimensions to be multiples of block size if needed\n pad_dim0 = (block_size0 - data_hp.shape[0] % block_size0) % block_size0\n pad_dim1 = (block_size1 - data_hp.shape[1] % block_size1) % block_size1\n\n if pad_dim0 > 0 or pad_dim1 > 0:\n logger.debug(\n f\"Padding weight from {data_hp.shape} to \"\n f\"({data_hp.shape[0] + pad_dim0}, {data_hp.shape[1] + pad_dim1}) \"\n f\"for blockwise FP8 quantization\"\n )\n data_hp = torch.nn.functional.pad(data_hp, (0, pad_dim1, 0, pad_dim0), mode=\"constant\", value=0)\n\n # FP8\n max_dtype = FP8_MAX\n\n padded_shape = data_hp.shape\n blk_m, blk_n = data_hp.shape[0] // block_size0, data_hp.shape[1] // block_size1\n\n # Reshape and permute - these are views, no memory allocation\n data_hp = data_hp.reshape(blk_m, block_size0, blk_n, block_size1)\n data_hp = data_hp.permute(0, 2, 1, 3).contiguous()\n\n # Flatten to (BLK_M, BLK_N, BLOCK_SIZE_M * BLOCK_SIZE_N) in float32 for precision\n data_hp = data_hp.to(torch.float32).flatten(start_dim=2)\n\n # Calculate max absolute value per block - use fused abs+amax\n max_abs = data_hp.abs().amax(dim=-1, keepdim=True)\n\n # Compute scale in-place where possible\n scale_fp = torch.empty_like(max_abs)\n torch.div(max_dtype, max_abs, out=scale_fp)\n # Handle edge cases: zero and inf\n scale_fp = torch.where(max_abs == 0, torch.ones_like(scale_fp), scale_fp)\n scale_fp = torch.where(max_abs == torch.inf, torch.ones_like(scale_fp), scale_fp)\n del max_abs # Free max_abs memory\n\n # Compute descale before modifying data\n descale_fp = torch.reciprocal(scale_fp)\n\n # Scale and clamp in a memory-efficient way\n data_hp.mul_(scale_fp)\n del scale_fp # Free scale memory\n data_hp.clamp_(min=-max_dtype, max=max_dtype)\n\n # Convert to FP8\n fp_data = data_hp.to(FP8_DTYPE)\n del data_hp # Free float32 data\n\n # Reshape back to original layout\n fp_data = fp_data.reshape(blk_m, blk_n, block_size0, block_size1).permute(0, 2, 1, 3).reshape(padded_shape)\n\n # Remove padding to restore original shape\n if original_shape[0] != padded_shape[0] or original_shape[1] != padded_shape[1]:\n fp_data = fp_data[: original_shape[0], : original_shape[1]].contiguous()\n\n return fp_data, descale_fp\n\n\ndef scaled_fp8_blockwise(\n data_hp: torch.Tensor,\n weight_block_size: list[int] | tuple[int, int],\n) -> tuple[torch.Tensor, torch.Tensor]:\n \"\"\"Cast tensor from high precision to FP8 with blockwise quantization.\n\n This function automatically selects the best available implementation:\n 1. Triton kernel (if available): Highest performance, minimal memory overhead\n 2. PyTorch fallback: Memory-optimized implementation using in-place operations\n\n To disable Triton and force PyTorch fallback, set environment variable:\n VERL_DISABLE_TRITON_FP8=1\n\n Args:\n data_hp: Input tensor of shape (M, N) in high precision (bf16/fp16/fp32).\n weight_block_size: Block size for quantization as [BLOCK_M, BLOCK_N].\n\n Returns:\n Tuple of (fp8_data, descale):\n - fp8_data: FP8 quantized tensor\n - descale: Per-block descale factors for dequantization\n \"\"\"\n assert len(data_hp.shape) == 2, \"Only 2d input tensor is supported\"\n\n # Use Triton kernel if available and not disabled\n if _TRITON_AVAILABLE and not _DISABLE_TRITON_FP8:\n return scaled_fp8_blockwise_triton(data_hp, weight_block_size)\n\n # PyTorch fallback implementation (memory-optimized)\n return _scaled_fp8_blockwise_pytorch(data_hp, weight_block_size)\n"}83{"file_name": "verl__utils__kernel__kernels.py", "text": "#\n# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.\n# SPDX-License-Identifier: Apache-2.0\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n#\n\n# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nImplementations of the linear cross entropy with token entropy kernel.\n\"\"\"\n\nimport typing\nfrom dataclasses import dataclass\n\nimport torch\nimport torch.distributed as dist\n\nfrom verl.utils.device import get_device_capability, get_device_name, is_cuda_available\n\ntry:\n import triton\n import triton.language as tl\n\n HAVE_TRITON = True\n SUPPORT_CUDA_TMA = is_cuda_available and get_device_capability()[0] >= 9 and hasattr(tl, \"make_tensor_descriptor\")\n\nexcept ImportError:\n HAVE_TRITON = False\n SUPPORT_CUDA_TMA = False\n\nfrom verl.utils.device import get_torch_device\n\nif not HAVE_TRITON:\n from contextlib import contextmanager\n from unittest.mock import MagicMock\n\n @contextmanager\n def null_decorator(*args, **kwargs):\n if len(kwargs) == 0 and len(args) == 1 and callable(args[0]):\n return args[0]\n else:\n\n def inner(func):\n return func\n\n return inner\n\n triton = MagicMock()\n triton.jit = null_decorator\n triton.autotune = null_decorator\n tl = MagicMock()\n\nelif SUPPORT_CUDA_TMA:\n # TMA descriptors require a global memory allocation\n def alloc_fn(size: int, alignment: int, stream: typing.Optional[int]):\n return torch.empty(size, device=get_device_name(), dtype=torch.int8)\n\n # https://github.com/triton-lang/triton/commit/43625fc968b693ab51884ca95adbcf3e43483fd0\n # Triton 3.5.0 stores allocators in ContextVar; values do not propagate to new\n # threads by default. Some execution paths in verl use thread pools (e.g.,\n # concurrent.futures), so we set a ContextVar *default* to avoid falling\n # back to NullAllocator in worker threads.\n try:\n import contextvars\n\n import triton.runtime._allocation as _triton_allocation\n\n if isinstance(getattr(_triton_allocation, \"_allocator\", None), contextvars.ContextVar):\n _triton_allocation._allocator = contextvars.ContextVar(\n _triton_allocation._allocator.name,\n default=alloc_fn,\n )\n except (ImportError, AttributeError):\n pass\n\n triton.set_allocator(alloc_fn)\n\n\n@dataclass\nclass EntropyReductionEnum:\n \"\"\"\n Enum for the reduction method of cross entropy.\n \"\"\"\n\n _None = 0\n _Sum = 1\n _Mean = 2\n\n\ndef get_entropy_reduction_enum_number(reduction: str) -> int:\n \"\"\"\n Get the enum number for the reduction method of cross entropy.\n \"\"\"\n _enum = EntropyReductionEnum._None\n if reduction == \"none\":\n _enum = EntropyReductionEnum._None\n elif reduction == \"sum\":\n _enum = EntropyReductionEnum._Sum\n elif reduction == \"mean\":\n _enum = EntropyReductionEnum._Mean\n else:\n raise ValueError(f\"Invalid reduction: {reduction}\")\n return _enum\n\n\ndef get_entropy_reduction_enum(ce_reduction: int) -> EntropyReductionEnum:\n \"\"\"\n Get the enum for the reduction method of cross entropy.\n \"\"\"\n _enum = EntropyReductionEnum._None\n if ce_reduction == 0:\n _enum = EntropyReductionEnum._None\n elif ce_reduction == 1:\n _enum = EntropyReductionEnum._Sum\n elif ce_reduction == 2:\n _enum = EntropyReductionEnum._Mean\n else:\n raise ValueError(f\"Invalid ce_reduction: {ce_reduction}\")\n return _enum\n\n\n@dataclass\nclass BackwardEnum:\n \"\"\"\n Enum for the backward method.\n \"\"\"\n\n _Total_Fuse_MN = (\n 0 # Fuse d_logits & d_hidden & d_weight, no intermediate storage, requires fp32 for d_hidden & d_weight\n )\n _Total_Separate = 1 # Store d_logits, no special requirements for d_hidden & d_weight\n _Split_Dlogits_N = 2 # split d_logits along its N dimension, aka. vocab_size\n _Split_Dlogits_M = 3 # split d_logits along its M dimension, aka. num_tokens\n\n\n@dataclass\nclass Config:\n \"\"\"Configuration for efficient entropy kernel operations.\n\n Args:\n _backward (BackwardEnum): Backward computation method. Defaults to BackwardEnum._Split_Dlogits_N.\n _use_triton (bool): Whether to use Triton kernels for computation. Defaults to True.\n \"\"\"\n\n _backward: BackwardEnum = BackwardEnum._Split_Dlogits_N\n _use_triton: bool = True\n\n\n_config = Config()\n\n\ndef set_backward_method(backward_method: BackwardEnum):\n \"\"\"\n Set the backward method.\n \"\"\"\n global _config\n _config._backward = backward_method\n\n\n@triton.autotune(\n configs=[triton.Config({\"BLOCK_SIZE_M\": 128, \"BLOCK_SIZE_N\": 256, \"BLOCK_SIZE_K\": 32}, num_stages=3, num_warps=8)],\n key=[\"num_tokens\", \"hidden_size\", \"vocab_size\"],\n)\n@triton.jit\ndef efficient_entropy_kernel_general_mainloop(\n rank,\n hidden_ptr,\n weight_ptr,\n labels_ptr,\n num_tokens,\n hidden_size,\n vocab_size,\n vocab_per_split,\n stride_hidden_m: tl.int64,\n stride_hidden_k: tl.int64,\n stride_weight_n: tl.int64,\n stride_weight_k: tl.int64,\n max_ptr,\n stride_max_m: tl.int64,\n stride_max_n: tl.int64,\n accu_ptr,\n stride_accu_m: tl.int64,\n stride_accu_n: tl.int64,\n entropy_b_ptr,\n stride_entropy_b_m: tl.int64,\n stride_entropy_b_n: tl.int64,\n global_logprobs_ptr,\n stride_global_logprobs: tl.int64,\n global_logprobs_scalar_ptr,\n rcp_temperature: tl.float32,\n # Meta-parameters\n BLOCK_SIZE_M: tl.constexpr,\n BLOCK_SIZE_N: tl.constexpr,\n BLOCK_SIZE_K: tl.constexpr,\n USE_TMA: tl.constexpr,\n):\n \"\"\"\n forward mainloop\n \"\"\"\n pid = tl.program_id(axis=0)\n num_splits = (vocab_size + vocab_per_split - 1) // vocab_per_split\n num_pid_m = tl.cdiv(num_tokens, BLOCK_SIZE_M)\n num_pid_n = tl.cdiv(vocab_per_split, BLOCK_SIZE_N)\n pid_m = pid % num_pid_m\n pid_n = pid // num_pid_m\n\n if pid_m == 0 and pid_n == 0:\n tl.store(global_logprobs_scalar_ptr, 0.0)\n\n # create pointers for the first blocks of hidden\n start_offs_am = pid_m * BLOCK_SIZE_M\n offs_am = start_offs_am + tl.arange(0, BLOCK_SIZE_M)\n offs_k = tl.arange(0, BLOCK_SIZE_K)\n\n if USE_TMA:\n # using TMA and device-side descriptor creation\n hidden_desc = tl.make_tensor_descriptor(\n hidden_ptr,\n shape=[num_tokens, hidden_size],\n strides=[stride_hidden_m, 1],\n block_shape=[BLOCK_SIZE_M, BLOCK_SIZE_K],\n )\n\n weight_desc = tl.make_tensor_descriptor(\n weight_ptr,\n shape=[vocab_size, hidden_size],\n strides=[stride_weight_n, 1],\n block_shape=[BLOCK_SIZE_N, BLOCK_SIZE_K],\n )\n\n else:\n hidden_ptrs = hidden_ptr + (offs_am[:, None] * stride_hidden_m + offs_k[None, :] * stride_hidden_k)\n\n # load labels for this block\n labels = tl.load(labels_ptr + offs_am, mask=offs_am < num_tokens)\n\n # traverse over N dimension\n # _max = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n _max = tl.full((BLOCK_SIZE_M,), -float(\"inf\"), dtype=tl.float32)\n _accu = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n _entropy_b = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n _logprobs = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n for n in range(0, num_pid_n):\n start_offs_bn = pid_n * vocab_per_split + n * BLOCK_SIZE_N\n offs_bn = start_offs_bn + tl.arange(0, BLOCK_SIZE_N)\n\n logits = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)\n if not USE_TMA:\n # weight_ptrs = weight_ptr + (offs_k[:, None] * stride_weight_k + offs_bn[None, :] * stride_weight_n)\n weight_ptrs = weight_ptr + (offs_bn[:, None] * stride_weight_n + offs_k[None, :] * stride_weight_k)\n\n # iterate over K dimension\n for k in range(0, tl.cdiv(hidden_size, BLOCK_SIZE_K)):\n if USE_TMA:\n # load the next block of hidden and weight\n start_offs_k = k * BLOCK_SIZE_K\n _hidden = hidden_desc.load([start_offs_am, start_offs_k])\n _weight = weight_desc.load([start_offs_bn, start_offs_k])\n else:\n # load the next block of hidden and weight\n _hidden = tl.load(\n hidden_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_am[:, None] < num_tokens),\n other=0.0,\n )\n\n _weight = tl.load(\n weight_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K)\n & (offs_bn[:, None] < (min((pid_n + 1) * vocab_per_split, vocab_size))),\n other=0.0,\n )\n\n # advance the ptrs to the next K block\n hidden_ptrs += BLOCK_SIZE_K * stride_hidden_k\n weight_ptrs += BLOCK_SIZE_K * stride_weight_k\n\n # GEMM\n logits = tl.dot(_hidden, _weight.trans(), logits)\n\n if not USE_TMA:\n # reset hidden_ptrs for next iteration\n hidden_ptrs -= hidden_size * stride_hidden_k\n\n # scale logits by temperature\n logits *= rcp_temperature\n\n # update global maximum\n _max_old = _max\n m_pid_n = tl.max(logits, axis=1)\n _max = tl.maximum(_max_old, m_pid_n)\n\n exp_logits = tl.exp(logits - _max[:, None])\n coeff = tl.exp(_max_old - _max)\n _accu = coeff * _accu + tl.sum(exp_logits, axis=1)\n\n _entropy_b = _entropy_b * coeff + tl.sum(logits * exp_logits, axis=1)\n\n label_mask = (offs_bn + rank * vocab_size)[None, :] == labels[:, None]\n _logprobs += tl.sum(logits * label_mask, axis=1)\n\n # store maximum\n offs_max_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)\n offs_max_n = pid_n\n maximum_ptrs = max_ptr + offs_max_n * stride_max_n + offs_max_m * stride_max_m\n tl.store(maximum_ptrs, _max, mask=(offs_max_m < num_tokens) & (offs_max_n < num_splits))\n\n # store entropy\n accu_ptrs = accu_ptr + offs_max_n * stride_accu_n + offs_max_m * stride_accu_m\n tl.store(accu_ptrs, _accu, mask=(offs_max_m < num_tokens) & (offs_max_n[None] < num_splits))\n entropy_b_ptrs = entropy_b_ptr + offs_max_n * stride_entropy_b_n + offs_max_m * stride_entropy_b_m\n tl.store(entropy_b_ptrs, _entropy_b, mask=(offs_max_m < num_tokens) & (offs_max_n < num_splits))\n # store logprobs\n vocab_left_idx = pid_n * vocab_per_split + rank * vocab_size\n vocab_right_idx = min((pid_n + 1) * vocab_per_split, vocab_size) + rank * vocab_size\n mask = (labels >= vocab_left_idx) & (labels < vocab_right_idx)\n mask &= offs_am < num_tokens\n global_logprobs_ptrs = global_logprobs_ptr + offs_am * stride_global_logprobs\n # tl.atomic_add(global_logprobs_ptrs, _logprobs, mask=mask)\n tl.store(global_logprobs_ptrs, _logprobs, mask=mask)\n\n\n@triton.autotune(configs=[triton.Config({\"BLOCK_SIZE_M\": 16, \"BLOCK_SIZE_N\": 64})], key=[\"num_tokens\", \"num_splits\"])\n@triton.jit\ndef efficient_entropy_triton_kernel_epilogue(\n max_ptr,\n stride_max_m: tl.int64,\n stride_max_n: tl.int64,\n num_tokens,\n num_splits,\n global_max_ptr,\n stride_global_max: tl.int64,\n accu_ptr,\n stride_accu_m: tl.int64,\n stride_accu_n: tl.int64,\n global_accu_ptr,\n stride_global_accu: tl.int64,\n entropy_b_ptr,\n stride_entropy_b_m: tl.int64,\n stride_entropy_b_n: tl.int64,\n global_entropy_b_ptr,\n stride_global_entropy_b: tl.int64,\n global_entropy_ptr,\n stride_global_entropy: tl.int64,\n global_logprobs_ptr,\n stride_global_logprobs: tl.int64,\n global_logprobs_scalar_ptr,\n reduction: int,\n BLOCK_SIZE_M: tl.constexpr,\n BLOCK_SIZE_N: tl.constexpr,\n):\n \"\"\"\n foward epilogue\n \"\"\"\n pid_m = tl.program_id(axis=0)\n\n offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)\n global_max = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n global_accu = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n global_entropy_b = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n for pid_n in range(0, tl.cdiv(num_splits, BLOCK_SIZE_N)):\n offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)\n max_ptrs = max_ptr + offs_m[:, None] * stride_max_m + offs_n[None, :] * stride_max_n\n\n _max = tl.load(max_ptrs, mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits), other=0.0)\n\n accu_ptrs = accu_ptr + offs_m[:, None] * stride_accu_m + offs_n[None, :] * stride_accu_n\n _accu = tl.load(accu_ptrs, mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits), other=0.0)\n\n entropy_b_ptrs = entropy_b_ptr + offs_m[:, None] * stride_entropy_b_m + offs_n[None, :] * stride_entropy_b_n\n _entropy_b = tl.load(\n entropy_b_ptrs, mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits), other=0.0\n )\n\n # local reduction\n _max_old = global_max\n _local_max = tl.max(_max, axis=1)\n global_max = tl.maximum(global_max, _local_max)\n\n _scale = tl.exp(_max - global_max[:, None])\n _coeff = tl.exp(_max_old - global_max)\n global_accu = _coeff * global_accu + tl.sum(_scale * _accu, axis=1)\n global_entropy_b = _coeff * global_entropy_b + tl.sum(_scale * _entropy_b, axis=1)\n\n # store\n maximum_ptrs = global_max_ptr + offs_m * stride_global_max\n tl.store(maximum_ptrs, global_max, mask=offs_m < num_tokens)\n\n # store entropy_b\n global_entropy_b = tl.fdiv(global_entropy_b, global_accu) # entropy_b\n tl.store(global_entropy_b_ptr + offs_m * stride_global_entropy_b, global_entropy_b, mask=offs_m < num_tokens)\n\n # store entropy\n global_accu_ptrs = global_accu_ptr + offs_m * stride_global_accu\n tl.store(global_accu_ptrs, global_accu, mask=offs_m < num_tokens)\n global_entropy = tl.log(global_accu) + global_max - global_entropy_b # entropy_a\n global_entropy_ptrs = global_entropy_ptr + offs_m * stride_global_entropy\n tl.store(global_entropy_ptrs, global_entropy, mask=offs_m < num_tokens)\n # update logprobs\n global_logprobs_ptrs = global_logprobs_ptr + offs_m * stride_global_logprobs\n global_logprobs = tl.load(global_logprobs_ptrs, mask=offs_m < num_tokens)\n global_logprobs = global_max + tl.log(global_accu) - global_logprobs\n\n global_logprobs = -1 * global_logprobs\n if reduction == 0:\n tl.store(global_logprobs_ptrs, global_logprobs, mask=offs_m < num_tokens)\n elif reduction == 1:\n global_logprobs_scalar = tl.sum(global_logprobs, axis=0)\n tl.atomic_add(global_logprobs_scalar_ptr, global_logprobs_scalar)\n elif reduction == 2:\n global_logprobs_scalar = tl.sum(global_logprobs, axis=0) / num_tokens.to(tl.float32)\n tl.atomic_add(global_logprobs_scalar_ptr, global_logprobs_scalar)\n\n\n@triton.autotune(configs=[triton.Config({\"BLOCK_SIZE_M\": 16, \"BLOCK_SIZE_N\": 64})], key=[\"num_tokens\", \"num_splits\"])\n@triton.jit\ndef efficient_entropy_triton_kernel_epilogue_tp(\n num_tokens,\n num_splits,\n reduced_max_ptr,\n stride_reduced_max_m: tl.int64,\n stride_reduced_max_n: tl.int64,\n original_max_ptr,\n stride_original_max_m: tl.int64,\n stride_original_max_n: tl.int64,\n accu_ptr,\n stride_accu_m: tl.int64,\n stride_accu_n: tl.int64,\n entropy_b_ptr,\n stride_entropy_b_m: tl.int64,\n stride_entropy_b_n: tl.int64,\n global_max_ptr,\n stride_global_max: tl.int64,\n global_accu_ptr,\n stride_global_accu: tl.int64,\n global_entropy_b_ptr,\n stride_global_entropy_b: tl.int64,\n BLOCK_SIZE_M: tl.constexpr,\n BLOCK_SIZE_N: tl.constexpr,\n):\n pid_m = tl.program_id(axis=0)\n\n offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)\n\n global_max = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n global_accu = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n global_entropy_b = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32)\n for pid_n in range(0, tl.cdiv(num_splits, BLOCK_SIZE_N)):\n offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)\n\n _reduced_max = tl.load(\n reduced_max_ptr + offs_m[:, None] * stride_reduced_max_m + offs_n[None, :] * stride_reduced_max_n,\n mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits),\n other=0.0,\n )\n _original_max = tl.load(\n original_max_ptr + offs_m[:, None] * stride_original_max_m + offs_n[None, :] * stride_original_max_n,\n mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits),\n other=0.0,\n )\n _accu = tl.load(\n accu_ptr + offs_m[:, None] * stride_accu_m + offs_n[None, :] * stride_accu_n,\n mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits),\n other=0.0,\n )\n\n # local reduce-max\n _max_old = global_max\n _local_max = tl.max(_reduced_max, axis=1)\n global_max = tl.maximum(global_max, _local_max)\n\n # update accumulate\n _coeff = tl.exp(_max_old - global_max)\n _scale = tl.exp(_original_max - global_max[:, None])\n global_accu = _coeff * global_accu + tl.sum(_scale * _accu, axis=1)\n\n # update entropy_b\n _entropy_b = tl.load(\n entropy_b_ptr + offs_m[:, None] * stride_entropy_b_m + offs_n[None, :] * stride_entropy_b_n,\n mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits),\n other=0.0,\n )\n global_entropy_b = _coeff * global_entropy_b + tl.sum(_scale * _entropy_b, axis=1)\n\n # store\n tl.store(global_max_ptr + offs_m * stride_global_max, global_max, mask=offs_m < num_tokens)\n tl.store(global_accu_ptr + offs_m * stride_global_accu, global_accu, mask=offs_m < num_tokens)\n tl.store(global_entropy_b_ptr + offs_m * stride_global_entropy_b, global_entropy_b, mask=offs_m < num_tokens)\n\n\n@triton.autotune(configs=[triton.Config({\"BLOCK_SIZE_M\": 16})], key=[\"num_tokens\"])\n@triton.jit\ndef efficient_entropy_triton_epilogue_tp_update(\n num_tokens,\n logprobs_ptr,\n stride_logprobs: tl.int64,\n maximum_ptr,\n stride_maximum: tl.int64,\n accumulate_ptr,\n stride_accumulate: tl.int64,\n entropy_b_ptr,\n stride_entropy_b: tl.int64,\n entropy_ptr,\n stride_entropy: tl.int64,\n logprobs_scalar_ptr,\n reduction: int,\n BLOCK_SIZE_M: tl.constexpr,\n):\n pid_m = tl.program_id(axis=0)\n\n offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)\n\n maximum = tl.load(maximum_ptr + offs_m * stride_maximum, mask=offs_m < num_tokens)\n accumulate = tl.load(accumulate_ptr + offs_m * stride_accumulate, mask=offs_m < num_tokens)\n\n entropy_b = tl.load(entropy_b_ptr + offs_m * stride_entropy_b, mask=offs_m < num_tokens)\n entropy_b = tl.fdiv(entropy_b, accumulate)\n tl.store(entropy_b_ptr + offs_m * stride_entropy_b, entropy_b, mask=offs_m < num_tokens)\n\n entropy = tl.log(accumulate) + maximum - entropy_b\n tl.store(entropy_ptr + offs_m * stride_entropy, entropy, mask=offs_m < num_tokens)\n\n logprobs = tl.load(logprobs_ptr + offs_m * stride_logprobs, mask=offs_m < num_tokens)\n logprobs = maximum + tl.log(accumulate) - logprobs\n\n logprobs = -1 * logprobs\n if reduction == 0:\n tl.store(logprobs_ptr + offs_m * stride_logprobs, logprobs, mask=offs_m < num_tokens)\n elif reduction == 1:\n logprobs_scalar = tl.sum(logprobs, axis=0)\n tl.atomic_add(logprobs_scalar_ptr, logprobs_scalar)\n elif reduction == 2:\n logprobs_scalar = tl.sum(logprobs, axis=0) / num_tokens.to(tl.float32)\n tl.atomic_add(logprobs_scalar_ptr, logprobs_scalar)\n\n\n_dedicated_stream, _dedicated_events = None, None\n\n\ndef efficient_entropy_forward(\n hidden: torch.Tensor,\n weight: torch.Tensor,\n labels: torch.Tensor,\n reduction: typing.Optional[int] = 2,\n temperature: typing.Optional[float] = 1.0,\n dist_process_group: typing.Optional[dist.ProcessGroup] = None,\n) -> list[torch.Tensor]:\n \"\"\"\n forward host function\n \"\"\"\n assert hidden.is_cuda and weight.is_cuda and labels.is_cuda\n assert weight.device == hidden.device and labels.device == hidden.device\n assert hidden.dim() == 2 and weight.dim() == 2 and labels.dim() == 1\n assert hidden.is_contiguous() and weight.is_contiguous() and labels.is_contiguous()\n\n assert hidden.shape[0] == labels.shape[0] and hidden.shape[1] == weight.shape[1]\n\n _rank = 0 if dist_process_group is None else dist.get_rank(dist_process_group)\n _world_size = 1 if dist_process_group is None else dist.get_world_size(dist_process_group)\n\n if dist_process_group is not None and not hasattr(efficient_entropy_forward, \"_initialized\"):\n global _dedicated_stream, _dedicated_events\n _dedicated_stream = get_torch_device().Stream(hidden.device)\n _dedicated_events = [get_torch_device().Event() for _ in range(2)]\n efficient_entropy_forward._initialized = True\n\n num_tokens, hidden_size = hidden.shape\n num_tokens = labels.shape[0]\n vocab_size, hidden_size = weight.shape\n assert hidden_size % 128 == 0\n\n REDUCTION = get_entropy_reduction_enum(reduction)\n\n if REDUCTION == EntropyReductionEnum._None:\n if dist_process_group is None:\n logprobs = torch.empty((num_tokens,), device=hidden.device, dtype=torch.float32)\n else:\n logprobs = torch.zeros((num_tokens,), device=hidden.device, dtype=torch.float32)\n elif REDUCTION in (EntropyReductionEnum._Sum, EntropyReductionEnum._Mean):\n logprobs = torch.empty((), device=hidden.device, dtype=torch.float32)\n else:\n raise ValueError(f\"Invalid reduction: {reduction}\")\n\n entropy = torch.empty((num_tokens,), device=hidden.device, dtype=torch.float32)\n assert logprobs.is_contiguous() and entropy.is_contiguous()\n\n maximum = torch.empty_like(entropy)\n accumulate_and_entropy_b = torch.empty((num_tokens * 2,), device=hidden.device, dtype=torch.float32)\n accumulate_and_entropy_b_view = accumulate_and_entropy_b.view(2, num_tokens)\n accumulate = accumulate_and_entropy_b_view[0, :]\n entropy_b = accumulate_and_entropy_b_view[1, :]\n assert maximum.is_contiguous() and accumulate.is_contiguous() and entropy_b.is_contiguous()\n\n vocab_per_split = 1024\n assert vocab_per_split % 128 == 0\n num_splits = (vocab_size + vocab_per_split - 1) // vocab_per_split\n\n _max = torch.empty((num_tokens, num_splits), device=hidden.device, dtype=torch.float32)\n _accu = torch.empty((num_tokens, num_splits), device=hidden.device, dtype=torch.float32)\n _entropy_b = torch.empty((num_tokens, num_splits), device=hidden.device, dtype=torch.float32)\n\n if REDUCTION == EntropyReductionEnum._None:\n _logprobs = logprobs\n else:\n _logprobs = torch.empty((num_tokens,), device=hidden.device, dtype=torch.float32)\n\n assert _accu.is_contiguous() and _entropy_b.is_contiguous() and _max.is_contiguous()\n assert _accu.is_cuda and _entropy_b.is_cuda and _max.is_cuda\n\n if _config._use_triton:\n # 1D kernel launch, then split the tile\n def mainloop_grid(meta):\n return (triton.cdiv(num_tokens, meta[\"BLOCK_SIZE_M\"]) * num_splits,)\n\n efficient_entropy_kernel_general_mainloop[mainloop_grid](\n _rank,\n hidden,\n weight,\n labels,\n num_tokens,\n hidden_size,\n vocab_size,\n vocab_per_split,\n hidden.stride(0),\n hidden.stride(1),\n weight.stride(0),\n weight.stride(1),\n _max,\n _max.stride(0),\n _max.stride(1),\n _accu,\n _accu.stride(0),\n _accu.stride(1),\n _entropy_b,\n _entropy_b.stride(0),\n _entropy_b.stride(1),\n _logprobs,\n _logprobs.stride(0),\n logprobs,\n 1.0 / temperature,\n USE_TMA=SUPPORT_CUDA_TMA and hidden.stride(1) == 1 and weight.stride(1) == 1,\n )\n else:\n raise AssertionError(\"Triton is required for efficient entropy kernel\")\n\n # reduction on maximum and maximum_indices\n def epilogue_grid(meta):\n return (triton.cdiv(num_tokens, meta[\"BLOCK_SIZE_M\"]),)\n\n if dist_process_group is None:\n efficient_entropy_triton_kernel_epilogue[epilogue_grid](\n _max,\n _max.stride(0),\n _max.stride(1),\n num_tokens,\n num_splits,\n maximum,\n maximum.stride(0),\n _accu,\n _accu.stride(0),\n _accu.stride(1),\n accumulate,\n accumulate.stride(0),\n _entropy_b,\n _entropy_b.stride(0),\n _entropy_b.stride(1),\n entropy_b,\n entropy_b.stride(0),\n entropy,\n entropy.stride(0),\n _logprobs,\n _logprobs.stride(0),\n logprobs,\n REDUCTION,\n )\n else:\n # tensor-parallel\n _max_backup = _max.clone()\n dist.all_reduce(_max, op=dist.ReduceOp.MAX, group=dist_process_group)\n\n get_torch_device().current_stream().record_event(_dedicated_events[0])\n with get_torch_device().stream(_dedicated_stream):\n _dedicated_stream.wait_event(_dedicated_events[0])\n dist.all_reduce(_logprobs, op=dist.ReduceOp.SUM, group=dist_process_group)\n _dedicated_stream.record_event(_dedicated_events[1])\n\n efficient_entropy_triton_kernel_epilogue_tp[epilogue_grid](\n num_tokens,\n num_splits,\n _max,\n _max.stride(0),\n _max.stride(1),\n _max_backup,\n _max_backup.stride(0),\n _max_backup.stride(1),\n _accu,\n _accu.stride(0),\n _accu.stride(1),\n _entropy_b,\n _entropy_b.stride(0),\n _entropy_b.stride(1),\n maximum,\n maximum.stride(0),\n accumulate,\n accumulate.stride(0),\n entropy_b,\n entropy_b.stride(0),\n )\n get_torch_device().current_stream().wait_event(_dedicated_events[1])\n\n dist.all_reduce(accumulate_and_entropy_b, op=dist.ReduceOp.SUM, group=dist_process_group)\n\n # update logprobs & entropy\n efficient_entropy_triton_epilogue_tp_update[epilogue_grid](\n num_tokens,\n _logprobs,\n _logprobs.stride(0),\n maximum,\n maximum.stride(0),\n accumulate,\n accumulate.stride(0),\n entropy_b,\n entropy_b.stride(0),\n entropy,\n entropy.stride(0),\n logprobs,\n REDUCTION,\n )\n\n return (logprobs, entropy, maximum, accumulate, entropy_b)\n\n\n# NOTE: merge d_weight & d_hidden here, split along M & N\n@triton.autotune(\n configs=[\n triton.Config(\n {\"BLOCK_SIZE_M\": 128, \"BLOCK_SIZE_N\": 128, \"BLOCK_SIZE_K\": 32, \"GROUP_SIZE_M\": 16},\n num_stages=3,\n num_warps=8,\n )\n ],\n key=[\"num_tokens\", \"hidden_size\", \"vocab_size\"],\n)\n@triton.jit\ndef efficient_entropy_backward_kernel_general_mainloop_MN(\n num_tokens: int,\n hidden_size: int,\n vocab_size: int,\n rank: int,\n hidden_ptr,\n stride_hidden_m: tl.int64,\n stride_hidden_k: tl.int64,\n weight_ptr,\n stride_weight_n: tl.int64,\n stride_weight_k: tl.int64,\n labels_ptr,\n stride_labels: tl.int64,\n maximum_ptr,\n stride_maximum: tl.int64,\n accu_ptr,\n stride_accu: tl.int64,\n d_entropy_ptr,\n stride_d_entropy: tl.int64,\n d_logprobs_ptr,\n stride_d_logprobs: tl.int64,\n reduction: int,\n entropy_b_ptr,\n stride_entropy_b: tl.int64,\n d_hidden_ptr,\n stride_d_hidden_m: tl.int64,\n stride_d_hidden_k: tl.int64,\n d_weight_ptr,\n stride_d_weight_n: tl.int64,\n stride_d_weight_k: tl.int64,\n rcp_temperature: tl.float32,\n BLOCK_SIZE_M: tl.constexpr,\n BLOCK_SIZE_N: tl.constexpr,\n BLOCK_SIZE_K: tl.constexpr,\n GROUP_SIZE_M: tl.constexpr,\n USE_TMA: tl.constexpr,\n):\n \"\"\"\n backward mainloop, where d_logits & d_hidden & d_weight are fused\n \"\"\"\n # block swizzling\n # pid = tl.program_id(axis=0)\n # num_pid_m = tl.cdiv(num_tokens, BLOCK_SIZE_M)\n # pid_m = pid % num_pid_m\n # pid_n = pid // num_pid_m\n\n pid = tl.program_id(axis=0)\n num_pid_m = tl.cdiv(num_tokens, BLOCK_SIZE_M)\n num_pid_n = tl.cdiv(vocab_size, BLOCK_SIZE_N)\n num_pid_in_group = GROUP_SIZE_M * num_pid_n\n group_id = pid // num_pid_in_group\n first_pid_m = group_id * GROUP_SIZE_M\n group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)\n pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)\n pid_n = (pid % num_pid_in_group) // group_size_m\n\n start_offs_am = pid_m * BLOCK_SIZE_M\n offs_am = start_offs_am + tl.arange(0, BLOCK_SIZE_M)\n start_offs_bn = pid_n * BLOCK_SIZE_N\n offs_bn = start_offs_bn + tl.arange(0, BLOCK_SIZE_N)\n offs_k = tl.arange(0, BLOCK_SIZE_K)\n if USE_TMA:\n # using TMA and device-side descriptor creation\n hidden_desc = tl.make_tensor_descriptor(\n hidden_ptr,\n shape=[num_tokens, hidden_size],\n strides=[stride_hidden_m, 1],\n block_shape=[BLOCK_SIZE_M, BLOCK_SIZE_K],\n )\n\n weight_desc = tl.make_tensor_descriptor(\n weight_ptr,\n shape=[vocab_size, hidden_size],\n strides=[stride_weight_n, 1],\n block_shape=[BLOCK_SIZE_N, BLOCK_SIZE_K],\n )\n\n maximum_ptrs = maximum_ptr + offs_am * stride_maximum\n maximum = tl.load(maximum_ptrs, mask=offs_am < num_tokens, other=0.0)\n accu_ptrs = accu_ptr + offs_am * stride_accu\n accu = tl.load(accu_ptrs, mask=offs_am < num_tokens, other=1e-6) # epsilon to avoid division by zero\n accu_rcp = tl.fdiv(1.0, accu)\n\n d_entropy_ptrs = d_entropy_ptr + offs_am * stride_d_entropy\n d_entropy = tl.load(d_entropy_ptrs, mask=offs_am < num_tokens, other=0.0)\n if reduction == 0: # none\n d_logprobs_ptrs = d_logprobs_ptr + offs_am * stride_d_logprobs\n d_logprobs = tl.load(d_logprobs_ptrs, mask=offs_am < num_tokens, other=0.0)\n elif reduction == 1: # sum\n d_logprobs = tl.load(d_logprobs_ptr)\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n else: # mean\n d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n d_logprobs = -1 * d_logprobs\n\n entropy_b_ptrs = entropy_b_ptr + offs_am * stride_entropy_b\n entropy_b = tl.load(entropy_b_ptrs, mask=offs_am < num_tokens, other=0.0)\n\n if not USE_TMA:\n hidden_ptrs = hidden_ptr + (offs_am[:, None] * stride_hidden_m + offs_k[None, :] * stride_hidden_k)\n # weight_ptrs = weight_ptr + (offs_k[:, None] * stride_weight_k + offs_bn[None, :] * stride_weight_n)\n weight_ptrs = weight_ptr + (offs_bn[:, None] * stride_weight_n + offs_k[None, :] * stride_weight_k)\n labels_ptrs = labels_ptr + offs_am * stride_labels\n labels = tl.load(labels_ptrs, mask=offs_am < num_tokens, other=0)\n\n d_hidden_ptrs = d_hidden_ptr + offs_am[:, None] * stride_d_hidden_m + offs_k[None, :] * stride_d_hidden_k\n # d_weight_ptrs = d_weight_ptr + offs_k[:, None] * stride_d_weight_k + offs_bn[None, :] * stride_d_weight_n\n d_weight_ptrs = d_weight_ptr + offs_bn[:, None] * stride_d_weight_n + offs_k[None, :] * stride_d_weight_k\n\n logits = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)\n for k in range(0, tl.cdiv(hidden_size, BLOCK_SIZE_K)):\n if USE_TMA:\n start_offs_k = k * BLOCK_SIZE_K\n _hidden = hidden_desc.load([start_offs_am, start_offs_k])\n _weight = weight_desc.load([start_offs_bn, start_offs_k])\n else:\n _hidden = tl.load(\n hidden_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_am[:, None] < num_tokens),\n other=0.0,\n )\n _weight = tl.load(\n weight_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_bn[:, None] < vocab_size),\n other=0.0,\n )\n hidden_ptrs += BLOCK_SIZE_K * stride_hidden_k\n weight_ptrs += BLOCK_SIZE_K * stride_weight_k\n\n logits = tl.dot(_hidden, _weight.T, logits)\n\n if not USE_TMA:\n hidden_ptrs -= hidden_size * stride_hidden_k\n weight_ptrs -= hidden_size * stride_weight_k\n\n # scale logits by temperature\n logits *= rcp_temperature\n\n exp_logits = tl.exp(logits - maximum[:, None])\n\n mask = (offs_bn + rank * vocab_size)[None, :] == labels[:, None]\n d_logits = d_logprobs[:, None] * (exp_logits * accu_rcp[:, None] - mask)\n d_logits += d_entropy[:, None] * (-exp_logits * accu_rcp[:, None]) * (logits - entropy_b[:, None])\n\n # scale d_logits by temperature\n d_logits *= rcp_temperature\n\n # loop for d_weight & d_hidden\n for k in range(0, tl.cdiv(hidden_size, BLOCK_SIZE_K)):\n start_offs_k = k * BLOCK_SIZE_K\n if USE_TMA:\n _hidden = hidden_desc.load([start_offs_am, start_offs_k])\n else:\n _hidden = tl.load(\n hidden_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_am[:, None] < num_tokens),\n other=0.0,\n )\n # _d_weight = tl.dot(tl.trans(_hidden).to(tl.float32), d_logits)\n # tl.atomic_add(d_weight_ptrs,\n # _d_weight,\n # mask=(offs_k[:, None] < hidden_size - k * BLOCK_SIZE_K) & (offs_bn[None, :] < vocab_size))\n _d_weight = tl.dot(d_logits.trans(), _hidden.to(tl.float32))\n tl.atomic_add(\n d_weight_ptrs,\n _d_weight,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_bn[:, None] < vocab_size),\n )\n\n if USE_TMA:\n _weight = weight_desc.load([start_offs_bn, start_offs_k])\n else:\n # _weight = tl.load(\n # weight_ptrs,\n # mask=(offs_k[:, None] < hidden_size - k * BLOCK_SIZE_K) & (offs_bn[None, :] < vocab_size),\n # other=0.0\n # )\n # _d_hidden = tl.dot(d_logits, tl.trans(_weight).to(tl.float32))\n _weight = tl.load(\n weight_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_bn[:, None] < vocab_size),\n other=0.0,\n )\n _d_hidden = tl.dot(d_logits, _weight.to(tl.float32))\n tl.atomic_add(\n d_hidden_ptrs,\n _d_hidden,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_am[:, None] < num_tokens),\n )\n\n if not USE_TMA:\n hidden_ptrs += BLOCK_SIZE_K * stride_hidden_k\n weight_ptrs += BLOCK_SIZE_K * stride_weight_k\n d_hidden_ptrs += BLOCK_SIZE_K * stride_d_hidden_k\n d_weight_ptrs += BLOCK_SIZE_K * stride_d_weight_k\n\n\n@triton.autotune(\n configs=[\n triton.Config(\n {\"BLOCK_SIZE_M\": 128, \"BLOCK_SIZE_N\": 128, \"BLOCK_SIZE_K\": 32, \"GROUP_SIZE_M\": 16},\n num_stages=3,\n num_warps=8,\n ),\n ],\n key=[\"num_tokens\", \"hidden_size\", \"vocab_size\"],\n)\n@triton.jit\ndef efficient_entropy_backward_kernel_d_hidden(\n num_tokens: int,\n hidden_size: int,\n vocab_size: int,\n rank: int,\n hidden_ptr,\n stride_hidden_m: tl.int64,\n stride_hidden_k: tl.int64,\n weight_ptr,\n stride_weight_n: tl.int64,\n stride_weight_k: tl.int64,\n labels_ptr,\n stride_labels: tl.int64,\n maximum_ptr,\n stride_maximum: tl.int64,\n accu_ptr,\n stride_accu: tl.int64,\n d_entropy_ptr,\n stride_d_entropy: tl.int64,\n d_logprobs_ptr,\n stride_d_logprobs: tl.int64,\n reduction: int,\n entropy_b_ptr,\n stride_entropy_b: tl.int64,\n d_hidden_ptr,\n stride_d_hidden_m: tl.int64,\n stride_d_hidden_k: tl.int64,\n rcp_temperature: tl.float32,\n BLOCK_SIZE_M: tl.constexpr,\n BLOCK_SIZE_N: tl.constexpr,\n BLOCK_SIZE_K: tl.constexpr,\n GROUP_SIZE_M: tl.constexpr,\n):\n \"\"\"\n backward d_hidden\n \"\"\"\n pid = tl.program_id(axis=0)\n num_pid_m = tl.cdiv(num_tokens, BLOCK_SIZE_M)\n pid_m = pid % num_pid_m\n pid_k = pid // num_pid_m\n\n offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)\n offs_k = tl.arange(0, BLOCK_SIZE_K)\n result_offs_k = pid_k * BLOCK_SIZE_K + offs_k\n\n maximum = tl.load(maximum_ptr + offs_m * stride_maximum, mask=offs_m < num_tokens, other=0.0)\n accu = tl.load(accu_ptr + offs_m * stride_accu, mask=offs_m < num_tokens, other=1e-6)\n accu_rcp = tl.fdiv(1.0, accu)\n d_entropy = tl.load(d_entropy_ptr + offs_m * stride_d_entropy, mask=offs_m < num_tokens, other=0.0)\n if reduction == 0:\n d_logprobs = tl.load(d_logprobs_ptr + offs_m * stride_d_logprobs, mask=offs_m < num_tokens, other=0.0)\n elif reduction == 1:\n d_logprobs = tl.load(d_logprobs_ptr)\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n else:\n d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n d_logprobs = -1 * d_logprobs\n\n entropy_b = tl.load(entropy_b_ptr + offs_m * stride_entropy_b, mask=offs_m < num_tokens, other=0.0)\n labels = tl.load(labels_ptr + offs_m * stride_labels, mask=offs_m < num_tokens, other=0)\n\n # iterate over vocab_size\n d_hidden = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_K), dtype=tl.float32)\n for n in range(0, tl.cdiv(vocab_size, BLOCK_SIZE_N)):\n offs_n = n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)\n\n hidden_ptrs = hidden_ptr + (offs_m[:, None] * stride_hidden_m + offs_k[None, :] * stride_hidden_k)\n weight_ptrs = weight_ptr + (offs_n[:, None] * stride_weight_n + offs_k[None, :] * stride_weight_k)\n\n # iterate over hidden_size to get logits\n logits = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)\n for k in range(0, tl.cdiv(hidden_size, BLOCK_SIZE_K)):\n _hidden = tl.load(\n hidden_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_m[:, None] < num_tokens),\n other=0.0,\n )\n _weight = tl.load(\n weight_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_n[:, None] < vocab_size),\n other=0.0,\n )\n\n logits = tl.dot(_hidden, _weight.trans(), logits)\n\n hidden_ptrs += BLOCK_SIZE_K * stride_hidden_k\n weight_ptrs += BLOCK_SIZE_K * stride_weight_k\n\n # scale logits by temperature\n logits *= rcp_temperature\n\n exp_logits = tl.exp(logits - maximum[:, None])\n\n mask = (offs_n + rank * vocab_size)[None, :] == labels[:, None]\n d_logits = d_logprobs[:, None] * (exp_logits * accu_rcp[:, None] - mask)\n d_logits += d_entropy[:, None] * (-exp_logits * accu_rcp[:, None]) * (logits - entropy_b[:, None])\n\n # scale d_logits\n d_logits *= rcp_temperature\n\n # calculate d_hidden\n weight_ptrs = weight_ptr + (offs_n[:, None] * stride_weight_n + result_offs_k[None, :] * stride_weight_k)\n _weight = tl.load(\n weight_ptrs, mask=(result_offs_k[None, :] < hidden_size) & (offs_n[:, None] < vocab_size), other=0.0\n )\n d_hidden = tl.dot(d_logits.to(weight_ptr.dtype.element_ty), _weight, d_hidden)\n\n # write back\n tl.store(\n d_hidden_ptr + offs_m[:, None] * stride_d_hidden_m + result_offs_k[None, :] * stride_d_hidden_k,\n d_hidden,\n mask=(offs_m[:, None] < num_tokens) & (result_offs_k[None, :] < hidden_size),\n )\n\n\n@triton.autotune(\n configs=[\n triton.Config(\n {\"BLOCK_SIZE_M\": 128, \"BLOCK_SIZE_N\": 128, \"BLOCK_SIZE_K\": 32, \"GROUP_SIZE_M\": 16},\n num_stages=3,\n num_warps=8,\n ),\n ],\n key=[\"num_tokens\", \"hidden_size\", \"vocab_size\"],\n)\n@triton.jit\ndef efficient_entropy_backward_kernel_d_weight(\n num_tokens: int,\n hidden_size: int,\n vocab_size: int,\n rank: int,\n hidden_ptr,\n stride_hidden_m: tl.int64,\n stride_hidden_k: tl.int64,\n weight_ptr,\n stride_weight_n: tl.int64,\n stride_weight_k: tl.int64,\n labels_ptr,\n stride_labels: tl.int64,\n maximum_ptr,\n stride_maximum: tl.int64,\n accu_ptr,\n stride_accu: tl.int64,\n d_entropy_ptr,\n stride_d_entropy: tl.int64,\n d_logprobs_ptr,\n stride_d_logprobs: tl.int64,\n reduction: int,\n entropy_b_ptr,\n stride_entropy_b: tl.int64,\n d_weight_ptr,\n stride_d_weight_n: tl.int64,\n stride_d_weight_k: tl.int64,\n rcp_temperature: tl.float32,\n BLOCK_SIZE_M: tl.constexpr,\n BLOCK_SIZE_N: tl.constexpr,\n BLOCK_SIZE_K: tl.constexpr,\n GROUP_SIZE_M: tl.constexpr,\n):\n pid = tl.program_id(axis=0)\n num_pid_n = tl.cdiv(vocab_size, BLOCK_SIZE_N)\n pid_n = pid % num_pid_n\n pid_k = pid // num_pid_n\n\n offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)\n offs_k = tl.arange(0, BLOCK_SIZE_K)\n result_offs_k = pid_k * BLOCK_SIZE_K + offs_k\n\n d_weight = tl.zeros((BLOCK_SIZE_N, BLOCK_SIZE_K), dtype=tl.float32)\n for m in range(0, tl.cdiv(num_tokens, BLOCK_SIZE_M)):\n offs_m = m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)\n\n maximum = tl.load(maximum_ptr + offs_m * stride_maximum, mask=offs_m < num_tokens, other=0.0)\n accu = tl.load(accu_ptr + offs_m * stride_accu, mask=offs_m < num_tokens, other=1e-6)\n accu_rcp = tl.fdiv(1.0, accu)\n d_entropy = tl.load(d_entropy_ptr + offs_m * stride_d_entropy, mask=offs_m < num_tokens, other=0.0)\n if reduction == 0:\n d_logprobs = tl.load(d_logprobs_ptr + offs_m * stride_d_logprobs, mask=offs_m < num_tokens, other=0.0)\n elif reduction == 1:\n d_logprobs = tl.load(d_logprobs_ptr)\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n else:\n d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n d_logprobs = -1 * d_logprobs\n\n entropy_b = tl.load(entropy_b_ptr + offs_m * stride_entropy_b, mask=offs_m < num_tokens, other=0.0)\n labels = tl.load(labels_ptr + offs_m * stride_labels, mask=offs_m < num_tokens, other=0)\n\n hidden_ptrs = hidden_ptr + (offs_m[:, None] * stride_hidden_m + offs_k[None, :] * stride_hidden_k)\n weight_ptrs = weight_ptr + (offs_n[:, None] * stride_weight_n + offs_k[None, :] * stride_weight_k)\n\n logits = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)\n for k in range(0, tl.cdiv(hidden_size, BLOCK_SIZE_K)):\n _hidden = tl.load(\n hidden_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_m[:, None] < num_tokens),\n other=0.0,\n )\n _weight = tl.load(\n weight_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_n[:, None] < vocab_size),\n other=0.0,\n )\n\n logits = tl.dot(_hidden, _weight.trans(), logits)\n\n hidden_ptrs += BLOCK_SIZE_K * stride_hidden_k\n weight_ptrs += BLOCK_SIZE_K * stride_weight_k\n\n logits *= rcp_temperature\n\n exp_logits = tl.exp(logits - maximum[:, None])\n\n mask = (offs_n + rank * vocab_size)[None, :] == labels[:, None]\n d_logits = d_logprobs[:, None] * (exp_logits * accu_rcp[:, None] - mask)\n d_logits += d_entropy[:, None] * (-exp_logits * accu_rcp[:, None]) * (logits - entropy_b[:, None])\n\n d_logits *= rcp_temperature\n\n hidden_ptrs = hidden_ptr + (offs_m[:, None] * stride_hidden_m + result_offs_k[None, :] * stride_hidden_k)\n _hidden = tl.load(\n hidden_ptrs, mask=(result_offs_k[None, :] < hidden_size) & (offs_m[:, None] < num_tokens), other=0.0\n )\n d_weight = tl.dot(d_logits.to(d_weight_ptr.dtype.element_ty).trans(), _hidden, d_weight)\n\n # write back\n tl.store(\n d_weight_ptr + offs_n[:, None] * stride_d_weight_n + result_offs_k[None, :] * stride_d_weight_k,\n d_weight,\n mask=(offs_n[:, None] < vocab_size) & (result_offs_k[None, :] < hidden_size),\n )\n\n\n# NOTE: split tile from d_logits' perspective\n@triton.autotune(\n configs=[\n triton.Config(\n {\"BLOCK_SIZE_M\": 128, \"BLOCK_SIZE_N\": 256, \"BLOCK_SIZE_K\": 32, \"GROUP_SIZE_M\": 16},\n num_stages=3,\n num_warps=8,\n ),\n ],\n key=[\"num_tokens\", \"hidden_size\", \"vocab_size\"],\n)\n@triton.jit\ndef efficient_entropy_backward_kernel_general_d_logits(\n num_tokens: int,\n hidden_size: int,\n vocab_size: int,\n rank: int,\n hidden_ptr,\n stride_hidden_m: tl.int64,\n stride_hidden_k: tl.int64,\n weight_ptr,\n stride_weight_n: tl.int64,\n stride_weight_k: tl.int64,\n labels_ptr,\n stride_labels: tl.int64,\n maximum_ptr,\n stride_maximum: tl.int64,\n accu_ptr,\n stride_accu: tl.int64,\n d_entropy_ptr,\n stride_d_entropy: tl.int64,\n d_logprobs_ptr,\n stride_d_logprobs: tl.int64,\n reduction: int,\n entropy_b_ptr,\n stride_entropy_b,\n d_logits_ptr,\n stride_d_logits_m: tl.int64,\n stride_d_logits_n: tl.int64,\n rcp_temperature: tl.float32,\n BLOCK_SIZE_M: tl.constexpr,\n BLOCK_SIZE_N: tl.constexpr,\n BLOCK_SIZE_K: tl.constexpr,\n GROUP_SIZE_M: tl.constexpr,\n USE_TMA: tl.constexpr,\n):\n \"\"\"\n backward d_logits\n \"\"\"\n # block swizzling\n # pid = tl.program_id(axis=0)\n # num_pid_m = tl.cdiv(num_tokens, BLOCK_SIZE_M)\n # pid_m = pid % num_pid_m\n # pid_n = pid // num_pid_m\n\n pid = tl.program_id(axis=0)\n num_pid_m = tl.cdiv(num_tokens, BLOCK_SIZE_M)\n num_pid_n = tl.cdiv(vocab_size, BLOCK_SIZE_N)\n num_pid_in_group = GROUP_SIZE_M * num_pid_n\n group_id = pid // num_pid_in_group\n first_pid_m = group_id * GROUP_SIZE_M\n group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)\n pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)\n pid_n = (pid % num_pid_in_group) // group_size_m\n\n start_offs_am = pid_m * BLOCK_SIZE_M\n offs_am = start_offs_am + tl.arange(0, BLOCK_SIZE_M)\n start_offs_bn = pid_n * BLOCK_SIZE_N\n offs_bn = start_offs_bn + tl.arange(0, BLOCK_SIZE_N)\n offs_k = tl.arange(0, BLOCK_SIZE_K)\n\n maximum_ptrs = maximum_ptr + offs_am * stride_maximum\n maximum = tl.load(maximum_ptrs, mask=offs_am < num_tokens, other=0.0)\n accu_ptrs = accu_ptr + offs_am * stride_accu\n accu = tl.load(accu_ptrs, mask=offs_am < num_tokens, other=1e-6) # epsilon to avoid division by zero\n accu_rcp = tl.fdiv(1.0, accu)\n\n d_entropy_ptrs = d_entropy_ptr + offs_am * stride_d_entropy\n d_entropy = tl.load(d_entropy_ptrs, mask=offs_am < num_tokens, other=0.0)\n if reduction == 0: # none\n d_logprobs_ptrs = d_logprobs_ptr + offs_am * stride_d_logprobs\n d_logprobs = tl.load(d_logprobs_ptrs, mask=offs_am < num_tokens, other=0.0)\n elif reduction == 1: # sum\n d_logprobs = tl.load(d_logprobs_ptr)\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n else: # mean\n d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n d_logprobs = -1 * d_logprobs\n\n entropy_b_ptrs = entropy_b_ptr + offs_am * stride_entropy_b\n entropy_b = tl.load(entropy_b_ptrs, mask=offs_am < num_tokens, other=0.0)\n\n labels_ptrs = labels_ptr + offs_am * stride_labels\n labels = tl.load(labels_ptrs, mask=offs_am < num_tokens, other=0)\n\n logits = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)\n\n if USE_TMA:\n # using TMA and device-side descriptor creation\n hidden_desc = tl.make_tensor_descriptor(\n hidden_ptr,\n shape=[num_tokens, hidden_size],\n strides=[stride_hidden_m, 1],\n block_shape=[BLOCK_SIZE_M, BLOCK_SIZE_K],\n )\n weight_desc = tl.make_tensor_descriptor(\n weight_ptr,\n shape=[vocab_size, hidden_size],\n strides=[stride_weight_n, 1],\n block_shape=[BLOCK_SIZE_N, BLOCK_SIZE_K],\n )\n else:\n hidden_ptrs = hidden_ptr + (offs_am[:, None] * stride_hidden_m + offs_k[None, :] * stride_hidden_k)\n # weight_ptrs = weight_ptr + (offs_k[:, None] * stride_weight_k + offs_bn[None, :] * stride_weight_n)\n weight_ptrs = weight_ptr + (offs_bn[:, None] * stride_weight_n + offs_k[None, :] * stride_weight_k)\n\n for k in range(0, tl.cdiv(hidden_size, BLOCK_SIZE_K)):\n if USE_TMA:\n start_offs_k = k * BLOCK_SIZE_K\n _hidden = hidden_desc.load([start_offs_am, start_offs_k])\n _weight = weight_desc.load([start_offs_bn, start_offs_k])\n else:\n _hidden = tl.load(\n hidden_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_am[:, None] < num_tokens),\n other=0.0,\n )\n _weight = tl.load(\n weight_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_bn[:, None] < vocab_size),\n other=0.0,\n )\n hidden_ptrs += BLOCK_SIZE_K * stride_hidden_k\n weight_ptrs += BLOCK_SIZE_K * stride_weight_k\n logits = tl.dot(_hidden, _weight.T, logits)\n\n if not USE_TMA:\n hidden_ptrs -= hidden_size * stride_hidden_k\n weight_ptrs -= hidden_size * stride_weight_k\n\n # scale logits by temperature\n logits *= rcp_temperature\n\n exp_logits = tl.exp(logits - maximum[:, None])\n\n mask = (offs_bn + rank * vocab_size)[None, :] == labels[:, None]\n d_logits = d_logprobs[:, None] * (exp_logits * accu_rcp[:, None] - mask)\n d_logits += d_entropy[:, None] * (-exp_logits * accu_rcp[:, None]) * (logits - entropy_b[:, None])\n\n # scale d_logits by temperature\n d_logits *= rcp_temperature\n\n # store d_logits\n d_logits_ptrs = d_logits_ptr + offs_am[:, None] * stride_d_logits_m + offs_bn[None, :] * stride_d_logits_n\n tl.store(\n d_logits_ptrs,\n d_logits, # will be implicitly converted to d_logits_ptrs.dtype.element_ty\n mask=(offs_am[:, None] < num_tokens) & (offs_bn[None, :] < vocab_size),\n )\n\n\n@triton.autotune(\n configs=[\n triton.Config(\n {\"BLOCK_SIZE_M\": 128, \"BLOCK_SIZE_N\": 256, \"BLOCK_SIZE_K\": 32, \"GROUP_SIZE_M\": 16},\n num_stages=3,\n num_warps=8,\n ),\n ],\n key=[\"num_tokens\", \"hidden_size\", \"vocab_size\"],\n)\n@triton.jit\ndef efficient_entropy_backward_kernel_general_d_logits_split_N(\n split_idx: int,\n num_tokens: int,\n hidden_size: int,\n vocab_size: int,\n vocab_per_split: int,\n rank: int,\n hidden_ptr,\n stride_hidden_m: tl.int64,\n stride_hidden_k: tl.int64,\n weight_ptr,\n stride_weight_n: tl.int64,\n stride_weight_k: tl.int64,\n labels_ptr,\n stride_labels: tl.int64,\n maximum_ptr,\n stride_maximum: tl.int64,\n accu_ptr,\n stride_accu: tl.int64,\n d_entropy_ptr,\n stride_d_entropy: tl.int64,\n d_logprobs_ptr,\n stride_d_logprobs: tl.int64,\n reduction: int,\n entropy_b_ptr,\n stride_entropy_b,\n d_logits_ptr,\n stride_d_logits_m: tl.int64,\n stride_d_logits_n: tl.int64,\n rcp_temperature: tl.float32,\n BLOCK_SIZE_M: tl.constexpr,\n BLOCK_SIZE_N: tl.constexpr,\n BLOCK_SIZE_K: tl.constexpr,\n GROUP_SIZE_M: tl.constexpr,\n USE_TMA: tl.constexpr,\n):\n pid = tl.program_id(axis=0)\n num_pid_m = tl.cdiv(num_tokens, BLOCK_SIZE_M)\n num_pid_n = tl.cdiv(vocab_per_split, BLOCK_SIZE_N)\n num_pid_in_group = GROUP_SIZE_M * num_pid_n\n group_id = pid // num_pid_in_group\n first_pid_m = group_id * GROUP_SIZE_M\n group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)\n pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)\n pid_n = (pid % num_pid_in_group) // group_size_m\n\n start_offs_am = pid_m * BLOCK_SIZE_M\n offs_am = start_offs_am + tl.arange(0, BLOCK_SIZE_M)\n start_offs_bn = split_idx * vocab_per_split + pid_n * BLOCK_SIZE_N\n offs_bn = start_offs_bn + tl.arange(0, BLOCK_SIZE_N)\n offs_k = tl.arange(0, BLOCK_SIZE_K)\n\n maximum = tl.load(maximum_ptr + offs_am * stride_maximum, mask=offs_am < num_tokens, other=0.0)\n accu = tl.load(accu_ptr + offs_am * stride_accu, mask=offs_am < num_tokens, other=1e-6)\n accu_rcp = tl.fdiv(1.0, accu)\n d_entropy = tl.load(d_entropy_ptr + offs_am * stride_d_entropy, mask=offs_am < num_tokens, other=0.0)\n if reduction == 0:\n d_logprobs = tl.load(d_logprobs_ptr + offs_am * stride_d_logprobs, mask=offs_am < num_tokens, other=0.0)\n elif reduction == 1:\n d_logprobs = tl.load(d_logprobs_ptr)\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n else:\n d_logprobs = tl.fdiv(tl.load(d_logprobs_ptr), num_tokens.to(tl.float32))\n d_logprobs = tl.broadcast_to(d_logprobs, (BLOCK_SIZE_M,))\n d_logprobs = -1 * d_logprobs\n entropy_b = tl.load(entropy_b_ptr + offs_am * stride_entropy_b, mask=offs_am < num_tokens, other=0.0)\n labels = tl.load(labels_ptr + offs_am * stride_labels, mask=offs_am < num_tokens, other=0)\n\n logits = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)\n\n if USE_TMA:\n # using TMA and device-side descriptor creation\n hidden_desc = tl.make_tensor_descriptor(\n hidden_ptr,\n shape=[num_tokens, hidden_size],\n strides=[stride_hidden_m, 1],\n block_shape=[BLOCK_SIZE_M, BLOCK_SIZE_K],\n )\n weight_desc = tl.make_tensor_descriptor(\n weight_ptr,\n shape=[vocab_size, hidden_size],\n strides=[stride_weight_n, 1],\n block_shape=[BLOCK_SIZE_N, BLOCK_SIZE_K],\n )\n else:\n hidden_ptrs = hidden_ptr + (offs_am[:, None] * stride_hidden_m + offs_k[None, :] * stride_hidden_k)\n weight_ptrs = weight_ptr + (offs_bn[:, None] * stride_weight_n + offs_k[None, :] * stride_weight_k)\n vocab_right_bound = min((split_idx + 1) * vocab_per_split, vocab_size)\n\n for k in range(0, tl.cdiv(hidden_size, BLOCK_SIZE_K)):\n if USE_TMA:\n start_offs_k = k * BLOCK_SIZE_K\n _hidden = hidden_desc.load([start_offs_am, start_offs_k])\n _weight = weight_desc.load([start_offs_bn, start_offs_k])\n else:\n _hidden = tl.load(\n hidden_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_am[:, None] < num_tokens),\n other=0.0,\n )\n _weight = tl.load(\n weight_ptrs,\n mask=(offs_k[None, :] < hidden_size - k * BLOCK_SIZE_K) & (offs_bn[:, None] < vocab_right_bound),\n other=0.0,\n )\n hidden_ptrs += BLOCK_SIZE_K * stride_hidden_k\n weight_ptrs += BLOCK_SIZE_K * stride_weight_k\n logits = tl.dot(_hidden, _weight.T, logits)\n\n logits *= rcp_temperature\n exp_logits = tl.exp(logits - maximum[:, None])\n\n mask = (offs_bn + rank * vocab_size)[None, :] == labels[:, None]\n d_logits = d_logprobs[:, None] * (exp_logits * accu_rcp[:, None] - mask)\n d_logits += d_entropy[:, None] * (-exp_logits * accu_rcp[:, None]) * (logits - entropy_b[:, None])\n\n d_logits *= rcp_temperature\n\n # filter d_logits with mask\n result_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)\n mask = (offs_am[:, None] < num_tokens) & (result_offs_n[None, :] < vocab_per_split)\n\n tl.store(\n d_logits_ptr + offs_am[:, None] * stride_d_logits_m + result_offs_n[None, :] * stride_d_logits_n, d_logits, mask\n )\n\n\ndef efficient_entropy_backward(\n dlogprobs: torch.Tensor,\n dentropy: torch.Tensor,\n hidden: torch.Tensor,\n weight: torch.Tensor,\n labels: torch.Tensor,\n maximum: torch.Tensor,\n acc: torch.Tensor,\n entropy_b: torch.Tensor,\n reduction: typing.Optional[int] = 2,\n should_return_fp32_grad: bool = False,\n temperature: typing.Optional[float] = 1.0,\n dist_process_group: typing.Optional[dist.ProcessGroup] = None,\n) -> list[torch.Tensor]:\n \"\"\"\n backward host function\n \"\"\"\n assert hidden.is_cuda and weight.is_cuda and labels.is_cuda\n assert weight.device == hidden.device and labels.device == hidden.device\n assert hidden.dim() == 2 and weight.dim() == 2 and labels.dim() == 1\n assert hidden.is_contiguous() and weight.is_contiguous() and labels.is_contiguous()\n assert hidden.shape[0] == labels.shape[0] and hidden.shape[1] == weight.shape[1]\n\n _rank = 0 if dist_process_group is None else dist.get_rank(dist_process_group)\n _world_size = 1 if dist_process_group is None else dist.get_world_size(dist_process_group)\n\n num_tokens, hidden_size = hidden.shape\n num_tokens = labels.shape[0]\n vocab_size, hidden_size = weight.shape\n assert hidden_size % 128 == 0\n\n REDUCTION = get_entropy_reduction_enum(reduction)\n\n if REDUCTION == EntropyReductionEnum._None:\n assert dlogprobs.shape == (num_tokens,)\n else:\n assert dlogprobs.dim() == 0\n\n assert dlogprobs.is_contiguous() and dentropy.is_contiguous()\n assert dlogprobs.is_cuda and dentropy.is_cuda\n assert dlogprobs.device == hidden.device and dlogprobs.device == dentropy.device\n assert dentropy.shape == (num_tokens,)\n\n d_hidden, d_weight = None, None\n if _config._backward == BackwardEnum._Total_Fuse_MN or should_return_fp32_grad:\n d_hidden = torch.zeros_like(hidden, dtype=torch.float32, device=hidden.device)\n d_weight = torch.zeros_like(weight, dtype=torch.float32, device=weight.device)\n else:\n d_hidden = torch.empty_like(hidden, dtype=hidden.dtype, device=hidden.device)\n d_weight = torch.empty_like(weight, dtype=hidden.dtype, device=weight.device)\n assert d_hidden.is_contiguous() and d_weight.is_contiguous()\n\n assert maximum.is_contiguous() and acc.is_contiguous()\n assert maximum.device == hidden.device and acc.device == hidden.device\n assert maximum.shape == labels.shape == acc.shape\n assert maximum.is_cuda and acc.is_cuda\n\n vocab_per_split = 1024\n assert vocab_per_split % 128 == 0\n num_splits = (vocab_size + vocab_per_split - 1) // vocab_per_split\n\n assert entropy_b.is_contiguous() and entropy_b.is_cuda\n assert entropy_b.shape == (num_tokens,)\n\n if _config._backward == BackwardEnum._Total_Fuse_MN:\n # --- Triton doesn't materialize d_logits at all. Split tiles at the perspective of d_logits.\n def mainloop_grid(meta):\n return (triton.cdiv(num_tokens, meta[\"BLOCK_SIZE_M\"]) * triton.cdiv(vocab_size, meta[\"BLOCK_SIZE_N\"]),)\n\n efficient_entropy_backward_kernel_general_mainloop_MN[mainloop_grid](\n num_tokens,\n hidden_size,\n vocab_size,\n _rank,\n hidden,\n hidden.stride(0),\n hidden.stride(1),\n weight,\n weight.stride(0),\n weight.stride(1),\n labels,\n labels.stride(0),\n maximum,\n maximum.stride(0),\n acc,\n acc.stride(0),\n dentropy,\n dentropy.stride(0),\n dlogprobs,\n dlogprobs.stride(0) if REDUCTION == EntropyReductionEnum._None else 0,\n REDUCTION,\n entropy_b,\n entropy_b.stride(0),\n d_hidden,\n d_hidden.stride(0),\n d_hidden.stride(1),\n d_weight,\n d_weight.stride(0),\n d_weight.stride(1),\n 1.0 / temperature,\n USE_TMA=SUPPORT_CUDA_TMA and hidden.stride(1) == 1 and weight.stride(1) == 1,\n )\n\n elif _config._backward == BackwardEnum._Total_Separate:\n _d_logits = torch.empty((num_tokens, vocab_size), device=hidden.device, dtype=hidden.dtype).contiguous()\n assert _d_logits.is_contiguous()\n\n if _config._use_triton:\n\n def d_logits_grid(meta):\n return (triton.cdiv(num_tokens, meta[\"BLOCK_SIZE_M\"]) * triton.cdiv(vocab_size, meta[\"BLOCK_SIZE_N\"]),)\n\n efficient_entropy_backward_kernel_general_d_logits[d_logits_grid](\n num_tokens,\n hidden_size,\n vocab_size,\n _rank,\n hidden,\n hidden.stride(0),\n hidden.stride(1),\n weight,\n weight.stride(0),\n weight.stride(1),\n labels,\n labels.stride(0),\n maximum,\n maximum.stride(0),\n acc,\n acc.stride(0),\n dentropy,\n dentropy.stride(0),\n dlogprobs,\n dlogprobs.stride(0) if REDUCTION == EntropyReductionEnum._None else 0,\n REDUCTION,\n entropy_b,\n entropy_b.stride(0),\n _d_logits,\n _d_logits.stride(0),\n _d_logits.stride(1),\n 1.0 / temperature,\n USE_TMA=SUPPORT_CUDA_TMA and hidden.stride(1) == 1 and weight.stride(1) == 1,\n )\n\n torch.matmul(_d_logits, weight, out=d_hidden)\n torch.matmul(_d_logits.T, hidden, out=d_weight)\n else:\n raise AssertionError(\"Triton is required for efficient entropy kernel\")\n\n elif _config._backward == BackwardEnum._Split_Dlogits_N:\n vocab_per_split = 9504\n num_splits = (vocab_size + vocab_per_split - 1) // vocab_per_split\n\n _d_logits = torch.empty((num_tokens, vocab_per_split), device=hidden.device, dtype=hidden.dtype).contiguous()\n assert _d_logits.is_contiguous()\n\n def d_logits_grid(meta):\n return (triton.cdiv(num_tokens, meta[\"BLOCK_SIZE_M\"]) * triton.cdiv(vocab_per_split, meta[\"BLOCK_SIZE_N\"]),)\n\n for split_idx in range(num_splits):\n efficient_entropy_backward_kernel_general_d_logits_split_N[d_logits_grid](\n split_idx,\n num_tokens,\n hidden_size,\n vocab_size,\n vocab_per_split,\n _rank,\n hidden,\n hidden.stride(0),\n hidden.stride(1),\n weight,\n weight.stride(0),\n weight.stride(1),\n labels,\n labels.stride(0),\n maximum,\n maximum.stride(0),\n acc,\n acc.stride(0),\n dentropy,\n dentropy.stride(0),\n dlogprobs,\n dlogprobs.stride(0) if REDUCTION == EntropyReductionEnum._None else 0,\n REDUCTION,\n entropy_b,\n entropy_b.stride(0),\n _d_logits,\n _d_logits.stride(0),\n _d_logits.stride(1),\n 1.0 / temperature,\n USE_TMA=SUPPORT_CUDA_TMA and hidden.stride(1) == 1 and weight.stride(1) == 1,\n )\n\n if split_idx == (num_splits - 1):\n vocab_right_bound = min((split_idx + 1) * vocab_per_split, vocab_size) - split_idx * vocab_per_split\n _d_logits = _d_logits[:, :vocab_right_bound].contiguous()\n\n if split_idx == 0:\n torch.matmul(\n _d_logits, weight[split_idx * vocab_per_split : (split_idx + 1) * vocab_per_split, :], out=d_hidden\n )\n else:\n d_hidden += torch.matmul(\n _d_logits, weight[split_idx * vocab_per_split : (split_idx + 1) * vocab_per_split, :]\n )\n torch.matmul(\n _d_logits.T, hidden, out=d_weight[split_idx * vocab_per_split : (split_idx + 1) * vocab_per_split, :]\n )\n\n elif _config._backward == BackwardEnum._Split_Dlogits_M:\n raise NotImplementedError(\"BackwardEnum._Split_Dlogits_M is not implemented yet\")\n\n return d_hidden, d_weight\n"}84{"file_name": "verl__utils__kernel__linear_cross_entropy.py", "text": "#\n# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.\n# SPDX-License-Identifier: Apache-2.0\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n#\n\n# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport typing\n\nimport torch\nimport torch.distributed as dist\n\n\nclass LinearCrossEntropy(torch.autograd.Function):\n @staticmethod\n def forward(\n ctx,\n hidden: torch.Tensor,\n weight: torch.Tensor,\n labels: torch.Tensor,\n temperature: typing.Optional[float] = 1.0,\n reduction: typing.Optional[str] = \"none\",\n dist_process_group: typing.Optional[dist.ProcessGroup] = None,\n ) -> list[torch.Tensor]:\n \"\"\"_summary_\n\n Args:\n ctx (_type_): _description_\n hidden (torch.Tensor): (batch_size, num_tokens, hidden_size) -> (batch_size * num_tokens, hidden_size)\n weight (torch.Tensor): (vocab_size, hidden_size)\n labels (torch.Tensor): (batch_size, num_tokens) -> (batch_size * num_tokens, )\n temperature (typing.Optional[float], optional): _description_. Defaults to 1.0.\n reduction (typing.Optional[str], optional): _description_. Defaults to \"none\".\n dist_process_group (typing.Optional[dist.ProcessGroup], optional): _description_. Defaults to None.\n\n Returns:\n typing.List[torch.Tensor]: _description_\n \"\"\"\n\n assert isinstance(temperature, float), f\"temperature must be a float, but got {type(temperature)}\"\n assert isinstance(reduction, str), f\"reduction must be a str, but got {type(reduction)}\"\n with torch.cuda.nvtx.range(\"LinearCrossEntropy-forward\"):\n from . import kernels\n\n REDUCTION = kernels.get_entropy_reduction_enum_number(reduction.lower())\n\n original_hidden_shape = hidden.shape\n if len(hidden.shape) != 2:\n hidden = hidden.view(-1, hidden.shape[-1]) # (batch_size * num_tokens, hidden_size)\n if len(labels.shape) != 1:\n labels = labels.view(-1)\n\n logprobs, entropy, _maximum, _accumulate, _entropy_b = kernels.efficient_entropy_forward(\n hidden, weight, labels, REDUCTION, temperature, dist_process_group\n )\n\n ctx.save_for_backward(hidden, weight, labels, _maximum, _accumulate, _entropy_b)\n ctx.original_hidden_shape = original_hidden_shape\n ctx.REDUCTION = REDUCTION\n ctx.dist_process_group = dist_process_group\n ctx.should_return_fp32_grad = False\n ctx.temperature = temperature\n return logprobs, entropy\n\n @staticmethod\n def backward(ctx, dlogprobs: torch.Tensor, dentropy: torch.Tensor) -> list[torch.Tensor]:\n from . import kernels\n\n with torch.cuda.nvtx.range(\"LinearCrossEntropy-backward\"):\n (hidden, weight, labels, _maximum, _accumulate, _entropy_b) = ctx.saved_tensors\n REDUCTION = ctx.REDUCTION\n dist_process_group = ctx.dist_process_group\n should_return_fp32_grad = ctx.should_return_fp32_grad\n temperature = ctx.temperature\n\n d_hidden, d_weight = kernels.efficient_entropy_backward(\n dlogprobs,\n dentropy,\n hidden,\n weight,\n labels,\n _maximum,\n _accumulate,\n _entropy_b,\n REDUCTION,\n should_return_fp32_grad,\n temperature,\n dist_process_group,\n )\n d_hidden = d_hidden.view(ctx.original_hidden_shape)\n\n return (d_hidden, d_weight, None, None, None, None)\n\n\nlinear_cross_entropy = LinearCrossEntropy.apply\n"}85{"file_name": "verl__utils__logger__aggregate_logger.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nA Ray logger will receive logging info from different processes.\n\"\"\"\n\nimport datetime\nimport logging\nimport numbers\nimport pprint\n\nimport torch\n\n\ndef concat_dict_to_str(dict: dict, step):\n output = [f\"step:{step}\"]\n for k, v in dict.items():\n if isinstance(v, numbers.Number):\n output.append(f\"{k}:{pprint.pformat(v)}\")\n output_str = \" - \".join(output)\n return output_str\n\n\nclass LocalLogger:\n \"\"\"\n A local logger that logs messages to the console.\n\n Args:\n print_to_console (bool): Whether to print to the console.\n \"\"\"\n\n def __init__(self, print_to_console=True):\n self.print_to_console = print_to_console\n\n def flush(self):\n pass\n\n def log(self, data, step):\n if self.print_to_console:\n print(concat_dict_to_str(data, step=step), flush=True)\n\n\nclass DecoratorLoggerBase:\n \"\"\"\n Base class for all decorators that log messages.\n\n Args:\n role (str): The role (the name) of the logger.\n logger (logging.Logger): The logger instance to use for logging.\n level (int): The logging level.\n rank (int): The rank of the process.\n log_only_rank_0 (bool): If True, only log for rank 0.\n \"\"\"\n\n def __init__(\n self, role: str, logger: logging.Logger = None, level=logging.DEBUG, rank: int = 0, log_only_rank_0: bool = True\n ):\n self.role = role\n self.logger = logger\n self.level = level\n self.rank = rank\n self.log_only_rank_0 = log_only_rank_0\n self.logging_function = self.log_by_logging\n if logger is None:\n self.logging_function = self.log_by_print\n\n def log_by_print(self, log_str):\n if not self.log_only_rank_0 or self.rank == 0:\n print(f\"{self.role} {log_str}\", flush=True)\n\n def log_by_logging(self, log_str):\n if self.logger is None:\n raise ValueError(\"Logger is not initialized\")\n if not self.log_only_rank_0 or self.rank == 0:\n self.logger.log(self.level, f\"{self.role} {log_str}\")\n\n\ndef print_rank_0(message):\n \"\"\"If distributed is initialized, print only on rank 0.\"\"\"\n if torch.distributed.is_initialized():\n if torch.distributed.get_rank() == 0:\n print(message, flush=True)\n else:\n print(message, flush=True)\n\n\ndef print_with_rank(message: str, rank: int = 0, log_only_rank_0: bool = False):\n \"\"\"_summary_\n Print a message with rank information.\n This function prints the message only if `log_only_rank_0` is False or if the rank is 0.\n\n Args:\n message (str): _description_\n rank (int, optional): _description_. Defaults to 0.\n log_only_rank_0 (bool, optional): _description_. Defaults to False.\n \"\"\"\n if not log_only_rank_0 or rank == 0:\n print(f\"[Rank {rank}] {message}\", flush=True)\n\n\ndef print_with_rank_and_timer(message: str, rank: int = 0, log_only_rank_0: bool = False):\n \"\"\"_summary_\n Print a message with rank information and a timestamp.\n This function prints the message only if `log_only_rank_0` is False or if the rank is 0.\n\n Args:\n message (str): _description_\n rank (int, optional): _description_. Defaults to 0.\n log_only_rank_0 (bool, optional): _description_. Defaults to False.\n \"\"\"\n now = datetime.datetime.now()\n message = f\"[{now.strftime('%Y-%m-%d %H:%M:%S')}] [Rank {rank}] {message}\"\n if not log_only_rank_0 or rank == 0:\n print(message, flush=True)\n\n\ndef log_with_rank(message: str, rank, logger: logging.Logger, level=logging.INFO, log_only_rank_0: bool = False):\n \"\"\"_summary_\n Log a message with rank information using a logger.\n This function logs the message only if `log_only_rank_0` is False or if the rank is 0.\n Args:\n message (str): The message to log.\n rank (int): The rank of the process.\n logger (logging.Logger): The logger instance to use for logging.\n level (int, optional): The logging level. Defaults to logging.INFO.\n log_only_rank_0 (bool, optional): If True, only log for rank 0. Defaults to False.\n \"\"\"\n if not log_only_rank_0 or rank == 0:\n logger.log(level, f\"[Rank {rank}] {message}\")\n"}86{"file_name": "verl__utils__megatron__memory.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport torch\n\nfrom verl.utils.device import get_device_id\n\n\nclass MemoryBuffer:\n def __init__(self, numel, numel_padded, dtype):\n self.numel = numel\n self.numel_padded = numel_padded\n self.dtype = dtype\n self.data = torch.zeros(self.numel_padded, dtype=self.dtype, device=get_device_id(), requires_grad=False)\n\n def zero(self):\n \"\"\"Reset the buffer to zero.\"\"\"\n self.data.zero_()\n\n def get(self, shape, start_index):\n \"\"\"Return a tensor with the input `shape` as a view into the\n 1-D data starting at `start_index`.\"\"\"\n end_index = start_index + shape.numel()\n assert end_index <= self.numel, \"requested tensor is out of the buffer range.\"\n buffer_tensor = self.data[start_index:end_index]\n buffer_tensor = buffer_tensor.view(shape)\n return buffer_tensor\n"}87{"file_name": "verl__utils__megatron__pipeline_parallel.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport torch\nfrom megatron.core import parallel_state as mpu\n\nfrom .sequence_parallel import pad_to_sequence_parallel\n\n\ndef compute_transformers_input_shapes(batches, meta_info):\n from flash_attn.bert_padding import unpad_input # flash 2 is a must for Megatron\n\n # pre-compute input shapes for each micro-batch at each pp stage\n input_shapes = []\n for model_inputs in batches:\n input_ids = model_inputs[\"input_ids\"]\n attention_mask = model_inputs[\"attention_mask\"]\n input_ids_rmpad = unpad_input(input_ids.unsqueeze(dim=-1), attention_mask)[0] # (total_nnz, 1)\n if meta_info[\"sequence_parallel\"]:\n input_ids_rmpad = pad_to_sequence_parallel(input_ids_rmpad)\n # compute shapes for model_inputs\n input_shapes.append(\n torch.Size(\n [\n input_ids_rmpad.shape[0] // mpu.get_tensor_model_parallel_world_size(),\n 1,\n meta_info[\"hidden_size\"],\n ]\n )\n )\n else:\n # compute shapes for model_inputs\n input_shapes.append(torch.Size([input_ids_rmpad.shape[0], 1, meta_info[\"hidden_size\"]]))\n return input_shapes\n\n\ndef make_batch_generator(batches, vpp_size):\n \"\"\"\n Creates a batch generator suitable for Megatron pipeline parallelism,\n handling virtual pipeline parallelism (VPP).\n\n If VPP is used (vpp_size > 1), it duplicates the batch iterator for each\n virtual pipeline stage. Otherwise, it returns a single iterator.\n\n Args:\n batches: An iterable (e.g., list) of micro-batches.\n vpp_size (int): The virtual pipeline model parallel size.\n\n Returns:\n An iterator or a list of iterators over the micro-batches.\n \"\"\"\n if vpp_size > 1:\n # has vpp\n batch_generator = [batches] * vpp_size # number of vpp chunks\n batch_generator = [iter(b) for b in batch_generator]\n else:\n # no vpp\n batch_generator = iter(batches)\n return batch_generator\n"}88{"file_name": "verl__utils__megatron__router_replay_patch.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport warnings\nfrom enum import Enum\n\nimport torch\n\ntry:\n from megatron.core.transformer.moe.moe_utils import (\n apply_router_token_dropping,\n compute_routing_scores_for_aux_loss,\n group_limited_topk,\n )\n from megatron.core.transformer.moe.token_dispatcher import MoEAlltoAllTokenDispatcher\nexcept ImportError:\n warnings.warn(\"NPU not support router replay for now.\", stacklevel=2)\n MoEAlltoAllTokenDispatcher = None\nfrom megatron.core.transformer.moe.router import TopKRouter\nfrom megatron.core.transformer.transformer_config import TransformerConfig\n\n# https://github.com/THUDM/slime/blob/main/slime/utils/routing_replay.py\n\n\nclass RouterReplayAction(Enum):\n RECORD = \"record\"\n REPLAY_FORWARD = \"replay_forward\"\n REPLAY_BACKWARD = \"replay_backward\"\n\n\nclass RouterReplay:\n \"\"\"\n A class to manage the recording and replaying of MoE routing decisions.\n It holds all router instances and provides static methods to globally\n control recording and replaying.\n \"\"\"\n\n # Static variable to hold all router instances, one per MoE layer.\n router_instances = []\n\n @staticmethod\n def set_replay_data(all_layers_topk_indices: list):\n \"\"\"\n Distributes the topk indices for all layers to their respective RouterReplay instances.\n :param all_layers_topk_indices: A list of tensors, where each tensor contains the\n topk indices for a specific layer. The order\n must match the instantiation order of the routers.\n \"\"\"\n if len(all_layers_topk_indices) != len(RouterReplay.router_instances):\n raise ValueError(\n f\"The number of replay tensors ({len(all_layers_topk_indices)}) \"\n f\"does not match the number of router instances ({len(RouterReplay.router_instances)}).\"\n )\n for i, router_instance in enumerate(RouterReplay.router_instances):\n router_instance.set_target_indices(all_layers_topk_indices[i])\n\n @staticmethod\n def get_recorded_data() -> list:\n \"\"\"\n Collects the recorded topk indices from all RouterReplay instances.\n :return: A list of tensors, each containing the recorded topk indices for a layer.\n \"\"\"\n return [router.get_recorded_indices() for router in RouterReplay.router_instances]\n\n @staticmethod\n def clear_global_indices():\n \"\"\"Clears the recorded and target topk indices in all instances.\"\"\"\n for router in RouterReplay.router_instances:\n router.clear_indices()\n\n def __init__(self):\n \"\"\"Initializes a RouterReplay instance for a specific layer.\"\"\"\n self.target_topk_idx = None # For replay\n self.recorded_topk_idx = None # For recording\n self.router_replay_action = None # Router replay action for this layer\n self.replay_backward_list = [] # List of tensors for backward pass replay\n self.layer_number = None # Global layer index if available\n RouterReplay.router_instances.append(self)\n\n def set_target_indices(self, topk_indices: torch.Tensor):\n \"\"\"Sets the target topk indices for replay.\"\"\"\n self.target_topk_idx = topk_indices\n self.replay_backward_list.append(topk_indices)\n\n def get_recorded_indices(self):\n \"\"\"Returns the recorded topk indices.\"\"\"\n return self.recorded_topk_idx\n\n def record_indices(self, topk_indices: torch.Tensor):\n \"\"\"Records the topk indices.\"\"\"\n self.recorded_topk_idx = topk_indices\n\n def clear_indices(self):\n \"\"\"Clears the recorded and target topk indices.\"\"\"\n self.recorded_topk_idx = None\n self.target_topk_idx = None\n self.replay_backward_list = []\n\n def set_router_replay_action(self, router_replay_action: RouterReplayAction):\n \"\"\"Sets the router replay action for this layer.\"\"\"\n self.router_replay_action = router_replay_action\n\n def clear_router_replay_action(self):\n \"\"\"Clears the router replay action for this layer.\"\"\"\n self.router_replay_action = None\n\n @staticmethod\n def set_global_router_replay_action(router_replay_action: RouterReplayAction):\n \"\"\"Sets the router replay action for all router instances.\"\"\"\n for router in RouterReplay.router_instances:\n router.set_router_replay_action(router_replay_action)\n\n @staticmethod\n def clear_global_router_replay_action():\n \"\"\"Clears the router replay action for all router instances.\"\"\"\n for router in RouterReplay.router_instances:\n router.clear_router_replay_action()\n\n\ndef _patched_topk_routing_with_score_function(\n logits: torch.Tensor,\n topk: int,\n use_pre_softmax: bool,\n num_groups: int,\n group_topk: int,\n score_function: str,\n expert_bias: torch.Tensor,\n fused: bool,\n router_replay: RouterReplay,\n scaling_factor: float,\n):\n \"\"\"\n Patched version of topk_routing_with_score_function that supports router replay.\n \"\"\"\n num_tokens, num_experts = logits.shape\n\n def _compute_topk(scores, topk, num_groups=None, group_topk=None):\n if group_topk:\n return group_limited_topk(\n scores=scores,\n topk=topk,\n num_tokens=num_tokens,\n num_experts=num_experts,\n num_groups=num_groups,\n group_topk=group_topk,\n )\n else:\n return torch.topk(scores, k=topk, dim=1)\n\n def compute_topk(scores, topk, num_groups=None, group_topk=None):\n # Default behavior if no replay is active\n\n routing_action = router_replay.router_replay_action if router_replay is not None else None\n\n if routing_action is None:\n return _compute_topk(scores, topk, num_groups=num_groups, group_topk=group_topk)\n\n if routing_action == RouterReplayAction.RECORD:\n probs, top_indices = _compute_topk(scores, topk, num_groups=num_groups, group_topk=group_topk)\n if router_replay is not None:\n router_replay.record_indices(top_indices)\n return probs, top_indices\n\n elif routing_action == RouterReplayAction.REPLAY_FORWARD:\n if router_replay is None or router_replay.target_topk_idx is None:\n # Fallback if replay data is not available\n return _compute_topk(scores, topk, num_groups=num_groups, group_topk=group_topk)\n\n # Use the provided indices for replay\n top_indices = router_replay.target_topk_idx\n # Ensure indices are on the correct device\n top_indices = top_indices.to(scores.device)\n # Gather the scores for the replayed indices to get the probabilities\n probs = scores.gather(1, top_indices)\n return probs, top_indices\n elif routing_action == RouterReplayAction.REPLAY_BACKWARD:\n if router_replay is None or not router_replay.replay_backward_list:\n # Fallback if replay data is not available\n return _compute_topk(scores, topk, num_groups=num_groups, group_topk=group_topk)\n\n # Use the last recorded indices for backward replay\n top_indices = router_replay.replay_backward_list.pop(0)\n # Ensure indices are on the correct device\n top_indices = top_indices.to(scores.device)\n # Gather the scores for the replayed indices to get the probabilities\n probs = scores.gather(1, top_indices)\n return probs, top_indices\n else: # Unknown action, fallback\n return _compute_topk(scores, topk, num_groups=num_groups, group_topk=group_topk)\n\n if score_function == \"softmax\":\n if use_pre_softmax:\n scores = torch.softmax(logits, dim=-1, dtype=torch.float32).type_as(logits)\n probs, top_indices = compute_topk(scores, topk, num_groups, group_topk)\n else:\n scores, top_indices = compute_topk(logits, topk, num_groups, group_topk)\n probs = torch.softmax(scores, dim=-1, dtype=torch.float32).type_as(logits)\n elif score_function == \"sigmoid\":\n scores = torch.sigmoid(logits.float()).type_as(logits)\n if expert_bias is not None:\n scores_for_routing = scores + expert_bias\n _, top_indices = compute_topk(scores_for_routing, topk, num_groups, group_topk)\n scores = torch.gather(scores, dim=1, index=top_indices).type_as(logits)\n else:\n scores, top_indices = compute_topk(scores, topk, num_groups, group_topk)\n probs = scores / (scores.sum(dim=-1, keepdim=True) + 1e-20) if topk > 1 else scores\n else:\n raise ValueError(f\"Invalid score_function: {score_function}\")\n\n if scaling_factor:\n probs = probs * scaling_factor\n\n if torch.are_deterministic_algorithms_enabled():\n # build [num_tokens, num_experts] from [num_tokens, topk]\n routing_probs = torch.zeros_like(logits)\n rows = torch.arange(num_tokens, device=logits.device).unsqueeze(1)\n routing_probs.index_put_((rows, top_indices), probs, accumulate=False)\n\n routing_map = torch.zeros_like(logits, dtype=logits.dtype)\n routing_map.index_put_((rows, top_indices), torch.ones_like(probs, dtype=routing_map.dtype), accumulate=False)\n routing_map = routing_map.bool()\n else:\n # TODO Try using element-wise operations instead of scatter?\n routing_probs = torch.zeros_like(logits).scatter(1, top_indices, probs)\n routing_map = torch.zeros_like(logits).int().scatter(1, top_indices, 1).bool()\n\n return routing_probs, routing_map\n\n\ndef patched_routing(self, logits: torch.Tensor, *args, **kwargs):\n \"\"\"Top-k routing function\n\n Args:\n logits (torch.Tensor): Logits tensor after gating.\n\n Returns:\n probs (torch.Tensor): The probabilities of token to experts assignment.\n routing_map (torch.Tensor): The mapping of token to experts assignment,\n with shape [num_tokens, num_experts].\n \"\"\"\n seq_length, bsz = logits.shape[:2]\n logits = logits.view(-1, self.config.num_moe_experts)\n\n # Apply Z-Loss\n logits = self.apply_z_loss(logits)\n\n # Calculate probs and routing_map for token dispatching\n if self.routing_type == \"sinkhorn\":\n probs, routing_map = self.sinkhorn_load_balancing(logits)\n else:\n probs, routing_map = _patched_topk_routing_with_score_function(\n logits=logits,\n topk=self.topk,\n use_pre_softmax=self.config.moe_router_pre_softmax,\n num_groups=self.config.moe_router_num_groups,\n group_topk=self.config.moe_router_group_topk,\n scaling_factor=self.config.moe_router_topk_scaling_factor,\n score_function=self.score_function,\n expert_bias=self.expert_bias,\n fused=self.config.moe_router_fusion,\n router_replay=self.router_replay,\n )\n\n # Apply token dropping to probs and routing_map.\n if self.config.moe_expert_capacity_factor is not None:\n probs, routing_map = apply_router_token_dropping(\n probs,\n routing_map,\n router_topk=self.topk,\n capacity_factor=self.config.moe_expert_capacity_factor,\n drop_policy=self.config.moe_token_drop_policy,\n pad_to_capacity=self.config.moe_pad_expert_input_to_capacity,\n )\n\n # Apply each aux loss type and attach aux loss autograd function to probs\n if self.training and torch.is_grad_enabled() and self.is_aux_loss_enabled():\n # Calculate scores and routing_map for aux loss\n routing_map_for_aux_loss, scores_for_aux_loss = compute_routing_scores_for_aux_loss(\n logits, self.topk, self.score_function, fused=self.config.moe_router_fusion\n )\n probs = self._apply_aux_loss(probs, scores_for_aux_loss, routing_map_for_aux_loss)\n probs = self._apply_seq_aux_loss(probs, scores_for_aux_loss, routing_map_for_aux_loss, seq_length, bsz)\n probs = self._apply_global_aux_loss(probs, scores_for_aux_loss, routing_map_for_aux_loss)\n\n # Update expert bias and tokens_per_expert\n # Prevent extra local tokens accumulation on evaluation or activation recomputation\n if self.enable_expert_bias and torch.is_grad_enabled():\n with torch.no_grad():\n self.local_tokens_per_expert += routing_map.sum(dim=0)\n\n return probs, routing_map\n\n\ndef apply_router_replay_patch():\n \"\"\"\n Applies the monkey patch for MoE Router Replay functionality.\n This patch dynamically adds the 'enable_routing_replay' attribute to TransformerConfig\n and modifies the TopKRouter to support recording and replaying of routing decisions.\n \"\"\"\n print(\"Applying Router Replay Patch...\")\n # Clear router instances to avoid state leakage between model initializations.\n RouterReplay.router_instances.clear()\n # Step 1: Patch TransformerConfig to include the feature flag\n if not hasattr(TransformerConfig, \"enable_routing_replay\"):\n # Add class attribute with default value\n TransformerConfig.enable_routing_replay = False\n\n # Store original __init__ method\n original_tf_config_init = TransformerConfig.__init__\n\n # Define new __init__ method that safely handles enable_routing_replay parameter\n def patched_tf_config_init(self, *args, **kwargs):\n # Simple solution: remove the unknown parameter before calling original constructor\n enable_routing_replay = kwargs.pop(\"enable_routing_replay\", TransformerConfig.enable_routing_replay)\n\n # Call original constructor with remaining kwargs\n original_tf_config_init(self, *args, **kwargs)\n\n # Set the instance attribute\n self.enable_routing_replay = enable_routing_replay\n\n # Apply the patch\n TransformerConfig.__init__ = patched_tf_config_init\n\n # Step 2: Patch TopKRouter only once to ensure idempotency.\n if hasattr(TopKRouter, \"_router_replay_patched\"):\n return\n\n original_init = TopKRouter.__init__\n original_set_layer_number = TopKRouter.set_layer_number\n\n def patched_set_layer_number(self, layer_number: int):\n original_set_layer_number(self, layer_number)\n if self.router_replay is not None:\n self.router_replay.layer_number = layer_number\n\n # Step 3: Define the new __init__ method\n def patched_init(self, *args, **kwargs):\n original_init(self, *args, **kwargs)\n self.router_replay = None\n if self.config.enable_routing_replay:\n self.router_replay = RouterReplay()\n\n # Step 4: Patch MoEAlltoAllTokenDispatcher.preprocess to handle router replay\n # When router replay is enabled, duplicate indices in top_indices can cause\n # routing_map.sum() < num_tokens * topk, leading to split size mismatch in alltoall.\n if MoEAlltoAllTokenDispatcher is not None and not hasattr(MoEAlltoAllTokenDispatcher, \"_preprocess_patched\"):\n original_preprocess = MoEAlltoAllTokenDispatcher.preprocess\n\n def patched_preprocess(self, routing_map):\n \"\"\"Patched preprocess that handles router replay correctly for alltoall dispatcher.\"\"\"\n # Call original preprocess\n result = original_preprocess(self, routing_map)\n\n # Fix num_out_tokens when router replay is enabled\n if (\n getattr(self.config, \"enable_routing_replay\", False)\n and not self.drop_and_pad\n and self.config.moe_expert_capacity_factor is None\n and not (\n getattr(self.config, \"moe_router_padding_for_quantization\", None)\n or getattr(self.config, \"moe_router_padding_for_fp8\", None)\n )\n ):\n # With router replay, duplicate indices can reduce the actual routed\n # token count, so derive it from the routing map instead.\n self.num_out_tokens = int(routing_map.sum().item())\n\n return result\n\n MoEAlltoAllTokenDispatcher.preprocess = patched_preprocess\n MoEAlltoAllTokenDispatcher._preprocess_patched = True\n\n # Step 5: Apply the patches\n TopKRouter.__init__ = patched_init\n TopKRouter.routing = patched_routing\n TopKRouter.set_layer_number = patched_set_layer_number\n TopKRouter._router_replay_patched = True\n"}89{"file_name": "verl__utils__megatron__router_replay_utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\"\"\"\nRouter Replay Utilities\nUtilities for handling router replay functionality in Megatron models.\n\"\"\"\n\nimport warnings\nfrom typing import Optional\n\nimport torch\n\ntry:\n from megatron.core.pipeline_parallel.utils import is_vp_first_stage, is_vp_last_stage\nexcept ImportError:\n warnings.warn(\"NPU not support router replay for now.\", stacklevel=2)\n pass\n\nfrom megatron.core import parallel_state as mpu\nfrom megatron.core.pipeline_parallel.schedules import get_schedule_table\nfrom megatron.core.tensor_parallel import gather_from_sequence_parallel_region, scatter_to_sequence_parallel_region\nfrom megatron.core.transformer.transformer_config import TransformerConfig\nfrom megatron.core.transformer.transformer_layer import get_transformer_layer_offset\n\nfrom verl.models.mcore.util import (\n postprocess_packed_seqs,\n preprocess_packed_seqs,\n preprocess_thd_no_padding,\n)\nfrom verl.utils.device import get_device_name\nfrom verl.utils.megatron.router_replay_patch import RouterReplay, RouterReplayAction\n\ndevice_name = get_device_name()\n\n\n# from megatron.core.transformer.transformer_block import get_num_layers_to_build\ndef get_num_layers_to_build(\n config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None\n) -> int:\n \"\"\"\n Determine the number of transformer layers to build for the current pipeline stage.\n Args:\n config (TransformerConfig): Configuration object containing transformer model parameters.\n vp_stage (Optional[int]): Virtual pipeline stage number.\n pp_rank (Optional[int]): Pipeline parallel rank.\n\n Returns:\n int: The number of layers to be built for the current pipeline stage.\n \"\"\"\n # If we have a custom PP layout, straightforwardly\n # return the number of decoders in the layout array.\n if hasattr(config, \"pipeline_model_parallel_layout\") and config.pipeline_model_parallel_layout is not None:\n from megatron.core.transformer.enums import LayerType\n\n return config.pipeline_model_parallel_layout.get_num_layers_to_build(\n layer_type=LayerType.decoder, vp_stage=vp_stage\n )\n\n # Fallback for legacy tests.\n if pp_rank is None:\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n\n is_first_pp_stage = pp_rank == 0\n is_last_pp_stage = pp_rank == config.pipeline_model_parallel_size - 1\n\n if config.num_layers_in_first_pipeline_stage is not None or config.num_layers_in_last_pipeline_stage is not None:\n assert not (config.account_for_embedding_in_pipeline_split or config.account_for_loss_in_pipeline_split), (\n \" \\\n Does not support standalone embedding stage and standalone loss stage with uneven pp\"\n )\n # Number of layers to distribute over rest of pipeline stages\n layers_to_distribute = config.num_layers\n # Number of pipeline stages left for distributing transformer layers\n pipeline_stages_left = config.pipeline_model_parallel_size\n\n # If the uneven first (last) pipeline stage is enabled, remove the specified number\n # of layers to calculate the number of layers on each middle pipeline stage.\n if config.num_layers_in_first_pipeline_stage is not None:\n layers_to_distribute -= config.num_layers_in_first_pipeline_stage\n pipeline_stages_left -= 1\n\n if config.num_layers_in_last_pipeline_stage is not None:\n layers_to_distribute -= config.num_layers_in_last_pipeline_stage\n pipeline_stages_left -= 1\n\n # If pp_size <= 2, we do not have any intermediate pipeline stages, and we do not\n # need to check if the left over layers are divisible by the left over stages.\n if pipeline_stages_left > 0:\n assert layers_to_distribute % pipeline_stages_left == 0, (\n \"With uneven pipelineing the left over layers must be divisible by left over stages\"\n )\n num_layers_per_pipeline_rank = layers_to_distribute // pipeline_stages_left\n else:\n num_layers_per_pipeline_rank = 0\n\n # If the uneven first (last) pipeline stage is enabled, return the specified number\n # of layers for all virtual pipeline parallel stages within the first (last) pipeline\n # parallel stage.\n\n if is_first_pp_stage and config.num_layers_in_first_pipeline_stage is not None:\n num_layers_per_pipeline_rank = config.num_layers_in_first_pipeline_stage\n\n if is_last_pp_stage and config.num_layers_in_last_pipeline_stage is not None:\n num_layers_per_pipeline_rank = config.num_layers_in_last_pipeline_stage\n else:\n # Include the embedding layer and loss layer into pipeline parallelism partition\n num_layers = config.num_layers\n if config.account_for_embedding_in_pipeline_split:\n num_layers += 1\n\n if config.account_for_loss_in_pipeline_split:\n num_layers += 1\n\n assert num_layers % config.pipeline_model_parallel_size == 0, (\n \"num_layers should be divisible by pipeline_model_parallel_size\"\n )\n num_layers_per_pipeline_rank = num_layers // config.pipeline_model_parallel_size\n\n vp_size = config.virtual_pipeline_model_parallel_size\n if vp_size is not None and config.pipeline_model_parallel_size > 1:\n # Interleaved pipeline parallelism:\n # Number of layers in each model chunk is the number of layers in the stage,\n # divided by the number of model chunks in a stage.\n # With 8 layers, 2 stages, and 4 model chunks, we want an assignment of\n # layers to stages like (each list is a model chunk):\n # Stage 0: [0] [2] [4] [6]\n # Stage 1: [1] [3] [5] [7]\n # With 8 layers, 2 stages, and 2 virtual stages, we want an assignment of\n # layers to stages like (each list is a model chunk):\n # Stage 0: [0, 1] [4, 5]\n # Stage 1: [2, 3] [6, 7]\n\n assert num_layers_per_pipeline_rank % vp_size == 0, (\n f\"num_layers_per_pipeline_rank {num_layers_per_pipeline_rank} \\\n should be divisible by vp_size {vp_size}\"\n )\n num_layers_per_virtual_stage = num_layers_per_pipeline_rank // vp_size\n\n num_layers_to_build = num_layers_per_virtual_stage\n\n else:\n # Non-interleaved pipeline parallelism:\n # Each stage gets a contiguous set of layers.\n num_layers_to_build = num_layers_per_pipeline_rank\n\n # The embedding (or loss) layer cannot function as a standalone transformer layer\n # Reduce the number of layers to construct by 1 on the first (or last) stage if the\n # embedding (or loss) layer is included in the pipeline parallelism partition and placement.\n if config.account_for_embedding_in_pipeline_split:\n if is_vp_first_stage(vp_stage, vp_size) and is_first_pp_stage:\n num_layers_to_build -= 1\n assert num_layers_to_build >= 0, \"Not enough layers in the first virtual pipeline stage\"\n\n if config.account_for_loss_in_pipeline_split:\n if is_vp_last_stage(vp_stage, vp_size) and is_last_pp_stage:\n num_layers_to_build -= 1\n assert num_layers_to_build >= 0, \"Not enough layers in the last virtual pipeline stage\"\n\n return num_layers_to_build\n\n\ndef merge_router_topk_indices(attention_mask, input_ids, mini_layer_topk_idx_list, tf_config, vp_rank=None):\n \"\"\"\n Merge recorded router top-k indices across sequence-parallel ranks for all router instances,\n then pack/unpack them to align with the original (batch, seq_len) layout and append the result.\n\n Args:\n attention_mask (torch.Tensor): Attention mask of shape [batch_size, seq_len]. Used to determine\n the valid token positions during pack/unpack.\n input_ids (torch.Tensor): Input token IDs of shape [batch_size, seq_len]. Used together with\n attention_mask for sequence packing/unpacking.\n mini_layer_topk_idx_list (list): A Python list to which the merged top-k indices tensor will be appended.\n tf_config: Megatron/Transformer engine configuration object. Used to locate router instances for\n the current micro-batch.\n vp_rank (Optional[int]): Virtual pipeline stage rank override. If None, the current VP rank from\n Megatron parallel state will be used.\n\n Returns:\n None: The function has side effects only; it appends a tensor of shape\n [1, dynamic_bs_all, layer_num, topk] to mini_layer_topk_idx_list.\n \"\"\"\n with torch.no_grad():\n router_instances_list = RouterReplayHelper.get_micro_batch_router_list(tf_config, vp_rank)\n layers_topk_idx = []\n for router in router_instances_list:\n layers_topk_idx.append(router.recorded_topk_idx.to(torch.uint8)) # dynamic_bs, topk\n\n # layer_num, dynamic_bs, topk -> dynamic_bs, layer_num, topk\n layers_topk_idx = torch.stack(layers_topk_idx).permute(1, 0, 2).to(device_name)\n # dynamic_bs, layer_num, topk -> 1, dynamic_bs_all, layer_num, topk\n layers_topk_idx = (\n gather_from_sequence_parallel_region(layers_topk_idx, tensor_parallel_output_grad=False)\n .unsqueeze(0)\n .contiguous()\n )\n\n batch_size, seq_len = attention_mask.shape[:2]\n _, packed_seq_params = preprocess_packed_seqs(input_ids, attention_mask, pre_process=True)\n layers_topk_idx = postprocess_packed_seqs(\n layers_topk_idx, packed_seq_params, attention_mask, batch_size, seq_len, post_process=True\n )\n mini_layer_topk_idx_list.append(layers_topk_idx.cpu())\n\n\ndef set_router_replay_data(layers_topk_idx, attention_mask, tf_config, vp_rank=None):\n \"\"\"\n Scatter the packed router top-k indices back to sequence-parallel ranks and update each local\n RouterReplay instance with target indices for replay mode.\n\n This function prepares the per-layer, per-sample top-k routing decisions (recorded during an earlier\n forward) so that subsequent replay passes can follow exactly the same routing.\n\n Args:\n layers_topk_idx (torch.Tensor): Router top-k indices with shape [bs, max_seq_len, layer_num, topk].\n This should be the merged output produced by merge_router_topk_indices.\n attention_mask (torch.Tensor): Attention mask [batch_size, seq_len] used for pack/unpack alignment.\n tf_config: Megatron/Transformer engine configuration object.\n vp_rank (Optional[int]): Virtual pipeline stage rank override. If None, the current VP rank from\n Megatron parallel state will be used.\n\n Returns:\n None: The function updates internal RouterReplay instances in-place.\n \"\"\"\n with torch.no_grad():\n if layers_topk_idx.is_nested:\n layers_topk_idx_rmpad, _ = preprocess_thd_no_padding(layers_topk_idx, pre_process=True)\n else:\n layers_topk_idx_rmpad, _ = preprocess_packed_seqs(layers_topk_idx, attention_mask, pre_process=True)\n layers_topk_idx_rmpad = layers_topk_idx_rmpad.contiguous() # 1, dynamic_bs_all, layer_num, topk\n\n # 1, dynamic_bs_split, layer_num, topk\n layers_topk_idx_rmpad_split = scatter_to_sequence_parallel_region(\n layers_topk_idx_rmpad.to(device_name).squeeze(dim=0)\n ).unsqueeze(dim=0)\n\n # dynamic_bs_split, layer_num, topk -> layer_num, dynamic_bs_split, topk\n layers_topk_idx_reshape = layers_topk_idx_rmpad_split.permute(0, 2, 1, 3).squeeze(\n dim=0\n ) # layer_num, dynamic_bs_all, topk\n num_layers_in_data = layers_topk_idx_reshape.shape[0]\n use_global_layer_index = getattr(tf_config, \"num_layers\", None) == num_layers_in_data\n local_rank_info = get_current_rank_layer_info(tf_config, vp_rank)\n offset, _ = local_rank_info[\"start\"], local_rank_info[\"end\"]\n router_instances_list = RouterReplayHelper.get_micro_batch_router_list(tf_config, vp_rank)\n for i, router in enumerate(router_instances_list):\n layer_idx = None\n if use_global_layer_index:\n layer_number = getattr(router, \"layer_number\", None)\n if layer_number is not None:\n layer_idx = layer_number - 1\n if layer_idx is None:\n layer_idx = i + offset\n if layer_idx < 0 or layer_idx >= num_layers_in_data:\n raise ValueError(\n f\"router replay layer index {layer_idx} out of range for data with {num_layers_in_data} layers\"\n )\n router.set_target_indices(layers_topk_idx_reshape[layer_idx].to(torch.int64))\n\n\ndef reorder_and_merge_vpp_layers(\n micro_batch_tensor_list,\n num_microbatches: int,\n vpp_size: int,\n microbatch_group_size_per_vp_stage: int,\n) -> torch.Tensor:\n \"\"\"\n Reorder and merge per-VPP layer blocks into a contiguous layer dimension.\n\n Given a tensor shaped as [bs*vpp_size, max_token_len, layer_num_per_vpp, topk], this function:\n 1) Builds the schedule table for virtual microbatches and reorders the first dimension so that entries\n belonging to the same model chunk (VPP stage) become contiguous.\n 2) Reshapes and merges the (vpp_size, layer_num_per_vpp) into a single layer dimension, producing\n [bs, max_token_len, layer_num, topk].\n\n Args:\n micro_batch_tensor_list : the list of Input tensor.\n num_microbatches (int): Number of microbatches per pipeline stage (bs).\n vpp_size (int): Virtual pipeline parallel size (number of model chunks).\n microbatch_group_size_per_vp_stage (int): Number of consecutive microbatches processed per VPP stage.\n\n Returns:\n torch.Tensor: Output tensor of shape [bs, max_token_len, layer_num, topk].\n\n Raises:\n ValueError: If input tensor dimensionality or expected sizes do not match.\n RuntimeError: If the computed output shape is unexpected or the schedule length mismatches.\n \"\"\"\n # 1) Build schedule table: map each virtual_microbatch_id -> (microbatch_id, model_chunk_id)\n schedule_table = get_schedule_table(num_microbatches, vpp_size, microbatch_group_size_per_vp_stage)\n\n # 2) Group by model_chunk_id to build reorder indices so entries of the same chunk become contiguous along dim 0\n tensor_by_chunk = [[] for _ in range(vpp_size)]\n mini_tensor_list = []\n\n for vidx, (_mb, chunk_id) in enumerate(schedule_table):\n tensor_by_chunk[chunk_id].append(micro_batch_tensor_list[vidx])\n\n for chunk_id in range(vpp_size):\n mini_tensor_list.append(torch.cat(tensor_by_chunk[chunk_id], dim=0))\n\n out = torch.cat(mini_tensor_list, dim=2)\n return out\n\n\ndef get_current_rank_layer_info(tf_config, vp_rank=None):\n # When vp_rank is None, default to the current VP rank (or 0 if VP is disabled).\n \"\"\"Return the local layer range/count for the current process and the full assignment table.\n\n Args:\n tf_config: Configuration object used by compute_pipeline_layer_assignment.\n vp_rank (Optional[int]): Explicit virtual pipeline stage rank to query. If None, uses\n mpu.get_virtual_pipeline_model_parallel_rank() when VP is enabled; otherwise 0.\n\n Returns:\n Tuple[dict, dict]: A tuple of (local_assignment, all_assignments) where local_assignment contains\n keys {\"start\", \"end\", \"count\"} for the current (pp_rank, vp_stage).\n \"\"\"\n if vp_rank is None:\n vp_rank = 0\n num_layers_to_build = get_num_layers_to_build(tf_config, vp_stage=vp_rank)\n offset = get_transformer_layer_offset(tf_config, vp_stage=vp_rank)\n local = {}\n local[\"start\"] = offset\n local[\"end\"] = offset + num_layers_to_build\n local[\"count\"] = num_layers_to_build\n return local\n\n\ndef pp_gather(local_layers_router_map, tf_config):\n # TODO: Consider non-uniform layer allocation cases.\n \"\"\"\n Gather local router maps from all PP ranks into a global router map.\n\n Args:\n local_layers_router_map (torch.Tensor): Local router map of shape\n [bs, max_seq_len, local_num_layers, topk].\n tf_config: Configuration providing pipeline_model_parallel_size.\n\n Returns:\n torch.Tensor: Global router map of shape [bs, max_seq_len, num_layers, topk] placed on CPU.\n \"\"\"\n pp_size = tf_config.pipeline_model_parallel_size\n if pp_size <= 1:\n return local_layers_router_map\n\n pp_group = mpu.get_pipeline_model_parallel_group()\n world_size = torch.distributed.get_world_size(pp_group)\n local_layers_router_map = local_layers_router_map.to(device_name)\n layers_topk_idx_global_list = [\n torch.empty(\n size=local_layers_router_map.shape,\n dtype=local_layers_router_map.dtype,\n device=local_layers_router_map.device,\n )\n for _ in range(world_size)\n ]\n torch.distributed.all_gather(\n tensor=local_layers_router_map,\n tensor_list=layers_topk_idx_global_list,\n group=pp_group,\n async_op=False,\n )\n vp_size = tf_config.virtual_pipeline_model_parallel_size\n if vp_size is not None:\n vpp_router_map_offset = [[] for _ in range(pp_size)]\n for pp_stage in range(pp_size):\n vpp_router_map_offset[pp_stage].append(0)\n for vp_stage in range(vp_size):\n num_layers_to_build = get_num_layers_to_build(tf_config, vp_stage, pp_stage)\n vpp_router_map_offset[pp_stage].append(num_layers_to_build + vpp_router_map_offset[pp_stage][-1])\n layers_topk_idx_global = []\n for vp_stage in range(vp_size):\n for pp_stage in range(pp_size):\n piece = slice(vpp_router_map_offset[pp_stage][vp_stage], vpp_router_map_offset[pp_stage][vp_stage + 1])\n layers_topk_idx_global.append(layers_topk_idx_global_list[pp_stage][:, :, piece, :])\n global_router_map = torch.cat(layers_topk_idx_global, dim=2).to(\"cpu\")\n else:\n global_router_map = torch.cat(layers_topk_idx_global_list, dim=2).to(\"cpu\")\n\n return global_router_map\n\n\nclass RouterReplayHelper:\n \"\"\"Helper class to query router replay state and locate local RouterReplay instances.\"\"\"\n\n @staticmethod\n def get_micro_batch_router_list(tf_config, vp_rank=None):\n \"\"\"\n Return the list of RouterReplay instances corresponding to the current micro-batch and local\n (pp_rank, vp_stage) layer range.\n\n When virtual pipeline (VPP) is enabled, the local range for the PP rank is expanded to include\n all VP stages by multiplying the per-VP count by vp_size. The returned slice is taken from the\n global RouterReplay.router_instances list.\n\n Args:\n tf_config: Configuration object used to compute layer assignments.\n vp_rank (Optional[int]): Explicit virtual pipeline stage to query. If None, the current VP\n rank from Megatron parallel state is used when available.\n Returns:\n list: A contiguous sublist of RouterReplay.router_instances for the local layer range.\n \"\"\"\n vp_size = tf_config.virtual_pipeline_model_parallel_size\n if vp_size is not None:\n vp_rank = 0 if vp_rank is None else vp_rank\n offset = 0\n for pre_vp_stage in range(vp_size):\n if pre_vp_stage == vp_rank:\n break\n num_layers_to_build = get_num_layers_to_build(tf_config, pre_vp_stage)\n offset += num_layers_to_build\n else:\n offset = 0\n\n num_layers_to_build = get_num_layers_to_build(tf_config, vp_rank)\n router_instances_list = RouterReplay.router_instances[offset : offset + num_layers_to_build]\n return router_instances_list\n\n @staticmethod\n def is_r2_record_action(tf_config, vp_rank=None) -> bool:\n \"\"\"Return True if the current router_replay_action is RECORD (R2) for the local router instances.\n\n This inspects the first local RouterReplay instance's router_replay_action and compares it to\n RouterReplayAction.RECORD.\n \"\"\"\n router_instances_list = RouterReplayHelper.get_micro_batch_router_list(tf_config, vp_rank)\n return router_instances_list and router_instances_list[0].router_replay_action == RouterReplayAction.RECORD\n\n @staticmethod\n def is_replay_forward_action(tf_config, vp_rank=None) -> bool:\n \"\"\"Return True if the current router_replay_action is REPLAY_FORWARD for the local router instances.\n\n This inspects the first local RouterReplay instance's router_replay_action and compares it to\n RouterReplayAction.REPLAY_FORWARD.\n \"\"\"\n router_instances_list = RouterReplayHelper.get_micro_batch_router_list(tf_config, vp_rank)\n return (\n router_instances_list and router_instances_list[0].router_replay_action == RouterReplayAction.REPLAY_FORWARD\n )\n\n @staticmethod\n def is_replay_backward_action(tf_config, vp_rank=None) -> bool:\n \"\"\"Return True if the current router_replay_action is REPLAY_BACKWARD for the local router instances.\n\n This inspects the first local RouterReplay instance's router_replay_action and compares it to\n RouterReplayAction.REPLAY_BACKWARD.\n \"\"\"\n router_instances_list = RouterReplayHelper.get_micro_batch_router_list(tf_config, vp_rank)\n return (\n router_instances_list\n and router_instances_list[0].router_replay_action == RouterReplayAction.REPLAY_BACKWARD\n )\n"}90{"file_name": "verl__utils__megatron__sequence_parallel.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport torch\nimport torch.nn.functional as F\nfrom megatron.core import parallel_state as mpu\n\n\ndef mark_parameter_as_sequence_parallel(parameter):\n parameter.sequence_parallel = True\n\n\ndef is_sequence_parallel_param(param):\n return hasattr(param, \"sequence_parallel\") and param.sequence_parallel\n\n\ndef pad_to_sequence_parallel(unpad_tokens: torch.Tensor):\n \"\"\"pad the tokens such that the total length is a multiple of sp world size\n\n Args:\n unpad_tokens: (total_nnz, ...). Tokens after removing padding\n\n Returns:\n the padded tokens: (total_nnz + pad_size,...)\n\n \"\"\"\n total_nnz = unpad_tokens.shape[0]\n sp_world_size = mpu.get_tensor_model_parallel_world_size()\n\n pad_size = 0 if total_nnz % sp_world_size == 0 else sp_world_size - total_nnz % sp_world_size\n\n if pad_size > 0:\n if unpad_tokens.ndim == 1:\n unpad_tokens = F.pad(unpad_tokens, (0, pad_size))\n elif unpad_tokens.ndim == 2:\n unpad_tokens = F.pad(unpad_tokens, (0, 0, 0, pad_size))\n else:\n raise NotImplementedError(f\"Padding dim {unpad_tokens.ndim()} is not supported\")\n\n return unpad_tokens\n"}91{"file_name": "verl__utils__megatron__tensor_parallel.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nUtilities for using tensor_parallel in megatron\n\"\"\"\n\nfrom typing import TYPE_CHECKING\n\nimport torch\nimport torch.distributed as dist\nfrom megatron.core import parallel_state as mpu\nfrom torch.nn import init\n\nif TYPE_CHECKING:\n from megatron.core import ModelParallelConfig\n\n\ndef update_kwargs_with_config(dictionary: dict, config: \"ModelParallelConfig\"):\n dictionary[\"config\"] = config\n return dictionary\n\n\ndef get_default_kwargs_for_model_parallel_config():\n model_parallel_config_kwargs = {\n \"params_dtype\": torch.float32,\n \"use_cpu_initialization\": False,\n \"perform_initialization\": True,\n \"gradient_accumulation_fusion\": False,\n \"sequence_parallel\": False,\n }\n return model_parallel_config_kwargs\n\n\ndef get_default_model_parallel_config():\n from megatron.core import ModelParallelConfig\n\n return ModelParallelConfig(**get_default_kwargs_for_model_parallel_config())\n\n\ndef get_common_default_kwargs_for_parallel_linear():\n default_model_parallel_config = get_default_model_parallel_config()\n common_default_kwargs = {\n \"init_method\": init.xavier_normal_,\n \"stride\": 1,\n \"keep_master_weight_for_test\": False,\n \"config\": default_model_parallel_config,\n }\n return common_default_kwargs\n\n\ndef get_default_kwargs_for_column_parallel_linear():\n from megatron.core import ModelParallelConfig\n\n model_parallel_config_kwargs = get_default_kwargs_for_model_parallel_config()\n column_parallel_config_kwargs = {\n \"async_tensor_model_parallel_allreduce\": False,\n }\n model_parallel_config_kwargs.update(column_parallel_config_kwargs)\n column_default_kwargs = {\n \"config\": ModelParallelConfig(**model_parallel_config_kwargs),\n }\n common_default_kwargs = get_common_default_kwargs_for_parallel_linear()\n common_default_kwargs.update(column_default_kwargs)\n return common_default_kwargs\n\n\ndef get_default_kwargs_for_row_parallel_linear():\n common_default_kwargs = get_common_default_kwargs_for_parallel_linear()\n return common_default_kwargs\n\n\ndef get_default_kwargs_for_parallel_embedding():\n from megatron.core import ModelParallelConfig\n\n model_parallel_config_kwargs = get_default_kwargs_for_model_parallel_config()\n embedding_default_kwargs = {\n \"init_method\": init.xavier_normal_,\n \"config\": ModelParallelConfig(**model_parallel_config_kwargs),\n }\n return embedding_default_kwargs\n\n\ndef is_tensor_parallel_param(param):\n return hasattr(param, \"tensor_model_parallel\") and param.tensor_model_parallel\n\n\ndef get_tensor_parallel_partition_dim(param):\n assert is_tensor_parallel_param(param)\n return param.partition_dim\n\n\ndef get_tensor_parallel_partition_stride(param):\n assert is_tensor_parallel_param(param)\n return param.partition_stride\n\n\nclass _VocabParallelEntropy(torch.autograd.Function):\n @staticmethod\n def forward(ctx, vocab_parallel_logits: torch.Tensor) -> torch.Tensor:\n @torch.compile(dynamic=True)\n def mul_reduce(a, b):\n return (a * b).sum(dim=-1, keepdim=True)\n\n logits_max = vocab_parallel_logits.max(dim=-1, keepdim=True).values\n dist.all_reduce(logits_max, op=dist.ReduceOp.MAX, group=mpu.get_tensor_model_parallel_group())\n normalized_vocab_parallel_logits = vocab_parallel_logits - logits_max\n normalized_exp_logits = normalized_vocab_parallel_logits.exp_()\n normalized_sum_exp_logits = normalized_exp_logits.sum(dim=-1, keepdim=True)\n dist.all_reduce(normalized_sum_exp_logits, group=mpu.get_tensor_model_parallel_group())\n softmax_logits = normalized_exp_logits.div_(normalized_sum_exp_logits)\n sum_softmax_times_logits = mul_reduce(softmax_logits, vocab_parallel_logits)\n dist.all_reduce(sum_softmax_times_logits, group=mpu.get_tensor_model_parallel_group())\n entropy = logits_max + normalized_sum_exp_logits.log() - sum_softmax_times_logits\n ctx.save_for_backward(vocab_parallel_logits, softmax_logits, sum_softmax_times_logits)\n return entropy.squeeze(dim=-1)\n\n @staticmethod\n def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor:\n vocab_parallel_logits, softmax_logits, sum_softmax_times_logits = ctx.saved_tensors\n # reuse softmax_logits as grad\n vocab_parallel_logits.sub_(sum_softmax_times_logits)\n softmax_logits.mul_(vocab_parallel_logits)\n softmax_logits.mul_(grad_output.unsqueeze(dim=-1))\n # recover vocab_parallel_logits\n vocab_parallel_logits.add_(sum_softmax_times_logits)\n softmax_logits.mul_(-1)\n return softmax_logits\n\n\ndef vocab_parallel_entropy(vocab_parallel_logits: torch.Tensor) -> torch.Tensor:\n \"\"\"Compute entropy when the logits are sharded in tp ranks\n\n Args:\n vocab_parallel_logits: (total_nnz, vocab_size // tp_size)\n\n Returns: (total_nnz,)\n\n \"\"\"\n return _VocabParallelEntropy.apply(vocab_parallel_logits)\n\n\ndef vocab_parallel_log_probs_from_logits(logits, labels):\n \"\"\"TODO(zhangchi.usc1992): We may change the implementation later\"\"\"\n from megatron.core import tensor_parallel\n\n return -tensor_parallel.vocab_parallel_cross_entropy(vocab_parallel_logits=logits, target=labels)\n\n\ndef vocab_parallel_log_probs_from_logits_response_rmpad(input_ids, attention_mask, logits_rmpad, response_length):\n \"\"\"Similar to log_probs_from_logits_response_rmpad, but the logits_rmpad is now spliited across tensor parallel\n region.\n This will further reduce the peak memory usage during training\n\n Args:\n input_ids: [batch_size, seqlen]\n attention_mask: [batch_size, seqlen]\n logits_rmpad: [total_nnz, vocab_size // tp_size]\n response_length: int\n\n \"\"\"\n from flash_attn.bert_padding import pad_input, unpad_input\n\n batch_size, seqlen = input_ids.shape\n input_ids_rmpad, indices, *_ = unpad_input(input_ids.unsqueeze(-1), attention_mask=attention_mask)\n input_ids_rmpad = input_ids_rmpad.squeeze(-1)\n input_ids_rmpad_rolled = torch.roll(input_ids_rmpad, shifts=-1, dims=0)\n full_log_probs_rmpad = vocab_parallel_log_probs_from_logits(\n logits=logits_rmpad, labels=input_ids_rmpad_rolled\n ) # (total_nnz,)\n full_output = pad_input(\n hidden_states=full_log_probs_rmpad.unsqueeze(-1), indices=indices, batch=batch_size, seqlen=seqlen\n )\n output = full_output.squeeze(-1)[:, -response_length - 1 : -1] # [batch_size, response_length]\n return output\n"}92{"file_name": "verl__utils__megatron_peft_utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"Utilities for PEFT (Parameter-Efficient Fine-Tuning) of Megatron in VERL.\"\"\"\n\nimport os\nfrom pathlib import Path\nfrom typing import Iterator\n\nimport torch\n\n# Map megatron lora target modules to HF-style module names for vLLM\nMEGATRON_TO_HF_MODULES = {\n \"linear_qkv\": [\"q_proj\", \"k_proj\", \"v_proj\"],\n \"linear_proj\": [\"o_proj\"],\n \"linear_fc1\": [\"gate_proj\", \"up_proj\"],\n \"linear_fc2\": [\"down_proj\"],\n \"router\": [\"gate\"],\n # Canonical LoRA mappings\n \"linear_q\": [\"q_proj\"],\n \"linear_k\": [\"k_proj\"],\n \"linear_v\": [\"v_proj\"],\n \"linear_fc1_up\": [\"up_proj\"],\n \"linear_fc1_gate\": [\"gate_proj\"],\n # MLA mappings\n \"linear_kv_down_proj\": [\"kv_a_proj_with_mqa\"],\n \"linear_kv_up_proj\": [\"kv_b_proj\"],\n \"linear_q_down_proj\": [\"q_a_proj\"],\n \"linear_q_up_proj\": [\"q_b_proj\"],\n \"linear_q_proj\": [\"q_proj\"],\n}\n\n# Modules with stacked parameters that need .base_layer suffix in vLLM\nSTACKED_PARAMS = [\n \".q_proj.weight\",\n \".q_proj.bias\",\n \".k_proj.weight\",\n \".k_proj.bias\",\n \".v_proj.weight\",\n \".v_proj.bias\",\n \".o_proj.weight\",\n \".o_proj.bias\",\n \".gate_proj.weight\",\n \".up_proj.weight\",\n \".down_proj.weight\",\n \".mlp.gate.weight\",\n \".mlp.gate.bias\",\n \".mlp.gate.e_score_correction_bias\",\n \".kv_a_proj_with_mqa.weight\",\n \".kv_b_proj.weight\",\n \".q_a_proj.weight\",\n \".q_b_proj.weight\",\n]\n\n\ndef _get_rank_checkpoint_path(base_path: str) -> str:\n \"\"\"Get rank-specific checkpoint path following Megatron's convention.\n\n Returns path like: base_path/mp_rank_{tp:02d}_{pp:03d}_{ep:03d}/\n\n Args:\n base_path: Base checkpoint directory\n\n Returns:\n Rank-specific subdirectory path\n \"\"\"\n from megatron.core import mpu\n\n tensor_rank = mpu.get_tensor_model_parallel_rank()\n pipeline_rank = mpu.get_pipeline_model_parallel_rank()\n expert_rank = mpu.get_expert_model_parallel_rank()\n\n pipeline_parallel = mpu.get_pipeline_model_parallel_world_size() > 1\n expert_parallel = mpu.get_expert_model_parallel_world_size() > 1\n\n if not pipeline_parallel:\n rank_path = os.path.join(base_path, f\"mp_rank_{tensor_rank:02d}\")\n else:\n rank_path = os.path.join(base_path, f\"mp_rank_{tensor_rank:02d}_{pipeline_rank:03d}\")\n\n if expert_parallel:\n rank_path = rank_path + f\"_{expert_rank:03d}\"\n\n return rank_path\n\n\ndef get_adapter_state_dict(model):\n \"\"\"Extract only adapter parameters from a model.\n\n Args:\n model: PyTorch model (possibly wrapped in DDP/Float16Module)\n\n Returns:\n Dict of adapter parameter names to tensors\n \"\"\"\n from verl.utils.megatron_utils import unwrap_model\n\n # Unwrap model from DDP/Float16Module\n unwrapped = unwrap_model(model)\n if isinstance(unwrapped, list):\n unwrapped = unwrapped[0]\n\n adapter_state = {}\n for name, param in unwrapped.named_parameters():\n if \".adapter.\" in name.lower():\n adapter_state[name] = param.data.clone()\n\n return adapter_state\n\n\ndef save_adapter_checkpoint(\n model: torch.nn.Module | list[torch.nn.Module],\n checkpoint_path: str,\n rank: int = 0,\n):\n \"\"\"Save only adapter parameters to checkpoint.\n\n This is much more efficient than saving the full model when using PEFT,\n as adapters typically represent <1% of total parameters.\n\n Uses Megatron's distributed checkpoint structure: each rank saves to\n checkpoint_path/mp_rank_{tp:02d}_{pp:03d}/adapter.pt\n\n Args:\n model: Model or list of models\n checkpoint_path: Base path to save checkpoint (rank-specific subdirs created)\n rank: Process rank (used for logging only)\n \"\"\"\n\n if isinstance(model, list):\n models = model\n else:\n models = [model]\n\n # Get adapter state from first model\n adapter_state = get_adapter_state_dict(models[0])\n\n if not adapter_state:\n if rank == 0:\n print(\"Warning: No adapter parameters found to save\")\n return\n\n # Get rank-specific directory path\n Path(checkpoint_path).mkdir(parents=True, exist_ok=True)\n rank_path = _get_rank_checkpoint_path(checkpoint_path)\n adapter_file = rank_path + \"_adapter.pt\"\n\n torch.save(\n {\n \"adapter_state_dict\": adapter_state,\n },\n adapter_file,\n )\n\n if rank == 0:\n print(f\"Saved {len(adapter_state)} adapter parameters to {checkpoint_path} (distributed)\")\n\n\ndef load_adapter_checkpoint(\n model: torch.nn.Module | list[torch.nn.Module],\n checkpoint_path: str,\n strict: bool = True,\n):\n \"\"\"Load adapter parameters from checkpoint.\n\n Loads from Megatron's distributed checkpoint structure: reads from\n checkpoint_path/mp_rank_{tp:02d}_{pp:03d}/adapter.pt for each rank.\n\n Args:\n model: Model or list of models\n checkpoint_path: Base path to checkpoint directory\n strict: Whether to strictly enforce parameter name matching\n \"\"\"\n from megatron.core import mpu\n\n from verl.utils.megatron_utils import unwrap_model\n\n # Get rank-specific path\n rank_path = _get_rank_checkpoint_path(checkpoint_path)\n adapter_file = rank_path + \"_adapter.pt\"\n\n if not os.path.isfile(adapter_file):\n raise FileNotFoundError(f\"Adapter checkpoint not found: {adapter_file}\")\n\n checkpoint = torch.load(adapter_file, map_location=\"cpu\")\n adapter_state = checkpoint.get(\"adapter_state_dict\", {})\n\n if not adapter_state:\n print(\"Warning: No adapter parameters found in checkpoint\")\n return\n\n if isinstance(model, list):\n models = model\n else:\n models = [model]\n\n # Load adapter parameters into each model (for VPP, models may have multiple chunks)\n loaded_count = 0\n for m in models:\n unwrapped = unwrap_model(m)\n if isinstance(unwrapped, list):\n unwrapped = unwrapped[0]\n\n # Load parameters\n _, unexpected = unwrapped.load_state_dict(adapter_state, strict=False)\n\n if strict and unexpected:\n raise RuntimeError(f\"Error loading adapter checkpoint:\\nUnexpected keys: {unexpected}\")\n\n loaded_count += len(adapter_state)\n\n if (\n mpu.get_data_parallel_rank() == 0\n and mpu.get_tensor_model_parallel_rank() == 0\n and mpu.get_pipeline_model_parallel_rank() == 0\n ):\n print(f\"Loaded {len(adapter_state)} adapter parameters from {checkpoint_path}\")\n\n\ndef count_adapter_parameters(model):\n \"\"\"Count the number of trainable adapter parameters.\n\n Args:\n model: PyTorch model\n\n Returns:\n Tuple of (adapter_params, total_params, percentage)\n \"\"\"\n from verl.utils.megatron_utils import unwrap_model\n\n unwrapped = unwrap_model(model)\n if isinstance(unwrapped, list):\n unwrapped = unwrapped[0]\n\n adapter_params = 0\n total_params = 0\n\n for name, param in unwrapped.named_parameters():\n total_params += param.numel()\n if \"lora\" in name.lower() or \"adapter\" in name.lower():\n if param.requires_grad:\n adapter_params += param.numel()\n\n percentage = 100 * adapter_params / total_params if total_params > 0 else 0\n\n return adapter_params, total_params, percentage\n\n\ndef print_adapter_info(model):\n \"\"\"Print information about adapter parameters in the model.\"\"\"\n adapter_params, total_params, percentage = count_adapter_parameters(model)\n\n print(f\"\\n{'=' * 60}\")\n print(\"PEFT Adapter Information:\")\n print(f\" Total parameters: {total_params:,}\")\n print(f\" Adapter parameters: {adapter_params:,}\")\n print(f\" Trainable percentage: {percentage:.2f}%\")\n print(f\"{'=' * 60}\\n\")\n\n\ndef convert_megatron_to_hf_target_modules(megatron_modules: list[str]) -> list[str]:\n \"\"\"Convert megatron lora target modules to HF-style module names.\n\n Args:\n megatron_modules: List of megatron-style module names.\n\n Returns:\n List of HF-style module names with duplicates removed.\n \"\"\"\n hf_target_modules = []\n for module in megatron_modules:\n if module in MEGATRON_TO_HF_MODULES:\n hf_target_modules.extend(MEGATRON_TO_HF_MODULES[module])\n else:\n hf_target_modules.append(module)\n # Remove duplicates while preserving order\n return list(dict.fromkeys(hf_target_modules))\n\n\ndef build_peft_config_for_vllm(lora_config: dict) -> dict:\n \"\"\"Build a peft_config dict compatible with vLLM's PEFTHelper from megatron lora config.\n\n Args:\n lora_config: Megatron lora configuration dictionary.\n\n Returns:\n A dictionary compatible with vLLM's PEFTHelper.from_dict().\n \"\"\"\n from peft import TaskType\n\n target_modules = lora_config.get(\"target_modules\", [\"linear_qkv\", \"linear_proj\", \"linear_fc1\", \"linear_fc2\"])\n exclude_modules = lora_config.get(\"exclude_modules\", [])\n hf_target_modules = convert_megatron_to_hf_target_modules(target_modules)\n hf_exclude_modules = convert_megatron_to_hf_target_modules(exclude_modules)\n\n return {\n \"task_type\": TaskType.CAUSAL_LM,\n \"r\": lora_config.get(\"rank\", 0),\n \"lora_alpha\": lora_config.get(\"alpha\", 32),\n \"target_modules\": hf_target_modules,\n \"exclude_modules\": hf_exclude_modules,\n \"bias\": \"none\",\n \"lora_dropout\": lora_config.get(\"dropout\", 0.0),\n }\n\n\n# vLLM needs to target all-linear no matter about specific LoRA config\ndef add_base_layer_suffix(\n params: Iterator[tuple[str, torch.Tensor]],\n model_type: str,\n) -> Iterator[tuple[str, torch.Tensor]]:\n \"\"\"Yield param pairs with a base-layer suffix added to the param name.\n\n Args:\n params: Iterator of (param_name, tensor)\n model_type: The type of the model (e.g., \"llama\").\n \"\"\"\n stacked_params = STACKED_PARAMS\n # TODO: other models may have more special treatment, or integrate this into Megatron-Bridge\n if model_type == \"llama\":\n stacked_params = [\".embed_tokens.weight\", *STACKED_PARAMS]\n for name, param in params:\n ending_suffix = \"\"\n for suffix in stacked_params:\n if name.endswith(suffix):\n ending_suffix = suffix\n break\n if ending_suffix:\n suffix = ending_suffix.rsplit(\".\", 1)[-1]\n name = f\"{name[: -len(suffix)]}base_layer.{suffix}\"\n yield name, param\n\n\n__all__ = [\n \"get_adapter_state_dict\",\n \"save_adapter_checkpoint\",\n \"load_adapter_checkpoint\",\n \"count_adapter_parameters\",\n \"print_adapter_info\",\n \"convert_megatron_to_hf_target_modules\",\n \"build_peft_config_for_vllm\",\n \"add_base_layer_suffix\",\n]\n"}93{"file_name": "verl__utils__megatron_utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.\n# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"Pretrain utilities.\"\"\"\n\nimport gc\nimport inspect\nimport logging\nimport os\nimport warnings\nfrom dataclasses import dataclass\nfrom typing import Any\n\nimport torch\nimport torch.nn.functional as F\nfrom megatron.core import ModelParallelConfig, mpu, parallel_state, tensor_parallel\nfrom megatron.core.distributed import DistributedDataParallel as DDP\nfrom megatron.core.distributed import DistributedDataParallelConfig\nfrom megatron.core.enums import ModelType\nfrom megatron.core.optimizer import ChainedOptimizer\nfrom megatron.core.parallel_state import get_global_memory_buffer\nfrom megatron.core.transformer import MLATransformerConfig, TransformerConfig\nfrom megatron.core.transformer.module import Float16Module\nfrom megatron.core.transformer.multi_token_prediction import MTPLossLoggingHelper\nfrom megatron.core.utils import get_attr_wrapped_model\nfrom transformers import PretrainedConfig\n\nimport verl.utils.megatron.tensor_parallel as tp_utils\nfrom verl.utils.device import get_device_id, get_device_name, get_torch_device\nfrom verl.utils.fs import local_mkdir_safe\nfrom verl.utils.model import normalize_model_name\nfrom verl.utils.torch_dtypes import PrecisionType\nfrom verl.workers.config import HFModelConfig, McoreEngineConfig\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef get_model_config(model):\n return get_attr_wrapped_model(model, \"config\", allow_none=False)\n\n\ndef get_model(\n model_provider_func,\n model_type=ModelType.encoder_or_decoder,\n wrap_with_ddp=True,\n use_distributed_optimizer=True,\n transformer_config=None,\n override_ddp_config=None,\n):\n \"\"\"Build the model.\"\"\"\n # Build model.\n if (\n mpu.get_pipeline_model_parallel_world_size() > 1\n and mpu.get_virtual_pipeline_model_parallel_world_size() is not None\n ):\n assert model_type != ModelType.encoder_and_decoder, (\n \"Interleaved schedule not supported for model with both encoder and decoder\"\n )\n model = []\n has_vp_stage = inspect.signature(mpu.is_pipeline_first_stage).parameters.get(\"vp_stage\", None) is not None\n for i in range(mpu.get_virtual_pipeline_model_parallel_world_size()):\n mpu.set_virtual_pipeline_model_parallel_rank(i)\n # Set pre_process and post_process only after virtual rank is set.\n extra_kwargs = {} if not has_vp_stage else {\"ignore_virtual\": False, \"vp_stage\": i}\n pre_process = mpu.is_pipeline_first_stage(**extra_kwargs)\n post_process = mpu.is_pipeline_last_stage(**extra_kwargs)\n this_model = model_provider_func(pre_process=pre_process, post_process=post_process, vp_stage=i)\n this_model.model_type = model_type\n model.append(this_model)\n mpu.set_virtual_pipeline_model_parallel_rank(0)\n else:\n pre_process = mpu.is_pipeline_first_stage()\n post_process = mpu.is_pipeline_last_stage()\n add_encoder = True\n add_decoder = True\n assert model_type != ModelType.encoder_and_decoder, \"Model type encoder_and_decoder is not supported\"\n if model_type == ModelType.encoder_and_decoder:\n if mpu.get_pipeline_model_parallel_world_size() > 1:\n assert mpu.get_pipeline_model_parallel_split_rank() is not None, (\n \"Split rank needs to be specified for model with both encoder and decoder\"\n )\n rank = mpu.get_pipeline_model_parallel_rank()\n split_rank = mpu.get_pipeline_model_parallel_split_rank()\n world_size = mpu.get_pipeline_model_parallel_world_size()\n pre_process = rank == 0 or rank == split_rank\n post_process = (rank == (split_rank - 1)) or (rank == (world_size - 1))\n add_encoder = mpu.is_pipeline_stage_before_split()\n add_decoder = mpu.is_pipeline_stage_after_split()\n model = model_provider_func(\n pre_process=pre_process, post_process=post_process, add_encoder=add_encoder, add_decoder=add_decoder\n )\n else:\n model = model_provider_func(pre_process=pre_process, post_process=post_process)\n model.model_type = model_type\n\n if not isinstance(model, list):\n model = [model]\n\n # Set tensor model parallel attributes if not set.\n # Only parameters that are already tensor model parallel have these\n # attributes set for them. We should make sure the default attributes\n # are set for all params so the optimizer can use them.\n for model_module in model:\n for param in model_module.parameters():\n tensor_parallel.set_defaults_if_not_set_tensor_model_parallel_attributes(param)\n\n # Print number of parameters.\n if mpu.get_data_parallel_rank() == 0:\n print(\n \" > number of parameters on (tensor, pipeline) model parallel rank ({}, {}): {}\".format(\n mpu.get_tensor_model_parallel_rank(),\n mpu.get_pipeline_model_parallel_rank(),\n sum([sum([p.nelement() for p in model_module.parameters()]) for model_module in model]),\n ),\n flush=True,\n )\n\n # GPU allocation.\n if transformer_config is None or (not transformer_config.use_cpu_initialization):\n for model_module in model:\n model_module.to(f\"{get_device_name()}:{get_device_id()}\")\n\n # Fp16 conversion.\n config: TransformerConfig = get_model_config(model[0])\n config.fp8 = None\n tfconfig: TransformerConfig = model[0].config\n if config.fp16 or config.bf16: # the ModelParallelConfig in GPTModel\n model = [Float16Module(config, model_module) for model_module in model]\n\n if wrap_with_ddp:\n ddp_models = []\n ddp_config_dict = {\n \"use_distributed_optimizer\": use_distributed_optimizer,\n \"grad_reduce_in_fp32\": True,\n \"overlap_grad_reduce\": False,\n }\n if override_ddp_config is not None:\n ddp_config_dict.update(override_ddp_config)\n ddp_config = DistributedDataParallelConfig(**ddp_config_dict)\n for model_chunk_idx, model_chunk in enumerate(model):\n ddp_model = DDP(\n config=tfconfig,\n module=model_chunk,\n disable_bucketing=(model_chunk_idx > 0),\n ddp_config=ddp_config,\n )\n ddp_models.append(ddp_model)\n model = ddp_models\n # # Broadcast params from data parallel src rank to other data parallel ranks.\n # # if args.data_parallel_random_init:\n for model_module in model:\n model_module.broadcast_params()\n return model\n\n\n@dataclass\nclass McoreModuleWrapperConfig:\n \"\"\"Configuration for Mcore module wrapper.\"\"\"\n\n is_value_model: bool = False\n share_embeddings_and_output_weights: bool = False\n wrap_with_ddp: bool = True\n use_distributed_optimizer: bool = True\n\n\ndef make_megatron_module(\n wrap_config: McoreModuleWrapperConfig,\n tf_config: TransformerConfig,\n hf_config: PretrainedConfig,\n bridge: Any = None,\n provider: Any = None,\n override_model_config: dict[str, Any] = None,\n override_ddp_config: dict[str, Any] = None,\n peft_cls: Any = None,\n peft_config: Any = None,\n):\n if override_model_config is None:\n override_model_config = {}\n\n if bridge is not None:\n if provider is None:\n from verl.models.mcore.mbridge import freeze_moe_router, make_value_model\n\n value_model_hook = make_value_model\n else:\n from verl.models.mcore.bridge import freeze_moe_router, make_value_model\n\n hidden_size = (\n hf_config.text_config.hidden_size if hasattr(hf_config, \"text_config\") else hf_config.hidden_size\n )\n value_model_hook = make_value_model(hidden_size, provider.sequence_parallel)\n\n post_model_creation_callbacks = []\n if wrap_config.is_value_model:\n post_model_creation_callbacks.append(value_model_hook)\n if override_model_config.get(\"moe_config\", {}).get(\"freeze_moe_router\", False):\n post_model_creation_callbacks.append(freeze_moe_router)\n if provider is not None:\n # When using PEFT with Megatron-Bridge, we must apply PEFT transformation\n # BEFORE wrapping the model in DDP. This is required because:\n # 1. PEFT freezes base model parameters (requires_grad=False)\n # 2. DDP must be aware of which parameters are trainable when building gradient buckets\n # 3. The distributed optimizer must only track trainable (adapter) parameters\n # See Megatron-Bridge docs: training/peft.md\n\n # Register PEFT transformation as pre-wrap hook if peft_cls is specified\n # This must happen BEFORE DDP wrapping to avoid KeyError with frozen parameters\n if peft_cls is not None:\n from verl.utils.megatron_peft_utils import load_adapter_checkpoint, print_adapter_info\n\n def peft_pre_wrap_hook(model):\n \"\"\"Pre-wrap hook that applies PEFT transformation.\"\"\"\n # Apply PEFT transformation - this will freeze base model and add adapters\n # The PEFT callable handles both freezing and transformation\n transformed_model = peft_cls(model, training=True)\n\n # Set parameters to save (adapter-only checkpointing)\n peft_cls.set_params_to_save(transformed_model)\n\n # Load adapter weights if adapter_path is specified\n adapter_path = getattr(peft_config, \"adapter_path\", None)\n if adapter_path is not None and adapter_path:\n print(f\"Loading adapter weights from: {adapter_path}\")\n load_adapter_checkpoint(transformed_model, adapter_path)\n\n # Print PEFT statistics\n if torch.distributed.get_rank() == 0:\n print_adapter_info(transformed_model)\n\n return transformed_model\n\n provider.register_pre_wrap_hook(peft_pre_wrap_hook)\n\n # Register post-creation callbacks (make_value_model, freeze_moe_router) as pre-wrap hooks\n for callback in post_model_creation_callbacks:\n provider.register_pre_wrap_hook(callback)\n\n # Create DDP config if needed\n ddp_config = None\n if wrap_config.wrap_with_ddp:\n from megatron.bridge.training.config import DistributedDataParallelConfig\n\n ddp_config_dict = {\n \"use_distributed_optimizer\": wrap_config.use_distributed_optimizer,\n }\n # Apply any DDP config overrides\n if override_ddp_config is not None:\n ddp_config_dict.update(override_ddp_config)\n\n ddp_config = DistributedDataParallelConfig(**ddp_config_dict)\n ddp_config.finalize()\n\n # Now call provide_distributed_model with all hooks registered\n # Hooks will be applied automatically before DDP wrapping\n model = provider.provide_distributed_model(\n wrap_with_ddp=wrap_config.wrap_with_ddp,\n ddp_config=ddp_config,\n fp16=provider.fp16,\n bf16=provider.bf16,\n )\n\n # Extract TransformerConfig from the created model\n tf_config = get_model_config(model[0] if isinstance(model, list) else model)\n else:\n model = bridge.get_model(\n post_model_creation_callbacks=post_model_creation_callbacks,\n wrap_with_ddp=wrap_config.wrap_with_ddp,\n fp16=tf_config.fp16,\n bf16=tf_config.bf16,\n ddp_config=override_ddp_config,\n )\n\n if isinstance(tf_config, MLATransformerConfig):\n # Keep the same behavior as hf_to_mcore_config_dpskv3\n from verl.models.mcore.patch import apply_patch\n\n apply_patch()\n else:\n\n def megatron_model_provider(pre_process, post_process, vp_stage=None):\n from verl.models.mcore import init_mcore_model\n\n parallel_model = init_mcore_model(\n tf_config,\n hf_config,\n pre_process,\n post_process,\n share_embeddings_and_output_weights=wrap_config.share_embeddings_and_output_weights,\n value=wrap_config.is_value_model,\n freeze_moe_router=override_model_config.get(\"moe_config\", {}).get(\"freeze_moe_router\", False),\n vp_stage=vp_stage,\n )\n parallel_model.to(get_device_name())\n return parallel_model\n\n model = get_model(\n megatron_model_provider,\n wrap_with_ddp=wrap_config.wrap_with_ddp,\n use_distributed_optimizer=wrap_config.use_distributed_optimizer,\n override_ddp_config=override_ddp_config,\n )\n return model, tf_config\n\n\nALL_MODULE_WRAPPER_CLASSNAMES = (DDP, Float16Module)\n\n\ndef unwrap_model(model, module_instances=ALL_MODULE_WRAPPER_CLASSNAMES):\n return_list = True\n if not isinstance(model, list):\n model = [model]\n return_list = False\n unwrapped_model = []\n for model_module in model:\n while isinstance(model_module, module_instances):\n model_module = model_module.module\n unwrapped_model.append(model_module)\n if not return_list:\n return unwrapped_model[0]\n return unwrapped_model\n\n\ndef convert_config(hf_config: PretrainedConfig, megatron_config) -> TransformerConfig:\n \"\"\"[Deprecated] convert config\n\n Args:\n hf_config (PretrainedConfig): _description_\n megatron_config (_type_): _description_\n\n Returns:\n TransformerConfig: _description_\n \"\"\"\n\n warnings.warn(\"[deprecated] use config converter for more model support\", stacklevel=2)\n print(f\"megatron config {megatron_config}\")\n dt = PrecisionType.to_dtype(megatron_config.params_dtype)\n print(f\"pipeline_dtype=megatron_config {dt}\")\n qkv_bias = True if \"Qwen2ForCausalLM\" in hf_config.architectures else getattr(hf_config, \"attention_bias\", False)\n overlap_p2p_comm = (\n mpu.get_virtual_pipeline_model_parallel_world_size() is not None\n and mpu.get_virtual_pipeline_model_parallel_world_size() > 1\n )\n batch_p2p_comm = False\n transformer_config = TransformerConfig(\n num_layers=hf_config.num_hidden_layers,\n hidden_size=hf_config.hidden_size,\n num_attention_heads=hf_config.num_attention_heads,\n num_query_groups=hf_config.num_key_value_heads,\n ffn_hidden_size=hf_config.intermediate_size,\n # max_position_embeddings=hf_config.max_position_embeddings,\n activation_func=F.silu,\n normalization=\"RMSNorm\",\n # rotary_percent=False, # default,\n gated_linear_unit=True, # for llama\n use_cpu_initialization=True,\n apply_residual_connection_post_layernorm=False, # check what's this mean\n add_bias_linear=False,\n tensor_model_parallel_size=mpu.get_tensor_model_parallel_world_size(),\n pipeline_model_parallel_size=mpu.get_pipeline_model_parallel_world_size(),\n virtual_pipeline_model_parallel_size=mpu.get_virtual_pipeline_model_parallel_world_size(),\n context_parallel_size=mpu.get_context_parallel_world_size(),\n overlap_p2p_comm=overlap_p2p_comm,\n batch_p2p_comm=batch_p2p_comm,\n pipeline_dtype=dt,\n params_dtype=dt,\n sequence_parallel=mpu.get_tensor_model_parallel_world_size() > 1,\n variable_seq_lengths=True,\n masked_softmax_fusion=True,\n moe_token_dispatcher_type=\"alltoall\",\n attention_dropout=hf_config.attention_dropout,\n hidden_dropout=getattr(hf_config, \"hidden_dropout\", 0.0),\n add_qkv_bias=qkv_bias,\n bf16=dt is torch.bfloat16,\n )\n\n return transformer_config\n\n\ndef mcore_model_parallel_config(\n sequence_parallel: bool,\n params_dtype: torch.dtype,\n) -> ModelParallelConfig:\n # WARNING: Code should not reach this point. This function is deprecated and will be removed.\n # Please use hf_to_mcore_config_dense() from verl.models.mcore.config_converter instead.\n warnings.warn(\n \"Code should not reach this point. This function is deprecated and will be removed. Please use \"\n \"hf_to_mcore_config_dense() from verl.models.mcore.config_converter instead.\",\n DeprecationWarning,\n stacklevel=2,\n )\n return ModelParallelConfig(\n tensor_model_parallel_size=mpu.get_tensor_model_parallel_world_size(),\n pipeline_model_parallel_size=mpu.get_pipeline_model_parallel_world_size(),\n virtual_pipeline_model_parallel_size=mpu.get_virtual_pipeline_model_parallel_world_size(),\n context_parallel_size=mpu.get_context_parallel_world_size(),\n sequence_parallel=sequence_parallel,\n params_dtype=params_dtype,\n pipeline_dtype=params_dtype,\n bf16=True,\n fp16=False,\n timers=None,\n )\n\n\n@torch.no_grad()\ndef offload_megatron_model_to_cpu(models):\n \"\"\"\n In megatron, the model and optimizer storage are:\n - bf16 parameter data chunked in model parallel group\n - fp32 grad chunked in model parallel group\n - fp32 main_parameter chunked in model and dp group\n - fp32 optimizer state chunked in model and dp group\n \"\"\"\n for model_chunk in models:\n if isinstance(model_chunk, DDP):\n model_chunk_all_buffers = [model_chunk.buffers, model_chunk.expert_parallel_buffers]\n for buffers in model_chunk_all_buffers:\n for buffer in buffers:\n # offload parameters\n if buffer.param_data.storage().size() > 0:\n buffer.param_data.cpu_data = buffer.param_data.data.cpu().pin_memory()\n buffer.param_data_size = buffer.param_data.storage().size()\n buffer.param_data.storage().resize_(0)\n\n assert buffer.param_data_size == buffer.param_data.cpu_data.storage().size()\n\n if buffer.grad_data.storage().size() > 0:\n # if the grad_data size is already zero, we assume that it is already offloaded\n buffer.grad_data_size = buffer.grad_data.storage().size()\n buffer.grad_data.storage().resize_(0)\n else:\n # we need this for ref module\n for _, param in model_chunk.named_parameters():\n param.data = param.data.to(\"cpu\", non_blocking=True)\n if param.grad is not None:\n param.grad = param.grad.to(\"cpu\", non_blocking=True)\n gc.collect()\n get_torch_device().empty_cache()\n\n\n@torch.no_grad()\ndef load_megatron_model_to_gpu(models, load_grad=True):\n for model_chunk in models:\n if isinstance(model_chunk, DDP):\n model_chunk_all_buffers = [model_chunk.buffers, model_chunk.expert_parallel_buffers]\n for buffers in model_chunk_all_buffers:\n for buffer in buffers:\n # sometimes, we don't want to load grad for pure inference\n if load_grad and hasattr(buffer, \"grad_data_size\"):\n buffer.grad_data.storage().resize_(buffer.grad_data_size)\n buffer.grad_data.zero_()\n\n if buffer.param_data.storage().size() == 0:\n buffer.param_data.storage().resize_(buffer.param_data_size)\n # copy data from cpu to cuda\n buffer.param_data.copy_(buffer.param_data.cpu_data, non_blocking=True)\n else:\n # we need this for ref module\n device_id = get_device_id()\n for _, param in model_chunk.named_parameters():\n param.data = param.data.to(device_id, non_blocking=True)\n if param.grad is not None:\n param.grad = param.grad.to(device_id, non_blocking=True)\n gc.collect()\n get_torch_device().empty_cache()\n\n\n@torch.no_grad()\ndef offload_megatron_copy_params(optimizers):\n \"\"\"\n Offload optimizer parameters to CPU. Supports both Megatron optimizers\n and `ChainedOptimizer`, which wraps a list of underlying optimizers.\n\n Args:\n optimizers: The optimizer or ChainedOptimizer instance.\n \"\"\"\n\n def _iter_opts(opt):\n if isinstance(opt, ChainedOptimizer):\n return opt.chained_optimizers\n return [opt]\n\n def offload_tensor_to_cpu(tensor):\n if tensor is None:\n return\n tensor.data = tensor.data.to(\"cpu\", non_blocking=True)\n\n def offload_group_to_cpu(group):\n if group is None:\n return\n\n if isinstance(group, list):\n for param_group in group:\n if isinstance(param_group, list):\n for param in param_group:\n offload_tensor_to_cpu(param)\n else:\n offload_tensor_to_cpu(param_group)\n else:\n offload_tensor_to_cpu(group)\n\n # Offload all parameter groups to CPU for each underlying optimizer\n\n for _opt in _iter_opts(optimizers):\n if hasattr(_opt, \"shard_fp32_from_float16_groups\"):\n offload_group_to_cpu(_opt.shard_fp32_from_float16_groups)\n\n\n@torch.no_grad()\ndef load_megatron_copy_params(optimizers):\n \"\"\"\n Load optimizer parameters back to GPU. Handles ChainedOptimizer.\n\n Args:\n optimizers: Optimizer or ChainedOptimizer instance.\n \"\"\"\n\n def _iter_opts(opt):\n if isinstance(opt, ChainedOptimizer):\n return opt.chained_optimizers\n return [opt]\n\n def load_tensor_to_gpu(tensor):\n if tensor is None:\n return\n device_id = get_device_id()\n tensor.data = tensor.data.to(device_id, non_blocking=True)\n\n def load_group_to_gpu(group):\n if group is None:\n return\n\n if isinstance(group, list):\n for param_group in group:\n if isinstance(param_group, list):\n for param in param_group:\n load_tensor_to_gpu(param)\n else:\n load_tensor_to_gpu(param_group)\n else:\n load_tensor_to_gpu(group)\n\n # Load all parameter groups to GPU for each underlying optimizer\n\n for _opt in _iter_opts(optimizers):\n if hasattr(_opt, \"shard_fp32_from_float16_groups\"):\n load_group_to_gpu(_opt.shard_fp32_from_float16_groups)\n\n\n@torch.no_grad()\ndef offload_megatron_optimizer(optimizers):\n def _iter_opts(opt):\n if isinstance(opt, ChainedOptimizer):\n return opt.chained_optimizers\n return [opt]\n\n for _opt in _iter_opts(optimizers):\n offload_megatron_copy_params(_opt)\n ## worker may hold zero parameter when enabling custom pipeline layout\n if _opt.optimizer is not None:\n # HybridDeviceOptimizer: offload all sub-optimizer states to CPU\n # TODO: this should be a method in Megatron-LM's HybridDeviceOptimizer\n hdo = _opt.optimizer\n if all(hasattr(hdo, attr) for attr in (\"sub_optimizers\", \"inner_param_to_orig_param\", \"state\")):\n for optimizer in hdo.sub_optimizers:\n for param, state in optimizer.state.items():\n for k, v in state.items():\n if not isinstance(v, torch.Tensor):\n continue\n orig_param = hdo.inner_param_to_orig_param.get(param, param)\n hdo.state[orig_param][k] = state[k] = v.to(\"cpu\")\n else:\n opt_state_dict_values = _opt.optimizer.state.values()\n for v in opt_state_dict_values:\n if \"exp_avg\" in v:\n v[\"exp_avg\"] = v[\"exp_avg\"].to(\"cpu\", non_blocking=True)\n if \"exp_avg_sq\" in v:\n v[\"exp_avg_sq\"] = v[\"exp_avg_sq\"].to(\"cpu\", non_blocking=True)\n\n try:\n # Free TransformerEngine's dummy weight gradients cache\n # https://github.com/NVIDIA/TransformerEngine/blob/release_v2.10/transformer_engine/pytorch/module/base.py#L64\n from transformer_engine.pytorch.module.base import _dummy_wgrads\n\n _dummy_wgrads.clear()\n except ImportError:\n pass\n\n # Free Megatron-LM's global memory buffer\n get_global_memory_buffer().buffer.clear()\n\n gc.collect()\n get_torch_device().empty_cache()\n\n\n@torch.no_grad()\ndef load_megatron_optimizer(optimizers):\n def _iter_opts(opt):\n if isinstance(opt, ChainedOptimizer):\n return opt.chained_optimizers\n return [opt]\n\n for _opt in _iter_opts(optimizers):\n load_megatron_copy_params(_opt)\n ## worker may hold zero parameter when enabling custom pipeline layout\n if _opt.optimizer is not None:\n # if we are using HybridDeviceOptimizer, we need to only move gpu optimizer state to gpu\n if hasattr(_opt.optimizer, \"_move_new_state_to_right_device\"):\n _opt.optimizer._move_new_state_to_right_device()\n else:\n opt_state_dict_values = _opt.optimizer.state.values()\n for v in opt_state_dict_values:\n if \"exp_avg\" in v:\n v[\"exp_avg\"] = v[\"exp_avg\"].to(get_device_id(), non_blocking=True)\n if \"exp_avg_sq\" in v:\n v[\"exp_avg_sq\"] = v[\"exp_avg_sq\"].to(get_device_id(), non_blocking=True)\n gc.collect()\n get_torch_device().empty_cache()\n\n\ndef get_dist_checkpoint_path(checkpoint_path):\n local_mkdir_safe(checkpoint_path)\n local_mkdir_safe(os.path.join(checkpoint_path, \"dist_ckpt\"))\n return os.path.join(checkpoint_path, \"dist_ckpt\")\n\n\ndef get_hf_model_checkpoint_path(checkpoint_path):\n local_mkdir_safe(checkpoint_path)\n local_mkdir_safe(os.path.join(checkpoint_path, \"huggingface\"))\n return os.path.join(checkpoint_path, \"huggingface\")\n\n\ndef get_transformer_config_checkpoint_path(checkpoint_path):\n os.makedirs(checkpoint_path, exist_ok=True)\n return os.path.join(checkpoint_path, \"transformer_config.json\")\n\n\ndef convert_megatron_model_to_transformers_model(\n name,\n param,\n config: PretrainedConfig,\n tp_size: int,\n num_query_groups: int,\n convert_qkv_gate_up_by_trunk_concat=False,\n):\n \"\"\"Convert megatron model to transformers model.\"\"\"\n new_params = {}\n\n def convert_qkv_shard(full_tensor, q_name, k_name, v_name):\n nonlocal config\n nonlocal tp_size\n nonlocal num_query_groups\n\n q_shard_list = []\n k_shard_list = []\n v_shard_list = []\n hidden_size_per_head = getattr(config, \"head_dim\", config.hidden_size // config.num_attention_heads)\n\n if config.num_key_value_heads >= tp_size:\n q_size_tp = hidden_size_per_head * config.num_attention_heads // tp_size\n kv_size_tp = hidden_size_per_head * config.num_key_value_heads // tp_size\n total_size = q_size_tp + 2 * kv_size_tp\n for i in range(tp_size):\n num_query_groups_per_partition = num_query_groups // tp_size\n qkv_part = full_tensor[i * total_size : (i + 1) * total_size]\n q_size_chunk = q_size_tp // num_query_groups_per_partition\n kv_size_chunk = kv_size_tp // num_query_groups_per_partition\n for qkv_part_chunk in qkv_part.chunk(num_query_groups_per_partition):\n q_part = qkv_part_chunk[:q_size_chunk]\n k_part = qkv_part_chunk[q_size_chunk : q_size_chunk + kv_size_chunk]\n v_part = qkv_part_chunk[q_size_chunk + kv_size_chunk :]\n q_shard_list.append(q_part)\n k_shard_list.append(k_part)\n v_shard_list.append(v_part)\n else:\n q_size_tp = hidden_size_per_head * config.num_attention_heads // tp_size\n kv_size_tp = hidden_size_per_head\n total_size = q_size_tp + 2 * kv_size_tp\n for i in range(tp_size):\n num_query_groups_per_partition = num_query_groups // tp_size\n qkv_part = full_tensor[i * total_size : (i + 1) * total_size]\n q_size_chunk = q_size_tp // num_query_groups_per_partition\n kv_size_chunk = kv_size_tp // num_query_groups_per_partition\n for qkv_part_chunk in qkv_part.chunk(num_query_groups_per_partition):\n q_part = qkv_part_chunk[:q_size_chunk]\n k_part = qkv_part_chunk[q_size_chunk : q_size_chunk + kv_size_chunk]\n v_part = qkv_part_chunk[q_size_chunk + kv_size_chunk :]\n q_shard_list.append(q_part)\n if i * config.num_key_value_heads % tp_size == 0:\n k_shard_list.append(k_part)\n v_shard_list.append(v_part)\n\n new_params[q_name] = torch.cat(q_shard_list, dim=0)\n new_params[k_name] = torch.cat(k_shard_list, dim=0)\n new_params[v_name] = torch.cat(v_shard_list, dim=0)\n\n def convert_gate_up_shard(full_tensor, gate_name, up_name):\n nonlocal config\n nonlocal tp_size\n\n intermediate_size_tp = config.intermediate_size // tp_size\n gate_weight_list = []\n up_weight_list = []\n for i in range(tp_size):\n gate_up_weight_tp = full_tensor[intermediate_size_tp * 2 * i : intermediate_size_tp * 2 * (i + 1)]\n gate_weight_tp = gate_up_weight_tp[:intermediate_size_tp]\n up_weight_tp = gate_up_weight_tp[intermediate_size_tp:]\n gate_weight_list.append(gate_weight_tp)\n up_weight_list.append(up_weight_tp)\n\n new_params[gate_name] = torch.cat(gate_weight_list, dim=0)\n new_params[up_name] = torch.cat(up_weight_list, dim=0)\n\n if name == \"embedding.word_embeddings.weight\":\n new_params[\"model.embed_tokens.weight\"] = param\n elif \"self_attention\" in name:\n splitted_name = name.split(\".\")\n layer_number = splitted_name[2]\n component = splitted_name[4]\n param_type = splitted_name[5]\n if component == \"linear_proj\":\n new_params[f\"model.layers.{layer_number}.self_attn.o_proj.weight\"] = param\n elif component == \"linear_qkv\" and not isinstance(param, list):\n if param_type == \"layer_norm_weight\":\n new_params[f\"model.layers.{layer_number}.input_layernorm.weight\"] = param\n else:\n if convert_qkv_gate_up_by_trunk_concat:\n convert_qkv_shard(\n param,\n f\"model.layers.{layer_number}.self_attn.q_proj.{param_type}\",\n f\"model.layers.{layer_number}.self_attn.k_proj.{param_type}\",\n f\"model.layers.{layer_number}.self_attn.v_proj.{param_type}\",\n )\n else:\n new_params[f\"model.layers.{layer_number}.self_attn.qkv_proj.{param_type}\"] = param\n elif component == \"q_layernorm\" or component == \"k_layernorm\":\n hf_component = component.replace(\"layer\", \"\")\n new_params[f\"model.layers.{layer_number}.self_attn.{hf_component}.weight\"] = param\n else:\n assert isinstance(param, list) and len(param) == 3\n assert param_type == \"weight\" or param_type == \"bias\"\n new_params[f\"model.layers.{layer_number}.self_attn.q_proj.{param_type}\"] = param[0]\n new_params[f\"model.layers.{layer_number}.self_attn.k_proj.{param_type}\"] = param[1]\n new_params[f\"model.layers.{layer_number}.self_attn.v_proj.{param_type}\"] = param[2]\n elif \"mlp\" in name:\n splitted_name = name.split(\".\")\n layer_number = splitted_name[2]\n component = splitted_name[4]\n param_type = splitted_name[5]\n if component == \"linear_fc1\" and not isinstance(param, list):\n if param_type == \"layer_norm_weight\":\n new_params[f\"model.layers.{layer_number}.post_attention_layernorm.weight\"] = param\n elif param_type == \"weight\":\n if convert_qkv_gate_up_by_trunk_concat:\n convert_gate_up_shard(\n param,\n f\"model.layers.{layer_number}.mlp.gate_proj.weight\",\n f\"model.layers.{layer_number}.mlp.up_proj.weight\",\n )\n else:\n new_params[f\"model.layers.{layer_number}.mlp.gate_up_proj.weight\"] = param\n elif component == \"linear_fc1\" and isinstance(param, list):\n assert len(param) == 2\n assert param_type == \"weight\" or param_type == \"bias\"\n new_params[f\"model.layers.{layer_number}.mlp.gate_proj.weight\"] = param[0]\n new_params[f\"model.layers.{layer_number}.mlp.up_proj.weight\"] = param[1]\n elif component == \"linear_fc2\":\n new_params[f\"model.layers.{layer_number}.mlp.down_proj.weight\"] = param\n elif name == \"decoder.final_layernorm.weight\":\n new_params[\"model.norm.weight\"] = param\n elif name == \"output_layer.weight\":\n new_params[\"lm_head.weight\"] = param\n else:\n raise ValueError(f\"Unknown param name: {name}\")\n return new_params.keys(), new_params.values()\n\n\ndef broadcast_from_megatron_pp(tensor: torch.Tensor):\n # tensor is not None only in one of the pp ranks\n if tensor is not None:\n shape = tensor.shape\n dtype = tensor.dtype\n tensor_parallel = getattr(tensor, \"tensor_model_parallel\", None)\n partition_dim = getattr(tensor, \"partition_dim\", None)\n tensor_spec = (shape, dtype, tensor_parallel, partition_dim)\n else:\n tensor_spec = None\n tensor_spec_output = [None] * mpu.get_pipeline_model_parallel_world_size()\n torch.distributed.all_gather_object(\n object_list=tensor_spec_output, obj=tensor_spec, group=mpu.get_pipeline_model_parallel_group()\n )\n # find the src rank\n target_tensor_spec = None\n src_rank = None\n for rank, tensor_spec in enumerate(tensor_spec_output):\n if tensor_spec is not None:\n if target_tensor_spec is None:\n target_tensor_spec = tensor_spec\n else:\n raise ValueError(\"A tensor exists on two pp ranks\")\n src_rank = rank\n assert target_tensor_spec is not None\n if tensor is None:\n tensor = torch.empty(size=target_tensor_spec[0], dtype=target_tensor_spec[1], device=get_device_id())\n if target_tensor_spec[2] is not None:\n tensor.tensor_model_parallel = target_tensor_spec[2]\n if target_tensor_spec[3] is not None:\n tensor.partition_dim = target_tensor_spec[3]\n\n global_rank = torch.distributed.get_global_rank(group=mpu.get_pipeline_model_parallel_group(), group_rank=src_rank)\n torch.distributed.broadcast(tensor=tensor, src=global_rank, group=mpu.get_pipeline_model_parallel_group())\n return tensor\n\n\ndef broadcast_str_from_megatron_pp(obj: Any):\n obj_output = [None] * mpu.get_pipeline_model_parallel_world_size()\n torch.distributed.all_gather_object(object_list=obj_output, obj=obj, group=mpu.get_pipeline_model_parallel_group())\n\n src_rank = None\n target_obj = None\n for rank, item in enumerate(obj_output):\n if item is not None:\n if target_obj is not None:\n raise ValueError(\"An object exists on two pp ranks\")\n target_obj = item\n src_rank = rank\n\n assert target_obj is not None, \"No valid object found to broadcast.\"\n\n global_rank = torch.distributed.get_global_rank(group=mpu.get_pipeline_model_parallel_group(), group_rank=src_rank)\n\n obj_output = [None] * torch.distributed.get_world_size(group=mpu.get_pipeline_model_parallel_group())\n obj_output[0] = target_obj\n torch.distributed.broadcast_object_list(\n object_list=obj_output, src=global_rank, group=mpu.get_pipeline_model_parallel_group()\n )\n\n return obj_output[0]\n\n\ndef default_tp_concat_fn(\n layer_name_mapping,\n name,\n train_params,\n infer_params,\n model_config,\n hf_config=None,\n convert_qkv_gate_up_by_simple_split=False,\n):\n \"\"\"\n name: name of the parameter\n train_params: training parameters\n infer_params (Iterable[torch.Tensor]): a iterator towards list of parameters all-gathered from micro_dp_group\n model_config: huggingface model_config\n TODO(zhangchi.usc1992): currently, the implementation is adhoc. We can move this function to the model\n definition so that it is model-agnostic. If the model doesn't implement this function,\n we can throw an error to force user disable TP HybridEngine.\n \"\"\"\n from megatron.core import mpu\n\n train_tp_size = mpu.get_tensor_model_parallel_world_size()\n if layer_name_mapping.get(\"qkv_layer_name\") in name and \"layer_norm\" not in name:\n # if the tensor is qkv, for each param on tp, split into q, k, v\n # concat q, k, v separately.\n q_lst = []\n k_lst = []\n v_lst = []\n num_attention_heads = model_config.num_attention_heads\n num_key_value_heads = model_config.num_key_value_heads\n if \"vision_model\" in name:\n num_attention_heads = hf_config.vision_config.num_heads\n num_key_value_heads = hf_config.vision_config.num_heads\n assert num_attention_heads % num_key_value_heads == 0\n num_q_per_kv = num_attention_heads // num_key_value_heads\n assert infer_params[0].shape[0] % (num_q_per_kv + 2) == 0, (\n f\"param '{name}' shape '{infer_params[0].shape}' dim0 is not divisible by {num_q_per_kv + 2}\"\n )\n kv_size_per_tp = infer_params[0].shape[0] // (num_q_per_kv + 2)\n split_size = [kv_size_per_tp * num_q_per_kv, kv_size_per_tp, kv_size_per_tp]\n for infer_param in infer_params:\n num_query_groups_per_partition = num_key_value_heads // train_tp_size\n for chunk in infer_param.chunk(num_query_groups_per_partition):\n split_size = [\n kv_size_per_tp * num_q_per_kv // num_query_groups_per_partition,\n kv_size_per_tp // num_query_groups_per_partition,\n kv_size_per_tp // num_query_groups_per_partition,\n ]\n q, k, v = chunk.split(split_size)\n q_lst.append(q)\n k_lst.append(k)\n v_lst.append(v)\n q = torch.cat(q_lst, dim=0)\n k = torch.cat(k_lst, dim=0)\n v = torch.cat(v_lst, dim=0)\n infer_params = torch.cat((q, k, v), dim=0) if not convert_qkv_gate_up_by_simple_split else [q, k, v]\n\n elif (\n layer_name_mapping.get(\"gate_proj_layer_name\") in name\n and \"layer_norm\" not in name\n and \"vision_model.projection\" not in name\n ):\n # if the tensor is gate and proj\n gate_lst = []\n up_lst = []\n for infer_param in infer_params:\n gate, up = infer_param.chunk(2)\n gate_lst.append(gate)\n up_lst.append(up)\n gate = torch.cat(gate_lst, dim=0)\n up = torch.cat(up_lst, dim=0)\n infer_params = torch.cat((gate, up), dim=0) if not convert_qkv_gate_up_by_simple_split else [gate, up]\n\n elif \"mlp.experts.linear_fc2.weight\" in name: # moe\n infer_params = torch.cat(infer_params, dim=1)\n\n else:\n # concat tensor\n infer_params = torch.cat(infer_params, dim=tp_utils.get_tensor_parallel_partition_dim(train_params))\n\n return infer_params\n\n\ndef per_tensor_generator(\n actor_module,\n model_config,\n weight_converter,\n transformer_config,\n layer_name_mapping,\n convert_qkv_gate_up_by_simple_split=True,\n):\n from megatron.core import parallel_state as mpu\n\n pp_rank = mpu.get_pipeline_model_parallel_rank()\n ep_size = mpu.get_expert_model_parallel_world_size()\n etp_size = mpu.get_expert_tensor_parallel_world_size()\n ep_group = mpu.get_expert_model_parallel_group()\n etp_group = mpu.get_expert_tensor_parallel_group()\n vpp_size = len(actor_module)\n all_gather_group = mpu.get_tensor_model_parallel_group()\n all_gather_group_size = torch.distributed.get_world_size(group=all_gather_group)\n\n def tensor_generator():\n for scan_vpp_idx in range(vpp_size):\n existing_keys = set()\n model = unwrap_model(actor_module[scan_vpp_idx])\n for name, param in model.named_parameters():\n existing_keys.add(name)\n yield name, param\n # note\n # there is a bug in megatron GPTModel\n # decoder.layers[n].mlp.router.expert_bias\" in GPTModel is not registered in named_parameter, but in\n # state_dict(). for now we patch it by adding those keys to extra_keys.\n extra_keys = [x for x in model.state_dict().keys() if \"_extra_state\" not in x and x not in existing_keys]\n for name in extra_keys:\n yield name, model.state_dict()[name].to(get_device_id())\n\n # we need first make all rank get full model information\n meta_info = []\n for scan_vpp_idx in range(vpp_size):\n existing_keys = set()\n model = unwrap_model(actor_module[scan_vpp_idx])\n for idx, (name, _) in enumerate(model.named_parameters()):\n existing_keys.add(name)\n meta_info.append((pp_rank, scan_vpp_idx, idx, name))\n extra_keys = [x for x in model.state_dict().keys() if \"_extra_state\" not in x and x not in existing_keys]\n for name in extra_keys:\n meta_info.append((pp_rank, scan_vpp_idx, idx, name))\n\n obj_spec_output = [None] * mpu.get_pipeline_model_parallel_world_size()\n torch.distributed.all_gather_object(\n object_list=obj_spec_output, obj=meta_info, group=mpu.get_pipeline_model_parallel_group()\n )\n layer_list_meta = [item for sublist in obj_spec_output for item in sublist]\n\n gen_func = tensor_generator()\n\n # lazy load tensor for full model\n for cur_pp_rank, scan_vpp_idx, idx, name in layer_list_meta:\n if model_config.tie_word_embeddings and (\"output_layers\" in name):\n import warnings\n\n warnings.warn(\n \"Current model sharing word and embedding weights, skip output layer conversion\", stacklevel=2\n )\n continue\n\n if cur_pp_rank == pp_rank:\n try:\n cur_name, cur_tensor = next(gen_func)\n except StopIteration:\n cur_name, cur_tensor = None, None\n cur_name = normalize_model_name(name, cur_pp_rank, scan_vpp_idx, transformer_config)\n else:\n cur_tensor, cur_name = None, None\n\n # pp broadcast model tensor and name\n cur_name = broadcast_str_from_megatron_pp(cur_name)\n broad_pp_tensor = broadcast_from_megatron_pp(cur_tensor)\n\n # (xya): this is a hack to fix the name of the parameters\n while cur_name.startswith(\"module.\"):\n cur_name = cur_name[len(\"module.\") :]\n\n # EP\n if \".mlp.experts.linear_fc\" in cur_name and ep_size > 1:\n num_experts = weight_converter.mcore_config.num_moe_experts\n num_experts_per_rank = num_experts // ep_size\n infer_params = [torch.empty_like(broad_pp_tensor) for _ in range(ep_size)]\n torch.distributed.all_gather(infer_params, broad_pp_tensor, group=ep_group)\n\n name_prefix, local_expert_id = cur_name.split(\".weight\")\n local_expert_id = int(local_expert_id)\n global_expert_ids = [num_experts_per_rank * ep_rank + local_expert_id for ep_rank in range(ep_size)]\n global_expert_names = [f\"{name_prefix}.weight{expert_id}\" for expert_id in global_expert_ids]\n\n for name, param in zip(global_expert_names, infer_params, strict=True):\n if etp_size > 1:\n # gather etp\n etp_params = [torch.empty_like(param) for _ in range(etp_size)]\n torch.distributed.all_gather(etp_params, param, group=etp_group)\n params = etp_params\n else:\n params = [param]\n\n merge_params = default_tp_concat_fn(\n layer_name_mapping,\n name,\n broad_pp_tensor,\n params,\n model_config,\n weight_converter.hf_config,\n convert_qkv_gate_up_by_simple_split,\n )\n if not isinstance(merge_params, list):\n merge_params = [merge_params]\n converted_names, converted_params = weight_converter.convert_param(name, merge_params)\n\n yield from zip(converted_names, [param.detach() for param in converted_params], strict=True)\n continue\n\n # tp all gather\n if tp_utils.is_tensor_parallel_param(broad_pp_tensor):\n # allocate a new tensor with proper size\n if all_gather_group_size <= 1:\n infer_params = [broad_pp_tensor]\n else:\n infer_params = [torch.empty_like(broad_pp_tensor) for _ in range(all_gather_group_size)]\n torch.distributed.all_gather(infer_params, broad_pp_tensor, group=mpu.get_tensor_model_parallel_group())\n infer_params = default_tp_concat_fn(\n layer_name_mapping,\n cur_name,\n broad_pp_tensor,\n infer_params,\n model_config,\n weight_converter.hf_config,\n convert_qkv_gate_up_by_simple_split,\n )\n else:\n infer_params = broad_pp_tensor\n\n if not isinstance(infer_params, list):\n infer_params = [infer_params]\n converted_names, converted_params = weight_converter.convert_param(cur_name, infer_params)\n\n yield from zip(converted_names, [param.detach() for param in converted_params], strict=True)\n\n\ndef get_transformer_layer_offset(pipeline_rank, vp_stage, config: TransformerConfig):\n \"\"\"\n Get the index offset of any pipeline stage, given the level of pipelining.\n\n Make pipeline_rank and vp_stage as two arguments to make it more flexible,\n which is able to fetch layer offset for any pipeline stage.\n The original function only returns the layer offset for current pipeline stage.\n\n Extension to https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/core/transformer/transformer_layer.py::get_transformer_layer_offset\n \"\"\"\n\n has_vp_stage = (\n inspect.signature(parallel_state.is_pipeline_first_stage).parameters.get(\"vp_stage\", None) is not None\n )\n extra_kwargs = {} if not has_vp_stage else {\"ignore_virtual\": False, \"vp_stage\": vp_stage}\n\n if config.pipeline_model_parallel_size > 1:\n if hasattr(config, \"pipeline_model_parallel_layout\") and config.pipeline_model_parallel_layout:\n from megatron.core.transformer.enums import LayerType\n\n offset = config.pipeline_model_parallel_layout.get_layer_offset(\n layer_type=LayerType.decoder, vp_stage=vp_stage\n )\n elif (\n config.num_layers_in_first_pipeline_stage is not None\n or config.num_layers_in_last_pipeline_stage is not None\n ):\n # Calculate number of pipeline stages to distribute the remaining Transformer\n # layers after deducting the Transformer layers in the first or the last stages\n middle_pipeline_stages = config.pipeline_model_parallel_size\n middle_pipeline_stages -= sum(\n [\n 1 if x is not None else 0\n for x in (\n config.num_layers_in_first_pipeline_stage,\n config.num_layers_in_last_pipeline_stage,\n )\n ]\n )\n\n # Calculate layers to distribute in each pipeline stage. If the\n # num_layers_in_first_pipeline_stage and num_layers_in_last_pipeline_stage\n # are not set, we will not enable uneven pipeline. All layers will be treated\n # as middle layers.\n num_layers_in_first_pipeline_stage = (\n 0 if config.num_layers_in_first_pipeline_stage is None else config.num_layers_in_first_pipeline_stage\n )\n num_layers_in_last_pipeline_stage = (\n 0 if config.num_layers_in_last_pipeline_stage is None else config.num_layers_in_last_pipeline_stage\n )\n\n middle_num_layers = (\n config.num_layers - num_layers_in_first_pipeline_stage - num_layers_in_last_pipeline_stage\n )\n\n if (vp_size := config.virtual_pipeline_model_parallel_size) is not None:\n assert vp_stage is not None, \"vp_stage must be provided if virtual pipeline model parallel size is set\"\n\n # Calculate number of layers in each virtual model chunk\n # If the num_layers_in_first_pipeline_stage and\n # num_layers_in_last_pipeline_stage are not set, all pipeline stages\n # will be treated as middle pipeline stages in the calculation\n num_layers_per_virtual_model_chunk_in_first_pipeline_stage = (\n 0\n if config.num_layers_in_first_pipeline_stage is None\n else config.num_layers_in_first_pipeline_stage // vp_size\n )\n\n num_layers_per_virtual_model_chunk_in_last_pipeline_stage = (\n 0\n if config.num_layers_in_last_pipeline_stage is None\n else config.num_layers_in_last_pipeline_stage // vp_size\n )\n\n num_layers_per_vritual_model_chunk_in_middle_pipeline_stage = middle_num_layers // vp_size\n\n # First stage + middle stage + last stage\n total_virtual_chunks = (\n num_layers_per_virtual_model_chunk_in_first_pipeline_stage\n + num_layers_per_vritual_model_chunk_in_middle_pipeline_stage\n + num_layers_per_virtual_model_chunk_in_last_pipeline_stage\n )\n\n # Calculate the layer offset with interleaved uneven pipeline parallelism\n if pipeline_rank == 0:\n offset = vp_stage * total_virtual_chunks\n else:\n offset = (\n vp_stage * total_virtual_chunks\n + num_layers_per_virtual_model_chunk_in_first_pipeline_stage\n + (pipeline_rank - 1)\n * (num_layers_per_vritual_model_chunk_in_middle_pipeline_stage // middle_pipeline_stages)\n )\n else:\n if middle_pipeline_stages > 0:\n num_layers_per_pipeline_rank = middle_num_layers // middle_pipeline_stages\n else:\n num_layers_per_pipeline_rank = 0\n\n middle_pipeline_rank = (\n pipeline_rank if config.num_layers_in_first_pipeline_stage is None else pipeline_rank - 1\n )\n\n if pipeline_rank == 0:\n offset = 0\n else:\n offset = (middle_pipeline_rank * num_layers_per_pipeline_rank) + num_layers_in_first_pipeline_stage\n else:\n num_layers = config.num_layers\n\n # Increase the number of layers by one if we include the embedding (loss)\n # layer into pipeline parallelism partition and placement\n if config.account_for_embedding_in_pipeline_split:\n num_layers += 1\n\n if config.account_for_loss_in_pipeline_split:\n num_layers += 1\n\n num_layers_per_pipeline_rank = num_layers // config.pipeline_model_parallel_size\n\n if (vp_size := config.virtual_pipeline_model_parallel_size) is not None:\n assert vp_stage is not None, \"vp_stage must be provided if virtual pipeline model parallel size is set\"\n\n num_layers_per_virtual_rank = num_layers_per_pipeline_rank // vp_size\n total_virtual_chunks = num_layers // vp_size\n offset = vp_stage * total_virtual_chunks + (pipeline_rank * num_layers_per_virtual_rank)\n\n # Reduce the offset of embedding layer from the total layer number\n if config.account_for_embedding_in_pipeline_split and not parallel_state.is_pipeline_first_stage(\n **extra_kwargs\n ):\n offset -= 1\n else:\n offset = pipeline_rank * num_layers_per_pipeline_rank\n\n # Reduce the offset of embedding layer from the total layer number\n if config.account_for_embedding_in_pipeline_split and not parallel_state.is_pipeline_first_stage(\n **extra_kwargs\n ):\n offset -= 1\n else:\n offset = 0\n return offset\n\n\ndef register_megatron_training_hooks(model: list[torch.nn.Module], optimizer):\n from megatron.core.distributed import finalize_model_grads\n from megatron.core.utils import get_model_config\n\n try:\n from megatron.core.distributed.fsdp.mcore_fsdp_adapter import FullyShardedDataParallel as megatron_FSDP\n except ImportError:\n megatron_FSDP = DDP\n\n # register some callbacks for megatron training, following https://github.com/NVIDIA/Megatron-LM/blob/core_v0.15.0rc7/megatron/training/training.py#L2039-L2057\n for one_model in model:\n config = get_model_config(one_model)\n config.grad_scale_func = optimizer.scale_loss\n config.finalize_model_grads_func = finalize_model_grads\n\n overlap_param_gather = getattr(optimizer.config, \"overlap_param_gather\", False)\n overlap_grad_reduce = getattr(one_model.ddp_config, \"overlap_grad_reduce\", False)\n align_grad_reduce = True # default to True, seldom to be false\n align_param_gather = getattr(one_model.ddp_config, \"align_param_gather\", False)\n\n if isinstance(model[0], megatron_FSDP | DDP) and overlap_grad_reduce:\n assert config.no_sync_func is None, (\n \"When overlap_grad_reduce is True, config.no_sync_func must be None; \"\n \"a custom no_sync_func is not supported when overlapping grad-reduce\"\n )\n config.no_sync_func = [model_chunk.no_sync for model_chunk in model]\n if len(model) == 1:\n config.no_sync_func = config.no_sync_func[0]\n if align_grad_reduce:\n config.grad_sync_func = [model_chunk.start_grad_sync for model_chunk in model]\n if len(model) == 1:\n config.grad_sync_func = config.grad_sync_func[0]\n if overlap_param_gather and align_param_gather:\n config.param_sync_func = [model_chunk.start_param_sync for model_chunk in model]\n if len(model) == 1:\n config.param_sync_func = config.param_sync_func[0]\n\n\ndef mapping_string_to_attn_backend(args: dict) -> dict:\n if \"attention_backend\" in args and isinstance(args[\"attention_backend\"], str):\n from megatron.core.transformer.enums import AttnBackend\n\n args[\"attention_backend\"] = AttnBackend[args[\"attention_backend\"]]\n return args\n\n\ndef get_megatron_mtp_loss(n_micro_batch):\n # Calculate MTP loss scale similar to Megatron-LM implementation\n mtp_loss_scale = 1.0 / n_micro_batch\n\n # Create a dummy total_loss_dict to collect MTP metrics\n total_loss_dict = {}\n\n # Track MTP metrics - this will populate total_loss_dict with MTP losses\n MTPLossLoggingHelper.track_mtp_metrics(\n loss_scale=mtp_loss_scale, iteration=0, writer=None, wandb_writer=None, total_loss_dict=total_loss_dict\n )\n # Add MTP metrics to losses_reduced if any were collected\n # total_loss_dict: {'mtp_1 loss': tensor(value, device='cuda:0')}\n output = {}\n if total_loss_dict:\n for key, value in total_loss_dict.items():\n # Convert key to have proper prefix and format\n formatted_key = f\"mtp_losses/{key.replace(' ', '_')}\"\n # only added to the 0th batch, the MTP loss obtained is a global value, and will be the same for every batch\n output[formatted_key] = value.cpu().item()\n return output\n\n\ndef get_megatron_module_device(models: list[Any]) -> str:\n if not models:\n return \"cpu\"\n\n model_chunk = models[0]\n if not model_chunk.buffers:\n try:\n return next(model_chunk.module.parameters()).device.type\n except StopIteration:\n return \"cpu\"\n\n buffer = model_chunk.buffers[0]\n if buffer.param_data.storage().size() == 0:\n return \"cpu\"\n else:\n return get_device_name()\n\n\ndef check_mtp_config(model_config: HFModelConfig, engine_config: McoreEngineConfig):\n has_mtp = (\n model_config.hf_config.num_nextn_predict_layers > 0\n if hasattr(model_config.hf_config, \"num_nextn_predict_layers\")\n else False\n )\n enable_mtp = model_config.mtp.enable\n\n if \"mtp_loss_scaling_factor\" not in engine_config.override_transformer_config:\n engine_config.override_transformer_config[\"mtp_loss_scaling_factor\"] = model_config.mtp.mtp_loss_scaling_factor\n\n if enable_mtp and not model_config.mtp.enable_train:\n # disable parameter update by configure the loss scale to 0\n engine_config.override_transformer_config[\"mtp_loss_scaling_factor\"] = 0\n\n # Modify the hf_config before initialization, and apply patch after innitialization\n if enable_mtp and not has_mtp:\n logger.error(\"enable mtp while model has no mtp layer, ignore model.mtp.enable\")\n model_config.mtp.enable = False\n model_config.mtp.enable_train = False\n elif has_mtp and not enable_mtp:\n model_config.hf_config.num_nextn_predict_layers = 0\n\n\ndef patch_engine_mtp(module, model_config):\n logger.warning(\"Applying mtp patch...\")\n from verl.models.mcore.mtp_patch import patch_mtp_layer_get_embeddings, patch_postprocess\n\n print(module)\n if isinstance(module, list):\n for m in module:\n patch_postprocess(m)\n if model_config.mtp.detach_encoder:\n patch_mtp_layer_get_embeddings(m)\n else:\n patch_postprocess(module)\n if model_config.mtp.detach_encoder:\n patch_mtp_layer_get_embeddings(module)\n"}94{"file_name": "verl__utils__metric__utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nMetrics utils.\n\"\"\"\n\nfrom enum import Enum\nfrom typing import Any, Optional, Union\n\nimport numpy as np\nimport torch\n\n\ndef reduce_metrics(metrics: dict[str, Union[\"Metric\", list[Any]]]) -> dict[str, Any]:\n \"\"\"\n Reduces a dictionary of metric lists by computing the mean, max, or min of each list.\n The reduce operation is determined by the key name:\n - If the key contains \"max\", np.max is used\n - If the key contains \"min\", np.min is used\n - Otherwise, np.mean is used\n\n Args:\n metrics: A dictionary mapping metric names to lists of metric values.\n\n Returns:\n A dictionary with the same keys but with each list replaced by its reduced value.\n\n Example:\n >>> metrics = {\n ... \"loss\": [1.0, 2.0, 3.0],\n ... \"accuracy\": [0.8, 0.9, 0.7],\n ... \"max_reward\": [5.0, 8.0, 6.0],\n ... \"min_error\": [0.1, 0.05, 0.2]\n ... }\n >>> reduce_metrics(metrics)\n {\"loss\": 2.0, \"accuracy\": 0.8, \"max_reward\": 8.0, \"min_error\": 0.05}\n \"\"\"\n for key, val in metrics.items():\n if isinstance(val, Metric):\n metrics[key] = val.aggregate()\n elif \"max\" in key:\n metrics[key] = np.max(val)\n elif \"min\" in key:\n metrics[key] = np.min(val)\n else:\n metrics[key] = np.mean(val)\n return metrics\n\n\nclass AggregationType(Enum):\n MEAN = \"mean\"\n SUM = \"sum\"\n MIN = \"min\"\n MAX = \"max\"\n\n\nNumericType = int, float, torch.Tensor, np.ndarray\nNumeric = int | float | torch.Tensor | np.ndarray\n\n\nclass Metric:\n \"\"\"\n A metric aggregator for collecting and aggregating numeric values.\n\n This class accumulates numeric values (int, float, or scalar tensors) and computes\n an aggregate statistic based on the specified aggregation type (MEAN, SUM, MIN, or MAX).\n\n Args:\n aggregation: The aggregation method to use. Can be a string (\"mean\", \"sum\", \"min\", \"max\")\n or an AggregationType enum value.\n value: Optional initial value(s) to add. Can be a single numeric value or a list of values.\n\n Example:\n >>> metric = Metric(aggregation=\"mean\", value=1.0)\n >>> metric.append(2.0)\n >>> metric.append(3.0)\n >>> metric.aggregate()\n 2.0\n \"\"\"\n\n def __init__(self, aggregation: str | AggregationType, value: Optional[Numeric | list[Numeric]] = None) -> None:\n if isinstance(aggregation, str):\n self.aggregation = AggregationType(aggregation)\n else:\n self.aggregation = aggregation\n if not isinstance(self.aggregation, AggregationType):\n raise ValueError(f\"Unsupported aggregation type: {aggregation}\")\n self.values = []\n if value is not None:\n self.append(value)\n\n def append(self, value: Union[Numeric, \"Metric\"]) -> None:\n if isinstance(value, Metric):\n self.extend(value)\n return\n if isinstance(value, torch.Tensor):\n if value.numel() != 1:\n raise ValueError(\"Only scalar tensors can be converted to float\")\n value = value.detach().item()\n if not isinstance(value, NumericType):\n raise ValueError(f\"Unsupported value type: {type(value)}\")\n self.values.append(value)\n\n def extend(self, values: Union[\"Metric\", list[Numeric]]) -> None:\n if isinstance(values, Metric):\n if values.aggregation != self.aggregation:\n raise ValueError(f\"Aggregation type mismatch: {self.aggregation} != {values.aggregation}\")\n values = values.values\n for value in values:\n self.append(value)\n\n def aggregate(self) -> float:\n return self._aggregate(self.values, self.aggregation)\n\n @classmethod\n def _aggregate(cls, values: list[Numeric], aggregation: AggregationType) -> float:\n match aggregation:\n case AggregationType.MEAN:\n return np.mean(values)\n case AggregationType.SUM:\n return np.sum(values)\n case AggregationType.MIN:\n return np.min(values)\n case AggregationType.MAX:\n return np.max(values)\n\n @classmethod\n def aggregate_dp(cls, metric_lists: list[\"Metric\"]) -> float:\n if not metric_lists:\n raise ValueError(\"Cannot aggregate an empty list of metrics.\")\n value_lists = [ml.values for ml in metric_lists]\n if not all(len(ls) == len(value_lists[0]) for ls in value_lists):\n raise ValueError(\n f\"All Metric instances must have the same number of values \"\n f\"for dp aggregation: {[len(ls) for ls in value_lists]}\"\n )\n value_arrays = np.array(value_lists) # [num_dp, num_grad_accumulation]\n aggregation = metric_lists[0].aggregation\n match aggregation:\n case AggregationType.SUM | AggregationType.MEAN:\n return cls._aggregate(\n values=np.mean(value_arrays, axis=0), aggregation=aggregation\n ) # mean over dp ranks\n case AggregationType.MIN | AggregationType.MAX:\n return cls._aggregate(values=value_arrays.flatten(), aggregation=aggregation) # min/max over all values\n\n @classmethod\n def from_dict(cls, data: dict[str, Numeric], aggregation: str | AggregationType) -> dict[str, \"Metric\"]:\n return {key: cls(value=value, aggregation=aggregation) for key, value in data.items()}\n\n def init_list(self) -> \"Metric\":\n return Metric(aggregation=self.aggregation)\n"}95{"file_name": "verl__utils__model.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nUtilities to create common models from huggingface\n\"\"\"\n\nimport json\nimport os\nimport re\nimport warnings\nfrom dataclasses import dataclass\nfrom typing import Optional\n\nimport numpy as np\nimport torch\nfrom tensordict.tensorclass import NonTensorData\nfrom torch import nn\nfrom transformers import (\n AutoConfig,\n AutoModel,\n AutoModelForCausalLM,\n AutoModelForImageTextToText,\n AutoModelForSequenceClassification,\n AutoModelForTokenClassification,\n AutoModelForVision2Seq,\n GenerationConfig,\n MistralForSequenceClassification,\n PretrainedConfig,\n PreTrainedModel,\n)\nfrom transformers.modeling_outputs import CausalLMOutputWithPast\n\nfrom verl.models.registry import ModelRegistry\nfrom verl.utils.import_utils import is_trl_available\n\n\nclass LambdaLayer(nn.Module):\n def __init__(self, fn):\n super().__init__()\n self.fn = fn\n\n def forward(self, *args, **kwargs):\n return self.fn(*args, **kwargs)\n\n\ndef squeeze(x):\n return torch.squeeze(x, dim=-1)\n\n\ndef update_model_config(module_config, override_config_kwargs):\n \"\"\"Update the module config with the override_config_kwargs.\n Args:\n module_config: The module config from Huggingface Transformers.\n override_config_kwargs: The kwargs to override the module config.\n \"\"\"\n for key, val in override_config_kwargs.items():\n if isinstance(val, dict):\n update_model_config(getattr(module_config, key), val)\n else:\n setattr(module_config, key, val)\n\n\ndef get_huggingface_actor_config(model_name: str, override_config_kwargs=None, trust_remote_code=False) -> dict:\n if override_config_kwargs is None:\n override_config_kwargs = {}\n assert isinstance(override_config_kwargs, dict), (\n f\"override_config_kwargs must be a dict, got {type(override_config_kwargs)}\"\n )\n module_config = AutoConfig.from_pretrained(model_name, trust_remote_code=trust_remote_code)\n update_model_config(module_config, override_config_kwargs)\n\n return module_config\n\n\ndef get_generation_config(\n model: str,\n trust_remote_code: bool = False,\n) -> Optional[GenerationConfig]:\n try:\n return GenerationConfig.from_pretrained(model)\n except OSError: # Not found\n try:\n config = get_huggingface_actor_config(\n model,\n trust_remote_code=trust_remote_code,\n )\n return GenerationConfig.from_model_config(config)\n except OSError: # Not found\n return None\n\n\ndef create_huggingface_actor(model_name: str, override_config_kwargs=None, automodel_kwargs=None) -> nn.Module:\n \"\"\"\n\n Args:\n model_name:\n override_config_kwargs:\n\n Returns:\n\n \"\"\"\n if override_config_kwargs is None:\n override_config_kwargs = {}\n if automodel_kwargs is None:\n automodel_kwargs = {}\n assert isinstance(override_config_kwargs, dict), (\n f\"override_config_kwargs must be a dict, got {type(override_config_kwargs)}\"\n )\n module_config = get_huggingface_actor_config(\n model_name, override_config_kwargs, trust_remote_code=automodel_kwargs.get(\"trust_remote_code\", False)\n )\n module: nn.Module = AutoModelForCausalLM.from_config(module_config, **automodel_kwargs)\n return module\n\n\ndef create_huggingface_critic(model_name: str, override_config_kwargs=None, automodel_kwargs=None) -> nn.Module:\n \"\"\"\n\n Args:\n model_name:\n override_config_kwargs:\n\n Returns:\n\n \"\"\"\n critic_module: nn.Module = create_huggingface_actor(\n model_name, override_config_kwargs=override_config_kwargs, automodel_kwargs=automodel_kwargs\n )\n if automodel_kwargs is None:\n automodel_kwargs = {}\n torch_dtype = automodel_kwargs.get(\"torch_dtype\", torch.float32)\n critic_module.lm_head = nn.Sequential(\n nn.Linear(critic_module.config.hidden_size, 1, dtype=torch_dtype), LambdaLayer(fn=squeeze)\n )\n return critic_module\n\n\ndef get_model_size(model: nn.Module, scale=\"auto\"):\n n_params = sum(p.numel() for p in model.parameters())\n\n if scale == \"auto\":\n if n_params > 1e9:\n scale = \"B\"\n elif n_params > 1e6:\n scale = \"M\"\n elif n_params > 1e3:\n scale = \"K\"\n else:\n scale = \"\"\n\n if scale == \"B\":\n n_params = n_params / 1e9\n elif scale == \"M\":\n n_params = n_params / 1e6\n elif scale == \"K\":\n n_params = n_params / 1e3\n elif scale == \"\":\n pass\n else:\n raise NotImplementedError(f\"Unknown scale {scale}\")\n\n return n_params, scale\n\n\ndef print_model_size(model: nn.Module, name: str = None):\n n_params, scale = get_model_size(model, scale=\"auto\")\n if name is None:\n name = model.__class__.__name__\n print(f\"{name} contains {n_params:.2f}{scale} parameters\")\n\n\ndef create_random_mask(\n input_ids: torch.Tensor,\n max_ratio_of_valid_token: float,\n max_ratio_of_left_padding: float,\n min_ratio_of_valid_token: float = 0,\n):\n \"\"\"Create a random mask given input_ids. Support left padding and right padding.\n Process:\n - Sample valid token length\n - Sample left_padding length\n - Generate padding\n\n Args:\n input_ids:\n shape (batch_size, seq_len)\n\n Returns:\n\n \"\"\"\n assert max_ratio_of_valid_token > 0 and max_ratio_of_valid_token <= 1.0\n assert max_ratio_of_left_padding >= 0 and max_ratio_of_left_padding < 1.0\n assert min_ratio_of_valid_token <= max_ratio_of_valid_token\n\n batch_size, sequence_length = input_ids.shape\n max_num_valid_tokens = int(sequence_length * max_ratio_of_valid_token)\n min_num_valid_tokens = max(1, int(sequence_length * min_ratio_of_valid_token))\n max_left_padding = int(sequence_length * max_ratio_of_left_padding)\n assert max_num_valid_tokens + max_left_padding <= sequence_length\n assert max_num_valid_tokens > 0 and max_ratio_of_valid_token <= sequence_length\n masks = torch.ones_like(input_ids, dtype=torch.int64)\n # TODO: we can make this faster\n for i in range(batch_size):\n num_left_padding = np.random.randint(low=0, high=max_left_padding + 1, dtype=np.int64)\n num_valid = np.random.randint(low=min_num_valid_tokens, high=max_num_valid_tokens + 1, dtype=np.int64)\n\n for index in range(num_left_padding):\n masks[i, index] = 0\n\n for index in range(num_left_padding + num_valid, sequence_length):\n masks[i, index] = 0\n return masks\n\n\ndef compute_position_id_with_mask(mask):\n return torch.clip(torch.cumsum(mask, dim=-1) - 1, min=0, max=None)\n\n\ndef convert_weight_keys(state_dict: dict[str, torch.Tensor], model: PreTrainedModel):\n # convert state dict keys: https://github.com/huggingface/transformers/pull/38385\n if not hasattr(model, \"_checkpoint_conversion_mapping\"):\n return state_dict\n\n reverse_key_mapping = {v: k for k, v in model._checkpoint_conversion_mapping.items()}\n original_weights = {}\n for key, value in state_dict.items():\n for pattern, replacement in reverse_key_mapping.items():\n replacement = replacement.lstrip(\"^\") # strip off un-needed chars and patterns\n replacement = re.sub(r\"\\(.*\\)\", \"\", replacement)\n key, n_replace = re.subn(pattern, replacement, key)\n # Early exit of the loop\n if n_replace > 0:\n break\n\n original_weights[key] = value\n\n return original_weights\n\n\ndef check_exclude_modules(config, key: str) -> bool:\n \"\"\"\n A helper method to check if the passed module's key name matches any of the exclude modules in the adapter_config.\n Adapted from https://github.com/huggingface/peft/blob/main/src/peft/tuners/tuners_utils.py\n\n Args:\n config (`LoraConfig` | `LycorisConfig`): A config to match exclude modules from\n key (`str`): A key to search any matches in config\n\n Returns:\n True of match object if key matches any exclude modules from config, False if no match found\n \"\"\"\n if hasattr(config, \"exclude_modules\") and config.exclude_modules:\n if isinstance(config.exclude_modules, str):\n if re.fullmatch(config.exclude_modules, key):\n return True\n elif key in config.exclude_modules:\n return True\n elif any(key.endswith(f\".{exclude_key}\") for exclude_key in config.exclude_modules):\n return True\n return False\n\n\ndef check_target_modules(config, key: str) -> bool:\n \"\"\"\n A helper method to check if the passed module's key name matches any of the target modules in the adapter_config.\n Adapted from https://github.com/huggingface/peft/blob/main/src/peft/tuners/tuners_utils.py\n\n Args:\n config (`LoraConfig` | `LycorisConfig`): A config to match target modules from\n key (`str`): A key to search any matches in config\n\n Returns:\n True of match object if key matches any target modules from config, False if no match found\n \"\"\"\n if isinstance(config.target_modules, str):\n target_module_found = re.fullmatch(config.target_modules, key)\n elif key in config.target_modules:\n # this module is specified directly in target_modules\n target_module_found = True\n else:\n target_module_found = any(key.endswith(f\".{target_key}\") for target_key in config.target_modules)\n\n layer_indexes = getattr(config, \"layers_to_transform\", None)\n layers_pattern = getattr(config, \"layers_pattern\", None)\n\n is_using_layer_indexes = layer_indexes is not None and (\n len(layer_indexes) != 0 if isinstance(layer_indexes, list) else True\n )\n if is_using_layer_indexes and target_module_found:\n layer_index = None\n # TODO: It's still unclear how empty layers_pattern (None, [], or \"\") should behave\n # For now, empty layers_pattern means any layer pattern is ok\n if layers_pattern is None or len(layers_pattern) == 0:\n layer_index = re.match(r\".*\\.[^.]*\\.(\\d+)\\.\", key)\n else:\n layers_pattern = [layers_pattern] if isinstance(layers_pattern, str) else layers_pattern\n for pattern in layers_pattern:\n layer_index = re.match(rf\".*\\.{pattern}\\.(\\d+)\\.\", key)\n if layer_index is not None:\n break\n\n if layer_index is None:\n target_module_found = False\n else:\n layer_index = int(layer_index.group(1))\n if isinstance(layer_indexes, int):\n target_module_found = layer_index == layer_indexes\n else:\n target_module_found = layer_index in layer_indexes\n\n return target_module_found\n\n\ndef normalize_model_name(name, pp_rank, vpp_rank, transformer_config, layer_name=\"layers\"):\n \"\"\"\n Transform the model name in each model_chunk in each pp stage into the name in inference engine\n \"\"\"\n from verl.utils.megatron_utils import get_transformer_layer_offset\n\n layer_offset = get_transformer_layer_offset(pp_rank, vpp_rank, transformer_config)\n\n if layer_name in name: # belong to an intermediate layer\n split_name = name.split(\".\")\n # find the num next to split_name\n for i, name in enumerate(split_name):\n if name == layer_name:\n break\n layer_num_idx = i + 1\n # check the name\n assert len(split_name) >= layer_num_idx + 1, f\"split_name = {split_name}\"\n assert split_name[layer_num_idx].isdigit(), f\"split_name = {split_name}\"\n # increment layer_num_idx by layer_offset\n split_name[layer_num_idx] = str(int(split_name[layer_num_idx]) + layer_offset)\n name = \".\".join(split_name) # weight name in inference_tp_model\n return name\n\n\ndef normalize_pp_vpp_params(params, num_hidden_layers, layer_name=\"layers\"):\n \"\"\"\n Normalize the pp vpp params into a complete named parameters.\n This is useful when gather parameters from pp ranks and passed to a model without pp\n\n params: Iterable[List[Dict[str, param]]]\n params contains a list of pp, with a list of vpp named_parameters in each vpp chunk.\n output: Dict[str, param]\n\n \"\"\"\n pp_size = len(params)\n for pp_rank in range(len(params)):\n vpp_size = len(params[pp_rank])\n for vpp_rank in range(vpp_size):\n for name, param in params[pp_rank][vpp_rank].items():\n normalized_name = normalize_model_name(\n name, pp_rank, vpp_rank, pp_size, vpp_size, num_hidden_layers, layer_name=layer_name\n )\n yield normalized_name, param\n\n\ndef get_parallel_model_from_config(\n config, megatron_config, pre_process=None, post_process=None, share_embeddings_and_output_weights=False, value=False\n):\n from megatron.core import ModelParallelConfig\n\n assert isinstance(megatron_config, ModelParallelConfig)\n model_class = _get_parallel_model_architecture_from_config(config, value)\n\n model = model_class(\n config,\n megatron_config,\n pre_process=pre_process,\n post_process=post_process,\n share_embeddings_and_output_weights=share_embeddings_and_output_weights,\n )\n return model\n\n\ndef _get_parallel_model_architecture_from_config(config: PretrainedConfig, value=False) -> type[nn.Module]:\n architectures = getattr(config, \"architectures\", [])\n for arch in architectures:\n model_cls = ModelRegistry.load_model_cls(arch, value)\n print(\"after load model cls\")\n if model_cls is not None:\n return model_cls\n raise ValueError(\n f\"Model architectures {architectures} are not supported for now. Supported architectures: \"\n f\"{ModelRegistry.get_supported_archs()}\"\n )\n\n\ndef _load_hf_model(config, model_config, is_value_model):\n \"\"\"Helper function containing the loading hf model logic\"\"\"\n from accelerate import init_empty_weights\n from megatron.core import parallel_state as mpu\n\n from verl.models.mcore.saver import _megatron_calc_global_rank\n\n assert hasattr(model_config, \"architectures\"), \"architectures cannot be empty when load weight!\"\n architectures = getattr(model_config, \"architectures\", [])\n\n # get auto class\n auto_cls = get_hf_auto_model_class(model_config)\n\n if config.model.path.startswith(\"hdfs:\"):\n from verl.utils.fs import copy_to_local\n\n print(f\"start download from {config.model.path}\")\n local_model_path = copy_to_local(src=config.model.path, use_shm=config.model.get(\"use_shm\", False))\n print(\"finish download\")\n else:\n local_model_path = config.model.path\n print(f\"load from local dir {local_model_path}\")\n\n src_rank = _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=0, cp_rank=mpu.get_context_parallel_rank())\n cpu_init_weights = lambda: torch.device(\"cpu\")\n init_context = init_empty_weights if torch.distributed.get_rank() != src_rank else cpu_init_weights\n with init_context(), warnings.catch_warnings():\n warnings.simplefilter(\"ignore\")\n # TODO: to find a better way to load mistral7b-rm lm_head\n if \"mistral7b-rm\" in config.model.path:\n model = MistralForSequenceClassification.from_pretrained(\n local_model_path,\n torch_dtype=\"auto\",\n # device_map=\"auto\", # disable auto device_map, the HF weight is only loaded to CPU in src_rank\n # low_cpu_mem_usage=True\n ) # use score head instead of lm_head\n state_dict = model.state_dict()\n state_dict[\"lm_head.weight\"] = state_dict[\"score.weight\"]\n state_dict[\"model.embed_tokens.weight\"] = state_dict[\"model.embed_tokens.weight\"][\n :32000\n ] # workaround, 32001 -> 32000\n is_value_model = True\n else:\n model = auto_cls.from_pretrained(\n local_model_path,\n torch_dtype=\"auto\",\n # device_map=\"auto\", # disable auto device_map, the HF weight is only loaded to CPU in src_rank\n # low_cpu_mem_usage=True\n )\n state_dict = model.state_dict()\n\n return architectures, model, state_dict, is_value_model\n\n\ndef get_hf_model_path(config):\n if config.model.path.startswith(\"hdfs:\"):\n from verl.utils.fs import copy_to_local\n\n local_model_path = copy_to_local(src=config.model.path, use_shm=config.model.get(\"use_shm\", False))\n else:\n local_model_path = config.model.path\n return local_model_path\n\n\ndef load_megatron_model_weights(config, model_config, parallel_model, params_dtype, is_value_model=False):\n \"\"\"Load weights for verl customized model.\"\"\"\n architectures, model, state_dict, is_value_model = _load_hf_model(config, model_config, is_value_model)\n\n from verl.models.weight_loader_registry import get_weight_loader\n\n print(f\"before weight loader: architectures = {architectures}...\")\n for arch in architectures:\n print(f\"call weight loader arch = {arch}, model config = {model.config}\")\n weight_loader = get_weight_loader(arch)\n weight_loader(\n state_dict=state_dict,\n wrapped_models=parallel_model,\n config=model.config,\n params_dtype=params_dtype,\n is_value_model=is_value_model,\n tie_word_embeddings=model_config.tie_word_embeddings,\n )\n return model.config\n\n\ndef load_megatron_gptmodel_weights(config, model_config, parallel_model, params_dtype, is_value_model=False):\n \"\"\"Load weights for mcore GPT model.\"\"\"\n _, model, state_dict, is_value_model = _load_hf_model(config, model_config, is_value_model)\n\n from verl.models.mcore.loader import load_state_dict_to_megatron_gptmodel\n\n load_state_dict_to_megatron_gptmodel(\n state_dict=state_dict,\n wrapped_models=parallel_model,\n config=model.config,\n params_dtype=params_dtype,\n is_value_model=is_value_model,\n )\n del state_dict, model\n\n\n# pad input_ids_rmpad, cu_seqlens and max_seqlen_in_batch to be divisible by tp\ndef pad_packed_inputs(unpad_tokens: torch.Tensor, cu_seqlens, max_seqlen_in_batch, size):\n \"\"\"pad the tokens such that the total length is a multiple of size.\n This function is useful when applying sequence parallel and context parallel\n\n Args:\n unpad_tokens: (total_nnz, ...). Tokens after removing padding\n cu_seqlens: (total_nnz + 1,)\n max_seqlen_in_batch: int\n\n Returns:\n\n \"\"\"\n F = nn.functional\n\n total_nnz = unpad_tokens.shape[0]\n\n pad_size = 0 if total_nnz % size == 0 else size - total_nnz % size\n\n # we assume adding a new data in the batch with seqlen pad_size\n if pad_size > 0:\n if unpad_tokens.ndim == 1:\n unpad_tokens = F.pad(unpad_tokens, (0, pad_size))\n elif unpad_tokens.ndim == 2:\n unpad_tokens = F.pad(unpad_tokens, (0, 0, 0, pad_size))\n else:\n raise NotImplementedError(f\"Padding dim {unpad_tokens.ndim()} is not supported\")\n\n cu_seqlens = F.pad(cu_seqlens, (0, 1), value=pad_size + cu_seqlens[-1])\n max_seqlen_in_batch = max(max_seqlen_in_batch, pad_size)\n\n return unpad_tokens, cu_seqlens, max_seqlen_in_batch\n\n\ndef load_mcore_dist_weights(parallel_model, dist_weight_path, is_value_model=False, prefix=\"\"):\n from megatron.core import dist_checkpointing\n from megatron.core.dist_checkpointing.serialization import StrictHandling\n\n from verl.utils.megatron_utils import unwrap_model\n\n # strict = StrictHandling.IGNORE_ALL if is_value_model else StrictHandling.ASSUME_OK_UNEXPECTED\n strict = StrictHandling.ASSUME_OK_UNEXPECTED\n for model in parallel_model:\n ssd = unwrap_model(model).sharded_state_dict(prefix=prefix)\n if is_value_model:\n for k in list(ssd.keys()):\n if \"output_layer\" in k:\n ssd.pop(k)\n dist_checkpointing.load(ssd, dist_weight_path, strict=strict)\n\n return\n\n\ndef get_parallel_gptmodel_from_config(\n tfconfig, hf_config, pre_process=None, post_process=None, share_embeddings_and_output_weights=False, value=False\n):\n from megatron.core.models.gpt.gpt_layer_specs import get_gpt_decoder_block_spec\n from megatron.core.models.gpt.gpt_model import GPTModel\n\n use_te = True\n assert tfconfig.normalization == \"RMSNorm\", \"only RMSNorm is supported for now\"\n transformer_layer_spec = get_gpt_decoder_block_spec(tfconfig, use_transformer_engine=use_te)\n rope_scaling_args = {}\n if hf_config.rope_scaling is not None:\n assert hf_config.rope_scaling[\"type\"] == \"linear\", \"only linear scaling is supported for now\"\n rope_scaling_args[\"seq_len_interpolation_factor\"] = hf_config.rope_scaling[\"factor\"]\n parallel_model = GPTModel(\n config=tfconfig,\n transformer_layer_spec=transformer_layer_spec,\n vocab_size=hf_config.vocab_size,\n max_sequence_length=hf_config.max_position_embeddings,\n pre_process=pre_process,\n post_process=post_process,\n share_embeddings_and_output_weights=share_embeddings_and_output_weights,\n position_embedding_type=\"rope\",\n rotary_base=hf_config.rope_theta,\n **rope_scaling_args,\n )\n # # for layer in parallel_model.decoder.layers:\n # layer.self_attention.core_attention.flash_attention.softmax_scale = None\n if post_process and value:\n from verl.models.llama.megatron.layers.parallel_linear import LinearForLastLayer\n\n parallel_model.output_layer = LinearForLastLayer(\n input_size=tfconfig.hidden_size, output_size=1, config=tfconfig\n )\n return parallel_model\n\n\ndef patch_valuehead_model(model) -> None:\n from types import MethodType\n\n from transformers import PreTrainedModel\n from trl import AutoModelForCausalLMWithValueHead\n\n def tie_weights(self: \"AutoModelForCausalLMWithValueHead\") -> None:\n if isinstance(self.pretrained_model, PreTrainedModel):\n self.pretrained_model.tie_weights()\n\n def get_input_embeddings(self: \"AutoModelForCausalLMWithValueHead\") -> torch.nn.Module:\n if isinstance(self.pretrained_model, PreTrainedModel):\n return self.pretrained_model.get_input_embeddings()\n\n def get_output_embeddings(self: \"AutoModelForCausalLMWithValueHead\") -> torch.nn.Module:\n if isinstance(self.pretrained_model, PreTrainedModel):\n return self.pretrained_model.get_output_embeddings()\n\n def can_generate(self):\n return False\n\n ignore_modules = [name for name, _ in model.named_parameters() if \"pretrained_model\" in name]\n model._keys_to_ignore_on_save = ignore_modules\n model.tie_weights = MethodType(tie_weights, model)\n model.get_input_embeddings = MethodType(get_input_embeddings, model)\n model.get_output_embeddings = MethodType(get_output_embeddings, model)\n model.can_generate = MethodType(can_generate, model)\n model._no_split_modules = getattr(model.pretrained_model, \"_no_split_modules\", [])\n\n\ndef load_valuehead_model(local_path, torch_dtype, model_config, trust_remote_code):\n from transformers import AutoModelForCausalLM, AutoModelForTokenClassification, AutoModelForVision2Seq\n\n try:\n model = AutoModelForTokenClassification.from_pretrained(\n pretrained_model_name_or_path=local_path,\n torch_dtype=torch_dtype,\n config=model_config,\n attn_implementation=\"flash_attention_2\",\n trust_remote_code=trust_remote_code,\n )\n return model\n except BaseException as e:\n if not is_trl_available():\n raise RuntimeError(\n f\"model({local_path}) is not a value head model, please install trl to make it valid\"\n ) from e\n\n assert is_trl_available()\n\n from trl import AutoModelForCausalLMWithValueHead\n\n if type(model_config) in AutoModelForVision2Seq._model_mapping.keys():\n module_class = AutoModelForVision2Seq\n else:\n module_class = AutoModelForCausalLM\n ori_model = module_class.from_pretrained(\n pretrained_model_name_or_path=local_path,\n torch_dtype=torch_dtype,\n config=model_config,\n attn_implementation=\"flash_attention_2\",\n trust_remote_code=trust_remote_code,\n )\n model = AutoModelForCausalLMWithValueHead.from_pretrained(ori_model)\n patch_valuehead_model(model)\n return model\n\n\n_architecture_to_auto_class = {\n \"ForCausalLM\": AutoModelForCausalLM,\n \"ForVision2Seq\": AutoModelForVision2Seq,\n \"ForTokenClassification\": AutoModelForTokenClassification,\n \"ForSequenceClassification\": AutoModelForSequenceClassification,\n}\n\n\ndef get_hf_auto_model_class(hf_config):\n has_remote_code = hasattr(hf_config, \"auto_map\") and any(\n hf_config.architectures[0] in val for val in hf_config.auto_map.values()\n )\n if has_remote_code:\n auto_class = next(k for k, v in hf_config.auto_map.items() if hf_config.architectures[0] in v)\n match auto_class:\n case \"AutoModelForVision2Seq\":\n actor_module_class = AutoModelForVision2Seq\n case \"AutoModelForCausalLM\":\n actor_module_class = AutoModelForCausalLM\n case \"AutoModelForImageTextToText\":\n actor_module_class = AutoModelForImageTextToText\n case _:\n actor_module_class = AutoModel\n else:\n actor_module_class = AutoModel\n # For VLM models, we use type to check instead of architecture\n if type(hf_config) in AutoModelForImageTextToText._model_mapping.keys():\n actor_module_class = AutoModelForImageTextToText\n else:\n for key, cls in _architecture_to_auto_class.items():\n if key in hf_config.architectures[0]:\n actor_module_class = cls\n break\n\n return actor_module_class\n\n\ndef extract_multi_modal_inputs(\n batch_data: list[dict[str, torch.Tensor]],\n indices: Optional[list[int]] = None,\n) -> dict[str, torch.Tensor | list[torch.Tensor]]:\n \"\"\"\n Extract and process multi-modal inputs from a batch.\n\n Args:\n batch_data (list[dict[str, torch.Tensor]]): The batch containing potential multi-modal inputs\n indices (Optional[list[int]]): If provided, only extract inputs at these indices\n\n Returns:\n dict[str, torch.Tensor | list[torch.Tensor]]: Processed multi-modal inputs ready for model consumption\n\n \"\"\"\n multi_modal_inputs = {}\n multi_modal_inputs_collected = {}\n has_image_bound = False\n\n selected_batch_data = batch_data\n if indices is not None:\n selected_batch_data = [batch_data[i] for i in indices if i < len(batch_data)]\n\n for inputs in selected_batch_data:\n inputs = inputs.data if isinstance(inputs, NonTensorData) else inputs\n # Mixed pure text and multi-modal dataset.\n if inputs is None:\n continue\n if \"image_bound\" in inputs:\n has_image_bound = True\n for key, value in inputs.items():\n if value is not None:\n if key not in multi_modal_inputs_collected:\n multi_modal_inputs_collected[key] = []\n multi_modal_inputs_collected[key].append(value)\n\n for key, values in multi_modal_inputs_collected.items():\n if has_image_bound: # minicpm-o logic\n multi_modal_inputs[key] = values\n else:\n multi_modal_inputs[key] = torch.cat(values, dim=0)\n\n return multi_modal_inputs\n\n\ndef get_lora_rank_from_adapter(adapter_path: str | os.PathLike) -> int:\n \"\"\"\n Extract LoRA rank from adapter configuration file.\n\n Args:\n adapter_path: Path to LoRA adapter directory\n\n Returns:\n LoRA rank value from adapter_config.json\n\n Raises:\n FileNotFoundError: If adapter path or config file doesn't exist\n ValueError: If config file is invalid or missing rank\n \"\"\"\n adapter_path = os.path.abspath(os.path.expanduser(str(adapter_path)))\n\n if not os.path.exists(adapter_path):\n raise FileNotFoundError(f\"LoRA adapter path not found: {adapter_path}\")\n\n config_path = os.path.join(adapter_path, \"adapter_config.json\")\n if not os.path.exists(config_path):\n raise FileNotFoundError(f\"adapter_config.json not found in {adapter_path}\")\n\n try:\n with open(config_path, encoding=\"utf-8\") as f:\n config = json.load(f)\n if \"r\" not in config:\n raise ValueError(f\"LoRA rank 'r' not found in {config_path}\")\n return int(config[\"r\"])\n except json.JSONDecodeError as e:\n raise ValueError(f\"Invalid JSON in {config_path}: {e}\") from e\n except (KeyError, ValueError) as e:\n raise ValueError(f\"Cannot parse LoRA rank from {config_path}: {e}\") from e\n\n\n@dataclass\nclass CausalLMOutputForPPO(CausalLMOutputWithPast):\n log_probs: Optional[torch.FloatTensor] = None\n entropy: Optional[torch.FloatTensor] = None\n"}96{"file_name": "verl__utils__npu_flash_attn_utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nimport torch\nimport torch.nn.functional as F\nfrom einops import rearrange, repeat\n\n\n# Copied from https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/bert_padding.py\nclass IndexFirstAxis(torch.autograd.Function):\n @staticmethod\n def forward(ctx, input, indices):\n ctx.save_for_backward(indices)\n assert input.ndim >= 2\n ctx.first_axis_dim, other_shape = input.shape[0], input.shape[1:]\n second_dim = other_shape.numel()\n # TD [2022-03-04] For some reason torch.gather is a bit faster than indexing.\n # return input[indices]\n return torch.gather(rearrange(input, \"b ... -> b (...)\"), 0, repeat(indices, \"z -> z d\", d=second_dim)).reshape(\n -1, *other_shape\n )\n\n @staticmethod\n def backward(ctx, grad_output):\n (indices,) = ctx.saved_tensors\n assert grad_output.ndim >= 2\n other_shape = grad_output.shape[1:]\n grad_output = rearrange(grad_output, \"b ... -> b (...)\")\n grad_input = torch.zeros(\n [ctx.first_axis_dim, grad_output.shape[1]],\n device=grad_output.device,\n dtype=grad_output.dtype,\n )\n # TD [2022-03-04] For some reason torch.scatter is a bit faster than indexing.\n # grad_input[indices] = grad_output\n grad_input.scatter_(0, repeat(indices, \"z -> z d\", d=grad_output.shape[1]), grad_output)\n return grad_input.reshape(ctx.first_axis_dim, *other_shape), None\n\n\nindex_first_axis = IndexFirstAxis.apply\n\n\n# Copied from https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/bert_padding.py\nclass IndexPutFirstAxis(torch.autograd.Function):\n @staticmethod\n def forward(ctx, values, indices, first_axis_dim):\n ctx.save_for_backward(indices)\n assert indices.ndim == 1\n assert values.ndim >= 2\n output = torch.zeros(first_axis_dim, *values.shape[1:], device=values.device, dtype=values.dtype)\n # TD [2022-03-04] For some reason torch.scatter is a bit faster than indexing.\n output[indices] = values\n # output.scatter_(0, repeat(indices, 'z -> z d', d=values.shape[1]), values)\n return output\n\n @staticmethod\n def backward(ctx, grad_output):\n (indices,) = ctx.saved_tensors\n # TD [2022-03-04] For some reason torch.gather is a bit faster than indexing.\n grad_values = grad_output[indices]\n # grad_values = torch.gather(grad_output, 0, repeat(indices, 'z -> z d', d=grad_output.shape[1]))\n return grad_values, None, None\n\n\nindex_put_first_axis = IndexPutFirstAxis.apply\n\n\n# Copied from https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/bert_padding.py\ndef pad_input(hidden_states, indices, batch, seqlen):\n \"\"\"\n Arguments:\n hidden_states: (total_nnz, ...), where total_nnz = number of tokens in selected in attention_mask.\n indices: (total_nnz), the indices that represent the non-masked tokens of the original padded input sequence.\n batch: int, batch size for the padded sequence.\n seqlen: int, maximum sequence length for the padded sequence.\n Return:\n hidden_states: (batch, seqlen, ...)\n \"\"\"\n # dim = hidden_states.shape[-1]\n # output = torch.zeros((batch * seqlen), dim, device=hidden_states.device, dtype=hidden_states.dtype)\n # output[indices] = hidden_states\n output = index_put_first_axis(hidden_states, indices, batch * seqlen)\n return rearrange(output, \"(b s) ... -> b s ...\", b=batch)\n\n\n# Copied from https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/bert_padding.py\ndef unpad_input(hidden_states, attention_mask, unused_mask=None):\n \"\"\"\n Arguments:\n hidden_states: (batch, seqlen, ...)\n attention_mask: (batch, seqlen), bool / int, 1 means valid and 0 means not valid.\n unused_mask: (batch, seqlen), bool / int, 1 means the element is allocated but unused.\n Return:\n hidden_states: (total_nnz, ...), where total_nnz = number of tokens selected in attention_mask + unused_mask.\n indices: (total_nnz), the indices of masked tokens from the flattened input sequence.\n cu_seqlens: (batch + 1), the cumulative sequence lengths, used to index into hidden_states.\n max_seqlen_in_batch: int\n seqused: (batch), returns the number of tokens selected in attention_mask + unused_mask.\n \"\"\"\n all_masks = (attention_mask + unused_mask) if unused_mask is not None else attention_mask\n seqlens_in_batch = all_masks.sum(dim=-1, dtype=torch.int32)\n used_seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)\n indices = torch.nonzero(all_masks.flatten(), as_tuple=False).flatten()\n max_seqlen_in_batch = seqlens_in_batch.max().item()\n cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))\n # TD [2022-03-04] We don't want to index with a bool mask, because Pytorch will expand the\n # bool mask, then call nonzero to get the indices, then index with those. The indices is @dim\n # times larger than it needs to be, wasting memory. It's faster and more memory-efficient to\n # index with integer indices. Moreover, torch's index is a bit slower than it needs to be,\n # so we write custom forward and backward to make it a bit faster.\n return (\n index_first_axis(rearrange(hidden_states, \"b s ... -> (b s) ...\"), indices),\n indices,\n cu_seqlens,\n max_seqlen_in_batch,\n used_seqlens_in_batch,\n )\n"}97{"file_name": "verl__utils__profiler__empty_annotations.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Callable, Optional\n\n\ndef mark_start_range(\n message: Optional[str] = None,\n color: Optional[str] = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n) -> None:\n pass\n\n\ndef mark_end_range(range_id: str) -> None:\n pass\n\n\ndef mark_annotate(\n message: Optional[str] = None,\n color: Optional[str] = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n) -> Callable:\n def decorator(func):\n return func\n\n return decorator\n"}98{"file_name": "verl__utils__profiler__mstx_profile.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n# Inspired from https://gitee.com/ascend/MindSpeed-RL/blob/master/mindspeed_rl/utils/utils.py\nimport functools\nimport logging\nimport os\nfrom contextlib import contextmanager\nfrom typing import Any, Callable, Optional\n\nimport torch_npu\nfrom packaging import version\nfrom torch_npu.npu import mstx\n\nfrom .config import NPUToolConfig\nfrom .profile import DistProfiler, ProfilerConfig\n\n\ndef mark_start_range(message: Optional[str] = None) -> None:\n \"\"\"Start a mark range in the profiler.\n\n Args:\n message (str, optional):\n The message to be displayed in the profiler. Defaults to None.\n \"\"\"\n return mstx.range_start(message=message)\n\n\ndef mark_end_range(range_id: str) -> None:\n \"\"\"End a mark range in the profiler.\n\n Args:\n range_id (str):\n The id of the mark range to end.\n \"\"\"\n return mstx.range_end(range_id)\n\n\ndef mark_annotate(message: Optional[str] = None) -> Callable:\n \"\"\"Decorate a function to annotate a mark range along with the function life cycle.\n\n Args:\n message (str, optional):\n The message to be displayed in the profiler. Defaults to None.\n \"\"\"\n\n def decorator(func):\n profile_message = message or func.__name__\n return mstx.mstx_range(profile_message)(func)\n\n return decorator\n\n\n@contextmanager\ndef marked_timer(name: str, timing_raw: dict[str, float], *args: Any, **kwargs: Any) -> None:\n \"\"\"Context manager for timing with MSTX markers.\n\n This utility function measures the execution time of code within its context,\n accumulates the timing information, and adds MSTX markers for profiling.\n\n Args:\n name (str): The name/identifier for this timing measurement.\n timing_raw (Dict[str, float]): Dictionary to store timing information.\n\n Yields:\n None: This is a context manager that yields control back to the code block.\n \"\"\"\n if args:\n logging.warning(f\"Args are not supported in mstx_profile, but received: {args}\")\n if kwargs:\n logging.warning(f\"Kwargs are not supported in mstx_profile, but received: {kwargs}\")\n mark_range = mark_start_range(message=name)\n from .performance import _timer\n\n yield from _timer(name, timing_raw)\n mark_end_range(mark_range)\n\n\ndef get_npu_profiler(\n contents: list[str],\n profile_level: str,\n profile_save_path: str,\n analysis: bool,\n role: Optional[str] = None,\n profile_step: Optional[str] = None,\n):\n \"\"\"Generate and return an NPU profiler object.\n\n Args:\n contents (list[str]):\n A list of options to control the collection content,\n such as npu, cpu, memory, shapes, module, stack.\n profile_level (str):\n The collection level, which can be set to level_none,\n level0, level1 and level2.\n profile_save_path (str):\n The path to save the collected data.\n analysis (bool):\n Whether to enables automatic data parsing.\n role (str, optional):\n The role of the current data collection. Defaults to None.\n profile_step(str, optional):\n The current training step. Defaults to None.\n \"\"\"\n if profile_level == \"level_none\":\n level = torch_npu.profiler.ProfilerLevel.Level_none\n elif profile_level == \"level0\":\n level = torch_npu.profiler.ProfilerLevel.Level0\n elif profile_level == \"level1\":\n level = torch_npu.profiler.ProfilerLevel.Level1\n elif profile_level == \"level2\":\n level = torch_npu.profiler.ProfilerLevel.Level2\n else:\n raise ValueError(f\"level only supports level0, 1, 2, and level_none, but gets {profile_level}\")\n\n if profile_step:\n profile_save_path = os.path.join(profile_save_path, profile_step)\n if role:\n profile_save_path = os.path.join(profile_save_path, role)\n\n # The ability to filter communication via mstx_domain_exclude requires torch_npu==2.1 or higher.\n if version.parse(torch_npu.__version__) < version.parse(\"2.1\"):\n raise RuntimeError(\"torch_npu==2.1 or higher is required to use mstx_domain_exclude\")\n\n experimental_config = torch_npu.profiler._ExperimentalConfig(\n profiler_level=level,\n export_type=torch_npu.profiler.ExportType.Db,\n data_simplification=True,\n msprof_tx=True,\n mstx_domain_exclude=[\"communication\"],\n )\n\n activites = []\n if contents is None or \"npu\" in contents:\n activites.append(torch_npu.profiler.ProfilerActivity.NPU)\n if contents is None or \"cpu\" in contents:\n activites.append(torch_npu.profiler.ProfilerActivity.CPU)\n\n prof = torch_npu.profiler.profile(\n with_modules=contents is None or \"module\" in contents,\n with_stack=contents is None or \"stack\" in contents,\n record_shapes=contents is None or \"shapes\" in contents,\n profile_memory=contents is None or \"memory\" in contents,\n activities=activites,\n on_trace_ready=torch_npu.profiler.tensorboard_trace_handler(profile_save_path, analyse_flag=analysis),\n experimental_config=experimental_config,\n )\n return prof\n\n\nclass NPUProfiler(DistProfiler):\n \"\"\"\n NPU profiler. Initialized in a worker to control the NPU profiler.\n \"\"\"\n\n _define_count = 0\n\n def __init__(self, rank: int, config: ProfilerConfig, tool_config: NPUToolConfig, **kwargs):\n \"\"\"Initialize the NsightSystemsProfiler.\n\n Args:\n rank (int): The rank of the current process.\n config (Optional[ProfilerConfig]): Configuration for the profiler. If None, a default configuration is used.\n tool_config (NPUToolConfig): The config to control npu profiler behavior.\n \"\"\"\n if not config:\n config = ProfilerConfig(ranks=[], enable=False)\n if not tool_config:\n assert not config.enable, \"tool_config must be set when profiler is enabled\"\n self.discrete: bool = tool_config.discrete\n self.profile_npu = None\n self.profile_contents = tool_config.contents\n self.profile_level = tool_config.level\n self.profile_save_path = config.save_path\n self.analysis = tool_config.analysis\n\n def start(self, **kwargs):\n role = kwargs.get(\"role\", None)\n if not self.discrete and NPUProfiler._define_count == 0:\n self.profile_npu = get_npu_profiler(\n contents=self.profile_contents,\n profile_level=self.profile_level,\n profile_save_path=self.profile_save_path,\n analysis=self.analysis,\n role=role,\n )\n self.profile_npu.start()\n NPUProfiler._define_count += 1\n\n def stop(self):\n if not self.discrete and NPUProfiler._define_count == 1:\n self.profile_npu.step()\n self.profile_npu.stop()\n NPUProfiler._define_count -= 1\n\n def annotate(self, message: Optional[str] = None, role: Optional[str] = None, **kwargs_outer) -> Callable:\n \"\"\"Decorate a Worker member function to profile the current rank in the current training step.\n\n Requires the target function to be a member function of a Worker,\n which has a member field `profiler` with NPUProfiler type.\n\n Args:\n message (str, optional):\n The message to be displayed in the profiler. Defaults to None.\n role (str, optional):\n The role of the current data collection. Defaults to None.\n \"\"\"\n\n def decorator(func):\n @functools.wraps(func)\n def wrapper(*args, **kwargs_inner):\n profile_name = message or func.__name__\n discrete_mode = self.discrete\n\n if not discrete_mode:\n mark_range = mark_start_range(message=profile_name)\n else:\n profile_npu = get_npu_profiler(\n contents=self.profile_contents,\n profile_level=self.profile_level,\n profile_save_path=self.profile_save_path,\n analysis=self.analysis,\n role=role,\n )\n profile_npu.start()\n mark_range = mark_start_range(message=profile_name)\n\n result = func(*args, **kwargs_inner)\n\n if not discrete_mode:\n mark_end_range(mark_range)\n else:\n mark_end_range(mark_range)\n profile_npu.step()\n profile_npu.stop()\n\n return result\n\n return wrapper\n\n return decorator\n"}99{"file_name": "verl__utils__profiler__nvtx_profile.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport functools\nfrom contextlib import contextmanager\nfrom typing import Callable, Optional\n\nimport nvtx\nimport torch\n\nfrom .config import NsightToolConfig\nfrom .profile import DistProfiler, ProfilerConfig\n\n\ndef mark_start_range(\n message: Optional[str] = None,\n color: Optional[str] = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n) -> None:\n \"\"\"Start a mark range in the profiler.\n\n Args:\n message (str, optional):\n The message to be displayed in the profiler. Defaults to None.\n color (str, optional):\n The color of the range. Defaults to None.\n domain (str, optional):\n The domain of the range. Defaults to None.\n category (str, optional):\n The category of the range. Defaults to None.\n \"\"\"\n return nvtx.start_range(message=message, color=color, domain=domain, category=category)\n\n\ndef mark_end_range(range_id: str) -> None:\n \"\"\"End a mark range in the profiler.\n\n Args:\n range_id (str):\n The id of the mark range to end.\n \"\"\"\n return nvtx.end_range(range_id)\n\n\ndef mark_annotate(\n message: Optional[str] = None,\n color: Optional[str] = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n) -> Callable:\n \"\"\"Decorate a function to annotate a mark range along with the function life cycle.\n\n Args:\n message (str, optional):\n The message to be displayed in the profiler. Defaults to None.\n color (str, optional):\n The color of the range. Defaults to None.\n domain (str, optional):\n The domain of the range. Defaults to None.\n category (str, optional):\n The category of the range. Defaults to None.\n \"\"\"\n\n def decorator(func):\n profile_message = message or func.__name__\n return nvtx.annotate(profile_message, color=color, domain=domain, category=category)(func)\n\n return decorator\n\n\n@contextmanager\ndef marked_timer(\n name: str,\n timing_raw: dict[str, float],\n color: str = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n):\n \"\"\"Context manager for timing with NVTX markers.\n\n This utility function measures the execution time of code within its context,\n accumulates the timing information, and adds NVTX markers for profiling.\n\n Args:\n name (str): The name/identifier for this timing measurement.\n timing_raw (Dict[str, float]): Dictionary to store timing information.\n color (Optional[str]): Color for the NVTX marker. Defaults to None.\n domain (Optional[str]): Domain for the NVTX marker. Defaults to None.\n category (Optional[str]): Category for the NVTX marker. Defaults to None.\n\n Yields:\n None: This is a context manager that yields control back to the code block.\n \"\"\"\n mark_range = mark_start_range(message=name, color=color, domain=domain, category=category)\n from .performance import _timer\n\n yield from _timer(name, timing_raw)\n mark_end_range(mark_range)\n\n\nclass NsightSystemsProfiler(DistProfiler):\n \"\"\"Nsight system profiler. Installed in a worker to control the Nsight system profiler.\"\"\"\n\n def __init__(self, rank: int, config: Optional[ProfilerConfig], tool_config: Optional[NsightToolConfig], **kwargs):\n \"\"\"Initialize the NsightSystemsProfiler.\n\n Args:\n rank (int): The rank of the current process.\n config (Optional[ProfilerConfig]): Configuration for the profiler. If None, a default configuration is used.\n \"\"\"\n # If no configuration is provided, create a default ProfilerConfig with an empty list of ranks\n if not config:\n config = ProfilerConfig(ranks=[])\n if not tool_config:\n assert not config.enable, \"tool_config must be provided when profiler is enabled\"\n self.discrete: bool = tool_config.discrete\n\n def start(self, **kwargs):\n if not self.discrete:\n torch.cuda.profiler.start()\n\n def stop(self):\n if not self.discrete:\n torch.cuda.profiler.stop()\n\n def annotate(\n self,\n message: Optional[str] = None,\n color: Optional[str] = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n **kwargs_outer,\n ) -> Callable:\n \"\"\"Decorate a Worker member function to profile the current rank in the current training step.\n\n Requires the target function to be a member function of a Worker, which has a member field `profiler` with\n NightSystemsProfiler type.\n\n Args:\n message (str, optional):\n The message to be displayed in the profiler. Defaults to None.\n color (str, optional):\n The color of the range. Defaults to None.\n domain (str, optional):\n The domain of the range. Defaults to None.\n category (str, optional):\n The category of the range. Defaults to None.\n \"\"\"\n\n def decorator(func):\n @functools.wraps(func)\n def wrapper(*args, **kwargs_inner):\n profile_name = message or func.__name__\n\n if self.discrete:\n torch.cuda.profiler.start()\n mark_range = mark_start_range(message=profile_name, color=color, domain=domain, category=category)\n\n result = func(*args, **kwargs_inner)\n\n mark_end_range(mark_range)\n if self.discrete:\n torch.cuda.profiler.stop()\n\n return result\n\n return wrapper\n\n return decorator\n"}100{"file_name": "verl__utils__profiler__performance.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport datetime\nimport inspect\nimport logging\nfrom contextlib import contextmanager\nfrom typing import Any, Optional\n\nimport torch\nimport torch.distributed as dist\nfrom codetiming import Timer\n\nfrom verl.utils.device import get_device_id, get_torch_device\nfrom verl.utils.logger import DecoratorLoggerBase\n\n\ndef _get_current_mem_info(unit: str = \"GB\", precision: int = 2) -> tuple[str]:\n \"\"\"Get current memory usage.\n\n Note that CPU device memory info is always 0.\n\n Args:\n unit (str, optional): The unit of memory measurement. Defaults to \"GB\".\n precision (int, optional): The number of decimal places to round memory values. Defaults to 2.\n\n Returns:\n tuple[str]: A tuple containing memory allocated, memory reserved, memory used, and memory total\n in the specified unit.\n \"\"\"\n assert unit in [\"GB\", \"MB\", \"KB\"]\n device = get_torch_device()\n # torch.cpu.memory_allocated() does not exist\n if device == torch.cpu:\n return \"0.00\", \"0.00\", \"0.00\", \"0.00\"\n\n divisor = 1024**3 if unit == \"GB\" else 1024**2 if unit == \"MB\" else 1024\n mem_allocated = get_torch_device().memory_allocated()\n mem_reserved = get_torch_device().memory_reserved()\n # use get_torch_device().mem_get_info to profile device memory\n # since vllm's sleep mode works below pytorch\n # see https://github.com/vllm-project/vllm/pull/11743#issuecomment-2754338119\n mem_free, mem_total = get_torch_device().mem_get_info()\n mem_used = mem_total - mem_free\n mem_allocated = f\"{mem_allocated / divisor:.{precision}f}\"\n mem_reserved = f\"{mem_reserved / divisor:.{precision}f}\"\n mem_used = f\"{mem_used / divisor:.{precision}f}\"\n mem_total = f\"{mem_total / divisor:.{precision}f}\"\n return mem_allocated, mem_reserved, mem_used, mem_total\n\n\ndef log_gpu_memory_usage(head: str, logger: logging.Logger = None, level=logging.DEBUG, rank: int = 0):\n \"\"\"Log GPU memory usage information.\n\n Args:\n head (str): A descriptive header for the memory usage log message.\n logger (logging.Logger, optional): Logger instance to use for logging. If None, prints to stdout.\n level: Logging level to use. Defaults to logging.DEBUG.\n rank (int): The rank of the process to log memory for. Defaults to 0.\n \"\"\"\n if (not dist.is_initialized()) or (rank is None) or (dist.get_rank() == rank):\n mem_allocated, mem_reserved, mem_used, mem_total = _get_current_mem_info()\n message = (\n f\"{head}, memory allocated (GB): {mem_allocated}, memory reserved (GB): {mem_reserved}, \"\n f\"device memory used/total (GB): {mem_used}/{mem_total}\"\n )\n\n if logger is None:\n print(message)\n else:\n logger.log(msg=message, level=level)\n\n\nclass GPUMemoryLogger(DecoratorLoggerBase):\n \"\"\"A decorator class to log GPU memory usage.\n\n Example:\n >>> from verl.utils.profiler.performance import GPUMemoryLogger\n >>> @GPUMemoryLogger(role=\"actor\")\n >>> def update_actor(self, batch):\n ... # real actor update logics\n ... return\n \"\"\"\n\n def __init__(self, role: str, logger: logging.Logger = None, level=logging.DEBUG, log_only_rank_0: bool = True):\n if dist.is_initialized() and dist.get_world_size() > 1:\n rank = dist.get_rank()\n else:\n rank = 0\n super().__init__(role, logger, level, rank, log_only_rank_0)\n\n def __call__(self, decorated_function: callable):\n def f(*args, **kwargs):\n return self.log(decorated_function, *args, **kwargs)\n\n return f\n\n def log(self, func, *args, **kwargs):\n name = func.__name__\n mem_allocated, mem_reserved, mem_used, mem_total = _get_current_mem_info()\n message = (\n f\"Before {name}, memory allocated (GB): {mem_allocated}, memory reserved (GB): {mem_reserved}, \"\n f\"device memory used/total (GB): {mem_used}/{mem_total}\"\n )\n self.logging_function(message)\n\n output = func(*args, **kwargs)\n\n mem_allocated, mem_reserved, mem_used, mem_total = _get_current_mem_info()\n message = (\n f\"After {name}, memory allocated (GB): {mem_allocated}, memory reserved (GB): {mem_reserved}, \"\n f\"device memory used/total (GB): {mem_used}/{mem_total}\"\n )\n\n self.logging_function(message)\n return output\n\n\ndef log_print(ctn: Any):\n current_time = datetime.datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\")\n\n frame = inspect.currentframe().f_back\n function_name = frame.f_code.co_name\n line_number = frame.f_lineno\n file_name = frame.f_code.co_filename.split(\"/\")[-1]\n print(f\"[{current_time}-{file_name}:{line_number}:{function_name}]: {ctn}\")\n\n\ndef _timer(name: str, timing_raw: dict[str, float]):\n \"\"\"Inner function that handles the core timing logic.\n\n Args:\n name (str): The name/identifier for this timing measurement.\n timing_raw (Dict[str, float]): Dictionary to store timing information.\n \"\"\"\n with Timer(name=name, logger=None) as timer:\n yield\n if name not in timing_raw:\n timing_raw[name] = 0\n timing_raw[name] += timer.last\n\n\n@contextmanager\ndef simple_timer(name: str, timing_raw: dict[str, float]):\n \"\"\"Context manager for basic timing without NVTX markers.\n\n This utility function measures the execution time of code within its context\n and accumulates the timing information in the provided dictionary.\n\n Args:\n name (str): The name/identifier for this timing measurement.\n timing_raw (Dict[str, float]): Dictionary to store timing information.\n\n Yields:\n None: This is a context manager that yields control back to the code block.\n \"\"\"\n yield from _timer(name, timing_raw)\n\n\n@contextmanager\ndef marked_timer(\n name: str,\n timing_raw: dict[str, float],\n color: str = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n):\n \"\"\"Context manager for timing with platform markers.\n\n This utility function measures the execution time of code within its context,\n accumulates the timing information, and adds platform markers for profiling.\n This function is a default implementation when hardware profiler is not available.\n\n Args:\n name (str): The name/identifier for this timing measurement.\n timing_raw (Dict[str, float]): Dictionary to store timing information.\n color (Optional[str]): Color for the marker. Defaults to None.\n domain (Optional[str]): Domain for the marker. Defaults to None.\n category (Optional[str]): Category for the marker. Defaults to None.\n\n Yields:\n None: This is a context manager that yields control back to the code block.\n \"\"\"\n yield from _timer(name, timing_raw)\n\n\ndef reduce_timing(\n timing_raw: dict[str, float], reduce_op: torch.distributed.ReduceOp = torch.distributed.ReduceOp.AVG\n) -> dict[str, float]:\n \"\"\"Reduce timing information across all processes.\n\n This function uses distributed communication to gather and sum the timing\n information from all processes in a distributed environment.\n\n Args:\n timing_raw (Dict[str, float]): Dictionary containing timing information.\n\n Returns:\n Dict[str, float]: Reduced timing information.\n \"\"\"\n if not dist.is_initialized():\n return timing_raw\n\n key_list, timing_list = [], []\n for key in sorted(timing_raw.keys()):\n key_list.append(key)\n timing_list.append(timing_raw[key])\n timing_list = torch.tensor(timing_list, dtype=torch.float32, device=get_device_id())\n torch.distributed.all_reduce(timing_list, op=reduce_op)\n timing_list = [tensor.item() for tensor in timing_list.to(\"cpu\")]\n timing_generate = {key_list[i]: timing_list[i] for i in range(len(key_list))}\n return timing_generate\n\n\ndef topk_reduce_ratio_min_max(timing: float, k: int = 10) -> tuple[float, float, float]:\n \"\"\"Calculate topk items take-up ratio, and min/max timing across all ranks.\"\"\"\n if not dist.is_initialized():\n return -1.0, -1.0, -1.0\n\n world_size = dist.get_world_size()\n timing_tensor = torch.tensor(timing, dtype=torch.float32, device=get_device_id())\n tensor_list = [torch.zeros(1, dtype=torch.float32, device=get_device_id()) for _ in range(world_size)]\n torch.distributed.all_gather(tensor_list, timing_tensor)\n tensor_stack = torch.stack(tensor_list)\n timing_min = tensor_stack.min().cpu().item()\n timing_max = tensor_stack.max().cpu().item()\n top_k_percentile = torch.quantile(tensor_stack, 1 - k / 100)\n tail_ratio = torch.mean((tensor_stack > top_k_percentile).float()).cpu().item()\n return tail_ratio, timing_min, timing_max\n\n\ndef gather_timing(timing_raw: dict[str, float]) -> dict[str, list[float]]:\n if not dist.is_initialized():\n return {k: [v] for k, v in timing_raw.items()}\n\n key_list, timing_list = [], []\n for key in sorted(timing_raw.keys()):\n key_list.append(key)\n timing_list.append(timing_raw[key])\n\n world_size = torch.distributed.get_world_size()\n\n object_gather_list = [None] * world_size\n\n torch.distributed.all_gather_object(object_gather_list, timing_list)\n\n timing_generate = {\n key_list[i]: [timing_list[i] for timing_list in object_gather_list] for i in range(len(key_list))\n }\n\n return timing_generate\n"}101{"file_name": "verl__utils__profiler__profile.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport functools\nfrom typing import Callable, Optional\n\nfrom ..memory_utils import MemorySnapshotSampler, enable_memory_visualize\nfrom .config import ProfilerConfig, TorchMemoryToolConfig\n\n\ndef mark_start_range(\n message: Optional[str] = None,\n color: Optional[str] = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n) -> None:\n \"\"\"Start a profiling range marker (no-op implementation).\n\n Args:\n message (Optional[str]): Message to associate with the range marker.\n color (Optional[str]): Color for the marker visualization.\n domain (Optional[str]): Domain for the marker.\n category (Optional[str]): Category for the marker.\n \"\"\"\n pass\n\n\ndef mark_end_range(range_id: str) -> None:\n \"\"\"End a profiling range marker (no-op implementation).\n\n Args:\n range_id (str): Identifier of the range to end.\n \"\"\"\n pass\n\n\ndef mark_annotate(\n message: Optional[str] = None,\n color: Optional[str] = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n) -> Callable:\n \"\"\"Decorator to annotate a function with profiling markers (no-op implementation).\n\n Args:\n message (Optional[str]): Message to associate with the annotation.\n color (Optional[str]): Color for the marker visualization.\n domain (Optional[str]): Domain for the marker.\n category (Optional[str]): Category for the marker.\n\n Returns:\n Callable: Decorator function that returns the original function unchanged.\n \"\"\"\n\n def decorator(func):\n return func\n\n return decorator\n\n\nclass DistProfiler:\n \"\"\"A dispatcher that delegates to specific profilers based on config.tool.\n\n Supported tools:\n - nsys: NsightSystemsProfiler\n - npu: NPUProfiler (Ascend)\n - torch: PyTorch torch.profiler wrapper\n - torch_memory: Torch CUDA memory snapshot dump\n \"\"\"\n\n def __init__(\n self, rank: int, config: Optional[ProfilerConfig] = None, tool_config: Optional[object] = None, **kwargs\n ):\n # Default config\n if not config:\n config = ProfilerConfig(ranks=[], enable=False, tool_config=None)\n\n if tool_config is None:\n tool_config = config.tool_config\n\n self.config = config\n self.tool_config = tool_config\n\n self._impl = None\n self._tool = getattr(config, \"tool\", None)\n self._enable = config.enable\n self._this_step = False\n\n # Normalize rank selection\n self._this_rank = False\n if config.all_ranks:\n self._this_rank = True\n elif config.ranks:\n self._this_rank = rank in config.ranks\n else:\n # default rank 0 if enabled but ranks unspecified\n self._this_rank = (rank == 0) if self._enable else False\n\n # TorchMemoryProfiler currently do not support discrete mode.\n self._discrete = getattr(tool_config, \"discrete\", False) if tool_config else False\n\n # Lazy import to avoid circular deps\n if self._tool == \"nsys\":\n from .nvtx_profile import NsightSystemsProfiler as _Nsight\n\n self._impl = _Nsight(rank=rank, config=config, tool_config=tool_config, **kwargs)\n elif self._tool == \"npu\":\n from .mstx_profile import NPUProfiler as _Npu\n\n self._impl = _Npu(rank=rank, config=config, tool_config=tool_config, **kwargs)\n elif self._tool == \"torch\":\n from .torch_profile import Profiler as _Torch\n\n self._impl = _Torch(rank=rank, config=config, tool_config=tool_config)\n elif self._tool == \"torch_memory\":\n self._impl = TorchMemoryProfiler(rank=rank, config=config, tool_config=tool_config)\n else:\n # Fallback to a no-op impl\n self._impl = _NoOpProfiler()\n\n def check_enable(self):\n return self._enable\n\n def check_this_rank(self):\n return self._this_rank\n\n def check_this_step(self):\n return self._this_step\n\n def is_discrete_mode(self):\n return self._discrete\n\n def start(self, **kwargs):\n if self.check_enable() and self.check_this_rank():\n self._this_step = True\n return getattr(self._impl, \"start\", lambda **_: None)(**kwargs)\n\n def stop(self):\n if self.check_enable() and self.check_this_rank():\n self._this_step = False\n return getattr(self._impl, \"stop\", lambda: None)()\n\n @classmethod\n def annotate(\n cls,\n message: Optional[str] = None,\n color: Optional[str] = None,\n domain: Optional[str] = None,\n category: Optional[str] = None,\n **kwargs_outer,\n ) -> Callable:\n def decorator(func):\n @functools.wraps(func)\n def wrapper(self_instance, *args, **kwargs_inner):\n profiler = getattr(self_instance, \"profiler\", None)\n if (\n not profiler\n or not profiler.check_enable()\n or not profiler.check_this_step()\n or not profiler.check_this_rank()\n ):\n return func(self_instance, *args, **kwargs_inner)\n\n impl = profiler._impl\n if hasattr(impl, \"annotate\"):\n try:\n actual_decorator = impl.annotate(\n message=message, color=color, domain=domain, category=category, **kwargs_outer\n )\n\n return actual_decorator(func)(self_instance, *args, **kwargs_inner)\n except Exception:\n return func(self_instance, *args, **kwargs_inner)\n return func(self_instance, *args, **kwargs_inner)\n\n return wrapper\n\n return decorator\n\n\nclass _NoOpProfiler:\n def start(self, **kwargs):\n return\n\n def stop(self):\n return\n\n\nclass TorchMemoryProfiler:\n \"\"\"Profiler that dumps CUDA memory snapshots at step boundaries.\n\n Behavior:\n - On first construction (per process), enable memory history recording if CUDA is available\n - On start(step=X), remember sub_dir for this step\n - On stop(), dump a memory snapshot into config.save_path under the remembered sub_dir\n \"\"\"\n\n _memory_history_enabled: bool = False\n\n def __init__(\n self, rank: int, config: Optional[ProfilerConfig], tool_config: Optional[TorchMemoryToolConfig] = None\n ):\n # Always respond to explicit start/stop calls for torch_memory tool,\n # regardless of per-role enable flag, to align with global step control.\n self.enable = True\n if not config:\n config = ProfilerConfig(ranks=[])\n self.config = config\n self.rank = rank\n self.this_step = False\n self.sub_dir = None\n self.sampler = MemorySnapshotSampler()\n\n # Get parameters from tool_config, with fallback to defaults\n if tool_config:\n trace_alloc_max_entries = tool_config.trace_alloc_max_entries\n stack_depth = tool_config.stack_depth\n else:\n trace_alloc_max_entries = 100_000\n stack_depth = 32\n\n # Best-effort enable memory history once\n if not TorchMemoryProfiler._memory_history_enabled:\n try:\n enable_memory_visualize(trace_alloc_max_entries=trace_alloc_max_entries, stack_depth=stack_depth)\n except Exception:\n # silently ignore if not supported\n pass\n TorchMemoryProfiler._memory_history_enabled = True\n\n def start(self, **kwargs):\n if not self.enable:\n return\n if not self._should_profile_this_rank():\n return\n profile_step = kwargs.get(\"profile_step\", None)\n # Keep ranks aligned under same folder name\n self.sub_dir = f\"step{profile_step}\" if profile_step is not None else None\n self.this_step = True\n\n def stop(self):\n if not self.enable or not self.this_step:\n return\n self.this_step = False\n if not self._should_profile_this_rank():\n return\n out_dir = self.config.save_path or \"outputs/profile\"\n tag = \"torch_memory\"\n # Dump snapshot; all ranks write into same sub_dir\n try:\n self.sampler.dump_memory_snapshot(out_dir=out_dir, tag=tag, sub_dir=self.sub_dir)\n except Exception:\n pass\n\n def _should_profile_this_rank(self) -> bool:\n if self.config.all_ranks:\n return True\n if self.config.ranks:\n return self.rank in self.config.ranks\n # default rank 0\n return self.rank == 0\n\n\nclass DistProfilerExtension:\n \"\"\"An extension class for DistProfiler that provides distributed profiling capabilities.\n\n It is intended for workers in verl that single controller invokes.\n\n This class wraps a DistProfiler instance and provides methods to start/stop profiling\n that can be dispatched across multiple ranks in a distributed training environment.\n\n Args:\n profiler (DistProfiler): The base distributed profiler instance to extend\n \"\"\"\n\n def __init__(self, profiler: DistProfiler):\n self.profiler = profiler\n\n from verl.single_controller.base.decorator import Dispatch, register\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def start_profile(self, **kwargs) -> None:\n \"\"\"Start profiling for the current rank in the current training step.\"\"\"\n self.profiler.start(**kwargs)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def stop_profile(self) -> None:\n \"\"\"Stop profiling for the current rank in the current training step.\"\"\"\n self.profiler.stop()\n"}102{"file_name": "verl__utils__profiler__torch_profile.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport functools\nimport os\nfrom typing import Callable, Optional\n\nimport torch\n\nfrom .config import ProfilerConfig, TorchProfilerToolConfig\nfrom .profile import DistProfiler\n\n\ndef get_torch_profiler(\n contents: list[str],\n save_path: str,\n role: Optional[str] = None,\n save_file_prefix: Optional[str] = None,\n rank: int = 0,\n):\n if role:\n save_path = os.path.join(save_path, role)\n\n os.makedirs(save_path, exist_ok=True)\n\n save_file_name = f\"prof_rank-{rank}.json.gz\"\n if save_file_prefix:\n save_file_name = f\"{save_file_prefix}_{save_file_name}\"\n save_path = os.path.join(save_path, save_file_name)\n\n def _trace_handler(prof):\n print(f\"[Profiler] Saving trace to {save_path}\")\n prof.export_chrome_trace(save_path)\n\n contents = set(contents) if contents else set()\n activities = []\n if not contents or \"cpu\" in contents:\n activities.append(torch.profiler.ProfilerActivity.CPU)\n if not contents or \"cuda\" in contents:\n activities.append(torch.profiler.ProfilerActivity.CUDA)\n\n return torch.profiler.profile(\n activities=activities,\n with_stack=\"stack\" in contents,\n record_shapes=\"shapes\" in contents,\n profile_memory=\"memory\" in contents,\n on_trace_ready=_trace_handler,\n )\n\n\nclass Profiler(DistProfiler):\n \"\"\"A PyTorch profiler wrapper class for collecting performance metrics.\n\n This profiler provides a convenient interface for profiling PyTorch operations,\n with support for:\n\n - CPU and CUDA activity profiling\n - Configurable profiling schedule (wait/warmup/active steps)\n - Multi-rank profiling support\n - Chrome trace export\n\n Args:\n config: Configuration object containing profiling parameters\n \"\"\"\n\n _define_count = 0\n\n def __init__(\n self,\n rank,\n config: ProfilerConfig,\n tool_config: Optional[TorchProfilerToolConfig] = None,\n save_file_prefix=None,\n ):\n # note : if we do not set use_profile, it will be set as None, so that all function will be skip\n config = config or ProfilerConfig(ranks=[], enable=False)\n self.save_file_prefix = save_file_prefix\n\n if not tool_config:\n assert not config.enable, \"tool_config must be provided when profiler is enabled\"\n\n self.prof = None\n self.rank = rank\n self.config = config\n self.tool_config = tool_config\n self.contents = self.tool_config.contents\n self.save_path = self.config.save_path\n # Align with other profilers: read discrete mode, default to False for torch profiler\n self.discrete = getattr(self.tool_config, \"discrete\", False)\n\n def check(self):\n return self.prof is not None\n\n def start(self, **kwargs):\n role = kwargs.get(\"role\", None)\n if not self.discrete and Profiler._define_count == 0:\n self.prof = get_torch_profiler(\n contents=self.contents,\n save_path=self.save_path,\n role=role,\n save_file_prefix=self.save_file_prefix,\n rank=self.rank,\n )\n print(f\"[Profiler] started for rank {self.rank}\")\n self.prof.start()\n Profiler._define_count += 1\n\n def step(self):\n if self.check():\n self.prof.step()\n\n def stop(self):\n if not self.discrete and Profiler._define_count == 1:\n self.step()\n print(f\"[Profiler] stopped for rank {self.rank}\")\n self.prof.stop()\n Profiler._define_count -= 1\n\n def annotate(self, message: Optional[str] = None, role: Optional[str] = None, **kwargs_outer) -> Callable:\n \"\"\"Decorate a Worker member function to profile the current rank in the current training step.\n\n Requires the target function to be a member function of a Worker,\n which has a member field `profiler` with Profiler type.\n\n Args:\n message (str, optional):\n The message to be displayed in the profiler. Defaults to None.\n role (str, optional):\n The role of the current data collection. Defaults to None.\n \"\"\"\n\n def decorator(func):\n @functools.wraps(func)\n def wrapper(*args, **kwargs_inner):\n profile_name = message or func.__name__\n\n if not self.discrete:\n # In continuous mode, we just record function, profiler started globally\n with torch.profiler.record_function(profile_name):\n return func(*args, **kwargs_inner)\n\n # In discrete mode, we start/stop profiler around the function\n prof = get_torch_profiler(\n contents=self.contents,\n save_path=self.save_path,\n role=role,\n save_file_prefix=self.save_file_prefix,\n rank=self.rank,\n )\n prof.start()\n with torch.profiler.record_function(profile_name):\n result = func(*args, **kwargs_inner)\n prof.stop()\n return result\n\n return wrapper\n\n return decorator\n"}103{"file_name": "verl__utils__qat__core.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\"\"\"QAT (Quantization-Aware Training) utilities for verl FSDP training.\"\"\"\n\nimport json\nimport logging\nimport re\nfrom dataclasses import dataclass, field\nfrom typing import Any, Optional\n\nimport torch.nn as nn\n\nfrom verl.base_config import BaseConfig\n\nlogger = logging.getLogger(__name__)\n\n\n@dataclass\nclass QATConfig(BaseConfig):\n \"\"\"Unified configuration for QAT (Quantization-Aware Training).\"\"\"\n\n enable: bool = False\n mode: str = \"w4a16\"\n group_size: int = 16\n ignore_patterns: list[str] = field(default_factory=lambda: [\"lm_head\", \"embed_tokens\", \"re:.*mlp.gate$\"])\n activation_observer: str = \"static_minmax\"\n quantization_config_path: Optional[str] = None\n\n\ndef load_quantization_config(qat_config: QATConfig) -> dict[str, Any]:\n \"\"\"Load quantization config JSON file from QATConfig.\"\"\"\n if not qat_config.quantization_config_path:\n raise ValueError(\"quantization_config_path is required when QAT is enabled\")\n\n logger.info(f\"Loading QAT quantization config from: {qat_config.quantization_config_path}\")\n\n with open(qat_config.quantization_config_path) as f:\n quant_config = json.load(f)\n\n if qat_config.ignore_patterns:\n original_ignore = quant_config.get(\"ignore\", [])\n quant_config[\"ignore\"] = qat_config.ignore_patterns\n if original_ignore != qat_config.ignore_patterns:\n logger.info(f\"Overriding JSON 'ignore' field: {original_ignore} -> {qat_config.ignore_patterns}\")\n\n logger.info(\"Successfully loaded QAT quantization config\")\n return quant_config\n\n\ndef _should_quantize(name: str, module: nn.Module, config: QATConfig) -> bool:\n \"\"\"Check if a module should be quantized.\"\"\"\n if not isinstance(module, nn.Linear):\n return False\n\n for pattern in config.ignore_patterns:\n if pattern.startswith(\"re:\"):\n regex = pattern[3:]\n if re.match(regex, name):\n logger.debug(f\"Ignoring {name} due to regex pattern: {regex}\")\n return False\n else:\n if pattern in name:\n logger.debug(f\"Ignoring {name} due to pattern: {pattern}\")\n return False\n\n if module.in_features % config.group_size != 0:\n logger.warning(\n f\"Skipping {name}: in_features={module.in_features} not divisible by group_size={config.group_size}\"\n )\n return False\n\n return True\n\n\ndef apply_qat(\n model: nn.Module,\n config: QATConfig | dict[str, Any],\n) -> nn.Module:\n \"\"\"Apply QAT to a model by replacing nn.Linear with QATLinear.\"\"\"\n from verl.utils.qat.linear import QATLinear, QATMode\n\n if not isinstance(config, QATConfig):\n config = QATConfig(**config)\n\n if not config.enable:\n logger.info(\"QAT is disabled, returning original model\")\n return model\n\n mode = QATMode(config.mode.lower())\n logger.info(f\"Applying QAT with mode={mode.value}, group_size={config.group_size}\")\n\n modules_to_replace = []\n for name, module in model.named_modules():\n if _should_quantize(name, module, config):\n modules_to_replace.append((name, module))\n\n logger.info(f\"Found {len(modules_to_replace)} Linear layers to convert to QAT\")\n\n converted_count = 0\n for name, module in modules_to_replace:\n if isinstance(module, QATLinear):\n continue\n\n fake_quant_module = QATLinear.from_linear(\n module,\n mode=mode,\n group_size=config.group_size,\n activation_observer=config.activation_observer,\n )\n\n _set_module(model, name, fake_quant_module)\n converted_count += 1\n\n logger.info(f\"Successfully applied QAT to {converted_count} layers\")\n\n return model\n\n\ndef _set_module(model: nn.Module, name: str, new_module: nn.Module):\n \"\"\"Set a module in the model by its full name.\"\"\"\n parts = name.split(\".\")\n parent = model\n for part in parts[:-1]:\n parent = getattr(parent, part)\n setattr(parent, parts[-1], new_module)\n\n\nFUSION_PATTERNS = {\n \"qkv\": [\"q_proj\", \"k_proj\", \"v_proj\"],\n \"gate_up\": [\"gate_proj\", \"up_proj\"],\n}\n\n\ndef setup_fusion_siblings(model: nn.Module):\n \"\"\"Setup fusion siblings for QKV and GateUp layers.\"\"\"\n import weakref\n\n from verl.utils.qat.linear import QATLinear\n\n qat_modules = {name: m for name, m in model.named_modules() if isinstance(m, QATLinear)}\n\n counts = {}\n for group_name, suffixes in FUSION_PATTERNS.items():\n groups: dict[str, dict[str, nn.Module]] = {}\n for name, module in qat_modules.items():\n for suffix in suffixes:\n if name.endswith(suffix):\n parent = name.rsplit(\".\", 1)[0]\n groups.setdefault(parent, {})[suffix] = module\n\n count = 0\n for parent, projs in groups.items():\n if len(projs) >= 2:\n modules = list(projs.values())\n for i, m in enumerate(modules):\n siblings = modules[:i] + modules[i + 1 :]\n m._fusion_siblings_ref = [weakref.ref(s) for s in siblings]\n count += 1\n counts[group_name] = count\n\n logger.info(f\"[QAT Fuse] Setup fusion siblings: {counts}\")\n return counts\n\n\ndef enable_qat_fuse(model: nn.Module):\n \"\"\"Enable QAT fuse mode: sets up fusion siblings for weight scale fusion.\"\"\"\n setup_fusion_siblings(model)\n model._qat_fuse_enabled = True\n logger.info(\"[QAT Fuse] Enabled QAT fuse mode\")\n\n\ndef invalidate_all_scales(model: nn.Module):\n \"\"\"Clear all cached weight scales after optimizer.step().\"\"\"\n from verl.utils.qat.linear import QATLinear\n\n count = 0\n for module in model.modules():\n if isinstance(module, QATLinear):\n module._weight_blockwise_scale = None\n module._weight_global_scale = None\n module._cached_weight_amax = None\n count += 1\n\n logger.debug(f\"[QAT Fuse] Invalidated scales for {count} QATLinear layers\")\n"}104{"file_name": "verl__utils__qat__linear.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\"\"\"QAT FakeQuantized Linear module for NVFP4 (W4A4/W4A16) with FSDP compatibility.\n\nIncludes Triton kernels for high-performance FP4 quantization.\n\"\"\"\n\nfrom enum import Enum\nfrom typing import Optional\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n__all__ = [\"QATLinear\", \"QATMode\"]\n\n\nimport triton\nimport triton.language as tl\n\n_TORCH_TO_TL_DTYPE = {\n torch.float32: tl.float32,\n torch.float16: tl.float16,\n torch.bfloat16: tl.bfloat16,\n}\nFP4_E2M1_MAX: float = 6.0\nFP8_E4M3_MAX: float = 448.0\n\n\n@triton.jit\ndef _fp4_fake_quant_kernel(\n x_ptr,\n y_ptr,\n M,\n N,\n global_scale_ptr,\n stride_xm,\n stride_xn,\n stride_ym,\n stride_yn,\n BLOCK_SIZE: tl.constexpr,\n TILE_M: tl.constexpr,\n TILE_N: tl.constexpr,\n NUM_FP4_BLOCKS: tl.constexpr,\n OUT_DTYPE: tl.constexpr,\n FP4_MAX: tl.constexpr,\n FP8_MAX: tl.constexpr,\n):\n pid_m = tl.program_id(axis=0)\n pid_n = tl.program_id(axis=1)\n row_start = pid_m * TILE_M\n col_start = pid_n * TILE_N\n\n x_block_ptr = tl.make_block_ptr(\n base=x_ptr,\n shape=(M, N),\n strides=(stride_xm, stride_xn),\n offsets=(row_start, col_start),\n block_shape=(TILE_M, TILE_N),\n order=(1, 0),\n )\n y_block_ptr = tl.make_block_ptr(\n base=y_ptr,\n shape=(M, N),\n strides=(stride_ym, stride_yn),\n offsets=(row_start, col_start),\n block_shape=(TILE_M, TILE_N),\n order=(1, 0),\n )\n\n global_scale = tl.load(global_scale_ptr).to(tl.float32)\n global_scale_safe = tl.where(global_scale > 0.0, global_scale, 1e-12)\n\n tile = tl.load(x_block_ptr, boundary_check=(0, 1), padding_option=\"zero\").to(tl.float32)\n tile_reshaped = tl.reshape(tile, (TILE_M, NUM_FP4_BLOCKS, BLOCK_SIZE))\n x_abs = tl.abs(tile_reshaped)\n\n block_max = tl.max(x_abs, axis=2, keep_dims=True)\n block_max_scaled = block_max / (FP4_MAX * global_scale_safe)\n block_max_scaled = tl.minimum(block_max_scaled, FP8_MAX)\n block_max_quant = block_max_scaled.to(tl.float8e4nv).to(tl.float32) * global_scale\n block_max_quant = tl.where(block_max_quant >= 1e-5, block_max_quant, 1.0)\n\n block_max_quant_broadcast = tl.broadcast_to(block_max_quant, (TILE_M, NUM_FP4_BLOCKS, BLOCK_SIZE))\n abs_scaled = x_abs / block_max_quant_broadcast\n\n q_val = tl.where(\n abs_scaled <= 0.25,\n 0.0,\n tl.where(\n abs_scaled < 0.75,\n 0.5,\n tl.where(\n abs_scaled <= 1.25,\n 1.0,\n tl.where(\n abs_scaled < 1.75,\n 1.5,\n tl.where(\n abs_scaled <= 2.5,\n 2.0,\n tl.where(abs_scaled < 3.5, 3.0, tl.where(abs_scaled <= 5.0, 4.0, FP4_MAX)),\n ),\n ),\n ),\n ),\n )\n\n x_rescaled = q_val * block_max_quant_broadcast\n x_rescaled = tl.where(tile_reshaped >= 0, x_rescaled, -x_rescaled)\n tile_quant = tl.reshape(x_rescaled, (TILE_M, TILE_N))\n\n tl.store(y_block_ptr, tile_quant.to(OUT_DTYPE), boundary_check=(0, 1))\n\n\ndef fp4_fake_quant_weight(\n weight: torch.Tensor,\n global_amax: torch.Tensor = None,\n block_size: int = 16,\n tile_rows: int = 16,\n tile_cols: int = 64,\n) -> torch.Tensor:\n \"\"\"Apply FP4 fake quantization using Triton kernel.\"\"\"\n x_shape = weight.shape\n x_dtype = weight.dtype\n x = weight.reshape(-1, x_shape[-1]).contiguous()\n M, N = x.shape\n y = torch.empty_like(x)\n\n stride_xm, stride_xn = x.stride()\n stride_ym, stride_yn = y.stride()\n\n tile_cols = max(tile_cols, block_size)\n tile_cols_aligned = ((tile_cols + block_size - 1) // block_size) * block_size\n num_fp4_blocks = tile_cols_aligned // block_size\n\n if global_amax is None:\n global_amax = weight.abs().max().to(torch.float32)\n global_scale = global_amax.float() / (FP4_E2M1_MAX * FP8_E4M3_MAX)\n\n grid = (triton.cdiv(M, tile_rows), triton.cdiv(N, tile_cols_aligned))\n\n _fp4_fake_quant_kernel[grid](\n x,\n y,\n M,\n N,\n global_scale,\n stride_xm,\n stride_xn,\n stride_ym,\n stride_yn,\n BLOCK_SIZE=block_size,\n TILE_M=tile_rows,\n TILE_N=tile_cols_aligned,\n NUM_FP4_BLOCKS=num_fp4_blocks,\n OUT_DTYPE=_TORCH_TO_TL_DTYPE[x_dtype],\n FP4_MAX=FP4_E2M1_MAX,\n FP8_MAX=FP8_E4M3_MAX,\n )\n return y.view(*x_shape)\n\n\nclass STEFP4QuantTriton(torch.autograd.Function):\n \"\"\"Straight-Through Estimator wrapper for Triton FP4 quantization kernel.\"\"\"\n\n @staticmethod\n def forward(ctx, x: torch.Tensor, global_amax: torch.Tensor, block_size: int) -> torch.Tensor:\n return fp4_fake_quant_weight(x, global_amax=global_amax, block_size=block_size)\n\n @staticmethod\n def backward(ctx, grad_output: torch.Tensor) -> tuple:\n return grad_output, None, None\n\n\nclass QATMode(str, Enum):\n \"\"\"QAT quantization mode.\"\"\"\n\n W4A4 = \"w4a4\" # Weight 4-bit, Activation 4-bit (dynamic)\n W4A16 = \"w4a16\" # Weight 4-bit, Activation 16-bit (weight only)\n\n\nclass QATLinear(nn.Linear):\n \"\"\"QAT FakeQuantized Linear layer with FSDP compatibility.\"\"\"\n\n _UNINITIALIZED_SCALE = -1.0\n\n def __init__(\n self,\n in_features: int,\n out_features: int,\n bias: bool = True,\n mode: QATMode = QATMode.W4A4,\n group_size: int = 16,\n activation_observer: str = \"static_minmax\", # Observer strategy for activation global_scale\n device: Optional[torch.device] = None,\n dtype: Optional[torch.dtype] = None,\n ):\n super().__init__(in_features, out_features, bias, device=device, dtype=dtype)\n\n self.mode = mode\n self.group_size = group_size\n self.activation_observer = activation_observer\n\n self._weight_blockwise_scale: Optional[torch.Tensor] = None\n self._weight_global_scale: Optional[torch.Tensor] = None\n self._cached_weight_amax: Optional[torch.Tensor] = None\n self._fusion_siblings_ref = None\n\n if mode == QATMode.W4A4:\n self.register_buffer(\n \"input_global_scale\", torch.tensor([self._UNINITIALIZED_SCALE], dtype=torch.float32), persistent=True\n )\n\n self.register_buffer(\n \"input_amax\", torch.tensor([self._UNINITIALIZED_SCALE], dtype=torch.float32), persistent=True\n )\n\n self._ema_decay: float = 0.01\n\n self.fake_quant_enabled = True\n\n @classmethod\n def from_linear(\n cls,\n linear: nn.Linear,\n mode: QATMode = QATMode.W4A4,\n group_size: int = 16,\n activation_observer: str = \"static_minmax\",\n ) -> \"QATLinear\":\n \"\"\"Create QATLinear from an existing nn.Linear.\"\"\"\n has_bias = linear.bias is not None\n\n new_linear = cls(\n in_features=linear.in_features,\n out_features=linear.out_features,\n bias=has_bias,\n mode=mode,\n group_size=group_size,\n activation_observer=activation_observer,\n device=linear.weight.device,\n dtype=linear.weight.dtype,\n )\n\n if linear.weight.device != torch.device(\"meta\"):\n new_linear.weight = nn.Parameter(linear.weight.clone())\n if has_bias:\n new_linear.bias = nn.Parameter(linear.bias.clone())\n\n return new_linear\n\n def _is_amax_initialized(self) -> bool:\n \"\"\"Check if input_amax has been initialized.\"\"\"\n if not hasattr(self, \"input_amax\"):\n return False\n return self.input_amax.item() != self._UNINITIALIZED_SCALE\n\n def _update_input_global_scale(self, x: torch.Tensor):\n \"\"\"Update static input_global_scale based on observer strategy.\"\"\"\n assert self.mode == QATMode.W4A4, \"_update_input_global_scale should only be called in W4A4 mode\"\n\n current_amax = torch.amax(torch.abs(x)).detach().to(torch.float32)\n\n if torch.distributed.is_initialized() and torch.distributed.get_world_size() > 1:\n torch.distributed.all_reduce(current_amax, op=torch.distributed.ReduceOp.MAX)\n\n scale_factor = FP8_E4M3_MAX * FP4_E2M1_MAX\n\n if self.activation_observer == \"memoryless_minmax\":\n new_scale = (scale_factor / (current_amax + 1e-12)).view(1)\n self.input_global_scale.copy_(new_scale.to(self.input_global_scale.device))\n\n elif self.activation_observer == \"static_minmax\":\n if not self._is_amax_initialized():\n self.input_amax.copy_(current_amax.view(1).to(self.input_amax.device))\n else:\n new_amax = torch.maximum(self.input_amax, current_amax.view(1).to(self.input_amax.device))\n self.input_amax.copy_(new_amax)\n amax_f32 = self.input_amax.to(torch.float32)\n new_scale = (scale_factor / (amax_f32 + 1e-12)).float().view(1)\n self.input_global_scale.copy_(new_scale.to(self.input_global_scale.device))\n\n elif self.activation_observer == \"minmax\":\n if not self._is_amax_initialized():\n self.input_amax.copy_(current_amax.view(1).to(self.input_amax.device))\n else:\n new_amax = (1 - self._ema_decay) * self.input_amax + self._ema_decay * current_amax.view(1).to(\n self.input_amax.device\n )\n self.input_amax.copy_(new_amax)\n amax_f32 = self.input_amax.to(torch.float32)\n new_scale = (scale_factor / (amax_f32 + 1e-12)).float().view(1)\n self.input_global_scale.copy_(new_scale.to(self.input_global_scale.device))\n\n else:\n raise ValueError(f\"Unknown activation_observer: {self.activation_observer}\")\n\n def _fake_quantize_weight(self, weight: torch.Tensor) -> torch.Tensor:\n \"\"\"Apply fake quantization to weight tensor using Triton kernel.\"\"\"\n with torch.no_grad():\n if self._cached_weight_amax is not None:\n global_amax = self._cached_weight_amax\n else:\n siblings_ref = getattr(self, \"_fusion_siblings_ref\", None)\n\n if siblings_ref is not None:\n siblings = [ref() for ref in siblings_ref if ref() is not None]\n siblings = [s for s in siblings if s.weight.device != torch.device(\"meta\")]\n\n for sibling in siblings:\n sibling_amax = getattr(sibling, \"_cached_weight_amax\", None)\n if sibling_amax is not None:\n global_amax = sibling_amax\n self._cached_weight_amax = global_amax\n break\n else:\n all_modules = [self] + siblings\n amaxes = [m.weight.abs().max().to(torch.float32) for m in all_modules]\n global_amax = torch.max(torch.stack(amaxes))\n\n self._cached_weight_amax = global_amax\n for sibling in siblings:\n sibling._cached_weight_amax = global_amax\n else:\n global_amax = weight.abs().max().to(torch.float32)\n self._cached_weight_amax = global_amax\n\n if self._weight_global_scale is None:\n self._weight_global_scale = global_amax.float() / (FP4_E2M1_MAX * FP8_E4M3_MAX)\n\n result = STEFP4QuantTriton.apply(weight, global_amax, self.group_size)\n\n return result\n\n def _fake_quantize_activation(self, x: torch.Tensor) -> torch.Tensor:\n \"\"\"Apply fake quantization to activation tensor (W4A4 mode only).\"\"\"\n original_shape = x.shape\n\n if x.dim() == 3:\n x_2d = x.view(-1, x.shape[-1])\n else:\n x_2d = x\n\n if self.training:\n self._update_input_global_scale(x_2d)\n\n if self.input_global_scale.item() == self._UNINITIALIZED_SCALE:\n raise RuntimeError(\"W4A4 input_global_scale uninitialized. Load PTQ model first.\")\n\n global_amax = (FP4_E2M1_MAX * FP8_E4M3_MAX) / self.input_global_scale.to(x.device)\n result = STEFP4QuantTriton.apply(x_2d, global_amax, self.group_size)\n return result.view(original_shape)\n\n def forward(self, x: torch.Tensor) -> torch.Tensor:\n \"\"\"Forward pass with fake quantization.\"\"\"\n if not self.fake_quant_enabled:\n return F.linear(x, self.weight, self.bias)\n\n weight_fq = self._fake_quantize_weight(self.weight)\n\n if self.mode == QATMode.W4A4:\n x_fq = self._fake_quantize_activation(x)\n else:\n x_fq = x\n\n return F.linear(x_fq, weight_fq, self.bias)\n\n def extra_repr(self) -> str:\n return (\n f\"in_features={self.in_features}, out_features={self.out_features}, \"\n f\"bias={self.bias is not None}, mode={self.mode.value}, \"\n f\"group_size={self.group_size}, fake_quant_enabled={self.fake_quant_enabled}\"\n )\n"}105{"file_name": "verl__utils__qat__quantizer.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\"\"\"\nFast NVFP4 Quantizer for verl FSDP training.\n\nDirectly computes scales and quantizes weights using compressed_tensors APIs.\nIncludes scale computation utilities for weight quantization.\n\"\"\"\n\nimport logging\nimport os\nimport re\nfrom typing import Generator, Iterable, Optional\n\nimport torch\nfrom compressed_tensors.compressors.quantized_compressors.fp4_quantized import NVFP4PackedCompressor\nfrom compressed_tensors.quantization.quant_args import (\n FP4_E2M1_DATA,\n FP8_E4M3_DATA,\n QuantizationArgs,\n QuantizationStrategy,\n QuantizationType,\n)\nfrom compressed_tensors.quantization.utils.helpers import generate_gparam\n\nfrom verl.utils.device import get_device_name, get_torch_device\n\nlogger = logging.getLogger(__name__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n_LAYER_IDX_RE = re.compile(r\"layers\\.(\\d+)\\.\")\n\n\ndef compute_blockwise_scale(\n weight: torch.Tensor,\n global_scale: torch.Tensor,\n group_size: int = 16,\n) -> torch.Tensor:\n \"\"\"Compute blockwise scale using pre-computed global_scale (for fusion).\n Returns FP8 E4M3 blockwise scale tensor.\n \"\"\"\n out_features, in_features = weight.shape\n num_groups = in_features // group_size\n weight_reshaped = weight.view(out_features, num_groups, group_size)\n block_max = torch.amax(torch.abs(weight_reshaped), dim=-1).to(torch.float32)\n\n local_scale = block_max / FP4_E2M1_DATA.max\n blockwise_scale_f32 = torch.clamp(\n global_scale * local_scale,\n min=-FP8_E4M3_DATA.max,\n max=FP8_E4M3_DATA.max,\n )\n\n blockwise_scale = blockwise_scale_f32.to(torch.float8_e4m3fn)\n eps = torch.finfo(torch.float8_e4m3fn).eps\n blockwise_scale = torch.where(\n blockwise_scale == 0,\n torch.tensor(eps, dtype=blockwise_scale.dtype, device=weight.device),\n blockwise_scale,\n )\n\n return blockwise_scale\n\n\n# Fusion patterns for transformer models\nFUSE_PATTERNS = {\n \"qkv\": [\"q_proj\", \"k_proj\", \"v_proj\"],\n \"gate_up\": [\"gate_proj\", \"up_proj\"],\n}\n\n\ndef fuse_global_scales(\n layer_global_scales: dict[str, torch.Tensor],\n strategy: str = \"min\",\n) -> dict[str, torch.Tensor]:\n \"\"\"Fuse global scales for QKV/GateUp groups (take min across group).\"\"\"\n if not layer_global_scales:\n return {}\n\n # Group by parent module\n parent_to_children: dict[str, dict[str, str]] = {}\n for name in layer_global_scales:\n parent, child = name.rsplit(\".\", 1) if \".\" in name else (\"\", name)\n parent_to_children.setdefault(parent, {})[child] = name\n\n fused_scales = {}\n processed = set()\n\n for parent, children in parent_to_children.items():\n for _, patterns in FUSE_PATTERNS.items():\n matched = [children[p] for p in patterns if p in children]\n if len(matched) == len(patterns):\n group_scales = [layer_global_scales[n] for n in matched]\n if strategy == \"min\":\n fused_scale = torch.min(torch.cat(group_scales)).reshape([1])\n else:\n raise ValueError(f\"Unknown fuse strategy: {strategy}\")\n for layer_name in matched:\n fused_scales[layer_name] = fused_scale.clone()\n processed.add(layer_name)\n\n for name, scale in layer_global_scales.items():\n if name not in processed:\n fused_scales[name] = scale\n\n return fused_scales\n\n\nclass QATQuantizer:\n \"\"\"Quantizer for QAT-trained weights using compressed_tensors APIs.\"\"\"\n\n def __init__(\n self,\n mode: str = \"w4a16\",\n group_size: int = 16,\n ignore_patterns: Optional[list] = None,\n device: Optional[torch.device] = None,\n param_dtype: Optional[torch.dtype] = None,\n ):\n self.mode = mode.lower()\n self._is_w4a4 = self.mode == \"w4a4\" # W4A4 needs input_global_scale\n self.group_size = group_size\n self.ignore_patterns = ignore_patterns or [\"lm_head\", \"embed_tokens\", \"re:.*mlp.gate$\"]\n self.device = device or torch.device(get_device_name())\n self.param_dtype = param_dtype\n\n self._compressor = NVFP4PackedCompressor()\n self._quant_args = QuantizationArgs(\n num_bits=4,\n type=QuantizationType.FLOAT,\n symmetric=True,\n strategy=QuantizationStrategy.TENSOR_GROUP,\n group_size=group_size,\n scale_dtype=FP8_E4M3_DATA.dtype,\n )\n\n def _should_quantize(self, name: str, tensor: torch.Tensor) -> bool:\n \"\"\"Check if parameter should be quantized.\"\"\"\n if not name.endswith(\".weight\"):\n return False\n if tensor.dim() != 2:\n return False\n if tensor.shape[1] % self.group_size != 0:\n return False\n\n module_name = name.rsplit(\".weight\", 1)[0]\n\n for pattern in self.ignore_patterns:\n if pattern.startswith(\"re:\"):\n # Regex pattern - use re.match like vLLM does\n regex = pattern[3:]\n if re.match(regex, module_name):\n return False\n else:\n if pattern in module_name:\n return False\n return True\n\n @staticmethod\n def _extract_layer_idx(name: str) -> Optional[int]:\n \"\"\"Extract decoder layer index from parameter name.\"\"\"\n match = _LAYER_IDX_RE.search(name)\n return int(match.group(1)) if match else None\n\n def _process_layer_group(\n self,\n layer_idx: Optional[int],\n layer_params: dict[str, torch.Tensor],\n input_global_scales: dict[str, torch.Tensor],\n output_device: torch.device,\n ) -> list[tuple[str, torch.Tensor]]:\n \"\"\"Quantize one decoder layer's buffered params. Returns list of (name, tensor).\"\"\"\n layer_weights = {}\n layer_passthrough = {}\n\n for name, tensor in layer_params.items():\n if \"input_global_scale\" in name or \"input_amax\" in name:\n continue\n\n if self._should_quantize(name, tensor):\n layer_name = name.rsplit(\".weight\", 1)[0]\n layer_weights[layer_name] = (name, tensor)\n else:\n layer_passthrough[name] = tensor\n\n if layer_idx is None and layer_weights:\n raise RuntimeError(\n f\"[QAT Quantizer] Unexpected quantizable weights outside decoder layers: \"\n f\"{list(layer_weights.keys())}. These should be in ignore_patterns.\"\n )\n\n if not layer_weights:\n return [(name, tensor.to(output_device)) for name, tensor in layer_passthrough.items()]\n\n # Move weights to GPU, compute global scales\n weights_on_gpu = {}\n layer_global_scales = {}\n\n for layer_name, (_, tensor) in layer_weights.items():\n weight_gpu = tensor.to(device=self.device, dtype=self.param_dtype)\n weights_on_gpu[layer_name] = weight_gpu\n amax = torch.amax(torch.abs(weight_gpu)).to(torch.float32)\n layer_global_scales[layer_name] = generate_gparam(\n -amax.unsqueeze(0),\n amax.unsqueeze(0),\n scale_data=FP8_E4M3_DATA,\n quant_data=FP4_E2M1_DATA,\n dtype=torch.float32,\n )\n\n fused_global_scales = fuse_global_scales(layer_global_scales, strategy=\"min\")\n\n results = []\n\n for layer_name, weight_gpu in weights_on_gpu.items():\n fused_global_scale = fused_global_scales[layer_name]\n weight_scale = compute_blockwise_scale(weight_gpu, fused_global_scale, self.group_size)\n weight_packed = self._compressor.compress_weight(\n weight=weight_gpu,\n scale=weight_scale.float(),\n global_scale=fused_global_scale,\n quantization_args=self._quant_args,\n )[\"weight_packed\"]\n\n results.append((f\"{layer_name}.weight_packed\", weight_packed.to(output_device)))\n results.append((f\"{layer_name}.weight_scale\", weight_scale.to(output_device)))\n results.append((f\"{layer_name}.weight_global_scale\", fused_global_scale.to(output_device)))\n\n if self._is_w4a4:\n if layer_name in input_global_scales:\n results.append(\n (\n f\"{layer_name}.input_global_scale\",\n input_global_scales[layer_name].float().to(output_device),\n )\n )\n else:\n raise ValueError(\n f\"W4A4 mode requires input_global_scale for layer '{layer_name}', \"\n f\"but it's not found or uninitialized (-1.0).\"\n )\n\n del weights_on_gpu, layer_global_scales, fused_global_scales\n\n for name, tensor in layer_passthrough.items():\n results.append((name, tensor.to(output_device)))\n\n return results\n\n def quantize_with_fusion(\n self,\n params: dict[str, torch.Tensor] | Iterable[tuple[str, torch.Tensor]],\n target_device: Optional[torch.device] = None,\n ) -> Generator[tuple[str, torch.Tensor], None, None]:\n \"\"\"Streaming quantize: consume input layer by layer, yield (name, tensor) pairs.\"\"\"\n if isinstance(params, dict):\n params = params.items()\n\n output_device = target_device or torch.device(\"cpu\")\n\n _sentinel = object()\n current_layer_idx = _sentinel\n layer_buffer: dict[str, torch.Tensor] = {}\n input_global_scales: dict[str, torch.Tensor] = {}\n for name, tensor in params:\n tensor_cpu = tensor.to(\"cpu\") if tensor.is_cuda else tensor\n layer_idx = self._extract_layer_idx(name)\n\n # Collect input_global_scales for W4A4 as we go\n if self._is_w4a4 and \"input_global_scale\" in name:\n scale_layer_name = name.replace(\".input_global_scale\", \"\")\n if tensor_cpu.numel() == 1 and tensor_cpu.item() == -1.0:\n logger.warning(f\"W4A4: {scale_layer_name} input_global_scale is uninitialized\")\n else:\n input_global_scales[scale_layer_name] = tensor_cpu\n\n # Layer boundary: flush previous layer\n if layer_idx != current_layer_idx and current_layer_idx is not _sentinel and layer_buffer:\n yield from self._process_layer_group(\n current_layer_idx, layer_buffer, input_global_scales, output_device\n )\n layer_buffer = {}\n\n current_layer_idx = layer_idx\n layer_buffer[name] = tensor_cpu\n\n # Flush last buffered layer\n if layer_buffer:\n yield from self._process_layer_group(current_layer_idx, layer_buffer, input_global_scales, output_device)\n\n get_torch_device().empty_cache()\n\n\n__all__ = [\n \"QATQuantizer\",\n]\n"}106{"file_name": "verl__utils__qat__vllm_patch.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\"\"\"\nvLLM NVFP4 Patches for Dynamic Weight Updates.\n\nEnables dynamic weight reloading for NVFP4 quantized models in vLLM.\n\nSupported schemes:\n- Dense: W4A16-FP4, W4A4-FP4\n- MoE: NVFP4-MoE\n\"\"\"\n\nimport logging\nimport os\nfrom typing import Optional\nfrom unittest.mock import patch\n\nimport torch\nfrom torch.nn import Parameter\n\nfrom verl.utils.device import get_device_name\n\nlogger = logging.getLogger(__name__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\nclass ParamMetaDict(dict):\n \"\"\"\n Dict-like class for parameter management with metadata-based rebuild and tensor swap.\n\n Supports:\n - Rebuild of deleted parameters from saved metadata\n - Tensor Swap for parameters with shape changes (address stability for CUDA Graph)\n \"\"\"\n\n def __init__(self, model: torch.nn.Module, device: Optional[torch.device] = None):\n \"\"\"\n Initialize ParamMetaDict from a model.\n\n Args:\n model: vLLM model (may be wrapped in ModelRunner)\n device: Device for created parameters\n \"\"\"\n super().__init__()\n self.device = device\n\n # Get the actual model (handle vLLM's wrapper structure)\n actual_model = model\n if hasattr(model, \"model\"):\n actual_model = model.model\n self._model = actual_model\n\n # Build mappings by scanning all modules\n self._layer_meta_cache: dict[str, dict] = {} # Cache of _hf_param_meta\n self._tensor_swap_layers: dict[str, dict] = {} # Layers needing tensor swap\n\n self._build_mappings()\n\n # Initialize with current parameters\n for name, param in actual_model.named_parameters():\n self[name] = param\n\n def _build_mappings(self):\n \"\"\"Build layer metadata cache for rebuild and tensor swap.\"\"\"\n for layer_name, module in self._model.named_modules():\n # Check for _hf_param_meta which indicates this layer has HF format params\n if hasattr(module, \"_hf_param_meta\"):\n self._layer_meta_cache[layer_name] = {\n \"module\": module,\n \"meta\": module._hf_param_meta,\n }\n\n # Check for tensor swap layers (weight_scale with shape change)\n if \"weight_scale\" in module._hf_param_meta:\n marlin_refs = getattr(module, \"_marlin_tensor_refs\", {})\n if \"weight_scale\" in marlin_refs:\n self._tensor_swap_layers[layer_name] = {\n \"module\": module,\n \"marlin_ref\": marlin_refs[\"weight_scale\"],\n \"hf_meta\": module._hf_param_meta[\"weight_scale\"],\n }\n\n # MoE layers (w13_weight_scale, w2_weight_scale)\n if \"w13_weight_scale\" in module._hf_param_meta:\n marlin_refs = getattr(module, \"_marlin_tensor_refs\", {})\n if \"w13_weight_scale\" in marlin_refs:\n self._tensor_swap_layers[f\"{layer_name}.w13\"] = {\n \"module\": module,\n \"param_name\": \"w13_weight_scale\",\n \"marlin_ref\": marlin_refs[\"w13_weight_scale\"],\n \"hf_meta\": module._hf_param_meta[\"w13_weight_scale\"],\n }\n if \"w2_weight_scale\" in marlin_refs:\n self._tensor_swap_layers[f\"{layer_name}.w2\"] = {\n \"module\": module,\n \"param_name\": \"w2_weight_scale\",\n \"marlin_ref\": marlin_refs[\"w2_weight_scale\"],\n \"hf_meta\": module._hf_param_meta[\"w2_weight_scale\"],\n }\n\n def _try_rebuild(self, key: str) -> Optional[Parameter]:\n \"\"\"\n Try to rebuild a parameter from metadata if it was deleted.\n\n Args:\n key: Full parameter name\n\n Returns:\n Rebuilt parameter or None if cannot rebuild\n \"\"\"\n # Extract layer name and param name\n parts = key.rsplit(\".\", 1)\n if len(parts) != 2:\n return None\n\n layer_name, param_name = parts\n\n # Check if we have metadata for this layer\n if layer_name not in self._layer_meta_cache:\n return None\n\n cache_entry = self._layer_meta_cache[layer_name]\n module = cache_entry[\"module\"]\n meta = cache_entry[\"meta\"]\n\n # Check if this param needs rebuild\n if param_name not in meta:\n return None\n\n # Already exists on module?\n if hasattr(module, param_name):\n param = getattr(module, param_name)\n if param is not None:\n return param\n\n # Rebuild from metadata\n new_param = _create_param_from_meta(module, param_name, meta[param_name], self.device)\n module.register_parameter(param_name, new_param)\n return new_param\n\n def prepare_for_reload(self) -> None:\n \"\"\"Replace Marlin-format tensors with HF-shape tensors for reload.\"\"\"\n for layer_name, swap_info in self._tensor_swap_layers.items():\n module = swap_info[\"module\"]\n param_name = swap_info.get(\"param_name\", \"weight_scale\")\n hf_meta = swap_info[\"hf_meta\"]\n if hasattr(module, param_name):\n new_param = _create_param_from_meta(module, param_name, hf_meta, self.device)\n setattr(module, param_name, new_param)\n\n def __getitem__(self, key: str) -> Parameter:\n \"\"\"Get parameter with rebuild support.\"\"\"\n # Try standard lookup first\n if key in dict.keys(self):\n return super().__getitem__(key)\n\n # Try rebuild from metadata\n param = self._try_rebuild(key)\n if param is not None:\n self[key] = param\n return param\n\n raise KeyError(f\"Parameter not found: {key}\")\n\n def __contains__(self, key: str) -> bool:\n \"\"\"Check if parameter exists (with rebuild check).\"\"\"\n if super().__contains__(key):\n return True\n\n # Check if can rebuild from metadata\n parts = key.rsplit(\".\", 1)\n if len(parts) == 2:\n layer_name, param_name = parts\n if layer_name in self._layer_meta_cache:\n meta = self._layer_meta_cache[layer_name][\"meta\"]\n if param_name in meta:\n return True\n\n return False\n\n def get(self, key: str, default=None):\n \"\"\"Get parameter with default.\"\"\"\n try:\n return self[key]\n except KeyError:\n return default\n\n\ndef _create_param_from_meta(\n module: torch.nn.Module,\n param_name: str,\n meta: dict,\n device: Optional[torch.device] = None,\n) -> Parameter:\n \"\"\"Create a Parameter from saved metadata. Used by rebuild and tensor swap.\"\"\"\n shape = meta[\"shape\"]\n dtype = meta[\"dtype\"]\n dev = device or meta.get(\"device\", get_device_name())\n param_class = meta.get(\"param_class\", Parameter)\n\n weight_loaders = getattr(module, \"_weight_loaders\", {})\n weight_loader = weight_loaders.get(param_name)\n\n data = torch.empty(shape, dtype=dtype, device=dev)\n\n try:\n if param_class is not Parameter and weight_loader is not None:\n kwargs = {\"data\": data, \"weight_loader\": weight_loader}\n if \"input_dim\" in meta:\n kwargs[\"input_dim\"] = meta[\"input_dim\"]\n if \"output_dim\" in meta:\n kwargs[\"output_dim\"] = meta[\"output_dim\"]\n new_param = param_class(**kwargs)\n else:\n new_param = Parameter(data, requires_grad=False)\n if weight_loader is not None:\n new_param.weight_loader = weight_loader\n except Exception as e:\n logger.warning(f\"Failed to create param {param_name} with class {param_class}: {e}, using Parameter\")\n new_param = Parameter(data, requires_grad=False)\n if weight_loader is not None:\n new_param.weight_loader = weight_loader\n\n if \"quant_method\" in meta:\n new_param.quant_method = meta[\"quant_method\"]\n\n return new_param\n\n\ndef save_param_meta(layer: torch.nn.Module, param_name: str):\n \"\"\"Save parameter metadata for rebuild.\"\"\"\n if not hasattr(layer, \"_hf_param_meta\"):\n layer._hf_param_meta = {}\n\n param = getattr(layer, param_name, None)\n if param is None:\n return\n\n meta = {\n \"shape\": tuple(param.shape),\n \"dtype\": param.dtype,\n \"device\": str(param.device),\n \"param_class\": type(param), # Save the actual parameter class\n }\n\n # Save vLLM-specific attributes needed for reconstruction\n if hasattr(param, \"_input_dim\"):\n meta[\"input_dim\"] = param._input_dim\n if hasattr(param, \"_output_dim\"):\n meta[\"output_dim\"] = param._output_dim\n\n # Save MoE-specific attributes (quant_method is required by weight_loader)\n if hasattr(param, \"quant_method\"):\n meta[\"quant_method\"] = param.quant_method\n\n layer._hf_param_meta[param_name] = meta\n\n\ndef _check_first_call(layer: torch.nn.Module) -> bool:\n \"\"\"Check if this is the first process_weights call, and increment counter.\"\"\"\n count = getattr(layer, \"_process_weights_call_count\", 0)\n layer._process_weights_call_count = count + 1\n return count == 0\n\n\n# Dense W4A16 Patches\ndef patched_w4a16_process_weights_after_loading(self, layer: torch.nn.Module) -> None:\n \"\"\"Patched process_weights_after_loading for W4A16 Dense layer.\"\"\"\n import vllm._custom_ops as ops\n from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import (\n marlin_make_workspace_new,\n marlin_permute_scales,\n nvfp4_marlin_process_global_scale,\n nvfp4_marlin_process_scales,\n )\n\n is_first_call = _check_first_call(layer)\n\n group_size = 16\n part_size_n = layer.output_size_per_partition\n part_size_k = layer.input_size_per_partition\n device = layer.weight_packed.device\n param_dtype = getattr(layer, \"params_dtype\", torch.float16)\n\n # Save metadata (first call only)\n if is_first_call:\n save_param_meta(layer, \"weight_packed\")\n save_param_meta(layer, \"weight_global_scale\")\n save_param_meta(layer, \"weight_scale\")\n if not hasattr(layer, \"_weight_loaders\"):\n layer._weight_loaders = {}\n for pname in [\"weight_packed\", \"weight_global_scale\", \"weight_scale\"]:\n param = getattr(layer, pname, None)\n if param is not None and hasattr(param, \"weight_loader\"):\n layer._weight_loaders[pname] = param.weight_loader\n\n # Get HF format data\n weight_packed_hf = layer.weight_packed.data\n weight_global_scale_hf = layer.weight_global_scale.data\n weight_scale_hf = layer.weight_scale.data\n\n # Create workspace (first call only)\n if is_first_call:\n layer.workspace = marlin_make_workspace_new(device)\n\n # Convert to Marlin format\n perm = torch.empty(0, dtype=torch.int, device=device)\n qweight = weight_packed_hf.view(torch.int32).T.contiguous()\n marlin_weight = ops.gptq_marlin_repack(\n b_q_weight=qweight,\n perm=perm,\n size_k=part_size_k,\n size_n=part_size_n,\n num_bits=4,\n is_a_8bit=False,\n )\n\n weight_scale = weight_scale_hf.T.contiguous().to(param_dtype)\n weight_scale_permuted = marlin_permute_scales(\n s=weight_scale,\n size_k=part_size_k,\n size_n=part_size_n,\n group_size=group_size,\n is_a_8bit=False,\n )\n marlin_weight_scale = nvfp4_marlin_process_scales(weight_scale_permuted)\n\n weight_scale_2_raw = (1.0 / weight_global_scale_hf.max()).to(param_dtype)\n marlin_weight_scale_2 = nvfp4_marlin_process_global_scale(weight_scale_2_raw)\n\n # Update compute parameters\n if is_first_call:\n layer.weight = Parameter(marlin_weight, requires_grad=False)\n layer.weight_scale = Parameter(marlin_weight_scale, requires_grad=False)\n layer.weight_scale_2 = Parameter(marlin_weight_scale_2, requires_grad=False)\n if not hasattr(layer, \"_marlin_tensor_refs\"):\n layer._marlin_tensor_refs = {}\n layer._marlin_tensor_refs[\"weight_scale\"] = layer.weight_scale.data\n else:\n layer.weight.data.copy_(marlin_weight)\n layer.weight_scale_2.data.copy_(marlin_weight_scale_2)\n marlin_scale_ref = layer._marlin_tensor_refs.get(\"weight_scale\")\n if marlin_scale_ref is not None:\n marlin_scale_ref.copy_(marlin_weight_scale)\n layer.weight_scale = Parameter(marlin_scale_ref, requires_grad=False)\n else:\n logger.warning(\"W4A16: _marlin_tensor_refs['weight_scale'] not found\")\n layer.weight_scale = Parameter(marlin_weight_scale, requires_grad=False)\n\n # Delete HF parameters\n if hasattr(layer, \"weight_packed\"):\n delattr(layer, \"weight_packed\")\n if hasattr(layer, \"weight_global_scale\"):\n delattr(layer, \"weight_global_scale\")\n\n\ndef patched_w4a4_process_weights_after_loading(self, layer: torch.nn.Module) -> None:\n \"\"\"Patched process_weights_after_loading for W4A4 Dense (all backends).\"\"\"\n from vllm.model_executor.layers.quantization.utils.quant_utils import swizzle_blockscale\n\n is_first_call = _check_first_call(layer)\n\n _W4A4_HF_PARAMS = [\"weight_packed\", \"weight_scale\", \"weight_global_scale\", \"input_global_scale\"]\n\n if is_first_call:\n for pname in _W4A4_HF_PARAMS:\n save_param_meta(layer, pname)\n if not hasattr(layer, \"_weight_loaders\"):\n layer._weight_loaders = {}\n for pname in _W4A4_HF_PARAMS:\n param = getattr(layer, pname, None)\n if param is not None and hasattr(param, \"weight_loader\"):\n layer._weight_loaders[pname] = param.weight_loader\n\n weight_packed_data = layer.weight_packed.data\n weight_scale_data = layer.weight_scale.data\n input_global_scale_data = layer.input_global_scale.data\n weight_global_scale_data = layer.weight_global_scale.data\n\n global_input_scale = input_global_scale_data.max().to(torch.float32)\n global_weight_scale = weight_global_scale_data.max().to(torch.float32)\n\n if self.backend == \"flashinfer-trtllm\":\n from flashinfer import shuffle_matrix_a, shuffle_matrix_sf_a\n\n epilogue_tile_m = 128\n processed_weight = shuffle_matrix_a(weight_packed_data.view(torch.uint8), epilogue_tile_m)\n processed_weight_scale = (\n shuffle_matrix_sf_a(weight_scale_data.view(torch.uint8), epilogue_tile_m)\n .reshape(weight_scale_data.shape)\n .view(torch.float8_e4m3fn)\n )\n elif self.backend == \"fbgemm\":\n processed_weight_scale = swizzle_blockscale(weight_scale_data).view(-1).view(torch.uint8)\n processed_weight = weight_packed_data\n else:\n # cutlass / flashinfer-cutlass\n processed_weight_scale = swizzle_blockscale(weight_scale_data)\n processed_weight = weight_packed_data\n\n alpha = 1.0 / (global_input_scale * global_weight_scale)\n\n if is_first_call:\n layer.weight_packed = Parameter(processed_weight, requires_grad=False)\n layer.weight_scale = Parameter(processed_weight_scale, requires_grad=False)\n layer.input_global_scale = Parameter(global_input_scale, requires_grad=False)\n layer.weight_global_scale = Parameter(global_weight_scale, requires_grad=False)\n layer.alpha = Parameter(alpha, requires_grad=False)\n\n if not hasattr(layer, \"_marlin_tensor_refs\"):\n layer._marlin_tensor_refs = {}\n layer._marlin_tensor_refs[\"weight_packed\"] = layer.weight_packed.data\n layer._marlin_tensor_refs[\"weight_scale\"] = layer.weight_scale.data\n layer._marlin_tensor_refs[\"input_global_scale\"] = layer.input_global_scale.data\n layer._marlin_tensor_refs[\"weight_global_scale\"] = layer.weight_global_scale.data\n layer._marlin_tensor_refs[\"alpha\"] = layer.alpha.data\n else:\n refs = layer._marlin_tensor_refs\n for ref_name, new_data in [\n (\"weight_packed\", processed_weight),\n (\"weight_scale\", processed_weight_scale),\n (\"input_global_scale\", global_input_scale),\n (\"weight_global_scale\", global_weight_scale),\n (\"alpha\", alpha),\n ]:\n ref = refs.get(ref_name)\n if ref is not None:\n ref.copy_(new_data)\n setattr(layer, ref_name, Parameter(ref, requires_grad=False))\n else:\n logger.warning(f\"W4A4: _marlin_tensor_refs['{ref_name}'] not found, creating new Parameter\")\n setattr(\n layer,\n ref_name,\n Parameter(\n new_data.clone() if isinstance(new_data, torch.Tensor) else torch.tensor(new_data),\n requires_grad=False,\n ),\n )\n\n\ndef _marlin_repack_experts(packed, perm, size_k, size_n, num_experts):\n \"\"\"Repack weight for each expert into Marlin format and stack.\"\"\"\n import vllm._custom_ops as ops\n\n result = []\n for i in range(num_experts):\n qweight = packed[i].view(torch.int32).T.contiguous()\n result.append(\n ops.gptq_marlin_repack(\n b_q_weight=qweight,\n perm=perm,\n size_k=size_k,\n size_n=size_n,\n num_bits=4,\n is_a_8bit=False,\n )\n )\n return torch.stack(result)\n\n\ndef _marlin_process_scales_experts(scale_hf, param_dtype, size_k, size_n, group_size, num_experts):\n \"\"\"Process scales for each expert into Marlin format and stack.\"\"\"\n from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import (\n marlin_permute_scales,\n nvfp4_marlin_process_scales,\n )\n\n result = []\n scales = scale_hf.to(param_dtype)\n for i in range(num_experts):\n s = marlin_permute_scales(\n s=scales[i].T,\n size_k=size_k,\n size_n=size_n,\n group_size=group_size,\n is_a_8bit=False,\n )\n result.append(nvfp4_marlin_process_scales(s))\n return torch.stack(result)\n\n\ndef _process_nvfp4_moe_marlin(self, layer: torch.nn.Module, is_first_call: bool) -> None:\n \"\"\"Process MoE layer with MARLIN backend (W4A16).\"\"\"\n from vllm.model_executor.layers.fused_moe.oracle.nvfp4 import make_nvfp4_moe_kernel\n from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import (\n marlin_make_workspace_new,\n nvfp4_marlin_process_global_scale,\n )\n\n group_size = 16\n e = layer.num_experts\n k = layer.hidden_size\n n = layer.intermediate_size_per_partition\n device = layer.w13_weight_packed.device\n param_dtype = layer.params_dtype\n w13_num_shards = 2 if self.moe.is_act_and_mul else 1\n\n if is_first_call:\n layer.workspace = marlin_make_workspace_new(device, 4)\n\n perm = torch.empty(0, dtype=torch.int, device=device)\n\n if self.moe.is_act_and_mul and not torch.allclose(\n layer.w13_weight_global_scale[:, 0], layer.w13_weight_global_scale[:, 1]\n ):\n logger.warning(\"w1_weight_global_scale must match w3_weight_global_scale. Accuracy may be affected.\")\n\n size_n_w13, size_k_w13 = n * w13_num_shards, k\n size_n_w2, size_k_w2 = k, n\n\n w13_weight_marlin = _marlin_repack_experts(layer.w13_weight_packed.data, perm, size_k_w13, size_n_w13, e)\n w2_weight_marlin = _marlin_repack_experts(layer.w2_weight_packed.data, perm, size_k_w2, size_n_w2, e)\n w13_weight_scale_marlin = _marlin_process_scales_experts(\n layer.w13_weight_scale.data, param_dtype, size_k_w13, size_n_w13, group_size, e\n )\n w2_weight_scale_marlin = _marlin_process_scales_experts(\n layer.w2_weight_scale.data, param_dtype, size_k_w2, size_n_w2, group_size, e\n )\n\n # Process global scales\n w13_scale_2 = 1.0 / layer.w13_weight_global_scale[:, 0]\n w2_scale_2 = 1.0 / layer.w2_weight_global_scale.data\n w13_scale_2_processed = nvfp4_marlin_process_global_scale(w13_scale_2.to(param_dtype))\n w2_scale_2_processed = nvfp4_marlin_process_global_scale(w2_scale_2.to(param_dtype))\n\n # Update parameters\n if is_first_call:\n layer.w13_weight = Parameter(w13_weight_marlin, requires_grad=False)\n layer.w2_weight = Parameter(w2_weight_marlin, requires_grad=False)\n layer.w13_weight_scale = Parameter(w13_weight_scale_marlin, requires_grad=False)\n layer.w2_weight_scale = Parameter(w2_weight_scale_marlin, requires_grad=False)\n layer.w13_weight_scale_2 = Parameter(w13_scale_2_processed, requires_grad=False)\n layer.w2_weight_scale_2 = Parameter(w2_scale_2_processed, requires_grad=False)\n if not hasattr(layer, \"_marlin_tensor_refs\"):\n layer._marlin_tensor_refs = {}\n layer._marlin_tensor_refs[\"w13_weight_scale\"] = layer.w13_weight_scale.data\n layer._marlin_tensor_refs[\"w2_weight_scale\"] = layer.w2_weight_scale.data\n else:\n layer.w13_weight.data.copy_(w13_weight_marlin)\n layer.w2_weight.data.copy_(w2_weight_marlin)\n layer.w13_weight_scale_2.data.copy_(w13_scale_2_processed)\n layer.w2_weight_scale_2.data.copy_(w2_scale_2_processed)\n w13_marlin_ref = layer._marlin_tensor_refs.get(\"w13_weight_scale\")\n w2_marlin_ref = layer._marlin_tensor_refs.get(\"w2_weight_scale\")\n if w13_marlin_ref is not None:\n w13_marlin_ref.copy_(w13_weight_scale_marlin)\n layer.w13_weight_scale = Parameter(w13_marlin_ref, requires_grad=False)\n else:\n logger.warning(\"MoE: _marlin_tensor_refs['w13_weight_scale'] not found\")\n layer.w13_weight_scale.data.copy_(w13_weight_scale_marlin)\n if w2_marlin_ref is not None:\n w2_marlin_ref.copy_(w2_weight_scale_marlin)\n layer.w2_weight_scale = Parameter(w2_marlin_ref, requires_grad=False)\n else:\n logger.warning(\"MoE: _marlin_tensor_refs['w2_weight_scale'] not found\")\n layer.w2_weight_scale.data.copy_(w2_weight_scale_marlin)\n\n layer.w13_input_scale = None\n layer.w2_input_scale = None\n\n # Initialize kernel\n self.moe_quant_config = self.get_fused_moe_quant_config(layer)\n if self.moe_quant_config is not None and (\n (not self.moe.moe_parallel_config.use_all2all_kernels) or self.moe.moe_parallel_config.use_naive_all2all_kernels\n ):\n self.kernel = make_nvfp4_moe_kernel(\n moe_quant_config=self.moe_quant_config,\n moe_config=self.moe,\n experts_cls=self.experts_cls,\n )\n\n\ndef _process_nvfp4_moe_flashinfer_cutlass(self, layer: torch.nn.Module, is_first_call: bool) -> None:\n \"\"\"Process MoE layer with FlashInfer/CUTLASS backend (W4A4).\"\"\"\n from vllm.model_executor.layers.fused_moe.oracle.nvfp4 import (\n convert_to_nvfp4_moe_kernel_format,\n make_nvfp4_moe_kernel,\n )\n from vllm.model_executor.utils import replace_parameter\n\n w13_packed = layer.w13_weight_packed.data\n w2_packed = layer.w2_weight_packed.data\n w13_scale_hf = layer.w13_weight_scale.data\n w2_scale_hf = layer.w2_weight_scale.data\n\n if self.moe.is_act_and_mul and not torch.allclose(\n layer.w13_weight_global_scale[:, 0], layer.w13_weight_global_scale[:, 1]\n ):\n logger.warning(\"w1_weight_global_scale must match w3_weight_global_scale. Accuracy may be affected.\")\n w13_weight_global_scale = layer.w13_weight_global_scale[:, 0].contiguous()\n\n w13_temp = Parameter(w13_packed.clone(), requires_grad=False)\n w2_temp = Parameter(w2_packed.clone(), requires_grad=False)\n\n if is_first_call:\n layer.w13_weight = w13_temp\n layer.w2_weight = w2_temp\n\n (\n w13,\n w13_scale,\n w13_scale_2,\n a13_scale,\n w2,\n w2_scale,\n w2_scale_2,\n a2_scale,\n ) = convert_to_nvfp4_moe_kernel_format(\n nvfp4_backend=self.nvfp4_backend,\n layer=layer,\n w13=w13_temp,\n w13_scale=w13_scale_hf,\n w13_scale_2=(1.0 / w13_weight_global_scale),\n a13_scale=(1.0 / layer.w13_input_global_scale),\n w2=w2_temp,\n w2_scale=w2_scale_hf,\n w2_scale_2=(1.0 / layer.w2_weight_global_scale),\n a2_scale=(1.0 / layer.w2_input_global_scale),\n is_act_and_mul=self.moe.is_act_and_mul,\n )\n\n # Update parameters\n if is_first_call:\n replace_parameter(layer, \"w13_weight\", w13)\n replace_parameter(layer, \"w2_weight\", w2)\n layer.w13_weight_scale = Parameter(w13_scale, requires_grad=False)\n layer.w2_weight_scale = Parameter(w2_scale, requires_grad=False)\n if not hasattr(layer, \"_marlin_tensor_refs\"):\n layer._marlin_tensor_refs = {}\n layer._marlin_tensor_refs[\"w13_weight_scale\"] = layer.w13_weight_scale.data\n layer._marlin_tensor_refs[\"w2_weight_scale\"] = layer.w2_weight_scale.data\n else:\n layer.w13_weight.data.copy_(w13.data)\n layer.w2_weight.data.copy_(w2.data)\n w13_scale_ref = layer._marlin_tensor_refs.get(\"w13_weight_scale\")\n w2_scale_ref = layer._marlin_tensor_refs.get(\"w2_weight_scale\")\n if w13_scale_ref is not None:\n w13_scale_ref.copy_(w13_scale)\n layer.w13_weight_scale = Parameter(w13_scale_ref, requires_grad=False)\n else:\n logger.warning(\"MoE W4A4: _marlin_tensor_refs['w13_weight_scale'] not found\")\n layer.w13_weight_scale.data.copy_(w13_scale)\n if w2_scale_ref is not None:\n w2_scale_ref.copy_(w2_scale)\n layer.w2_weight_scale = Parameter(w2_scale_ref, requires_grad=False)\n else:\n logger.warning(\"MoE W4A4: _marlin_tensor_refs['w2_weight_scale'] not found\")\n layer.w2_weight_scale.data.copy_(w2_scale)\n\n layer.w13_weight_scale_2 = w13_scale_2\n layer.w2_weight_scale_2 = w2_scale_2\n layer.w13_input_scale = a13_scale\n layer.w2_input_scale = a2_scale\n\n # Initialize kernel\n self.moe_quant_config = self.get_fused_moe_quant_config(layer)\n if self.moe_quant_config is not None and (\n (not self.moe.moe_parallel_config.use_all2all_kernels) or self.moe.moe_parallel_config.use_naive_all2all_kernels\n ):\n self.kernel = make_nvfp4_moe_kernel(\n moe_quant_config=self.moe_quant_config,\n moe_config=self.moe,\n experts_cls=self.experts_cls,\n )\n\n\n# MoE NVFP4 Patches (entry points)\ndef patched_nvfp4_moe_process_weights_after_loading(self, layer: torch.nn.Module) -> None:\n \"\"\"Patched process_weights_after_loading for NVFP4 MoE layer.\"\"\"\n from vllm.model_executor.layers.fused_moe.oracle.nvfp4 import NvFp4MoeBackend\n\n is_first_call = _check_first_call(layer)\n\n # Save metadata (first call only)\n if is_first_call:\n save_param_meta(layer, \"w13_weight_packed\")\n save_param_meta(layer, \"w2_weight_packed\")\n save_param_meta(layer, \"w13_weight_scale\")\n save_param_meta(layer, \"w2_weight_scale\")\n if not hasattr(layer, \"_weight_loaders\"):\n layer._weight_loaders = {}\n for pname in [\"w13_weight_packed\", \"w2_weight_packed\", \"w13_weight_scale\", \"w2_weight_scale\"]:\n param = getattr(layer, pname, None)\n if param is not None and hasattr(param, \"weight_loader\"):\n layer._weight_loaders[pname] = param.weight_loader\n\n is_marlin = self.nvfp4_backend == NvFp4MoeBackend.MARLIN\n if is_marlin:\n _process_nvfp4_moe_marlin(self, layer, is_first_call)\n else:\n _process_nvfp4_moe_flashinfer_cutlass(self, layer, is_first_call)\n\n # Delete HF parameters\n if hasattr(layer, \"w13_weight_packed\"):\n delattr(layer, \"w13_weight_packed\")\n if hasattr(layer, \"w2_weight_packed\"):\n delattr(layer, \"w2_weight_packed\")\n\n\n_PATCH_TARGETS = [\n # Dense W4A16\n (\n \"vllm.model_executor.layers.quantization.compressed_tensors.schemes.\"\n \"compressed_tensors_w4a16_nvfp4.CompressedTensorsW4A16Fp4.process_weights_after_loading\",\n patched_w4a16_process_weights_after_loading,\n ),\n # Dense W4A4\n (\n \"vllm.model_executor.layers.quantization.compressed_tensors.schemes.\"\n \"compressed_tensors_w4a4_nvfp4.CompressedTensorsW4A4Fp4.process_weights_after_loading\",\n patched_w4a4_process_weights_after_loading,\n ),\n # MoE NVFP4\n (\n \"vllm.model_executor.layers.quantization.compressed_tensors.\"\n \"compressed_tensors_moe.CompressedTensorsW4A4Nvfp4MoEMethod.process_weights_after_loading\",\n patched_nvfp4_moe_process_weights_after_loading,\n ),\n]\n\n_applied_patches = []\n\n\ndef apply_qat_patches():\n \"\"\"Apply NVFP4 patches to support dynamic weight updates. Call before model loading.\"\"\"\n global _applied_patches\n\n if _applied_patches:\n logger.warning(\"QAT patches already applied, skipping\")\n return _applied_patches\n\n logger.info(\"Applying NVFP4 patches for dynamic weight loading...\")\n\n for target, replacement in _PATCH_TARGETS:\n p = patch(target, replacement)\n _applied_patches.append(p)\n p.start()\n\n logger.info(f\"Applied {len(_applied_patches)} NVFP4 patches for dynamic weight loading\")\n return _applied_patches\n\n\ndef prepare_qat_for_load_weights(model, device=None):\n \"\"\"\n Prepare QAT model for weight loading. Call ONCE before multi-bucket weight loading.\n\n Args:\n model: vLLM model\n device: Device for created parameters\n \"\"\"\n inner_model = model\n if hasattr(model, \"model\"):\n inner_model = model.model\n\n param_meta = ParamMetaDict(inner_model, device=device)\n\n param_meta.prepare_for_reload()\n logger.info(f\"[prepare_qat] Tensor swap prepared for {len(param_meta._tensor_swap_layers)} layers\")\n\n # Rebuild deleted (W4A16) or overwritten (W4A4) params back to HF format\n rebuilt_count = 0\n for layer_name, cache_entry in param_meta._layer_meta_cache.items():\n module = cache_entry[\"module\"]\n for param_name, pm in cache_entry[\"meta\"].items():\n existing = getattr(module, param_name, None)\n if existing is not None:\n hf_shape = tuple(pm[\"shape\"])\n hf_dtype = pm[\"dtype\"]\n if (\n tuple(existing.shape) == hf_shape\n and existing.dtype == hf_dtype\n and hasattr(existing, \"weight_loader\")\n ):\n continue\n new_param = _create_param_from_meta(module, param_name, pm, device)\n module.register_parameter(param_name, new_param)\n rebuilt_count += 1\n\n logger.info(f\"[prepare_qat] Rebuilt {rebuilt_count} parameters\")\n inner_model._param_meta_for_restore = param_meta\n return param_meta\n\n\ndef manual_process_weights_after_loading(model):\n \"\"\"Trigger weight post-processing for all quantized layers after load_weights.\"\"\"\n dense_count = 0\n moe_count = 0\n\n actual_model = model\n if hasattr(model, \"model\"):\n actual_model = model.model\n\n for module in actual_model.modules():\n if hasattr(module, \"scheme\"):\n module.scheme.process_weights_after_loading(module)\n dense_count += 1\n\n quant_method = getattr(module, \"quant_method\", None)\n if quant_method is not None and not hasattr(module, \"scheme\"):\n if hasattr(quant_method, \"process_weights_after_loading\"):\n # Skip KV cache quantization methods\n if \"KVCache\" in quant_method.__class__.__name__:\n continue\n quant_method.process_weights_after_loading(module)\n moe_count += 1\n\n logger.debug(f\"Processed {dense_count} dense layers, {moe_count} MoE layers\")\n return dense_count + moe_count\n\n\n__all__ = [\n \"apply_qat_patches\",\n \"prepare_qat_for_load_weights\",\n \"manual_process_weights_after_loading\",\n]\n"}107{"file_name": "verl__utils__ray_utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nContains commonly used utilities for ray\n\"\"\"\n\nimport asyncio\nimport concurrent.futures\nimport functools\nimport inspect\nimport os\nfrom typing import Any, Optional\n\nimport ray\n\n\ndef ray_noset_visible_devices(env_vars=os.environ):\n # Refer to\n # https://github.com/ray-project/ray/blob/161849364a784442cc659fb9780f1a6adee85fce/python/ray/_private/accelerators/nvidia_gpu.py#L95-L96\n # https://github.com/ray-project/ray/blob/161849364a784442cc659fb9780f1a6adee85fce/python/ray/_private/accelerators/amd_gpu.py#L102-L103\n # https://github.com/ray-project/ray/blob/3b9e729f6a669ffd85190f901f5e262af79771b0/python/ray/_private/accelerators/amd_gpu.py#L114-L115\n # https://github.com/ray-project/ray/blob/161849364a784442cc659fb9780f1a6adee85fce/python/ray/_private/accelerators/npu.py#L94-L95\n # https://github.com/ray-project/ray/blob/161849364a784442cc659fb9780f1a6adee85fce/python/ray/_private/accelerators/hpu.py#L116-L117\n # https://github.com/ray-project/ray/blob/161849364a784442cc659fb9780f1a6adee85fce/python/ray/_private/accelerators/neuron.py#L108-L109\n # https://github.com/ray-project/ray/blob/161849364a784442cc659fb9780f1a6adee85fce/python/ray/_private/accelerators/tpu.py#L171-L172\n # https://github.com/ray-project/ray/blob/161849364a784442cc659fb9780f1a6adee85fce/python/ray/_private/accelerators/intel_gpu.py#L97-L98\n NOSET_VISIBLE_DEVICES_ENV_VARS_LIST = [\n \"RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES\",\n \"RAY_EXPERIMENTAL_NOSET_ROCR_VISIBLE_DEVICES\",\n \"RAY_EXPERIMENTAL_NOSET_HIP_VISIBLE_DEVICES\",\n \"RAY_EXPERIMENTAL_NOSET_ASCEND_RT_VISIBLE_DEVICES\",\n \"RAY_EXPERIMENTAL_NOSET_HABANA_VISIBLE_MODULES\",\n \"RAY_EXPERIMENTAL_NOSET_NEURON_RT_VISIBLE_CORES\",\n \"RAY_EXPERIMENTAL_NOSET_TPU_VISIBLE_CHIPS\",\n \"RAY_EXPERIMENTAL_NOSET_ONEAPI_DEVICE_SELECTOR\",\n ]\n return any(env_vars.get(env_var) for env_var in NOSET_VISIBLE_DEVICES_ENV_VARS_LIST)\n\n\ndef parallel_put(data_list: list[Any], max_workers: Optional[int] = None):\n \"\"\"\n Puts a list of data into the Ray object store in parallel using a thread pool.\n\n Args:\n data_list (List[Any]): A list of Python objects to be put into the Ray object store.\n max_workers (int, optional): The maximum number of worker threads to use.\n Defaults to min(len(data_list), 16).\n\n Returns:\n List[ray.ObjectRef]: A list of Ray object references corresponding to the input data_list,\n maintaining the original order.\n \"\"\"\n assert len(data_list) > 0, \"data_list must not be empty\"\n\n def put_data(index, data):\n return index, ray.put(data)\n\n if max_workers is None:\n max_workers = min(len(data_list), 16)\n\n with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:\n data_list_f = [executor.submit(put_data, i, data) for i, data in enumerate(data_list)]\n res_lst = []\n for future in concurrent.futures.as_completed(data_list_f):\n res_lst.append(future.result())\n\n # reorder based on index\n output = [None for _ in range(len(data_list))]\n for res in res_lst:\n index, data_ref = res\n output[index] = data_ref\n\n return output\n\n\ndef get_event_loop():\n try:\n loop = asyncio.get_event_loop()\n except RuntimeError:\n loop = asyncio.new_event_loop()\n asyncio.set_event_loop(loop)\n\n return loop\n\n\ndef auto_await(func):\n \"\"\"Auto await a coroutine function.\n\n If the function is called in an async context (with a running event loop),\n it will return the coroutine object. Otherwise, it will block the current thread\n and run the coroutine until completion.\n \"\"\"\n\n @functools.wraps(func)\n def wrapper(*args, **kwargs):\n coro = func(*args, **kwargs)\n\n if not inspect.iscoroutine(coro):\n return coro\n\n try:\n loop = asyncio.get_running_loop()\n except RuntimeError:\n loop = None\n\n if loop and loop.is_running():\n return coro\n else:\n return asyncio.run(coro)\n\n return wrapper\n"}108{"file_name": "verl__utils__rendezvous__ray_backend.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nimport time\n\nimport ray\nfrom cupy.cuda.nccl import NcclCommunicator, get_unique_id\nfrom ray.util import list_named_actors\n\n\n@ray.remote\nclass NCCLIDStore:\n def __init__(self, nccl_id):\n self._nccl_id = nccl_id\n\n def get(self):\n return self._nccl_id\n\n\ndef get_nccl_id_store_by_name(name):\n all_actors = list_named_actors(all_namespaces=True)\n matched_actors = [actor for actor in all_actors if actor.get(\"name\", None) == name]\n if len(matched_actors) == 1:\n actor = matched_actors[0]\n return ray.get_actor(**actor)\n elif len(matched_actors) > 1:\n logging.warning(\"multiple actors with same name found: %s\", matched_actors)\n elif len(matched_actors) == 0:\n logging.info(\"failed to get any actor named %s\", name)\n return None\n\n\ndef create_nccl_communicator_in_ray(\n rank: int, world_size: int, group_name: str, max_retries: int = 100, interval_s: int = 5\n):\n if rank == 0:\n nccl_id = get_unique_id()\n nccl_id_store = NCCLIDStore.options(name=group_name).remote(nccl_id)\n\n assert ray.get(nccl_id_store.get.remote()) == nccl_id\n communicator = NcclCommunicator(\n ndev=world_size,\n commId=nccl_id,\n rank=0,\n )\n return communicator\n else:\n for i in range(max_retries):\n nccl_id_store = get_nccl_id_store_by_name(group_name)\n if nccl_id_store is not None:\n logging.info(\"nccl_id_store %s got\", group_name)\n nccl_id = ray.get(nccl_id_store.get.remote())\n logging.info(\"nccl id for %s got: %s\", group_name, nccl_id)\n communicator = NcclCommunicator(\n ndev=world_size,\n commId=nccl_id,\n rank=rank,\n )\n return communicator\n logging.info(\"failed to get nccl_id for %d time, sleep for %d seconds\", i + 1, interval_s)\n time.sleep(interval_s)\n"}109{"file_name": "verl__utils__reward_score____init__.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n# from . import gsm8k, math, prime_math, prime_code\n\nfrom verl.utils.import_utils import deprecated\n\n\ndef default_compute_score(\n data_source,\n solution_str,\n ground_truth,\n extra_info=None,\n sandbox_fusion_url=None,\n concurrent_semaphore=None,\n memory_limit_mb=None,\n **kwargs,\n):\n \"\"\"Compute the score for a given solution based on the data source.\n\n Args:\n data_source (str): The source dataset identifier which determines the scoring method.\n solution_str (str): The solution string to be evaluated.\n ground_truth (str): The ground truth answer for comparison.\n extra_info (dict, optional): Additional information that might be needed for scoring. Defaults to None.\n\n Returns:\n float: The computed score as a floating point number. If the result is a dictionary,\n it returns the dictionary instead.\n\n Raises:\n NotImplementedError: If the reward function is not implemented for the given data source.\n \"\"\"\n if data_source == \"openai/gsm8k\":\n from . import gsm8k\n\n res = gsm8k.compute_score(solution_str, ground_truth)\n elif data_source in [\"lighteval/MATH\", \"DigitalLearningGmbH/MATH-lighteval\", \"HuggingFaceH4/MATH-500\"]:\n from . import math_reward\n\n res = math_reward.compute_score(solution_str, ground_truth)\n # [Optional] Math-Verify Integration\n # For enhanced accuracy, consider utilizing Math-Verify (https://github.com/huggingface/Math-Verify).\n # Note: Math-Verify needs to be manually installed via pip: `pip install math-verify`.\n # To use it, override the `compute_score` function with the following implementation:\n\n # from . import math_verify\n # res = math_verify.compute_score(solution_str, ground_truth)\n elif data_source in [\"math_dapo\", \"math\", \"math_dapo_reasoning\"] or data_source.startswith(\"aime\"):\n from . import math_dapo\n\n res = math_dapo.compute_score(solution_str, ground_truth)\n elif data_source in [\n \"numina_aops_forum\",\n \"numina_synthetic_math\",\n \"numina_amc_aime\",\n \"numina_synthetic_amc\",\n \"numina_cn_k12\",\n \"numina_olympiads\",\n ]:\n from . import prime_math\n\n res = prime_math.compute_score(solution_str, ground_truth)\n elif data_source in [\"codecontests\", \"apps\", \"codeforces\", \"taco\"]:\n # Use the passed sandbox_fusion_url if available\n if sandbox_fusion_url:\n from . import sandbox_fusion\n\n # Pass the URL directly, ground_truth likely contains test cases here\n res = sandbox_fusion.compute_score(\n sandbox_fusion_url, concurrent_semaphore, memory_limit_mb, solution_str, ground_truth, continuous=True\n )\n else:\n # If no sandbox URL is provided, fall back to prime_code or raise error\n from . import prime_code\n\n # Assuming prime_code doesn't need the URL\n res = prime_code.compute_score(solution_str, ground_truth, continuous=True)\n elif data_source in [\"hiyouga/geometry3k\"]:\n from . import geo3k\n\n res = geo3k.compute_score(solution_str, ground_truth)\n elif data_source in [\n \"searchR1_nq\",\n \"searchR1_triviaqa\",\n \"searchR1_popqa\",\n \"searchR1_hotpotqa\",\n \"searchR1_2wikimultihopqa\",\n \"searchR1_musique\",\n \"searchR1_bamboogle\",\n ]:\n from . import search_r1_like_qa_em\n\n res = search_r1_like_qa_em.compute_score(solution_str, ground_truth)\n\n else:\n raise NotImplementedError(f\"Reward function is not implemented for {data_source=}\")\n\n if isinstance(res, dict):\n return res\n elif isinstance(res, int | float | bool):\n return float(res)\n else:\n return float(res[0])\n\n\n@deprecated(\"verl.utils.reward_score.default_compute_score\")\ndef _default_compute_score(\n data_source,\n solution_str,\n ground_truth,\n extra_info=None,\n sandbox_fusion_url=None,\n concurrent_semaphore=None,\n memory_limit_mb=None,\n):\n \"\"\"\n Legacy function API to be deprecated. Please use `default_compute_score` instead.\n \"\"\"\n return default_compute_score(\n data_source, solution_str, ground_truth, extra_info, sandbox_fusion_url, concurrent_semaphore, memory_limit_mb\n )\n\n\n__all__ = [\"default_compute_score\"]\n"}110{"file_name": "verl__utils__reward_score__geo3k.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport re\n\nfrom mathruler.grader import extract_boxed_content, grade_answer\n\n\ndef format_reward(predict_str: str) -> float:\n pattern = re.compile(r\"<think>.*</think>.*\\\\boxed\\{.*\\}.*\", re.DOTALL)\n match_result = re.fullmatch(pattern, predict_str)\n return 1.0 if match_result else 0.0\n\n\ndef acc_reward(predict_str: str, ground_truth: str, use_boxed: bool = True) -> float:\n if use_boxed:\n answer = extract_boxed_content(predict_str)\n else:\n answer = predict_str\n return 1.0 if grade_answer(answer, ground_truth) else 0.0\n\n\ndef compute_score(predict_str: str, ground_truth: str, use_boxed: bool = True, format_score: float = 0.1) -> float:\n return (1.0 - format_score) * acc_reward(predict_str, ground_truth, use_boxed) + format_score * format_reward(\n predict_str\n )\n"}111{"file_name": "verl__utils__reward_score__gsm8k.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport re\n\n_SOLUTION_CLIP_CHARS = 300\n\n\ndef extract_solution(solution_str, method=\"strict\"):\n assert method in [\"strict\", \"flexible\"]\n\n # Optimization: Regular expression matching on very long strings can be slow.\n # For math problems, the final answer is usually at the end.\n # We only match on the last 300 characters, which is a safe approximation for 300 tokens.\n if len(solution_str) > _SOLUTION_CLIP_CHARS:\n solution_str = solution_str[-_SOLUTION_CLIP_CHARS:]\n\n if method == \"strict\":\n # this also tests the formatting of the model\n solutions = re.findall(\"#### (\\\\-?[0-9\\\\.\\\\,]+)\", solution_str)\n if len(solutions) == 0:\n final_answer = None\n else:\n # take the last solution\n final_answer = solutions[-1].replace(\",\", \"\").replace(\"$\", \"\")\n elif method == \"flexible\":\n answer = re.findall(\"(\\\\-?[0-9\\\\.\\\\,]+)\", solution_str)\n final_answer = None\n if len(answer) == 0:\n # no reward is there is no answer\n pass\n else:\n invalid_str = [\"\", \".\"]\n # find the last number that is not '.'\n for final_answer in reversed(answer):\n if final_answer not in invalid_str:\n break\n return final_answer\n\n\ndef compute_score(solution_str, ground_truth, method=\"strict\", format_score=0.0, score=1.0):\n \"\"\"The scoring function for GSM8k.\n\n Reference: Trung, Luong, et al. \"Reft: Reasoning with reinforced fine-tuning.\" Proceedings of the 62nd Annual\n Meeting of the Association for Computational Linguistics (Volume 1: Long Papers). 2024.\n\n Args:\n solution_str: the solution text\n ground_truth: the ground truth\n method: the method to extract the solution, choices are 'strict' and 'flexible'\n format_score: the score for the format\n score: the score for the correct answer\n \"\"\"\n answer = extract_solution(solution_str=solution_str, method=method)\n if answer is None:\n return 0\n else:\n if answer == ground_truth:\n return score\n else:\n return format_score\n"}112{"file_name": "verl__utils__reward_score__math_dapo.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n# Adapted from https://github.com/EleutherAI/lm-evaluation-harness/blob/main/lm_eval/tasks/hendrycks_math/utils.py\n\nimport re\nfrom typing import Optional\n\n\ndef last_boxed_only_string(string: str) -> Optional[str]:\n \"\"\"Extract the last LaTeX boxed expression from a string.\n\n Args:\n string: Input string containing LaTeX code\n\n Returns:\n The last boxed expression or None if not found\n \"\"\"\n idx = string.rfind(\"\\\\boxed{\")\n if idx < 0:\n return None\n\n i = idx\n right_brace_idx = None\n num_left_braces_open = 0\n\n while i < len(string):\n if string[i] == \"{\":\n num_left_braces_open += 1\n if string[i] == \"}\":\n num_left_braces_open -= 1\n if num_left_braces_open == 0:\n right_brace_idx = i\n break\n i += 1\n\n return string[idx : right_brace_idx + 1] if right_brace_idx is not None else None\n\n\ndef remove_boxed(s: str) -> str:\n \"\"\"Remove the LaTeX boxed command from a string.\n\n Args:\n s: String with format \"\\\\boxed{content}\"\n\n Returns:\n The content inside the boxed command\n \"\"\"\n left = \"\\\\boxed{\"\n assert s[: len(left)] == left, f\"box error: {s}\"\n assert s[-1] == \"}\", f\"box error: {s}\"\n return s[len(left) : -1]\n\n\n# Constants for normalization\nSUBSTITUTIONS = [\n (\"an \", \"\"),\n (\"a \", \"\"),\n (\".$\", \"$\"),\n (\"\\\\$\", \"\"),\n (r\"\\ \", \"\"),\n (\" \", \"\"),\n (\"mbox\", \"text\"),\n (\",\\\\text{and}\", \",\"),\n (\"\\\\text{and}\", \",\"),\n (\"\\\\text{m}\", \"\\\\text{}\"),\n]\n\nREMOVED_EXPRESSIONS = [\n \"square\",\n \"ways\",\n \"integers\",\n \"dollars\",\n \"mph\",\n \"inches\",\n \"hours\",\n \"km\",\n \"units\",\n \"\\\\ldots\",\n \"sue\",\n \"points\",\n \"feet\",\n \"minutes\",\n \"digits\",\n \"cents\",\n \"degrees\",\n \"cm\",\n \"gm\",\n \"pounds\",\n \"meters\",\n \"meals\",\n \"edges\",\n \"students\",\n \"childrentickets\",\n \"multiples\",\n \"\\\\text{s}\",\n \"\\\\text{.}\",\n \"\\\\text{\\ns}\",\n \"\\\\text{}^2\",\n \"\\\\text{}^3\",\n \"\\\\text{\\n}\",\n \"\\\\text{}\",\n r\"\\mathrm{th}\",\n r\"^\\circ\",\n r\"^{\\circ}\",\n r\"\\;\",\n r\",\\!\",\n \"{,}\",\n '\"',\n \"\\\\dots\",\n]\n\n\ndef normalize_final_answer(final_answer: str) -> str:\n \"\"\"Normalize a final answer to a quantitative reasoning question.\n\n Args:\n final_answer: The answer string to normalize\n\n Returns:\n Normalized answer string\n \"\"\"\n final_answer = final_answer.split(\"=\")[-1]\n\n # Apply substitutions and removals\n for before, after in SUBSTITUTIONS:\n final_answer = final_answer.replace(before, after)\n for expr in REMOVED_EXPRESSIONS:\n final_answer = final_answer.replace(expr, \"\")\n\n # Extract and normalize LaTeX math\n final_answer = re.sub(r\"(.*?)(\\$)(.*?)(\\$)(.*)\", \"$\\\\3$\", final_answer)\n final_answer = re.sub(r\"(\\\\text\\{)(.*?)(\\})\", \"\\\\2\", final_answer)\n final_answer = re.sub(r\"(\\\\textbf\\{)(.*?)(\\})\", \"\\\\2\", final_answer)\n final_answer = re.sub(r\"(\\\\overline\\{)(.*?)(\\})\", \"\\\\2\", final_answer)\n final_answer = re.sub(r\"(\\\\boxed\\{)(.*)(\\})\", \"\\\\2\", final_answer)\n\n # Normalize shorthand TeX:\n # \\fracab -> \\frac{a}{b}\n # \\frac{abc}{bef} -> \\frac{abc}{bef}\n # \\fracabc -> \\frac{a}{b}c\n # \\sqrta -> \\sqrt{a}\n # \\sqrtab -> sqrt{a}b\n final_answer = re.sub(r\"(frac)([^{])(.)\", \"frac{\\\\2}{\\\\3}\", final_answer)\n final_answer = re.sub(r\"(sqrt)([^{])\", \"sqrt{\\\\2}\", final_answer)\n final_answer = final_answer.replace(\"$\", \"\")\n\n # Normalize numbers\n if final_answer.replace(\",\", \"\").isdigit():\n final_answer = final_answer.replace(\",\", \"\")\n\n return final_answer.strip()\n\n\ndef is_correct_minerva(\n solution_str: str, gt: str, gt_need_extract: bool = False, answer_pattern: str = r\"(?i)Answer\\s*:\\s*([^\\n]+)\"\n) -> tuple[bool, str]:\n \"\"\"Check if the solution is correct according to Minerva criteria.\n\n Args:\n solution_str: The solution string to check\n gt: The ground truth answer\n gt_need_extract: Whether the ground truth needs extraction\n answer_pattern: Regex pattern to extract the answer\n\n Returns:\n Tuple of (is_correct, normalized_prediction)\n \"\"\"\n # Extract answer from solution\n match = re.findall(answer_pattern, solution_str)\n extracted_answer = match[-1] if match else \"[INVALID]\"\n pred = normalize_final_answer(extracted_answer)\n\n # Process ground truth\n if gt_need_extract:\n gt = normalize_final_answer(remove_boxed(last_boxed_only_string(gt)))\n else:\n gt = normalize_final_answer(gt)\n\n return (pred == gt), pred\n\n\ndef is_correct_strict_box(\n pred: str, gt: str, pause_tokens_index: Optional[list[int]] = None\n) -> tuple[int, Optional[str]]:\n \"\"\"Check if the prediction is correct using strict boxed answer criteria.\n\n Args:\n pred: The prediction string\n gt: The ground truth answer\n pause_tokens_index: Indices of pause tokens\n\n Returns:\n Tuple of (score, extracted_prediction)\n \"\"\"\n # Extract the relevant part of the prediction\n if pause_tokens_index is not None:\n assert len(pause_tokens_index) == 4\n pred = pred[pause_tokens_index[-1] - 100 :]\n else:\n pred = pred[-100:]\n\n # Extract and check the boxed answer\n boxed_pred = last_boxed_only_string(pred)\n extracted_pred = remove_boxed(boxed_pred) if boxed_pred is not None else None\n\n return 1 if (extracted_pred == gt) else -1, extracted_pred\n\n\ndef verify(\n solution_str: str, answer: str, strict_box_verify: bool = False, pause_tokens_index: Optional[list[int]] = None\n) -> bool:\n \"\"\"Verify if the solution is correct.\n\n Args:\n solution_str: The solution string to verify\n answer: The ground truth answer\n strict_box_verify: Whether to use strict box verification\n pause_tokens_index: Indices of pause tokens\n\n Returns:\n True if the solution is correct, False otherwise\n \"\"\"\n if strict_box_verify:\n correct, pred = is_correct_strict_box(solution_str, answer, pause_tokens_index)\n return correct == 1, pred\n\n correct, pred = is_correct_minerva(solution_str, answer)\n return correct, pred\n\n\ndef compute_score(\n solution_str: str,\n ground_truth: str,\n strict_box_verify: bool = False,\n pause_tokens_index: Optional[list[int]] = None,\n) -> float:\n \"\"\"Compute the reward score for a solution.\n\n Args:\n solution_str: The solution string\n ground_truth: The ground truth answer\n strict_box_verify: Whether to use strict box verification\n pause_tokens_index: Indices of pause tokens\n\n Returns:\n Reward score (1.0 for correct, -1.0 for incorrect)\n \"\"\"\n # Limit solution length for efficiency\n solution_str = solution_str[-300:] # The longest answer in MATH-500 has 159 characters\n\n # Verify the solution\n correct, pred = verify(solution_str, ground_truth, strict_box_verify, pause_tokens_index)\n\n reward = 1.0 if correct else -1.0\n acc = correct\n\n return {\n \"score\": reward,\n \"acc\": acc,\n \"pred\": pred,\n }\n"}113{"file_name": "verl__utils__reward_score__math_reward.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n# Adapted from https://github.com/EleutherAI/lm-evaluation-harness/blob/main/lm_eval/tasks/hendrycks_math/utils.py\n\n\ndef compute_score(solution_str, ground_truth) -> float:\n retval = 0.0\n try:\n string_in_last_boxed = last_boxed_only_string(solution_str)\n if string_in_last_boxed is not None:\n answer = remove_boxed(string_in_last_boxed)\n if is_equiv(answer, ground_truth):\n retval = 1.0\n except Exception as e:\n print(e)\n\n return retval\n\n\n# string normalization from https://github.com/EleutherAI/lm-evaluation-harness/blob/master/lm_eval/tasks/hendrycks_math.py\ndef is_equiv(str1, str2, verbose=False):\n if str1 is None and str2 is None:\n print(\"WARNING: Both None\")\n return True\n if str1 is None or str2 is None:\n return False\n\n try:\n ss1 = strip_string(str1)\n ss2 = strip_string(str2)\n if verbose:\n print(ss1, ss2)\n return ss1 == ss2\n except Exception:\n return str1 == str2\n\n\ndef remove_boxed(s):\n if \"\\\\boxed \" in s:\n left = \"\\\\boxed \"\n assert s[: len(left)] == left\n return s[len(left) :]\n\n left = \"\\\\boxed{\"\n\n assert s[: len(left)] == left\n assert s[-1] == \"}\"\n\n return s[len(left) : -1]\n\n\ndef last_boxed_only_string(string):\n idx = string.rfind(\"\\\\boxed\")\n if \"\\\\boxed \" in string:\n return \"\\\\boxed \" + string.split(\"\\\\boxed \")[-1].split(\"$\")[0]\n if idx < 0:\n idx = string.rfind(\"\\\\fbox\")\n if idx < 0:\n return None\n\n i = idx\n right_brace_idx = None\n num_left_braces_open = 0\n while i < len(string):\n if string[i] == \"{\":\n num_left_braces_open += 1\n if string[i] == \"}\":\n num_left_braces_open -= 1\n if num_left_braces_open == 0:\n right_brace_idx = i\n break\n i += 1\n\n retval = None if right_brace_idx is None else string[idx : right_brace_idx + 1]\n\n return retval\n\n\ndef fix_fracs(string):\n substrs = string.split(\"\\\\frac\")\n new_str = substrs[0]\n if len(substrs) > 1:\n substrs = substrs[1:]\n for substr in substrs:\n new_str += \"\\\\frac\"\n if substr[0] == \"{\":\n new_str += substr\n else:\n try:\n assert len(substr) >= 2\n except Exception:\n return string\n a = substr[0]\n b = substr[1]\n if b != \"{\":\n if len(substr) > 2:\n post_substr = substr[2:]\n new_str += \"{\" + a + \"}{\" + b + \"}\" + post_substr\n else:\n new_str += \"{\" + a + \"}{\" + b + \"}\"\n else:\n if len(substr) > 2:\n post_substr = substr[2:]\n new_str += \"{\" + a + \"}\" + b + post_substr\n else:\n new_str += \"{\" + a + \"}\" + b\n string = new_str\n return string\n\n\ndef fix_a_slash_b(string):\n if len(string.split(\"/\")) != 2:\n return string\n a = string.split(\"/\")[0]\n b = string.split(\"/\")[1]\n try:\n a = int(a)\n b = int(b)\n assert string == \"{}/{}\".format(a, b)\n new_string = \"\\\\frac{\" + str(a) + \"}{\" + str(b) + \"}\"\n return new_string\n except Exception:\n return string\n\n\ndef remove_right_units(string):\n # \"\\\\text{ \" only ever occurs (at least in the val set) when describing units\n if \"\\\\text{ \" in string:\n splits = string.split(\"\\\\text{ \")\n assert len(splits) == 2\n return splits[0]\n else:\n return string\n\n\ndef fix_sqrt(string):\n if \"\\\\sqrt\" not in string:\n return string\n splits = string.split(\"\\\\sqrt\")\n new_string = splits[0]\n for split in splits[1:]:\n if split[0] != \"{\":\n a = split[0]\n new_substr = \"\\\\sqrt{\" + a + \"}\" + split[1:]\n else:\n new_substr = \"\\\\sqrt\" + split\n new_string += new_substr\n return new_string\n\n\ndef strip_string(string):\n # linebreaks\n string = string.replace(\"\\n\", \"\")\n\n # remove inverse spaces\n string = string.replace(\"\\\\!\", \"\")\n\n # replace \\\\ with \\\n string = string.replace(\"\\\\\\\\\", \"\\\\\")\n\n # replace tfrac and dfrac with frac\n string = string.replace(\"tfrac\", \"frac\")\n string = string.replace(\"dfrac\", \"frac\")\n\n # remove \\left and \\right\n string = string.replace(\"\\\\left\", \"\")\n string = string.replace(\"\\\\right\", \"\")\n\n # Remove circ (degrees)\n string = string.replace(\"^{\\\\circ}\", \"\")\n string = string.replace(\"^\\\\circ\", \"\")\n\n # remove dollar signs\n string = string.replace(\"\\\\$\", \"\")\n\n # remove units (on the right)\n string = remove_right_units(string)\n\n # remove percentage\n string = string.replace(\"\\\\\\\\%\", \"\")\n string = string.replace(\"\\\\%\", \"\")\n\n # \" 0.\" equivalent to \" .\" and \"{0.\" equivalent to \"{.\" Alternatively, add \"0\" if \".\" is the start of the string\n string = string.replace(\" .\", \" 0.\")\n string = string.replace(\"{.\", \"{0.\")\n # if empty, return empty string\n if len(string) == 0:\n return string\n if string[0] == \".\":\n string = \"0\" + string\n\n # to consider: get rid of e.g. \"k = \" or \"q = \" at beginning\n if len(string.split(\"=\")) == 2 and len(string.split(\"=\")[0]) <= 2:\n string = string.split(\"=\")[1]\n\n # fix sqrt3 --> sqrt{3}\n string = fix_sqrt(string)\n\n # remove spaces\n string = string.replace(\" \", \"\")\n\n # \\frac1b or \\frac12 --> \\frac{1}{b} and \\frac{1}{2}, etc. Even works with \\frac1{72} (but not \\frac{72}1).\n # Also does a/b --> \\\\frac{a}{b}\n string = fix_fracs(string)\n\n # manually change 0.5 --> \\frac{1}{2}\n if string == \"0.5\":\n string = \"\\\\frac{1}{2}\"\n\n # NOTE: X/Y changed to \\frac{X}{Y} in dataset, but in simple cases fix in case the model output is X/Y\n string = fix_a_slash_b(string)\n\n return string\n"}114{"file_name": "verl__utils__reward_score__prime_code__utils.py", "text": "# Copyright 2024 PRIME team and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n# Borrowed from: https://huggingface.co/spaces/codeparrot/apps_metric/blob/main/utils.py\n\nimport multiprocessing\nimport os\nimport sys\nimport traceback\nfrom typing import Optional\n\nfrom .testing_util import run_test\n\n\ndef _temp_run(sample, generation, debug, result, metadata_list, timeout):\n with open(os.devnull, \"w\") as devnull:\n sys.stdout = devnull\n sys.stderr = devnull\n try:\n res, metadata = run_test(in_outs=sample, test=generation, debug=debug, timeout=timeout)\n result.append(res)\n metadata_list.append(metadata)\n except Exception:\n # print(e) # some tracebacks are extremely long.\n traceback.print_exc(10)\n result.append([-1 for i in range(len(sample[\"inputs\"]))])\n metadata_list.append({})\n\n\ndef check_correctness(in_outs: Optional[dict], generation, timeout=10, debug=True):\n \"\"\"Check correctness of code generation with a global timeout.\n The global timeout is to catch some extreme/rare cases not handled by the timeouts\n inside `run_test`\"\"\"\n\n manager = multiprocessing.Manager()\n result = manager.list()\n metadata_list = manager.list()\n p = multiprocessing.Process(target=_temp_run, args=(in_outs, generation, debug, result, metadata_list, timeout))\n p.start()\n p.join(timeout=timeout + 1)\n if p.is_alive():\n p.kill()\n # p.terminate()\n if not result:\n # consider that all tests failed\n result = [[-1 for i in range(len(in_outs[\"inputs\"]))]]\n if debug:\n print(\"global timeout\")\n return result[0], metadata_list\n"}115{"file_name": "verl__utils__reward_score__prime_math__grader.py", "text": "# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n# Copyright (c) Microsoft Corporation.\n#\n# Permission is hereby granted, free of charge, to any person obtaining a copy\n# of this software and associated documentation files (the \"Software\"), to deal\n# in the Software without restriction, including without limitation the rights\n# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell\n# copies of the Software, and to permit persons to whom the Software is\n# furnished to do so, subject to the following conditions:\n#\n# The above copyright notice and this permission notice shall be included in all\n# copies or substantial portions of the Software.\n#\n# THE SOFTWARE IS PROVIDED \"AS IS\", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR\n# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,\n# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE\n# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER\n# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,\n# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE\n# SOFTWARE\n\n# Copyright (c) 2023 OpenAI\n#\n# Permission is hereby granted, free of charge, to any person obtaining a copy\n# of this software and associated documentation files (the \"Software\"), to deal\n# in the Software without restriction, including without limitation the rights\n# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell\n# copies of the Software, and to permit persons to whom the Software is\n# furnished to do so, subject to the following conditions:\n\n# The above copyright notice and this permission notice shall be included in all\n# copies or substantial portions of the Software.\n#\n# THE SOFTWARE IS PROVIDED \"AS IS\", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR\n# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,\n# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE\n# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER\n# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,\n# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE\n# SOFTWARE.\n\n# Copyright (c) 2021 Dan Hendrycks\n#\n# Permission is hereby granted, free of charge, to any person obtaining a copy\n# of this software and associated documentation files (the \"Software\"), to deal\n# in the Software without restriction, including without limitation the rights\n# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell\n# copies of the Software, and to permit persons to whom the Software is\n# furnished to do so, subject to the following conditions:\n#\n# The above copyright notice and this permission notice shall be included in all\n# copies or substantial portions of the Software.\n#\n# THE SOFTWARE IS PROVIDED \"AS IS\", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR\n# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,\n# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE\n# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER\n# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,\n# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE\n# SOFTWARE.\n\n# Copyright 2024 PRIME team and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nThis logic is largely copied from the Hendrycks' MATH release (math_equivalence), and borrowed from:\n- https://github.com/microsoft/ToRA/blob/main/src/eval/grader.py\n- https://github.com/microsoft/ProphetNet/tree/master/CRITIC\n- https://github.com/openai/prm800k\n\"\"\"\n\nimport contextlib\nimport math\nimport re\nfrom math import isclose\n\n# sympy related\nfrom sympy import N, simplify\nfrom sympy.parsing.latex import parse_latex\nfrom sympy.parsing.sympy_parser import parse_expr\n\n# verl related\nfrom verl.utils.py_functional import timeout_limit\n\n\ndef is_digit(s):\n try:\n if \"{,}\" in str(s):\n num = float(str(s).replace(\"{,}\", \"\"))\n return True, num\n\n num = float(str(s).replace(\",\", \"\"))\n return True, num\n except ValueError:\n return False, None\n\n\ndef normalize(answer, pi) -> str:\n # checking if answer is $<number> and removing $ in that case to compare\n if isinstance(answer, str) and bool(re.match(r\"\\$\\d+(\\.\\d+)?\", answer)):\n return answer[1:]\n\n # checking if answer is <number>% or <number>\\\\% and removing %\n if isinstance(answer, str) and (\n bool(re.match(r\"^\\d+(\\.\\d+)?%$\", answer)) or bool(re.match(r\"^\\d+(\\.\\d+)?\\\\%$\", answer))\n ):\n return answer.replace(\"\\\\%\", \"\").replace(\"%\", \"\")\n\n # handle base\n answer = handle_base(answer)\n\n # handle pi\n answer = handle_pi(answer, pi)\n\n return answer\n\n\ndef handle_base(x) -> str:\n if isinstance(x, str) and \"_\" in x:\n # Due to base\n x = x.split(\"_\")[0]\n x = float(x)\n return int(x)\n return x\n\n\ndef handle_pi(string, pi):\n if isinstance(string, str) and \"\\\\pi\" in string:\n # Find the first occurrence of \"\\pi\"\n idx = string.find(\"\\\\pi\")\n\n # Iterate over the string and find all occurrences of \"\\pi\" with a valid previous character\n while idx != -1:\n if idx > 0 and string[idx - 1].isdigit():\n # Replace \"\\pi\" with \"*math.pi\" if the previous character is a digit\n string = string[:idx] + f\"*{pi}\" + string[idx + 3 :]\n else:\n # Replace \"\\pi\" with \"1*math.pi\" if the previous character is not a digit\n string = string[:idx] + f\"1*{pi}\" + string[idx + 3 :]\n\n # Find the next occurrence of \"\\pi\"\n idx = string.find(\"\\\\pi\", idx + 1)\n\n # Evaluate the expression using eval() function\n with contextlib.suppress(Exception):\n string = eval(string)\n\n return string\n\n\ndef math_equal(\n prediction: bool | float | str,\n reference: float | str,\n include_percentage: bool = True,\n tolerance: float = 1e-4,\n timeout: float = 10.0,\n pi: float = math.pi,\n) -> bool:\n \"\"\"\n Exact match of math if and only if:\n 1. numerical equal: both can convert to float and are equal\n 2. symbolic equal: both can convert to sympy expression and are equal\n \"\"\"\n\n prediction = normalize(prediction, pi)\n reference = normalize(reference, pi)\n\n if isinstance(prediction, str) and len(prediction) > 1000: # handling weird corner-cases\n prediction = prediction[:1000]\n\n # 0. string comparison\n if isinstance(prediction, str) and isinstance(reference, str):\n if prediction.strip().lower() == reference.strip().lower():\n return True\n if prediction.replace(\" \", \"\") == reference.replace(\" \", \"\"):\n return True\n\n try: # 1. numerical equal\n if is_digit(prediction)[0] and is_digit(reference)[0]:\n prediction = is_digit(prediction)[1]\n reference = is_digit(reference)[1]\n # number questions\n gt_result = [reference / 100, reference, reference * 100] if include_percentage else [reference]\n for item in gt_result:\n try:\n if isclose(item, prediction, rel_tol=tolerance):\n return True\n except Exception:\n continue\n return False\n except Exception:\n pass\n\n if not prediction and prediction not in [0, False]:\n return False\n\n # 2. symbolic equal\n reference = str(reference).strip()\n prediction = str(prediction).strip()\n\n ## deal with [], (), {}\n prediction = format_intervals(prediction)\n\n pred_str, ref_str = prediction, reference\n if (prediction.startswith(\"[\") and prediction.endswith(\"]\") and not reference.startswith(\"(\")) or (\n prediction.startswith(\"(\") and prediction.endswith(\")\") and not reference.startswith(\"[\")\n ):\n pred_str = pred_str.strip(\"[]()\")\n ref_str = ref_str.strip(\"[]()\")\n for s in [\"{\", \"}\", \"(\", \")\"]:\n ref_str = ref_str.replace(s, \"\")\n pred_str = pred_str.replace(s, \"\")\n if pred_str == ref_str:\n return True\n\n ## [a, b] vs. [c, d], return a==c and b==d\n if (\n prediction\n and reference\n and prediction[0] in \"([\"\n and prediction[-1] in \")]\"\n and prediction[0] == reference[0]\n and prediction[-1] == reference[-1]\n ):\n pred_parts = prediction[1:-1].split(\",\")\n ref_parts = reference[1:-1].split(\",\")\n if len(pred_parts) == len(ref_parts) and all(\n [\n math_equal(pred_pt, ref_pt, include_percentage, tolerance)\n for pred_pt, ref_pt in zip(pred_parts, ref_parts, strict=True)\n ]\n ):\n return True\n\n if \",\" in prediction and \",\" in reference:\n pred_parts = [item.strip() for item in prediction.split(\",\")]\n ref_parts = [item.strip() for item in reference.split(\",\")]\n\n if len(pred_parts) == len(ref_parts):\n return bool(\n all(\n [\n math_equal(pred_parts[i], ref_parts[i], include_percentage, tolerance)\n for i in range(len(pred_parts))\n ]\n )\n )\n\n # if we have point == tuple of values\n if prediction.startswith(\"Point\") and reference[0] == \"(\" and reference[-1] == \")\":\n pred_parts = prediction[prediction.find(\"(\") + 1 : -1].split(\",\")\n ref_parts = reference[1:-1].split(\",\")\n if len(pred_parts) == len(ref_parts) and all(\n [\n math_equal(pred_pt, ref_pt, include_percentage, tolerance)\n for pred_pt, ref_pt in zip(pred_parts, ref_parts, strict=False)\n ]\n ):\n return True\n\n # if reference is a matrix\n if r\"\\begin{pmatrix}\" in reference and prediction.startswith(\"Matrix\"):\n try:\n pred_matrix = parse_expr(prediction)\n ref_matrix_items = reference.split()[1:-1:2]\n if len(pred_matrix) == len(ref_matrix_items) and all(\n [\n math_equal(pred, ref, include_percentage, tolerance)\n for ref, pred in zip(ref_matrix_items, pred_matrix, strict=False)\n ]\n ):\n return True\n except Exception:\n pass\n elif r\"\\begin{pmatrix}\" in reference and prediction.startswith(\"[\") and prediction.endswith(\"]\"):\n if isinstance(eval(prediction), list):\n try:\n pred_matrix = eval(prediction)\n # ref_matrix_items = reference.split()[1:-1:2]\n ref_matrix_items = (\n reference.removeprefix(r\"\\\\begin{pmatrix}\")\n .removeprefix(r\"\\begin{pmatrix}\")\n .removesuffix(r\"\\\\end{pmatrix}\")\n .removesuffix(r\"\\end{pmatrix}\")\n )\n ref_matrix_items = ref_matrix_items.split(\"\\\\\")\n ref_matrix_items = [row.split(\"&\") if \"&\" in row else row for row in ref_matrix_items]\n if len(pred_matrix) == len(ref_matrix_items) and all(\n [\n math_equal(pred, ref, include_percentage, tolerance)\n for ref, pred in zip(ref_matrix_items, pred_matrix, strict=False)\n ]\n ):\n return True\n except Exception:\n pass\n\n return symbolic_equal(prediction, reference, tolerance, timeout)\n\n\ndef symbolic_equal(a, b, tolerance, timeout=10.0):\n def _parse(s):\n for f in [parse_expr, parse_latex]:\n try:\n with timeout_limit(seconds=timeout):\n return f(s)\n except TimeoutError:\n print(f\"Parsing timed out for {s}\")\n continue\n except Exception:\n continue\n return s\n\n a = _parse(a)\n b = _parse(b)\n\n try:\n with timeout_limit(seconds=timeout):\n if simplify(a - b) == 0:\n return True\n except TimeoutError:\n print(f\"Simplification timed out for {a} - {b}\")\n pass\n except Exception:\n pass\n\n try:\n with timeout_limit(seconds=timeout):\n if isclose(N(a), N(b), rel_tol=tolerance):\n return True\n except TimeoutError:\n print(f\"Numerical evaluation timed out for {a}, {b}\")\n pass\n except Exception:\n pass\n return False\n\n\ndef format_intervals(prediction):\n patterns = {\n \"Interval(\": r\"^Interval\\((.*)\\)$\",\n \"Interval.Ropen(\": r\"^Interval\\.Ropen\\((.*)\\)$\",\n \"Interval.Lopen(\": r\"^Interval\\.Lopen\\((.*)\\)$\",\n \"Interval.open(\": r\"^Interval\\.open\\((.*)\\)$\",\n }\n\n for key, pattern in patterns.items():\n match = re.match(pattern, prediction)\n if match:\n inner_content = match.group(1)\n\n if key == \"Interval(\": # Intarval(a, b) == [a, b]\n return f\"[{inner_content}]\"\n elif key == \"Interval.Ropen(\": # Intarval.Ropen(a, b) == [a, b)\n return f\"[{inner_content})\"\n elif key == \"Interval.Lopen(\": # Intarval.Lopen(a, b) == (a, b]\n return f\"({inner_content}]\"\n elif key == \"Interval.open(\": # Intarval.open(a, b) == (a, b)\n return f\"({inner_content})\"\n\n return prediction\n"}116{"file_name": "verl__utils__reward_score__prime_math__math_normalize.py", "text": "# Copyright 2024 PRIME team and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n# Copyright (c) 2021 Dan Hendrycks\n#\n# Permission is hereby granted, free of charge, to any person obtaining a copy\n# of this software and associated documentation files (the \"Software\"), to deal\n# in the Software without restriction, including without limitation the rights\n# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell\n# copies of the Software, and to permit persons to whom the Software is\n# furnished to do so, subject to the following conditions:\n#\n# The above copyright notice and this permission notice shall be included in all\n# copies or substantial portions of the Software.\n#\n# THE SOFTWARE IS PROVIDED \"AS IS\", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR\n# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,\n# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE\n# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER\n# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,\n# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE\n# SOFTWARE.\n\"\"\"\nThis logic is largely copied from the Hendrycks' MATH release (math_equivalence).\n\nFrom: https://github.com/openai/prm800k/blob/main/prm800k/grading/math_normalize.py\n\"\"\"\n\nimport re\nfrom typing import Optional\n\n\ndef normalize_answer(answer: Optional[str]) -> Optional[str]:\n if answer is None:\n return None\n answer = answer.strip()\n try:\n # Remove enclosing `\\text{}`.\n m = re.search(r\"^\\\\text\\{(?P<text>.+?)\\}$\", answer)\n if m is not None:\n answer = m.group(\"text\").strip()\n return _strip_string(answer)\n except Exception:\n return answer\n\n\ndef _fix_fracs(string):\n substrs = string.split(\"\\\\frac\")\n new_str = substrs[0]\n if len(substrs) > 1:\n substrs = substrs[1:]\n for substr in substrs:\n new_str += \"\\\\frac\"\n if substr[0] == \"{\":\n new_str += substr\n else:\n try:\n assert len(substr) >= 2\n except Exception:\n return string\n a = substr[0]\n b = substr[1]\n if b != \"{\":\n if len(substr) > 2:\n post_substr = substr[2:]\n new_str += \"{\" + a + \"}{\" + b + \"}\" + post_substr\n else:\n new_str += \"{\" + a + \"}{\" + b + \"}\"\n else:\n if len(substr) > 2:\n post_substr = substr[2:]\n new_str += \"{\" + a + \"}\" + b + post_substr\n else:\n new_str += \"{\" + a + \"}\" + b\n string = new_str\n return string\n\n\ndef _fix_a_slash_b(string):\n if len(string.split(\"/\")) != 2:\n return string\n a = string.split(\"/\")[0]\n b = string.split(\"/\")[1]\n try:\n a = int(a)\n b = int(b)\n assert string == \"{}/{}\".format(a, b)\n new_string = \"\\\\frac{\" + str(a) + \"}{\" + str(b) + \"}\"\n return new_string\n except Exception:\n return string\n\n\ndef _remove_right_units(string):\n # \"\\\\text{ \" only ever occurs (at least in the val set) when describing units\n if \"\\\\text{ \" in string:\n splits = string.split(\"\\\\text{ \")\n assert len(splits) == 2\n return splits[0]\n else:\n return string\n\n\ndef _fix_sqrt(string):\n if \"\\\\sqrt\" not in string:\n return string\n splits = string.split(\"\\\\sqrt\")\n new_string = splits[0]\n for split in splits[1:]:\n if split[0] != \"{\":\n a = split[0]\n new_substr = \"\\\\sqrt{\" + a + \"}\" + split[1:]\n else:\n new_substr = \"\\\\sqrt\" + split\n new_string += new_substr\n return new_string\n\n\ndef _strip_string(string):\n # linebreaks\n string = string.replace(\"\\n\", \"\")\n\n # remove inverse spaces\n string = string.replace(\"\\\\!\", \"\")\n\n # replace \\\\ with \\\n string = string.replace(\"\\\\\\\\\", \"\\\\\")\n\n # replace tfrac and dfrac with frac\n string = string.replace(\"tfrac\", \"frac\")\n string = string.replace(\"dfrac\", \"frac\")\n\n # remove \\left and \\right\n string = string.replace(\"\\\\left\", \"\")\n string = string.replace(\"\\\\right\", \"\")\n\n # Remove circ (degrees)\n string = string.replace(\"^{\\\\circ}\", \"\")\n string = string.replace(\"^\\\\circ\", \"\")\n\n # remove dollar signs\n string = string.replace(\"\\\\$\", \"\")\n\n # remove units (on the right)\n string = _remove_right_units(string)\n\n # remove percentage\n string = string.replace(\"\\\\\\\\%\", \"\")\n string = string.replace(\"\\\\%\", \"\")\n\n # \" 0.\" equivalent to \" .\" and \"{0.\" equivalent to \"{.\" Alternatively, add \"0\" if \".\" is the start of the string\n string = string.replace(\" .\", \" 0.\")\n string = string.replace(\"{.\", \"{0.\")\n # if empty, return empty string\n if len(string) == 0:\n return string\n if string[0] == \".\":\n string = \"0\" + string\n\n # to consider: get rid of e.g. \"k = \" or \"q = \" at beginning\n if len(string.split(\"=\")) == 2 and len(string.split(\"=\")[0]) <= 2:\n string = string.split(\"=\")[1]\n\n # fix sqrt3 --> sqrt{3}\n string = _fix_sqrt(string)\n\n # remove spaces\n string = string.replace(\" \", \"\")\n\n # \\frac1b or \\frac12 --> \\frac{1}{b} and \\frac{1}{2}, etc. Even works with \\frac1{72} (but not \\frac{72}1).\n # Also does a/b --> \\\\frac{a}{b}\n string = _fix_fracs(string)\n\n # manually change 0.5 --> \\frac{1}{2}\n if string == \"0.5\":\n string = \"\\\\frac{1}{2}\"\n\n # NOTE: X/Y changed to \\frac{X}{Y} in dataset, but in simple cases fix in case the model output is X/Y\n string = _fix_a_slash_b(string)\n\n return string\n"}117{"file_name": "verl__utils__reward_score__sandbox_fusion__utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport concurrent.futures # <-- Import concurrent.futures\nimport json\nimport logging\nimport os\nimport threading\nimport time\nimport traceback\nimport uuid\nfrom typing import Any, Optional\n\nimport requests\n\nDEFAULT_TIMEOUT = 10 # Default compile and run timeout\nMAX_RETRIES = 3\nINITIAL_RETRY_DELAY = 1\nAPI_TIMEOUT = 10\n\nlogger = logging.getLogger(__name__)\n\n# Define supported languages list (optional, for documentation or validation)\nSUPPORTED_LANGUAGES = [\n \"python\",\n \"cpp\",\n \"nodejs\",\n \"go\",\n \"go_test\",\n \"java\",\n \"php\",\n \"csharp\",\n \"bash\",\n \"typescript\",\n \"sql\",\n \"rust\",\n \"cuda\",\n \"lua\",\n \"R\",\n \"perl\",\n \"D_ut\",\n \"ruby\",\n \"scala\",\n \"julia\",\n \"pytest\",\n \"junit\",\n \"kotlin_script\",\n \"jest\",\n \"verilog\",\n \"python_gpu\",\n \"lean\",\n \"swift\",\n \"racket\",\n]\n\n\ndef call_sandbox_api(\n sandbox_fusion_url: str,\n code: str,\n stdin: Optional[str],\n compile_timeout: int,\n run_timeout: int,\n memory_limit_mb: int,\n language: str = \"python\",\n) -> tuple[Optional[dict[str, Any]], Optional[str]]: # <-- Remove request_id parameter\n \"\"\"\n Calls the remote sandbox API to execute code with retry logic for Gateway Timeout,\n using increasing delay between retries. Logs internal calls with a unique ID.\n\n Args:\n sandbox_fusion_url: The URL of the sandbox fusion API.\n code: The code string to execute.\n stdin: The standard input string.\n compile_timeout: Compile timeout in seconds.\n run_timeout: Run timeout in seconds.\n language: The programming language of the code (e.g., \"python\", \"cpp\", \"java\"). Defaults to \"python\".\n\n Returns:\n A tuple (response_json, error_message).\n If successful, response_json is the API's returned JSON object, error_message is None.\n If failed after retries, response_json is None, error_message contains the error information.\n \"\"\"\n request_id = str(uuid.uuid4()) # <-- Generate request_id internally\n log_prefix = f\"[Request ID: {request_id}] \" # <-- Create log prefix\n\n if language not in SUPPORTED_LANGUAGES:\n error_msg = f\"{log_prefix}Unsupported language: {language}\"\n logger.error(error_msg)\n return None, error_msg\n\n payload = json.dumps(\n {\n \"compile_timeout\": compile_timeout,\n \"run_timeout\": run_timeout,\n \"code\": code,\n \"stdin\": stdin,\n \"memory_limit_MB\": memory_limit_mb,\n \"language\": language, # Use the passed language parameter\n \"files\": {},\n \"fetch_files\": [],\n }\n )\n headers = {\"Content-Type\": \"application/json\", \"Accept\": \"application/json\"}\n # Calculate a reasonable request timeout based on compile/run timeouts plus a buffer\n request_timeout = compile_timeout + run_timeout + API_TIMEOUT\n\n last_error = None # Store the last error encountered\n\n for attempt in range(MAX_RETRIES):\n try:\n logger.info(\n f\"{log_prefix}Attempt {attempt + 1}/{MAX_RETRIES}: Calling sandbox API at {sandbox_fusion_url}\"\n ) # <-- Use internal log_prefix\n response = requests.post(\n sandbox_fusion_url,\n headers=headers,\n data=payload,\n timeout=request_timeout, # Use the calculated timeout\n )\n\n # Check for Gateway Timeout (504) specifically for retrying\n if response.status_code == 504:\n last_error = (\n f\"{log_prefix}API Request Error: Gateway Timeout (504) on attempt \"\n f\"{attempt + 1}/{MAX_RETRIES}\"\n ) # <-- Use internal log_prefix\n logger.warning(last_error)\n if attempt < MAX_RETRIES - 1: # Don't sleep after the last attempt\n # Calculate increasing delay (e.g., 1s, 2s, 4s, ...) or (1s, 2s, 3s, ...)\n # Simple linear increase: delay = INITIAL_RETRY_DELAY * (attempt + 1)\n # Exponential backoff: delay = INITIAL_RETRY_DELAY * (2 ** attempt)\n delay = INITIAL_RETRY_DELAY * (attempt + 1) # Using linear increase for simplicity\n logger.info(f\"{log_prefix}Retrying after {delay} seconds...\") # <-- Use internal log_prefix\n time.sleep(delay)\n continue # Go to the next retry attempt\n\n # Check for other HTTP errors (e.g., 4xx, other 5xx)\n response.raise_for_status()\n\n # If successful (status code 2xx)\n logger.info(\n f\"{log_prefix}Sandbox API call successful on attempt {attempt + 1}\"\n ) # <-- Use internal log_prefix\n return response.json(), None\n\n except requests.exceptions.RequestException as e:\n last_error = f\"{log_prefix}API Request Error: {e}\" # <-- Use internal log_prefix\n break # Exit retry loop on non-504 request errors\n except json.JSONDecodeError as e:\n raw_response_text = response.text if \"response\" in locals() else \"N/A\"\n last_error = f\"{log_prefix}API Response JSON Decode Error: {e}\" # <-- Use internal log_prefix\n break # Exit retry loop on JSON decode errors\n except Exception as e:\n last_error = f\"{log_prefix}Unexpected Error: {e}\" # <-- Use internal log_prefix\n break # Exit retry loop on other unexpected errors\n\n # If loop finishes without returning success, return the last recorded error\n logger.error(f\"{log_prefix}Sandbox API call failed. Last error: {last_error}\") # <-- Use internal log_prefix\n # Return the error message without the prefix, as the caller doesn't need the internal ID\n # Ensure API call failure returns error message, leading to -1 in check_correctness\n return None, last_error.replace(log_prefix, \"API Call Failed: \") if last_error else \"API Call Failed after retries\"\n\n\ndef _process_single_case(\n case_index: int,\n stdin_data: Any,\n expected_output: Any,\n sandbox_fusion_url: str,\n generation: str,\n timeout: int,\n memory_limit_mb: int,\n language: str,\n concurrent_semaphore: Optional[threading.Semaphore] = None,\n fn_name: Optional[str] = None,\n) -> tuple[int, dict[str, Any]]:\n \"\"\"Helper function to process a single test case.\"\"\"\n api_response = None\n error_msg = None\n logger.info(f\"Processing test case {case_index + 1}.\")\n\n current_generation_code = generation\n\n if fn_name and language == \"python\":\n # Wrapper assumes stdin_data is a JSON string for function arguments.\n wrapper_code = f\"\"\"\nimport traceback\nfrom string import *\nfrom re import *\nfrom datetime import *\nfrom collections import *\nfrom heapq import *\nfrom bisect import *\nfrom copy import *\nfrom math import *\nfrom random import *\nfrom statistics import *\nfrom itertools import *\nfrom functools import *\nfrom operator import *\nfrom io import *\nfrom sys import *\nfrom json import *\nfrom builtins import *\nfrom typing import *\nimport string\nimport re\nimport datetime\nimport collections\nimport heapq\nimport bisect\nimport copy\nimport math\nimport random\nimport statistics\nimport itertools\nimport functools\nimport operator\nimport io\nimport sys\nimport json\n\n# === User's Original Code START ===\n{generation}\n# === User's Original Code END ===\n\n_SANDBOX_FN_NAME = \"{fn_name}\"\n\ndef _execute_user_function():\n # --- Input Parsing ---\n _raw_input_str = sys.stdin.read()\n _args = []\n if _raw_input_str.strip(): # If there's input\n try:\n _args = [json.loads(line) for line in _raw_input_str.split('\\\\n')]\n except json.JSONDecodeError as _je:\n sys.stderr.write(f\"WrapperError: Invalid JSON input for '{{_SANDBOX_FN_NAME}}': {{_je}}\\\\nInput was: \"\n f\"{{_raw_input_str[:200]}}\\\\n\")\n return None, True # result, error_occurred\n\n # --- Function Location and Execution ---\n try:\n _target_callable = None\n # Try global scope first\n if _SANDBOX_FN_NAME in globals():\n _target_callable = globals()[_SANDBOX_FN_NAME]\n # Else, if 'Solution' class exists, try to get its method\n elif 'Solution' in globals():\n _Solution_class = globals()['Solution']\n # Attempt to instantiate and get method.\n # Errors (e.g., Solution not a class, instantiation fails, method missing)\n # will be caught by the broad except block below.\n _solution_instance = _Solution_class()\n _target_callable = getattr(_solution_instance, _SANDBOX_FN_NAME)\n\n if not _target_callable:\n sys.stderr.write(f\"WrapperError: Function or method '{{_SANDBOX_FN_NAME}}' not found.\\\\n\")\n return None, True # result, error_occurred\n\n _fn_result = _target_callable(*_args)\n return _fn_result, False # result, no_error\n except Exception: # Catches errors from Solution instantiation, getattr, or function call\n sys.stderr.write(f\"Error during setup or execution of '{{_SANDBOX_FN_NAME}}':\\\\n{{traceback.format_exc()}}\\\\n\")\n return None, True # result, error_occurred\n\nif __name__ == '__main__':\n _result, _error_occurred = _execute_user_function()\n\n if not _error_occurred:\n # Serialize result to stdout\n if isinstance(_result, (dict, list, tuple)) or _result is None or isinstance(_result, bool):\n print(json.dumps(_result))\n elif isinstance(_result, (int, float, str)):\n print(str(_result)) # Ensure string conversion for print\n else:\n # For other types, default to string representation.\n print(str(_result))\n # Optional: To explicitly exit with an error code if the sandbox relies on it\n # else:\n # sys.exit(1)\n\"\"\"\n current_generation_code = wrapper_code\n\n stdin = None if stdin_data is None else str(stdin_data)\n try:\n if concurrent_semaphore:\n # logger.debug(f\"Case {case_index + 1}: Attempting to acquire semaphore.\")\n with concurrent_semaphore:\n # logger.debug(f\"Case {case_index + 1}: Semaphore acquired. Calling API.\")\n api_response, error_msg = call_sandbox_api(\n sandbox_fusion_url=sandbox_fusion_url,\n code=current_generation_code,\n stdin=stdin,\n compile_timeout=timeout,\n run_timeout=timeout,\n memory_limit_mb=memory_limit_mb,\n language=language,\n )\n # logger.debug(f\"Case {case_index + 1}: Semaphore released.\")\n else:\n api_response, error_msg = call_sandbox_api(\n sandbox_fusion_url=sandbox_fusion_url,\n code=current_generation_code,\n stdin=stdin,\n compile_timeout=timeout,\n run_timeout=timeout,\n memory_limit_mb=memory_limit_mb,\n language=language,\n )\n except Exception as e:\n error_msg = f\"API Request Exception during check_correctness for case {case_index + 1}: {e}\"\n logger.error(f\"Case {case_index + 1}: {error_msg}\")\n traceback.print_exc()\n\n metadata = {\n \"case_index\": case_index,\n \"input\": stdin,\n \"expected_output\": str(expected_output) if expected_output else None,\n \"api_request_error\": error_msg,\n \"api_response\": None,\n \"status\": \"unknown\",\n \"stdout\": None,\n \"stderr\": None,\n \"exit_code\": None,\n \"duration\": None,\n \"compile_duration\": None,\n \"compile_stderr\": None,\n \"api_status\": None,\n \"compile_status\": None,\n \"run_status\": None,\n }\n result_status = -1 # Default error: API request error or unknown sandbox error\n\n if error_msg:\n metadata[\"status\"] = \"api_error\"\n result_status = -1 # API request itself failed (includes timeout after retries)\n logger.error(f\"Case {case_index}: API error occurred: {error_msg}\")\n # Log code and input only on error for brevity\n generation_to_log = generation[:200] + \"...\" if len(generation) > 200 else generation\n logger.error(f\"Case {case_index}: code: {generation_to_log}\")\n logger.error(f\"Case {case_index}: input: {stdin}\")\n elif api_response:\n # --- Add debug logging ---\n logger.debug(f\"Case {case_index}: API Response: {api_response}\")\n metadata[\"api_response\"] = api_response\n metadata[\"api_status\"] = api_response.get(\"status\")\n compile_result = api_response.get(\"compile_result\")\n run_result = api_response.get(\"run_result\")\n\n # Extract compile information\n if compile_result:\n metadata[\"compile_status\"] = compile_result.get(\"status\")\n metadata[\"compile_duration\"] = compile_result.get(\"execution_time\")\n metadata[\"compile_stderr\"] = compile_result.get(\"stderr\")\n\n # Extract run information\n if run_result:\n metadata[\"run_status\"] = run_result.get(\"status\")\n metadata[\"stdout\"] = run_result.get(\"stdout\")\n metadata[\"stderr\"] = run_result.get(\"stderr\") # stderr during runtime\n metadata[\"exit_code\"] = run_result.get(\"return_code\")\n metadata[\"duration\"] = run_result.get(\"execution_time\")\n\n # --- Determine status based on API response ---\n api_status = metadata[\"api_status\"]\n\n if api_status == \"SandboxError\":\n metadata[\"status\"] = \"sandbox_error\"\n result_status = -1 # Internal sandbox error\n elif api_status == \"Failed\":\n # --- Add debug logging ---\n logger.debug(f\"API returned Failed status. Response: {api_response}\")\n logger.debug(f\"Compile Result: {compile_result}\")\n logger.debug(f\"Run Result: {run_result}\")\n # --- Check the logic here ---\n # Compile failed or timed out\n is_compile_error = compile_result and (\n metadata[\"compile_status\"] in [\"Error\", \"TimeLimitExceeded\"]\n or (metadata[\"compile_status\"] == \"Finished\" and compile_result.get(\"return_code\") != 0)\n )\n if is_compile_error:\n # Differentiate between compile_error and compile_timeout based on specific status\n if metadata[\"compile_status\"] == \"TimeLimitExceeded\":\n metadata[\"status\"] = \"compile_timeout\"\n else: # Includes Error and Finished but return_code != 0 cases\n metadata[\"status\"] = \"compile_error\"\n result_status = -4\n # Run failed or timed out\n elif run_result:\n # Modified condition: Check for TimeLimitExceeded OR (Finished with non-zero exit code) OR Error status\n is_runtime_error = (\n metadata[\"run_status\"] == \"TimeLimitExceeded\"\n or metadata[\"run_status\"] == \"Error\"\n or (metadata[\"run_status\"] == \"Finished\" and run_result.get(\"return_code\") != 0)\n )\n if is_runtime_error:\n if metadata[\"run_status\"] == \"TimeLimitExceeded\":\n metadata[\"status\"] = \"timeout\" # Runtime timeout\n result_status = -3\n else: # Includes Error and Finished with non-zero return_code\n metadata[\"status\"] = \"runtime_error\"\n result_status = -2\n else:\n # Other Failed status with run_result, classify as unknown failure\n logger.warning(f\"Unknown run_status '{metadata['run_status']}' or state within Failed API status.\")\n metadata[\"status\"] = \"unknown_failure\"\n result_status = -1 # Default to -1\n else:\n # Status is Failed but neither a clear compile error nor run_result exists\n logger.warning(\"API status Failed but cannot determine specific error type (compile/run).\")\n metadata[\"status\"] = \"unknown_failure_state\"\n result_status = -1 # Default to -1\n elif api_status == \"Success\":\n # Run completed successfully, now check the answer\n if run_result and metadata[\"run_status\"] == \"Finished\":\n actual_output = metadata[\"stdout\"] if metadata[\"stdout\"] is not None else \"\"\n # Note: Output might contain trailing newlines, need normalization\n if expected_output is None or str(actual_output).rstrip(\"\\n\") == str(expected_output).rstrip(\"\\n\"):\n result_status = True\n metadata[\"status\"] = \"success\"\n else:\n result_status = False\n metadata[\"status\"] = \"wrong_answer\"\n else:\n # Status is Success but run_result status is not Finished, this is unexpected\n metadata[\"status\"] = \"unexpected_success_state\"\n result_status = -1 # Classify as unknown error\n else:\n # API returned an unknown top-level status\n logger.warning(f\"Unknown API status received: {api_status}\")\n metadata[\"status\"] = f\"unknown_api_status_{api_status}\"\n result_status = -1 # Default to -1\n else: # api_response is None and no error_msg (Should not happen with current call_sandbox_api logic)\n metadata[\"status\"] = \"unknown_api_state\"\n result_status = -1\n logger.error(f\"Case {case_index}: Unknown API state (no response and no error message).\")\n return result_status, metadata\n\n\ndef check_correctness(\n sandbox_fusion_url: str,\n in_outs: Optional[dict],\n generation: str,\n timeout: int = DEFAULT_TIMEOUT,\n memory_limit_mb: int = 1024,\n language: str = \"python\",\n concurrent_semaphore: Optional[threading.Semaphore] = None,\n) -> tuple[list[Any], list[dict[str, Any]]]:\n \"\"\"\n Checks the correctness of code generation using the remote sandbox API,\n processing test cases concurrently.\n\n Args:\n sandbox_fusion_url: The URL of the sandbox fusion API.\n in_outs: Dictionary containing \"inputs\" and \"outputs\" lists.\n generation: The generated code string.\n timeout: Timeout for each test case (compile and run share this timeout).\n language: The programming language of the code.\n\n Returns:\n A tuple (results, metadata_list).\n results: A list containing the test result for each input/output pair\n (True/False/-1 api/sandbox err, -2 runtime err, -3 timeout, -4 compile err).\n Results are ordered corresponding to the inputs.\n metadata_list: A list containing metadata dictionaries for each test case,\n ordered corresponding to the inputs.\n \"\"\"\n logger.info(\"Starting correctness check for generation.\")\n\n if not in_outs or \"inputs\" not in in_outs or \"outputs\" not in in_outs:\n logger.warning(\"Invalid in_outs format provided.\")\n return [-1], [{\"error\": \"Invalid input/output data\"}]\n\n inputs = in_outs[\"inputs\"]\n expected_outputs = in_outs[\"outputs\"]\n fn_name = in_outs.get(\"fn_name\")\n num_cases = len(inputs)\n assert_cases = in_outs.get(\"assert_case\", [\"\"] * num_cases) # Default to empty strings if not provided\n results = [None] * num_cases # Initialize with placeholders\n metadata_list = [None] * num_cases # Initialize with placeholders\n\n if num_cases == 0:\n logger.warning(\"Empty inputs provided.\")\n return [], []\n\n if len(inputs) != len(expected_outputs):\n logger.warning(f\"Mismatch between number of inputs ({len(inputs)}) and outputs ({len(expected_outputs)}).\")\n # Return error based on the number of inputs provided\n return [-1] * num_cases, [{\"error\": \"Input/output count mismatch\", \"case_index\": i} for i in range(num_cases)]\n\n # If assert_cases is provided, it overrides inputs and outputs\n if len(assert_cases) != num_cases:\n logger.warning(\n f\"Mismatch between number of assert cases ({len(assert_cases)}) and inputs/outputs ({num_cases}).\"\n )\n return [-1] * num_cases, [{\"error\": \"Input/output count mismatch\", \"case_index\": i} for i in range(num_cases)]\n\n first_compile_error_index = -1\n\n # max_workers is limited by sandbox_fusion_max_concurrent from concurrent_semaphore\n with concurrent.futures.ThreadPoolExecutor(max_workers=max(32, os.cpu_count() * 5)) as executor:\n # Submit all tasks, passing the concurrent_semaphore to _process_single_case\n future_to_index = {\n executor.submit(\n _process_single_case,\n i,\n stdin_data,\n expected_outputs[i],\n sandbox_fusion_url,\n generation + \"\\n\\n\" + assert_cases[i], # Append assert case to generation\n timeout,\n memory_limit_mb,\n language,\n concurrent_semaphore,\n fn_name,\n ): i\n for i, stdin_data in enumerate(inputs)\n }\n\n # Process results as they complete\n for future in concurrent.futures.as_completed(future_to_index):\n index = future_to_index[future]\n try:\n result_status, metadata = future.result()\n results[index] = result_status\n metadata_list[index] = metadata\n\n # Check for compile error (-4)\n if result_status == -4:\n if first_compile_error_index == -1 or index < first_compile_error_index:\n first_compile_error_index = index\n # Optimization: could potentially cancel futures for index > first_compile_error_index\n # However, cancellation is not guaranteed. Post-processing is safer.\n\n except Exception as exc:\n logger.error(f\"Test case {index} generated an exception: {exc}\")\n traceback.print_exc()\n results[index] = -1 # Mark as API/internal error\n metadata_list[index] = {\n \"case_index\": index,\n \"input\": str(inputs[index]),\n \"expected_output\": str(expected_outputs[index]) if expected_outputs[index] else None,\n \"api_request_error\": f\"Internal execution error: {exc}\",\n \"status\": \"internal_error\",\n }\n\n # Post-processing for compile errors\n if first_compile_error_index != -1:\n logger.warning(\n f\"Compile error detected in case {first_compile_error_index}. Marking subsequent cases as compile errors.\"\n )\n for i in range(first_compile_error_index + 1, num_cases):\n # Only update if not already processed (though it should be None or have a result)\n if results[i] != -4: # Avoid overwriting if it somehow already got -4\n results[i] = -4\n # Update or create metadata for skipped cases due to compile error\n if metadata_list[i] is None: # If future failed before returning metadata\n metadata_list[i] = {\n \"case_index\": i,\n \"input\": str(inputs[i]),\n \"expected_output\": str(expected_outputs[i]) if expected_outputs[i] else None,\n \"api_request_error\": None,\n \"status\": \"compile_error_skipped\", # Indicate skipped due to prior compile error\n }\n else: # If future completed but result is overridden\n metadata_list[i][\"status\"] = \"compile_error_skipped\"\n\n logger.info(f\"Correctness check finished. Results: {results}\")\n return results, metadata_list\n"}118{"file_name": "verl__utils__reward_score__search_r1_like_qa_em.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n# Copyright 2023-2024 SGLang Team\n# Copyright 2025 Search-R1 Contributors\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n# Adapted from https://github.com/PeterGriffinJin/Search-R1/blob/main/verl/utils/reward_score/qa_em.py\n\nimport random\nimport re\nimport string\n\n\ndef normalize_answer(s):\n def remove_articles(text):\n return re.sub(r\"\\b(a|an|the)\\b\", \" \", text)\n\n def white_space_fix(text):\n return \" \".join(text.split())\n\n def remove_punc(text):\n exclude = set(string.punctuation)\n return \"\".join(ch for ch in text if ch not in exclude)\n\n def lower(text):\n return text.lower()\n\n return white_space_fix(remove_articles(remove_punc(lower(s))))\n\n\ndef em_check(prediction, golden_answers):\n if isinstance(golden_answers, str):\n golden_answers = [golden_answers]\n normalized_prediction = normalize_answer(prediction)\n score = 0\n for golden_answer in golden_answers:\n golden_answer = normalize_answer(golden_answer)\n if golden_answer == normalized_prediction:\n score = 1\n break\n return score\n\n\ndef subem_check(prediction, golden_answers):\n if isinstance(golden_answers, str):\n golden_answers = [golden_answers]\n normalized_prediction = normalize_answer(prediction)\n score = 0\n for golden_answer in golden_answers:\n golden_answer = normalize_answer(golden_answer)\n if golden_answer in normalized_prediction:\n score = 1\n break\n return score\n\n\ndef extract_solution(solution_str):\n \"\"\"Extract the equation from the solution string.\"\"\"\n # Remove everything before the first \"Assistant:\"\n # if \"Assistant:\" in solution_str:\n # solution_str = solution_str.split(\"Assistant:\", 1)[1]\n # elif \"<|im_start|>assistant\" in solution_str:\n # solution_str = solution_str.split(\"<|im_start|>assistant\", 1)[1]\n # else:\n # return None\n # solution_str = solution_str.split('\\n')[-1]\n\n answer_pattern = r\"<answer>(.*?)</answer>\"\n match = re.finditer(answer_pattern, solution_str, re.DOTALL)\n matches = list(match)\n\n # If there are 0 matches, return None\n if len(matches) < 1:\n return None\n\n # If there are 2 or more matches, return the last one\n return matches[-1].group(1).strip()\n\n\ndef count_answer_tags(text):\n opening_tags = text.count(\"<answer>\")\n closing_tags = text.count(\"</answer>\")\n\n return opening_tags, closing_tags\n\n\ndef compute_score(solution_str, ground_truth, method=\"strict\", format_score=0.0, score=1.0):\n \"\"\"The scoring function for exact match (EM).\n\n Args:\n solution_str: the solution text\n ground_truth: the ground truth\n method: the method to extract the solution, choices are 'strict' and 'flexible'\n format_score: the score for the format\n score: the score for the correct answer\n \"\"\"\n answer = extract_solution(solution_str=solution_str)\n open_count, close_count = count_answer_tags(solution_str)\n do_print = random.randint(1, 64) == 1\n\n if do_print:\n print(\"--------------------------------\")\n print(f\"Golden answers: {ground_truth['target']}\")\n if answer is not None:\n print(f\"Extracted answer is not None: {answer}\")\n else:\n print(\"Extracted answer: None!\")\n print(f\"Solution string: {solution_str}\")\n\n if answer is None:\n return 0\n else:\n if em_check(answer, ground_truth[\"target\"]):\n if open_count > 10 or close_count > 10: # prevent output a lot of </answer>\n score = score / 4\n return score\n return score\n else:\n return format_score\n\n\ndef compute_score_subem(solution_str, ground_truth, method=\"strict\", format_score=0.0, score=1.0):\n \"\"\"The scoring function for substring exact match (EM).\n\n Args:\n solution_str: the solution text\n ground_truth: the ground truth\n method: the method to extract the solution, choices are 'strict' and 'flexible'\n format_score: the score for the format\n score: the score for the correct answer\n \"\"\"\n answer = extract_solution(solution_str=solution_str)\n do_print = random.randint(1, 64) == 1\n\n if do_print:\n print(\"--------------------------------\")\n print(f\"Golden answers: {ground_truth['target']}\")\n print(f\"Extracted answer: {answer}\")\n print(f\"Solution string: {solution_str}\")\n\n if answer is None:\n return 0\n else:\n if subem_check(answer, ground_truth[\"target\"]):\n return score\n else:\n return format_score\n"}119{"file_name": "verl__utils__rollout_skip.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nfrom pathlib import Path\n\nfrom verl.protocol import DataProto\n\n\nclass RolloutSkip:\n \"\"\"\n RolloutSkip skips sequence generation during rollout by attempting to load previously dumped data.\n If no dumped data is found, it generates new sequences and saves them to disk.\n\n Args:\n config: The configuration object containing rollout settings.\n rollout_wg: The worker group that handles the rollout process.\n\n Note:\n When rollout.n or rollout.gen_batch_size differ from previous runs,\n new sequences will be generated and saved with different filenames.\n \"\"\"\n\n print_mark = \"[RolloutSkip()]\"\n\n def __init__(self, config, rollout_wg):\n self.rollout_config = config.actor_rollout_ref.rollout\n self.exp_name = config.data.get(\"experiment_name\", \"\")\n self.project_name = config.data.get(\"project_name\", \"\")\n\n self.n = int(self.rollout_config.get(\"n\", 0))\n self.gbs = int(config.data.get(\"gen_batch_size\", config.data.get(\"train_batch_size\", 0)))\n\n self.dumped_dir = Path(self.rollout_config.get(\"skip_dump_dir\", \"/tmp/verl/rollout_dump\"))\n self.dumped_dir.mkdir(parents=True, exist_ok=True)\n\n # Check if path is in Ray temporary directory\n if str(self.dumped_dir.absolute()).startswith(\"/tmp/ray/session\"):\n print(\n f\"\\033[33m{self.print_mark} Warning: \\nUsing dump path \",\n f\"'{self.dumped_dir.absolute()}' is not recommended \",\n \"as it's located in /tmp/ray/session*\\033[0m\",\n flush=True,\n )\n\n print(\n f\"{self.print_mark} Rollout skip dump path set to: \",\n f\"{self.dumped_dir.absolute()}\",\n flush=True,\n )\n\n self._rollout_wg = rollout_wg\n\n @property\n def curr_path_dump(self):\n return self.dumped_dir.joinpath(f\"{self.exp_name}_{self.project_name}_GBS{self.gbs}__N{self.n}\").absolute()\n\n def wrap_generate_sequences(self):\n try:\n self._rollout_wg.generate_sequences = wrap_generate_sequences(self, self._rollout_wg)\n print(\n f\"{self.print_mark} Successfully patched `actor_rollout_wg.generate_sequences()`\",\n flush=True,\n )\n except Exception as e:\n raise RuntimeError(\n \"{self.print_mark} Failed to patch `actor_rollout_wg.generate_sequences()`\",\n flush=True,\n ) from e\n\n def try_load(self):\n if not self.curr_path_dump.exists():\n print(\n f\"{self.print_mark} No data dump found at {self.curr_path_dump}.\",\n \"The trainer will generate and automatically dump the data for this first run.\",\n flush=True,\n )\n return None\n\n try:\n # * Load\n ret_batch = DataProto.load_from_disk(self.curr_path_dump)\n print(\n f\"\\033[32m{self.print_mark} Successfully load pre-generated data from {self.curr_path_dump}\\033[0m\",\n flush=True,\n )\n return ret_batch\n except Exception as e:\n print(\n f\"\\033[31m{self.print_mark} Failed to load pre-generated data from {self.curr_path_dump}\",\n f\"Error: {str(e)}\\033[0m\",\n flush=True,\n )\n return None\n\n def dump(self, outputs: DataProto):\n try:\n outputs.save_to_disk(self.curr_path_dump)\n print(\n f\"\\033[32m{self.print_mark} Successfully dump data in {self.curr_path_dump}\\033[0m\",\n flush=True,\n )\n except Exception as e:\n print(\n f\"\\033[31m{self.print_mark} Failed to dump data in {self.curr_path_dump}: {e}\\033[0m\",\n flush=True,\n )\n\n\ndef wrap_generate_sequences(rolloutskip: RolloutSkip, rollout_wg):\n generate_sequences = rollout_wg.generate_sequences\n\n def warp_fn(batch, **kwargs):\n gen_batch_output = rolloutskip.try_load()\n\n if gen_batch_output is None:\n # * 1. Generation\n gen_batch_output = generate_sequences(batch, **kwargs)\n # * 2. Dump\n rolloutskip.dump(gen_batch_output)\n return gen_batch_output\n\n return warp_fn\n"}120{"file_name": "verl__utils__rollout_trace.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport contextlib\nimport functools\nimport inspect\nimport os\nfrom contextvars import ContextVar\nfrom typing import Optional\n\nfrom pydantic import BaseModel\n\nfrom verl.utils.ray_utils import get_event_loop\n\n_trace_enabled: ContextVar[bool] = ContextVar(\"_trace_enabled\", default=True)\n\n\nclass RolloutTraceConfig:\n \"\"\"Configuration for rollout tracing with various backends.\n\n Singleton configuration class for managing rollout trace settings across different\n tracing backends like Weave and MLflow.\n\n Args:\n backend (Optional[str]): Tracing backend to use ('weave', 'mlflow', or None).\n client (Optional[object]): Client instance for the selected backend.\n token2text (bool): Whether to convert tokens to text in traces. Defaults to False.\n project_name (str): Name of the project for tracing.\n experiment_name (str): Name of the experiment for tracing.\n max_samples_per_step_per_worker (Optional[int]): Maximum number of unique samples to trace\n per worker per step. If None, all samples are traced. If set, each worker will randomly\n select up to this many unique samples to trace (including all their rollouts for GRPO).\n Total traces = max_samples_per_step_per_worker * num_workers * n_rollouts_per_sample.\n \"\"\"\n\n _instance: Optional[\"RolloutTraceConfig\"] = None\n backend: Optional[str] = None\n client: Optional[object] = None\n token2text: bool = False\n _initialized: bool = False\n project_name: str = None\n experiment_name: str = None\n max_samples_per_step_per_worker: Optional[int] = None\n\n def __new__(cls, *args, **kwargs):\n if cls._instance is None:\n cls._instance = super().__new__(cls)\n cls._instance._initialized = False\n return cls._instance\n\n @classmethod\n def get_instance(cls) -> \"RolloutTraceConfig\":\n if cls._instance is None:\n cls._instance = cls()\n return cls._instance\n\n @classmethod\n def init(\n cls,\n project_name: str,\n experiment_name: str,\n backend: str,\n token2text: bool = False,\n max_samples_per_step_per_worker: Optional[int] = None,\n ):\n config = cls.get_instance()\n if config._initialized:\n return\n\n config.backend = backend\n config.token2text = token2text\n config.project_name = project_name\n config.experiment_name = experiment_name\n config.max_samples_per_step_per_worker = max_samples_per_step_per_worker\n\n if backend == \"weave\":\n import weave\n\n config.client = weave.init(project_name)\n elif backend == \"mlflow\":\n import mlflow\n\n mlflow.config.enable_async_logging()\n config.client = mlflow\n\n MLFLOW_TRACKING_URI = os.environ.get(\"MLFLOW_TRACKING_URI\", \"sqlite:////tmp/mlruns.db\")\n mlflow.set_tracking_uri(MLFLOW_TRACKING_URI)\n\n mlflow.set_experiment(project_name)\n else:\n config.client = None\n\n config._initialized = True\n\n @classmethod\n def get_backend(cls) -> Optional[str]:\n return cls.get_instance().backend\n\n @classmethod\n def get_client(cls) -> Optional[object]:\n return cls.get_instance().client\n\n @classmethod\n def enable_token2text(cls) -> Optional[bool]:\n return cls.get_instance().token2text\n\n @classmethod\n def reset(cls):\n cls._instance = None\n\n\n@contextlib.contextmanager\ndef rollout_trace_attr(\n sample_index=None, step=None, rollout_n=None, name=\"rollout_trace\", validate=False, trace: bool = True\n):\n \"\"\"A context manager to add attributes to a trace for the configured backend.\n\n Args:\n sample_index: Sample index for the trace.\n step: Training step number.\n rollout_n: Rollout number (for GRPO with multiple rollouts per sample).\n name: Name for the trace span (used by mlflow backend).\n validate: Whether this is a validation run.\n trace: If False, disables tracing for the duration of the context.\n \"\"\"\n backend = RolloutTraceConfig.get_backend()\n\n should_skip = backend is not None and not trace\n\n if should_skip:\n token = _trace_enabled.set(False)\n try:\n yield\n finally:\n _trace_enabled.reset(token)\n return\n\n # Build attributes for the trace\n attributes = {}\n if backend:\n if sample_index is not None:\n attributes[\"sample_index\"] = sample_index\n if step is not None:\n attributes[\"step\"] = step\n if rollout_n is not None:\n attributes[\"rollout_n\"] = rollout_n\n attributes[\"validate\"] = validate\n attributes[\"experiment_name\"] = RolloutTraceConfig.get_instance().experiment_name\n\n if not attributes or backend is None:\n yield\n return\n\n if backend == \"weave\":\n import weave\n\n with weave.attributes(attributes):\n yield\n elif backend == \"mlflow\":\n import mlflow\n\n with mlflow.start_span(name=name) as span:\n trace_id = span.trace_id\n for key, value in attributes.items():\n mlflow.set_trace_tag(trace_id, str(key), str(value))\n yield\n else:\n yield\n\n\ndef rollout_trace_op(func):\n @functools.wraps(func)\n async def async_wrapper(self, *args, **kwargs):\n if not _trace_enabled.get():\n return await func(self, *args, **kwargs)\n\n backend = RolloutTraceConfig.get_backend()\n enable_token2text = RolloutTraceConfig.enable_token2text()\n if backend is None:\n return await func(self, *args, **kwargs)\n\n sig = inspect.signature(func)\n bound_args = sig.bind(self, *args, **kwargs)\n bound_args.apply_defaults()\n inputs = dict(bound_args.arguments)\n del inputs[\"self\"]\n\n async def add_token2text(self, result):\n if hasattr(result, \"prompt_ids\") and hasattr(self, \"tokenizer\") and hasattr(self.tokenizer, \"decode\"):\n # Use model_dump() for Pydantic models to get a proper copy,\n # otherwise vars() returns a reference to internal __dict__ which\n # can cause serialization issues with MLflow\n if isinstance(result, BaseModel):\n _result = result.model_dump()\n else:\n _result = dict(vars(result))\n loop = get_event_loop()\n if hasattr(result, \"prompt_ids\"):\n prompt_text = await loop.run_in_executor(None, self.tokenizer.decode, result.prompt_ids)\n _result[\"prompt_text\"] = prompt_text\n\n if hasattr(result, \"response_ids\"):\n response_text = await loop.run_in_executor(None, self.tokenizer.decode, result.response_ids)\n _result[\"response_text\"] = response_text\n return _result\n return result\n\n if backend == \"weave\":\n tracer = RolloutTraceConfig.get_client()\n from weave.trace.context import call_context\n\n cur_attributes = {**call_context.call_attributes.get()}\n call = tracer.create_call(op=func.__qualname__, inputs=inputs, attributes=cur_attributes)\n try:\n result = await func(self, *args, **kwargs)\n\n if enable_token2text:\n _result = await add_token2text(self, result)\n tracer.finish_call(call, output=_result)\n else:\n tracer.finish_call(call, output=result)\n\n return result\n\n except Exception as e:\n tracer.finish_call(call, exception=e)\n raise e\n elif backend == \"mlflow\":\n import mlflow\n\n with mlflow.start_span(name=func.__qualname__) as span:\n span.set_inputs(inputs)\n result = await func(self, *args, **kwargs)\n if enable_token2text:\n _result = await add_token2text(self, result)\n span.set_outputs(_result)\n else:\n span.set_outputs(result)\n\n return result\n\n else:\n return await func(self, *args, **kwargs)\n\n @functools.wraps(func)\n def wrapper(self, *args, **kwargs):\n if not _trace_enabled.get():\n return func(self, *args, **kwargs)\n\n backend = RolloutTraceConfig.get_backend()\n if backend is None:\n return func(self, *args, **kwargs)\n\n sig = inspect.signature(func)\n bound_args = sig.bind(self, *args, **kwargs)\n bound_args.apply_defaults()\n inputs = dict(bound_args.arguments)\n del inputs[\"self\"]\n\n if backend == \"weave\":\n tracer = RolloutTraceConfig.get_client()\n from weave.trace.context import call_context\n\n cur_attributes = {**call_context.call_attributes.get()}\n call = tracer.create_call(op=func.__qualname__, inputs=inputs, attributes=cur_attributes)\n try:\n result = func(self, *args, **kwargs)\n tracer.finish_call(call, output=result)\n return result\n except Exception as e:\n tracer.finish_call(call, exception=e)\n raise e\n elif backend == \"mlflow\":\n import mlflow\n\n return mlflow.trace(func)(self, *args, **kwargs)\n else:\n return func(self, *args, **kwargs)\n\n return async_wrapper if inspect.iscoroutinefunction(func) else wrapper\n"}121{"file_name": "verl__utils__seqlen_balancing.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport copy\nimport heapq\nfrom itertools import chain\n\nimport torch\nfrom torch import distributed as dist\n\nfrom verl.protocol import DataProto\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.device import get_device_name\n\n\ndef calculate_workload(seqlen_list: torch.Tensor) -> torch.Tensor:\n \"\"\"Calculate approximate computational workload for transformer attention.\n\n Estimates FLOPs for dense transformer blocks based on sequence length using\n the formula: FLOPs ≈ 12 * hidden_size² * seqlen + 2 * hidden_size * seqlen²\n\n The constants are calibrated for a 7B model (hidden_size=4096), yielding:\n workload ∝ 24576 * seqlen + seqlen²\n\n Args:\n seqlen_list: Sequence lengths as a tensor.\n\n Returns:\n torch.Tensor: Estimated workload values proportional to actual FLOPs.\n\n Note:\n The returned values are relative workloads, not actual FLOP counts.\n Useful for balancing computation across data parallel ranks.\n \"\"\"\n return 24576 * seqlen_list + seqlen_list**2\n\n\ndef karmarkar_karp(seqlen_list: list[int], k_partitions: int, equal_size: bool) -> list[list[int]]:\n \"\"\"Partition items into k groups using the Karmarkar-Karp differencing method.\n\n Implements the Largest Differencing Method (LDM) algorithm for balanced\n multi-way number partitioning. This heuristic produces near-optimal partitions\n by iteratively combining the sets with the largest difference.\n\n Args:\n seqlen_list: Values to partition (typically sequence lengths or workloads).\n k_partitions: Number of partitions to create.\n equal_size: If True, each partition will have exactly len(seqlen_list) / k_partitions\n items. If False, partitions may have different sizes.\n\n Returns:\n list[list[int]]: List of k partitions, each containing indices into seqlen_list.\n\n See Also:\n https://en.wikipedia.org/wiki/Largest_differencing_method\n\n Note:\n When equal_size=True, len(seqlen_list) must be divisible by k_partitions.\n \"\"\"\n\n # see: https://en.wikipedia.org/wiki/Largest_differencing_method\n class Set:\n def __init__(self) -> None:\n self.sum = 0\n self.items = []\n\n def add(self, idx: int, val: int):\n self.items.append((idx, val))\n self.sum += val\n\n def merge(self, other):\n for idx, val in other.items:\n self.items.append((idx, val))\n self.sum += val\n\n def __lt__(self, other):\n if self.sum != other.sum:\n return self.sum < other.sum\n if len(self.items) != len(other.items):\n return len(self.items) < len(other.items)\n return self.items < other.items\n\n class State:\n def __init__(self, items: list[tuple[int, int]], k: int) -> None:\n self.k = k\n # sets should always be decreasing order\n self.sets = [Set() for _ in range(k)]\n assert len(items) in [1, k], f\"{len(items)} not in [1, {k}]\"\n for i, (idx, seqlen) in enumerate(items):\n self.sets[i].add(idx=idx, val=seqlen)\n self.sets = sorted(self.sets, reverse=True)\n\n def get_partitions(self):\n partitions = []\n for i in range(len(self.sets)):\n cur_partition = []\n for idx, _ in self.sets[i].items:\n cur_partition.append(idx)\n partitions.append(cur_partition)\n return partitions\n\n def merge(self, other):\n for i in range(self.k):\n self.sets[i].merge(other.sets[self.k - 1 - i])\n self.sets = sorted(self.sets, reverse=True)\n\n @property\n def spread(self) -> int:\n return self.sets[0].sum - self.sets[-1].sum\n\n def __lt__(self, other):\n # least heap, let the state with largest spread to be popped first,\n # if the spread is the same, let the state who has the largest set\n # to be popped first.\n if self.spread != other.spread:\n return self.spread > other.spread\n return self.sets[0] > other.sets[0]\n\n def __repr__(self) -> str:\n repr_str = \"[\"\n for i in range(self.k):\n if i > 0:\n repr_str += \",\"\n repr_str += \"{\"\n for j, (_, seqlen) in enumerate(self.sets[i].items):\n if j > 0:\n repr_str += \",\"\n repr_str += str(seqlen)\n repr_str += \"}\"\n repr_str += \"]\"\n return repr_str\n\n sorted_seqlen_list = sorted([(seqlen, i) for i, seqlen in enumerate(seqlen_list)])\n states_pq = []\n if equal_size:\n assert len(seqlen_list) % k_partitions == 0, f\"{len(seqlen_list)} % {k_partitions} != 0\"\n for offset in range(0, len(sorted_seqlen_list), k_partitions):\n items = []\n for i in range(k_partitions):\n seqlen, idx = sorted_seqlen_list[offset + i]\n items.append((idx, seqlen))\n heapq.heappush(states_pq, State(items=items, k=k_partitions))\n else:\n for seqlen, idx in sorted_seqlen_list:\n heapq.heappush(states_pq, State(items=[(idx, seqlen)], k=k_partitions))\n\n while len(states_pq) > 1:\n state0 = heapq.heappop(states_pq)\n state1 = heapq.heappop(states_pq)\n # merge states\n state0.merge(state1)\n heapq.heappush(states_pq, state0)\n\n final_state = states_pq[0]\n partitions = final_state.get_partitions()\n if equal_size:\n for i, partition in enumerate(partitions):\n assert len(partition) * k_partitions == len(seqlen_list), (\n f\"{len(partition)} * {k_partitions} != {len(seqlen_list)}\"\n )\n return partitions\n\n\ndef greedy_partition(seqlen_list: list[int], k_partitions: int, equal_size: bool) -> list[list[int]]:\n \"\"\"Partition items into k groups using a greedy assignment strategy.\n\n Assigns each item to the partition with the smallest current sum, iterating\n through items in order. Simpler but typically less optimal than Karmarkar-Karp.\n\n Args:\n seqlen_list: Values to partition (typically sequence lengths or workloads).\n k_partitions: Number of partitions to create.\n equal_size: If True, adds a bias to ensure equal partition sizes.\n Requires len(seqlen_list) to be divisible by k_partitions.\n\n Returns:\n list[list[int]]: List of k partitions, each containing indices into seqlen_list.\n\n Note:\n When equal_size=True, a large bias is added to encourage equal distribution\n of items before considering the actual values.\n \"\"\"\n bias = sum(seqlen_list) + 1 if equal_size else 0\n sorted_seqlen = [(seqlen + bias, i) for i, seqlen in enumerate(seqlen_list)]\n partitions = [[] for _ in range(k_partitions)]\n partition_sums = [0 for _ in range(k_partitions)]\n for seqlen, i in sorted_seqlen:\n min_idx = None\n for j in range(k_partitions):\n if min_idx is None or partition_sums[j] < partition_sums[min_idx]:\n min_idx = j\n partitions[min_idx].append(i)\n partition_sums[min_idx] += seqlen\n if equal_size:\n for i, partition in enumerate(partitions):\n assert len(partition) * k_partitions == len(seqlen_list), (\n f\"{len(partition)} * {k_partitions} != {len(seqlen_list)}\"\n )\n return partitions\n\n\ndef get_seqlen_balanced_partitions(seqlen_list: list[int], k_partitions: int, equal_size: bool):\n \"\"\"\n Calculates partitions of indices from seqlen_list such that the sum of sequence lengths\n in each partition is balanced. Uses the Karmarkar-Karp differencing method.\n\n This is useful for balancing workload across devices or batches, especially when\n dealing with variable sequence lengths.\n\n Args:\n seqlen_list (List[int]): A list of sequence lengths for each item.\n k_partitions (int): The desired number of partitions.\n equal_size (bool): If True, ensures that each partition has the same number of items.\n Requires len(seqlen_list) to be divisible by k_partitions.\n If False, partitions can have varying numbers of items, focusing\n only on balancing the sum of sequence lengths.\n\n Returns:\n List[List[int]]: A list containing k_partitions lists. Each inner list contains the\n original indices of the items assigned to that partition. The indices\n within each partition list are sorted.\n\n Raises:\n AssertionError: If len(seqlen_list) < k_partitions.\n AssertionError: If equal_size is True and len(seqlen_list) is not divisible by k_partitions.\n AssertionError: If any resulting partition is empty.\n \"\"\"\n assert len(seqlen_list) >= k_partitions, f\"number of items:[{len(seqlen_list)}] < k_partitions:[{k_partitions}]\"\n\n def _check_and_sort_partitions(partitions):\n assert len(partitions) == k_partitions, f\"{len(partitions)} != {k_partitions}\"\n seen_idx = set()\n sorted_partitions = [None] * k_partitions\n for i, partition in enumerate(partitions):\n assert len(partition) > 0, f\"the {i}-th partition is empty\"\n for idx in partition:\n seen_idx.add(idx)\n sorted_partitions[i] = sorted(partition)\n assert seen_idx == set(range(len(seqlen_list)))\n return sorted_partitions\n\n partitions = karmarkar_karp(seqlen_list=seqlen_list, k_partitions=k_partitions, equal_size=equal_size)\n return _check_and_sort_partitions(partitions)\n\n\ndef log_seqlen_unbalance(seqlen_list: list[int], partitions: list[list[int]], prefix):\n \"\"\"\n Calculate and log metrics related to sequence length imbalance before and after partitioning.\n\n Args:\n seqlen_list (List[int]): A list of sequence lengths for each item.\n partitions (List[List[int]]): A list of partitions, where each inner list contains indices\n from seqlen_list assigned to that partition.\n prefix (str): A prefix to be added to each metric key in the returned dictionary.\n\n Returns:\n dict: A dictionary containing metrics related to sequence length imbalance.\n \"\"\"\n # Get the number of partitions\n k_partition = len(partitions)\n # assert len(seqlen_list) % k_partition == 0\n batch_size = len(seqlen_list) // k_partition\n min_sum_seqlen = None\n max_sum_seqlen = None\n total_sum_seqlen = 0\n\n # Iterate over each batch of sequence lengths\n for offset in range(0, len(seqlen_list), batch_size):\n cur_sum_seqlen = sum(seqlen_list[offset : offset + batch_size])\n if min_sum_seqlen is None or cur_sum_seqlen < min_sum_seqlen:\n min_sum_seqlen = cur_sum_seqlen\n if max_sum_seqlen is None or cur_sum_seqlen > max_sum_seqlen:\n max_sum_seqlen = cur_sum_seqlen\n total_sum_seqlen += cur_sum_seqlen\n\n balanced_sum_seqlen_list = []\n for partition in partitions:\n cur_sum_seqlen_balanced = sum([seqlen_list[i] for i in partition])\n balanced_sum_seqlen_list.append(cur_sum_seqlen_balanced)\n # print(\"balanced_sum_seqlen_list: \", balanced_sum_seqlen_list)\n min_sum_seqlen_balanced = min(balanced_sum_seqlen_list)\n max_sum_seqlen_balanced = max(balanced_sum_seqlen_list)\n\n return {\n f\"{prefix}/min\": min_sum_seqlen,\n f\"{prefix}/max\": max_sum_seqlen,\n f\"{prefix}/minmax_diff\": max_sum_seqlen - min_sum_seqlen,\n f\"{prefix}/balanced_min\": min_sum_seqlen_balanced,\n f\"{prefix}/balanced_max\": max_sum_seqlen_balanced,\n f\"{prefix}/mean\": total_sum_seqlen / len(partitions),\n }\n\n\ndef ceildiv(a: int, b: int) -> int:\n \"\"\"Compute ceiling division of a by b.\n\n Returns the smallest integer greater than or equal to a/b.\n Uses the identity: ceil(a/b) = floor((a + b - 1) / b) = -(-a // b)\n\n Args:\n a: Dividend (numerator).\n b: Divisor (denominator), must be non-zero.\n\n Returns:\n int: Ceiling of a divided by b.\n\n Example:\n >>> ceildiv(7, 3) # ceil(7/3) = ceil(2.33) = 3\n 3\n >>> ceildiv(6, 3) # ceil(6/3) = ceil(2.0) = 2\n 2\n \"\"\"\n return -(a // -b)\n\n\ndef roundup_divisible(a: int, b: int) -> int:\n \"\"\"Round up a to the nearest multiple of b.\n\n Returns the smallest multiple of b that is >= a.\n\n Args:\n a: Value to round up.\n b: Divisor to round to (must be positive).\n\n Returns:\n int: Smallest multiple of b that is >= a.\n\n Example:\n >>> roundup_divisible(7, 4) # nearest multiple of 4 >= 7 is 8\n 8\n >>> roundup_divisible(8, 4) # 8 is already a multiple of 4\n 8\n \"\"\"\n return ((a + b - 1) // b) * b\n\n\ndef rearrange_micro_batches(\n batch,\n max_token_len,\n dp_group=None,\n num_batches_divided_by=None,\n same_micro_num_in_dp=True,\n min_num_micro_batch=None,\n use_dynamic_bsz_balance=True,\n):\n \"\"\"\n Split a batch into micro-batches by total token count, with optional DP sync and padding.\n\n Args:\n batch (TensorDict): must include \"attention_mask\" (B*S); other fields are sliced similarly.\n max_token_len (int): max sum of attention_mask per micro-batch.\n dp_group (optional): torch.distributed group for data-parallel sync.\n num_batches_divided_by (optional): virtual pipeline parallel size, for megatron.\n same_micro_num_in_dp (bool): if True and dp_group set, pad all ranks to the same count.\n min_num_micro_batch (int, optional): force at least this many splits (pads empty ones).\n use_dynamic_bsz_balance (bool, optional): balance the computational workload between micro-batches\n\n Returns:\n List[TensorDict]: the micro-batches.\n List[List[int]]: index lists mapping each micro-batch back to original positions.\n \"\"\"\n # this is per local micro_bsz\n input_ids = batch[\"input_ids\"]\n if input_ids.is_nested:\n seq_len_effective: torch.Tensor = input_ids.offsets().diff()\n max_seq_len = max(seq_len_effective)\n else:\n max_seq_len = batch[\"attention_mask\"].shape[-1]\n seq_len_effective: torch.Tensor = batch[\"attention_mask\"].sum(dim=1)\n\n assert max_token_len >= max_seq_len, (\n f\"max_token_len must be greater than the sequence length. Got {max_token_len=} and {max_seq_len=}\"\n )\n total_seqlen = seq_len_effective.sum().item()\n # NOTE: num_microbatches <= batch_size, so take the min of this two.\n num_micro_batches = min(len(seq_len_effective), ceildiv(total_seqlen, max_token_len))\n if min_num_micro_batch is not None:\n # used to support pp\n num_micro_batches = max(min_num_micro_batch, num_micro_batches)\n if dist.is_initialized() and same_micro_num_in_dp:\n num_micro_batches = torch.tensor([num_micro_batches], device=get_device_name())\n dist.all_reduce(num_micro_batches, op=dist.ReduceOp.MAX, group=dp_group)\n num_micro_batches = num_micro_batches.cpu().item()\n if num_batches_divided_by is not None:\n num_micro_batches = roundup_divisible(num_micro_batches, num_batches_divided_by)\n\n assert num_micro_batches <= len(seq_len_effective)\n\n # upcast to int64 to avoid potential overflow im `calculate_workload` computation.\n seq_len_effective = seq_len_effective.long()\n # note that seq_len_effective is a GPU tensor. We need to make it a list to avoid D2H!\n workloads = calculate_workload(seq_len_effective).cpu().tolist()\n micro_bsz_idx = get_seqlen_balanced_partitions(workloads, num_micro_batches, equal_size=False)\n\n if use_dynamic_bsz_balance:\n # Use the sum of squared sequence lengths to approximate attention computation workload\n micro_bsz_idx.sort(\n key=lambda partition: (\n sum(workloads[idx] for idx in partition),\n partition[0] if partition else 0,\n ),\n reverse=True,\n )\n # Place smaller micro-batches at both ends to reduce the bubbles exposed during the warm-up and cool-down.\n micro_bsz_idx = micro_bsz_idx[::2][::-1] + micro_bsz_idx[1::2]\n\n micro_batches = []\n\n for partition in micro_bsz_idx:\n curr_micro_batch = tu.index_select_tensor_dict(batch, partition)\n micro_batches.append(curr_micro_batch)\n\n return micro_batches, micro_bsz_idx\n\n\ndef get_reverse_idx(idx_map):\n \"\"\"\n Build the inverse of an index mapping.\n\n Args:\n idx_map (Sequence[int]): Sequence where idx_map[i] = j.\n\n Returns:\n List[int]: Inverse mapping list such that output[j] = i for each i.\n \"\"\"\n reverse_idx_map = copy.deepcopy(idx_map)\n\n for i, idx in enumerate(idx_map):\n reverse_idx_map[idx] = i\n\n return reverse_idx_map\n\n\ndef prepare_dynamic_batch(\n data: DataProto,\n max_token_len: int,\n dp_group=None,\n num_batches_divided_by=None,\n same_micro_num_in_dp=True,\n min_num_micro_batch=None,\n use_dynamic_bsz_balance=True,\n) -> tuple[list[DataProto], list[list[int]]]:\n \"\"\"\n Prepare a batch for dynamic batching.\n\n Args:\n data (DataProto): The input data.\n max_token_len (int): The maximum token length for dynamic batching.\n\n Returns:\n Tuple[List[DataProto], List[List[int]]]: A tuple containing a list of DataProto objects\n and a list of index lists.\n \"\"\"\n batch, batch_idx_list = rearrange_micro_batches(\n data.batch,\n max_token_len=max_token_len,\n dp_group=dp_group,\n num_batches_divided_by=num_batches_divided_by,\n same_micro_num_in_dp=same_micro_num_in_dp,\n min_num_micro_batch=min_num_micro_batch,\n use_dynamic_bsz_balance=use_dynamic_bsz_balance,\n )\n micro_batches = []\n for i, batch_idx in enumerate(batch_idx_list):\n tensors = dict(batch[i])\n non_tensors = {key: value[batch_idx] for key, value in data.non_tensor_batch.items()}\n meta_info = copy.deepcopy(data.meta_info)\n micro_batches.append(DataProto.from_dict(tensors, non_tensors, meta_info=meta_info))\n\n return micro_batches, batch_idx_list\n\n\ndef restore_dynamic_batch(data: torch.Tensor, batch_idx_list: list[list[int]]) -> torch.Tensor:\n \"\"\"\n Restore a batch from dynamic batching.\n\n Args:\n data (torch.Tensor): The input data.\n batch_idx_list (List[List[int]]): The list of index lists.\n\n Returns:\n torch.Tensor: The restored data.\n \"\"\"\n indices = list(chain.from_iterable(batch_idx_list))\n batch_size = data.shape[0]\n assert len(indices) == batch_size, f\"{len(indices)} vs. {batch_size}\"\n revert_indices = torch.tensor(get_reverse_idx(indices), dtype=torch.long)\n\n if data.is_nested:\n data_lst = data.unbind()\n tensors = [data_lst[i] for i in revert_indices]\n reverted_data = torch.nested.as_nested_tensor(tensors, layout=torch.jagged)\n else:\n reverted_data = data[revert_indices]\n\n return reverted_data\n\n\ndef get_group_balanced_partitions(\n seqlen_list: list[int],\n uid_list: list,\n k_partitions: int,\n) -> list[list[int]]:\n \"\"\"\n Partition samples into k groups while keeping samples with the same uid together.\n\n Args:\n seqlen_list: List of sequence lengths for each sample.\n uid_list: List of uids identifying which samples share the same prefix.\n Samples with the same uid will be kept together.\n k_partitions: Number of partitions (typically world_size).\n\n Returns:\n List of k lists, each containing sample indices assigned to that partition.\n Samples with the same uid are guaranteed to be in the same partition.\n \"\"\"\n assert len(seqlen_list) == len(uid_list), \"seqlen_list and uid_list must have same length\"\n\n # Build groups: each group contains indices of samples with the same uid\n # Assumes samples with same uid are contiguous\n groups = [] # List of (group_indices, group_total_seqlen)\n current_uid = None\n current_indices = []\n current_seqlen = 0\n\n for i, (seqlen, uid) in enumerate(zip(seqlen_list, uid_list, strict=False)):\n if uid != current_uid:\n if current_indices:\n groups.append((current_indices, current_seqlen))\n current_uid = uid\n current_indices = [i]\n current_seqlen = seqlen\n else:\n current_indices.append(i)\n current_seqlen += seqlen\n\n # Don't forget the last group\n if current_indices:\n groups.append((current_indices, current_seqlen))\n\n num_groups = len(groups)\n assert num_groups >= k_partitions, (\n f\"Number of uid groups ({num_groups}) must be >= k_partitions ({k_partitions}). \"\n f\"Consider reducing world_size or increasing batch_size.\"\n )\n\n # Calculate workload for each group (as integers for partitioning)\n group_workloads = []\n for indices, total_seqlen in groups:\n # Use sum of individual workloads for more accurate estimation\n workload = sum(int(calculate_workload(torch.tensor([seqlen_list[i]])).item()) for i in indices)\n group_workloads.append(workload)\n\n # Use Karmarkar-Karp to partition groups\n # equal_size=True ensures each partition gets the same number of groups,\n # which is required when each group has the same number of samples (rollout.n)\n group_partitions = get_seqlen_balanced_partitions(\n seqlen_list=group_workloads,\n k_partitions=k_partitions,\n equal_size=True,\n )\n\n # Convert group partitions to sample partitions\n sample_partitions = []\n for group_partition in group_partitions:\n sample_indices = []\n for group_idx in group_partition:\n sample_indices.extend(groups[group_idx][0])\n sample_partitions.append(sorted(sample_indices))\n\n return sample_partitions\n"}122{"file_name": "verl__utils__sglang__sglang_fp8_utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nimport os\n\nimport torch\n\nfrom verl.utils.kernel.fp8_kernel import scaled_fp8_blockwise\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"INFO\"))\n\n\ndef should_quantize_param(param_name: str) -> bool:\n \"\"\"Determine whether to quantize to FP8 based on parameter name\n\n Quantization rules:\n - Must end with .weight (exclude bias)\n - Exclude embedding layers\n - Exclude normalization layers\n - Exclude output layer (lm_head)\n \"\"\"\n # Must be a weight parameter\n if not param_name.endswith(\".weight\"):\n return False\n\n # Layer types to exclude\n exclude_patterns = [\n \"embed_tokens\", # Embedding layer\n \"lm_head\", # Output layer\n \"layernorm\", # LayerNorm\n \"norm\", # Various Norm layers\n \"ln_\", # LayerNorm variants\n \"embeddings\", # Embeddings\n \"mlp.gate.weight\", # MoE router\n ]\n\n # Check if matches exclude patterns\n param_lower = param_name.lower()\n for pattern in exclude_patterns:\n if pattern in param_lower:\n return False\n\n # Layer types to include (Linear layers)\n include_patterns = [\n \"q_proj\", # Query projection\n \"k_proj\", # Key projection\n \"v_proj\", # Value projection\n \"o_proj\", # Output projection\n \"gate_proj\", # Gate projection (for MLP)\n \"up_proj\", # Up projection (for MLP)\n \"down_proj\", # Down projection (for MLP)\n \"fc1\", # Fully connected 1\n \"fc2\", # Fully connected 2\n \"mlp\", # MLP layers\n ]\n\n # Check if matches include patterns\n for pattern in include_patterns:\n if pattern in param_lower:\n logger.debug(f\"Will quantize FP8: {param_name}\")\n return True\n\n # Do not quantize by default\n logger.debug(f\"Skip quantization: {param_name}\")\n return False\n\n\ndef quant_weights_by_name(weights, quant_config, dtype=torch.bfloat16):\n \"\"\"FP8 quantization based on parameter name using a memory-efficient generator.\n\n\n Args:\n weights: Generator or iterable of (name, tensor) pairs\n quant_config: Quantization configuration\n dtype: Data type for intermediate computation\n\n Yields:\n Tuples of (name, tensor) for each weight and its scale\n \"\"\"\n if isinstance(quant_config, dict):\n weight_block_size = quant_config.get(\"weight_block_size\")\n else:\n weight_block_size = getattr(quant_config, \"weight_block_size\", None)\n\n if weight_block_size is None:\n raise ValueError(\"weight_block_size not found in quant_config\")\n\n for k, v in weights:\n # Check if quantization is needed\n if not should_quantize_param(k):\n yield (k, v)\n continue\n\n # Quantize to FP8\n try:\n if torch.distributed.get_rank() == 0:\n logger.debug(f\"Quantizing to FP8 blockwise: {k}\")\n\n param_lp, param_scale = scaled_fp8_blockwise(\n v.to(dtype),\n weight_block_size=weight_block_size,\n )\n param_scale = param_scale.squeeze(-1)\n\n # Yield the quantized weight and scale\n yield (k, param_lp)\n yield (k + \"_scale_inv\", param_scale)\n\n # Explicitly delete to help GC\n del param_lp, param_scale\n\n except Exception as e:\n logger.error(f\"Failed to quantize {k}: {e}\")\n # If quantization fails, use original weights\n yield (k, v)\n"}123{"file_name": "verl__utils__tensordict_utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nfrom typing import Any, Iterable\n\nimport torch\nfrom tensordict import TensorDict\nfrom tensordict.tensorclass import NonTensorData, NonTensorStack\n\n\ndef assign_non_tensor_data(tensor_dict: TensorDict, key, val):\n \"\"\"Assign a single non-tensor value to a TensorDict.\n\n Wraps the value in NonTensorData so it can be stored alongside tensors\n in the TensorDict. Use this for scalar metadata or simple non-tensor values.\n\n Args:\n tensor_dict: The TensorDict to assign to.\n key: The key under which to store the value.\n val: Any non-tensor value to store (e.g., string, int, dict).\n\n Raises:\n AssertionError: If tensor_dict is not a TensorDict.\n\n Example:\n >>> td = TensorDict({\"obs\": torch.randn(3, 4)}, batch_size=[3])\n >>> assign_non_tensor_data(td, \"experiment_name\", \"run_001\")\n \"\"\"\n assert isinstance(tensor_dict, TensorDict), \"input dict must be a TensorDict\"\n tensor_dict[key] = NonTensorData(val)\n\n\ndef assign_non_tensor_stack(tensor_dict: TensorDict, key, val: list):\n \"\"\"Assign a list with potentially nested structures (lists, dicts, etc.) to TensorDict.\n\n This function handles complex nested data structures like:\n - Lists of lists: [[], [0.5, 0.8], [0.9]]\n - Lists of dicts: [{\"acc\": 1.0}, {\"acc\": 0.0}]\n - Lists of lists of dicts: [[{\"content\": \"...\", \"role\": \"user\"}]]\n\n These structures are wrapped in NonTensorStack so TensorDict can handle them correctly.\n\n Args:\n tensor_dict: The TensorDict to assign to\n key: The key to assign the value under\n val: A list containing potentially nested structures\n\n Example:\n >>> td = TensorDict({}, batch_size=[])\n >>> turn_scores = [[], [0.5, 0.8], [0.9]]\n >>> assign_non_tensor_stack(td, \"turn_scores\", turn_scores)\n >>> # Now td[\"turn_scores\"] contains the nested data\n \"\"\"\n # Convert list to NonTensorStack to handle nested structures\n # This wraps each item in NonTensorData to preserve complex objects\n # TODO(petersh6): can convert back to val directly if we are not accessing .data from the NonTensorStack\n assert isinstance(tensor_dict, TensorDict), \"input dict must be a TensorDict\"\n tensor_dict[key] = NonTensorStack.from_list([NonTensorData(item) for item in val])\n\n\ndef assign_non_tensor(tensor_dict: TensorDict, **kwargs):\n \"\"\"Assign non-tensor data to a TensorDict.\n\n Automatically detects if the value is a list with nested structures and uses\n the appropriate assignment method (NonTensorData for simple values,\n NonTensorStack for lists with nested structures).\n\n Args:\n tensor_dict: The TensorDict to assign to\n **kwargs: Key-value pairs where values can be:\n - Simple values (stored as NonTensorData)\n - Lists with nested structures (stored as NonTensorStack)\n\n Example:\n >>> td = TensorDict({\"obs\": torch.randn(3, 4)}, batch_size=[3])\n >>> assign_non_tensor(\n ... tensor_dict=td,\n ... metadata=\"experiment_1\", # Simple value\n ... turn_scores=[[], [0.5, 0.8], [0.9]] # Nested list\n ... )\n \"\"\"\n assert isinstance(tensor_dict, TensorDict), \"input dict must be a TensorDict\"\n for key, val in kwargs.items():\n if isinstance(val, (NonTensorData | NonTensorStack)):\n tensor_dict[key] = val\n elif isinstance(val, list):\n # For lists, use NonTensorStack\n assign_non_tensor_stack(tensor_dict=tensor_dict, key=key, val=val)\n else:\n # For non-list values, use NonTensorData\n assign_non_tensor_data(tensor_dict=tensor_dict, key=key, val=val)\n return tensor_dict\n\n\ndef unwrap_non_tensor_data(data):\n \"\"\"Unwrap a NonTensorData object to get the underlying value.\n\n If the input is a NonTensorData wrapper, extracts and returns the\n underlying data. Otherwise, returns the input unchanged.\n\n Args:\n data: Either a NonTensorData object or any other value.\n\n Returns:\n The unwrapped data if input was NonTensorData, otherwise the\n original input unchanged.\n\n Example:\n >>> wrapped = NonTensorData(\"hello\")\n >>> unwrap_non_tensor_data(wrapped)\n 'hello'\n >>> unwrap_non_tensor_data(42) # Non-wrapped value\n 42\n \"\"\"\n if isinstance(data, NonTensorData):\n return data.data\n return data\n\n\ndef get_non_tensor_data(data: TensorDict, key: str, default):\n \"\"\"Retrieve and unwrap non-tensor data from a TensorDict.\n\n Fetches the value for the given key from the TensorDict and automatically\n unwraps it if it's stored as NonTensorData.\n\n Args:\n data: The TensorDict to retrieve from.\n key: The key to look up.\n default: Value to return if the key is not found.\n\n Returns:\n The unwrapped value if the key exists and was wrapped in NonTensorData,\n the raw value if it wasn't wrapped, or the default if key not found.\n\n Example:\n >>> td = TensorDict({}, batch_size=[])\n >>> assign_non_tensor_data(td, \"config\", {\"lr\": 0.01})\n >>> get_non_tensor_data(td, \"config\", None)\n {'lr': 0.01}\n >>> get_non_tensor_data(td, \"missing\", \"default_value\")\n 'default_value'\n \"\"\"\n output = data.get(key, default)\n return unwrap_non_tensor_data(output)\n\n\ndef concat_nested_tensors(tensors: list[torch.Tensor]) -> torch.Tensor:\n \"\"\"Concatenate multiple nested tensors along the batch dimension.\n\n Takes a list of nested tensors with jagged layout and concatenates them\n into a single nested tensor. Each input tensor must have 2 or more dimensions and be contiguous.\n\n Args:\n tensors: List of nested tensors to concatenate. All tensors must\n be nested, contiguous, and have 2 or more dimensions.\n\n Returns:\n A new nested tensor with jagged layout containing all rows from\n the input tensors concatenated along dimension 0.\n\n Raises:\n AssertionError: If any tensor is not nested, not contiguous, or\n doesn't have 2 or more dimensions.\n\n Example:\n >>> t1 = torch.nested.as_nested_tensor([torch.randn(3), torch.randn(5)], layout=torch.jagged)\n >>> t2 = torch.nested.as_nested_tensor([torch.randn(2), torch.randn(4)], layout=torch.jagged)\n >>> result = concat_nested_tensors([t1, t2])\n >>> # result contains 4 rows: lengths [3, 5, 2, 4]\n \"\"\"\n for tensor in tensors:\n assert tensor.is_nested and tensor.is_contiguous()\n unbind_tensors = []\n for tensor in tensors:\n assert len(tensor.shape) >= 2, f\"nested tensor must have 2 or more dimensions. Got {tensor.shape}\"\n unbind_tensor = tensor.unbind(0)\n unbind_tensors.extend(list(unbind_tensor))\n\n tensor = torch.nested.as_nested_tensor(unbind_tensors, layout=torch.jagged)\n return tensor\n\n\ndef concat_tensordict_with_none_bsz(data: list[TensorDict]):\n \"\"\"Handle concatenation of TensorDicts with empty batch size.\n\n For TensorDicts that contain only metadata (NonTensorData) with no batch\n dimension, returns the first TensorDict as the concatenation result.\n\n Args:\n data: List of TensorDicts, each with empty batch_size (batch_size=[]).\n\n Returns:\n The first TensorDict from the list, as metadata concatenation\n simply preserves the first instance.\n\n Raises:\n AssertionError: If any TensorDict has a non-empty batch_size.\n\n Note:\n This is used internally by concat_tensordict when handling\n TensorDicts that contain only non-tensor metadata.\n \"\"\"\n for d in data:\n assert len(d.batch_size) == 0\n # directly return the first meta info\n return data[0]\n\n\ndef concat_tensordict(data: list[TensorDict]) -> TensorDict:\n \"\"\"Concatenate multiple TensorDicts along dimension zero.\n\n Combines a list of TensorDicts into a single TensorDict by concatenating\n all tensors along the batch dimension (dim=0). Handles nested tensors\n specially by unbinding and rebinding them.\n\n Args:\n data: List of TensorDicts to concatenate. All TensorDicts must have\n the same keys and the same set of nested tensor keys.\n\n Returns:\n A new TensorDict containing concatenated tensors from all inputs.\n\n Raises:\n AssertionError: If data is empty or if TensorDicts have inconsistent\n nested tensor keys.\n\n Note:\n - For TensorDicts with empty batch_size, returns the first one\n - Nested tensors are handled specially via concat_nested_tensors\n - Regular tensors use TensorDict.cat for efficient concatenation\n \"\"\"\n assert len(data) > 0, \"Must have at least one tensordict\"\n\n # Find nested tensor keys from the first tensordict\n nested_tensor_keys = {key for key, value in data[0].items() if isinstance(value, torch.Tensor) and value.is_nested}\n\n if not nested_tensor_keys:\n if len(data[0].batch_size) == 0:\n return concat_tensordict_with_none_bsz(data)\n # if batch size is None (only contain NonTensorData)\n return TensorDict.cat(data, dim=0)\n\n # Create a list of tensordicts containing only non-nested tensors for concatenation\n regular_tds = []\n for td in data:\n current_nested_keys = {k for k, v in td.items() if isinstance(v, torch.Tensor) and v.is_nested}\n assert current_nested_keys == nested_tensor_keys, \"All tensordicts must have the same set of nested tensors.\"\n\n # Create a new TensorDict with non-nested items without modifying the original\n regular_items = {k: v for k, v in td.items() if k not in nested_tensor_keys}\n regular_tds.append(TensorDict(regular_items, batch_size=td.batch_size, device=td.device))\n\n # Concatenate the regular tensordicts\n output = TensorDict.cat(regular_tds, dim=0)\n\n # Concatenate and add nested tensors to the output\n for key in nested_tensor_keys:\n nested_tensors_to_concat = [td[key] for td in data]\n output[key] = concat_nested_tensors(nested_tensors_to_concat)\n\n return output\n\n\ndef chunk_tensordict(td: TensorDict, chunks: int) -> list[TensorDict]:\n \"\"\"Split a TensorDict into equal-sized chunks with special nested tensor handling.\n\n Divides a TensorDict into the specified number of chunks along the batch\n dimension. Handles 3D+ nested tensors specially since torch.chunk() doesn't\n support jagged tensors with 3 or more dimensions.\n\n Args:\n td: The TensorDict to split.\n chunks: Number of chunks to create. Must evenly divide len(td).\n\n Returns:\n List of TensorDicts, each containing a portion of the original data.\n\n Raises:\n AssertionError: If td is not a TensorDict or if its length is not\n evenly divisible by chunks.\n\n Note:\n This is a workaround for PyTorch issue #153238 where torch.chunk()\n doesn't support 3D jagged tensors (e.g., MRoPE position_ids).\n See: https://github.com/pytorch/pytorch/issues/153238\n \"\"\"\n assert isinstance(td, TensorDict) and len(td) % chunks == 0, (\n f\"expecting td with length divisible by chunks, but got {len(td)} and {chunks}\"\n )\n chunk_size = len(td) // chunks\n keys = {key for key, val in td.items() if isinstance(val, torch.Tensor) and val.is_nested and val.dim() >= 3}\n new_td = TensorDict({k: v for k, v in td.items() if k not in keys}, batch_size=td.batch_size, device=td.device)\n\n tds = new_td.chunk(chunks=chunks)\n for key in keys:\n tensors = td[key].unbind(dim=0)\n for i, chunk_td in enumerate(tds):\n chunk_td[key] = torch.nested.as_nested_tensor(\n tensors[i * chunk_size : (i + 1) * chunk_size], layout=torch.jagged\n )\n\n return tds\n\n\ndef get_tensordict(tensor_dict: dict[str, torch.Tensor | list], non_tensor_dict: dict = None) -> TensorDict:\n \"\"\"Create a TensorDict from tensors and non-tensor data.\n\n Automatically handles nested structures in lists by converting them to NonTensorStack.\n This enables support for:\n - Lists of lists: [[], [0.5, 0.8], [0.9]]\n - Lists of dicts: [{\"acc\": 1.0}, {\"acc\": 0.0}]\n - Lists of lists of dicts: [[{\"content\": \"...\", \"role\": \"user\"}]]\n\n Args:\n tensor_dict: Dictionary of tensors and lists to include in the TensorDict\n non_tensor_dict: Dictionary of metadata to store as NonTensorData\n\n Returns:\n TensorDict with proper handling of nested structures\n\n Example:\n >>> td = get_tensordict(\n ... tensor_dict={\n ... \"obs\": torch.randn(3, 4),\n ... \"turn_scores\": [[], [0.5, 0.8], [0.9]] # Nested list\n ... },\n ... non_tensor_dict={\"experiment\": \"test\"}\n ... )\n \"\"\"\n tensor_dict = tensor_dict.copy()\n if non_tensor_dict is None:\n non_tensor_dict = {}\n\n batch_size = None\n\n for key, val in tensor_dict.items():\n if isinstance(val, torch.Tensor) and val.is_nested:\n assert val.is_contiguous(), \"Nested tensors must be contiguous. Try setting layout=torch.jagged\"\n assert val.layout == torch.jagged, \"Nested tensors must be jagged.\"\n\n # Skip validation for NonTensorStack as it's already properly formatted\n if isinstance(val, NonTensorStack):\n if batch_size is None:\n batch_size = len(val)\n else:\n assert len(val) == batch_size, (\n f\"Batch size of NonTensorStack {key} is not consistent with other tensors. \"\n f\"Expected {batch_size}, got {len(val)}\"\n )\n continue\n\n if isinstance(val, list):\n for v in val:\n assert not isinstance(v, torch.Tensor), (\n \"Passing a list makes the data NonTensorStack, \"\n \"which doesn't support torch.Tensor. Please convert to numpy first\"\n )\n # Convert to NonTensorStack to handle nested structures\n tensor_dict[key] = NonTensorStack.from_list([NonTensorData(item) for item in val])\n\n assert isinstance(val, torch.Tensor | list)\n\n if batch_size is None:\n batch_size = val.size(0) if isinstance(val, torch.Tensor) else len(val)\n else:\n val_batch_size = val.size(0) if isinstance(val, torch.Tensor) else len(val)\n assert val_batch_size == batch_size, (\n f\"Batch size of tensor {key} is not consistent with other tensors. \"\n f\"Expected {batch_size}, got {val_batch_size}\"\n )\n\n if batch_size is None:\n batch_size = []\n else:\n batch_size = [batch_size]\n\n for key, val in non_tensor_dict.items():\n assert key not in tensor_dict\n tensor_dict[key] = NonTensorData(val)\n\n return TensorDict(source=tensor_dict, batch_size=batch_size)\n\n\ndef index_select_tensor_dict(batch: TensorDict, indices: torch.Tensor | list[int]) -> TensorDict:\n \"\"\"Select rows from a TensorDict using indices.\n\n Creates a new TensorDict containing only the rows specified by indices.\n Handles regular tensors, nested tensors, NonTensorStack, and NonTensorData\n appropriately.\n\n Args:\n batch: The TensorDict to index into. Can be None.\n indices: 1D tensor or list of integers specifying which rows to select.\n\n Returns:\n A new TensorDict containing only the selected rows, or None if\n batch was None.\n\n Raises:\n AssertionError: If indices is not 1-dimensional.\n\n Note:\n - Regular tensors are indexed directly\n - Nested tensors are unbound, indexed, and rebound\n - NonTensorStack is indexed by batch dimension\n - NonTensorData (scalar metadata) is preserved unchanged\n \"\"\"\n if isinstance(indices, list):\n indices = torch.tensor(indices)\n\n assert indices.dim() == 1, \"indices must be a 1D tensor\"\n\n data_dict = {}\n batch_size = indices.shape[0]\n\n if batch is not None:\n for key, tensor in batch.items():\n if isinstance(tensor, torch.Tensor) and not tensor.is_nested:\n data_dict[key] = tensor[indices]\n elif isinstance(tensor, torch.Tensor) and tensor.is_nested:\n tensor_lst = tensor.unbind() # for performance\n data_dict[key] = torch.nested.as_nested_tensor(\n [tensor_lst[idx] for idx in indices], layout=torch.jagged\n )\n else:\n # This handles NonTensorStack (indexable by batch dim) and NonTensorData (scalar metadata).\n if tensor.shape:\n data_dict[key] = tensor[indices]\n else:\n data_dict[key] = tensor\n selected_batch = TensorDict(source=data_dict, batch_size=batch_size)\n else:\n selected_batch = None\n\n return selected_batch\n\n\ndef union_tensor_dict(tensor_dict1: TensorDict, tensor_dict2: TensorDict) -> TensorDict:\n \"\"\"Merge two TensorDicts, adding keys from the second to the first.\n\n Performs an in-place union of two TensorDicts. Keys from tensor_dict2\n that don't exist in tensor_dict1 are added. Keys that exist in both\n must have identical values.\n\n Args:\n tensor_dict1: The base TensorDict to merge into (modified in-place).\n tensor_dict2: The TensorDict whose keys will be added to tensor_dict1.\n\n Returns:\n The modified tensor_dict1 containing the union of both TensorDicts.\n\n Raises:\n AssertionError: If batch sizes don't match, or if a key exists in\n both TensorDicts with different values.\n\n Example:\n >>> td1 = TensorDict({\"a\": torch.tensor([1, 2])}, batch_size=[2])\n >>> td2 = TensorDict({\"b\": torch.tensor([3, 4])}, batch_size=[2])\n >>> result = union_tensor_dict(td1, td2)\n >>> list(result.keys())\n ['a', 'b']\n \"\"\"\n assert tensor_dict1.batch_size == tensor_dict2.batch_size, (\n f\"Two tensor dict must have identical batch size. Got {tensor_dict1.batch_size} and {tensor_dict2.batch_size}\"\n )\n for key in tensor_dict2.keys():\n if key not in tensor_dict1.keys():\n # Note that there is a difference between tensor_dict2[key] and tensor_dict2.get(key)\n tensor_dict1[key] = tensor_dict2.get(key)\n else:\n if isinstance(tensor_dict2[key], torch.Tensor):\n assert tensor_dict1[key].equal(tensor_dict2[key]), (\n f\"{key} in tensor_dict1 and tensor_dict2 are not the same object\"\n )\n else:\n # non-tensor\n assert tensor_dict1[key] == tensor_dict2[key], (\n f\"{key} in tensor_dict1 and tensor_dict2 are not the same object\"\n )\n\n return tensor_dict1\n\n\ndef make_iterator(tensordict: TensorDict, mini_batch_size, epochs, seed=None, dataloader_kwargs=None):\n \"\"\"Create an iterator that yields mini-batches from a TensorDict.\n\n Wraps a TensorDict in a DataLoader-style iterator that yields mini-batches\n for the specified number of epochs. Useful for training loops.\n\n Args:\n tensordict: The TensorDict to iterate over.\n mini_batch_size: Size of each mini-batch. Must evenly divide the\n TensorDict's batch size.\n epochs: Number of times to iterate through the entire dataset.\n seed: Optional random seed for reproducible shuffling.\n dataloader_kwargs: Optional dict of additional kwargs to pass to\n the underlying DataLoader (e.g., shuffle=True, num_workers=4).\n\n Returns:\n An iterator that yields TensorDict mini-batches.\n\n Raises:\n AssertionError: If batch size is not divisible by mini_batch_size.\n\n Example:\n >>> td = TensorDict({\"obs\": torch.randn(100, 4)}, batch_size=[100])\n >>> for batch in make_iterator(td, mini_batch_size=10, epochs=2):\n ... # batch is a TensorDict with batch_size=[10]\n ... pass\n \"\"\"\n from torch.utils.data import DataLoader\n\n assert tensordict.batch_size[0] % mini_batch_size == 0, f\"{tensordict.batch_size[0]} % {mini_batch_size} != 0\"\n # we can directly create a dataloader from TensorDict\n if dataloader_kwargs is None:\n dataloader_kwargs = {}\n\n if seed is not None:\n generator = torch.Generator()\n generator.manual_seed(seed)\n else:\n generator = None\n\n assert isinstance(dataloader_kwargs, dict)\n\n idx_lst = torch.arange(tensordict.shape[0])\n\n train_dataloader = DataLoader(\n dataset=idx_lst, batch_size=mini_batch_size, collate_fn=lambda x: x, generator=generator, **dataloader_kwargs\n )\n\n def get_data():\n for _ in range(epochs):\n for idx in train_dataloader:\n yield index_select_tensor_dict(tensordict, idx)\n\n return iter(get_data())\n\n\ndef assert_tensordict_eq(tensordict1: TensorDict, tensordict2: TensorDict):\n \"\"\"Assert that two TensorDicts are equal.\n\n Performs a deep equality check between two TensorDicts, verifying that\n they have the same keys with identical values. Handles nested tensors\n by comparing their unbound components.\n\n Args:\n tensordict1: First TensorDict to compare.\n tensordict2: Second TensorDict to compare.\n\n Raises:\n AssertionError: If the TensorDicts differ in keys, value types, or\n value contents. The error message indicates what differs.\n\n Note:\n - Regular tensors are compared element-wise\n - Nested tensors are unbound and compared component by component\n - Non-tensor values are compared with standard equality\n \"\"\"\n tensordict1_key_set = set(tensordict1.keys())\n tensordict2_key_set = set(tensordict2.keys())\n assert tensordict1_key_set == tensordict2_key_set, (\n f\"key set diffs. Got {tensordict2_key_set=} vs {tensordict1_key_set=}\"\n )\n\n for key in tensordict1.keys():\n val = tensordict1[key]\n val2 = tensordict2[key]\n\n assert type(val) is type(val2), f\"The type of {key} must be the same. Got {type(val)} vs {type(val2)}\"\n\n if isinstance(val, torch.Tensor):\n if val.is_nested:\n assert val.is_nested and val2.is_nested, (\n f\"Both tensors must be nested tensors. {val.is_nested=}, {val2.is_nested=}\"\n )\n t1, t2 = val.unbind(), val2.unbind()\n assert len(t1) == len(t2), f\"Nested tensor should have the same lengths. {len(t1)=} vs {len(t2)=}\"\n for c1, c2 in zip(t1, t2, strict=True):\n assert torch.equal(c1, c2), f\"Nested tensor components have different values. {c1=} vs {c2=}\"\n else:\n assert torch.all(torch.eq(val, val2)).item()\n else:\n assert val == val2\n\n\ndef get(tensordict: TensorDict, key: str, default=None) -> Any:\n \"\"\"Get a value from a TensorDict with automatic unwrapping.\n\n Retrieves a value from the TensorDict and automatically converts it\n to a Python-native format:\n - Tensors are returned as-is\n - NonTensorStack is converted to a Python list\n - NonTensorData is unwrapped to its underlying value\n\n Args:\n tensordict: The TensorDict to retrieve from.\n key: The key to look up.\n default: Value to return if the key doesn't exist. Defaults to None.\n\n Returns:\n The value for the key in its native format, or default if not found.\n\n Example:\n >>> td = get_tensordict({\"obs\": torch.randn(3, 4), \"labels\": [\"a\", \"b\", \"c\"]})\n >>> get(td, \"obs\") # Returns torch.Tensor\n >>> get(td, \"labels\") # Returns [\"a\", \"b\", \"c\"] as a list\n >>> get(td, \"missing\", \"default\") # Returns \"default\"\n \"\"\"\n if key not in tensordict:\n return default\n\n output = tensordict.get(key)\n if isinstance(output, torch.Tensor):\n return output\n elif isinstance(output, NonTensorStack):\n return output.tolist()\n else:\n assert isinstance(output, NonTensorData)\n return output.data\n\n\ndef get_keys(tensordict: TensorDict, keys: Iterable[str]) -> TensorDict:\n \"\"\"Extract a subset of keys from a TensorDict into a new TensorDict.\n\n Creates a new TensorDict containing only the specified keys. Values\n are properly categorized as tensor or non-tensor data.\n\n Args:\n tensordict: The source TensorDict.\n keys: Iterable of key names to extract.\n\n Returns:\n A new TensorDict containing only the specified keys with their values.\n\n Raises:\n KeyError: If any key in keys doesn't exist in the tensordict.\n\n Example:\n >>> td = get_tensordict({\"a\": torch.randn(3), \"b\": torch.randn(3), \"c\": torch.randn(3)})\n >>> subset = get_keys(td, [\"a\", \"c\"])\n >>> list(subset.keys())\n ['a', 'c']\n \"\"\"\n tensor_output = {}\n non_tensor_output = {}\n for key in keys:\n if key not in tensordict.keys():\n raise KeyError(f\"key {key} not in tensordict\")\n output = tensordict.get(key)\n if isinstance(output, torch.Tensor):\n tensor_output[key] = output\n elif isinstance(output, NonTensorStack):\n tensor_output[key] = output.tolist()\n else:\n assert isinstance(output, NonTensorData)\n non_tensor_output[key] = output.data\n\n return get_tensordict(tensor_output, non_tensor_output)\n\n\ndef pop(tensordict: TensorDict, key: str, default=None) -> Any:\n \"\"\"Remove and return a value from a TensorDict with automatic unwrapping.\n\n Removes the specified key from the TensorDict and returns its value,\n automatically converting to Python-native format (same as get()).\n\n Args:\n tensordict: The TensorDict to pop from.\n key: The key to remove and return.\n default: Value to return if the key doesn't exist. Defaults to None.\n\n Returns:\n The value for the key in its native format, or default if not found.\n The key is removed from the TensorDict.\n\n Example:\n >>> td = get_tensordict({\"obs\": torch.randn(3, 4), \"labels\": [\"a\", \"b\", \"c\"]})\n >>> labels = pop(td, \"labels\") # Returns [\"a\", \"b\", \"c\"], removes from td\n >>> \"labels\" in td.keys()\n False\n \"\"\"\n _sentinel = object()\n output = tensordict.pop(key, _sentinel)\n if output is _sentinel:\n return default\n\n if isinstance(output, torch.Tensor):\n return output\n elif isinstance(output, NonTensorStack):\n return output.tolist()\n else:\n assert isinstance(output, NonTensorData)\n return output.data\n\n\ndef pop_keys(tensordict: TensorDict, keys: Iterable[str]) -> TensorDict:\n \"\"\"Remove multiple keys from a TensorDict and return them as a new TensorDict.\n\n Removes the specified keys from the source TensorDict and creates a new\n TensorDict containing those keys and their values.\n\n Args:\n tensordict: The source TensorDict to pop from (modified in-place).\n keys: Iterable of key names to remove and return.\n\n Returns:\n A new TensorDict containing the popped keys and their values.\n\n Raises:\n KeyError: If any key in keys doesn't exist in the tensordict.\n\n Example:\n >>> td = get_tensordict({\"a\": torch.randn(3), \"b\": torch.randn(3), \"c\": torch.randn(3)})\n >>> popped = pop_keys(td, [\"a\", \"c\"])\n >>> list(td.keys()) # Only 'b' remains\n ['b']\n >>> list(popped.keys())\n ['a', 'c']\n \"\"\"\n tensor_output = {}\n non_tensor_output = {}\n for key in keys:\n if key not in tensordict.keys():\n raise KeyError(f\"key {key} not in tensordict\")\n output = tensordict.get(key)\n if isinstance(output, torch.Tensor):\n tensor_output[key] = tensordict.pop(key)\n elif isinstance(output, NonTensorStack):\n tensor_output[key] = tensordict.pop(key).tolist()\n else:\n assert isinstance(output, NonTensorData)\n non_tensor_output[key] = tensordict.pop(key)\n\n return get_tensordict(tensor_output, non_tensor_output)\n\n\ndef pad_to_divisor(data: TensorDict, size_divisor: int):\n \"\"\"Pad a TensorDict's batch dimension to be divisible by a given divisor.\n\n If the TensorDict's length is not evenly divisible by size_divisor,\n pads the batch dimension by repeating elements from the beginning.\n Useful for ensuring even distribution across workers in distributed training.\n\n Args:\n data: The TensorDict to pad.\n size_divisor: The divisor that the padded length must be divisible by.\n\n Returns:\n tuple: A tuple containing:\n - data (TensorDict): The padded TensorDict (or original if no padding needed)\n - pad_size (int): Number of elements added as padding (0 if none)\n\n Raises:\n AssertionError: If data is not a TensorDict.\n\n Example:\n >>> td = TensorDict({\"obs\": torch.randn(10, 4)}, batch_size=[10])\n >>> padded, pad_size = pad_to_divisor(td, 4)\n >>> len(padded) # 12 (next multiple of 4 after 10)\n 12\n >>> pad_size\n 2\n \"\"\"\n assert isinstance(data, TensorDict), \"data must be a TensorDict\"\n if len(data) % size_divisor != 0:\n pad_size = size_divisor - len(data) % size_divisor\n padding_protos = []\n remaining_pad = pad_size\n while remaining_pad > 0:\n take_size = min(remaining_pad, len(data))\n padding_protos.append(data[:take_size])\n remaining_pad -= take_size\n data_padded = torch.cat([data] + padding_protos)\n else:\n if len(data) == 0:\n logging.warning(\"padding a DataProto with no item, no changed made\")\n pad_size = 0\n data_padded = data\n return data_padded, pad_size\n\n\ndef unpad(data: TensorDict, pad_size):\n \"\"\"Remove padding from a TensorDict.\n\n Reverses the effect of pad_to_divisor by removing the specified number\n of elements from the end of the TensorDict.\n\n Args:\n data: The padded TensorDict.\n pad_size: Number of padding elements to remove. If 0, returns\n data unchanged.\n\n Returns:\n The TensorDict with padding removed, equivalent to data[:-pad_size].\n\n Example:\n >>> td = TensorDict({\"obs\": torch.randn(12, 4)}, batch_size=[12])\n >>> unpadded = unpad(td, pad_size=2)\n >>> len(unpadded)\n 10\n \"\"\"\n if pad_size != 0:\n data = data[:-pad_size]\n return data\n\n\ndef contiguous(data: TensorDict) -> TensorDict:\n \"\"\"Call contiguous on a tensor dict. The contiguous function of tensordict lib will make NonTensorStack.\n This function will always return a new tensordict\n\n Args:\n data: The input tensordict\n\n Returns:\n a tensordict that is contiguous\n\n \"\"\"\n tensor_dict = {}\n non_tensor_dict = {}\n\n for key in data.keys():\n val = data.get(key)\n if isinstance(val, NonTensorData):\n non_tensor_dict[key] = val\n elif isinstance(val, NonTensorStack):\n tensor_dict[key] = val\n else:\n assert isinstance(val, torch.Tensor), f\"Expect val to be a torch.Tensor. Got {type(val)}\"\n tensor_dict[key] = val.contiguous()\n\n return get_tensordict(tensor_dict=tensor_dict, non_tensor_dict=non_tensor_dict)\n\n\ndef maybe_fix_3d_position_ids(data: TensorDict):\n # note for tensordict with pickle/unpickle. nested tensor in tensordict after consolidate and pickle/unpickle\n # will incur indexing error for ragged tensor. This only happens when using 3D position ids in VLMs.\n # This is likely a bug in tensordict. As a workaround, we manually set _ragged_index.\n if \"position_ids\" in data.keys() and data[\"position_ids\"].dim() == 3 and data[\"position_ids\"].is_nested:\n data[\"position_ids\"]._ragged_idx = 2\n"}124{"file_name": "verl__utils__tokenizer.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"Utils for tokenization.\"\"\"\n\nimport types\nimport warnings\n\n__all__ = [\"hf_tokenizer\", \"hf_processor\"]\n\n\ndef set_pad_token_id(tokenizer):\n \"\"\"Set pad_token_id to eos_token_id if it is None.\n\n Args:\n tokenizer (transformers.PreTrainedTokenizer): The tokenizer to be set.\n\n \"\"\"\n if tokenizer.pad_token_id is None:\n tokenizer.pad_token_id = tokenizer.eos_token_id\n warnings.warn(f\"tokenizer.pad_token_id is None. Now set to {tokenizer.eos_token_id}\", stacklevel=1)\n if tokenizer.pad_token is None:\n tokenizer.pad_token = tokenizer.eos_token\n warnings.warn(f\"tokenizer.pad_token is None. Now set to {tokenizer.eos_token}\", stacklevel=1)\n\n\ndef hf_tokenizer(name_or_path, correct_pad_token=True, correct_gemma2=True, **kwargs):\n \"\"\"Create a huggingface pretrained tokenizer which correctness handles eos and pad tokens.\n\n Args:\n\n name (str): The name of the tokenizer.\n correct_pad_token (bool): Whether to correct the pad token id.\n correct_gemma2 (bool): Whether to correct the gemma2 tokenizer.\n\n Returns:\n\n transformers.PreTrainedTokenizer: The pretrained tokenizer.\n\n \"\"\"\n from transformers import AutoTokenizer\n\n if correct_gemma2 and isinstance(name_or_path, str) and \"gemma-2-2b-it\" in name_or_path:\n # the EOS token in gemma2 is ambiguious, which may worsen RL performance.\n # https://huggingface.co/google/gemma-2-2b-it/commit/17a01657f5c87135bcdd0ec7abb4b2dece04408a\n warnings.warn(\n \"Found gemma-2-2b-it tokenizer. Set eos_token and eos_token_id to <end_of_turn> and 107.\", stacklevel=1\n )\n kwargs[\"eos_token\"] = \"<end_of_turn>\"\n kwargs[\"eos_token_id\"] = 107\n tokenizer = AutoTokenizer.from_pretrained(name_or_path, **kwargs)\n if correct_pad_token:\n set_pad_token_id(tokenizer)\n return tokenizer\n\n\ndef hf_processor(name_or_path, **kwargs):\n \"\"\"Create a huggingface processor to process multimodal data.\n\n Args:\n name_or_path (str): The name of the processor.\n\n Returns:\n transformers.ProcessorMixin: The pretrained processor.\n \"\"\"\n from transformers import AutoConfig, AutoProcessor\n\n try:\n processor = AutoProcessor.from_pretrained(name_or_path, **kwargs)\n config = AutoConfig.from_pretrained(name_or_path, **kwargs)\n\n # Bind vlm model's get_rope_index method to processor\n processor.config = config\n match processor.__class__.__name__:\n case \"Qwen2VLProcessor\":\n from transformers.models.qwen2_vl import Qwen2VLModel\n\n processor.get_rope_index = types.MethodType(Qwen2VLModel.get_rope_index, processor)\n case \"Qwen2_5_VLProcessor\":\n from transformers.models.qwen2_5_vl import Qwen2_5_VLModel\n\n processor.get_rope_index = types.MethodType(Qwen2_5_VLModel.get_rope_index, processor)\n case \"Qwen3VLProcessor\":\n from transformers.models.qwen3_vl import Qwen3VLModel\n\n processor.get_rope_index = types.MethodType(Qwen3VLModel.get_rope_index, processor)\n case \"Glm4vImageProcessor\":\n from transformers.models.glm4v import Glm4vModel\n\n processor.get_rope_index = types.MethodType(Glm4vModel.get_rope_index, processor)\n case \"MllamaProcessor\":\n pass # MllamaProcessor and MllamaModel doesn't have get_rope_index property\n case _:\n raise ValueError(f\"Unsupported processor type: {processor.__class__.__name__}\")\n except Exception as e:\n processor = None\n # TODO(haibin.lin): try-catch should be removed after adding transformer version req to setup.py to avoid\n # silent failure\n warnings.warn(f\"Failed to create processor: {e}. This may affect multimodal processing\", stacklevel=1)\n # Avoid load tokenizer, see:\n # https://github.com/huggingface/transformers/blob/v4.49.0/src/transformers/models/auto/processing_auto.py#L344\n if processor is not None and \"Processor\" not in processor.__class__.__name__:\n processor = None\n return processor\n"}125{"file_name": "verl__utils__torch_dtypes.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nAdapted from Cruise.\n\"\"\"\n\nimport torch\n\nHALF_LIST = [16, \"16\", \"fp16\", \"float16\", torch.float16]\nFLOAT_LIST = [32, \"32\", \"fp32\", \"float32\", torch.float32]\nBFLOAT_LIST = [\"bf16\", \"bfloat16\", torch.bfloat16]\n\n\nclass PrecisionType:\n \"\"\"Type of precision used.\n\n >>> PrecisionType.HALF == 16\n True\n >>> PrecisionType.HALF in (16, \"16\")\n True\n \"\"\"\n\n HALF = \"16\"\n FLOAT = \"32\"\n FULL = \"64\"\n BFLOAT = \"bf16\"\n MIXED = \"mixed\"\n\n @staticmethod\n def supported_type(precision: str | int) -> bool:\n return any(x == precision for x in PrecisionType)\n\n @staticmethod\n def supported_types() -> list[str]:\n return [x.value for x in PrecisionType]\n\n @staticmethod\n def is_fp16(precision):\n return precision in HALF_LIST\n\n @staticmethod\n def is_fp32(precision):\n return precision in FLOAT_LIST\n\n @staticmethod\n def is_bf16(precision):\n return precision in BFLOAT_LIST\n\n @staticmethod\n def to_dtype(precision):\n if precision in HALF_LIST:\n return torch.float16\n elif precision in FLOAT_LIST:\n return torch.float32\n elif precision in BFLOAT_LIST:\n return torch.bfloat16\n else:\n raise RuntimeError(f\"unexpected precision: {precision}\")\n\n @staticmethod\n def to_str(precision):\n if precision == torch.float16:\n return \"fp16\"\n elif precision == torch.float32:\n return \"fp32\"\n elif precision == torch.bfloat16:\n return \"bf16\"\n else:\n raise RuntimeError(f\"unexpected precision: {precision}\")\n"}126{"file_name": "verl__utils__torch_functional.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nContain small torch utilities\n\"\"\"\n\nimport math\nfrom contextlib import contextmanager\nfrom typing import Optional\n\nimport torch\nimport torch.distributed\nimport torch.nn.functional as F\nfrom tensordict import TensorDict\nfrom torch import nn\nfrom torch.optim import Optimizer\nfrom torch.optim.lr_scheduler import LambdaLR\nfrom transformers import PreTrainedTokenizer\n\nfrom verl.utils.device import get_device_name, get_torch_device\n\ntry:\n from flash_attn.ops.triton.cross_entropy import cross_entropy_loss\n\n FLAH_ATTN_CROSS_ENTROPY_LOSS_AVAILABLE = True\nexcept ImportError:\n FLAH_ATTN_CROSS_ENTROPY_LOSS_AVAILABLE = False\n\n\ntry:\n import torch_npu\n\n NPU_CROSS_ENTROPY_LOSS_AVAILABLE = hasattr(torch_npu, \"npu_cross_entropy_loss\")\nexcept ImportError:\n NPU_CROSS_ENTROPY_LOSS_AVAILABLE = False\n\n\ndef gather_from_labels(data: torch.Tensor, label: torch.Tensor) -> torch.Tensor:\n \"\"\"Gather values from data tensor at positions specified by label indices.\n\n Selects elements from the last dimension of `data` based on indices in `label`.\n Commonly used to extract log-probabilities for specific token IDs from a\n vocabulary distribution.\n\n Args:\n data: Input tensor of shape (..., vocab_size) containing values to gather from.\n label: Index tensor of shape (...,) with values in range [0, vocab_size).\n\n Returns:\n torch.Tensor: Gathered values with shape (...,), same as label shape.\n\n Example:\n >>> logits = torch.randn(2, 3, 100) # [batch, seq, vocab]\n >>> labels = torch.randint(0, 100, (2, 3)) # [batch, seq]\n >>> gathered = gather_from_labels(logits, labels) # [batch, seq]\n \"\"\"\n output = torch.gather(data, -1, label.unsqueeze(-1)).squeeze(-1)\n return output\n\n\ndef logprobs_from_logits(logits, labels, inplace_backward=True):\n \"\"\"\n Compute per-token log-probabilities for the given labels.\n\n Uses a Flash-Attention–based cross-entropy (if available) for efficient backward,\n otherwise falls back to a standard log-softmax+gather approach.\n\n See: https://github.com/pytorch/pytorch/issues/563#issuecomment-330103591\n\n Args:\n logits (Tensor): Model outputs of shape (..., vocab_size).\n labels (LongTensor): True class indices of shape matching logits[..., :-1].\n inplace_backward (bool): If True and Flash-Attn is available, perform backward in-place.\n\n Returns:\n Tensor: Log-probabilities of the target labels, shape logits.shape[:-1].\n \"\"\"\n if FLAH_ATTN_CROSS_ENTROPY_LOSS_AVAILABLE:\n batch_dim = logits.shape[:-1]\n last_dim = logits.shape[-1]\n logits = logits.reshape(-1, last_dim)\n labels = labels.reshape(-1)\n output = logprobs_from_logits_flash_attn(logits, labels, inplace_backward=inplace_backward)\n output = output.view(*batch_dim)\n elif NPU_CROSS_ENTROPY_LOSS_AVAILABLE:\n output = logprobs_from_logits_torch_npu(logits, labels)\n else:\n output = logprobs_from_logits_v2(logits, labels)\n return output\n\n\ndef logprobs_from_logits_flash_attn(\n logits: torch.Tensor, labels: torch.Tensor, inplace_backward: bool = True\n) -> torch.Tensor:\n \"\"\"Compute log-probabilities using Flash Attention's optimized cross-entropy.\n\n Uses the Flash Attention library's Triton-based cross-entropy implementation\n for efficient computation on NVIDIA GPUs.\n\n Args:\n logits: Model output logits of shape (batch_size, vocab_size).\n labels: Target token indices of shape (batch_size,).\n inplace_backward: If True, perform backward pass in-place for memory efficiency.\n\n Returns:\n torch.Tensor: Log-probabilities for target labels, shape (batch_size,).\n\n Raises:\n AssertionError: If flash-attn version < 2.4.3 (different return format).\n \"\"\"\n output = cross_entropy_loss(logits, labels, inplace_backward=inplace_backward)\n assert isinstance(output, tuple), (\n \"please make sure flash-attn>=2.4.3 where cross_entropy_loss returns Tuple[losses, z_losses].\"\n )\n return -output[0]\n\n\ndef logprobs_from_logits_torch_npu(logits: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:\n \"\"\"Compute log-probabilities using Ascend NPU's optimized cross-entropy.\n\n Uses torch_npu's native cross-entropy implementation for efficient\n computation on Huawei Ascend NPU devices.\n\n Args:\n logits: Model output logits of shape (..., vocab_size).\n labels: Target token indices of shape (...,).\n\n Returns:\n torch.Tensor: Log-probabilities for target labels, same shape as labels.\n \"\"\"\n batch_dim = logits.shape[:-1]\n logits = logits.reshape(-1, logits.shape[-1])\n loss, _, _, _ = torch_npu.npu_cross_entropy_loss(logits, labels.reshape(-1), reduction=\"none\")\n return -loss.view(*batch_dim)\n\n\ndef logprobs_from_logits_naive(logits: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:\n \"\"\"Compute log-probabilities using standard log-softmax approach.\n\n Simple implementation using PyTorch's log_softmax followed by gathering.\n Less memory-efficient than specialized implementations but works on all devices.\n\n Args:\n logits: Model output logits of shape (..., vocab_size).\n labels: Target token indices of shape (...,).\n\n Returns:\n torch.Tensor: Log-probabilities for target labels, same shape as labels.\n \"\"\"\n logp = F.log_softmax(logits, dim=-1)\n logpy = gather_from_labels(logp, labels)\n return logpy\n\n\ndef logprobs_from_logits_v2(logits: torch.FloatTensor, labels: torch.Tensor) -> torch.Tensor:\n \"\"\"Memory-efficient log-probability computation using row-wise processing.\n\n Computes log-probabilities by processing one row at a time to reduce peak\n memory consumption. Uses logsumexp for float32/float64, falls back to\n log_softmax for bfloat16 due to numerical stability concerns.\n\n The mathematical identity used is: log_softmax(x_i) = x_i - logsumexp(x)\n\n Args:\n logits: Model output logits of shape (batch_size, seq_len, vocab_size)\n or (batch_size, vocab_size).\n labels: Target token indices matching logits shape without vocab dimension.\n\n Returns:\n torch.Tensor: Log-probabilities for target labels.\n\n Note:\n This implementation trades compute for memory by iterating over batch\n dimension, making it suitable for large vocabulary sizes.\n \"\"\"\n if logits.dtype in [torch.float32, torch.float64]:\n logits_labels = torch.gather(logits, dim=-1, index=labels.unsqueeze(-1)).squeeze(-1)\n # loop to reduce peak mem consumption\n logsumexp_values = torch.stack([torch.logsumexp(logit, dim=-1) for logit in logits])\n logprobs_labels = logits_labels - logsumexp_values # log_softmax(x_i) = x_i - logsumexp(x)\n else:\n # logsumexp approach is unstable with bfloat16, fall back to slightly less efficent approach\n logprobs_labels = []\n for row_logits, row_labels in zip(logits, labels, strict=True): # loop to reduce peak mem consumption\n row_logprobs = F.log_softmax(row_logits, dim=-1)\n row_logprobs_labels = row_logprobs.gather(dim=-1, index=row_labels.unsqueeze(-1)).squeeze(-1)\n logprobs_labels.append(row_logprobs_labels)\n logprobs_labels = torch.stack(logprobs_labels)\n return logprobs_labels\n\n\ndef clip_by_value(x: torch.Tensor, tensor_min: torch.Tensor, tensor_max: torch.Tensor) -> torch.Tensor:\n \"\"\"Clip tensor values to a range defined by tensor bounds.\n\n Extension of torch.clamp that supports tensor-valued min/max bounds\n instead of only scalar bounds.\n\n Args:\n x: Input tensor to clip.\n tensor_min: Minimum bound tensor (broadcastable to x).\n tensor_max: Maximum bound tensor (broadcastable to x).\n\n Returns:\n torch.Tensor: Clipped tensor with values in [tensor_min, tensor_max].\n\n See Also:\n https://github.com/pytorch/pytorch/issues/2793#issuecomment-428784713\n \"\"\"\n clipped = torch.max(torch.min(x, tensor_max), tensor_min)\n return clipped\n\n\ndef entropy_from_logits(logits: torch.Tensor) -> torch.Tensor:\n \"\"\"Calculate Shannon entropy from unnormalized logits.\n\n Computes H(p) = -sum(p * log(p)) using the numerically stable formula:\n entropy = logsumexp(logits) - sum(softmax(logits) * logits)\n\n Args:\n logits: Unnormalized log-probabilities of shape (..., vocab_size).\n\n Returns:\n torch.Tensor: Entropy values with shape (...,), one per distribution.\n \"\"\"\n pd = torch.nn.functional.softmax(logits, dim=-1)\n entropy = torch.logsumexp(logits, dim=-1) - torch.sum(pd * logits, dim=-1)\n return entropy\n\n\ndef entropy_from_logits_with_chunking(logits: torch.Tensor, chunk_size: int = 2048) -> torch.Tensor:\n \"\"\"Memory-efficient entropy calculation using chunked processing.\n\n Computes entropy by processing the batch in chunks to reduce peak memory\n usage. Useful for large batch sizes or when memory is constrained.\n\n Args:\n logits: Unnormalized log-probabilities of shape (batch_size, vocab_size).\n chunk_size: Number of samples to process at once. Defaults to 2048.\n\n Returns:\n torch.Tensor: Entropy values with shape (batch_size,).\n\n Note:\n Converts chunks to float32 for numerical stability during computation.\n \"\"\"\n entropy = torch.zeros(logits.shape[0], device=logits.device)\n for i in range(0, logits.shape[0], chunk_size):\n logits_chunk = logits[i : i + chunk_size].float()\n pd_chunk = torch.nn.functional.softmax(logits_chunk, dim=-1)\n entropy_chunk = torch.logsumexp(logits_chunk, dim=-1) - torch.sum(pd_chunk * logits_chunk, dim=-1)\n entropy[i : i + chunk_size] = entropy_chunk\n return entropy\n\n\ndef masked_sum(values: torch.Tensor, mask: torch.Tensor, axis: int | tuple[int, ...] | None = None) -> torch.Tensor:\n \"\"\"Compute sum of tensor values where mask is True.\n\n NaN values outside the mask are replaced with zeros to prevent\n contaminating the sum.\n\n Args:\n values: Input tensor containing values to sum.\n mask: Boolean or numeric mask tensor (same shape as values).\n Non-zero values indicate elements to include.\n axis: Dimension(s) along which to sum. None sums all elements.\n\n Returns:\n torch.Tensor: Sum of masked values, reduced along specified axis.\n \"\"\"\n # If NaNs exist out of mask, replace NaNs in values with a value that\n # won't affect the sum (e.g., 0 for masked regions)\n valid_values = torch.where(mask.bool(), values, 0.0)\n return (valid_values * mask).sum(axis=axis)\n\n\ndef masked_mean(values, mask, axis=None):\n \"\"\"\n Compute the mean of `values` over elements selected by `mask`.\n\n Args:\n values (Tensor): Input tensor.\n mask (Tensor): Boolean or numeric mask of the same shape as `values`.\n axis (int or tuple of int, optional): Dimension(s) along which to compute the mean.\n Defaults to None (over all elements).\n\n Returns:\n Tensor: Masked mean, with shape equal to `values` reduced over `axis`.\n \"\"\"\n s = masked_sum(values, mask, axis)\n return s / (mask.sum(axis=axis) + 1e-8)\n\n\ndef masked_var(values, mask, unbiased=True):\n \"\"\"Compute variance of tensor with masked values.\"\"\"\n mean = masked_mean(values, mask)\n centered_values = values - mean\n variance = masked_mean(centered_values**2, mask)\n if unbiased:\n mask_sum = mask.sum()\n if mask_sum == 0:\n raise ValueError(\"At least one element in the mask has to be 1.\")\n # note that if mask_sum == 1, then there is a division by zero issue\n # to avoid it you just need to use a larger minibatch_size\n if mask_sum == 1:\n raise ValueError(\"The sum of the mask is one, which can cause a division by zero.\")\n bessel_correction = mask_sum / (mask_sum - 1)\n variance = variance * bessel_correction\n return variance\n\n\ndef masked_whiten(values, mask, shift_mean=True):\n \"\"\"\n Whiten `values` by normalizing with mean and variance computed over `mask`.\n\n Args:\n values (torch.Tensor): Input tensor.\n mask (torch.Tensor): Boolean tensor of same shape, selects elements for stats.\n shift_mean (bool): If True (default), output is zero-mean;\n if False, the original mean is re-added after scaling.\n\n Returns:\n torch.Tensor: Whitened tensor of same shape as `values`.\n \"\"\"\n mean, var = masked_mean(values, mask), masked_var(values, mask)\n whitened = (values - mean) * torch.rsqrt(var + 1e-8)\n if not shift_mean:\n whitened += mean\n return whitened\n\n\ndef get_response_mask(response_id: torch.Tensor, eos_token: int | list[int] = 2, dtype=torch.int64):\n \"\"\"\n end of sentence token can be int or list: 1 or [1, 2]\n e.g.\n response_id = torch.tensor([[20, 10, 34, 1, 0, 0, 0],\n [78, 0, 76, 2, 1, 0, 0],\n [23, 98, 1, 0, 0, 0, 0],\n [33, 3, 98, 45, 1, 0, 0]])\n #eos_token=1\n response_mask: tensor([[1, 1, 1, 1, 0, 0, 0],\n [1, 1, 1, 1, 1, 0, 0],\n [1, 1, 1, 0, 0, 0, 0],\n [1, 1, 1, 1, 1, 0, 0]])\n #eos_token=[1,2]\n response_mask: tensor([[1, 1, 1, 1, 0, 0, 0],\n [1, 1, 1, 1, 0, 0, 0],\n [1, 1, 1, 0, 0, 0, 0],\n [1, 1, 1, 1, 1, 0, 0]])\n \"\"\"\n eos_mask = torch.isin(response_id, torch.tensor(eos_token, device=response_id.device)).int()\n return (eos_mask.cumsum(dim=1) - eos_mask).eq(0).to(dtype)\n\n\ndef compute_grad_norm(model: nn.Module) -> float:\n \"\"\"Compute the squared L2 norm of all gradients in a model.\n\n Sums the squared values of all gradient tensors across all parameters.\n Useful for monitoring gradient magnitudes during training.\n\n Args:\n model: PyTorch model with computed gradients.\n\n Returns:\n float: Sum of squared gradient values (not the square root).\n\n Note:\n Returns the squared norm, not the norm itself. To get the actual\n L2 norm, take the square root of the returned value.\n \"\"\"\n total_grad_square = 0\n for param in model.parameters():\n if param.grad is not None:\n total_grad_square += torch.sum(torch.square(param.grad.detach())).item()\n return total_grad_square\n\n\ndef broadcast_dict_tensor(tensors: dict[str, torch.Tensor] | TensorDict, src: int, group) -> None:\n \"\"\"Broadcast all tensors in a dictionary from source rank to all ranks.\n\n Iterates over all tensors in the dictionary and broadcasts each one\n from the source rank to all other ranks in the process group.\n\n Args:\n tensors: Dictionary or TensorDict containing tensors to broadcast.\n src: Source rank from which to broadcast.\n group: Process group for the broadcast operation.\n\n Note:\n This implementation broadcasts tensors one at a time. Could be optimized\n to use a single broadcast with packed tensors.\n \"\"\"\n for key in tensors.sorted_keys:\n torch.distributed.broadcast(tensors[key], src=src, group=group, async_op=False)\n\n\ndef allgather_dict_tensors(\n tensors: dict[str, torch.Tensor] | TensorDict, size: int, group, dim: int = 0\n) -> dict[str, torch.Tensor] | TensorDict:\n \"\"\"Gather tensors from all ranks and concatenate them.\n\n Performs all_gather on each tensor in the dictionary and concatenates\n the results along the specified dimension.\n\n Args:\n tensors: Dictionary or TensorDict containing tensors to gather.\n size: Number of ranks in the process group.\n group: Process group for the all_gather operation.\n dim: Dimension along which to concatenate gathered tensors. Defaults to 0.\n\n Returns:\n Dictionary or TensorDict (matching input type) with gathered and\n concatenated tensors. Each tensor's size along `dim` is multiplied by `size`.\n\n Note:\n This implementation gathers tensors one at a time synchronously.\n Could be optimized using async ops or packed all_gather.\n \"\"\"\n if isinstance(tensors, TensorDict):\n is_tensor_dict = True\n tensors_as_dict = tensors.to_dict()\n else:\n tensors_as_dict = tensors\n is_tensor_dict = False\n\n output = {}\n sorted_keys = sorted(tensors_as_dict.keys())\n for key in sorted_keys:\n val = tensors_as_dict[key]\n output[key] = [torch.empty_like(val) for _ in range(size)]\n torch.distributed.all_gather(output[key], val, group=group, async_op=False)\n output[key] = torch.cat(output[key], dim=dim)\n\n if is_tensor_dict:\n output = TensorDict(source=output, batch_size=tensors.batch_size[0] * size)\n\n return output\n\n\ndef allgather_dict_into_dict(data: dict, group=None) -> dict:\n \"\"\"allgather a dict into a dict of list\n\n Args:\n data: a dict\n group: the process group to allgather\n\n Returns: dict containing a list of the results from allgather\n\n \"\"\"\n assert isinstance(data, dict), f\"Expect data to be a dictionary, Got {type(data)}\"\n\n group_size = torch.distributed.get_world_size(group=group)\n\n final_metrics = {}\n all_metrics_lst = [None for _ in range(group_size)]\n torch.distributed.all_gather_object(all_metrics_lst, data, group=group)\n\n for all_metrics in all_metrics_lst:\n for key, val in all_metrics.items():\n if key not in final_metrics:\n final_metrics[key] = []\n final_metrics[key].append(val)\n return final_metrics\n\n\ndef split_dict_tensor_into_batches(tensors: TensorDict, batch_size) -> list[TensorDict]:\n assert tensors.batch_size[0] % batch_size == 0, (\n f\"input data batch size: {tensors.batch_size[0]}, split batch size: {batch_size}\"\n )\n return tensors.split(batch_size)\n\n\ndef pad_2d_list_to_length(response, pad_token_id, max_length=None):\n \"\"\"\n pad a 2D list (e.g. responses, logprobs) to a 2D tensor.\n \"\"\"\n response_length = max(len(sub_list) for sub_list in response)\n target_length = max_length if max_length is not None and max_length > response_length else response_length\n padded_response = [tuple(sub_list) + (pad_token_id,) * (target_length - len(sub_list)) for sub_list in response]\n tensor = torch.tensor(padded_response)\n return tensor\n\n\ndef pad_sequence_to_length(tensors, max_seq_len, pad_token_id, left_pad=False):\n \"\"\"\n pad a 2D tensors (e.g. responses, logprobs) in the last dim to max_seq_length.\n input shape: [bs, seq_length]\n output shape: [bs, max_seq_length]\n \"\"\"\n if tensors.shape[-1] >= max_seq_len:\n return tensors\n # (0, max_seq_len - tensors.shape[-1]) means right pad to max_seq_length and no left pad\n pad_tuple = (max_seq_len - tensors.shape[-1], 0) if left_pad else (0, max_seq_len - tensors.shape[-1])\n return F.pad(tensors, pad_tuple, \"constant\", pad_token_id)\n\n\ndef postprocess_data(\n input_ids: torch.Tensor,\n attention_mask: torch.Tensor,\n max_length: int,\n pad_token_id: int,\n left_pad=True,\n truncation=\"error\",\n):\n \"\"\"Process tokenizer outputs to consistent shapes via padding/truncation.\n\n Args:\n input_ids: Token indices [batch_size, seq_len]\n attention_mask: Mask [batch_size, seq_len]\n max_length: Target sequence length\n pad_token_id: Padding token ID\n left_pad: Pad left if True\n truncation: \"left\", \"right\", \"middle\" or \"error\"\n\n Returns:\n (input_ids, attention_mask) padded/truncated to max_length\n \"\"\"\n assert truncation in [\"left\", \"right\", \"middle\", \"error\"]\n assert input_ids.ndim == 2\n\n sequence_length = input_ids.shape[-1]\n if sequence_length < max_length:\n input_ids = pad_sequence_to_length(\n input_ids, max_seq_len=max_length, pad_token_id=pad_token_id, left_pad=left_pad\n )\n attention_mask = pad_sequence_to_length(\n attention_mask, max_seq_len=max_length, pad_token_id=0, left_pad=left_pad\n )\n elif sequence_length > max_length:\n if truncation == \"left\":\n # actually, left truncation may not be reasonable\n input_ids = input_ids[:, -max_length:]\n attention_mask = attention_mask[:, -max_length:]\n elif truncation == \"right\":\n input_ids = input_ids[:, :max_length]\n attention_mask = attention_mask[:, :max_length]\n elif truncation == \"middle\":\n left_half = max_length // 2\n right_half = max_length - left_half\n input_ids = torch.cat([input_ids[:, :left_half], input_ids[:, -right_half:]], dim=-1)\n attention_mask = torch.cat([attention_mask[:, :left_half], attention_mask[:, -right_half:]], dim=-1)\n elif truncation == \"error\":\n raise NotImplementedError(f\"{sequence_length=} is larger than {max_length=}\")\n else:\n raise NotImplementedError(f\"Unknown truncation method {truncation}\")\n\n return input_ids, attention_mask\n\n\ndef tokenize_and_postprocess_data(\n prompt: str, tokenizer: PreTrainedTokenizer, max_length: int, pad_token_id: int, left_pad=True, truncation=\"error\"\n):\n \"\"\"Tokenize text and process outputs to consistent tensor shapes.\n\n Args:\n prompt: Input text to tokenize\n tokenizer: HuggingFace tokenizer instance\n max_length: Target sequence length\n pad_token_id: Padding token ID\n left_pad: Pad left if True\n truncation: Truncation strategy (\"left\"/\"right\"/\"error\")\n\n Returns:\n Tuple of (input_ids, attention_mask) from postprocess_data\n \"\"\"\n input_data = tokenizer(prompt, return_tensors=\"pt\", add_special_tokens=False)\n input_ids = input_data[\"input_ids\"]\n attention_mask = input_data[\"attention_mask\"]\n\n return postprocess_data(input_ids, attention_mask, max_length, pad_token_id, left_pad, truncation)\n\n\ndef remove_pad_token(input_ids: torch.Tensor, attention_mask: torch.Tensor):\n \"\"\"Remove the pad token.\n\n Args:\n input_ids shape: [bs, seq_length]\n attention_mask shape: [bs, seq_length]\n Returns:\n no_padding_batch(List[List[int]]): contains the rmpad token ids per query.\n \"\"\"\n no_padding_batch = []\n for ids, mask in zip(input_ids, attention_mask, strict=True):\n no_padding_batch.append((ids[len(ids) - mask.sum() :]).cpu().numpy().tolist())\n return no_padding_batch\n\n\ndef log_probs_from_logits_response(input_ids, logits, response_length):\n \"\"\"Compute the response log_probs from full logits. Note that logits = model(input_ids)\n\n Args:\n input_ids: [batch_size, seqlen]\n logits: [batch_size, seqlen, vocab_size]\n\n Returns:\n response_log_prob:\n \"\"\"\n response_logits = logits[:, -response_length - 1 : -1]\n response = input_ids[:, -response_length:]\n response_log_prob = logprobs_from_logits(logits=response_logits, labels=response)\n return response_log_prob\n\n\ndef log_probs_from_logits_response_rmpad(input_ids, attention_mask, logits_rmpad, response_length):\n \"\"\"Compute the log_probs from logits with rmpad logits and pad input. Note that\n logits_rmpad = model(input_ids_rmpad). For each sentences, there is a shift between\n logits and input_ids.\n The reason for this function to is to compute logprobs_from_logits in rmpad mode because it is memory-intensive\n for large vocab_size\n\n Args:\n input_ids: [batch_size, seqlen]\n attention_mask: [batch_size, seqlen]\n logits_rmpad: [total_nnz, vocab_size]\n response_length: int\n \"\"\"\n from flash_attn.bert_padding import pad_input, unpad_input\n\n batch_size, seqlen = input_ids.shape\n input_ids_rmpad, indices, *_ = unpad_input(input_ids.unsqueeze(-1), attention_mask=attention_mask)\n input_ids_rmpad = input_ids_rmpad.squeeze(-1)\n input_ids_rmpad_rolled = torch.roll(input_ids_rmpad, shifts=-1, dims=0)\n full_log_probs_rmpad = logprobs_from_logits(logits=logits_rmpad, labels=input_ids_rmpad_rolled) # (total_nnz,)\n full_output = pad_input(\n hidden_states=full_log_probs_rmpad.unsqueeze(-1), indices=indices, batch=batch_size, seqlen=seqlen\n )\n output = full_output.squeeze(-1)[:, -response_length - 1 : -1] # [batch_size, response_length]\n return output\n\n\ndef log_probs_from_logits_all_rmpad(input_ids_rmpad, logits_rmpad, indices, batch_size, seqlen, response_length):\n \"\"\"Compute the log_probs from logits with rmpad input_ids and logits. Note that\n logits_rmpad = model(input_ids_rmpad). For each sentences, there is a shift between\n logits and input_ids.\n The reason for this function to is to compute logprobs_from_logits in rmpad mode because it is memory-intensive\n for large vocab_size\n\n Args:\n input_ids_rmpad: [1, total_nnz]\n logits_rmpad: [total_nnz, vocab_size]\n indices: [total_nnz]\n batch_size: int\n seqlen: int\n response_length: int\n \"\"\"\n if get_device_name() == \"cuda\":\n from flash_attn.bert_padding import pad_input\n elif get_device_name() == \"npu\":\n from verl.utils.attention_utils import pad_input\n\n input_ids_rmpad = input_ids_rmpad.transpose(0, 1) # transpose back to [total_nnz, 1]\n input_ids_rmpad = input_ids_rmpad.squeeze(-1)\n input_ids_rmpad_rolled = torch.roll(input_ids_rmpad, shifts=-1, dims=0)\n full_log_probs_rmpad = logprobs_from_logits(logits=logits_rmpad, labels=input_ids_rmpad_rolled) # (total_nnz,)\n full_output = pad_input(\n hidden_states=full_log_probs_rmpad.unsqueeze(-1), indices=indices, batch=batch_size, seqlen=seqlen\n )\n output = full_output.squeeze(-1)[:, -response_length - 1 : -1] # [batch_size, response_length]\n return output\n\n\ndef post_process_logits(input_ids, logits, temperature, top_k, top_p):\n if temperature != 1.0:\n logits = logits.div_(temperature) # inplace operation to avoid OOM\n # TODO: add them back\n # if top_k is not None and top_k > 0:\n # logits = TopKLogitsWarper(top_k=top_k)(input_ids, logits)\n # if top_p is not None and top_p < 1.0 and top_p > 0.0:\n # logits = TopPLogitsWarper(top_p=top_p)(input_ids, logits)\n return logits\n\n\ndef calculate_sum_pi_squared_from_logits(logits: torch.Tensor):\n \"\"\"\n Compute exact sum of squared probabilities from logits.\n Formula: Σπ² = exp(logsumexp(2*logits) - 2*logsumexp(logits))\n\n Used for optimal baseline variance reduction as described in\n \"What Matters for Model Merging at Scale?\" (arXiv:2410.03617)\n\n Args:\n logits: Logits tensor (..., vocab_size).\n\n Returns:\n Sum of squared probabilities tensor (...).\n \"\"\"\n return torch.exp(torch.logsumexp(2.0 * logits, dim=-1) - 2.0 * torch.logsumexp(logits, dim=-1))\n\n\n\"\"\"\nOptimizer related\n\"\"\"\n\n\ndef get_cosine_schedule_with_warmup(\n optimizer: Optimizer,\n num_warmup_steps: int,\n num_training_steps: int,\n min_lr_ratio: float = 0.0,\n num_cycles: float = 0.5,\n last_epoch: int = -1,\n init_lr_ratio: float = None,\n):\n \"\"\"\n Create a schedule with a learning rate that decreases following the values of the cosine function between the\n initial lr set in the optimizer to 0, after a warmup period during which it increases linearly between 0 and the\n initial lr set in the optimizer.\n Args:\n optimizer (:class:`~torch.optim.Optimizer`):\n The optimizer for which to schedule the learning rate.\n num_warmup_steps (:obj:`int`):\n The number of steps for the warmup phase.\n num_training_steps (:obj:`int`):\n The total number of training steps.\n min_lr_ratio (:obj:`float`, `optional`, defaults to 0.0):\n The minimum lr ratio w.r.t the maximum.\n num_cycles (:obj:`float`, `optional`, defaults to 0.5):\n The number of waves in the cosine schedule (the defaults is to just decrease from the max value to 0\n following a half-cosine).\n last_epoch (:obj:`int`, `optional`, defaults to -1):\n The index of the last epoch when resuming training.\n init_lr_ratio (:obj:`float`, `optional`, defaults to None):\n The initial lr ratio w.r.t the maximum.\n Return:\n :obj:`torch.optim.lr_scheduler.LambdaLR` with the appropriate schedule.\n \"\"\"\n min_lr_ratio = 0.0 if min_lr_ratio is None else min_lr_ratio\n assert min_lr_ratio >= 0 and min_lr_ratio <= 1.0\n coef = (1 - min_lr_ratio) * 0.5\n intercept = (1 + min_lr_ratio) * 0.5\n\n init_lr_ratio = 0.0 if init_lr_ratio is None else init_lr_ratio\n assert init_lr_ratio >= 0 and init_lr_ratio <= 1.0\n\n def lr_lambda(current_step):\n if current_step < num_warmup_steps:\n return init_lr_ratio + (1.0 - init_lr_ratio) * (float(current_step) / float(max(1, num_warmup_steps)))\n progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))\n x = math.cos(math.pi * float(num_cycles) * 2.0 * progress)\n return max(min_lr_ratio, x * coef + intercept)\n\n return LambdaLR(optimizer, lr_lambda, last_epoch)\n\n\ndef get_constant_schedule_with_warmup(\n optimizer: Optimizer,\n num_warmup_steps: int,\n last_epoch: int = -1,\n):\n \"\"\"\n Create a constant LR schedule with a linear warmup phase.\n\n Args:\n optimizer (Optimizer): Wrapped optimizer.\n num_warmup_steps (int): Number of steps to ramp up the LR from 0 to initial value.\n last_epoch (int, optional): The index of the last epoch when resuming training. Defaults to -1.\n\n Returns:\n LambdaLR: Scheduler that increases LR linearly during warmup, then holds it constant.\n \"\"\"\n\n def lr_lambda(current_step):\n if current_step < num_warmup_steps:\n return float(current_step) / float(max(1.0, num_warmup_steps))\n return 1.0\n\n return LambdaLR(optimizer, lr_lambda, last_epoch)\n\n\ndef prepare_decoder_attention_mask(attention_mask, input_shape, inputs_embeds):\n # create causal mask\n # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]\n combined_attention_mask = None\n if input_shape[-1] > 1:\n combined_attention_mask = _make_causal_mask(\n input_shape,\n inputs_embeds.dtype,\n device=inputs_embeds.device,\n )\n\n if attention_mask is not None:\n # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]\n expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to(\n inputs_embeds.device\n )\n combined_attention_mask = (\n expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask\n )\n\n return combined_attention_mask\n\n\n# Copied from transformers.models.bart.modeling_bart._make_causal_mask\ndef _make_causal_mask(input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device):\n \"\"\"\n Make causal mask used for bi-directional self-attention.\n \"\"\"\n bsz, tgt_len = input_ids_shape\n mask = torch.full((tgt_len, tgt_len), torch.finfo(dtype).min, device=device)\n mask_cond = torch.arange(mask.size(-1), device=device)\n mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)\n mask = mask.to(dtype)\n return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len)\n\n\n# Copied from transformers.models.bart.modeling_bart._expand_mask\ndef _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):\n \"\"\"\n Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.\n \"\"\"\n bsz, src_len = mask.size()\n tgt_len = tgt_len if tgt_len is not None else src_len\n\n expanded_mask = mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)\n\n inverted_mask = 1.0 - expanded_mask\n\n return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)\n\n\ndef get_unpad_data(attention_mask):\n seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)\n indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()\n max_seqlen_in_batch = seqlens_in_batch.max().item()\n cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))\n return (\n indices,\n cu_seqlens,\n max_seqlen_in_batch,\n )\n\n\ndef get_wsd_schedule_with_warmup(\n optimizer: Optimizer,\n num_warmup_steps: int,\n num_training_steps: int,\n min_lr_ratio: float = 0.0,\n num_cycles: float = 0.5,\n last_epoch: int = -1,\n stable_ratio: float = 0.9,\n):\n \"\"\"\n Create a Warmup-Stable-Decay learning rate scheduler.\n\n The schedule follows three phases:\n 1. Warmup: Learning rate increases linearly from 0 to the initial LR\n 2. Stable: Learning rate remains constant at the initial LR\n 3. Decay: Learning rate decreases following a cosine curve to min_lr_ratio * initial LR\n\n Args:\n optimizer (:class:`~torch.optim.Optimizer`):\n The optimizer for which to schedule the learning rate.\n num_warmup_steps (:obj:`int`):\n The number of steps for the warmup phase.\n num_training_steps (:obj:`int`):\n The total number of training steps.\n min_lr_ratio (:obj:`float`, `optional`, defaults to 0.0):\n The minimum learning rate ratio w.r.t the initial learning rate.\n num_cycles (:obj:`float`, `optional`, defaults to 0.5):\n The number of waves in the cosine schedule during decay phase.\n last_epoch (:obj:`int`, `optional`, defaults to -1):\n The index of the last epoch when resuming training.\n stable_ratio (:obj:`float`, `optional`, defaults to 0.0):\n The ratio of non-warmup steps that should maintain a constant learning rate.\n Set to 0.0 to behave exactly like cosine schedule.\n\n Return:\n :obj:`torch.optim.lr_scheduler.LambdaLR` with the appropriate schedule.\n \"\"\"\n remaining_steps = max(0, num_training_steps - num_warmup_steps)\n num_stable_steps = int(remaining_steps * stable_ratio)\n num_decay_steps = remaining_steps - num_stable_steps\n\n def lr_lambda(current_step):\n if current_step < num_warmup_steps:\n return float(current_step) / float(max(1, num_warmup_steps))\n if current_step < num_warmup_steps + num_stable_steps:\n return 1.0\n if current_step < num_training_steps:\n progress = float(current_step - num_warmup_steps - num_stable_steps) / float(max(1, num_decay_steps))\n value = max(0.0, 0.5 * (1.0 + math.cos(math.pi * float(num_cycles) * 2.0 * progress)))\n return (1.0 - min_lr_ratio) * value + min_lr_ratio\n return min_lr_ratio\n\n return LambdaLR(optimizer, lr_lambda, last_epoch)\n\n\n@contextmanager\ndef check_device_is_available():\n \"\"\"\n Some modules must be imported after CUDA is initialized. Such as sglang's sharding manager.\n\n This context manager checks if CUDA is available and raises an error if it is not.\n \"\"\"\n if not get_torch_device().is_available():\n raise RuntimeError(\"Device {} must be initialized before importing this module.\".format(get_device_name()))\n\n yield\n\n\ndef distributed_mean_max_min_std(local_tensor, compute_max=True, compute_min=True, compute_std=True):\n \"\"\"Compute distributed statistics across all processes.\n\n Args:\n local_tensor: Tensor containing local values\n compute_max: Include maximum value calculation\n compute_min: Include minimum value calculation\n compute_std: Include standard deviation calculation\n\n Returns:\n Tuple containing (mean, max, min, std) in this order. None for disabled metrics.\n \"\"\"\n # Sum the local tensor across all processes\n local_sum = torch.sum(local_tensor)\n local_num = torch.tensor(torch.numel(local_tensor), device=get_device_name())\n\n torch.distributed.all_reduce(local_sum, op=torch.distributed.ReduceOp.SUM)\n torch.distributed.all_reduce(local_num, op=torch.distributed.ReduceOp.SUM)\n\n global_mean = local_sum / local_num\n\n if compute_max:\n local_max = torch.max(local_tensor)\n torch.distributed.all_reduce(local_max, op=torch.distributed.ReduceOp.MAX)\n else:\n local_max = None\n\n if compute_min:\n local_min = torch.min(local_tensor)\n torch.distributed.all_reduce(local_min, op=torch.distributed.ReduceOp.MIN)\n else:\n local_min = None\n\n if compute_std:\n square_diff = torch.sum(torch.pow(local_tensor - global_mean, 2))\n torch.distributed.all_reduce(square_diff, op=torch.distributed.ReduceOp.SUM)\n global_std = torch.sqrt(square_diff / (local_num - 1))\n else:\n global_std = None\n\n return global_mean, local_max, local_min, global_std\n\n\ndef distributed_masked_mean(local_tensor, local_mask):\n \"\"\"Compute global mean of non-masked elements across distributed processes.\n\n Args:\n local_tensor (torch.Tensor): Input tensor with local values\n local_mask (torch.Tensor): Binary mask (1=valid, 0=ignore) matching local_tensor shape\n\n Returns:\n torch.Tensor: Global mean of all valid elements across processes\n \"\"\"\n local_tensor = local_tensor * local_mask\n\n local_sum = torch.sum(local_tensor)\n local_num = torch.sum(local_mask)\n\n torch.distributed.all_reduce(local_sum, op=torch.distributed.ReduceOp.SUM)\n torch.distributed.all_reduce(local_num, op=torch.distributed.ReduceOp.SUM)\n\n global_mean = local_sum / local_num\n return global_mean\n\n\ndef expand_as_nested(tensor: torch.Tensor, nested_tensor: torch.Tensor) -> torch.Tensor:\n \"\"\"\n\n Args:\n tensor: a tensor with shape (bsz,)\n nested_tensor: a nested tensor with shape (bsz, xxx)\n\n Returns:\n a tensor with the same shape as nested_tensor\n\n \"\"\"\n assert nested_tensor.is_nested, \"nested_tensor must be nested\"\n assert tensor.shape[0] == nested_tensor.shape[0], (\n f\"The batch shape must be the same. Got {tensor.shape[0]} vs {nested_tensor.shape[0]}\"\n )\n assert len(tensor.shape) == 1, \"The ndim of tensor must be 1\"\n assert len(nested_tensor.shape) == 2, \"The ndim of nested_tensor must be 2\"\n\n offsets = nested_tensor.offsets()\n seqlens = offsets.diff()\n output = torch.repeat_interleave(tensor, seqlens, dim=0)\n output = torch.nested.nested_tensor_from_jagged(values=output, offsets=offsets)\n return output\n\n\n@contextmanager\ndef use_original_torch_compile():\n \"\"\"torch.compile might be replaced by mindspeed on NPU, this contextmanager\n can revert torch.compile temporarily.\n \"\"\"\n try:\n from mindspeed.patch_utils import MindSpeedPatchesManager\n\n compile_patch = None\n for patch in MindSpeedPatchesManager.patches_info.values():\n if patch.orig_module_name == \"torch\" and patch.orig_func_name == \"compile\":\n if patch.is_applied():\n compile_patch = patch\n break\n if compile_patch is not None:\n compile_patch.remove_patch()\n yield\n compile_patch.apply_patch()\n else:\n yield\n except Exception:\n yield\n"}127{"file_name": "verl__utils__transferqueue_utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport asyncio\nimport functools\nimport inspect\nimport logging\nimport os\nimport threading\nfrom functools import wraps\nfrom typing import TYPE_CHECKING, Any, Callable\n\nif TYPE_CHECKING:\n from verl.single_controller.base.decorator import Dispatch\n\nfrom tensordict import TensorDict\n\ntry:\n from transfer_queue import (\n AsyncTransferQueueClient,\n BatchMeta,\n TransferQueueClient,\n )\n\nexcept ImportError:\n # TODO: Use a hacky workaround for ImportError since\n # transfer_queue isn't a default verl dependency.\n class BatchMeta:\n pass\n\n\nfrom verl.protocol import DataProto\n\nlogger = logging.getLogger(__name__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n_TRANSFER_QUEUE_CLIENT = None\n\nis_transferqueue_enabled = os.environ.get(\"TRANSFER_QUEUE_ENABLE\", False)\n\n\ndef create_transferqueue_client(\n client_id: str,\n config,\n sync: bool = False,\n) -> \"AsyncTransferQueueClient | TransferQueueClient\":\n global _TRANSFER_QUEUE_CLIENT\n if _TRANSFER_QUEUE_CLIENT is None:\n if sync:\n _TRANSFER_QUEUE_CLIENT = TransferQueueClient(client_id, config.controller_info)\n else:\n _TRANSFER_QUEUE_CLIENT = AsyncTransferQueueClient(client_id, config.controller_info)\n _TRANSFER_QUEUE_CLIENT.initialize_storage_manager(manager_type=config.storage_backend, config=config)\n\n return _TRANSFER_QUEUE_CLIENT\n\n\ndef get_transferqueue_client() -> \"AsyncTransferQueueClient | TransferQueueClient\":\n return _TRANSFER_QUEUE_CLIENT\n\n\n# TODO (TQ): verl will make all actor async, so this can be cleanup later.\ndef _run_async_in_temp_loop(async_func: Callable[..., Any], *args, **kwargs) -> Any:\n # Use a temporary event loop in a new thread because event\n # loop may already exist in server mode\n tmp_event_loop = asyncio.new_event_loop()\n thread = threading.Thread(\n target=tmp_event_loop.run_forever,\n name=\"batchmeta dataproto converter\",\n daemon=True,\n )\n\n def run_coroutine(coroutine):\n if not thread.is_alive():\n thread.start()\n future = asyncio.run_coroutine_threadsafe(coroutine, tmp_event_loop)\n return future.result()\n\n async def stop_loop():\n tmp_event_loop.stop()\n\n try:\n return run_coroutine(async_func(*args, **kwargs))\n finally:\n if thread.is_alive():\n asyncio.run_coroutine_threadsafe(stop_loop(), tmp_event_loop)\n thread.join()\n\n\ndef _find_batchmeta(*args, **kwargs):\n for arg in args:\n if isinstance(arg, BatchMeta):\n return arg\n for v in kwargs.values():\n if isinstance(v, BatchMeta):\n return v\n return None\n\n\nasync def _async_batchmeta_to_dataproto(batchmeta: \"BatchMeta\") -> DataProto:\n if batchmeta.samples == [] or batchmeta.samples is None:\n return DataProto(\n batch=TensorDict({}, batch_size=(0,)),\n non_tensor_batch={},\n meta_info=batchmeta.extra_info.copy(),\n )\n\n tensordict = await _TRANSFER_QUEUE_CLIENT.async_get_data(batchmeta)\n return DataProto.from_tensordict(tensordict, meta_info=batchmeta.extra_info.copy())\n\n\ndef _batchmeta_to_dataproto(batchmeta: \"BatchMeta\") -> DataProto:\n return _run_async_in_temp_loop(_async_batchmeta_to_dataproto, batchmeta)\n\n\nasync def _async_update_batchmeta_with_output(output: DataProto, batchmeta: \"BatchMeta\", func_name=None) -> \"BatchMeta\":\n pid = os.getpid()\n\n for k, v in output.meta_info.items():\n batchmeta.set_extra_info(k, v)\n\n if len(output) > 0:\n tensordict = output.to_tensordict()\n # pop meta_info\n for key in output.meta_info.keys():\n tensordict.pop(key)\n\n logger.info(\n f\"Task {func_name} (pid={pid}) putting output data to TransferQueue with \"\n f\"batch_size={tensordict.batch_size},\\n\"\n f\"tensordict keys={list(tensordict.keys())}\"\n )\n\n updated_batch_meta = await _TRANSFER_QUEUE_CLIENT.async_put(data=tensordict, metadata=batchmeta)\n return updated_batch_meta\n else:\n return batchmeta\n\n\ndef _update_batchmeta_with_output(output: DataProto, batchmeta: \"BatchMeta\", func_name=None) -> \"BatchMeta\":\n updated_batch_meta = _run_async_in_temp_loop(_async_update_batchmeta_with_output, output, batchmeta, func_name)\n return updated_batch_meta\n\n\ndef _compute_need_collect(dispatch_mode: \"dict | Dispatch\", args: list) -> bool:\n \"\"\"Compute whether data collection is needed for the current worker.\n\n This function determines whether the current worker should collect data based on\n the dispatch mode configuration and worker parameters. It's used to optimize\n distributed data collection by ensuring only the appropriate rank collects data.\n\n Args:\n dispatch_mode: Controls data collection logic for the current worker. Can be None,\n a Dispatch instance, or a dict with 'collect_fn' key. If None or Dispatch,\n always returns True (current worker should collect). If dict, checks\n collect_fn for lazy compute optimization.\n args: List of arguments passed to the function. Should contain a Worker instance\n as the first argument when using lazy compute mode.\n\n Returns:\n bool: True if data collection is needed, False otherwise.\n\n Note:\n Only checks worker attributes when dispatch_mode is a dict with 'collect_fn',\n the collect_fn is 'collect_lazy_compute_data_proto', and args[0] is a Worker.\n Otherwise, returns True. For the lazy compute case, checks the worker's\n data parallel rank for the mesh specified in collect_fn.args[0] to determine\n if this worker should collect data.\n \"\"\"\n from verl.single_controller.base.decorator import Dispatch\n from verl.single_controller.base.worker import Worker\n\n if dispatch_mode is None or isinstance(dispatch_mode, Dispatch):\n return True\n\n assert \"collect_fn\" in dispatch_mode.keys(), \"collect_fn should be in dispatch_mode.\"\n\n collect_fn = dispatch_mode[\"collect_fn\"]\n\n # Check if collect_fn is a functools.partial and handle gracefully\n if isinstance(collect_fn, functools.partial):\n collect_fn_name = collect_fn.func.__name__\n if collect_fn_name != \"collect_lazy_compute_data_proto\" or len(args) < 1 or not isinstance(args[0], Worker):\n return True\n\n collect_mesh_name = collect_fn.args[0] if collect_fn.args else None\n if collect_mesh_name is None:\n return True\n\n return args[0].query_collect_info(collect_mesh_name)\n else:\n # If collect_fn is not a partial, we can't extract mesh_name information\n # Fall back to default behavior (collect data)\n return True\n\n\ndef _postprocess_common(output, put_data, need_collect):\n \"\"\"Common post-processing logic for function outputs in TransferQueue bridge.\n\n This function handles the final return value based on whether data should be\n put into storage (put_data) and whether collection is needed (need_collect).\n It ensures proper return types based on the execution context.\n\n Args:\n output: The original output from the decorated function. Can be any type.\n put_data: bool, indicating whether the output should be put into TransferQueue.\n If True, output will be put to TQ and return the corresponding BatchMeta;\n if False, output will not be put into TQ.\n need_collect: bool, indicating whether this process needs to collect data.\n If False, the output will be replaced by an empty BatchMeta or DataProto\n to avoid redundant communication.\n\n Returns:\n - BatchMeta.empty(): When put_data=True but need_collect=False, indicating\n no data should be stored but BatchMeta structure is expected.\n - DataProto(): When put_data=False, need_collect=False, and output is DataProto,\n returning an empty DataProto.\n - output: In all other cases, returns the original output unchanged.\n\n Note:\n This function is used in the tqbridge decorator to normalize return values\n across different execution paths and avoid redundant data operations in\n distributed scenarios.\n \"\"\"\n if put_data and not need_collect:\n return BatchMeta.empty()\n elif not put_data and not need_collect and isinstance(output, DataProto):\n return DataProto()\n else:\n return output\n\n\ndef tqbridge(dispatch_mode: \"dict | Dispatch\" = None, put_data: bool = True):\n \"\"\"Creates a decorator for bridging BatchMeta and DataProto.\n\n This decorator automatically handles conversions between `BatchMeta` and\n `DataProto` in function parameters, and decides whether to sync function\n output back to `BatchMeta` based on configuration(`put_data`). It supports\n both synchronous and asynchronous functions (async def), and can control\n whether to enable enhanced logic via the global `HAS_TQ` variable (when disabled,\n simply calls the original function as-is).\n\n Args:\n dispatch_mode: Controls data collection behavior for the current worker. Passed to\n _compute_need_collect to determine if current worker should collect data.\n If None, _compute_need_collect will return True to fallback default logics.\n put_data: Whether put the DataProto into Storage after func return.\n If True, after function execution, the output result will be\n updated to `BatchMeta` and `BatchMeta` will be returned;\n If False, the function output result will be returned directly.\n Defaults to True.\n\n Returns:\n A decorator function used to decorate target functions (synchronous or asynchronous).\n \"\"\"\n\n def decorator(func):\n pid = os.getpid()\n\n @wraps(func)\n def inner(*args, **kwargs):\n batchmeta = _find_batchmeta(*args, **kwargs)\n if batchmeta is None:\n return func(*args, **kwargs)\n else:\n logger.info(\n f\"Task {func.__name__} (pid={pid}) is getting len_samples={batchmeta.size}, \"\n f\"global_idx={batchmeta.global_indexes}\"\n )\n args = [_batchmeta_to_dataproto(arg) if isinstance(arg, BatchMeta) else arg for arg in args]\n kwargs = {k: _batchmeta_to_dataproto(v) if isinstance(v, BatchMeta) else v for k, v in kwargs.items()}\n output = func(*args, **kwargs)\n need_collect = _compute_need_collect(dispatch_mode, args)\n if put_data and need_collect:\n updated_batch_meta = _update_batchmeta_with_output(output, batchmeta, func.__name__)\n return updated_batch_meta\n return _postprocess_common(output, put_data, need_collect)\n\n @wraps(func)\n async def async_inner(*args, **kwargs):\n batchmeta = _find_batchmeta(*args, **kwargs)\n if batchmeta is None:\n return await func(*args, **kwargs)\n else:\n logger.info(\n f\"Task {func.__name__} (pid={pid}) is getting len_samples={batchmeta.size}, \"\n f\"global_idx={batchmeta.global_indexes}\"\n )\n args = [await _async_batchmeta_to_dataproto(arg) if isinstance(arg, BatchMeta) else arg for arg in args]\n kwargs = {\n k: await _async_batchmeta_to_dataproto(v) if isinstance(v, BatchMeta) else v\n for k, v in kwargs.items()\n }\n output = await func(*args, **kwargs)\n need_collect = _compute_need_collect(dispatch_mode, args)\n if put_data and need_collect:\n updated_batchmeta = await _async_update_batchmeta_with_output(output, batchmeta, func.__name__)\n return updated_batchmeta\n return _postprocess_common(output, put_data, need_collect)\n\n @wraps(func)\n def dummy_inner(*args, **kwargs):\n output = func(*args, **kwargs)\n return output\n\n @wraps(func)\n async def dummy_async_inner(*args, **kwargs):\n output = await func(*args, **kwargs)\n return output\n\n wrapper_inner = inner if is_transferqueue_enabled else dummy_inner\n wrapper_async_inner = async_inner if is_transferqueue_enabled else dummy_async_inner\n\n wrapper = wrapper_async_inner if inspect.iscoroutinefunction(func) else wrapper_inner\n return wrapper\n\n return decorator\n"}128{"file_name": "verl__utils__ulysses.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nUtilities for DeepSpeed Ulysses Sequence Parallelism.\nDeepSpeed Ulysses Paper: https://arxiv.org/abs/2309.14509\nInspired from: https://github.com/deepspeedai/DeepSpeed/blob/master/deepspeed/sequence/layer.py\n\"\"\"\n\nfrom typing import Any, Optional\n\nimport torch\nimport torch.distributed as dist\nfrom torch import Tensor\nfrom torch.distributed import ProcessGroup\n\n_ULYSSES_SEQUENCE_PARALLEL_GROUP = None\n\n\ndef set_ulysses_sequence_parallel_group(group: dist.ProcessGroup):\n \"\"\"\n Set ulysses sequence parallel process group.\n \"\"\"\n global _ULYSSES_SEQUENCE_PARALLEL_GROUP\n _ULYSSES_SEQUENCE_PARALLEL_GROUP = group\n\n\ndef get_ulysses_sequence_parallel_group() -> Optional[dist.ProcessGroup]:\n \"\"\"\n Get ulysses sequence parallel process group.\n \"\"\"\n global _ULYSSES_SEQUENCE_PARALLEL_GROUP\n return _ULYSSES_SEQUENCE_PARALLEL_GROUP\n\n\ndef get_ulysses_sequence_parallel_world_size(group: ProcessGroup = None) -> int:\n \"\"\"\n Get ulysses sequence parallel world size.\n \"\"\"\n group = get_ulysses_sequence_parallel_group() if group is None else group\n return dist.get_world_size(group) if group else 1\n\n\ndef get_ulysses_sequence_parallel_rank(group: ProcessGroup = None) -> int:\n \"\"\"\n Get ulysses sequence parallel rank.\n \"\"\"\n group = get_ulysses_sequence_parallel_group() if group is None else group\n return dist.get_rank(group) if group else 0\n\n\ndef gather_seq_scatter_heads(\n x: Tensor,\n seq_dim: int,\n head_dim: int,\n unpadded_dim_size: int = 0,\n group: ProcessGroup = None,\n) -> Tensor:\n \"\"\"\n A func to sync embedding input with alltoall in sequence parallel\n gather sequence dimension and scatter head dim:\n e.g. seq_dim: 1, head_dim: 2\n [bsz, seq/n, h, ...] -> [bsz, seq, h/n, ...]\n \"\"\"\n group = get_ulysses_sequence_parallel_group() if group is None else group\n if not group:\n return x\n sp_world = get_ulysses_sequence_parallel_world_size(group)\n x = SeqAllToAll.apply(group, x, head_dim, seq_dim)\n if unpadded_dim_size and unpadded_dim_size % sp_world != 0:\n padding_size = x.size(seq_dim) - unpadded_dim_size\n x = _unpad_tensor(x, seq_dim, padding_size)\n return x\n\n\ndef gather_heads_scatter_seq(x: Tensor, head_dim: int, seq_dim: int, group: ProcessGroup = None) -> Tensor:\n \"\"\"\n A func to sync attention result with alltoall in sequence parallel\n gather head dimension and scatter seq dim:\n e.g. seq_dim: 1, head_dim: 2\n [bsz, seq, h/n, ...] -> [bsz, seq/n, h, ...]\n \"\"\"\n group = get_ulysses_sequence_parallel_group() if group is None else group\n if not group:\n return x\n dim_size = x.size(seq_dim)\n sp_world = get_ulysses_sequence_parallel_world_size(group)\n if dim_size % sp_world != 0:\n padding_size = sp_world - (dim_size % sp_world)\n x = _pad_tensor(x, seq_dim, padding_size)\n return SeqAllToAll.apply(group, x, seq_dim, head_dim, False)\n\n\ndef _pad_tensor(x: Tensor, dim: int, padding_size: int) -> Tensor:\n shape = list(x.shape)\n shape[dim] = padding_size\n pad = torch.zeros(shape, dtype=x.dtype, device=x.device)\n return torch.cat([x, pad], dim=dim)\n\n\ndef _unpad_tensor(x: Tensor, dim: int, padding_size: int) -> Tensor:\n slc = [slice(None)] * len(x.shape)\n slc[dim] = slice(0, -padding_size)\n return x[tuple(slc)]\n\n\ndef slice_input_tensor(x: Tensor, dim: int, padding: bool = True, group: ProcessGroup = None) -> Tensor:\n group = get_ulysses_sequence_parallel_group() if group is None else group\n sp_world_size = dist.get_world_size(group)\n sp_rank = get_ulysses_sequence_parallel_rank()\n dim_size = x.size(dim)\n # pad before slice\n if padding and dim_size % sp_world_size:\n padding_size = sp_world_size - (dim_size % sp_world_size)\n x = _pad_tensor(x, dim, padding_size)\n # slice the input tensor\n parts = x.size(dim) // sp_world_size\n slc = [slice(None)] * len(x.shape)\n slc[dim] = slice(sp_rank * parts, (sp_rank + 1) * parts)\n return x[tuple(slc)].contiguous()\n\n\ndef all_to_all_tensor(\n local_input: Tensor,\n scatter_dim: int,\n gather_dim: int,\n group: Optional[dist.ProcessGroup] = None,\n async_op: bool = False,\n):\n group = get_ulysses_sequence_parallel_group() if group is None else group\n seq_world_size = dist.get_world_size(group)\n input_list = [t.contiguous() for t in torch.tensor_split(local_input, seq_world_size, scatter_dim)]\n output_list = [torch.empty_like(input_list[0]) for _ in range(seq_world_size)]\n comm = dist.all_to_all(output_list, input_list, group=group, async_op=async_op)\n if async_op:\n\n def wait():\n comm.wait()\n return torch.cat(output_list, dim=gather_dim).contiguous()\n\n return wait\n return torch.cat(output_list, dim=gather_dim).contiguous()\n\n\ndef all_gather_tensor(local_tensor: Tensor, group: Optional[dist.ProcessGroup] = None, async_op: bool = False):\n group = get_ulysses_sequence_parallel_group() if group is None else group\n sp_world_size = dist.get_world_size(group=group)\n output_shape = list(local_tensor.shape)\n output_shape[0] = output_shape[0] * sp_world_size\n output = torch.empty(output_shape, dtype=local_tensor.dtype, device=local_tensor.device)\n dist.all_gather_into_tensor(output, local_tensor, group=group, async_op=async_op)\n return output\n\n\nclass SeqAllToAll(torch.autograd.Function):\n @staticmethod\n def forward(\n ctx: Any,\n group: dist.ProcessGroup,\n local_input: Tensor,\n scatter_dim: int,\n gather_dim: int,\n async_op: bool = False,\n ) -> Tensor:\n ctx.group = group\n ctx.scatter_dim = scatter_dim\n ctx.gather_dim = gather_dim\n ctx.async_op = async_op\n return all_to_all_tensor(local_input, scatter_dim, gather_dim, group, async_op)\n\n @staticmethod\n def backward(ctx: Any, *grad_output: Tensor) -> tuple[None, Tensor, None, None]:\n input_t = torch.cat(grad_output[1:], dim=ctx.gather_dim).contiguous() if ctx.async_op else grad_output[0]\n return (\n None,\n all_to_all_tensor(input_t, ctx.gather_dim, ctx.scatter_dim, ctx.group, False),\n None,\n None,\n None,\n None,\n )\n\n\nclass Gather(torch.autograd.Function):\n @staticmethod\n def forward(\n ctx: Any,\n group: dist.ProcessGroup,\n local_tensor: Tensor,\n gather_dim: int,\n grad_scaler: bool = True,\n async_op=False,\n ) -> Tensor:\n ctx.group = group\n ctx.gather_dim = gather_dim\n ctx.grad_scaler = grad_scaler\n ctx.async_op = async_op\n\n sp_world_size = dist.get_world_size(group=group)\n ctx.sp_world_size = sp_world_size\n\n sp_rank = dist.get_rank(group=group)\n ctx.sp_rank = sp_rank\n\n local_shape = list(local_tensor.size())\n split_size = local_shape[0]\n part_size = local_shape[gather_dim] # store original size\n ctx.part_size = part_size\n\n output = all_gather_tensor(local_tensor, group, async_op)\n return torch.cat(output.split(split_size, dim=0), dim=gather_dim)\n\n @staticmethod\n def backward(ctx: Any, grad_output: Tensor) -> Any:\n if ctx.grad_scaler:\n grad_output = grad_output * ctx.sp_world_size\n return (\n None,\n grad_output.split(ctx.part_size, dim=ctx.gather_dim)[ctx.sp_rank].contiguous(),\n None,\n None,\n None,\n None,\n )\n\n\ndef gather_outpus_and_unpad(*args, **kwargs):\n raise RuntimeError(\n \"please use verl.utils.ulysses.gather_outputs_and_unpad instead of verl.utils.ulysses.gather_outpus_and_unpad\"\n )\n\n\ndef gather_outputs_and_unpad(\n x: Tensor,\n gather_dim: int,\n unpad_dim: int = None,\n padding_size: int = 0,\n grad_scaler: bool = True,\n group: Optional[dist.ProcessGroup] = None,\n):\n \"\"\"\n Gather a tensor across a process group and optionally unpad its padded elements.\n\n Args:\n x (Tensor): Input tensor to gather.\n gather_dim (int): Dimension along which to gather across ranks.\n unpad_dim (int, optional): Dimension from which to remove padding. If None, no unpadding.\n padding_size (int): Number of padding elements to remove on `unpad_dim`. Defaults to 0.\n grad_scaler (bool): Whether to apply gradient scaling during gather. Defaults to True.\n group (ProcessGroup, optional): Process group for gathering. If None, uses\n `get_ulysses_sequence_parallel_group()`. If still None, returns `x` unchanged.\n\n Returns:\n Tensor: The gathered tensor, with padding removed if requested.\n \"\"\"\n group = get_ulysses_sequence_parallel_group() if group is None else group\n if group is None:\n return x\n x = Gather.apply(group, x, gather_dim, grad_scaler)\n if unpad_dim is not None:\n assert isinstance(padding_size, int), \"padding size is not given or is not an integer\"\n if padding_size == 0:\n return x\n x = _unpad_tensor(x, unpad_dim, padding_size)\n return x\n\n\ndef ulysses_pad(\n input_ids_rmpad: torch.Tensor, position_ids_rmpad: Optional[torch.Tensor] = None, sp_size: int = 1, pad_value=0\n):\n if position_ids_rmpad is not None:\n assert position_ids_rmpad.size(-2) == 1\n assert input_ids_rmpad.size(-1) == position_ids_rmpad.size(-1)\n if sp_size <= 1:\n return input_ids_rmpad, position_ids_rmpad, 0\n _, total_seq_len = input_ids_rmpad.shape\n pad_size = (sp_size - total_seq_len % sp_size) % sp_size\n if pad_size > 0:\n input_ids_rmpad = torch.nn.functional.pad(input_ids_rmpad, (0, pad_size), value=pad_value)\n if position_ids_rmpad is not None:\n pad_pos_ids = torch.arange(pad_size, device=position_ids_rmpad.device).unsqueeze(0)\n if position_ids_rmpad.dim() == 3:\n pad_pos_ids = pad_pos_ids.unsqueeze(0).repeat(position_ids_rmpad.size(0), 1, 1)\n position_ids_rmpad = torch.cat((position_ids_rmpad, pad_pos_ids), dim=-1)\n return input_ids_rmpad, position_ids_rmpad, pad_size\n\n\ndef ulysses_pad_and_slice_inputs(\n input_ids_rmpad: torch.Tensor,\n position_ids_rmpad: Optional[torch.Tensor] = None,\n sp_size: int = 1,\n skip_position_ids_rmpad: bool = False,\n pad_value=0,\n):\n \"\"\"\n Pad and slice input_ids to be divisible by sp_size\n Pad position_ids to be divisible by sp_size.\n\n Note both input_ids_rmpad and position_ids_rmpad will be padded and sliced.\n\n The is the utility of pre-forward for ulysses sequence parallelism\n\n Args:\n input_ids_rmpad: shape of [bsz, seqlen]\n position_ids_rmpad: shape of [bsz, seqlen], where bsz must be 1\n sp_size (int): ulysses sequence parallelism size\n skip_position_ids_rmpad: whether to skip position_ids_rmpad for VeOmniEngine\n\n Returns:\n torch.Tensor: padded and sliced input_ids\n torch.Tensor: padded and sliced position_ids\n int: pad size\n \"\"\"\n input_ids_rmpad, position_ids_rmpad, pad_size = ulysses_pad(\n input_ids_rmpad, position_ids_rmpad, sp_size, pad_value=pad_value\n )\n input_ids_rmpad = slice_input_tensor(input_ids_rmpad, dim=1, padding=False)\n if position_ids_rmpad is not None and not skip_position_ids_rmpad:\n position_ids_rmpad = slice_input_tensor(position_ids_rmpad, dim=1, padding=False)\n return input_ids_rmpad, position_ids_rmpad, pad_size\n\n\ndef validate_ulysses_config(num_heads, ulysses_sequence_size):\n if ulysses_sequence_size > 1:\n assert num_heads % ulysses_sequence_size == 0, (\n f\"num_heads ({num_heads}) must be divisible by ulysses sequence size({ulysses_sequence_size})\"\n )\n"}129{"file_name": "verl__utils__vllm__utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nfrom msgspec import field\nfrom packaging import version as vs\n\ntry:\n from vllm.lora.lora_model import LoRAModel\nexcept ImportError:\n from vllm.lora.models import LoRAModel\n\nfrom vllm.lora.request import LoRARequest\nfrom vllm.lora.utils import get_adapter_absolute_path\nfrom vllm.lora.worker_manager import LRUCacheWorkerLoRAManager\n\nfrom verl.third_party.vllm import get_version\n\n\nclass TensorLoRARequest(LoRARequest):\n peft_config: dict = field(default=None)\n lora_tensors: dict = field(default=None)\n\n\nclass VLLMHijack:\n @staticmethod\n def hijack():\n def hijack__load_adapter(self, lora_request: TensorLoRARequest) -> LoRAModel:\n \"\"\"\n based on vllm.lora.worker_manager.WorkerLoRAManager._load_adapter, support load adapter with lora tensors\n\n Reason:\n VLLM does not support adding LoRA from tensors directly. It only supports adding LoRA via file paths.\n To synchronize the LoRA tensors of the actor model, we need to find a workaround to enable VLLM to\n load memory-based LoRA tensors.\n \"\"\"\n try:\n supported_lora_modules = self._adapter_manager.supported_lora_modules\n packed_modules_mapping = self._adapter_manager.packed_modules_mapping\n expected_lora_modules: list[str] = []\n for module in supported_lora_modules:\n if module in packed_modules_mapping:\n expected_lora_modules.extend(packed_modules_mapping[module])\n else:\n expected_lora_modules.append(module)\n\n expected_lora_modules = list(set(expected_lora_modules))\n\n lora_tensors = None\n from vllm.lora.peft_helper import PEFTHelper\n\n if isinstance(lora_request, TensorLoRARequest):\n peft_config = lora_request.peft_config\n lora_tensors = lora_request.lora_tensors\n peft_helper = PEFTHelper.from_dict(peft_config)\n else:\n lora_path = get_adapter_absolute_path(lora_request.lora_path)\n\n peft_helper = PEFTHelper.from_local_dir(lora_path, self.max_position_embeddings)\n\n # Validates the LoRA configuration against requirements before\n # loading weights, throwing an exception if validation fails.\n peft_helper.validate_legal(self.lora_config)\n\n # For some models like Qwen2VL, we need to use hf_to_vllm_mapper\n # to ensure correct loading of lora weights.\n model = self._adapter_manager.model\n hf_to_vllm_mapper = None\n if hasattr(model, \"hf_to_vllm_mapper\") and model.hf_to_vllm_mapper is not None:\n hf_to_vllm_mapper = model.hf_to_vllm_mapper\n\n lora_request_kwargs = {\n \"peft_helper\": peft_helper,\n \"lora_model_id\": lora_request.lora_int_id,\n \"device\": \"cpu\",\n \"dtype\": self.lora_config.lora_dtype,\n \"weights_mapper\": hf_to_vllm_mapper,\n }\n if hasattr(self, \"embedding_padding_modules\"):\n lora_request_kwargs[\"embedding_modules\"] = self.embedding_modules\n lora_request_kwargs[\"embedding_padding_modules\"] = self.embedding_padding_modules\n else:\n lora_request_kwargs[\"model_vocab_size\"] = self.vocab_size\n if hasattr(self.lora_config, \"lora_extra_vocab_size\"):\n lora_request_kwargs[\"target_embedding_padding\"] = (\n self.vocab_size + self.lora_config.lora_extra_vocab_size\n )\n if isinstance(lora_request, TensorLoRARequest):\n lora = self._lora_model_cls.from_lora_tensors(\n tensors=lora_tensors,\n **lora_request_kwargs,\n )\n else:\n lora = self._lora_model_cls.from_local_checkpoint(\n lora_path,\n expected_lora_modules,\n **lora_request_kwargs,\n )\n except Exception:\n raise\n\n if getattr(lora, \"extra_vocab_size\", 0) > getattr(self.lora_config, \"lora_extra_vocab_size\", 0):\n raise ValueError(\n f\"LoRA added vocab size {lora.extra_vocab_size} is greater than lora_extra_vocab_size \"\n f\"{self.lora_config.lora_extra_vocab_size}.\"\n )\n return lora\n\n def do_hijack(target_cls, target_method_name, hooking_method):\n setattr(target_cls, target_method_name, hooking_method)\n\n do_hijack(LRUCacheWorkerLoRAManager, \"_load_adapter\", hijack__load_adapter)\n\n\ndef is_version_ge(pkg: str = \"vllm\", minver: str = \"0.7.3\"):\n \"\"\"check if the package version is greater than or equal to the minimum version\"\"\"\n return vs.parse(get_version(pkg)) >= vs.parse(minver)\n"}130{"file_name": "verl__utils__vllm__vllm_fp8_utils.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nfrom dataclasses import dataclass, field\nfrom unittest.mock import patch\n\nimport torch\nimport vllm\nfrom packaging import version\n\ntry:\n from vllm.model_executor.layers.fused_moe.layer import FusedMoE\n from vllm.model_executor.layers.linear import LinearBase\nexcept ImportError as e:\n raise ImportError(\"FP8 quantization not available\") from e\n\nfrom verl.utils.kernel.fp8_kernel import scaled_fp8_blockwise\n\nlogger = logging.getLogger(__name__)\n\n\n# Ref: https://github.com/NVIDIA-NeMo/RL/commit/bc24887c72a6e1b2699a228bc87c588546dfe6b7\n@dataclass()\nclass FP8State:\n # A cache of fp8 parameter names, we can check this cache to see if a\n # param name corresponds to a fp8 weight\n seen_params: set = field(default_factory=lambda: set())\n fp8_param_names: set = field(default_factory=lambda: set())\n vllm_patches: list = field(default_factory=lambda: [])\n\n\nfp8_state: FP8State = FP8State()\n\n\ndef is_fp8_model(vllm_config):\n from vllm.model_executor.layers.quantization.fp8 import Fp8Config\n\n if hasattr(vllm_config, \"quant_config\") and isinstance(vllm_config.quant_config, Fp8Config):\n return True\n\n return False\n\n\ndef get_module_from_param_name(model, name: str):\n # Split the name into parts (e.g., 'layers', '0', 'self_attn', 'q_proj', 'weight')\n # The module path is all but the last part (the parameter's own name)\n path_parts = name.split(\".\")\n module_path = path_parts[:-1]\n # Replace with the fused model name\n packed_modules_mapping = model.packed_modules_mapping\n reversed_mapping = {\n original_name: fused_name\n for fused_name, original_names_list in packed_modules_mapping.items()\n for original_name in original_names_list\n }\n if module_path[-1] in reversed_mapping.keys():\n module_path[-1] = reversed_mapping[module_path[-1]]\n\n current_module = model\n try:\n # Traverse the model hierarchy\n for part in module_path:\n if isinstance(current_module, FusedMoE):\n return current_module\n elif isinstance(current_module, torch.nn.ModuleList):\n current_module = current_module[int(part)]\n else:\n current_module = getattr(current_module, part)\n except (AttributeError, IndexError, ValueError) as e:\n print(f\"Warning: Could not find module for parameter '{name}'. Error: {e}\")\n return current_module\n\n\ndef is_fp8_weight(name, model):\n if name not in fp8_state.seen_params:\n fp8_state.seen_params.add(name)\n # Filter out bias params\n if name.endswith(\"weight\"):\n module = get_module_from_param_name(model, name)\n # We currently only quantize linear layers\n\n if (isinstance(module, LinearBase) and module.weight.dtype == torch.float8_e4m3fn) or (\n isinstance(module, FusedMoE)\n and module.w13_weight.dtype == torch.float8_e4m3fn\n and module.w2_weight.dtype == torch.float8_e4m3fn\n ):\n fp8_state.fp8_param_names.add(name)\n return name in fp8_state.fp8_param_names\n\n\ndef quant_weights(weights, model, quant_config, dtype=torch.bfloat16):\n \"\"\"Quantize weights to FP8 format using a memory-efficient generator.\n\n\n Args:\n weights: Generator or iterable of (name, tensor) pairs\n model: The model to check for FP8 weight names\n quant_config: Quantization configuration with weight_block_size\n dtype: Data type for intermediate computation (default: bfloat16)\n\n Yields:\n Tuples of (name, tensor) for each weight and its scale\n \"\"\"\n if quant_config.weight_block_size is None:\n raise ValueError(\"Currently only support blockwise quantization, please set weight_block_size in quant_config\")\n\n is_vllm_11_or_later = version.parse(vllm.__version__) >= version.parse(\"0.11.0\")\n\n for k, v in weights:\n if not is_fp8_weight(k, model):\n yield (k, v)\n continue\n\n # Cast the weight into fp8 and its scale factor\n if torch.distributed.get_rank() == 0:\n logger.debug(f\"Quantizing to FP8 blockwise: {k}\")\n\n param_lp, param_scale = scaled_fp8_blockwise(\n v.to(dtype),\n weight_block_size=quant_config.weight_block_size,\n )\n param_scale = param_scale.squeeze(-1)\n\n # Yield the quantized weight\n yield (k, param_lp)\n\n # Yield the scale with appropriate naming based on vLLM version\n if is_vllm_11_or_later:\n if \"expert\" in k:\n yield (k + \"_scale_inv\", param_scale)\n else:\n yield (k + \"_scale\", param_scale)\n else:\n yield (k + \"_scale_inv\", param_scale)\n\n # Explicitly delete original tensor reference to help GC\n del v, param_lp, param_scale\n\n\ndef load_quanted_weights(weights, model_runner):\n model = model_runner.model\n quant_config = model_runner.vllm_config.quant_config\n vllm_dtype = model_runner.vllm_config.model_config.dtype\n\n weights_quantized = quant_weights(weights, model, quant_config, dtype=vllm_dtype)\n\n # Monkey patch the param class to their subclass, as certain models\n # will check the param type to call the proper weightloader\n for name, param in model.named_parameters():\n if hasattr(param, \"subclass_type\"):\n param.orig_type = param.__class__\n param.__class__ = param.subclass_type\n # Finally load the weights into vllm\n loaded_params = model.load_weights(weights_quantized)\n # Undo the type change above to the original type\n for name, param in model.named_parameters():\n if hasattr(param, \"subclass_type\"):\n param.__class__ = param.orig_type\n return loaded_params\n\n\ndef process_weights_after_loading_for_vllm10(self, layer) -> None:\n \"\"\"This function is used to process the weights after loading for a Linear layer, it is used for vllm v0.10\n\n Compared to the original process_weights_after_loading in vllm, we just avoid creation of\n new torch.nn.Parameter objects, because that removes the weight_loader attribute which we need for refit.\n \"\"\"\n logger.debug(\"Applying patch process_weights_after_loading\")\n try:\n from vllm.model_executor.parameter import (\n BlockQuantScaleParameter,\n ModelWeightParameter,\n )\n except Exception:\n print(\"error\")\n from torch.nn import Parameter\n\n def _create_param_from_subclass_attributes(custom_param):\n param = Parameter(custom_param.data, requires_grad=False)\n base_param_dir = dir(torch.nn.Parameter)\n custom_param_dir = dir(custom_param)\n # Find the attributes that are unique to the custom parameter\n custom_attributes = [\n attr for attr in custom_param_dir if attr not in base_param_dir and not attr.startswith(\"__\")\n ]\n # Set the custom attributes into the base parameter object\n for attr in custom_attributes:\n setattr(param, attr, getattr(custom_param, attr))\n\n param.subclass_type = type(custom_param)\n return param\n\n assert self.block_quant and self.quant_config.is_checkpoint_fp8_serialized\n assert self.quant_config.activation_scheme == \"dynamic\"\n weight = layer.weight.data\n weight_scale_inv = layer.weight_scale_inv.data\n weight = self._maybe_pad_weight(weight)\n\n layer.weight = _create_param_from_subclass_attributes(\n ModelWeightParameter(\n data=weight,\n output_dim=0,\n input_dim=1,\n weight_loader=layer.weight.weight_loader,\n )\n )\n layer.weight_scale_inv = _create_param_from_subclass_attributes(\n BlockQuantScaleParameter(\n data=weight_scale_inv,\n output_dim=0,\n input_dim=1,\n weight_loader=layer.weight_scale_inv.weight_loader,\n )\n )\n\n\ndef process_weights_after_loading_for_vllm11(self, layer) -> None:\n \"\"\"This function is used to process the weights after loading for a Linear layer, it is used for vllm 0.11\n\n Compared to the original process_weights_after_loading in vllm, we just avoid creation of\n new torch.nn.Parameter objects, because that removes the weight_loader attribute which we need for refit.\n \"\"\"\n from torch.nn import Parameter\n from vllm.model_executor.layers.quantization.utils.fp8_utils import (\n maybe_post_process_fp8_weight_block,\n process_fp8_weight_block_strategy,\n )\n from vllm.model_executor.parameter import (\n BlockQuantScaleParameter,\n ModelWeightParameter,\n )\n\n assert self.block_quant and self.quant_config.is_checkpoint_fp8_serialized\n assert self.quant_config.activation_scheme == \"dynamic\"\n\n def _create_param_from_subclass_attributes(custom_param):\n param = Parameter(custom_param.data, requires_grad=False)\n base_param_dir = dir(torch.nn.Parameter)\n custom_param_dir = dir(custom_param)\n # Find the attributes that are unique to the custom parameter\n custom_attributes = [\n attr for attr in custom_param_dir if attr not in base_param_dir and not attr.startswith(\"__\")\n ]\n # Set the custom attributes into the base parameter object\n for attr in custom_attributes:\n setattr(param, attr, getattr(custom_param, attr))\n\n param.subclass_type = type(custom_param)\n return param\n\n weight_scale = layer.weight_scale_inv if hasattr(layer, \"weight_scale_inv\") else layer.weight_scale\n weight, weight_scale = process_fp8_weight_block_strategy(layer.weight, weight_scale)\n\n layer.weight = _create_param_from_subclass_attributes(\n ModelWeightParameter(\n data=weight.data,\n output_dim=0,\n input_dim=1,\n weight_loader=layer.weight.weight_loader,\n )\n )\n layer.weight_scale = _create_param_from_subclass_attributes(\n BlockQuantScaleParameter(\n data=weight_scale.data,\n output_dim=0,\n input_dim=1,\n weight_loader=layer.weight_scale_inv.weight_loader,\n )\n )\n\n del layer.weight_scale_inv\n\n if version.parse(vllm.__version__) == version.parse(\"0.11.0\"):\n maybe_post_process_fp8_weight_block(layer, self.cutlass_block_fp8_supported)\n else:\n maybe_post_process_fp8_weight_block(layer)\n\n\ndef process_weights_after_loading_moe_for_vllm10(self, layer) -> None:\n \"\"\"This function is used to process the weights after loading for a FusedMoE layer, it is used for vllm v0.10\"\"\"\n from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import is_rocm_aiter_moe_enabled\n from vllm.model_executor.layers.quantization.fp8 import _is_col_major, _swap_w13_to_w31\n from vllm.model_executor.layers.quantization.utils.fp8_utils import (\n get_col_major_tma_aligned_tensor,\n requant_weight_ue8m0_inplace,\n )\n from vllm.utils.deep_gemm import is_blackwell_deep_gemm_used\n\n self.rocm_aiter_moe_enabled = is_rocm_aiter_moe_enabled()\n assert self.quant_config.activation_scheme == \"dynamic\"\n if self.flashinfer_moe_enabled:\n w13_weight = _swap_w13_to_w31(layer.w13_weight.data)\n w13_weight_scale_inv = _swap_w13_to_w31(layer.w13_weight_scale_inv.data)\n w2_weight = layer.w2_weight.data\n w2_weight_scale_inv = layer.w2_weight_scale_inv.data\n else:\n w13_weight = layer.w13_weight.data\n w13_weight_scale_inv = layer.w13_weight_scale_inv.data\n w2_weight = layer.w2_weight\n w2_weight_scale_inv = layer.w2_weight_scale_inv\n\n from torch.nn import Parameter\n\n def _create_param_from_subclass_attributes(custom_data, custom_weight):\n param = Parameter(custom_data, requires_grad=False)\n base_param_dir = dir(torch.nn.Parameter)\n custom_weight_dir = dir(custom_weight)\n # Find the attributes that are unique to the custom parameter\n custom_attributes = [\n attr for attr in custom_weight_dir if attr not in base_param_dir and not attr.startswith(\"__\")\n ]\n # Set the custom attributes into the base parameter object\n for attr in custom_attributes:\n setattr(param, attr, getattr(custom_weight, attr))\n\n return param\n\n layer.w13_weight = _create_param_from_subclass_attributes(w13_weight, layer.w13_weight)\n layer.w13_weight_scale_inv = _create_param_from_subclass_attributes(\n w13_weight_scale_inv, layer.w13_weight_scale_inv\n )\n layer.w2_weight = _create_param_from_subclass_attributes(w2_weight, layer.w2_weight)\n layer.w2_weight_scale_inv = _create_param_from_subclass_attributes(w2_weight_scale_inv, layer.w2_weight_scale_inv)\n\n # DeepGemm scales need to be transposed and aligned. We try to do\n # it ahead of time for performance reasons.\n if self.allow_deep_gemm and not is_blackwell_deep_gemm_used():\n # Lazy import to avoid CUDA initialization problems.\n if _is_col_major(layer.w13_weight_scale_inv):\n layer.w13_weight_scale_inv = get_col_major_tma_aligned_tensor(layer.w13_weight_scale_inv).contiguous()\n if _is_col_major(layer.w2_weight_scale_inv):\n layer.w2_weight_scale_inv = get_col_major_tma_aligned_tensor(layer.w2_weight_scale_inv).contiguous()\n\n if is_blackwell_deep_gemm_used():\n assert layer.weight_block_size is not None\n # Re-quantise the expert weights so their scales are UE8M0.\n block_sz = tuple(layer.weight_block_size)\n requant_weight_ue8m0_inplace(\n layer.w13_weight.data,\n layer.w13_weight_scale_inv.data,\n block_sz,\n )\n requant_weight_ue8m0_inplace(\n layer.w2_weight.data,\n layer.w2_weight_scale_inv.data,\n block_sz,\n )\n\n if _is_col_major(layer.w13_weight_scale_inv):\n layer.w13_weight_scale_inv = get_col_major_tma_aligned_tensor(layer.w13_weight_scale_inv).contiguous()\n if _is_col_major(layer.w2_weight_scale_inv):\n layer.w2_weight_scale_inv = get_col_major_tma_aligned_tensor(layer.w2_weight_scale_inv).contiguous()\n\n\ndef process_weights_after_loading_moe_for_vllm11(self, layer) -> None:\n \"\"\"This function is used to process the weights after loading for a FusedMoE layer, it is used for vllm 0.11\"\"\"\n from vllm.model_executor.layers.quantization.utils.flashinfer_utils import (\n swap_w13_to_w31,\n )\n from vllm.model_executor.layers.quantization.utils.fp8_utils import (\n expert_weight_is_col_major,\n requant_weight_ue8m0_inplace,\n )\n from vllm.utils.deep_gemm import (\n get_col_major_tma_aligned_tensor,\n is_deep_gemm_e8m0_used,\n )\n\n try:\n from vllm.model_executor.layers.fused_moe.rocm_aiter_fused_moe import is_rocm_aiter_moe_enabled\n\n self.rocm_aiter_moe_enabled = is_rocm_aiter_moe_enabled()\n except ImportError:\n from vllm._aiter_ops import rocm_aiter_ops\n\n self.rocm_aiter_moe_enabled = rocm_aiter_ops.is_fused_moe_enabled()\n\n assert self.block_quant and self.quant_config.is_checkpoint_fp8_serialized\n assert self.quant_config.activation_scheme == \"dynamic\"\n\n if self.flashinfer_moe_backend is not None:\n layer.w13_weight.data = swap_w13_to_w31(layer.w13_weight.data)\n layer.w13_weight_scale_inv.data = swap_w13_to_w31(layer.w13_weight_scale_inv.data)\n\n if self.allow_deep_gemm and not is_deep_gemm_e8m0_used():\n if expert_weight_is_col_major(layer.w13_weight_scale_inv):\n layer.w13_weight_scale_inv = get_col_major_tma_aligned_tensor(layer.w13_weight_scale_inv)\n if expert_weight_is_col_major(layer.w2_weight_scale_inv):\n layer.w2_weight_scale_inv = get_col_major_tma_aligned_tensor(layer.w2_weight_scale_inv)\n\n if is_deep_gemm_e8m0_used():\n assert layer.weight_block_size is not None\n # Re-quantise the expert weights so their scales are UE8M0.\n block_sz = tuple(layer.weight_block_size)\n requant_weight_ue8m0_inplace(\n layer.w13_weight.data,\n layer.w13_weight_scale_inv.data,\n block_sz,\n )\n requant_weight_ue8m0_inplace(\n layer.w2_weight.data,\n layer.w2_weight_scale_inv.data,\n block_sz,\n )\n\n # Ensure column-major TMA alignment expected by DeepGEMM.\n if expert_weight_is_col_major(layer.w13_weight_scale_inv):\n layer.w13_weight_scale_inv = get_col_major_tma_aligned_tensor(layer.w13_weight_scale_inv)\n if expert_weight_is_col_major(layer.w2_weight_scale_inv):\n layer.w2_weight_scale_inv = get_col_major_tma_aligned_tensor(layer.w2_weight_scale_inv)\n\n\ndef apply_vllm_fp8_patches():\n logger.info(\"Applying vllm fp8 patches for blockwise quantization\")\n func1_path = \"vllm.model_executor.layers.quantization.fp8.Fp8LinearMethod.process_weights_after_loading\"\n patcher1 = patch(\n func1_path,\n process_weights_after_loading_for_vllm11\n if version.parse(vllm.__version__) >= version.parse(\"0.11.0\")\n else process_weights_after_loading_for_vllm10,\n )\n patcher1.start()\n func2_path = \"vllm.model_executor.layers.quantization.fp8.Fp8MoEMethod.process_weights_after_loading\"\n patcher2 = patch(\n func2_path,\n process_weights_after_loading_moe_for_vllm11\n if version.parse(vllm.__version__) >= version.parse(\"0.11.0\")\n else process_weights_after_loading_moe_for_vllm10,\n )\n patcher2.start()\n"}131{"file_name": "verl__workers__actor__base.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nThe base class for Actor\n\"\"\"\n\nfrom abc import ABC, abstractmethod\n\nimport torch\n\nfrom verl import DataProto\n\n__all__ = [\"BasePPOActor\"]\n\n\nclass BasePPOActor(ABC):\n def __init__(self, config):\n \"\"\"The base class for PPO actor\n\n Args:\n config (DictConfig): a config passed to the PPOActor. We expect the type to be\n DictConfig (https://omegaconf.readthedocs.io/), but it can be any namedtuple in general.\n \"\"\"\n super().__init__()\n self.config = config\n\n @abstractmethod\n def compute_log_prob(self, data: DataProto) -> torch.Tensor:\n \"\"\"Compute logits given a batch of data.\n\n Args:\n data (DataProto): a batch of data represented by DataProto. It must contain key ```input_ids```,\n ```attention_mask``` and ```position_ids```.\n\n Returns:\n DataProto: a DataProto containing the key ```log_probs```\n\n\n \"\"\"\n pass\n\n @abstractmethod\n def update_policy(self, data: DataProto) -> dict:\n \"\"\"Update the policy with an iterator of DataProto\n\n Args:\n data (DataProto): an iterator over the DataProto that returns by\n ```make_minibatch_iterator```\n\n Returns:\n Dict: a dictionary contains anything. Typically, it contains the statistics during updating the model\n such as ```loss```, ```grad_norm```, etc,.\n\n \"\"\"\n pass\n"}132{"file_name": "verl__workers__actor__megatron_actor.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nMegatron Actor.\nIn megatron actor, the differences are:\n1. We only make minibatch\n\nNote that our model doesn't have to be `MegatronModule` because we don't share embedding in the last layer\n\"\"\"\n\nimport itertools\nimport logging\nimport os\nfrom functools import partial\nfrom typing import Iterable\n\nimport torch\nimport torch.distributed\nfrom megatron.core import parallel_state as mpu\nfrom megatron.core.distributed import finalize_model_grads\n\n# from megatron.core.optimizer import DistributedOptimizer\nfrom megatron.core.optimizer import DistributedOptimizer\nfrom megatron.core.pipeline_parallel import get_forward_backward_func\nfrom omegaconf import OmegaConf\nfrom torch import nn\n\nfrom verl import DataProto\nfrom verl.trainer.ppo.core_algos import agg_loss, get_policy_loss_fn, kl_penalty\nfrom verl.utils.device import get_device_id, get_torch_device\nfrom verl.utils.megatron.pipeline_parallel import make_batch_generator\nfrom verl.utils.megatron.router_replay_patch import RouterReplay, RouterReplayAction\nfrom verl.utils.megatron.router_replay_utils import (\n RouterReplayHelper,\n merge_router_topk_indices,\n pp_gather,\n reorder_and_merge_vpp_layers,\n set_router_replay_data,\n)\nfrom verl.utils.megatron.tensor_parallel import vocab_parallel_entropy, vocab_parallel_log_probs_from_logits\nfrom verl.utils.megatron_utils import get_megatron_mtp_loss, get_model_config, unwrap_model\nfrom verl.utils.profiler import GPUMemoryLogger\nfrom verl.utils.py_functional import append_to_dict\nfrom verl.utils.seqlen_balancing import get_reverse_idx, rearrange_micro_batches\nfrom verl.utils.torch_functional import broadcast_dict_tensor\nfrom verl.workers.actor import BasePPOActor\nfrom verl.workers.config import MtpConfig\n\n__all__ = [\"MegatronPPOActor\"]\n\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\nclass MegatronPPOActor(BasePPOActor):\n def __init__(\n self,\n config,\n model_config,\n hf_config,\n tf_config,\n actor_module: nn.ModuleList,\n actor_optimizer: DistributedOptimizer,\n mtp_config: MtpConfig = None,\n ):\n \"\"\"MeagtronPPOActor class. This class implements the simple PPO logics when the model is built with Megatron.\n\n Args:\n config (OmegaConf): the basic config that contains the hyper-parameters of PPO Actor. It must contain\n\n ``ppo_micro_batch_size_per_gpu``: micro batch size when updating ppo.\n\n ``ppo_mini_batch_size``: minibatch size when updating ppo using the batch data.\n\n ``ppo_epochs``: number of epochs to update the actor using the batch data.\n\n ``shuffle``: whether to shuffle the data after each ppo epoch.\n\n ``clip_ratio``: clip ratio of the ppo algorithm. See https://arxiv.org/abs/1707.06347.\n\n ``entropy_coeff``: entropy coefficient of the PPO loss. See https://arxiv.org/abs/1707.06347.\n model_config (OmegaConf): model configuration. It must contains ``model_config.vocab_size`` and\n ``model_config.hidden_size``\n hf_config (PretrainedConfig): huggingface config\n tf_config (TransformerConfig): mcore transformer config\n mtp_config (MtpConfig): mtp config, default None\n actor_module (nn.ModuleList): actor module is a ModuleList that contains a list of nn.Module in this\n pp stage.\n each nn.Module in this rank holds a vpp module chunk. See https://arxiv.org/pdf/2104.04473.pdf for\n more details.\n The actor module has some constraints to follow in order to use the updating logics implemented here\n\n 1. It must implement unpad_input before any computation and pad_input after all the computation.\n Remove padding is an\n optimization that removes the padding tokens. See unpad_input and pad_input function in flash-attn\n (https://github.com/Dao-AILab/flash-attention/blob/main/flash_attn/bert_padding.py).\n\n 2. Each pp stage must return the hidden state with the same shape [total_nnz, 1, hidden_size],\n where total_nnz is the number of valid tokens in this batch. If sequence parallel is enabled, the size\n of the hidden state is [total_nnz // tp, 1, hidden_size].\n actor_optimizer (DistributedOptimizer): currently, we only support DistributedOptimizer in Megatron.\n It implements\n zero1 optimizer that shards the optimizer state across dp ranks.\n\n >>> from megatron.training import get_model\n >>> from megatron.optimizer import get_megatron_optimizer\n >>> actor_module = get_model(megatron_actor_model_provider, wrap_with_ddp=True)\n >>> actor_module = nn.ModuleList(actor_module)\n >>> actor_optimizer = get_megatron_optimizer(actor_module)\n >>> actor = MegatronPPOActor(config=config,\n >>> model_config=actor_model_config,\n >>> hf_config=hf_config,\n >>> tf_config=tf_config,\n >>> actor_module=actor_module,\n >>> actor_optimizer=actor_optimizer)\n \"\"\"\n super().__init__(config)\n self._validate_config(config)\n self.model_config = model_config\n self.hf_config = hf_config\n self.tf_config = tf_config\n self.mtp_config = mtp_config\n self.actor_module = actor_module\n self.actor_optimizer: DistributedOptimizer = actor_optimizer\n\n if self.mtp_config:\n assert self.mtp_config.enable, \"MTP requires mtp_config.enable to be True\"\n\n self.use_fused_kernels = self.config.get(\"use_fused_kernels\", False)\n if self.use_fused_kernels and not getattr(self.config, \"overlap_moe_expert_parallel_comm\", False):\n # do not patch if overlap_moe_expert_parallel_comm is enabled\n logger.warning_once(\n \"Recommend to disable use_fused_kernels since the fused kernel's performance is broken for triton>=3.3\"\n \"Unless you are using a very old version of triton < 3.3\"\n )\n from verl.models.mcore.model_forward_fused import patch_fused_forward\n\n for model in self.actor_module:\n patch_fused_forward(model)\n else:\n from verl.models.mcore.mtp_patch import patch_postprocess\n\n for model in self.actor_module:\n if self.mtp_config:\n from verl.models.mcore.mtp_patch import patch_mtp_layer_get_embeddings\n\n patch_postprocess(model)\n\n if self.mtp_config.detach_encoder:\n patch_mtp_layer_get_embeddings(model)\n\n self.optimizer_step_args = OmegaConf.create(\n {\n \"skip_grad\": None,\n \"overlap_dp_param_comm\": False,\n \"overlap_dp_grad_comm\": False,\n \"gradient_accumulation_steps\": 1,\n \"sequence_parallel\": self.tf_config.sequence_parallel,\n \"DDP_impl\": \"local\",\n \"layernorm_allreduce_bucket_threshold\": 0,\n \"reduce_grads_use_alltoall\": False,\n }\n )\n\n self.router_replay = self.config.router_replay\n self.enable_routing_replay = self.router_replay.mode != \"disabled\"\n if self.enable_routing_replay:\n self.mini_layer_topk_idx_list = []\n\n config = get_model_config(self.actor_module[0])\n print(config)\n config.finalize_model_grads_func = finalize_model_grads\n\n def _validate_config(self, config) -> None:\n \"\"\"Validate config options not implemented for Megatron backend\"\"\"\n assert config.get(\"ulysses_sequence_parallel_size\", 1) == 1\n if config.get(\"shuffle\", False):\n assert config.data_loader_seed is not None, \"If shuffle dataloader, seed must be manually set\"\n if config.megatron.tensor_model_parallel_size == 1:\n print(\"[Warining] Because actor tp size == 1, set sp to False\")\n config.megatron.sequence_parallel = False\n self.config = config\n\n @GPUMemoryLogger(role=\"megatron actor\", logger=logger)\n def compute_log_prob(self, data: DataProto, calculate_entropy=False) -> torch.Tensor:\n \"\"\"Compute the log probability of the responses given input_ids, attention_mask and position_ids\n\n Args:\n data (DataProto): a DataProto containing keys\n\n ``input_ids``: tensor of shape [batch_size, sequence_length]. torch.int64. Note that input_ids is the\n concatenation of prompt and response. Note that ``sequence_length = prompt_length + response_length``.\n\n ``attention_mask``: tensor of shape [batch_size, sequence_length]. torch.int64.\n\n ``position_ids``: tensor of shape [batch_size, sequence_length]. torch.int64.\n\n ``responses``: tensor of shape [batch_size, response_length]. torch.int64.\n\n Returns:\n DataProto: torch.Tensor: the log_prob tensor\n \"\"\"\n prev_modes = [m.training for m in self.actor_module]\n for module in self.actor_module:\n module.eval()\n use_dynamic_bsz = data.meta_info.get(\"use_dynamic_bsz\", False)\n micro_batch_size = data.meta_info.get(\"micro_batch_size\", None)\n max_token_len = data.meta_info.get(\"max_token_len\", None)\n if use_dynamic_bsz:\n assert max_token_len is not None, \"max_token_len must be set when use_dynamic_bsz is True\"\n max_token_len = max_token_len * self.config.megatron.context_parallel_size\n else:\n assert micro_batch_size is not None, (\n \"micro batch size is needed for forward compute when use_dynamic_bsz is False\"\n )\n\n def compute_logprobs_fn(output, data, use_dynamic_bsz=False, indices=None):\n response = data[\"responses\"]\n response_length = response.size(1)\n log_probs = output[\"log_probs\"][:, -response_length - 1 : -1].contiguous()\n return {\"log_probs\": log_probs}\n\n # We make recompute_old_log_prob by default here.\n # TODO (zhangchi.usc1992): actually, this function should only return log_prob and this logic should be\n # handled by user outside\n recompute_old_log_prob = self.config.get(\"recompute_old_log_prob\", True)\n\n entropys = torch.Tensor()\n if recompute_old_log_prob:\n select_keys = [\"responses\", \"input_ids\", \"attention_mask\", \"position_ids\"]\n\n if self.enable_routing_replay and self.config.router_replay.mode == \"R3\":\n assert \"routed_experts\" in data.batch.keys(), \"routed_experts must be in data.batch.keys()\"\n select_keys.append(\"routed_experts\")\n\n batch = data.select(batch_keys=select_keys).batch\n input_ids = batch[\"input_ids\"]\n batch_size = input_ids.size(0)\n response = batch[\"responses\"]\n response_length = response.size(1)\n with torch.no_grad():\n output = self.forward_backward_batch(\n data,\n forward_only=True,\n post_process_fn=compute_logprobs_fn,\n calculate_entropy=calculate_entropy,\n use_dynamic_bsz=use_dynamic_bsz,\n micro_batch_size=micro_batch_size,\n max_token_len=max_token_len,\n )\n if mpu.is_pipeline_last_stage(ignore_virtual=True):\n # only on last rank. It should be on every tp rank\n if calculate_entropy:\n log_probs = [o[0][\"log_probs\"] for o in output[\"output\"]] # (bs, seq_size)\n else:\n log_probs = [o[\"log_probs\"] for o in output[\"output\"]] # (bs, seq_size)\n log_probs = torch.cat(log_probs, dim=0).to(torch.float32)\n if use_dynamic_bsz:\n indices = output[\"indices\"]\n indices = list(itertools.chain.from_iterable(indices))\n assert len(indices) == log_probs.size(0), f\"{len(indices)} vs. {log_probs.size()}\"\n revert_indices = torch.tensor(get_reverse_idx(indices), dtype=torch.long)\n log_probs = log_probs[revert_indices]\n else:\n log_probs = torch.empty(\n size=(batch_size, response_length), dtype=torch.float32, device=input_ids.device\n )\n log_probs = log_probs.to(get_device_id())\n # broadcast across pp ranks\n torch.distributed.broadcast(\n tensor=log_probs,\n src=mpu.get_pipeline_model_parallel_last_rank(),\n group=mpu.get_pipeline_model_parallel_group(),\n async_op=False,\n )\n log_probs = log_probs.to(\"cpu\")\n if calculate_entropy:\n # Note that o[0] is metrics, o[1] is entropy\n if mpu.is_pipeline_last_stage(ignore_virtual=True):\n entropys = torch.cat([o[1] for o in output[\"output\"]], dim=0)\n entropys = entropys.to(torch.float32)\n if use_dynamic_bsz:\n indices = output[\"indices\"]\n indices = list(itertools.chain.from_iterable(indices))\n assert len(indices) == entropys.size(0), f\"{len(indices)} vs. {entropys.size()}\"\n revert_indices = torch.tensor(get_reverse_idx(indices), dtype=torch.long)\n entropys = entropys[revert_indices]\n else:\n entropys = torch.empty(\n size=(batch_size, response_length), dtype=torch.float32, device=input_ids.device\n )\n # broadcast across pp ranks\n entropys = entropys.to(get_device_id())\n torch.distributed.broadcast(\n tensor=entropys,\n src=mpu.get_pipeline_model_parallel_last_rank(),\n group=mpu.get_pipeline_model_parallel_group(),\n async_op=False,\n )\n entropys = entropys.to(\"cpu\")\n layers_topk_idx = None\n\n if RouterReplayHelper.is_r2_record_action(self.tf_config):\n # (bs, max_seq_len/response_len,local_layer_num,topk)\n layers_topk_idx = output[\"mini_layer_topk_idx_tensor\"].to(torch.uint8)\n if use_dynamic_bsz:\n indices = output[\"indices\"]\n indices = list(itertools.chain.from_iterable(indices))\n assert len(indices) == layers_topk_idx.size(0), f\"{len(indices)} vs. {layers_topk_idx.size()}\"\n revert_indices = torch.tensor(get_reverse_idx(indices), dtype=torch.long)\n layers_topk_idx = layers_topk_idx[revert_indices]\n layers_topk_idx = pp_gather(layers_topk_idx, self.tf_config)\n # add empty cache after each compute\n get_torch_device().empty_cache()\n\n for module, mode in zip(self.actor_module, prev_modes, strict=False):\n module.train(mode)\n return log_probs, entropys, layers_topk_idx\n\n def make_minibatch_iterator(self, data: DataProto) -> Iterable[DataProto]:\n \"\"\"Make minibatch iterator for updating the actor\n\n Args:\n data (DataProto): a DataProto containing keys\n\n ``input_ids``: tensor of shape [batch_size, sequence_length]. torch.int64, where\n ``sequence_length = prompt_length + response_length``\n\n ``attention_mask``: tensor of shape [batch_size, sequence_length]. torch.int64\n\n ``position_ids``: tensor of shape [batch_size, sequence_length]. torch.int64\n\n ``responses``: tensor of shape [batch_size, response_length]. torch.int64. Note that\n responses = input_ids[:, -response_length:]\n\n ``old_log_probs``: tensor of shape [batch_size, response_length]. torch.float32. The log probability\n of responses.\n\n ``advantages``: tensor of shape [batch_size, response_length]. torch.float32. The advantages of\n responses.\n See PPO paper for details. https://arxiv.org/abs/1707.06347\n\n Returns:\n\n \"\"\"\n select_keys = [\n \"responses\",\n \"input_ids\",\n \"attention_mask\",\n \"response_mask\",\n \"position_ids\",\n \"old_log_probs\",\n \"advantages\",\n ]\n if self.config.use_kl_loss:\n select_keys.append(\"ref_log_prob\")\n # Include pre-computed IS weights if present in batch\n # Weights are computed centrally in trainer and added to batch when algorithm.rollout_is=True\n if \"rollout_is_weights\" in data.batch.keys():\n select_keys.append(\"rollout_is_weights\")\n # Include rollout_log_probs for computing rollout_corr metrics in bypass mode\n if \"rollout_log_probs\" in data.batch.keys():\n select_keys.append(\"rollout_log_probs\")\n self.has_multi_modal_inputs = \"multi_modal_inputs\" in data.non_tensor_batch.keys()\n # router replay\n if self.enable_routing_replay:\n select_keys.append(\"routed_experts\")\n if self.has_multi_modal_inputs:\n data = data.select(select_keys, [\"multi_modal_inputs\"])\n else:\n data = data.select(batch_keys=select_keys)\n\n return data.make_iterator(\n mini_batch_size=self.config.ppo_mini_batch_size,\n epochs=self.config.ppo_epochs,\n seed=self.config.data_loader_seed,\n dataloader_kwargs={\"shuffle\": self.config.shuffle},\n )\n\n def forward_backward_batch(\n self,\n data: DataProto,\n forward_only=False,\n post_process_fn=None,\n calculate_entropy=False,\n use_dynamic_bsz=False,\n micro_batch_size=None,\n max_token_len=None,\n mini_batch_size=None,\n ):\n \"\"\"\n We assume:\n - The model takes input: (input_ids, attention_mask, position_ids). No rmpad for the input\n - The communication shape is (total_nnz_pad_to_sp // tp_size, 1, hidden_size) if sequence parallel is enabled\n \"\"\"\n # broadcast from last pp rank to all other pp ranks\n # TODO: actually, we just need to control the sampling order.\n data.to(get_device_id())\n data.batch = data.batch.contiguous()\n mini_batch = data\n broadcast_dict_tensor(\n mini_batch.batch,\n src=mpu.get_pipeline_model_parallel_last_rank(),\n group=mpu.get_pipeline_model_parallel_group(),\n )\n mini_batch.to(\"cpu\")\n # split into micro-batches\n mini_batch.batch[\"attention_mask\"] = mini_batch.batch[\"attention_mask\"].to(bool)\n self.has_multi_modal_inputs = \"multi_modal_inputs\" in mini_batch.non_tensor_batch.keys()\n if self.has_multi_modal_inputs:\n mini_batch.batch[\"multi_modal_inputs\"] = mini_batch.non_tensor_batch[\"multi_modal_inputs\"]\n mini_batch.batch[\"multi_modal_inputs_idx\"] = torch.Tensor(\n list(range(len(mini_batch.non_tensor_batch[\"multi_modal_inputs\"])))\n ).to(torch.int64)\n\n if mini_batch.batch[\"position_ids\"].dim() == 3: # qwen2vl mrope [bs, 3, seq_len]\n mini_batch.batch[\"position_ids\"] = mini_batch.batch[\"position_ids\"][\n :, 0\n ] # mcore patch recompute qwen2vl's pos ids during forward\n\n indices = None\n temperature = data.meta_info[\"temperature\"]\n if use_dynamic_bsz:\n assert max_token_len is not None, \"max_token_len must be set when use_dynamic_bsz is True\"\n vpp_size = mpu.get_virtual_pipeline_model_parallel_world_size()\n if vpp_size is not None and vpp_size > 1:\n microbatch_group_size_per_vp_stage = self.tf_config.microbatch_group_size_per_vp_stage\n micro_batches, indices = rearrange_micro_batches(\n batch=mini_batch.batch,\n num_batches_divided_by=microbatch_group_size_per_vp_stage,\n max_token_len=max_token_len,\n )\n assert len(micro_batches) % self.tf_config.microbatch_group_size_per_vp_stage == 0, (\n f\"micro_batches {micro_batches} must be divisible by microbatch_group_size_per_vp_stage \"\n f\"{microbatch_group_size_per_vp_stage} for megatron backend\"\n )\n else:\n micro_batches, indices = rearrange_micro_batches(batch=mini_batch.batch, max_token_len=max_token_len)\n total_seqlen = max_token_len\n else:\n assert micro_batch_size is not None, (\n \"micro_batch_size is needed to be passed in when not using dynamic batch size\"\n )\n micro_batches = mini_batch.batch.split(micro_batch_size)\n seq_len = micro_batches[0][\"input_ids\"].shape[1]\n total_seqlen = micro_batch_size * seq_len\n # compute input shapes for pp stages\n n_micro_batch = len(micro_batches)\n\n forward_backward_func = get_forward_backward_func()\n\n def loss_func(output, data, meta_info):\n # For memory efficiency\n # We move calculation of entropy to compute_log_probs, forward_only == True\n log_probs = None\n entropy = None\n if isinstance(output, dict):\n log_probs = output[\"log_probs\"]\n if \"entropy\" in output:\n entropy = output[\"entropy\"]\n else:\n assert isinstance(output, torch.Tensor)\n log_probs = output\n\n device = log_probs.device\n metrics = {}\n if forward_only:\n if post_process_fn is None:\n pass\n # metrics[\"logits\"] = output\n else:\n stats = post_process_fn(output, data)\n metrics.update(stats)\n if not calculate_entropy:\n return torch.tensor(1.0, device=device), metrics\n\n responses = data[\"responses\"]\n response_length = responses.size(1)\n response_mask = data[\"response_mask\"].to(bool)\n loss_agg_mode = self.config.loss_agg_mode\n # compute policy loss\n log_prob = log_probs[:, -response_length - 1 : -1].contiguous()\n ret_entropy = None\n stats = {}\n if not forward_only:\n old_log_prob = data[\"old_log_probs\"]\n advantages = data[\"advantages\"]\n\n entropy_coeff = self.config.entropy_coeff\n loss_agg_mode = self.config.loss_agg_mode\n\n loss_mode = self.config.policy_loss.get(\"loss_mode\", \"vanilla\")\n\n policy_loss_fn = get_policy_loss_fn(loss_mode)\n\n # Extract pre-computed rollout correction weights if present\n # Weights are computed centrally in trainer and added when algorithm.rollout_is=True\n rollout_is_weights = data.get(\"rollout_is_weights\", None)\n pg_loss, pg_metrics = policy_loss_fn(\n old_log_prob=old_log_prob,\n log_prob=log_prob,\n advantages=advantages,\n response_mask=response_mask,\n loss_agg_mode=loss_agg_mode,\n config=self.config,\n rollout_is_weights=rollout_is_weights,\n )\n stats.update(pg_metrics)\n\n # Skip if using bypass_mode loss (metrics already computed in pg_metrics)\n rollout_log_prob = data.get(\"rollout_log_probs\", None)\n if loss_mode != \"bypass_mode\" and rollout_log_prob is not None:\n # Compute metrics using CURRENT policy π_θ vs π_rollout\n # Tracks evolving off-policy gap as π_θ updates during mini-batch training\n from verl.trainer.ppo.rollout_corr_helper import compute_rollout_corr_metrics_from_logprobs\n\n rollout_corr_metrics = compute_rollout_corr_metrics_from_logprobs(\n log_prob=log_prob,\n rollout_log_prob=rollout_log_prob,\n response_mask=response_mask,\n )\n stats.update(rollout_corr_metrics)\n\n stats[\"actor/pg_loss\"] = pg_loss.detach().item()\n policy_loss = pg_loss\n\n if calculate_entropy:\n entropy = output[\"entropy\"][:, -response_length - 1 : -1].contiguous()\n if not forward_only:\n entropy_loss = agg_loss(loss_mat=entropy, loss_mask=response_mask, loss_agg_mode=loss_agg_mode)\n entropy_coeff = meta_info[\"entropy_coeff\"]\n policy_loss = pg_loss - entropy_coeff * entropy_loss\n else:\n ret_entropy = entropy\n\n if forward_only:\n policy_loss = torch.tensor(1.0, device=device)\n else:\n if self.config.use_kl_loss:\n ref_log_prob = data[\"ref_log_prob\"]\n # compute kl loss\n kld = kl_penalty(logprob=log_prob, ref_logprob=ref_log_prob, kl_penalty=self.config.kl_loss_type)\n kl_loss = agg_loss(loss_mat=kld, loss_mask=response_mask, loss_agg_mode=self.config.loss_agg_mode)\n\n policy_loss = policy_loss + kl_loss * self.config.kl_loss_coef\n metrics[\"actor/kl_loss\"] = kl_loss.detach().item()\n metrics[\"actor/kl_coef\"] = self.config.kl_loss_coef\n\n # return loss and stats\n\n append_to_dict(metrics, stats)\n return policy_loss, [metrics, ret_entropy]\n\n def forward_step(batch_iter, model, return_schedule_plan: bool = False):\n \"\"\"\n Args:\n batch_iter: the batch iterator\n model: the model\n return_schedule_plan: whether to return the schedule plan, for 1f1b overlap\n \"\"\"\n if return_schedule_plan:\n assert self.tf_config.overlap_moe_expert_parallel_comm, (\n \"overlap_moe_expert_parallel_comm must be enabled to return the schedule plan\"\n )\n # TODO: Fix this\n assert not calculate_entropy, \"calculate_entropy must be disabled to return the schedule plan\"\n from megatron.core.models.gpt.gpt_model import GPTModel\n\n assert isinstance(model, GPTModel), \"model must be a GPTModel\"\n assert self.use_fused_kernels, \"use_fused_kernels must be enabled to return the schedule plan\"\n # TODO: support VLM with MoE\n from verl.models.mcore.model_forward_1f1b_overlap import gptmodel_forward_1f1b_overlap\n\n batch = next(batch_iter)\n batch = batch.to(get_device_id())\n batch = batch.contiguous()\n\n input_ids = batch[\"input_ids\"]\n attention_mask = batch[\"attention_mask\"].to(bool)\n position_ids = batch[\"position_ids\"]\n\n unwrapped_model = unwrap_model(model)\n if hasattr(unwrapped_model, \"vp_stage\"):\n vp_rank = unwrapped_model.vp_stage\n else:\n vp_rank = 0\n\n multi_modal_inputs = {}\n if \"multi_modal_inputs\" in batch:\n from verl.utils.model import extract_multi_modal_inputs\n\n indices = batch.get(\"multi_modal_inputs_idx\", None)\n multi_modal_inputs = extract_multi_modal_inputs(batch[\"multi_modal_inputs\"], indices)\n responses = batch[\"responses\"]\n response_length = responses.size(1)\n label = position_ids.clone()\n label[:, -response_length - 1 : -1] = responses\n label_mask = attention_mask.clone()\n label_mask[:, : -response_length - 1] = False\n label_mask[:, -1] = False\n\n if RouterReplayHelper.is_replay_backward_action(self.tf_config, vp_rank):\n router_instance_list = RouterReplayHelper.get_micro_batch_router_list(self.tf_config, vp_rank)\n for router in router_instance_list:\n router.set_router_replay_action(RouterReplayAction.REPLAY_FORWARD)\n\n if RouterReplayHelper.is_replay_forward_action(self.tf_config, vp_rank):\n layers_topk_idx = batch[\"routed_experts\"]\n set_router_replay_data(layers_topk_idx, attention_mask, self.tf_config, vp_rank)\n\n from verl.models.mcore import get_mcore_forward_fn, get_mcore_forward_fused_fn\n\n if self.use_fused_kernels:\n forward_fn = get_mcore_forward_fused_fn(self.hf_config)\n if return_schedule_plan:\n forward_fn = gptmodel_forward_1f1b_overlap\n # return dict of [logits, entropy]\n output = forward_fn(\n model=model,\n input_ids=input_ids,\n position_ids=position_ids,\n attention_mask=attention_mask,\n labels=label,\n labels_mask=label_mask,\n temperature=temperature,\n multi_modal_inputs=multi_modal_inputs,\n )\n else:\n forward_fn = get_mcore_forward_fn(self.hf_config)\n\n def logits_processor(logits, label, label_mask):\n assert logits.shape[:2] == label.shape[:2]\n assert label.shape == label_mask.shape\n logits.div_(temperature)\n ret = {}\n if calculate_entropy:\n logits_bak = logits.clone()\n # # disable the hint until the fused_kernel is optimized for triton>=3.3\n # logger.warning_once(\n # \"For memory-efficient computation, enable fused kernels via \"\n # \"`actor_rollout_ref.model.use_fused_kernels=True`. \"\n # \"The current `clone()` operation ensures correctness but increases memory usage.\"\n # )\n entropy = vocab_parallel_entropy(logits)\n ret[\"entropy\"] = entropy\n else:\n logits_bak = logits\n log_probs = vocab_parallel_log_probs_from_logits(logits_bak, label)\n log_probs = log_probs.masked_fill(~label_mask, 0.0)\n ret[\"log_probs\"] = log_probs\n return ret\n\n logits_processor_args = {\"label\": label, \"label_mask\": label_mask}\n output = forward_fn(\n model=model,\n input_ids=input_ids,\n attention_mask=attention_mask,\n position_ids=position_ids,\n multi_modal_inputs=multi_modal_inputs,\n logits_processor=logits_processor,\n logits_processor_args=logits_processor_args,\n data_format=\"thd\" if self.config.megatron.use_remove_padding else \"bshd\",\n mtp_config=None if forward_only else self.mtp_config,\n )\n\n if forward_only:\n meta_info = None\n else:\n clip_ratio_c = self.config.get(\"clip_ratio_c\", 3.0)\n meta_info = {\n \"clip_ratio\": self.config.clip_ratio,\n \"entropy_coeff\": self.config.entropy_coeff,\n \"clip_ratio_c\": clip_ratio_c,\n }\n\n if RouterReplayHelper.is_r2_record_action(self.tf_config, vp_rank):\n merge_router_topk_indices(\n attention_mask, input_ids, self.mini_layer_topk_idx_list, self.tf_config, vp_rank\n )\n\n if RouterReplayHelper.is_replay_forward_action(self.tf_config, vp_rank):\n router_instance_list = RouterReplayHelper.get_micro_batch_router_list(self.tf_config, vp_rank)\n for router in router_instance_list:\n router.set_router_replay_action(RouterReplayAction.REPLAY_BACKWARD)\n\n return output, partial(loss_func, data=batch, meta_info=meta_info)\n\n # batch should be a list of batches inside micro-batches\n batch_generator = make_batch_generator(micro_batches, vpp_size=len(self.actor_module))\n\n # TODO: we may use the new schedule instead\n # for flash-attn: (seq_len, batch_size, hidden_size) = (mbs*seq_len, 1, hidden_size)\n if mpu.get_pipeline_model_parallel_world_size() > 1:\n losses_reduced = forward_backward_func(\n forward_step_func=forward_step,\n data_iterator=batch_generator,\n model=self.actor_module,\n num_microbatches=n_micro_batch,\n seq_length=total_seqlen, # no use when input_shapes was set\n micro_batch_size=1, # no use when input_shapes was set\n forward_only=forward_only,\n )\n else:\n losses_reduced = forward_backward_func(\n forward_step_func=forward_step,\n data_iterator=batch_generator,\n model=self.actor_module,\n num_microbatches=n_micro_batch,\n seq_length=total_seqlen, # in use for pp = 1\n micro_batch_size=1, # in use for pp = 1\n forward_only=forward_only,\n )\n # loss_reduces contains the stats returned from loss_func\n\n if self.has_multi_modal_inputs:\n data.batch.pop(\"multi_modal_inputs\")\n data.batch.pop(\"multi_modal_inputs_idx\")\n data.non_tensor_batch.pop(\"multi_modal_inputs\")\n\n losses_reduced = {\"output\": losses_reduced}\n if use_dynamic_bsz:\n losses_reduced[\"indices\"] = indices\n if RouterReplayHelper.is_r2_record_action(self.tf_config):\n if self.tf_config.virtual_pipeline_model_parallel_size is not None:\n # config = self.actor_module[0].module.module.config\n vp_size = len(self.actor_module)\n microbatch_group_size_per_vp_stage = self.tf_config.microbatch_group_size_per_vp_stage\n bs = n_micro_batch\n losses_reduced[\"mini_layer_topk_idx_tensor\"] = reorder_and_merge_vpp_layers(\n self.mini_layer_topk_idx_list, bs, vp_size, microbatch_group_size_per_vp_stage\n )\n else:\n losses_reduced[\"mini_layer_topk_idx_tensor\"] = torch.cat(self.mini_layer_topk_idx_list, dim=0)\n self.mini_layer_topk_idx_list = []\n\n # Collect and pass MTP metrics to losses_reduced\n if not forward_only and self.mtp_config and self.mtp_config.enable_train:\n metrics = get_megatron_mtp_loss(n_micro_batch)\n losses_reduced[\"mtp_losses\"] = [metrics]\n\n return losses_reduced\n\n @GPUMemoryLogger(role=\"megatron actor\", logger=logger)\n def update_policy(self, dataloader: Iterable[DataProto], enable_mtp: bool = False) -> dict:\n \"\"\"Update the policy with an iterator of DataProto\n\n Args:\n dataloader (Iterable[DataProto]): an iterator over the DataProto that returns by ``make_minibatch_iterator``\n The keys of each data batch is described in the make_minibatch_iterator.\n\n enable_mtp (bool, optional): whether to enable MTP communication\n\n Returns:\n Dict: a dictionary containing the statistics. Note that the statistics are only valid in the last pp stage\n and users have to combine the output in each dp rank manually.\n\n \"\"\"\n metrics = {}\n for data in dataloader:\n if self.config.router_replay.mode in [\"R2\", \"R3\"]:\n RouterReplay.set_global_router_replay_action(RouterReplayAction.REPLAY_FORWARD)\n self.actor_optimizer.zero_grad()\n # use use_contiguous_buffers_in_local_ddp and no overlap_dp_param_comm\n for chunk in self.actor_module:\n # if use distributed optimizer, zero grad buffer will be handled by optimizer\n chunk.zero_grad_buffer()\n\n calculate_entropy = self.config.entropy_coeff != 0\n if data.meta_info.get(\"micro_batch_size\", None) is not None:\n micro_batch_size = data.meta_info[\"micro_batch_size\"]\n else:\n micro_batch_size = self.config.ppo_micro_batch_size_per_gpu\n max_token_len = None\n if self.config.use_dynamic_bsz:\n max_token_len = self.config.ppo_max_token_len_per_gpu * self.config.megatron.context_parallel_size\n metric_micro_batch = self.forward_backward_batch(\n data,\n calculate_entropy=calculate_entropy,\n use_dynamic_bsz=self.config.use_dynamic_bsz,\n micro_batch_size=micro_batch_size,\n max_token_len=max_token_len,\n mini_batch_size=self.config.ppo_mini_batch_size,\n )\n\n mtp_losses = metric_micro_batch.get(\"mtp_losses\", None)\n if mtp_losses is not None:\n # mtp_losses is now in format: [{\"mtp_losses/mtp_1_loss\": [value1], \"mtp_losses/mtp_2_loss\": [value2]}]\n for mtp_metrics_dict in mtp_losses:\n append_to_dict(metrics, mtp_metrics_dict)\n\n metric_micro_batch = metric_micro_batch[\"output\"]\n for metric in metric_micro_batch:\n # Note that o[0] is metrics, o[1] is entropy, o[2] is response_mask\n append_to_dict(metrics, metric[0]) # append the metric from this micro-batch to global metrics.\n\n update_successful, grad_norm, num_zeros_in_grad = self.actor_optimizer.step()\n data = {\"actor/grad_norm\": grad_norm}\n append_to_dict(metrics, data)\n\n if update_successful:\n # allgather already execute in optimizer.step in new megatron\n pass\n else:\n raise NotImplementedError\n\n if self.config.router_replay.mode in [\"R2\", \"R3\"]:\n RouterReplay.clear_global_router_replay_action()\n RouterReplay.clear_global_indices()\n\n self.actor_optimizer.zero_grad()\n get_torch_device().empty_cache()\n return metrics\n"}133{"file_name": "verl__workers__config__critic.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport warnings\nfrom dataclasses import dataclass, field\nfrom typing import Optional\n\nfrom omegaconf import MISSING\n\nfrom verl.base_config import BaseConfig\nfrom verl.trainer.config import BaseModelConfig, CheckpointConfig\nfrom verl.utils.profiler import ProfilerConfig\n\nfrom .engine import FSDPEngineConfig, McoreEngineConfig\nfrom .model import HFModelConfig\nfrom .optimizer import OptimizerConfig\n\n__all__ = [\"CriticConfig\", \"FSDPCriticConfig\", \"McoreCriticConfig\", \"FSDPCriticModelCfg\"]\n\n\n@dataclass\nclass CriticConfig(BaseConfig):\n \"\"\"Configuration for critic model training.\n\n The inheritance from BaseConfig provides omegaconf.DictConfig-like interface for a dataclass config.\n\n Args:\n strategy (str): Strategy used for critic model training (fsdp, fsdp2, megatron).\n ppo_micro_batch_size_per_gpu (int): Local per-GPU micro batch size.\n rollout_n (int): Number of rollouts per update (mirrors actor rollout_n).\n optim (Dict[str, Any]): Optimizer configuration including lr, weight_decay, etc.\n model (Dict[str, Any]): Model configuration including path, tokenizer_path, etc.\n ppo_mini_batch_size (int): PPO mini-batch size per update.\n ppo_micro_batch_size (Optional[int]): Global micro batch size (deprecated).\n use_dynamic_bsz (bool): Whether to automatically adjust batch size at runtime.\n ppo_max_token_len_per_gpu (int): Max tokens per GPU in one PPO batch.\n forward_max_token_len_per_gpu (int): Max token length per GPU in forward pass.\n ppo_epochs (int): Number of PPO epochs per batch.\n shuffle (bool): Shuffle training data across PPO epochs.\n cliprange_value (float): PPO value function clipping range.\n loss_agg_mode (str): Loss aggregation mode.\n checkpoint (Dict[str, Any]): Checkpoint configuration.\n profiler (Dict[str, Any]): Profiler configuration.\n enable (Optional[bool]): Whether to enable the critic.\n \"\"\"\n\n _mutable_fields = BaseConfig._mutable_fields | {\n \"ppo_micro_batch_size_per_gpu\",\n \"ppo_mini_batch_size\",\n \"ppo_micro_batch_size\",\n \"model_config\",\n }\n\n strategy: str = MISSING\n ppo_micro_batch_size_per_gpu: Optional[int] = None\n enable: Optional[bool] = None\n rollout_n: int = 1\n ppo_mini_batch_size: int = 1\n use_dynamic_bsz: bool = False\n ppo_max_token_len_per_gpu: int = 32768\n # deprecate this\n forward_max_token_len_per_gpu: int = 32768\n ppo_infer_micro_batch_size_per_gpu: Optional[int] = None\n ppo_infer_max_token_len_per_gpu: int = 32768\n ppo_epochs: int = 1\n data_loader_seed: int = 1\n shuffle: bool = True\n cliprange_value: float = 0.5\n loss_agg_mode: str = \"token-mean\"\n ppo_micro_batch_size: Optional[int] = None\n engine: BaseConfig = field(default_factory=BaseConfig)\n optim: OptimizerConfig = field(default_factory=OptimizerConfig)\n # deprecate model to favor model_config\n model: BaseModelConfig = field(default_factory=BaseModelConfig)\n model_config: HFModelConfig = None\n checkpoint: CheckpointConfig = field(default_factory=CheckpointConfig)\n profiler: ProfilerConfig = field(default_factory=ProfilerConfig)\n\n def __post_init__(self):\n \"\"\"Validate critic configuration parameters.\"\"\"\n assert self.strategy != MISSING\n\n if self.model_config is None:\n warnings.warn(\"using model in Critic Config is deprecated, please use model_config instead\", stacklevel=2)\n self.model_config = HFModelConfig(\n path=self.model.path,\n tokenizer_path=self.model.tokenizer_path,\n override_config=self.model.override_config,\n external_lib=self.model.external_lib,\n trust_remote_code=self.model.trust_remote_code,\n )\n\n if not self.use_dynamic_bsz:\n self._check_mutually_exclusive(self.ppo_micro_batch_size, self.ppo_micro_batch_size_per_gpu, \"critic\")\n\n if self.ppo_micro_batch_size is not None:\n if self.ppo_mini_batch_size % self.ppo_micro_batch_size != 0:\n raise ValueError(\n f\"[critic] ppo_mini_batch_size ({self.ppo_mini_batch_size}) must be divisible by \"\n f\"ppo_micro_batch_size ({self.ppo_micro_batch_size})\"\n )\n\n def validate(self, n_gpus: int, train_batch_size: int):\n \"\"\"Validate critic configuration with runtime parameters.\n\n Args:\n n_gpus: Total number of GPUs available\n train_batch_size: Training batch size from data config\n \"\"\"\n if not self.use_dynamic_bsz:\n if train_batch_size < self.ppo_mini_batch_size:\n raise ValueError(\n f\"train_batch_size ({train_batch_size}) must be >= \"\n f\"critic.ppo_mini_batch_size ({self.ppo_mini_batch_size})\"\n )\n\n @staticmethod\n def _check_mutually_exclusive(mbs, mbs_per_gpu, name: str):\n \"\"\"Validate mutually exclusive micro batch size configuration options.\n\n Ensures that users don't set both deprecated micro_batch_size and\n the new micro_batch_size_per_gpu parameters simultaneously.\n\n Args:\n mbs: Deprecated micro batch size parameter value.\n mbs_per_gpu: New micro batch size per GPU parameter value.\n name (str): Configuration section name for error messages.\n\n Raises:\n ValueError: If both parameters are set or neither is set.\n \"\"\"\n param = \"micro_batch_size\"\n param_per_gpu = f\"{param}_per_gpu\"\n\n if mbs is None and mbs_per_gpu is None:\n raise ValueError(f\"[{name}] Please set at least one of '{name}.{param}' or '{name}.{param_per_gpu}'.\")\n\n if mbs is not None and mbs_per_gpu is not None:\n raise ValueError(\n f\"[{name}] You have set both '{name}.{param}' AND '{name}.{param_per_gpu}'. Please remove \"\n f\"'{name}.{param}' because only '*_{param_per_gpu}' is supported (the former is deprecated).\"\n )\n\n\n@dataclass\nclass McoreCriticConfig(CriticConfig):\n \"\"\"Configuration for Megatron-based critic model training.\n\n The inheritance from CriticConfig provides all base critic configuration plus Megatron-specific settings.\n\n Args:\n nccl_timeout (int): NCCL timeout in seconds for distributed operations.\n megatron (Dict[str, Any]): Megatron-specific parallelism settings.\n load_weight (bool): Whether to load initial weights.\n \"\"\"\n\n strategy: str = \"megatron\"\n nccl_timeout: int = 600\n megatron: McoreEngineConfig = field(default_factory=McoreEngineConfig)\n load_weight: bool = True\n\n def validate(self, n_gpus: int, train_batch_size: int):\n \"\"\"Validate Megatron critic configuration with runtime parameters.\"\"\"\n super().validate(n_gpus, train_batch_size)\n\n\n@dataclass\nclass FSDPCriticConfig(CriticConfig):\n \"\"\"Configuration for FSDP-based critic model training.\n\n The inheritance from CriticConfig provides all base critic configuration plus FSDP-specific settings.\n\n Args:\n forward_micro_batch_size (int): Forward-only batch size during inference (global).\n forward_micro_batch_size_per_gpu (int): Forward-only batch size during inference (per GPU).\n ulysses_sequence_parallel_size (int): [DEPRECATED] Ulysses sequence parallel size for long sequences.\n grad_clip (float): Gradient clipping for critic updates.\n \"\"\"\n\n _mutable_fields = CriticConfig._mutable_fields | {\n \"forward_micro_batch_size\",\n \"forward_micro_batch_size_per_gpu\",\n }\n\n strategy: str = \"fsdp\"\n forward_micro_batch_size: int = 1\n forward_micro_batch_size_per_gpu: int = 1\n ulysses_sequence_parallel_size: int = 1\n grad_clip: float = 1.0\n\n def __post_init__(self):\n \"\"\"Validate FSDP critic configuration parameters.\"\"\"\n super().__post_init__()\n\n if self.strategy in {\"fsdp\", \"fsdp2\"}:\n if self.ulysses_sequence_parallel_size > 1:\n if not self.model.get(\"use_remove_padding\", False):\n raise ValueError(\n \"When using sequence parallelism for critic, you must enable `use_remove_padding`.\"\n )\n\n def validate(self, n_gpus: int, train_batch_size: int):\n \"\"\"Validate FSDP critic configuration with runtime parameters.\"\"\"\n super().validate(n_gpus, train_batch_size)\n\n if not self.use_dynamic_bsz:\n sp_size = self.ulysses_sequence_parallel_size\n if self.ppo_micro_batch_size is not None:\n if self.ppo_micro_batch_size * sp_size < n_gpus:\n raise ValueError(\n f\"critic.ppo_micro_batch_size ({self.ppo_micro_batch_size}) * \"\n f\"ulysses_sequence_parallel_size ({sp_size}) must be >= n_gpus ({n_gpus})\"\n )\n\n\n@dataclass\nclass FSDPCriticModelCfg(BaseModelConfig):\n \"\"\"FSDP-enabled critic model configuration.\n Inherits base critic settings and adds distributed-memory and LoRA options.\n\n Args:\n use_shm (bool): Whether to use shared memory for loading the model.\n enable_activation_offload (bool): Offload activations to CPU to reduce GPU memory usage.\n use_remove_padding (bool): Use remove-padding optimization (saves compute).\n enable_gradient_checkpointing (bool): Enable gradient checkpointing for memory efficiency.\n fsdp_config (FSDPEngineConfig): FSDP-specific configuration block.\n lora_rank (int): Set to positive value to enable LoRA (e.g., 32).\n lora_alpha (int): LoRA scaling factor.\n target_modules (Union[str, List[str]]): LoRA target modules: \"all-linear\" or list of layer names.\n \"\"\"\n\n use_shm: bool = False\n enable_activation_offload: bool = False\n use_remove_padding: bool = False\n enable_gradient_checkpointing: bool = True\n fsdp_config: FSDPEngineConfig = field(default_factory=FSDPEngineConfig)\n lora_rank: int = 0\n lora_alpha: int = 16\n target_modules: str | list[str] = \"all-linear\"\n # TiledMLP configuration for memory-efficient MLP computation\n tiled_mlp: dict = field(default_factory=lambda: {\"enabled\": False, \"num_shards\": 4})\n"}134{"file_name": "verl__workers__config__engine.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport warnings\nfrom dataclasses import dataclass, field\nfrom typing import Any, Callable, Literal, Optional\n\nfrom verl.base_config import BaseConfig\nfrom verl.trainer.config import CheckpointConfig\n\nfrom ...utils.profiler import ProfilerConfig\nfrom .model import HFModelConfig\nfrom .optimizer import OptimizerConfig\n\n__all__ = [\n \"FSDPEngineConfig\",\n \"McoreEngineConfig\",\n \"TrainingWorkerConfig\",\n \"VeOmniEngineConfig\",\n \"EngineConfig\",\n \"EngineRouterReplayConfig\",\n]\n\n\n# TODO: rename to RouterReplayConfig after removing the legacy implementation\n@dataclass\nclass EngineRouterReplayConfig(BaseConfig):\n \"\"\"Configuration for router replay in MoE models.\n\n This configuration controls the routing behavior for Mixture of Experts (MoE) models,\n allowing for deterministic training through route recording and replay.\n\n Args:\n mode (str): Router replay mode. Options: 'disabled', 'R2', 'R3'.\n - 'disabled': No router replay functionality\n - 'R2': Use Router Replay routing strategy\n - 'R3': Use Rollout Router Replay routing strategy\n record_file (Optional[str]): File path to save recorded routing decisions.\n Required when mode is 'record', 'R2', or 'R3'.\n replay_file (Optional[str]): File path to load recorded routing decisions for replay.\n Required when mode is 'replay'.\n \"\"\"\n\n mode: str = \"disabled\"\n record_file: Optional[str] = None\n replay_file: Optional[str] = None\n\n def __post_init__(self):\n \"\"\"Validate router replay configuration.\"\"\"\n valid_modes = [\"disabled\", \"R2\", \"R3\"]\n if self.mode not in valid_modes:\n raise ValueError(f\"Invalid router_replay mode: {self.mode}. Must be one of {valid_modes}\")\n\n\n@dataclass\nclass EngineConfig(BaseConfig):\n _mutable_fields = BaseConfig._mutable_fields | {\n \"use_dynamic_bsz\",\n \"max_token_len_per_gpu\",\n \"micro_batch_size_per_gpu\",\n \"infer_max_token_len_per_gpu\",\n \"infer_micro_batch_size_per_gpu\",\n \"use_fused_kernels\",\n \"use_remove_padding\",\n }\n # whether to offload param\n param_offload: bool = False\n # whether to offload optimizer\n optimizer_offload: bool = False\n # whether to offload grad\n grad_offload: bool = False\n # whether the engine is forward only (e.g., ref policy)\n forward_only: bool = False\n # the strategy (backend)\n strategy: str = None\n # model dtype\n dtype: str = \"bfloat16\" # [\"bfloat16\", \"float16\"]\n # whether to use dynamic bsz\n use_dynamic_bsz: bool = True\n # for training\n max_token_len_per_gpu: int = None\n micro_batch_size_per_gpu: int = None\n # for inference\n infer_max_token_len_per_gpu: int = None\n infer_micro_batch_size_per_gpu: int = None\n # whether use fuse lm head kernel\n use_fused_kernels: bool = False\n # TODO (this may conflict with the one in model config)\n use_remove_padding: bool = True\n\n seed: int = 42\n\n full_determinism: bool = False\n router_replay: EngineRouterReplayConfig = field(default_factory=EngineRouterReplayConfig)\n\n def __post_init__(self):\n pass\n # TODO: turn on this check after we reorg config\n # if self.use_dynamic_bsz:\n # assert self.max_token_len_per_gpu is not None\n # else:\n # assert self.micro_batch_size_per_gpu is not None\n\n\n@dataclass\nclass McoreEngineConfig(EngineConfig):\n \"\"\"Configuration for Megatron parallelism.\n\n The inheritance from BaseConfig provides omegaconf.DictConfig-like interface for a dataclass config.\n\n Args:\n param_offload (bool): Whether to offload parameters to CPU.\n grad_offload (bool): Whether to offload gradients to CPU.\n optimizer_offload (bool): Whether to offload optimizer states to CPU.\n tensor_model_parallel_size (int): Tensor model parallel size.\n expert_model_parallel_size (int): Expert model parallel size for MoE models.\n expert_tensor_parallel_size (Optional[int]): Expert tensor parallel size for MoE models.\n pipeline_model_parallel_size (int): Pipeline model parallel size.\n virtual_pipeline_model_parallel_size (Optional[int]): Virtual pipeline model parallel size\n for interleaved scheduling.\n context_parallel_size (int): Context parallel size for long sequences.\n sequence_parallel (bool): Whether to enable sequence parallelism.\n use_distributed_optimizer (bool): Whether to use distributed optimizer.\n use_dist_checkpointing (bool): Whether to use distributed checkpointing.\n dist_checkpointing_path (Optional[str]): Path for distributed checkpointing.\n dist_ckpt_optim_fully_reshardable (bool): Use fully reshardable optimizer checkpoints.\n distrib_optim_fully_reshardable_mem_efficient (bool): Use memory-efficient fully reshardable format.\n seed (int): Random seed for reproducibility.\n override_ddp_config (dict[str, Any]): Override configuration for DDP.\n override_transformer_config (dict[str, Any]): Override configuration for transformer.\n use_mbridge (bool): Whether to use MBridge for communication.\n dtype (str): Mixed precision training param dtype, default \"bfloat16\"\n \"\"\"\n\n # sequence_parallel is not listed as a frozen field for auto-correction purpose\n _mutable_fields = EngineConfig._mutable_fields | {\"sequence_parallel\"}\n # mcore parallelism\n tensor_model_parallel_size: int = 1\n expert_model_parallel_size: int = 1\n expert_tensor_parallel_size: Optional[int] = None\n pipeline_model_parallel_size: int = 1\n virtual_pipeline_model_parallel_size: Optional[int] = None\n context_parallel_size: int = 1\n sequence_parallel: bool = True\n use_distributed_optimizer: bool = True\n use_dist_checkpointing: bool = False\n dist_checkpointing_path: Optional[str] = None\n dist_checkpointing_prefix: str = \"\"\n dist_ckpt_optim_fully_reshardable: bool = False\n distrib_optim_fully_reshardable_mem_efficient: bool = False\n override_ddp_config: dict[str, Any] = field(default_factory=dict)\n override_transformer_config: dict[str, Any] = field(default_factory=dict)\n override_mcore_model_config: dict[str, Any] = field(default_factory=dict)\n use_mbridge: bool = True\n vanilla_mbridge: bool = True\n strategy: str = \"megatron\"\n\n def __post_init__(self) -> None:\n super().__post_init__()\n \"\"\"config validation logics go here\"\"\"\n assert self.strategy == \"megatron\"\n assert self.dtype in [\"bfloat16\", \"float16\"], f\"dtype {self.dtype} not supported\"\n if self.tensor_model_parallel_size == 1:\n warnings.warn(\"set sequence parallel to false as TP size is 1\", stacklevel=2)\n self.sequence_parallel = False\n\n\n@dataclass\nclass FSDPEngineConfig(EngineConfig):\n \"\"\"Configuration for FSDP (Fully Sharded Data Parallel).\n\n The inheritance from BaseConfig provides omegaconf.DictConfig-like interface for a dataclass config.\n\n Args:\n wrap_policy (Dict[str, Any]): Configuration for FSDP wrap policy.\n param_offload (bool): Whether to offload parameters to CPU, default False\n optimizer_offload (bool): Whether to offload optimizer states to CPU, default False\n offload_policy (bool): Whether to offload policy model parameters, default False\n reshard_after_forward (bool): Whether to reshard parameters after forward pass, default True\n fsdp_size (int): FSDP group size. -1 means use all available GPUs.\n forward_prefetch (bool): Whether to prefetch parameters for next forward pass, default False\n model_dtype (str): Model data type used to initialize the transformers model. default \"fp32\"\n use_orig_params (bool): Whether to use original parameters when initialize FSDP1, default False\n seed (int): Random seed for reproducibility.\n full_determinism (bool): If true, enable_full_determinism is called to ensure reproducible results\n in distributed training. Important: this will negatively impact performance, so only use it for\n debugging.\n mixed_precision (Optional[dict[str, Any]]): Mixed precision configuration for FSDP, default None\n dtype (str): Mixed precision training param dtype, default \"bfloat16\"\n \"\"\"\n\n # ulysses_sequence_parallel_size is mutable for backward compatibility\n _mutable_fields = EngineConfig._mutable_fields | {\"ulysses_sequence_parallel_size\"}\n\n # fsdp specific flags\n wrap_policy: dict[str, Any] = field(default_factory=dict)\n offload_policy: bool = False\n reshard_after_forward: bool = True\n fsdp_size: int = -1\n forward_prefetch: bool = False\n model_dtype: str = \"fp32\"\n use_orig_params: bool = False\n mixed_precision: Optional[dict[str, Any]] = None\n ulysses_sequence_parallel_size: int = 1\n entropy_from_logits_with_chunking: bool = False\n use_torch_compile: bool = True\n entropy_checkpointing: bool = False\n strategy: str = \"fsdp\"\n\n def __post_init__(self):\n super().__post_init__()\n assert self.strategy in [\"fsdp\", \"fsdp2\"], f\"strategy {self.strategy} not supported\"\n\n\n@dataclass\nclass VeOmniEngineConfig(EngineConfig):\n \"\"\"Configuration for VeOmni.\n\n The inheritance from BaseConfig provides omegaconf.DictConfig-like interface for a dataclass config.\n\n Args:\n wrap_policy (Dict[str, Any]): Configuration for FSDP wrap policy.\n param_offload (bool): Whether to offload parameters to CPU, default False\n optimizer_offload (bool): Whether to offload optimizer states to CPU, default False\n offload_policy (bool): Whether to offload policy model parameters, default False\n reshard_after_forward (bool): Whether to reshard parameters after forward pass, default True\n fsdp_size (int): FSDP group size. -1 means use all available GPUs, default -1\n ulysses_parallel_size (int): Ulysses sequence parallel size, default 1\n expert_parallel_size (int): Expert parallel size, default 1\n init_device (str): Device to initialize model weights.\n 1. `cpu`: Init parameters on CPU in rank0 only.\n 2. `cuda`: Init parameters on GPU.\n 3. `meta`: Init parameters on meta.\n 4. `npu`: Init parameters on Ascend NPU.\n default \"meta\"\n enable_full_shard (bool): Enable fully shard for FSDP training (ZeRO-3), default False\n enable_fsdp_offload (bool): Enable CPU offload for FSDP1, default False\n enable_reentrant (bool): Use reentrant gradient checkpointing, default False\n attn_implementation (str): Attention implementation to use.\n 1. `eager`\n 2. `sdpa`\n 3. `flash_attention_2`\n 4. `flash_attention_3`\n 5. `veomni_flash_attention_2_with_sp`\n 6. `veomni_flash_attention_3_with_sp`\n 7. `native-sparse`\n default \"flash_attention_2\"\n Note: In case VeOmni add more attn_implementation, please check https://github.com/ByteDance-Seed/VeOmni/\n moe_implementation (str): MoE implementation to use.\n 1. `eager`\n 2. `fused`\n default \"fused\"\n Note: In case VeOmni add more moe_implementation, please check https://github.com/ByteDance-Seed/VeOmni/\n force_use_huggingface (bool): Force loading model from huggingface, default False\n activation_gpu_limit (float): When enabling activation offload, `activation_gpu_limit` GB\n activations are allowed to reserve on GPU, default 0.0\n basic_modules (list[str]): List of basic modules to use, default None\n forward_prefetch (bool): Whether to prefetch parameters for next forward pass, default False\n model_dtype (str): Model data type used to initialize the transformers model. default \"fp32\"\n use_orig_params (bool): Whether to use original parameters when initialize FSDP1, default False\n seed (int): Random seed for reproducibility.\n full_determinism (bool): If true, enable_full_determinism is called to ensure reproducible results\n in distributed training. Important: this will negatively impact performance, so only use it for\n debugging.\n mixed_precision (Optional[dict[str, Any]]): Mixed precision configuration for FSDP, default None\n\n \"\"\"\n\n wrap_policy: dict[str, Any] = field(default_factory=dict)\n offload_policy: bool = False\n reshard_after_forward: bool = True\n forward_prefetch: bool = False\n use_orig_params: bool = False\n entropy_from_logits_with_chunking: bool = False\n use_torch_compile: bool = True\n entropy_checkpointing: bool = False\n strategy: str = \"veomni\"\n fsdp_size: int = -1\n ulysses_parallel_size: int = 1\n expert_parallel_size: int = 1\n seed: int = 42\n full_determinism: bool = False\n mixed_precision: bool = False\n init_device: str = \"meta\"\n enable_full_shard: bool = False\n ckpt_manager: Literal[\"dcp\"] = \"dcp\"\n load_checkpoint_path: Optional[str] = None\n enable_fsdp_offload: bool = False\n enable_reentrant: bool = False\n attn_implementation: str = \"flash_attention_2\"\n moe_implementation: str = \"fused\"\n force_use_huggingface: bool = False\n activation_gpu_limit: float = 0.0\n basic_modules: Optional[list[str]] = field(default_factory=list)\n\n def __post_init__(self):\n super().__post_init__()\n assert self.strategy in [\"veomni\"], f\"strategy {self.strategy} not supported\"\n\n\n@dataclass\nclass TrainingWorkerConfig(BaseConfig):\n model_type: str = None # model type (language_model/value_model)\n model_config: HFModelConfig = None\n engine_config: EngineConfig = None\n optimizer_config: OptimizerConfig = None\n checkpoint_config: CheckpointConfig = None\n profiler_config: ProfilerConfig = None\n # automatically select engine and optimizer function.\n # This function takes model config and the device name as parameter.\n # Users can pass in a higher-order function to take more parameters\n auto_select_engine_optim_fn: Callable[[\"HFModelConfig\", str], tuple[\"EngineConfig\", \"OptimizerConfig\"]] = None\n"}135{"file_name": "verl__workers__config__optimizer.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport warnings\nfrom dataclasses import dataclass\nfrom typing import Optional\n\nfrom omegaconf import MISSING\n\nfrom verl.base_config import BaseConfig\n\n__all__ = [\"OptimizerConfig\", \"FSDPOptimizerConfig\", \"McoreOptimizerConfig\", \"build_optimizer\", \"VeOmniOptimizerConfig\"]\n\n\n@dataclass\nclass OptimizerConfig(BaseConfig):\n \"\"\"Base optimizer configuration.\n\n Args:\n lr (float): learning rate. Must be specified.\n lr_warmup_steps_ratio (float): Warmup steps ratio; total steps will be injected at runtime.\n total_training_steps (int): Total training steps (must be overridden at runtime).\n weight_decay (float): Weight decay factor.\n lr_warmup_steps (Optional[int]): Number of warmup steps; None delegates to lr_warmup_steps_ratio.\n \"\"\"\n\n _mutable_fields = {\"clip_grad\", \"total_training_steps\", \"lr_warmup_steps\"}\n\n lr: float = 1e-3\n lr_warmup_steps_ratio: float = 0.0\n total_training_steps: int = -1\n weight_decay: float = 0.01\n lr_warmup_steps: Optional[int] = -1\n betas: tuple[float, float] = (0.9, 0.999)\n clip_grad: float = 1.0\n # deprecate grad_clip\n grad_clip: Optional[float] = None\n\n def __post_init__(self):\n assert self.lr != MISSING\n if self.grad_clip is not None:\n warnings.warn(\"`grad_clip` is deprecated, use `clip_grad` instead.\", DeprecationWarning, stacklevel=2)\n self.clip_grad = self.grad_clip\n\n\n@dataclass\nclass VeOmniOptimizerConfig(OptimizerConfig):\n \"\"\"VeOmni optimizer configuration extending base OptimizerConfig.\n\n Args:\n optimizer (str): Optimizer name; default is \"adamw\".\n lr (float): Learning rate.\n lr_min (float): Minimum learning rate.\n lr_start (float): Starting learning rate for warmup.\n lr_decay_ratio (float): LR decay ratio.\n lr_scheduler_type (str): LR scheduler type: \"constant\" or \"cosine\".\n \"\"\"\n\n _mutable_fields = OptimizerConfig._mutable_fields.copy()\n\n optimizer: str = \"adamw\"\n lr_min: float = 0.0\n lr_start: float = 0.0\n lr_decay_ratio: float = 1.0\n lr_scheduler_type: str = \"constant\"\n override_optimizer_config: Optional[dict] = None\n\n\n@dataclass\nclass FSDPOptimizerConfig(OptimizerConfig):\n \"\"\"FSDP optimizer configuration extending base OptimizerConfig.\n\n Args:\n optimizer (str): Optimizer class name (e.g., \"AdamW\", \"AdamW8bit\", \"_AdamW\").\n optimizer_impl (str): Module path to import optimizer from (e.g., \"torch.optim\", \"torchao.optim\",\n \"bitsandbytes.optim\").\n lr (float): Learning rate.\n min_lr_ratio (Optional[float]): Minimum LR ratio for cosine schedule.\n lr_scheduler_type (str): LR scheduler type: \"constant\" or \"cosine\".\n num_cycles (float): Number of cosine cycles in LR schedule.\n \"\"\"\n\n _mutable_fields = OptimizerConfig._mutable_fields.copy()\n _mutable_fields.add(\"lr_scheduler_type\")\n\n optimizer: str = \"AdamW\"\n optimizer_impl: str = \"torch.optim\"\n min_lr_ratio: Optional[float] = None\n # deprecate warmup_style\n warmup_style: Optional[str] = None\n lr_scheduler_type: str = \"constant\"\n num_cycles: float = 0.5\n override_optimizer_config: Optional[dict] = None\n\n def __post_init__(self):\n if self.warmup_style is not None:\n assert self.warmup_style in [\"constant\", \"cosine\"]\n warnings.warn(\n \"`warmup_style` is deprecated, use `lr_scheduler_type` instead.\", DeprecationWarning, stacklevel=2\n )\n self.lr_scheduler_type = self.warmup_style\n assert self.lr_scheduler_type in [\"constant\", \"cosine\"]\n return super().__post_init__()\n\n\n@dataclass\nclass McoreOptimizerConfig(OptimizerConfig):\n \"\"\"Mcore optimizer configuration extending base OptimizerConfig.\n\n Args:\n optimizer (str): Optimizer name; default is \"adam\".\n lr (float): Learning rate.\n clip_grad (float): Gradient clipping norm.\n lr_warmup_init (float): Initial learning rate for warmup; defaults to 0.0.\n lr_decay_steps (Optional[int]): Number of decay steps.\n lr_decay_style (str): LR decay style: \"constant\", \"linear\", \"cosine\", or \"inverse_square_root\".\n min_lr (float): Minimum learning rate.\n weight_decay_incr_style (str): Weight decay increment style: \"constant\" or \"cosine\".\n lr_wsd_decay_style (str): Weight-standard-deviation decay style: \"constant\", \"exponential\", or \"cosine\".\n lr_wsd_decay_steps (Optional[int]): Number of steps for weight-standard-deviation decay.\n use_checkpoint_opt_param_scheduler (bool): Whether to use checkpoint optimizer parameter scheduler.\n \"\"\"\n\n optimizer: str = \"adam\"\n lr_warmup_init: float = 0.0\n lr_decay_steps: Optional[int] = None\n lr_decay_style: str = \"linear\"\n min_lr: float = 0.0\n weight_decay_incr_style: str = \"constant\"\n lr_wsd_decay_style: str = \"exponential\"\n lr_wsd_decay_steps: Optional[int] = None\n use_checkpoint_opt_param_scheduler: bool = False\n override_optimizer_config: Optional[dict] = None\n\n\ndef build_optimizer(parameters, config: FSDPOptimizerConfig):\n \"\"\"Build an optimizer based on the configuration.\n\n Dynamically imports and instantiates an optimizer class from the specified module.\n\n Args:\n parameters: Model parameters to optimize\n config: FSDPOptimizerConfig with optimizer settings\n\n Returns:\n Optimizer instance\n\n Examples:\n # PyTorch AdamW\n config.optimizer_impl = \"torch.optim\"\n config.optimizer = \"AdamW\"\n\n # TorchAO AdamW with bf16 stochastic rounding\n config.optimizer_impl = \"torchao.optim\"\n config.optimizer = \"_AdamW\"\n config.override_optimizer_config = {\"bf16_stochastic_round\": True}\n\n # BitsAndBytes AdamW 8bit\n config.optimizer_impl = \"bitsandbytes.optim\"\n config.optimizer = \"AdamW8bit\"\n \"\"\"\n import importlib\n\n optimizer_args = {\n \"lr\": config.lr,\n \"weight_decay\": config.weight_decay,\n }\n\n optimizer_name_lower = config.optimizer.lower()\n if \"adam\" in optimizer_name_lower or \"ademamix\" in optimizer_name_lower:\n optimizer_args[\"betas\"] = config.betas\n\n if config.override_optimizer_config is not None:\n optimizer_args.update(config.override_optimizer_config)\n\n try:\n module = importlib.import_module(config.optimizer_impl)\n optimizer_cls = getattr(module, config.optimizer)\n except ImportError as e:\n raise ImportError(\n f\"Failed to import module '{config.optimizer_impl}'. Make sure the package is installed. Error: {e}\"\n ) from e\n except AttributeError as e:\n raise AttributeError(\n f\"Optimizer '{config.optimizer}' not found in module '{config.optimizer_impl}'. \"\n f\"Available optimizers: {dir(module)}\"\n ) from e\n\n return optimizer_cls(parameters, **optimizer_args)\n"}136{"file_name": "verl__workers__config__reward.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nimport os\nfrom dataclasses import dataclass, field\nfrom typing import Optional\n\nfrom verl.base_config import BaseConfig\nfrom verl.trainer.config.config import ModuleConfig\n\nfrom .rollout import RolloutConfig\n\n__all__ = [\"SandboxFusionConfig\", \"RewardConfig\", \"RewardModelConfig\"]\n\nlogger = logging.getLogger(__name__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\n@dataclass\nclass RewardManagerConfig(BaseConfig):\n \"\"\"Configuration for reward manager.\n\n A reward manager defines the mechanism of computing rule-based reward and handling different reward sources.\n\n Args:\n source (str): Source of the reward manager. Options: ``\"register\"``, ``\"importlib\"``. Default: ``\"register\"``.\n name (str, optional):\n - When ``source`` is ``\"register\"``, the name is used in `get_reward_manager_cls(name)``.\n See ``verl/experimental/reward/reward_manager.py`` for options. Default: ``\"naive\"``.\n - When ``source`` is ``\"importlib\"``, the name is used in ``getattr(module, name)``,\n e.g., ``\"DAPORewardManager\"``.\n module (ModuleConfig, optional): Optional configuration for the external module defining the reward manager,\n \"\"\"\n\n source: str = \"register\"\n name: str = \"naive\"\n module: Optional[ModuleConfig] = field(default_factory=ModuleConfig)\n\n def __post_init__(self):\n super().__post_init__()\n if self.source == \"register\":\n from verl.experimental.reward_loop.reward_manager.registry import REWARD_MANAGER\n\n assert self.name in REWARD_MANAGER, (\n f\"Reward manager is not registered: {self.name=} ,{REWARD_MANAGER.keys()=}\"\n )\n elif self.source == \"importlib\":\n # NOTE: The existence is not checked since it depends on which machine the config is initialized on.\n assert self.module is not None and self.module.path is not None, (\n \"When source is importlib, module.path should be set.\"\n )\n\n\n@dataclass\nclass SandboxFusionConfig(BaseConfig):\n \"\"\"Configuration for cloud/local sandbox fusion.\n\n Args:\n url (Optional[str]): Cloud/local function URL for sandbox execution.\n max_concurrent (int): Max concurrent requests allowed to sandbox.\n memory_limit_mb (int): Max memory limit for each sandbox process in MB.\n \"\"\"\n\n url: Optional[str] = None\n max_concurrent: int = 64\n memory_limit_mb: int = 1024\n\n\n@dataclass\nclass RewardModelConfig(BaseConfig):\n _mutable_fields = BaseConfig._mutable_fields\n\n enable: bool = False\n enable_resource_pool: bool = False\n n_gpus_per_node: int = 0\n nnodes: int = 0\n model_path: Optional[str] = None\n inference: RolloutConfig = field(default_factory=RolloutConfig)\n\n\n@dataclass\nclass RewardConfig(BaseConfig):\n _mutable_fields = BaseConfig._mutable_fields\n\n # reward manager args\n num_workers: int = 8\n reward_manager: RewardManagerConfig = field(default_factory=RewardManagerConfig)\n\n # reward model args\n reward_model: RewardModelConfig = field(default_factory=RewardModelConfig)\n\n # sandbox fusion args\n sandbox_fusion: SandboxFusionConfig = field(default_factory=SandboxFusionConfig)\n"}137{"file_name": "verl__workers__config__rollout.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport warnings\nfrom dataclasses import dataclass, field\nfrom typing import Optional\n\nfrom omegaconf import MISSING\n\nfrom verl.base_config import BaseConfig\nfrom verl.utils.profiler import ProfilerConfig\nfrom verl.workers.config.model import MtpConfig\n\n__all__ = [\n \"SamplingConfig\",\n \"MultiTurnConfig\",\n \"CustomAsyncServerConfig\",\n \"AgentLoopConfig\",\n \"TraceConfig\",\n \"ServerConfig\",\n \"PrometheusConfig\",\n \"RolloutConfig\",\n \"CheckpointEngineConfig\",\n]\n\n\n@dataclass\nclass SamplingConfig(BaseConfig):\n temperature: float = 1.0\n top_k: int = -1\n top_p: float = 1.0\n do_sample: bool = True\n n: int = 1\n\n\n@dataclass\nclass MultiTurnConfig(BaseConfig):\n _mutable_fields = {\"max_assistant_turns\", \"max_user_turns\"}\n\n enable: bool = False\n max_assistant_turns: Optional[int] = None\n tool_config_path: Optional[str] = None\n max_user_turns: Optional[int] = None\n max_parallel_calls: int = 1\n max_tool_response_length: int = 256\n tool_response_truncate_side: str = \"middle\"\n interaction_config_path: Optional[str] = None\n use_inference_chat_template: bool = False\n tokenization_sanity_check_mode: str = \"strict\"\n format: str = \"hermes\"\n num_repeat_rollouts: Optional[int] = None\n\n\n@dataclass\nclass CustomAsyncServerConfig(BaseConfig):\n path: Optional[str] = None\n name: Optional[str] = None\n\n\n@dataclass\nclass AgentLoopConfig(BaseConfig):\n num_workers: int = 8\n default_agent_loop: str = \"single_turn_agent\"\n agent_loop_config_path: Optional[str] = None\n custom_async_server: CustomAsyncServerConfig = field(default_factory=CustomAsyncServerConfig)\n # Fully qualified class name for custom AgentLoopManager (e.g., \"mypackage.module.MyManager\").\n # Security: This class will be dynamically imported via importlib. Only use trusted class paths.\n agent_loop_manager_class: Optional[str] = None\n\n\n@dataclass\nclass TraceConfig(BaseConfig):\n backend: Optional[str] = None\n token2text: bool = False\n max_samples_per_step_per_worker: Optional[int] = None\n\n def __post_init__(self):\n if self.max_samples_per_step_per_worker is not None and self.max_samples_per_step_per_worker < 0:\n raise ValueError(\"`max_samples_per_step_per_worker` must be a non-negative integer or null.\")\n\n\n@dataclass\nclass ServerConfig(BaseConfig):\n \"\"\"\n Configuration for SGLang server when running in server mode\n \"\"\"\n\n timeout: float = 60.0\n max_attempts: int = 3\n retry_delay: float = 2.0\n max_connections: int = 1000\n max_start_wait_time: float = 300.0\n\n\n@dataclass\nclass PrometheusConfig(BaseConfig):\n \"\"\"\n Configuration for Prometheus server\n \"\"\"\n\n # whether enable prometheus on server mode rollout\n enable: bool = False\n # Port number that Prometheus listens on, default is 9090\n port: int = 9090\n # Path to Prometheus configuration file\n file: str = \"/tmp/ray/session_latest/metrics/prometheus/prometheus.yml\"\n # Specify served_model_name to avoid displaying overly long model paths in Grafana\n served_model_name: Optional[str] = None\n\n\n@dataclass\nclass CheckpointEngineConfig(BaseConfig):\n \"\"\"\n Configuration for checkpoint engine to update weights from trainer to rollout\n \"\"\"\n\n # Backend for checkpoint engine: naive, nccl, nixl, hccl\n backend: Optional[str] = MISSING\n # Bucket size in MB to transfer multiple weights at one time\n update_weights_bucket_megabytes: int = 2048\n # Additional keyword arguments for checkpoint engine\n engine_kwargs: dict = field(default_factory=dict)\n\n\n@dataclass\nclass RolloutConfig(BaseConfig):\n _mutable_fields = {\"max_model_len\", \"load_format\"}\n\n name: Optional[str] = MISSING\n mode: str = \"async\"\n\n temperature: float = 1.0\n top_k: int = -1\n top_p: float = 1.0\n do_sample: bool = True\n n: int = 1\n repetition_penalty: float = 1.0\n\n # Early termination threshold for multi-turn rollout in sglang.\n # Abort remaining requests when (1 - over_sample_rate) * total_requests are completed.\n over_sample_rate: float = 0.0\n\n prompt_length: int = 512\n response_length: int = 512\n\n dtype: str = \"bfloat16\"\n gpu_memory_utilization: float = 0.5\n ignore_eos: bool = False\n enforce_eager: bool = True\n cudagraph_capture_sizes: Optional[list] = None\n free_cache_engine: bool = True\n data_parallel_size: int = 1\n expert_parallel_size: int = 1\n tensor_model_parallel_size: int = 2\n pipeline_model_parallel_size: int = 1\n max_num_batched_tokens: int = 8192\n logprobs_mode: Optional[str] = \"processed_logprobs\"\n scheduling_policy: Optional[str] = \"fcfs\"\n\n # TODO: enable train_kwargs\n # train_sampling_config: SamplingConfig = field(default_factory=SamplingConfig)\n\n val_kwargs: SamplingConfig = field(default_factory=SamplingConfig)\n\n max_model_len: Optional[int] = None\n max_num_seqs: int = 1024\n\n # note that the logprob computation should belong to the actor\n log_prob_micro_batch_size: Optional[int] = None\n log_prob_micro_batch_size_per_gpu: Optional[int] = None\n log_prob_use_dynamic_bsz: bool = False\n log_prob_max_token_len_per_gpu: int = 16384\n\n disable_log_stats: bool = True\n\n multi_stage_wake_up: bool = False\n engine_kwargs: dict = field(default_factory=dict)\n\n calculate_log_probs: bool = False\n\n agent: AgentLoopConfig = field(default_factory=AgentLoopConfig)\n\n trace: TraceConfig = field(default_factory=TraceConfig)\n\n multi_turn: MultiTurnConfig = field(default_factory=MultiTurnConfig)\n\n # Server configuration for sglang server mode\n server: ServerConfig = field(default_factory=ServerConfig)\n\n # Use Prometheus to collect and monitor rollout statistics\n prometheus: PrometheusConfig = field(default_factory=PrometheusConfig)\n\n # Extension point for custom configurations\n custom: Optional[dict] = None\n\n # Checkpoint Engine config for update weights from trainer to rollout\n checkpoint_engine: CheckpointEngineConfig = field(default_factory=CheckpointEngineConfig)\n\n skip_rollout: bool = False\n\n skip_dump_dir: str = \"/tmp/rollout_dump\"\n\n profiler: Optional[ProfilerConfig] = None\n\n enable_chunked_prefill: bool = True\n\n enable_prefix_caching: bool = True\n\n load_format: str = \"dummy\"\n\n layered_summon: bool = False\n\n layer_name_map: dict = field(default_factory=dict)\n\n sglang_engine_mode: str = \"local\"\n\n limit_images: Optional[int] = None\n\n skip_tokenizer_init: bool = False\n\n quantization: Optional[str] = None\n\n quantization_config_file: Optional[str] = None\n\n enable_rollout_routing_replay: bool = False\n\n enable_sleep_mode: bool = True\n\n mtp: MtpConfig = field(default_factory=MtpConfig)\n\n qat: Optional[dict] = None\n\n def __post_init__(self):\n \"\"\"Validate the rollout config\"\"\"\n # Deprecation warning for mode field - only async mode is supported\n if self.mode == \"sync\":\n raise ValueError(\n \"Rollout mode 'sync' has been removed. Please set \"\n \"`actor_rollout_ref.rollout.mode=async` or remove the mode setting entirely.\"\n )\n if self.mode != \"async\":\n warnings.warn(\n f\"Unknown rollout mode '{self.mode}'. Only 'async' mode is supported. \"\n \"The 'mode' field is deprecated and will be removed in a future version.\",\n DeprecationWarning,\n stacklevel=2,\n )\n\n if self.expert_parallel_size > 1:\n assert self.expert_parallel_size == (self.tensor_model_parallel_size * self.data_parallel_size), (\n \"expert_parallel_size must be equal to tensor_model_parallel_size * data_parallel_size\"\n )\n\n if self.pipeline_model_parallel_size > 1:\n if self.name == \"vllm\" or self.name == \"sglang\" or self.name == \"trtllm\":\n raise NotImplementedError(\n f\"Current rollout {self.name=} not implemented pipeline_model_parallel_size > 1 yet.\"\n )\n"}138{"file_name": "verl__workers__critic__base.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nBase class for a critic\n\"\"\"\n\nfrom abc import ABC, abstractmethod\n\nimport torch\n\nfrom verl import DataProto\n\n__all__ = [\"BasePPOCritic\"]\n\n\nclass BasePPOCritic(ABC):\n def __init__(self, config):\n super().__init__()\n self.config = config\n\n @abstractmethod\n def compute_values(self, data: DataProto) -> torch.Tensor:\n \"\"\"Compute values\"\"\"\n pass\n\n @abstractmethod\n def update_critic(self, data: DataProto):\n \"\"\"Update the critic\"\"\"\n pass\n"}139{"file_name": "verl__workers__critic__megatron_critic.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nImplement a multiprocess PPOCritic\n\"\"\"\n\nimport itertools\nimport logging\nimport os\nfrom functools import partial\nfrom typing import Iterable\n\nimport torch\nimport torch.distributed\nfrom megatron.core import parallel_state as mpu\nfrom megatron.core.optimizer import DistributedOptimizer, OptimizerConfig\nfrom megatron.core.pipeline_parallel import get_forward_backward_func\nfrom omegaconf import OmegaConf\nfrom torch import nn\n\nfrom verl import DataProto\nfrom verl.trainer.ppo import core_algos\nfrom verl.utils.device import get_device_id, get_torch_device\nfrom verl.utils.megatron.pipeline_parallel import make_batch_generator\nfrom verl.utils.profiler import GPUMemoryLogger\nfrom verl.utils.py_functional import append_to_dict\nfrom verl.utils.seqlen_balancing import get_reverse_idx, rearrange_micro_batches\nfrom verl.utils.torch_functional import broadcast_dict_tensor, masked_mean\nfrom verl.workers.critic import BasePPOCritic\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\nclass MegatronPPOCritic(BasePPOCritic):\n def __init__(\n self,\n config,\n model_config,\n hf_config,\n tf_config,\n critic_module: nn.ModuleList,\n critic_optimizer: DistributedOptimizer,\n critic_optimizer_config: OptimizerConfig,\n ):\n super().__init__(config=config)\n self._validate_config(config)\n self.model_config = model_config\n self.hf_config = hf_config # huggingface config\n self.tf_config = tf_config # mcore transformer config\n\n self.critic_module = critic_module\n self.critic_optimizer = critic_optimizer\n self.critic_optimizer_config = critic_optimizer_config\n\n # we create a separate nametuple for optimizer step so that global args won't affect it.\n self.optimizer_step_args = OmegaConf.create(\n {\n \"skip_grad\": None,\n \"overlap_dp_param_comm\": False,\n \"overlap_dp_grad_comm\": False,\n \"gradient_accumulation_steps\": 1,\n \"sequence_parallel\": self.tf_config.sequence_parallel,\n \"DDP_impl\": \"local\",\n \"layernorm_allreduce_bucket_threshold\": 0,\n \"reduce_grads_use_alltoall\": False,\n }\n )\n\n def _validate_config(self, config) -> None:\n \"\"\"Validate config options not implemented for Megatron backend\"\"\"\n assert config.get(\"ulysses_sequence_parallel_size\", 1) == 1\n if config.shuffle:\n assert config.data_loader_seed is not None, \"If shuffle dataloader, seed must be manually set\"\n self.config = config\n\n @GPUMemoryLogger(\"megatron critic\", logger=logger)\n def compute_values(self, data: DataProto) -> DataProto:\n prev_modes = [m.training for m in self.critic_module]\n for module in self.critic_module:\n module.eval()\n responses = data.batch[\"responses\"]\n attention_mask = data.batch[\"attention_mask\"]\n use_dynamic_bsz = data.meta_info.get(\"use_dynamic_bsz\", False)\n micro_batch_size = data.meta_info.get(\"micro_batch_size\", None)\n max_token_len = data.meta_info.get(\"max_token_len\", None)\n assert micro_batch_size is not None, \"micro batch size is needed for forward compute\"\n if use_dynamic_bsz:\n assert max_token_len is not None, \"max_token_len must be set when use_dynamic_bsz is True\"\n max_token_len = max_token_len * self.config.megatron.context_parallel_size\n response_length = responses.size(1)\n with torch.no_grad():\n output = self.forward_backward_batch(\n data=data,\n forward_only=True,\n use_dynamic_bsz=use_dynamic_bsz,\n micro_batch_size=micro_batch_size,\n max_token_len=max_token_len,\n mini_batch_size=None,\n )\n if mpu.is_pipeline_last_stage(ignore_virtual=True):\n # only on last rank. It should be on every tp rank\n values = [o[\"vpreds\"] for o in output[\"output\"]] # (bs, seq_size, vocal_size)\n values = torch.cat(values, dim=0).to(torch.float32)\n if use_dynamic_bsz:\n indices = output[\"indices\"]\n indices = list(itertools.chain.from_iterable(indices))\n assert len(indices) == values.size(0), f\"{len(indices)} vs. {values.size()}\"\n revert_indices = torch.tensor(get_reverse_idx(indices), dtype=torch.long)\n values = values[revert_indices]\n else:\n values = torch.empty_like(attention_mask, dtype=torch.float32)\n\n # each tp ranks should contain the same value\n values = values[\n :, -response_length - 1 : -1\n ] # Values are predicted at the ends of prefixes, e.g., the last prompt token\n response_mask = attention_mask[:, -response_length:]\n values = values * response_mask # Only action tokens have values\n values = values.contiguous()\n\n # sync among pp ranks\n values = values.to(get_device_id())\n torch.distributed.broadcast(\n tensor=values,\n src=mpu.get_pipeline_model_parallel_last_rank(),\n group=mpu.get_pipeline_model_parallel_group(),\n )\n values = values.to(\"cpu\")\n\n # add empty cache after each compute\n get_torch_device().empty_cache()\n\n for module, mode in zip(self.critic_module, prev_modes, strict=False):\n module.train(mode)\n return values\n\n def make_minibatch_iterator(self, data: DataProto) -> Iterable[DataProto]:\n select_keys = [\"input_ids\", \"responses\", \"attention_mask\", \"position_ids\", \"values\", \"returns\"]\n data = data.select(batch_keys=select_keys)\n return data.make_iterator(\n mini_batch_size=self.config.ppo_mini_batch_size,\n epochs=self.config.ppo_epochs,\n seed=self.config.data_loader_seed,\n dataloader_kwargs={\"shuffle\": self.config.shuffle},\n )\n\n def forward_backward_batch(\n self,\n data: DataProto,\n forward_only=False,\n use_dynamic_bsz=False,\n micro_batch_size=None,\n max_token_len=None,\n mini_batch_size=None,\n ):\n # broadcast from last pp rank to all other pp ranks\n data.to(get_device_id())\n mini_batch = data\n mini_batch.batch = mini_batch.batch.contiguous()\n broadcast_dict_tensor(\n mini_batch.batch,\n src=mpu.get_pipeline_model_parallel_last_rank(),\n group=mpu.get_pipeline_model_parallel_group(),\n )\n mini_batch.to(\"cpu\")\n # split into micro-batches\n mini_batch.batch[\"attention_mask\"] = mini_batch.batch[\"attention_mask\"].to(bool)\n\n indices = None\n if use_dynamic_bsz:\n assert max_token_len is not None, \"max_token_len must be set when use_dynamic_bsz is True\"\n vpp_size = mpu.get_virtual_pipeline_model_parallel_world_size()\n if vpp_size is not None and vpp_size > 1:\n microbatch_group_size_per_vp_stage = self.tf_config.microbatch_group_size_per_vp_stage\n micro_batches, indices = rearrange_micro_batches(\n batch=mini_batch.batch,\n num_batches_divided_by=microbatch_group_size_per_vp_stage,\n max_token_len=max_token_len,\n )\n assert len(micro_batches) % self.tf_config.microbatch_group_size_per_vp_stage == 0, (\n f\"micro_batches {micro_batches} must be divisible by microbatch_group_size_per_vp_stage \"\n f\"{microbatch_group_size_per_vp_stage} for megatron backend\"\n )\n else:\n micro_batches, indices = rearrange_micro_batches(batch=mini_batch.batch, max_token_len=max_token_len)\n total_seqlen = max_token_len\n else:\n assert micro_batch_size is not None, (\n \"micro_batch_size is needed to be passed in when not using dynamic batch size\"\n )\n micro_batches = mini_batch.batch.split(micro_batch_size)\n seq_len = micro_batches[0][\"input_ids\"].shape[1]\n total_seqlen = micro_batch_size * seq_len\n n_micro_batch = len(micro_batches)\n\n forward_backward_func = get_forward_backward_func()\n\n def loss_func(output, data, meta_info):\n nonlocal use_dynamic_bsz\n\n if forward_only:\n return torch.tensor(1.0, device=output.device), {\"vpreds\": output}\n\n responses = data[\"responses\"]\n attention_mask = data[\"attention_mask\"]\n values = data[\"values\"]\n returns = data[\"returns\"]\n response_length = responses.size(1)\n\n response_mask = attention_mask[:, -response_length:]\n\n cliprange_value = self.config.cliprange_value\n\n vpreds = output # (bs, sequence_length)\n vpreds = vpreds[:, -response_length - 1 : -1]\n\n vf_loss, vf_clipfrac = core_algos.compute_value_loss(\n vpreds=vpreds,\n values=values,\n returns=returns,\n response_mask=response_mask,\n cliprange_value=cliprange_value,\n loss_agg_mode=self.config.loss_agg_mode,\n )\n\n stats = {\n \"critic/vf_loss\": vf_loss.detach().item(),\n \"critic/vf_clipfrac\": vf_clipfrac.detach().item(),\n \"critic/vpred_mean\": masked_mean(vpreds, response_mask).detach().item(),\n }\n\n return vf_loss, stats\n\n def forward_step(batch_iter, model):\n batch = next(batch_iter)\n batch = batch.to(get_device_id())\n batch = batch.contiguous()\n\n input_ids = batch[\"input_ids\"]\n attention_mask = batch[\"attention_mask\"]\n position_ids = batch[\"position_ids\"]\n from verl.models.mcore import get_mcore_forward_fn\n\n forward_fn = get_mcore_forward_fn(self.hf_config)\n\n output = forward_fn(\n model,\n input_ids,\n attention_mask,\n position_ids,\n {}, # multi_modal_inputs\n value_model=True,\n )\n\n return output, partial(loss_func, data=batch, meta_info={})\n\n # batch should be a list of batches inside micro-batches\n batch_generator = make_batch_generator(micro_batches, vpp_size=len(self.critic_module))\n\n # TODO: we may use the new schedule instead\n # for flash-attn: (seq_len, batch_size, hidden_size) = (mbs*seq_len, 1, hidden_size)\n if mpu.get_pipeline_model_parallel_world_size() > 1:\n losses_reduced = forward_backward_func(\n forward_step_func=forward_step,\n data_iterator=batch_generator,\n model=self.critic_module,\n num_microbatches=n_micro_batch,\n seq_length=total_seqlen, # no use when input_shapes was set\n micro_batch_size=1, # no use when input_shapes was set\n forward_only=forward_only,\n )\n else:\n losses_reduced = forward_backward_func(\n forward_step_func=forward_step,\n data_iterator=batch_generator,\n model=self.critic_module,\n num_microbatches=n_micro_batch,\n seq_length=total_seqlen, # in use for pp = 1\n micro_batch_size=1, # in use for pp = 1\n forward_only=forward_only,\n )\n # loss_reduces contains the stats returned from loss_func\n losses_reduced = {\"output\": losses_reduced}\n if use_dynamic_bsz:\n losses_reduced[\"indices\"] = indices\n return losses_reduced\n\n @GPUMemoryLogger(\"megatron critic\", logger=logger)\n def update_critic(self, dataloader: Iterable[DataProto]):\n metrics = {}\n\n for data in dataloader:\n self.critic_optimizer.zero_grad()\n # use use_contiguous_buffers_in_local_ddp and no overlap_dp_param_comm\n for chunk in self.critic_module:\n chunk.zero_grad_buffer()\n\n micro_batch_size = self.config.ppo_micro_batch_size_per_gpu\n max_token_len = None\n if self.config.use_dynamic_bsz:\n max_token_len = self.config.ppo_max_token_len_per_gpu * self.config.megatron.context_parallel_size\n metric_micro_batch = self.forward_backward_batch(\n data,\n forward_only=False,\n use_dynamic_bsz=self.config.use_dynamic_bsz,\n micro_batch_size=micro_batch_size,\n max_token_len=max_token_len,\n mini_batch_size=self.config.ppo_mini_batch_size,\n )\n metric_micro_batch = metric_micro_batch[\"output\"]\n update_successful, grad_norm, num_zeros_in_grad = self.critic_optimizer.step()\n learning_rate = self.critic_optimizer.param_groups[-1][\"lr\"]\n data = {\"critic/grad_norm\": grad_norm, \"critic/lr\": learning_rate}\n append_to_dict(metrics, data)\n\n if update_successful:\n # allgather already execute in optimizer.step in new megatron\n pass\n else:\n raise NotImplementedError\n\n for metric in metric_micro_batch:\n append_to_dict(metrics, metric) # append the metric from this micro-batch to global metrics.\n\n # add empty cache after each compute\n get_torch_device().empty_cache()\n return metrics\n"}140{"file_name": "verl__workers__engine__base.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nThe abstract base class defining the interface for model training engines.\n\"\"\"\n\nfrom abc import abstractmethod\nfrom contextlib import nullcontext\nfrom typing import Any, Callable, ContextManager, Generator, Optional\n\nimport torch\nfrom tensordict import TensorDict\n\nfrom verl.utils.device import get_device_name\nfrom verl.utils.tensordict_utils import maybe_fix_3d_position_ids\n\n\nclass BaseEngine:\n \"\"\"\n Abstract base class defining the interface for model training engines. Interface is subject to\n change before release.\n\n Engine implementations must subclass BaseEngine and provide concrete behavior for all methods.\n \"\"\"\n\n def initialize(self):\n \"\"\"\n Instantiate or load the model, optimizer, and learning rate scheduler.\n\n Should prepare all components necessary for training or evaluation.\n \"\"\"\n raise NotImplementedError\n\n @property\n @abstractmethod\n def is_param_offload_enabled(self) -> bool:\n \"\"\"Whether parameter offloading is enabled.\"\"\"\n raise NotImplementedError\n\n @property\n @abstractmethod\n def is_optimizer_offload_enabled(self) -> bool:\n \"\"\"Whether optimizer offloading is enabled.\"\"\"\n raise NotImplementedError\n\n def train_mode(self, **kwargs):\n \"\"\"\n Context manager entry for switching the engine and model into training mode.\n\n Usage:\n with engine.train_mode():\n # runs in training mode\n \"\"\"\n raise NotImplementedError\n\n def eval_mode(self, **kwargs):\n \"\"\"\n Context manager entry for switching the engine and model into evaluation mode.\n\n Usage:\n with engine.eval_mode():\n # runs in evaluation mode\n \"\"\"\n raise NotImplementedError\n\n def optimizer_zero_grad(self):\n \"\"\"\n Zero the gradients of the optimizer.\n \"\"\"\n raise NotImplementedError\n\n def optimizer_step(self):\n \"\"\"\n Perform an optimization step using the optimizer.\n \"\"\"\n raise NotImplementedError\n\n def lr_scheduler_step(self):\n \"\"\"\n Advance the learning rate scheduler by one step.\n\n Returns:\n current_lr (float or list[float]): Updated learning rate(s).\n \"\"\"\n raise NotImplementedError\n\n def forward_backward_batch(self, data: TensorDict, loss_function: Callable, forward_only=False) -> Any:\n \"\"\"\n Perform a forward pass and optionally a backward pass on a batch of data.\n\n Args:\n data: The input data for the forward pass, typically containing tensors and metadata.\n loss_function: The loss function to optimize. See `verl.workers.roles.utils.losses` for examples.\n forward_only: If True, perform only the forward pass. If False, perform forward and backward pass.\n\n Returns:\n Any: The output of the forward pass, which can be used for loss computation or other purposes.\n \"\"\"\n raise NotImplementedError\n\n def train_batch(self, data: TensorDict, loss_function: Callable) -> Any:\n \"\"\"\n Perform a training step on a batch of data.\n\n Args:\n data: The input data for training, typically containing tensors and metadata.\n loss_function: A function that computes the loss and metrics given a batch and predictions.\n\n Returns:\n dict[str, torch.Tensor]: A dictionary containing the aggregated training metrics for the batch.\n \"\"\"\n maybe_fix_3d_position_ids(data)\n\n self.optimizer_zero_grad()\n outputs = self.forward_backward_batch(data, loss_function, forward_only=False)\n grad_norm = self.optimizer_step()\n if self.is_mp_src_rank_with_outputs():\n assert \"grad_norm\" not in outputs[\"metrics\"]\n outputs[\"metrics\"][\"grad_norm\"] = grad_norm\n return outputs\n\n def infer_batch(self, data: TensorDict, loss_function: Optional[Callable] = None) -> Any:\n \"\"\"\n Perform inference on a batch of data.\n\n Args:\n data: The input data for inference, typically containing tensors and metadata.\n\n Returns:\n Any: The output of the inference, which can be used for predictions or other purposes.\n \"\"\"\n # see comments from train_batch\n maybe_fix_3d_position_ids(data)\n\n with torch.no_grad():\n outputs = self.forward_backward_batch(data, loss_function, forward_only=True)\n return outputs\n\n def get_per_tensor_param(self) -> tuple[Generator[tuple[str, torch.Tensor], None, None], Optional[dict]]:\n \"\"\"\n Get a generator that yields per-tensor parameters and optional peft config.\n\n Returns:\n Generator[tuple[str, torch.Tensor]]: A generator that yields tuples of parameter names and tensors.\n Optional[dict]: Optional peft config.\n \"\"\"\n raise NotImplementedError\n\n def get_data_parallel_size(self):\n raise NotImplementedError\n\n def get_data_parallel_rank(self):\n raise NotImplementedError\n\n def get_data_parallel_group(self):\n raise NotImplementedError\n\n def to(self, device: str, model: bool = True, optimizer: bool = True, grad: bool = True):\n \"\"\"\n Move model parameters, optimizer states, or both to the specified device.\n\n Args:\n device: Target device identifier.\n model: If True, move the model.\n optimizer: If True, move the optimizer states.\n grad: If True, move the gradient buffer.\n \"\"\"\n if not model:\n assert not optimizer and not grad, \"Model must be moved to device along with optimizer and grad\"\n\n def save_checkpoint(\n self,\n local_path: str,\n hdfs_path: Optional[str] = None,\n global_step: int = 0,\n max_ckpt_to_keep: Optional[int] = None,\n **kwargs,\n ) -> None:\n \"\"\"\n Save model, optimizer, and scheduler states to a checkpoint.\n\n Args:\n local_path: Local filesystem path to save checkpoint.\n hdfs_path: Optional HDFS path to copy checkpoint.\n global_step: Integer training step number for naming.\n max_ckpt_to_keep: Maximum number of recent checkpoints to retain.\n **kwargs: Arbitrary keyword arguments.\n \"\"\"\n raise NotImplementedError\n\n def load_checkpoint(\n self, local_path: str, hdfs_path: Optional[str] = None, del_local_after_load: bool = True, **kwargs\n ) -> None:\n \"\"\"\n Load model, optimizer, and scheduler states from a checkpoint.\n\n Args:\n local_path: Local filesystem path of the checkpoint.\n hdfs_path: Optional HDFS path where checkpoint is stored.\n del_local_after_load: Whether to delete local copy after loading.\n **kwargs: Arbitrary keyword arguments.\n \"\"\"\n raise NotImplementedError\n\n def is_mp_src_rank_with_outputs(self):\n \"\"\"\n Whether the current rank is the first rank in model parallel group that contains model outputs\n \"\"\"\n raise NotImplementedError\n\n def disable_adapter(self) -> ContextManager:\n \"\"\"\n Disable all adapters temporarily under the context in the model for LoRA\n \"\"\"\n return nullcontext()\n\n\nclass BaseEngineCtx:\n def __init__(self, engine: BaseEngine, mode, **kwargs):\n \"\"\"Base Engine context that handles load and offload\n\n Args:\n engine:\n **kwargs:\n \"\"\"\n self.engine = engine\n self.mode = mode\n assert self.mode in (\"train\", \"eval\")\n self.disable_auto_offload = kwargs.pop(\"disable_auto_offload\", False)\n\n def _context_switch(self, device):\n if self.disable_auto_offload:\n return\n should_move_model = self.engine.is_param_offload_enabled if device == \"cpu\" else True\n should_move_optimizer = self.engine.is_optimizer_offload_enabled if device == \"cpu\" else True\n if self.mode == \"eval\":\n self.engine.to(device=device, model=should_move_model, optimizer=False, grad=False)\n elif self.mode == \"train\":\n self.engine.to(\n device=device,\n model=should_move_model,\n optimizer=should_move_optimizer,\n grad=should_move_model,\n )\n\n def __enter__(self):\n self._context_switch(get_device_name())\n self.engine.mode = self.mode\n\n def __exit__(self, exc_type, exc_val, exc_tb):\n self._context_switch(\"cpu\")\n self.engine.mode = None\n\n\nclass EngineRegistry:\n \"\"\"\n A registry for managing and instantiating different types of training engines.\n\n This class uses a dictionary to store engine classes, mapping a string key to each class.\n It provides a decorator `register` to add new engines to the registry and a `new` method\n to create an instance of a registered engine.\n \"\"\"\n\n _engines = {}\n\n @classmethod\n def register(cls, model_type: str, backend: list[str] | str, device: list[str] | str = \"cuda\"):\n \"\"\"\n A class method decorator that registers an engine class with a given key.\n\n This allows for dynamic instantiation of engine classes by their registered key.\n\n Args:\n model_type (str): The type of the model\n backend (list[str] | str): The backend to use for the model type\n device (list[str] | str): The device type (e.g., \"cuda\", \"npu\", \"cpu\") this engine supports,\n default is \"cuda\"\n\n Returns:\n A decorator function that takes an engine class and registers it.\n \"\"\"\n\n def decorator(engine_class):\n assert issubclass(engine_class, BaseEngine)\n if model_type not in cls._engines:\n cls._engines[model_type] = {}\n\n backends = backend if isinstance(backend, list) else [backend]\n devices = device if isinstance(device, list) else [device]\n for current_backend in backends:\n for current_device in devices:\n if current_backend not in cls._engines[model_type]:\n cls._engines[model_type][current_backend] = {}\n if current_device not in cls._engines[model_type][current_backend]:\n cls._engines[model_type][current_backend][current_device] = engine_class\n\n return engine_class\n\n return decorator\n\n @classmethod\n def get_engine_cls(cls, model_type: str, backend: str):\n assert model_type in cls._engines, f\"Unknown model_type: {model_type}\"\n assert backend in cls._engines[model_type], f\"Unknown backend: {backend}\"\n device = get_device_name()\n assert device in cls._engines[model_type][backend], (\n f\"Unknown device: {device} for model_type: {model_type} and backend: {backend}\"\n )\n return cls._engines[model_type][backend][device]\n\n @classmethod\n def new(cls, model_type, backend, *args, **kwargs):\n \"\"\"\n Function to create a new training engine instance based on the provided config.\n Args:\n key: A configuration object containing the engine key and other settings.\n *args: Variable length argument list.\n **kwargs: Arbitrary keyword arguments.\n Returns:\n engine: An instance of the training engine corresponding to the config.\n Raises:\n NotImplementedError: If the engine key in the config does not match any known engines.\n \"\"\"\n engine_cls = cls.get_engine_cls(model_type, backend)\n return engine_cls(*args, **kwargs)\n"}141{"file_name": "verl__workers__engine__fsdp__transformer_impl.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nThe concrete Engine implementation using PyTorch FullyShardedDataParallel (FSDP)\n\"\"\"\n\nimport gc\nimport logging\nimport os\nimport warnings\nfrom contextlib import nullcontext\nfrom typing import Callable, ContextManager, Optional\n\nimport torch\nimport torch.distributed\nfrom peft import LoraConfig, TaskType, get_peft_model\nfrom tensordict import TensorDict\nfrom torch.distributed.fsdp import FullyShardedDataParallel as FSDP\nfrom torch.distributed.fsdp.api import FullStateDictConfig, ShardedStateDictConfig, StateDictType\nfrom torch.distributed.tensor import DTensor\n\nimport verl.utils.torch_functional as verl_F\nfrom verl.models.transformers.monkey_patch import apply_monkey_patch\nfrom verl.trainer.config import CheckpointConfig\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.activation_offload import enable_activation_offloading\nfrom verl.utils.checkpoint.fsdp_checkpoint_manager import FSDPCheckpointManager\nfrom verl.utils.dataset.dataset_utils import DatasetPadMode\nfrom verl.utils.debug import log_gpu_memory_usage\nfrom verl.utils.device import get_device_id, get_device_name\nfrom verl.utils.fsdp_utils import (\n CPUOffloadPolicy,\n FSDPModule,\n MixedPrecisionPolicy,\n apply_fsdp2,\n collect_lora_params,\n fsdp2_clip_grad_norm_,\n fsdp2_load_full_state_dict,\n fsdp_version,\n get_fsdp_wrap_policy,\n get_init_weight_context_manager,\n init_fn,\n load_fsdp_model_to_gpu,\n load_fsdp_optimizer,\n merged_lora_context,\n normalize_peft_param_name,\n offload_fsdp_model_to_cpu,\n offload_fsdp_optimizer,\n replace_lora_wrapper,\n)\nfrom verl.utils.model import convert_weight_keys, extract_multi_modal_inputs\nfrom verl.utils.py_functional import convert_to_regular_types\nfrom verl.utils.torch_functional import logprobs_from_logits\nfrom verl.utils.ulysses import (\n gather_outputs_and_unpad,\n get_ulysses_sequence_parallel_group,\n set_ulysses_sequence_parallel_group,\n ulysses_pad,\n ulysses_pad_and_slice_inputs,\n)\nfrom verl.workers.config import FSDPEngineConfig, FSDPOptimizerConfig, HFModelConfig\n\nfrom ..base import BaseEngine, BaseEngineCtx, EngineRegistry\nfrom ..utils import enable_full_determinism, postprocess_batch_func, prepare_micro_batches\nfrom .utils import create_device_mesh, get_sharding_strategy\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\ndevice_name = get_device_name()\n\n\nclass FSDPEngine(BaseEngine):\n \"\"\"\n Concrete Engine implementation using PyTorch FullyShardedDataParallel (FSDP).\n\n Supports model sharding, activation/optimizer offloading, LoRA, and sequence parallelism.\n \"\"\"\n\n def __init__(\n self,\n model_config: HFModelConfig,\n engine_config: FSDPEngineConfig,\n optimizer_config: FSDPOptimizerConfig,\n checkpoint_config: CheckpointConfig,\n ):\n \"\"\"\n Initialize the FSDPEngine.\n\n Sets up distributed device meshes, LoRA, and offload policies based on config.\n\n Args:\n config: Configuration object with FSDP and model settings.\n \"\"\"\n super().__init__()\n\n self.model_config = model_config\n self.engine_config = engine_config\n self.optimizer_config = optimizer_config\n self.checkpoint_config = checkpoint_config\n\n self.mode = None\n\n self.rank = torch.distributed.get_rank()\n\n # Apply NPU patches for FSDP backend\n from .utils import apply_npu_fsdp_patches\n\n apply_npu_fsdp_patches()\n\n # build device mesh for Ulysses Sequence Parallel\n\n self.use_remove_padding = self.model_config.use_remove_padding\n\n self._init_device_mesh()\n\n if self.engine_config.full_determinism:\n enable_full_determinism(seed=self.engine_config.seed)\n\n # set FSDP offload params\n self._is_offload_param = self.engine_config.param_offload\n self._is_offload_optimizer = self.engine_config.optimizer_offload\n self._is_lora = self.model_config.lora_rank > 0\n\n if self.engine_config.entropy_from_logits_with_chunking:\n entropy_from_logits = verl_F.entropy_from_logits_with_chunking\n else:\n entropy_from_logits = verl_F.entropy_from_logits\n\n self.compute_entropy_from_logits = (\n torch.compile(entropy_from_logits, dynamic=True)\n if self.engine_config.use_torch_compile # use torch compile by default\n else entropy_from_logits\n )\n\n @property\n def is_param_offload_enabled(self) -> bool:\n return self._is_offload_param\n\n @property\n def is_optimizer_offload_enabled(self) -> bool:\n return self._is_offload_optimizer\n\n def is_mp_src_rank_with_outputs(self):\n if self.ulysses_device_mesh is not None:\n is_collect = self.ulysses_device_mesh[\"sp\"].get_local_rank() == 0\n else:\n is_collect = True\n return is_collect\n\n def initialize(self):\n \"\"\"\n Build the model, optimizer, and learning rate scheduler under FSDP.\n\n Applies device, dtype, and precision configurations, including mixed precision.\n Sets up checkpoint manager and FLOPs counter.\n \"\"\"\n # This is used to import external_lib into the huggingface systems\n self._build_model_optimizer()\n\n self.checkpoint_manager = FSDPCheckpointManager(\n model=self.module,\n optimizer=self.optimizer,\n lr_scheduler=self.lr_scheduler,\n processing_class=self.model_config.get_processor(),\n checkpoint_config=self.checkpoint_config,\n trust_remote_code=self.model_config.trust_remote_code,\n )\n\n self.to(\n device=\"cpu\",\n model=self._is_offload_param,\n optimizer=self._is_offload_optimizer,\n grad=self._is_offload_param,\n )\n\n log_gpu_memory_usage(\"After offload model/optimizer/grad during init\", logger=logger)\n\n def _init_device_mesh(self):\n world_size = torch.distributed.get_world_size()\n from torch.distributed.device_mesh import init_device_mesh\n\n fsdp_size = self.engine_config.fsdp_size\n\n self.device_mesh = create_device_mesh(world_size=world_size, fsdp_size=fsdp_size)\n self.ulysses_device_mesh = None\n self.ulysses_parallel_group = None\n self.ulysses_sequence_parallel_size = self.engine_config.ulysses_sequence_parallel_size\n dp_size = self.get_data_parallel_size()\n if self.ulysses_sequence_parallel_size > 1:\n self.ulysses_device_mesh = init_device_mesh(\n device_name, mesh_shape=(dp_size, self.ulysses_sequence_parallel_size), mesh_dim_names=[\"dp\", \"sp\"]\n )\n self.ulysses_parallel_group = self.ulysses_device_mesh[\"sp\"].get_group()\n\n self.use_ulysses_sp = self.ulysses_sequence_parallel_size > 1\n\n def _build_module(self):\n from verl.utils.model import get_hf_auto_model_class\n from verl.utils.torch_dtypes import PrecisionType\n\n torch_dtype = self.engine_config.model_dtype\n\n if torch_dtype is None:\n # if it is training, we force torch_dtype to fp32\n torch_dtype = torch.float32 if not self.engine_config.forward_only else torch.bfloat16\n\n torch_dtype = PrecisionType.to_dtype(torch_dtype)\n\n init_context = get_init_weight_context_manager(\n use_meta_tensor=not self.model_config.hf_config.tie_word_embeddings, mesh=self.device_mesh\n )\n\n with init_context(), warnings.catch_warnings():\n warnings.simplefilter(\"ignore\")\n\n auto_class = get_hf_auto_model_class(hf_config=self.model_config.hf_config)\n\n module = auto_class.from_pretrained(\n pretrained_model_name_or_path=self.model_config.local_path,\n torch_dtype=torch_dtype,\n config=self.model_config.hf_config,\n trust_remote_code=self.model_config.trust_remote_code,\n )\n\n use_liger = self.model_config.use_liger\n # Apply Liger kernel to the model if use_liger is set to True\n if use_liger:\n from liger_kernel.transformers.monkey_patch import _apply_liger_kernel_to_instance\n\n _apply_liger_kernel_to_instance(model=module)\n\n fused_kernel_options = self.model_config.fused_kernel_options\n fused_kernels_backend = (\n fused_kernel_options.get(\"impl_backend\", None) if fused_kernel_options is not None else None\n )\n\n use_fused_kernels = self.model_config.use_fused_kernels\n apply_monkey_patch(\n model=module,\n use_remove_padding=self.use_remove_padding,\n ulysses_sp_size=self.ulysses_sequence_parallel_size,\n use_fused_kernels=use_fused_kernels,\n fused_kernels_backend=fused_kernels_backend,\n )\n\n # some parameters may not in torch_dtype\n module.to(torch_dtype)\n\n if self.model_config.enable_gradient_checkpointing:\n module.gradient_checkpointing_enable(gradient_checkpointing_kwargs={\"use_reentrant\": False})\n return module\n\n def _build_lora_module(self, module):\n module.enable_input_require_grads()\n\n lora_adapter_path = getattr(self.model_config, \"lora_adapter_path\", None)\n if lora_adapter_path is not None:\n from peft import PeftModel\n\n from verl.utils.fs import copy_to_local\n\n print(f\"Loading pre-trained LoRA adapter to from: {lora_adapter_path}\")\n # Copy adapter to local if needed\n local_adapter_path = copy_to_local(lora_adapter_path, use_shm=self.model_config.use_shm)\n\n module = PeftModel.from_pretrained(module, local_adapter_path, is_trainable=True)\n peft_config = module.peft_config[\"default\"]\n # Ensure task_type is TaskType enum, not string\n if isinstance(peft_config.task_type, str):\n peft_config.task_type = TaskType.CAUSAL_LM\n else:\n # Convert config to regular Python types before creating PEFT model\n lora_config = {\n \"task_type\": TaskType.CAUSAL_LM,\n \"r\": self.model_config.lora_rank,\n \"lora_alpha\": self.model_config.lora_alpha,\n \"target_modules\": convert_to_regular_types(self.model_config.target_modules),\n \"target_parameters\": convert_to_regular_types(self.model_config.target_parameters),\n \"exclude_modules\": convert_to_regular_types(self.model_config.exclude_modules),\n \"bias\": \"none\",\n }\n module = get_peft_model(module, LoraConfig(**lora_config))\n\n return module\n\n def _build_fsdp_module(self, module):\n # TODO(ziheng): need to improve\n from torch.distributed.fsdp import CPUOffload, MixedPrecision\n\n from verl.utils.torch_dtypes import PrecisionType\n\n mixed_precision_config = self.engine_config.mixed_precision\n if mixed_precision_config is not None:\n param_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"param_dtype\", \"bf16\"))\n reduce_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"reduce_dtype\", \"fp32\"))\n buffer_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"buffer_dtype\", \"fp32\"))\n else:\n param_dtype = torch.bfloat16\n reduce_dtype = torch.float32\n buffer_dtype = torch.float32\n\n mixed_precision = MixedPrecision(param_dtype=param_dtype, reduce_dtype=reduce_dtype, buffer_dtype=buffer_dtype)\n\n auto_wrap_policy = get_fsdp_wrap_policy(\n module=module,\n config=self.engine_config.wrap_policy,\n is_lora=self.model_config.lora_rank > 0,\n )\n\n fsdp_mesh = self.device_mesh\n sharding_strategy = get_sharding_strategy(fsdp_mesh)\n\n # Note: We force turn off CPUOffload because it causes incorrect results when using grad accumulation\n if self.engine_config.strategy == \"fsdp\":\n # cpu_offload:\n # - actor: None\n # - critic: None\n # - ref: CPUOffload(offload_params=True)\n\n # We force reference policy to use CPUOffload to save memory.\n # We force turn off CPUOffload for actor because it causes incorrect results when using grad accumulation\n cpu_offload = None\n if self.engine_config.forward_only:\n cpu_offload = CPUOffload(offload_params=True)\n self._is_offload_param = False\n self._is_offload_optimizer = False\n\n module = FSDP(\n module,\n param_init_fn=init_fn,\n auto_wrap_policy=auto_wrap_policy,\n device_id=get_device_id(),\n sharding_strategy=sharding_strategy,\n mixed_precision=mixed_precision,\n sync_module_states=True,\n device_mesh=self.device_mesh,\n forward_prefetch=self.engine_config.forward_prefetch,\n use_orig_params=self.engine_config.use_orig_params,\n cpu_offload=cpu_offload,\n )\n elif self.engine_config.strategy == \"fsdp2\":\n # - actor: offload_policy\n # - critic: offload_policy\n # - ref: CPUOffloadPolicy(pin_memory=True)\n assert CPUOffloadPolicy is not None, \"PyTorch version >= 2.4 is required for using fully_shard API (FSDP2)\"\n mp_policy = MixedPrecisionPolicy(\n param_dtype=param_dtype, reduce_dtype=reduce_dtype, cast_forward_inputs=True\n )\n offload_policy = None\n if self.engine_config.offload_policy or self.engine_config.forward_only:\n self._is_offload_param = False\n self._is_offload_optimizer = False\n offload_policy = CPUOffloadPolicy(pin_memory=True)\n\n fsdp_kwargs = {\n \"mesh\": fsdp_mesh,\n \"mp_policy\": mp_policy,\n \"offload_policy\": offload_policy,\n \"reshard_after_forward\": self.engine_config.reshard_after_forward,\n }\n full_state = module.state_dict()\n apply_fsdp2(module, fsdp_kwargs, self.engine_config)\n fsdp2_load_full_state_dict(module, full_state, fsdp_mesh, offload_policy)\n else:\n raise NotImplementedError(f\"Unknown strategy {self.engine_config.strategy}\")\n\n if self.model_config.enable_activation_offload:\n enable_gradient_checkpointing = self.model_config.enable_gradient_checkpointing\n enable_activation_offloading(module, self.engine_config.strategy, enable_gradient_checkpointing)\n\n if torch.distributed.get_world_size() == 1 and fsdp_version(module) == 1:\n FSDP.set_state_dict_type(\n module,\n state_dict_type=StateDictType.FULL_STATE_DICT,\n state_dict_config=FullStateDictConfig(),\n )\n elif fsdp_version(module) == 1:\n FSDP.set_state_dict_type(\n module,\n state_dict_type=StateDictType.SHARDED_STATE_DICT,\n state_dict_config=ShardedStateDictConfig(),\n )\n\n return module\n\n def _build_optimizer(self, module):\n from verl.workers.config.optimizer import build_optimizer\n\n optimizer = build_optimizer(module.parameters(), self.optimizer_config)\n\n return optimizer\n\n def _build_lr_scheduler(self, optimizer):\n from verl.utils.torch_functional import get_constant_schedule_with_warmup, get_cosine_schedule_with_warmup\n\n optim_config = self.optimizer_config\n\n total_steps = optim_config.total_training_steps\n num_warmup_steps = optim_config.lr_warmup_steps\n lr_scheduler_type = optim_config.lr_scheduler_type\n min_lr_ratio = optim_config.min_lr_ratio\n num_cycles = optim_config.num_cycles\n if num_warmup_steps <= 0:\n num_warmup_steps_ratio = optim_config.lr_warmup_steps_ratio\n num_warmup_steps = int(num_warmup_steps_ratio * total_steps)\n\n if self.rank == 0:\n print(f\"Total steps: {total_steps}, num_warmup_steps: {num_warmup_steps}\")\n\n if lr_scheduler_type == \"constant\":\n lr_scheduler = get_constant_schedule_with_warmup(optimizer=optimizer, num_warmup_steps=num_warmup_steps)\n elif lr_scheduler_type == \"cosine\":\n lr_scheduler = get_cosine_schedule_with_warmup(\n optimizer=optimizer,\n num_warmup_steps=num_warmup_steps,\n num_training_steps=total_steps,\n min_lr_ratio=min_lr_ratio,\n num_cycles=num_cycles,\n )\n else:\n raise NotImplementedError(f\"LR scheduler type {lr_scheduler_type} is not supported\")\n return lr_scheduler\n\n def _build_model_optimizer(self):\n from verl.utils.model import print_model_size\n\n # Load base model with specified configuration and dtype\n module = self._build_module()\n # Apply LoRA adapters if low-rank adaptation is enabled\n if self._is_lora:\n module = self._build_lora_module(module)\n\n # Synchronize all distributed processes before proceeding\n torch.distributed.barrier()\n if self.rank == 0:\n print_model_size(module)\n log_gpu_memory_usage(\"After init model from HF AutoModel\", logger=logger)\n\n # Wrap model with FSDP for distributed training (sharding, mixed precision, etc.)\n log_gpu_memory_usage(\"Before FSDP\", logger=None)\n module = self._build_fsdp_module(module)\n log_gpu_memory_usage(\"After FSDP\", logger=None)\n\n if not self.engine_config.forward_only:\n # Initialize optimizer with model parameters and config settings\n optimizer = self._build_optimizer(module)\n # Create learning rate scheduler with warmup and decay settings\n lr_scheduler = self._build_lr_scheduler(optimizer)\n else:\n optimizer = None\n lr_scheduler = None\n\n self.module = module\n self.optimizer = optimizer\n self.lr_scheduler = lr_scheduler\n\n def train_mode(self, **kwargs):\n \"\"\"\n Return a context manager that switches to training mode with FSDP-specific handling.\n\n Includes parameter and optimizer offload entry/exit.\n \"\"\"\n return EngineTrainModeCtx(self, **kwargs)\n\n def eval_mode(self, **kwargs):\n \"\"\"\n Return a context manager that switches to evaluation mode with FSDP-specific handling.\n\n Includes activation offload entry/exit.\n \"\"\"\n return EngineEvalModeCtx(self, **kwargs)\n\n def get_data_parallel_rank(self):\n if self.ulysses_device_mesh is not None:\n return self.ulysses_device_mesh[\"dp\"].get_local_rank()\n else:\n return torch.distributed.get_rank()\n\n def get_data_parallel_size(self):\n return torch.distributed.get_world_size() // self.ulysses_sequence_parallel_size\n\n def get_data_parallel_group(self):\n if self.ulysses_device_mesh is not None:\n return self.ulysses_device_mesh.get_group(mesh_dim=\"dp\")\n else:\n return torch.distributed.group.WORLD\n\n def forward_backward_batch(self, data: TensorDict, loss_function: Callable, forward_only=False) -> list[TensorDict]:\n # note that the global_batch_size should include data on all the dp\n tu.assign_non_tensor(data, sp_size=self.ulysses_sequence_parallel_size)\n\n # compute num_tokens in global batch for loss normalization\n batch_num_tokens = data[\"loss_mask\"].sum().to(get_device_id())\n torch.distributed.all_reduce(\n batch_num_tokens, op=torch.distributed.ReduceOp.SUM, group=self.get_data_parallel_group()\n )\n tu.assign_non_tensor(data, batch_num_tokens=batch_num_tokens.item())\n tu.assign_non_tensor(data, dp_size=self.get_data_parallel_size())\n\n micro_batches, indices = prepare_micro_batches(\n data=data, dp_group=self.get_data_parallel_group(), same_micro_num_in_dp=True\n )\n\n output_lst = []\n\n ctx = torch.no_grad() if forward_only else nullcontext()\n\n for micro_batch in micro_batches:\n with ctx:\n loss, meta_info = self.forward_step(micro_batch, loss_function=loss_function, forward_only=forward_only)\n\n if not forward_only:\n loss.backward()\n\n output_lst.append(meta_info)\n\n # postprocess and return\n return postprocess_batch_func(output_lst=output_lst, indices=indices, data=data)\n\n def forward_step(self, micro_batch: TensorDict, loss_function, forward_only):\n raise NotImplementedError(\"forward_step must be implemented in subclass\")\n\n def optimizer_zero_grad(self):\n \"\"\"\n Zero gradients and enforce FSDP grad-clipping logic.\n \"\"\"\n self.optimizer.zero_grad()\n\n def optimizer_step(self):\n \"\"\"\n Clip gradients, skip update if non-finite, and step optimizer.\n\n Returns:\n grad_norm (float): Norm of gradients before clipping.\n \"\"\"\n assert self.optimizer_config.clip_grad is not None\n\n if isinstance(self.module, FSDP):\n grad_norm = self.module.clip_grad_norm_(self.optimizer_config.clip_grad)\n elif isinstance(self.module, FSDPModule):\n grad_norm = fsdp2_clip_grad_norm_(self.module.parameters(), max_norm=self.optimizer_config.clip_grad)\n else:\n grad_norm = torch.nn.utils.clip_grad_norm_(\n self.module.parameters(), max_norm=self.optimizer_config.clip_grad\n )\n\n if isinstance(grad_norm, DTensor):\n grad_norm = grad_norm.full_tensor()\n\n # if grad_norm is not finite, skip the update\n if not torch.isfinite(grad_norm):\n print(f\"WARN: grad_norm is not finite: {grad_norm}\")\n self.optimizer.zero_grad()\n else:\n self.optimizer.step()\n return grad_norm.item()\n\n def lr_scheduler_step(self):\n \"\"\"\n Advance FSDP scheduler and return updated learning rate.\n \"\"\"\n self.lr_scheduler.step()\n lr = self.lr_scheduler.get_last_lr()[0] # only return the first group\n return lr\n\n def to(self, device: str, model: bool = True, optimizer: bool = True, grad: bool = True):\n \"\"\"\n Move FSDP model and/or optimizer to CPU or GPU with offload support.\n Note that this function executes irrespective of offload config. It serves as manual control\n \"\"\"\n super().to(device=device, model=model, optimizer=optimizer, grad=grad)\n\n if self.engine_config.forward_only:\n # force cpu_offload\n return\n\n device_name = get_device_name()\n\n assert device in (device_name, \"cpu\")\n if device == device_name:\n if model:\n load_fsdp_model_to_gpu(self.module)\n if optimizer and self.optimizer is not None:\n load_fsdp_optimizer(self.optimizer, device)\n gc.collect()\n elif device == \"cpu\":\n if model:\n offload_fsdp_model_to_cpu(self.module)\n if optimizer and self.optimizer is not None:\n offload_fsdp_optimizer(self.optimizer)\n else:\n raise ValueError(f\"Invalid device type: {device}\")\n\n def save_checkpoint(\n self,\n local_path: str,\n hdfs_path: Optional[str] = None,\n global_step: int = 0,\n max_ckpt_to_keep: Optional[int] = None,\n **kwargs,\n ) -> None:\n \"\"\"\n Save FSDP checkpoint, handling parameter offload as needed.\n \"\"\"\n origin_module_device = next(self.module.parameters()).device.type\n if self._is_offload_param or origin_module_device == \"cpu\":\n load_fsdp_model_to_gpu(self.module)\n\n self.checkpoint_manager.save_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, global_step=global_step, max_ckpt_to_keep=max_ckpt_to_keep\n )\n\n torch.distributed.barrier()\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.module)\n\n def load_checkpoint(\n self, local_path: str, hdfs_path: Optional[str] = None, del_local_after_load: int = True, **kwargs\n ) -> None:\n \"\"\"\n Load FSDP checkpoint, restoring parameters and optimizer state.\n \"\"\"\n import torch\n\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.module)\n\n self.checkpoint_manager.load_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, del_local_after_load=del_local_after_load\n )\n\n torch.distributed.barrier()\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.module)\n\n if self._is_offload_optimizer:\n offload_fsdp_optimizer(self.optimizer)\n\n def get_per_tensor_param(self, layered_summon=False, base_sync_done=False, **kwargs):\n log_gpu_memory_usage(\"Before load_fsdp_model_to_gpu\", logger=logger)\n\n load_fsdp_model_to_gpu(self.module)\n\n log_gpu_memory_usage(\"After load_fsdp_model_to_gpu\", logger=logger)\n\n peft_config = None\n merge_lora = self.model_config.lora.get(\"merge\", False)\n\n peft_model = getattr(self.module, \"_fsdp_wrapped_module\", self.module)\n if hasattr(peft_model, \"peft_config\"): # LoRA\n if not merge_lora:\n peft_config = peft_model.peft_config.get(\"default\", None)\n params = collect_lora_params(\n module=self.module,\n layered_summon=layered_summon,\n base_sync_done=base_sync_done,\n )\n if not base_sync_done:\n params = {replace_lora_wrapper(k, peft_config): v for k, v in params.items()}\n else: # merge lora\n with merged_lora_context(self.module, backup_adapters=True):\n params = self.module.state_dict()\n params = normalize_peft_param_name(params)\n else:\n params = self.module.state_dict()\n\n params = convert_weight_keys(params, getattr(self.module, \"_fsdp_wrapped_module\", self.module))\n\n log_gpu_memory_usage(\"Before offload_fsdp_model_to_cpu\", logger=logger)\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.module)\n log_gpu_memory_usage(\"After offload_fsdp_model_to_cpu\", logger=logger)\n\n if peft_config is not None and base_sync_done:\n per_tensor_param = params.items()\n else:\n device = get_device_id() # used when fsdp2 set cpu_offload_policy\n # TODO: cast fp32 to bf16 to reduce weight sync overhead, need more fine-grained control, e.g MoE gate\n per_tensor_param = (\n (\n name,\n param.to(device, non_blocking=True).full_tensor().to(torch.bfloat16, non_blocking=True)\n if isinstance(param, DTensor)\n else param,\n )\n for name, param in params.items()\n )\n # return per_tensor_param, peft_config\n # Convert peft_config to dict for vLLM compatibility (PEFTHelper.from_dict expects dict)\n peft_config_dict = peft_config.to_dict() if peft_config is not None else None\n return per_tensor_param, peft_config_dict\n\n def disable_adapter(self) -> ContextManager:\n return self.module.disable_adapter()\n\n\nclass EngineEvalModeCtx(BaseEngineCtx):\n def __init__(self, engine: FSDPEngine, **kwargs):\n super().__init__(engine=engine, mode=\"eval\", **kwargs)\n\n def __enter__(self):\n assert isinstance(self.engine, FSDPEngine)\n super().__enter__()\n self.prev_sp_group = get_ulysses_sequence_parallel_group()\n set_ulysses_sequence_parallel_group(self.engine.ulysses_parallel_group)\n self.engine.module.eval()\n\n def __exit__(self, exc_type, exc_value, traceback):\n assert isinstance(self.engine, FSDPEngine)\n set_ulysses_sequence_parallel_group(self.prev_sp_group)\n\n # https://pytorch.org/docs/stable/notes/fsdp.html#fsdp-notes\n # unshard the root FSDP module\n if self.engine.engine_config.fsdp_size > 1:\n if fsdp_version(self.engine.module) == 1:\n self.engine.module._handle.reshard(True)\n elif fsdp_version(self.engine.module) == 2:\n self.engine.module.reshard()\n\n super().__exit__(exc_type, exc_value, traceback)\n\n\nclass EngineTrainModeCtx(BaseEngineCtx):\n def __init__(self, engine: FSDPEngine, **kwargs):\n super().__init__(engine=engine, mode=\"train\", **kwargs)\n\n def __enter__(self):\n assert isinstance(self.engine, FSDPEngine)\n super().__enter__()\n self.prev_sp_group = get_ulysses_sequence_parallel_group()\n set_ulysses_sequence_parallel_group(self.engine.ulysses_parallel_group)\n self.engine.module.train()\n\n def __exit__(self, exc_type, exc_value, traceback):\n assert isinstance(self.engine, FSDPEngine)\n set_ulysses_sequence_parallel_group(self.prev_sp_group)\n self.engine.optimizer_zero_grad()\n super().__exit__(exc_type, exc_value, traceback)\n\n\n@EngineRegistry.register(model_type=\"language_model\", backend=[\"fsdp\", \"fsdp2\"], device=[\"cuda\", \"npu\"])\nclass FSDPEngineWithLMHead(FSDPEngine):\n def prepare_model_inputs(self, micro_batch: TensorDict):\n use_remove_padding = tu.get_non_tensor_data(data=micro_batch, key=\"use_remove_padding\", default=True)\n pad_mode = tu.get_non_tensor_data(data=micro_batch, key=\"pad_mode\", default=DatasetPadMode.NO_PADDING)\n use_fused_kernels = tu.get_non_tensor_data(data=micro_batch, key=\"use_fused_kernels\", default=False)\n temperature = micro_batch[\"temperature\"]\n temperature_item = temperature\n if use_fused_kernels:\n assert not isinstance(temperature, torch.Tensor), (\n \"use_fused_kernels does not support per sample temperature yet\"\n )\n assert pad_mode == DatasetPadMode.NO_PADDING, f\"pad_mode {pad_mode} not supported\"\n\n multi_modal_inputs = extract_multi_modal_inputs(micro_batch.get(\"multi_modal_inputs\", []))\n input_ids = micro_batch[\"input_ids\"]\n position_ids = micro_batch[\"position_ids\"]\n\n if not isinstance(temperature, torch.Tensor):\n temperature = torch.tensor([temperature] * input_ids.shape[0], device=input_ids.device)\n\n temperature = temperature.to(torch.float32)\n assert temperature.shape[0] == input_ids.shape[0]\n\n # args used to get outputs\n output_args = {}\n\n if use_remove_padding:\n # support per sample temperature\n # temperature (bsz,)\n # input_ids (bsz, j1)\n temperature_rmpad = verl_F.expand_as_nested(temperature, input_ids).values() # (total_nnz,)\n temperature_rmpad = temperature_rmpad.unsqueeze(0) # (1, total_nnz)\n\n if pad_mode == DatasetPadMode.NO_PADDING:\n input_ids_rmpad = input_ids.values().unsqueeze(0) # (1, total_nnz)\n if position_ids.dim() == 3:\n position_ids_rmpad = position_ids.values().unsqueeze(1) # (4, 1, total_nnz)\n else:\n position_ids_rmpad = position_ids.values().unsqueeze(0) # (1, total_nnz)\n else:\n raise NotImplementedError(f\"pad_mode {pad_mode} not implemented\")\n\n # for compute the log_prob\n input_ids_rmpad_rolled = torch.roll(input_ids_rmpad, shifts=-1, dims=1) # (1, total_nnz)\n\n # pad and slice the inputs if sp > 1\n if self.use_ulysses_sp:\n is_vlm_model = hasattr(getattr(self.module, \"module\", self.module).config, \"vision_config\")\n if is_vlm_model:\n # vlm model's inputs will be sliced after embedding\n input_ids_rmpad, position_ids_rmpad, pad_size = ulysses_pad(\n input_ids_rmpad,\n position_ids_rmpad=position_ids_rmpad,\n sp_size=self.ulysses_sequence_parallel_size,\n )\n else:\n input_ids_rmpad, position_ids_rmpad, pad_size = ulysses_pad_and_slice_inputs(\n input_ids_rmpad,\n position_ids_rmpad=position_ids_rmpad,\n sp_size=self.ulysses_sequence_parallel_size,\n skip_position_ids_rmpad=True if self.__class__.__name__ == \"VeOmniEngineWithLMHead\" else False,\n )\n input_ids_rmpad_rolled, _, _ = ulysses_pad_and_slice_inputs(\n input_ids_rmpad_rolled,\n position_ids_rmpad=None,\n sp_size=self.ulysses_sequence_parallel_size,\n )\n\n temperature_rmpad, _, _ = ulysses_pad_and_slice_inputs(\n temperature_rmpad, position_ids_rmpad=None, sp_size=self.ulysses_sequence_parallel_size, pad_value=1\n )\n\n output_args[\"pad_size\"] = pad_size\n\n input_ids_rmpad_rolled = input_ids_rmpad_rolled.squeeze(0) # ((total_nnz / sp) + pad)\n temperature_rmpad = temperature_rmpad.squeeze(0)\n output_args[\"input_ids_rmpad_rolled\"] = input_ids_rmpad_rolled\n output_args[\"temperature_rmpad\"] = temperature_rmpad\n\n # only pass input_ids and position_ids to enable flash_attn_varlen\n\n model_inputs = {\n \"input_ids\": input_ids_rmpad,\n \"attention_mask\": None,\n \"position_ids\": position_ids_rmpad,\n }\n\n else:\n if pad_mode == DatasetPadMode.NO_PADDING:\n input_ids = micro_batch[\"input_ids\"]\n position_ids = micro_batch[\"position_ids\"]\n loss_mask = micro_batch[\"loss_mask\"]\n\n pad_token_id = tu.get_non_tensor_data(data=micro_batch, key=\"pad_token_id\", default=0)\n batch_size = micro_batch.batch_size[0]\n seq_len_effective = input_ids.offsets().diff()\n max_seq_len = max(seq_len_effective)\n\n input_ids_rmpad_rolled = torch.roll(input_ids.values(), shifts=-1, dims=0)\n output_args[\"input_ids_rmpad_rolled\"] = input_ids_rmpad_rolled\n # we store the per sample temperature\n output_args[\"temperature\"] = temperature\n\n input_ids = torch.nested.to_padded_tensor(\n input_ids, padding=pad_token_id, output_size=(batch_size, max_seq_len)\n )\n\n if position_ids.dim() == 3:\n position_ids = torch.nested.to_padded_tensor(\n position_ids, padding=0, output_size=(batch_size, 4, max_seq_len)\n ).transpose(0, 1) # (4, batch_size, max_seq_len)\n else:\n position_ids = torch.nested.to_padded_tensor(\n position_ids, padding=0, output_size=(batch_size, max_seq_len)\n )\n\n attention_mask_list = [torch.ones_like(t, dtype=torch.int32) for t in loss_mask]\n attention_mask = torch.nested.as_nested_tensor(attention_mask_list, layout=torch.jagged)\n attention_mask = torch.nested.to_padded_tensor(\n attention_mask, padding=0, output_size=(batch_size, max_seq_len)\n )\n\n model_inputs = {\n \"input_ids\": input_ids,\n \"attention_mask\": attention_mask,\n \"position_ids\": position_ids,\n }\n\n else:\n raise NotImplementedError(f\"pad_mode {pad_mode} not implemented\")\n\n extra_args = {}\n if use_fused_kernels:\n extra_args[\"temperature\"] = temperature_item\n extra_args[\"return_dict\"] = True\n\n model_inputs.update(multi_modal_inputs)\n model_inputs.update(extra_args)\n\n return model_inputs, output_args\n\n def prepare_model_outputs(self, output, output_args, micro_batch: TensorDict):\n use_remove_padding = tu.get_non_tensor_data(data=micro_batch, key=\"use_remove_padding\", default=True)\n pad_mode = tu.get_non_tensor_data(data=micro_batch, key=\"pad_mode\", default=DatasetPadMode.NO_PADDING)\n use_fused_kernels = tu.get_non_tensor_data(data=micro_batch, key=\"use_fused_kernels\", default=False)\n calculate_entropy = tu.get_non_tensor_data(data=micro_batch, key=\"calculate_entropy\", default=False)\n\n model_output = {}\n\n input_ids = micro_batch[\"input_ids\"]\n\n if use_remove_padding:\n input_ids_rmpad_rolled = output_args[\"input_ids_rmpad_rolled\"]\n temperature_rmpad = output_args[\"temperature_rmpad\"]\n\n if use_fused_kernels:\n # temperature is singleton\n log_probs = output.log_probs.squeeze(0) # (total_nnz,)\n entropy_rmpad = output.entropy.squeeze(0) # (total_nnz,)\n else:\n logits_rmpad = output.logits.squeeze(0) # (total_nnz, vocab_size)\n logits_rmpad.div_(temperature_rmpad.clamp(min=1e-8).unsqueeze(-1).to(logits_rmpad.dtype))\n\n # if use_sp: ((total_nnz / sp) + pad) ; if not use_sp: (batch, seqlen)\n inplace_backward = True\n if calculate_entropy:\n inplace_backward = False\n log_probs = logprobs_from_logits(\n logits=logits_rmpad,\n labels=input_ids_rmpad_rolled,\n inplace_backward=inplace_backward,\n )\n\n # compute entropy\n if calculate_entropy:\n if not self.engine_config.entropy_checkpointing:\n entropy_rmpad = self.compute_entropy_from_logits(logits_rmpad) # ((total_nnz / sp) + pad)\n else:\n entropy_rmpad = torch.utils.checkpoint.checkpoint(\n self.compute_entropy_from_logits, logits_rmpad\n )\n\n # gather log_prob if sp > 1\n if self.use_ulysses_sp:\n pad_size = output_args[\"pad_size\"]\n\n # gather and unpad for the ulysses sp\n log_probs = gather_outputs_and_unpad(\n log_probs,\n gather_dim=0,\n unpad_dim=0,\n padding_size=pad_size,\n )\n if calculate_entropy:\n entropy_rmpad = gather_outputs_and_unpad(\n entropy_rmpad,\n gather_dim=0,\n unpad_dim=0,\n padding_size=pad_size,\n )\n\n if pad_mode == DatasetPadMode.NO_PADDING:\n cu_seqlens = input_ids.offsets()\n # (bsz, j1), for each sample, is the length of each sample: [real_prompt length + real_response length]\n log_probs = torch.nested.nested_tensor_from_jagged(log_probs, cu_seqlens)\n if calculate_entropy:\n entropy = torch.nested.nested_tensor_from_jagged(entropy_rmpad, cu_seqlens)\n else:\n raise NotImplementedError(f\"pad_mode {pad_mode} not implemented\")\n\n else: # not using rmpad and no ulysses sp\n response_length = tu.get_non_tensor_data(data=micro_batch, key=\"max_response_length\", default=1024)\n if use_fused_kernels:\n log_probs = output.log_probs[:, -response_length - 1 : -1]\n entropy = output.entropy[:, -response_length - 1 : -1] # (bsz, response_length)\n\n else:\n logits = output.logits # (bsz, response_length, vocab_size)\n temperature = output_args[\"temperature\"] # (bsz,)\n temperature = temperature.unsqueeze(-1).unsqueeze(-1)\n logits.div_(temperature.clamp(min=1e-8).to(logits.dtype))\n\n if calculate_entropy:\n if not self.engine_config.entropy_checkpointing:\n entropy = verl_F.entropy_from_logits(logits)\n else:\n entropy = torch.utils.checkpoint.checkpoint(verl_F.entropy_from_logits, logits)\n\n if pad_mode == DatasetPadMode.NO_PADDING:\n cu_seqlens = input_ids.offsets()\n seq_lengths = cu_seqlens.diff()\n starts = torch.zeros_like(seq_lengths, dtype=torch.int64)\n logits = torch.nested.narrow(logits, 1, starts, seq_lengths, layout=torch.jagged)\n logits_rmpad = torch.cat([t for t in logits.unbind()])\n input_ids_rmpad_rolled = output_args[\"input_ids_rmpad_rolled\"]\n log_probs = logprobs_from_logits(logits=logits_rmpad, labels=input_ids_rmpad_rolled)\n # (bsz, j1), for each sample, length of each sample: [real_prompt_length + real_response_length]\n log_probs = torch.nested.nested_tensor_from_jagged(log_probs, cu_seqlens)\n if calculate_entropy:\n entropy = torch.nested.narrow(entropy, 1, starts, seq_lengths, layout=torch.jagged)\n entropy_rmpad = torch.cat([t for t in entropy.unbind()])\n entropy = torch.nested.nested_tensor_from_jagged(entropy_rmpad, cu_seqlens)\n else:\n raise NotImplementedError(f\"pad_mode {pad_mode} not implemented\")\n\n model_output[\"log_probs\"] = log_probs\n if calculate_entropy:\n model_output[\"entropy\"] = entropy\n\n return model_output\n\n def forward_step(self, micro_batch: TensorDict, loss_function, forward_only):\n device_name = get_device_name()\n # actually, we should avoid assigning like this...\n micro_batch = micro_batch.to(get_device_id())\n model_inputs, output_args = self.prepare_model_inputs(micro_batch=micro_batch)\n\n with torch.autocast(device_type=device_name, dtype=torch.bfloat16):\n raw_output = self.module(\n **model_inputs,\n use_cache=False,\n ) # prevent model thinks we are generating\n\n model_output = self.prepare_model_outputs(\n output=raw_output, output_args=output_args, micro_batch=micro_batch\n )\n\n if loss_function is not None:\n loss, metrics = loss_function(\n model_output=model_output, data=micro_batch, dp_group=self.get_data_parallel_group()\n )\n else:\n assert forward_only, \"forward_only must be True when loss_function is None\"\n loss = torch.tensor(1.0, device=device_name)\n metrics = {}\n\n output = {\n \"model_output\": model_output,\n \"loss\": loss.detach().item(),\n \"metrics\": metrics,\n }\n\n return loss, output\n\n\n@EngineRegistry.register(model_type=\"value_model\", backend=[\"fsdp\", \"fsdp2\"], device=[\"cuda\", \"npu\"])\nclass FSDPEngineWithValueHead(FSDPEngineWithLMHead):\n \"\"\"\n The only difference between critic and actor is how the raw model output is processed\n \"\"\"\n\n def prepare_model_outputs(self, output, output_args, micro_batch: TensorDict):\n use_remove_padding = tu.get_non_tensor_data(data=micro_batch, key=\"use_remove_padding\", default=True)\n pad_mode = tu.get_non_tensor_data(data=micro_batch, key=\"pad_mode\", default=DatasetPadMode.NO_PADDING)\n\n input_ids = micro_batch[\"input_ids\"]\n if use_remove_padding:\n if hasattr(self.module, \"v_head\"):\n # For trl.AutoModelForCausalLMWithValueHead\n values_rmpad = output[2].squeeze(0).unsqueeze(-1)\n else:\n values_rmpad = output.logits\n values_rmpad = values_rmpad.squeeze(0) # (total_nnz, 1)\n # critic model arch is like Qwen3ForTokenClassfication and num_labels=1\n # so we squeeze the last dimension here to get the value for each token\n values_rmpad = values_rmpad.squeeze(-1)\n\n # gather output if sp > 1\n if self.use_ulysses_sp:\n pad_size = output_args[\"pad_size\"]\n values_rmpad = gather_outputs_and_unpad(values_rmpad, gather_dim=0, unpad_dim=0, padding_size=pad_size)\n\n if pad_mode == DatasetPadMode.NO_PADDING:\n cu_seqlens = input_ids.offsets()\n # (bsz, j1), for each sample, is the length of each sample: [real_prompt length + real_response length]\n values = torch.nested.nested_tensor_from_jagged(values_rmpad, cu_seqlens)\n else:\n raise NotImplementedError(f\"pad_mode {pad_mode} not implemented\")\n\n else:\n if hasattr(self.module, \"v_head\"):\n # For trl.AutoModelForCausalLMWithValueHead\n values = output[2]\n else:\n values = output.logits\n\n if pad_mode == DatasetPadMode.NO_PADDING:\n cu_seqlens = input_ids.offsets()\n seq_lengths = cu_seqlens.diff()\n starts = torch.zeros_like(seq_lengths, dtype=torch.int64)\n values = torch.nested.narrow(values, 1, starts, seq_lengths, layout=torch.jagged)\n values_rmpad = torch.cat([t for t in values.unbind()])\n # (bsz, j1), for each sample, length of each sample: [real_prompt_length + real_response_length]\n values = torch.nested.nested_tensor_from_jagged(values_rmpad, cu_seqlens)\n else:\n raise NotImplementedError(f\"pad_mode {pad_mode} not implemented\")\n\n return {\"values\": values}\n"}142{"file_name": "verl__workers__engine__fsdp__utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport logging\nimport os\n\nimport torch\nfrom torch.distributed.device_mesh import init_device_mesh\n\nfrom verl.utils.device import get_device_name, is_npu_available\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef apply_npu_fsdp_patches():\n \"\"\"Apply NPU patches for FSDP backend if NPU is available.\"\"\"\n if is_npu_available:\n try:\n import verl.models.transformers.npu_patch # noqa\n\n if torch.distributed.is_initialized() and torch.distributed.get_rank() == 0:\n logger.info(\"Applied NPU patches for FSDP backend\")\n except Exception as e:\n logger.warning(f\"Failed to apply NPU patches: {e}\")\n\n\ndef create_device_mesh(world_size, fsdp_size):\n \"\"\"\n Create a device mesh for distributed training based on the world size and FSDP size.\n\n Args:\n world_size (int): Total number of processes in the distributed training setup.\n fsdp_size (int): Size of the Fully Sharded Data Parallel (FSDP) group.\n\n Returns:\n torch.distributed.device_mesh.DeviceMesh: The initialized device mesh.\n \"\"\"\n device_name = get_device_name()\n if fsdp_size < 0 or fsdp_size >= world_size:\n device_mesh = init_device_mesh(device_name, mesh_shape=(world_size,), mesh_dim_names=[\"fsdp\"])\n else:\n device_mesh = init_device_mesh(\n device_name, mesh_shape=(world_size // fsdp_size, fsdp_size), mesh_dim_names=[\"ddp\", \"fsdp\"]\n )\n return device_mesh\n\n\ndef get_sharding_strategy(device_mesh):\n \"\"\"\n Determine the appropriate sharding strategy based on the number of dimensions of the device mesh.\n\n Args:\n device_mesh (torch.distributed.device_mesh.DeviceMesh): The device mesh used for distributed training.\n\n Returns:\n torch.distributed.fsdp.ShardingStrategy: The sharding strategy to be used with FSDP.\n\n Raises:\n NotImplementedError: If the number of dimensions of the device mesh is neither 1 nor 2.\n \"\"\"\n from torch.distributed.fsdp import ShardingStrategy\n\n if device_mesh.ndim == 1:\n sharding_strategy = ShardingStrategy.FULL_SHARD\n elif device_mesh.ndim == 2:\n sharding_strategy = ShardingStrategy.HYBRID_SHARD\n else:\n raise NotImplementedError(f\"Get device mesh ndim={device_mesh.ndim}, but only support 1 or 2\")\n return sharding_strategy\n"}143{"file_name": "verl__workers__engine__megatron__transformer_impl.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport logging\nimport os\nfrom functools import partial\nfrom typing import Any, Callable, ContextManager, Iterator, Optional\n\nimport torch\nimport torch.distributed\nfrom megatron.core import parallel_state as mpu\nfrom megatron.core.pipeline_parallel import get_forward_backward_func\nfrom omegaconf import OmegaConf\nfrom tensordict import TensorDict\n\nimport verl.utils.torch_functional as verl_F\nfrom verl.models.mcore import get_mcore_forward_fused_no_padding_fn, get_mcore_weight_converter\nfrom verl.trainer.config import CheckpointConfig\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.checkpoint.megatron_checkpoint_manager import MegatronCheckpointManager\nfrom verl.utils.dataset.dataset_utils import DatasetPadMode\nfrom verl.utils.debug import log_gpu_memory_usage\nfrom verl.utils.device import get_device_id, get_device_name\nfrom verl.utils.megatron.pipeline_parallel import make_batch_generator\nfrom verl.utils.megatron.router_replay_patch import RouterReplay, RouterReplayAction, apply_router_replay_patch\nfrom verl.utils.megatron.router_replay_utils import (\n RouterReplayHelper,\n set_router_replay_data,\n)\nfrom verl.utils.megatron.tensor_parallel import vocab_parallel_entropy, vocab_parallel_log_probs_from_logits\nfrom verl.utils.megatron_peft_utils import add_base_layer_suffix, build_peft_config_for_vllm\nfrom verl.utils.megatron_utils import (\n check_mtp_config,\n get_megatron_module_device,\n get_megatron_mtp_loss,\n load_megatron_model_to_gpu,\n load_megatron_optimizer,\n offload_megatron_model_to_cpu,\n offload_megatron_optimizer,\n patch_engine_mtp,\n register_megatron_training_hooks,\n unwrap_model,\n)\nfrom verl.utils.model import extract_multi_modal_inputs, load_mcore_dist_weights\nfrom verl.workers.config import HFModelConfig, McoreEngineConfig, McoreOptimizerConfig\n\nfrom ..base import BaseEngine, BaseEngineCtx, EngineRegistry\nfrom ..utils import postprocess_batch_func, prepare_micro_batches\nfrom .utils import set_random_seed\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\nclass MegatronEngine(BaseEngine):\n def __init__(\n self,\n model_config: HFModelConfig,\n engine_config: McoreEngineConfig,\n optimizer_config: McoreOptimizerConfig,\n checkpoint_config: CheckpointConfig,\n ):\n super().__init__()\n\n self.model_config = model_config\n self.engine_config = engine_config\n self.optimizer_config = optimizer_config\n self.checkpoint_config = checkpoint_config\n assert self.engine_config.use_mbridge, \"use_mbridge must be True\"\n self._init_device_mesh()\n\n set_random_seed(seed=self.engine_config.seed)\n\n self._is_offload_param = self.engine_config.param_offload\n self._is_offload_grad = self.engine_config.grad_offload\n self._is_offload_optimizer = self.engine_config.optimizer_offload\n\n self.mode = None\n\n self.layer_name_mapping = {\n \"qkv_layer_name\": \"self_attention.linear_qkv.\",\n \"gate_proj_layer_name\": \"linear_fc1.\",\n }\n self.weight_converter = None\n\n # Router replay configuration for MoE models\n self.enable_routing_replay = self.engine_config.router_replay.mode != \"disabled\"\n logger.info(f\"enable_routing_replay in MegatronEngine: {self.enable_routing_replay}\")\n if self.enable_routing_replay:\n apply_router_replay_patch()\n\n def _init_device_mesh(self):\n # TODO: set different parallelism for actor, critic, ref\n if mpu.is_initialized():\n return\n\n mpu.initialize_model_parallel(\n tensor_model_parallel_size=self.engine_config.tensor_model_parallel_size,\n pipeline_model_parallel_size=self.engine_config.pipeline_model_parallel_size,\n virtual_pipeline_model_parallel_size=self.engine_config.virtual_pipeline_model_parallel_size,\n use_sharp=False,\n context_parallel_size=self.engine_config.context_parallel_size,\n expert_model_parallel_size=self.engine_config.expert_model_parallel_size,\n expert_tensor_parallel_size=self.engine_config.expert_tensor_parallel_size,\n nccl_communicator_config_path=None,\n )\n\n def _build_tf_config(self):\n from verl.utils.megatron_utils import mapping_string_to_attn_backend\n from verl.utils.torch_dtypes import PrecisionType\n\n check_mtp_config(self.model_config, self.engine_config)\n\n self.param_dtype = PrecisionType.to_dtype(self.engine_config.dtype)\n self.dtype = PrecisionType.to_dtype(self.param_dtype)\n\n override_transformer_config = mapping_string_to_attn_backend({**self.engine_config.override_transformer_config})\n if self.enable_routing_replay:\n override_transformer_config[\"enable_routing_replay\"] = True\n\n self.provider = None\n self.vanilla_bridge = self.engine_config.vanilla_mbridge\n\n if self.vanilla_bridge:\n from verl.models.mcore.mbridge import AutoBridge\n\n bridge = AutoBridge.from_config(self.model_config.hf_config, dtype=self.param_dtype)\n bridge.set_extra_args(**override_transformer_config)\n tf_config = bridge.config\n tf_config.fp16 = self.param_dtype == torch.float16\n tf_config.bf16 = self.param_dtype == torch.bfloat16\n else:\n from verl.models.mcore.bridge import AutoBridge\n\n # Use Megatron-Bridge to convert HF config to Megatron config\n bridge = AutoBridge.from_hf_pretrained(\n self.model_config.local_path, trust_remote_code=self.model_config.trust_remote_code\n )\n # Get Megatron provider and configure it\n provider = bridge.to_megatron_provider(load_weights=False)\n\n # In case of invalid overrides, we need to make sure some critical params are set correctly\n provider.params_dtype = self.param_dtype\n\n # Ensure dtype settings propagate to Megatron-Bridge/TE\n provider.fp16 = self.param_dtype == torch.float16\n provider.bf16 = self.param_dtype == torch.bfloat16\n\n # Pass distributed info\n provider.tensor_model_parallel_size = self.engine_config.tensor_model_parallel_size\n provider.pipeline_model_parallel_size = self.engine_config.pipeline_model_parallel_size\n provider.expert_model_parallel_size = self.engine_config.expert_model_parallel_size\n provider.expert_tensor_parallel_size = self.engine_config.expert_tensor_parallel_size\n provider.virtual_pipeline_model_parallel_size = self.engine_config.virtual_pipeline_model_parallel_size\n provider.context_parallel_size = self.engine_config.context_parallel_size\n provider.sequence_parallel = self.engine_config.sequence_parallel\n\n # Match verl implementation (need variable_seq_lengths)\n from megatron.core.transformer.enums import AttnBackend\n\n provider.attention_backend = AttnBackend.flash\n provider.variable_seq_lengths = True\n provider.moe_token_dispatcher_type = \"alltoall\"\n provider.moe_router_load_balancing_type = \"none\"\n\n # Apply transformer config overrides\n for key, value in override_transformer_config.items():\n setattr(provider, key, value)\n\n provider.finalize()\n self.provider = provider\n tf_config = None # Will be set after model creation\n self.bridge = bridge\n\n if not self.bridge:\n self.weight_converter = get_mcore_weight_converter(self.model_config.hf_config, self.dtype)\n\n if torch.distributed.get_rank() == 0:\n if tf_config is not None:\n print(f\"TF config: {tf_config}\")\n self.tf_config = tf_config\n\n from verl.workers.config.megatron_peft import get_peft_cls\n\n self.peft_cls = get_peft_cls(\n model_config=self.model_config, bridge=self.bridge, provider=self.provider, dtype=self.param_dtype\n )\n\n def _build_megatron_module(self):\n from verl.utils.megatron_utils import McoreModuleWrapperConfig, make_megatron_module\n from verl.utils.model import print_model_size\n\n # TODO: add more cases\n is_value_model = (\n \"ForTokenClassification\" in self.model_config.architectures[0]\n or \"ForSequenceClassification\" in self.model_config.architectures[0]\n )\n\n self.is_value_model = is_value_model\n\n if self.engine_config.forward_only:\n wrap_with_ddp = False\n else:\n wrap_with_ddp = True\n\n wrap_config = McoreModuleWrapperConfig(\n is_value_model=is_value_model, # actor is not value model\n share_embeddings_and_output_weights=self.model_config.share_embeddings_and_output_weights,\n wrap_with_ddp=wrap_with_ddp,\n use_distributed_optimizer=self.engine_config.use_distributed_optimizer,\n )\n module, updated_tf_config = make_megatron_module(\n wrap_config=wrap_config,\n tf_config=self.tf_config,\n hf_config=self.model_config.hf_config,\n bridge=self.bridge,\n provider=self.provider,\n override_model_config=self.engine_config.override_mcore_model_config,\n override_ddp_config=self.engine_config.override_ddp_config,\n peft_cls=self.peft_cls,\n peft_config=self.model_config.get(\"lora\", None),\n )\n self.tf_config = updated_tf_config\n print(f\"module: {len(module)}\")\n\n if self.engine_config.use_dist_checkpointing:\n load_mcore_dist_weights(module, self.engine_config.dist_checkpointing_path, is_value_model=is_value_model)\n else:\n if self.vanilla_bridge:\n self.bridge.load_weights(module, self.model_config.local_path)\n else:\n allowed_mismatched_params = []\n if self.is_value_model:\n allowed_mismatched_params = [\"output_layer.weight\"]\n self.bridge.load_hf_weights(\n module, self.model_config.local_path, allowed_mismatched_params=allowed_mismatched_params\n )\n\n if torch.distributed.get_rank() == 0:\n print_model_size(module[0])\n\n if self.enable_routing_replay:\n print(f\"routing replay layers: {len(RouterReplay.router_instances)}\")\n\n return module\n\n def _maybe_enable_fused_kernels(self):\n if not self.engine_config.use_fused_kernels:\n return\n\n if self.is_value_model or self.model_config.mtp.enable:\n logger.warning_once(\n \"Fused kernels are not supported for value models or when MTP is enabled in Megatron engine; disabling.\"\n )\n self.engine_config.use_fused_kernels = False\n return\n\n from verl.models.mcore.model_forward_fused import patch_fused_forward\n\n for model in self.module:\n patch_fused_forward(model)\n\n def _build_optimizer(self):\n from verl.utils.megatron.optimizer import get_megatron_optimizer, init_megatron_optim_config\n\n optim_config_megatron = init_megatron_optim_config(\n self.optimizer_config,\n use_distributed_optimizer=self.engine_config.use_distributed_optimizer,\n fp16=self.param_dtype == torch.float16,\n )\n optimizer = get_megatron_optimizer(model=self.module, config=optim_config_megatron)\n register_megatron_training_hooks(self.module, optimizer)\n return optimizer\n\n def _build_lr_scheduler(self):\n from verl.utils.megatron.optimizer import get_megatron_optimizer_param_scheduler\n\n optimizer_scheduler = get_megatron_optimizer_param_scheduler(\n optimizer=self.optimizer, config=self.optimizer_config\n )\n return optimizer_scheduler\n\n @property\n def is_param_offload_enabled(self) -> bool:\n return self._is_offload_param\n\n @property\n def is_optimizer_offload_enabled(self) -> bool:\n return self._is_offload_optimizer\n\n def is_mp_src_rank_with_outputs(self):\n return (\n mpu.get_tensor_model_parallel_rank() == 0\n and mpu.get_pipeline_model_parallel_rank() == mpu.get_pipeline_model_parallel_world_size() - 1\n and mpu.get_context_parallel_rank() == 0\n )\n\n def initialize(self):\n self._build_tf_config()\n\n self.module = self._build_megatron_module()\n\n self._maybe_enable_fused_kernels()\n\n if self.model_config.mtp.enable:\n patch_engine_mtp(self.module, self.model_config)\n\n # For forward_only, we don't need optimizer, lr_scheduler, checkpoint_mananager\n if self.engine_config.forward_only:\n self.optimizer = None\n self.lr_scheduler = None\n return\n\n self.optimizer = self._build_optimizer()\n self.lr_scheduler = self._build_lr_scheduler()\n\n full_reshardable = self.engine_config.dist_ckpt_optim_fully_reshardable\n mem_eff = self.engine_config.distrib_optim_fully_reshardable_mem_efficient\n\n tmp_config = OmegaConf.create(\n {\n \"model\": {\"path\": self.model_config.local_path},\n \"megatron\": {\n \"dist_ckpt_optim_fully_reshardable\": full_reshardable,\n \"distrib_optim_fully_reshardable_mem_efficient\": mem_eff,\n },\n }\n )\n\n role = \"actor\" if not self.is_value_model else \"critic\"\n\n self.checkpoint_mananager = MegatronCheckpointManager(\n config=tmp_config,\n checkpoint_config=self.checkpoint_config,\n model_config=self.model_config.hf_config,\n transformer_config=self.tf_config,\n role=role,\n model=self.module,\n arch=self.model_config.architectures[0],\n hf_config=self.model_config.hf_config,\n param_dtype=self.param_dtype,\n share_embeddings_and_output_weights=self.model_config.share_embeddings_and_output_weights,\n processing_class=self.model_config.get_processor(),\n optimizer=self.optimizer,\n optimizer_scheduler=self.lr_scheduler,\n use_distributed_optimizer=self.engine_config.use_distributed_optimizer,\n use_checkpoint_opt_param_scheduler=self.optimizer_config.use_checkpoint_opt_param_scheduler,\n bridge=self.bridge,\n provider=self.provider,\n peft_cls=self.peft_cls,\n use_dist_checkpointing=self.engine_config.use_dist_checkpointing,\n )\n\n self.to(\n device=\"cpu\",\n model=self._is_offload_param,\n optimizer=self._is_offload_optimizer,\n grad=self._is_offload_param,\n )\n\n log_gpu_memory_usage(\"After offload model/optimizer/grad during init\", logger=logger)\n\n def train_mode(self, **kwargs):\n \"\"\"\n Context manager entry for switching the engine and model into training mode.\n\n Usage:\n with engine.train_mode():\n # runs in training mode\n \"\"\"\n return EngineTrainModeCtx(self, **kwargs)\n\n def eval_mode(self, **kwargs):\n \"\"\"\n Context manager entry for switching the engine and model into evaluation mode.\n\n Usage:\n with engine.eval_mode():\n # runs in evaluation mode\n \"\"\"\n return EngineEvalModeCtx(self, **kwargs)\n\n def optimizer_zero_grad(self):\n \"\"\"\n Zero out gradients of all parameters before starting a new backward pass.\n \"\"\"\n self.optimizer.zero_grad()\n # use use_contiguous_buffers_in_local_ddp and no overlap_dp_param_comm\n for chunk in self.module:\n # if use distributed optimizer, zero grad buffer will be handled by optimizer\n chunk.zero_grad_buffer()\n\n def optimizer_step(self):\n \"\"\"\n Perform an optimization step to update model parameters based on accumulated gradients.\n\n Returns:\n grad_norm (float): The norm of the gradients before clipping or update.\n \"\"\"\n update_successful, grad_norm, num_zeros_in_grad = self.optimizer.step()\n\n if update_successful:\n # allgather already execute in optimizer.step in new megatron\n pass\n else:\n raise NotImplementedError(\"Megatron optimizer step failed. This should not happen\")\n\n return grad_norm\n\n def lr_scheduler_step(self):\n \"\"\"\n Advance the learning rate scheduler by one step.\n\n Returns:\n current_lr (float or list[float]): Updated learning rate(s).\n \"\"\"\n from verl.utils.megatron.optimizer import get_megatron_last_lr\n\n self.lr_scheduler.step(1)\n return get_megatron_last_lr(self.optimizer)\n\n def to(self, device: str, model: bool = True, optimizer: bool = True, grad: bool = True):\n \"\"\"\n Move model parameters, optimizer states, or both to the specified device.\n Note that this function executes irrespective of offload config. It serves as manual control\n\n Args:\n device: Target device identifier.\n model: If True, move the model.\n optimizer: If True, move the optimizer states.\n \"\"\"\n super().to(device=device, model=model, optimizer=optimizer, grad=grad)\n\n device_name = get_device_name()\n\n assert device in (device_name, \"cpu\")\n if device == device_name:\n if model:\n load_megatron_model_to_gpu(self.module, load_grad=grad)\n if optimizer and self.optimizer is not None:\n load_megatron_optimizer(self.optimizer)\n elif device == \"cpu\":\n if model:\n offload_megatron_model_to_cpu(self.module)\n if optimizer and self.optimizer is not None:\n offload_megatron_optimizer(self.optimizer)\n else:\n raise ValueError(f\"Invalid device type: {device}\")\n\n def get_data_parallel_rank(self):\n return mpu.get_data_parallel_rank()\n\n def get_data_parallel_size(self):\n return mpu.get_data_parallel_world_size()\n\n def get_data_parallel_group(self):\n return mpu.get_data_parallel_group()\n\n def save_checkpoint(\n self,\n local_path: str,\n hdfs_path: Optional[str] = None,\n global_step: int = 0,\n max_ckpt_to_keep: Optional[int] = None,\n **kwargs,\n ) -> None:\n \"\"\"\n Save model, optimizer, and scheduler states to a checkpoint.\n\n Args:\n local_path: Local filesystem path to save checkpoint.\n hdfs_path: Optional HDFS path to copy checkpoint.\n global_step: Integer training step number for naming.\n max_ckpt_to_keep: Maximum number of recent checkpoints to retain.\n \"\"\"\n origin_module_device = get_megatron_module_device(self.module)\n if self._is_offload_param or origin_module_device == \"cpu\":\n load_megatron_model_to_gpu(self.module, load_grad=True)\n self.checkpoint_mananager.save_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, global_step=global_step, max_ckpt_to_keep=max_ckpt_to_keep\n )\n torch.distributed.barrier()\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.module)\n\n def load_checkpoint(\n self, local_path: str, hdfs_path: Optional[str] = None, del_local_after_load: bool = True, **kwargs\n ) -> None:\n \"\"\"\n Load model, optimizer, and scheduler states from a checkpoint.\n\n Args:\n local_path: Local filesystem path of the checkpoint.\n hdfs_path: Optional HDFS path where checkpoint is stored.\n del_local_after_load: Whether to delete local copy after loading.\n \"\"\"\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.module)\n self.checkpoint_mananager.load_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, del_local_after_load=del_local_after_load\n )\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.module)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.optimizer)\n\n def forward_backward_batch(self, data: TensorDict, loss_function: Callable, forward_only=False) -> Any:\n tu.assign_non_tensor(data, sp_size=self.engine_config.context_parallel_size)\n\n # compute num_tokens in global batch for loss normalization\n batch_num_tokens = data[\"loss_mask\"].sum().to(get_device_id())\n torch.distributed.all_reduce(\n batch_num_tokens, op=torch.distributed.ReduceOp.SUM, group=self.get_data_parallel_group()\n )\n tu.assign_non_tensor(data, batch_num_tokens=batch_num_tokens.item())\n tu.assign_non_tensor(data, dp_size=self.get_data_parallel_size())\n\n vpp_size = mpu.get_virtual_pipeline_model_parallel_world_size()\n if vpp_size is not None and vpp_size > 1:\n num_batches_divided_by = self.tf_config.microbatch_group_size_per_vp_stage\n else:\n num_batches_divided_by = None\n\n micro_batches, indices = prepare_micro_batches(\n data=data,\n dp_group=self.get_data_parallel_group(),\n num_batches_divided_by=num_batches_divided_by,\n same_micro_num_in_dp=True,\n min_num_micro_batch=None,\n )\n\n if num_batches_divided_by is not None:\n assert len(micro_batches) % num_batches_divided_by == 0, (\n f\"micro_batches {micro_batches} must be divisible by num_batches_divided_by \"\n f\"{num_batches_divided_by} for megatron backend\"\n )\n\n # compute input shapes for pp stages\n n_micro_batch = len(micro_batches)\n\n for micro_batch in micro_batches:\n tu.assign_non_tensor(micro_batch, num_micro_batch=n_micro_batch)\n\n forward_backward_func = get_forward_backward_func()\n\n postprocess_micro_batch_func = partial(\n self.postprocess_micro_batch_func,\n forward_only=forward_only,\n loss_function=loss_function,\n )\n\n tu.assign_non_tensor(data, num_micro_batch=n_micro_batch)\n\n forward_step = partial(self.forward_step, postprocess_micro_batch_func=postprocess_micro_batch_func)\n\n enable_routing_replay = tu.get_non_tensor_data(data, key=\"enable_routing_replay\", default=False)\n\n if enable_routing_replay:\n RouterReplay.set_global_router_replay_action(RouterReplayAction.REPLAY_FORWARD)\n\n # batch should be a list of batches inside micro-batches\n batch_generator = make_batch_generator(micro_batches, vpp_size=len(self.module))\n\n # TODO: we may use the new schedule instead\n # for flash-attn: (seq_len, batch_size, hidden_size) = (mbs*seq_len, 1, hidden_size)\n losses_reduced = forward_backward_func(\n forward_step_func=forward_step,\n data_iterator=batch_generator,\n model=self.module,\n num_microbatches=n_micro_batch,\n seq_length=1, # the communication shape is obtained via p2p comm\n micro_batch_size=1, # the communication shape is obtained via p2p comm\n forward_only=forward_only,\n )\n\n if enable_routing_replay:\n if self.engine_config.router_replay.mode in [\"R3\"]:\n RouterReplay.clear_global_indices()\n RouterReplay.clear_global_router_replay_action()\n\n if self.model_config.mtp.enable and self.is_mp_src_rank_with_outputs():\n # add mtp_losses\n metrics = get_megatron_mtp_loss(n_micro_batch)\n if \"metrics\" not in losses_reduced[0]:\n losses_reduced[0][\"metrics\"] = {}\n losses_reduced[0][\"metrics\"].update(metrics)\n\n if mpu.is_pipeline_last_stage(ignore_virtual=True):\n output = postprocess_batch_func(output_lst=losses_reduced, indices=indices, data=data)\n return output\n else:\n return {}\n\n def get_per_tensor_param(self, base_sync_done=False, **kwargs):\n load_megatron_model_to_gpu(self.module, load_grad=False)\n peft_config = None\n non_merge_lora_sync = self.peft_cls is not None and not self.model_config.lora.get(\"merge\", False)\n if self.vanilla_bridge:\n per_tensor_param = self.bridge.export_weights(self.module)\n elif base_sync_done and non_merge_lora_sync:\n # Only export adapter weights\n peft_config = build_peft_config_for_vllm(self.model_config.lora)\n per_tensor_param = self.bridge.export_adapter_weights(self.module)\n else:\n per_tensor_param = self.bridge.export_hf_weights(self.module)\n if non_merge_lora_sync:\n per_tensor_param = add_base_layer_suffix(\n per_tensor_param, model_type=self.model_config.hf_config.model_type\n )\n return per_tensor_param, peft_config\n\n def disable_adapter(self) -> ContextManager:\n return self.peft_cls.disable_adapter(self.module)\n\n def forward_step(self, batch_iter, model, postprocess_micro_batch_func):\n raise NotImplementedError(\"forward_step must be implemented in subclass\")\n\n def postprocess_micro_batch_func(self, output, data: TensorDict, forward_only: bool, loss_function):\n raise NotImplementedError(\"postprocess_micro_batch_func must be implemented in subclass\")\n\n\nclass EngineEvalModeCtx(BaseEngineCtx):\n def __init__(self, engine: MegatronEngine, **kwargs):\n super().__init__(engine=engine, mode=\"eval\", **kwargs)\n\n def __enter__(self):\n assert isinstance(self.engine, MegatronEngine)\n super().__enter__()\n # mcore module is a list of model chunk in each vpp stage\n for module in self.engine.module:\n module.eval()\n\n def __exit__(self, exc_type, exc_value, traceback):\n assert isinstance(self.engine, MegatronEngine)\n super().__exit__(exc_type, exc_value, traceback)\n\n\nclass EngineTrainModeCtx(BaseEngineCtx):\n def __init__(self, engine: MegatronEngine, **kwargs):\n super().__init__(engine=engine, mode=\"train\", **kwargs)\n\n def __enter__(self):\n assert isinstance(self.engine, MegatronEngine)\n super().__enter__()\n # mcore module is a list of model chunk in each vpp stage\n for module in self.engine.module:\n module.train()\n\n def __exit__(self, exc_type, exc_value, traceback):\n assert isinstance(self.engine, MegatronEngine)\n self.engine.optimizer_zero_grad()\n super().__exit__(exc_type, exc_value, traceback)\n\n\n@EngineRegistry.register(model_type=\"language_model\", backend=\"megatron\")\nclass MegatronEngineWithLMHead(MegatronEngine):\n def prepare_model_inputs(self, batch: TensorDict):\n input_ids = batch[\"input_ids\"]\n loss_mask = batch[\"loss_mask\"].to(bool)\n multi_modal_inputs = extract_multi_modal_inputs(batch.get(\"multi_modal_inputs\", []))\n\n routed_experts = batch.get(\"routed_experts\", [])\n\n return {\n \"input_ids\": input_ids,\n \"loss_mask\": loss_mask,\n \"multi_modal_inputs\": multi_modal_inputs,\n \"routed_experts\": routed_experts,\n }\n\n def prepare_model_outputs(self, output: dict, data: TensorDict):\n calculate_entropy = tu.get_non_tensor_data(data, key=\"calculate_entropy\", default=False)\n\n log_prob = output[\"log_probs\"]\n model_output = {\"log_probs\": log_prob}\n if calculate_entropy:\n entropy = output[\"entropy\"]\n model_output[\"entropy\"] = entropy\n\n return model_output\n\n def forward_step(self, batch_iter: Iterator[TensorDict], model, postprocess_micro_batch_func):\n batch: TensorDict = next(batch_iter)\n batch = batch.to(get_device_id())\n use_fused_kernels = tu.get_non_tensor_data(batch, key=\"use_fused_kernels\", default=False)\n calculate_entropy = tu.get_non_tensor_data(batch, key=\"calculate_entropy\", default=False)\n pad_mode = tu.get_non_tensor_data(batch, key=\"pad_mode\", default=DatasetPadMode.NO_PADDING)\n temperature = batch[\"temperature\"]\n model_inputs = self.prepare_model_inputs(batch)\n input_ids = model_inputs[\"input_ids\"]\n multi_modal_inputs = model_inputs[\"multi_modal_inputs\"]\n loss_mask = model_inputs[\"loss_mask\"]\n\n unwrapped_model = unwrap_model(model)\n if hasattr(unwrapped_model, \"vp_stage\"):\n vp_rank = unwrapped_model.vp_stage\n else:\n vp_rank = 0\n\n if RouterReplayHelper.is_replay_backward_action(self.tf_config, vp_rank):\n router_instance_list = RouterReplayHelper.get_micro_batch_router_list(self.tf_config, vp_rank)\n for router in router_instance_list:\n router.set_router_replay_action(RouterReplayAction.REPLAY_FORWARD)\n\n if RouterReplayHelper.is_replay_forward_action(self.tf_config, vp_rank):\n layers_topk_idx = model_inputs[\"routed_experts\"]\n set_router_replay_data(layers_topk_idx, None, self.tf_config, vp_rank)\n\n if pad_mode == DatasetPadMode.NO_PADDING:\n label = input_ids.clone()\n else:\n raise NotImplementedError(f\"Pad mode {pad_mode} is not supported for megatron engine\")\n\n from verl.models.mcore import get_mcore_forward_no_padding_fn\n\n if use_fused_kernels:\n if not self.engine_config.use_remove_padding:\n logger.warning_once(\n \"Fused kernels require `use_remove_padding=True` for Megatron engine. Falling back to non-fused.\"\n )\n use_fused_kernels = False\n elif isinstance(temperature, torch.Tensor):\n if temperature.numel() != 1:\n logger.warning_once(\n \"Fused kernels do not support per-sample temperature. Falling back to non-fused.\"\n )\n use_fused_kernels = False\n else:\n temperature_value = float(temperature.item())\n else:\n temperature_value = float(temperature)\n\n if use_fused_kernels:\n fused_forward_fn = get_mcore_forward_fused_no_padding_fn(self.model_config.hf_config)\n output = fused_forward_fn(\n model=model,\n input_ids=input_ids,\n labels=label,\n multi_modal_inputs=multi_modal_inputs,\n temperature=temperature_value,\n calculate_entropy=calculate_entropy,\n pad_token_id=self.model_config.tokenizer.pad_token_id,\n )\n else:\n if not isinstance(temperature, torch.Tensor):\n temperature = torch.tensor([temperature] * input_ids.shape[0], device=input_ids.device)\n\n temperature = temperature.to(torch.float32)\n assert temperature.shape[0] == input_ids.shape[0]\n temperature = verl_F.expand_as_nested(temperature, input_ids) # (bsz, j1)\n\n forward_fn = get_mcore_forward_no_padding_fn(self.model_config.hf_config)\n\n def logits_processor(logits, label, temperature):\n assert logits.shape[:2] == label.shape[:2]\n # avoid non-positive temperature such as padding\n temperature[temperature <= 0] = 1e-8\n assert torch.all(temperature > 0).item(), f\"temperature tensor must be positive. Got {temperature}\"\n logits.div_(temperature.unsqueeze(dim=-1).to(logits.dtype))\n ret = {}\n if calculate_entropy:\n logits_bak = logits.clone()\n # # disable the hint until the fused_kernel is optimized for triton>=3.3\n # if torch.distributed.get_rank() == 0:\n # logger.warning_once(\n # \"For memory-efficient computation, enable fused kernels via \"\n # \"`actor_rollout_ref.model.use_fused_kernels=True`. \"\n # \"The current `clone()` operation ensures correctness but increases memory usage.\"\n # )\n entropy = vocab_parallel_entropy(logits)\n ret[\"entropy\"] = entropy\n else:\n logits_bak = logits\n\n log_probs = vocab_parallel_log_probs_from_logits(logits_bak, label)\n ret[\"log_probs\"] = log_probs\n return ret\n\n logits_processor_args = {\"label\": label, \"temperature\": temperature, \"loss_mask\": loss_mask}\n\n output = forward_fn(\n model,\n input_ids,\n multi_modal_inputs,\n logits_processor=logits_processor,\n logits_processor_args=logits_processor_args,\n vision_model=hasattr(self.model_config.hf_config, \"vision_config\"),\n pad_token_id=self.model_config.tokenizer.pad_token_id,\n data_format=\"thd\" if self.engine_config.use_remove_padding else \"bshd\",\n enable_mtp=self.model_config.mtp.enable_train,\n )\n\n # Router replay: switch to backward replay mode for next backward pass\n if RouterReplayHelper.is_replay_forward_action(self.tf_config, vp_rank):\n router_instance_list = RouterReplayHelper.get_micro_batch_router_list(self.tf_config, vp_rank)\n for router in router_instance_list:\n router.set_router_replay_action(RouterReplayAction.REPLAY_BACKWARD)\n\n return output, partial(postprocess_micro_batch_func, data=batch)\n\n def postprocess_micro_batch_func(self, output, data: TensorDict, forward_only: bool, loss_function):\n # For memory efficiency\n # We move calculation of entropy to compute_log_probs, forward_only == True\n device = data[\"input_ids\"].device\n model_output = self.prepare_model_outputs(output, data)\n\n if loss_function is not None:\n loss, metrics = loss_function(model_output=model_output, data=data, dp_group=self.get_data_parallel_group())\n # scale loss by num_micro_batch because megatron will scale loss\n # by n_micro_batch inside pp schedule\n scaled_loss = loss * data[\"num_micro_batch\"]\n else:\n assert forward_only, \"forward_only must be True when loss_function is None\"\n loss = torch.tensor(1.0, device=device)\n scaled_loss = loss\n metrics = {}\n\n output = {\n \"model_output\": model_output,\n \"loss\": loss.detach().item(),\n \"metrics\": metrics,\n }\n\n # return loss and stats\n return scaled_loss, output\n\n\n@EngineRegistry.register(model_type=\"value_model\", backend=\"megatron\")\nclass MegatronEngineWithValueHead(MegatronEngineWithLMHead):\n # for value head\n def forward_step(self, batch_iter, model, postprocess_micro_batch_func):\n batch: TensorDict = next(batch_iter)\n batch = batch.to(get_device_id())\n model_inputs = self.prepare_model_inputs(batch)\n input_ids = model_inputs[\"input_ids\"]\n multi_modal_inputs = model_inputs[\"multi_modal_inputs\"]\n\n from verl.models.mcore import get_mcore_forward_no_padding_fn\n\n forward_fn = get_mcore_forward_no_padding_fn(self.model_config.hf_config)\n\n output = forward_fn(\n model,\n input_ids,\n multi_modal_inputs,\n value_model=True,\n vision_model=hasattr(self.model_config.hf_config, \"vision_config\"),\n pad_token_id=self.model_config.tokenizer.pad_token_id,\n enable_mtp=self.model_config.mtp.enable_train,\n )\n\n return output, partial(postprocess_micro_batch_func, data=batch)\n\n def prepare_model_outputs(self, output: dict | torch.Tensor, data: TensorDict):\n return {\"values\": output}\n"}144{"file_name": "verl__workers__engine__mindspeed__transformer_impl.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport logging\nimport os\n\ntry:\n from mindspeed.megatron_adaptor import repatch\nexcept ImportError:\n repatch = None\n\nfrom verl.trainer.config import CheckpointConfig\nfrom verl.workers.config import HFModelConfig, McoreEngineConfig, McoreOptimizerConfig\n\nfrom ..base import EngineRegistry\nfrom ..megatron import MegatronEngineWithLMHead\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\n@EngineRegistry.register(model_type=\"language_model\", backend=\"megatron\", device=\"npu\")\nclass MindspeedEngineWithLMHead(MegatronEngineWithLMHead):\n def __init__(\n self,\n model_config: HFModelConfig,\n engine_config: McoreEngineConfig,\n optimizer_config: McoreOptimizerConfig,\n checkpoint_config: CheckpointConfig,\n ):\n super().__init__(model_config, engine_config, optimizer_config, checkpoint_config)\n\n repatch_config = {\"use_flash_attn\": True}\n if self.engine_config.context_parallel_size > 1:\n repatch_config[\"context_parallel_size\"] = self.engine_config.context_parallel_size\n\n repatch(repatch_config)\n"}145{"file_name": "verl__workers__engine__utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport os\nimport random\n\nimport numpy as np\nimport torch\nfrom tensordict import TensorDict\n\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.dataset.dataset_utils import DatasetPadMode\nfrom verl.utils.device import is_npu_available\nfrom verl.utils.py_functional import append_to_dict\nfrom verl.utils.seqlen_balancing import rearrange_micro_batches, restore_dynamic_batch\n\n\ndef enable_full_determinism(seed: int):\n \"\"\"\n Helper function for reproducibility in distributed training.\n See https://pytorch.org/docs/stable/notes/randomness.html for details.\n \"\"\"\n\n os.environ[\"PYTHONHASHSEED\"] = str(seed)\n os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":16:8\"\n os.environ[\"NCCL_DETERMINISTIC\"] = \"1\"\n os.environ[\"FLASH_ATTENTION_DETERMINISTIC\"] = \"1\"\n if is_npu_available:\n # The environment variable required to enable deterministic mode on Ascend NPUs.\n os.environ[\"NCCL_DETERMINISTIC\"] = \"true\"\n os.environ[\"CLOSE_MATMUL_K_SHIFT\"] = \"1\"\n\n random.seed(seed)\n np.random.seed(seed)\n torch.manual_seed(seed)\n torch.cuda.manual_seed(seed)\n torch.cuda.manual_seed_all(seed)\n torch.use_deterministic_algorithms(True, warn_only=True)\n # Enable CUDNN deterministic mode\n torch.backends.cudnn.deterministic = True\n torch.backends.cudnn.benchmark = False\n torch.backends.cudnn.enabled = False\n if is_npu_available:\n torch.npu.manual_seed(seed)\n torch.npu.manual_seed_all(seed)\n\n\ndef prepare_micro_batches(\n data: TensorDict,\n dp_group=None,\n num_batches_divided_by=None,\n same_micro_num_in_dp=True,\n min_num_micro_batch=None,\n use_dynamic_bsz_balance=True,\n):\n \"\"\"\n Prepare micro batches from data.\n \"\"\"\n use_dynamic_bsz = tu.get_non_tensor_data(data=data, key=\"use_dynamic_bsz\", default=True)\n sp_size = tu.get_non_tensor_data(data=data, key=\"sp_size\", default=1)\n\n if use_dynamic_bsz:\n assert \"max_token_len_per_gpu\" in data.keys(), \"max_token_len_per_gpu must be set when use_dynamic_bsz is True\"\n max_token_len_per_gpu = data[\"max_token_len_per_gpu\"]\n max_token_len = max_token_len_per_gpu * sp_size\n micro_batches, batch_idx_list = rearrange_micro_batches(\n data,\n max_token_len=max_token_len,\n dp_group=dp_group,\n num_batches_divided_by=num_batches_divided_by,\n same_micro_num_in_dp=same_micro_num_in_dp,\n min_num_micro_batch=min_num_micro_batch,\n use_dynamic_bsz_balance=use_dynamic_bsz_balance,\n )\n else:\n micro_batch_size_per_gpu = data[\"micro_batch_size_per_gpu\"]\n micro_batches = tu.chunk_tensordict(data, len(data) // micro_batch_size_per_gpu)\n batch_idx_list = None\n return micro_batches, batch_idx_list\n\n\ndef postprocess_batch_func(output_lst, indices, data: TensorDict):\n \"\"\"postprocess the output of a forward_backward_batch.\n output_lst is a list of dict containing outputs for each micro-batch\n reorder entropy and outputs. Return None for other pp ranks\n only on last rank. It should be on every tp rank\n\n each losses_reduced contains 1. model_output, 2. loss, 3. metrics.\n \"\"\"\n\n use_dynamic_bsz = tu.get_non_tensor_data(data=data, key=\"use_dynamic_bsz\", default=True)\n pad_mode = tu.get_non_tensor_data(data=data, key=\"pad_mode\", default=DatasetPadMode.NO_PADDING)\n assert pad_mode == DatasetPadMode.NO_PADDING, \"postprocess_batch_func only support NO_PADDING pad_mode\"\n\n # losses_reduced is a list of dict containing outputs for each micro-batch\n # reorder entropy and outputs. Return None for other pp ranks\n # only on last rank. It should be on every tp rank\n\n # losses_reduced contains 1. model_output, 2. loss, 3. metrics.\n # We perform reverse\n\n model_output = {}\n losses = []\n aggregated_metrics = {}\n\n # model output\n for o in output_lst:\n if \"model_output\" in o:\n for key, val in o[\"model_output\"].items():\n if key not in model_output:\n model_output[key] = []\n model_output[key].append(val)\n\n # concat results from micro batches\n for key, val in model_output.items():\n if pad_mode == DatasetPadMode.NO_PADDING:\n tensors = [tensor for nt in model_output[key] for tensor in nt.unbind()]\n model_output[key] = torch.nested.as_nested_tensor(tensors, layout=torch.jagged)\n else:\n raise NotImplementedError(f\"pad_mode {pad_mode} not implemented\")\n\n # reverse with dynamic bsz\n if use_dynamic_bsz:\n model_output[key] = restore_dynamic_batch(model_output[key], indices)\n\n # loss\n for o in output_lst:\n if \"loss\" in o:\n losses.append(o[\"loss\"])\n\n # metrics\n for o in output_lst:\n if \"metrics\" in o:\n metrics = o[\"metrics\"]\n append_to_dict(aggregated_metrics, metrics)\n\n output = {\n \"model_output\": model_output,\n \"loss\": losses,\n \"metrics\": aggregated_metrics,\n }\n\n return output\n"}146{"file_name": "verl__workers__engine__veomni__transformer_impl.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nimport logging\nfrom dataclasses import dataclass, field\nfrom typing import Any, Callable, Optional, Sequence\n\nimport torch\nimport torch.distributed as dist\nfrom tensordict import TensorDict\nfrom torch.distributed.tensor import DTensor\nfrom veomni.distributed import parallel_state\nfrom veomni.distributed.offloading import build_activation_offloading_context\nfrom veomni.distributed.torch_parallelize import build_parallelize_model\nfrom veomni.models.auto import build_foundation_model\nfrom veomni.optim import build_lr_scheduler, build_optimizer\n\nimport verl.utils.torch_functional as verl_F\nfrom verl.trainer.config import CheckpointConfig\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.checkpoint.fsdp_checkpoint_manager import FSDPCheckpointManager\nfrom verl.utils.device import get_device_id, get_device_name\nfrom verl.utils.fsdp_utils import fsdp_version\nfrom verl.utils.model import convert_weight_keys\nfrom verl.utils.profiler import log_gpu_memory_usage\nfrom verl.utils.ulysses import (\n get_ulysses_sequence_parallel_group,\n set_ulysses_sequence_parallel_group,\n)\nfrom verl.workers.config import HFModelConfig, VeOmniEngineConfig, VeOmniOptimizerConfig\n\nfrom ..base import BaseEngineCtx, EngineRegistry\nfrom ..fsdp.transformer_impl import FSDPEngine, FSDPEngineWithLMHead\nfrom ..utils import enable_full_determinism, postprocess_batch_func, prepare_micro_batches\nfrom .utils import (\n MOE_PARAM_HANDERS,\n VL_TYPE2INDEX,\n load_veomni_model_to_gpu,\n load_veomni_optimizer,\n offload_veomni_model_to_cpu,\n offload_veomni_optimizer,\n)\n\nlogger = logging.getLogger(__file__)\n\n\nclass VeOmniEngine(FSDPEngine):\n def __init__(\n self,\n model_config: HFModelConfig,\n engine_config: VeOmniEngineConfig,\n optimizer_config: VeOmniOptimizerConfig,\n checkpoint_config: CheckpointConfig,\n **kwargs,\n ):\n \"\"\"\n Initialize the VeOmniEngine.\n\n Sets up distributed device meshes, LoRA, and offload policies based on config.\n\n Args:\n config: Configuration object with VeOmni and model settings.\n \"\"\"\n\n self.model_config = model_config\n self.engine_config = engine_config\n self.optimizer_config = optimizer_config\n self.checkpoint_config = checkpoint_config\n # VeOmniEngine only supports fsdp2.\n self.data_parallel_mode = \"fsdp2\"\n self.rank = dist.get_rank()\n\n fsdp_size = self.engine_config.fsdp_size\n world_size = dist.get_world_size()\n dp_size = world_size // self.engine_config.ulysses_parallel_size\n\n if fsdp_size < 0 or fsdp_size >= dp_size:\n data_parallel_replicate_size = 1\n data_parallel_shard_size = dp_size\n else:\n if dp_size % fsdp_size != 0:\n raise ValueError(\n f\"Data parallel size ({dp_size}) must be divisible by fsdp_size ({fsdp_size}). \"\n \"Please adjust your parallel configuration.\"\n )\n data_parallel_replicate_size = dp_size // fsdp_size\n data_parallel_shard_size = fsdp_size\n\n parallel_state.init_parallel_state(\n dp_size=dp_size,\n dp_replicate_size=data_parallel_replicate_size,\n dp_shard_size=data_parallel_shard_size,\n ep_size=self.engine_config.expert_parallel_size,\n ulysses_size=self.engine_config.ulysses_parallel_size,\n dp_mode=self.data_parallel_mode,\n )\n\n if self.engine_config.full_determinism:\n enable_full_determinism(seed=self.engine_config.seed)\n\n self.use_remove_padding = self.model_config.use_remove_padding\n\n self._is_offload_param = self.engine_config.param_offload\n self._is_offload_optimizer = self.engine_config.optimizer_offload\n self._is_lora = self.model_config.lora_rank > 0\n\n self.use_ulysses_sp = parallel_state.get_parallel_state().sp_enabled\n self.ulysses_sequence_parallel_size = self.engine_config.ulysses_parallel_size\n\n if self.use_ulysses_sp:\n self.ulysses_parallel_group = parallel_state.get_parallel_state().device_mesh[\"sp\"].get_group()\n else:\n self.ulysses_parallel_group = None\n\n if self.engine_config.entropy_from_logits_with_chunking:\n entropy_from_logits = verl_F.entropy_from_logits_with_chunking\n else:\n entropy_from_logits = verl_F.entropy_from_logits\n\n self.compute_entropy_from_logits = (\n torch.compile(entropy_from_logits, dynamic=True)\n if self.engine_config.use_torch_compile # use torch compile by default\n else entropy_from_logits\n )\n\n def initialize(self):\n \"\"\"\n Build the model, optimizer, and learning rate scheduler under VeOmni.\n\n Applies device, dtype, and precision configurations, including mixed precision.\n Sets up checkpoint manager and FLOPs counter.\n \"\"\"\n self._build_model_optimizer()\n\n self.checkpoint_manager = FSDPCheckpointManager(\n model=self.module,\n optimizer=self.optimizer,\n lr_scheduler=self.lr_scheduler,\n processing_class=self.model_config.get_processor(),\n checkpoint_config=self.checkpoint_config,\n trust_remote_code=self.model_config.trust_remote_code,\n )\n\n self.to(\n device=\"cpu\",\n model=self._is_offload_param,\n optimizer=self._is_offload_optimizer,\n grad=self._is_offload_optimizer,\n )\n\n log_gpu_memory_usage(\"After offload model/optimizer/grad during init\", logger=logger)\n\n def _build_optimizer(self, module):\n optimizer = build_optimizer(\n module,\n lr=self.optimizer_config.lr,\n betas=self.optimizer_config.betas,\n weight_decay=self.optimizer_config.weight_decay,\n optimizer_type=self.optimizer_config.optimizer,\n )\n get_optimizer_pre_hook = getattr(module, \"get_optimizer_pre_hook\", None)\n if get_optimizer_pre_hook is not None:\n optimizer_pre_hook = get_optimizer_pre_hook(module, module.config, self.data_parallel_mode)\n optimizer.register_step_pre_hook(optimizer_pre_hook)\n\n return optimizer\n\n def _build_lr_scheduler(self, optimizer):\n optim_config = self.optimizer_config\n lr_scheduler = build_lr_scheduler(\n optimizer,\n train_steps=optim_config.total_training_steps,\n lr=optim_config.lr,\n lr_min=optim_config.lr_min,\n lr_decay_style=optim_config.lr_scheduler_type,\n lr_decay_ratio=optim_config.lr_decay_ratio,\n lr_warmup_ratio=optim_config.lr_warmup_steps_ratio,\n lr_start=optim_config.lr_start,\n )\n\n return lr_scheduler\n\n def _build_model_optimizer(self):\n # Load base model with specified configuration and dtype\n module = build_foundation_model(\n config_path=self.model_config.hf_config_path,\n weights_path=self.model_config.path,\n torch_dtype=\"float32\" if self.engine_config.mixed_precision else \"bfloat16\",\n attn_implementation=self.engine_config.attn_implementation,\n moe_implementation=self.engine_config.moe_implementation,\n init_device=self.engine_config.init_device,\n )\n log_gpu_memory_usage(\"After load base model\", logger=logger)\n\n # Applies parallel strategies to the model.\n log_gpu_memory_usage(\"Before parallelize model\", logger=logger)\n module = build_parallelize_model(\n module,\n init_device=self.engine_config.init_device,\n weights_path=self.model_config.path,\n enable_full_shard=self.engine_config.enable_full_shard,\n enable_mixed_precision=self.engine_config.mixed_precision,\n enable_gradient_checkpointing=self.model_config.enable_gradient_checkpointing,\n enable_fsdp_offload=self.engine_config.enable_fsdp_offload,\n basic_modules=module._no_split_modules + self.engine_config.basic_modules,\n enable_reentrant=self.engine_config.enable_reentrant,\n enable_forward_prefetch=self.engine_config.forward_prefetch,\n )\n log_gpu_memory_usage(\"After parallelize model\", logger=logger)\n\n if not self.engine_config.forward_only:\n # Initialize optimizer with model parameters and config settings\n optimizer = self._build_optimizer(module)\n # Create learning rate scheduler with warmup and decay settings\n lr_scheduler = self._build_lr_scheduler(optimizer)\n else:\n optimizer = None\n lr_scheduler = None\n\n self.module = module\n self.optimizer = optimizer\n self.lr_scheduler = lr_scheduler\n self.model_fwd_context, self.model_bwd_context = build_activation_offloading_context(\n self.model_config.enable_activation_offload,\n self.model_config.enable_gradient_checkpointing,\n self.engine_config.activation_gpu_limit,\n )\n\n def optimizer_step(self):\n \"\"\"\n Perform an optimization step using the optimizer.\n \"\"\"\n if hasattr(self.module, \"clip_grad_norm_\"):\n grad_norm = self.module.clip_grad_norm_(self.optimizer_config.clip_grad)\n else:\n grad_norm = torch.nn.utils.clip_grad_norm_(self.module.parameters(), self.optimizer_config.clip_grad)\n\n if isinstance(grad_norm, DTensor):\n grad_norm = grad_norm.full_tensor()\n\n # if grad_norm is not finite, skip the update\n if not torch.isfinite(grad_norm):\n print(f\"WARN: grad_norm is not finite: {grad_norm}\")\n self.optimizer.zero_grad()\n else:\n self.optimizer.step()\n return grad_norm.item()\n\n def forward_backward_batch(self, data: TensorDict, loss_function: Callable, forward_only=False) -> Any:\n \"\"\"\n Perform a forward pass and optionally a backward pass on a batch of data.\n\n Args:\n data: The input data for the forward pass, typically containing tensors and metadata.\n loss_function: The loss function to optimize. See `verl.workers.roles.utils.losses` for examples.\n forward_only: If True, perform only the forward pass. If False, perform forward and backward pass.\n\n Returns:\n Any: The output of the forward pass, which can be used for loss computation or other purposes.\n \"\"\"\n tu.assign_non_tensor(data, sp_size=parallel_state.get_parallel_state().ulysses_size)\n\n # compute num_tokens in global batch for loss normalization\n batch_num_tokens = data[\"loss_mask\"].sum().to(get_device_id())\n torch.distributed.all_reduce(\n batch_num_tokens, op=torch.distributed.ReduceOp.SUM, group=self.get_data_parallel_group()\n )\n tu.assign_non_tensor(data, batch_num_tokens=batch_num_tokens.item())\n tu.assign_non_tensor(data, dp_size=self.get_data_parallel_size())\n\n micro_batches, indices = prepare_micro_batches(\n data=data, dp_group=self.get_data_parallel_group(), same_micro_num_in_dp=True\n )\n\n output_lst = []\n\n for micro_batch in micro_batches:\n with self.model_fwd_context:\n loss, meta_info = self.forward_step(micro_batch, loss_function=loss_function, forward_only=forward_only)\n if not forward_only:\n with self.model_bwd_context:\n loss.backward()\n\n output_lst.append(meta_info)\n\n return postprocess_batch_func(output_lst=output_lst, indices=indices, data=data)\n\n def get_data_parallel_rank(self):\n return parallel_state.get_parallel_state().device_mesh.get_local_rank(\"dp\")\n\n def get_data_parallel_size(self):\n return torch.distributed.get_world_size() // parallel_state.get_parallel_state().ulysses_size\n\n def get_data_parallel_group(self):\n if parallel_state.get_parallel_state().ulysses_size > 1:\n return parallel_state.get_parallel_state().device_mesh.get_group(mesh_dim=\"dp\")\n else:\n return torch.distributed.group.WORLD\n\n def is_mp_src_rank_with_outputs(self):\n \"\"\"\n Whether the current rank is the first rank in model parallel group that contains model outputs\n \"\"\"\n if parallel_state.get_parallel_state().ulysses_size > 1:\n is_collect = parallel_state.get_parallel_state().device_mesh[\"ulysses\"].get_local_rank() == 0\n else:\n is_collect = True\n return is_collect\n\n def train_mode(self, **kwargs):\n \"\"\"\n Return a context manager that switches to training mode with VeOmni-specific handling.\n\n Includes parameter and optimizer offload entry/exit.\n \"\"\"\n return EngineTrainModeCtx(self, **kwargs)\n\n def eval_mode(self, **kwargs):\n \"\"\"\n Return a context manager that switches to evaluation mode with VeOmni-specific handling.\n\n Includes activation offload entry/exit.\n \"\"\"\n return EngineEvalModeCtx(self, **kwargs)\n\n def to(self, device: str, model: bool = True, optimizer: bool = True, grad: bool = True):\n \"\"\"\n Move model parameters, optimizer states, or both to the specified device.\n Note that this function executes irrespective of offload config. It serves as manual control.\n\n Args:\n device: Target device identifier.\n model: If True, move the model.\n optimizer: If True, move the optimizer states.\n \"\"\"\n super(FSDPEngine, self).to(device=device, model=model, optimizer=optimizer, grad=grad)\n\n device_name = get_device_name()\n\n assert device in (device_name, \"cpu\")\n if device == device_name:\n if model:\n load_veomni_model_to_gpu(self.module)\n if optimizer and self.optimizer is not None:\n load_veomni_optimizer(self.optimizer, device)\n elif device == \"cpu\":\n if model:\n offload_veomni_model_to_cpu(self.module)\n if optimizer and self.optimizer is not None:\n offload_veomni_optimizer(self.optimizer)\n else:\n raise ValueError(f\"Invalid device type: {device}\")\n\n def save_checkpoint(\n self,\n local_path: str,\n hdfs_path: Optional[str] = None,\n global_step: int = 0,\n max_ckpt_to_keep: Optional[int] = None,\n **kwargs,\n ) -> None:\n \"\"\"\n Save VeOmni checkpoint, handling parameter offload as needed.\n \"\"\"\n origin_module_device = next(self.module.parameters()).device.type\n if self._is_offload_param or origin_module_device == \"cpu\":\n load_veomni_model_to_gpu(self.module)\n\n self.checkpoint_manager.save_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, global_step=global_step, max_ckpt_to_keep=max_ckpt_to_keep\n )\n\n torch.distributed.barrier()\n if self._is_offload_param:\n offload_veomni_model_to_cpu(self.module)\n\n def load_checkpoint(\n self, local_path: str, hdfs_path: Optional[str] = None, del_local_after_load: int = True, **kwargs\n ) -> None:\n \"\"\"\n Load VeOmni checkpoint, restoring parameters and optimizer state.\n \"\"\"\n if self._is_offload_param:\n load_veomni_model_to_gpu(self.module)\n\n self.checkpoint_manager.load_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, del_local_after_load=del_local_after_load\n )\n\n torch.distributed.barrier()\n if self._is_offload_param:\n offload_veomni_model_to_cpu(self.module)\n\n if self._is_offload_optimizer:\n offload_veomni_optimizer(self.optimizer)\n\n def get_per_tensor_param(self, **kwargs):\n load_veomni_model_to_gpu(self.module)\n\n params = self.module.state_dict()\n params = convert_weight_keys(params, getattr(self.module, \"_fsdp_wrapped_module\", self.module))\n\n if self._is_offload_param:\n offload_veomni_model_to_cpu(self.module)\n\n device = get_device_id()\n ps = parallel_state.get_parallel_state()\n model_type = getattr(self.module.config, \"model_type\", \"default\")\n process_func = MOE_PARAM_HANDERS.get(model_type, lambda n, t: iter([(n, t)]))\n\n def param_generator():\n for name, param in params.items():\n unsharded_tensor = param.full_tensor() if isinstance(param, DTensor) else param\n\n is_expert_layer = \"mlp.experts.\" in name\n is_proj = any(p in name for p in [\"down_proj\", \"gate_proj\", \"up_proj\", \"gate_up_proj\"])\n\n if is_expert_layer and is_proj and ps.ep_enabled:\n output_shape = list(unsharded_tensor.shape)\n output_shape[0] *= ps.ep_size\n stacked_tensor = torch.empty(output_shape, dtype=unsharded_tensor.dtype, device=device)\n\n # all gather expert tensors [32, H, I] -> [128, H, I]\n torch.distributed.all_gather_into_tensor(stacked_tensor, unsharded_tensor, group=ps.ep_group)\n yield from process_func(name, stacked_tensor)\n\n del stacked_tensor\n else:\n if is_expert_layer:\n yield from process_func(name, unsharded_tensor)\n else:\n yield name, unsharded_tensor\n\n # TODO: support VeOmni LoRA\n return param_generator(), None\n\n\nclass EngineEvalModeCtx(BaseEngineCtx):\n def __init__(self, engine: VeOmniEngine, **kwargs):\n super().__init__(engine=engine, mode=\"eval\", **kwargs)\n\n def __enter__(self):\n assert isinstance(self.engine, VeOmniEngine)\n super().__enter__()\n self.prev_sp_group = get_ulysses_sequence_parallel_group()\n set_ulysses_sequence_parallel_group(self.engine.ulysses_parallel_group)\n self.engine.module.train()\n\n def __exit__(self, exc_type, exc_value, traceback):\n assert isinstance(self.engine, VeOmniEngine)\n set_ulysses_sequence_parallel_group(self.prev_sp_group)\n\n # https://pytorch.org/docs/stable/notes/fsdp.html#fsdp-notes\n # unshard the root FSDP module\n if parallel_state.get_parallel_state().dp_shard_size > 1:\n if fsdp_version(self.engine.module) == 1:\n self.engine.module._handle.reshard(True)\n elif fsdp_version(self.engine.module) == 2:\n self.engine.module.reshard()\n\n super().__exit__(exc_type, exc_value, traceback)\n\n\nclass EngineTrainModeCtx(BaseEngineCtx):\n def __init__(self, engine: VeOmniEngine, **kwargs):\n super().__init__(engine=engine, mode=\"train\", **kwargs)\n\n def __enter__(self):\n assert isinstance(self.engine, VeOmniEngine)\n super().__enter__()\n self.prev_sp_group = get_ulysses_sequence_parallel_group()\n set_ulysses_sequence_parallel_group(self.engine.ulysses_parallel_group)\n # TODO: Switch to eval mode after Integrating the CI environment\n # VeOmni (ref: https://github.com/ByteDance-Seed/VeOmni/pull/421)\n self.engine.module.train()\n\n def __exit__(self, exc_type, exc_value, traceback):\n assert isinstance(self.engine, VeOmniEngine)\n set_ulysses_sequence_parallel_group(self.prev_sp_group)\n self.engine.optimizer_zero_grad()\n super().__exit__(exc_type, exc_value, traceback)\n\n\n@dataclass\nclass OmniSequenceShardCollator:\n \"\"\"\n Data collator to chunk inputs along the sequence length.\n \"\"\"\n\n # features to slice sequence dimension\n sp_slice_features: dict[str, int] = field(\n default_factory=lambda: {\n \"input_ids\": -1,\n \"labels\": -1,\n \"pixel_values\": 0,\n \"pixel_values_videos\": 0,\n },\n metadata={\"help\": \"features to slice sequence dimension.\"},\n )\n\n # features to padding sequence dimension\n padding_features: dict[str, int] = field(\n default_factory=lambda: {\n \"pixel_values\": 0,\n },\n metadata={\"help\": \"features to padding sequence dimension.\"},\n )\n\n # padding scale for padding features\n padding_scale: dict[str, int] = field(\n default_factory=lambda: {\"pixel_values\": 4}, metadata={\"help\": \"padding scale for padding features.\"}\n )\n\n def __post_init__(self):\n self.sp_size = parallel_state.get_parallel_state().sp_size\n self.sp_rank = parallel_state.get_parallel_state().sp_rank\n\n def sp_slice(self, feature: torch.Tensor, dim: int = -1) -> dict[str, \"torch.Tensor\"]:\n seq_length = feature.size(dim)\n sp_chunk_size = (seq_length + self.sp_size - 1) // self.sp_size\n return feature.narrow(dim, self.sp_rank * sp_chunk_size, sp_chunk_size)\n\n def sp_padding(\n self, tensor: \"torch.Tensor\", dim: int = -1, pad_value: int = 0, pad_scale: int = 1\n ) -> \"torch.Tensor\":\n \"\"\"\n Pads a tensor with pad_length to aligns tensor with sp size.\n \"\"\"\n seq_length = tensor.size(dim)\n scale_sp_size = self.sp_size * pad_scale\n\n sp_chunk_size = (seq_length + scale_sp_size - 1) // scale_sp_size\n pad_size = sp_chunk_size * scale_sp_size - seq_length\n if pad_size == 0:\n return tensor\n\n pad_shape = list(tensor.shape)\n pad_shape[dim] = pad_size\n pad = torch.full(pad_shape, fill_value=pad_value, dtype=tensor.dtype, device=tensor.device)\n return torch.cat((tensor, pad), dim=dim)\n\n def __call__(self, batch: Sequence[dict[str, \"torch.Tensor\"]]) -> dict[str, \"torch.Tensor\"]:\n for key in batch.keys():\n if key in self.padding_features.keys():\n batch[key] = self.sp_padding(\n batch[key],\n dim=self.sp_slice_features.get(key, -1),\n pad_value=self.padding_features[key],\n pad_scale=self.padding_scale.get(key, 1),\n )\n\n # sp slice\n for key in batch.keys():\n if key in self.sp_slice_features.keys():\n batch[key] = self.sp_slice(batch[key], dim=self.sp_slice_features[key])\n\n return batch\n\n\n@EngineRegistry.register(model_type=\"language_model\", backend=[\"veomni\"], device=[\"cuda\", \"npu\"])\nclass VeOmniEngineWithLMHead(VeOmniEngine, FSDPEngineWithLMHead):\n def prepare_model_inputs(self, micro_batch: TensorDict):\n # TODO: Cannot work properly for qwen_vl ulysses\n model_inputs, output_args = super().prepare_model_inputs(micro_batch)\n input_ids_rmpad = model_inputs[\"input_ids\"]\n if self.module.config.model_type in VL_TYPE2INDEX.keys():\n image_mask = input_ids_rmpad == VL_TYPE2INDEX[self.module.config.model_type][\"IMAGE_INPUT_INDEX\"]\n video_mask = input_ids_rmpad == VL_TYPE2INDEX[self.module.config.model_type][\"VIDEO_INPUT_INDEX\"]\n model_inputs.update({\"image_mask\": image_mask, \"video_mask\": video_mask})\n\n if parallel_state.get_parallel_state().sp_enabled:\n omni_sequence_shard_collator = OmniSequenceShardCollator()\n omni_sequence_shard_collator(model_inputs)\n\n return model_inputs, output_args\n"}147{"file_name": "verl__workers__engine__veomni__utils.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport torch\n\nfrom verl.utils.device import get_device_id, get_torch_device\n\nVL_TYPE2INDEX = {\n \"qwen2_5_vl\": {\n \"IMAGE_INPUT_INDEX\": 151655,\n \"VIDEO_INPUT_INDEX\": 151656,\n },\n \"qwen3_vl\": {\n \"IMAGE_INPUT_INDEX\": 151655,\n \"VIDEO_INPUT_INDEX\": 151656,\n },\n \"qwen3_vl_moe\": {\n \"IMAGE_INPUT_INDEX\": 151655,\n \"VIDEO_INPUT_INDEX\": 151656,\n },\n}\n\n\n@torch.no_grad()\ndef offload_veomni_model_to_cpu(model, empty_cache: bool = True):\n from torch.distributed.fsdp._fully_shard._fsdp_common import TrainingState\n from torch.distributed.fsdp._fully_shard._fsdp_state import _get_module_fsdp_state\n\n for module in model.modules():\n state = _get_module_fsdp_state(module)\n if state is None:\n continue\n fsdp_param_group = state._fsdp_param_group\n\n if fsdp_param_group is None:\n continue\n\n fsdp_param_group._training_state = TrainingState.IDLE\n\n model.reshard()\n model.cpu()\n if empty_cache:\n get_torch_device().empty_cache()\n\n\n@torch.no_grad()\ndef load_veomni_model_to_gpu(model):\n device = get_device_id()\n model.to(device)\n\n\n@torch.no_grad()\ndef offload_veomni_optimizer(optimizer):\n optimizers = []\n # Check if this is a MultiOptimizer (for ep and non-ep parameters when ep+fsdp2 is enabled)\n if hasattr(optimizer, \"_is_multi_optimizer\") and optimizer._is_multi_optimizer:\n optimizers.extend(optimizer.optimizers_dict.values())\n else:\n optimizers.append(optimizer)\n\n for opt in optimizers:\n if not opt.state:\n continue\n for param_group in opt.param_groups:\n for param in param_group[\"params\"]:\n state = opt.state[param]\n for key, value in state.items():\n if isinstance(value, torch.Tensor):\n state[key] = value.to(\"cpu\", non_blocking=True)\n\n\n@torch.no_grad()\ndef load_veomni_optimizer(optimizer, device_id):\n optimizers = []\n # Check if this is a MultiOptimizer (for ep and non-ep parameters when ep+fsdp2 is enabled)\n if hasattr(optimizer, \"_is_multi_optimizer\") and optimizer._is_multi_optimizer:\n optimizers.extend(optimizer.optimizers_dict.values())\n else:\n optimizers.append(optimizer)\n\n for opt in optimizers:\n if not opt.state:\n continue\n for param_group in opt.param_groups:\n for param in param_group[\"params\"]:\n state = opt.state[param]\n for key, value in state.items():\n if isinstance(value, torch.Tensor):\n state[key] = value.to(device_id, non_blocking=True)\n\n\ndef _map_moe_params_qwen3_moe(name, tensor):\n for i in range(tensor.size(0)):\n new_key = name.replace(\"mlp.experts.\", f\"mlp.experts.{i}.\") + \".weight\"\n yield new_key, tensor[i].to(get_device_id(), non_blocking=True)\n\n\nMOE_PARAM_HANDERS = {\n \"qwen3_moe\": _map_moe_params_qwen3_moe,\n}\n"}148{"file_name": "verl__workers__engine_workers.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport functools\nimport logging\nimport os\nfrom contextlib import nullcontext\nfrom functools import partial\nfrom itertools import chain\n\nimport torch\nfrom codetiming import Timer\nfrom omegaconf import DictConfig, open_dict\nfrom tensordict import NonTensorData, TensorDict\nfrom torch.distributed.device_mesh import init_device_mesh\n\ntry:\n from verl.workers.engine.mindspeed.transformer_impl import repatch\nexcept ImportError:\n repatch = None\nfrom verl.checkpoint_engine import CheckpointEngineRegistry\nfrom verl.single_controller.base import Worker\nfrom verl.single_controller.base.decorator import Dispatch, make_nd_compute_dataproto_dispatch_fn, register\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.utils.device import get_device_name, set_expandable_segments\nfrom verl.utils.distributed import initialize_global_process_group_ray\nfrom verl.utils.flops_counter import FlopsCounter\nfrom verl.utils.memory_utils import aggressive_empty_cache\nfrom verl.utils.metric.utils import Metric\nfrom verl.utils.profiler import DistProfiler, DistProfilerExtension, ProfilerConfig, log_gpu_memory_usage\nfrom verl.utils.py_functional import append_to_dict\nfrom verl.utils.tensordict_utils import maybe_fix_3d_position_ids\nfrom verl.utils.torch_functional import allgather_dict_into_dict\nfrom verl.workers.config import ActorConfig, HFModelConfig, RolloutConfig, TrainingWorkerConfig\nfrom verl.workers.rollout.base import BaseRollout, get_rollout_class\nfrom verl.workers.utils.losses import ppo_loss\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef _with_routing_replay_flag(enabled: bool):\n \"\"\"Decorator to set 'enable_routing_replay' flag on the data TensorDict.\"\"\"\n\n def decorator(func):\n @functools.wraps(func)\n def wrapper(self, data: TensorDict, *args, **kwargs):\n if self.enable_routing_replay:\n tu.assign_non_tensor_data(data, \"enable_routing_replay\", enabled)\n return func(self, data, *args, **kwargs)\n\n return wrapper\n\n return decorator\n\n\nclass TrainingWorker(Worker, DistProfilerExtension):\n \"\"\"\n TrainingWorker provides a Tinker-like API (https://thinkingmachines.ai/tinker/) as a RayWorkerGroup\n to a single controller. Currently, we only provide more coarse grained APIs,\n and do not provide exact APIs as Tinker does. But this can be added in the future.\n \"\"\"\n\n def __init__(self, config: TrainingWorkerConfig):\n Worker.__init__(self)\n\n from verl.workers.engine import BaseEngine, EngineRegistry\n\n initialize_global_process_group_ray(timeout_second=None)\n\n self.config = config\n self.model_config = self.config.model_config\n self.engine_config = self.config.engine_config\n self.optimizer_config = self.config.optimizer_config\n self.checkpoint_config = self.config.checkpoint_config\n self.device_name = get_device_name()\n\n if self.engine_config is None:\n assert self.optimizer_config is None\n if self.config.auto_select_engine_optim_fn is None:\n raise ValueError(\n \"engine_config is not provided and auto_select_engine_optim_fn is not set. \"\n \"Cannot determine engine backend.\"\n )\n # Support automatically select engine backend given model config\n self.engine_config, self.optimizer_config = self.config.auto_select_engine_optim_fn(\n self.model_config, self.device_name\n )\n\n # we use the one defined in model\n # TODO: this is not elegant and should refactor later\n self.engine_config.use_remove_padding = self.model_config.use_remove_padding\n self.engine_config.use_fused_kernels = self.model_config.use_fused_kernels\n\n if repatch is not None:\n # NPU MindSpeed patch, will be refactored with MindSpeedEngine.\n repatch(self.engine_config.get(\"override_transformer_config\", {}))\n\n # TODO: add DistProfilerExtension\n self.profiler_config = self.config.profiler_config\n if self.profiler_config is not None:\n self.profiler_tool_config = self.profiler_config.tool_config.get(self.profiler_config.tool, {})\n else:\n self.profiler_tool_config = None\n\n DistProfilerExtension.__init__(\n self, DistProfiler(rank=self.rank, config=self.profiler_config, tool_config=self.profiler_tool_config)\n )\n\n self.engine: BaseEngine = EngineRegistry.new(\n model_type=self.config.model_type,\n backend=self.engine_config.strategy,\n model_config=self.model_config,\n engine_config=self.engine_config,\n optimizer_config=self.optimizer_config,\n checkpoint_config=self.checkpoint_config,\n )\n\n # build dispatch info\n self._register_dispatch_collect_info(\n mesh_name=\"train\",\n dp_rank=self.engine.get_data_parallel_rank(),\n is_collect=self.engine.is_mp_src_rank_with_outputs(),\n )\n\n self.flops_counter = FlopsCounter(self.model_config.hf_config)\n\n self.loss_fn = None\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def to(self, device, model=True, optimizer=True, grad=True):\n \"\"\"Manual control of load/offload\"\"\"\n assert device in [\"cpu\", \"device\"]\n\n if device == \"device\":\n device = get_device_name()\n\n self.engine.to(device=device, model=model, optimizer=optimizer, grad=grad)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def set_loss_fn(self, loss_fn):\n self.loss_fn = loss_fn\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def reset(self):\n \"\"\"\n Reset the model engine to the initial state. If the engine is not initialized,\n we initialize it. Otherwise, reload ckpt and reset states\n \"\"\"\n self.engine.initialize()\n\n def _postprocess_output(self, output, *, global_token_num, delta_time, forward_only, images_seqlens):\n \"\"\"\n\n Args:\n output: a dictionary containing loss, model_outputs and metrics\n\n Returns:\n\n \"\"\"\n # TODO: whether to log memory\n # metrics[\"perf/max_memory_allocated_gb\"] = get_torch_device().max_memory_allocated() / (1024 ** 3)\n # metrics[\"perf/max_memory_reserved_gb\"] = get_torch_device().max_memory_reserved() / (1024 ** 3)\n # metrics[\"perf/cpu_memory_used_gb\"] = psutil.virtual_memory().used / (1024 ** 3)\n\n metrics: dict = output.pop(\"metrics\")\n # perform all gather in dp group to ensure that it's correct.\n # Here each metric in metrics can be a list (micro-batch metrics) or a singleton\n # we should always sum the loss of each micro-batch as we scale by global_bsz/global_token\n loss = torch.sum(torch.tensor(output.pop(\"loss\"), device=self.device_name))\n torch.distributed.all_reduce(\n loss, op=torch.distributed.ReduceOp.AVG, group=self.engine.get_data_parallel_group()\n )\n loss = loss.item()\n\n # For grad_norm, we do not perform all reduce because it is already been done when clipping grad\n grad_norm = metrics.pop(\"grad_norm\", None)\n lr = metrics.pop(\"lr\", None)\n\n # For other metrics, we perform all gather in dp group\n final_metrics = allgather_dict_into_dict(data=metrics, group=self.engine.get_data_parallel_group())\n final_metrics[\"loss\"] = loss\n if grad_norm is not None:\n final_metrics[\"grad_norm\"] = grad_norm\n if lr is not None:\n final_metrics[\"lr\"] = lr\n\n # TODO: confirm the mtp loss IS same across dp\n for k, v in final_metrics.items():\n if k.startswith(\"mtp_losses\"):\n flatten_v = [sublist[0] for sublist in v] # sublist should be single element\n final_metrics[k] = sum(flatten_v) / len(flatten_v)\n # compute mfu\n if global_token_num is not None:\n estimated_flops, promised_flops = self.flops_counter.estimate_flops(\n global_token_num, delta_time, images_seqlens=images_seqlens\n )\n final_metrics[\"mfu\"] = estimated_flops / promised_flops / torch.distributed.get_world_size()\n if forward_only:\n final_metrics[\"mfu\"] /= 3.0\n # model outputs\n model_output = output.pop(\"model_output\", {})\n # We only return final_metrics\n final_output = tu.get_tensordict(tensor_dict=model_output, non_tensor_dict={\"metrics\": final_metrics})\n return final_output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"train\"), blocking=False)\n def train_mini_batch(self, data: TensorDict) -> TensorDict:\n \"\"\"Split a batch into N mini-batches run for multiple epochs\n\n Args:\n data:\n\n Returns:\n\n \"\"\"\n maybe_fix_3d_position_ids(data)\n batch_size_per_dp = data.shape[0]\n disable_auto_offload = tu.pop(data, key=\"disable_auto_offload\", default=False)\n mini_batch_size = tu.pop(data, key=\"mini_batch_size\", default=None)\n num_mini_batch = tu.pop(data, key=\"num_mini_batch\", default=None)\n epochs = tu.pop(data, key=\"epochs\", default=1)\n seed = tu.pop(data, key=\"seed\", default=42)\n dataloader_kwargs = tu.pop(data, key=\"dataloader_kwargs\", default={})\n\n assert mini_batch_size is not None or num_mini_batch is not None\n\n if mini_batch_size is None:\n assert batch_size_per_dp % num_mini_batch == 0, f\"Got {batch_size_per_dp=} and {num_mini_batch=}\"\n mini_batch_size_per_gpu = batch_size_per_dp // num_mini_batch\n else:\n assert mini_batch_size % self.engine.get_data_parallel_size() == 0, (\n f\"Got {mini_batch_size=} and {self.engine.get_data_parallel_size()=}\"\n )\n mini_batch_size_per_gpu = mini_batch_size // self.engine.get_data_parallel_size()\n\n # make iterator\n dataloader = tu.make_iterator(\n data,\n mini_batch_size=mini_batch_size_per_gpu,\n epochs=epochs,\n seed=seed + self.engine.get_data_parallel_rank(),\n dataloader_kwargs=dataloader_kwargs,\n )\n\n with (\n self.engine.train_mode(disable_auto_offload=disable_auto_offload),\n Timer(name=\"train_batch\", logger=None),\n ):\n # update\n output_lst = []\n total_num_iterations = data.shape[0] // mini_batch_size_per_gpu * epochs\n\n for batch_idx, mini_batch_td in enumerate(dataloader):\n # add global token num\n global_token_num = mini_batch_td[\"input_ids\"].offsets().diff().tolist() # (total_nnz,)\n # allgather from dp rank\n global_token_num_output = [None] * self.engine.get_data_parallel_size()\n torch.distributed.all_gather_object(\n global_token_num_output, global_token_num, self.engine.get_data_parallel_group()\n )\n global_token_num = [x for xs in global_token_num_output for x in xs]\n tu.assign_non_tensor(\n mini_batch_td,\n global_token_num=NonTensorData(global_token_num),\n update_lr_scheduler=batch_idx == total_num_iterations - 1,\n disable_auto_offload=True,\n )\n actor_output = self.train_batch(mini_batch_td)\n output_lst.append(actor_output)\n\n if self.engine.is_mp_src_rank_with_outputs():\n actor_output = [tu.get(output, \"metrics\") for output in output_lst]\n metrics = {}\n for output in actor_output:\n for key, val in output.items():\n # flattn dp and micro batch\n if isinstance(val, list):\n output[key] = (\n Metric.aggregate_dp(val)\n if isinstance(val[0], Metric)\n else list(chain.from_iterable(val))\n )\n append_to_dict(metrics, output)\n\n output = tu.get_tensordict(tensor_dict={}, non_tensor_dict={\"metrics\": metrics}).cpu()\n else:\n output = None\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"train\"), blocking=False)\n def train_batch(self, data: TensorDict) -> TensorDict:\n assert self.loss_fn is not None, \"loss function can't be None when calling train_batch\"\n assert not self.engine_config.forward_only, \"Can't run `train_batch` when forward_only is in the engine config.\"\n # global_token_num should be a list of number of tokens of each seq in this batch\n global_token_num = tu.get(data, key=\"global_token_num\")\n disable_auto_offload = tu.get(data, key=\"disable_auto_offload\", default=False)\n images_seqlens = tu.get(data, key=\"images_seqlens\", default=None)\n\n # inject engineering parameters if not specified\n default_keys = dict(\n use_remove_padding=self.model_config.use_remove_padding,\n use_dynamic_bsz=self.engine_config.use_dynamic_bsz,\n max_token_len_per_gpu=self.engine_config.max_token_len_per_gpu,\n micro_batch_size_per_gpu=self.engine_config.micro_batch_size_per_gpu,\n use_fused_kernels=self.engine_config.use_fused_kernels,\n )\n\n for key, val in default_keys.items():\n if key not in data.keys():\n tu.assign_non_tensor(data, **{key: val})\n\n with (\n self.engine.train_mode(disable_auto_offload=disable_auto_offload),\n Timer(name=\"train_batch\", logger=None) as timer,\n ):\n output = self.engine.train_batch(data, loss_function=self.loss_fn)\n # containing loss, model_output and metrics\n # for training, we only care about loss and metrics\n delta_time = timer.last\n\n update_lr_scheduler = tu.get(data, key=\"update_lr_scheduler\", default=False)\n # update lr scheduler\n if update_lr_scheduler:\n lr = self.engine.lr_scheduler_step()\n else:\n lr = None\n\n if self.engine.is_mp_src_rank_with_outputs():\n # we don't need model_output in training. Maybe we change out mind later\n output.pop(\"model_output\")\n if lr is not None:\n output[\"metrics\"][\"lr\"] = lr\n final_output = self._postprocess_output(\n output,\n global_token_num=global_token_num,\n delta_time=delta_time,\n forward_only=False,\n images_seqlens=images_seqlens,\n ).cpu()\n else:\n final_output = None\n\n return final_output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"train\"), blocking=False)\n def infer_batch(self, data: TensorDict) -> TensorDict:\n # add mfu calculator\n global_token_num = tu.get(data, key=\"global_token_num\")\n compute_loss = tu.get(data, key=\"compute_loss\", default=True)\n disable_auto_offload = tu.get(data, key=\"disable_auto_offload\", default=False)\n no_lora_adapter = tu.pop(data, key=\"no_lora_adapter\", default=False)\n images_seqlens = tu.get(data, key=\"images_seqlens\", default=None)\n\n default_keys = dict(\n use_remove_padding=self.model_config.use_remove_padding,\n use_dynamic_bsz=self.engine_config.use_dynamic_bsz,\n max_token_len_per_gpu=self.engine_config.infer_max_token_len_per_gpu,\n micro_batch_size_per_gpu=self.engine_config.infer_micro_batch_size_per_gpu,\n use_fused_kernels=self.engine_config.use_fused_kernels,\n )\n\n for key, val in default_keys.items():\n if key not in data.keys():\n tu.assign_non_tensor(data, **{key: val})\n\n # for sft training, we need to compute loss in eval\n loss_function = self.loss_fn if compute_loss else None\n\n with (\n self.engine.eval_mode(disable_auto_offload=disable_auto_offload),\n Timer(name=\"eval_batch\", logger=None) as timer,\n ):\n adapter_ctx = self.engine.disable_adapter() if no_lora_adapter else nullcontext()\n with adapter_ctx:\n output = self.engine.infer_batch(data, loss_function=loss_function)\n delta_time = timer.last\n\n if self.engine.is_mp_src_rank_with_outputs():\n final_output = self._postprocess_output(\n output,\n global_token_num=global_token_num,\n delta_time=delta_time,\n forward_only=True,\n images_seqlens=images_seqlens,\n ).cpu()\n else:\n final_output = None\n\n return final_output\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def save_checkpoint(self, local_path, hdfs_path=None, global_step=0, max_ckpt_to_keep=None):\n return self.engine.save_checkpoint(local_path, hdfs_path, global_step, max_ckpt_to_keep)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def load_checkpoint(self, local_path, hdfs_path=None, del_local_after_load=False):\n return self.engine.load_checkpoint(local_path, hdfs_path, del_local_after_load)\n\n\nclass ActorRolloutRefWorker(Worker, DistProfilerExtension):\n \"\"\"Hybrid worker that includes actor model, rollout and optional ref model.\n For standalone actor or rollout, use ActorWorker or BaseRollout respectively.\n\n NOTE: ActorRolloutRefWorker no longer support spmd mode and run native server mode.\n \"\"\"\n\n def __init__(self, config: DictConfig, role: str, **kwargs):\n Worker.__init__(self)\n self.config = config\n self.role = role\n self.actor: TrainingWorker = None\n self.ref: TrainingWorker = None\n self.rollout: BaseRollout = None\n assert self.role in [\"actor\", \"rollout\", \"ref\", \"actor_rollout\", \"actor_rollout_ref\"]\n self._is_actor = self.role in [\"actor\", \"actor_rollout\", \"actor_rollout_ref\"]\n self._is_rollout = self.role in [\"rollout\", \"actor_rollout\", \"actor_rollout_ref\"]\n self._is_ref = self.role in [\"ref\", \"actor_rollout_ref\"]\n\n if self._is_actor:\n omega_profiler_config = config.actor.get(\"profiler\", {})\n elif self._is_rollout:\n # NOTE: In colocation mode, rollout config may not take effect (follow the actor config)\n # This is for extendability in AsyncRL cases\n omega_profiler_config = config.rollout.get(\"profiler\", {})\n else:\n omega_profiler_config = config.ref.get(\"profiler\", {})\n\n profiler_config = omega_conf_to_dataclass(omega_profiler_config, dataclass_type=ProfilerConfig)\n if omega_profiler_config.get(\"tool\", None) in [\"npu\", \"nsys\", \"torch\", \"torch_memory\"]:\n tool_config = omega_conf_to_dataclass(\n omega_profiler_config.get(\"tool_config\", {}).get(omega_profiler_config.get(\"tool\"))\n )\n else:\n tool_config = None\n\n self.enable_routing_replay = (\n self.config.actor.strategy == \"megatron\" and self.config.actor.megatron.router_replay.mode != \"disabled\"\n )\n\n DistProfilerExtension.__init__(\n self, DistProfiler(rank=self.rank, config=profiler_config, tool_config=tool_config)\n )\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def set_loss_fn(self, loss_fn):\n self.actor.set_loss_fn(loss_fn=loss_fn)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def to(self, device, model=True, optimizer=True, grad=True):\n \"\"\"Manual control of load/offload\"\"\"\n self.actor.to(device=device, model=model, optimizer=optimizer, grad=grad)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def init_model(self):\n model_config: HFModelConfig = omega_conf_to_dataclass(self.config.model)\n\n # 1. build reference model\n if \"ref\" in self.role:\n # TODO: align ref config with actor config\n with open_dict(self.config.ref):\n self.config.ref.ppo_mini_batch_size = self.config.actor.ppo_mini_batch_size\n self.config.ref.ppo_micro_batch_size = self.config.ref.pop(\"log_prob_micro_batch_size\", None)\n self.config.ref.ppo_micro_batch_size_per_gpu = self.config.ref.pop(\n \"log_prob_micro_batch_size_per_gpu\", None\n )\n self.config.ref.use_dynamic_bsz = self.config.ref.pop(\"log_prob_use_dynamic_bsz\", False)\n self.config.ref.ppo_max_token_len_per_gpu = self.config.ref.pop(\"log_prob_max_token_len_per_gpu\", None)\n ref_config: ActorConfig = omega_conf_to_dataclass(self.config.ref)\n ref_config.model_config = model_config\n\n # construct TrainingWorkerConfig\n ref_training_config = TrainingWorkerConfig(\n model_type=\"language_model\",\n model_config=ref_config.model_config,\n engine_config=ref_config.engine,\n optimizer_config=ref_config.optim,\n checkpoint_config=ref_config.checkpoint,\n )\n\n # assign engine configs\n ref_training_config.engine_config.use_dynamic_bsz = self.config.ref.use_dynamic_bsz\n ref_training_config.engine_config.infer_max_token_len_per_gpu = self.config.ref.ppo_max_token_len_per_gpu\n ref_training_config.engine_config.infer_micro_batch_size_per_gpu = (\n self.config.ref.ppo_micro_batch_size_per_gpu\n )\n ref_training_config.engine_config.use_remove_padding = model_config.use_remove_padding\n\n self.ref = TrainingWorker(config=ref_training_config)\n self.ref.reset()\n self.set_dispatch_collect(mesh_name=\"ref\", **self.ref.get_dispatch_collect())\n\n # 2. build actor model\n if \"actor\" in self.role:\n actor_config: ActorConfig = omega_conf_to_dataclass(self.config.actor)\n actor_config.model_config = model_config\n actor_training_config = TrainingWorkerConfig(\n model_type=\"language_model\",\n model_config=actor_config.model_config,\n engine_config=actor_config.engine,\n optimizer_config=actor_config.optim,\n checkpoint_config=actor_config.checkpoint,\n )\n\n assert self.config.actor.use_dynamic_bsz == self.config.rollout.log_prob_use_dynamic_bsz\n\n # assign engine configs\n actor_training_config.engine_config.use_dynamic_bsz = self.config.actor.use_dynamic_bsz\n actor_training_config.engine_config.infer_max_token_len_per_gpu = (\n self.config.rollout.log_prob_max_token_len_per_gpu\n )\n actor_training_config.engine_config.infer_micro_batch_size_per_gpu = (\n self.config.rollout.log_prob_micro_batch_size_per_gpu\n )\n actor_training_config.engine_config.max_token_len_per_gpu = self.config.actor.ppo_max_token_len_per_gpu\n actor_training_config.engine_config.micro_batch_size_per_gpu = (\n self.config.actor.ppo_micro_batch_size_per_gpu\n )\n actor_training_config.engine_config.use_remove_padding = model_config.use_remove_padding\n\n if self.config.actor.use_dynamic_bsz:\n assert self.config.rollout.log_prob_max_token_len_per_gpu is not None\n assert self.config.actor.ppo_max_token_len_per_gpu is not None\n else:\n assert self.config.rollout.log_prob_micro_batch_size_per_gpu is not None\n assert self.config.actor.ppo_micro_batch_size_per_gpu is not None\n\n self.loss_fn = partial(ppo_loss, config=actor_config)\n self.actor = TrainingWorker(config=actor_training_config)\n self.actor.reset()\n self.actor.set_loss_fn(self.loss_fn)\n self.set_dispatch_collect(mesh_name=\"actor\", **self.actor.get_dispatch_collect())\n\n # 3. build rollout engine\n if \"rollout\" in self.role:\n rollout_config: RolloutConfig = omega_conf_to_dataclass(self.config.rollout)\n\n # TODO: move rollout_device_mesh into ServerAdapter\n # 3.1 build rollout device mesh (sglang need only)\n infer_tp = rollout_config.tensor_model_parallel_size * rollout_config.data_parallel_size\n infer_pp = rollout_config.pipeline_model_parallel_size\n infer_world_size = infer_tp * infer_pp\n dp = self.world_size // infer_world_size\n assert self.world_size % infer_world_size == 0, (\n f\"rollout world_size: {self.world_size} is not divisible by infer_world_size: {infer_world_size}\"\n )\n rollout_device_mesh = init_device_mesh(\n get_device_name(), mesh_shape=(dp, infer_tp, infer_pp), mesh_dim_names=[\"dp\", \"infer_tp\", \"infer_pp\"]\n )\n\n # 3.2 initialize rollout engine\n rollout_cls: type[BaseRollout] = get_rollout_class(rollout_config.name, rollout_config.mode)\n self.rollout = rollout_cls(\n config=rollout_config, model_config=model_config, device_mesh=rollout_device_mesh\n )\n\n # used for LoRA\n self.base_sync_done: bool = \"dummy\" not in self.config.rollout.load_format\n self.layered_summon = self.config.rollout.get(\"layered_summon\", False)\n self.peft_merge: bool = model_config.lora.get(\"merge\", False)\n\n # 4. build checkpoint engine\n if \"actor\" in self.role:\n checkpoint_engine_config = omega_conf_to_dataclass(self.config.rollout.checkpoint_engine)\n backend = checkpoint_engine_config.backend\n bucket_size = checkpoint_engine_config.update_weights_bucket_megabytes << 20\n engine_kwargs = checkpoint_engine_config.engine_kwargs.get(backend, {})\n self.checkpoint_engine = CheckpointEngineRegistry.new(\n backend, is_master=(torch.distributed.get_rank() == 0), bucket_size=bucket_size, **engine_kwargs\n )\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"ref\"))\n @DistProfiler.annotate(color=\"olive\", role=\"ref_compute_log_prob\")\n @_with_routing_replay_flag(enabled=False)\n def compute_ref_log_prob(self, data: TensorDict) -> TensorDict:\n output = self.ref.infer_batch(data=data)\n return output.cpu() if output is not None else None\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"actor\"))\n @DistProfiler.annotate(color=\"blue\", role=\"actor_compute_log_prob\")\n @_with_routing_replay_flag(enabled=True)\n def compute_log_prob(self, data: TensorDict) -> TensorDict:\n output = self.actor.infer_batch(data)\n\n return output.cpu() if output is not None else None\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"actor\"))\n @DistProfiler.annotate(color=\"red\", role=\"actor_update\")\n @_with_routing_replay_flag(enabled=True)\n def update_actor(self, data: TensorDict) -> TensorDict:\n output = self.actor.train_mini_batch(data=data)\n return output.cpu() if output is not None else None\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def load_checkpoint(self, local_path, hdfs_path=None, del_local_after_load=False):\n assert \"actor\" in self.role, \"load_checkpoint only support actor role\"\n self.actor.load_checkpoint(local_path, hdfs_path, del_local_after_load)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def save_checkpoint(self, local_path, hdfs_path=None, global_step=0, max_ckpt_to_keep=None):\n assert \"actor\" in self.role, \"save_checkpoint only support actor role\"\n self.actor.save_checkpoint(local_path, hdfs_path, global_step, max_ckpt_to_keep)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=False)\n async def update_weights(self):\n \"\"\"Update weights from trainer to rollout.\n\n 1. For sync training with colocated trainer and rollout, update rollout directly from model engine.\n - before update_weights: rollout should be in sleep mode.\n - after update_weights: rollout should be in wake_up mode.\n 2. For async training with disaggregated trainer and rollout, send_weights only by checkpoint engine.\n \"\"\"\n assert self.checkpoint_engine is not None\n\n # 0. send_weights only for async training with disaggregated trainer and rollout\n if self.config.rollout.checkpoint_engine.backend != \"naive\":\n per_tensor_param, _ = self.engine.get_per_tensor_param()\n await self.checkpoint_engine.send_weights(per_tensor_param)\n return\n\n set_expandable_segments(False)\n log_gpu_memory_usage(\"Before resume weights\", logger=logger)\n\n # 1. resume weights and update weights\n if self.config.rollout.free_cache_engine:\n await self.rollout.resume(tags=[\"weights\"])\n log_gpu_memory_usage(\"After resume weights\", logger=logger)\n\n # 2. get per tensor generator from engine, this will load model to gpu\n per_tensor_param, peft_config = self.actor.engine.get_per_tensor_param(\n layered_summon=self.layered_summon, base_sync_done=True\n )\n\n await self.rollout.update_weights(per_tensor_param, peft_config=peft_config, base_sync_done=True)\n\n do_lora_base_sync = False\n if not self.peft_merge and peft_config is not None:\n # set sleep level for LoRA adapter weights only sync\n # TODO: make this configurable so that users with small\n # main memory can trade sync time to avoid OOM\n self.rollout.sleep_level = 1\n\n do_lora_base_sync = (not self.base_sync_done) or (\n self.rollout.sleep_level != 1 and self.config.rollout.free_cache_engine\n )\n\n if do_lora_base_sync:\n per_tensor_base_params, _ = self.actor.engine.get_per_tensor_param(\n layered_summon=self.layered_summon, base_sync_done=False\n )\n await self.rollout.update_weights(per_tensor_base_params, peft_config=peft_config, base_sync_done=False)\n\n log_gpu_memory_usage(\"After update_weights\", logger=logger)\n\n # 3. offload model to cpu\n self.actor.engine.to(\"cpu\", model=True, optimizer=False, grad=False)\n aggressive_empty_cache(force_sync=True)\n\n # 4. resume kv_cache\n if self.config.rollout.free_cache_engine:\n await self.rollout.resume(tags=[\"kv_cache\"])\n log_gpu_memory_usage(\"After resume kv_cache\", logger=logger)\n\n self.base_sync_done = True\n set_expandable_segments(True)\n\n @register(dispatch_mode=Dispatch.DP_COMPUTE, blocking=False)\n def execute_checkpoint_engine(self, method: str, *args, **kwargs):\n \"\"\"Execute checkpoint engine method.\n\n Args:\n method (str): Checkpoint engine method name.\n *args: Variable length argument list.\n **kwargs: Arbitrary keyword arguments.\n\n \"\"\"\n return getattr(self.checkpoint_engine, method)(*args, **kwargs)\n"}149{"file_name": "verl__workers__fsdp_workers.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nThe main entry point to run the PPO algorithm\n\"\"\"\n\nimport datetime\nimport json\nimport logging\nimport os\nimport warnings\nfrom dataclasses import asdict\n\nimport psutil\nimport torch\nimport torch.distributed\nimport torch.distributed as dist\nfrom codetiming import Timer\nfrom omegaconf import DictConfig, OmegaConf, open_dict\nfrom omegaconf.errors import ConfigAttributeError\nfrom peft import LoraConfig, TaskType, get_peft_model\nfrom safetensors.torch import save_file\nfrom torch.distributed.device_mesh import init_device_mesh\nfrom torch.distributed.fsdp import FullyShardedDataParallel as FSDP\nfrom torch.distributed.fsdp.api import FullStateDictConfig, ShardedStateDictConfig, StateDictType\n\ntry:\n # for torch 2.5+\n from torch.distributed.tensor import DTensor\nexcept ImportError:\n from torch.distributed._tensor import DTensor\n\nfrom verl import DataProto\nfrom verl.models.transformers.monkey_patch import apply_monkey_patch\nfrom verl.single_controller.base import Worker\nfrom verl.single_controller.base.decorator import Dispatch, make_nd_compute_dataproto_dispatch_fn, register\nfrom verl.utils import hf_processor, hf_tokenizer\nfrom verl.utils.activation_offload import enable_activation_offloading\nfrom verl.utils.checkpoint.fsdp_checkpoint_manager import FSDPCheckpointManager\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.utils.device import (\n get_device_id,\n get_device_name,\n get_nccl_backend,\n get_torch_device,\n set_expandable_segments,\n)\nfrom verl.utils.flops_counter import FlopsCounter\nfrom verl.utils.fs import copy_to_local\nfrom verl.utils.fsdp_utils import (\n CPUOffloadPolicy,\n MixedPrecisionPolicy,\n apply_fsdp2,\n collect_lora_params,\n fsdp2_load_full_state_dict,\n fsdp_version,\n get_fsdp_wrap_policy,\n get_init_weight_context_manager,\n get_shard_placement_fn,\n init_fn,\n layered_summon_lora_params,\n load_fsdp_model_to_gpu,\n load_fsdp_optimizer,\n offload_fsdp_model_to_cpu,\n offload_fsdp_optimizer,\n replace_lora_wrapper,\n)\nfrom verl.utils.import_utils import import_external_libs\nfrom verl.utils.memory_utils import aggressive_empty_cache\nfrom verl.utils.model import convert_weight_keys\nfrom verl.utils.profiler import DistProfiler, DistProfilerExtension, ProfilerConfig, log_gpu_memory_usage, simple_timer\nfrom verl.utils.profiler.performance import reduce_timing, topk_reduce_ratio_min_max\nfrom verl.utils.py_functional import convert_to_regular_types\n\n# QAT support\nfrom verl.utils.qat import apply_qat, enable_qat_fuse\nfrom verl.utils.ray_utils import get_event_loop\nfrom verl.workers.config import FSDPCriticConfig, FSDPEngineConfig, HFModelConfig, RolloutConfig\nfrom verl.workers.config.optimizer import build_optimizer\nfrom verl.workers.rollout import get_rollout_class\nfrom verl.workers.sharding_manager.fsdp_ulysses import FSDPUlyssesShardingManager\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\ndevice_name = get_device_name()\n\n\ndef create_device_mesh(world_size, fsdp_size):\n if fsdp_size < 0 or fsdp_size >= world_size:\n device_mesh = init_device_mesh(device_name, mesh_shape=(world_size,), mesh_dim_names=[\"fsdp\"])\n else:\n device_mesh = init_device_mesh(\n device_name, mesh_shape=(world_size // fsdp_size, fsdp_size), mesh_dim_names=[\"ddp\", \"fsdp\"]\n )\n return device_mesh\n\n\ndef get_sharding_strategy(device_mesh, zero3_enable=True):\n from torch.distributed.fsdp import ShardingStrategy\n\n if zero3_enable:\n fsdp_strategy = ShardingStrategy.FULL_SHARD\n hsdp_strategy = ShardingStrategy.HYBRID_SHARD\n else:\n fsdp_strategy = ShardingStrategy.SHARD_GRAD_OP\n hsdp_strategy = ShardingStrategy._HYBRID_SHARD_ZERO2\n\n if device_mesh.ndim == 1:\n sharding_strategy = fsdp_strategy\n elif device_mesh.ndim == 2:\n sharding_strategy = hsdp_strategy\n else:\n raise NotImplementedError(f\"Get device mesh ndim={device_mesh.ndim}, but only support 1 or 2\")\n return sharding_strategy\n\n\ndef get_vl_model_vision_tower(vl_model_instance):\n \"\"\"\n Util to extract Vision Tower from a VL model instance\n \"\"\"\n if hasattr(vl_model_instance, \"model\") and hasattr(vl_model_instance.model, \"visual\"):\n # transformers >= 4.52.0\n return vl_model_instance.model.visual\n elif hasattr(vl_model_instance, \"visual\"):\n # transformers < 4.52.0\n return vl_model_instance.visual\n return None\n\n\nclass ActorRolloutRefWorker(Worker, DistProfilerExtension):\n \"\"\"\n This worker can be instantiated as a standalone actor or a standalone rollout or a standalone reference policy\n or a hybrid engine based on the config.rollout\n \"\"\"\n\n def __init__(self, config: DictConfig, role: str, **kwargs):\n Worker.__init__(self)\n\n self.config = config\n import torch.distributed\n\n if not torch.distributed.is_initialized():\n rank = int(os.environ.get(\"RANK\", 0))\n world_size = int(os.environ.get(\"WORLD_SIZE\", 1))\n torch.distributed.init_process_group(\n backend=f\"cpu:gloo,{get_device_name()}:{get_nccl_backend()}\",\n rank=rank,\n world_size=world_size,\n timeout=datetime.timedelta(seconds=self.config.get(\"nccl_timeout\", 600)),\n init_method=os.environ.get(\"DIST_INIT_METHOD\", None),\n )\n\n # Apply NPU patches for FSDP backend\n from verl.workers.engine.fsdp.utils import apply_npu_fsdp_patches\n\n apply_npu_fsdp_patches()\n\n # build device mesh for FSDP\n world_size = torch.distributed.get_world_size()\n # TODO(sgm): support FSDP hybrid shard for larger model\n self.device_mesh = create_device_mesh(world_size=world_size, fsdp_size=self.config.actor.fsdp_config.fsdp_size)\n\n # build device mesh for Ulysses Sequence Parallel\n self.ulysses_device_mesh = None\n self.ulysses_sequence_parallel_size = self.config.actor.get(\"ulysses_sequence_parallel_size\", 1)\n dp = world_size // self.ulysses_sequence_parallel_size\n if self.ulysses_sequence_parallel_size > 1:\n self.ulysses_device_mesh = init_device_mesh(\n device_name, mesh_shape=(dp, self.ulysses_sequence_parallel_size), mesh_dim_names=[\"dp\", \"sp\"]\n )\n\n # create training dispatch\n if self.ulysses_device_mesh is not None:\n is_collect = self.ulysses_device_mesh[\"sp\"].get_local_rank() == 0\n self._register_dispatch_collect_info(\n \"actor\", dp_rank=self.ulysses_device_mesh[\"dp\"].get_local_rank(), is_collect=is_collect\n )\n else:\n self._register_dispatch_collect_info(\"actor\", dp_rank=self.rank, is_collect=True)\n\n self.ulysses_sharding_manager = FSDPUlyssesShardingManager(self.ulysses_device_mesh)\n self._lora_rank = self.config.model.get(\"lora_rank\", 0)\n self._is_lora = self.config.model.get(\"lora_adapter_path\") is not None or self._lora_rank > 0\n\n self.role = role\n assert self.role in [\"actor\", \"rollout\", \"ref\", \"actor_rollout\", \"actor_rollout_ref\"]\n\n self._is_actor = self.role in [\"actor\", \"actor_rollout\", \"actor_rollout_ref\"]\n self._is_rollout = self.role in [\"rollout\", \"actor_rollout\", \"actor_rollout_ref\"]\n self._is_ref = self.role in [\"ref\", \"actor_rollout_ref\"]\n self.use_orig_params = self.config.actor.fsdp_config.get(\"use_orig_params\", False)\n\n # TODO(haibin.lin):\n # As of now the type of config is DictConfig, if we assign config.profiler with ProfilerConfig,\n # it will actually convert the ProfilerConfig dataclass back to a DictConfig.\n # We can still use ProfilerConfig for testing purpose (tests/utils/test_nvtx_profile.py)\n # as they provides DictConfig-like interface\n # The benefit of creating the dataclass config is to perform validation during __post_init__\n if self._is_actor:\n omega_profiler_config = config.actor.get(\"profiler\", {})\n elif self._is_rollout:\n # NOTE: In colocation mode, rollout config may not take effect (follow the actor config)\n # This is for extendability in AsyncRL cases\n omega_profiler_config = config.rollout.get(\"profiler\", {})\n elif self._is_ref:\n omega_profiler_config = config.ref.get(\"profiler\", {})\n else:\n raise ValueError(\n f\"Invalid role {self.role}, should be one of \"\n \"['actor', 'rollout', 'ref', 'actor_rollout', 'actor_rollout_ref']\"\n )\n # omega_profiler_config is DictConfig\n # profiler_config is a ProfilerConfig dataclass\n profiler_config = omega_conf_to_dataclass(omega_profiler_config, dataclass_type=ProfilerConfig)\n if omega_profiler_config.get(\"tool\", None) in [\"npu\", \"nsys\", \"torch\", \"torch_memory\"]:\n tool_config = omega_conf_to_dataclass(\n omega_profiler_config.get(\"tool_config\", {}).get(omega_profiler_config.get(\"tool\"))\n )\n else:\n tool_config = None\n DistProfilerExtension.__init__(\n self, DistProfiler(rank=self.rank, config=profiler_config, tool_config=tool_config)\n )\n\n self._is_offload_param = False\n self._is_offload_optimizer = False\n if self._is_actor:\n self._is_offload_param = self.config.actor.fsdp_config.get(\"param_offload\", False)\n self._is_offload_optimizer = self.config.actor.fsdp_config.get(\"optimizer_offload\", False)\n elif self._is_ref:\n # TODO: it seems that manual offload is slowly than FSDP offload\n self._is_offload_param = self.config.ref.fsdp_config.get(\"param_offload\", False)\n\n # normalize config\n if self._is_actor:\n self.config.actor.ppo_mini_batch_size *= self.config.rollout.n\n self.config.actor.ppo_mini_batch_size //= self.device_mesh.size() // self.ulysses_sequence_parallel_size\n assert self.config.actor.ppo_mini_batch_size > 0, (\n f\"ppo_mini_batch_size {self.config.actor.ppo_mini_batch_size} should be larger than 0 after \"\n f\"normalization\"\n )\n # micro bsz\n if self.config.actor.ppo_micro_batch_size is not None:\n self.config.actor.ppo_micro_batch_size //= (\n self.device_mesh.size() // self.ulysses_sequence_parallel_size\n )\n self.config.actor.ppo_micro_batch_size_per_gpu = self.config.actor.ppo_micro_batch_size\n\n if self.config.actor.ppo_micro_batch_size_per_gpu is not None:\n assert self.config.actor.ppo_mini_batch_size % self.config.actor.ppo_micro_batch_size_per_gpu == 0, (\n f\"normalized ppo_mini_batch_size {self.config.actor.ppo_mini_batch_size} should be divisible by \"\n f\"ppo_micro_batch_size_per_gpu {self.config.actor.ppo_micro_batch_size_per_gpu}\"\n )\n assert self.config.actor.ppo_mini_batch_size // self.config.actor.ppo_micro_batch_size_per_gpu > 0, (\n f\"normalized ppo_mini_batch_size {self.config.actor.ppo_mini_batch_size} should be larger than \"\n f\"ppo_micro_batch_size_per_gpu {self.config.actor.ppo_micro_batch_size_per_gpu}\"\n )\n\n # normalize rollout config\n if self._is_rollout and self.config.rollout.log_prob_micro_batch_size is not None:\n self.config.rollout.log_prob_micro_batch_size //= (\n self.device_mesh.size() // self.ulysses_sequence_parallel_size\n )\n self.config.rollout.log_prob_micro_batch_size_per_gpu = self.config.rollout.log_prob_micro_batch_size\n # normalize ref config\n if self._is_ref and self.config.ref.log_prob_micro_batch_size is not None:\n self.config.ref.log_prob_micro_batch_size //= self.device_mesh.size() // self.ulysses_sequence_parallel_size\n self.config.ref.log_prob_micro_batch_size_per_gpu = self.config.ref.log_prob_micro_batch_size\n\n def _init_qat_config(self):\n \"\"\"Initialize QAT configuration from actor.qat.\"\"\"\n try:\n self.qat_config = self.config.actor.qat\n self._qat_enabled = self.qat_config.enable\n if self._qat_enabled:\n logger.info(\n f\"QAT enabled: mode={self.qat_config.mode}, config_path={self.qat_config.quantization_config_path}\"\n )\n except (AttributeError, KeyError, ConfigAttributeError):\n # QAT config not provided, disable QAT\n self._qat_enabled = False\n self.qat_config = None\n\n def _restore_w4a4_input_scales(self, model, model_path):\n \"\"\"Restore input_global_scale and input_amax from checkpoint for W4A4 mode.\"\"\"\n import glob\n\n from safetensors import safe_open\n\n safetensor_files = glob.glob(f\"{model_path}/model*.safetensors\")\n loaded_count = 0\n\n for sf_path in safetensor_files:\n with safe_open(sf_path, framework=\"pt\") as f:\n for key in f.keys():\n if \"input_global_scale\" in key:\n module_path = key.replace(\".input_global_scale\", \"\")\n amax_key = f\"{module_path}.input_amax\"\n\n module = model\n for part in module_path.split(\".\"):\n module = getattr(module, part)\n\n scale_val = f.get_tensor(key)\n val = scale_val.item() if scale_val.numel() == 1 else scale_val.max().item()\n module.input_global_scale.fill_(val)\n\n amax_val = f.get_tensor(amax_key)\n amax = amax_val.item() if amax_val.numel() == 1 else amax_val.max().item()\n module.input_amax.fill_(amax)\n loaded_count += 1\n\n if self.rank == 0:\n logger.info(f\"[W4A4] Loaded {loaded_count} input scales from checkpoint\")\n\n def _build_model_optimizer(\n self,\n model_path,\n fsdp_config: FSDPEngineConfig,\n optim_config,\n override_model_config,\n use_remove_padding=False,\n use_fused_kernels=False,\n enable_gradient_checkpointing=False,\n trust_remote_code=False,\n use_liger=False,\n role=\"actor\",\n enable_activation_offload=False,\n use_prefix_grouper=False,\n use_tiled_mlp=False,\n tiled_mlp_shards=4,\n ):\n from torch.distributed.fsdp import CPUOffload, MixedPrecision\n from transformers import (\n AutoConfig,\n AutoModel,\n AutoModelForCausalLM,\n AutoModelForImageTextToText,\n AutoModelForVision2Seq,\n )\n\n from verl.utils.model import get_generation_config, print_model_size, update_model_config\n from verl.utils.torch_dtypes import PrecisionType\n\n assert role in [\"actor\", \"ref\"]\n\n # TiledMLP requires FSDP2 for correct gradient computation\n if use_tiled_mlp and self.config.actor.strategy == \"fsdp\":\n raise ValueError(\"TiledMLP requires FSDP2. Set `actor_rollout_ref.actor.strategy=fsdp2`.\")\n\n log_gpu_memory_usage(f\"Before init {role} from HF AutoModel\", logger=logger)\n local_path = model_path\n\n # note that we have to create model in fp32. Otherwise, the optimizer is in bf16, which is incorrect\n # TODO(zhangchi.usc1992): 1. support create from random initialized model. 2. Support init with FSDP directly\n self.tokenizer = hf_tokenizer(local_path, trust_remote_code=trust_remote_code)\n self.processor = hf_processor(local_path, trust_remote_code=trust_remote_code)\n\n if self.config.model.get(\"custom_chat_template\", None) is not None:\n if self.processor is not None:\n self.processor.chat_template = self.config.model.custom_chat_template\n else:\n self.tokenizer.chat_template = self.config.model.custom_chat_template\n\n torch_dtype = fsdp_config.get(\"model_dtype\", None)\n if torch_dtype is None:\n torch_dtype = torch.float32 if self._is_actor else torch.bfloat16\n else:\n torch_dtype = PrecisionType.to_dtype(torch_dtype)\n\n # override model kwargs\n attn_implementation = override_model_config.get(\"attn_implementation\", \"flash_attention_2\")\n actor_model_config = AutoConfig.from_pretrained(\n local_path, trust_remote_code=trust_remote_code, attn_implementation=attn_implementation\n )\n # TODO: VL models use VisionAttention, which directly uses flash_attention in transformers>=4.53\n # which will be patched by _ulysses_flash_attention_forward, but errorly misses position_ids\n # Maybe support Ulysses in VisionAttention in the future and remove this patch\n if self.ulysses_sequence_parallel_size > 1 and hasattr(actor_model_config, \"vision_config\"):\n actor_model_config.vision_config._attn_implementation = \"eager\"\n\n # patch for qwen2.5-vl: when using flash_attention_3, set vision tower to use flash_attention_2\n # because the vision tower does not support flash_attention_3\n if (\n getattr(actor_model_config, \"model_type\", None) == \"qwen2_5_vl\"\n and attn_implementation == \"flash_attention_3\"\n and hasattr(actor_model_config, \"vision_config\")\n ):\n actor_model_config.vision_config._attn_implementation = \"flash_attention_2\"\n\n # patch for kimi-vl\n if getattr(actor_model_config, \"model_type\", None) == \"kimi_vl\":\n actor_model_config.text_config.topk_method = \"greedy\"\n\n self.generation_config = get_generation_config(local_path, trust_remote_code=trust_remote_code)\n\n override_config_kwargs = {\n \"bos_token_id\": self.tokenizer.bos_token_id,\n \"eos_token_id\": self.tokenizer.eos_token_id,\n \"pad_token_id\": self.tokenizer.pad_token_id,\n }\n\n if self.config.model.get(\"mtp\", {}).get(\"enable\", False):\n raise NotImplementedError(\"Right now, MTP is not supported in FSDP\")\n else:\n if hasattr(actor_model_config, \"num_nextn_predict_layers\"):\n actor_model_config.num_nextn_predict_layers = 0\n\n override_config_kwargs.update(override_model_config)\n update_model_config(actor_model_config, override_config_kwargs=override_config_kwargs)\n if self.rank == 0:\n print(f\"Model config after override: {actor_model_config}\")\n\n # NOTE(fix me): tie_word_embedding causes meta_tensor init to hang\n init_context = get_init_weight_context_manager(\n use_meta_tensor=not actor_model_config.tie_word_embeddings, mesh=self.device_mesh\n )\n\n with init_context(), warnings.catch_warnings():\n warnings.simplefilter(\"ignore\")\n has_remote_code = hasattr(actor_model_config, \"auto_map\") and any(\n actor_model_config.architectures[0] in val for val in actor_model_config.auto_map.values()\n )\n if has_remote_code:\n auto_class = next(\n k for k, v in actor_model_config.auto_map.items() if actor_model_config.architectures[0] in v\n )\n match auto_class:\n case \"AutoModelForVision2Seq\":\n actor_module_class = AutoModelForVision2Seq\n case \"AutoModelForCausalLM\":\n actor_module_class = AutoModelForCausalLM\n case \"AutoModelForImageTextToText\":\n actor_module_class = AutoModelForImageTextToText\n case _:\n actor_module_class = AutoModel\n else:\n if type(actor_model_config) in AutoModelForVision2Seq._model_mapping.keys():\n actor_module_class = AutoModelForVision2Seq\n elif type(actor_model_config) in AutoModelForCausalLM._model_mapping.keys():\n actor_module_class = AutoModelForCausalLM\n elif type(actor_model_config) in AutoModelForImageTextToText._model_mapping.keys():\n actor_module_class = AutoModelForImageTextToText\n else:\n actor_module_class = AutoModel\n\n actor_module = actor_module_class.from_pretrained(\n pretrained_model_name_or_path=local_path,\n torch_dtype=torch_dtype,\n config=actor_model_config,\n trust_remote_code=trust_remote_code,\n attn_implementation=attn_implementation,\n )\n\n # Apply Liger kernel to the model if use_liger is set to True\n if use_liger:\n from liger_kernel.transformers.monkey_patch import _apply_liger_kernel_to_instance\n\n _apply_liger_kernel_to_instance(model=actor_module)\n\n fused_kernel_options = self.config.model.get(\"fused_kernel_options\", None)\n fused_kernels_backend = (\n fused_kernel_options.get(\"impl_backend\", None) if fused_kernel_options is not None else None\n )\n\n apply_monkey_patch(\n model=actor_module,\n use_remove_padding=use_remove_padding,\n ulysses_sp_size=self.ulysses_sequence_parallel_size,\n use_fused_kernels=use_fused_kernels,\n fused_kernels_backend=fused_kernels_backend,\n use_prefix_grouper=use_prefix_grouper,\n use_tiled_mlp=use_tiled_mlp,\n tiled_mlp_shards=tiled_mlp_shards,\n )\n\n # some parameters may not in torch_dtype. TODO(zhangchi.usc1992) remove this after we switch to fsdp2\n actor_module.to(torch_dtype)\n\n if enable_gradient_checkpointing:\n actor_module.gradient_checkpointing_enable(gradient_checkpointing_kwargs={\"use_reentrant\": False})\n\n if self._is_lora:\n print(\"Applying LoRA to actor module\")\n actor_module.enable_input_require_grads()\n\n lora_adapter_path = self.config.model.get(\"lora_adapter_path\")\n if lora_adapter_path is not None:\n from peft import PeftModel\n\n print(f\"Loading pre-trained LoRA adapter to {role} from: {lora_adapter_path}\")\n\n # Copy adapter to local if needed\n local_adapter_path = copy_to_local(lora_adapter_path, use_shm=self.config.model.get(\"use_shm\", False))\n\n actor_module = PeftModel.from_pretrained(actor_module, local_adapter_path, is_trainable=True)\n peft_config = actor_module.peft_config[\"default\"]\n # Ensure task_type is TaskType enum, not string\n if isinstance(peft_config.task_type, str):\n peft_config.task_type = TaskType.CAUSAL_LM\n\n else:\n # Convert config to regular Python types before creating PEFT model\n lora_config = {\n \"task_type\": TaskType.CAUSAL_LM,\n \"r\": self.config.model.lora_rank,\n \"lora_alpha\": self.config.model.lora_alpha,\n \"target_modules\": convert_to_regular_types(self.config.model.target_modules),\n \"exclude_modules\": convert_to_regular_types(self.config.model.exclude_modules),\n \"bias\": \"none\",\n }\n actor_module = get_peft_model(actor_module, LoraConfig(**lora_config))\n\n self.use_orig_params = fsdp_config.get(\"use_orig_params\", False)\n if self.config.actor.get(\"freeze_vision_tower\", False):\n vision_tower = get_vl_model_vision_tower(actor_module)\n if vision_tower is not None:\n vision_tower.requires_grad_(False)\n self.use_orig_params = True\n if self.rank == 0:\n print(\"[actor model] Vision tower is set to not trainable.\")\n else:\n if self.rank == 0:\n print(\"[actor model] No vision tower found.\")\n\n # Apply QAT before FSDP wrapping (actor only)\n if role == \"actor\" and self._qat_enabled:\n actor_module = apply_qat(actor_module, self.qat_config)\n enable_qat_fuse(actor_module)\n if self.qat_config.mode == \"w4a4\":\n self._restore_w4a4_input_scales(actor_module, self.config.model.path)\n\n torch.distributed.barrier()\n\n if self.rank == 0:\n print_model_size(actor_module)\n\n log_gpu_memory_usage(f\"After init {role} from HF AutoModel\", logger=logger)\n\n # We wrap FSDP for rollout as well\n mixed_precision_config = fsdp_config.get(\"mixed_precision\", None)\n if mixed_precision_config is not None:\n param_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"param_dtype\", \"bf16\"))\n reduce_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"reduce_dtype\", \"fp32\"))\n buffer_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"buffer_dtype\", \"fp32\"))\n else:\n param_dtype = PrecisionType.to_dtype(fsdp_config.dtype)\n reduce_dtype = torch.float32\n buffer_dtype = torch.float32\n\n mixed_precision = MixedPrecision(param_dtype=param_dtype, reduce_dtype=reduce_dtype, buffer_dtype=buffer_dtype)\n\n # Store param_dtype for QAT quantizer\n self._param_dtype = param_dtype\n\n auto_wrap_policy = get_fsdp_wrap_policy(\n module=actor_module,\n config=fsdp_config.get(\"wrap_policy\", None),\n is_lora=self._is_lora,\n )\n\n # if self._is_rollout and self.config.rollout.name == \"hf\":\n # # TODO(zhangchi.usc1992, shengguangming) fix me.\n # Current, auto_wrap_policy causes HFRollout to hang in Gemma\n # auto_wrap_policy = None\n\n if self.rank == 0:\n print(f\"wrap_policy: {auto_wrap_policy}\")\n\n fsdp_mesh = self.device_mesh\n fsdp_enable_zero3 = fsdp_config.reshard_after_forward\n sharding_strategy = get_sharding_strategy(fsdp_mesh, fsdp_enable_zero3)\n\n # TODO: add transformer policy\n # We force reference policy to use CPUOffload to save memory.\n # We force turn off CPUOffload for actor because it causes incorrect results when using grad accumulation\n cpu_offload = None if role == \"actor\" else CPUOffload(offload_params=True)\n fsdp_strategy = self.config.actor.strategy\n if fsdp_strategy == \"fsdp\":\n actor_module_fsdp = FSDP(\n actor_module,\n cpu_offload=cpu_offload,\n param_init_fn=init_fn,\n auto_wrap_policy=auto_wrap_policy,\n device_id=get_device_id(),\n sharding_strategy=sharding_strategy, # zero3\n mixed_precision=mixed_precision,\n sync_module_states=True,\n device_mesh=self.device_mesh,\n use_orig_params=self.use_orig_params,\n forward_prefetch=fsdp_config.get(\"forward_prefetch\", False),\n )\n elif fsdp_strategy == \"fsdp2\":\n assert CPUOffloadPolicy is not None, \"PyTorch version >= 2.4 is required for using fully_shard API (FSDP2)\"\n mp_policy = MixedPrecisionPolicy(\n param_dtype=param_dtype, reduce_dtype=reduce_dtype, cast_forward_inputs=True\n )\n if role == \"actor\" and fsdp_config.offload_policy:\n cpu_offload = CPUOffloadPolicy(pin_memory=True)\n self._is_offload_param = False\n self._is_offload_optimizer = False\n else:\n cpu_offload = None if role == \"actor\" else CPUOffloadPolicy(pin_memory=True)\n\n fsdp_kwargs = {\n \"mesh\": fsdp_mesh,\n \"mp_policy\": mp_policy,\n \"offload_policy\": cpu_offload,\n \"reshard_after_forward\": fsdp_config.reshard_after_forward,\n \"shard_placement_fn\": get_shard_placement_fn(fsdp_size=self.device_mesh.shape[-1]),\n }\n full_state = actor_module.state_dict()\n apply_fsdp2(actor_module, fsdp_kwargs, fsdp_config)\n fsdp2_load_full_state_dict(actor_module, full_state, fsdp_mesh, cpu_offload)\n actor_module_fsdp = actor_module\n else:\n raise NotImplementedError(f\"not implement {fsdp_strategy}\")\n\n if enable_activation_offload:\n enable_activation_offloading(actor_module_fsdp, fsdp_strategy, enable_gradient_checkpointing)\n\n log_gpu_memory_usage(f\"After {role} FSDP init\", logger=logger)\n\n # TODO: add more optimizer args into config\n if role == \"actor\" and optim_config is not None:\n from verl.utils.torch_functional import get_constant_schedule_with_warmup, get_cosine_schedule_with_warmup\n\n actor_optimizer = build_optimizer(actor_module_fsdp.parameters(), optim_config)\n\n total_steps = optim_config.get(\"total_training_steps\", 0)\n num_warmup_steps = int(optim_config.get(\"lr_warmup_steps\", -1))\n lr_scheduler_type = optim_config.get(\"lr_scheduler_type\", \"constant\")\n min_lr_ratio = optim_config.get(\"min_lr_ratio\", 0.0)\n num_cycles = optim_config.get(\"num_cycles\", 0.5)\n if num_warmup_steps < 0:\n num_warmup_steps_ratio = optim_config.get(\"lr_warmup_steps_ratio\", 0.0)\n num_warmup_steps = int(num_warmup_steps_ratio * total_steps)\n\n if self.rank == 0:\n print(f\"Total steps: {total_steps}, num_warmup_steps: {num_warmup_steps}\")\n\n if lr_scheduler_type == \"constant\":\n actor_lr_scheduler = get_constant_schedule_with_warmup(\n optimizer=actor_optimizer, num_warmup_steps=num_warmup_steps\n )\n elif lr_scheduler_type == \"cosine\":\n actor_lr_scheduler = get_cosine_schedule_with_warmup(\n optimizer=actor_optimizer,\n num_warmup_steps=num_warmup_steps,\n num_training_steps=total_steps,\n min_lr_ratio=min_lr_ratio,\n num_cycles=num_cycles,\n )\n else:\n raise NotImplementedError(f\"LR scheduler type {lr_scheduler_type} is not supported\")\n\n log_gpu_memory_usage(f\"After {role} optimizer init\", logger=logger)\n else:\n actor_optimizer = None\n actor_lr_scheduler = None\n\n return actor_module_fsdp, actor_optimizer, actor_lr_scheduler, actor_model_config\n\n def _build_rollout(self, trust_remote_code=False):\n from torch.distributed.device_mesh import init_device_mesh\n\n # 1. parse rollout and huggingface model config\n rollout_config: RolloutConfig = omega_conf_to_dataclass(self.config.rollout)\n model_config: HFModelConfig = omega_conf_to_dataclass(self.config.model, dataclass_type=HFModelConfig)\n self.model_config = model_config\n\n # 2. build rollout device mesh\n infer_tp = self.config.rollout.tensor_model_parallel_size * self.config.rollout.data_parallel_size\n infer_pp = self.config.rollout.pipeline_model_parallel_size\n infer_world_size = infer_tp * infer_pp\n dp = self.world_size // infer_world_size\n assert self.world_size % infer_world_size == 0, (\n f\"rollout world_size: {self.world_size} is not divisible by infer_world_size: {infer_world_size}\"\n )\n rollout_device_mesh = init_device_mesh(\n device_name, mesh_shape=(dp, infer_tp, infer_pp), mesh_dim_names=[\"dp\", \"infer_tp\", \"infer_pp\"]\n )\n rollout_name = self.config.rollout.name\n\n self.rollout_device_mesh = rollout_device_mesh\n\n if rollout_name == \"hf\":\n self._register_dispatch_collect_info(\"rollout\", dp_rank=self.rank, is_collect=True)\n else:\n is_collect = (\n rollout_device_mesh[\"infer_tp\"].get_local_rank() == 0\n and rollout_device_mesh[\"infer_pp\"].get_local_rank() == 0\n )\n self._register_dispatch_collect_info(\n \"rollout\", dp_rank=rollout_device_mesh[\"dp\"].get_local_rank(), is_collect=is_collect\n )\n\n # 4. build rollout model\n log_gpu_memory_usage(f\"Before building {self.config.rollout.name} rollout\", logger=logger)\n self.rollout = get_rollout_class(rollout_config.name, rollout_config.mode)(\n config=rollout_config, model_config=model_config, device_mesh=rollout_device_mesh\n )\n log_gpu_memory_usage(f\"After building {self.config.rollout.name} rollout\", logger=logger)\n\n # Full params\n if torch.distributed.get_world_size() == 1 and fsdp_version(self.actor_module_fsdp) == 1:\n FSDP.set_state_dict_type(\n self.actor_module_fsdp,\n state_dict_type=StateDictType.FULL_STATE_DICT,\n state_dict_config=FullStateDictConfig(),\n )\n elif fsdp_version(self.actor_module_fsdp) == 1:\n FSDP.set_state_dict_type(\n self.actor_module_fsdp,\n state_dict_type=StateDictType.SHARDED_STATE_DICT,\n state_dict_config=ShardedStateDictConfig(),\n )\n\n # used for LoRA\n self.base_sync_done: bool = \"dummy\" not in self.config.rollout.load_format\n self.layered_summon = self.config.rollout.get(\"layered_summon\", False)\n\n # 5. switch to trainer mode\n # NOTE: It's critical that hybrid engine in trainer mode initially to load checkpoint.\n # For async mode, we can't call run_until_complete here, so we will switch to trainer mode in AgentLoopManager.\n # Note: sync mode is deprecated and rejected in RolloutConfig.__post_init__\n\n async def rollout_mode(self):\n \"\"\"Context switch hybridengine to rollout mode.\"\"\"\n aggressive_empty_cache(force_sync=True)\n\n log_gpu_memory_usage(\"Before load_fsdp_model_to_gpu\", logger=logger)\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.actor_module_fsdp)\n log_gpu_memory_usage(\"After load_fsdp_model_to_gpu\", logger=logger)\n\n peft_config = None\n peft_model = getattr(self.actor_module_fsdp, \"_fsdp_wrapped_module\", self.actor_module_fsdp)\n if hasattr(peft_model, \"peft_config\"): # LoRA\n peft_config = peft_model.peft_config.get(\"default\", None)\n params = collect_lora_params(\n module=self.actor_module_fsdp,\n layered_summon=self.config.rollout.get(\"layered_summon\", False),\n base_sync_done=self.base_sync_done,\n )\n if not self.base_sync_done:\n params = {replace_lora_wrapper(k, peft_config): v for k, v in params.items()}\n else:\n params = self.actor_module_fsdp.state_dict()\n\n params = convert_weight_keys(\n params, getattr(self.actor_module_fsdp, \"_fsdp_wrapped_module\", self.actor_module_fsdp)\n )\n\n # Special handling for LoRA with sleep_level=2:\n # When sleep_level=2, base model weights are destroyed during each sleep cycle.\n # separately collect and update LoRA weights and base model weights through their respective interfaces.\n # Here: params contains LoRA weights, base_model_params contains base model weights.\n # Only needed if the rollout engine actually sleeps/frees weights (free_cache_engine=True).\n if (\n peft_config is not None\n and getattr(self.rollout, \"sleep_level\", None) == 2\n and self.config.rollout.free_cache_engine\n ):\n base_model_params = collect_lora_params(\n module=self.actor_module_fsdp,\n layered_summon=self.layered_summon,\n base_sync_done=False,\n )\n base_model_params = {replace_lora_wrapper(k, peft_config): v for k, v in base_model_params.items()}\n base_model_params = convert_weight_keys(\n base_model_params, getattr(self.actor_module_fsdp, \"_fsdp_wrapped_module\", self.actor_module_fsdp)\n )\n\n log_gpu_memory_usage(\"Before offload_fsdp_model_to_cpu\", logger=logger)\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.actor_module_fsdp)\n log_gpu_memory_usage(\"After offload_fsdp_model_to_cpu\", logger=logger)\n\n set_expandable_segments(False)\n\n if peft_config is not None and self.base_sync_done:\n per_tensor_param = params.items() if isinstance(params, dict) else params # Fixed: handle dict case\n else:\n device = get_device_id() # used when fsdp2 set cpu_offload_policy\n per_tensor_param = (\n (name, param.to(device, non_blocking=True).full_tensor() if isinstance(param, DTensor) else param)\n for name, param in params.items()\n )\n\n # QAT: quantize weights before sending to vLLM\n if self._qat_enabled:\n from verl.utils.qat.quantizer import QATQuantizer\n\n quantizer = QATQuantizer(\n mode=self.qat_config.mode,\n group_size=self.qat_config.group_size,\n ignore_patterns=self.qat_config.ignore_patterns,\n device=torch.device(get_device_id()),\n param_dtype=self._param_dtype,\n )\n per_tensor_param = quantizer.quantize_with_fusion(\n per_tensor_param,\n target_device=torch.device(\"cpu\"),\n )\n aggressive_empty_cache(force_sync=True)\n\n if self.config.rollout.free_cache_engine:\n await self.rollout.resume(tags=[\"weights\"])\n log_gpu_memory_usage(\"After resume weights\", logger=logger)\n\n if (\n peft_config is not None\n and getattr(self.rollout, \"sleep_level\", None) == 2\n and self.config.rollout.free_cache_engine\n ):\n per_tensor_base_params = (\n (name, param.to(device, non_blocking=True).full_tensor() if isinstance(param, DTensor) else param)\n for name, param in base_model_params.items()\n )\n await self.rollout.update_weights(per_tensor_base_params, base_sync_done=False)\n del base_model_params, per_tensor_base_params\n\n await self.rollout.update_weights(per_tensor_param, peft_config=peft_config, base_sync_done=self.base_sync_done)\n log_gpu_memory_usage(\"After update_weights\", logger=logger)\n del params, per_tensor_param\n aggressive_empty_cache(force_sync=True)\n if self.config.rollout.free_cache_engine:\n await self.rollout.resume(tags=[\"kv_cache\"])\n log_gpu_memory_usage(\"After resume kv_cache\", logger=logger)\n\n self.base_sync_done = True\n set_expandable_segments(True)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def init_model(self):\n from verl.workers.actor import DataParallelPPOActor\n\n # This is used to import external_lib into the huggingface systems\n import_external_libs(self.config.model.get(\"external_lib\", None))\n\n # Initialize QAT config before _build_model_optimizer\n self._init_qat_config()\n\n override_model_config = OmegaConf.to_container(OmegaConf.create(self.config.model.get(\"override_config\", {})))\n use_remove_padding = self.config.model.get(\"use_remove_padding\", False)\n use_shm = self.config.model.get(\"use_shm\", False)\n use_fused_kernels = self.config.model.get(\"use_fused_kernels\", False)\n\n if self._is_actor or self._is_rollout:\n # we need the model for actor and rollout\n if self._is_actor:\n optim_config = self.config.actor.optim\n fsdp_config = omega_conf_to_dataclass(self.config.actor.fsdp_config)\n else:\n optim_config = None\n fsdp_config = FSDPEngineConfig()\n\n local_path = copy_to_local(self.config.model.path, use_shm=use_shm)\n # TiledMLP configuration for memory-efficient MLP computation\n tiled_mlp_config = self.config.model.get(\"tiled_mlp\", {})\n use_tiled_mlp = tiled_mlp_config.get(\"enabled\", False)\n tiled_mlp_shards = tiled_mlp_config.get(\"num_shards\", 4)\n\n (\n self.actor_module_fsdp,\n self.actor_optimizer,\n self.actor_lr_scheduler,\n self.actor_model_config,\n ) = self._build_model_optimizer(\n model_path=local_path,\n fsdp_config=fsdp_config,\n optim_config=optim_config,\n override_model_config=override_model_config,\n use_remove_padding=use_remove_padding,\n use_fused_kernels=use_fused_kernels,\n enable_gradient_checkpointing=self.config.model.get(\"enable_gradient_checkpointing\", False),\n trust_remote_code=self.config.model.get(\"trust_remote_code\", False),\n use_liger=self.config.model.get(\"use_liger\", False),\n role=\"actor\",\n enable_activation_offload=self.config.model.get(\"enable_activation_offload\", False),\n use_prefix_grouper=self.config.actor.get(\"use_prefix_grouper\", False),\n use_tiled_mlp=use_tiled_mlp,\n tiled_mlp_shards=tiled_mlp_shards,\n )\n\n # get the original unwrapped module\n if fsdp_version(self.actor_module_fsdp) == 1:\n self.actor_module = self.actor_module_fsdp._fsdp_wrapped_module\n\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.actor_module_fsdp)\n log_gpu_memory_usage(\"After offload actor model during init\", logger=logger)\n\n if self._is_offload_optimizer:\n offload_fsdp_optimizer(optimizer=self.actor_optimizer)\n log_gpu_memory_usage(\"After offload actor optimizer during init\", logger=logger)\n\n if self._is_actor:\n actor_cfg = omega_conf_to_dataclass(self.config.actor)\n self.actor = DataParallelPPOActor(\n config=actor_cfg, actor_module=self.actor_module_fsdp, actor_optimizer=self.actor_optimizer\n )\n\n if self._is_rollout:\n self._build_rollout(trust_remote_code=self.config.model.get(\"trust_remote_code\", False))\n\n if self._is_ref:\n ref_model_path = self.config.model.path\n ref_model = self.config.ref.get(\"model\", None)\n if ref_model is not None:\n ref_model_path = ref_model.get(\"path\", self.config.model.path)\n\n if self.rank == 0:\n print(\"reference model:\", ref_model_path)\n local_path = copy_to_local(ref_model_path, use_shm=use_shm)\n use_prefix_grouper = hasattr(self.config, \"actor\") and self.config.actor.get(\"use_prefix_grouper\", False)\n\n # TiledMLP for ref model: use ref config if specified, otherwise use actor config\n ref_tiled_mlp_config = self.config.ref.get(\"tiled_mlp\", None)\n if ref_tiled_mlp_config is None:\n ref_tiled_mlp_config = self.config.model.get(\"tiled_mlp\", {})\n ref_use_tiled_mlp = ref_tiled_mlp_config.get(\"enabled\", False)\n ref_tiled_mlp_shards = ref_tiled_mlp_config.get(\"num_shards\", 4)\n\n self.ref_module_fsdp = self._build_model_optimizer(\n model_path=local_path,\n fsdp_config=omega_conf_to_dataclass(self.config.ref.fsdp_config),\n optim_config=None,\n override_model_config=override_model_config,\n use_remove_padding=use_remove_padding,\n use_fused_kernels=use_fused_kernels,\n trust_remote_code=self.config.model.get(\"trust_remote_code\", False),\n use_liger=self.config.model.get(\"use_liger\", False),\n role=\"ref\",\n use_prefix_grouper=use_prefix_grouper,\n use_tiled_mlp=ref_use_tiled_mlp,\n tiled_mlp_shards=ref_tiled_mlp_shards,\n )[0]\n OmegaConf.set_struct(self.config.ref, True)\n with open_dict(self.config.ref):\n self.config.ref.use_remove_padding = use_remove_padding\n self.config.ref.use_fused_kernels = use_fused_kernels\n if use_prefix_grouper:\n self.config.ref.use_prefix_grouper = use_prefix_grouper\n self.ref_policy = DataParallelPPOActor(config=self.config.ref, actor_module=self.ref_module_fsdp)\n\n if self._is_actor:\n self.flops_counter = FlopsCounter(self.actor_model_config)\n self.checkpoint_manager = FSDPCheckpointManager(\n model=self.actor_module_fsdp,\n optimizer=self.actor.actor_optimizer,\n lr_scheduler=self.actor_lr_scheduler,\n processing_class=self.processor if self.processor is not None else self.tokenizer,\n checkpoint_config=self.config.actor.checkpoint,\n trust_remote_code=self.config.model.get(\"trust_remote_code\", False),\n )\n\n if not self._is_actor and self._is_rollout:\n # If ActorRolloutRefWorker is initialized as a standalone rollout,\n # create a checkpoint manager for FSDP model to allow loading FSDP checkpoints for rollout.\n\n checkpoint_contents = OmegaConf.create({\"load_contents\": [\"model\"], \"save_contents\": []})\n self.checkpoint_manager = FSDPCheckpointManager(\n model=self.actor_module_fsdp,\n optimizer=None,\n lr_scheduler=None,\n processing_class=self.processor if self.processor is not None else self.tokenizer,\n checkpoint_config=checkpoint_contents,\n )\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"actor\"))\n @DistProfiler.annotate(color=\"red\", role=\"actor_update\")\n def update_actor(self, data: DataProto):\n assert self._is_actor\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.actor_module_fsdp)\n if self._is_offload_optimizer:\n load_fsdp_optimizer(optimizer=self.actor_optimizer, device_id=get_device_id())\n\n with self.ulysses_sharding_manager:\n data = data.to(\"cpu\") # data will to device with each micro batch on actor.update_policy\n data.meta_info.setdefault(\"pad_token_id\", self.tokenizer.pad_token_id)\n # perform training\n with Timer(name=\"update_policy\", logger=None) as timer:\n metrics = self.actor.update_policy(data=data)\n delta_time = timer.last\n global_num_tokens = data.meta_info[\"global_token_num\"]\n images_seqlens = data.meta_info.get(\"images_seqlens\", None)\n estimated_flops, promised_flops = self.flops_counter.estimate_flops(\n global_num_tokens, delta_time, images_seqlens=images_seqlens\n )\n metrics[\"perf/mfu/actor\"] = (\n estimated_flops * self.config.actor.ppo_epochs / promised_flops / self.world_size\n )\n metrics[\"perf/max_memory_allocated_gb\"] = get_torch_device().max_memory_allocated() / (1024**3)\n metrics[\"perf/max_memory_reserved_gb\"] = get_torch_device().max_memory_reserved() / (1024**3)\n metrics[\"perf/cpu_memory_used_gb\"] = psutil.virtual_memory().used / (1024**3)\n\n lr = self.actor_lr_scheduler.get_last_lr()[0]\n metrics[\"actor/lr\"] = lr.item() if torch.is_tensor(lr) else lr\n self.actor_lr_scheduler.step()\n\n # TODO: here, we should return all metrics\n output = DataProto(meta_info={\"metrics\": metrics})\n\n output = output.to(\"cpu\")\n\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.actor_module_fsdp)\n log_gpu_memory_usage(\"After offload actor model during update_actor\", logger=logger)\n if self._is_offload_optimizer:\n offload_fsdp_optimizer(optimizer=self.actor_optimizer)\n log_gpu_memory_usage(\"After offload actor optimizer during update_actor\", logger=logger)\n\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"rollout\"))\n @DistProfiler.annotate(color=\"red\", role=\"rollout_generate\")\n def generate_sequences(self, prompts: DataProto):\n # Support all hardwares\n assert self._is_rollout\n prompts = prompts.to(get_device_id())\n\n meta_info = {\n \"eos_token_id\": self.generation_config.eos_token_id\n if self.generation_config is not None\n else self.tokenizer.eos_token_id,\n \"pad_token_id\": self.generation_config.pad_token_id\n if self.generation_config is not None\n else self.tokenizer.pad_token_id,\n }\n prompts.meta_info.update(meta_info)\n\n timing_generate = {}\n if self._is_actor: # For rollout only, we do not switch context.\n loop = get_event_loop()\n loop.run_until_complete(self.rollout_mode())\n log_gpu_memory_usage(\"After switch to rollout mode\", logger=logger)\n\n with simple_timer(\"generate_sequences\", timing_generate):\n output = self.rollout.generate_sequences(prompts=prompts)\n\n if self._is_actor:\n loop.run_until_complete(self.trainer_mode())\n log_gpu_memory_usage(\"After switch to trainer mode\", logger=logger)\n\n # We calculate the average timing across all ranks\n # to make sure meta_info[\"timing\"] is the same\n timing_generate_topk_ratio, timing_generate_min, timing_generate_max = topk_reduce_ratio_min_max(\n timing_generate[\"generate_sequences\"]\n )\n timing_generate = reduce_timing(timing_generate)\n timing_generate.update(\n {\n \"generation_timing/max\": timing_generate_max,\n \"generation_timing/min\": timing_generate_min,\n \"generation_timing/topk_ratio\": timing_generate_topk_ratio,\n }\n )\n output.meta_info[\"timing\"] = timing_generate\n output = output.to(\"cpu\")\n\n # clear kv cache\n get_torch_device().empty_cache()\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"actor\"))\n @DistProfiler.annotate(color=\"blue\", role=\"actor_compute_log_prob\")\n def compute_log_prob(self, data: DataProto):\n # when is_lora is True, we use the actor without lora applied to calculate the log_prob\n # which is mostly used for ref log_prob calculation\n assert self._is_actor\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.actor_module_fsdp)\n\n # Support all hardwares\n from contextlib import nullcontext\n\n is_lora = data.meta_info.pop(\"is_lora\", False)\n adapter_ctx = self.actor.actor_module.disable_adapter() if is_lora else nullcontext()\n # we should always recompute old_log_probs when it is HybridEngine\n config_source = self.config.ref if is_lora else self.config.rollout\n data.meta_info[\"micro_batch_size\"] = config_source.log_prob_micro_batch_size_per_gpu\n data.meta_info[\"max_token_len\"] = config_source.log_prob_max_token_len_per_gpu\n data.meta_info[\"use_dynamic_bsz\"] = config_source.log_prob_use_dynamic_bsz\n data.meta_info[\"temperature\"] = self.config.rollout.temperature\n data.meta_info.setdefault(\"pad_token_id\", self.tokenizer.pad_token_id)\n # perform recompute log_prob\n calculate_entropy = not is_lora\n with self.ulysses_sharding_manager:\n with adapter_ctx:\n outputs = self.actor.compute_log_prob(data=data, calculate_entropy=calculate_entropy)\n if not is_lora:\n tensors = {\"old_log_probs\": outputs[\"log_probs\"]}\n else:\n tensors = {\"ref_log_prob\": outputs[\"log_probs\"]}\n if calculate_entropy:\n tensors[\"entropys\"] = outputs[\"entropys\"]\n if \"sum_pi_squared\" in outputs:\n tensors[\"sum_pi_squared\"] = outputs[\"sum_pi_squared\"]\n output = DataProto.from_dict(\n tensors=tensors,\n meta_info={\"temperature\": self.config.rollout.temperature},\n )\n\n output = output.to(\"cpu\")\n\n # https://pytorch.org/docs/stable/notes/fsdp.html#fsdp-notes\n # unshard the root FSDP module\n if self.world_size > 1 and fsdp_version(self.actor.actor_module) == 1:\n self.actor.actor_module._handle.reshard(True)\n\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.actor_module_fsdp)\n log_gpu_memory_usage(\"After offload actor model during compute_log_prob\", logger=logger)\n\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"actor\"))\n @DistProfiler.annotate(color=\"olive\", role=\"ref_compute_log_prob\")\n def compute_ref_log_prob(self, data: DataProto):\n if self._is_lora:\n # if _is_lora, actor without lora applied is the ref\n data.meta_info[\"is_lora\"] = True\n return self.compute_log_prob(data)\n assert self._is_ref\n # else:\n # otherwise, the class have a standalone ref model\n\n micro_batch_size = self.config.ref.log_prob_micro_batch_size_per_gpu\n data.meta_info[\"micro_batch_size\"] = micro_batch_size\n data.meta_info[\"temperature\"] = self.config.rollout.temperature\n data.meta_info[\"max_token_len\"] = self.config.ref.log_prob_max_token_len_per_gpu\n data.meta_info[\"use_dynamic_bsz\"] = self.config.ref.log_prob_use_dynamic_bsz\n data.meta_info.setdefault(\"pad_token_id\", self.tokenizer.pad_token_id)\n with self.ulysses_sharding_manager:\n data = data.to(\"cpu\") # data will to device with each micro batch on ref.compute_log_prob\n outputs = self.ref_policy.compute_log_prob(data=data, calculate_entropy=False)\n output = DataProto.from_dict(tensors={\"ref_log_prob\": outputs[\"log_probs\"]})\n\n output = output.to(\"cpu\")\n\n # https://pytorch.org/docs/stable/notes/fsdp.html#fsdp-notes\n # unshard the root FSDP module\n if self.world_size > 1:\n if fsdp_version(self.ref_policy.actor_module) == 1:\n self.ref_policy.actor_module._handle.reshard(True)\n elif fsdp_version(self.ref_policy.actor_module) == 2:\n self.ref_policy.actor_module.reshard()\n\n return output\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def save_checkpoint(self, local_path, hdfs_path=None, global_step=0, max_ckpt_to_keep=None):\n from verl.utils.logger import log_with_rank\n\n # only support save and load ckpt for actor\n assert self._is_actor\n\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.actor_module_fsdp)\n\n self.checkpoint_manager.save_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, global_step=global_step, max_ckpt_to_keep=max_ckpt_to_keep\n )\n dist.barrier()\n\n if self._is_lora and hasattr(getattr(self, \"actor_module\", self.actor_module_fsdp), \"peft_config\"):\n lora_save_path = os.path.join(local_path, \"lora_adapter\")\n peft_model = getattr(self, \"actor_module\", self.actor_module_fsdp)\n peft_config = {}\n if dist.get_rank() == 0:\n os.makedirs(lora_save_path, exist_ok=True)\n peft_config = asdict(peft_model.peft_config.get(\"default\", {}))\n peft_config[\"task_type\"] = peft_config[\"task_type\"].value\n peft_config[\"peft_type\"] = peft_config[\"peft_type\"].value\n peft_config[\"target_modules\"] = list(peft_config[\"target_modules\"])\n try:\n if fsdp_version(self.actor_module_fsdp) > 0:\n self.actor_module_fsdp = self.actor_module_fsdp.to(get_device_name())\n lora_params = layered_summon_lora_params(self.actor_module_fsdp)\n if dist.get_rank() == 0:\n save_file(lora_params, os.path.join(lora_save_path, \"adapter_model.safetensors\"))\n with open(os.path.join(lora_save_path, \"adapter_config.json\"), \"w\", encoding=\"utf-8\") as f:\n json.dump(peft_config, f, ensure_ascii=False, indent=4)\n except Exception as e:\n log_with_rank(\n f\"Save LoRA Adapter Error ({e})\", rank=dist.get_rank(), logger=logger, log_only_rank_0=True\n )\n\n dist.barrier()\n log_with_rank(\n f\"[rank-{self.rank}]: Saved LoRA adapter to: {lora_save_path}\",\n rank=dist.get_rank(),\n logger=logger,\n log_only_rank_0=True,\n )\n\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.actor_module_fsdp)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def load_checkpoint(self, local_path, hdfs_path=None, del_local_after_load=False):\n assert self._is_actor or (not self._is_actor and self._is_rollout), (\n f\"Checkpoint loading is only supported for Actor or standalone Rollout Workers, but got \"\n f\"{self._is_actor} and {self._is_rollout}\"\n )\n\n # No checkpoint to load, just offload the model and optimizer to CPU\n if local_path is None:\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.actor_module_fsdp)\n if self._is_offload_optimizer:\n offload_fsdp_optimizer(self.actor_optimizer)\n return\n\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.actor_module_fsdp)\n\n self.checkpoint_manager.load_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, del_local_after_load=del_local_after_load\n )\n\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.actor_module_fsdp)\n\n if self._is_offload_optimizer:\n offload_fsdp_optimizer(self.actor_optimizer)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def start_profile(self, **kwargs) -> None:\n \"\"\"Start profiling for the current rank in the current training step.\"\"\"\n self.profiler.start(**kwargs)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def stop_profile(self) -> None:\n \"\"\"Stop profiling for the current rank in the current training step.\"\"\"\n self.profiler.stop()\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def dump_memory_snapshot(self, tag: str = \"manual\", sub_dir: str = None) -> None:\n \"\"\"Manually trigger a CUDA memory snapshot dump on all ranks.\"\"\"\n # Memory snapshot is now handled by the profiler system\n # This method is kept for backward compatibility but delegates to profiler\n if hasattr(self, \"profiler\") and hasattr(self.profiler, \"_impl\"):\n try:\n # Try to use the profiler's memory snapshot functionality\n if hasattr(self.profiler._impl, \"sampler\"):\n out_dir = OmegaConf.select(self.config, \"actor.profiler.save_path\") or \".\"\n self.profiler._impl.sampler.dump_memory_snapshot(out_dir=out_dir, tag=tag, sub_dir=sub_dir)\n except Exception:\n # silently ignore if profiler doesn't support memory snapshots\n pass\n\n\nclass CriticWorker(Worker, DistProfilerExtension):\n def __init__(self, config: FSDPCriticConfig):\n Worker.__init__(self)\n omega_profiler_config = config.get(\"profiler\", {})\n profiler_config = omega_conf_to_dataclass(omega_profiler_config, dataclass_type=ProfilerConfig)\n if omega_profiler_config.get(\"tool\", None) in [\"npu\", \"nsys\", \"torch\", \"torch_memory\"]:\n tool_config = omega_conf_to_dataclass(\n omega_profiler_config.get(\"tool_config\", {}).get(omega_profiler_config.get(\"tool\"))\n )\n else:\n tool_config = None\n DistProfilerExtension.__init__(\n self, DistProfiler(rank=self.rank, config=profiler_config, tool_config=tool_config)\n )\n import torch.distributed\n\n self.config = config\n if not torch.distributed.is_initialized():\n torch.distributed.init_process_group(\n backend=get_nccl_backend(),\n timeout=datetime.timedelta(seconds=self.config.get(\"nccl_timeout\", 600)),\n init_method=os.environ.get(\"DIST_INIT_METHOD\", None),\n )\n self.config: FSDPCriticConfig = config\n\n # build device mesh for Ulysses Sequence Parallel\n world_size = torch.distributed.get_world_size()\n from torch.distributed.device_mesh import init_device_mesh\n\n fsdp_size = self.config.model.fsdp_config.fsdp_size\n self.device_mesh = create_device_mesh(world_size=world_size, fsdp_size=fsdp_size)\n\n self.ulysses_device_mesh = None\n self.ulysses_sequence_parallel_size = self.config.get(\"ulysses_sequence_parallel_size\", 1)\n dp = world_size // self.ulysses_sequence_parallel_size\n if self.ulysses_sequence_parallel_size > 1:\n self.ulysses_device_mesh = init_device_mesh(\n device_name, mesh_shape=(dp, self.ulysses_sequence_parallel_size), mesh_dim_names=[\"dp\", \"sp\"]\n )\n\n # create training dispatch\n if self.ulysses_device_mesh is not None:\n is_collect = self.ulysses_device_mesh[\"sp\"].get_local_rank() == 0\n self._register_dispatch_collect_info(\n \"critic\", dp_rank=self.ulysses_device_mesh[\"dp\"].get_local_rank(), is_collect=is_collect\n )\n else:\n self._register_dispatch_collect_info(\"critic\", dp_rank=self.rank, is_collect=True)\n\n self.ulysses_sharding_manager = FSDPUlyssesShardingManager(self.ulysses_device_mesh)\n\n # set FSDP offload params\n self._is_offload_param = self.config.model.fsdp_config.param_offload\n self._is_offload_optimizer = self.config.model.fsdp_config.optimizer_offload\n\n # normalize config\n self.config.ppo_mini_batch_size *= self.config.rollout_n\n self.config.ppo_mini_batch_size //= torch.distributed.get_world_size() // self.ulysses_sequence_parallel_size\n if self.config.ppo_micro_batch_size is not None:\n self.config.ppo_micro_batch_size //= (\n torch.distributed.get_world_size() // self.ulysses_sequence_parallel_size\n )\n self.config.forward_micro_batch_size //= (\n torch.distributed.get_world_size() // self.ulysses_sequence_parallel_size\n )\n self.config.ppo_micro_batch_size_per_gpu = self.config.ppo_micro_batch_size\n self.config.forward_micro_batch_size_per_gpu = self.config.forward_micro_batch_size\n\n if self.config.ppo_micro_batch_size_per_gpu is not None:\n assert self.config.ppo_mini_batch_size % self.config.ppo_micro_batch_size_per_gpu == 0, (\n f\"normalized ppo_mini_batch_size {self.config.ppo_mini_batch_size} should be divisible by \"\n f\"ppo_micro_batch_size_per_gpu {self.config.ppo_micro_batch_size_per_gpu}\"\n )\n assert self.config.ppo_mini_batch_size // self.config.ppo_micro_batch_size_per_gpu > 0, (\n f\"normalized ppo_mini_batch_size {self.config.ppo_mini_batch_size} should be larger than \"\n f\"ppo_micro_batch_size_per_gpu {self.config.ppo_micro_batch_size_per_gpu}\"\n )\n self._is_lora = (\n self.config.model.get(\"lora_adapter_path\") is not None or self.config.model.get(\"lora_rank\", 0) > 0\n )\n self.use_orig_params = self.config.model.fsdp_config.get(\"use_orig_params\", False)\n\n def _build_critic_model_optimizer(self, config: FSDPCriticConfig):\n # the following line is necessary\n from torch.distributed.fsdp import MixedPrecision\n\n from verl.utils.model import load_valuehead_model, print_model_size\n from verl.utils.torch_dtypes import PrecisionType\n\n use_shm = config.model.get(\"use_shm\", False)\n local_path = copy_to_local(config.model.path, use_shm=use_shm)\n # note that the tokenizer between actor and critic may be different. So override tokenizer info with actor info\n # using random initialized model from any architecture. May not be the same as Actor.\n\n tokenizer_path = copy_to_local(config.model.tokenizer_path, use_shm=use_shm)\n self.tokenizer = hf_tokenizer(tokenizer_path, trust_remote_code=config.model.get(\"trust_remote_code\", False))\n self.processor = hf_processor(tokenizer_path, trust_remote_code=config.model.get(\"trust_remote_code\", False))\n\n if self.config.model.get(\"custom_chat_template\", None) is not None:\n if self.processor is not None:\n self.processor.chat_template = self.config.model.custom_chat_template\n else:\n self.tokenizer.chat_template = self.config.model.custom_chat_template\n override_config = OmegaConf.to_container(OmegaConf.create(self.config.model.get(\"override_config\", {})))\n override_config_kwargs = {\n \"bos_token_id\": self.tokenizer.bos_token_id,\n \"eos_token_id\": self.tokenizer.eos_token_id,\n \"pad_token_id\": self.tokenizer.pad_token_id,\n }\n override_config_kwargs.update(override_config)\n if self.rank == 0:\n print(f\"Critic overriding config {override_config_kwargs}\")\n\n torch_dtype = self.config.model.fsdp_config.get(\"model_dtype\", \"fp32\")\n torch_dtype = PrecisionType.to_dtype(torch_dtype)\n\n from transformers import AutoConfig\n\n # override model kwargs\n attn_implementation = override_config.get(\"attn_implementation\", \"flash_attention_2\")\n critic_model_config = AutoConfig.from_pretrained(\n local_path,\n attn_implementation=attn_implementation,\n trust_remote_code=config.model.get(\"trust_remote_code\", False),\n )\n # TODO: VL models use VisionAttention, which directly uses flash_attention in transformers>=4.53\n # which will be patched by _ulysses_flash_attention_forward, but errorly misses position_ids\n # Maybe support Ulysses in VisionAttention in the future and remove this patch\n if self.ulysses_sequence_parallel_size > 1 and hasattr(critic_model_config, \"vision_config\"):\n critic_model_config.vision_config._attn_implementation = \"eager\"\n\n critic_model_config.num_labels = 1\n # patch for kimi-vl\n if getattr(critic_model_config, \"model_type\", None) == \"kimi_vl\":\n critic_model_config.text_config.topk_method = \"greedy\"\n\n init_context = get_init_weight_context_manager(\n use_meta_tensor=not critic_model_config.tie_word_embeddings, mesh=self.device_mesh\n )\n\n # TiledMLP configuration for memory-efficient MLP computation\n tiled_mlp_config = config.model.get(\"tiled_mlp\", {})\n use_tiled_mlp = tiled_mlp_config.get(\"enabled\", False)\n tiled_mlp_shards = tiled_mlp_config.get(\"num_shards\", 4)\n\n # TiledMLP requires FSDP2 for correct gradient computation\n if use_tiled_mlp and config.strategy == \"fsdp\":\n raise ValueError(\"TiledMLP requires FSDP2. Set `critic.strategy=fsdp2`.\")\n\n with init_context(), warnings.catch_warnings():\n warnings.simplefilter(\"ignore\")\n critic_model_config.classifier_dropout = 0.0\n critic_model_config.hidden_dropout = \"0\"\n critic_model_config.summary_dropout_prob = 0.0\n\n critic_module = load_valuehead_model(\n local_path,\n torch_dtype,\n critic_model_config,\n config.model.get(\"trust_remote_code\", False),\n )\n\n use_remove_padding = config.model.get(\"use_remove_padding\", False)\n\n apply_monkey_patch(\n model=critic_module,\n use_remove_padding=use_remove_padding,\n ulysses_sp_size=self.ulysses_sequence_parallel_size,\n use_tiled_mlp=use_tiled_mlp,\n tiled_mlp_shards=tiled_mlp_shards,\n )\n\n # some parameters may not in torch_dtype\n critic_module.to(torch_dtype)\n\n if config.model.get(\"enable_gradient_checkpointing\", False):\n critic_module.gradient_checkpointing_enable(gradient_checkpointing_kwargs={\"use_reentrant\": False})\n\n if self._is_lora:\n print(\"Applying LoRA to critic module\")\n critic_module.enable_input_require_grads()\n\n # Check if we should load a pre-trained LoRA adapter\n lora_adapter_path = self.config.model.get(\"lora_adapter_path\")\n if lora_adapter_path is not None:\n from peft import PeftModel\n\n print(f\"Loading pre-trained LoRA adapter to critic from: {lora_adapter_path}\")\n\n # Copy adapter to local if needed\n local_adapter_path = copy_to_local(lora_adapter_path, use_shm=self.config.model.get(\"use_shm\", False))\n\n critic_module = PeftModel.from_pretrained(critic_module, local_adapter_path, is_trainable=True)\n peft_config = critic_module.peft_config[\"default\"]\n # Ensure task_type is TaskType enum, not string\n # Use TOKEN_CLS for Critic since it's loaded as AutoModelForTokenClassification\n if isinstance(peft_config.task_type, str):\n peft_config.task_type = TaskType.TOKEN_CLS\n\n else:\n # Convert config to regular Python types before creating PEFT model\n # Use TOKEN_CLS for Critic since it's loaded as AutoModelForTokenClassification\n lora_config = {\n \"task_type\": TaskType.TOKEN_CLS,\n \"r\": self.config.model.lora_rank,\n \"lora_alpha\": self.config.model.lora_alpha,\n \"target_modules\": convert_to_regular_types(self.config.model.target_modules),\n \"bias\": \"none\",\n }\n critic_module = get_peft_model(critic_module, LoraConfig(**lora_config))\n\n if self.rank == 0:\n print_model_size(critic_module)\n\n self.critic_model_config = critic_model_config\n\n fsdp_config = self.config.model.fsdp_config\n mixed_precision_config = fsdp_config.get(\"mixed_precision\", None)\n if mixed_precision_config is not None:\n param_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"param_dtype\", \"bf16\"))\n reduce_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"reduce_dtype\", \"fp32\"))\n buffer_dtype = PrecisionType.to_dtype(mixed_precision_config.get(\"buffer_dtype\", \"fp32\"))\n else:\n param_dtype = torch.bfloat16\n reduce_dtype = torch.float32\n buffer_dtype = torch.float32\n\n mixed_precision = MixedPrecision(param_dtype=param_dtype, reduce_dtype=reduce_dtype, buffer_dtype=buffer_dtype)\n\n auto_wrap_policy = get_fsdp_wrap_policy(\n module=critic_module,\n config=self.config.model.fsdp_config.wrap_policy,\n is_lora=self._is_lora,\n )\n\n log_gpu_memory_usage(\"Before critic FSDP\", logger=None)\n\n fsdp_mesh = self.device_mesh\n sharding_strategy = get_sharding_strategy(fsdp_mesh)\n\n self.use_orig_params = fsdp_config.get(\"use_orig_params\", False)\n if self.config.model.get(\"freeze_vision_tower\", False):\n vision_tower = get_vl_model_vision_tower(critic_module)\n if vision_tower is not None:\n vision_tower.requires_grad_(False)\n self.use_orig_params = True\n if self.rank == 0:\n print(\"[critic model] Vision tower is set to not trainable.\")\n else:\n if self.rank == 0:\n print(\"[critic model] No vision tower found.\")\n\n # Note: We force turn off CPUOffload for critic because it causes incorrect results when using grad accumulation\n if config.strategy == \"fsdp\":\n critic_module = FSDP(\n critic_module,\n param_init_fn=init_fn,\n use_orig_params=self.use_orig_params,\n auto_wrap_policy=auto_wrap_policy,\n device_id=get_device_id(),\n sharding_strategy=sharding_strategy,\n mixed_precision=mixed_precision,\n sync_module_states=True,\n forward_prefetch=self.config.model.fsdp_config.forward_prefetch,\n device_mesh=self.device_mesh,\n cpu_offload=None,\n )\n elif config.strategy == \"fsdp2\":\n assert CPUOffloadPolicy is not None, \"PyTorch version >= 2.4 is required for using fully_shard API (FSDP2)\"\n mp_policy = MixedPrecisionPolicy(\n param_dtype=param_dtype, reduce_dtype=reduce_dtype, cast_forward_inputs=True\n )\n offload_policy = None\n if fsdp_config.offload_policy:\n self._is_offload_param = False\n self._is_offload_optimizer = False\n offload_policy = CPUOffloadPolicy(pin_memory=True)\n\n fsdp_kwargs = {\n \"mesh\": fsdp_mesh,\n \"mp_policy\": mp_policy,\n \"offload_policy\": offload_policy,\n \"reshard_after_forward\": fsdp_config.reshard_after_forward,\n \"shard_placement_fn\": get_shard_placement_fn(fsdp_size=self.device_mesh.shape[-1]),\n }\n full_state = critic_module.state_dict()\n apply_fsdp2(critic_module, fsdp_kwargs, fsdp_config)\n fsdp2_load_full_state_dict(critic_module, full_state, fsdp_mesh, offload_policy)\n else:\n raise NotImplementedError(f\"Unknown strategy {config.strategy}\")\n\n if config.model.get(\"enable_activation_offload\", False):\n enable_gradient_checkpointing = config.model.get(\"enable_gradient_checkpointing\", False)\n enable_activation_offloading(critic_module, config.strategy, enable_gradient_checkpointing)\n\n log_gpu_memory_usage(\"After critic FSDP\", logger=None)\n\n critic_optimizer = build_optimizer(critic_module.parameters(), config.optim)\n\n total_steps = config.optim.get(\"total_training_steps\", 0)\n num_warmup_steps = int(config.optim.get(\"lr_warmup_steps\", -1))\n\n lr_scheduler_type = config.optim.get(\"lr_scheduler_type\", \"constant\")\n if num_warmup_steps < 0:\n num_warmup_steps_ratio = config.optim.get(\"lr_warmup_steps_ratio\", 0.0)\n num_warmup_steps = int(num_warmup_steps_ratio * total_steps)\n\n if self.rank == 0:\n print(f\"Total steps: {total_steps}, num_warmup_steps: {num_warmup_steps}\")\n\n from verl.utils.torch_functional import get_constant_schedule_with_warmup, get_cosine_schedule_with_warmup\n\n if lr_scheduler_type == \"constant\":\n critic_lr_scheduler = get_constant_schedule_with_warmup(\n optimizer=critic_optimizer, num_warmup_steps=num_warmup_steps\n )\n elif lr_scheduler_type == \"cosine\":\n min_lr_ratio = config.optim.get(\"min_lr_ratio\", 0.0)\n num_cycles = config.optim.get(\"num_cycles\", 0.5)\n critic_lr_scheduler = get_cosine_schedule_with_warmup(\n optimizer=critic_optimizer,\n num_warmup_steps=num_warmup_steps,\n num_training_steps=total_steps,\n min_lr_ratio=min_lr_ratio,\n num_cycles=num_cycles,\n )\n else:\n raise NotImplementedError(f\"LR scheduler type {lr_scheduler_type} is not supported\")\n\n return critic_module, critic_optimizer, critic_lr_scheduler\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def init_model(self):\n # This is used to import external_lib into the huggingface systems\n import_external_libs(self.config.model.get(\"external_lib\", None))\n\n from verl.workers.critic import DataParallelPPOCritic\n\n self.critic_module, self.critic_optimizer, self.critic_lr_scheduler = self._build_critic_model_optimizer(\n self.config\n )\n\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.critic_module)\n log_gpu_memory_usage(\"After offload critic model during init\", logger=logger)\n if self._is_offload_optimizer:\n offload_fsdp_optimizer(optimizer=self.critic_optimizer)\n log_gpu_memory_usage(\"After offload critic optimizer during init\", logger=logger)\n\n self.critic = DataParallelPPOCritic(\n config=self.config, critic_module=self.critic_module, critic_optimizer=self.critic_optimizer\n )\n\n self.flops_counter = FlopsCounter(self.critic_model_config)\n self.checkpoint_manager = FSDPCheckpointManager(\n model=self.critic_module,\n optimizer=self.critic_optimizer,\n lr_scheduler=self.critic_lr_scheduler,\n processing_class=self.processor if self.processor is not None else self.tokenizer,\n checkpoint_config=self.config.checkpoint,\n trust_remote_code=self.config.model.get(\"trust_remote_code\", False),\n )\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"critic\"))\n @DistProfiler.annotate(color=\"cyan\", role=\"compute_values\")\n def compute_values(self, data: DataProto):\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.critic_module)\n micro_batch_size = self.config.forward_micro_batch_size_per_gpu\n data.meta_info[\"micro_batch_size\"] = micro_batch_size\n data.meta_info[\"max_token_len\"] = self.config.forward_max_token_len_per_gpu\n data.meta_info[\"use_dynamic_bsz\"] = self.config.use_dynamic_bsz\n # perform forward computation\n with self.ulysses_sharding_manager:\n data = data.to(\"cpu\") # data will to device with each micro batch on critic.compute_values\n values = self.critic.compute_values(data=data)\n output = DataProto.from_dict(tensors={\"values\": values})\n\n output = output.to(\"cpu\")\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.critic_module)\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"critic\"))\n @DistProfiler.annotate(color=\"pink\", role=\"critic_update\")\n def update_critic(self, data: DataProto):\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.critic_module)\n if self._is_offload_optimizer:\n load_fsdp_optimizer(optimizer=self.critic_optimizer, device_id=get_device_id())\n\n # perform forward computation\n with self.ulysses_sharding_manager:\n data = data.to(\"cpu\") # data will to device with each micro batch on critic.update_critic\n with Timer(name=\"update_critic\", logger=None) as timer:\n metrics = self.critic.update_critic(data=data)\n delta_time = timer.last\n\n global_num_tokens = data.meta_info[\"global_token_num\"]\n estimated_flops, promised_flops = self.flops_counter.estimate_flops(global_num_tokens, delta_time)\n metrics[\"perf/mfu/critic\"] = estimated_flops * self.config.ppo_epochs / promised_flops / self.world_size\n\n lr = self.critic_lr_scheduler.get_last_lr()[0]\n metrics[\"critic/lr\"] = lr\n self.critic_lr_scheduler.step()\n\n output = DataProto(batch=None, meta_info={\"metrics\": metrics})\n\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.critic_module)\n if self._is_offload_optimizer:\n offload_fsdp_optimizer(optimizer=self.critic_optimizer)\n\n output = output.to(\"cpu\")\n return output\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def save_checkpoint(self, local_path, hdfs_path=None, global_step=0, max_ckpt_to_keep=None):\n import torch\n\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.critic_module)\n\n self.checkpoint_manager.save_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, global_step=global_step, max_ckpt_to_keep=max_ckpt_to_keep\n )\n\n torch.distributed.barrier()\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.critic_module)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def load_checkpoint(self, local_path, hdfs_path=None, del_local_after_load=True):\n import torch\n\n if self._is_offload_param:\n load_fsdp_model_to_gpu(self.critic_module)\n\n self.checkpoint_manager.load_checkpoint(\n local_path=local_path, hdfs_path=hdfs_path, del_local_after_load=del_local_after_load\n )\n\n torch.distributed.barrier()\n if self._is_offload_param:\n offload_fsdp_model_to_cpu(self.critic_module)\n\n if self._is_offload_optimizer:\n offload_fsdp_optimizer(self.critic_optimizer)\n\n\n# ================================= Async related workers =================================\nclass AsyncActorRolloutRefWorker(ActorRolloutRefWorker):\n @register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=False)\n async def update_weights(self):\n await self.rollout_mode()\n return True\n"}150{"file_name": "verl__workers__megatron_workers.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nThe main entry point to run the PPO algorithm\n\"\"\"\n\nimport datetime\nimport logging\nimport os\nimport time\n\nimport psutil\nimport torch\nimport torch.distributed\nfrom codetiming import Timer\nfrom omegaconf import DictConfig, OmegaConf\n\ntry:\n from verl.workers.engine.mindspeed.transformer_impl import repatch\nexcept ImportError:\n repatch = None\n\nfrom contextlib import nullcontext\n\nfrom megatron.core import parallel_state as mpu\n\nfrom verl import DataProto\nfrom verl.models.mcore import get_mcore_weight_converter\nfrom verl.single_controller.base import Worker\nfrom verl.single_controller.base.decorator import Dispatch, make_nd_compute_dataproto_dispatch_fn, register\nfrom verl.utils import hf_tokenizer\nfrom verl.utils.checkpoint.megatron_checkpoint_manager import MegatronCheckpointManager\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.utils.device import (\n get_device_id,\n get_device_name,\n get_nccl_backend,\n get_torch_device,\n set_expandable_segments,\n)\nfrom verl.utils.distributed import set_numa_affinity\nfrom verl.utils.flops_counter import FlopsCounter\nfrom verl.utils.fs import copy_to_local\nfrom verl.utils.megatron.router_replay_patch import RouterReplay, RouterReplayAction, apply_router_replay_patch\nfrom verl.utils.megatron_peft_utils import add_base_layer_suffix, build_peft_config_for_vllm\nfrom verl.utils.megatron_utils import (\n load_megatron_model_to_gpu,\n load_megatron_optimizer,\n offload_megatron_model_to_cpu,\n offload_megatron_optimizer,\n per_tensor_generator,\n register_megatron_training_hooks,\n)\nfrom verl.utils.memory_utils import aggressive_empty_cache\nfrom verl.utils.model import get_hf_model_path, load_mcore_dist_weights, load_megatron_gptmodel_weights\nfrom verl.utils.profiler import (\n DistProfiler,\n DistProfilerExtension,\n GPUMemoryLogger,\n ProfilerConfig,\n log_gpu_memory_usage,\n simple_timer,\n)\nfrom verl.utils.profiler.performance import reduce_timing, topk_reduce_ratio_min_max\nfrom verl.utils.ray_utils import get_event_loop\nfrom verl.utils.torch_functional import use_original_torch_compile\nfrom verl.workers.actor.megatron_actor import MegatronPPOActor\nfrom verl.workers.config import HFModelConfig, McoreCriticConfig, RolloutConfig\nfrom verl.workers.critic.megatron_critic import MegatronPPOCritic\nfrom verl.workers.rollout import get_rollout_class\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\ndef set_random_seed(seed, only_rollout=False):\n import random\n\n import numpy as np\n import torch\n\n torch.manual_seed(seed)\n np.random.seed(seed)\n random.seed(seed)\n if not only_rollout and get_torch_device().device_count() > 0:\n from megatron.core import tensor_parallel\n\n tensor_parallel.model_parallel_cuda_manual_seed(seed)\n # FIXME: torch cumsum not support deterministic (used in vllm sampler),\n # https://github.com/pytorch/pytorch/issues/89492\n # torch.use_deterministic_algorithms(True, warn_only=True)\n # os.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'\n\n\nclass MegatronWorker(Worker):\n def _init_hf_config_and_tf_config(\n self,\n model_path,\n tokenizer_or_path,\n dtype,\n override_model_config,\n override_transformer_config,\n trust_remote_code=False,\n megatron_config=None,\n enable_mtp=False,\n ):\n from transformers import AutoConfig\n\n from verl.models.mcore import hf_to_mcore_config\n from verl.utils import hf_processor\n from verl.utils.model import update_model_config\n\n # Step 1: initialize the tokenizer\n self.local_path = copy_to_local(model_path)\n if tokenizer_or_path is None:\n self.tokenizer = hf_tokenizer(self.local_path, trust_remote_code=trust_remote_code)\n self.processor = hf_processor(self.local_path, trust_remote_code=trust_remote_code)\n elif isinstance(tokenizer_or_path, str):\n self.tokenizer = hf_tokenizer(copy_to_local(tokenizer_or_path), trust_remote_code=trust_remote_code)\n self.processor = hf_processor(copy_to_local(tokenizer_or_path), trust_remote_code=trust_remote_code)\n else:\n self.tokenizer = tokenizer_or_path\n self.processor = tokenizer_or_path\n\n if self.config.model.get(\"custom_chat_template\", None) is not None:\n if self.processor is not None:\n self.processor.chat_template = self.config.model.custom_chat_template\n else:\n self.tokenizer.chat_template = self.config.model.custom_chat_template\n\n # Step 2: get the hf\n hf_config = AutoConfig.from_pretrained(self.local_path, trust_remote_code=trust_remote_code)\n\n # Step 3: override the hf config\n override_config_kwargs = {\n \"bos_token_id\": self.tokenizer.bos_token_id,\n \"eos_token_id\": self.tokenizer.eos_token_id,\n \"pad_token_id\": self.tokenizer.pad_token_id,\n }\n override_config_kwargs.update(override_model_config.get(\"model_config\", {}))\n self.share_embeddings_and_output_weights = getattr(hf_config, \"tie_word_embeddings\", False)\n\n # only actor need enable mtp\n if enable_mtp:\n assert hf_config.num_nextn_predict_layers > 0, \"MTP requires at least one nextn_predict_layer\"\n assert megatron_config.use_mbridge, \"MTP requires use_mbridge to be True\"\n assert megatron_config.vanilla_mbridge, \"MTP requires vanilla_mbridge to be True\"\n override_transformer_config[\"mtp_loss_scaling_factor\"] = self.config.model.mtp.mtp_loss_scaling_factor\n else:\n if hasattr(hf_config, \"num_nextn_predict_layers\"):\n hf_config.num_nextn_predict_layers = 0\n\n self.enable_mtp = enable_mtp\n\n update_model_config(hf_config, override_config_kwargs=override_config_kwargs)\n self.architectures = getattr(hf_config, \"architectures\", None)\n if self.rank == 0:\n print(f\"Model config after override: {hf_config}\")\n\n from verl.models.mcore.config_converter import mapping_string_to_attn_backend\n\n # todo: remove this line after mcore adopt mbridge 0.15, now for compatibility\n override_transformer_config = mapping_string_to_attn_backend(override_transformer_config)\n fp16 = dtype == torch.float16\n bf16 = dtype == torch.bfloat16\n if fp16:\n assert megatron_config.use_mbridge, \"fp16 mode requires use_mbridge to be True\"\n\n self.provider = None\n self.vanilla_bridge = megatron_config.get(\"vanilla_mbridge\", True)\n if megatron_config.use_mbridge:\n if self.vanilla_bridge:\n from verl.models.mcore.mbridge import AutoBridge\n\n bridge = AutoBridge.from_config(hf_config, dtype=dtype)\n bridge.set_extra_args(**override_transformer_config)\n tf_config = bridge.config\n tf_config.fp16 = fp16\n tf_config.bf16 = bf16\n else:\n from verl.models.mcore.bridge import AutoBridge\n\n # Use Megatron-Bridge to convert HF config to Megatron config\n bridge = AutoBridge.from_hf_pretrained(self.local_path, trust_remote_code=trust_remote_code)\n # Get Megatron provider and configure it\n provider = bridge.to_megatron_provider(load_weights=False)\n\n # In case of invalid overrides, we need to make sure some critical params are set correctly\n provider.params_dtype = dtype\n\n # Ensure dtype settings propagate to Megatron-Bridge/TE\n provider.fp16 = fp16\n provider.bf16 = bf16\n\n # Pass distributed info\n provider.tensor_model_parallel_size = megatron_config.tensor_model_parallel_size\n provider.pipeline_model_parallel_size = megatron_config.pipeline_model_parallel_size\n provider.expert_model_parallel_size = megatron_config.expert_model_parallel_size\n provider.expert_tensor_parallel_size = megatron_config.expert_tensor_parallel_size\n provider.virtual_pipeline_model_parallel_size = megatron_config.virtual_pipeline_model_parallel_size\n provider.context_parallel_size = megatron_config.context_parallel_size\n provider.sequence_parallel = megatron_config.sequence_parallel\n\n # Match verl implementation (need variable_seq_lengths)\n from megatron.core.transformer.enums import AttnBackend\n\n provider.attention_backend = AttnBackend.flash\n provider.variable_seq_lengths = True\n provider.moe_token_dispatcher_type = \"alltoall\"\n provider.moe_router_load_balancing_type = \"none\"\n\n # Apply transformer config overrides\n for key, value in override_transformer_config.items():\n setattr(provider, key, value)\n\n provider.finalize()\n self.provider = provider\n tf_config = None # Will be set after model creation\n self.bridge = bridge\n else:\n tf_config = hf_to_mcore_config(hf_config, dtype, **override_transformer_config)\n self.bridge = None\n\n if torch.distributed.get_rank() == 0:\n if tf_config is not None:\n print(f\"TF config: {tf_config}\")\n self.hf_config = hf_config\n self.tf_config = tf_config\n\n # Get PEFT config from model.lora if specified\n from verl.workers.config.megatron_peft import get_peft_cls\n\n self.peft_cls = get_peft_cls(\n model_config=self.config.model, bridge=self.bridge, provider=self.provider, dtype=dtype\n )\n\n\nclass ActorRolloutRefWorker(MegatronWorker, DistProfilerExtension):\n \"\"\"\n This worker can be instantiated as a standalone actor or a standalone rollout or a standalone reference policy\n or a hybrid engine based on the config.rollout\n \"\"\"\n\n def __init__(self, config: DictConfig, role: str, **kwargs):\n Worker.__init__(self)\n self.config = config\n if repatch is not None:\n # NPU MindSpeed patch, will be refactored with MindSpeedEngine.\n repatch(self.config.actor.megatron.get(\"override_transformer_config\", {}))\n\n self.role = role\n assert self.role in [\"actor\", \"rollout\", \"ref\", \"actor_rollout\", \"actor_rollout_ref\"]\n\n self._is_actor = self.role in [\"actor\", \"actor_rollout\", \"actor_rollout_ref\"]\n self._is_rollout = self.role in [\"rollout\", \"actor_rollout\", \"actor_rollout_ref\"]\n self._is_ref = self.role in [\"ref\", \"actor_rollout_ref\"]\n\n # NOTE(sgm): We utilize colocate WorkerGroup by default.\n # As a result, Workers for different model share the same process.\n # Therefore, we only require one distribute initialization.\n # To utilize different parallel strategy in different models:\n # 1, users should disable WorkerDict; 2.assign different ResourcePool to different models,\n # 3. and apply the following patch in ray==2.10, https://github.com/ray-project/ray/pull/44385\n if not torch.distributed.is_initialized():\n set_numa_affinity()\n rank = int(os.environ[\"LOCAL_RANK\"])\n torch.distributed.init_process_group(\n backend=f\"cpu:gloo,{get_device_name()}:{get_nccl_backend()}\",\n timeout=datetime.timedelta(seconds=self.config.get(\"nccl_timeout\", 600)),\n init_method=os.environ.get(\"DIST_INIT_METHOD\", None),\n )\n get_torch_device().set_device(rank)\n\n if self._is_actor or self._is_ref:\n mpu.initialize_model_parallel(\n tensor_model_parallel_size=self.config.actor.megatron.tensor_model_parallel_size,\n pipeline_model_parallel_size=self.config.actor.megatron.pipeline_model_parallel_size,\n virtual_pipeline_model_parallel_size=self.config.actor.megatron.virtual_pipeline_model_parallel_size,\n use_sharp=False,\n context_parallel_size=self.config.actor.megatron.context_parallel_size,\n expert_model_parallel_size=self.config.actor.megatron.expert_model_parallel_size,\n expert_tensor_parallel_size=self.config.actor.megatron.expert_tensor_parallel_size,\n nccl_communicator_config_path=None,\n )\n\n if self._is_actor or self._is_ref:\n is_collect = (\n mpu.get_tensor_model_parallel_rank() == 0\n and mpu.get_pipeline_model_parallel_rank() == mpu.get_pipeline_model_parallel_world_size() - 1\n and mpu.get_context_parallel_rank() == 0\n )\n self._register_dispatch_collect_info(\n mesh_name=\"actor\", dp_rank=mpu.get_data_parallel_rank(), is_collect=is_collect\n )\n only_rollout = self._is_rollout and not self._is_actor\n\n self.enable_routing_replay = False\n if self._is_actor:\n self.router_replay = self.config.actor.router_replay\n self.enable_routing_replay = self.router_replay.mode != \"disabled\"\n\n if self.enable_routing_replay:\n apply_router_replay_patch()\n\n set_random_seed(seed=self.config.actor.megatron.seed, only_rollout=only_rollout)\n\n if self._is_actor:\n omega_profiler_config = config.actor.get(\"profiler\", {})\n elif self._is_rollout:\n # NOTE: In colocation mode, rollout config may not take effect (follow the actor config)\n # This is for extendability in AsyncRL cases\n omega_profiler_config = config.rollout.get(\"profiler\", {})\n elif self._is_ref:\n omega_profiler_config = config.ref.get(\"profiler\", {})\n else:\n raise ValueError(\n f\"Invalid role {self.role}, should be one of \"\n \"['actor', 'rollout', 'ref', 'actor_rollout', 'actor_rollout_ref']\"\n )\n # omega_profiler_config is DictConfig\n # profiler_config is a ProfilerConfig dataclass\n profiler_config = omega_conf_to_dataclass(omega_profiler_config, dataclass_type=ProfilerConfig)\n if omega_profiler_config.get(\"tool\", None) in [\"npu\", \"nsys\", \"torch\", \"torch_memory\"]:\n tool_config = omega_conf_to_dataclass(\n omega_profiler_config.get(\"tool_config\", {}).get(omega_profiler_config.get(\"tool\"))\n )\n else:\n tool_config = None\n DistProfilerExtension.__init__(\n self, DistProfiler(rank=self.rank, config=profiler_config, tool_config=tool_config)\n )\n\n # TODO(sgm): Currently, we only support reference model param offload\n # will support other offload later\n self._is_offload_param = False\n self._is_offload_grad = False\n self._is_offload_optimizer = False\n\n # Initialize LoRA-related attributes (will be updated in _build_rollout if needed)\n self.base_sync_done = False\n self.peft_merge = False\n\n # normalize config\n if self._is_actor:\n self.config.actor.ppo_mini_batch_size *= self.config.rollout.n\n self.config.actor.ppo_mini_batch_size //= mpu.get_data_parallel_world_size()\n if self.config.actor.get(\"ppo_micro_batch_size\", None):\n self.config.actor.ppo_micro_batch_size //= mpu.get_data_parallel_world_size()\n self.config.rollout.log_prob_micro_batch_size //= mpu.get_data_parallel_world_size()\n self.config.actor.ppo_micro_batch_size_per_gpu = self.config.actor.ppo_micro_batch_size\n self.config.rollout.log_prob_micro_batch_size_per_gpu = self.config.rollout.log_prob_micro_batch_size\n\n self._is_offload_param = self.config.actor.megatron.get(\"param_offload\", False)\n self._is_offload_grad = self.config.actor.megatron.get(\"grad_offload\", False)\n self._is_offload_optimizer = self.config.actor.megatron.get(\"optimizer_offload\", False)\n elif self._is_ref:\n if self.config.ref.get(\"log_prob_micro_batch_size\", None):\n self.config.ref.log_prob_micro_batch_size //= mpu.get_data_parallel_world_size()\n self.config.ref.log_prob_micro_batch_size_per_gpu = self.config.ref.log_prob_micro_batch_size\n else:\n assert self.config.ref.get(\"log_prob_micro_batch_size_per_gpu\", None) is not None, (\n \"Please note that in the ref policy configuration, `log_prob_micro_batch_size_per_gpu` and \"\n \"`log_prob_micro_batch_size` should not be None at the same time.\"\n )\n self._ref_is_offload_param = self.config.ref.megatron.get(\"param_offload\", False)\n\n def _build_model_optimizer(\n self, model_path, optim_config, override_model_config, override_transformer_config, override_ddp_config=None\n ):\n from verl.utils.megatron.optimizer import (\n get_megatron_optimizer,\n get_megatron_optimizer_param_scheduler,\n init_megatron_optim_config,\n )\n from verl.utils.megatron_utils import McoreModuleWrapperConfig, make_megatron_module\n from verl.utils.model import get_generation_config, print_model_size\n\n self._init_hf_config_and_tf_config(\n model_path,\n self.config.model.get(\"tokenizer_path\") or model_path,\n self.dtype,\n override_model_config,\n override_transformer_config,\n self.config.model.get(\"trust_remote_code\", False),\n self.config.actor.megatron if not self._is_ref else self.config.ref.megatron,\n self.config.model.get(\"mtp\", {}).get(\"enable\", False),\n )\n self.generation_config = get_generation_config(\n self.local_path,\n self.config.model.get(\"trust_remote_code\", False),\n )\n\n if self._is_actor or self._is_rollout:\n wrap_config = McoreModuleWrapperConfig(\n is_value_model=False, # actor is not value model\n share_embeddings_and_output_weights=self.share_embeddings_and_output_weights,\n wrap_with_ddp=True,\n use_distributed_optimizer=self.config.actor.megatron.use_distributed_optimizer,\n )\n actor_module, updated_tf_config = make_megatron_module(\n wrap_config=wrap_config,\n tf_config=self.tf_config,\n hf_config=self.hf_config,\n bridge=self.bridge,\n provider=self.provider,\n override_model_config=override_model_config,\n override_ddp_config=override_ddp_config,\n peft_cls=self.peft_cls,\n peft_config=self.config.model.get(\"lora\", None),\n )\n self.tf_config = updated_tf_config\n print(f\"actor_module: {len(actor_module)}\")\n if self.config.actor.load_weight:\n if self.config.actor.megatron.use_dist_checkpointing:\n load_mcore_dist_weights(\n actor_module,\n self.config.actor.megatron.dist_checkpointing_path,\n is_value_model=False,\n prefix=self.config.actor.megatron.dist_checkpointing_prefix,\n )\n else:\n if self.bridge is not None:\n local_model_path = get_hf_model_path(self.config)\n if self.vanilla_bridge:\n self.bridge.load_weights(actor_module, local_model_path)\n else:\n self.bridge.load_hf_weights(actor_module, local_model_path)\n else:\n load_megatron_gptmodel_weights(\n self.config, self.hf_config, actor_module, params_dtype=self.dtype, is_value_model=False\n )\n\n if self.rank == 0:\n print_model_size(actor_module[0])\n log_gpu_memory_usage(\"After MegatronPPOActor init\", logger=logger)\n elif self._is_ref:\n wrap_config = McoreModuleWrapperConfig(\n is_value_model=False, # ref is not value model\n share_embeddings_and_output_weights=self.share_embeddings_and_output_weights,\n wrap_with_ddp=False,\n use_distributed_optimizer=self.config.ref.megatron.use_distributed_optimizer,\n )\n ref_module, updated_tf_config = make_megatron_module(\n wrap_config=wrap_config,\n tf_config=self.tf_config,\n hf_config=self.hf_config,\n bridge=self.bridge,\n provider=self.provider,\n override_model_config=override_model_config,\n )\n self.tf_config = updated_tf_config\n if self.config.ref.load_weight: # should align with the actor:\n assert self.config.actor.load_weight == self.config.ref.load_weight\n print(\"load ref weight start\")\n if self.config.ref.megatron.use_dist_checkpointing:\n load_mcore_dist_weights(\n ref_module,\n self.config.ref.megatron.dist_checkpointing_path,\n is_value_model=False,\n prefix=self.config.ref.megatron.dist_checkpointing_prefix,\n )\n else:\n if self.bridge is not None:\n local_model_path = get_hf_model_path(self.config)\n if self.vanilla_bridge:\n self.bridge.load_weights(ref_module, local_model_path)\n else:\n self.bridge.load_hf_weights(ref_module, local_model_path)\n else:\n load_megatron_gptmodel_weights(\n self.config, self.hf_config, ref_module, params_dtype=self.dtype, is_value_model=False\n )\n log_gpu_memory_usage(\"After ref module init\", logger=logger)\n return ref_module, self.hf_config\n\n # TODO: add more optimizer args into config\n if self._is_actor:\n optim_config_megatron = init_megatron_optim_config(\n optim_config,\n use_distributed_optimizer=wrap_config.use_distributed_optimizer,\n fp16=self.dtype == torch.float16,\n )\n actor_optimizer = get_megatron_optimizer(model=actor_module, config=optim_config_megatron)\n actor_optimizer_scheduler = get_megatron_optimizer_param_scheduler(\n optimizer=actor_optimizer, config=optim_config\n )\n else:\n optim_config = None\n actor_optimizer = None\n actor_optimizer_scheduler = None\n\n log_gpu_memory_usage(\"After actor optimizer init\", logger=logger)\n\n register_megatron_training_hooks(actor_module, actor_optimizer)\n\n return actor_module, actor_optimizer, actor_optimizer_scheduler, self.hf_config, optim_config\n\n def _build_rollout(self, trust_remote_code=False):\n from torch.distributed.device_mesh import init_device_mesh\n\n # 1. parse rollout and huggingface model config\n rollout_config: RolloutConfig = omega_conf_to_dataclass(self.config.rollout)\n model_config: HFModelConfig = omega_conf_to_dataclass(self.config.model)\n\n # 2. build rollout device mesh\n infer_tp = self.config.rollout.tensor_model_parallel_size * self.config.rollout.data_parallel_size\n infer_pp = self.config.rollout.pipeline_model_parallel_size\n infer_world_size = infer_tp * infer_pp\n dp = self.world_size // infer_world_size\n assert self.world_size % infer_world_size == 0, (\n f\"rollout world_size: {self.world_size} is not divisible by infer_world_size: {infer_world_size}\"\n )\n rollout_device_mesh = init_device_mesh(\n get_device_name(), mesh_shape=(dp, infer_tp, infer_pp), mesh_dim_names=[\"dp\", \"infer_tp\", \"infer_pp\"]\n )\n\n self.rollout_device_mesh = rollout_device_mesh\n\n is_collect = (\n rollout_device_mesh[\"infer_tp\"].get_local_rank() == 0\n and rollout_device_mesh[\"infer_pp\"].get_local_rank() == 0\n )\n self._register_dispatch_collect_info(\n \"rollout\", dp_rank=rollout_device_mesh[\"dp\"].get_local_rank(), is_collect=is_collect\n )\n\n # 4. build rollout model\n log_gpu_memory_usage(f\"Before building {self.config.rollout.name} rollout\", logger=logger)\n self.rollout = get_rollout_class(rollout_config.name, rollout_config.mode)(\n config=rollout_config, model_config=model_config, device_mesh=rollout_device_mesh\n )\n log_gpu_memory_usage(f\"After building {self.config.rollout.name} rollout\", logger=logger)\n\n # Initialize base_sync_done for LoRA\n self.base_sync_done: bool = \"dummy\" not in self.config.rollout.load_format\n self.peft_merge: bool = model_config.lora.get(\"merge\", False)\n\n # 5. switch to trainer mode\n # NOTE: It's critical that hybrid engine in trainer mode initially to load checkpoint.\n # For async mode, we can't call run_until_complete here, so we will switch to trainer mode in AgentLoopManager.\n # Note: sync mode is deprecated and rejected in RolloutConfig.__post_init__\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def init_model(self):\n if self.config.model.get(\"external_lib\", None) is not None:\n # This is used to import external_lib into the huggingface systems\n import importlib\n\n importlib.import_module(self.config.model.external_lib)\n\n from verl.utils.torch_dtypes import PrecisionType\n\n override_model_config = OmegaConf.to_container(OmegaConf.create(self.config.model.get(\"override_config\", {})))\n if self._is_actor:\n override_transformer_config = OmegaConf.to_container(\n OmegaConf.create(self.config.actor.megatron.get(\"override_transformer_config\", {}))\n )\n if self.enable_routing_replay:\n override_transformer_config[\"enable_routing_replay\"] = True\n override_ddp_config = OmegaConf.to_container(\n OmegaConf.create(self.config.actor.megatron.get(\"override_ddp_config\", {}))\n )\n elif self._is_ref:\n override_transformer_config = OmegaConf.to_container(\n OmegaConf.create(self.config.ref.megatron.get(\"override_transformer_config\", {}))\n )\n else:\n override_transformer_config = {}\n self.param_dtype = PrecisionType.to_dtype(self.config.actor.megatron.dtype)\n log_gpu_memory_usage(\"Before init actor model and optimizer\", logger=logger)\n self.dtype = PrecisionType.to_dtype(self.param_dtype)\n if self._is_actor:\n # we need the model for actor and rollout\n optim_config = self.config.actor.optim if self._is_actor else None\n (\n self.actor_module,\n self.actor_optimizer,\n self.actor_optimizer_scheduler,\n self.actor_model_config,\n self.actor_optim_config,\n ) = self._build_model_optimizer(\n model_path=self.config.model.path,\n optim_config=optim_config,\n override_model_config=override_model_config,\n override_transformer_config=override_transformer_config,\n override_ddp_config=override_ddp_config,\n )\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.actor_module)\n log_gpu_memory_usage(\"After offload actor params and grad during init\", logger=logger)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.actor_optimizer)\n log_gpu_memory_usage(\"After offload actor optimizer during init\", logger=logger)\n\n if self._is_actor:\n actor_cfg = omega_conf_to_dataclass(self.config.actor)\n self.actor = MegatronPPOActor(\n config=actor_cfg,\n model_config=self.actor_model_config,\n hf_config=self.hf_config,\n tf_config=self.tf_config,\n actor_module=self.actor_module,\n actor_optimizer=self.actor_optimizer,\n mtp_config=self.config.model.mtp if self.config.model.mtp.enable else None,\n )\n print(f\"routing replay layers: {len(RouterReplay.router_instances)}\")\n log_gpu_memory_usage(\"After MegatronPPOActor init\", logger=logger)\n\n if self._is_rollout:\n with use_original_torch_compile():\n self._build_rollout(trust_remote_code=self.config.model.get(\"trust_remote_code\", False))\n log_gpu_memory_usage(\"After rollout init\", logger=logger)\n\n if self._is_ref:\n self.ref_module, self.ref_model_config = self._build_model_optimizer(\n model_path=self.config.model.path,\n optim_config=None,\n override_model_config=override_model_config,\n override_transformer_config=override_transformer_config,\n )\n log_gpu_memory_usage(\"After ref model init\", logger=logger)\n self.ref_policy = MegatronPPOActor(\n config=self.config.ref,\n model_config=self.ref_model_config,\n hf_config=self.hf_config,\n tf_config=self.tf_config,\n actor_module=self.ref_module,\n actor_optimizer=None,\n )\n if self._ref_is_offload_param:\n offload_megatron_model_to_cpu(self.ref_module)\n log_gpu_memory_usage(\"After offload ref params during init\", logger=logger)\n\n if self._is_actor:\n self.flops_counter = FlopsCounter(self.actor_model_config)\n self.checkpoint_mananager = MegatronCheckpointManager(\n config=self.config,\n checkpoint_config=self.config.actor.checkpoint,\n model_config=self.actor_model_config,\n transformer_config=self.tf_config,\n role=\"actor\",\n model=self.actor_module,\n arch=self.architectures[0],\n hf_config=self.hf_config,\n param_dtype=self.param_dtype,\n share_embeddings_and_output_weights=self.share_embeddings_and_output_weights,\n processing_class=self.processor if self.processor is not None else self.tokenizer,\n optimizer=self.actor_optimizer,\n optimizer_scheduler=self.actor_optimizer_scheduler,\n use_distributed_optimizer=self.config.actor.megatron.use_distributed_optimizer,\n use_checkpoint_opt_param_scheduler=self.config.actor.optim.use_checkpoint_opt_param_scheduler,\n bridge=self.bridge,\n provider=self.provider,\n use_dist_checkpointing=self.config.actor.megatron.use_dist_checkpointing,\n peft_cls=self.peft_cls,\n )\n\n self.layer_name_mapping = {\n \"qkv_layer_name\": \"self_attention.linear_qkv.\",\n \"gate_proj_layer_name\": \"linear_fc1.\",\n }\n self.weight_converter = None\n if not self.config.actor.megatron.use_mbridge:\n self.weight_converter = get_mcore_weight_converter(self.actor_model_config, self.dtype)\n\n get_torch_device().empty_cache()\n log_gpu_memory_usage(\"After init_model finish\", logger=logger)\n\n async def rollout_mode(self):\n \"\"\"Context switch hybridengine to rollout mode.\"\"\"\n aggressive_empty_cache(force_sync=True)\n set_expandable_segments(False)\n\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.actor.actor_module, load_grad=False)\n log_gpu_memory_usage(\"After load actor params during rollout_mode\", logger=logger)\n\n # Build peft_config for vLLM LoRA support\n peft_config = None\n do_lora_base_sync = False\n if not self.peft_merge and self.peft_cls is not None:\n peft_config = build_peft_config_for_vllm(self.config.model.get(\"lora\", {}))\n # set sleep level for LoRA adapter weights only sync\n # TODO: make this configurable so that users with small\n # main memory can trade sync time to avoid OOM\n self.rollout.sleep_level = 1\n\n do_lora_base_sync = (not self.base_sync_done) or (\n self.rollout.sleep_level != 1 and self.config.rollout.free_cache_engine\n )\n\n if self.bridge is not None:\n if self.vanilla_bridge:\n per_tensor_param = self.bridge.export_weights(self.actor.actor_module)\n elif not self.peft_merge and self.peft_cls is not None:\n # Only export adapter weights\n per_tensor_param = self.bridge.export_adapter_weights(self.actor.actor_module)\n else:\n per_tensor_param = self.bridge.export_hf_weights(self.actor.actor_module)\n else:\n per_tensor_param = per_tensor_generator(\n self.actor.actor_module,\n self.actor_model_config,\n self.weight_converter,\n self.tf_config,\n self.layer_name_mapping,\n )\n\n if self.config.rollout.free_cache_engine:\n await self.rollout.resume(tags=[\"weights\"])\n if do_lora_base_sync:\n # Base layer sync\n per_tensor_param_lora_base = self.bridge.export_hf_weights(\n self.actor.actor_module, merge_adapter_weights=False\n )\n await self.rollout.update_weights(\n add_base_layer_suffix(per_tensor_param_lora_base, model_type=self.hf_config.model_type),\n peft_config=peft_config,\n base_sync_done=False,\n )\n\n # Mark base sync as done after first successful sync\n self.base_sync_done = True\n\n await self.rollout.update_weights(per_tensor_param, peft_config=peft_config, base_sync_done=True)\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.actor.actor_module)\n aggressive_empty_cache(force_sync=True)\n if self.config.rollout.free_cache_engine:\n await self.rollout.resume(tags=[\"kv_cache\"])\n\n set_expandable_segments(True)\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"actor\"))\n @GPUMemoryLogger(role=\"update_actor\", logger=logger)\n @DistProfiler.annotate(color=\"red\", role=\"actor_update\")\n def update_actor(self, data: DataProto):\n assert self._is_actor\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.actor_module)\n log_gpu_memory_usage(\"After load actor params and grad during update_actor\", logger=logger)\n if self._is_offload_optimizer:\n load_megatron_optimizer(self.actor_optimizer)\n log_gpu_memory_usage(\"After load actor optimizer during update_actor\", logger=logger)\n\n micro_batch_size = self.config.actor.ppo_micro_batch_size_per_gpu\n data.meta_info[\"micro_batch_size\"] = micro_batch_size\n dataloader = self.actor.make_minibatch_iterator(data=data)\n with Timer(name=\"update_policy\", logger=None) as timer:\n metrics = self.actor.update_policy(dataloader=dataloader)\n delta_time = timer.last\n global_num_tokens = data.meta_info[\"global_token_num\"]\n images_seqlens = data.meta_info.get(\"images_seqlens\", None)\n estimated_flops, promised_flops = self.flops_counter.estimate_flops(\n global_num_tokens, delta_time, images_seqlens=images_seqlens\n )\n metrics[\"perf/mfu/actor\"] = estimated_flops * self.config.actor.ppo_epochs / promised_flops / self.world_size\n metrics[\"perf/max_memory_allocated_gb\"] = get_torch_device().max_memory_allocated() / (1024**3)\n metrics[\"perf/max_memory_reserved_gb\"] = get_torch_device().max_memory_reserved() / (1024**3)\n metrics[\"perf/cpu_memory_used_gb\"] = psutil.virtual_memory().used / (1024**3)\n from verl.utils.megatron.optimizer import get_megatron_last_lr\n\n metrics[\"actor/lr\"] = get_megatron_last_lr(self.actor_optimizer)\n self.actor_optimizer_scheduler.step(1)\n\n # TODO: here, we should return all metrics\n output = DataProto(meta_info={\"metrics\": metrics})\n output = output.to(\"cpu\")\n\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.actor_module)\n log_gpu_memory_usage(\"After offload actor params and grad during update_actor\", logger=logger)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.actor_optimizer)\n log_gpu_memory_usage(\"After offload actor optimizer during update_actor\", logger=logger)\n\n aggressive_empty_cache(force_sync=True)\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"rollout\"))\n @GPUMemoryLogger(role=\"generate_sequences\", logger=logger)\n @DistProfiler.annotate(color=\"red\", role=\"rollout_generate\")\n def generate_sequences(self, prompts: DataProto):\n assert self._is_rollout\n prompts = prompts.to(get_device_name())\n meta_info = {\n \"eos_token_id\": self.generation_config.eos_token_id\n if self.generation_config is not None\n else self.tokenizer.eos_token_id,\n \"pad_token_id\": self.generation_config.pad_token_id\n if self.generation_config is not None\n else self.tokenizer.pad_token_id,\n }\n prompts.meta_info.update(meta_info)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.actor_optimizer)\n\n timing_generate = {}\n if self._is_actor: # For rollout only, we do not switch context.\n loop = get_event_loop()\n loop.run_until_complete(self.rollout_mode())\n log_gpu_memory_usage(\"After switch to rollout mode\", logger=logger)\n\n with simple_timer(\"generate_sequences\", timing_generate):\n output = self.rollout.generate_sequences(prompts=prompts)\n\n if self._is_actor:\n loop.run_until_complete(self.trainer_mode())\n log_gpu_memory_usage(\"After switch to trainer mode\", logger=logger)\n\n # We calculate the average timing across all ranks\n # to make sure meta_info[\"timing\"] is the same\n timing_generate_topk_ratio, timing_generate_min, timing_generate_max = topk_reduce_ratio_min_max(\n timing_generate[\"generate_sequences\"]\n )\n timing_generate = reduce_timing(timing_generate)\n timing_generate.update(\n {\n \"generation_timing/max\": timing_generate_max,\n \"generation_timing/min\": timing_generate_min,\n \"generation_timing/topk_ratio\": timing_generate_topk_ratio,\n }\n )\n output.meta_info[\"timing\"] = timing_generate\n output = output.to(\"cpu\")\n # clear kv cache\n aggressive_empty_cache(force_sync=True)\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"actor\"))\n @GPUMemoryLogger(role=\"compute_ref_log_prob\", logger=logger)\n @DistProfiler.annotate(color=\"olive\", role=\"ref_compute_log_prob\")\n def compute_ref_log_prob(self, data: DataProto):\n if self.peft_cls is not None:\n # if is lora, actor without lora applied is the ref\n data.meta_info[\"is_lora\"] = True\n return self.compute_log_prob(data)\n assert self._is_ref\n if self._ref_is_offload_param:\n load_megatron_model_to_gpu(self.ref_module, load_grad=False)\n log_gpu_memory_usage(\"After load ref params and grad during compute_ref_log_prob\", logger=logger)\n micro_batch_size = self.config.ref.log_prob_micro_batch_size_per_gpu\n data.meta_info[\"micro_batch_size\"] = micro_batch_size\n data.meta_info[\"max_token_len\"] = self.config.ref.log_prob_max_token_len_per_gpu\n data.meta_info[\"use_dynamic_bsz\"] = self.config.ref.log_prob_use_dynamic_bsz\n data.meta_info[\"temperature\"] = self.config.rollout.temperature\n output, _, _ = self.ref_policy.compute_log_prob(data=data, calculate_entropy=False)\n output = DataProto.from_dict(tensors={\"ref_log_prob\": output})\n output = output.to(\"cpu\")\n if self._ref_is_offload_param:\n offload_megatron_model_to_cpu(self.ref_module)\n log_gpu_memory_usage(\"After offload ref params and grad during compute_ref_log_prob\", logger=logger)\n aggressive_empty_cache(force_sync=True)\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"actor\"))\n @GPUMemoryLogger(role=\"compute_log_prob\", logger=logger)\n @DistProfiler.annotate(color=\"blue\", role=\"actor_compute_log_prob\")\n def compute_log_prob(self, data: DataProto):\n assert self._is_actor\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.actor_module, load_grad=False)\n log_gpu_memory_usage(\"After load actor params and grad during compute_log_prob\", logger=logger)\n is_lora = data.meta_info.pop(\"is_lora\", False)\n adapter_ctx = self.peft_cls.disable_adapter(self.actor_module) if is_lora else nullcontext()\n # we should always recompute old_log_probs when it is HybridEngine\n config_source = self.config.ref if is_lora else self.config.rollout\n data.meta_info[\"micro_batch_size\"] = config_source.log_prob_micro_batch_size_per_gpu\n data.meta_info[\"max_token_len\"] = config_source.log_prob_max_token_len_per_gpu\n data.meta_info[\"use_dynamic_bsz\"] = config_source.log_prob_use_dynamic_bsz\n data.meta_info[\"temperature\"] = self.config.rollout.temperature\n\n if self.enable_routing_replay and self.config.actor.router_replay.mode == \"R2\":\n RouterReplay.set_global_router_replay_action(RouterReplayAction.RECORD)\n\n if self.enable_routing_replay and self.config.actor.router_replay.mode == \"R3\":\n RouterReplay.set_global_router_replay_action(RouterReplayAction.REPLAY_FORWARD)\n\n with adapter_ctx:\n output, entropys, layers_topk_idx = self.actor.compute_log_prob(data=data, calculate_entropy=not is_lora)\n tensors = {\"ref_log_prob\": output} if is_lora else {\"old_log_probs\": output}\n if not is_lora:\n tensors[\"entropys\"] = entropys\n output = DataProto.from_dict(\n tensors=tensors,\n meta_info={\"temperature\": self.config.rollout.temperature},\n )\n if self.config.actor.router_replay.mode == \"R2\":\n output.batch[\"routed_experts\"] = layers_topk_idx\n\n if self.config.actor.router_replay.mode in [\"R2\", \"R3\"]:\n RouterReplay.clear_global_indices()\n RouterReplay.clear_global_router_replay_action()\n\n output = output.to(\"cpu\")\n # clear kv cache\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.actor_module)\n log_gpu_memory_usage(\"After offload actor params and grad during compute_log_prob\", logger=logger)\n aggressive_empty_cache(force_sync=True)\n return output\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def load_checkpoint(self, checkpoint_path, hdfs_path=None, del_local_after_load=True):\n # No checkpoint to load, just offload the model and optimizer to CPU\n if checkpoint_path is None:\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.actor_module)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.actor_optimizer)\n log_gpu_memory_usage(\"After offload actor params and optimizer during load_checkpoint\", logger=logger)\n return\n\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.actor_module)\n self.checkpoint_mananager.load_checkpoint(\n local_path=checkpoint_path, hdfs_path=hdfs_path, del_local_after_load=del_local_after_load\n )\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.actor_module)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.actor_optimizer)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def load_pretrained_model(self, checkpoint_path, del_local_after_load=True):\n pass\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def save_checkpoint(self, checkpoint_path, hdfs_path=None, global_step=0, max_ckpt_to_keep=None):\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.actor_module)\n if self.checkpoint_mananager.checkpoint_config.async_save and self._is_offload_optimizer:\n load_megatron_optimizer(self.actor_optimizer)\n self.checkpoint_mananager.save_checkpoint(\n local_path=checkpoint_path, hdfs_path=hdfs_path, global_step=global_step, max_ckpt_to_keep=max_ckpt_to_keep\n )\n torch.distributed.barrier()\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.actor_module)\n if self.checkpoint_mananager.checkpoint_config.async_save and self._is_offload_optimizer:\n offload_megatron_optimizer(self.actor_optimizer)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def async_calls_finalize_fn_exec(self, blocking=False):\n from megatron.core.dist_checkpointing.strategies.base import async_calls\n\n async_calls.maybe_finalize_async_calls(blocking=blocking)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def start_profile(self, **kwargs) -> None:\n \"\"\"Start profiling for the current rank in the current training step.\"\"\"\n self.profiler.start(**kwargs)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def stop_profile(self) -> None:\n \"\"\"Stop profiling for the current rank in the current training step.\"\"\"\n self.profiler.stop()\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def dump_memory_snapshot(self, tag: str = \"manual\", sub_dir: str = None) -> None:\n \"\"\"Manually trigger a CUDA memory snapshot dump on all ranks.\"\"\"\n # Memory snapshot is now handled by the profiler system\n # This method is kept for backward compatibility but delegates to profiler\n if hasattr(self, \"profiler\") and hasattr(self.profiler, \"_impl\"):\n try:\n # Try to use the profiler's memory snapshot functionality\n if hasattr(self.profiler._impl, \"sampler\"):\n out_dir = OmegaConf.select(self.config, \"actor.profiler.save_path\") or \".\"\n self.profiler._impl.sampler.dump_memory_snapshot(out_dir=out_dir, tag=tag, sub_dir=sub_dir)\n except Exception as e:\n # Log a warning if memory snapshot fails. This might be expected if the profiler doesn't support it.\n logger.warning(f\"Failed to dump memory snapshot: {e}\")\n\n\nclass AsyncActorRolloutRefWorker(ActorRolloutRefWorker):\n @register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=False)\n async def update_weights(self):\n await self.rollout_mode()\n return True\n\n\nclass CriticWorker(MegatronWorker, DistProfilerExtension):\n def __init__(self, config: McoreCriticConfig):\n Worker.__init__(self)\n\n omega_profiler_config = config.get(\"profiler\", {})\n profiler_config = omega_conf_to_dataclass(omega_profiler_config, dataclass_type=ProfilerConfig)\n if omega_profiler_config.get(\"tool\", None) in [\"npu\", \"nsys\", \"torch\", \"torch_memory\"]:\n tool_config = omega_conf_to_dataclass(\n omega_profiler_config.get(\"tool_config\", {}).get(omega_profiler_config.get(\"tool\"))\n )\n else:\n tool_config = None\n DistProfilerExtension.__init__(\n self, DistProfiler(rank=self.rank, config=profiler_config, tool_config=tool_config)\n )\n self.config: McoreCriticConfig = config\n\n # NOTE(sgm): We utilize colocate WorkerGroup by default.\n # As a result, Workers for different model share the same process.\n # Therefore, we only require one distribute initialization.\n # To utilize different parallel strategy in different models:\n # 1, users should disable WorkerDict; 2.assign different ResourcePool to different models,\n # 3. and apply the following patch in ray==2.10, https://github.com/ray-project/ray/pull/44385\n if not torch.distributed.is_initialized():\n set_numa_affinity()\n rank = int(os.environ[\"LOCAL_RANK\"])\n torch.distributed.init_process_group(\n backend=get_nccl_backend(),\n timeout=datetime.timedelta(seconds=self.config.get(\"nccl_timeout\", 600)),\n init_method=os.environ.get(\"DIST_INIT_METHOD\", None),\n )\n get_torch_device().set_device(rank)\n\n mpu.initialize_model_parallel(\n tensor_model_parallel_size=self.config.megatron.tensor_model_parallel_size,\n pipeline_model_parallel_size=self.config.megatron.pipeline_model_parallel_size,\n virtual_pipeline_model_parallel_size=self.config.megatron.virtual_pipeline_model_parallel_size,\n use_sharp=False,\n context_parallel_size=self.config.megatron.context_parallel_size,\n expert_model_parallel_size=self.config.megatron.expert_model_parallel_size,\n expert_tensor_parallel_size=self.config.megatron.expert_tensor_parallel_size,\n nccl_communicator_config_path=None,\n )\n\n is_collect = (\n mpu.get_tensor_model_parallel_rank() == 0\n and mpu.get_pipeline_model_parallel_rank() == mpu.get_pipeline_model_parallel_world_size() - 1\n and mpu.get_context_parallel_rank() == 0\n )\n self._register_dispatch_collect_info(\n mesh_name=\"critic\", dp_rank=mpu.get_data_parallel_rank(), is_collect=is_collect\n )\n\n set_random_seed(seed=self.config.megatron.seed)\n\n # set FSDP offload params\n self._is_offload_param = self.config.megatron.param_offload\n self._is_offload_optimizer = self.config.megatron.optimizer_offload\n\n # normalize config\n self.config.ppo_mini_batch_size *= self.config.rollout_n\n self.config.ppo_mini_batch_size //= mpu.get_data_parallel_world_size()\n if self.config.get(\"ppo_micro_batch_size\", None):\n self.config.ppo_micro_batch_size //= mpu.get_data_parallel_world_size()\n self.config.ppo_micro_batch_size_per_gpu = self.config.ppo_micro_batch_size\n\n # TODO(sgm): support critic model offload\n\n def _build_critic_model_optimizer(\n self, model_path, optim_config, override_model_config, override_transformer_config, override_ddp_config\n ):\n from verl.utils.megatron.optimizer import (\n get_megatron_optimizer,\n get_megatron_optimizer_param_scheduler,\n init_megatron_optim_config,\n )\n from verl.utils.megatron_utils import McoreModuleWrapperConfig, make_megatron_module\n from verl.utils.model import print_model_size\n\n self._init_hf_config_and_tf_config(\n model_path,\n self.config.model.get(\"tokenizer_path\") or model_path,\n self.dtype,\n override_model_config,\n override_transformer_config,\n self.config.model.get(\"trust_remote_code\", False),\n self.config.megatron,\n )\n\n wrap_config = McoreModuleWrapperConfig(\n is_value_model=True, # critic is value model\n share_embeddings_and_output_weights=False,\n wrap_with_ddp=True,\n use_distributed_optimizer=self.config.megatron.use_distributed_optimizer,\n )\n critic_module, updated_tf_config = make_megatron_module(\n wrap_config=wrap_config,\n tf_config=self.tf_config,\n hf_config=self.hf_config,\n bridge=self.bridge,\n provider=self.provider,\n override_model_config=override_model_config,\n override_ddp_config=override_ddp_config,\n peft_cls=self.peft_cls,\n peft_config=self.config.model.get(\"lora\", None),\n )\n self.tf_config = updated_tf_config\n # note that here critic_module will be a list to be compatible with the construction of interleaved pp (vpp).\n # but here, we do not use pp (vpp) yet. For simplicity, we remove the list\n # critic_module = nn.ModuleList(critic_module)\n\n if self.config.load_weight:\n t0 = time.time()\n if self.config.megatron.use_dist_checkpointing:\n load_mcore_dist_weights(\n critic_module,\n self.config.megatron.dist_checkpointing_path,\n is_value_model=True,\n prefix=self.config.megatron.dist_checkpointing_prefix,\n )\n else:\n if self.bridge is not None:\n local_model_path = get_hf_model_path(self.config)\n if self.vanilla_bridge:\n self.bridge.load_weights(critic_module, local_model_path)\n else:\n self.bridge.load_hf_weights(\n critic_module, local_model_path, allowed_mismatched_params=[\"output_layer.weight\"]\n )\n else:\n load_megatron_gptmodel_weights(\n self.config, self.hf_config, critic_module, params_dtype=self.dtype, is_value_model=True\n )\n t1 = time.time()\n if torch.distributed.get_rank() == 0:\n print(f\"critic load_weight time: {t1 - t0}\")\n if self.rank == 0:\n print_model_size(critic_module[0])\n\n # TODO: add more optimizer args into config\n optim_config_megatron = init_megatron_optim_config(\n optim_config,\n use_distributed_optimizer=wrap_config.use_distributed_optimizer,\n fp16=self.dtype == torch.float16,\n )\n critic_optimizer = get_megatron_optimizer(model=critic_module, config=optim_config_megatron)\n critic_optimizer_scheduler = get_megatron_optimizer_param_scheduler(\n optimizer=critic_optimizer, config=optim_config\n )\n get_torch_device().empty_cache()\n\n register_megatron_training_hooks(critic_module, critic_optimizer)\n\n return critic_module, critic_optimizer, critic_optimizer_scheduler, self.hf_config, optim_config\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def init_model(self):\n # create critic\n\n from verl.utils.torch_dtypes import PrecisionType\n\n if self.config.model.get(\"external_lib\", None) is not None:\n # This is used to import external_lib into the huggingface systems\n import importlib\n\n importlib.import_module(self.config.model.external_lib)\n override_model_config = OmegaConf.to_container(OmegaConf.create(self.config.model.get(\"override_config\", {})))\n override_transformer_config = OmegaConf.to_container(\n OmegaConf.create(self.config.megatron.get(\"override_transformer_config\", {}))\n )\n override_ddp_config = OmegaConf.to_container(\n OmegaConf.create(self.config.megatron.get(\"override_ddp_config\", {}))\n )\n self.param_dtype = PrecisionType.to_dtype(self.config.megatron.dtype)\n self.dtype = PrecisionType.to_dtype(self.param_dtype)\n (\n self.critic_module,\n self.critic_optimizer,\n self.critic_optimizer_scheduler,\n self.critic_model_config,\n critic_optimizer_config,\n ) = self._build_critic_model_optimizer(\n model_path=self.config.model.path,\n optim_config=self.config.optim,\n override_model_config=override_model_config,\n override_transformer_config=override_transformer_config,\n override_ddp_config=override_ddp_config,\n )\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.critic_module)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.critic_optimizer)\n\n self.critic = MegatronPPOCritic(\n config=self.config,\n model_config=self.critic_model_config,\n hf_config=self.hf_config,\n tf_config=self.tf_config,\n critic_module=self.critic_module,\n critic_optimizer=self.critic_optimizer,\n critic_optimizer_config=critic_optimizer_config,\n )\n self.flops_counter = FlopsCounter(self.critic_model_config)\n self.checkpoint_mananager = MegatronCheckpointManager(\n config=self.config,\n checkpoint_config=self.config.checkpoint,\n model_config=self.critic_model_config,\n transformer_config=self.tf_config,\n role=\"critic\",\n model=self.critic_module,\n arch=self.architectures[0],\n hf_config=self.hf_config,\n param_dtype=self.param_dtype,\n share_embeddings_and_output_weights=False,\n processing_class=self.processor if self.processor is not None else self.tokenizer,\n optimizer=self.critic_optimizer,\n optimizer_scheduler=self.critic_optimizer_scheduler,\n use_distributed_optimizer=self.config.megatron.use_distributed_optimizer,\n use_checkpoint_opt_param_scheduler=self.config.optim.use_checkpoint_opt_param_scheduler,\n bridge=self.bridge,\n provider=self.provider,\n use_dist_checkpointing=self.config.megatron.use_dist_checkpointing,\n peft_cls=self.peft_cls,\n )\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"critic\"))\n @DistProfiler.annotate(color=\"cyan\", role=\"compute_values\")\n def compute_values(self, data: DataProto):\n micro_batch_size = self.config.ppo_micro_batch_size_per_gpu\n data.meta_info[\"micro_batch_size\"] = micro_batch_size\n data.meta_info[\"max_token_len\"] = self.config.forward_max_token_len_per_gpu\n data.meta_info[\"use_dynamic_bsz\"] = self.config.use_dynamic_bsz\n data = data.to(get_device_id())\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.critic_module)\n values = self.critic.compute_values(data=data)\n output = DataProto.from_dict(tensors={\"values\": values})\n output = output.to(\"cpu\")\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.critic_module)\n return output\n\n @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name=\"critic\"))\n @DistProfiler.annotate(color=\"pink\", role=\"critic_update\")\n def update_critic(self, data: DataProto):\n data = data.to(get_device_id())\n\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.critic_module)\n if self._is_offload_optimizer:\n load_megatron_optimizer(self.critic_optimizer)\n\n dataloader = self.critic.make_minibatch_iterator(data)\n with Timer(name=\"update_critic\", logger=None) as timer:\n metrics = self.critic.update_critic(dataloader=dataloader)\n delta_time = timer.last\n global_num_tokens = data.meta_info[\"global_token_num\"]\n estimated_flops, promised_flops = self.flops_counter.estimate_flops(global_num_tokens, delta_time)\n metrics[\"perf/mfu/critic\"] = estimated_flops * self.config.ppo_epochs / promised_flops / self.world_size\n from verl.utils.megatron.optimizer import get_megatron_last_lr\n\n metrics[\"critic/lr\"] = get_megatron_last_lr(self.critic_optimizer)\n self.critic_optimizer_scheduler.step(1)\n\n output = DataProto(batch=None, meta_info={\"metrics\": metrics})\n\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.critic_module)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.critic_optimizer)\n output = output.to(\"cpu\")\n return output\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def load_checkpoint(self, checkpoint_path, hdfs_path=None, del_local_after_load=True):\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.critic_module)\n self.checkpoint_mananager.load_checkpoint(\n local_path=checkpoint_path, hdfs_path=hdfs_path, del_local_after_load=del_local_after_load\n )\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.critic_module)\n if self._is_offload_optimizer:\n offload_megatron_optimizer(self.critic_optimizer)\n\n @register(dispatch_mode=Dispatch.ONE_TO_ALL)\n def save_checkpoint(self, checkpoint_path, hdfs_path=None, global_steps=0, max_ckpt_to_keep=None):\n if self._is_offload_param:\n load_megatron_model_to_gpu(self.critic_module)\n self.checkpoint_mananager.save_checkpoint(\n local_path=checkpoint_path, hdfs_path=hdfs_path, global_step=global_steps, max_ckpt_to_keep=max_ckpt_to_keep\n )\n if self._is_offload_param:\n offload_megatron_model_to_cpu(self.critic_module)\n"}151{"file_name": "verl__workers__reward_manager__abstract.py", "text": "# Copyright 2023-2025 SGLang Team\n# Copyright Amazon.com, Inc. or its affiliates.\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom abc import ABC, abstractmethod\nfrom typing import Any, Callable\n\nimport torch\n\nfrom verl.protocol import DataProto\n\nRawRewardFn = Callable[..., Any]\n\n\nclass AbstractRewardManager(ABC):\n @abstractmethod\n def __init__(\n self,\n tokenizer: Any,\n num_examine: int,\n compute_score: RawRewardFn | None,\n reward_fn_key: str = \"data_source\",\n **kwargs: Any,\n ):\n pass\n\n @abstractmethod\n def __call__(\n self,\n data: DataProto,\n return_dict: bool = False,\n ) -> torch.Tensor | dict[str, Any]:\n pass\n\n def _extract_reward_from_rm_scores(\n self, data: DataProto, return_dict: bool = False\n ) -> torch.Tensor | dict[str, Any] | None:\n \"\"\"\n Extract reward from already-computed rm_scores if available.\n This has been deprecated.\n\n Args:\n data: DataProto object containing the batch data\n return_dict: Whether to return a dictionary with reward_tensor and reward_extra_info\n\n Returns:\n If rm_scores exists:\n - If return_dict=True: dict with \"reward_tensor\" and \"reward_extra_info\"\n - If return_dict=False: torch.Tensor of rm_scores\n If rm_scores doesn't exist: None\n \"\"\"\n if \"rm_scores\" not in data.batch.keys():\n return None\n\n if return_dict:\n reward_extra_keys = data.meta_info.get(\"reward_extra_keys\", [])\n reward_extra_info = {key: data.non_tensor_batch[key] for key in reward_extra_keys}\n return {\"reward_tensor\": data.batch[\"rm_scores\"], \"reward_extra_info\": reward_extra_info}\n else:\n return data.batch[\"rm_scores\"]\n"}152{"file_name": "verl__workers__reward_manager__batch.py", "text": "# Copyright 2025 Individual Contributor: Mert Unsal\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom collections import defaultdict\nfrom typing import Any\n\nimport torch\n\nfrom verl import DataProto\nfrom verl.workers.reward_manager import register\nfrom verl.workers.reward_manager.abstract import AbstractRewardManager, RawRewardFn\n\n\n@register(\"batch\")\nclass BatchRewardManager(AbstractRewardManager):\n \"\"\"\n A batch reward manager that computes rewards for a batch of data.\n\n Args:\n tokenizer (Tokenizer): The tokenizer to use for decoding the responses.\n num_examine (int): The number of responses to examine.\n compute_score (callable): The function to compute the rewards.\n reward_fn_key (str): The key to use for the reward function.\n reward_kwargs (dict): The keyword arguments to pass to the reward function.\n \"\"\"\n\n def __init__(\n self, tokenizer, num_examine, compute_score: RawRewardFn, reward_fn_key=\"data_source\", **reward_kwargs\n ):\n self.tokenizer = tokenizer\n self.num_examine = num_examine\n self.compute_score = compute_score\n self.reward_fn_key = reward_fn_key\n self.reward_kwargs = reward_kwargs\n\n def verify(self, data):\n prompt_ids = data.batch[\"prompts\"]\n response_ids = data.batch[\"responses\"]\n attention_mask = data.batch[\"attention_mask\"]\n\n prompt_len = prompt_ids.shape[-1]\n valid_response_lengths = attention_mask[:, prompt_len:].sum(dim=-1)\n\n responses_str = []\n for i in range(len(data)):\n valid_len = valid_response_lengths[i]\n valid_response_ids = response_ids[i][:valid_len]\n response_str = self.tokenizer.decode(valid_response_ids, skip_special_tokens=True)\n responses_str.append(response_str)\n\n ground_truths = [item.non_tensor_batch[\"reward_model\"].get(\"ground_truth\", None) for item in data]\n data_sources = data.non_tensor_batch[self.reward_fn_key]\n rollout_reward_scores = data.non_tensor_batch.get(\"reward_scores\", [{} for _ in range(len(data))])\n extras = data.non_tensor_batch.get(\"extra_info\", [{} for _ in range(len(data))])\n\n for i in range(len(data)):\n extras[i][\"rollout_reward_scores\"] = rollout_reward_scores[i]\n\n scores = self.compute_score(\n data_sources=data_sources,\n solution_strs=responses_str,\n ground_truths=ground_truths,\n extra_infos=extras,\n **self.reward_kwargs,\n )\n\n return scores\n\n def __call__(self, data: DataProto, return_dict: bool = False) -> torch.Tensor | dict[str, Any]:\n # If there is rm score, we directly return rm score. Otherwise, we compute via rm_score_fn\n reward_from_rm_scores = self._extract_reward_from_rm_scores(data, return_dict)\n if reward_from_rm_scores is not None:\n return reward_from_rm_scores\n\n reward_tensor = torch.zeros_like(data.batch[\"responses\"], dtype=torch.float32)\n reward_extra_info = defaultdict(list)\n prompt_ids = data.batch[\"prompts\"]\n prompt_len = prompt_ids.shape[-1]\n attention_mask = data.batch[\"attention_mask\"]\n valid_response_lengths = attention_mask[:, prompt_len:].sum(dim=-1)\n data_sources = data.non_tensor_batch[self.reward_fn_key]\n\n scores = self.verify(data)\n rewards = []\n already_printed: dict[str, Any] = {}\n\n for i in range(len(data)):\n length = valid_response_lengths[i].item()\n score = scores[i]\n\n if isinstance(score, dict):\n reward = score[\"score\"]\n for key, value in score.items():\n reward_extra_info[key].append(value)\n else:\n reward = score\n\n rewards.append(reward)\n reward_tensor[i, length - 1] = reward\n\n data_source = data_sources[i]\n if already_printed.get(data_source, 0) < self.num_examine:\n response_str = self.tokenizer.decode(data.batch[\"responses\"][i][:length], skip_special_tokens=True)\n prompt_str = self.tokenizer.decode(data.batch[\"prompts\"][i], skip_special_tokens=True)\n ground_truth = data[i].non_tensor_batch[\"reward_model\"].get(\"ground_truth\", None)\n print(\"[prompt]\", prompt_str)\n print(\"[response]\", response_str)\n print(\"[ground_truth]\", ground_truth)\n print(\"[score]\", scores[i])\n already_printed[data_source] = already_printed.get(data_source, 0) + 1\n\n data.batch[\"acc\"] = torch.tensor(rewards, dtype=torch.float32, device=prompt_ids.device)\n\n if return_dict:\n return {\"reward_tensor\": reward_tensor, \"reward_extra_info\": reward_extra_info}\n else:\n return reward_tensor\n"}153{"file_name": "verl__workers__reward_manager__dapo.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom collections import defaultdict\n\nimport torch\n\nfrom verl import DataProto\nfrom verl.utils.reward_score import default_compute_score\nfrom verl.workers.reward_manager import register\nfrom verl.workers.reward_manager.abstract import AbstractRewardManager\n\n\n@register(\"dapo\")\nclass DAPORewardManager(AbstractRewardManager):\n \"\"\"The reward manager.\"\"\"\n\n def __init__(\n self,\n tokenizer,\n num_examine,\n compute_score=None,\n reward_fn_key=\"data_source\",\n max_resp_len=None,\n overlong_buffer_cfg=None,\n ) -> None:\n self.tokenizer = tokenizer\n self.num_examine = num_examine # the number of batches of decoded responses to print to the console\n self.compute_score = compute_score or default_compute_score\n self.reward_fn_key = reward_fn_key\n self.overlong_buffer_cfg = overlong_buffer_cfg\n self.max_resp_len = max_resp_len\n\n if self.overlong_buffer_cfg is not None:\n assert self.max_resp_len is not None, (\n f\"max_resp_len must be provided if {overlong_buffer_cfg=}, but got None\"\n )\n assert self.max_resp_len >= self.overlong_buffer_cfg.len, (\n \"max_resp_len must be larger than overlong_buffer.len\"\n )\n assert not self.overlong_buffer_cfg.enable or self.overlong_buffer_cfg.len > 0, (\n \"overlong_buffer.len must be positive when overlong penalty is enabled,\"\n f\"but got {self.overlong_buffer_cfg.len}.\"\n \"To disable the overlong penalty, set overlong_buffer.enable = False\"\n )\n\n def __call__(self, data: DataProto, return_dict: bool = False):\n \"\"\"We will expand this function gradually based on the available datasets\"\"\"\n\n # If there is rm score, we directly return rm score. Otherwise, we compute via rm_score_fn\n reward_from_rm_scores = self._extract_reward_from_rm_scores(data, return_dict)\n if reward_from_rm_scores is not None:\n return reward_from_rm_scores\n\n reward_tensor = torch.zeros_like(data.batch[\"responses\"], dtype=torch.float32)\n reward_extra_info = defaultdict(list)\n\n already_print_data_sources = {}\n\n for i in range(len(data)):\n data_item = data[i] # DataProtoItem\n\n prompt_ids = data_item.batch[\"prompts\"]\n\n prompt_length = prompt_ids.shape[-1]\n\n valid_prompt_length = data_item.batch[\"attention_mask\"][:prompt_length].sum()\n valid_prompt_ids = prompt_ids[-valid_prompt_length:]\n\n response_ids = data_item.batch[\"responses\"]\n valid_response_length = data_item.batch[\"attention_mask\"][prompt_length:].sum()\n valid_response_ids = response_ids[:valid_response_length]\n\n # decode\n prompt_str = self.tokenizer.decode(valid_prompt_ids, skip_special_tokens=True)\n response_str = self.tokenizer.decode(valid_response_ids, skip_special_tokens=True)\n eos_token = self.tokenizer.eos_token\n if response_str.endswith(eos_token):\n response_str = response_str[: -len(eos_token)]\n\n ground_truth = data_item.non_tensor_batch[\"reward_model\"][\"ground_truth\"]\n\n data_source = data_item.non_tensor_batch[self.reward_fn_key]\n\n extra_info = data_item.non_tensor_batch.get(\"extra_info\", {})\n\n rollout_reward_scores = data_item.non_tensor_batch.get(\"reward_scores\", {})\n\n extra_info[\"rollout_reward_scores\"] = rollout_reward_scores\n\n result = self.compute_score(\n data_source=data_source,\n solution_str=response_str,\n ground_truth=ground_truth,\n extra_info=extra_info,\n )\n\n score: float\n if isinstance(result, dict):\n score = result[\"score\"]\n # Store the information including original reward\n for key, value in result.items():\n reward_extra_info[key].append(value)\n else:\n score = result\n reward_extra_info[\"acc\"].append(score)\n\n reward = score\n\n if self.overlong_buffer_cfg.enable:\n overlong_buffer_len = self.overlong_buffer_cfg.len\n expected_len = self.max_resp_len - overlong_buffer_len\n exceed_len = valid_response_length - expected_len\n overlong_penalty_factor = self.overlong_buffer_cfg.penalty_factor\n overlong_reward = min(-exceed_len / overlong_buffer_len * overlong_penalty_factor, 0)\n reward += overlong_reward\n if self.overlong_buffer_cfg.log:\n reward_extra_info[\"overlong_reward\"].append(overlong_reward)\n reward_extra_info[\"overlong\"].append(overlong_reward < 0)\n\n reward_tensor[i, valid_response_length - 1] = reward\n\n if data_source not in already_print_data_sources:\n already_print_data_sources[data_source] = 0\n\n if already_print_data_sources[data_source] < self.num_examine:\n already_print_data_sources[data_source] += 1\n print(\"[prompt]\", prompt_str)\n print(\"[response]\", response_str)\n print(\"[ground_truth]\", ground_truth)\n if isinstance(result, dict):\n for key, value in result.items():\n print(f\"[{key}]\", value)\n else:\n print(\"[score]\", score)\n\n if return_dict:\n return {\n \"reward_tensor\": reward_tensor,\n \"reward_extra_info\": reward_extra_info,\n }\n else:\n return reward_tensor\n"}154{"file_name": "verl__workers__reward_manager__naive.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom collections import defaultdict\nfrom typing import Any\n\nimport torch\n\nfrom verl import DataProto\nfrom verl.utils.reward_score import default_compute_score\nfrom verl.workers.reward_manager import register\nfrom verl.workers.reward_manager.abstract import AbstractRewardManager\n\n\n@register(\"naive\")\nclass NaiveRewardManager(AbstractRewardManager):\n \"\"\"The reward manager.\"\"\"\n\n def __init__(self, tokenizer, num_examine, compute_score=None, reward_fn_key=\"data_source\") -> None:\n \"\"\"\n Initialize the NaiveRewardManager instance.\n\n Args:\n tokenizer: The tokenizer used to decode token IDs into text.\n num_examine: The number of batches of decoded responses to print to the console for debugging purpose.\n compute_score: A function to compute the reward score. If None, `default_compute_score` will be used.\n reward_fn_key: The key used to access the data source in the non-tensor batch data. Defaults to\n \"data_source\".\n \"\"\"\n self.tokenizer = tokenizer # Store the tokenizer for decoding token IDs\n self.num_examine = num_examine # the number of batches of decoded responses to print to the console\n self.compute_score = compute_score or default_compute_score\n self.reward_fn_key = reward_fn_key # Store the key for accessing the data source\n\n def __call__(self, data: DataProto, return_dict: bool = False) -> torch.Tensor | dict[str, Any]:\n \"\"\"We will expand this function gradually based on the available datasets\"\"\"\n\n # If there is rm score, we directly return rm score. Otherwise, we compute via rm_score_fn\n reward_from_rm_scores = self._extract_reward_from_rm_scores(data, return_dict)\n if reward_from_rm_scores is not None:\n return reward_from_rm_scores\n\n reward_tensor = torch.zeros_like(data.batch[\"responses\"], dtype=torch.float32)\n reward_extra_info = defaultdict(list)\n\n already_print_data_sources = {}\n\n for i in range(len(data)):\n data_item = data[i] # DataProtoItem\n\n prompt_ids = data_item.batch[\"prompts\"]\n\n prompt_length = prompt_ids.shape[-1]\n\n valid_prompt_length = data_item.batch[\"attention_mask\"][:prompt_length].sum()\n valid_prompt_ids = prompt_ids[-valid_prompt_length:]\n\n response_ids = data_item.batch[\"responses\"]\n valid_response_length = data_item.batch[\"attention_mask\"][prompt_length:].sum()\n valid_response_ids = response_ids[:valid_response_length]\n\n # decode\n prompt_str = self.tokenizer.decode(valid_prompt_ids, skip_special_tokens=True)\n response_str = self.tokenizer.decode(valid_response_ids, skip_special_tokens=True)\n\n ground_truth = data_item.non_tensor_batch[\"reward_model\"][\"ground_truth\"]\n data_source = data_item.non_tensor_batch[self.reward_fn_key]\n extra_info = data_item.non_tensor_batch.get(\"extra_info\", {})\n num_turns = data_item.non_tensor_batch.get(\"__num_turns__\", None)\n rollout_reward_scores = data_item.non_tensor_batch.get(\"reward_scores\", {})\n extra_info[\"num_turns\"] = num_turns\n extra_info[\"rollout_reward_scores\"] = rollout_reward_scores\n\n score = self.compute_score(\n data_source=data_source,\n solution_str=response_str,\n ground_truth=ground_truth,\n extra_info=extra_info,\n )\n\n if isinstance(score, dict):\n reward = score[\"score\"]\n # Store the information including original reward\n for key, value in score.items():\n reward_extra_info[key].append(value)\n else:\n reward = score\n\n reward_tensor[i, valid_response_length - 1] = reward\n\n if data_source not in already_print_data_sources:\n already_print_data_sources[data_source] = 0\n\n if already_print_data_sources[data_source] < self.num_examine:\n already_print_data_sources[data_source] += 1\n print(\"[prompt]\", prompt_str)\n print(\"[response]\", response_str)\n print(\"[ground_truth]\", ground_truth)\n if isinstance(score, dict):\n for key, value in score.items():\n print(f\"[{key}]\", value)\n else:\n print(\"[score]\", score)\n\n if return_dict:\n return {\n \"reward_tensor\": reward_tensor,\n \"reward_extra_info\": reward_extra_info,\n }\n else:\n return reward_tensor\n"}155{"file_name": "verl__workers__reward_manager__prime.py", "text": "# Copyright 2024 PRIME team and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport asyncio\nfrom concurrent.futures import ProcessPoolExecutor\nfrom functools import partial\nfrom typing import Any, Callable, Optional\n\nimport psutil\nimport torch\nfrom transformers import PreTrainedTokenizer\n\nfrom verl import DataProto\nfrom verl.utils.ray_utils import get_event_loop\nfrom verl.utils.reward_score import default_compute_score\nfrom verl.workers.reward_manager import register\nfrom verl.workers.reward_manager.abstract import AbstractRewardManager\n\n\nasync def single_compute_score(evaluation_func, completion, reference, task, task_extra_info, executor, timeout=300.0):\n loop = get_event_loop()\n try:\n # Ensure process_completion is called properly\n future = loop.run_in_executor(executor, partial(evaluation_func, task, completion, reference, task_extra_info))\n return await asyncio.wait_for(future, timeout=timeout)\n except asyncio.TimeoutError:\n print(f\"[Timeout] Task timeout: {completion}\")\n return None # Default value for timed-out rows\n except Exception as e:\n print(f\"[Error] Task failed: {e}, completion: {completion[:80]}\")\n return None # Default value for failed rows\n\n\nasync def parallel_compute_score_async(\n evaluation_func, completions, references, tasks, extra_info=None, num_processes=64\n):\n if extra_info is None:\n extra_info = [None] * len(tasks)\n scores = []\n with ProcessPoolExecutor(max_workers=num_processes) as executor:\n # to prevent very occasional starvation caused by some anomalous programs ( like infinite loop ), the\n # exceptions in async programs will instantly halt the evaluation, and all summoned processes will be killed.\n try:\n # Create tasks for all rows\n tasks_async = [\n single_compute_score(evaluation_func, c, r, t, ei, executor, timeout=300.0)\n for c, r, t, ei in zip(completions, references, tasks, extra_info, strict=True)\n ]\n results = await asyncio.gather(*tasks_async, return_exceptions=False)\n except Exception as e:\n print(f\"[Exception] async gather failed: {e}\")\n raise\n finally:\n terminated_count = 0\n for pid, proc in executor._processes.items():\n try:\n p = psutil.Process(pid)\n p.terminate()\n try:\n p.wait(timeout=5)\n except psutil.TimeoutExpired:\n p.kill()\n terminated_count += 1\n except Exception:\n pass\n print(f\"[Shutdown] {terminated_count} subprocess(es) terminated.\")\n\n # Process results\n for result, completion, reference, task in zip(results, completions, references, tasks, strict=True):\n if isinstance(result, Exception) or result is None:\n # Handle failed or timed-out tasks\n scores.append(0.0)\n elif isinstance(result, int | float | bool):\n scores.append(float(result))\n else:\n scores.append(float(result[0]))\n return scores\n\n\ndef run_reward_scoring(evaluation_func, completions, references, tasks, extra_info=None, num_processes=64):\n loop = asyncio.new_event_loop()\n asyncio.set_event_loop(loop)\n try:\n return loop.run_until_complete(\n parallel_compute_score_async(evaluation_func, completions, references, tasks, extra_info, num_processes)\n )\n finally:\n loop.close()\n\n\n@register(\"prime\")\nclass PrimeRewardManager(AbstractRewardManager):\n \"\"\"\n The Reward Manager used in https://github.com/PRIME-RL/PRIME\n \"\"\"\n\n def __init__(\n self,\n tokenizer: PreTrainedTokenizer,\n num_examine: int,\n compute_score: Optional[Callable] = None,\n reward_fn_key: str = \"data_source\",\n ) -> None:\n self.tokenizer = tokenizer\n self.num_examine = num_examine # the number of batches of decoded responses to print to the console\n self.compute_score = compute_score or default_compute_score\n self.reward_fn_key = reward_fn_key\n\n def verify(self, data):\n \"\"\"\n verify the batch and save as ``acc`` tensor\n \"\"\"\n # batched scoring\n prompt_ids = data.batch[\"prompts\"]\n\n response_ids = data.batch[\"responses\"]\n sequences_str = self.tokenizer.batch_decode(response_ids, skip_special_tokens=True)\n ground_truth = [data_item.non_tensor_batch[\"reward_model\"][\"ground_truth\"] for data_item in data]\n data_sources = data.non_tensor_batch[self.reward_fn_key]\n extra_info = data.non_tensor_batch.get(\"extra_info\", None)\n\n assert len(sequences_str) == len(ground_truth) == len(data_sources)\n try:\n scores = run_reward_scoring(\n self.compute_score,\n completions=sequences_str,\n references=ground_truth,\n tasks=data_sources,\n extra_info=extra_info,\n num_processes=64,\n )\n except asyncio.TimeoutError:\n print(\"[Timeout] Global reward scoring timed out. Setting all as 0.\")\n scores = [0.0 for _ in range(len(sequences_str))]\n except Exception as e:\n print(f\"[Error] Unexpected error during scoring. Setting all as 0. {e}\")\n scores = [0.0 for _ in range(len(sequences_str))]\n data.batch[\"acc\"] = torch.tensor(scores, dtype=torch.float32, device=prompt_ids.device)\n return scores\n\n def __call__(self, data: DataProto, return_dict: bool = False) -> torch.Tensor | dict[str, Any]:\n \"\"\"We will expand this function gradually based on the available datasets\"\"\"\n\n # If there is rm score, we directly return rm score. Otherwise, we compute via rm_score_fn\n reward_from_rm_scores = self._extract_reward_from_rm_scores(data, return_dict)\n if reward_from_rm_scores is not None:\n return reward_from_rm_scores\n\n reward_tensor = torch.zeros_like(data.batch[\"responses\"], dtype=torch.float32)\n\n already_print_data_sources = {}\n\n # batched scoring\n prompt_ids = data.batch[\"prompts\"]\n prompt_length = prompt_ids.shape[-1]\n\n response_ids = data.batch[\"responses\"]\n valid_response_length = data.batch[\"attention_mask\"][:, prompt_length:].sum(dim=-1)\n sequences_str = self.tokenizer.batch_decode(response_ids, skip_special_tokens=True)\n data_sources = data.non_tensor_batch[\"data_source\"]\n\n scores = self.verify(data)\n\n for i in range(len(data)):\n data_source = data_sources[i]\n reward_tensor[i, valid_response_length[i].item() - 1] = scores[i]\n\n if data_source not in already_print_data_sources:\n already_print_data_sources[data_source] = 0\n\n if already_print_data_sources[data_source] < self.num_examine:\n already_print_data_sources[data_source] += 1\n print(sequences_str)\n\n if return_dict:\n return {\"reward_tensor\": reward_tensor}\n else:\n return reward_tensor\n"}156{"file_name": "verl__workers__reward_manager__registry.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nfrom typing import Callable\n\nfrom verl.workers.reward_manager.abstract import AbstractRewardManager\n\n__all__ = [\"register\", \"get_reward_manager_cls\"]\n\nREWARD_MANAGER_REGISTRY: dict[str, type[AbstractRewardManager]] = {}\n\n\ndef register(name: str) -> Callable[[type[AbstractRewardManager]], type[AbstractRewardManager]]:\n \"\"\"Decorator to register a reward manager class with a given name.\n\n Args:\n name: `(str)`\n The name of the reward manager.\n \"\"\"\n\n def decorator(cls: type[AbstractRewardManager]) -> type[AbstractRewardManager]:\n if name in REWARD_MANAGER_REGISTRY and REWARD_MANAGER_REGISTRY[name] != cls:\n raise ValueError(\n f\"Reward manager {name} has already been registered: {REWARD_MANAGER_REGISTRY[name]} vs {cls}\"\n )\n REWARD_MANAGER_REGISTRY[name] = cls\n return cls\n\n return decorator\n\n\ndef get_reward_manager_cls(name: str) -> type[AbstractRewardManager]:\n \"\"\"Get the reward manager class with a given name.\n\n Args:\n name: `(str)`\n The name of the reward manager.\n\n Returns:\n `(type)`: The reward manager class.\n \"\"\"\n if name not in REWARD_MANAGER_REGISTRY:\n raise ValueError(f\"Unknown reward manager: {name}\")\n return REWARD_MANAGER_REGISTRY[name]\n"}157{"file_name": "verl__workers__rollout__base.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport importlib\nfrom abc import ABC, abstractmethod\nfrom typing import Generator\n\nimport torch\nfrom torch.distributed.device_mesh import DeviceMesh\n\nfrom verl import DataProto\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.workers.config import HFModelConfig, RolloutConfig\n\n__all__ = [\"BaseRollout\"]\n\n\nclass BaseRollout(ABC):\n \"\"\"Base class for rollout.\"\"\"\n\n def __init__(\n self,\n config: RolloutConfig,\n model_config: HFModelConfig,\n device_mesh: DeviceMesh,\n ):\n self.config = omega_conf_to_dataclass(config)\n self.model_config: HFModelConfig = omega_conf_to_dataclass(model_config, dataclass_type=HFModelConfig)\n self.device_mesh = device_mesh\n\n @abstractmethod\n async def resume(self, tags: list[str]):\n \"\"\"Resume rollout weights or kv cache in GPU memory.\n\n Args:\n tags: weights or kv_cache.\n \"\"\"\n pass\n\n @abstractmethod\n async def update_weights(\n self,\n weights: Generator[tuple[str, torch.Tensor], None, None],\n **kwargs,\n ):\n \"\"\"Update the weights of the rollout model.\n\n Args:\n weights: A generator that yields the name of the weight tensor and the tensor itself.\n \"\"\"\n pass\n\n @abstractmethod\n async def release(self):\n \"\"\"Release weights and kv cache in GPU memory.\"\"\"\n pass\n\n def generate_sequences(self, prompts: DataProto) -> DataProto:\n \"\"\"Batch generate sequences in sync mode.\n\n Args:\n prompts: The input prompts.\n\n Returns:\n The output sequences.\n \"\"\"\n raise NotImplementedError\n\n\n_ROLLOUT_REGISTRY = {\n (\"vllm\", \"async\"): \"verl.workers.rollout.vllm_rollout.ServerAdapter\",\n (\"sglang\", \"async\"): \"verl.workers.rollout.sglang_rollout.sglang_rollout.ServerAdapter\",\n (\"trtllm\", \"async\"): \"verl.workers.rollout.trtllm_rollout.trtllm_rollout.ServerAdapter\",\n}\n\n\ndef get_rollout_class(rollout_name: str, mode: str = \"async\") -> type[BaseRollout]:\n \"\"\"Get the rollout class by name.\n\n Args:\n rollout_name: The name of the rollout.\n mode: The mode of the rollout, async: server mode.\n\n Returns:\n The rollout class.\n \"\"\"\n assert (rollout_name, mode) in _ROLLOUT_REGISTRY, f\"Rollout {rollout_name} with mode {mode} not found\"\n fqdn = _ROLLOUT_REGISTRY[(rollout_name, mode)]\n module_name, class_name = fqdn.rsplit(\".\", 1)\n rollout_module = importlib.import_module(module_name)\n return getattr(rollout_module, class_name)\n"}158{"file_name": "verl__workers__rollout__hf_rollout.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nRollout with huggingface models.\nTODO: refactor this class. Currently, it will hang when using FSDP HybridShard. We should actually create a single\nGPU model. Then, get full state_dict and bind the state_dict to the single GPU model. Then, use the single GPU model\nto perform generation.\n\"\"\"\n\nimport contextlib\n\nimport torch\nimport torch.distributed\nfrom tensordict import TensorDict\nfrom torch import nn\nfrom torch.distributed.fsdp import FullyShardedDataParallel as FSDP\nfrom transformers import GenerationConfig\n\nfrom verl import DataProto\nfrom verl.utils.device import get_device_name, get_torch_device\nfrom verl.utils.torch_functional import get_response_mask\n\nfrom .base import BaseRollout\n\n__all__ = [\"HFRollout\"]\n\n\nclass HFRollout(BaseRollout):\n def __init__(self, module: nn.Module, config):\n super().__init__()\n self.config = config\n self.module = module\n\n def generate_sequences(self, prompts: DataProto) -> DataProto:\n batch_size = prompts.batch.batch_size[0]\n num_chunks = max(batch_size // self.config.get(\"micro_batch_size\", batch_size), 1)\n batch_prompts = prompts.chunk(chunks=num_chunks)\n output = [self._generate_minibatch(p) for p in batch_prompts]\n output = DataProto.concat(output)\n return output\n\n @torch.no_grad()\n def _generate_minibatch(self, prompts: DataProto) -> DataProto:\n # make sampling args can be overridden by inputs\n do_sample = prompts.meta_info.get(\"do_sample\", self.config.do_sample)\n is_validate = prompts.meta_info.get(\"validate\", False)\n\n temperature = prompts.meta_info.get(\"temperature\", self.config.temperature)\n response_length = prompts.meta_info.get(\"response_length\", self.config.response_length)\n top_p = prompts.meta_info.get(\"top_p\", self.config.get(\"top_p\", 1.0))\n top_k = max(0, prompts.meta_info.get(\"top_k\", self.config.get(\"top_k\", 0))) # to be compatible with vllm\n\n if not do_sample:\n # do_sample==False -> greedy decoding\n kwargs = {\n \"do_sample\": False,\n \"num_beams\": 1,\n }\n elif is_validate:\n # do validate and do sample -> use val_kwargs\n kwargs = {\n \"do_sample\": True,\n \"num_beams\": 1,\n \"top_k\": max(0, self.config.val_kwargs.top_k), # to be compatible with vllm\n \"top_p\": self.config.val_kwargs.top_p,\n \"temperature\": self.config.val_kwargs.temperature,\n \"num_return_sequences\": 1, # if validate, already repeat in ray_trainer\n }\n else:\n # do_sample -> use rollout config\n kwargs = {\n \"do_sample\": True,\n \"num_beams\": 1,\n \"top_p\": top_p,\n \"top_k\": top_k,\n \"temperature\": temperature,\n # already repeat in ray_trainer\n # https://github.com/volcengine/verl/blob/2fdfbdcba6f2e076f64bc47922d8fe6cf7dc7da5/verl/trainer/ppo/ray_trainer.py#L1117\n \"num_return_sequences\": 1,\n }\n\n # make config according to generate mode\n generation_config = GenerationConfig(**kwargs)\n\n idx = prompts.batch[\"input_ids\"] # (bs, prompt_length)\n prompt_length = idx.size(1)\n attention_mask = prompts.batch[\"attention_mask\"] # left-padded attention_mask\n position_ids = prompts.batch[\"position_ids\"]\n\n # used to construct attention_mask\n eos_token_id = prompts.meta_info[\"eos_token_id\"]\n pad_token_id = prompts.meta_info[\"pad_token_id\"]\n\n self.module.eval()\n param_ctx = contextlib.nullcontext()\n\n if isinstance(self.module, FSDP):\n # recurse need to set to False according to https://github.com/pytorch/pytorch/issues/100069\n param_ctx = FSDP.summon_full_params(self.module, writeback=False, recurse=False)\n with param_ctx, torch.autocast(device_type=get_device_name(), dtype=torch.bfloat16):\n output = self.module.generate(\n input_ids=idx,\n attention_mask=attention_mask,\n position_ids=position_ids,\n do_sample=do_sample,\n max_new_tokens=response_length,\n eos_token_id=eos_token_id,\n pad_token_id=pad_token_id,\n generation_config=generation_config,\n output_scores=False, # this is potentially very large\n return_dict_in_generate=True,\n use_cache=True,\n )\n\n # TODO: filter out the seq with no answers like ds-chat\n seq = output.sequences\n generated_batch_size = seq.size(0) # bs * num_return_sequences\n\n # huggingface generate will stop generating when all the batch reaches [EOS].\n # We have to pad to response_length\n sequence_length = prompt_length + self.config.response_length\n delta_length = sequence_length - seq.shape[1]\n\n if delta_length > 0:\n delta_tokens = torch.ones(size=(generated_batch_size, delta_length), device=seq.device, dtype=seq.dtype)\n delta_tokens = pad_token_id * delta_tokens\n seq = torch.cat((seq, delta_tokens), dim=1)\n assert seq.shape[1] == sequence_length\n\n # make necessary reputations if num_return_sequences > 1\n num_return_sequences = kwargs.get(\"num_return_sequences\", 1)\n if num_return_sequences > 1:\n position_ids = position_ids.repeat_interleave(num_return_sequences, dim=0)\n attention_mask = attention_mask.repeat_interleave(num_return_sequences, dim=0)\n\n prompt = seq[:, :prompt_length] # (generated_batch_size, prompt_length)\n response = seq[:, prompt_length:] # (generated_batch_size, response_length)\n\n response_length = response.size(1)\n delta_position_id = torch.arange(1, response_length + 1, device=position_ids.device)\n delta_position_id = delta_position_id.unsqueeze(0).repeat(generated_batch_size, 1)\n\n response_position_ids = position_ids[:, -1:] + delta_position_id\n position_ids = torch.cat([position_ids, response_position_ids], dim=-1)\n\n response_attention_mask = get_response_mask(\n response_id=response, eos_token=eos_token_id, dtype=attention_mask.dtype\n )\n attention_mask = torch.cat((attention_mask, response_attention_mask), dim=-1)\n\n batch = TensorDict(\n {\n \"prompts\": prompt,\n \"responses\": response,\n \"input_ids\": seq,\n \"attention_mask\": attention_mask,\n \"position_ids\": position_ids,\n },\n batch_size=generated_batch_size,\n )\n\n # empty cache before compute old_log_prob\n get_torch_device().empty_cache()\n\n self.module.train()\n return DataProto(batch=batch)\n"}159{"file_name": "verl__workers__rollout__naive__naive_rollout.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nIn single GPU rollout, the sequences are generated directly by sampling from the model.\nThe output will contain\n1. output_ids\n2. attention_masks (left padding)\n3. eos_masks\n4. log_probs\n\"\"\"\n\nimport torch\nimport torch.nn.functional as F\nfrom tensordict import TensorDict\nfrom torch import nn\n\nfrom verl import DataProto\nfrom verl.utils.torch_functional import logprobs_from_logits\n\nfrom ..base import BaseRollout\n\n__all__ = [\"NaiveRollout\"]\n\n\nclass NaiveRollout(BaseRollout):\n def __init__(self, module: nn.Module, config):\n \"\"\"A naive rollout. It requires the module to be compatible with huggingface APIs. That is:\n The module should define __call__ to receive input_ids, attention_mask and position_ids.\n It outputs a structure that contains logits field.\n\n Args:\n module: module here follows huggingface APIs\n config: DictConfig\n \"\"\"\n super().__init__()\n self.config = config\n self.module = module\n\n @torch.no_grad()\n def generate_sequences(self, prompts: DataProto) -> DataProto:\n \"\"\"Generate sequences\"\"\"\n idx = prompts.batch[\"input_ids\"] # (bs, prompt_length)\n attention_mask = prompts.batch[\"attention_mask\"] # left-padded attention_mask\n position_ids = prompts.batch[\"position_ids\"]\n\n # used to construct attention_mask\n eos_token_id = prompts.meta_info[\"eos_token_id\"]\n\n batch_size = idx.size(0)\n prompt_length = idx.size(1)\n\n self.module.eval()\n\n prev_attention_mask = torch.ones(size=(batch_size, 1), dtype=attention_mask.dtype, device=attention_mask.device)\n\n logits_lst = []\n for _ in range(self.config.response_length):\n # if the sequence context is growing too long we must crop it at block_size\n # idx_cond = idx if idx.size(1) <= self.config.block_size else idx[:, -self.config.block_size:]\n idx_cond = idx\n # forward the model to get the logits for the index in the sequence\n # we use huggingface APIs here\n output = self.module(input_ids=idx_cond, attention_mask=attention_mask, position_ids=position_ids)\n logits = output.logits\n # pluck the logits at the final step and scale by desired temperature\n logits = logits[:, -1, :] / self.config.temperature # (bs, vocab_size)\n # optionally crop the logits to only the top k options\n if self.config.top_k is not None:\n v, _ = torch.topk(logits, min(self.config.top_k, logits.size(-1)))\n logits[logits < v[:, [-1]]] = -float(\"Inf\")\n # apply softmax to convert logits to (normalized) probabilities\n probs = F.softmax(logits, dim=-1)\n # sample from the distribution\n if self.config.do_sample:\n idx_next = torch.multinomial(probs, num_samples=1)\n else:\n idx_next = torch.argmax(probs, dim=-1, keepdim=True)\n\n attention_mask = torch.cat((attention_mask, prev_attention_mask), dim=-1)\n\n for token_id in eos_token_id:\n prev_attention_mask = torch.logical_and(idx_next != token_id, prev_attention_mask.bool())\n prev_attention_mask.to(attention_mask.dtype)\n\n position_ids = torch.cat((position_ids, position_ids[:, -1:] + 1), dim=-1)\n\n # append sampled index to the running sequence and continue\n idx = torch.cat((idx, idx_next), dim=1)\n logits_lst.append(logits)\n\n logits = torch.stack(logits_lst, dim=1) # (bs, response_length, vocab_size)\n prompts = idx[:, :prompt_length] # (bs, prompt_length)\n response = idx[:, prompt_length:] # (bs, response_length)\n log_probs = logprobs_from_logits(logits=logits, labels=response)\n batch = TensorDict(\n {\n \"input_ids\": prompts,\n \"responses\": response,\n \"sequences\": idx,\n \"old_log_probs\": log_probs,\n \"attention_mask\": attention_mask,\n \"position_ids\": position_ids,\n },\n batch_size=batch_size,\n )\n\n self.module.train()\n\n return DataProto(batch=batch)\n"}160{"file_name": "verl__workers__rollout__replica.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport asyncio\nimport logging\nimport os\nfrom abc import ABC, abstractmethod\nfrom enum import Enum\nfrom typing import Any, Callable, Optional\n\nfrom omegaconf import DictConfig\nfrom pydantic import BaseModel\nfrom ray.actor import ActorHandle\n\nfrom verl.single_controller.ray import RayClassWithInitArgs, RayResourcePool, RayWorkerGroup, ResourcePoolManager\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.utils.device import is_torch_npu_available\nfrom verl.workers.config import HFModelConfig, RolloutConfig\n\nlogger = logging.getLogger(__file__)\n\n\nclass TokenOutput(BaseModel):\n token_ids: list[int]\n \"\"\"response token ids\"\"\"\n log_probs: Optional[list[float]] = None\n \"\"\"logprobs of response token ids\"\"\"\n routed_experts: Optional[Any] = None\n \"\"\"routed experts of response token ids\"\"\"\n stop_reason: Optional[str] = None\n \"\"\"stop reason: 'completed', 'aborted', or None for unknown\"\"\"\n num_preempted: Optional[int] = None\n \"\"\"number of preempted times for metric calculation\"\"\"\n\n\nclass RolloutMode(Enum):\n # Rollout engine and training engine(fsdp/megatron) fused in same process\n # Rollout and trainer share GPUs, switch context with weight synchronization.\n # Usage scenarios: on-policy training.\n HYBRID = \"hybrid\"\n\n # Rollout engine colocated with hybrid engine in same ray placement group but in separate process.\n # Rollout and hybrid processes share GPUs, switch context without weight synchronization.\n # Usage scenarios: GRM (LLM as a judge).\n COLOCATED = \"colocated\"\n\n # Standalone rollout server with separate GPU resource, disaggregated architecture.\n # Usage scenarios: off-policy training.\n STANDALONE = \"standalone\"\n\n\nclass RolloutReplica(ABC):\n \"\"\"Rollout replica is an individual server instance, which may be deployed on single or multiple nodes.\n It is equivalent to launch server in each node with command line:\n\n SGLang:\n ```\n python -m sglang.launch_server --node-rank 0 --nnode 2 ...\n python -m sglang.launch_server --node-rank 1 --nnode 2 ...\n ```\n\n vLLM:\n ```\n vllm serve --data-parallel-size 16 --data-parallel-size-local 8 --data-parallel-start-rank 0 ...\n vllm serve --data-parallel-size 16 --data-parallel-size-local 8 --data-parallel-start-rank 8 ...\n ```\n\n Args:\n replica_rank: int, rank of this rollout replica.\n config: RolloutConfig, full config.\n model_config: DictConfig, model config.\n gpus_per_node: int, number of gpus per node.\n \"\"\"\n\n def __init__(\n self,\n replica_rank: int,\n config: RolloutConfig,\n model_config: DictConfig,\n gpus_per_node: int = 8,\n is_reward_model: bool = False,\n ) -> None:\n self.replica_rank = replica_rank\n self.config = omega_conf_to_dataclass(config)\n self.model_config: HFModelConfig = model_config\n\n self.world_size = (\n self.config.tensor_model_parallel_size\n * self.config.data_parallel_size\n * self.config.pipeline_model_parallel_size\n )\n self.gpus_per_node = gpus_per_node\n self.gpus_per_replica_node = min(gpus_per_node, self.world_size)\n assert self.world_size % self.gpus_per_replica_node == 0, (\n f\"world_size {self.world_size} must be divisible by gpus_per_node {self.gpus_per_replica_node}\"\n )\n self.nnodes = self.world_size // self.gpus_per_replica_node\n self.is_reward_model = is_reward_model\n\n self.rollout_mode: RolloutMode = None\n self.workers: list[ActorHandle] = []\n self.resource_pool: RayResourcePool = None\n self.bundle_indices: list[int] = []\n\n self.servers: list[ActorHandle] = []\n self._server_address: str = None\n self._server_handle: ActorHandle = None\n\n async def init_hybrid(self, worker_group: RayWorkerGroup):\n \"\"\"Init hybrid rollout server, rollout engine and training engine(fsdp/megatron) fused in same process.\n\n Args:\n worker_group: RayWorkerGroup, fused workers where training engine(fsdp/megatron) have been initialized.\n \"\"\"\n self.rollout_mode = RolloutMode.HYBRID\n self.workers = worker_group.workers[\n self.world_size * self.replica_rank : self.world_size * (self.replica_rank + 1)\n ]\n await self.launch_servers()\n\n async def init_hybrid_colocated(self, worker_group: RayWorkerGroup, resource_pool: RayResourcePool):\n \"\"\"Init hybrid rollout server, rollout engine and training engine(fsdp/megatron) fused in same process.\n\n Args:\n worker_group: RayWorkerGroup, fused workers where training engine(fsdp/megatron) have been initialized.\n resource_pool: RayResourcePool, ray placement group where hybrid engine processes have been launched.\n bundle_indices: list[int], bundle indices for this rollout replica.\n \"\"\"\n self.rollout_mode = RolloutMode.HYBRID\n self.workers = worker_group.workers[\n self.world_size * self.replica_rank : self.world_size * (self.replica_rank + 1)\n ]\n self.resource_pool = resource_pool\n self.bundle_indices = [self.replica_rank * self.world_size + idx for idx in range(self.world_size)]\n await self.launch_servers()\n\n # TODO(sgm): this should be the default solution, but need to make the RolloutMode more clear.\n async def init_colocated(self, resource_pool: RayResourcePool):\n \"\"\"Init colocated rollout server, rollout engine and hybrid engine colocated in same ray placement group\n but in separate processes.\n\n Args:\n resource_pool: RayResourcePool, ray placement group where hybrid engine processes have been launched.\n \"\"\"\n self.rollout_mode = RolloutMode.COLOCATED\n self.resource_pool = resource_pool\n use_gpu = self.rollout_worker_use_gpu()\n\n worker_group = RayWorkerGroup(\n resource_pool=self.resource_pool,\n ray_cls_with_init=self.get_ray_class_with_init_args(),\n bin_pack=False,\n name_prefix=f\"rollout_colocate_{self.replica_rank}\"\n if not self.is_reward_model\n else f\"rollout_reward_colocate_{self.replica_rank}\",\n use_gpu=use_gpu,\n device_name=\"cuda\" if not is_torch_npu_available(check_device=False) else \"npu\",\n )\n self.workers = worker_group.workers\n await self.launch_servers()\n\n async def init_standalone(self):\n \"\"\"Init standalone rollout server, create new resource pool for this rollout.\"\"\"\n # create resource pool for this rollout\n self.rollout_mode = RolloutMode.STANDALONE\n resource_pool_name = (\n f\"rollout_pool_{self.replica_rank}\"\n if not self.is_reward_model\n else f\"rollout_pool_reward_{self.replica_rank}\"\n )\n resource_pool_spec = {\n resource_pool_name: [self.gpus_per_replica_node] * self.nnodes,\n }\n resource_pool_manager = ResourcePoolManager(resource_pool_spec=resource_pool_spec, mapping=None)\n resource_pool_manager.create_resource_pool()\n self.resource_pool = resource_pool_manager.resource_pool_dict[resource_pool_name]\n\n # create worker group for this rollout\n use_gpu = self.rollout_worker_use_gpu()\n worker_group = RayWorkerGroup(\n resource_pool=self.resource_pool,\n ray_cls_with_init=self.get_ray_class_with_init_args(),\n bin_pack=False,\n name_prefix=f\"rollout_standalone_{self.replica_rank}\"\n if not self.is_reward_model\n else f\"rollout_reward_standalone_{self.replica_rank}\",\n use_gpu=use_gpu,\n device_name=\"cuda\" if not is_torch_npu_available(check_device=False) else \"npu\",\n )\n self.workers = worker_group.workers\n await self.launch_servers()\n\n @abstractmethod\n def get_ray_class_with_init_args(self) -> RayClassWithInitArgs:\n \"\"\"Get rollout worker actor class for colocated and standalone mode.\"\"\"\n raise NotImplementedError\n\n @abstractmethod\n async def launch_servers(self):\n \"\"\"Launch http server in each node.\"\"\"\n raise NotImplementedError\n\n @property\n def server_address(self) -> str:\n \"\"\"Get rollout server address for OpenAI chat completion.\"\"\"\n return self._server_address\n\n @property\n def server_handle(self) -> ActorHandle:\n \"\"\"Get rollout server handle for Token-in-token-out generation.\"\"\"\n return self._server_handle\n\n def rollout_worker_use_gpu(self) -> bool:\n return True\n\n async def wake_up(self):\n \"\"\"Wake up each rollout server.\"\"\"\n await asyncio.gather(*[server.wake_up.remote() for server in self.servers])\n\n async def sleep(self):\n \"\"\"Sleep each rollout server.\"\"\"\n await asyncio.gather(*[server.sleep.remote() for server in self.servers])\n\n async def abort_all_requests(self):\n \"\"\"Partial rollout: abort and save all unfinished requests in each rollout server.\"\"\"\n # TODO(wuxibin)\n # await asyncio.gather(*[server.abort_all_requests.remote() for server in self.servers])\n print(f\"abort all requests in rollout replica {self.replica_rank}\")\n\n async def resume_all_requests(self):\n \"\"\"Partial rollout: resume all unfinished requests in each rollout server.\"\"\"\n # TODO(wuxibin)\n # await asyncio.gather(*[server.resume_all_requests.remote() for server in self.servers])\n print(f\"resume all requests in rollout replica {self.replica_rank}\")\n\n async def clear_kv_cache(self):\n \"\"\"reset kv cache in each rollout server.\"\"\"\n await asyncio.gather(*[server.clear_kv_cache.remote() for server in self.servers])\n\n async def start_profile(self, **kwargs):\n \"\"\"Start profiling on the replica.\"\"\"\n await asyncio.gather(*[server.start_profile.remote(**kwargs) for server in self.servers])\n\n async def stop_profile(self):\n \"\"\"Stop profiling on the replica.\"\"\"\n await asyncio.gather(*[server.stop_profile.remote() for server in self.servers])\n\n\nclass RolloutReplicaRegistry:\n \"\"\"Factory for managing rollout replica implementations.\"\"\"\n\n _registry: dict[str, Callable[[], type[RolloutReplica]]] = {}\n\n @classmethod\n def register(cls, name: str, loader: Callable[[], type[RolloutReplica]]) -> None:\n \"\"\"Register a new rollout replica type.\"\"\"\n cls._registry[name] = loader\n\n @classmethod\n def get(cls, name: str) -> type[RolloutReplica]:\n \"\"\"Get a rollout replica class by name.\"\"\"\n if name not in cls._registry:\n raise ValueError(f\"Unknown rollout mode: {name}. Available: {list(cls._registry.keys())}\")\n return cls._registry[name]()\n\n\n# Loader functions for built-in types\ndef _load_vllm():\n from verl.workers.rollout.vllm_rollout.vllm_async_server import vLLMReplica\n\n return vLLMReplica\n\n\ndef _load_sglang():\n os.environ[\"SGLANG_USE_CPU_ENGINE\"] = \"1\"\n\n try:\n import vllm # noqa: F401\n except ImportError:\n import sys\n import types\n from unittest.mock import Mock\n\n mock_vllm = types.ModuleType(\"vllm\")\n\n mock_custom_ops = types.ModuleType(\"vllm._custom_ops\")\n mock_custom_ops.scaled_fp8_quant = Mock()\n mock_vllm._custom_ops = mock_custom_ops\n\n mock_model_executor = types.ModuleType(\"vllm.model_executor\")\n mock_layers = types.ModuleType(\"vllm.model_executor.layers\")\n mock_activation = types.ModuleType(\"vllm.model_executor.layers.activation\")\n\n class GeluAndMul: # noqa: N801\n pass\n\n class SiluAndMul: # noqa: N801\n pass\n\n mock_activation.GeluAndMul = GeluAndMul\n mock_activation.SiluAndMul = SiluAndMul\n mock_layers.activation = mock_activation\n mock_model_executor.layers = mock_layers\n mock_vllm.model_executor = mock_model_executor\n\n sys.modules[\"vllm\"] = mock_vllm\n sys.modules[\"vllm._custom_ops\"] = mock_custom_ops\n sys.modules[\"vllm.model_executor\"] = mock_model_executor\n sys.modules[\"vllm.model_executor.layers\"] = mock_layers\n sys.modules[\"vllm.model_executor.layers.activation\"] = mock_activation\n\n from verl.workers.rollout.sglang_rollout.async_sglang_server import SGLangReplica\n\n del os.environ[\"SGLANG_USE_CPU_ENGINE\"]\n return SGLangReplica\n\n\ndef _load_trtllm():\n from verl.workers.rollout.trtllm_rollout.trtllm_async_server import TRTLLMReplica\n\n return TRTLLMReplica\n\n\n# Register built-in types\nRolloutReplicaRegistry.register(\"vllm\", _load_vllm)\nRolloutReplicaRegistry.register(\"sglang\", _load_sglang)\nRolloutReplicaRegistry.register(\"trtllm\", _load_trtllm)\n\n\n# Original function for backward compatibility\ndef get_rollout_replica_class(rollout: str) -> type[RolloutReplica]:\n return RolloutReplicaRegistry.get(rollout)\n"}161{"file_name": "verl__workers__rollout__schemas.py", "text": "# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport difflib\nimport logging\nimport os\nfrom enum import Enum\nfrom typing import Any, Optional\n\nimport torch\nfrom pydantic import BaseModel, ConfigDict, model_validator\nfrom transformers import PreTrainedTokenizer, PreTrainedTokenizerFast, ProcessorMixin\n\nfrom verl.tools.schemas import OpenAIFunctionToolCall, OpenAIFunctionToolSchema, ToolResponse\nfrom verl.utils.model import compute_position_id_with_mask\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\nBASE_CHAT_HISTORY = [\n {\"role\": \"system\", \"content\": \"You are a helpful assistant.\"},\n {\"role\": \"user\", \"content\": \"I am a user.\"},\n]\n\n\nclass FinishReasonTypeEnum(str, Enum):\n \"\"\"The enum for finish reason type.\"\"\"\n\n LENGTH = \"length\"\n STOP = \"stop\"\n TOOL_CALL = \"tool_calls\"\n\n @classmethod\n def from_str(cls, value: str) -> \"FinishReasonTypeEnum\":\n if value == \"stop\":\n return cls.STOP\n elif value == \"length\":\n return cls.LENGTH\n elif value == \"tool_calls\":\n return cls.TOOL_CALL\n else:\n raise ValueError(f\"Unsupported finish reason type: {value}\")\n\n\nclass Message(BaseModel):\n role: str\n content: str | dict[str, Any] | list[dict[str, Any]] | ToolResponse\n tool_calls: Optional[list[OpenAIFunctionToolCall]] = None\n\n\nclass AsyncRolloutRequestStateEnum(str, Enum):\n \"\"\"The enum for async rollout request state.\"\"\"\n\n PENDING = \"pending\"\n RUNNING = \"running\"\n COMPLETED = \"completed\"\n FAILED = \"failed\"\n TOOL_CALLING = \"tool_calling\"\n INTERACTING = \"interacting\"\n\n\nclass TokenizationSanityCheckModeEnum(str, Enum):\n \"\"\"The enum for tokenization sanity check mode.\"\"\"\n\n DISABLE = \"disable\"\n STRICT = \"strict\"\n IGNORE_STRIPPABLE = \"ignore_strippable\"\n\n\nclass AsyncRolloutRequest(BaseModel):\n \"\"\"The data model for async rollout.\"\"\"\n\n model_config = ConfigDict(arbitrary_types_allowed=True)\n\n batch_data_id: int = 0\n rollout_offset: int = 0\n request_id: str\n state: AsyncRolloutRequestStateEnum\n messages: list[Message]\n multi_modal_keys: Optional[list[str]] = None\n multi_modal_data: Optional[dict[str, Any]] = None\n multi_modal_inputs: Optional[dict[str, torch.Tensor]] = None\n tool_schemas: Optional[list[OpenAIFunctionToolSchema]] = None\n tools_kwargs: dict[str, Any] = {}\n interaction_kwargs: dict[str, Any] = {}\n input_ids: Optional[torch.Tensor] = None\n prompt_ids: Optional[torch.Tensor] = None\n response_ids: Optional[torch.Tensor] = None\n attention_mask: Optional[torch.Tensor] = None\n prompt_attention_mask: Optional[torch.Tensor] = None\n response_attention_mask: Optional[torch.Tensor] = None\n position_ids: Optional[torch.Tensor] = None\n prompt_position_ids: Optional[torch.Tensor] = None\n response_position_ids: Optional[torch.Tensor] = None\n loss_mask: Optional[torch.Tensor] = None\n prompt_loss_mask: Optional[torch.Tensor] = None\n response_loss_mask: Optional[torch.Tensor] = None\n reward_scores: dict[str, float]\n max_prompt_len: int\n max_response_len: int = 8192\n max_model_len: int = 32768\n metrics: dict[str, list[Any]] = {}\n output_token_ids: torch.Tensor | None = None\n rollout_log_probs: torch.Tensor | None = None\n\n use_inference_chat_template: bool\n tokenization_sanity_check_mode: TokenizationSanityCheckModeEnum\n generation_prompt_ids: Optional[torch.Tensor] = None\n base_conv_wo_gen_prompt_end_pos: int\n base_conv_with_gen_prompt_end_pos: int\n\n @model_validator(mode=\"before\")\n @classmethod\n def initialize_request(cls, values):\n if not (messages := values.get(\"messages\")):\n raise ValueError(\"messages is required for AsyncRolloutRequest initialization\")\n if not (max_prompt_len := values.get(\"max_prompt_len\")):\n raise ValueError(\"max_prompt_len is required for AsyncRolloutRequest initialization\")\n if not (processing_class := values.pop(\"processing_class\", None)):\n raise ValueError(\"processing_class is required for AsyncRolloutRequest initialization\")\n\n values[\"messages\"] = [Message.model_validate(msg) for msg in messages]\n\n # If there is no multi_modal_keys, we assume the multi-modal data is image and video.\n if not values.get(\"multi_modal_keys\"):\n values[\"multi_modal_keys\"] = [\"image\", \"video\"]\n if not values.get(\"multi_modal_data\"):\n values[\"multi_modal_data\"] = {key: [] for key in values[\"multi_modal_keys\"]}\n else:\n # check if all multi_modal_keys are in multi_modal_data\n for key in values[\"multi_modal_keys\"]:\n if key not in values[\"multi_modal_data\"]:\n values[\"multi_modal_data\"][key] = []\n if not values.get(\"multi_modal_inputs\"):\n values[\"multi_modal_inputs\"] = {}\n\n tools = (\n [tool.model_dump() for tool in tool_schemas] if (tool_schemas := values.get(\"tool_schemas\", [])) else None\n )\n\n multi_modal_data = values[\"multi_modal_data\"]\n tokens_without_prompt = cls._handle_apply_chat_template(\n processing_class,\n messages,\n multi_modal_data=multi_modal_data,\n tools=tools,\n add_generation_prompt=False,\n tokenize=True,\n )\n if (\n values.get(\"input_ids\") is None\n or values.get(\"attention_mask\") is None\n or values.get(\"position_ids\") is None\n ):\n tokenization_dict_with_prompt = cls._handle_apply_chat_template(\n processing_class,\n messages,\n multi_modal_data=multi_modal_data,\n tools=tools,\n add_generation_prompt=True,\n tokenize=True,\n return_dict=True,\n )\n\n values[\"input_ids\"], values[\"attention_mask\"] = (\n tokenization_dict_with_prompt[\"input_ids\"],\n tokenization_dict_with_prompt[\"attention_mask\"],\n )\n if values[\"input_ids\"].shape[-1] > max_prompt_len:\n # Only log the warning to avoid truncating in the middle of generation prompt. Consider raising an\n # error for this case in the future.\n # Ensure batch_data_id exists with default value if not provided\n if \"batch_data_id\" not in values:\n values[\"batch_data_id\"] = cls.model_fields[\"batch_data_id\"].default\n logger.warning(\n f\"Prompt {values['batch_data_id']} has length {values['input_ids'].shape[-1]} \"\n f\"which is greater than max_prompt_len {max_prompt_len} after applied chat template with tools.\"\n )\n\n # Process multi_modal_inputs\n multi_modal_inputs = tokenization_dict_with_prompt.copy()\n multi_modal_inputs.pop(\"input_ids\", None)\n multi_modal_inputs.pop(\"attention_mask\", None)\n values[\"multi_modal_inputs\"] = multi_modal_inputs\n\n values[\"position_ids\"] = values[\"prompt_position_ids\"] = cls._get_position_ids(\n processing_class, values[\"input_ids\"], values[\"attention_mask\"], multi_modal_inputs\n )\n\n values[\"prompt_ids\"], values[\"prompt_attention_mask\"] = values[\"input_ids\"], values[\"attention_mask\"]\n values[\"loss_mask\"] = values[\"prompt_loss_mask\"] = torch.zeros_like(values[\"input_ids\"], dtype=torch.bool)\n values[\"generation_prompt_ids\"] = values[\"input_ids\"][..., tokens_without_prompt.shape[-1] :]\n values[\"base_conv_wo_gen_prompt_end_pos\"] = cls._handle_apply_chat_template(\n processing_class,\n BASE_CHAT_HISTORY,\n multi_modal_data=multi_modal_data,\n tools=tools,\n add_generation_prompt=False,\n tokenize=True,\n ).shape[-1]\n\n values[\"base_conv_with_gen_prompt_end_pos\"] = cls._handle_apply_chat_template(\n processing_class,\n BASE_CHAT_HISTORY,\n multi_modal_data=multi_modal_data,\n tools=tools,\n add_generation_prompt=True,\n tokenize=True,\n ).shape[-1]\n\n return values\n\n @staticmethod\n def _handle_apply_chat_template(\n processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin,\n messages: list[Message],\n multi_modal_data: dict[str, Any],\n tools: Optional[list[OpenAIFunctionToolSchema]] = None,\n add_generation_prompt: bool = False,\n tokenize: bool = False,\n return_dict: bool = False,\n ):\n raw_prompt = processing_class.apply_chat_template(\n messages, tools=tools, add_generation_prompt=add_generation_prompt, tokenize=False\n )\n if not tokenize:\n return raw_prompt\n\n if isinstance(processing_class, PreTrainedTokenizer) or isinstance(processing_class, PreTrainedTokenizerFast):\n if any(len(values) > 0 for values in multi_modal_data.values()):\n logger.warning(\n \"There is multi_modal_data but you are not using a processor. Multi-modal data will be ignored.\"\n )\n model_inputs = processing_class(text=[raw_prompt], return_tensors=\"pt\")\n elif isinstance(processing_class, ProcessorMixin):\n # When we update multi_model_keys, we also need to update this logic\n images = images if len(images := multi_modal_data.get(\"image\", [])) > 0 else None\n videos = videos if len(videos := multi_modal_data.get(\"video\", [])) > 0 else None\n model_inputs = processing_class(text=[raw_prompt], images=images, videos=videos, return_tensors=\"pt\")\n else:\n raise ValueError(f\"Unsupported processing class type: {type(processing_class)}\")\n\n model_inputs = dict(model_inputs)\n if return_dict:\n return model_inputs\n else:\n return model_inputs[\"input_ids\"]\n\n @staticmethod\n def _get_position_ids(\n processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin,\n input_ids: torch.Tensor,\n attention_mask: torch.Tensor,\n multi_modal_inputs: Optional[dict[str, torch.Tensor]] = None,\n ) -> torch.Tensor:\n # special case for qwen2vl\n is_qwen2vl = (\n hasattr(processing_class, \"image_processor\")\n and \"Qwen2VLImageProcessor\" in processing_class.image_processor.__class__.__name__\n )\n if is_qwen2vl:\n from verl.models.transformers.qwen2_vl import get_rope_index\n\n image_grid_thw = video_grid_thw = second_per_grid_ts = None\n if multi_modal_inputs:\n image_grid_thw = multi_modal_inputs.get(\"image_grid_thw\")\n video_grid_thw = multi_modal_inputs.get(\"video_grid_thw\")\n second_per_grid_ts = multi_modal_inputs.get(\"second_per_grid_ts\")\n\n assert input_ids.dim() == 2 and input_ids.shape[0] == 1, (\n f\"input_ids should be 2D with batch size 1, but got shape {input_ids.shape}\"\n )\n assert attention_mask.dim() == 2 and attention_mask.shape[0] == 1, (\n f\"attention_mask should be 2D with batch size 1, but got shape {attention_mask.shape}\"\n )\n new_position_ids = get_rope_index(\n processing_class,\n input_ids=input_ids.squeeze(0),\n image_grid_thw=image_grid_thw,\n video_grid_thw=video_grid_thw,\n second_per_grid_ts=second_per_grid_ts,\n attention_mask=attention_mask.squeeze(0),\n )\n return new_position_ids # (3, seq_len)\n else:\n return compute_position_id_with_mask(attention_mask) # (1, seq_len)\n\n def _update_input_ids(\n self,\n processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin,\n new_input_ids: torch.Tensor,\n attention_mask: bool,\n loss_mask: bool,\n new_multi_modal_inputs: Optional[dict[str, torch.Tensor]] = None,\n ) -> None:\n \"\"\"\n Update the input_ids, attention_mask, position_ids, and loss_mask of the request in additive manner.\n \"\"\"\n self.input_ids = torch.cat([self.input_ids, new_input_ids], dim=-1)\n attention_mask = torch.ones_like(new_input_ids) * int(attention_mask)\n self.attention_mask = torch.cat([self.attention_mask, attention_mask], dim=-1)\n loss_mask = torch.ones_like(new_input_ids) * int(loss_mask)\n self.loss_mask = torch.cat([self.loss_mask, loss_mask], dim=-1)\n\n if new_multi_modal_inputs:\n self._update_multi_modal_inputs(new_multi_modal_inputs)\n\n new_position_ids = self._get_position_ids(\n processing_class, new_input_ids, attention_mask, new_multi_modal_inputs\n )\n\n last_pos = self.position_ids[..., -1:]\n new_position_ids = new_position_ids + (last_pos + 1)\n\n self.position_ids = torch.cat([self.position_ids, new_position_ids], dim=-1)\n\n assert (\n self.input_ids.shape[-1]\n == self.attention_mask.shape[-1]\n == self.position_ids.shape[-1]\n == self.loss_mask.shape[-1]\n ), f\"\"\"Request {self.request_id} has different length of {self.input_ids.shape[-1]=}, \n {self.attention_mask.shape[-1]=}, {self.position_ids.shape[-1]=}, {self.loss_mask.shape[-1]=}\"\"\"\n\n def _update_multi_modal_inputs(self, new_multi_modal_inputs: dict[str, torch.Tensor]) -> None:\n \"\"\"\n Update the multi_modal_inputs of the request in additive manner.\n \"\"\"\n for key in new_multi_modal_inputs:\n input_tensor = new_multi_modal_inputs[key]\n self.multi_modal_inputs[key] = (\n torch.cat([self.multi_modal_inputs[key], input_tensor], dim=0)\n if key in self.multi_modal_inputs\n else input_tensor\n )\n\n def get_generation_prompt_ids(\n self, processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin\n ) -> list[int]:\n \"\"\"\n Get the generation prompt ids for rollout engine.\n\n Because rollout engine(SGLang) requires the ids to be a list, we need to convert the tensor to a list.\n \"\"\"\n generation_prompt_ids = (\n None\n if self.input_ids[..., -self.generation_prompt_ids.shape[-1] :].eq(self.generation_prompt_ids).all()\n else self.generation_prompt_ids\n )\n if generation_prompt_ids is not None:\n self._update_input_ids(processing_class, generation_prompt_ids, attention_mask=True, loss_mask=False)\n\n if self.use_inference_chat_template:\n messages = [msg.model_dump() for msg in self.messages]\n tools = [tool.model_dump() for tool in self.tool_schemas] if self.tool_schemas else None\n generation_prompt_ids = self._handle_apply_chat_template(\n processing_class,\n messages,\n multi_modal_data=self.multi_modal_data,\n tools=tools,\n add_generation_prompt=True,\n tokenize=True,\n )\n return generation_prompt_ids.squeeze(0).tolist()\n else:\n return self.input_ids.squeeze(0).tolist()\n\n def add_user_message(\n self,\n processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin,\n content: str,\n ) -> None:\n self.messages.append(Message(role=\"user\", content=content))\n messages = [*BASE_CHAT_HISTORY, self.messages[-1]]\n tools = [tool.model_dump() for tool in self.tool_schemas] if self.tool_schemas else None\n\n # We don't need to pass multi_modal_data here because we don't have any multi-modal data from Engine\n # Inference, it is pure text.\n content_ids = self._handle_apply_chat_template(\n processing_class, messages, multi_modal_data={}, tools=tools, add_generation_prompt=False, tokenize=True\n )[..., self.base_conv_wo_gen_prompt_end_pos :]\n self._update_input_ids(processing_class, content_ids, attention_mask=True, loss_mask=False)\n\n def add_assistant_message(\n self,\n processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin,\n content: str,\n content_ids: Optional[torch.Tensor] = None,\n tool_calls: Optional[list[OpenAIFunctionToolCall]] = None,\n ) -> None:\n self.messages.append(Message(role=\"assistant\", content=content, tool_calls=tool_calls))\n if content_ids is None:\n messages = [*BASE_CHAT_HISTORY, self.messages[-1]]\n tools = [tool.model_dump() for tool in self.tool_schemas] if self.tool_schemas else None\n\n # We don't need to pass multi_modal_data here because we don't have any multi-modal data from Engine\n # Inference, it is pure text.\n content_ids = self._handle_apply_chat_template(\n processing_class, messages, multi_modal_data={}, tools=tools, add_generation_prompt=False, tokenize=True\n )[..., self.base_conv_with_gen_prompt_end_pos :]\n self._update_input_ids(processing_class, content_ids, attention_mask=True, loss_mask=True)\n\n def add_tool_response_messages(\n self,\n processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin,\n contents: list[ToolResponse],\n ) -> None:\n if not contents or all(content.is_empty() for content in contents):\n return\n # We also handle the case when tool returns image\n # We require the processing of the image and video to be done at tool.execute() level\n delta_multi_modal_data = {key: [] for key in self.multi_modal_keys}\n for content in contents:\n if content.is_text_only():\n self.messages.append(Message(role=\"tool\", content=content.text))\n else:\n content_list = []\n # When we update multi_model_keys, we also need to update this logic\n if content.image:\n content_list.extend([{\"type\": \"image\"} for _ in content.image])\n delta_multi_modal_data[\"image\"].extend(content.image)\n if content.video:\n content_list.extend([{\"type\": \"video\"} for _ in content.video])\n delta_multi_modal_data[\"video\"].extend(content.video)\n if content.text:\n content_list.append({\"type\": \"text\", \"text\": content.text})\n self.messages.append(Message(role=\"tool\", content=content_list))\n\n messages = [*BASE_CHAT_HISTORY, *self.messages[-len(contents) :]]\n tools = [tool.model_dump() for tool in self.tool_schemas] if self.tool_schemas else None\n\n for key in self.multi_modal_keys:\n if len(delta_multi_modal_data[key]) > 0:\n self.multi_modal_data[key].extend(delta_multi_modal_data[key])\n\n # We just passed the new multi-modal data to the chat template to update the input_ids.\n content_info = self._handle_apply_chat_template(\n processing_class,\n messages,\n multi_modal_data=delta_multi_modal_data,\n tools=tools,\n add_generation_prompt=False,\n tokenize=True,\n return_dict=True,\n )\n content_ids = content_info[\"input_ids\"][..., self.base_conv_wo_gen_prompt_end_pos :]\n\n # process multi_modal_inputs\n multi_modal_inputs = content_info.copy()\n multi_modal_inputs.pop(\"input_ids\", None)\n multi_modal_inputs.pop(\"attention_mask\", None)\n\n # chat templates include generation prompt tokens (e.g., \"<im_start>assistant\\n\")\n # So when tool response is added, we need to explicitly remove these tokens.\n self._remove_generation_prompt_ids_if_present()\n\n self._update_input_ids(\n processing_class,\n content_ids,\n attention_mask=True,\n loss_mask=False,\n new_multi_modal_inputs=multi_modal_inputs,\n )\n\n def update_metrics(self, metrics: Any, tool_id: str) -> None:\n \"\"\"\n metrics: should be a dict of tools_name -> Any\n \"\"\"\n if self.metrics.get(tool_id) is None:\n self.metrics[tool_id] = []\n self.metrics[tool_id].append(metrics)\n\n def _get_prompt_diffs(\n self,\n processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin,\n full_prompt_ids: torch.Tensor,\n current_prompt_ids: torch.Tensor,\n diff_surrounding_chars: int = 10,\n ) -> list[dict[str, Any]]:\n \"\"\"Get differences between full prompt and current prompt with surrounding context.\n\n This function helps debug tokenization mismatches by showing the differences between\n full prompt and current prompt with surrounding context. Instead of just showing\n the exact diff, it includes additional tokens before and after to help locate\n the issue in the chat template.\n\n For example, if the actual diff is a newline change from \"\\n\\n\" to \"\\n\", with\n diff_surrounding_chars the output might look like:\n\n full_prompt_chunk: \"<|im_start|>assistant\\n\\nI think...\"\n current_prompt_chunk: \"<|im_start|>assistant\\nI think...\"\n\n This context makes it much easier to identify where in the chat template the\n mismatch occurs.\n\n Args:\n processing_class: The processing class to use for decoding the token IDs\n full_prompt_ids: Token IDs from applying chat template to all messages at once\n current_prompt_ids: Token IDs from incremental chat template application\n diff_surrounding_chars: Number of surrounding characters to include for context (default: 10)\n\n Returns:\n List of dicts containing the differing chunks with context and their indices\n \"\"\"\n full_prompt_ids = full_prompt_ids.squeeze(0)\n current_prompt_ids = current_prompt_ids.squeeze(0)\n full_prompt = processing_class.decode(full_prompt_ids, skip_special_tokens=False)\n current_prompt = processing_class.decode(current_prompt_ids, skip_special_tokens=False)\n s = difflib.SequenceMatcher(None, full_prompt, current_prompt, autojunk=False)\n diffs = []\n for tag, i1, i2, j1, j2 in s.get_opcodes():\n if tag == \"equal\":\n continue\n\n # Get the surrounding context for better readability\n start_i = max(0, i1 - diff_surrounding_chars)\n end_i = min(len(full_prompt), i2 + diff_surrounding_chars)\n start_j = max(0, j1 - diff_surrounding_chars)\n end_j = min(len(current_prompt), j2 + diff_surrounding_chars)\n\n diffs.append(\n {\n \"full_prompt_chunk\": full_prompt[start_i:end_i],\n \"current_prompt_chunk\": current_prompt[start_j:end_j],\n \"indices\": (start_i, end_i, start_j, end_j),\n }\n )\n return diffs\n\n def _remove_generation_prompt_ids_if_present(self) -> None:\n \"\"\"\n Remove generation prompt IDs from input tensors if they are present at the end.\n \"\"\"\n if self.input_ids[..., -self.generation_prompt_ids.shape[-1] :].eq(self.generation_prompt_ids).all():\n self.input_ids = self.input_ids[..., : -self.generation_prompt_ids.shape[-1]]\n self.attention_mask = self.attention_mask[..., : -self.generation_prompt_ids.shape[-1]]\n self.position_ids = self.position_ids[..., : -self.generation_prompt_ids.shape[-1]]\n self.loss_mask = self.loss_mask[..., : -self.generation_prompt_ids.shape[-1]]\n\n def finalize(\n self,\n processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin,\n reward_scores: dict[str, list[float]],\n finish_reason_type: FinishReasonTypeEnum = FinishReasonTypeEnum.STOP,\n ) -> None:\n self.state = AsyncRolloutRequestStateEnum.COMPLETED\n self.reward_scores = reward_scores\n\n # In case we failed to generate the assistant message and the generation prompt ids were already added to\n # input_ids, remove them from the end of input_ids\n self._remove_generation_prompt_ids_if_present()\n\n self.response_ids = self.input_ids[..., self.prompt_ids.shape[-1] :]\n\n if self.tokenization_sanity_check_mode != TokenizationSanityCheckModeEnum.DISABLE:\n # When there is a diff, we log the diffs with diff_surrounding_chars context\n diff_surrounding_chars = 10\n\n messages = [msg.model_dump() for msg in self.messages]\n tools = [tool.model_dump() for tool in self.tool_schemas] if self.tool_schemas else None\n full_prompt_info = self._handle_apply_chat_template(\n processing_class,\n messages,\n multi_modal_data=self.multi_modal_data,\n tools=tools,\n add_generation_prompt=False,\n tokenize=True,\n return_dict=True,\n )\n full_prompt_ids = full_prompt_info[\"input_ids\"]\n\n # We must use dict(full_prompt_info) to convert BatchFeature values to a new dict\n # because np.array() only keeps the keys for BatchFeature.\n full_prompt_multi_modal_inputs = full_prompt_info.copy()\n full_prompt_multi_modal_inputs.pop(\"input_ids\", None)\n full_prompt_multi_modal_inputs.pop(\"attention_mask\", None)\n\n for multi_modal_inputs_key in self.multi_modal_inputs:\n if multi_modal_inputs_key in full_prompt_multi_modal_inputs:\n if (\n not self.multi_modal_inputs[multi_modal_inputs_key]\n .eq(full_prompt_multi_modal_inputs[multi_modal_inputs_key])\n .all()\n ):\n logger.warning(\n f\"Multi-modal data {multi_modal_inputs_key} is not consistent. \"\n f\"This may lead to unexpected behavior during training. \"\n f\"Please review your multi_modal_inputs logic.\"\n )\n else:\n logger.warning(\n f\"Multi-modal inputs key {multi_modal_inputs_key} is not found in the multi_modal_inputs. \"\n f\"This may lead to unexpected behavior during training.\"\n f\"Please review your multi_modal_inputs logic.\"\n )\n\n if diffs := self._get_prompt_diffs(\n processing_class, full_prompt_ids, self.input_ids, diff_surrounding_chars=diff_surrounding_chars\n ):\n log_warning = False\n if self.tokenization_sanity_check_mode == TokenizationSanityCheckModeEnum.STRICT:\n log_warning = True\n elif self.tokenization_sanity_check_mode == TokenizationSanityCheckModeEnum.IGNORE_STRIPPABLE:\n non_strippable_diffs_exist = any(\n d[\"full_prompt_chunk\"].strip() or d[\"current_prompt_chunk\"].strip() for d in diffs\n )\n if non_strippable_diffs_exist:\n log_warning = True\n\n if log_warning:\n mode_str = f\" ({self.tokenization_sanity_check_mode.value})\"\n logger.warning(\n f\"Inconsistent training and inference tokenization detected{mode_str}. This may lead to \"\n f\"unexpected behavior during training. Please review your chat template to determine if this \"\n f\"is intentional. For more information, refer to the multiturn README.md.\"\n )\n logger.warning(\n f\"Showing {diff_surrounding_chars} characters before and after the diffs for context and \"\n f\"better readability.\"\n )\n diff_details_list = []\n for d in diffs:\n i1, i2, j1, j2 = d[\"indices\"]\n diff_details_list.append(\n f\"idx {i1}:{i2} -> {j1}:{j2} | full_prompt_chunk: {repr(d['full_prompt_chunk'])} | \"\n f\"current_prompt_chunk: {repr(d['current_prompt_chunk'])}\"\n )\n diff_details = \"\\n\".join(diff_details_list)\n logger.warning(f\"Found differences:\\n{diff_details}\")\n\n if finish_reason_type == FinishReasonTypeEnum.STOP:\n pass\n elif finish_reason_type == FinishReasonTypeEnum.LENGTH:\n pass\n else:\n raise ValueError(f\"Unsupported finalize finish reason type: {finish_reason_type}\")\n self.truncate_output_ids(processing_class)\n\n assert (\n self.input_ids.shape[-1]\n == self.attention_mask.shape[-1]\n == self.position_ids.shape[-1]\n == self.loss_mask.shape[-1]\n ), f\"\"\"Request {self.request_id} has different length of {self.input_ids.shape[-1]=}, \n {self.attention_mask.shape[-1]=}, {self.position_ids.shape[-1]=}, {self.loss_mask.shape[-1]=}\"\"\"\n\n def truncate_output_ids(\n self, processing_class: PreTrainedTokenizer | PreTrainedTokenizerFast | ProcessorMixin\n ) -> None:\n self.input_ids = self.input_ids[..., : self.max_model_len]\n self.attention_mask = self.attention_mask[..., : self.max_model_len]\n self.position_ids = self.position_ids[..., : self.max_model_len]\n self.loss_mask = self.loss_mask[..., : self.max_model_len]\n self.response_ids = self.input_ids[..., self.prompt_ids.shape[-1] :][..., : self.max_response_len]\n self.response_attention_mask = self.attention_mask[..., self.prompt_attention_mask.shape[-1] :][\n ..., : self.max_response_len\n ]\n self.response_position_ids = self.position_ids[..., self.prompt_position_ids.shape[-1] :][\n ..., : self.max_response_len\n ]\n self.response_loss_mask = self.loss_mask[..., self.prompt_loss_mask.shape[-1] :][..., : self.max_response_len]\n"}162{"file_name": "verl__workers__rollout__sglang_rollout__async_sglang_server.py", "text": "# Copyright 2023-2024 SGLang Team\n# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport asyncio\nimport dataclasses\nimport json\nimport logging\nimport os\nfrom typing import Any, Optional\n\nimport ray\nimport sglang\nimport sglang.srt.entrypoints.engine\nimport torch\nfrom packaging import version\nfrom ray.actor import ActorHandle\nfrom sglang.srt.entrypoints.http_server import (\n ServerArgs,\n _GlobalState,\n _launch_subprocesses,\n app,\n set_global_state,\n)\nfrom sglang.srt.managers.io_struct import (\n GenerateReqInput,\n ReleaseMemoryOccupationReqInput,\n ResumeMemoryOccupationReqInput,\n)\nfrom sglang.srt.managers.tokenizer_manager import ServerStatus\n\nfrom verl.single_controller.ray import RayClassWithInitArgs\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.utils.device import get_visible_devices_keyword\nfrom verl.utils.net_utils import get_free_port, is_valid_ipv6_address\nfrom verl.utils.profiler import DistProfiler, build_sglang_profiler_args\nfrom verl.workers.config import HFModelConfig, RolloutConfig\nfrom verl.workers.rollout.replica import RolloutMode, RolloutReplica, TokenOutput\nfrom verl.workers.rollout.sglang_rollout.sglang_rollout import ServerAdapter, _set_envs_and_config\nfrom verl.workers.rollout.utils import get_max_position_embeddings, run_unvicorn\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(logging.INFO)\n\nvisible_devices_keyword = get_visible_devices_keyword()\n\n\nclass SGLangHttpServer:\n \"\"\"SGLang http server in single node, this is equivalent to launch server with command line:\n ```\n python -m sglang.launch_server --node-rank 0 --nnode 1 ...\n ```\n\n Args:\n config (DictConfig): full config.\n rollout_mode (RolloutMode): rollout mode.\n replica_rank (int): replica rank, a replica may contain multiple nodes.\n node_rank (int): node rank.\n nnodes (int): number of nodes.\n cuda_visible_devices (str): cuda visible devices.\n \"\"\"\n\n def __init__(\n self,\n config: RolloutConfig,\n model_config: HFModelConfig,\n rollout_mode: RolloutMode,\n workers: list[ActorHandle],\n replica_rank: int,\n node_rank: int,\n nnodes: int,\n cuda_visible_devices: str,\n base_gpu_id: int,\n ):\n print(f\"SGLang http server: {rollout_mode=}, {replica_rank=}, {node_rank=}, {nnodes=}, {cuda_visible_devices=}\")\n os.environ[visible_devices_keyword] = cuda_visible_devices\n\n self.config: RolloutConfig = omega_conf_to_dataclass(config)\n self.model_config: HFModelConfig = omega_conf_to_dataclass(model_config, dataclass_type=HFModelConfig)\n max_position_embeddings = get_max_position_embeddings(self.model_config.hf_config)\n if self.config.max_model_len is None:\n self.config.max_model_len = max_position_embeddings\n else:\n if self.config.max_model_len > max_position_embeddings:\n raise ValueError(\n f\"max_model_len ({self.config.max_model_len}) should be less than or equal to \"\n f\"max_position_embeddings ({max_position_embeddings})\"\n )\n self.rollout_mode = rollout_mode\n self.workers = workers\n\n self.replica_rank = replica_rank\n self.node_rank = node_rank\n self.nnodes = nnodes\n self.base_gpu_id = base_gpu_id\n\n if self.rollout_mode != RolloutMode.HYBRID and self.config.load_format == \"dummy\":\n logger.warning(f\"rollout mode is {self.rollout_mode}, load_format is dummy, set to auto\")\n self.config.load_format = \"auto\"\n\n # used for http server\n self._server_address = ray.util.get_node_ip_address().strip(\"[]\")\n self._server_port = None\n\n # used for controlling sglang server profiler\n profiler_config = self.config.profiler\n tool_config = None\n if profiler_config is not None:\n if profiler_config.tool in [\"torch\", \"npu\"]:\n tool_config = omega_conf_to_dataclass((profiler_config.tool_config or {}).get(profiler_config.tool))\n else:\n logger.warning(f\"agent loop only support torch and npu profiler, got {profiler_config.tool}\")\n profiler_config = None\n self.profiler_controller = DistProfiler(self.replica_rank, config=profiler_config, tool_config=tool_config)\n\n # used for NCCL process group\n if self.node_rank == 0:\n self._master_address = self._server_address\n self._master_port, self._master_sock = get_free_port(self._server_address)\n logger.info(\n f\"SGLangHttpServer, replica_rank: {self.replica_rank}, \"\n f\"master address: {self._master_address}, port: {self._master_port}\"\n )\n else:\n self._master_address = None\n self._master_port = None\n\n def get_master_address(self):\n \"\"\"Get master address and port for init NCCL process group.\"\"\"\n return self._master_address, self._master_port\n\n def get_server_address(self):\n \"\"\"Get http server address and port.\"\"\"\n assert self._server_port is not None, \"http server is not launched, port is None\"\n return self._server_address, self._server_port\n\n async def launch_server(self, master_address: str = None, master_port: int = None):\n if self.node_rank != 0:\n assert master_address and master_port, \"non-master node should provide master address and port\"\n self._master_address = master_address\n self._master_port = master_port\n\n engine_kwargs = self.config.get(\"engine_kwargs\", {}).get(\"sglang\", {}) or {}\n attention_backend = engine_kwargs.pop(\"attention_backend\", None)\n quantization = self.config.get(\"quantization\", None)\n if quantization is not None:\n if quantization == \"fp8\":\n assert version.parse(sglang.__version__) >= version.parse(\"0.5.5\"), (\n \"sglang>=0.5.5 is required for FP8 quantization\"\n )\n FP8_BLOCK_QUANT_KWARGS = {\n \"activation_scheme\": \"dynamic\",\n \"fmt\": \"e4m3\",\n \"quant_method\": \"fp8\",\n \"weight_block_size\": [128, 128],\n }\n fp8_block_quant_kwargs = dict(FP8_BLOCK_QUANT_KWARGS)\n else:\n raise ValueError(f\"Currently only support fp8 quantization, got: {quantization}\")\n dist_init_addr = (\n f\"[{self._master_address}]:{self._master_port}\"\n if is_valid_ipv6_address(self._master_address)\n else f\"{self._master_address}:{self._master_port}\"\n )\n infer_tp = self.config.tensor_model_parallel_size * self.config.data_parallel_size\n args = {\n \"model_path\": self.model_config.local_path,\n \"dtype\": self.config.dtype,\n \"mem_fraction_static\": self.config.gpu_memory_utilization,\n \"disable_cuda_graph\": self.config.enforce_eager,\n \"enable_memory_saver\": True,\n \"base_gpu_id\": self.base_gpu_id,\n \"gpu_id_step\": 1,\n \"tp_size\": infer_tp,\n \"dp_size\": self.config.data_parallel_size,\n \"ep_size\": self.config.expert_parallel_size,\n \"node_rank\": self.node_rank,\n \"load_format\": self.config.load_format,\n \"dist_init_addr\": dist_init_addr,\n \"nnodes\": self.nnodes,\n \"trust_remote_code\": self.model_config.trust_remote_code,\n \"max_running_requests\": self.config.get(\"max_num_seqs\", None),\n \"log_level\": \"error\",\n \"mm_attention_backend\": \"fa3\",\n \"attention_backend\": attention_backend if attention_backend is not None else \"fa3\",\n \"skip_tokenizer_init\": self.config.skip_tokenizer_init,\n \"skip_server_warmup\": True,\n \"quantization\": quantization,\n \"json_model_override_args\": json.dumps({\"quantization_config\": fp8_block_quant_kwargs})\n if quantization == \"fp8\"\n else json.dumps({}),\n **engine_kwargs,\n }\n\n if self.config.prometheus.enable:\n if self.config.prometheus.served_model_name:\n # Extract model name from path if it's a full path\n served_model_name = self.config.prometheus.served_model_name\n if \"/\" in served_model_name:\n # If it's a full path, extract the last part as model name\n served_model_name = served_model_name.split(\"/\")[-1]\n args[\"served_model_name\"] = served_model_name\n\n # start sglang metrics\n args[\"enable_metrics\"] = True\n\n # enable_weights_cpu_backup is supported in sglang>=0.5.3\n if \"enable_weights_cpu_backup\" in [f.name for f in dataclasses.fields(ServerArgs)]:\n enable_weights_cpu_backup = True if self.rollout_mode == RolloutMode.COLOCATED else False\n args[\"enable_weights_cpu_backup\"] = enable_weights_cpu_backup\n\n if self.config.enable_rollout_routing_replay:\n args.update({\"enable_return_routed_experts\": True})\n\n # mtp\n if self.config.mtp.enable and self.config.mtp.enable_rollout:\n # Enable weights CPU backup for sglang >= 0.5.6\n if sglang.__version__ < \"0.5.6\":\n raise ValueError(f\"sglang version {sglang.__version__} is not supported for MTP rollout\")\n\n args[\"speculative_algorithm\"] = self.config.mtp.speculative_algorithm\n args[\"speculative_num_steps\"] = self.config.mtp.speculative_num_steps\n args[\"speculative_eagle_topk\"] = self.config.mtp.speculative_eagle_topk\n args[\"speculative_num_draft_tokens\"] = self.config.mtp.speculative_num_draft_tokens\n\n args[\"enable_weights_cpu_backup\"] = True\n args[\"enable_draft_weights_cpu_backup\"] = True\n\n # NOTE: We can't directly call SGLang's launch_server since it's not an async function.\n # https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/entrypoints/http_server.py\n sglang.srt.entrypoints.engine._set_envs_and_config = _set_envs_and_config\n os.environ[\"SGLANG_BLOCK_NONZERO_RANK_CHILDREN\"] = \"0\"\n server_args = ServerArgs(**args)\n if version.parse(sglang.__version__) >= version.parse(\"0.5.7\"):\n self.tokenizer_manager, self.template_manager, self.scheduler_info, *_ = _launch_subprocesses(\n server_args=server_args,\n init_tokenizer_manager_func=sglang.srt.entrypoints.engine.init_tokenizer_manager,\n run_scheduler_process_func=sglang.srt.entrypoints.engine.run_scheduler_process,\n run_detokenizer_process_func=sglang.srt.entrypoints.engine.run_detokenizer_process,\n )\n else:\n self.tokenizer_manager, self.template_manager, self.scheduler_info, *_ = _launch_subprocesses(\n server_args=server_args\n )\n\n # In multi-node cases, non-zero rank nodes should not launch http server.\n if self.node_rank > 0:\n return\n\n set_global_state(\n _GlobalState(\n tokenizer_manager=self.tokenizer_manager,\n template_manager=self.template_manager,\n scheduler_info=self.scheduler_info,\n )\n )\n app.is_single_tokenizer_mode = True\n\n # Set warmup_thread_{kw}args to avoid AttributeError in lifespan function\n app.server_args = server_args\n app.warmup_thread_kwargs = {\"server_args\": server_args}\n app.warmup_thread_args = (server_args, None, None)\n\n # Manually add Prometheus middleware before starting server\n # This ensures /metrics endpoint is available immediately\n if server_args.enable_metrics:\n from sglang.srt.utils.common import add_prometheus_middleware\n\n add_prometheus_middleware(app)\n\n self._server_port, self._server_task = await run_unvicorn(app, server_args, self._server_address)\n self.tokenizer_manager.server_status = ServerStatus.Up\n\n async def wake_up(self):\n if self.node_rank != 0:\n return\n\n if self.rollout_mode == RolloutMode.HYBRID:\n # In hybrid mode, rollout is wake up in `update_weights`\n raise ValueError(f\"wake_up not support rollout_mode {self.rollout_mode}\")\n elif self.rollout_mode == RolloutMode.COLOCATED:\n # Directly call engine to wake up without sync weights.\n obj = ResumeMemoryOccupationReqInput(tags=[\"kv_cache\", \"weights\"])\n await self.tokenizer_manager.resume_memory_occupation(obj, None)\n await self.tokenizer_manager.flush_cache()\n elif self.rollout_mode == RolloutMode.STANDALONE:\n logger.info(\"skip wake_up in standalone mode\")\n\n async def sleep(self):\n if self.node_rank != 0 or not self.config.free_cache_engine:\n return\n\n if self.rollout_mode == RolloutMode.HYBRID:\n obj = ReleaseMemoryOccupationReqInput(tags=[\"kv_cache\", \"weights\"])\n await self.tokenizer_manager.release_memory_occupation(obj, None)\n elif self.rollout_mode == RolloutMode.COLOCATED:\n obj = ReleaseMemoryOccupationReqInput(tags=[\"kv_cache\", \"weights\"])\n await self.tokenizer_manager.release_memory_occupation(obj, None)\n elif self.rollout_mode == RolloutMode.STANDALONE:\n logger.info(\"skip sleep in standalone mode\")\n\n async def clear_kv_cache(self):\n if self.node_rank == 0:\n await self.tokenizer_manager.flush_cache()\n\n async def generate(\n self,\n prompt_ids: torch.Tensor,\n sampling_params: dict[str, Any],\n request_id: str,\n image_data: Optional[list[Any]] = None,\n video_data: Optional[list[Any]] = None,\n ) -> TokenOutput:\n \"\"\"Generate sequence with token-in-token-out.\"\"\"\n # TODO(@wuxibin): switch to `/generate` http endpoint once multi-modal support ready.\n max_possible_tokens = self.config.max_model_len - len(prompt_ids)\n\n if max_possible_tokens < 0:\n raise ValueError(\n f\"Prompt length ({len(prompt_ids)}) exceeds the model's maximum context length \"\n f\"({self.config.max_model_len}).\"\n )\n\n if \"max_new_tokens\" in sampling_params:\n max_new_tokens = sampling_params.pop(\"max_new_tokens\")\n elif \"max_tokens\" in sampling_params:\n # support vllm-style 'max_tokens' param\n max_new_tokens = sampling_params.pop(\"max_tokens\")\n else:\n max_new_tokens = self.config.response_length + self.config.prompt_length - len(prompt_ids)\n\n # Clamp max_new_tokens to the valid range [0, max_possible_tokens]\n max_new_tokens = max(0, min(max_new_tokens, max_possible_tokens))\n\n assert max_new_tokens <= max_possible_tokens, (\n f\"max_new_tokens {max_new_tokens} exceeds available context space {max_possible_tokens}\"\n )\n sampling_params[\"max_new_tokens\"] = max_new_tokens\n return_logprob = sampling_params.pop(\"logprobs\", False)\n\n request = {\n \"rid\": request_id,\n \"input_ids\": prompt_ids,\n \"sampling_params\": sampling_params,\n \"return_logprob\": return_logprob,\n \"image_data\": image_data,\n # TODO: support video input for sglang\n # video_data=video_data,\n }\n\n if self.config.enable_rollout_routing_replay:\n request.update({\"return_routed_experts\": True})\n\n generate_request = GenerateReqInput(**request)\n\n output = await self.tokenizer_manager.generate_request(generate_request, None).__anext__()\n if return_logprob:\n output_token_logprobs = output[\"meta_info\"][\"output_token_logprobs\"]\n log_probs, token_ids = zip(\n *[(log_prob, token_ids) for log_prob, token_ids, _ in output_token_logprobs], strict=True\n )\n else:\n token_ids = output[\"output_ids\"]\n log_probs = None\n\n routed_experts = None\n if self.config.enable_rollout_routing_replay:\n if self.config.skip_tokenizer_init:\n routed_experts = output.get(\"meta_info\", {}).get(\"routed_experts\", None)\n else:\n from sglang.srt.layers.moe.routed_experts_capturer import extract_routed_experts_from_meta_info\n\n hf_config = self.model_config.hf_config\n if not hasattr(hf_config, \"num_hidden_layers\") or not hasattr(hf_config, \"num_experts_per_tok\"):\n raise AttributeError(\n \"enable_rollout_routing_replay is set, but hf_config is missing \"\n \"'num_hidden_layers' or 'num_experts_per_tok'. This feature requires an MoE model \"\n \"configuration that defines these attributes.\"\n )\n routed_experts = extract_routed_experts_from_meta_info(output).reshape(\n -1, hf_config.num_hidden_layers, hf_config.num_experts_per_tok\n )\n\n return TokenOutput(token_ids=token_ids, log_probs=log_probs, routed_experts=routed_experts)\n\n async def start_profile(self, **kwargs):\n if (\n self.profiler_controller.check_enable()\n and self.profiler_controller.check_this_rank()\n and self.profiler_controller.is_discrete_mode()\n ):\n profile_args = build_sglang_profiler_args(\n self.profiler_controller.config, self.profiler_controller.tool_config, self.replica_rank\n )\n await self.tokenizer_manager.start_profile(**profile_args)\n\n async def stop_profile(self):\n if (\n self.profiler_controller.check_enable()\n and self.profiler_controller.check_this_rank()\n and self.profiler_controller.is_discrete_mode()\n ):\n await self.tokenizer_manager.stop_profile()\n\n\n_rollout_worker_actor_cls = ray.remote(ServerAdapter)\n\n\nclass SGLangReplica(RolloutReplica):\n def __init__(\n self,\n replica_rank: int,\n config: RolloutConfig,\n model_config: HFModelConfig,\n gpus_per_node: int = 8,\n is_reward_model: bool = False,\n ):\n super().__init__(replica_rank, config, model_config, gpus_per_node, is_reward_model)\n self.server_class = ray.remote(SGLangHttpServer)\n\n def get_ray_class_with_init_args(self) -> RayClassWithInitArgs:\n \"\"\"Get rollout worker actor class for colocated and standalone mode.\"\"\"\n worker_dict_cls = RayClassWithInitArgs(\n cls=_rollout_worker_actor_cls,\n config=self.config,\n model_config=self.model_config,\n device_mesh=None,\n )\n return worker_dict_cls\n\n async def launch_servers(self):\n \"\"\"Launch http server in each node.\"\"\"\n assert len(self.workers) == self.world_size, (\n f\"worker number {len(self.workers)} not equal to world size {self.world_size}\"\n )\n\n # get (node_id, CUDA_VISIBLE_DEVICES) of all workers\n worker_infos = await asyncio.gather(\n *[\n worker.__ray_call__.remote(\n lambda self: (ray.get_runtime_context().get_node_id(), os.environ[visible_devices_keyword])\n )\n for worker in self.workers\n ]\n )\n worker_cuda_visible_devices = [worker_info[1] for worker_info in worker_infos]\n worker_node_ids = [worker_info[0] for worker_info in worker_infos]\n base_gpu_id = 0\n infer_tp = self.config.tensor_model_parallel_size * self.config.data_parallel_size\n replica_world_size = infer_tp * self.config.pipeline_model_parallel_size\n if os.environ.get(f\"RAY_EXPERIMENTAL_NOSET_{visible_devices_keyword}\", None):\n logger.warning(f\"RAY_EXPERIMENTAL_NOSET_{visible_devices_keyword} is set True!\")\n base_gpu_id = (0 + self.replica_rank * replica_world_size) % self.gpus_per_node\n # create server actor in each node with node affinity and cuda visible devices\n for node_rank in range(self.nnodes):\n workers = self.workers[\n node_rank * self.gpus_per_replica_node : (node_rank + 1) * self.gpus_per_replica_node\n ]\n node_cuda_visible_devices_set = worker_cuda_visible_devices[\n node_rank * self.gpus_per_replica_node : (node_rank + 1) * self.gpus_per_replica_node\n ]\n node_cuda_visible_devices = \",\".join(\n map(\n str,\n sorted(\n set(\n int(device)\n for worker_devices_set in node_cuda_visible_devices_set\n for device in worker_devices_set.split(\",\")\n if device.strip()\n )\n ),\n )\n )\n\n node_id = worker_node_ids[node_rank * self.gpus_per_replica_node]\n name = (\n f\"sglang_server_{self.replica_rank}_{node_rank}\"\n if not self.is_reward_model\n else f\"sglang_server_reward_{self.replica_rank}_{node_rank}\"\n )\n server = self.server_class.options(\n scheduling_strategy=ray.util.scheduling_strategies.NodeAffinitySchedulingStrategy(\n node_id=node_id,\n soft=False,\n ),\n runtime_env={\"env_vars\": {f\"RAY_EXPERIMENTAL_NOSET_{visible_devices_keyword}\": \"1\"}},\n name=name,\n ).remote(\n config=self.config,\n model_config=self.model_config,\n rollout_mode=self.rollout_mode,\n workers=workers,\n replica_rank=self.replica_rank,\n node_rank=node_rank,\n nnodes=self.nnodes,\n cuda_visible_devices=node_cuda_visible_devices,\n base_gpu_id=base_gpu_id,\n )\n self.servers.append(server)\n\n # launch http server in each node\n master_address, master_port = await self.servers[0].get_master_address.remote()\n await asyncio.gather(\n *[\n server.launch_server.remote(master_address=master_address, master_port=master_port)\n for server in self.servers\n ]\n )\n\n # get http server address from first server\n server_address, server_port = await self.servers[0].get_server_address.remote()\n self._server_handle = self.servers[0]\n self._server_address = (\n f\"[{server_address}]:{server_port}\"\n if is_valid_ipv6_address(server_address)\n else f\"{server_address}:{server_port}\"\n )\n"}163{"file_name": "verl__workers__rollout__sglang_rollout__sglang_rollout.py", "text": "# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nfrom __future__ import annotations\n\nimport logging\nimport multiprocessing as mp\nimport os\nfrom typing import Generator\n\nimport ray\nimport sglang.srt.entrypoints.engine\nimport torch\nfrom sglang.srt.server_args import ServerArgs\nfrom sglang.srt.utils import (\n assert_pkg_version,\n is_cuda,\n set_prometheus_multiproc_dir,\n set_ulimit,\n)\nfrom sglang.srt.weight_sync.utils import update_weights as sgl_update_weights\nfrom torch.distributed.device_mesh import DeviceMesh, init_device_mesh\n\nfrom verl.utils.net_utils import is_valid_ipv6_address\nfrom verl.workers.config import HFModelConfig, RolloutConfig\nfrom verl.workers.rollout.base import BaseRollout\nfrom verl.workers.rollout.sglang_rollout.http_server_engine import AsyncHttpServerAdapter\nfrom verl.workers.rollout.sglang_rollout.utils import get_named_tensor_buckets\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n\n# patch to avoid issue https://github.com/sgl-project/sglang/issues/6723\ndef _set_envs_and_config(server_args: ServerArgs):\n # Set global environments\n os.environ[\"TF_CPP_MIN_LOG_LEVEL\"] = \"3\"\n os.environ[\"NCCL_CUMEM_ENABLE\"] = \"0\"\n os.environ[\"NCCL_NVLS_ENABLE\"] = str(int(server_args.enable_nccl_nvls))\n os.environ[\"TORCH_NCCL_AVOID_RECORD_STREAMS\"] = \"1\"\n os.environ[\"CUDA_DEVICE_MAX_CONNECTIONS\"] = \"4\"\n os.environ[\"CUDA_MODULE_LOADING\"] = \"AUTO\"\n # Enable faulthandler in subprocesses\n os.environ[\"PYTHONFAULTHANDLER\"] = \"1\"\n\n # Set prometheus env vars\n if server_args.enable_metrics:\n set_prometheus_multiproc_dir()\n\n # Set ulimit\n set_ulimit()\n\n # Check flashinfer version\n if server_args.attention_backend == \"flashinfer\":\n assert_pkg_version(\n \"flashinfer_python\",\n \"0.2.5\",\n \"Please uninstall the old version and reinstall the latest version by following the instructions at https://docs.flashinfer.ai/installation.html.\",\n )\n if is_cuda():\n assert_pkg_version(\n \"sgl-kernel\",\n \"0.1.1\",\n \"Please reinstall the latest version with `pip install sgl-kernel --force-reinstall`\",\n )\n\n # Set mp start method\n mp.set_start_method(\"spawn\", force=True)\n\n\nsglang.srt.entrypoints.engine._set_envs_and_config = _set_envs_and_config\n\n\n# because chatCompletion is an async method, it makes the whole ray actor be an async actor\n# which can not call loop.run_until_complete. So we need to make the engine to be an async class\nclass ServerAdapter(BaseRollout):\n \"\"\"SGLang server adapter used in native http server mode, serve as http client to request SGLang server\n to resume/release/update weights and kv_cache.\n\n - hybrid mode: reside in each hybrid worker to sync weights between training engine and SGLang server.\n - standalone/colocated mode: just a dummy placeholder to occupy the GPU to prevent ray scheduling new GPU actor.\n \"\"\"\n\n def __init__(\n self,\n config: RolloutConfig,\n model_config: HFModelConfig,\n device_mesh: DeviceMesh,\n ):\n if config.get(\"quantization\", None) == \"fp8\":\n import sglang\n from packaging import version\n\n assert version.parse(sglang.__version__) >= version.parse(\"0.5.5\"), (\n \"sglang>=0.5.5 is required for FP8 quantization\"\n )\n FP8_BLOCK_QUANT_KWARGS = {\n \"activation_scheme\": \"dynamic\",\n \"fmt\": \"e4m3\",\n \"quant_method\": \"fp8\",\n \"weight_block_size\": [128, 128],\n }\n fp8_block_quant_kwargs = dict(FP8_BLOCK_QUANT_KWARGS)\n model_config.hf_config.quantization_config = fp8_block_quant_kwargs\n super().__init__(config, model_config, device_mesh)\n self._engine: AsyncHttpServerAdapter = None\n\n rank = int(os.environ[\"RANK\"])\n local_world_size = int(os.environ[\"RAY_LOCAL_WORLD_SIZE\"])\n rollout_world_size = self.config.tensor_model_parallel_size * self.config.data_parallel_size\n self.replica_rank = rank // rollout_world_size\n self.rollout_rank = rank % rollout_world_size\n self.node_rank = self.rollout_rank // local_world_size\n self.local_rank = self.rollout_rank % local_world_size\n\n async def _init_server_adapter(self):\n if self._engine is not None:\n return\n\n # device_mesh is needed to gather cuda ipc handle to update weights\n if self.device_mesh is None:\n assert torch.distributed.is_initialized(), \"torch distributed must be initialized\"\n infer_tp = self.config.tensor_model_parallel_size * self.config.data_parallel_size\n infer_pp = self.config.pipeline_model_parallel_size\n infer_world_size = infer_tp * infer_pp\n dp = torch.distributed.get_world_size() // infer_world_size\n self.device_mesh = init_device_mesh(\n \"cpu\", mesh_shape=(dp, infer_tp, infer_pp), mesh_dim_names=[\"dp\", \"infer_tp\", \"infer_pp\"]\n )\n\n # Only init http server adapter in tp rank 0\n if self.device_mesh[\"infer_tp\"].get_local_rank() != 0:\n return\n\n # Lazy init http server adapter because http server is launched after hybrid engine.\n self.server_actor = ray.get_actor(f\"sglang_server_{self.replica_rank}_{self.node_rank}\")\n server_address, server_port = await self.server_actor.get_server_address.remote()\n logger.debug(\n f\"replica_rank={self.replica_rank} node_rank={self.node_rank}, \"\n f\"server address: {server_address}, port: {server_port}\"\n )\n host = f\"[{server_address}]\" if is_valid_ipv6_address(server_address) else server_address\n self._engine = AsyncHttpServerAdapter(\n model_path=self.model_config.local_path,\n host=host,\n port=server_port,\n launch_server=False,\n trust_remote_code=self.model_config.trust_remote_code,\n )\n\n async def resume(self, tags: list[str]):\n \"\"\"Resume rollout weights or kv cache in GPU memory.\n\n Args:\n tag: weights or kv_cache.\n \"\"\"\n await self._init_server_adapter()\n if self.device_mesh[\"infer_tp\"].get_local_rank() == 0 and self.config.free_cache_engine:\n await self._engine.resume_memory_occupation(tags=tags)\n\n async def release(self):\n \"\"\"Release weights and kv cache in GPU memory.\"\"\"\n await self._init_server_adapter()\n if self.device_mesh[\"infer_tp\"].get_local_rank() == 0 and self.config.free_cache_engine:\n await self._engine.release_memory_occupation(tags=[\"kv_cache\", \"weights\"])\n\n async def update_weights(self, weights: Generator[tuple[str, torch.Tensor], None, None], **kwargs):\n \"\"\"\n Update model weights using tensor buckets, similar to THUDM/slime's implementation.\n\n Notes:\n - For the best performance of `rebuild_cuda_tensor`, it is recommended to:\n 1. Enable `RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES`.\n 2. Manually set `CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7`\n when using Tensor Parallelism (TP >= 8).\n - See reference implementations in SLIME:\n - Main logic: https://github.com/THUDM/slime/blob/fb7605cc5fb09af0f9369d37f7192f12bddee577/slime/ray/ppo_actor.py#L452\n - runtime envs: https://github.com/THUDM/slime/blob/fb7605cc5fb09af0f9369d37f7192f12bddee577/slime/ray/ppo_actor.py#L39\n \"\"\"\n await self._init_server_adapter()\n\n update_weights_bucket_bytes = int(self.config.checkpoint_engine.update_weights_bucket_megabytes) << 20\n if self.config.get(\"quantization\", None) == \"fp8\":\n from verl.utils.sglang.sglang_fp8_utils import quant_weights_by_name\n\n logger.info(\"Convert bf16 weights to fp8 format before loading\")\n weights = quant_weights_by_name(\n weights,\n self.model_config.hf_config.quantization_config,\n dtype=self.model_config.hf_config.dtype,\n )\n else:\n weights = weights\n\n async for params_batch in get_named_tensor_buckets(weights, update_weights_bucket_bytes):\n await sgl_update_weights(\n engine=self._engine,\n params_batch=params_batch,\n device_mesh_key=\"infer_tp\",\n device_mesh=self.device_mesh,\n )\n\n if self.device_mesh[\"infer_tp\"].get_local_rank() == 0:\n await self._engine.flush_cache()\n"}164{"file_name": "verl__workers__rollout__sglang_rollout__utils.py", "text": "# Copyright 2023-2024 SGLang Team\n# Copyright 2025 ModelBest Inc. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport pickle\nfrom typing import Any, Iterator, Optional\n\nimport numpy as np\nimport torch\nimport torch.distributed as dist\n\nfrom verl.utils.device import get_device_name\nfrom verl.workers.rollout.utils import ensure_async_iterator\n\n\ndef broadcast_pyobj(\n data: list[Any],\n rank: int,\n dist_group: Optional[torch.distributed.ProcessGroup] = None,\n src: int = 0,\n force_cpu_device: bool = False,\n):\n \"\"\"from https://github.com/sgl-project/sglang/blob/844e2f227ab0cce6ef818a719170ce37b9eb1e1b/python/sglang/srt/utils.py#L905\n\n Broadcast inputs from src rank to all other ranks with torch.dist backend.\n The `rank` here refer to the source rank on global process group (regardless\n of dist_group argument).\n \"\"\"\n device = torch.device(get_device_name() if not force_cpu_device else \"cpu\")\n\n if rank == src:\n if len(data) == 0:\n tensor_size = torch.tensor([0], dtype=torch.long, device=device)\n dist.broadcast(tensor_size, src=src, group=dist_group)\n else:\n serialized_data = pickle.dumps(data)\n size = len(serialized_data)\n\n tensor_data = torch.ByteTensor(np.frombuffer(serialized_data, dtype=np.uint8)).to(device)\n tensor_size = torch.tensor([size], dtype=torch.long, device=device)\n\n dist.broadcast(tensor_size, src=src, group=dist_group)\n dist.broadcast(tensor_data, src=src, group=dist_group)\n return data\n else:\n tensor_size = torch.tensor([0], dtype=torch.long, device=device)\n dist.broadcast(tensor_size, src=src, group=dist_group)\n size = tensor_size.item()\n\n if size == 0:\n return []\n\n tensor_data = torch.empty(size, dtype=torch.uint8, device=device)\n dist.broadcast(tensor_data, src=src, group=dist_group)\n\n serialized_data = bytes(tensor_data.cpu().numpy())\n data = pickle.loads(serialized_data)\n return data\n\n\nasync def get_named_tensor_buckets(\n iterable: Iterator[tuple[str, torch.Tensor]], bucket_bytes: int\n) -> Iterator[list[tuple[str, torch.Tensor]]]:\n \"\"\"\n Group tensors into buckets based on a specified size in megabytes.\n\n Args:\n iterable: An iterator of tuples containing tensor names and tensors.\n bucket_bytes: The maximum size of each bucket in bytes.\n\n Yields:\n Lists of tuples, where each tuple contains a tensor name and its corresponding tensor.\n\n Example:\n >>> tensors = [('tensor1', torch.randn(1000, 1000)), ('tensor2', torch.randn(2000, 2000))]\n >>> for bucket in get_named_tensor_buckets(tensors, bucket_size_mb=10):\n ... print(bucket)\n [('tensor1', tensor(...)), ('tensor2', tensor(...))]\n\n \"\"\"\n if bucket_bytes <= 0:\n raise ValueError(f\"bucket_bytes must be greater than 0, got {bucket_bytes}\")\n\n current_bucket = []\n current_size = 0\n async for name, tensor in ensure_async_iterator(iterable):\n tensor_size = tensor.element_size() * tensor.numel()\n if current_size + tensor_size > bucket_bytes:\n if current_bucket:\n yield current_bucket\n current_bucket = [(name, tensor.clone())]\n current_size = tensor_size\n else:\n current_bucket.append((name, tensor.clone()))\n current_size += tensor_size\n\n if current_bucket:\n yield current_bucket\n"}165{"file_name": "verl__workers__rollout__trtllm_rollout__trtllm_async_server.py", "text": "# Copyright 2026 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport asyncio\nimport logging\nimport os\nfrom typing import Any, Optional\n\nimport ray\nimport torch\nfrom omegaconf import DictConfig\nfrom ray.actor import ActorHandle\nfrom ray.util import placement_group_table\nfrom ray.util.placement_group import PlacementGroup\n\nfrom verl.single_controller.ray import RayClassWithInitArgs, SubRayResourcePool\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.utils.net_utils import is_valid_ipv6_address\nfrom verl.workers.config import HFModelConfig, RolloutConfig\nfrom verl.workers.rollout.replica import RolloutMode, RolloutReplica, TokenOutput\nfrom verl.workers.rollout.trtllm_rollout.trtllm_rollout import ServerAdapter\nfrom verl.workers.rollout.utils import get_max_position_embeddings, run_unvicorn\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(logging.INFO)\n\n\n@ray.remote\nclass TRTLLMHttpServer:\n \"\"\"TensorRT LLM HTTP server in single node.\n\n Args:\n config (DictConfig): full config.\n model_config (HFModelConfig): model config.\n is_reward_model (bool): whether this is a reward model.\n rollout_mode (RolloutMode): rollout mode.\n workers (list[ActorHandle]): list of rollout workers.\n replica_rank (int): replica rank, a replica may contain multiple nodes.\n max_colocate_count (int): max colocate count.\n pgs (list[PlacementGroup]): placement groups.\n bundle_indices (list[list[int]]): bundle indices.\n \"\"\"\n\n def __init__(\n self,\n config: RolloutConfig,\n model_config: HFModelConfig,\n is_reward_model: bool,\n rollout_mode: RolloutMode,\n workers: list[ActorHandle],\n replica_rank: int,\n max_colocate_count: int,\n pgs: list[PlacementGroup] = None,\n bundle_indices: list[list[int]] = None,\n ):\n os.environ[\"TRT_LLM_DISABLE_LOAD_WEIGHTS_IN_PARALLEL\"] = \"1\"\n assert torch.cuda.is_available(), \"TRTLLM http server should run on GPU node\"\n\n self.config: RolloutConfig = omega_conf_to_dataclass(config)\n self.model_config: HFModelConfig = omega_conf_to_dataclass(model_config, dataclass_type=HFModelConfig)\n self.is_reward_model = is_reward_model\n max_position_embeddings = get_max_position_embeddings(self.model_config.hf_config)\n if self.config.max_model_len is None:\n self.config.max_model_len = max_position_embeddings\n else:\n if self.config.max_model_len > max_position_embeddings:\n raise ValueError(\n f\"max_model_len ({self.config.max_model_len}) should be less than or equal to \"\n f\"max_position_embeddings ({max_position_embeddings})\"\n )\n self.rollout_mode = rollout_mode\n self.workers = workers\n self.replica_rank = replica_rank\n self.max_colocate_count = max_colocate_count\n self.pgs = pgs\n self.bundle_indices = bundle_indices\n\n if self.rollout_mode != RolloutMode.HYBRID and self.config.load_format == \"dummy\":\n logger.warning(f\"rollout mode is {self.rollout_mode}, load_format is dummy, set to auto\")\n self.config.load_format = \"auto\"\n\n # used for http server\n self._server_address = ray.util.get_node_ip_address().strip(\"[]\")\n self._server_port = None\n\n logger.info(f\"TRTLLMHttpServer, replica_rank: {self.replica_rank}\")\n\n self.sampling_args = {\n \"detokenize\": False,\n \"end_id\": -1,\n \"pad_id\": self.model_config.hf_config.pad_token_id,\n \"stop_token_ids\": [self.model_config.hf_config.eos_token_id],\n \"include_stop_str_in_output\": True,\n }\n\n def get_server_address(self):\n \"\"\"Get http server address and port.\"\"\"\n assert self._server_port is not None, \"http server is not launched, port is None\"\n return self._server_address, self._server_port\n\n async def launch_server(self):\n from tensorrt_llm import AsyncLLM\n from tensorrt_llm.llmapi import CapacitySchedulerPolicy, CudaGraphConfig, KvCacheConfig, SchedulerConfig\n from tensorrt_llm.serve import OpenAIServer\n\n assert self.config.pipeline_model_parallel_size == 1, \"pipeline_model_parallel_size > 1 is not supported yet\"\n\n engine_kwargs = self.config.get(\"engine_kwargs\", {}).get(\"trtllm\", {}) or {}\n kv_cache_config = KvCacheConfig(\n enable_block_reuse=self.config.enable_prefix_caching,\n free_gpu_memory_fraction=self.config.gpu_memory_utilization,\n )\n\n per_worker_gpu_share = 1.0 / self.max_colocate_count\n\n llm_kwargs = {\n \"model\": self.model_config.local_path,\n \"backend\": \"pytorch\",\n \"dtype\": self.config.dtype,\n \"enable_chunked_prefill\": self.config.enable_chunked_prefill,\n \"skip_tokenizer_init\": self.config.skip_tokenizer_init,\n \"orchestrator_type\": \"ray\",\n \"ray_worker_extension_cls\": \"tensorrt_llm.llmapi.rlhf_utils.WorkerExtension\",\n \"kv_cache_config\": kv_cache_config,\n \"max_seq_len\": self.config.max_model_len,\n \"max_batch_size\": self.config.max_num_seqs,\n \"max_num_tokens\": self.config.max_num_batched_tokens,\n \"tensor_parallel_size\": self.config.tensor_model_parallel_size,\n \"pipeline_parallel_size\": self.config.pipeline_model_parallel_size,\n \"moe_expert_parallel_size\": self.config.expert_parallel_size,\n \"trust_remote_code\": self.model_config.trust_remote_code,\n \"placement_groups\": self.pgs,\n \"placement_bundle_indices\": self.bundle_indices,\n \"per_worker_gpu_share\": per_worker_gpu_share,\n \"enable_sleep\": self.config.enable_sleep_mode,\n \"allreduce_strategy\": \"NCCL\",\n \"sampler_type\": \"TRTLLMSampler\",\n **engine_kwargs,\n }\n\n if self.is_reward_model:\n llm_kwargs.update(\n {\n \"cuda_graph_config\": None,\n \"disable_overlap_scheduler\": True,\n }\n )\n else:\n llm_kwargs.update(\n {\n \"cuda_graph_config\": CudaGraphConfig(\n enable_padding=True,\n batch_sizes=self.config.cudagraph_capture_sizes,\n max_batch_size=0 if self.config.cudagraph_capture_sizes else self.config.max_num_seqs,\n ),\n \"scheduler_config\": SchedulerConfig(\n capacity_scheduler_policy=CapacitySchedulerPolicy.MAX_UTILIZATION,\n ),\n }\n )\n\n self.llm = await AsyncLLM(**llm_kwargs)\n\n trtllm_server = OpenAIServer(\n llm=self.llm,\n model=self.model_config.local_path,\n tool_parser=None,\n server_role=None,\n metadata_server_cfg=None,\n )\n app = trtllm_server.app\n self._server_port, self._server_task = await run_unvicorn(app, None, self._server_address)\n\n async def generate(\n self,\n prompt_ids: list[int],\n sampling_params: dict[str, Any],\n request_id: str,\n image_data: Optional[list[Any]] = None,\n video_data: Optional[list[Any]] = None,\n ) -> TokenOutput:\n \"\"\"Generate sequence with token-in-token-out.\"\"\"\n assert image_data is None and video_data is None, \"Multimodality is not yet supported in TRTLLMHttpServer.\"\n\n from tensorrt_llm.llmapi import SamplingParams\n\n max_tokens = min(self.config.response_length, self.config.max_model_len - len(prompt_ids))\n sampling_params[\"max_tokens\"] = max_tokens\n sampling_params[\"logprobs\"] = 1 if sampling_params.pop(\"logprobs\", False) else None\n if sampling_params[\"top_k\"] == -1:\n sampling_params[\"top_k\"] = 0\n sampling_params.update(self.sampling_args)\n\n trt_llm_sampling_params = SamplingParams(**sampling_params)\n outputs = await self.llm.generate_async(\n inputs=prompt_ids,\n sampling_params=trt_llm_sampling_params,\n )\n\n token_ids = outputs.outputs[0].token_ids\n log_probs = None\n if trt_llm_sampling_params.logprobs is not None:\n log_probs = [list(d.values())[0].logprob for d in outputs.outputs[0].logprobs]\n return TokenOutput(token_ids=token_ids, log_probs=log_probs)\n\n async def wake_up(self):\n if self.rollout_mode == RolloutMode.HYBRID:\n # In hybrid mode, rollout is wake up in `update_weights`\n raise ValueError(f\"wake_up not support rollout_mode {self.rollout_mode}\")\n if self.rollout_mode == RolloutMode.COLOCATED:\n await self.llm.resume(tags=ServerAdapter.get_full_tags())\n elif self.rollout_mode == RolloutMode.STANDALONE:\n logger.info(\"skip wake_up in standalone mode\")\n\n async def sleep(self):\n if not self.config.free_cache_engine:\n return\n\n if self.rollout_mode == RolloutMode.HYBRID:\n await self.llm.release(tags=ServerAdapter.get_full_tags())\n elif self.rollout_mode == RolloutMode.COLOCATED:\n await self.llm.release(tags=ServerAdapter.get_full_tags())\n elif self.rollout_mode == RolloutMode.STANDALONE:\n logger.info(\"skip sleep in standalone mode\")\n\n async def report_device_ids(self) -> list[str]:\n \"\"\"Report GPU device UUIDs from TRT-LLM workers.\"\"\"\n return await self.llm.collective_rpc(\n \"report_device_id\",\n unique_reply_rank=0,\n )\n\n\n_rollout_worker_actor_cls = ray.remote(ServerAdapter)\n\n\nclass TRTLLMReplica(RolloutReplica):\n def __init__(\n self,\n replica_rank: int,\n config: RolloutConfig,\n model_config: DictConfig,\n gpus_per_node: int = 8,\n is_reward_model: bool = False,\n ) -> None:\n super().__init__(replica_rank, config, model_config, gpus_per_node, is_reward_model)\n self.node_ip = ray.util.get_node_ip_address().strip(\"[]\")\n\n def get_ray_class_with_init_args(self) -> RayClassWithInitArgs:\n \"\"\"Get rollout worker actor class for colocated and standalone mode.\"\"\"\n worker_dict_cls = RayClassWithInitArgs(\n cls=_rollout_worker_actor_cls,\n config=self.config,\n model_config=self.model_config,\n device_mesh=None,\n replica_rank=self.replica_rank,\n )\n return worker_dict_cls\n\n def rollout_worker_use_gpu(self) -> bool:\n return False\n\n def get_pgs_and_bundle_indices(self) -> tuple[list[PlacementGroup], list[list[int]]]:\n \"\"\"Get placement groups and bundle indices for the replica.\"\"\"\n\n start_pg_index = 0\n local_bundle_index = 0\n\n # For SubRayResourcePool, the replica is assigned sub pool specific for this replica.\n if isinstance(self.resource_pool, SubRayResourcePool):\n assert self.resource_pool.subgroup_world_size == self.world_size, (\n \"Subgroup world size must be equal to world size\"\n )\n local_bundle_index = self.resource_pool.start_bundle_index\n # For RayResourcePool, the replica is assigned to entire resource pool.\n # We need to find start pg index and local bundle index based on replica rank.\n else:\n local_bundle_index = self.world_size * self.replica_rank\n\n while local_bundle_index >= self.resource_pool.pgs[start_pg_index].bundle_count:\n start_pg_index += 1\n local_bundle_index -= self.resource_pool.pgs[start_pg_index].bundle_count\n assert (\n start_pg_index < len(self.resource_pool.pgs)\n and local_bundle_index < self.resource_pool.pgs[start_pg_index].bundle_count\n ), \"Start pg index or local bundle index out of range\"\n\n # Global Bundle View for Replica x 2 & TP=4:\n # ┌───────────────────┬───────────────────┐\n # │ Placement Group 0 │ Placement Group 1 │\n # ├────┬────┬────┬────┼────┬────┬────┬────┤\n # │ 0 │ 1 │ 2 │ 3 │ 0 │ 1 │ 2 │ 3 │\n # └────┴────┴────┴────┴────┴────┴────┴────┘\n # └───────────────┘ └───────────────┘\n # Replica 0 Replica 1\n # (4 GPUs) (4 GPUs)\n\n left_bundle_count = self.world_size\n\n pgs = []\n bundle_indices = []\n\n for pg in self.resource_pool.pgs[start_pg_index:]:\n if left_bundle_count == 0:\n break\n\n left_bundle_count_in_pg = min(left_bundle_count, pg.bundle_count - local_bundle_index)\n pg_bundle_indices = [local_bundle_index + idx for idx in range(left_bundle_count_in_pg)]\n pgs.append(pg)\n bundle_indices.append(pg_bundle_indices)\n left_bundle_count -= left_bundle_count_in_pg\n local_bundle_index = 0\n\n assert left_bundle_count == 0, \"all bundle indices should be assigned\"\n\n return pgs, bundle_indices\n\n async def launch_servers(self):\n assert self.nnodes == 1, \"TRTLLMReplica doesn't support multiple nodes for single replica yet.\"\n assert self.resource_pool.pgs is not None, \"placement groups are not initialized\"\n\n pgs, bundle_indices = self.get_pgs_and_bundle_indices()\n\n # Check server process should be launched on the same node as first bundle of first pg.\n first_pg_data = placement_group_table(pgs[0])\n node_id = first_pg_data[\"bundles_to_node_id\"][bundle_indices[0][0]]\n print(f\"TRTLLMReplica: {self.replica_rank}\")\n print(f\"pg node_id: {node_id}\")\n print(f\"pgs: {pgs}\")\n print(f\"bundle_indices: {bundle_indices}\")\n\n # TRTLLMReplica is a 1:1 map from replica to TRTLLMHttpServer.\n name = (\n f\"trtllm_server_{self.replica_rank}\"\n if not self.is_reward_model\n else f\"trtllm_server_reward_{self.replica_rank}\"\n )\n\n server = TRTLLMHttpServer.options(\n scheduling_strategy=ray.util.scheduling_strategies.NodeAffinitySchedulingStrategy(\n node_id=node_id,\n soft=False,\n ),\n runtime_env={\"env_vars\": {\"RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES\": \"1\"}},\n name=name,\n ).remote(\n config=self.config,\n model_config=self.model_config,\n is_reward_model=self.is_reward_model,\n rollout_mode=self.rollout_mode,\n workers=self.workers,\n replica_rank=self.replica_rank,\n max_colocate_count=self.resource_pool.max_colocate_count,\n pgs=pgs,\n bundle_indices=bundle_indices,\n )\n self.servers.append(server)\n\n # launch http server in each node\n await asyncio.gather(*[server.launch_server.remote() for server in self.servers])\n\n # get http server address from first server\n server_address, server_port = await self.servers[0].get_server_address.remote()\n self._server_handle = self.servers[0]\n self._server_address = (\n f\"[{server_address}]:{server_port}\"\n if is_valid_ipv6_address(server_address)\n else f\"{server_address}:{server_port}\"\n )\n"}166{"file_name": "verl__workers__rollout__trtllm_rollout__trtllm_rollout.py", "text": "# Copyright 2026 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nfrom __future__ import annotations\n\nimport asyncio\nimport base64\nimport contextlib\nimport gc\nimport logging\nimport os\nimport pickle\nimport threading\nfrom contextlib import asynccontextmanager\nfrom typing import Any, Generator, Optional\n\nimport aiohttp\nimport pynvml\nimport ray\nimport torch\nimport torch.distributed as dist\nfrom torch.distributed.device_mesh import DeviceMesh, init_device_mesh\nfrom torch.multiprocessing.reductions import reduce_tensor\n\nfrom verl.utils.device import get_torch_device\nfrom verl.utils.net_utils import is_valid_ipv6_address\nfrom verl.workers.config import HFModelConfig, RolloutConfig\nfrom verl.workers.rollout.base import BaseRollout\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"WARN\"))\n\n# Default configuration constants\nDEFAULT_TIMEOUT = 60.0\nDEFAULT_MAX_ATTEMPTS = 3\nDEFAULT_RETRY_DELAY = 2.0\nDEFAULT_MAX_CONNECTIONS = 2000\nDEFAULT_MAX_WAIT_TIME = 300.0\n\n\n@contextlib.contextmanager\ndef nvml_context():\n \"\"\"Context manager for NVML initialization and shutdown.\n\n Raises:\n RuntimeError: If NVML initialization fails\n \"\"\"\n try:\n pynvml.nvmlInit()\n yield\n except pynvml.NVMLError as e:\n raise RuntimeError(f\"Failed to initialize NVML: {e}\") from e\n finally:\n try:\n pynvml.nvmlShutdown()\n except pynvml.NVMLError:\n pass\n\n\n_NVML_INITIALIZED = False\n_NVML_LOCK = threading.Lock()\n\n\ndef get_device_uuid(id: int) -> str:\n \"\"\"Get the UUID of a CUDA device using NVML.\"\"\"\n global _NVML_INITIALIZED\n with _NVML_LOCK:\n if not _NVML_INITIALIZED:\n try:\n pynvml.nvmlInit()\n _NVML_INITIALIZED = True\n except pynvml.NVMLError as e:\n raise RuntimeError(f\"Failed to initialize NVML: {e}\") from e\n\n # Get the device handle and UUID\n try:\n handle = pynvml.nvmlDeviceGetHandleByIndex(id)\n uuid = pynvml.nvmlDeviceGetUUID(handle)\n # Ensure the UUID is returned as a string, not bytes\n if isinstance(uuid, bytes):\n return uuid.decode(\"utf-8\")\n elif isinstance(uuid, str):\n return uuid\n else:\n raise RuntimeError(f\"Unexpected UUID type: {type(uuid)} for device {id} (global index: {id})\")\n except pynvml.NVMLError as e:\n raise RuntimeError(f\"Failed to get device UUID for device {id} (global index: {id}): {e}\") from e\n\n\nasync def _read_async_response(resp: aiohttp.ClientResponse) -> dict[str, Any]:\n if resp.status == 204 or (resp.content_length == 0):\n return {}\n\n try:\n return await resp.json(content_type=None)\n except Exception:\n try:\n text = await resp.text()\n except Exception:\n return {}\n return {\n \"content_type\": (resp.headers.get(\"Content-Type\") or \"\"),\n \"text\": text,\n }\n\n\nclass AsyncTRTLLMHttpAdapter:\n def __init__(\n self,\n host: str,\n port: int,\n timeout: float = DEFAULT_TIMEOUT,\n max_attempts: int = DEFAULT_MAX_ATTEMPTS,\n retry_delay: float = DEFAULT_RETRY_DELAY,\n max_connections: int = DEFAULT_MAX_CONNECTIONS,\n ):\n self.host = host\n self.port = port\n self.timeout = timeout\n self.max_attempts = max_attempts\n self.retry_delay = retry_delay\n self.max_connections = max_connections\n\n @asynccontextmanager\n async def _get_session(self) -> aiohttp.ClientSession:\n \"\"\"Context manager for safe session access with proper connection pooling.\n\n Yields:\n aiohttp.ClientSession: Session instance for making HTTP requests\n\n Note:\n This method creates a new session for each request to avoid resource competition\n while still maintaining proper connection pooling through the shared connector.\n \"\"\"\n # Create a new session for each request to avoid resource competition\n connector = aiohttp.TCPConnector(\n limit=self.max_connections,\n limit_per_host=self.max_connections // 4,\n ttl_dns_cache=300,\n use_dns_cache=True,\n )\n timeout = aiohttp.ClientTimeout(total=self.timeout)\n session = aiohttp.ClientSession(connector=connector, timeout=timeout)\n\n try:\n yield session\n finally:\n # Always close the session to free up resources\n if not session.closed:\n await session.close()\n\n async def _make_async_request(\n self,\n endpoint: str,\n payload: Optional[dict[str, Any]] = None,\n timeout: float = DEFAULT_TIMEOUT,\n method: str = \"POST\",\n return_status: bool = False,\n ) -> dict[str, Any] | int:\n \"\"\"Make an async HTTP request with retry logic and consistent error handling.\n\n Args:\n endpoint (str): The API endpoint to call (without leading slash)\n payload (Optional[Dict[str, Any]], optional): The JSON payload to send.\n Defaults to empty dict if None.\n method (str, optional): HTTP method to use. Defaults to \"POST\".\n\n Returns:\n Dict[str, Any]: The JSON response from the server\n\n Raises:\n aiohttp.ClientResponseError: If the HTTP request fails with a client/server error\n RuntimeError: If all retry attempts are exhausted\n\n Note:\n - Uses exponential backoff for retries\n - Logs warnings for timeout and connection errors, errors for HTTP errors\n \"\"\"\n\n url = f\"http://{self.host}:{self.port}/{endpoint}\"\n\n for attempt in range(self.max_attempts):\n try:\n async with self._get_session() as session:\n if method.upper() == \"GET\":\n async with session.get(url, timeout=timeout) as response:\n response.raise_for_status()\n return response.status if return_status else await _read_async_response(response)\n else:\n async with session.post(url, json=payload or {}, timeout=timeout) as response:\n response.raise_for_status()\n return response.status if return_status else await _read_async_response(response)\n\n except asyncio.TimeoutError:\n logger.warning(f\"Async request to {endpoint} timed out (attempt {attempt + 1})\")\n except aiohttp.ClientConnectorError:\n logger.warning(f\"Connection error for {endpoint} (attempt {attempt + 1})\")\n except aiohttp.ClientResponseError as e:\n logger.error(f\"HTTP error for {endpoint}: {e}\")\n raise\n except Exception as e:\n logger.error(f\"Unexpected error for {endpoint}: {e}\")\n if attempt == self.max_attempts - 1:\n raise\n\n if attempt < self.max_attempts - 1:\n await asyncio.sleep(self.retry_delay * (2**attempt))\n\n raise RuntimeError(f\"Failed to complete async request to {endpoint} after {self.max_attempts} attempts\")\n\n async def resume_memory_occupation(self, tags: list[str]):\n \"\"\"Resume GPU memory occupation (async version).\n\n Similar to AsyncEngine, this method handles first-time weight reloading\n by calling release_memory_occupation if needed.\n\n Args:\n tags (Optional[List[str]], optional): List of tags to specify which memory to resume.\n If None, resumes all memory. Defaults to None. [\"weights\", \"kv_cache\"]\n\n Returns:\n Dict[str, Any]: Server response indicating memory resume status\n \"\"\"\n return await self._make_async_request(\"resume_memory\", {\"tags\": tags})\n\n async def release_memory_occupation(self, tags: list[str]):\n \"\"\"Release GPU memory occupation temporarily (async version).\n\n Args:\n tags (Optional[List[str]], optional): List of tags to specify which memory to release.\n If None, releases all memory. Defaults to None. [\"weights\", \"kv_cache\"]\n\n Returns:\n Dict[str, Any]: Server response indicating memory release status\n \"\"\"\n return await self._make_async_request(\"release_memory\", {\"tags\": tags})\n\n async def update_weights(self, weights: dict[str, str]):\n \"\"\"Update model weights from tensor data asynchronously.\n\n Args:\n weights: A dictionary that maps the device uuid of the weight handles.\n\n Returns:\n Dict[str, Any]: Server response containing update status\n \"\"\"\n return await self._make_async_request(\"update_weights\", {\"weights\": weights})\n\n\nclass ServerAdapter(BaseRollout):\n _WEIGHTS_TAGS = [\n \"sampler\",\n \"drafter\",\n \"guided_decoder\",\n \"spec_resource_manager\",\n \"model_extra\",\n \"executor_extra\",\n \"model\",\n \"draft_model\",\n ]\n\n @staticmethod\n def get_full_tags() -> list[str]:\n return ServerAdapter._WEIGHTS_TAGS + [\"kv_cache\"]\n\n def __init__(\n self, config: RolloutConfig, model_config: HFModelConfig, device_mesh: DeviceMesh, replica_rank: int = -1\n ):\n super().__init__(config, model_config, device_mesh)\n self._adapter = None\n self.hybrid_device_mesh = None\n self.gpu_id = None\n self.is_leader_rank = None\n self.replica_rank = None\n self.is_dp_rank = None\n\n # hybrid mode\n if self.device_mesh is not None:\n assert device_mesh.mesh_dim_names.index(\"dp\") == 0, \"DP dim should always be the first dimension\"\n\n # Clone a new device mesh for CPU backend only (used for internal ranks communication)\n device_mesh_kwargs = dict(\n mesh_shape=device_mesh.mesh.shape,\n mesh_dim_names=device_mesh.mesh_dim_names,\n )\n self.hybrid_device_mesh = init_device_mesh(\"cpu\", **device_mesh_kwargs)\n\n self.hybrid_device_mesh[self.hybrid_device_mesh.mesh_dim_names[1:]]._flatten(mesh_dim_name=\"exclude_dp\")\n self.is_leader_rank = self.hybrid_device_mesh[\"exclude_dp\"].get_local_rank() == 0\n logger.info(f\"is_dp_leader: {self.is_leader_rank}\")\n logger.info(f\"exclude_dp_rank = {self.hybrid_device_mesh['exclude_dp'].get_local_rank()}\")\n logger.info(f\"exclude_dp_size = {self.hybrid_device_mesh['exclude_dp'].size()}\")\n self.gpu_id = ray.get_gpu_ids()[0]\n self.replica_rank = self.hybrid_device_mesh[\"dp\"].get_local_rank()\n assert len(ray.get_gpu_ids()) == 1, \"ServerAdapter should run on a single GPU node\"\n else:\n rank = int(os.environ[\"RANK\"])\n self.replica_rank = replica_rank\n self.is_leader_rank = rank == 0\n\n # Below is required for all modes.\n assert self.replica_rank >= 0, \"replica_rank is not set\"\n assert self.is_leader_rank is not None, \"is_leader_rank is not set\"\n\n self.node_ip = ray.util.get_node_ip_address().strip(\"[]\")\n\n async def _init_server_adapter(self):\n if self._adapter is not None:\n return\n\n # Lazy init http server adapter because http server is launched after hybrid engine.\n self.server_actor = ray.get_actor(f\"trtllm_server_{self.replica_rank}\")\n server_address, server_port = await self.server_actor.get_server_address.remote()\n assert server_address == self.node_ip, f\"server address: {server_address} != node_ip: {self.node_ip}\"\n\n logger.debug(f\"replica_rank={self.replica_rank}, server address: {server_address}, port: {server_port}\")\n host = f\"[{server_address}]\" if is_valid_ipv6_address(server_address) else server_address\n self._adapter = AsyncTRTLLMHttpAdapter(\n host=host,\n port=server_port,\n timeout=self.config.server.timeout,\n max_attempts=self.config.server.max_attempts,\n retry_delay=self.config.server.retry_delay,\n max_connections=self.config.server.max_connections,\n )\n\n async def resume(self, tags: list[str]):\n \"\"\"Resume rollout weights or kv cache in GPU memory.\n\n Args:\n tag: weights or kv_cache.\n \"\"\"\n # Synchronize all ranks before resuming KV cache to ensure non-leader ranks\n # have completed actor offloading to CPU, preventing OOM issue.\n if \"kv_cache\" in tags and self.config.free_cache_engine:\n await asyncio.to_thread(dist.barrier, group=self.hybrid_device_mesh[\"exclude_dp\"].get_group())\n if self.is_leader_rank and self.config.free_cache_engine:\n if \"weights\" in tags:\n tags = self._WEIGHTS_TAGS\n elif \"kv_cache\" in tags:\n tags = [\"kv_cache\"]\n else:\n raise ValueError(f\"Invalid tag: {tags}\")\n await self._init_server_adapter()\n await self._adapter.resume_memory_occupation(tags=tags)\n\n async def release(self):\n \"\"\"Release weights and kv cache in GPU memory.\"\"\"\n if self.is_leader_rank and self.config.free_cache_engine:\n await self._init_server_adapter()\n tags = self._WEIGHTS_TAGS + [\"kv_cache\"]\n await self._adapter.release_memory_occupation(tags=tags)\n\n async def update_weights_from_ipc_handles(self, device_handles):\n assert self.hybrid_device_mesh is not None, \"hybrid_device_mesh is not set\"\n\n \"\"\"Update weights from IPC handles.\"\"\"\n if self.is_leader_rank:\n gathered_handles = [None for _ in range(self.hybrid_device_mesh[\"exclude_dp\"].size())]\n else:\n gathered_handles = None\n\n await asyncio.to_thread(\n dist.gather_object,\n obj=device_handles,\n object_gather_list=gathered_handles,\n group_dst=0,\n group=self.hybrid_device_mesh[\"exclude_dp\"].get_group(),\n )\n\n if self.is_leader_rank:\n all_handles = {k: v for d in gathered_handles for k, v in d.items()}\n await self._adapter.update_weights(all_handles)\n\n await asyncio.to_thread(dist.barrier, group=self.hybrid_device_mesh[\"exclude_dp\"].get_group())\n\n async def update_weights(self, weights: Generator[tuple[str, torch.Tensor], None, None], **kwargs):\n assert self.hybrid_device_mesh is not None, \"hybrid_device_mesh is not set\"\n\n \"\"\"Update the weights of the rollout model.\n\n Args:\n weights: A generator that yields the name of the weight tensor and the tensor itself.\n \"\"\"\n if self.is_leader_rank:\n await self._init_server_adapter()\n\n total_available_bytes = int(self.config.checkpoint_engine.update_weights_bucket_megabytes) * 1024 * 1024\n\n try:\n device_uuid = get_device_uuid(self.gpu_id)\n except Exception as e:\n logger.error(f\"Failed to get device UUID in update_weights(): {e}\")\n device_uuid = None\n raise e\n\n cur_available_bytes = total_available_bytes\n cur_handles = []\n\n async def flush():\n nonlocal cur_available_bytes, cur_handles\n if not cur_handles:\n return\n serialized_device_handles = {device_uuid: base64.b64encode(pickle.dumps(cur_handles)).decode(\"utf-8\")}\n await self.update_weights_from_ipc_handles(serialized_device_handles)\n cur_available_bytes = total_available_bytes\n cur_handles = []\n\n for name, param in weights:\n size_in_bytes = param.element_size() * param.numel()\n if size_in_bytes > cur_available_bytes:\n await flush()\n\n assert cur_available_bytes >= size_in_bytes, (\n f\"cur_available_bytes: {cur_available_bytes:,} size_in_bytes: {size_in_bytes:,} name: {name}\"\n )\n cur_available_bytes -= size_in_bytes\n handle = reduce_tensor(param.detach())\n cur_handles.append((name, handle))\n\n await flush()\n\n if self.is_leader_rank:\n # Finalize update weights\n await self._adapter.update_weights(None)\n await asyncio.to_thread(dist.barrier, group=self.hybrid_device_mesh[\"exclude_dp\"].get_group())\n\n gc.collect()\n get_torch_device().empty_cache()\n\n def _get_attribute(self, name: str):\n return getattr(self, name)\n"}167{"file_name": "verl__workers__rollout__vllm_rollout__vllm_async_server.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\nimport argparse\nimport asyncio\nimport inspect\nimport json\nimport logging\nimport os\nfrom pprint import pprint\nfrom typing import Any, Callable, Optional\n\nimport numpy as np\nimport ray\nimport vllm.entrypoints.cli.serve\nfrom packaging import version\nfrom ray.actor import ActorHandle\nfrom vllm import SamplingParams\nfrom vllm.engine.arg_utils import AsyncEngineArgs\nfrom vllm.entrypoints.cli.serve import run_headless\nfrom vllm.entrypoints.openai.api_server import build_app, init_app_state\nfrom vllm.inputs import TokensPrompt\nfrom vllm.lora.request import LoRARequest\nfrom vllm.outputs import RequestOutput\nfrom vllm.usage.usage_lib import UsageContext\nfrom vllm.v1.engine.async_llm import AsyncLLM\n\nfrom verl.single_controller.ray import RayClassWithInitArgs\nfrom verl.utils.config import omega_conf_to_dataclass\nfrom verl.utils.device import get_resource_name, get_visible_devices_keyword\nfrom verl.utils.net_utils import get_free_port, is_valid_ipv6_address\nfrom verl.utils.profiler import DistProfiler, build_vllm_profiler_args\nfrom verl.utils.vllm.vllm_fp8_utils import apply_vllm_fp8_patches\nfrom verl.workers.config import HFModelConfig, RolloutConfig\nfrom verl.workers.rollout.replica import RolloutMode, RolloutReplica, TokenOutput\nfrom verl.workers.rollout.utils import get_max_position_embeddings, run_unvicorn\nfrom verl.workers.rollout.vllm_rollout import ServerAdapter\nfrom verl.workers.rollout.vllm_rollout.utils import (\n VLLM_LORA_INT_ID,\n VLLM_LORA_NAME,\n VLLM_LORA_PATH,\n SuppressSignalInThread,\n build_cli_args_from_config,\n get_vllm_max_lora_rank,\n)\n\n_VLLM_VERSION = version.parse(vllm.__version__)\n\nif _VLLM_VERSION > version.parse(\"0.11.0\"):\n from vllm.utils.argparse_utils import FlexibleArgumentParser\n\n if _VLLM_VERSION == version.parse(\"0.12.0\"):\n from vllm.entrypoints.harmony_utils import get_encoding\n\n elif _VLLM_VERSION >= version.parse(\"0.13.0\"):\n from vllm.entrypoints.openai.parser.harmony_utils import get_encoding\n\n else:\n get_encoding = None\n\n if get_encoding is not None and os.getenv(\"VERL_USE_GPT_OSS\", \"0\") == \"1\":\n get_encoding()\nelse:\n from vllm.utils import FlexibleArgumentParser\n\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(logging.INFO)\n\n\nclass vLLMHttpServer:\n \"\"\"vLLM http server in single node, this is equivalent to launch server with command line:\n ```\n vllm serve --tensor-parallel-size=8 ...\n ```\n \"\"\"\n\n def __init__(\n self,\n config: RolloutConfig,\n model_config: HFModelConfig,\n rollout_mode: RolloutMode,\n workers: list[ActorHandle],\n replica_rank: int,\n node_rank: int,\n gpus_per_node: int,\n nnodes: int,\n cuda_visible_devices: str,\n ):\n \"\"\"\n Args:\n config (RolloutConfig): full config.\n model_config (HFModelConfig): model config.\n rollout_mode (RolloutMode): rollout mode.\n replica_rank (int): replica rank, a replica may contain multiple nodes.\n node_rank (int): node rank.\n gpus_per_node (int): number of gpus per node.\n nnodes (int): number of nodes.\n cuda_visible_devices (str): cuda visible devices.\n \"\"\"\n os.environ[get_visible_devices_keyword()] = cuda_visible_devices\n\n self.config: RolloutConfig = omega_conf_to_dataclass(config)\n self.model_config: HFModelConfig = omega_conf_to_dataclass(model_config, dataclass_type=HFModelConfig)\n max_position_embeddings = get_max_position_embeddings(self.model_config.hf_config)\n if self.config.max_model_len is None:\n self.config.max_model_len = max_position_embeddings\n else:\n if self.config.max_model_len > max_position_embeddings:\n raise ValueError(\n f\"max_model_len ({self.config.max_model_len}) should be less than or equal to \"\n f\"max_position_embeddings ({max_position_embeddings})\"\n )\n\n self.rollout_mode = rollout_mode\n self.workers = workers\n\n self.replica_rank = replica_rank\n self.node_rank = node_rank\n self.gpus_per_node = gpus_per_node\n self.nnodes = nnodes\n\n if self.rollout_mode != RolloutMode.HYBRID and self.config.load_format == \"dummy\":\n logger.warning(f\"rollout mode is {self.rollout_mode}, load_format is dummy, set to auto\")\n self.config.load_format = \"auto\"\n\n # used for http server\n self._server_address = ray.util.get_node_ip_address().strip(\"[]\")\n self._server_port = None\n\n # used for controlling vllm server profiler\n profiler_config = self.config.profiler\n tool_config = None\n if profiler_config is not None:\n if profiler_config.tool in [\"torch\", \"npu\"]:\n tool_config = omega_conf_to_dataclass((profiler_config.tool_config or {}).get(profiler_config.tool))\n else:\n logger.warning(f\"agent loop only support torch and npu profiler, got {profiler_config.tool}\")\n profiler_config = None\n self.profiler_controller = DistProfiler(self.replica_rank, config=profiler_config, tool_config=tool_config)\n\n # used for data parallel: --data-parallel-address, --data-parallel-rpc-port\n if self.node_rank == 0:\n self._master_address = self._server_address\n # used for torch.distributed.init_process_group\n self._master_port, self._master_sock = get_free_port(self._server_address)\n # used for data parallel: --data-parallel-address, --data-parallel-rpc-port\n self._dp_rpc_port, self._dp_rpc_sock = get_free_port(self._server_address)\n self._dp_master_port, self._dp_master_sock = get_free_port(self._server_address)\n else:\n self._master_address = None\n self._master_port = None\n self._dp_rpc_port = None\n self._dp_master_port = None\n\n logger.info(\n f\"vLLMHttpServer, replica_rank: {self.replica_rank}, node_rank: {self.node_rank}, \"\n f\"{get_visible_devices_keyword()}: {cuda_visible_devices}, \"\n f\"master_address: {self._master_address}, master_port: {self._master_port}, \"\n f\"data_parallel_rpc_port: {self._dp_rpc_port}, data_parallel_master_port: {self._dp_master_port}\"\n )\n\n def get_master_address(self):\n \"\"\"Get master address and port for data parallel.\n Returns:\n tuple: (master_address, master_port, dp_rpc_port)\n \"\"\"\n return self._master_address, self._master_port, self._dp_rpc_port\n\n def get_server_address(self):\n \"\"\"Get http server address and port.\"\"\"\n assert self._server_port is not None, \"http server is not launched, port is None\"\n return self._server_address, self._server_port\n\n async def collective_rpc(\n self,\n method: str | Callable,\n timeout: float | None = None,\n args: tuple = (),\n kwargs: dict[str, Any] | None = None,\n ):\n await self.engine.collective_rpc(\n method=method,\n timeout=timeout,\n args=args,\n kwargs=kwargs,\n )\n\n async def launch_server(self, master_address: str = None, master_port: int = None, dp_rpc_port: int = None):\n if self.node_rank != 0:\n assert master_address and master_port and dp_rpc_port, (\n \"non-master node should provide master_address, master_port and dp_rpc_port\"\n )\n self._master_address = master_address\n self._master_port = master_port\n self._dp_rpc_port = dp_rpc_port\n\n # 1. setup vllm serve cli args\n engine_kwargs = self.config.get(\"engine_kwargs\", {}).get(\"vllm\", {}) or {}\n engine_kwargs = {key: val for key, val in engine_kwargs.items() if val is not None}\n if self.config.get(\"limit_images\", None): # support for multi-image data\n engine_kwargs[\"limit_mm_per_prompt\"] = {\"image\": self.config.get(\"limit_images\")}\n if self.config.cudagraph_capture_sizes:\n engine_kwargs[\"cuda_graph_sizes\"] = self.config.cudagraph_capture_sizes\n\n # Override default generation config from hugging face model config,\n # user can still override them by passing kwargs in each request.\n override_generation_config = dict(\n temperature=self.config.temperature,\n top_k=self.config.top_k,\n top_p=self.config.top_p,\n repetition_penalty=1.0,\n max_new_tokens=self.config.response_length,\n )\n logger.info(f\"override_generation_config: {override_generation_config}\")\n\n logger.info(f\"enable_sleep_mode: {self.config.enable_sleep_mode}\")\n if not self.config.enable_sleep_mode:\n from verl.utils.device import set_expandable_segments\n\n set_expandable_segments(True)\n\n quantization = self.config.quantization\n hf_overrides = {}\n\n # Handle QAT (Quantization-Aware Training) configuration\n qat_config_dict = getattr(self.config, \"qat\", {}) or {}\n if qat_config_dict.get(\"enable\", False):\n # QAT uses compressed-tensors quantization, apply patches for dynamic weight loading\n from verl.utils.qat import QATConfig, apply_qat_patches, load_quantization_config\n\n apply_qat_patches()\n\n # Load quantization config from JSON file\n qat_config = QATConfig(**qat_config_dict)\n quantization_config_dict = load_quantization_config(qat_config)\n hf_overrides[\"quantization_config\"] = quantization_config_dict\n quantization = \"compressed-tensors\"\n\n logger.info(\"QAT quantization config injected to vLLM async server\")\n elif quantization is not None:\n # Handle other quantization methods (fp8, torchao)\n _SUPPORTED_QUANTIZATION = [\"fp8\", \"torchao\"]\n if quantization not in _SUPPORTED_QUANTIZATION:\n raise ValueError(f\"Currently only support {_SUPPORTED_QUANTIZATION} quantization, got: {quantization}\")\n\n if quantization == \"fp8\":\n # Ignore MoE router layers for FP8 quantization\n all_mlp_gate_layers = []\n for layer in range(self.model_config.hf_config.num_hidden_layers):\n all_mlp_gate_layers.append(f\"model.layers.{layer}.mlp.gate\")\n\n FP8_BLOCK_QUANT_KWARGS = {\n \"activation_scheme\": \"dynamic\",\n \"fmt\": \"e4m3\",\n \"quant_method\": \"fp8\",\n \"weight_block_size\": [128, 128],\n \"ignored_layers\": all_mlp_gate_layers,\n }\n hf_overrides[\"quantization_config\"] = dict(FP8_BLOCK_QUANT_KWARGS)\n # Apply vllm fp8 patches\n # Will remove the patch after vllm support on-the-fly quant for rollout natively.\n apply_vllm_fp8_patches()\n # for subprocesses patching\n os.environ[\"VERL_VLLM_FP8_QUANT_ENABLED\"] = \"1\"\n\n if quantization is not None and self.config.quantization_config_file is not None:\n hf_overrides[\"quantization_config_file\"] = self.config.quantization_config_file\n\n compilation_config = engine_kwargs.pop(\"compilation_config\", None) or {}\n if isinstance(compilation_config, str):\n compilation_config = json.loads(compilation_config)\n compilation_config.setdefault(\"cudagraph_mode\", \"FULL_AND_PIECEWISE\")\n\n # FULL cuda graph is not yet supported with DCP, downgrade to PIECEWISE\n dcp_size = engine_kwargs.get(\"decode_context_parallel_size\", 1) or 1\n if dcp_size > 1 and compilation_config[\"cudagraph_mode\"] == \"FULL_AND_PIECEWISE\":\n logger.warning(\n \"FULL cuda graph is not supported with DCP (decode_context_parallel_size=%d), \"\n \"downgrading cudagraph_mode to PIECEWISE.\",\n dcp_size,\n )\n compilation_config[\"cudagraph_mode\"] = \"PIECEWISE\"\n\n compilation_config = json.dumps(compilation_config)\n args = {\n \"dtype\": self.config.dtype,\n \"load_format\": self.config.load_format,\n \"skip_tokenizer_init\": False,\n \"distributed_executor_backend\": \"mp\",\n \"worker_extension_cls\": \"verl.workers.rollout.vllm_rollout.utils.vLLMColocateWorkerExtension\",\n \"trust_remote_code\": self.model_config.trust_remote_code,\n \"max_model_len\": self.config.max_model_len,\n \"max_num_seqs\": self.config.max_num_seqs,\n \"enable_chunked_prefill\": self.config.enable_chunked_prefill,\n \"max_num_batched_tokens\": self.config.max_num_batched_tokens,\n \"enable_prefix_caching\": self.config.enable_prefix_caching,\n \"enable_sleep_mode\": self.config.enable_sleep_mode,\n \"logprobs_mode\": self.config.logprobs_mode,\n \"enforce_eager\": self.config.enforce_eager,\n \"gpu_memory_utilization\": self.config.gpu_memory_utilization,\n \"disable_log_stats\": self.config.disable_log_stats,\n \"tensor_parallel_size\": self.config.tensor_model_parallel_size,\n \"seed\": self.replica_rank + self.config.get(\"seed\", 0),\n \"override_generation_config\": json.dumps(override_generation_config),\n \"quantization\": quantization,\n \"hf_overrides\": hf_overrides,\n \"scheduling_policy\": self.config.scheduling_policy,\n \"compilation_config\": compilation_config,\n **engine_kwargs,\n }\n\n # update profiler args\n profiler_args = build_vllm_profiler_args(\n self.profiler_controller.config, self.profiler_controller.tool_config, self.replica_rank\n )\n if _VLLM_VERSION >= version.parse(\"0.13.0\"):\n # vLLM >= 0.13.0 supports profiler config via CLI args; env vars still work but will be deprecated\n args.update(profiler_args)\n\n if self.config.prometheus.enable:\n if self.config.prometheus.served_model_name:\n # Extract model name from path if it's a full path\n served_model_name = self.config.prometheus.served_model_name\n if \"/\" in served_model_name:\n # If it's a full path, extract the last part as model name\n served_model_name = served_model_name.split(\"/\")[-1]\n args[\"served_model_name\"] = served_model_name\n\n # mtp\n if self.config.mtp.enable and self.config.mtp.enable_rollout:\n speculative_config = {\n \"method\": self.config.mtp.method,\n \"num_speculative_tokens\": self.config.mtp.num_speculative_tokens,\n }\n args[\"speculative_config\"] = speculative_config\n\n if self.config.expert_parallel_size > 1:\n assert self.gpus_per_node % self.config.tensor_model_parallel_size == 0, (\n \"gpus_per_node should be divisible by tensor_model_parallel_size\"\n )\n data_parallel_size_local = self.gpus_per_node // self.config.tensor_model_parallel_size\n assert len(self.workers) == data_parallel_size_local * self.config.tensor_model_parallel_size, (\n f\"num workers ({len(self.workers)}) should be equal to dp_size_local \"\n )\n f\"({data_parallel_size_local}) * tp_size ({self.config.tensor_model_parallel_size})\"\n\n args.update(\n {\n \"enable_expert_parallel\": self.config.expert_parallel_size > 1,\n \"data_parallel_size\": self.config.data_parallel_size,\n \"data_parallel_size_local\": data_parallel_size_local,\n \"data_parallel_start_rank\": self.node_rank * data_parallel_size_local,\n \"data_parallel_address\": self._master_address,\n \"data_parallel_rpc_port\": self._dp_rpc_port,\n }\n )\n\n # used for torch.distributed.init_process_group\n if self.nnodes > 1:\n args.update(\n {\n \"master_addr\": self._master_address,\n \"master_port\": self._master_port,\n \"node_rank\": self.node_rank,\n \"nnodes\": self.nnodes,\n \"data_parallel_address\": self._master_address,\n \"data_parallel_rpc_port\": self._dp_rpc_port,\n }\n )\n\n # update lora-related args\n lora_rank = self.model_config.lora.get(\"rank\", 0)\n if lora_rank <= 0:\n lora_rank = (\n self.model_config.lora_rank\n ) # FIXME: fallback to lora_rank for now, we should unify lora settings.\n\n if self.model_config.lora.get(\"merge\", False):\n lora_rank = 0\n\n if lora_rank > 0:\n lora_args = {\n \"enable_lora\": True,\n \"max_loras\": 1,\n \"max_lora_rank\": get_vllm_max_lora_rank(lora_rank),\n }\n if self.model_config.lora.get(\"fully_sharded_loras\", False):\n lora_args[\"fully_sharded_loras\"] = True\n args.update(lora_args)\n\n if self.config.enable_rollout_routing_replay:\n args.update({\"enable_return_routed_experts\": True})\n\n server_args = [\"serve\", self.model_config.local_path] + build_cli_args_from_config(args)\n\n if self.replica_rank == 0:\n pprint(server_args)\n\n CMD_MODULES = [vllm.entrypoints.cli.serve]\n parser = FlexibleArgumentParser(description=\"vLLM CLI\")\n subparsers = parser.add_subparsers(required=False, dest=\"subparser\")\n cmds = {}\n for cmd_module in CMD_MODULES:\n new_cmds = cmd_module.cmd_init()\n for cmd in new_cmds:\n cmd.subparser_init(subparsers).set_defaults(dispatch_function=cmd.cmd)\n cmds[cmd.name] = cmd\n server_args = parser.parse_args(args=server_args)\n server_args.model = server_args.model_tag\n if server_args.subparser in cmds:\n cmds[server_args.subparser].validate(server_args)\n\n # 3. launch server\n if self.node_rank == 0:\n self._master_sock.close()\n await self.run_server(server_args)\n else:\n # TODO: avoid connect before master_sock close\n await asyncio.sleep(3)\n await self.run_headless(server_args)\n\n async def run_server(self, args: argparse.Namespace):\n engine_args = AsyncEngineArgs.from_cli_args(args)\n usage_context = UsageContext.OPENAI_API_SERVER\n vllm_config = engine_args.create_engine_config(usage_context=usage_context)\n vllm_config.parallel_config.data_parallel_master_port = self._dp_master_port\n\n fn_args = set(dict(inspect.signature(AsyncLLM.from_vllm_config).parameters).keys())\n kwargs = {}\n if \"enable_log_requests\" in fn_args:\n kwargs[\"enable_log_requests\"] = engine_args.enable_log_requests\n if \"disable_log_stats\" in fn_args:\n kwargs[\"disable_log_stats\"] = engine_args.disable_log_stats\n\n engine_client = AsyncLLM.from_vllm_config(vllm_config=vllm_config, usage_context=usage_context, **kwargs)\n\n # Don't keep the dummy data in memory\n await engine_client.reset_mm_cache()\n await engine_client.collective_rpc(\n method=\"monkey_patch_model\", kwargs={\"vocab_size\": len(self.model_config.tokenizer)}\n )\n\n build_app_sig = inspect.signature(build_app)\n supported_tasks: tuple[Any, ...] = ()\n if \"supported_tasks\" in build_app_sig.parameters:\n supported_tasks = await engine_client.get_supported_tasks()\n app = build_app(args, supported_tasks)\n else:\n app = build_app(args)\n\n init_app_sig = inspect.signature(init_app_state)\n if \"vllm_config\" in init_app_sig.parameters:\n await init_app_state(engine_client, vllm_config, app.state, args)\n elif \"supported_tasks\" in init_app_sig.parameters:\n await init_app_state(engine_client, app.state, args, supported_tasks)\n else:\n await init_app_state(engine_client, app.state, args)\n if self.replica_rank == 0 and self.node_rank == 0:\n logger.info(f\"Initializing a V1 LLM engine with config: {vllm_config}\")\n\n self.engine = engine_client\n self._server_port, self._server_task = await run_unvicorn(app, args, self._server_address)\n\n async def run_headless(self, args: argparse.Namespace):\n \"\"\"Run headless server in a separate thread.\"\"\"\n\n def run_headless_wrapper():\n with SuppressSignalInThread():\n run_headless(args)\n\n def on_run_headless_done(future: asyncio.Future):\n try:\n exc = future.exception()\n if exc:\n logger.exception(f\"run_headless failed with exception: {exc}\")\n else:\n logger.warning(\"run_headless completed successfully, but it's not expected.\")\n except Exception as e:\n logger.exception(f\"get result from run_headless failed: {e}\")\n finally:\n os._exit(1)\n\n self.task = asyncio.create_task(asyncio.to_thread(run_headless_wrapper))\n self.task.add_done_callback(on_run_headless_done)\n\n async def generate(\n self,\n prompt_ids: list[int],\n sampling_params: dict[str, Any],\n request_id: str,\n image_data: Optional[list[Any]] = None,\n video_data: Optional[list[Any]] = None,\n priority: int = 0,\n ) -> TokenOutput:\n \"\"\"Generate sequence with token-in-token-out.\"\"\"\n # Calculate the maximum possible new tokens based on available context space\n # This serves as a safety upper bound\n max_possible_tokens = self.config.max_model_len - len(prompt_ids)\n if max_possible_tokens < 0:\n raise ValueError(\n f\"Prompt length ({len(prompt_ids)}) exceeds the model's maximum context length \"\n f\"({self.config.max_model_len}).\"\n )\n\n # Determine max_tokens from sampling_params or use configured response_length as default\n if \"max_tokens\" in sampling_params:\n max_tokens = sampling_params.pop(\"max_tokens\")\n elif \"max_new_tokens\" in sampling_params:\n # support sglang-style 'max_new_tokens' param\n max_tokens = sampling_params.pop(\"max_new_tokens\")\n else:\n # Default to a calculation that considers configured lengths\n max_tokens = self.config.response_length + self.config.prompt_length - len(prompt_ids)\n\n # Clamp max_tokens to the valid range [0, max_possible_tokens]\n max_tokens = max(0, min(max_tokens, max_possible_tokens))\n\n assert max_tokens <= max_possible_tokens, (\n f\"max_tokens {max_tokens} exceeds available context space {max_possible_tokens}\"\n )\n sampling_params[\"logprobs\"] = 0 if sampling_params.pop(\"logprobs\", False) else None\n sampling_params.setdefault(\"repetition_penalty\", self.config.get(\"repetition_penalty\", 1.0))\n sampling_params = SamplingParams(max_tokens=max_tokens, **sampling_params)\n prompt_ids = _qwen2_5_vl_dedup_image_tokens(prompt_ids, self.model_config.processor)\n multi_modal_data = {}\n if image_data is not None:\n multi_modal_data[\"image\"] = image_data\n if video_data is not None:\n multi_modal_data[\"video\"] = video_data\n\n prompt = TokensPrompt(prompt_token_ids=prompt_ids, multi_modal_data=multi_modal_data)\n\n # Add lora request\n lora_request = None\n if (\n self.model_config.lora_rank > 0 or self.model_config.lora.get(\"rank\", 0) > 0\n ) and not self.model_config.lora.get(\"merge\", False):\n # Make sure we also check that the lora is already loaded in the engine\n lora_loaded = VLLM_LORA_INT_ID in await self.engine.list_loras()\n if lora_loaded:\n lora_request = LoRARequest(\n lora_name=VLLM_LORA_NAME, lora_int_id=VLLM_LORA_INT_ID, lora_path=VLLM_LORA_PATH\n )\n\n generator = self.engine.generate(\n prompt=prompt,\n sampling_params=sampling_params,\n request_id=request_id,\n lora_request=lora_request,\n priority=priority,\n )\n\n # Get final response\n final_res: Optional[RequestOutput] = None\n async for output in generator:\n final_res = output\n assert final_res is not None\n\n token_ids = final_res.outputs[0].token_ids\n log_probs = None\n if sampling_params.logprobs is not None:\n log_probs = [logprobs[token_ids[i]].logprob for i, logprobs in enumerate(final_res.outputs[0].logprobs)]\n\n routed_experts = None\n if self.config.enable_rollout_routing_replay:\n routed_experts = final_res.outputs[0].routed_experts\n\n # Determine stop reason from finish_reason\n finish_reason = final_res.outputs[0].finish_reason\n if finish_reason == \"abort\":\n stop_reason = \"aborted\"\n elif finish_reason in (\"stop\", \"length\"):\n stop_reason = \"completed\"\n else:\n stop_reason = finish_reason # for more stop reason in the future\n\n num_preempted = None\n\n if hasattr(final_res.outputs[0], \"num_preempted\"):\n num_preempted = final_res.outputs[0].num_preempted\n\n return TokenOutput(\n token_ids=token_ids,\n log_probs=log_probs,\n routed_experts=routed_experts,\n stop_reason=stop_reason,\n num_preempted=num_preempted,\n )\n\n async def wake_up(self):\n if self.node_rank != 0:\n return\n\n if self.rollout_mode == RolloutMode.HYBRID:\n # In hybrid mode, rollout is wake up in `update_weights`\n raise ValueError(f\"wake_up not support rollout_mode {self.rollout_mode}\")\n elif self.rollout_mode == RolloutMode.COLOCATED:\n # Directly call engine to wake up without sync weights.\n await self.engine.wake_up(tags=[\"kv_cache\", \"weights\"])\n await self.engine.reset_prefix_cache()\n elif self.rollout_mode == RolloutMode.STANDALONE:\n logger.info(\"skip wake_up in standalone mode\")\n\n async def sleep(self):\n if self.node_rank != 0 or not self.config.free_cache_engine:\n return\n\n if self.rollout_mode == RolloutMode.HYBRID:\n # Don't use engine.sleep(level=2) here\n await self.engine.collective_rpc(\"sleep\", kwargs={\"level\": 2})\n\n # clear encoder cache: https://github.com/vllm-project/vllm/pull/33452\n # await self.engine.reset_encoder_cache()\n elif self.rollout_mode == RolloutMode.COLOCATED:\n await self.engine.sleep(level=1)\n elif self.rollout_mode == RolloutMode.STANDALONE:\n logger.info(\"skip sleep in standalone mode\")\n\n async def start_profile(self, **kwargs):\n if (\n self.profiler_controller.check_enable()\n and self.profiler_controller.check_this_rank()\n and self.profiler_controller.is_discrete_mode()\n ):\n await self.engine.start_profile(**kwargs)\n\n async def stop_profile(self):\n if (\n self.profiler_controller.check_enable()\n and self.profiler_controller.check_this_rank()\n and self.profiler_controller.is_discrete_mode()\n ):\n await self.engine.stop_profile()\n\n async def clear_kv_cache(self):\n if self.node_rank == 0:\n await self.engine.reset_prefix_cache()\n\n async def wait_for_requests_to_drain(self):\n await self.engine.wait_for_requests_to_drain()\n\n async def abort_all_requests(self, reset_prefix_cache: bool = True) -> dict[str, Any]:\n \"\"\"Abort all ongoing generation requests.\n\n On vLLM >= 0.12.0, uses AsyncLLM.pause_generation() to abort in-flight\n requests, drain, and clear caches. The engine remains paused after this\n call — use resume_generation() to accept new requests (e.g. before\n validation).\n\n On vLLM < 0.12.0, manually aborts each request and resets prefix cache.\n\n Returns:\n dict[str, Any]: Dictionary containing:\n - aborted_count: Number of requests aborted\n - request_ids: List of aborted request IDs\n \"\"\"\n try:\n if _VLLM_VERSION >= version.parse(\"0.12.0\"):\n # Snapshot request IDs before pausing for reporting\n request_ids = list(self.engine.output_processor.request_states.keys())\n\n # pause_generation with wait_for_inflight_requests=False will:\n # 1. Set engine to paused state (blocks new generate calls)\n # 2. Abort all in-flight requests\n # 3. Wait for requests to drain\n # 4. Clear prefix and mm caches if clear_cache=True\n await self.engine.pause_generation(\n wait_for_inflight_requests=False,\n clear_cache=reset_prefix_cache,\n )\n else:\n # Take an atomic snapshot to avoid race conditions with the vLLM engine thread\n request_states_snapshot = list(self.engine.output_processor.request_states.items())\n request_ids = [req_id for req_id, _ in request_states_snapshot]\n\n if not request_ids:\n return {\"aborted_count\": 0, \"request_ids\": []}\n\n # For each request, create an abort output and put it to its queue\n # This allows the generator to receive the aborted result\n from vllm.v1.engine import FinishReason\n\n for _, req_state in request_states_snapshot:\n request_output = req_state.make_request_output(\n [], pooling_output=None, finish_reason=FinishReason.ABORT, stop_reason=None\n )\n req_state.queue.put(request_output)\n\n # Abort requests in the output processor and engine core\n self.engine.output_processor.abort_requests(request_ids)\n await self.engine.engine_core.abort_requests_async(request_ids)\n\n # Try to reset prefix cache to ensure clean state\n if reset_prefix_cache:\n await self.clear_kv_cache()\n logger.info(\"Prefix cache reset after abort\")\n\n logger.info(f\"Aborted {len(request_ids)} requests: {request_ids}\")\n return {\"aborted_count\": len(request_ids), \"request_ids\": request_ids}\n\n except Exception as e:\n logger.error(f\"Error aborting requests: {e}\")\n return {\"aborted_count\": 0, \"request_ids\": [], \"error\": str(e)}\n\n async def resume_generation(self):\n \"\"\"Resume generation after abort_all_requests (pause_generation).\n\n Only effective on vLLM >= 0.12.0 where pause_generation is used.\n No-op on older versions.\n \"\"\"\n if self.node_rank != 0:\n return\n if _VLLM_VERSION >= version.parse(\"0.12.0\"):\n await self.engine.resume_generation()\n\n async def abort_request(self, request_id: str, reset_prefix_cache: bool = True) -> dict[str, Any]:\n \"\"\"Abort a specific generation request.\n\n Args:\n request_id: The ID of the request to abort.\n\n Returns:\n dict[str, Any]: Dictionary containing abort result.\n \"\"\"\n try:\n request_states = self.engine.output_processor.request_states\n req_state = request_states.get(request_id)\n\n if req_state is None:\n return {\"aborted\": False, \"error\": f\"Request {request_id} not found\"}\n\n # Create abort output and put it to the queue\n from vllm.v1.engine import FinishReason\n\n request_output = req_state.make_request_output(\n [], pooling_output=None, finish_reason=FinishReason.ABORT, stop_reason=None\n )\n req_state.queue.put(request_output)\n\n # Abort in output processor and engine core\n self.engine.output_processor.abort_requests([request_id])\n await self.engine.engine_core.abort_requests_async([request_id])\n\n # Try to reset prefix cache to ensure clean state\n if reset_prefix_cache:\n await self.clear_kv_cache()\n logger.info(f\"Prefix cache reset after abort request {request_id}\")\n\n logger.info(f\"Aborted request: {request_id}\")\n return {\"aborted\": True, \"request_id\": request_id}\n\n except Exception as e:\n logger.error(f\"Error aborting request {request_id}: {e}\")\n return {\"aborted\": False, \"request_id\": request_id, \"error\": str(e)}\n\n\n_rollout_worker_actor_cls = ray.remote(ServerAdapter)\n\n\nclass vLLMReplica(RolloutReplica):\n def __init__(\n self,\n replica_rank: int,\n config: RolloutConfig,\n model_config: HFModelConfig,\n gpus_per_node: int = 8,\n is_reward_model: bool = False,\n ):\n super().__init__(replica_rank, config, model_config, gpus_per_node, is_reward_model)\n self.server_class = ray.remote(vLLMHttpServer)\n\n def get_ray_class_with_init_args(self) -> RayClassWithInitArgs:\n \"\"\"Get rollout worker actor class for colocated and standalone mode.\"\"\"\n worker_dict_cls = RayClassWithInitArgs(\n cls=_rollout_worker_actor_cls,\n config=self.config,\n model_config=self.model_config,\n device_mesh=None,\n )\n return worker_dict_cls\n\n async def launch_servers(self):\n \"\"\"Launch http server in each node.\"\"\"\n assert len(self.workers) == self.world_size, (\n f\"worker number {len(self.workers)} not equal to world size {self.world_size}\"\n )\n\n # NOTE: We always use MP Executor backend whether it's single-node or multi-node.\n # For multi-node without DP (e.g TP=16), need vllm>=0.11.1, https://github.com/vllm-project/vllm/pull/23691\n if self.config.data_parallel_size == 1 and self.nnodes > 1:\n assert _VLLM_VERSION >= version.parse(\"0.11.1\"), (\n \"For multi-node MP Executor, either (1) set data_parallel_size > 1 or (2) upgrade vLLM to >= 0.11.1\"\n )\n\n # get (node_id, CUDA_VISIBLE_DEVICES) of all workers\n worker_infos = await asyncio.gather(\n *[\n worker.__ray_call__.remote(\n lambda self: (\n ray.get_runtime_context().get_node_id(),\n ray.get_runtime_context().get_accelerator_ids()[get_resource_name()][0],\n )\n )\n for worker in self.workers\n ]\n )\n worker_cuda_visible_devices = [worker_info[1] for worker_info in worker_infos]\n worker_node_ids = [worker_info[0] for worker_info in worker_infos]\n\n # create server actor in each node with node affinity and cuda visible devices\n nnodes, gpus_per_replica_node = self.nnodes, self.gpus_per_replica_node\n for node_rank in range(nnodes):\n workers = self.workers[node_rank * gpus_per_replica_node : (node_rank + 1) * gpus_per_replica_node]\n node_cuda_visible_devices = \",\".join(\n worker_cuda_visible_devices[node_rank * gpus_per_replica_node : (node_rank + 1) * gpus_per_replica_node]\n )\n node_id = worker_node_ids[node_rank * gpus_per_replica_node]\n name = (\n f\"vllm_server_{self.replica_rank}_{node_rank}\"\n if not self.is_reward_model\n else f\"vllm_server_reward_{self.replica_rank}_{node_rank}\"\n )\n server = self.server_class.options(\n scheduling_strategy=ray.util.scheduling_strategies.NodeAffinitySchedulingStrategy(\n node_id=node_id,\n soft=False,\n ),\n runtime_env={\"env_vars\": {\"RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES\": \"1\"}},\n name=name,\n ).remote(\n config=self.config,\n model_config=self.model_config,\n rollout_mode=self.rollout_mode,\n workers=workers,\n replica_rank=self.replica_rank,\n node_rank=node_rank,\n gpus_per_node=gpus_per_replica_node,\n nnodes=nnodes,\n cuda_visible_devices=node_cuda_visible_devices,\n )\n self.servers.append(server)\n\n # launch http server in each node\n master_address, master_port, dp_rpc_port = await self.servers[0].get_master_address.remote()\n await asyncio.gather(\n *[\n server.launch_server.remote(\n master_address=master_address, master_port=master_port, dp_rpc_port=dp_rpc_port\n )\n for server in self.servers\n ]\n )\n\n # get http server address from first server\n server_address, server_port = await self.servers[0].get_server_address.remote()\n self._server_handle = self.servers[0]\n self._server_address = (\n f\"[{server_address}]:{server_port}\"\n if is_valid_ipv6_address(server_address)\n else f\"{server_address}:{server_port}\"\n )\n\n async def sleep(self):\n \"\"\"Sleep each rollout server.\"\"\"\n # Drain DP engines for safe sleep.\n await self.servers[0].wait_for_requests_to_drain.remote()\n await asyncio.gather(*[server.sleep.remote() for server in self.servers])\n\n async def abort_all_requests(self) -> dict[str, Any]:\n \"\"\"Abort all ongoing generation requests across all servers.\n\n Returns:\n dict[str, Any]: Combined abort results from all servers.\n \"\"\"\n results = await asyncio.gather(*[server.abort_all_requests.remote() for server in self.servers])\n\n total_aborted = sum(r.get(\"aborted_count\", 0) for r in results)\n all_request_ids = []\n for r in results:\n all_request_ids.extend(r.get(\"request_ids\", []))\n\n return {\n \"aborted_count\": total_aborted,\n \"request_ids\": all_request_ids,\n \"server_results\": results,\n }\n\n async def resume_generation(self):\n \"\"\"Resume generation on all servers after abort_all_requests.\"\"\"\n await asyncio.gather(*[server.resume_generation.remote() for server in self.servers])\n\n # TODO(petersh6): refact the checkpoint engine's update_weights and rename this method\n async def resume_all_requests(self):\n \"\"\"Resume all requests on all servers.\"\"\"\n await asyncio.gather(*[server.resume_generation.remote() for server in self.servers])\n\n async def abort_request(self, request_id: str) -> dict[str, Any]:\n \"\"\"Abort a specific request. Tries all servers since we don't know which one has it.\n\n Args:\n request_id: The ID of the request to abort.\n\n Returns:\n dict[str, Any]: Abort result.\n \"\"\"\n # TODO(petersh6): we should only abort on the server that has the request.\n results = await asyncio.gather(*[server.abort_request.remote(request_id) for server in self.servers])\n\n for r in results:\n if r.get(\"aborted\", False):\n return r\n\n return {\"aborted\": False, \"request_id\": request_id, \"error\": \"Request not found on any server\"}\n\n\ndef _qwen2_5_vl_dedup_image_tokens(prompt_ids: list[int], processor):\n \"\"\"Deduplicate consecutive image tokens in prompt_ids for Qwen2.5-VL, since vLLM will replicate the\n <|image_pad|> and <|video_pad|> token by image_data.\n\n For example,\n ```\n <|vision_start|><|image_pad|><|image_pad|>...<|image_pad|><|vision_end|>\n =>\n <|vision_start|><|image_pad|><|vision_end|>\n ```\n \"\"\"\n if processor is not None and \"Qwen2VLImageProcessor\" in processor.image_processor.__class__.__name__:\n prompt_ids = np.array(prompt_ids)\n\n # Create a mask where True indicates elements to keep\n mask = np.ones(len(prompt_ids), dtype=bool)\n\n # Find where the array equals the value\n is_value = (prompt_ids == processor.image_token_id) | (prompt_ids == processor.video_token_id)\n\n # Find consecutive duplicates by checking if previous element is also the value\n mask[1:] &= ~(is_value[1:] & is_value[:-1])\n\n return prompt_ids[mask].tolist()\n else:\n return prompt_ids\n"}168{"file_name": "verl__workers__rollout__vllm_rollout__vllm_rollout.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nThe vllm_rollout that can be applied in different backend\nWhen working with FSDP:\n- Use DTensor weight loader (recommended) or HF weight loader\n- Utilize state_dict from the FSDP to synchronize the weights among tp ranks in vLLM\nWhen working with Megatron:\n- Use Megatron weight loader\n- During training, only the current pp stage holds the parameters\n- Before inference, broadcast the parameters of the current pp rank\n to all other pp ranks (all pp ranks holds all the parameters)\n- Bind the parameters to the inference engine\n- Do inference in tp. pp is treated as additional dp\n- After inference, all the parameters that doesn't belong to this pp rank is freed.\n\"\"\"\n\nimport gc\nimport logging\nimport os\nimport time\nfrom typing import Any, Generator, Optional\n\nimport ray\nimport torch\nimport zmq\nfrom packaging import version as vs\nfrom torch.distributed.device_mesh import DeviceMesh\nfrom torch.multiprocessing.reductions import reduce_tensor\n\nfrom verl import DataProto\nfrom verl.third_party.vllm import VLLM_SLEEP_LEVEL, get_version\nfrom verl.utils.device import get_device_id, get_device_name, get_torch_device, is_support_ipc\nfrom verl.workers.config import HFModelConfig, RolloutConfig\nfrom verl.workers.rollout.base import BaseRollout\nfrom verl.workers.rollout.utils import ensure_async_iterator\nfrom verl.workers.rollout.vllm_rollout.utils import TensorMetadata, get_device_uuid\n\nlogger = logging.getLogger(__file__)\nlogger.setLevel(os.getenv(\"VERL_LOGGING_LEVEL\", \"INFO\"))\n\n\ndef _check_vllm_version_for_sleep_level():\n # https://github.com/vllm-project/vllm/issues/25171\n minver = \"0.11.0\"\n current_version = get_version(\"vllm\")\n if not current_version:\n logger.warning(\"Could not determine vLLM version, assuming an older version for sleep_level configuration.\")\n return False\n return vs.parse(current_version) >= vs.parse(minver)\n\n\nclass ServerAdapter(BaseRollout):\n \"\"\"\n vLLM server adapter used in native async mode, serve as a client to request vLLM server\n to resume/release/update weights and kv_cache.\n \"\"\"\n\n def __init__(\n self,\n config: RolloutConfig,\n model_config: HFModelConfig,\n device_mesh: DeviceMesh,\n ):\n super().__init__(config, model_config, device_mesh)\n self.server_handle: ray.actor.ActorHandle = None\n\n rank = int(os.environ[\"RANK\"])\n local_world_size = int(os.environ[\"RAY_LOCAL_WORLD_SIZE\"])\n rollout_world_size = (\n self.config.tensor_model_parallel_size\n * self.config.data_parallel_size\n * self.config.pipeline_model_parallel_size\n )\n self.replica_rank = rank // rollout_world_size\n self.rollout_rank = rank % rollout_world_size\n self.node_rank = self.rollout_rank // local_world_size\n\n if config.layered_summon or (config.expert_parallel_size > 1 and not _check_vllm_version_for_sleep_level()):\n logger.warning(\"Setting the sleep level to 1 may cause a memory overflow.\")\n self.sleep_level = 1\n else:\n self.sleep_level = VLLM_SLEEP_LEVEL\n\n self.device_uuid = get_device_uuid(get_device_id())\n self.zmq_context = zmq.Context()\n self.zmq_handle = f\"ipc:///tmp/rl-colocate-zmq-{self.device_uuid}.sock\"\n\n self.use_shm = not is_support_ipc()\n if self.use_shm:\n logger.warning(\n \"IPC is not supported on your devices. Falling back to shared memory for weight transfer, \"\n \"which may cause performance degradation. If you are using Ascend NPUs, please ensure that \"\n \"your software and CANN toolkit versions meet the requirements for IPC support. (Ascend HDK version \"\n \">= 25.3.rc1 and CANN toolkit version >= 8.3.RC1)\"\n )\n\n async def _execute_method(\n self,\n method: str,\n non_block: bool = False,\n timeout: Optional[float] = None,\n args: tuple = (),\n kwargs: Optional[dict] = None,\n ) -> Any:\n \"\"\"Execute method on inference engine via ray.\n\n Args:\n method: The method name to execute on the server.\n non_block: If True, execute the method asynchronously and return immediately.\n timeout: Timeout for the collective_rpc call.\n args: Positional arguments for the method.\n kwargs: Keyword arguments for the method.\n\n Returns:\n The result of the method execution, or None if non_block=True.\n \"\"\"\n if self.rollout_rank != 0:\n return None\n\n # Lazy init http server adapter because http server is launched after hybrid engine.\n if self.server_handle is None:\n self.server_handle = ray.get_actor(f\"vllm_server_{self.replica_rank}_{self.node_rank}\")\n\n future = self.server_handle.collective_rpc.remote(method, timeout=timeout, args=args, kwargs=kwargs)\n return future if non_block else await future\n\n async def resume(self, tags: list[str]):\n \"\"\"Resume rollout weights or kv cache in GPU memory.\n\n Args:\n tags: weights or kv_cache.\n \"\"\"\n if self.config.free_cache_engine:\n await self._execute_method(\"wake_up\", kwargs={\"tags\": tags})\n\n async def release(self):\n \"\"\"Release weights and kv cache in GPU memory.\"\"\"\n if self.config.free_cache_engine:\n await self._execute_method(\"sleep\", kwargs={\"level\": self.sleep_level})\n\n @torch.no_grad()\n async def update_weights(self, weights: Generator[tuple[str, torch.Tensor], None, None], **kwargs):\n \"\"\"Update model weights via CUDA IPC (fallback to shared memory if IPC not supported) to inference workers.\"\"\"\n start_time = time.time()\n\n future = await self._execute_method(\n \"update_weights_from_ipc\",\n non_block=True,\n kwargs={**kwargs, \"use_shm\": self.use_shm},\n )\n\n # build communication buffer\n bucket_size_mb = self.config.checkpoint_engine.update_weights_bucket_megabytes\n bucket_size = int(bucket_size_mb) << 20\n s = self.zmq_context.socket(zmq.REQ)\n s.bind(self.zmq_handle)\n\n buffer, shm = None, None\n if not self.use_shm:\n buffer = torch.empty(bucket_size, dtype=torch.uint8, device=f\"{get_device_name()}:0\")\n handle = reduce_tensor(buffer)\n s.send_pyobj(handle)\n else:\n import uuid\n from multiprocessing import shared_memory\n\n # Create unique name for shared memory\n shm_name = f\"verl_weights_{uuid.uuid4().hex}\"\n shm = shared_memory.SharedMemory(name=shm_name, create=True, size=bucket_size)\n buffer = torch.frombuffer(shm.buf, dtype=torch.uint8)\n\n comm_metadata = {\"name\": shm_name, \"size\": bucket_size}\n s.send_pyobj(comm_metadata)\n\n s.recv()\n\n # send bucket weights\n offset = 0\n bucket_meta: dict[str, TensorMetadata] = {}\n # dtype = PrecisionType.to_dtype(self.config.dtype)\n async for name, weight in ensure_async_iterator(weights):\n # model parameters are in fp32 full precision\n # (vermouth1992) we should not force cast weight here because some parameters\n # (such as moe gate) have to keep fp32 precision. If a weight is bf16 in the rollout side,\n # the rollout should automatically cast on demand. However, this would incur a higher weight\n # transfer volume.\n # weight = weight.to(dtype, non_blocking=True)\n\n # fill the tensor bucket\n if offset + weight.nbytes > bucket_size:\n get_torch_device().synchronize()\n s.send_pyobj({\"bucket_meta\": bucket_meta, \"is_last\": False})\n s.recv()\n bucket_meta = {}\n offset = 0\n\n # TODO: slice embedding layer weight into chunks\n assert offset + weight.nbytes <= bucket_size, (\n f\"Weight {name}({weight.shape}, {weight.dtype}) is too large to fit in the bucket.\"\n f\"Please increase rollout.update_weights_bucket_megabytes({bucket_size_mb} MB).\"\n )\n bucket_meta[name] = {\n \"name\": name,\n \"shape\": weight.shape,\n \"dtype\": weight.dtype,\n \"offset\": offset,\n }\n buffer[offset : offset + weight.nbytes].copy_(weight.view(-1).view(torch.uint8), non_blocking=True)\n offset += weight.nbytes\n\n # send the last bucket\n get_torch_device().synchronize()\n s.send_pyobj({\"bucket_meta\": bucket_meta, \"is_last\": True})\n s.recv()\n\n # clean up\n s.close()\n del buffer\n if shm is not None:\n shm.close()\n shm.unlink()\n del shm\n gc.collect()\n get_torch_device().ipc_collect()\n get_torch_device().empty_cache()\n if future is not None:\n await future\n\n # reset prefix cache after updating weights\n if self.rollout_rank == 0:\n await self.server_handle.clear_kv_cache.remote()\n\n if self.replica_rank == 0 and self.rollout_rank == 0:\n logger.info(f\"update_weights done, time cost: {time.time() - start_time:.2f}s\")\n\n def generate_sequences(self, prompts: DataProto) -> DataProto:\n \"\"\"Batch generate sequences in sync mode.\n\n Note: ServerAdapter uses async server mode and does not support synchronous\n generation. Since SPMD mode was retired (PR #4411), the generation workflow\n should use the async server interface instead.\n\n Raises:\n NotImplementedError: Always raised as sync generation is not supported.\n \"\"\"\n raise NotImplementedError(\n \"ServerAdapter does not support synchronous generate_sequences(). \"\n \"The vLLM SPMD mode was retired in PR #4411. For batch generation, \"\n \"please use the async server interface via vLLMReplica and AsyncLLMServerManager, \"\n \"or use HFRollout for synchronous generation. \"\n \"See https://github.com/volcengine/verl/issues/4682 for more details.\"\n )\n"}169{"file_name": "verl__workers__sharding_manager__base.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nSharding manager to implement HybridEngine\n\"\"\"\n\nfrom verl import DataProto\n\n\nclass BaseShardingManager:\n def __init__(self):\n self.timing = {}\n\n def __enter__(self):\n pass\n\n def __exit__(self, exc_type, exc_value, traceback):\n pass\n\n def preprocess_data(self, data: DataProto) -> DataProto:\n return data\n\n def postprocess_data(self, data: DataProto) -> DataProto:\n return data\n"}170{"file_name": "verl__workers__sharding_manager__fsdp_ulysses.py", "text": "# Copyright 2024 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"\nContains a resharding manager that binds weights from FSDP zero3 to XPerfGPT\n\"\"\"\n\nfrom torch.distributed.device_mesh import DeviceMesh\n\nfrom verl import DataProto\nfrom verl.protocol import all_gather_data_proto\nfrom verl.utils.ulysses import get_ulysses_sequence_parallel_group, set_ulysses_sequence_parallel_group\n\nfrom .base import BaseShardingManager\n\n\nclass FSDPUlyssesShardingManager(BaseShardingManager):\n \"\"\"\n Sharding manager to support data resharding when using FSDP + Ulysses\n \"\"\"\n\n def __init__(self, device_mesh: DeviceMesh):\n super().__init__()\n self.device_mesh = device_mesh\n self.seed_offset = 12345\n\n def __enter__(self):\n if self.device_mesh is not None:\n # We have a global SP group\n # so we have to change to use model-specific sp group\n self.prev_sp_group = get_ulysses_sequence_parallel_group()\n set_ulysses_sequence_parallel_group(self.device_mesh[\"sp\"].get_group())\n # TODO: check how to set seed for each model\n\n def __exit__(self, exc_type, exc_value, traceback):\n # restore random states\n if self.device_mesh is not None:\n # revert to previous sp group\n set_ulysses_sequence_parallel_group(self.prev_sp_group)\n # TODO: check how to set seed for each model\n\n def preprocess_data(self, data: DataProto) -> DataProto:\n \"\"\"\n AllGather data from sp region\n This is because the data is first sharded along the FSDP dimension as we utilize the DP_COMPUTE\n In Ulysses, we need to make sure the same data is used across a SP group\n \"\"\"\n if self.device_mesh is not None:\n group = self.device_mesh[\"sp\"].get_group()\n\n all_gather_data_proto(data=data, process_group=group)\n return data\n\n def postprocess_data(self, data: DataProto) -> DataProto:\n \"\"\"\n Split the data to follow FSDP partition\n \"\"\"\n if self.device_mesh is not None:\n sp_size = self.device_mesh[\"sp\"].size()\n sp_rank = self.device_mesh[\"sp\"].get_local_rank()\n data = data.chunk(chunks=sp_size)[sp_rank]\n return data\n"}171{"file_name": "verl__workers__utils__losses.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\n\nimport torch\nimport torch.nn.functional as F\nfrom tensordict import TensorDict\n\nfrom verl.trainer.ppo.core_algos import agg_loss, compute_value_loss, get_policy_loss_fn, kl_penalty\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.dataset.dataset_utils import DatasetPadMode\nfrom verl.utils.metric import AggregationType, Metric\nfrom verl.utils.torch_functional import masked_mean, masked_sum\nfrom verl.workers.config import ActorConfig, CriticConfig\nfrom verl.workers.utils.padding import no_padding_2_padding\n\n\ndef sft_loss(config: ActorConfig, model_output, data: TensorDict, dp_group=None):\n pad_mode = tu.get_non_tensor_data(data=data, key=\"pad_mode\", default=DatasetPadMode.NO_PADDING)\n dp_size = data[\"dp_size\"]\n batch_num_tokens = data[\"batch_num_tokens\"]\n\n log_prob = model_output[\"log_probs\"]\n\n if pad_mode == DatasetPadMode.NO_PADDING:\n # log_prob and loss mask are nested tensors of shape [bsz, j1]\n # for each sample, loss mask shape is [1, prompt_length + response_length]\n loss_mask = data[\"loss_mask\"]\n\n log_prob_flatten = log_prob.values()\n loss_mask_flatten = loss_mask.values()\n\n # left-shift the loss mask by one token to align with log_prob\n loss_mask_flatten = torch.roll(loss_mask_flatten, shifts=-1, dims=0)\n\n # NOTE: loss is averaged over all tokens in the batch across all data parallel groups,\n # For FSDP backend, the loss is directly used for backward; while for Megatron backend,\n # the loss should be scaled by `num_microbatches` for pp schedule.\n loss = -masked_sum(log_prob_flatten, loss_mask_flatten) / batch_num_tokens * dp_size\n else:\n response_mask = data[\"response_mask\"].to(bool)\n loss = -masked_sum(log_prob, response_mask) / batch_num_tokens * dp_size\n\n return loss, {}\n\n\ndef _slice_response_from_unpad_output(tensor: torch.Tensor, data: TensorDict) -> torch.Tensor:\n \"\"\"Slice response from unpad model output.\n\n Args:\n tensor: model output tensor of shape [bsz, 1]\n data: TensorDict with \"prompt_ids\", \"response_ids\", \"attention_mask\"\n\n Returns:\n tensor: sliced response tensor of shape [bsz, max_response_len]\n \"\"\"\n values = tensor.values() if tensor.is_nested else tensor\n prompt_ids = data[\"prompts\"]\n response_ids = data[\"responses\"]\n attention_mask = data[\"attention_mask\"]\n\n if prompt_ids.is_nested:\n prompt_lens = prompt_ids.offsets().diff()\n response_lens = response_ids.offsets().diff()\n max_response_len = response_ids.offsets().max().item()\n else:\n assert not attention_mask.is_nested\n prompt_lens = attention_mask[:, : prompt_ids.shape[1]].sum(dim=1)\n response_lens = attention_mask[:, prompt_ids.shape[1] :].sum(dim=1)\n max_response_len = response_ids.shape[1]\n\n sequence_lens = prompt_lens + response_lens\n sequence_offsets = sequence_lens.cumsum(dim=0)\n assert sequence_offsets[-1].item() == values.shape[0]\n\n response_list = []\n for resp_len, seq_offset in zip(response_lens, sequence_offsets, strict=True):\n pad_size = max_response_len - resp_len\n # left-shift model output by one token for log_probs/values\n response_list.append(F.pad(values[seq_offset - resp_len - 1 : seq_offset - 1], (0, pad_size)))\n\n output = torch.stack(response_list, dim=0)\n return output\n\n\ndef ppo_loss(config: ActorConfig, model_output, data: TensorDict, dp_group=None):\n \"\"\"Computes ppo loss from model output (log_prob, entropy, values, etc. ) and old_log_probs from data.\"\"\"\n log_prob = no_padding_2_padding(model_output[\"log_probs\"], data)\n entropy = model_output.get(\"entropy\", None)\n if entropy is not None:\n entropy = no_padding_2_padding(entropy, data)\n\n # global batch info for loss aggregation\n config.global_batch_info[\"dp_size\"] = data[\"dp_size\"]\n config.global_batch_info[\"batch_num_tokens\"] = data[\"batch_num_tokens\"]\n config.global_batch_info[\"global_batch_size\"] = data[\"global_batch_size\"]\n config.global_batch_info[\"loss_scale_factor\"] = config.loss_scale_factor\n\n # assumes that if any of the global batch info is set, the policy_loss_fn will\n # normalize using dp_size/global_bsz/global_token; in this case, metric aggregation should be SUM\n # to reflect the mean loss over the global batch\n if (\n data[\"dp_size\"] > 1\n or data[\"batch_num_tokens\"] is not None\n or data[\"global_batch_size\"] is not None\n or config.loss_scale_factor is not None\n ):\n metric_aggregation = AggregationType.SUM\n else:\n metric_aggregation = AggregationType.MEAN\n\n metrics = {}\n\n response_mask = data[\"response_mask\"].to(bool)\n # compute policy loss\n old_log_prob = data[\"old_log_probs\"]\n advantages = data[\"advantages\"]\n rollout_is_weights = data.get(\"rollout_is_weights\", None)\n\n loss_agg_mode = config.loss_agg_mode\n\n loss_mode = config.policy_loss.get(\"loss_mode\", \"vanilla\")\n\n policy_loss_fn = get_policy_loss_fn(loss_mode)\n pg_loss, pg_metrics = policy_loss_fn(\n old_log_prob=old_log_prob,\n log_prob=log_prob,\n advantages=advantages,\n response_mask=response_mask,\n loss_agg_mode=loss_agg_mode,\n config=config,\n rollout_is_weights=rollout_is_weights,\n )\n\n # AggregationType.MEAN for pg metrics: assumes policy_loss_fn normalizes by local_bsz/local_tokens\n # Ex: in compute_policy_loss_vanilla, pg_metrics are pg_clipfrac, ppo_kl, pg_clipfrac_lower\n pg_metrics = Metric.from_dict(pg_metrics, aggregation=AggregationType.MEAN)\n\n metrics.update(pg_metrics)\n metrics[\"actor/pg_loss\"] = Metric(value=pg_loss, aggregation=metric_aggregation)\n policy_loss = pg_loss\n\n # add entropy loss\n if entropy is not None:\n entropy_loss = agg_loss(\n loss_mat=entropy, loss_mask=response_mask, loss_agg_mode=loss_agg_mode, **config.global_batch_info\n )\n entropy_coeff = config.entropy_coeff\n policy_loss -= entropy_coeff * entropy_loss\n metrics[\"actor/entropy_loss\"] = Metric(value=entropy_loss, aggregation=metric_aggregation)\n\n # add kl loss\n if config.use_kl_loss:\n ref_log_prob = data[\"ref_log_prob\"]\n # compute kl loss\n kld = kl_penalty(logprob=log_prob, ref_logprob=ref_log_prob, kl_penalty=config.kl_loss_type)\n kl_loss = agg_loss(\n loss_mat=kld, loss_mask=response_mask, loss_agg_mode=config.loss_agg_mode, **config.global_batch_info\n )\n\n policy_loss += kl_loss * config.kl_loss_coef\n metrics[\"kl_loss\"] = Metric(value=kl_loss, aggregation=metric_aggregation)\n metrics[\"kl_coef\"] = config.kl_loss_coef\n\n return policy_loss, metrics\n\n\ndef value_loss(config: CriticConfig, model_output, data: TensorDict, dp_group=None):\n \"\"\"value loss\n\n Args:\n config: CriticConfig\n model_output: model output from the model\n data: the input to the model\n dp_group: data paralle group\n\n Returns:\n value loss\n \"\"\"\n vpreds = _slice_response_from_unpad_output(model_output[\"values\"], data) # (bsz, response_length)\n\n values = data[\"values\"]\n returns = data[\"returns\"]\n response_mask = data[\"response_mask\"].to(bool)\n\n vf_loss, vf_clipfrac = compute_value_loss(\n vpreds=vpreds,\n values=values,\n returns=returns,\n response_mask=response_mask,\n cliprange_value=config.cliprange_value,\n loss_agg_mode=config.loss_agg_mode,\n )\n\n metrics = {}\n\n metrics.update(\n {\n \"critic/vf_loss\": vf_loss.detach().item(),\n \"critic/vf_clipfrac\": vf_clipfrac.detach().item(),\n \"critic/vpred_mean\": masked_mean(vpreds, response_mask).detach().item(),\n }\n )\n\n return vf_loss, metrics\n"}172{"file_name": "verl__workers__utils__padding.py", "text": "# Copyright 2025 Bytedance Ltd. and/or its affiliates\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n# http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\nimport torch\nimport torch.nn.functional as F\nfrom tensordict import TensorDict\n\nfrom verl.utils import tensordict_utils as tu\nfrom verl.utils.attention_utils import index_first_axis, unpad_input\n\n\ndef left_right_2_no_padding(data: TensorDict) -> TensorDict:\n \"\"\"\n Convert TensorDict from left-right padding to no-padding format.\n\n Args:\n data: TensorDict with \"input_ids\", \"attention_mask\", \"response_mask\", \"position_ids\"\n\n Returns:\n data: TensorDict with\n - Tensor includes NestedTensors like \"input_ids\", \"loss_mask\", \"position_ids\"\n - NonTensorData includes \"max_seq_len\", \"max_response_len\", \"indices\"\n\n Note:\n 1. the return input_ids/position_ids/loss_mask are nested tensor.\n 2. we will remove \"attention_mask\", \"response\" in the return data, but \"response_mask\" is kept.\n \"\"\"\n assert \"input_ids\" in data, \"input_ids is required in left-right padding data\"\n assert \"attention_mask\" in data, \"attention_mask is required in left-right padding data\"\n assert \"response_mask\" in data, \"response_mask is required in left-right padding data\"\n assert \"position_ids\" in data, \"position_ids is required in left-right padding data\"\n\n input_ids = data.pop(\"input_ids\")\n attention_mask = data[\"attention_mask\"]\n response_mask = data[\"response_mask\"]\n position_ids = data[\"position_ids\"] # (bs, seq_len) or # (bs, 4, seq_len)\n\n max_seq_len, max_response_len = input_ids.shape[1], response_mask.shape[1]\n tu.assign_non_tensor_data(data, \"max_seq_len\", max_seq_len)\n tu.assign_non_tensor_data(data, \"max_response_len\", max_response_len)\n\n input_ids_rmpad, indices, cu_seqlens, *_ = unpad_input(input_ids.unsqueeze(-1), attention_mask)\n tu.assign_non_tensor_data(data, \"indices\", indices)\n\n input_ids_nested = torch.nested.nested_tensor_from_jagged(input_ids_rmpad.squeeze(-1), offsets=cu_seqlens)\n\n position_ids_list = []\n for i in range(attention_mask.shape[0]):\n curr_mask = attention_mask[i].bool()\n curr_pos_ids = position_ids[i]\n if curr_pos_ids.dim() == 1: # (seq_len,)\n valid_ids = curr_pos_ids[curr_mask]\n else: # (4, seq_len)\n valid_ids = curr_pos_ids[:, curr_mask]\n position_ids_list.append(valid_ids)\n position_ids_nested = torch.nested.as_nested_tensor(position_ids_list, layout=torch.jagged)\n\n data[\"input_ids\"] = input_ids_nested\n data[\"position_ids\"] = position_ids_nested\n data[\"loss_mask\"] = data[\"response_mask\"]\n\n routed_experts = data.get(\"routed_experts\", None)\n if routed_experts is not None and not routed_experts.is_nested:\n if routed_experts.max() <= 255:\n routed_experts = routed_experts.to(torch.uint8)\n routed_experts_rmpad = index_first_axis(routed_experts.unsqueeze(-1).flatten(0, 1), indices)\n routed_experts_nested = torch.nested.nested_tensor_from_jagged(\n routed_experts_rmpad.squeeze(-1), offsets=cu_seqlens\n )\n data[\"routed_experts\"] = routed_experts_nested\n\n return data\n\n\ndef no_padding_2_padding(tensor: torch.Tensor, data: TensorDict) -> torch.Tensor:\n \"\"\"Slice response from unpad model output.\n\n Args:\n tensor: a nested tensor or a 1D tensor in shape (total_nnz,),\n total_nnz is the total number of tokens across all sequences in the batch\n data: TensorDict with \"prompts\", \"responses\", \"attention_mask\"\n\n Returns:\n tensor: sliced response tensor of shape [bsz, max_response_len]\n \"\"\"\n values = tensor.values() if tensor.is_nested else tensor\n prompt_ids = data[\"prompts\"]\n response_ids = data[\"responses\"]\n attention_mask = data[\"attention_mask\"]\n\n max_response_len = tu.get_non_tensor_data(data=data, key=\"max_response_len\", default=-1)\n\n if prompt_ids.is_nested:\n prompt_lens = prompt_ids.offsets().diff()\n response_lens = response_ids.offsets().diff()\n if max_response_len < 0:\n max_response_len = response_ids.offsets().diff().max().item()\n else:\n assert not attention_mask.is_nested\n prompt_lens = attention_mask[:, : prompt_ids.shape[1]].sum(dim=1)\n response_lens = attention_mask[:, prompt_ids.shape[1] :].sum(dim=1)\n max_response_len = response_ids.shape[1]\n\n sequence_lens = prompt_lens + response_lens\n sequence_offsets = sequence_lens.cumsum(dim=0)\n assert sequence_offsets[-1].item() == values.shape[0]\n\n response_list = []\n for resp_len, seq_offset in zip(response_lens, sequence_offsets, strict=True):\n pad_size = max_response_len - resp_len\n # left-shift model output by one token for log_probs/values\n response_list.append(F.pad(values[seq_offset - resp_len - 1 : seq_offset - 1], (0, pad_size)))\n\n output = torch.stack(response_list, dim=0)\n return output\n"}173 