Files
enginex-ascend-910-vllm/vllm_ascend/model_executor/offloader/prefetch.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

264 lines
9.6 KiB
Python

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""NPU port of PrefetchOffloader — uses torch.npu.* APIs including is_current_stream_capturing."""
from collections.abc import Generator
import torch
import torch.nn as nn
import torch_npu # noqa: F401
import vllm.model_executor.offloader.prefetch_ops # noqa: F401
from vllm.logger import logger
from vllm.model_executor.offloader.base import BaseOffloader, should_pin_memory
from vllm.model_executor.offloader.prefetch import (
ParamInfo,
StaticBufferPool,
_BaseParamOffloader,
)
class NPUPrefetchOffloader(BaseOffloader):
"""NPU version of PrefetchOffloader — replaces torch.cuda.* with torch.npu.*."""
def __init__(
self,
group_size: int,
num_in_group: int,
prefetch_step: int,
offload_params: set[str] | None = None,
mode: str = "cpu",
):
self.group_size = group_size
self.num_in_group = num_in_group
self.prefetch_step = prefetch_step
self.offload_params = offload_params or set()
self.mode = mode
self.copy_stream = torch.npu.Stream()
self.module_offloaders: list[_NPUModuleOffloader] = []
self.buffer_pool: StaticBufferPool | None = None
self.total_offloaded_bytes = 0
def wrap_modules(
self,
modules_generator: Generator[nn.Module, None, None],
) -> list[nn.Module]:
assert len(self.module_offloaders) == 0
all_modules = []
offload_modules = []
for module_index, module in enumerate(modules_generator):
all_modules.append(module)
if module_index % self.group_size >= self.group_size - self.num_in_group:
if self.offload_params:
whitelist = [
name
for name, _ in module.named_parameters()
if any(f".{p}." in f".{name}." for p in self.offload_params)
]
else:
whitelist = [name for name, _ in module.named_parameters()]
if not whitelist:
continue
offload_modules.append(module)
self.module_offloaders.append(
_NPUModuleOffloader(
mode=self.mode,
module=module,
copy_stream=self.copy_stream,
whitelist_param_names=whitelist,
layer_idx=len(self.module_offloaders),
)
)
for index, module in enumerate(offload_modules):
self._hook_module_forward(index, module)
return all_modules
def _hook_module_forward(self, index: int, module: nn.Module):
original_forward = module.forward
def forward(*args, **kwargs):
module.forward = original_forward
input_tensor = args[0] if args else kwargs.get("hidden_states")
torch.ops.vllm.wait_prefetch(input_tensor, index)
output = original_forward(*args, **kwargs)
next_index = (index + self.prefetch_step) % len(self.module_offloaders)
if isinstance(output, tuple):
torch.ops.vllm.start_prefetch(output[0], next_index)
else:
torch.ops.vllm.start_prefetch(output, next_index)
module.forward = forward
return output
module.forward = forward
def _wait_for_layer(self, layer_idx: int):
offloader = self.module_offloaders[layer_idx]
if torch.npu.is_current_stream_capturing():
if not offloader._prefetch_in_capture:
return
torch.npu.current_stream().wait_event(offloader._copy_done_event)
offloader._prefetch_in_capture = False
else:
if offloader._event_valid_for_eager:
torch.npu.current_stream().wait_event(offloader._copy_done_event)
else:
torch.npu.current_stream().wait_stream(self.copy_stream)
def _start_prefetch(self, layer_idx: int):
self.module_offloaders[layer_idx].start_onload_to_static()
def sync_prev_onload(self):
torch.npu.current_stream().wait_stream(self.copy_stream)
def join_after_forward(self):
for offloader in self.module_offloaders:
if offloader._prefetch_in_capture:
torch.npu.current_stream().wait_event(offloader._copy_done_event)
offloader._prefetch_in_capture = False
def post_init(self):
for offloader in self.module_offloaders:
offloader.sync_cpu_storage()
param_infos: list[ParamInfo] = []
device: torch.device | None = None
for offloader in self.module_offloaders:
param_infos.extend(offloader.get_param_infos())
if device is None:
device = offloader.device
if device is None:
return
self.buffer_pool = StaticBufferPool(
param_infos=param_infos,
slot_capacity=self.prefetch_step,
device=device,
)
for idx, offloader in enumerate(self.module_offloaders):
slot_idx = idx % self.prefetch_step
offloader.assign_buffer_slot(self.buffer_pool, slot_idx)
for offloader in self.module_offloaders:
offloader.post_init()
self.total_offloaded_bytes += offloader.offloaded_bytes
logger.info_once(
f"[NPUPrefetchOffloader] Initialized {len(self.module_offloaders)} modules. "
f"Total NPU memory saved: {self.total_offloaded_bytes / 1e9:.4f} GB, "
f"Static buffer pool: {self.buffer_pool.total_bytes / 1e9:.4f} GB "
f"(group_size={self.group_size}, num_in_group={self.num_in_group}, "
f"prefetch_step={self.prefetch_step})"
)
for i in range(min(self.prefetch_step, len(self.module_offloaders))):
self.module_offloaders[i].start_onload_to_static()
class _NPUModuleOffloader:
"""NPU version of _ModuleOffloader: all torch.cuda.* → torch.npu.*."""
def __init__(
self,
mode: str,
module: nn.Module,
copy_stream: torch.npu.Stream,
whitelist_param_names: list[str],
layer_idx: int,
):
self.mode = mode
self.module = module
self.device = next(module.parameters()).device
self.copy_stream = copy_stream
self.layer_idx = layer_idx
self.offloaded_bytes = 0
self._copy_done_event = torch.npu.Event()
self._event_valid_for_eager = False
self._prefetch_in_capture = False
assert self.device != torch.device("cpu")
self._buffer_pool: StaticBufferPool | None = None
self._buffer_slot_idx: int = 0
param_dict = dict(self.module.named_parameters())
assert all(name in param_dict for name in whitelist_param_names)
self._param_offloaders = {
name: _BaseParamOffloader.create(mode, module=module, param_name=name) for name in whitelist_param_names
}
def post_init(self):
for param_offloader in self._param_offloaders.values():
param_offloader.post_init()
self.offloaded_bytes += param_offloader.offloaded_bytes
def sync_cpu_storage(self):
for param_offloader in self._param_offloaders.values():
param_offloader.sync_cpu_storage()
deleted = [
name for name, offloader in self._param_offloaders.items() if getattr(offloader, "_param_deleted", False)
]
for name in deleted:
del self._param_offloaders[name]
def get_param_infos(self) -> list[ParamInfo]:
infos = []
for name, offloader in self._param_offloaders.items():
cpu_storage = offloader._cpu_storage
assert cpu_storage is not None
infos.append(
ParamInfo(
name=name,
shape=tuple(cpu_storage.shape),
stride=tuple(cpu_storage.stride()),
dtype=cpu_storage.dtype,
)
)
return infos
def assign_buffer_slot(self, pool: StaticBufferPool, slot_idx: int):
self._buffer_pool = pool
self._buffer_slot_idx = slot_idx
for name, offloader in self._param_offloaders.items():
cpu_storage = offloader._cpu_storage
assert cpu_storage is not None
buffer = pool.get_buffer(
name=name,
shape=tuple(cpu_storage.shape),
stride=tuple(cpu_storage.stride()),
dtype=cpu_storage.dtype,
slot_idx=slot_idx,
)
offloader.assign_static_buffer(buffer)
def start_onload_to_static(self):
assert self._buffer_pool is not None
self._prefetch_in_capture = torch.npu.is_current_stream_capturing()
fork_event = torch.npu.Event()
torch.npu.current_stream().record_event(fork_event)
self.copy_stream.wait_event(fork_event)
with torch.npu.stream(self.copy_stream):
for name, offloader in self._param_offloaders.items():
cpu_storage = offloader._cpu_storage
gpu_buffer = offloader._gpu_buffer
assert cpu_storage is not None
assert gpu_buffer is not None
assert not should_pin_memory() or cpu_storage.is_pinned(), f"CPU storage for {name} is not pinned!"
gpu_buffer.copy_(cpu_storage, non_blocking=True)
self._copy_done_event.record(self.copy_stream)
self._event_valid_for_eager = not torch.npu.is_current_stream_capturing()