Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b870c8edb4 |
@@ -52,8 +52,8 @@ class AllPool(TokenPoolingMethod):
|
|||||||
# DispatchPooler passes the full hidden_states tensor.
|
# DispatchPooler passes the full hidden_states tensor.
|
||||||
# slice out the subgroup once, then split it by
|
# slice out the subgroup once, then split it by
|
||||||
# per-request token counts
|
# per-request token counts
|
||||||
group_start = int(pooling_cursor.first_token_indices_gpu[0].item())
|
group_start = int(pooling_cursor.first_token_indices_cpu[0])
|
||||||
group_end = int(pooling_cursor.last_token_indices_gpu[-1].item()) + 1
|
group_end = group_start + sum(split_sizes)
|
||||||
hidden_states_group = hidden_states[group_start:group_end]
|
hidden_states_group = hidden_states[group_start:group_end]
|
||||||
hidden_states_lst = list(hidden_states_group.split(split_sizes))
|
hidden_states_lst = list(hidden_states_group.split(split_sizes))
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ pin_memory = is_pin_memory_available()
|
|||||||
class PoolingCursor:
|
class PoolingCursor:
|
||||||
first_token_indices_gpu: torch.Tensor
|
first_token_indices_gpu: torch.Tensor
|
||||||
last_token_indices_gpu: torch.Tensor
|
last_token_indices_gpu: torch.Tensor
|
||||||
|
first_token_indices_cpu: torch.Tensor
|
||||||
prompt_lens_cpu: torch.Tensor
|
prompt_lens_cpu: torch.Tensor
|
||||||
seq_lens_cpu: torch.Tensor
|
seq_lens_cpu: torch.Tensor
|
||||||
num_scheduled_tokens_cpu: torch.Tensor
|
num_scheduled_tokens_cpu: torch.Tensor
|
||||||
@@ -24,6 +25,7 @@ class PoolingCursor:
|
|||||||
return PoolingCursor(
|
return PoolingCursor(
|
||||||
first_token_indices_gpu=self.first_token_indices_gpu[indices],
|
first_token_indices_gpu=self.first_token_indices_gpu[indices],
|
||||||
last_token_indices_gpu=self.last_token_indices_gpu[indices],
|
last_token_indices_gpu=self.last_token_indices_gpu[indices],
|
||||||
|
first_token_indices_cpu=self.first_token_indices_cpu[indices],
|
||||||
prompt_lens_cpu=self.prompt_lens_cpu[indices],
|
prompt_lens_cpu=self.prompt_lens_cpu[indices],
|
||||||
seq_lens_cpu=self.seq_lens_cpu[indices],
|
seq_lens_cpu=self.seq_lens_cpu[indices],
|
||||||
num_scheduled_tokens_cpu=self.num_scheduled_tokens_cpu[indices],
|
num_scheduled_tokens_cpu=self.num_scheduled_tokens_cpu[indices],
|
||||||
@@ -117,12 +119,13 @@ class PoolingMetadata:
|
|||||||
assert len(prompt_lens) == n_seq
|
assert len(prompt_lens) == n_seq
|
||||||
|
|
||||||
num_scheduled_tokens_cpu = torch.from_numpy(num_scheduled_tokens_np)
|
num_scheduled_tokens_cpu = torch.from_numpy(num_scheduled_tokens_np)
|
||||||
|
cumsum_cpu = torch.zeros(
|
||||||
|
n_seq + 1, dtype=torch.int64, pin_memory=pin_memory, device="cpu"
|
||||||
|
)
|
||||||
|
torch.cumsum(num_scheduled_tokens_cpu, dim=0, out=cumsum_cpu[1:])
|
||||||
|
|
||||||
if query_start_loc_gpu is None:
|
if query_start_loc_gpu is None:
|
||||||
cumsum = torch.zeros(
|
cumsum_gpu = cumsum_cpu.to(device, non_blocking=True)
|
||||||
n_seq + 1, dtype=torch.int64, pin_memory=pin_memory, device="cpu"
|
|
||||||
)
|
|
||||||
torch.cumsum(num_scheduled_tokens_cpu, dim=0, out=cumsum[1:])
|
|
||||||
cumsum = cumsum.to(device, non_blocking=True)
|
|
||||||
else:
|
else:
|
||||||
if query_start_loc_gpu.shape[0] != n_seq + 1:
|
if query_start_loc_gpu.shape[0] != n_seq + 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -135,10 +138,11 @@ class PoolingMetadata:
|
|||||||
"query_start_loc_gpu must be on the same device as the "
|
"query_start_loc_gpu must be on the same device as the "
|
||||||
f"hidden states: {query_start_loc_gpu.device} != {device}."
|
f"hidden states: {query_start_loc_gpu.device} != {device}."
|
||||||
)
|
)
|
||||||
cumsum = query_start_loc_gpu
|
cumsum_gpu = query_start_loc_gpu
|
||||||
self.pooling_cursor = PoolingCursor(
|
self.pooling_cursor = PoolingCursor(
|
||||||
first_token_indices_gpu=cumsum[:n_seq],
|
first_token_indices_gpu=cumsum_gpu[:n_seq],
|
||||||
last_token_indices_gpu=cumsum[1:] - 1,
|
last_token_indices_gpu=cumsum_gpu[1:] - 1,
|
||||||
|
first_token_indices_cpu=cumsum_cpu[:n_seq],
|
||||||
prompt_lens_cpu=prompt_lens,
|
prompt_lens_cpu=prompt_lens,
|
||||||
seq_lens_cpu=seq_lens_cpu,
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
num_scheduled_tokens_cpu=num_scheduled_tokens_cpu,
|
num_scheduled_tokens_cpu=num_scheduled_tokens_cpu,
|
||||||
|
|||||||
Reference in New Issue
Block a user