diff --git a/packages/sie_server/src/sie_server/adapters/_prompt_limit.py b/packages/sie_server/src/sie_server/adapters/_prompt_limit.py new file mode 100644 index 000000000..5f3de48b4 --- /dev/null +++ b/packages/sie_server/src/sie_server/adapters/_prompt_limit.py @@ -0,0 +1,131 @@ +"""Limits on the task prompt GLiNER-family adapters encode with each document. + +GLiNER, GLiNER2 and GLiREL encode a request's labels, relation types, class +labels and schema fields (its task prompt) with every document. The prompt is +not billed (input tokens count the document), so it is bounded instead, as +GLiFormer and GLiNER2.5-Decide bound theirs. A request is rejected with +``InvalidInputError`` (HTTP 400 ``INVALID_INPUT``) when + +* a label, relation type, class label, task name, field name or choice has + more than ``MAX_LABEL_CHARS`` characters (checked first, before anything is + tokenized), or +* the prompt takes more than ``max_prompt_tokens`` tokens, counted with the + model's tokenizer. + +The defaults sit far above the label sets these models are used with: the +largest label set among this repository's examples takes about 40 tokens, a +60-type PII list about 230, and a structured-extraction schema of 50 described +fields about 1,200, while ``DEFAULT_MAX_PROMPT_TOKENS`` holds several hundred +entity types and ``DEFAULT_MAX_SCHEMA_PROMPT_TOKENS`` a schema of dozens of +described fields with their choices. + +Descriptions are not limited one by one; the whole prompt is tokenized only +after its characters are checked, so counting costs at most +``MAX_PROMPT_CHARS_PER_TOKEN`` characters of tokenization per allowed token. +""" + +from __future__ import annotations + +import hashlib +from collections import OrderedDict +from collections.abc import Callable, Hashable, Iterable +from typing import Any + +from sie_server.types.inputs import InvalidInputError + +# Characters a label, relation type, class label, task name, field name or choice may have. +MAX_LABEL_CHARS = 128 +# Tokens a request's labels, relation types and class labels may take. +DEFAULT_MAX_PROMPT_TOKENS = 1024 +# Tokens a GLiNER2 request's labels or schema (field names, descriptions and choices) may take. +DEFAULT_MAX_SCHEMA_PROMPT_TOKENS = 2048 +# No token of these tokenizers covers more characters than this. +MAX_PROMPT_CHARS_PER_TOKEN = 32 +_PROMPT_CACHE_SIZE = 256 + + +def validate_max_prompt_tokens(value: object) -> int: + """A ``max_prompt_tokens`` adapter option, checked. + + Raises: + ValueError: The value is not a positive integer. + """ + if isinstance(value, bool) or not isinstance(value, int) or value < 1: + raise ValueError("max_prompt_tokens must be a positive integer") + return value + + +def check_label_chars(model: str, kind: str, values: Iterable[str]) -> None: + """Reject a label (or relation type, class label, field name, choice) of more than ``MAX_LABEL_CHARS`` characters. + + Raises: + InvalidInputError: A value is not a string or is too long. + """ + for value in values: + if not isinstance(value, str): + raise InvalidInputError(f"{model} {kind} must be strings") + if len(value) > MAX_LABEL_CHARS: + raise InvalidInputError(f"{model} {kind} may have at most {MAX_LABEL_CHARS} characters each") + + +class PromptLimit: + """Checks a request's task prompt against ``max_tokens``, remembering recent prompts' sizes.""" + + __slots__ = ("_counts", "max_tokens", "model") + + def __init__(self, model: str, max_tokens: int) -> None: + self.model = model + self.max_tokens = validate_max_prompt_tokens(max_tokens) + self._counts: OrderedDict[bytes, int] = OrderedDict() + + def check(self, texts: Iterable[str], count: Callable[[], int], key: Hashable) -> int: + """The prompt's tokens, from ``count()``, after checking ``texts`` (its strings) and the result. + + ``key`` identifies the prompt (its strings, in order, and anything else + its size depends on); a prompt seen recently is not counted again. + + Raises: + InvalidInputError: The prompt takes more than ``max_tokens`` tokens. + """ + digest = hashlib.sha256(repr(key).encode("utf-8", "surrogatepass")).digest() + tokens = self._counts.get(digest) + if tokens is None: + chars = sum(len(text) for text in texts) + if chars > self.max_tokens * MAX_PROMPT_CHARS_PER_TOKEN: + raise InvalidInputError(self._message(None, chars)) + tokens = int(count()) + self._counts[digest] = tokens + if len(self._counts) > _PROMPT_CACHE_SIZE: + self._counts.popitem(last=False) + else: + self._counts.move_to_end(digest) + if tokens > self.max_tokens: + raise InvalidInputError(self._message(tokens, None)) + return tokens + + def _message(self, tokens: int | None, chars: int | None) -> str: + size = f"{tokens} tokens" if tokens is not None else f"{chars} characters" + return ( + f"{self.model} labels, relation types, class labels and schema fields take {size}; " + f"a request may use at most {self.max_tokens} tokens for them" + ) + + +def gliner_prompt_counter(model: Any) -> Callable[[list[str], list[str]], int]: + """Tokens of the prompt a loaded ``gliner`` model builds for entity and relation types. + + The prompt is built by the model's own processor (``prepare_inputs``, with + no document words) and tokenized as the processor tokenizes it. + """ + processor = model.data_processor + tokenizer = processor.transformer_tokenizer + + def count(entity_types: list[str], relation_types: list[str]) -> int: + kwargs = {"relations": relation_types} if relation_types else {} + (words,), _ = processor.prepare_inputs([[]], entity_types, **kwargs) + if not words: + return 0 + encoding = tokenizer(list(words), is_split_into_words=True, add_special_tokens=False) + return len(encoding["input_ids"]) + + return count diff --git a/packages/sie_server/src/sie_server/adapters/gliner/__init__.py b/packages/sie_server/src/sie_server/adapters/gliner/__init__.py index 4bb803f7c..84a410836 100644 --- a/packages/sie_server/src/sie_server/adapters/gliner/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/gliner/__init__.py @@ -17,6 +17,11 @@ Joint entity-relation ("relex") models also extract relations between the entities they find when a request names relation types in ``options["relation_labels"]``. Without it they return entities only. + +A request's labels and relation types, which GLiNER encodes with every +document and does not bill, may have at most 128 characters each and take at +most ``max_prompt_tokens`` tokens together (default 1024); a longer prompt is +rejected with ``INVALID_INPUT``. """ import math @@ -27,6 +32,12 @@ import torch from sie_server.adapters._base_adapter import BaseAdapter +from sie_server.adapters._prompt_limit import ( + DEFAULT_MAX_PROMPT_TOKENS, + PromptLimit, + check_label_chars, + gliner_prompt_counter, +) from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_REQUIRES_TEXT, ComputePrecision from sie_server.adapters._word_window import bound_gliner_words, plan_forwards @@ -96,6 +107,7 @@ def __init__( multi_label: bool = False, merge_adjacent_entities: bool = False, relation_threshold: float | None = None, + max_prompt_tokens: int = DEFAULT_MAX_PROMPT_TOKENS, compute_precision: ComputePrecision = "float16", revision: str | None = None, **kwargs: Any, # Accept extra args from loader (e.g., pooling) @@ -112,6 +124,9 @@ def __init__( relation_threshold: Minimum relation score (0-1) for joint entity-relation models. None uses the entity threshold, as the gliner library does. + max_prompt_tokens: Most tokens a request's labels and relation + types may take in the prompt encoded with each document (see + ``_prompt_limit``). compute_precision: Compute precision for inference. revision: Optional HuggingFace revision/branch/commit SHA to pin when loading model artifacts. @@ -133,6 +148,9 @@ def __init__( self._extracts_relations = False # True when the encoder's attention memory grows with the square of a row (see ``_inference``). self._quadratic_attention = False + self._prompt_limit = PromptLimit("GLiNER", max_prompt_tokens) + # Tokens of the label prompt, as the loaded model builds it; None until loaded. + self._count_prompt: Any = None def load(self, device: str) -> None: """Load the model onto the specified device. @@ -173,6 +191,7 @@ def load(self, device: str) -> None: # gliner's max_len counts words, whatever their subwords: read at most a # bounded number of subwords too, with a long word in pieces. self._quadratic_attention = bound_gliner_words(self._model) + self._count_prompt = gliner_prompt_counter(self._model) def extract( self, @@ -227,6 +246,8 @@ def extract( if relation_labels and not self._extracts_relations: raise InvalidInputError(_ERR_NO_RELATIONS) + self._check_prompt(labels, relation_labels) + # Extract texts from all items texts = [self._extract_text(item) for item in items] if any(not text.strip() for text in texts): @@ -303,6 +324,27 @@ def extract( return ExtractOutput(entities=all_entities, relations=all_relations, input_token_counts=input_token_counts) + def _check_prompt(self, labels: list[str], relation_labels: list[str]) -> None: + """Reject a request whose labels and relation types take more than ``max_prompt_tokens``. + + gliner encodes the label prompt with every document, and only the + document is billed. + + Raises: + InvalidInputError: The prompt is too long, or a label is not a string. + """ + check_label_chars("GLiNER", "labels", labels) + check_label_chars("GLiNER", "relation_labels", relation_labels) + entity_types = list(dict.fromkeys(labels)) # gliner drops repeated labels + count = self._count_prompt + + def tokens() -> int: + return count(entity_types, relation_labels) if count is not None else 0 + + self._prompt_limit.check( + [*entity_types, *relation_labels], tokens, (tuple(entity_types), tuple(relation_labels)) + ) + @staticmethod def _validate_relation_labels(value: Any, labels: list[str]) -> list[str]: """Return the requested relation types (empty when none were asked for).""" diff --git a/packages/sie_server/src/sie_server/adapters/gliner2/adapter.py b/packages/sie_server/src/sie_server/adapters/gliner2/adapter.py index 911acb6d8..6623bc2d6 100644 --- a/packages/sie_server/src/sie_server/adapters/gliner2/adapter.py +++ b/packages/sie_server/src/sie_server/adapters/gliner2/adapter.py @@ -13,6 +13,7 @@ from huggingface_hub import snapshot_download from sie_server.adapters._base_adapter import BaseAdapter +from sie_server.adapters._prompt_limit import DEFAULT_MAX_SCHEMA_PROMPT_TOKENS, PromptLimit, check_label_chars from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_REQUIRES_TEXT, ComputePrecision from sie_server.adapters._word_window import ( @@ -76,6 +77,12 @@ class GLiNER2Adapter(BaseAdapter): - Batch methods cover entities, relations, structured data, and classification - Classification uses ``classify_text()`` / ``batch_classify_text()`` + A request's labels, class labels, relation types or schema fields (field + names, descriptions and choices), which gliner2 encodes with every document + and does not bill, may take at most ``max_prompt_tokens`` tokens (default + 2048), and each label, task name, field name or choice at most 128 + characters; a longer prompt is rejected with ``INVALID_INPUT``. + Reference models: - fastino/gliner2-base-v1 - fastino/gliner2-large-v1 @@ -98,6 +105,7 @@ def __init__( default_labels: list[str] | None = None, multi_label: bool = False, max_seq_length: int | None = None, + max_prompt_tokens: int = DEFAULT_MAX_SCHEMA_PROMPT_TOKENS, compute_precision: ComputePrecision = "float16", revision: str | None = None, **kwargs: Any, @@ -115,6 +123,9 @@ def __init__( multi_label: Whether the configured classification task may return multiple labels. max_seq_length: Maximum document and schema input length. + max_prompt_tokens: Most tokens a request's labels, class labels, + relation types or schema fields may take in the task prompt + encoded with each document (see ``_prompt_limit``). compute_precision: Compute precision for inference. revision: Optional HuggingFace revision/branch/commit SHA to pin when loading model artifacts. @@ -127,6 +138,7 @@ def __init__( self._default_labels = self._validate_labels(default_labels) if default_labels is not None else None self._multi_label = multi_label self._max_seq_length = max_seq_length + self._prompt_limit = PromptLimit("GLiNER2", max_prompt_tokens) self._compute_precision = compute_precision self._revision = revision @@ -254,8 +266,14 @@ def extract( raise ValueError("GLiNER2 structured extraction does not accept classification_task") structures = self._json_schema_to_structures(output_schema) specs = [spec for fields in structures.values() for spec in fields] - # A field's choices are read twice: in its structure and in a prefix before the document. - rows = self._row_tokens(windows, specs + specs) + # A field's choices are listed in its structure, and each again in a prefix before the document. + choices = [ + choice for definition in output_schema["properties"].values() for choice in definition.get("enum") or [] + ] + check_label_chars("GLiNER2", "output_schema property names", output_schema["properties"]) + check_label_chars("GLiNER2", "output_schema enum values", choices) + prompt = self._prompt_tokens(specs, key=("json", tuple(specs)), extra=len(choices)) + rows = self._row_tokens(windows, prompt) with torch.inference_mode(): raw_results = self._run_planned( model_texts, @@ -286,7 +304,11 @@ def extract( normalized_entities = [ self._normalize_input_entities(item, entities or []) for item, entities in zip(items, relation_entities) ] - rows = self._row_tokens(windows, normalized_labels, per_entry=_PROMPT_TOKENS_PER_RELATION) + check_label_chars("GLiNER2", "labels", normalized_labels) + prompt = self._prompt_tokens( + normalized_labels, per_entry=_PROMPT_TOKENS_PER_RELATION, key=("relations", tuple(normalized_labels)) + ) + rows = self._row_tokens(windows, prompt) with torch.inference_mode(): raw_results = self._run_planned( model_texts, @@ -314,6 +336,12 @@ def extract( if classification_task is not None: if not isinstance(classification_task, str) or not classification_task.strip(): raise ValueError("GLiNER2 classification_task must be a non-empty string") + check_label_chars("GLiNER2", "classification_task", [classification_task]) + check_label_chars("GLiNER2", "labels", normalized_labels) + prompt = self._prompt_tokens( + [classification_task, *normalized_labels], + key=("classification", classification_task, tuple(normalized_labels)), + ) return self._classify( model_texts, normalized_labels, @@ -321,7 +349,7 @@ def extract( multi_label=multi_label, threshold=effective_threshold, input_token_counts=input_token_counts, - rows=self._row_tokens(windows, [classification_task, *normalized_labels]), + rows=self._row_tokens(windows, prompt), ) def extract_entities(batch: list[str]) -> list[Any]: @@ -345,10 +373,12 @@ def extract_entities(batch: list[str]) -> list[Any]: max_len=self._max_seq_length, ) + check_label_chars("GLiNER2", "labels", normalized_labels) + prompt = self._prompt_tokens(normalized_labels, key=("entities", tuple(normalized_labels))) with torch.inference_mode(): raw_results = self._run_planned( model_texts, - self._row_tokens(windows, normalized_labels), + self._row_tokens(windows, prompt), extract_entities, rows_per_pass=1 if len(texts) == 1 else _PACKAGE_BATCH_SIZE, ) @@ -412,22 +442,43 @@ def classify(batch: list[str]) -> list[Any]: input_token_counts=input_token_counts, ) - def _row_tokens( + def _prompt_tokens( self, - windows: list[tuple[str, int | None]], - prompt_entries: Iterable[str], + entries: list[str], *, + key: tuple[Any, ...], per_entry: int = _PROMPT_TOKENS_PER_ENTRY, - ) -> list[int] | None: + extra: int = 0, + ) -> int | None: + """Estimated tokens of the task prompt gliner2 builds from ``entries``, checked against the limit. + + Each label, class label, relation type and schema field is counted + with the tokens gliner2 adds around it, plus ``extra`` (a token per + field choice, which gliner2 lists again before the document). None when words are + not counted (no bounded splitter is installed); the prompt's + characters are still checked. + + Raises: + InvalidInputError: The prompt takes more than ``max_prompt_tokens``. + """ + count = self._count_subwords + strings = [entry for entry in entries if isinstance(entry, str)] + + def tokens() -> int: + if count is None: + return 0 + return sum(count(strings)) + per_entry * len(strings) + extra + _ROW_OVERHEAD_TOKENS + + prompt = self._prompt_limit.check(strings, tokens, (per_entry, extra, *key)) + return prompt if count is not None else None + + def _row_tokens(self, windows: list[tuple[str, int | None]], prompt: int | None) -> list[int] | None: """Estimated tokens of each item's encoder row: the task prompt, then the words it reads. None when the words were not counted (no bounded splitter is installed). """ - count = self._count_subwords - if count is None or any(subwords is None for _, subwords in windows): + if prompt is None or any(subwords is None for _, subwords in windows): return None - entries = [entry for entry in prompt_entries if isinstance(entry, str)] - prompt = sum(count(entries)) + per_entry * len(entries) + _ROW_OVERHEAD_TOKENS return [prompt + (subwords or 0) for _, subwords in windows] def _run_planned( diff --git a/packages/sie_server/src/sie_server/adapters/glirel/__init__.py b/packages/sie_server/src/sie_server/adapters/glirel/__init__.py index 9d5df2d26..3d520e353 100644 --- a/packages/sie_server/src/sie_server/adapters/glirel/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/glirel/__init__.py @@ -7,6 +7,11 @@ Reference models: - jackboyla/glirel-large-v0 (zero-shot relation extraction) - jackboyla/glirel_re_large-v0 (relation-focused variant) + +A request's relation types, which GLiREL encodes with every text, may have at +most 128 characters each and take at most ``max_prompt_tokens`` tokens together +(default 1024), and an item may carry at most ``MAX_ENTITIES`` (256) entities; +other requests are rejected with ``INVALID_INPUT``. """ import re @@ -17,6 +22,7 @@ import torch from sie_server.adapters._base_adapter import BaseAdapter +from sie_server.adapters._prompt_limit import DEFAULT_MAX_PROMPT_TOKENS, PromptLimit, check_label_chars from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_REQUIRES_TEXT, ComputePrecision from sie_server.adapters._word_window import SubwordCounter, WindowedSplitter, split_word_counter, subword_budget @@ -27,6 +33,9 @@ # Error messages _ERR_REQUIRES_LABELS = "GLiREL requires labels parameter for relation extraction" _ERR_REQUIRES_ENTITIES = "GLiREL requires entities in item metadata for relation extraction" +# GLiREL scores every pair of an item's entities (in Python while preparing +# the batch, and on the GPU), so an item may carry at most this many. +MAX_ENTITIES = 256 _TOKEN_PATTERN = re.compile(r"\w+(?:[-_]\w+)*|\S") # Subword tokens a text may take per word GLiREL reads: its checkpoints read # English, whose prose, code and logs run at up to 2.4 subwords per word. @@ -72,6 +81,7 @@ def __init__( model_name_or_path: str | Path, *, threshold: float = 0.3, + max_prompt_tokens: int = DEFAULT_MAX_PROMPT_TOKENS, compute_precision: ComputePrecision = "float16", revision: str | None = None, **kwargs: Any, # Accept extra args from loader @@ -81,6 +91,8 @@ def __init__( Args: model_name_or_path: HuggingFace model ID or local path to GLiREL model. threshold: Minimum confidence score for relation extraction (0-1). + max_prompt_tokens: Most tokens a request's relation types may take + in the prompt encoded with each text (see ``_prompt_limit``). compute_precision: Compute precision for inference. revision: Optional HuggingFace revision/branch/commit SHA to pin when loading model artifacts. @@ -96,6 +108,9 @@ def __init__( self._device: str | None = None # The words of a text GLiREL reads (see ``_tokenize``); None until loaded. self._words: WindowedSplitter | None = None + self._prompt_limit = PromptLimit("GLiREL", max_prompt_tokens) + # Tokens of the relation-type prompt, as the loaded model builds it; None until loaded. + self._count_prompt: Any = None def load(self, device: str) -> None: """Load the model onto the specified device. @@ -126,6 +141,9 @@ def load(self, device: str) -> None: int(self._model.base_config.max_len), getattr(embeddings.model, "config", None), ) + self._count_prompt = _prompt_counter( + embeddings.tokenizer, str(self._model.rel_token), str(self._model.sep_token) + ) def _bound_words(self, tokenizer: Any, max_words: int, encoder_config: Any = None) -> None: """Read at most ``max_words`` words of a text, and a bounded number of their subwords. @@ -177,23 +195,25 @@ def extract( Raises: RuntimeError: If model not loaded. - ValueError: If labels not provided or items lack entities. - InvalidInputError: If an entity is not an object or has invalid offsets. + InvalidInputError: If labels are missing or too long, or an item lacks + text or entities, carries more than ``MAX_ENTITIES`` entities, or + has invalid entity offsets. """ self._check_loaded() if not labels: - raise ValueError(_ERR_REQUIRES_LABELS) + raise InvalidInputError(_ERR_REQUIRES_LABELS) + self._check_prompt(labels) - # Check every item's entities before running any item. + # Check every item before running any, so bad input fails as a 400 before model work. inputs: list[tuple[str, list[dict[str, Any]]]] = [] for item in items: text = self._extract_text(item) entities = self._extract_entities(item) - if not entities: - raise ValueError(_ERR_REQUIRES_ENTITIES) - + raise InvalidInputError(_ERR_REQUIRES_ENTITIES) + if len(entities) > MAX_ENTITIES: + raise InvalidInputError(f"GLiREL items may carry at most {MAX_ENTITIES} entities in metadata") for entity in entities: self._validate_entity_span(entity, text) inputs.append((text, entities)) @@ -256,10 +276,27 @@ def extract( return ExtractOutput(entities=all_entities, relations=all_relations) + def _check_prompt(self, labels: list[str]) -> None: + """Reject a request whose relation types take more than ``max_prompt_tokens``. + + GLiREL encodes the relation types with every text, and GLiREL output + carries no input token count. + + Raises: + InvalidInputError: The prompt is too long, or a relation type is not a string. + """ + check_label_chars("GLiREL", "labels", labels) + count = self._count_prompt + + def tokens() -> int: + return count(labels) if count is not None else 0 + + self._prompt_limit.check(labels, tokens, tuple(labels)) + def _extract_text(self, item: Item) -> str: """Extract text from an item.""" if item.text is None: - raise ValueError(ERR_REQUIRES_TEXT.format(adapter_name="GLiREL adapter")) + raise InvalidInputError(ERR_REQUIRES_TEXT.format(adapter_name="GLiREL adapter")) return item.text def _extract_entities(self, item: Item) -> list[dict[str, Any]]: @@ -267,7 +304,10 @@ def _extract_entities(self, item: Item) -> list[dict[str, Any]]: metadata = item.metadata if metadata is None: return [] - return metadata.get("entities", []) + entities = metadata.get("entities", []) + if not isinstance(entities, list): + raise InvalidInputError("GLiREL item metadata.entities must be a list") + return entities def _tokenize(self, text: str) -> tuple[list[str], list[tuple[int, int]], int]: """Tokenize text like GLiREL, keeping the words it reads: ``(words, offsets, end of the text read)``. @@ -376,6 +416,17 @@ def _relation_entity_text( return str(relation_text) +def _prompt_counter(tokenizer: Any, rel_token: str, sep_token: str) -> Any: + """Tokens of GLiREL's prompt for a list of relation types: ``[REL] type ... [REL] type [SEP]``.""" + + def count(labels: list[str]) -> int: + words = [word for label in labels for word in (rel_token, label)] + [sep_token] + encoding = tokenizer(words, is_split_into_words=True, add_special_tokens=False) + return len(encoding["input_ids"]) + + return count + + def _words(text: str) -> Iterator[tuple[str, int, int]]: """GLiREL's words of ``text`` with their character offsets.""" for match in _TOKEN_PATTERN.finditer(text): diff --git a/packages/sie_server/tests/adapters/test_gliner2_long_text.py b/packages/sie_server/tests/adapters/test_gliner2_long_text.py index 3a5e0b3ec..cb87b2921 100644 --- a/packages/sie_server/tests/adapters/test_gliner2_long_text.py +++ b/packages/sie_server/tests/adapters/test_gliner2_long_text.py @@ -373,14 +373,15 @@ def test_the_prefix_reads_what_gliner2_reads_with_its_sentence_end(text: str, ma def test_relation_rows_run_in_passes_within_the_attention_budget() -> None: # gliner2 builds a structure of about ten tokens around each relation type. - adapter, model = make_adapter(encoder_config={"model_type": "deberta-v2"}) + adapter, model = make_adapter(encoder_config={"model_type": "deberta-v2"}, max_prompt_tokens=8192) labels = [f"rel{index}" for index in range(150)] text = " ".join(["abcdefghi"] * 400) entities = [{"text": "abcdefghi", "label": "x", "start": 0, "end": 9}] adapter.extract([Item(text=text, metadata={"entities": entities}) for _ in range(8)], labels=labels) - estimated = adapter._row_tokens([adapter._window(text)] * 1, labels, per_entry=_PROMPT_TOKENS_PER_RELATION) + prompt = adapter._prompt_tokens(labels, per_entry=_PROMPT_TOKENS_PER_RELATION, key=("relations", tuple(labels))) + estimated = adapter._row_tokens([adapter._window(text)], prompt) assert estimated is not None for batch in model.inputs: rows, width = batch.input_ids.shape @@ -389,7 +390,7 @@ def test_relation_rows_run_in_passes_within_the_attention_budget() -> None: def test_structured_rows_are_not_underestimated() -> None: - adapter, model = make_adapter(encoder_config={"model_type": "deberta-v2"}) + adapter, model = make_adapter(encoder_config={"model_type": "deberta-v2"}, max_prompt_tokens=8192) schema = { "type": "object", "properties": { @@ -401,11 +402,11 @@ def test_structured_rows_are_not_underestimated() -> None: adapter.extract([Item(text=text)], output_schema=schema) - structures = adapter._json_schema_to_structures(schema) - specs = [spec for fields in structures.values() for spec in fields] - estimated = adapter._row_tokens([adapter._window(text)], specs + specs) + (prompt,) = adapter._prompt_limit._counts.values() # the estimate extract() planned with + estimated = adapter._row_tokens([adapter._window(text)], prompt) assert estimated is not None - assert model.inputs[-1].input_ids.shape[1] <= estimated[0] + # The schema is estimated from its field specs, not gliner2's own word split of it. + assert model.inputs[-1].input_ids.shape[1] <= estimated[0] * 1.02 @pytest.mark.parametrize("max_len", [1, 2]) diff --git a/packages/sie_server/tests/adapters/test_prompt_limit.py b/packages/sie_server/tests/adapters/test_prompt_limit.py new file mode 100644 index 000000000..266203cd5 --- /dev/null +++ b/packages/sie_server/tests/adapters/test_prompt_limit.py @@ -0,0 +1,323 @@ +"""The limit on the task prompt (labels, relation types, class labels, schema fields) of GLiNER-family adapters.""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import MagicMock + +import pytest +from sie_server.adapters._prompt_limit import ( + DEFAULT_MAX_PROMPT_TOKENS, + DEFAULT_MAX_SCHEMA_PROMPT_TOKENS, + MAX_LABEL_CHARS, + MAX_PROMPT_CHARS_PER_TOKEN, + PromptLimit, +) +from sie_server.adapters.gliner import GLiNERAdapter +from sie_server.adapters.gliner2.classification import GLiNER2ClassificationAdapter +from sie_server.adapters.glirel import MAX_ENTITIES, GLiRELAdapter +from sie_server.types.inputs import InvalidInputError, Item + +from .test_gliner2_long_text import make_adapter as make_gliner2_adapter +from .test_gliner_long_words import DEBERTA, FakeGLiNER, char_tokenizer, encoded, load + +# The largest label set among the repository's examples (12 labels). The stand-in +# tokenizer takes a token per character, so these take about 200 tokens. +EXAMPLE_LABELS = [ + "prompt injection", + "jailbreak", + "role play", + "system prompt leak", + "harmful request", + "data exfiltration", + "instruction override", + "obfuscation", + "social engineering", + "policy evasion", + "privilege escalation", + "benign", +] +# Sixty short types (about 300 tokens with the stand-in tokenizer). +PII_LABELS = [f"pii{index}" for index in range(60)] + + +class Counting: + def __init__(self, tokens: int) -> None: + self.tokens = tokens + self.calls = 0 + + def __call__(self) -> int: + self.calls += 1 + return self.tokens + + +def test_a_prompt_within_the_limit_passes_and_is_counted_once() -> None: + limit = PromptLimit("Model", 100) + count = Counting(40) + + assert limit.check(["a", "b"], count, ("a", "b")) == 40 + assert limit.check(["a", "b"], count, ("a", "b")) == 40 + assert count.calls == 1 + + +def test_a_prompt_over_the_limit_is_rejected_every_time() -> None: + limit = PromptLimit("Model", 100) + count = Counting(101) + + for _ in range(2): + with pytest.raises(InvalidInputError, match="take 101 tokens; a request may use at most 100"): + limit.check(["a"], count, ("a",)) + assert count.calls == 1 + + +def test_a_prompt_of_too_many_characters_is_rejected_before_it_is_tokenized() -> None: + limit = PromptLimit("Model", 10) + count = Counting(1) + + with pytest.raises(InvalidInputError, match="characters"): + limit.check(["x" * (10 * MAX_PROMPT_CHARS_PER_TOKEN + 1)], count, ("long",)) + assert count.calls == 0 + + +@pytest.mark.parametrize("value", [0, -1, True, 1.5, "1024"]) +def test_the_limit_must_be_a_positive_integer(value: Any) -> None: + with pytest.raises(ValueError, match="positive integer"): + PromptLimit("Model", value) + + +def test_gliner_counts_the_prompt_its_processor_builds() -> None: + model = FakeGLiNER(DEBERTA) + adapter = load(GLiNERAdapter, model) + labels = ["person", "organization"] + + tokens = adapter._count_prompt(labels, []) + + # The encoded row is [CLS], the prompt, the document's subwords, [SEP]. + (row,) = encoded(model, ["abc"], labels) + assert tokens == len(row) - 3 - 2 + + +@pytest.mark.parametrize("labels", [EXAMPLE_LABELS, PII_LABELS], ids=["example-labels", "pii-list"]) +def test_gliner_accepts_ordinary_label_sets(labels: list[str]) -> None: + model = FakeGLiNER(DEBERTA) + adapter = load(GLiNERAdapter, model) + + adapter.extract([Item(text="Priya Raman works at Novartis.")], labels=labels) + + assert model.calls == [["Priya Raman works at Novartis."]] + + +@pytest.mark.parametrize( + "labels", + [[f"{index}" + "x" * 120 for index in range(12)], [f"type number {index}" for index in range(100)]], + ids=["long-labels", "many-labels"], +) +def test_gliner_rejects_a_prompt_over_the_limit_before_inference(labels: list[str]) -> None: + model = FakeGLiNER(DEBERTA) + adapter = load(GLiNERAdapter, model) + + with pytest.raises(InvalidInputError, match=f"at most {DEFAULT_MAX_PROMPT_TOKENS} tokens"): + adapter.extract([Item(text="Priya Raman works at Novartis.")], labels=labels) + assert model.calls == [] + + +def test_gliner_limit_is_configurable() -> None: + model = FakeGLiNER(DEBERTA) + adapter = GLiNERAdapter("fake/gliner", max_prompt_tokens=20) + adapter._model = model + adapter._count_prompt = load(GLiNERAdapter, FakeGLiNER(DEBERTA))._count_prompt + + with pytest.raises(InvalidInputError, match="at most 20 tokens"): + adapter.extract([Item(text="Priya Raman works at Novartis.")], labels=["organization", "location"]) + + +def glirel_adapter() -> tuple[GLiRELAdapter, MagicMock]: + adapter = GLiRELAdapter("fake/glirel") + adapter._model = MagicMock() + adapter._model.predict_relations.return_value = [] + adapter._bound_words(char_tokenizer(), 64, DEBERTA) + from sie_server.adapters.glirel import _prompt_counter + + adapter._count_prompt = _prompt_counter(char_tokenizer(), "[", "]") + return adapter, adapter._model.predict_relations + + +def test_glirel_rejects_relation_types_over_the_limit() -> None: + adapter, predict = glirel_adapter() + item = Item( + text="Tim Cook leads Apple", + metadata={"entities": [{"text": "Tim Cook", "label": "PERSON", "start": 0, "end": 8}]}, + ) + + adapter.extract([item], labels=["ceo of", "works for"]) + with pytest.raises(InvalidInputError, match="GLiREL"): + adapter.extract([item], labels=[f"{i}" + "z" * 500 for i in range(3)]) + assert predict.call_count == 1 + + +def test_gliner2_accepts_ordinary_label_sets_and_schemas() -> None: + adapter, model = make_gliner2_adapter() + schema = { + "type": "object", + "properties": { + f"field{index}": {"type": "string", "description": "the value of a field on an insurance claim form"} + for index in range(30) + }, + } + + adapter.extract([Item(text="Priya Raman works at Novartis")], labels=PII_LABELS) + adapter._model.batch_extract_json = MagicMock(return_value=[{}]) + adapter.extract([Item(text="Priya Raman works at Novartis")], output_schema=schema) + + assert len(model.inputs) == 1 + adapter._model.batch_extract_json.assert_called_once() + + +@pytest.mark.parametrize( + ("labels", "output_schema"), + [ + ([f"{index:03d}" + "w" * 117 for index in range(60)], None), + (None, {"type": "object", "properties": {"f": {"type": "string", "description": "d" * 70_000}}}), + ( + None, + { + "type": "object", + "properties": { + f"field{index}": {"type": "string", "enum": [f"choice {index} {option}" for option in range(40)]} + for index in range(40) + }, + }, + ), + ], + ids=["long-labels", "long-description", "many-choices"], +) +def test_gliner2_rejects_a_prompt_over_the_limit_before_inference( + labels: list[str] | None, output_schema: dict[str, Any] | None +) -> None: + adapter, model = make_gliner2_adapter() + adapter._model.batch_extract_json = MagicMock(return_value=[{}]) + + with pytest.raises(InvalidInputError, match=f"at most {DEFAULT_MAX_SCHEMA_PROMPT_TOKENS} tokens"): + adapter.extract([Item(text="Priya Raman works at Novartis")], labels=labels, output_schema=output_schema) + assert model.inputs == [] + adapter._model.batch_extract_json.assert_not_called() + + +def test_gliner2_rejects_a_classification_prompt_over_the_limit() -> None: + adapter, model = make_gliner2_adapter(classification_task="topic") + + with pytest.raises(InvalidInputError, match="GLiNER2"): + adapter.extract([Item(text="Priya Raman works at Novartis")], labels=[f"t{index}" * 40 for index in range(80)]) + assert model.inputs == [] + + +class NoCounting: + def __call__(self, *args: Any, **kwargs: Any) -> Any: + raise AssertionError("a label over the character limit must be rejected before it is tokenized") + + +def test_gliner_rejects_a_long_label_before_tokenizing_it() -> None: + model = FakeGLiNER(DEBERTA) + adapter = load(GLiNERAdapter, model) + adapter._count_prompt = NoCounting() + + with pytest.raises(InvalidInputError, match=f"at most {MAX_LABEL_CHARS} characters"): + adapter.extract([Item(text="Priya Raman works at Novartis.")], labels=["q" * (4 * 1024 * 1024)]) + assert model.calls == [] + + +def test_gliner_rejects_long_relation_types() -> None: + adapter = load(GLiNERAdapter, FakeGLiNER(DEBERTA)) + adapter._extracts_relations = True + + with pytest.raises(InvalidInputError, match="relation_labels may have at most"): + adapter.extract( + [Item(text="Priya Raman works at Novartis.")], + labels=["person"], + options={"relation_labels": ["r" * (MAX_LABEL_CHARS + 1)]}, + ) + + +@pytest.mark.parametrize( + ("labels", "output_schema", "match"), + [ + (["x" * (MAX_LABEL_CHARS + 1)], None, "labels may have at most"), + (None, {"type": "object", "properties": {"n" * (MAX_LABEL_CHARS + 1): {"type": "string"}}}, "property names"), + ( + None, + {"type": "object", "properties": {"f": {"type": "string", "enum": ["c" * (MAX_LABEL_CHARS + 1)]}}}, + "enum values", + ), + ], + ids=["label", "field-name", "choice"], +) +def test_gliner2_rejects_long_labels_field_names_and_choices( + labels: list[str] | None, output_schema: dict[str, Any] | None, match: str +) -> None: + adapter, model = make_gliner2_adapter() + adapter._count_subwords = NoCounting() + + with pytest.raises(InvalidInputError, match=match): + adapter.extract([Item(text="Priya Raman works at Novartis")], labels=labels, output_schema=output_schema) + assert model.inputs == [] + + +def test_gliner2_classifier_rejects_a_long_task_name_and_long_class_labels() -> None: + # The GLiGuard route (gliner2 2.x in the transformers5 bundle) uses the same checks. + adapter, model = make_gliner2_adapter(GLiNER2ClassificationAdapter, classification_task="prompt_safety") + + with pytest.raises(InvalidInputError, match="classification_task may have at most"): + adapter.extract([Item(text="Hello")], labels=["safe", "unsafe"], options={"classification_task": "t" * 200}) + with pytest.raises(InvalidInputError, match="GLiNER2"): + adapter.extract([Item(text="Hello")], labels=[f"class {index:03d} " + "k" * 110 for index in range(60)]) + adapter.extract([Item(text="Hello")], labels=["safe", "unsafe"]) + assert len(model.inputs) == 1 + + +def glirel_item(entities: list[Any], text: str = "Tim Cook leads Apple") -> Item: + return Item(text=text, metadata={"entities": entities}) + + +def test_glirel_caps_the_entities_of_an_item() -> None: + adapter, predict = glirel_adapter() + text = "word " * (MAX_ENTITIES + 1) + entities = [ + {"text": "word", "label": "X", "start": 5 * index, "end": 5 * index + 4} for index in range(MAX_ENTITIES + 1) + ] + + with pytest.raises(InvalidInputError, match=f"at most {MAX_ENTITIES} entities"): + adapter.extract([glirel_item(entities, text)], labels=["ceo of"]) + adapter.extract([glirel_item(entities[:MAX_ENTITIES], text)], labels=["ceo of"]) + assert predict.call_count == 1 + + +@pytest.mark.parametrize( + "entity", + [ + {"text": "x", "label": "ORG", "start": 5, "end": 2}, + {"text": "x", "label": "ORG", "start": "0", "end": 3}, + {"text": "x", "label": "ORG", "start": 500, "end": 505}, + "Tim Cook", + ], + ids=["reversed", "non-integer", "past-the-text", "not-an-object"], +) +def test_glirel_rejects_bad_entities_as_invalid_input_before_any_model_work(entity: Any) -> None: + adapter, predict = glirel_adapter() + good = glirel_item([{"text": "Tim Cook", "label": "PERSON", "start": 0, "end": 8}]) + + with pytest.raises(InvalidInputError): + adapter.extract([good, glirel_item([entity])], labels=["ceo of"]) + predict.assert_not_called() + + +def test_glirel_rejects_missing_input_as_invalid_input() -> None: + adapter, _ = glirel_adapter() + + with pytest.raises(InvalidInputError, match="requires labels"): + adapter.extract([glirel_item([{"text": "Tim", "label": "P", "start": 0, "end": 3}])], labels=[]) + with pytest.raises(InvalidInputError, match="requires entities"): + adapter.extract([Item(text="Tim Cook leads Apple")], labels=["ceo of"]) + with pytest.raises(InvalidInputError, match="must be a list"): + adapter.extract([Item(text="Tim Cook", metadata={"entities": "Tim"})], labels=["ceo of"]) + with pytest.raises(InvalidInputError, match="labels may have at most"): + adapter.extract([glirel_item([{"text": "Tim", "label": "P", "start": 0, "end": 3}])], labels=["r" * 200])