Coverage for /opt/cloud/slavespace/usr1/096471637100f3de0fcfc01072822a80/dttest/api/python/llm_datadist_v1/llm_types.py: 97%

358 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-07-27 10:03 +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 

13__all__ = [ 

14 "CacheDesc", 

15 "CacheKey", 

16 "CacheKeyByIdAndIndex", 

17 "KvCache", 

18 "BlocksCacheKey", 

19 "Placement", 

20 "CacheTask", 

21 "TransferConfig", 

22 "LayerSynchronizer", 

23] 

24 

25from abc import ABC, abstractmethod 

26from enum import Enum, IntEnum 

27from typing import List, Optional, Tuple, Union 

28 

29from llm_datadist_v1 import llm_wrapper 

30 

31from .data_type import DataType 

32from .status import LLMException, LLMStatusCode, raise_if_false, raise_if_true 

33from .utils import log 

34from .utils.utils import ( 

35 check_int32, 

36 check_int64, 

37 check_isinstance, 

38 check_list_int64, 

39 check_list_uint64, 

40 check_type, 

41 check_uint32, 

42 check_uint64, 

43) 

44 

45_INVALID_ID = 2**64 - 1 

46 

47 

48class PushType(Enum): 

49 NO_CACHE_KEY = 0 

50 BLOCKS_CACHE_KEY = 1 

51 CACHE_KEY_BY_ID = 2 

52 

53 

54class RegisterMemStatus(Enum): 

55 OK = 0 

56 PREPARING = 1 

57 FAILED = 2 

58 

59 

60class Memtype(Enum): 

61 MEM_TYPE_DEVICE = 0 

62 MEM_TYPE_HOST = 1 

63 

64 

65class MemInfo(object): 

66 def __init__(self, mem_type: Memtype, addr: int, size: int): 

67 """ 

68 初始化 

69 Args: 

70 mem_type: 内存类型 

71 addr: 数据地址 

72 size: 数据大小 

73 """ 

74 check_isinstance("mem_type", mem_type, Memtype) 

75 check_isinstance("addr", addr, int) 

76 check_isinstance("size", size, int) 

77 self._mem_type = mem_type 

78 self._addr = addr 

79 self._size = size 

80 

81 @property 

82 def mem_type(self): 

83 return self._mem_type 

84 

85 @property 

86 def addr(self): 

87 return self._addr 

88 

89 @property 

90 def size(self): 

91 return self._size 

92 

93 def __str__(self): 

94 return f"MemInfo(mem_type={str(self.mem_type)}, addr={str(self.addr)}, size={str(self.size)})" 

95 

96 def __repr__(self): 

97 return self.__str__() 

98 

99 

100int_to_mem_status_dict = { 

101 0: RegisterMemStatus.OK, 

102 1: RegisterMemStatus.PREPARING, 

103 2: RegisterMemStatus.FAILED, 

104} 

105 

106 

107class Placement(IntEnum): 

108 HOST = 0 

109 DEVICE = 1 

110 

111 

112class CacheDesc(object): 

113 """ 

114 Cache描述 

115 

116 Args: 

117 num_tensors: Cache中tensor的个数, 操作Cache时, 所有tensor会做同样的操作 

118 shape: tensor shape 

119 data_type: data type 

120 placement: 标识kv cache所在的设备 

121 batch_dim_index (Optional): batch dim的index, 默认为0 

122 seq_len_dim_index (Optional): seq_len dim的index, 高阶API pull_kv时需要该字段判断是否能够按实际大小拉取 

123 kv_tensor_format (Optional): kv tensor的data format 

124 

125 Examples: 

126 >>> from llm_datadist import CacheDesc 

127 >>> tensor_num = 80 

128 >>> tensor_shape = [4, 256] 

129 >>> tensor_data_type = DataType.DT_FLOAT16 

130 >>> cache_desc_1 = CacheDesc(tensor_num, tensor_shape, tensor_data_type) 

131 >>> # 指定 batch_dim_index = 0, seq_len_dim_index = 1 

132 >>> cache_desc_2 = CacheDesc(tensor_num, tensor_shape, tensor_data_type, Placement.DEVICE, 0, 1) 

133 """ 

