Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 9 additions & 6 deletions python/sglang/srt/entrypoints/engine_score_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,8 +56,10 @@ def score(

Setwise scoring is expressed via ``score_extraction_token_id``: pass a single
item containing the whole candidate block with one extraction token per
candidate, and the head is pooled AT those positions, so ``scores`` becomes
the ``[N x num_labels]`` per-candidate matrix.
candidate, and the readout is taken AT those positions, so ``scores``
becomes the ``[N x num_labels]`` per-candidate matrix. SequenceClassification
pools the head there; CausalLM reads label-token logprobs there (batched
path only).

Args:
query: The query text or pre-tokenized token IDs.
Expand All @@ -69,10 +71,11 @@ def score(
embed_override_token_id: Placeholder token ID used to locate override positions.
query_embed_overrides: Embedding vectors replacing placeholder tokens in query.
item_embed_overrides: Per-item embedding vectors replacing placeholder tokens in items.
score_extraction_token_id: SequenceClassification-only. When set, pool the
head at every occurrence of this token per sequence instead of the last
token; ``scores`` becomes nested — one ``[Ni x num_labels]`` matrix per
item (``len(scores) == len(items)``), where ``Ni`` is the number of
score_extraction_token_id: When set, read the score head (SeqCls) or
label-token logprobs (CausalLM) at every occurrence of
this token per sequence instead of the last token; ``scores`` becomes
nested — one ``[Ni x num_labels]`` matrix per item
(``len(scores) == len(items)``), where ``Ni`` is the number of
extraction tokens (candidates) in item ``i``.
return_pooled_hidden_states: Whether to include raw pooled transformer
hidden states (before the task head) in the result. Only supported
Expand Down
9 changes: 5 additions & 4 deletions python/sglang/srt/entrypoints/openai/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -1408,10 +1408,11 @@ class ScoringRequest(BaseModel):
item_first: bool = False
return_pooled_hidden_states: bool = False

# Setwise readout (SequenceClassification-only): when set, the head is pooled
# AT every occurrence of this token in each `query + item` sequence instead of
# the last token, and `scores` is returned nested (one `[Nᵢ x num_labels]`
# matrix per item). --enable-mis fuses items; otherwise each is scored alone.
# Setwise readout: when set, the readout is taken AT every occurrence of this
# token in each `query + item` sequence instead of the last token, and `scores`
# is returned nested (one `[Nᵢ x num_labels]` matrix per item). SeqCls pools the
# head there; CausalLM reads label-token logprobs there. Both support batched
# and --enable-mis.
score_extraction_token: Optional[str] = None

model: str = DEFAULT_MODEL_NAME
Expand Down
99 changes: 98 additions & 1 deletion python/sglang/srt/layers/logits_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -505,10 +505,12 @@ def forward(
aux_hidden_states: Optional[AuxHiddenStates] = None,
hidden_states_before_norm: Optional[torch.Tensor] = None,
) -> LogitsProcessorOutput:
# Extract MIS indices before ForwardBatch → LogitsMetadata conversion
# Extract MIS / setwise indices before ForwardBatch → LogitsMetadata conversion
multi_item_delimiter_indices = None
token_indices_to_pool = None
if isinstance(logits_metadata, ForwardBatch):
multi_item_delimiter_indices = logits_metadata.multi_item_delimiter_indices
token_indices_to_pool = logits_metadata.token_indices_to_pool
logits_metadata = LogitsMetadata.from_forward_batch(logits_metadata)

# Autotune dummy run discards this output. `is False` not `not`: None
Expand All @@ -526,6 +528,18 @@ def forward(
forward_mode=logits_metadata.forward_mode,
)

# Setwise scoring (CausalLM): read label-token logprobs AT each anchor
# position instead of the last token. Takes precedence over the MIS
# delimiter path; the two are mutually exclusive for generation.
if token_indices_to_pool is not None and logits_metadata.is_prefill_only:
return self.compute_logprobs_at_positions(
input_ids,
hidden_states,
lm_head,
logits_metadata,
token_indices_to_pool,
)

# Multi-item scoring only for prefill-only requests with pre-computed indices.
if multi_item_delimiter_indices is not None and logits_metadata.is_prefill_only:
return self.compute_logprobs_for_multi_item_scoring(
Expand Down Expand Up @@ -1252,6 +1266,89 @@ def compute_logprobs_for_multi_item_scoring(
mm_input_embeds=logits_metadata.mm_input_embeds,
)

def compute_logprobs_at_positions(
self,
input_ids,
hidden_states,
lm_head: VocabParallelEmbedding,
logits_metadata: Union[LogitsMetadata, ForwardBatch],
token_indices_to_pool: List[torch.Tensor],
):
"""Compute label-token logprobs AT each requested position (setwise, CausalLM).

Mirrors ``compute_logprobs_for_multi_item_scoring`` but reads the LM head
AT ``token_indices_to_pool`` (no delimiter - 1 shift, no discarded row):
each anchor's logprobs are P(next token | prefix up to and including the
anchor), one row per anchor.
"""
device = input_ids.device
all_tensors = []
if logits_metadata.extend_seq_lens_cpu is not None:
offset = 0
for req_seq_len, indices_tensor in zip(
logits_metadata.extend_seq_lens_cpu, token_indices_to_pool
):
if len(indices_tensor) > 0:
all_tensors.append(indices_tensor + offset)
offset += req_seq_len
else:
all_tensors.append(token_indices_to_pool[0])
pooled_indices = torch.cat(all_tensors).to(device, non_blocking=True)

sliced_hidden = hidden_states[pooled_indices]
sliced_logits = self._get_logits(sliced_hidden, lm_head, logits_metadata)
sliced_logprobs = torch.nn.functional.log_softmax(sliced_logits, dim=-1)

input_token_ids_logprobs_val = []
input_token_ids_logprobs_idx = []
input_top_logprobs_val = None
input_top_logprobs_idx = None

if (
logits_metadata.token_ids_logprobs
or logits_metadata.extend_return_top_logprob
):
logits_metadata.extend_logprob_pruned_lens_cpu = [
len(t) for t in token_indices_to_pool
]

if logits_metadata.extend_token_ids_logprob:
(
input_token_ids_logprobs_val,
input_token_ids_logprobs_idx,
) = get_token_ids_logprobs_raw(
sliced_logprobs,
logits_metadata.token_ids_logprobs,
stage=LogprobStage.PREFILL,
extend_logprob_pruned_lens_cpu=logits_metadata.extend_logprob_pruned_lens_cpu,
no_copy_to_cpu=True,
)

if logits_metadata.extend_return_top_logprob:
(
input_top_logprobs_val,
input_top_logprobs_idx,
) = get_top_logprobs_raw(
sliced_logprobs,
logits_metadata.top_logprobs_nums,
stage=LogprobStage.PREFILL,
extend_logprob_pruned_lens_cpu=logits_metadata.extend_logprob_pruned_lens_cpu,
)

# Zeros to satisfy the shared logprob pipeline's non-None / length asserts;
# score_request() reads only input_token_ids_logprobs_val (see the MIS path).
input_token_logprobs = torch.zeros(pooled_indices.shape[0], device=device)

return LogitsProcessorOutput(
next_token_logits=None,
input_token_logprobs=input_token_logprobs,
input_top_logprobs_val=input_top_logprobs_val,
input_top_logprobs_idx=input_top_logprobs_idx,
input_token_ids_logprobs_val=input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=input_token_ids_logprobs_idx,
mm_input_embeds=logits_metadata.mm_input_embeds,
)


def _reassemble_tp_lm_head_all_to_all_output(
all_to_all_output: torch.Tensor,
Expand Down
13 changes: 13 additions & 0 deletions python/sglang/srt/managers/io_struct.py
Original file line number Diff line number Diff line change
Expand Up @@ -357,6 +357,11 @@ class GenerateReqInput:
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None

# Token positions for setwise pooling readout (CausalLM: label-token logprobs
# are read AT these positions instead of the last token).
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
token_indices_to_pool: Optional[Union[List[List[int]], List[int]]] = None

# Cache namespace used to isolate otherwise-identical prefixes.
cache_salt: Optional[Union[List[str], str]] = None

Expand Down Expand Up @@ -1023,6 +1028,11 @@ def __getitem__(self, i):
if self.multi_item_delimiter_indices is not None
else None
),
token_indices_to_pool=(
self.token_indices_to_pool[i]
if self.token_indices_to_pool is not None
else None
),
)
cache[i] = sub
return sub
Expand Down Expand Up @@ -1120,6 +1130,9 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True):
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None

# Token positions for setwise pooling readout (CausalLM)
token_indices_to_pool: Optional[List[int]] = None

# For observability
# Pickled Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]]
time_stats: Optional[PickleWrapper] = None
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/managers/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -2798,6 +2798,7 @@ def handle_generate_request(
dllm_config=self.dllm_config,
time_stats=recv_req.time_stats,
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
token_indices_to_pool=recv_req.token_indices_to_pool,
)
req.tokenizer = self.tokenizer

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,24 +27,24 @@ 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)
uses_scoring_positions = self._uses_scoring_positions(req)

