diff --git a/autotest/config-npu.yaml b/autotest/config-npu.yaml index d49dd843a6..2432c65bab 100644 --- a/autotest/config-npu.yaml +++ b/autotest/config-npu.yaml @@ -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 @@ -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: diff --git a/xtuner/v1/rl/rollout/vllm.py b/xtuner/v1/rl/rollout/vllm.py index 1c4ae7fd64..afa9dadd46 100644 --- a/xtuner/v1/rl/rollout/vllm.py +++ b/xtuner/v1/rl/rollout/vllm.py @@ -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 @@ -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, @@ -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"} @@ -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"