134 

135 def __init__( 

136 self, 

137 num_tensors: int, 

138 shape: Union[Tuple[int], List[int]], 

139 data_type: DataType, 

140 placement: Placement = Placement.DEVICE, 

141 batch_dim_index: int = 0, 

142 seq_len_dim_index: int = -1, 

143 kv_tensor_format: str = None, 

144 ): 

145 self._num_tensors = num_tensors 

146 self._shape = shape 

147 check_uint32("num_tensors", num_tensors) 

148 check_isinstance("shape", shape, [tuple, list], int) 

149 check_isinstance("data_type", data_type, DataType) 

150 check_isinstance("placement", placement, Placement) 

151 check_isinstance("batch_dim_index", batch_dim_index, int) 

152 check_int32("seq_len_dim_index", seq_len_dim_index) 

153 check_isinstance("kv_tensor_format", kv_tensor_format, str) 

154 raise_if_false( 

155 0 <= batch_dim_index < len(shape), 

156 "batch_dim_index {0} out of range, [0, {1})", 

157 batch_dim_index, 

158 len(shape), 

159 ) 

160 raise_if_false( 

161 seq_len_dim_index == -1 or 0 <= seq_len_dim_index < len(shape), 

162 "seq_len_dim_index {0} is invalid, should be -1 or in range [0, {1})", 

163 seq_len_dim_index, 

164 len(shape), 

165 ) 

166 check_list_int64("shape", shape) 

167 self._data_type = data_type 

168 self._batch_dim = batch_dim_index 

169 self._batch_size = shape[batch_dim_index] 

170 self._seq_len_dim_index = seq_len_dim_index 

171 self._size = -1 

172 self._kv_tensor_format = kv_tensor_format 

173 self._placement = placement 

174 self._is_blocks = False 

175 

176 def __repr__(self): 

177 return ( 

178 f"CacheDesc(num_tensors={self.num_tensors}, " 

179 f"shape={self.shape}, " 

180 f"data_type={self.data_type}, " 

181 f"placement={self.placement}, " 

182 f"batch_dim_index={self.batch_dim}, " 

183 f"seq_len_dim_index={self._seq_len_dim_index}," 

184 f"kv_tensor_format={self.kv_tensor_format})" 

185 ) 

186 

187 @property 

188 def num_tensors(self) -> int: 

189 return self._num_tensors 

190 

191 @property 

192 def shape(self) -> List[int]: 

193 return self._shape 

194 

195 def update_dim(self, dim_index: int, dim_value: int) -> None: 

196 self._shape[dim_index] = dim_value 

197 self._size = -1 

198 

199 @property 

200 def data_type(self) -> DataType: 

201 return self._data_type 

202 

203 @property 

204 def batch_dim(self) -> int: 

205 return self._batch_dim 

206 

207 @property 

208 def batch_size(self) -> int: 

209 return self._batch_size 

210 

211 @property 

212 def seq_len_dim_index(self) -> int: 

213 return self._seq_len_dim_index 

214 

215 @property 

216 def size(self) -> int: 

217 if self._size == -1: 

218 self._size = llm_wrapper.calc_tensor_size(self.shape, self._data_type.value) 

219 if self._size < 0: 

220 raise LLMException(f"Failed to calc tensor size, shape = {self.shape}, data_type = {self.data_type}") 

221 return self._size 

222 

223 @property 

224 def kv_tensor_format(self) -> str: 

225 return self._kv_tensor_format 

226 

227 @property 

228 def placement(self) -> Placement: 

229 return self._placement 

230 

231 

232class BlocksCacheKey(object): 