# Process logprob values - handle multi-item scoring vs regular requests
if is_multi_item_scoring:
# Multi-item scoring: use all logprobs as-is
# Process logprob values - handle position-based scoring vs regular requests
if uses_scoring_positions:
# Position-based scoring: use all logprobs as-is
req.logprob.input_token_logprobs_val = input_token_logprobs
else:
# Regular request: add None at start, remove last (sampling token)
req.logprob.input_token_logprobs_val = [None] + input_token_logprobs[:-1]

# Process logprob indices based on scoring type
if is_multi_item_scoring:
# MIS scores come from input_token_ids_logprobs, not input_token_logprobs.
if uses_scoring_positions:
# Position-based scores come from input_token_ids_logprobs, not input_token_logprobs.
# But the shared pipeline requires input_token_logprobs_idx to be the same
# length as input_token_logprobs_val (validated at line 816). We fill with
# MIS_DELIMITER_TOKEN_ID as a dummy — score_request() ignores this field.
delimiter_count = len(req.multi_item_delimiter_indices)
input_token_logprobs_idx = [MIS_DELIMITER_TOKEN_ID] * delimiter_count
position_count = len(self._scoring_positions(req))
input_token_logprobs_idx = [MIS_DELIMITER_TOKEN_ID] * position_count
else:
# Regular request: include all tokens from logprob_start_len onwards
input_token_logprobs_idx = req.origin_input_ids[req.logprob_start_len :]
Expand All @@ -60,11 +60,11 @@ def _process_input_top_logprobs(self, req: Req) -> None:
if req.logprob.top_logprobs_num <= 0:
return

