[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:
co-authored by
Claude Opus 4.7
parent
79dbfe4505
commit
850021378a
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user