233 """ 

234 BlocksCacheKey, PROMPT allocate blocks cache时用于建立索引, DECODER pull_kv时作为索引传入 

235 

236 Args: 

237 prompt_cluster_id or cluster_id: remote cluster id, 用于pull_kv时需要准确填写, allocate_cache时会忽略该字段 

238 model_id (Optional): model id, 默认为0, 在投机等同时加载多个model的场景需要按需设置 

239 

240 Raises: 

241 TypeError: if `prompt_cluster_id` or `cluster_id` is not int 

242 TypeError: if `model_id` is not int 

243 

244 Example: 

245 >>> from llm_datadist import BlocksCacheKey 

246 >>> cluster_id = 1 

247 >>> # model_id取默认值 

248 >>> cache_key = BlocksCacheKey(cluster_id) 

249 >>> # 指定model_id 

250 >>> model_id = 1 

251 >>> cache_key_1 = BlocksCacheKey(cluster_id, model_id) 

252 """ 

253 

254 def __init__(self, *args, **kwargs): 

255 raise_if_false((len(args) + len(kwargs)) <= 2, "Param num is over limit") 

256 if len(args) > 0: 

257 kwargs["cluster_id"] = args[0] 

258 if len(args) > 1: 

259 kwargs["model_id"] = args[1] 

260 raise_if_false( 

261 "prompt_cluster_id" in kwargs or "cluster_id" in kwargs, 

262 "Param prompt_cluster_id or cluster_id is required", 

263 ) 

264 valid_keys = ["prompt_cluster_id", "cluster_id", "model_id"] 

265 for k in kwargs.keys(): 

266 raise_if_false(k in valid_keys, "Unsupported param:{}", k) 

267 self._cluster_id = kwargs["cluster_id"] if "cluster_id" in kwargs else kwargs["prompt_cluster_id"] 

268 self._cluster_id_key = "cluster_id" if "cluster_id" in kwargs else "prompt_cluster_id" 

269 self._model_id = kwargs["model_id"] if "model_id" in kwargs else 0 

270 check_isinstance("cluster_id", self._cluster_id, int, extra_fmt="The first param ") 

271 check_isinstance("model_id", self._model_id, int, extra_fmt="The second param ") 

272 

273 @property 

274 def prompt_cluster_id(self) -> int: 

275 return self._cluster_id 

276 

277 @property 

278 def cluster_id(self) -> int: 

279 return self._cluster_id 

280 

281 @property 

282 def model_id(self) -> int: 

283 return self._model_id 

284 

285 def __repr__(self): 

286 return f"BlocksCacheKey({self._cluster_id_key}={self.cluster_id}, model_id={self.model_id})" 

287 

288 

289class CacheKey(object): 

290 """ 

291 CacheKey, REMOTE allocate_cache时用于建立索引, LOCAL pull_cache时作为索引传入 

292 

293 Args: 

294 prompt_cluster_id or cluster_id: remote cluster id, 用于pull_kv时需要准确填写, allocate_cache时会忽略该字段 

295 req_id: request id 

296 model_id (Optional): model id, 默认为0, 在投机等同时加载多个model的场景需要按需设置 

297 prefix_id (Optional): prefix id, 默认为2 ** 64 - 1 

298 

299 Raises: 

300 TypeError: if `prompt_cluster_id` or `cluster_id` is not int 

301 TypeError: if `req_id` is not int 

302 TypeError: if `model_id` is not int 

303 TypeError: if `prefix_id` is not int 

304 

305 Example: 

306 >>> from llm_datadist import CacheKey 

307 >>> cluster_id = 1 

308 >>> request_id = 1 

309 >>> # model_id取默认值 

310 >>> cache_key = CacheKey(cluster_id, request_id) 

311 >>> # 指定model_id 

312 >>> cache_key_1 = CacheKey(cluster_id, request_id, 2) 

313 """ 

314 

315 def __init__(self, *args, **kwargs): 

316 raise_if_false((len(args) + len(kwargs)) <= 4, "Param num is over limit") 