is_multi_item_scoring = self._is_multi_item_scoring(req)
uses_scoring_positions = self._uses_scoring_positions(req)

# Initialize arrays - multi-item scoring starts empty, others start with None
req.logprob.input_top_logprobs_val = [] if is_multi_item_scoring else [None]
req.logprob.input_top_logprobs_idx = [] if is_multi_item_scoring else [None]
# Initialize arrays - position-based scoring starts empty, others start with None
req.logprob.input_top_logprobs_val = [] if uses_scoring_positions else [None]
req.logprob.input_top_logprobs_idx = [] if uses_scoring_positions else [None]

# Extend arrays with temp values
for val, idx in zip(
Expand All @@ -75,8 +75,8 @@ def _process_input_top_logprobs(self, req: Req) -> None:
req.logprob.input_top_logprobs_val.extend(val)
req.logprob.input_top_logprobs_idx.extend(idx)

# Remove last token (sampling token) for non multi-item scoring requests
if not is_multi_item_scoring:
# Remove last token (sampling token) for non position-based scoring requests
if not uses_scoring_positions:
req.logprob.input_top_logprobs_val.pop()
req.logprob.input_top_logprobs_idx.pop()

Expand Down Expand Up @@ -118,14 +118,14 @@ def _process_input_token_ids_logprobs(self, req: Req) -> None:
if req.logprob.token_ids_logprob is None:
return

is_multi_item_scoring = self._is_multi_item_scoring(req)
uses_scoring_positions = self._uses_scoring_positions(req)

# Initialize arrays - multi-item scoring starts empty, others start with None
# Initialize arrays - position-based scoring starts empty, others start with None
req.logprob.input_token_ids_logprobs_val = (
[] if is_multi_item_scoring else [None]
[] if uses_scoring_positions else [None]
)
req.logprob.input_token_ids_logprobs_idx = (
[] if is_multi_item_scoring else [None]
[] if uses_scoring_positions else [None]
)

# Process temp values - convert tensors to lists and extend arrays
Expand All @@ -140,8 +140,8 @@ def _process_input_token_ids_logprobs(self, req: Req) -> None:
)
req.logprob.input_token_ids_logprobs_idx.extend(idx)

