Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions verl/trainer/constants_ppo.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
import json
import os

import torch
from ray._private.runtime_env.constants import RAY_JOB_CONFIG_JSON_ENV_VAR

from verl.utils.device import get_device_capability
Expand All @@ -27,6 +28,14 @@
if (_major or 0) >= 10 and os.environ.get("TLLM_DISABLE_NVLS_MNNVL", "0") == "1":
_gb200_nccl_env = {"NCCL_NVLS_ENABLE": "0", "NCCL_MNNVL_ENABLE": "0"}

# On ROCm, Ray 2.x force-clears accelerator visibility for num_gpus=0 actors
# (e.g. the SGLang server actor), leaving them unable to see any GPU. Disable
# that override so the actor keeps its HIP visibility. Scoped to ROCm to avoid
# changing Ray's default behavior on other platforms.
_rocm_ray_env = {}
if torch.version.hip is not None:
_rocm_ray_env = {"RAY_ACCEL_ENV_VAR_OVERRIDE_ON_ZERO": "0"}

PPO_RAY_RUNTIME_ENV = {
"env_vars": {
"TOKENIZERS_PARALLELISM": "true",
Expand All @@ -43,6 +52,7 @@
"HCCL_NPU_SOCKET_PORT_RANGE": "auto",
"HSA_NO_SCRATCH_RECLAIM": "1",
**_gb200_nccl_env,
**_rocm_ray_env,
},
}

Expand Down
16 changes: 12 additions & 4 deletions verl/workers/rollout/sglang_rollout/async_sglang_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@
)
from sglang.srt.managers.tokenizer_manager import ServerStatus

from verl.plugin.platform import get_platform
from verl.utils.config import omega_conf_to_dataclass
from verl.utils.device import get_visible_devices_keyword
from verl.utils.net_utils import get_free_port, is_valid_ipv6_address
Expand Down Expand Up @@ -256,9 +257,11 @@ async def launch_server(self, master_address: str = None, master_port: int = Non
attention_backend = engine_kwargs.pop("attention_backend", None)
mm_attention_backend = engine_kwargs.pop("mm_attention_backend", None)
if attention_backend is None:
# FA3 CUDA-graph capture is broken on sglang>=0.5.12 (#22800);
# default to flashinfer (users can opt into fa4 via engine_kwargs).
if version.parse(sglang.__version__) >= version.parse("0.5.12"):
if torch.version.hip is not None:
attention_backend = "aiter"
elif version.parse(sglang.__version__) >= version.parse("0.5.12"):
# FA3 CUDA-graph capture is broken on sglang>=0.5.12 (#22800);
# default to flashinfer (users can opt into fa4 via engine_kwargs).
attention_backend = "flashinfer"
else:
attention_backend = "fa3"
Expand Down Expand Up @@ -791,7 +794,12 @@ async def launch_servers(self):
node_id=node_id,
soft=False,
),
runtime_env={"env_vars": {f"RAY_EXPERIMENTAL_NOSET_{visible_devices_keyword}": "1"}},
runtime_env={
"env_vars": {
**{var: "1" for var in get_platform().ray_noset_envvars()},
**get_platform().rollout_env_vars(),
}
},
name=name,
max_concurrency=self.max_concurrency,
).remote(
Expand Down
Loading