[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]
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]}) != "
@@ -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