[Score API] Hoist query placeholder scan and specialize PositionalEmbeds stacking (#23513)

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
jsheng_Linkedin
2026-04-29 13:51:53 -07:00
committed by GitHub
co-authored by Claude Opus 4.7
parent 79dbfe4505
commit 850021378a
2 changed files with 52 additions and 33 deletions
+10 -4
View File
@@ -41,11 +41,17 @@ class PositionalEmbeds:
positions: List[int] positions: List[int]
def __post_init__(self): 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): if isinstance(self.embeds, list):
self.embeds = torch.cat( if not self.embeds:
[e if e.dim() == 2 else e.unsqueeze(0) for e in self.embeds], dim=0 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): if self.embeds.shape[0] != len(self.positions):
raise ValueError( raise ValueError(
f"embeds length ({self.embeds.shape[0]}) != " f"embeds length ({self.embeds.shape[0]}) != "
@@ -416,6 +416,16 @@ class TokenizerManagerScoreMixin:
query_embed_overrides is not None or item_embed_overrides is not None 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: if use_multi_item_scoring:
# Multi-item scoring: concatenate with placeholder delimiter token. # Multi-item scoring: concatenate with placeholder delimiter token.
# Positions are derived from item lengths (delimiter_indices), not # Positions are derived from item lengths (delimiter_indices), not
@@ -429,33 +439,27 @@ class TokenizerManagerScoreMixin:
if not has_embeds: if not has_embeds:
return None, input_ids, None, delimiter_indices return None, input_ids, None, delimiter_indices
# Resolve embed overrides across the combined multi-item-scoring sequence # Resolve embed overrides across the combined multi-item-scoring sequence.
all_embeds: List[torch.Tensor] = [] all_embeds: List[torch.Tensor] = list(q_embeds)
all_positions: List[int] = [] all_positions: List[int] = list(q_positions)
current_offset = len(query) + 1 # +1 for first delimiter current_offset = len(query) + 1 # +1 for first delimiter
for i, item in enumerate(items): for i, item in enumerate(items):
item_embs = item_embed_overrides[i] if item_embed_overrides else None item_embs = item_embed_overrides[i] if item_embed_overrides else None
pe = self._resolve_embed_overrides_for_request( i_embeds, i_positions = self._resolve_overrides_for_sequence(
query if i == 0 else [], # only resolve query overrides once
item, item,
embed_override_token_id,
query_embed_overrides if i == 0 else None,
item_embs, item_embs,
current_offset, embed_override_token_id,
f"items[{i}]", position_offset=current_offset,
label=f"items[{i}]",
) )
if pe is not None: all_embeds.extend(i_embeds)
# pe.embeds is a stacked tensor after PositionalEmbeds.__post_init__ all_positions.extend(i_positions)
all_embeds.append(pe.embeds)
all_positions.extend(pe.positions)
current_offset += len(item) + 1 # +1 for delimiter current_offset += len(item) + 1 # +1 for delimiter
if all_embeds: if all_embeds:
# PositionalEmbeds.__post_init__ does the single torch.cat stack.
positional_embed_overrides = [ positional_embed_overrides = [
PositionalEmbeds( PositionalEmbeds(embeds=all_embeds, positions=all_positions)
embeds=torch.cat(all_embeds, dim=0),
positions=all_positions,
)
] ]
else: else:
positional_embed_overrides = None positional_embed_overrides = None
@@ -472,25 +476,34 @@ class TokenizerManagerScoreMixin:
return None, input_ids, None, None return None, input_ids, None, None
positional_embed_overrides = [] positional_embed_overrides = []
any_overrides = False
for i, item in enumerate(items): for i, item in enumerate(items):
item_embs = item_embed_overrides[i] if item_embed_overrides else None item_embs = item_embed_overrides[i] if item_embed_overrides else None
pe = self._resolve_embed_overrides_for_request( i_embeds, i_positions = self._resolve_overrides_for_sequence(
query,
item, item,
embed_override_token_id,
query_embed_overrides,
item_embs, item_embs,
item_position_offset=len(query), embed_override_token_id,
item_label=f"items[{i}]", 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 = ( return (
positional_embed_overrides None,
if any(pe is not None for pe in positional_embed_overrides) input_ids,
else None positional_embed_overrides if any_overrides else None,
None,
) )
return None, input_ids, positional_embed_overrides, None
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Main entry point # Main entry point