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:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user