[NPU]fix:NPUMLATokenToKVPool object has no attribute "kv_buffer" (#25249)
This commit is contained in:
@@ -199,6 +199,64 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
||||
slot_indices=loc,
|
||||
)
|
||||
|
||||
def _chunk_copy_npu_to_cpu(self, buf_of_layers, indices):
|
||||
chunk_size = self.cpu_offloading_chunk_size
|
||||
out = []
|
||||
for tensors_per_layer in buf_of_layers: # [k_buf, v_buf]
|
||||
layer_chunks = []
|
||||
for i in range(0, len(indices), chunk_size):
|
||||
ci = indices[i : i + chunk_size]
|
||||
layer_chunks.append(
|
||||
[
|
||||
t[ci].to("cpu", non_blocking=True)
|
||||
for t in tensors_per_layer
|
||||
if t is not None
|
||||
]
|
||||
)
|
||||
out.append(layer_chunks)
|
||||
return out
|
||||
|
||||
# Parent MHATokenToKVPool.get_cpu_copy / load_cpu_copy use
|
||||
# `self.k_buffer[layer_id][chunk_indices]` which indexes the first dim.
|
||||
# NPUMHATokenToKVPool stores buffers as
|
||||
# (num_pages, page_size, head_num, head_dim) # use_fia=False
|
||||
# (num_pages*page_size, 1, head_num, head_dim) # use_fia=True
|
||||
def get_cpu_copy(self, indices):
|
||||
torch.npu.synchronize()
|
||||
buf_of_layers = []
|
||||
for local_layer_id in range(self.layer_num):
|
||||
k_layer = self.k_buffer[local_layer_id].view(
|
||||
-1, self.head_num, self.head_dim
|
||||
)
|
||||
v_layer = self.v_buffer[local_layer_id].view(
|
||||
-1, self.head_num, self.head_dim
|
||||
)
|
||||
buf_of_layers.append([k_layer, v_layer])
|
||||
kv_cache_cpu = self._chunk_copy_npu_to_cpu(buf_of_layers, indices)
|
||||
torch.npu.synchronize()
|
||||
return kv_cache_cpu
|
||||
|
||||
def load_cpu_copy(self, kv_cache_cpu, indices):
|
||||
torch.npu.synchronize()
|
||||
chunk_size = self.cpu_offloading_chunk_size
|
||||
for local_layer_id in range(self.layer_num):
|
||||
k_layer = self.k_buffer[local_layer_id].view(
|
||||
-1, self.head_num, self.head_dim
|
||||
)
|
||||
v_layer = self.v_buffer[local_layer_id].view(
|
||||
-1, self.head_num, self.head_dim
|
||||
)
|
||||
for i in range(0, len(indices), chunk_size):
|
||||
chunk_indices = indices[i : i + chunk_size]
|
||||
k_cpu, v_cpu = (
|
||||
kv_cache_cpu[local_layer_id][i // chunk_size][0],
|
||||
kv_cache_cpu[local_layer_id][i // chunk_size][1],
|
||||
)
|
||||
assert k_cpu.shape[0] == v_cpu.shape[0] == len(chunk_indices)
|
||||
k_layer[chunk_indices] = k_cpu.to(k_layer.device, non_blocking=True)
|
||||
v_layer[chunk_indices] = v_cpu.to(v_layer.device, non_blocking=True)
|
||||
torch.npu.synchronize()
|
||||
|
||||
|
||||
class NPUMLATokenToKVPool(MLATokenToKVPool):
|
||||
|
||||
@@ -402,3 +460,64 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
|
||||
loc.view(-1, 1),
|
||||
index_k.view(-1, 1, self.index_head_dim),
|
||||
)
|
||||
|
||||
def _chunk_copy_npu_to_cpu(self, buf_of_layers, indices):
|
||||
chunk_size = self.cpu_offloading_chunk_size
|
||||
out = []
|
||||
for tensors_per_layer in buf_of_layers: # [k_buf, v_buf, ik_buf/None]
|
||||
layer_chunks = []
|
||||
for i in range(0, len(indices), chunk_size):
|
||||
ci = indices[i : i + chunk_size]
|
||||
layer_chunks.append(
|
||||
[
|
||||
t[ci].to("cpu", non_blocking=True)
|
||||
for t in tensors_per_layer
|
||||
if t is not None
|
||||
]
|
||||
)
|
||||
out.append(layer_chunks)
|
||||
return out
|
||||
|
||||
def get_cpu_copy(self, indices):
|
||||
torch.npu.synchronize()
|
||||
buf_of_layers = []
|
||||
has_ik = self.index_head_dim is not None
|
||||
for local_layer_id in range(self.layer_num):
|
||||
k_layer = self.k_buffer[local_layer_id].view(-1, 1, self.kv_lora_rank)
|
||||
v_layer = self.v_buffer[local_layer_id].view(-1, 1, self.qk_rope_head_dim)
|
||||
ik_layer = (
|
||||
self.index_k_buffer[local_layer_id].view(-1, 1, self.index_head_dim)
|
||||
if has_ik
|
||||
else None
|
||||
)
|
||||
buf_of_layers.append([k_layer, v_layer, ik_layer])
|
||||
|
||||
kv_cache_cpu = self._chunk_copy_npu_to_cpu(buf_of_layers, indices)
|
||||
torch.npu.synchronize()
|
||||
return kv_cache_cpu
|
||||
|
||||
def load_cpu_copy(self, kv_cache_cpu, indices):
|
||||
torch.npu.synchronize()
|
||||
chunk_size = self.cpu_offloading_chunk_size
|
||||
has_ik = self.index_head_dim is not None
|
||||
for local_layer_id in range(self.layer_num):
|
||||
k_layer = self.k_buffer[local_layer_id].view(-1, 1, self.kv_lora_rank)
|
||||
v_layer = self.v_buffer[local_layer_id].view(-1, 1, self.qk_rope_head_dim)
|
||||
ik_layer = (
|
||||
self.index_k_buffer[local_layer_id].view(-1, 1, self.index_head_dim)
|
||||
if has_ik
|
||||
else None
|
||||
)
|
||||
for i in range(0, len(indices), chunk_size):
|
||||
chunk_indices = indices[i : i + chunk_size]
|
||||
chunk = kv_cache_cpu[local_layer_id][i // chunk_size]
|
||||
k_cpu, v_cpu = chunk[0], chunk[1]
|
||||
assert k_cpu.shape[0] == len(chunk_indices)
|
||||
k_layer[chunk_indices] = k_cpu.to(k_layer.device, non_blocking=True)
|
||||
v_layer[chunk_indices] = v_cpu.to(v_layer.device, non_blocking=True)
|
||||
if has_ik:
|
||||
ik_cpu = chunk[2]
|
||||
ik_layer[chunk_indices] = ik_cpu.to(
|
||||
ik_layer.device, non_blocking=True
|
||||
)
|
||||
torch.npu.synchronize()
|
||||
|
||||
Reference in New Issue
Block a user