Skip to content
Merged
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
6 changes: 5 additions & 1 deletion autotest/config-npu.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -217,7 +217,8 @@ case:
accelerator: NPU
output_path: /mnt/hwfile/llmrazor/qa-llm-cic/qa-llm-cicd/test_output
resource:
image: ccr-hw/910c_cann910_pta291_rl_s1_vllm:v0.1.0-20260826T114552Z
image: ccr-hw/910c_cann910_pta291_rl_s1_vllm:v0.5.0-20260917T114217Z
pip_package: unset ASCEND_WORK_PATH;pip install -e . --no-deps; pip install more-itertools triton
envs:
- MODEL_PATH=/mnt/hwfile/llmrazor/qa-llm-cicd/qa_test_models/Qwen3-30B-A3B
- DATA_PATH=/mnt/hwfile/llmrazor/qa-llm-cicd/xtuner_resource/datasets/gsm8k/train-mini.jsonl
Expand All @@ -230,6 +231,9 @@ case:
- XTUNER_ACTIVATION_OFFLOAD=1
- VLLM_VERSION=0.11.0
- VLLM_USE_V1=1
- TORCH_NPU_DEVICE_CAPABILITY=9.0
- PIP_INDEX_URL=http://pkg.pjlab.org.cn/repository/pypi-tsinghua/simple
- PIP_TRUSTED_HOST=pkg.pjlab.org.cn
assert_info:
base_metric: npu-qwen3-rl-vllm/tracker.jsonl
check_metrics:
Expand Down
59 changes: 52 additions & 7 deletions xtuner/v1/rl/rollout/vllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,9 +28,8 @@ def stateless_init_process_group(master_address, master_port, rank, world_size,
"""VLLM provides `StatelessProcessGroup` to create a process group without
considering the global process group in torch.distributed.

It is recommended to create `StatelessProcessGroup`, and then initialize
the data-plane communication (NCCL) between external (train processes)
and vLLM workers.
It is recommended to create `StatelessProcessGroup`, and then initialize the data-plane communication (NCCL)
between external (train processes) and vLLM workers.
"""
from vllm.distributed.utils import StatelessProcessGroup

Expand Down Expand Up @@ -210,6 +209,55 @@ def __init__(
)
self.tp_size = self.config.tensor_parallel_size // self.dp_size

def _get_request_payload(self, rollout_state: RolloutState) -> dict:
"""Build the chat-completions request body for one generation
request."""
sample_params = rollout_state.sample_params
payload: dict[str, Any] = {
"model": self.config.model_path,
"messages": rollout_state.message,
"stream": sample_params.stream,
}

if rollout_state.tools is not None:
payload["tools"] = rollout_state.tools
if rollout_state.tool_choice is not None:
payload["tool_choice"] = rollout_state.tool_choice

# vLLM-Ascend (USE_TOKEN_IN=1) accepts explicit input_ids; partial-rollout
# replay passes the request prefix via rollout_state.tokens.
if rollout_state.tokens is not None:
payload["input_ids"] = rollout_state.tokens
elif "train_prompt_ids" in rollout_state.extra_fields:
payload["input_ids"] = rollout_state.extra_fields["train_prompt_ids"]

if "image_data" in rollout_state.extra_fields:
image_data = rollout_state.extra_fields["image_data"]
assert isinstance(payload["messages"], list), "image_data requires messages to be a list"
image_index = 0
for message in payload["messages"]:
if not isinstance(message, dict) or message.get("role") != "user":
continue
new_content = []
for content_part in message.get("content", []):
if not isinstance(content_part, dict):
new_content.append(content_part)
continue
if content_part.get("type") == "image_url":
content_part["image_url"]["url"] = f"file://{image_data[image_index]}"
content_part["image_url"].pop("image_wh", None)
image_index += 1
new_content.append(content_part)
message["content"] = new_content
assert image_index == len(image_data), f"Expected {len(image_data)} images, but processed {image_index}."

vllm_sample_params = self._transform_sample_params(sample_params.model_dump())
vllm_sample_params["return_routed_experts"] = (
self.enable_return_routed_experts and sample_params.return_routed_experts
)
payload.update(vllm_sample_params)
return payload

async def _create_request(
self,
url: str,
Expand Down Expand Up @@ -285,9 +333,6 @@ def _transform_sample_params(self, sample_params: Dict, extra_params: Dict = {})
def get_logprobs(self, input_ids, sampling_params):
pass

def generate(self, input_ids, sampling_params):
pass

def sleep(self, level=1):
url = f"{self.server_url}/{self.endpoints['sleep']}"
headers = {"Content-Type": "application/json"}
Expand Down Expand Up @@ -368,7 +413,7 @@ def _transform_rollout_config_to_server_configs(self) -> Namespace:
"cudagraph_capture_sizes": [16, 12, 8, 4, 2, 1],
"cudagraph_mode": "FULL_DECODE_ONLY",
}
args["additional_config"] = {"enable_cpu_binding": True}
args["additional_config"] = {"enable_cpu_binding": True, "weight_nz_mode": 0}
args["limit_mm_per_prompt"] = {"image": 10, "video": 0}
args["enable_log_requests"] = False
args["uvicorn_log_level"] = "error"
Expand Down
Loading