ccf-vllm-ascend/vllm_ascend/model_executor/offloader/prefetch.py

275 lines
9.7 KiB
Python

from collections.abc import Generator
from dataclasses import dataclass
import torch
import torch.nn as nn
from vllm.logger import init_logger
from vllm.model_executor.offloader.base import BaseOffloader
from vllm.utils.platform_utils import is_pin_memory_available
from vllm_ascend.model_executor.offloader.selection import select_offload_layers
logger = init_logger(__name__)
@dataclass
class _ParamState:
name: str
param: nn.Parameter
cpu_tensor: torch.Tensor
original_device: torch.device
num_bytes: int
static_tensor: torch.Tensor | None = None
class _StaticBufferPool:
def __init__(self, *, slot_capacity: int, device: torch.device):
self.slot_capacity = slot_capacity
self.device = device
self.total_bytes = 0
self._buffers: dict[tuple[str, tuple[int, ...], tuple[int, ...], torch.dtype], list[torch.Tensor]] = {}
@staticmethod
def _num_bytes(tensor: torch.Tensor) -> int:
return tensor.numel() * tensor.element_size()
def get_buffer(self, state: _ParamState, slot_idx: int) -> torch.Tensor:
key = (
state.name,
tuple(state.cpu_tensor.shape),
tuple(state.cpu_tensor.stride()),
state.cpu_tensor.dtype,
)
if key not in self._buffers:
buffers = []
for _ in range(self.slot_capacity):
buffer = torch.empty_strided(
size=state.cpu_tensor.shape,
stride=state.cpu_tensor.stride(),
dtype=state.cpu_tensor.dtype,
device=self.device,
)
buffers.append(buffer)
self.total_bytes += self._num_bytes(buffer)
self._buffers[key] = buffers
return self._buffers[key][slot_idx % self.slot_capacity]
class _PrefetchModuleOffloader:
def __init__(
self,
*,
module: nn.Module,
layer_idx: int,
slot_idx: int,
offload_params: set[str],
):
self.module = module
self.layer_idx = layer_idx
self.slot_idx = slot_idx
self.offload_params = offload_params
self.param_states: list[_ParamState] = []
self.offloaded_bytes = 0
self.copy_done_event = None
self._capture_params()
def _matches_param(self, name: str) -> bool:
if not self.offload_params:
return True
return any(f".{part}." in f".{name}." for part in self.offload_params)
@staticmethod
def _tensor_nbytes(tensor: torch.Tensor) -> int:
return tensor.numel() * tensor.element_size()
@staticmethod
def _copy_to_cpu_storage(tensor: torch.Tensor) -> torch.Tensor:
cpu_tensor = torch.empty_strided(
size=tensor.shape,
stride=tensor.stride(),
dtype=tensor.dtype,
device="cpu",
pin_memory=is_pin_memory_available(),
)
cpu_tensor.copy_(tensor.detach())
return cpu_tensor
def _capture_params(self) -> None:
for name, param in self.module.named_parameters(recurse=True):
if not self._matches_param(name):
continue
cpu_tensor = self._copy_to_cpu_storage(param)
self.param_states.append(
_ParamState(
name=name,
param=param,
cpu_tensor=cpu_tensor,
original_device=param.device,
num_bytes=self._tensor_nbytes(cpu_tensor),
)
)
param.data = cpu_tensor
def post_init(self, pool: _StaticBufferPool | None) -> None:
self.offloaded_bytes = 0
for state in self.param_states:
state.cpu_tensor = self._copy_to_cpu_storage(state.param)
state.num_bytes = self._tensor_nbytes(state.cpu_tensor)
self.offloaded_bytes += state.num_bytes
if pool is None:
state.param.data = state.cpu_tensor
continue
state.static_tensor = pool.get_buffer(state, self.slot_idx)
state.param.data = state.static_tensor
def start_prefetch(self, copy_stream) -> None:
if copy_stream is None:
return
if not self.param_states or self.param_states[0].static_tensor is None:
return
fork_event = torch.npu.Event()
torch.npu.current_stream().record_event(fork_event)
copy_stream.wait_event(fork_event)
with torch.npu.stream(copy_stream):
for state in self.param_states:
assert state.static_tensor is not None
state.static_tensor.copy_(state.cpu_tensor, non_blocking=True)
self.copy_done_event = torch.npu.Event()
self.copy_done_event.record(copy_stream)
def wait_prefetch(self) -> None:
if self.copy_done_event is not None:
torch.npu.current_stream().wait_event(self.copy_done_event)
class AscendPrefetchOffloader(BaseOffloader):
def __init__(
self,
*,
group_size: int,
num_in_group: int,
prefetch_step: int,
offload_params: set[str] | None = None,
):
if prefetch_step <= 0:
raise ValueError("prefetch_step must be > 0")
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.selected_layer_indices: tuple[int, ...] = ()
self.module_offloaders: list[_PrefetchModuleOffloader] = []
self.total_offloaded_bytes = 0
self.static_buffer_bytes = 0
self.buffer_pool: _StaticBufferPool | None = None
self.copy_stream = None
@staticmethod
def _npu_available() -> bool:
return hasattr(torch, "npu") and torch.npu.is_available()
def _get_device(self) -> torch.device | None:
for module_offloader in self.module_offloaders:
for state in module_offloader.param_states:
if state.original_device.type != "cpu":
return state.original_device
if self._npu_available():
return torch.device("npu")
return None
def wrap_modules(
self,
modules_generator: Generator[nn.Module, None, None],
) -> list[nn.Module]:
modules = list(modules_generator)
selection = select_offload_layers(
num_layers=len(modules),
group_size=self.group_size,
num_in_group=self.num_in_group,
)
self.selected_layer_indices = selection.layer_indices
if not selection.enabled:
return modules
offload_position = 0
for layer_idx in selection.layer_indices:
module = modules[layer_idx]
if not any(True for _ in module.parameters(recurse=True)):
continue
module_offloader = _PrefetchModuleOffloader(
module=module,
layer_idx=layer_idx,
slot_idx=offload_position % self.prefetch_step,
offload_params=self.offload_params,
)
if not module_offloader.param_states:
continue
self.module_offloaders.append(module_offloader)
self._install_forward_hook(offload_position, module_offloader)
offload_position += 1
logger.info(
"[AscendPrefetchOffloader] selected_layers=%s, wrapped=%d, "
"group_size=%d, num_in_group=%d, prefetch_step=%d",
self.selected_layer_indices,
len(self.module_offloaders),
self.group_size,
self.num_in_group,
self.prefetch_step,
)
return modules
def _install_forward_hook(
self,
offload_position: int,
module_offloader: _PrefetchModuleOffloader,
) -> None:
original_forward = module_offloader.module.forward
def forward(*args, **kwargs):
module_offloader.wait_prefetch()
output = original_forward(*args, **kwargs)
if self.module_offloaders:
next_position = (offload_position + self.prefetch_step) % len(self.module_offloaders)
self.module_offloaders[next_position].start_prefetch(self.copy_stream)
return output
module_offloader.module.forward = forward
def post_init(self) -> None:
self.total_offloaded_bytes = 0
device = self._get_device()
if device is not None and self._npu_available():
self.copy_stream = torch.npu.Stream(device=device)
self.buffer_pool = _StaticBufferPool(
slot_capacity=self.prefetch_step,
device=device,
)
else:
self.copy_stream = None
self.buffer_pool = None
for module_offloader in self.module_offloaders:
module_offloader.post_init(self.buffer_pool)
self.total_offloaded_bytes += module_offloader.offloaded_bytes
self.static_buffer_bytes = 0 if self.buffer_pool is None else self.buffer_pool.total_bytes
logger.info(
"[AscendPrefetchOffloader] initialized wrapped=%d, saved=%.4f GB, "
"static_buffer=%.4f GB, selected_layers=%s",
len(self.module_offloaders),
self.total_offloaded_bytes / float(2**30),
self.static_buffer_bytes / float(2**30),
self.selected_layer_indices,
)
for module_offloader in self.module_offloaders[: self.prefetch_step]:
module_offloader.start_prefetch(self.copy_stream)
def sync_prev_onload(self) -> None:
if self.copy_stream is not None:
torch.npu.current_stream().wait_stream(self.copy_stream)
def join_after_forward(self) -> None:
self.sync_prev_onload()