Compare commits

...
Author SHA1 Message Date
yewentao256 b870c8edb4 optimize allpool forward
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-05-04 22:47:43 +00:00
2 changed files with 14 additions and 10 deletions
@@ -52,8 +52,8 @@ class AllPool(TokenPoolingMethod):
# DispatchPooler passes the full hidden_states tensor.
# slice out the subgroup once, then split it by
# per-request token counts
group_start = int(pooling_cursor.first_token_indices_gpu[0].item())
group_end = int(pooling_cursor.last_token_indices_gpu[-1].item()) + 1
group_start = int(pooling_cursor.first_token_indices_cpu[0])
group_end = group_start + sum(split_sizes)
hidden_states_group = hidden_states[group_start:group_end]
hidden_states_lst = list(hidden_states_group.split(split_sizes))
else:
+12 -8
View File
@@ -16,6 +16,7 @@ pin_memory = is_pin_memory_available()
class PoolingCursor:
first_token_indices_gpu: torch.Tensor
last_token_indices_gpu: torch.Tensor
first_token_indices_cpu: torch.Tensor
prompt_lens_cpu: torch.Tensor
seq_lens_cpu: torch.Tensor
num_scheduled_tokens_cpu: torch.Tensor
@@ -24,6 +25,7 @@ class PoolingCursor:
return PoolingCursor(
first_token_indices_gpu=self.first_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],
seq_lens_cpu=self.seq_lens_cpu[indices],
num_scheduled_tokens_cpu=self.num_scheduled_tokens_cpu[indices],
@@ -117,12 +119,13 @@ class PoolingMetadata:
assert len(prompt_lens) == n_seq
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:
cumsum = torch.zeros(
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)
cumsum_gpu = cumsum_cpu.to(device, non_blocking=True)
else:
if query_start_loc_gpu.shape[0] != n_seq + 1:
raise ValueError(
@@ -135,10 +138,11 @@ class PoolingMetadata:
"query_start_loc_gpu must be on the same device as the "
f"hidden states: {query_start_loc_gpu.device} != {device}."
)
cumsum = query_start_loc_gpu
cumsum_gpu = query_start_loc_gpu
self.pooling_cursor = PoolingCursor(
first_token_indices_gpu=cumsum[:n_seq],
last_token_indices_gpu=cumsum[1:] - 1,
first_token_indices_gpu=cumsum_gpu[:n_seq],
last_token_indices_gpu=cumsum_gpu[1:] - 1,
first_token_indices_cpu=cumsum_cpu[:n_seq],
prompt_lens_cpu=prompt_lens,
seq_lens_cpu=seq_lens_cpu,
num_scheduled_tokens_cpu=num_scheduled_tokens_cpu,