Support return_logprob for spec v2 (overlap safe) (#19801)

Co-authored-by: Ratish1 <ratish1501@gmail.com>
Co-authored-by: Ratish1 <formula733@gmail.com>
Co-authored-by: hnyls2002 <lsyincs@gmail.com>
This commit is contained in:
Qiaolin Yu
2026-03-10 15:38:27 -07:00
committed by GitHub
parent 76ee4bb98c
commit 09a118fafe
6 changed files with 314 additions and 38 deletions

View File

@@ -1055,7 +1055,7 @@ class LogitsProcessor(nn.Module):
input_token_ids_logprobs_val,
input_token_ids_logprobs_idx,
) = get_token_ids_logprobs_prefill(
sliced_logprobs, logits_metadata, delay_cpu_copy=True
sliced_logprobs, logits_metadata, no_copy_to_cpu=True
)
# Get the logprob of top-k tokens

View File

@@ -68,11 +68,13 @@ def get_top_logprobs_raw(
top_logprobs_nums: List[int],
stage: LogprobStage,
extend_logprob_pruned_lens_cpu: Optional[List[int]] = None,
no_copy_to_cpu: bool = False,
):
max_k = max(top_logprobs_nums)
values, indices = logprobs.topk(max_k, dim=-1)
values = values.tolist()
indices = indices.tolist()
if not no_copy_to_cpu:
values = values.tolist()
indices = indices.tolist()
top_logprobs_val = []
top_logprobs_idx = []
@@ -110,57 +112,73 @@ def get_top_logprobs_prefill(
def get_top_logprobs(
logprobs: torch.Tensor,
top_logprobs_nums: List[int],
no_copy_to_cpu: bool = False,
):
return get_top_logprobs_raw(logprobs, top_logprobs_nums, stage=LogprobStage.DECODE)
return get_top_logprobs_raw(
logprobs,
top_logprobs_nums,
stage=LogprobStage.DECODE,
no_copy_to_cpu=no_copy_to_cpu,
)
def get_token_ids_logprobs_raw(
logprobs: torch.Tensor,
token_ids_logprobs: List[Optional[List[int]]],
token_ids_logprobs_list: List[Optional[List[int]]],
stage: LogprobStage,
extend_logprob_pruned_lens_cpu: Optional[List[int]] = None,
delay_cpu_copy: bool = False,
no_copy_to_cpu: bool = False,
):
vals, idxs = [], []
if stage == LogprobStage.DECODE:
for i, token_ids in enumerate(token_ids_logprobs):
for i, token_ids in enumerate(token_ids_logprobs_list):
if token_ids is None:
vals.append([])
idxs.append([])
else:
vals.append(logprobs[i, token_ids].tolist())
token_ids_tensor = torch.tensor(token_ids, dtype=torch.long).to(
logprobs.device, non_blocking=True
)
row = logprobs[i, token_ids_tensor]
vals.append(row if no_copy_to_cpu else row.tolist())
idxs.append(token_ids)
else: # prefill
pt = 0
for token_ids, pruned_len in zip(
token_ids_logprobs, extend_logprob_pruned_lens_cpu
for i, (token_ids, pruned_len) in enumerate(
zip(token_ids_logprobs_list, extend_logprob_pruned_lens_cpu)
):
if pruned_len <= 0:
vals.append([])
idxs.append([])
continue
pos_logprobs = logprobs[pt : pt + pruned_len, token_ids]
vals.append(pos_logprobs if delay_cpu_copy else pos_logprobs.tolist())
token_ids_tensor = torch.tensor(token_ids, dtype=torch.long).to(
logprobs.device, non_blocking=True
)
pos_logprobs = logprobs[pt : pt + pruned_len, token_ids_tensor]
vals.append(pos_logprobs if no_copy_to_cpu else pos_logprobs.tolist())
idxs.append([token_ids for _ in range(pruned_len)])
pt += pruned_len
return vals, idxs
def get_token_ids_logprobs_prefill(
all_logprobs, logits_metadata: LogitsMetadata, delay_cpu_copy=False
all_logprobs, logits_metadata: LogitsMetadata, no_copy_to_cpu=False
):
return get_token_ids_logprobs_raw(
all_logprobs,
logits_metadata.token_ids_logprobs,
stage=LogprobStage.PREFILL,
extend_logprob_pruned_lens_cpu=logits_metadata.extend_logprob_pruned_lens_cpu,
delay_cpu_copy=delay_cpu_copy,
no_copy_to_cpu=no_copy_to_cpu,
)
def get_token_ids_logprobs(logprobs, token_ids_logprobs):
def get_token_ids_logprobs(logprobs, token_ids_logprobs, no_copy_to_cpu=False):
return get_token_ids_logprobs_raw(
logprobs, token_ids_logprobs, stage=LogprobStage.DECODE
logprobs,
token_ids_logprobs,
stage=LogprobStage.DECODE,
no_copy_to_cpu=no_copy_to_cpu,
)