[PP] Add pipeline parallelism (#5724)
This commit is contained in:
@@ -214,6 +214,8 @@ class MHATokenToKVPool(KVCache):
|
||||
layer_num: int,
|
||||
device: str,
|
||||
enable_memory_saver: bool,
|
||||
start_layer: Optional[int] = None,
|
||||
end_layer: Optional[int] = None,
|
||||
):
|
||||
self.size = size
|
||||
self.page_size = page_size
|
||||
@@ -232,6 +234,8 @@ class MHATokenToKVPool(KVCache):
|
||||
self.head_dim = head_dim
|
||||
self.layer_num = layer_num
|
||||
self._create_buffers()
|
||||
self.start_layer = start_layer or 0
|
||||
self.end_layer = end_layer or layer_num - 1
|
||||
|
||||
self.layer_transfer_counter = None
|
||||
self.capture_mode = False
|
||||
@@ -281,6 +285,8 @@ class MHATokenToKVPool(KVCache):
|
||||
|
||||
# for disagg
|
||||
def get_contiguous_buf_infos(self):
|
||||
# layer_num x [seq_len, head_num, head_dim]
|
||||
# layer_num x [page_num, page_size, head_num, head_dim]
|
||||
kv_data_ptrs = [
|
||||
self.get_key_buffer(i).data_ptr() for i in range(self.layer_num)
|
||||
] + [self.get_value_buffer(i).data_ptr() for i in range(self.layer_num)]
|
||||
@@ -320,24 +326,24 @@ class MHATokenToKVPool(KVCache):
|
||||
# transfer prepared data from host to device
|
||||
flat_data = flat_data.to(device=self.device, non_blocking=False)
|
||||
k_data, v_data = flat_data[0], flat_data[1]
|
||||
self.k_buffer[layer_id][indices] = k_data
|
||||
self.v_buffer[layer_id][indices] = v_data
|
||||
self.k_buffer[layer_id - self.start_layer][indices] = k_data
|
||||
self.v_buffer[layer_id - self.start_layer][indices] = v_data
|
||||
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id)
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
|
||||
if self.store_dtype != self.dtype:
|
||||
return self.k_buffer[layer_id].view(self.dtype)
|
||||
return self.k_buffer[layer_id]
|
||||
return self.k_buffer[layer_id - self.start_layer].view(self.dtype)
|
||||
return self.k_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_value_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id)
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
|
||||
if self.store_dtype != self.dtype:
|
||||
return self.v_buffer[layer_id].view(self.dtype)
|
||||
return self.v_buffer[layer_id]
|
||||
return self.v_buffer[layer_id - self.start_layer].view(self.dtype)
|
||||
return self.v_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_kv_buffer(self, layer_id: int):
|
||||
return self.get_key_buffer(layer_id), self.get_value_buffer(layer_id)
|
||||
@@ -369,12 +375,12 @@ class MHATokenToKVPool(KVCache):
|
||||
current_stream = self.device_module.current_stream()
|
||||
self.alt_stream.wait_stream(current_stream)
|
||||
with self.device_module.stream(self.alt_stream):
|
||||
self.k_buffer[layer_id][loc] = cache_k
|
||||
self.v_buffer[layer_id][loc] = cache_v
|
||||
self.k_buffer[layer_id - self.start_layer][loc] = cache_k
|
||||
self.v_buffer[layer_id - self.start_layer][loc] = cache_v
|
||||
current_stream.wait_stream(self.alt_stream)
|
||||
else:
|
||||
self.k_buffer[layer_id][loc] = cache_k
|
||||
self.v_buffer[layer_id][loc] = cache_v
|
||||
self.k_buffer[layer_id - self.start_layer][loc] = cache_k
|
||||
self.v_buffer[layer_id - self.start_layer][loc] = cache_v
|
||||
|
||||
|
||||
@torch.compile
|
||||
@@ -484,6 +490,8 @@ class MLATokenToKVPool(KVCache):
|
||||
layer_num: int,
|
||||
device: str,
|
||||
enable_memory_saver: bool,
|
||||
start_layer: Optional[int] = None,
|
||||
end_layer: Optional[int] = None,
|
||||
):
|
||||
self.size = size
|
||||
self.page_size = page_size
|
||||
@@ -497,6 +505,8 @@ class MLATokenToKVPool(KVCache):
|
||||
self.kv_lora_rank = kv_lora_rank
|
||||
self.qk_rope_head_dim = qk_rope_head_dim
|
||||
self.layer_num = layer_num
|
||||
self.start_layer = start_layer or 0
|
||||
self.end_layer = end_layer or layer_num - 1
|
||||
|
||||
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=enable_memory_saver
|
||||
@@ -540,19 +550,21 @@ class MLATokenToKVPool(KVCache):
|
||||
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id)
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
|
||||
if self.store_dtype != self.dtype:
|
||||
return self.kv_buffer[layer_id].view(self.dtype)
|
||||
return self.kv_buffer[layer_id]
|
||||
return self.kv_buffer[layer_id - self.start_layer].view(self.dtype)
|
||||
return self.kv_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_value_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id)
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
|
||||
if self.store_dtype != self.dtype:
|
||||
return self.kv_buffer[layer_id][..., : self.kv_lora_rank].view(self.dtype)
|
||||
return self.kv_buffer[layer_id][..., : self.kv_lora_rank]
|
||||
return self.kv_buffer[layer_id - self.start_layer][
|
||||
..., : self.kv_lora_rank
|
||||
].view(self.dtype)
|
||||
return self.kv_buffer[layer_id - self.start_layer][..., : self.kv_lora_rank]
|
||||
|
||||
def get_kv_buffer(self, layer_id: int):
|
||||
return self.get_key_buffer(layer_id), self.get_value_buffer(layer_id)
|
||||
@@ -568,9 +580,11 @@ class MLATokenToKVPool(KVCache):
|
||||
if cache_k.dtype != self.dtype:
|
||||
cache_k = cache_k.to(self.dtype)
|
||||
if self.store_dtype != self.dtype:
|
||||
self.kv_buffer[layer_id][loc] = cache_k.view(self.store_dtype)
|
||||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k.view(
|
||||
self.store_dtype
|
||||
)
|
||||
else:
|
||||
self.kv_buffer[layer_id][loc] = cache_k
|
||||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
|
||||
|
||||
def set_mla_kv_buffer(
|
||||
self,
|
||||
@@ -605,7 +619,7 @@ class MLATokenToKVPool(KVCache):
|
||||
def transfer_per_layer(self, indices, flat_data, layer_id):
|
||||
# transfer prepared data from host to device
|
||||
flat_data = flat_data.to(device=self.device, non_blocking=False)
|
||||
self.kv_buffer[layer_id][indices] = flat_data
|
||||
self.kv_buffer[layer_id - self.start_layer][indices] = flat_data
|
||||
|
||||
|
||||
class DoubleSparseTokenToKVPool(KVCache):
|
||||
@@ -620,6 +634,8 @@ class DoubleSparseTokenToKVPool(KVCache):
|
||||
device: str,
|
||||
heavy_channel_num: int,
|
||||
enable_memory_saver: bool,
|
||||
start_layer: Optional[int] = None,
|
||||
end_layer: Optional[int] = None,
|
||||
):
|
||||
self.size = size
|
||||
self.page_size = page_size
|
||||
@@ -657,17 +673,23 @@ class DoubleSparseTokenToKVPool(KVCache):
|
||||
for _ in range(layer_num)
|
||||
]
|
||||
|
||||
self.start_layer = start_layer or 0
|
||||
self.end_layer = end_layer or layer_num - 1
|
||||
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
return self.k_buffer[layer_id]
|
||||
return self.k_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_value_buffer(self, layer_id: int):
|
||||
return self.v_buffer[layer_id]
|
||||
return self.v_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_label_buffer(self, layer_id: int):
|
||||
return self.label_buffer[layer_id]
|
||||
return self.label_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_kv_buffer(self, layer_id: int):
|
||||
return self.k_buffer[layer_id], self.v_buffer[layer_id]
|
||||
return (
|
||||
self.k_buffer[layer_id - self.start_layer],
|
||||
self.v_buffer[layer_id - self.start_layer],
|
||||
)
|
||||
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
@@ -679,9 +701,9 @@ class DoubleSparseTokenToKVPool(KVCache):
|
||||
):
|
||||
# NOTE(Andy): ignore the dtype check
|
||||
layer_id = layer.layer_id
|
||||
self.k_buffer[layer_id][loc] = cache_k
|
||||
self.v_buffer[layer_id][loc] = cache_v
|
||||
self.label_buffer[layer_id][loc] = cache_label
|
||||
self.k_buffer[layer_id - self.start_layer][loc] = cache_k
|
||||
self.v_buffer[layer_id - self.start_layer][loc] = cache_v
|
||||
self.label_buffer[layer_id - self.start_layer][loc] = cache_label
|
||||
|
||||
def get_flat_data(self, indices):
|
||||
pass
|
||||
@@ -930,7 +952,7 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
return self.kv_buffer[:, :, indices]
|
||||
|
||||
def get_flat_data_by_layer(self, indices, layer_id):
|
||||
return self.kv_buffer[:, layer_id, indices]
|
||||
return self.kv_buffer[:, layer_id - self.start_layer, indices]
|
||||
|
||||
def assign_flat_data(self, indices, flat_data):
|
||||
self.kv_buffer[:, :, indices] = flat_data
|
||||
@@ -955,12 +977,20 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
for i in range(len(device_indices_cpu)):
|
||||
h_index = host_indices[i * self.page_size]
|
||||
d_index = device_indices_cpu[i]
|
||||
device_pool.k_buffer[layer_id][d_index : d_index + self.page_size].copy_(
|
||||
self.kv_buffer[0, layer_id, h_index : h_index + self.page_size],
|
||||
device_pool.k_buffer[layer_id - self.start_layer][
|
||||
d_index : d_index + self.page_size
|
||||
].copy_(
|
||||
self.kv_buffer[
|
||||
0, layer_id - self.start_layer, h_index : h_index + self.page_size
|
||||
],
|
||||
non_blocking=True,
|
||||
)
|
||||
device_pool.v_buffer[layer_id][d_index : d_index + self.page_size].copy_(
|
||||
self.kv_buffer[1, layer_id, h_index : h_index + self.page_size],
|
||||
device_pool.v_buffer[layer_id - self.start_layer][
|
||||
d_index : d_index + self.page_size
|
||||
].copy_(
|
||||
self.kv_buffer[
|
||||
1, layer_id - self.start_layer, h_index : h_index + self.page_size
|
||||
],
|
||||
non_blocking=True,
|
||||
)
|
||||
|
||||
@@ -1015,7 +1045,7 @@ class MLATokenToKVPoolHost(HostKVCache):
|
||||
return self.kv_buffer[:, indices]
|
||||
|
||||
def get_flat_data_by_layer(self, indices, layer_id):
|
||||
return self.kv_buffer[layer_id, indices]
|
||||
return self.kv_buffer[layer_id - self.start_layer, indices]
|
||||
|
||||
def assign_flat_data(self, indices, flat_data):
|
||||
self.kv_buffer[:, indices] = flat_data
|
||||
@@ -1036,7 +1066,11 @@ class MLATokenToKVPoolHost(HostKVCache):
|
||||
for i in range(len(device_indices_cpu)):
|
||||
h_index = host_indices[i * self.page_size]
|
||||
d_index = device_indices_cpu[i]
|
||||
device_pool.kv_buffer[layer_id][d_index : d_index + self.page_size].copy_(
|
||||
self.kv_buffer[layer_id, h_index : h_index + self.page_size],
|
||||
device_pool.kv_buffer[layer_id - self.start_layer][
|
||||
d_index : d_index + self.page_size
|
||||
].copy_(
|
||||
self.kv_buffer[
|
||||
layer_id - self.start_layer, h_index : h_index + self.page_size
|
||||
],
|
||||
non_blocking=True,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user