diff --git a/vllm/compilation/cuda_graph.py b/vllm/compilation/cuda_graph.py --- a/vllm/compilation/cuda_graph.py +++ b/vllm/compilation/cuda_graph.py @@ import dataclasses +import os import weakref @@ -from vllm.compilation.monitor import validate_cudagraph_capturing_enabled +from vllm.compilation.monitor import ( + set_cudagraph_capturing_enabled, + validate_cudagraph_capturing_enabled, +) @@ class CUDAGraphEntry: @@ cudagraph: torch.cuda.CUDAGraph | None = None output: Any | None = None + replay_count: int = 0 @@ if cudagraph_options is None: cudagraph_options = CUDAGraphOptions() + if os.environ.get("VLLM_XPU_CUDAGRAPH_STRONG_OUTPUT", "0") == "1": + cudagraph_options = dataclasses.replace( + cudagraph_options, weak_ref_output=False + ) self.cudagraph_options = cudagraph_options @@ def clear_graphs(self) -> None: self.concrete_cudagraph_entries.clear() + + @staticmethod + def _recapture_after_n_replays() -> int: + raw = os.environ.get("VLLM_XPU_CUDAGRAPH_RECAPTURE_AFTER_N_REPLAYS", "0") + try: + return max(0, int(raw)) + except ValueError: + return 0 @@ # Sync offloader before replay - ensures any external dependencies # from pre-capture prefetches are satisfied. get_offloader().sync_prev_onload() + recapture_after = self._recapture_after_n_replays() + if recapture_after and entry.replay_count >= recapture_after: + entry.cudagraph = None + entry.output = None + entry.input_addresses = None + entry.replay_count = 0 + if ( + os.environ.get( + "VLLM_XPU_CUDAGRAPH_ALLOW_RUNTIME_RECAPTURE", "0" + ) + == "1" + ): + set_cudagraph_capturing_enabled(True) + try: + return self.__call__(*args, **kwargs) + finally: + set_cudagraph_capturing_enabled(False) + return self.__call__(*args, **kwargs) + sync_replay = os.environ.get("VLLM_XPU_SYNC_CUDAGRAPH_REPLAY", "0") == "1" + if sync_replay and hasattr(torch, "xpu"): + torch.xpu.synchronize() entry.cudagraph.replay() + entry.replay_count += 1 + if sync_replay and hasattr(torch, "xpu"): + torch.xpu.synchronize() return entry.output diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ if self.vllm_config.parallel_config.data_parallel_size > 1: @@ # num_tokens_across_dp will no-longer be valid assert batch_descriptor.num_tokens == num_tokens_padded + + if ( + os.environ.get("VLLM_XPU_DISABLE_DECODE_CUDAGRAPH_REPLAY", "0") == "1" + and uniform_decode + and cudagraph_mode == CUDAGraphMode.PIECEWISE + ): + # Diagnostic guard for XPU graph replay corruption: keep the same + # compiled runnable and padding decision, but bypass decode replay. + cudagraph_mode = CUDAGraphMode.NONE cudagraph_stats = None if self.vllm_config.observability_config.cudagraph_metrics: @@ if self.execute_model_state is not None: raise RuntimeError( "State error: sample_tokens() must be called " "after execute_model() returns None." ) + if os.environ.get("VLLM_XPU_CUDAGRAPH_MARK_STEP_BEGIN", "0") == "1": + torch.compiler.cudagraph_mark_step_begin() if self.routed_experts_initialized: capturer = RoutedExpertsCapturer.get_instance() diff --git a/scripts/probe-fixed-chatml-completion-repeat.py b/scripts/probe-fixed-chatml-completion-repeat.py --- a/scripts/probe-fixed-chatml-completion-repeat.py +++ b/scripts/probe-fixed-chatml-completion-repeat.py @@ import json import re import time +import urllib.error import urllib.request @@ parser.add_argument("--request-delay-s", type=float, default=0.0) parser.add_argument("--stop-on-mismatch", action="store_true") + parser.add_argument( + "--logprobs", + type=int, + default=0, + help="Request completion logprobs/top_logprobs for mismatch diagnosis.", + ) parser.add_argument("--output-json", type=Path, required=True) @@ } + if args.logprobs > 0: + payload["logprobs"] = args.logprobs started = time.perf_counter() - data = post_json(f"{base_url}/v1/completions", payload, args.timeout) + try: + data = post_json(f"{base_url}/v1/completions", payload, args.timeout) + except urllib.error.HTTPError as exc: + body = exc.read().decode("utf-8", errors="replace") + row = { + "index": index, + "text": "", + "normalized": "", + "token_ids": [], + "finish_reason": None, + "elapsed_s": time.perf_counter() - started, + "usage": None, + "http_error": { + "code": exc.code, + "reason": exc.reason, + "body": body, + }, + "pass": False, + } + rows.append(row) + mismatches.append(row) + if args.stop_on_mismatch: + break + if args.request_delay_s > 0: + time.sleep(args.request_delay_s) + continue elapsed = time.perf_counter() - started @@ "request_delay_s": args.request_delay_s, "stop_on_mismatch": args.stop_on_mismatch, + "logprobs": args.logprobs, "expected_normalized": case.get("expected_normalized"),