[Generative Score API] Multi-Item scoring with custom attention mask. (#10979)
This commit is contained in:
@@ -104,7 +104,10 @@ class SchedulerOutputProcessorMixin:
|
||||
assert extend_input_len_per_req is not None
|
||||
extend_logprob_start_len = extend_logprob_start_len_per_req[i]
|
||||
extend_input_len = extend_input_len_per_req[i]
|
||||
num_input_logprobs = extend_input_len - extend_logprob_start_len
|
||||
|
||||
num_input_logprobs = self._calculate_num_input_logprobs(
|
||||
req, extend_input_len, extend_logprob_start_len
|
||||
)
|
||||
|
||||
if req.return_logprob:
|
||||
self.add_logprob_return_values(
|
||||
@@ -159,8 +162,8 @@ class SchedulerOutputProcessorMixin:
|
||||
extend_input_len = extend_input_len_per_req[i]
|
||||
if extend_logprob_start_len < extend_input_len:
|
||||
# Update input logprobs.
|
||||
num_input_logprobs = (
|
||||
extend_input_len - extend_logprob_start_len
|
||||
num_input_logprobs = self._calculate_num_input_logprobs(
|
||||
req, extend_input_len, extend_logprob_start_len
|
||||
)
|
||||
if req.return_logprob:
|
||||
self.add_input_logprob_return_values(
|
||||
@@ -303,6 +306,153 @@ class SchedulerOutputProcessorMixin:
|
||||
):
|
||||
self.log_decode_stats(can_run_cuda_graph, running_batch=batch)
|
||||
|
||||
def _process_input_token_logprobs(
|
||||
self, req: Req, input_token_logprobs: List
|
||||
) -> None:
|
||||
"""Process input token logprobs values and indices."""
|
||||
is_multi_item_scoring = self._is_multi_item_scoring(req)
|
||||
|
||||
# Process logprob values - handle multi-item scoring vs regular requests
|
||||
if is_multi_item_scoring:
|
||||
# Multi-item scoring: use all logprobs as-is
|
||||
req.input_token_logprobs_val = input_token_logprobs
|
||||
else:
|
||||
# Regular request: add None at start, remove last (sampling token)
|
||||
req.input_token_logprobs_val = [None] + input_token_logprobs[:-1]
|
||||
|
||||
# Process logprob indices based on scoring type
|
||||
if is_multi_item_scoring:
|
||||
# Multi-item scoring: only include delimiter token positions
|
||||
relevant_tokens = req.origin_input_ids[req.logprob_start_len :]
|
||||
input_token_logprobs_idx = [
|
||||
token_id
|
||||
for token_id in relevant_tokens
|
||||
if token_id == self.server_args.multi_item_scoring_delimiter
|
||||
]
|
||||
else:
|
||||
# Regular request: include all tokens from logprob_start_len onwards
|
||||
input_token_logprobs_idx = req.origin_input_ids[req.logprob_start_len :]
|
||||
|
||||
# Clip padded hash values from image tokens to prevent detokenization errors
|
||||
req.input_token_logprobs_idx = [
|
||||
x if x < self.model_config.vocab_size - 1 else 0
|
||||
for x in input_token_logprobs_idx
|
||||
]
|
||||
|
||||
def _process_input_top_logprobs(self, req: Req) -> None:
|
||||
"""Process input top logprobs."""
|
||||
if req.top_logprobs_num <= 0:
|
||||
return
|
||||
|
||||
is_multi_item_scoring = self._is_multi_item_scoring(req)
|
||||
|
||||
# Initialize arrays - multi-item scoring starts empty, others start with None
|
||||
req.input_top_logprobs_val = [] if is_multi_item_scoring else [None]
|
||||
req.input_top_logprobs_idx = [] if is_multi_item_scoring else [None]
|
||||
|
||||
# Extend arrays with temp values
|
||||
for val, idx in zip(
|
||||
req.temp_input_top_logprobs_val,
|
||||
req.temp_input_top_logprobs_idx,
|
||||
strict=True,
|
||||
):
|
||||
req.input_top_logprobs_val.extend(val)
|
||||
req.input_top_logprobs_idx.extend(idx)
|
||||
|
||||
# Remove last token (sampling token) for non multi-item scoring requests
|
||||
if not is_multi_item_scoring:
|
||||
req.input_top_logprobs_val.pop()
|
||||
req.input_top_logprobs_idx.pop()
|
||||
|
||||
# Clean up temp storage
|
||||
req.temp_input_top_logprobs_idx = None
|
||||
req.temp_input_top_logprobs_val = None
|
||||
|
||||
def _process_input_token_ids_logprobs(self, req: Req) -> None:
|
||||
"""Process input token IDs logprobs."""
|
||||
if req.token_ids_logprob is None:
|
||||
return
|
||||
|
||||
is_multi_item_scoring = self._is_multi_item_scoring(req)
|
||||
|
||||
# Initialize arrays - multi-item scoring starts empty, others start with None
|
||||
req.input_token_ids_logprobs_val = [] if is_multi_item_scoring else [None]
|
||||
req.input_token_ids_logprobs_idx = [] if is_multi_item_scoring else [None]
|
||||
|
||||
# Process temp values - convert tensors to lists and extend arrays
|
||||
for val, idx in zip(
|
||||
req.temp_input_token_ids_logprobs_val,
|
||||
req.temp_input_token_ids_logprobs_idx,
|
||||
strict=True,
|
||||
):
|
||||
val_list = val.tolist() if isinstance(val, torch.Tensor) else val
|
||||
req.input_token_ids_logprobs_val.extend(
|
||||
val_list if isinstance(val_list, list) else [val_list]
|
||||
)
|
||||
req.input_token_ids_logprobs_idx.extend(idx)
|
||||
|
||||
# Remove last token (sampling token) for non multi-item scoring requests
|
||||
if not is_multi_item_scoring:
|
||||
req.input_token_ids_logprobs_val.pop()
|
||||
req.input_token_ids_logprobs_idx.pop()
|
||||
|
||||
# Clean up temp storage
|
||||
req.temp_input_token_ids_logprobs_idx = None
|
||||
req.temp_input_token_ids_logprobs_val = None
|
||||
|
||||
def _calculate_relevant_tokens_len(self, req: Req) -> int:
|
||||
"""Calculate the expected length of logprob arrays based on whether multi-item scoring is enabled.
|
||||
|
||||
For multi-item scoring, only delimiter positions have logprobs.
|
||||
For regular requests, all positions from logprob_start_len onwards have logprobs.
|
||||
"""
|
||||
is_multi_item_scoring = self._is_multi_item_scoring(req)
|
||||
|
||||
if is_multi_item_scoring:
|
||||
# Multi-item scoring: count delimiter tokens from logprob_start_len onwards
|
||||
relevant_tokens = req.origin_input_ids[req.logprob_start_len :]
|
||||
return sum(
|
||||
1
|
||||
for token_id in relevant_tokens
|
||||
if token_id == self.server_args.multi_item_scoring_delimiter
|
||||
)
|
||||
else:
|
||||
# Regular request: all tokens from logprob_start_len onwards
|
||||
return len(req.origin_input_ids) - req.logprob_start_len
|
||||
|
||||
def _calculate_num_input_logprobs(
|
||||
self, req: Req, extend_input_len: int, extend_logprob_start_len: int
|
||||
) -> int:
|
||||
"""Calculate the number of input logprobs based on whether multi-item scoring is enabled.
|
||||
|
||||
For multi-item scoring, only delimiter positions have logprobs.
|
||||
For regular requests, all positions in the range have logprobs.
|
||||
"""
|
||||
is_multi_item_scoring = self._is_multi_item_scoring(req)
|
||||
|
||||
if is_multi_item_scoring:
|
||||
# Multi-item scoring: count delimiter tokens in the relevant portion
|
||||
relevant_tokens = req.origin_input_ids[
|
||||
extend_logprob_start_len:extend_input_len
|
||||
]
|
||||
return sum(
|
||||
1
|
||||
for token_id in relevant_tokens
|
||||
if token_id == self.server_args.multi_item_scoring_delimiter
|
||||
)
|
||||
else:
|
||||
# Regular request: all tokens in the range
|
||||
return extend_input_len - extend_logprob_start_len
|
||||
|
||||
def _is_multi_item_scoring(self, req: Req) -> bool:
|
||||
"""Check if request uses multi-item scoring.
|
||||
|
||||
Multi-item scoring applies to prefill-only requests when a delimiter
|
||||
token is configured. In this mode, only positions containing the
|
||||
delimiter token receive logprobs.
|
||||
"""
|
||||
return req.is_prefill_only and self.server_args.multi_item_scoring_delimiter
|
||||
|
||||
def add_input_logprob_return_values(
|
||||
self: Scheduler,
|
||||
i: int,
|
||||
@@ -371,63 +521,14 @@ class SchedulerOutputProcessorMixin:
|
||||
assert req.input_top_logprobs_val is None
|
||||
assert req.input_top_logprobs_idx is None
|
||||
|
||||
# Compute input_token_logprobs_val
|
||||
# Always pad the first one with None.
|
||||
req.input_token_logprobs_val = [None]
|
||||
req.input_token_logprobs_val.extend(input_token_logprobs)
|
||||
# The last input logprob is for sampling, so just pop it out.
|
||||
req.input_token_logprobs_val.pop()
|
||||
# Process all input logprob types using helper functions
|
||||
self._process_input_token_logprobs(req, input_token_logprobs)
|
||||
self._process_input_top_logprobs(req)
|
||||
|
||||
# Compute input_token_logprobs_idx
|
||||
input_token_logprobs_idx = req.origin_input_ids[req.logprob_start_len :]
|
||||
# Clip the padded hash values from image tokens.
|
||||
# Otherwise, it will lead to detokenization errors.
|
||||
input_token_logprobs_idx = [
|
||||
x if x < self.model_config.vocab_size - 1 else 0
|
||||
for x in input_token_logprobs_idx
|
||||
]
|
||||
req.input_token_logprobs_idx = input_token_logprobs_idx
|
||||
|
||||
if req.top_logprobs_num > 0:
|
||||
req.input_top_logprobs_val = [None]
|
||||
req.input_top_logprobs_idx = [None]
|
||||
assert len(req.temp_input_token_ids_logprobs_val) == len(
|
||||
req.temp_input_token_ids_logprobs_idx
|
||||
)
|
||||
for val, idx in zip(
|
||||
req.temp_input_top_logprobs_val,
|
||||
req.temp_input_top_logprobs_idx,
|
||||
strict=True,
|
||||
):
|
||||
req.input_top_logprobs_val.extend(val)
|
||||
req.input_top_logprobs_idx.extend(idx)
|
||||
|
||||
# Last token is a sample token.
|
||||
req.input_top_logprobs_val.pop()
|
||||
req.input_top_logprobs_idx.pop()
|
||||
req.temp_input_top_logprobs_idx = None
|
||||
req.temp_input_top_logprobs_val = None
|
||||
|
||||
if req.token_ids_logprob is not None:
|
||||
req.input_token_ids_logprobs_val = [None]
|
||||
req.input_token_ids_logprobs_idx = [None]
|
||||
|
||||
for val, idx in zip(
|
||||
req.temp_input_token_ids_logprobs_val,
|
||||
req.temp_input_token_ids_logprobs_idx,
|
||||
strict=True,
|
||||
):
|
||||
req.input_token_ids_logprobs_val.extend(val)
|
||||
req.input_token_ids_logprobs_idx.extend(idx)
|
||||
|
||||
# Last token is a sample token.
|
||||
req.input_token_ids_logprobs_val.pop()
|
||||
req.input_token_ids_logprobs_idx.pop()
|
||||
req.temp_input_token_ids_logprobs_idx = None
|
||||
req.temp_input_token_ids_logprobs_val = None
|
||||
self._process_input_token_ids_logprobs(req)
|
||||
|
||||
if req.return_logprob:
|
||||
relevant_tokens_len = len(req.origin_input_ids) - req.logprob_start_len
|
||||
relevant_tokens_len = self._calculate_relevant_tokens_len(req)
|
||||
assert len(req.input_token_logprobs_val) == relevant_tokens_len
|
||||
assert len(req.input_token_logprobs_idx) == relevant_tokens_len
|
||||
if req.top_logprobs_num > 0:
|
||||
|
||||
Reference in New Issue
Block a user