Coverage for /opt/cloud/slavespace/usr1/096471637100f3de0fcfc01072822a80/dttest/api/python/llm_datadist_v1/kv_cache_manager.py: 17%
202 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-27 10:02 +0800
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-27 10:02 +0800
1#!/usr/bin/env python3
2# -*- coding: utf-8 -*-
3# -----------------------------------------------------------------------------------------------------------
4# Copyright (c) 2025 Huawei Technologies Co., Ltd.
5# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
6# CANN Open Software License Agreement Version 2.0 (the "License").
7# Please refer to the License for details. You may not use this file except in compliance with the License.
8# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
9# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
10# See LICENSE in the root of the software repository for the full text of the License.
11# -----------------------------------------------------------------------------------------------------------
13from typing import Dict, List, Optional, Tuple, Union
15from .configs import LLMRole
16from .llm_types import (
17 BlocksCacheKey,
18 CacheDesc,
19 CacheKey,
20 CacheKeyByIdAndIndex,
21 CacheTask,
22 KvCache,
23 LayerSynchronizer,
24 Placement,
25 TransferConfig,
26)
27from .llm_utils import (
28 TransferCacheParameters,
29 is_invalid_id,
30 is_valid_id,
31 layer_range_to_tensor_indices,
32 pack_block_cache_key,
33 pack_cache_desc,
34 pack_cache_key,
35 pack_cache_key_by_id,
36 transfer_cache_async,
37)
38from .status import LLMStatusCode, handle_llm_status, raise_if_false, raise_if_true
39from .tensor import Tensor
40from .utils import log
41from .utils.utils import (
42 check_dict,
43 check_int32,
44 check_int64,
45 check_isinstance,
46 check_positive_or_set_default,
47 check_type,
48 check_uint32,
49 check_uint64,
50)
52_NUM_TENSORS_PER_LAYER = 2
53_INVALID_ID = 2**64 - 1
56class KvCacheManager(object):
57 """
58 提供了一组KvCache的操作函数, 在LLMEngine初始化后通过LLMEngine获取实例
60 Examples:
61 >>> from llm_datadist import LLMRole, LLMDataDist
62 >>> # init LLMDataDist
63 >>> cluster_id = 0
64 >>> datadist = LLMDataDist(LLMRole.PROMPT, cluster_id)
65 >>> engine_options = {}
66 >>> datadist.init(engine_options)
67 >>> # get KvCacheManager
68 >>> kv_cache_manager = datadist.kv_cache_manager
69 """
71 def __init__(self, llm_engine, role: LLMRole) -> None:
72 self._llm_engine = llm_engine
73 self._role = role
74 self._initialized = True
76 def is_initialized(self) -> bool:
77 return self._initialized
79 def allocate_blocks_cache(
80 self, cache_desc: CacheDesc, blocks_cache_key: Optional[BlocksCacheKey] = None
81 ) -> KvCache:
82 """
83 分配Blocks Cache, cache分配成功后
84 需通过deallocate_cache释放
86 Args:
87 cache_desc: Cache描述
88 blocks_cache_key(Optional): 仅当LLMRole为PROMPT时可设置, 用于在DECODER拉取KV
89 Returns:
90 KvCache
92 Examples:
93 >>> from llm_datadist import LLMRole, CacheDesc, DataType, LLMDataDist
94 >>> # init LLMDataDist
95 >>> cluster_id = 0
96 >>> datadist = LLMDataDist(LLMRole.PROMPT, cluster_id)
97 >>> engine_options = {} # 按需填写
98 >>> datadist.init(engine_options)
99 >>> # get KvCacheManager
100 >>> kv_cache_manager = datadist.kv_cache_manager
101 >>> # allocate_cache
102 >>> kv_cache_desc = CacheDesc(num_tensors=80, shape=[4, 256], data_type=DataType.DT_FLOAT16)
103 >>> # case 1: batch_size = 4, 只有第二个有效, 此时后两个cache_key可以省略
104 >>> block_cache_key = BlocksCacheKey(prompt_cluster_id=0, model_id=0)
105 >>> cache = kv_cache_manager.allocate_blocks_cache(kv_cache_desc, block_cache_key)
106 >>> kv_cache_manager.deallocate_cache(cache)
107 """
108 check_isinstance("cache_desc", cache_desc, CacheDesc)
109 check_isinstance("blocks_cache_key", blocks_cache_key, BlocksCacheKey)
110 raise_if_false(cache_desc.num_tensors > 0, "num_tensors should be bigger than zero.")
111 raise_if_false(
112 cache_desc.placement == Placement.DEVICE,
113 "Only support allocate device cache",
114 )
115 if self._role == LLMRole.DECODER:
116 raise_if_false(blocks_cache_key is None, "blocks_cache_key is not supported by DECODER")
117 wrapped_cache_keys = [pack_block_cache_key(blocks_cache_key)] if blocks_cache_key is not None else []
118 ret, cache_id_and_addr = self._llm_engine.allocate_cache(pack_cache_desc(cache_desc), wrapped_cache_keys)
119 handle_llm_status(ret, "[allocate_blocks_cache]", f"cache_desc = {cache_desc}")
120 kv_cache = KvCache(cache_id_and_addr[0], cache_desc, cache_id_and_addr[1], self)
121 log.info("[allocate_blocks_cache] success, cache_id = %d", kv_cache.cache_id)
122 return kv_cache
124 def allocate_cache(
125 self,
126 cache_desc: CacheDesc,
127 cache_keys: Union[Tuple[CacheKey], List[CacheKey]] = (),
128 ) -> KvCache:
129 """
130 分配Cache, cache分配成功后, 会同时被cache_id与cache_keys(如果传了)引用, 只有当这些引用都解除后, cache所占用的资源才会实际释放
131 cache_id的引用需通过deallocate_cache解除
132 cache_keys的引用则可以通过以下2种方式解除:
133 1. DECODER调用pull_kv接口, pull_kv成功后解除
134 2. PROMPT调用remove_cache_key接口
136 Args:
137 cache_desc: Cache描述
138 cache_keys(Optional): 仅当LLMRole为PROMPT时可设置, 用于在DECODER拉取KV
139 如果Cache的batch size > 1, 则需要提供相同数量的CacheKey, 分别引用一组kv tensor
140 如果当次推理的batch未占用满,即存在无效batch index,则需要插入特殊的CacheKey(req_id = UINT64_MAX)占位,
141 如果空闲的batch_index在末尾,则可以省略
142 Returns:
143 KvCache
145 Examples:
146 >>> from llm_datadist import LLMRole, CacheDesc, DataType, LLMDataDist
147 >>> # init LLMDataDist
148 >>> cluster_id = 0
149 >>> datadist = LLMDataDist(LLMRole.PROMPT, cluster_id)
150 >>> engine_options = {} # 按需填写
151 >>> datadist.init(engine_options)
152 >>> # get KvCacheManager
153 >>> kv_cache_manager = datadist.kv_cache_manager
154 >>> # allocate_cache
155 >>> kv_cache_desc = CacheDesc(num_tensors=80, shape=[4, 256], data_type=DataType.DT_FLOAT16)
156 >>> # case 1: batch_size = 4, 只有第二个有效, 此时后两个cache_key可以省略
157 >>> kv_cache_key_1 = CacheKey(prompt_cluster_id=0, req_id=1, model_id=0)
158 >>> padding_cache_key = CacheKey(prompt_cluster_id=0, req_id=2 ** 64 - 1, model_id=0)
159 >>> kv_cache_keys = [padding_cache_key, kv_cache_key_1]
160 >>> cache = kv_cache_manager.allocate_cache(kv_cache_desc, kv_cache_keys)
161 >>> # case 2: batch_size = 4, 只有最后一个有效, 此时所有cache_key都不能省略
162 >>> kv_cache_key_3 = CacheKey(prompt_cluster_id=0, req_id=3, model_id=0)
163 >>> kv_cache_keys = [padding_cache_key, padding_cache_key, padding_cache_key, kv_cache_key_3]
164 >>> cache_1 = kv_cache_manager.allocate_cache(kv_cache_desc, kv_cache_keys)
165 >>> # 释放cache_id对cache的引用
166 >>> kv_cache_manager.deallocate_cache(cache)
167 >>> kv_cache_manager.deallocate_cache(cache_1)
168 >>> # 释放cache_key对cache的引用
169 >>> kv_cache_manager.remove_cache_key(kv_cache_key_1)
170 >>> kv_cache_manager.remove_cache_key(kv_cache_key_3)
171 """
172 check_isinstance("cache_desc", cache_desc, CacheDesc)
173 check_isinstance("cache_keys", cache_keys, [list, tuple], CacheKey)
174 raise_if_false(cache_desc.num_tensors > 0, "num_tensors should be bigger than zero.")
175 raise_if_false(
176 cache_desc.placement == Placement.DEVICE,
177 "Only support allocate device cache",
178 )
179 log.info(
180 "[allocate_cache] start, cache_desc = %s, cache_keys = %s",
181 cache_desc,
182 cache_keys,
183 )
184 if self._role != LLMRole.PROMPT:
185 raise_if_false(
186 len(cache_keys) == 0,
187 "cache_keys is not supported by {0}",
188 self._role.name,
189 )
190 wrapped_cache_keys = [pack_cache_key(cache_key) for cache_key in cache_keys]
191 ret, cache_id_and_addr = self._llm_engine.allocate_cache(pack_cache_desc(cache_desc), wrapped_cache_keys)
192 handle_llm_status(ret, "[allocate_cache]", f"cache_desc = {cache_desc}")
193 kv_cache = KvCache(cache_id_and_addr[0], cache_desc, cache_id_and_addr[1], self)
194 log.info("[allocate_cache] success, cache_id = %d", kv_cache.cache_id)
195 return kv_cache
197 def deallocate_cache(self, cache: KvCache) -> None:
198 """
199 释放Cache, 如果该Cache在Allocate时关联了CacheKey, 则实际的释放会延后到所有的CacheKey被拉取或执行了remove_cache_key
200 释放之后,不应再对该KvCache做任何操作
202 Args:
203 cache: Cache
205 Examples:
206 see examples of allocate_cache
207 """
208 check_isinstance("cache", cache, KvCache)
209 raise_if_true(self._is_cpu_cache(cache), "Only support device cache")
210 log.info("[deallocate_cache] start, cache_id = %d", cache.cache_id)
211 ret = self._llm_engine.deallocate_cache(cache.cache_id)
212 handle_llm_status(ret, "[deallocate_cache]", f"cache_id = {cache.cache_id}")
213 cache._per_device_tensor_addrs = []
214 cache._valid = False
215 log.info("[deallocate_cache] success")
217 def remove_cache_key(self, cache_key: CacheKey) -> None:
218 """
219 移除CacheKey, 仅当LLMRole为PROMPT时可调用
220 移除CacheKey后, 该Cache将无法再被pull_cache拉取
222 Args:
223 cache_key: CacheKey
225 Examples:
226 see examples of allocate_cache
227 """
228 self._check_role("[remove_cache_key]", LLMRole.PROMPT)
229 check_isinstance("cache_key", cache_key, CacheKey)
230 log.info("[remove_cache_key] start, cache_key = %s", cache_key)
231 ret = self._llm_engine.remove_cache_key(pack_cache_key(cache_key))
232 handle_llm_status(ret, "[remove_cache_key]", f"cache_key = {cache_key}")
233 log.info("[remove_cache_key] success")
235 def pull_blocks(
236 self,
237 prompt_cache_key: BlocksCacheKey,
238 decoder_kv_cache: KvCache,
239 prompt_blocks: List[int],
240 decoder_blocks: List[int],
241 **kwargs,
242 ):
243 """
244 PA模式下拉取KV
245 Args:
246 prompt_cache_key: prompt缓存key
247 decoder_kv_cache: decoder目标缓存
248 prompt_blocks: prompt block列表
249 decoder_blocks: decoder block列表
250 **kwargs:
251 src_layer_range: 源层范围
252 dst_layer_range: 目标层范围
253 tensor_num_per_layer: 每层tensor数量
254 """
255 src_layer_range = kwargs.get("src_layer_range")
256 dst_layer_range = kwargs.get("dst_layer_range")
257 tensor_num_per_layer = kwargs.get("tensor_num_per_layer", _NUM_TENSORS_PER_LAYER)
258 self._check_role("[pull_blocks]", LLMRole.DECODER)
259 check_isinstance("prompt_cache_key", prompt_cache_key, BlocksCacheKey)
260 check_isinstance("decoder_kv_cache", decoder_kv_cache, KvCache)
261 check_isinstance("prompt_blocks", prompt_blocks, list, int)
262 check_isinstance("decoder_blocks", decoder_blocks, list, int)
263 raise_if_true(self._is_cpu_cache(decoder_kv_cache), "Only support device cache")
264 raise_if_false(len(prompt_blocks) > 0, "prompt_blocks can not be empty.")
265 raise_if_false(len(decoder_blocks) > 0, "decoder_blocks can not be empty.")
266 raise_if_false(
267 len(prompt_blocks) == len(decoder_blocks),
268 "Param prompt_blocks and decoder_blocks size should be same.",
269 )
270 check_uint32("tensor_num_per_layer", tensor_num_per_layer)
271 raise_if_false(
272 tensor_num_per_layer > 0,
273 "[pull_blocks] param check failed, tensor_num_per_layer ({0}) is invalid, should [1, {1}]",
274 tensor_num_per_layer,
275 decoder_kv_cache.cache_desc.num_tensors,
276 )
278 log.info(
279 "[pull_blocks] start, target cache_id = %d, cache_key = %s, "
280 "src_layer_range = %s, dst_layer_range = %s, tensor_num_per_layer = %d",
281 decoder_kv_cache.cache_id,
282 prompt_cache_key,
283 src_layer_range,
284 dst_layer_range,
285 tensor_num_per_layer,
286 )
287 src_tensor_indices, dst_tensor_indices = layer_range_to_tensor_indices(
288 src_layer_range, dst_layer_range, tensor_num_per_layer
289 )
290 param = (
291 -1,
292 0,
293 prompt_blocks,
294 decoder_blocks,
295 src_tensor_indices,
296 dst_tensor_indices,
297 -1,
298 -1,
299 tensor_num_per_layer,
300 )
301 ret = self._llm_engine.pull_cache(decoder_kv_cache.cache_id, pack_block_cache_key(prompt_cache_key), param)
302 handle_llm_status(ret, "[pull_blocks]", f"prompt_cache_key = {prompt_cache_key}")
303 log.info("[pull_blocks] success")
305 def pull_cache(
306 self,
307 cache_key: Union[CacheKey, CacheKeyByIdAndIndex],
308 kv_cache: KvCache,
309 batch_index: int = 0,
310 size: int = -1,
311 **kwargs,
312 ) -> None:
313 """
314 拉取KV, 仅当LLMRole为DECODER时可调用
316 Args:
317 cache_key: CacheKey或CacheKeyByIdAndIndex
318 kv_cache: 目标KvCache
319 batch_index: batch index
320 size: 拉取的tensor大小, -1表示拉取全部大小
321 **kwargs:
322 src_layer_range: 源层范围
323 dst_layer_range: 目标层范围
324 src_cache_offset: 源cache偏移
325 dst_cache_offset: 目标cache偏移
326 tensor_num_per_layer: 每层tensor数量
328 Examples:
329 >>> from llm_datadist import LLMRole, CacheDesc, DataType, LLMDataDist
330 >>> # init LLMDataDist
331 >>> cluster_id = 0
332 >>> datadist = LLMDataDist(LLMRole.DECODER, cluster_id)
333 >>> engine_options = {} # 按需填写
334 >>> datadist.init(engine_options)
335 >>> # get KvCacheManager
336 >>> kv_cache_manager = datadist.kv_cache_manager
337 >>> # allocate_cache
338 >>> kv_cache_desc = CacheDesc(num_tensors=80, shape=[4, 256], data_type=DataType.DT_FLOAT16)
339 >>> cache = kv_cache_manager.allocate_cache(kv_cache_desc)
340 >>> # pull prompt kv to allocated cache
341 >>> prompt_cache_key = CacheKey(prompt_cluster_id=0, req_id = 1, model_id = 0)
342 >>> kv_cache_manager.pull_cache(prompt_cache_key, cache, 0)
343 """
344 src_layer_range = kwargs.get("src_layer_range")
345 dst_layer_range = kwargs.get("dst_layer_range")
346 src_cache_offset = kwargs.get("src_cache_offset")
347 dst_cache_offset = kwargs.get("dst_cache_offset")
348 tensor_num_per_layer = kwargs.get("tensor_num_per_layer", _NUM_TENSORS_PER_LAYER)
349 self._check_role("[pull_cache]", LLMRole.DECODER)
350 check_isinstance("cache_key", cache_key, [CacheKey, CacheKeyByIdAndIndex])
351 check_isinstance("kv_cache", kv_cache, KvCache)
352 raise_if_true(self._is_cpu_cache(kv_cache), "Only support device cache")
353 check_int64("size", size)
354 check_uint32("batch_index", batch_index)
355 raise_if_false(
356 size == -1 or size > 0,
357 "[pull_cache] param check failed, size ({0}) is invalid, should be = -1 or > 0",
358 size,
359 )
360 src_cache_offset = check_positive_or_set_default("src_cache_offset", src_cache_offset)
361 dst_cache_offset = check_positive_or_set_default("dst_cache_offset", dst_cache_offset)
362 if check_type(cache_key, CacheKey):
363 KvCacheManager.check_cache_key(cache_key)
364 packed_cache_key = pack_cache_key(cache_key)
365 else:
366 packed_cache_key = pack_cache_key_by_id(cache_key)
367 check_uint32("tensor_num_per_layer", tensor_num_per_layer)
368 raise_if_false(
369 tensor_num_per_layer > 0,
370 "[pull_cache] param check failed, tensor_num_per_layer ({0}) is invalid, should [1, {1}]",
371 tensor_num_per_layer,
372 kv_cache.cache_desc.num_tensors,
373 )
375 log.info(
376 "[pull_cache] start, cache_id = %d, batch_index = %d, size = %d, cache_key = %s, "
377 "src_layer_range = %s, dst_layer_range = %s, src_cache_offset = %d, dst_cache_offset = %d, tensor_num_per_layer = %d",
378 kv_cache.cache_id,
379 batch_index,
380 size,
381 cache_key,
382 src_layer_range,
383 dst_layer_range,
384 src_cache_offset,
385 dst_cache_offset,
386 tensor_num_per_layer,
387 )
388 src_tensor_indices, dst_tensor_indices = layer_range_to_tensor_indices(
389 src_layer_range, dst_layer_range, tensor_num_per_layer
390 )
392 param = (
393 size,
394 batch_index,
395 [],
396 [],
397 src_tensor_indices,
398 dst_tensor_indices,
399 src_cache_offset,
400 dst_cache_offset,
401 tensor_num_per_layer,
402 )
403 ret = self._llm_engine.pull_cache(kv_cache.cache_id, packed_cache_key, param)
404 handle_llm_status(ret, "[pull_cache]", f"cache_key = {cache_key}")
405 log.info("[pull_cache] success")
407 @staticmethod
408 def check_cache_key(cache_key: CacheKey) -> None:
409 raise_if_true(
410 is_invalid_id(cache_key.req_id) and is_invalid_id(cache_key.prefix_id),
411 f"one of req id and prefix id should contain valid value:[0, 2**64-1), "
412 f"req id:{cache_key.req_id},prefix id{cache_key.prefix_id}.",
413 )
414 raise_if_true(
415 is_valid_id(cache_key.req_id) and is_valid_id(cache_key.prefix_id),
416 "only one of req id and prefix id should contain valid value:[0, 2**64-1), "
417 f"req id:{cache_key.req_id}, prefix id{cache_key.prefix_id}.",
418 )
420 @staticmethod
421 def _verify_caches(src: KvCache, dst: KvCache, src_to_dst: Dict[int, int]):
422 check_isinstance("src", src, KvCache)
423 check_isinstance("dst", dst, KvCache)
424 check_isinstance("src_to_dst", src_to_dst, dict, int)
426 src_block_size = src.cache_desc.size // src.cache_desc.batch_size
427 dst_block_size = dst.cache_desc.size // dst.cache_desc.batch_size
428 raise_if_false(
429 src_block_size == dst_block_size,
430 f"src block size:{src_block_size} and dst block size:{dst_block_size} must be equal",
431 )
432 log.info("src and dst cache block size:%d", src_block_size)
434 src_num_tensors = src.cache_desc.num_tensors
435 dst_num_tensors = dst.cache_desc.num_tensors
436 raise_if_false(
437 src_num_tensors == dst_num_tensors,
438 f"src num_tensors:{src_num_tensors} and dst num_tensors:{dst_num_tensors} must be equal",
439 )
440 log.info("src and dst cache num:%d", src_num_tensors)
442 src_block_num = src.cache_desc.batch_size
443 dst_block_num = dst.cache_desc.batch_size
444 log.info("src num block:%d, dst num block:%d", src_block_num, dst_block_num)
445 for src_block_index, dst_block_index in src_to_dst.items():
446 raise_if_false(
447 0 <= src_block_index < src_block_num,
448 f"src_block_index:{src_block_index} must be in [0, {src_block_num})",
449 )
450 raise_if_false(
451 0 <= dst_block_index < dst_block_num,
452 f"dst_block_index:{dst_block_index} must be in [0, {dst_block_num})",
453 )
455 @staticmethod
456 def _is_cpu_cache(cache: KvCache):
457 return cache.cache_desc.placement == Placement.HOST
459 def copy_blocks(self, cache: KvCache, copy_block_info: Dict[int, List[int]]):
460 """
461 Args:
462 cache: 目标缓存
463 copy_block_info: 拷贝信息, (int, List[int])代表(原始block index, 目标block index)
464 """
465 check_isinstance("cache", cache, KvCache)
466 check_isinstance("copy_block_info", copy_block_info, dict)
467 check_dict("copy_block_info", copy_block_info, int, list, int)
468 raise_if_true(self._is_cpu_cache(cache), "Only support device cache")
469 copy_block_infos = []
470 for src_block, dst_blocks in copy_block_info.items():
471 check_uint64("src_block", src_block)
472 for dst_block in dst_blocks:
473 check_uint64("dst_block", dst_block)
474 copy_block_infos.append((src_block, dst_block))
475 param = (
476 cache.cache_id,
477 cache.cache_id,
478 0,
479 0,
480 0,
481 -1,
482 _INVALID_ID,
483 copy_block_infos,
484 )
485 ret = self._llm_engine.copy_cache(param)
486 handle_llm_status(ret, "[copy_blocks]", "cache id is:%s" % cache.cache_id)
487 log.info("[copy_blocks] success")
489 def copy_cache(
490 self,
491 dst: KvCache,
492 src: KvCache,
493 dst_batch_index: int = 0,
494 src_batch_index: int = 0,
495 offset: int = 0,
496 size: int = -1,
497 req_id: Optional[int] = None,
498 ) -> None:
499 """
500 拷贝KV, src/dst的CacheDesc需要匹配
502 Args:
503 dst: 目标Cache
504 src: 源Cache
505 dst_batch_index: 目标Cache的batch_index
506 src_batch_index: 源Cache的batch_index
507 offset: 每个tensor的偏移
508 size: 每个tensor拷贝的大小
509 req_id(Optional): 本次操作关联的req_id, 仅用于维测
511 Examples:
512 >>> from llm_datadist import LLMRole, CacheDesc, DataType, LLMDataDist
513 >>> # init LLMDataDist
514 >>> cluster_id = 0
515 >>> datadist = LLMDataDist(LLMRole.DECODER, cluster_id)
516 >>> engine_options = {} # 按需填写
517 >>> datadist.init(engine_options)
518 >>> # get KvCacheManager
519 >>> kv_cache_manager = datadist.kv_cache_manager
520 >>> # allocate caches
521 >>> tmp_kv_cache_desc = CacheDesc(num_tensors=80, shape=[1, 256], data_type=DataType.DT_FLOAT16)
522 >>> model_kv_cache_desc = CacheDesc(num_tensors=80, shape=[4, 256], data_type=DataType.DT_FLOAT16)
523 >>> tmp_cache = kv_cache_manager.allocate_cache(tmp_kv_cache_desc)
524 >>> model_cache = kv_cache_manager.allocate_cache(model_kv_cache_desc)
525 >>> # pull prompt kv to tmp cache
526 >>> prompt_cache_key = CacheKey(prompt_cluster_id=0, req_id = 1, model_id = 0)
527 >>> kv_cache_manager.pull_cache(prompt_cache_key, tmp_cache)
528 >>> # copy cache from tmp cache to model cache
529 >>> kv_cache_manager.copy_cache(model_cache, tmp_cache)
530 """
531 check_isinstance("dst", dst, KvCache)
532 raise_if_true(self._is_cpu_cache(dst), "Only support device cache")
533 check_isinstance("src", src, KvCache)
534 raise_if_true(self._is_cpu_cache(src), "Only support device cache")
535 check_uint32("dst_batch_index", dst_batch_index)
536 check_uint32("src_batch_index", src_batch_index)
537 check_uint64("offset", offset)
538 check_int64("size", size)
539 raise_if_false(
540 size == -1 or size > 0,
541 "[copy_cache] param check failed, size ({0}) is invalid, should be = -1 or > 0",
542 size,
543 )
544 user_param = (
545 dst.cache_id,
546 src.cache_id,
547 dst_batch_index,
548 src_batch_index,
549 offset,
550 size,
551 req_id,
552 [],
553 )
554 log.info("[copy_cache] start, param = %s", user_param)
555 if req_id is not None:
556 check_isinstance("req_id", req_id, int)
557 else:
558 req_id = _INVALID_ID
559 param = (
560 dst.cache_id,
561 src.cache_id,
562 dst_batch_index,
563 src_batch_index,
564 offset,
565 size,
566 req_id,
567 [],
568 )
569 ret = self._llm_engine.copy_cache(param)
570 handle_llm_status(ret, "[copy_cache]", f"param = {param}")
571 log.info("[copy_cache] success")
573 def get_cache_tensors(self, cache: KvCache, tensor_index: int = 0) -> List:
574 """
575 获取cache tensor
577 Args:
578 cache: KvCache
579 tensor_index: tensor index
580 """
581 check_isinstance("cache", cache, KvCache)
582 check_int32("tensor_index", tensor_index)
583 raise_if_false(
584 0 <= tensor_index < cache.cache_desc.num_tensors,
585 "[get_cache_tensors] param check failed, tensor_index ({0}) out of range, [0, {1})",
586 tensor_index,
587 cache.cache_desc.num_tensors,
588 )
589 ret, outputs = self._llm_engine.get_tensor(cache.cache_id, tensor_index)
590 handle_llm_status(
591 ret,
592 "[get_cache_tensors]",
593 {"cache_id": cache.cache_id, "tensor_index": tensor_index},
594 )
595 tensors = [Tensor.from_tensor_tuple(output) for output in outputs]
596 return tensors
598 def swap_blocks(self, src: KvCache, dst: KvCache, src_to_dst: Dict[int, int]) -> None:
599 """
600 交换blocks
602 Args:
603 src: 源KvCache
604 dst: 目的KvCache
605 src_to_dst: block index的字典
606 """
607 self._verify_caches(src, dst, src_to_dst)
608 src_placement = src.cache_desc.placement
609 dst_placement = dst.cache_desc.placement
610 is_swap_in = (src_placement == Placement.HOST) and (dst_placement == Placement.DEVICE)
611 is_swap_out = (src_placement == Placement.DEVICE) and (dst_placement == Placement.HOST)
612 raise_if_false(
613 is_swap_in or is_swap_out,
614 f"swap src placement:{src_placement} to dst placement:{dst_placement} is not support",
615 )
617 # 0标识swap in,1标识swap out
618 swap_type = 0 if is_swap_in else 1
619 block_size = src.cache_desc.size // src.cache_desc.batch_size
620 default_cache_id = -1
621 src_cache = (default_cache_id, src.per_device_tensor_addrs)
622 dst_cache = (default_cache_id, dst.per_device_tensor_addrs)
623 ret = self._llm_engine.swap_blocks(
624 src_cache,
625 dst_cache,
626 block_size,
627 swap_type,
628 self._llm_engine.dict_to_vector(src_to_dst),
629 )
630 handle_llm_status(ret, "[swap_blocks]", "swap blocks failed")
631 log.info("[swap_blocks] success")
633 def transfer_cache_async(
634 self,
635 src_cache: KvCache,
636 layer_synchronizer: LayerSynchronizer,
637 transfer_configs: Union[List[TransferConfig], Tuple[TransferConfig]],
638 src_block_indices: Optional[Union[List[int], Tuple[int]]] = None,
639 dst_block_indices: Optional[Union[List[int], Tuple[int]]] = None,
640 dst_block_memory_size: Optional[int] = None,
641 ) -> CacheTask:
642 self._check_role("[transfer_cache_async]", LLMRole.PROMPT)
643 check_isinstance("src_cache", src_cache, KvCache, allow_none=False)
644 params = TransferCacheParameters(
645 src_cache,
646 transfer_configs,
647 src_block_indices,
648 dst_block_indices,
649 dst_block_memory_size,
650 )
651 log.info("[transfer_cache_async] start, params = %s", params)
652 return transfer_cache_async(
653 params,
654 layer_synchronizer,
655 self._llm_engine.transfer_cache,
656 LLMStatusCode.LLM_WAIT_PROCESS_TIMEOUT,
657 )
659 def _check_role(self, func_name, role: LLMRole) -> None:
660 raise_if_false(self._role == role, "{0} is not supported by {1}", func_name, self._role)
662 def _switch_role(self, role: LLMRole) -> None:
663 log.info(f"[switch_role] [{self._role.name}->{role.name}] success")
664 self._role = role