From 10cf1e482c4d4b0304b099f8a219ddfcecc4449b Mon Sep 17 00:00:00 2001 From: yewentao256 Date: Mon, 22 Dec 2025 21:12:46 +0000 Subject: [PATCH] reset --- vllm/compilation/cuda_graph.py | 86 +++++----------------------------- 1 file changed, 13 insertions(+), 73 deletions(-) diff --git a/vllm/compilation/cuda_graph.py b/vllm/compilation/cuda_graph.py index d82d8f2820c..4c1cfbfa848 100644 --- a/vllm/compilation/cuda_graph.py +++ b/vllm/compilation/cuda_graph.py @@ -126,11 +126,6 @@ class CUDAGraphEntry: # during capture, and check if they are the same during replay input_addresses: list[int] | None = None - # CUDA graphs will replay using the capture-time input buffers; - # keeping references here ensures they stay alive and enables - # optionally copying runtime inputs into these buffers before replay. - input_tensors: list[torch.Tensor] | None = None - @dataclasses.dataclass class CUDAGraphOptions: @@ -207,43 +202,6 @@ class CUDAGraphWrapper: # in case we need to access the original runnable. return self.runnable - @staticmethod - def _tensor_inputs(args) -> list[torch.Tensor]: - # preserve the same ordering as the original address debug check. - return [x for x in args if isinstance(x, torch.Tensor)] - - @staticmethod - def _copy_runtime_inputs_into_captured_inputs( - entry: CUDAGraphEntry, runtime_inputs: list[torch.Tensor] - ) -> None: - """Copy runtime inputs into capture-time inputs so replay sees new values""" - captured_inputs = entry.input_tensors - if captured_inputs is None: - return - if len(runtime_inputs) != len(captured_inputs): - raise RuntimeError( - "CUDAGraph replay input tensor count mismatch. " - f"Expected {len(captured_inputs)}, got {len(runtime_inputs)}." - ) - - for i, (captured, runtime) in enumerate(zip(captured_inputs, runtime_inputs)): - if runtime.data_ptr() == captured.data_ptr(): - continue - if ( - runtime.device != captured.device - or runtime.dtype != captured.dtype - or runtime.shape != captured.shape - ): - raise RuntimeError( - "CUDAGraph replay input tensor mismatch at index " - f"{i}: captured(device={captured.device}, dtype={captured.dtype}, " - f"shape={tuple(captured.shape)}), runtime(device={runtime.device}, " - f"dtype={runtime.dtype}, shape={tuple(runtime.shape)}). " - "Disable CUDA graphs (e.g. --enforce-eager / cudagraph_mode=NONE) " - "or ensure inputs are passed from stable, preallocated buffers." - ) - captured.copy_(runtime, non_blocking=True) - def __call__(self, *args, **kwargs): forward_context = get_forward_context() batch_descriptor = forward_context.batch_descriptor @@ -283,9 +241,9 @@ class CUDAGraphWrapper: # validate that cudagraph capturing is legal at this point. validate_cudagraph_capturing_enabled() - tensor_inputs = self._tensor_inputs(args) - entry.input_tensors = tensor_inputs - entry.input_addresses = [x.data_ptr() for x in tensor_inputs] + input_addresses = [ + x.data_ptr() for x in args if isinstance(x, torch.Tensor) + ] cudagraph = torch.cuda.CUDAGraph() with ExitStack() as stack: @@ -327,34 +285,16 @@ class CUDAGraphWrapper: # the weak ref of the output, so that pytorch can correctly # manage the memory during cuda graph capture return output - - # if runtime input tensor addresses differ from the capture-time addresses, - # CUDA graph replay would read stale values. - # we need tocopy runtime inputs into the capture-time buffers before replaying. - copy_inputs = self.compilation_config.cudagraph_copy_inputs - if self.is_debugging_mode or copy_inputs: - runtime_inputs = self._tensor_inputs(args) - new_input_addresses = [x.data_ptr() for x in runtime_inputs] - if new_input_addresses != entry.input_addresses: - if copy_inputs: - self._copy_runtime_inputs_into_captured_inputs( - entry, runtime_inputs - ) - if self.is_debugging_mode: - logger.warning_once( - "CUDAGraph replay saw different input addresses, s" - "copied runtime inputs into captured input buffers.", - scope="local", - ) - else: - assert new_input_addresses == entry.input_addresses, ( - f"Input addresses for cudagraphs are different " - f"during replay. Expected {entry.input_addresses}, " - f"got {new_input_addresses}. " - "Set compilation_config.cudagraph_copy_inputs=True to copy " - "runtime inputs into captured buffers before replay, or " - "disable CUDA graphs." - ) + if self.is_debugging_mode: + # check if the input addresses are the same + new_input_addresses = [ + x.data_ptr() for x in args if isinstance(x, torch.Tensor) + ] + assert new_input_addresses == entry.input_addresses, ( + f"Input addresses for cudagraphs are different " + f"during replay. Expected {entry.input_addresses}, " + f"got {new_input_addresses}" + ) entry.cudagraph.replay() return entry.output