317 if len(args) > 0: 

318 kwargs["cluster_id"] = args[0] 

319 if len(args) > 1: 

320 kwargs["req_id"] = args[1] 

321 if len(args) > 2: 

322 kwargs["model_id"] = args[2] 

323 if len(args) > 3: 

324 kwargs["prefix_id"] = args[3] 

325 raise_if_false( 

326 "prompt_cluster_id" in kwargs or "cluster_id" in kwargs, 

327 "Param prompt_cluster_id or cluster_id is required", 

328 ) 

329 raise_if_false("req_id" in kwargs or "req_id" in kwargs, "Param req_id is required") 

330 valid_keys = [ 

331 "prompt_cluster_id", 

332 "cluster_id", 

333 "req_id", 

334 "model_id", 

335 "prefix_id", 

336 ] 

337 for k in kwargs.keys(): 

338 raise_if_false(k in valid_keys, "Unsupported param:{}", k) 

339 self._cluster_id = kwargs["cluster_id"] if "cluster_id" in kwargs else kwargs["prompt_cluster_id"] 

340 self._req_id = kwargs["req_id"] 

341 self._model_id = kwargs["model_id"] if "model_id" in kwargs else 0 

342 self._prefix_id = kwargs["prefix_id"] if "prefix_id" in kwargs else _INVALID_ID 

343 check_uint64("cluster_id", self._cluster_id) 

344 check_uint64("req_id", self._req_id) 

345 check_uint64("model_id", self._model_id) 

346 check_uint64("prefix_id", self._prefix_id) 

347 

348 @property 

349 def prompt_cluster_id(self) -> int: 

350 return self._cluster_id 

351 

352 @property 

353 def cluster_id(self) -> int: 

354 return self._cluster_id 

355 

356 @property 

357 def req_id(self) -> int: 

358 return self._req_id 

359 

360 @property 

361 def prefix_id(self) -> int: 

362 return self._prefix_id 

363 

364 @property 

365 def model_id(self) -> int: 

366 return self._model_id 

367 

368 def __repr__(self): 

369 return ( 

370 f"CacheKey(cluster_id={self.cluster_id}, " 

371 f"req_id={self.req_id}, " 

372 f"prefix_id={self.prefix_id}, " 

373 f"model_id={self.model_id})" 

374 ) 

375 

376 

377class CacheKeyByIdAndIndex(object): 

378 """ 

379 索引Prompt中kv cache的单个batch index, 用于Decoder pull_cache时作为索引传入 

380 

381 Args: 

382 cluster_id: kv所在的prompt的cluster_id 

383 cache_id: kv cache的id 

384 batch_index (Optional): batch index, 默认为0, 在投机等同时加载多个model的场景需要按需设置 

385 

386 Raises: 

387 TypeError: if `cluster_id` is not int 

388 TypeError: if `cache_id` is not int 

389 TypeError: if `batch_index` is not int 

390 """ 

391 

392 def __init__(self, cluster_id: int, cache_id: int, batch_index=0): 

393 check_uint64("cluster_id", cluster_id) 

394 check_int64("cache_id", cache_id) 

395 check_uint32("batch_index", batch_index) 

396 self._prompt_cluster_id = cluster_id 

397 self._prompt_cache_id = cache_id 

398 self._prompt_batch_index = batch_index 

399 self._req_id = _INVALID_ID 

400 

401 @property 

402 def cluster_id(self) -> int: 

403 return self._prompt_cluster_id 

404 

405 @property 

406 def cache_id(self) -> int: 

407 return self._prompt_cache_id 

408 

409 @property 

410 def batch_index(self) -> int: 

411 return self._prompt_batch_index 

412 

413 def __repr__(self): 

414 return ( 

415 f"CacheKeyByIdAndIndex(cluster_id={self.cluster_id}, " 

416 f"cache_id={self.cache_id}, " 

417 f"batch_index={self.batch_index})" 

418 ) 

419 

420 

421class KvCache(object): 

422 """ 

423 Kv Cache, 由KvCacheManager.allocate_cache创建,用户不应直接构造 

424 """ 

425 

426 def __init__( 

427 self, 

428 cache_id: int, 

429 cache_desc: CacheDesc, 

430 per_device_tensor_addrs: List[List[int]], 

431 kv_cache_manager, 

432 ): 

433 self._cache_id = cache_id 

434 check_int64("cache_id", cache_id) 

435 self._cache_desc = cache_desc 

436 self._per_device_tensor_addrs = per_device_tensor_addrs 

437 self._valid = True 

438 self._kv_cache_manager = kv_cache_manager 

439 

440 def __del__(self): 

441 # 保底释放kv cache, 但更推荐主动通过调用kv_cache_manager.deallocate_cache来释放kv cache,而不应该遗留到此处自动释放 

442 if self._valid and self._kv_cache_manager is not None and self._kv_cache_manager.is_initialized(): 

443 log.info("auto deallocate kv cache: %s", self.cache_id) 

444 self._kv_cache_manager.deallocate_cache(self) 

445 

446 @property 

447 def cache_id(self) -> int: 

448 return self._cache_id 

449 

450 @property 

451 def cache_desc(self) -> CacheDesc: 

452 return self._cache_desc 

453 

454 @property 

455 def per_device_tensor_addrs(self) -> List[List[int]]: 

456 return self._per_device_tensor_addrs 

457 

458 def __str__(self): 

459 return f"KvCache(cache_id = {self._cache_id}, num_devices = {len(self._per_device_tensor_addrs)})" 

460 

461 def __repr__(self): 

462 return self.__str__() 

463 

464 @classmethod 

465 def create_cpu_cache(cls, cache_desc: CacheDesc, addrs: Union[List[int], List[List[int]]]): 

466 check_isinstance("cache_desc", cache_desc, CacheDesc) 

467 raise_if_false(cache_desc.placement == Placement.HOST, "cache_desc placement must be HOST") 

468 check_isinstance("addrs", addrs, list) 

469 raise_if_false(len(addrs) > 0, "addrs length should be bigger than zero.") 

470 last_type = type(addrs[0]) 

471 for addr in addrs: 

472 check_isinstance("the internal element of addrs", addr, [list, int]) 

473 if check_type(addr, list): 

474 check_isinstance("the internal element of addrs", addr, list, int) 

475 raise_if_false( 

476 cache_desc.num_tensors == len(addr), 

477 f"cache_desc num_tensors:{cache_desc.num_tensors} " 

478 f"should be equal to size of the internal element of addrs:{len(addr)}", 

479 ) 

480 raise_if_false( 

481 check_type(addr, last_type), 

482 "the type of the internal element of addrs should be consistent.", 

483 ) 

484 last_type = type(addr) 

485 if last_type == int: 

486 raise_if_false( 

487 cache_desc.num_tensors == len(addrs), 

488 f"cache_desc num_tensors:{cache_desc.num_tensors} should be equal to size of addrs:{len(addrs)}", 

489 ) 

490 return cls(-1, cache_desc, [addrs] if last_type == int else addrs, None) 

491 

492 

493class LayerSynchronizer(ABC): 

494 @abstractmethod 

495 def synchronize_layer(self, layer_index: int, timeout_in_millis: Optional[int]) -> bool: 

496 """ 

497 阻塞等待指定层计算完成 

498 

499 Args: 

500 layer_index(int): 要同步的layer的index 

501 timeout_in_millis(Optional[int]): 超时时间, 不配置则不超时 

502 

503 Returns: 

504 True: 同步成功, False: 同步失败 

505 """ 

506 

507 

508def check_layer_range(name, layer_range: Optional[range], allow_none=True): 

509 if allow_none and layer_range is None: 

510 return 

511 check_isinstance(name, layer_range, range, allow_none=allow_none) 

