0
vllm_ascend/model_executor/__init__.py
Normal file
0
vllm_ascend/model_executor/__init__.py
Normal file
0
vllm_ascend/model_executor/offloader/__init__.py
Normal file
0
vllm_ascend/model_executor/offloader/__init__.py
Normal file
263
vllm_ascend/model_executor/offloader/prefetch.py
Normal file
263
vllm_ascend/model_executor/offloader/prefetch.py
Normal file
@@ -0,0 +1,263 @@
|
||||
# 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()
|
||||
Reference in New Issue
Block a user