| 380 | |
| 381 | |
| 382 | class KVCacheUpdater: |
| 383 | |
| 384 | def __init__(self): |
| 385 | self.use_paged_kv_cache = None |
| 386 | self.num_layers = None |
| 387 | self.num_kv_heads = None |
| 388 | self.head_dim = None |
| 389 | self.elt_size = None |
| 390 | self.past_key_value_list = None |
| 391 | self.max_kv_cache_length = None |
| 392 | self.kv_cache_manager = None |
| 393 | self.host_kv_cache_pool_pointers = None |
| 394 | |
| 395 | def init_linear_kv_cache(self, num_layers, num_kv_heads, head_dim, |
| 396 | kv_cache_type, past_key_value_list): |
| 397 | self.use_paged_kv_cache = False |
| 398 | self.num_layers = num_layers |
| 399 | self.num_kv_heads = num_kv_heads |
| 400 | self.head_dim = head_dim |
| 401 | self.past_key_value_list = past_key_value_list |
| 402 | self.elt_size = torch.zeros(1, dtype=kv_cache_type).element_size() |
| 403 | self.max_kv_cache_length = past_key_value_list[0].shape[3] |
| 404 | |
| 405 | def init_paged_kv_cache(self, num_layers, num_kv_heads, head_dim, |
| 406 | kv_cache_type, kv_cache_manager, |
| 407 | host_kv_cache_pool_pointers): |
| 408 | self.use_paged_kv_cache = True |
| 409 | self.num_layers = num_layers |
| 410 | self.num_kv_heads = num_kv_heads |
| 411 | self.head_dim = head_dim |
| 412 | self.kv_cache_manager = kv_cache_manager |
| 413 | self.host_kv_cache_pool_pointers = host_kv_cache_pool_pointers |
| 414 | self.elt_size = torch.zeros(1, dtype=kv_cache_type).element_size() |
| 415 | |
| 416 | def update(self, accepted_draft_token_offsets, |
| 417 | packed_accepted_draft_tokens_indices, sequence_length_buffer, |
| 418 | rewind_tokens): |
| 419 | assert isinstance(rewind_tokens, torch.Tensor) or isinstance( |
| 420 | rewind_tokens, int) |
| 421 | rewind_tokens_tensor = rewind_tokens if isinstance( |
| 422 | rewind_tokens, torch.Tensor) else None |
| 423 | rewind_tokens_count = rewind_tokens if isinstance(rewind_tokens, |
| 424 | int) else 0 |
| 425 | assert self.use_paged_kv_cache is not None |
| 426 | if self.use_paged_kv_cache: |
| 427 | if self.kv_cache_manager.has_single_pool(): |
| 428 | kv_cache_manager = self.kv_cache_manager.get_single_kv_cache_manager( |
| 429 | ) |
| 430 | else: |
| 431 | raise RuntimeError( |
| 432 | "Currently, using KVCacheUpdater with more then single memory pool is not supported" |
| 433 | ) |
| 434 | |
| 435 | host_kv_cache_block_offsets = kv_cache_manager.get_block_offsets(1) |
| 436 | kv_cache_block_offsets = host_kv_cache_block_offsets.to('cuda') |
| 437 | torch.ops.tensorrt_llm.update_kv_cache_draft_token_location( |
| 438 | accepted_draft_token_offsets, |
| 439 | packed_accepted_draft_tokens_indices, |