refactor: use single-index access and update docs for cache architecture
- Replace all buffer[layer_id][loc] double indexing with buffer[layer_id, loc] single advanced indexing in cache.py and attention.py - Revert KVStorage buffers back to 4D [n_layers, size, n_kv_heads, head_dim], remove leftover 3D reshape/view in MLA path - Update docs/guides/inference.md, docs/developer/internals.md, docs/developer/architecture.md to reflect new PagePool/KVStorage/ReqToTokenPool/KVCache classes
This commit is contained in:
@@ -155,10 +155,10 @@ class ReqToTokenPool:
|
||||
|
||||
|
||||
class KVStorage:
|
||||
"""Token-level flat KV cache storage with NHD layout.
|
||||
"""Token-level KV cache storage.
|
||||
|
||||
Buffers: [n_layers, size, n_kv_heads, head_dim]. Each token occupies
|
||||
one contiguous row. Logical ordering is determined by ReqToTokenPool.
|
||||
one slot indexed by ReqToTokenPool.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -185,8 +185,8 @@ class KVStorage:
|
||||
return self.v_buffer[layer_id]
|
||||
|
||||
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
|
||||
self.k_buffer[layer_id][loc] = k
|
||||
self.v_buffer[layer_id][loc] = v
|
||||
self.k_buffer[layer_id, loc] = k
|
||||
self.v_buffer[layer_id, loc] = v
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -87,8 +87,8 @@ class GQA(nn.Module):
|
||||
q, k = self.q_norm(q), self.k_norm(k)
|
||||
|
||||
if kv_cache is not None:
|
||||
kv_cache.k_buffer[self.layer_id][kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[self.layer_id][kv_cache.out_cache_loc] = v
|
||||
kv_cache.k_buffer[self.layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[self.layer_id, kv_cache.out_cache_loc] = v
|
||||
|
||||
max_len = kv_cache.seq_lens.max()
|
||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||
@@ -97,8 +97,8 @@ class GQA(nn.Module):
|
||||
< kv_cache.seq_lens[:, None]
|
||||
)
|
||||
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||
k = kv_cache.k_buffer[self.layer_id][indices]
|
||||
v = kv_cache.v_buffer[self.layer_id][indices]
|
||||
k = kv_cache.k_buffer[self.layer_id, indices]
|
||||
v = kv_cache.v_buffer[self.layer_id, indices]
|
||||
|
||||
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
||||
|
||||
@@ -204,8 +204,8 @@ class MLA(nn.Module):
|
||||
k = self.k_norm(k)
|
||||
|
||||
if kv_cache is not None:
|
||||
kv_cache.k_buffer[self.layer_id][kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[self.layer_id][kv_cache.out_cache_loc] = v
|
||||
kv_cache.k_buffer[self.layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[self.layer_id, kv_cache.out_cache_loc] = v
|
||||
|
||||
max_len = kv_cache.seq_lens.max()
|
||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||
@@ -214,8 +214,8 @@ class MLA(nn.Module):
|
||||
< kv_cache.seq_lens[:, None]
|
||||
)
|
||||
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||
k = kv_cache.k_buffer[self.layer_id][indices]
|
||||
v = kv_cache.v_buffer[self.layer_id][indices]
|
||||
k = kv_cache.k_buffer[self.layer_id, indices]
|
||||
v = kv_cache.v_buffer[self.layer_id, indices]
|
||||
|
||||
q = q.permute(0, 2, 1, 3)
|
||||
k = k.permute(0, 2, 1, 3)
|
||||
|
||||
Reference in New Issue
Block a user