512 raise_if_false(layer_range.step == 1, f"check {name}.step == 1 failed") 

513 raise_if_false( 

514 0 <= layer_range.start < layer_range.stop, 

515 "check 0 <= range.start < range.stop failed, {0} = {1}", 

516 name, 

517 layer_range, 

518 ) 

519 

520 

521class TransferWithCacheKeyConfig: 

522 def __init__( 

523 self, 

524 cache_key: Union[BlocksCacheKey, CacheKeyByIdAndIndex], 

525 src_layer_range: range = None, 

526 dst_layer_range: range = None, 

527 src_batch_index: int = 0, 

528 ): 

529 check_isinstance( 

530 "cache_key", 

531 cache_key, 

532 [BlocksCacheKey, CacheKeyByIdAndIndex], 

533 allow_none=False, 

534 ) 

535 self._cache_key = cache_key 

536 check_layer_range("src_layer_range", src_layer_range, allow_none=False) 

537 self._src_layer_range = src_layer_range 

538 check_layer_range("dst_layer_range", dst_layer_range, allow_none=False) 

539 self._dst_layer_range = dst_layer_range 

540 raise_if_false( 

541 (src_layer_range.stop - src_layer_range.start) == (dst_layer_range.stop - src_layer_range.start), 

542 "src_layer_range size shoulde be equal to dst_layer_range size", 

543 ) 

544 check_uint32("src_batch_index", src_batch_index) 

545 

546 raise_if_true( 

547 check_type(cache_key, BlocksCacheKey) and src_batch_index != 0, 

548 "src_batch_index shoulde be 0 when cache_key is BlocksCacheKey.", 

549 ) 

550 self._src_batch_index = src_batch_index 

551 self.dst_cluster_id = cache_key.cluster_id 

552 

553 def __repr__(self) -> str: 

554 return ( 

555 f"TransferWithCacheKeyConfig(cache_key={self._cache_key}," 

556 f" src_layer_range={self._src_layer_range}," 

557 f" dst_layer_range={self._dst_layer_range}," 

558 f" src_batch_index={self._src_batch_index})" 

559 ) 

560 

561 @property 

562 def cache_key(self): 

563 return self._cache_key 

564 

565 @cache_key.setter 

566 def cache_key(self, cache_key: Union[BlocksCacheKey, CacheKeyByIdAndIndex]) -> None: 

567 check_isinstance("cache_key", cache_key, [BlocksCacheKey, CacheKeyByIdAndIndex]) 

568 self._cache_key = cache_key 

569 

570 @property 

571 def src_layer_range(self) -> Optional[range]: 

572 return self._src_layer_range 

573 

574 @src_layer_range.setter 

575 def src_layer_range(self, src_layer_range: Optional[range]) -> None: 

576 check_layer_range("src_layer_range", src_layer_range) 

577 self._src_layer_range = src_layer_range 

578 

579 @property 

580 def dst_layer_range(self) -> Optional[range]: 

581 return self._dst_layer_range 

582 

583 @dst_layer_range.setter 

584 def dst_layer_range(self, dst_layer_range: Optional[range]) -> None: 

585 check_layer_range("dst_layer_range", dst_layer_range) 

586 self._dst_layer_range = dst_layer_range 

587 

588 @property 

589 def src_batch_index(self) -> int: 

590 return self._src_batch_index 

591 

592 @src_batch_index.setter 

593 def src_batch_index(self, src_batch_index: int) -> None: 

594 raise_if_true( 

595 check_type(self.cache_key, BlocksCacheKey) and src_batch_index != 0, 

596 "src_batch_index shoulde be 0 when cache_key is BlocksCacheKey.", 

597 ) 

598 self._check_src_batch_index(src_batch_index) 

599 self._src_batch_index = src_batch_index 

600 

601 @staticmethod 

602 def _check_src_batch_index(src_batch_index: int): 

603 check_uint32("src_batch_index", src_batch_index) 

