From 850021378a85a95c39d7148724550a628935b5ec Mon Sep 17 00:00:00 2001 From: jsheng_Linkedin <162757399+fortunecookiee@users.noreply.github.com> Date: Wed, 29 Apr 2026 13:51:53 -0700 Subject: [PATCH] [Score API] Hoist query placeholder scan and specialize PositionalEmbeds stacking (#23513) Co-authored-by: Claude Opus 4.7 (1M context) --- python/sglang/srt/managers/embed_types.py | 14 ++-- .../managers/tokenizer_manager_score_mixin.py | 71 +++++++++++-------- 2 files changed, 52 insertions(+), 33 deletions(-) diff --git a/python/sglang/srt/managers/embed_types.py b/python/sglang/srt/managers/embed_types.py index 987102aed..48d6859d9 100644 --- a/python/sglang/srt/managers/embed_types.py +++ b/python/sglang/srt/managers/embed_types.py @@ -41,11 +41,17 @@ class PositionalEmbeds: positions: List[int] def __post_init__(self): - # Stack list of tensors into a single [N, hidden_dim] tensor + # Normalize list of tensors into a single [N, hidden_dim] tensor. + # Dispatch by element rank to avoid a per-element unsqueeze. if isinstance(self.embeds, list): - self.embeds = torch.cat( - [e if e.dim() == 2 else e.unsqueeze(0) for e in self.embeds], dim=0 - ) + if not self.embeds: + self.embeds = torch.cat(self.embeds, dim=0) # raises — empty is invalid + elif self.embeds[0].dim() == 1: + # [hidden_dim] elements → stack adds the leading dim. + self.embeds = torch.stack(self.embeds, dim=0) + else: + # [1, hidden_dim] (already has the leading dim) → plain concat. + self.embeds = torch.cat(self.embeds, dim=0) if self.embeds.shape[0] != len(self.positions): raise ValueError( f"embeds length ({self.embeds.shape[0]}) != " diff --git a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py index bb05a57c0..d8b0b1753 100644 --- a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py +++ b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py @@ -416,6 +416,16 @@ class TokenizerManagerScoreMixin: query_embed_overrides is not None or item_embed_overrides is not None ) + # Query placeholder positions are invariant across items — resolve once. + # (No-op returning ([], []) if has_embeds is False or query_embed_overrides is None.) + q_embeds, q_positions = self._resolve_overrides_for_sequence( + query, + query_embed_overrides, + embed_override_token_id, + position_offset=0, + label="query", + ) + if use_multi_item_scoring: # Multi-item scoring: concatenate with placeholder delimiter token. # Positions are derived from item lengths (delimiter_indices), not @@ -429,33 +439,27 @@ class TokenizerManagerScoreMixin: if not has_embeds: return None, input_ids, None, delimiter_indices - # Resolve embed overrides across the combined multi-item-scoring sequence - all_embeds: List[torch.Tensor] = [] - all_positions: List[int] = [] + # Resolve embed overrides across the combined multi-item-scoring sequence. + all_embeds: List[torch.Tensor] = list(q_embeds) + all_positions: List[int] = list(q_positions) current_offset = len(query) + 1 # +1 for first delimiter for i, item in enumerate(items): item_embs = item_embed_overrides[i] if item_embed_overrides else None - pe = self._resolve_embed_overrides_for_request( - query if i == 0 else [], # only resolve query overrides once + i_embeds, i_positions = self._resolve_overrides_for_sequence( item, - embed_override_token_id, - query_embed_overrides if i == 0 else None, item_embs, - current_offset, - f"items[{i}]", + embed_override_token_id, + position_offset=current_offset, + label=f"items[{i}]", ) - if pe is not None: - # pe.embeds is a stacked tensor after PositionalEmbeds.__post_init__ - all_embeds.append(pe.embeds) - all_positions.extend(pe.positions) + all_embeds.extend(i_embeds) + all_positions.extend(i_positions) current_offset += len(item) + 1 # +1 for delimiter if all_embeds: + # PositionalEmbeds.__post_init__ does the single torch.cat stack. positional_embed_overrides = [ - PositionalEmbeds( - embeds=torch.cat(all_embeds, dim=0), - positions=all_positions, - ) + PositionalEmbeds(embeds=all_embeds, positions=all_positions) ] else: positional_embed_overrides = None @@ -472,25 +476,34 @@ class TokenizerManagerScoreMixin: return None, input_ids, None, None positional_embed_overrides = [] + any_overrides = False for i, item in enumerate(items): item_embs = item_embed_overrides[i] if item_embed_overrides else None - pe = self._resolve_embed_overrides_for_request( - query, + i_embeds, i_positions = self._resolve_overrides_for_sequence( item, - embed_override_token_id, - query_embed_overrides, item_embs, - item_position_offset=len(query), - item_label=f"items[{i}]", + embed_override_token_id, + position_offset=len(query), + label=f"items[{i}]", ) - positional_embed_overrides.append(pe) + combined_embeds = q_embeds + i_embeds + if combined_embeds: + positional_embed_overrides.append( + PositionalEmbeds( + embeds=combined_embeds, + positions=q_positions + i_positions, + ) + ) + any_overrides = True + else: + positional_embed_overrides.append(None) - positional_embed_overrides = ( - positional_embed_overrides - if any(pe is not None for pe in positional_embed_overrides) - else None + return ( + None, + input_ids, + positional_embed_overrides if any_overrides else None, + None, ) - return None, input_ids, positional_embed_overrides, None # ------------------------------------------------------------------ # Main entry point