@@ -17,24 +17,24 @@
|
||||
# CANN-mem-based pytorch pluggable allocator to implement sleep mode.
|
||||
#
|
||||
import dataclasses
|
||||
import gc
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Callable, Dict, Optional, Tuple, Union
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from acl.rt import memcpy # type: ignore # noqa: F401
|
||||
from vllm.logger import logger
|
||||
|
||||
from vllm_ascend.platform import NPUPlatform
|
||||
|
||||
|
||||
def find_loaded_library(lib_name) -> Optional[str]:
|
||||
def find_loaded_library(lib_name) -> str | None:
|
||||
"""
|
||||
According to according to https://man7.org/linux/man-pages/man5/proc_pid_maps.5.html,
|
||||
the file `/proc/self/maps` contains the memory maps of the process, which includes the
|
||||
shared libraries loaded by the process. We can use this file to find the path of the
|
||||
a loaded library.
|
||||
""" # noqa
|
||||
""" # noqa
|
||||
found_line = None
|
||||
with open("/proc/self/maps") as f:
|
||||
for line in f:
|
||||
@@ -49,20 +49,22 @@ def find_loaded_library(lib_name) -> Optional[str]:
|
||||
start = found_line.index("/")
|
||||
path = found_line[start:].strip()
|
||||
filename = path.split("/")[-1]
|
||||
assert filename.rpartition(".so")[0].startswith(lib_name), \
|
||||
f"Unexpected filename: {filename} for library {lib_name}"
|
||||
assert filename.rpartition(".so")[0].startswith(lib_name), f"Unexpected filename: {filename} for library {lib_name}"
|
||||
return path
|
||||
|
||||
|
||||
camem_available = False
|
||||
try:
|
||||
from vllm_ascend.vllm_ascend_C import ( # type: ignore # noqa: F401
|
||||
init_module, python_create_and_map, python_unmap_and_release)
|
||||
init_module,
|
||||
python_create_and_map,
|
||||
python_unmap_and_release,
|
||||
)
|
||||
|
||||
lib_name = find_loaded_library("vllm_ascend_C")
|
||||
camem_available = True
|
||||
except ImportError as e:
|
||||
logger.warning(
|
||||
"Failed to import vllm_ascend_C:%s. Sleep mode will be disabled. ", e)
|
||||
logger.warning("Failed to import vllm_ascend_C:%s. Sleep mode will be disabled. ", e)
|
||||
init_module = None
|
||||
python_create_and_map = None
|
||||
python_unmap_and_release = None
|
||||
@@ -70,14 +72,14 @@ except ImportError as e:
|
||||
libcudart = None
|
||||
|
||||
# py_device, py_alignedSize, py_d_mem, py_p_memHandle
|
||||
HandleType = Tuple[int, int, int, int]
|
||||
HandleType = tuple[int, int, int, int]
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class AllocationData:
|
||||
handle: HandleType
|
||||
tag: str
|
||||
cpu_backup_tensor: Optional[torch.Tensor] = None
|
||||
cpu_backup_tensor: torch.Tensor | None = None
|
||||
|
||||
|
||||
def create_and_map(allocation_handle: HandleType) -> None:
|
||||
@@ -90,18 +92,18 @@ def unmap_and_release(allocation_handle: HandleType) -> None:
|
||||
|
||||
def get_pluggable_allocator(
|
||||
python_malloc_fn: Callable[[tuple[int, int, int, int]], None],
|
||||
python_free_func: Callable[[int], tuple[int, int, int, int]]
|
||||
python_free_func: Callable[[int], tuple[int, int, int, int]],
|
||||
) -> torch.npu.memory.NPUPluggableAllocator:
|
||||
init_module(python_malloc_fn, python_free_func)
|
||||
new_alloc = torch.npu.memory.NPUPluggableAllocator(lib_name, 'my_malloc',
|
||||
'my_free')
|
||||
new_alloc = torch.npu.memory.NPUPluggableAllocator(lib_name, "my_malloc", "my_free")
|
||||
return new_alloc
|
||||
|
||||
|
||||
@contextmanager
|
||||
def use_memory_pool_with_allocator(
|
||||
python_malloc_fn: Callable[[tuple[int, int, int, int]], None],
|
||||
python_free_func: Callable[[int], tuple[int, int, int, int]]):
|
||||
python_malloc_fn: Callable[[tuple[int, int, int, int]], None],
|
||||
python_free_func: Callable[[int], tuple[int, int, int, int]],
|
||||
):
|
||||
new_alloc = get_pluggable_allocator(python_malloc_fn, python_free_func)
|
||||
mem_pool = torch.npu.memory.MemPool(new_alloc._allocator)
|
||||
with torch.npu.memory.use_mem_pool(mem_pool):
|
||||
@@ -129,8 +131,11 @@ class CaMemAllocator:
|
||||
the global variable will be overwritten and the free callback will
|
||||
not work as expected.
|
||||
"""
|
||||
|
||||
instance = None
|
||||
default_tag: str = "default"
|
||||
# Allocations with this tag stay mapped across sleep/wake cycles.
|
||||
sleep_persistent_tag: str = "sleep_persistent"
|
||||
|
||||
@staticmethod
|
||||
def get_instance() -> "CaMemAllocator":
|
||||
@@ -145,22 +150,22 @@ class CaMemAllocator:
|
||||
|
||||
def __init__(self):
|
||||
conf = os.environ.get("PYTORCH_NPU_ALLOC_CONF", "")
|
||||
assert "expandable_segments:True" not in conf, \
|
||||
("Expandable segments are not compatible with memory pool. "
|
||||
assert "expandable_segments:True" not in conf, (
|
||||
"Expandable segments are not compatible with memory pool. "
|
||||
"Please track https://github.com/pytorch/pytorch/issues/147851 "
|
||||
"for the latest updates.")
|
||||
"for the latest updates."
|
||||
)
|
||||
|
||||
self.pointer_to_data: Dict[int, AllocationData] = {}
|
||||
self.pointer_to_data: dict[int, AllocationData] = {}
|
||||
self.current_tag: str = CaMemAllocator.default_tag
|
||||
self.allocator_and_pools: Dict[str, Any] = {}
|
||||
self.allocator_and_pools: dict[str, Any] = {}
|
||||
|
||||
def python_malloc_callback(self, allocation_handle: HandleType) -> None:
|
||||
"""
|
||||
Internal method to store the allocation data
|
||||
when memory is allocated in the memory pool."""
|
||||
py_d_mem = allocation_handle[2]
|
||||
self.pointer_to_data[py_d_mem] = AllocationData(
|
||||
allocation_handle, self.current_tag)
|
||||
self.pointer_to_data[py_d_mem] = AllocationData(allocation_handle, self.current_tag)
|
||||
return
|
||||
|
||||
def python_free_callback(self, ptr: int) -> HandleType:
|
||||
@@ -172,13 +177,10 @@ class CaMemAllocator:
|
||||
data.cpu_backup_tensor = None
|
||||
return data.handle
|
||||
|
||||
def sleep(
|
||||
self,
|
||||
offload_tags: Optional[Union[Tuple[str, ...],
|
||||
str]] = None) -> None:
|
||||
def sleep(self, offload_tags: tuple[str, ...] | str | None = None) -> None:
|
||||
"""
|
||||
Put the allocator in sleep mode.
|
||||
All data in the memory allocation with the specified tag will be
|
||||
All data in the memory allocation with the specified tag will be
|
||||
offloaded to CPU memory, and others will be discarded.
|
||||
:param offload_tags: The tags of the memory allocation that will be
|
||||
offloaded. The rest of the memory allocation will be discarded.
|
||||
@@ -186,55 +188,81 @@ class CaMemAllocator:
|
||||
if offload_tags is None:
|
||||
# by default, allocated tensors are offloaded
|
||||
# when the allocator sleeps
|
||||
offload_tags = (CaMemAllocator.default_tag, )
|
||||
offload_tags = (CaMemAllocator.default_tag,)
|
||||
elif isinstance(offload_tags, str):
|
||||
offload_tags = (offload_tags, )
|
||||
offload_tags = (offload_tags,)
|
||||
|
||||
assert isinstance(offload_tags, tuple)
|
||||
|
||||
offload_count = sum(1 for data in self.pointer_to_data.values() if data.tag in offload_tags)
|
||||
logger.info(
|
||||
"CaMem sleep: offloading %s/%s allocations (tags=%s)",
|
||||
offload_count,
|
||||
len(self.pointer_to_data),
|
||||
offload_tags,
|
||||
)
|
||||
for ptr, data in self.pointer_to_data.items():
|
||||
if data.tag == CaMemAllocator.sleep_persistent_tag:
|
||||
# This memory is not offloaded or released during sleep.
|
||||
continue
|
||||
handle = data.handle
|
||||
if data.tag in offload_tags:
|
||||
size_in_bytes = handle[1]
|
||||
cpu_backup_tensor = torch.empty(
|
||||
size_in_bytes,
|
||||
dtype=torch.uint8,
|
||||
device='cpu',
|
||||
pin_memory=NPUPlatform.is_pin_memory_available())
|
||||
cpu_backup_tensor = torch.empty(size_in_bytes, dtype=torch.uint8, device="cpu", pin_memory=True)
|
||||
cpu_ptr = cpu_backup_tensor.data_ptr()
|
||||
ACL_MEMCPY_DEVICE_TO_HOST = 2
|
||||
dest_max = cpu_ptr + size_in_bytes * 2
|
||||
memcpy(cpu_ptr, dest_max, ptr, size_in_bytes,
|
||||
ACL_MEMCPY_DEVICE_TO_HOST)
|
||||
memcpy(cpu_ptr, dest_max, ptr, size_in_bytes, ACL_MEMCPY_DEVICE_TO_HOST)
|
||||
data.cpu_backup_tensor = cpu_backup_tensor
|
||||
unmap_and_release(handle)
|
||||
|
||||
def wake_up(self, tags: Optional[list[str]] = None) -> None:
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
|
||||
def wake_up(self, tags: list[str] | None = None) -> None:
|
||||
"""
|
||||
Wake up the allocator from sleep mode.
|
||||
All data that is previously offloaded will be loaded back to GPU
|
||||
All data that is previously offloaded will be loaded back to GPU
|
||||
memory, and the rest of the data will have empty memory."""
|
||||
restore_count = sum(1 for data in self.pointer_to_data.values() if tags is None or data.tag in tags)
|
||||
logger.info(
|
||||
"CaMem wake_up: restoring %s/%s allocations (tags=%s)",
|
||||
restore_count,
|
||||
len(self.pointer_to_data),
|
||||
tags or "all",
|
||||
)
|
||||
for ptr, data in self.pointer_to_data.items():
|
||||
if data.tag == CaMemAllocator.sleep_persistent_tag:
|
||||
# It was never released in sleep(), so there is nothing to remap.
|
||||
continue
|
||||
if tags is None or data.tag in tags:
|
||||
handle = data.handle
|
||||
create_and_map(handle)
|
||||
if data.cpu_backup_tensor is not None:
|
||||
cpu_backup_tensor = data.cpu_backup_tensor
|
||||
if cpu_backup_tensor is not None:
|
||||
size_in_bytes = cpu_backup_tensor.numel(
|
||||
) * cpu_backup_tensor.element_size()
|
||||
size_in_bytes = cpu_backup_tensor.numel() * cpu_backup_tensor.element_size()
|
||||
cpu_ptr = cpu_backup_tensor.data_ptr()
|
||||
ACL_MEMCPY_HOST_TO_DEVICE = 1
|
||||
dest_max = ptr + size_in_bytes * 2
|
||||
memcpy(ptr, dest_max, cpu_ptr, size_in_bytes,
|
||||
ACL_MEMCPY_HOST_TO_DEVICE)
|
||||
memcpy(ptr, dest_max, cpu_ptr, size_in_bytes, ACL_MEMCPY_HOST_TO_DEVICE)
|
||||
data.cpu_backup_tensor = None
|
||||
|
||||
@contextmanager
|
||||
def use_memory_pool(self, tag: Optional[str] = None):
|
||||
def use_allocation_tag(self, tag: str):
|
||||
"""Temporarily override the tag assigned to new allocations."""
|
||||
old_tag = self.current_tag
|
||||
self.current_tag = tag
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self.current_tag = old_tag
|
||||
|
||||
@contextmanager
|
||||
def use_memory_pool(self, tag: str | None = None):
|
||||
"""
|
||||
A context manager to use the memory pool.
|
||||
All memory allocation created inside the context will be allocated
|
||||
All memory allocation created inside the context will be allocated
|
||||
in the memory pool, and has the specified tag.
|
||||
:param tag: The tag of the memory allocation. If None, the default tag
|
||||
will be used.
|
||||
@@ -246,8 +274,7 @@ class CaMemAllocator:
|
||||
|
||||
old_tag = self.current_tag
|
||||
self.current_tag = tag
|
||||
with use_memory_pool_with_allocator(self.python_malloc_callback,
|
||||
self.python_free_callback) as data:
|
||||
with use_memory_pool_with_allocator(self.python_malloc_callback, self.python_free_callback) as data:
|
||||
# start to hit another PyTorch bug in PyTorch 2.6,
|
||||
# possibly because of gc-related issue w.r.t. the allocator and
|
||||
# the memory pool.
|
||||
|
||||
186
vllm_ascend/device_allocator/sleep_mem_optimized.py
Normal file
186
vllm_ascend/device_allocator/sleep_mem_optimized.py
Normal file
@@ -0,0 +1,186 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# Copyright 2023 The vLLM team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, MutableMapping
|
||||
from dataclasses import fields
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from vllm.config import VllmConfig, set_current_vllm_config
|
||||
from vllm.distributed.parallel_state import _groups
|
||||
from vllm.logger import logger
|
||||
from vllm.utils.mem_constants import GiB_bytes
|
||||
|
||||
from vllm_ascend.compilation import acl_graph
|
||||
|
||||
|
||||
class SleepWakeupManager:
|
||||
def __init__(self, vllm_config: VllmConfig, worker: Any, model_runner_getter: Callable[[], Any]):
|
||||
self.acl_graph = AclGraphSleepWakeupManager(vllm_config, model_runner_getter)
|
||||
self.hccl = HcclSleepWakeupManager(vllm_config, worker)
|
||||
self._model_runner_getter = model_runner_getter
|
||||
|
||||
@staticmethod
|
||||
def _measure_memory_released(cleanup: Callable[[], None]) -> int:
|
||||
free_bytes_before_cleanup = torch.npu.mem_get_info()[0]
|
||||
cleanup()
|
||||
free_bytes_after_cleanup = torch.npu.mem_get_info()[0]
|
||||
return max(free_bytes_after_cleanup - free_bytes_before_cleanup, 0)
|
||||
|
||||
def sleep(self) -> None:
|
||||
model_runner = self._model_runner_getter()
|
||||
free_bytes_before_cleanup = torch.npu.mem_get_info()[0]
|
||||
if model_runner.use_aclgraph:
|
||||
self.acl_graph.sleep()
|
||||
self.hccl.sleep()
|
||||
free_bytes_after_cleanup = torch.npu.mem_get_info()[0]
|
||||
free_mem = free_bytes_after_cleanup - free_bytes_before_cleanup
|
||||
logger.info(
|
||||
"Sleep mode released HCCL and attention workspace memory: %.3f GiB.",
|
||||
free_mem / GiB_bytes,
|
||||
)
|
||||
|
||||
def wakeup(self, tags: list[str] | None = None) -> None:
|
||||
self.hccl.wakeup()
|
||||
model_runner = self._model_runner_getter()
|
||||
if model_runner.use_aclgraph:
|
||||
self.acl_graph.wakeup(tags)
|
||||
|
||||
|
||||
class AclGraphSleepWakeupManager:
|
||||
def __init__(self, vllm_config: VllmConfig, model_runner_getter: Callable[[], Any]):
|
||||
self.vllm_config = vllm_config
|
||||
self._model_runner_getter = model_runner_getter
|
||||
|
||||
@staticmethod
|
||||
def clear_attention_workspaces(params) -> None:
|
||||
if params is None:
|
||||
return
|
||||
for num_tokens in params.workspaces:
|
||||
params.workspaces[num_tokens] = None
|
||||
|
||||
@classmethod
|
||||
def clear_all_attention_workspaces(cls) -> None:
|
||||
cls.clear_attention_workspaces(acl_graph._graph_params)
|
||||
cls.clear_attention_workspaces(acl_graph._draft_graph_params)
|
||||
cls.clear_attention_workspaces(acl_graph._draft_graph_prefill_params)
|
||||
|
||||
@staticmethod
|
||||
def reset_graph_params(params) -> None:
|
||||
if params is None:
|
||||
return
|
||||
for graph_field in fields(params):
|
||||
attr_dict = getattr(params, graph_field.name, None)
|
||||
if not isinstance(attr_dict, MutableMapping):
|
||||
continue
|
||||
for num_tokens, value in attr_dict.items():
|
||||
if isinstance(value, list):
|
||||
attr_dict[num_tokens] = []
|
||||
|
||||
@classmethod
|
||||
def reset_all_graph_params(cls) -> None:
|
||||
cls.reset_graph_params(acl_graph._graph_params)
|
||||
cls.reset_graph_params(acl_graph._draft_graph_params)
|
||||
cls.reset_graph_params(acl_graph._draft_graph_prefill_params)
|
||||
for wrapper in list(acl_graph._acl_graph_wrappers):
|
||||
wrapper.concrete_aclgraph_entries.clear()
|
||||
wrapper.first_run_finished = False
|
||||
|
||||
@staticmethod
|
||||
def reset_model_runner_graph_manager(model_runner: Any) -> None:
|
||||
manager = getattr(model_runner, "cudagraph_manager", None)
|
||||
if manager is None:
|
||||
return
|
||||
if hasattr(manager, "graphs"):
|
||||
manager.graphs.clear()
|
||||
if hasattr(manager, "_graphs_captured"):
|
||||
manager._graphs_captured = False
|
||||
|
||||
def sleep(self) -> None:
|
||||
self.clear_all_attention_workspaces()
|
||||
self.reset_all_graph_params()
|
||||
self.reset_model_runner_graph_manager(self._model_runner_getter())
|
||||
|
||||
def wakeup(self, tags: list[str] | None = None) -> None:
|
||||
if tags is not None and "kv_cache" not in tags:
|
||||
# Level-2 wakeup restores weights before external weight loading;
|
||||
# recapture graphs only after KV cache is restored.
|
||||
return
|
||||
model_runner = self._model_runner_getter()
|
||||
with set_current_vllm_config(self.vllm_config):
|
||||
model_runner.capture_model()
|
||||
|
||||
|
||||
class HcclSleepWakeupManager:
|
||||
def __init__(self, vllm_config: VllmConfig, worker: Any):
|
||||
self.vllm_config = vllm_config
|
||||
self.worker = worker
|
||||
|
||||
@staticmethod
|
||||
def iter_alive_group_coordinators():
|
||||
seen: set[int] = set()
|
||||
for group_ref in list(_groups.values()):
|
||||
group = group_ref()
|
||||
if group is None or id(group) in seen:
|
||||
continue
|
||||
seen.add(id(group))
|
||||
yield group
|
||||
|
||||
@classmethod
|
||||
def destroy_hccl(cls) -> int:
|
||||
num_destroyed = 0
|
||||
for group in cls.iter_alive_group_coordinators():
|
||||
if group.destroy_hccl():
|
||||
num_destroyed += 1
|
||||
return num_destroyed
|
||||
|
||||
@classmethod
|
||||
def restore_hccl(cls) -> int:
|
||||
num_restored = 0
|
||||
for group in cls.iter_alive_group_coordinators():
|
||||
if group.restore_hccl():
|
||||
num_restored += 1
|
||||
return num_restored
|
||||
|
||||
@staticmethod
|
||||
def refresh_moe_hccl_groups() -> None:
|
||||
from vllm_ascend.ops.fused_moe.moe_comm_method import _MoECommMethods
|
||||
|
||||
for comm_method in _MoECommMethods.values():
|
||||
dispatcher = getattr(comm_method, "token_dispatcher", None)
|
||||
refresh_fn = getattr(dispatcher, "refresh_hccl_group", None)
|
||||
if callable(refresh_fn):
|
||||
refresh_fn()
|
||||
|
||||
def sleep(self) -> None:
|
||||
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
||||
for handle in getattr(self.worker, "_pp_send_work", []):
|
||||
handle.wait()
|
||||
self.worker._pp_send_work = []
|
||||
torch.npu.synchronize()
|
||||
num_destroyed = self.destroy_hccl()
|
||||
if num_destroyed > 0:
|
||||
logger.info("Destroyed %d HCCL process groups for sleep mode.", num_destroyed)
|
||||
|
||||
def wakeup(self) -> None:
|
||||
with set_current_vllm_config(self.vllm_config):
|
||||
num_restored = self.restore_hccl()
|
||||
self.refresh_moe_hccl_groups()
|
||||
logger.info("Restored %d HCCL process groups after sleep mode.", num_restored)
|
||||
Reference in New Issue
Block a user