# Remove last token (sampling token) for non multi-item scoring requests
if not is_multi_item_scoring:
# Remove last token (sampling token) for non position-based scoring requests
if not uses_scoring_positions:
req.logprob.input_token_ids_logprobs_val.pop()
req.logprob.input_token_ids_logprobs_idx.pop()

Expand All @@ -150,15 +150,16 @@ def _process_input_token_ids_logprobs(self, req: Req) -> 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.
"""Calculate the expected length of logprob arrays based on whether position-based scoring is used.

For multi-item scoring, only delimiter positions have logprobs.
For regular requests, all positions from logprob_start_len onwards have logprobs.
For position-based scoring (MIS delimiters or setwise anchors), only those
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)
positions = self._scoring_positions(req)

if is_multi_item_scoring:
return len(req.multi_item_delimiter_indices)
if positions is not None:
return len(positions)
else:
return len(req.origin_input_ids[req.logprob_start_len :])

Expand All @@ -168,36 +169,53 @@ def calculate_num_input_logprobs(
extend_input_len: int,
extend_logprob_start_len: int,
) -> int:
"""Calculate the number of input logprobs based on whether multi-item scoring is enabled.
"""Calculate the number of input logprobs based on whether position-based scoring is used.

For multi-item scoring, only delimiter positions have logprobs.
For regular requests, all positions in the range have logprobs.
For position-based scoring (MIS delimiters or setwise anchors), only those
positions have logprobs. For regular requests, all positions in the range
have logprobs.
"""
is_multi_item_scoring = self._is_multi_item_scoring(req)
positions = self._scoring_positions(req)

if is_multi_item_scoring:
# Count pre-computed delimiter indices within the extend range
if positions is not None:
# Count pre-computed scoring positions within the extend range
return sum(
1
for idx in req.multi_item_delimiter_indices
for idx in positions
if extend_logprob_start_len <= idx < extend_input_len
)
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.
def _scoring_positions(self, req: Req):
"""Position-based scoring readout indices for a prefill-only request, else None.

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.
Both MIS (delimiter positions, --enable-mis) and setwise scoring
(token_indices_to_pool, per-anchor readout) compute label-token logprobs
at a fixed set of positions instead of at every input token; the logprob
bookkeeping below is identical for the two.
"""
return (
if not req.is_prefill_only:
return None
# Setwise anchors take precedence over MIS delimiters: a fused setwise
# request carries both, and its logprobs are read at the anchors.
if req.token_indices_to_pool is not None:
return req.token_indices_to_pool
if (
get_exec().features.enable_mis
and req.is_prefill_only
and req.multi_item_delimiter_indices is not None
)
):
return req.multi_item_delimiter_indices
return None

def _uses_scoring_positions(self, req: Req) -> bool:
"""Whether the request uses position-based scoring (MIS or setwise).

In this mode only the scoring positions receive logprobs, so the shared
logprob pipeline skips the None prefix and last-token pop.
"""
return self._scoring_positions(req) is not None

def add_input_logprob_return_values(
self,
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/managers/tokenizer_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -1533,6 +1533,7 @@ def _create_tokenized_object(
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
num_items_assigned=obj.num_items_assigned,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
token_indices_to_pool=obj.token_indices_to_pool,
encoder_urls=obj.encoder_urls,
)
elif isinstance(obj, EmbeddingReqInput):
Expand Down
Loading
Loading