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
7 changes: 2 additions & 5 deletions lightllm/common/kv_cache_mem_manager/allocator.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
import torch
from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt
from lightllm.utils.dist_utils import get_current_rank_in_node
from lightllm.utils.envs_utils import get_unique_server_name
from lightllm.utils.log_utils import init_logger
from typing import Union, List

Expand All @@ -25,10 +24,8 @@ def __init__(self, size: int) -> None:
self.can_use_mem_size = self.size

rank_in_node = get_current_rank_in_node()
# 用共享内存进行共享,router 模块读取进行精确的调度估计, nccl port 作为一个单机中单实列的标记。防止冲突。
self.shared_can_use_token_num = SharedInt(
f"{get_unique_server_name()}_mem_manger_can_use_token_num_{rank_in_node}"
)
# 用共享内存进行共享,router 模块读取进行精确的调度估计;基础层会统一添加服务前缀以防止实例冲突。
self.shared_can_use_token_num = SharedInt(f"mem_manger_can_use_token_num_{rank_in_node}")
self.shared_can_use_token_num.set_value(self.can_use_mem_size)
return

Expand Down
13 changes: 7 additions & 6 deletions lightllm/common/kv_cache_mem_manager/mem_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,8 @@
get_current_rank_in_node,
get_node_world_size,
)
from lightllm.utils.envs_utils import get_unique_server_name, get_env_start_args
from lightllm.utils.envs_utils import get_env_start_args
from lightllm.utils.shm_utils import get_service_shm_name
from lightllm.utils.config_utils import get_num_key_value_heads
from lightllm.common.kv_trans_kernel.nixl_kv_trans import page_io
from lightllm.utils.device_utils import kv_trans_use_p2p
Expand Down Expand Up @@ -237,10 +238,10 @@ def write_to_shm(self, req_manager):
# 避免过多无用的数据复制和传输开销。
self.req_to_token_indexs: torch.Tensor = req_manager.req_to_token_indexs

