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

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# ----------------------------------------------------------------------------------------------------------- 

12 

13from typing import Dict, List, Optional, Tuple, Union 

14 

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) 

51 

52_NUM_TENSORS_PER_LAYER = 2 

53_INVALID_ID = 2**64 - 1 

54 

55 

56class KvCacheManager(object): 

57 """ 

58 提供了一组KvCache的操作函数, 在LLMEngine初始化后通过LLMEngine获取实例 

59 

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 """ 

70 

71 def __init__(self, llm_engine, role: LLMRole) -> None: 

72 self._llm_engine = llm_engine 

73 self._role = role 

74 self._initialized = True 

75 

76 def is_initialized(self) -> bool: 

77 return self._initialized 

78 

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释放 

85 

86 Args: 

87 cache_desc: Cache描述 

88 blocks_cache_key(Optional): 仅当LLMRole为PROMPT时可设置, 用于在DECODER拉取KV 

89 Returns: 

90 KvCache 

91 

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 

123 

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接口 

135 

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 

144 

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 

196 

197 def deallocate_cache(self, cache: KvCache) -> None: 

198 """ 

199 释放Cache, 如果该Cache在Allocate时关联了CacheKey, 则实际的释放会延后到所有的CacheKey被拉取或执行了remove_cache_key 

200 释放之后,不应再对该KvCache做任何操作 

201 

202 Args: 

203 cache: Cache 

204 

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") 

216 

217 def remove_cache_key(self, cache_key: CacheKey) -> None: 

218 """ 

219 移除CacheKey, 仅当LLMRole为PROMPT时可调用 

220 移除CacheKey后, 该Cache将无法再被pull_cache拉取 

221 

222 Args: 

223 cache_key: CacheKey 

224 

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") 

234 

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 ) 

277 

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") 

304 

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时可调用 

315 

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数量 

327 

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 ) 

374 

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 ) 

391 

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") 

406 

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 ) 

419 

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) 

425 

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) 

433 

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) 

441 

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 ) 

454 

455 @staticmethod 

456 def _is_cpu_cache(cache: KvCache): 

457 return cache.cache_desc.placement == Placement.HOST 

458 

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") 

488 

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需要匹配 

501 

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, 仅用于维测 

510 

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") 

572 

573 def get_cache_tensors(self, cache: KvCache, tensor_index: int = 0) -> List: 

574 """ 

575 获取cache tensor 

576 

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 

597 

598 def swap_blocks(self, src: KvCache, dst: KvCache, src_to_dst: Dict[int, int]) -> None: 

599 """ 

600 交换blocks 

601 

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 ) 

616 

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") 

632 

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 ) 

658 

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) 

661 

662 def _switch_role(self, role: LLMRole) -> None: 

663 log.info(f"[switch_role] [{self._role.name}->{role.name}] success") 

664 self._role = role