Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 2 additions & 9 deletions lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,6 @@


class Deepseek2MemoryManager(MemoryManager):

operator_class = Deepseek2MemOperator

def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False, mem_fraction=0.9):
Expand All @@ -28,14 +27,8 @@ def get_cell_size(self):
def _init_buffers(self, size, dtype, head_num, head_dim, layer_num):
self.kv_buffer = torch.empty((layer_num, size + 1, head_num, head_dim), dtype=dtype, device="cuda")

def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor:
self.kv_move_buffer = torch.empty(
(page_num, page_size, self.layer_num, self.head_num, self.head_dim), dtype=self.dtype, device="cuda"
)
self._buffer_mem_indexes_tensors = [
torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) for _ in range(page_num)
]
return self.kv_move_buffer
def get_paged_kv_move_buffer_shape(self, page_num, page_size):
return (page_num, page_size, self.layer_num, self.head_num, self.head_dim)

def write_mem_to_page_kv_move_buffer(
self,
Expand Down
35 changes: 31 additions & 4 deletions lightllm/common/kv_cache_mem_manager/mem_manager.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import re
import os
import math
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
Expand All @@ -23,7 +24,6 @@


class MemoryManager:

operator_class = NormalMemOperator

def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False, mem_fraction=0.9):
Expand Down Expand Up @@ -58,21 +58,49 @@ def get_att_input_params(self, layer_index: int) -> Tuple[Any, Any]:
def get_cell_size(self):
return 2 * self.head_num * self.head_dim * self.layer_num * torch._utils._element_size(self.dtype)

def get_paged_kv_move_buffer_shape(self, page_num, page_size):
num_kv_head = get_num_key_value_heads(get_env_start_args().model_dir)
return (page_num, page_size, self.layer_num, 2 * num_kv_head, self.head_dim)

def get_paged_kv_move_buffer_size_in_bytes(self, page_num, page_size):
shape = self.get_paged_kv_move_buffer_shape(page_num, page_size)
return math.prod(shape) * torch._utils._element_size(self.dtype)

def profile_size(self, mem_fraction):
if self.size is not None:
return

torch.cuda.empty_cache()
world_size = dist.get_world_size()
available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction)
args = get_env_start_args()
pd_kv_move_buffer_size_in_bytes = 0
if args.run_mode in ["prefill", "decode"]:
pd_kv_move_buffer_size_in_bytes = self.get_paged_kv_move_buffer_size_in_bytes(
page_num=args.pd_kv_page_num,
page_size=args.pd_kv_page_size,
)
available_memory -= pd_kv_move_buffer_size_in_bytes / 1024 ** 3
cell_size = self.get_cell_size()
self.size = int(available_memory * 1024 ** 3 / cell_size)
if world_size > 1:
tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}")
dist.all_reduce(tensor, op=dist.ReduceOp.MIN)
self.size = tensor.item()
if pd_kv_move_buffer_size_in_bytes > 0 and self.size <= 0:
raise RuntimeError(
"PD KV transfer page buffer reservation leaves no memory for the token KV cache; "
"reduce --pd_kv_page_size or --pd_kv_page_num, or increase --mem_fraction"
)
pd_kv_move_buffer_log = ""
if pd_kv_move_buffer_size_in_bytes > 0:
pd_kv_move_buffer_log = (
f"{str(pd_kv_move_buffer_size_in_bytes / 1024 ** 3)} GB is reserved "
"for the PD KV transfer page buffer\n"
)
logger.info(
f"{str(available_memory)} GB space is available after load the model weight\n"
f"{str(available_memory)} GB space is available for the token KV cache\n"
f"{pd_kv_move_buffer_log}"
f"{str(cell_size / 1024 ** 2)} MB is the size of one token kv cache\n"
f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n"
)
Expand All @@ -86,9 +114,8 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num):
self.kv_buffer = torch.empty((layer_num, size + 1, 2 * head_num, head_dim), dtype=dtype, device="cuda")

def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor:
num_kv_head = get_num_key_value_heads(get_env_start_args().model_dir)
self.kv_move_buffer = torch.empty(
(page_num, page_size, self.layer_num, 2 * num_kv_head, self.head_dim), dtype=self.dtype, device="cuda"
self.get_paged_kv_move_buffer_shape(page_num, page_size), dtype=self.dtype, device="cuda"
)
self._buffer_mem_indexes_tensors = [
torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) for _ in range(page_num)
Expand Down
Loading