lock = FileLock(f"/tmp/{get_unique_server_name()}_mem_manager_lock")
lock = FileLock(f"/tmp/{get_service_shm_name('mem_manager_lock')}")
with lock:
node_world_size = get_node_world_size()
shm_name = f"{get_unique_server_name()}_mem_manager_{get_current_rank_in_node()}"
shm_name = f"mem_manager_{get_current_rank_in_node()}"
obj_bytes_array = [ForkingPickler.dumps(self).tobytes() for _ in range(node_world_size * 2)]
obj_size = len(obj_bytes_array[0])
shm = create_or_link_shm(
Expand All @@ -256,8 +257,8 @@ def write_to_shm(self, req_manager):

@staticmethod
def loads_from_shm(rank_in_node: int) -> "MemoryManager":
shm_name = f"{get_unique_server_name()}_mem_manager_{rank_in_node}"
lock = FileLock(f"/tmp/{get_unique_server_name()}_mem_manager_lock")
shm_name = f"mem_manager_{rank_in_node}"
lock = FileLock(f"/tmp/{get_service_shm_name('mem_manager_lock')}")
logger.info(f"get memmanager from shm {shm_name}")
with lock:
shm = create_or_link_shm(name=shm_name, expected_size=-1, force_mode="link")
Expand Down Expand Up @@ -285,7 +286,7 @@ def __init__(self) -> None:
# 兼容多机 dp size=1 纯 tp 模式的情况
self.is_multinode_tp = args.dp == 1 and args.nnodes > 1
self.shared_tp_infos = [
SharedInt(f"{get_unique_server_name()}_mem_manger_can_use_token_num_{rank_in_node}")
SharedInt(f"mem_manger_can_use_token_num_{rank_in_node}")
for rank_in_node in range(0, self.node_world_size, self.dp_world_size)
]

Expand Down
2 changes: 1 addition & 1 deletion lightllm/server/api_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,7 +111,7 @@ def set_args(self, args: StartArgs):
self.metric_client = MetricClient(get_shm_port_args().metric_port)
self.httpserver_manager = HttpServerManager(args=args)
dp_size_in_node = max(1, args.dp // args.nnodes) # 兼容多机纯tp的运行模式,这时候 1 // 2 == 0, 需要兼容
self.shared_token_load = TokenLoad(f"{get_unique_server_name()}_shared_token_load", dp_size_in_node)
self.shared_token_load = TokenLoad("shared_token_load", dp_size_in_node)


g_objs = G_Objs()
Expand Down
4 changes: 4 additions & 0 deletions lightllm/server/api_start.py
Original file line number Diff line number Diff line change
Expand Up @@ -338,6 +338,7 @@ def _launch_subprocesses(args: StartArgs):
validate_ports(ports_to_check)

set_env_start_args(args)
process_manager.setup_exit_controller()
get_shm_port_args(create=True)
# 多机用于收发node ip, 这个地方修改了args env,所以需要重新设置一下。
send_and_receive_node_ip(args)
Expand Down Expand Up @@ -480,6 +481,7 @@ def pd_master_start(args: StartArgs):

validate_ports([args.port])
set_env_start_args(args)
process_manager.setup_exit_controller()
get_shm_port_args(create=True)
logger.info(f"all start args:{args}")

Expand Down Expand Up @@ -545,6 +547,7 @@ def visual_only_start(args):
ports_to_check.append(args.visual_rpyc_port)
validate_ports(ports_to_check)
set_env_start_args(args)
process_manager.setup_exit_controller()
get_shm_port_args(create=True)
logger.info(f"all start args:{args}")

Expand Down Expand Up @@ -572,6 +575,7 @@ def config_server_start(args):
ports_to_check.append(args.config_server_visual_redis_port)
validate_ports(ports_to_check)
set_env_start_args(args)
process_manager.setup_exit_controller()
get_shm_port_args(create=True)
logger.info(f"all start args:{args}")

Expand Down
13 changes: 4 additions & 9 deletions lightllm/server/core/objs/req.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@
from .shm_array import ShmArray
from .token_chunck_hash_list import TokenHashList, CpuCachePageList, TokenPageLenList
from lightllm.server.req_id_generator import convert_sub_id_to_group_id
from lightllm.utils.envs_utils import get_unique_server_name
from lightllm.utils.envs_utils import get_env_start_args
from lightllm.utils.config_utils import is_linear_att_mixed_model
from lightllm.utils.kv_cache_utils import compute_token_list_hash
Expand Down Expand Up @@ -280,22 +279,19 @@ def _fill_linear_att_token_hash(self):
return

def create_prompt_ids_shm_array(self):
service_uni_name = get_unique_server_name()
name = f"{service_uni_name}_shm_prompts_{self.index_in_shm_mem}"
name = f"shm_prompts_{self.index_in_shm_mem}"
self.shm_prompt_ids = ShmArray(name, (self.alloc_shm_numpy_len,), dtype=np.int64)
self.shm_prompt_ids.create_shm()
return

def link_prompt_ids_shm_array(self):
service_uni_name = get_unique_server_name()
name = f"{service_uni_name}_shm_prompts_{self.index_in_shm_mem}"
name = f"shm_prompts_{self.index_in_shm_mem}"
self.shm_prompt_ids = ShmArray(name, (self.alloc_shm_numpy_len,), dtype=np.int64)
self.shm_prompt_ids.link_shm()
return

def create_logprobs_shm_array(self):
service_uni_name = get_unique_server_name()
name = f"{service_uni_name}_shm_logprobs_{self.index_in_shm_mem}"
name = f"shm_logprobs_{self.index_in_shm_mem}"
self.shm_logprobs = ShmArray(
name,
(self.alloc_shm_numpy_len,),
Expand All @@ -308,8 +304,7 @@ def create_logprobs_shm_array(self):
return

def link_logprobs_shm_array(self):
service_uni_name = get_unique_server_name()
name = f"{service_uni_name}_shm_logprobs_{self.index_in_shm_mem}"
name = f"shm_logprobs_{self.index_in_shm_mem}"
self.shm_logprobs = ShmArray(
name,
(self.alloc_shm_numpy_len,),
Expand Down
5 changes: 2 additions & 3 deletions lightllm/server/core/objs/shm_objs_io_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@
import pickle
from lightllm.server.core.objs.atomic_lock import AtomicShmLock
from lightllm.utils.envs_utils import get_env_start_args
from lightllm.utils.envs_utils import get_unique_server_name
from lightllm.utils.log_utils import init_logger
from lightllm.utils.shm_utils import create_or_link_shm

Expand All @@ -14,8 +13,8 @@
class ShmObjsIOBuffer:
def __init__(self, tail_str=""):
self.args = get_env_start_args()
self.name = f"{get_unique_server_name()}_ShmReqsBufferParams_{tail_str}"
self.lock = AtomicShmLock(lock_name=f"{get_unique_server_name()}_ShmReqsBufferParams_atomlock_{tail_str}")
self.name = f"ShmReqsBufferParams_{tail_str}"
self.lock = AtomicShmLock(lock_name=f"ShmReqsBufferParams_atomlock_{tail_str}")
self._create_or_link_shm()
self.node_world_size = self.args.tp // self.args.nnodes

Expand Down
11 changes: 5 additions & 6 deletions lightllm/server/core/objs/shm_req_manager.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
import ctypes
import numpy as np
from lightllm.utils.envs_utils import get_unique_server_name
from multiprocessing import shared_memory
from lightllm.utils.log_utils import init_logger
from .req import Req, ChunkedPrefillReq
Expand Down Expand Up @@ -44,7 +43,7 @@ def init_reqs_shm(self):
self._init_reqs_shm()

def _init_reqs_shm(self):
shm_name = f"{get_unique_server_name()}_req_shm_total"
shm_name = "req_shm_total"
self.reqs_shm = create_or_link_shm(shm_name, self.req_shm_byte_size)
return

Expand All @@ -56,21 +55,21 @@ def init_to_req_objs(self):
return

def init_to_req_locks(self):
array_lock_name = f"{get_unique_server_name()}_array_reqs_lock"
array_lock_name = "array_reqs_lock"
self.reqs_lock = AtomicShmArrayLock(array_lock_name, self.max_req_num)
return

def get_req_lock_by_index(self, req_index_in_mem: int) -> AtomicLockItem:
return self.reqs_lock.get_lock_context(req_index_in_mem)

def init_manager_lock(self):
lock_name = f"{get_unique_server_name()}_shm_reqs_manager_lock"
lock_name = "shm_reqs_manager_lock"
self.manager_lock = AtomicShmLock(lock_name)
return

def init_alloc_state_shm(self):
shm_name = f"{get_unique_server_name()}_req_alloc_states"
req_link_list_name = f"{get_unique_server_name()}_req_linked_states"
shm_name = "req_alloc_states"
req_link_list_name = "req_linked_states"
self.linked_req_manager = ReqLinkedListManager(req_link_list_name, self.max_req_num)
self.alloc_state_shm = ShmArray(shm_name, (self.max_req_num,), np.int32)
self.alloc_state_shm.create_shm()
Expand Down
4 changes: 1 addition & 3 deletions lightllm/server/core/objs/token_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,6 @@

import numpy as np

from lightllm.utils.envs_utils import get_unique_server_name
from lightllm.utils.log_utils import init_logger
from lightllm.utils.shm_utils import create_or_link_shm

Expand Down Expand Up @@ -175,5 +174,4 @@ def _build_routed_experts_response(self, packed_routed: Optional[np.ndarray]) ->
}

def _shm_name(self) -> str:
service_uni_name = get_unique_server_name()
return f"{service_uni_name}_shm_final_token_metadata_{self.req.index_in_shm_mem}"
return f"shm_final_token_metadata_{self.req.index_in_shm_mem}"
7 changes: 6 additions & 1 deletion lightllm/server/embed_cache/utils.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,10 @@
import multiprocessing.shared_memory as shm

from lightllm.utils.shm_utils import get_service_shm_name


def create_shm(name, data):
name = get_service_shm_name(name)
try:
data_size = len(data)
shared_memory = shm.SharedMemory(name=name, create=True, size=data_size)
Expand All @@ -12,16 +15,18 @@ def create_shm(name, data):


def read_shm(name):
name = get_service_shm_name(name)
shared_memory = shm.SharedMemory(name=name)
data = shared_memory.buf.tobytes()
return data


def free_shm(name):
name = get_service_shm_name(name)
shared_memory = shm.SharedMemory(name=name)
shared_memory.close()
shared_memory.unlink()


def get_shm_name_data(uid):
return str(uid) + "-data"
return f"{uid}-data"
9 changes: 4 additions & 5 deletions lightllm/server/httpserver/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,6 @@
from .manager_ext import HttpRlManagerHelper
from lightllm.utils.statics_utils import MovingAverage
from lightllm.utils.config_utils import get_vocab_size
from lightllm.utils.envs_utils import get_unique_server_name
from lightllm.utils.shm_port_args import get_shm_port_args
from lightllm.utils.error_utils import (
ClientDisconnected,
Expand All @@ -62,7 +61,7 @@ def __init__(

self.multinode_req_manager = None
self.nnodes = args.nnodes
self._shm_lock_pool = AtomicShmArrayLock(f"{get_unique_server_name()}_lightllm_resource_lock", 2)
self._shm_lock_pool = AtomicShmArrayLock("lightllm_resource_lock", 2)
self._resource_lock = AsyncLock(self._shm_lock_pool.get_lock_context(0))
self._run_reqs_count_lock = AsyncLock(self._shm_lock_pool.get_lock_context(1))
self.node_rank = args.node_rank
Expand Down Expand Up @@ -129,19 +128,19 @@ def __init__(
self.vocab_size = max(get_vocab_size(args.model_dir), self.tokenizer.vocab_size)

# Timemark of the latest successful inference, used by passive /health checks.
self.latest_success_infer_time_mark = SharedInt(f"{get_unique_server_name()}_latest_success_infer_time_mark")
self.latest_success_infer_time_mark = SharedInt("latest_success_infer_time_mark")
self.latest_success_infer_time_mark.set_value(int(time.time()))

self.rl_controller: Optional[HttpRlController] = HttpRlController(self) if args.enable_rl else None

self.run_reqs_count_mark = SharedInt(f"{get_unique_server_name()}_run_reqs_count_mark")
self.run_reqs_count_mark = SharedInt("run_reqs_count_mark")
self.run_reqs_count_mark.set_value(0)

# 用于记录真实的--max_total_token_num 参数,当这个参数在启动参数中没有设置的时候,其是在推理进程中被分析出来的,
# 这个时候如果 --max_req_total_len > --max_total_token_num 时,如果httpserver放过一些非法的输入进入后续的模块可能
# 会触发整个系统崩溃,所以httpserver需要知道真实的 max_total_token_num的数据,用于提前拦截非法请求等参数。
# router 进程会在启动后向这个共享内存写入正确的max_total_token_num 参数,用于后续的请求控制。
self.shm_max_total_token_num = SharedInt(f"{get_unique_server_name()}_shm_max_total_token_num")
self.shm_max_total_token_num = SharedInt("shm_max_total_token_num")
return

def _log_stage_timing(self, group_request_id: int, start_time: float, stage: str, **kwargs):
Expand Down
10 changes: 5 additions & 5 deletions lightllm/server/multi_level_kv_cache/cpu_cache_client.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import ctypes
from lightllm.utils.envs_utils import get_env_start_args, get_unique_server_name, get_disk_cache_prompt_limit_length
from lightllm.utils.envs_utils import get_env_start_args, get_disk_cache_prompt_limit_length
from typing import List, Optional, Tuple
from lightllm.utils.log_utils import init_logger
from lightllm.common.cpu_cache import CpuCacheCreator, CpuCacheTensorSpec
Expand All @@ -20,7 +20,7 @@ def __init__(self, only_create_meta_data: bool, init_shm_data: bool):
# to do here need calcu from from settings.
self.kv_cache_tensor_meta = calcu_cpu_cache_meta()
self.page_num: int = self.kv_cache_tensor_meta.page_num
self.lock = AtomicShmLock(lock_name=f"{get_unique_server_name()}_cpu_kv_cache_client_lock")
self.lock = AtomicShmLock(lock_name="cpu_kv_cache_client_lock")
self._create_cpu_status_list(init_shm_data)

if not only_create_meta_data:
Expand Down Expand Up @@ -268,18 +268,18 @@ def recycle_pages(self, page_list: List[int]):

def _create_cpu_status_list(self, init_shm_data: bool):
self.page_items = ShmLinkedList(
name=f"{get_unique_server_name()}_cpu_kv_cache_page_items",
name="cpu_kv_cache_page_items",
item_class=_CpuPageStatus,
capacity=self.page_num,
init_shm_data=init_shm_data,
)
self.page_hash_dict = ShmDict(
name=f"{get_unique_server_name()}_cpu_kv_cache_hash",
name="cpu_kv_cache_hash",
capacity=self.page_num * 2,
init_shm_data=init_shm_data,
)
self.offload_page_indexes = IntList(
name=f"{get_unique_server_name()}_cpu_kv_cache_offload_page_indexes",
name="cpu_kv_cache_offload_page_indexes",
capacity=self.page_num * 2,
init_shm_data=init_shm_data,
)
Expand Down
7 changes: 3 additions & 4 deletions lightllm/server/multi_level_kv_cache/shm_objs.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from multiprocessing import shared_memory
from typing import List, Optional
from lightllm.utils.log_utils import init_logger
from lightllm.utils.auto_shm_cleanup import register_posix_shm_for_cleanup
from lightllm.utils.shm_utils import get_service_shm_name

logger = init_logger(__name__)

Expand Down Expand Up @@ -290,11 +290,10 @@ def key(self, value: int):
self.key_high = (value >> 64) & 0xFFFFFFFFFFFFFFFF


def _create_shm(name: str, byte_size: int, auto_cleanup: bool = False):
def _create_shm(name: str, byte_size: int):
name = get_service_shm_name(name)
try:
shm = shared_memory.SharedMemory(name=name, create=True, size=byte_size)
if auto_cleanup:
register_posix_shm_for_cleanup(name)
logger.info(f"create lock shm {name}")
except:
shm = shared_memory.SharedMemory(name=name, create=False, size=byte_size)
Expand Down
11 changes: 4 additions & 7 deletions lightllm/server/req_id_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,30 +19,27 @@ class ReqIDGenerator:
def __init__(self):
from lightllm.server.core.objs.atomic_lock import AtomicShmLock
from lightllm.server.core.objs.shm_array import ShmArray
from lightllm.utils.envs_utils import get_unique_server_name, get_env_start_args
from lightllm.utils.envs_utils import get_env_start_args

self.args = get_env_start_args()
self.use_config_server = (
self.args.config_server_host and self.args.config_server_port and self.args.run_mode == "pd_master"
)
self.current_id = ShmArray(f"{get_unique_server_name()}_req_id_gen", (2,), dtype=np.int64)
self.current_id = ShmArray("req_id_gen", (2,), dtype=np.int64)
self.current_id.create_shm()
self.current_id.arr[0] = 0
self.current_id.arr[1] = 0
self.lock = AtomicShmLock(f"{get_unique_server_name()}_req_id_gen_lock")
self.lock = AtomicShmLock("req_id_gen_lock")
self._wait_all_workers_ready()
logger.info("ReqIDGenerator init finished")

def _wait_all_workers_ready(self):
if self.args.httpserver_workers == 1:
return

from lightllm.utils.envs_utils import get_unique_server_name
from lightllm.server.core.objs.shm_array import ShmArray

_sync_shm = ShmArray(
f"{get_unique_server_name()}_httpworker_start_sync", (self.args.httpserver_workers,), dtype=np.int64
)
_sync_shm = ShmArray("httpworker_start_sync", (self.args.httpserver_workers,), dtype=np.int64)
_sync_shm.create_shm()
# 等待所有 httpserver 的 worker 启动完成,防止重新初始化对应的请求id 对应的shm
try_count = 0
Expand Down
Loading
Loading