[VLM] Simplify post-processing of replacement info (#12269)
Signed-off-by: DarkLight1337 <tlleungac@connect.ust.hk>
This commit is contained in:
@@ -1,7 +1,8 @@
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import defaultdict
|
||||
from collections.abc import Callable, ItemsView, Iterable, Mapping, Sequence
|
||||
from collections.abc import (Callable, Generator, ItemsView, Iterable, Mapping,
|
||||
Sequence)
|
||||
from dataclasses import dataclass, field
|
||||
from functools import lru_cache
|
||||
from typing import (TYPE_CHECKING, Generic, NamedTuple, Optional, Protocol,
|
||||
@@ -31,6 +32,24 @@ _S = TypeVar("_S", str, list[int])
|
||||
_PromptSeq = Union[str, list[int]]
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptReplacementDetails:
|
||||
full: _PromptSeq
|
||||
"""The full replacement."""
|
||||
|
||||
features: _PromptSeq
|
||||
"""
|
||||
The part of the replacement that corresponds to placeholder feature tokens.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def from_seq(seq: _PromptSeq):
|
||||
return PromptReplacementDetails(full=seq, features=seq)
|
||||
|
||||
|
||||
_PromptRepl = Union[_PromptSeq, PromptReplacementDetails]
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptReplacement:
|
||||
"""
|
||||
@@ -43,8 +62,8 @@ class PromptReplacement:
|
||||
target: _PromptSeq
|
||||
"""The token sequence (or text) to find and replace."""
|
||||
|
||||
replacement: Union[Callable[[int], _PromptSeq],
|
||||
_PromptSeq] = field(repr=False)
|
||||
replacement: Union[Callable[[int], _PromptRepl],
|
||||
_PromptRepl] = field(repr=False)
|
||||
"""
|
||||
Given the index of the processed item within :attr:`modality`,
|
||||
output the replacement token sequence (or text).
|
||||
@@ -112,6 +131,14 @@ class _BoundPromptSequence:
|
||||
_text: Optional[str]
|
||||
_token_ids: Optional[list[int]]
|
||||
|
||||
@staticmethod
|
||||
def from_seq(tokenizer: AnyTokenizer, seq: _PromptSeq):
|
||||
return _BoundPromptSequence(
|
||||
tokenizer=tokenizer,
|
||||
_text=seq if isinstance(seq, str) else None,
|
||||
_token_ids=seq if isinstance(seq, list) else None,
|
||||
)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self._text is None and self._token_ids is None:
|
||||
raise ValueError("At least one of 'text' and 'token_ids' must be "
|
||||
@@ -134,6 +161,12 @@ class _BoundPromptSequence:
|
||||
return self._token_ids
|
||||
|
||||
|
||||
@dataclass
|
||||
class _BoundPromptReplacementGroup:
|
||||
full: _BoundPromptSequence
|
||||
features: _BoundPromptSequence
|
||||
|
||||
|
||||
@dataclass
|
||||
class BoundPromptReplacement:
|
||||
"""
|
||||
@@ -145,24 +178,18 @@ class BoundPromptReplacement:
|
||||
modality: str
|
||||
|
||||
_target: _PromptSeq
|
||||
_replacement: Union[Callable[[int], _PromptSeq],
|
||||
_PromptSeq] = field(repr=False)
|
||||
_replacement: Union[Callable[[int], _PromptRepl],
|
||||
_PromptRepl] = field(repr=False)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self._replacement_cache = dict[int, _BoundPromptSequence]()
|
||||
self._replacement_cache = dict[int, _BoundPromptReplacementGroup]()
|
||||
|
||||
@property
|
||||
def target(self) -> _BoundPromptSequence:
|
||||
"""The token sequence (or text) to find and replace."""
|
||||
target = self._target
|
||||
return _BoundPromptSequence.from_seq(self.tokenizer, self._target)
|
||||
|
||||
return _BoundPromptSequence(
|
||||
tokenizer=self.tokenizer,
|
||||
_text=target if isinstance(target, str) else None,
|
||||
_token_ids=target if isinstance(target, list) else None,
|
||||
)
|
||||
|
||||
def get_replacement(self, item_idx: int) -> _BoundPromptSequence:
|
||||
def get_replacement(self, item_idx: int) -> _BoundPromptReplacementGroup:
|
||||
"""
|
||||
Given the index of the processed item within :attr:`modality`,
|
||||
output the replacement token sequence (or text).
|
||||
@@ -177,10 +204,16 @@ class BoundPromptReplacement:
|
||||
else:
|
||||
cache_key = None
|
||||
|
||||
bound_replacement = _BoundPromptSequence(
|
||||
tokenizer=self.tokenizer,
|
||||
_text=replacement if isinstance(replacement, str) else None,
|
||||
_token_ids=replacement if isinstance(replacement, list) else None,
|
||||
if not isinstance(replacement, PromptReplacementDetails):
|
||||
replacement = PromptReplacementDetails.from_seq(replacement)
|
||||
|
||||
bound_full = _BoundPromptSequence.from_seq(self.tokenizer,
|
||||
replacement.full)
|
||||
bound_features = _BoundPromptSequence.from_seq(self.tokenizer,
|
||||
replacement.features)
|
||||
bound_replacement = _BoundPromptReplacementGroup(
|
||||
full=bound_full,
|
||||
features=bound_features,
|
||||
)
|
||||
|
||||
if cache_key is not None:
|
||||
@@ -197,7 +230,7 @@ class _TokenMatch(NamedTuple):
|
||||
def iter_token_matches(
|
||||
token_ids: list[int],
|
||||
match_ids: list[int],
|
||||
) -> Iterable[_TokenMatch]:
|
||||
) -> Generator[_TokenMatch]:
|
||||
"""
|
||||
Yield each occurrence of :code:`match_ids` in :code:`token_ids`.
|
||||
|
||||
@@ -272,15 +305,15 @@ class _PromptReplacementTextMatch(_PromptReplacementMatch):
|
||||
|
||||
|
||||
@dataclass
|
||||
class PlaceholderInfo:
|
||||
class PlaceholderFeaturesInfo:
|
||||
modality: str
|
||||
item_idx: int
|
||||
start_idx: int
|
||||
replacement: list[int]
|
||||
tokens: list[int]
|
||||
|
||||
@property
|
||||
def length(self) -> int:
|
||||
return len(self.replacement)
|
||||
return len(self.tokens)
|
||||
|
||||
def to_range(self) -> PlaceholderRange:
|
||||
return PlaceholderRange(
|
||||
@@ -362,10 +395,10 @@ def _replace_matches(
|
||||
replacement = repl_info.get_replacement(item_idx)
|
||||
|
||||
if isinstance(prompt, str):
|
||||
repl_seq = replacement.text
|
||||
repl_seq = replacement.full.text
|
||||
out_seqs.append(prompt[prev_end_idx:start_idx] + repl_seq)
|
||||
else:
|
||||
repl_seq = replacement.token_ids
|
||||
repl_seq = replacement.full.token_ids
|
||||
out_seqs.append(prompt[prev_end_idx:start_idx] + repl_seq)
|
||||
|
||||
prev_end_idx = end_idx
|
||||
@@ -408,7 +441,7 @@ def _iter_placeholders(
|
||||
mm_prompt_repls: Mapping[str, Sequence[BoundPromptReplacement]],
|
||||
prompt: list[int],
|
||||
mm_item_counts: Mapping[str, int],
|
||||
) -> Iterable[PlaceholderInfo]:
|
||||
) -> Iterable[PlaceholderFeaturesInfo]:
|
||||
"""
|
||||
Yield each set of placeholder tokens found in :code:`prompt`.
|
||||
|
||||
@@ -432,23 +465,33 @@ def _iter_placeholders(
|
||||
|
||||
for repl_info in modality_repls:
|
||||
replacement = repl_info.get_replacement(item_idx)
|
||||
repl_tokens = replacement.token_ids
|
||||
repl_len = len(repl_tokens)
|
||||
end_idx = start_idx + repl_len
|
||||
repl_tokens_full = replacement.full.token_ids
|
||||
repl_len_full = len(repl_tokens_full)
|
||||
end_idx_full = start_idx + repl_len_full
|
||||
|
||||
if repl_len == 0 or end_idx > prompt_len:
|
||||
if repl_len_full == 0 or end_idx_full > prompt_len:
|
||||
continue
|
||||
|
||||
if prompt[start_idx:end_idx] == repl_tokens:
|
||||
yield PlaceholderInfo(
|
||||
modality=modality,
|
||||
item_idx=item_idx,
|
||||
start_idx=start_idx,
|
||||
replacement=repl_tokens,
|
||||
)
|
||||
if prompt[start_idx:end_idx_full] == repl_tokens_full:
|
||||
repl_tokens_feat = replacement.features.token_ids
|
||||
|
||||
try:
|
||||
match = next(
|
||||
iter_token_matches(repl_tokens_full,
|
||||
repl_tokens_feat))
|
||||
yield PlaceholderFeaturesInfo(
|
||||
modality=modality,
|
||||
item_idx=item_idx,
|
||||
start_idx=start_idx + match.start_idx,
|
||||
tokens=repl_tokens_feat,
|
||||
)
|
||||
except StopIteration:
|
||||
raise AssertionError(
|
||||
f"{repl_tokens_feat=} should be a "
|
||||
f"subsequence of {repl_tokens_full=}") from None
|
||||
|
||||
# Exclude overlapping matches
|
||||
start_idx = end_idx
|
||||
start_idx = end_idx_full
|
||||
item_idx_by_modality[modality] += 1
|
||||
found = True
|
||||
break
|
||||
@@ -464,7 +507,7 @@ def find_mm_placeholders(
|
||||
mm_prompt_repls: Mapping[str, Sequence[BoundPromptReplacement]],
|
||||
prompt: list[int],
|
||||
mm_item_counts: Mapping[str, int],
|
||||
) -> Mapping[str, list[PlaceholderInfo]]:
|
||||
) -> Mapping[str, list[PlaceholderFeaturesInfo]]:
|
||||
it = _iter_placeholders(mm_prompt_repls, prompt, mm_item_counts)
|
||||
return dict(full_groupby_modality(it))
|
||||
|
||||
@@ -679,7 +722,7 @@ class BaseMultiModalProcessor(ABC, Generic[_I]):
|
||||
mm_prompt_repls: Mapping[str, Sequence[BoundPromptReplacement]],
|
||||
new_token_ids: list[int],
|
||||
mm_item_counts: Mapping[str, int],
|
||||
) -> Mapping[str, list[PlaceholderInfo]]:
|
||||
) -> Mapping[str, list[PlaceholderFeaturesInfo]]:
|
||||
return find_mm_placeholders(mm_prompt_repls, new_token_ids,
|
||||
mm_item_counts)
|
||||
|
||||
@@ -948,7 +991,7 @@ class BaseMultiModalProcessor(ABC, Generic[_I]):
|
||||
token_ids: list[int],
|
||||
mm_prompt_repls: Mapping[str, Sequence[BoundPromptReplacement]],
|
||||
mm_item_counts: Mapping[str, int],
|
||||
) -> tuple[list[int], str, Mapping[str, list[PlaceholderInfo]]]:
|
||||
) -> tuple[list[int], str, Mapping[str, list[PlaceholderFeaturesInfo]]]:
|
||||
tokenizer = self.info.get_tokenizer()
|
||||
|
||||
mm_token_matches = {
|
||||
@@ -1037,7 +1080,7 @@ class BaseMultiModalProcessor(ABC, Generic[_I]):
|
||||
|
||||
def _validate_mm_placeholders(
|
||||
self,
|
||||
mm_placeholders: Mapping[str, list[PlaceholderInfo]],
|
||||
mm_placeholders: Mapping[str, list[PlaceholderFeaturesInfo]],
|
||||
mm_item_counts: Mapping[str, int],
|
||||
*,
|
||||
allow_missing: bool = False,
|
||||
|
||||
Reference in New Issue
Block a user