MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT-LLM / KVCacheUpdater

Class KVCacheUpdater

tensorrt_llm/runtime/kv_cache_manager.py:382–473  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

380
381
382class 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,

Callers 1

decodeMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected