import unittest
import time
from llm_datadist import (
CacheDesc,
DataType,
LLMClusterInfo,
LLMDataDist,
LLMRole,
LLMStatusCode,
LlmConfig,
MemInfo,
Memtype,
Placement,
)
class LlmLocalCommResSt(unittest.TestCase):
_LOCAL_IP = "127.0.0.1"
_LISTEN_PORT = 26008
_REMOTE_CLUSTER_ID = 1
def setUp(self) -> None:
print(f"Begin {self.__class__.__name__}.{self._testMethodName}")
config = LlmConfig()
config.device_id = 0
config.rdma_service_level = 100
config.rdma_traffic_class = 100
config.listen_ip_info = f"{self._LOCAL_IP}:{self._LISTEN_PORT}"
config.local_comm_res = """
{
"server_count": "1",
"server_list": [{
"device": [{
"device_id": "0",
"device_ip": "1.1.1.1"
}],
"server_id": "127.0.0.1"
}],
"status": "completed",
"version": "1.0"
}
"""
engine_options = config.generate_options()
self.llm_datadist = LLMDataDist(LLMRole.PROMPT, 1)
self.llm_datadist.init(engine_options)
time.sleep(1)
self._cleanup_cluster_link()
self.has_exception = False
def tearDown(self) -> None:
print(f"End {self.__class__.__name__}.{self._testMethodName}")
self._cleanup_cluster_link()
self.llm_datadist.finalize()
def create_link_cluster(self):
ret, rets = self.llm_datadist.link_clusters([self._make_cluster_info()], 5000)
self.assertEqual(ret, LLMStatusCode.LLM_SUCCESS)
def test_local_comm_res_not_init(self):
cluster = self._make_cluster_info()
llm_datadist = LLMDataDist(LLMRole.PROMPT, 2)
try:
ret, rets = llm_datadist.link_clusters([cluster], 5000)
except Exception as e:
print(f"{type(e).__name__} - {str(e)}")
import traceback
print(traceback.format_exc())
self.has_exception = True
self.assertEqual(self.has_exception, True)
def test_local_comm_res_register_cache(self):
cache_mgr = self.llm_datadist.cache_manager
cache_desc = CacheDesc(1, [2, 4], DataType.DT_INT8, Placement.DEVICE)
cache = cache_mgr.register_cache(cache_desc, [1])
print(cache.cache_desc)
print(cache.tensor_addrs)
self.assertEqual(cache.cache_id, 1)
try:
cache_mgr.unregister_cache(1)
except Exception as e:
print(f"{type(e).__name__} - {str(e)}")
import traceback
print(traceback.format_exc())
self.has_exception = True
self.assertEqual(self.has_exception, False)
def test_local_comm_res_switch_role(self):
try:
self.llm_datadist.switch_role(LLMRole.DECODER)
options = {"llm.listenIpInfo": "127.0.0.1:26008"}
self.llm_datadist.switch_role(LLMRole.PROMPT, options)
options = {"llm.listenIpInfo": "127.0.0.1:26009"}
self.llm_datadist.switch_role(LLMRole.PROMPT, options)
except Exception as e:
print(f"{type(e).__name__} - {str(e)}")
import traceback
print(traceback.format_exc())
self.has_exception = True
self.assertEqual(self.has_exception, False)
def test_remap_registered_memory(self):
self.create_link_cluster()
cache_mgr = self.llm_datadist.cache_manager
mem_info = MemInfo(Memtype.MEM_TYPE_DEVICE, 1234, 1)
mem_infos = [mem_info]
print(f"mem_info={mem_info}")
try:
cache_mgr.remap_registered_memory(mem_info)
cache_mgr.remap_registered_memory(mem_infos)
except Exception as e:
print(f"{type(e).__name__} - {str(e)}")
import traceback
print(traceback.format_exc())
self.has_exception = True
self.assertEqual(self.has_exception, False)
def test_unlink_cluster(self):
self.create_link_cluster()
ret, rets = self.llm_datadist.unlink_clusters([self._make_cluster_info()], 5000)
self.assertEqual(ret, LLMStatusCode.LLM_SUCCESS)
def _make_cluster_info(self):
cluster = LLMClusterInfo()
cluster.remote_cluster_id = self._REMOTE_CLUSTER_ID
cluster.append_local_ip_info(self._LOCAL_IP, self._LISTEN_PORT)
cluster.append_remote_ip_info(self._LOCAL_IP, self._LISTEN_PORT)
return cluster
def _cleanup_cluster_link(self):
self.llm_datadist.unlink_clusters([self._make_cluster_info()], 5000)