604 

605 

606class TransferConfig: 

607 def __init__( 

608 self, 

609 dst_cluster_id: int, 

610 dst_addrs: List[int], 

611 src_layer_range: Optional[range] = None, 

612 src_batch_index: int = 0, 

613 ): 

614 self._check_dst_cluster_id(dst_cluster_id) 

615 self._check_dst_addrs(dst_addrs) 

616 check_layer_range("src_layer_range", src_layer_range) 

617 self._check_src_batch_index(src_batch_index) 

618 self._dst_cluster_id = dst_cluster_id 

619 self._dst_addrs = dst_addrs 

620 self._src_layer_range = src_layer_range 

621 self._src_batch_index = src_batch_index 

622 self.dst_layer_range = None 

623 

624 def __repr__(self) -> str: 

625 return ( 

626 f"TransferConfig(src_cluster_id={self._dst_cluster_id}," 

627 f" dst_addrs={self._dst_addrs}," 

628 f" src_layer_range={self._src_layer_range}," 

629 f" src_batch_index={self._src_batch_index})" 

630 ) 

631 

632 def __str__(self): 

633 return self.__repr__() 

634 

635 @property 

636 def dst_cluster_id(self) -> int: 

637 return self._dst_cluster_id 

638 

639 @property 

640 def dst_addrs(self) -> List[int]: 

641 return self._dst_addrs 

642 

643 @property 

644 def src_layer_range(self) -> Optional[range]: 

645 return self._src_layer_range 

646 

647 @property 

648 def src_batch_index(self) -> int: 

649 return self._src_batch_index 

650 

651 @dst_cluster_id.setter 

652 def dst_cluster_id(self, cluster_id: int) -> None: 

653 self._check_dst_cluster_id(cluster_id) 

654 self._dst_cluster_id = cluster_id 

655 

656 @dst_addrs.setter 

657 def dst_addrs(self, dst_addrs: List[int]) -> None: 

658 self._check_dst_addrs(dst_addrs) 

659 self._dst_addrs = dst_addrs 

660 

661 @src_layer_range.setter 

662 def src_layer_range(self, layer_range: Optional[range]) -> None: 

663 check_layer_range("src_layer_range", layer_range) 

664 self._src_layer_range = layer_range 

665 

666 @src_batch_index.setter 

667 def src_batch_index(self, batch_index: int) -> None: 

668 self._check_src_batch_index(batch_index) 

669 self._src_batch_index = batch_index 

670 

671 @staticmethod 

672 def _check_dst_cluster_id(dst_cluster_id: int): 

673 check_uint64("dst_cluster_id", dst_cluster_id) 

674 

675 @staticmethod 

676 def _check_dst_addrs(dst_addrs: List[int]): 

677 check_isinstance("dst_addrs", dst_addrs, [list, tuple], int, allow_none=False) 

678 check_list_uint64("dst_addrs", dst_addrs) 

679 

680 @staticmethod 

681 def _check_src_batch_index(src_batch_index: int): 

682 check_uint32("src_batch_index", src_batch_index) 

683 

684 

685class CacheTask: 

686 def __init__(self, transfer_async_task): 

687 self._transfer_async_task = transfer_async_task 

688 

689 def synchronize(self, timeout_in_millis: Optional[int] = None) -> LLMStatusCode: 

690 timeout = None 

691 if timeout_in_millis is not None: 

692 check_uint32("timeout_in_millis", timeout_in_millis) 

693 timeout = timeout_in_millis / 1000 

694 return self._transfer_async_task.get(timeout) 

695 

696 def get_results(self, timeout_in_millis: Optional[int] = None) -> List[LLMStatusCode]: 

697 timeout = None 

698 if timeout_in_millis is not None: 

699 check_uint32("timeout_in_millis", timeout_in_millis) 

700 timeout = timeout_in_millis / 1000 

701 return self._transfer_async_task.get_results(timeout)