forked from vllm-ascend/ccf-vllm-ascend
275 lines
9.7 KiB
Python
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()
|