[Generative Score API] Multi-Item scoring with custom attention mask. (#10979)

This commit is contained in:
Sundara Raman Ramachandran
2025-10-08 18:47:32 -07:00
committed by GitHub
parent e22b13c569
commit 53bd00d975
10 changed files with 1121 additions and 129 deletions

View File

@@ -382,6 +382,12 @@ class ServerArgs:
offload_prefetch_step: int = 1
offload_mode: str = "cpu"
# Scoring configuration
# Delimiter token ID used to combine Query and Items into a single sequence for multi-item scoring.
# Format: Query<delimiter>Item1<delimiter>Item2<delimiter>...
# This enables efficient batch processing of multiple items against a single query.
multi_item_scoring_delimiter: Optional[Union[int]] = None
# Optimization/debug options
disable_radix_cache: bool = False
cuda_graph_max_bs: Optional[int] = None
@@ -2334,7 +2340,13 @@ class ServerArgs:
choices=["float32", "bfloat16"],
help="The data type of the SSM states in mamba cache.",
)
# Args for multi-item-scoring
parser.add_argument(
"--multi-item-scoring-delimiter",
type=int,
default=ServerArgs.multi_item_scoring_delimiter,
help="Delimiter token ID for multi-item scoring. Used to combine Query and Items into a single sequence: Query<delimiter>Item1<delimiter>Item2<delimiter>... This enables efficient batch processing of multiple items against a single query.",
)
# Hierarchical cache
parser.add_argument(
"--enable-hierarchical-cache",
@@ -3004,6 +3016,17 @@ class ServerArgs:
"lof",
], f"To use priority scheduling, schedule_policy must be 'fcfs' or 'lof'. '{self.schedule_policy}' is not supported."
# Check multi-item scoring
if self.multi_item_scoring_delimiter is not None:
assert self.disable_radix_cache, (
"Multi-item scoring requires radix cache to be disabled. "
"Please set --disable-radix-cache when using --multi-item-scoring-delimiter."
)
assert self.chunked_prefill_size == -1, (
"Multi-item scoring requires chunked prefill to be disabled. "
"Please set --chunked-prefill-size -1 when using --multi-item-scoring-delimiter."
)
def check_lora_server_args(self):
assert self.max_loras_per_batch > 0, "max_loras_per_batch must be positive"