1313from huggingface_hub import snapshot_download
1414
1515from sie_server .adapters ._base_adapter import BaseAdapter
16+ from sie_server .adapters ._prompt_limit import DEFAULT_MAX_SCHEMA_PROMPT_TOKENS , PromptLimit , check_label_chars
1617from sie_server .adapters ._spec import AdapterSpec
1718from sie_server .adapters ._types import ERR_REQUIRES_TEXT , ComputePrecision
1819from sie_server .adapters ._word_window import (
@@ -76,6 +77,12 @@ class GLiNER2Adapter(BaseAdapter):
7677 - Batch methods cover entities, relations, structured data, and classification
7778 - Classification uses ``classify_text()`` / ``batch_classify_text()``
7879
80+ A request's labels, class labels, relation types or schema fields (field
81+ names, descriptions and choices), which gliner2 encodes with every document
82+ and does not bill, may take at most ``max_prompt_tokens`` tokens (default
83+ 2048), and each label, task name, field name or choice at most 128
84+ characters; a longer prompt is rejected with ``INVALID_INPUT``.
85+
7986 Reference models:
8087 - fastino/gliner2-base-v1
8188 - fastino/gliner2-large-v1
@@ -98,6 +105,7 @@ def __init__(
98105 default_labels : list [str ] | None = None ,
99106 multi_label : bool = False ,
100107 max_seq_length : int | None = None ,
108+ max_prompt_tokens : int = DEFAULT_MAX_SCHEMA_PROMPT_TOKENS ,
101109 compute_precision : ComputePrecision = "float16" ,
102110 revision : str | None = None ,
103111 ** kwargs : Any ,
@@ -115,6 +123,9 @@ def __init__(
115123 multi_label: Whether the configured classification task may return
116124 multiple labels.
117125 max_seq_length: Maximum document and schema input length.
126+ max_prompt_tokens: Most tokens a request's labels, class labels,
127+ relation types or schema fields may take in the task prompt
128+ encoded with each document (see ``_prompt_limit``).
118129 compute_precision: Compute precision for inference.
119130 revision: Optional HuggingFace revision/branch/commit SHA to pin when
120131 loading model artifacts.
@@ -127,6 +138,7 @@ def __init__(
127138 self ._default_labels = self ._validate_labels (default_labels ) if default_labels is not None else None
128139 self ._multi_label = multi_label
129140 self ._max_seq_length = max_seq_length
141+ self ._prompt_limit = PromptLimit ("GLiNER2" , max_prompt_tokens )
130142 self ._compute_precision = compute_precision
131143 self ._revision = revision
132144
@@ -254,8 +266,14 @@ def extract(
254266 raise ValueError ("GLiNER2 structured extraction does not accept classification_task" )
255267 structures = self ._json_schema_to_structures (output_schema )
256268 specs = [spec for fields in structures .values () for spec in fields ]
257- # A field's choices are read twice: in its structure and in a prefix before the document.
258- rows = self ._row_tokens (windows , specs + specs )
269+ # A field's choices are listed in its structure, and each again in a prefix before the document.
270+ choices = [
271+ choice for definition in output_schema ["properties" ].values () for choice in definition .get ("enum" ) or []
272+ ]
273+ check_label_chars ("GLiNER2" , "output_schema property names" , output_schema ["properties" ])
274+ check_label_chars ("GLiNER2" , "output_schema enum values" , choices )
275+ prompt = self ._prompt_tokens (specs , key = ("json" , tuple (specs )), extra = len (choices ))
276+ rows = self ._row_tokens (windows , prompt )
259277 with torch .inference_mode ():
260278 raw_results = self ._run_planned (
261279 model_texts ,
@@ -286,7 +304,11 @@ def extract(
286304 normalized_entities = [
287305 self ._normalize_input_entities (item , entities or []) for item , entities in zip (items , relation_entities )
288306 ]
289- rows = self ._row_tokens (windows , normalized_labels , per_entry = _PROMPT_TOKENS_PER_RELATION )
307+ check_label_chars ("GLiNER2" , "labels" , normalized_labels )
308+ prompt = self ._prompt_tokens (
309+ normalized_labels , per_entry = _PROMPT_TOKENS_PER_RELATION , key = ("relations" , tuple (normalized_labels ))
310+ )
311+ rows = self ._row_tokens (windows , prompt )
290312 with torch .inference_mode ():
291313 raw_results = self ._run_planned (
292314 model_texts ,
@@ -314,14 +336,20 @@ def extract(
314336 if classification_task is not None :
315337 if not isinstance (classification_task , str ) or not classification_task .strip ():
316338 raise ValueError ("GLiNER2 classification_task must be a non-empty string" )
339+ check_label_chars ("GLiNER2" , "classification_task" , [classification_task ])
340+ check_label_chars ("GLiNER2" , "labels" , normalized_labels )
341+ prompt = self ._prompt_tokens (
342+ [classification_task , * normalized_labels ],
343+ key = ("classification" , classification_task , tuple (normalized_labels )),
344+ )
317345 return self ._classify (
318346 model_texts ,
319347 normalized_labels ,
320348 task = classification_task ,
321349 multi_label = multi_label ,
322350 threshold = effective_threshold ,
323351 input_token_counts = input_token_counts ,
324- rows = self ._row_tokens (windows , [ classification_task , * normalized_labels ] ),
352+ rows = self ._row_tokens (windows , prompt ),
325353 )
326354
327355 def extract_entities (batch : list [str ]) -> list [Any ]:
@@ -345,10 +373,12 @@ def extract_entities(batch: list[str]) -> list[Any]:
345373 max_len = self ._max_seq_length ,
346374 )
347375
376+ check_label_chars ("GLiNER2" , "labels" , normalized_labels )
377+ prompt = self ._prompt_tokens (normalized_labels , key = ("entities" , tuple (normalized_labels )))
348378 with torch .inference_mode ():
349379 raw_results = self ._run_planned (
350380 model_texts ,
351- self ._row_tokens (windows , normalized_labels ),
381+ self ._row_tokens (windows , prompt ),
352382 extract_entities ,
353383 rows_per_pass = 1 if len (texts ) == 1 else _PACKAGE_BATCH_SIZE ,
354384 )
@@ -412,22 +442,43 @@ def classify(batch: list[str]) -> list[Any]:
412442 input_token_counts = input_token_counts ,
413443 )
414444
415- def _row_tokens (
445+ def _prompt_tokens (
416446 self ,
417- windows : list [tuple [str , int | None ]],
418- prompt_entries : Iterable [str ],
447+ entries : list [str ],
419448 * ,
449+ key : tuple [Any , ...],
420450 per_entry : int = _PROMPT_TOKENS_PER_ENTRY ,
421- ) -> list [int ] | None :
451+ extra : int = 0 ,
452+ ) -> int | None :
453+ """Estimated tokens of the task prompt gliner2 builds from ``entries``, checked against the limit.
454+
455+ Each label, class label, relation type and schema field is counted
456+ with the tokens gliner2 adds around it, plus ``extra`` (a token per
457+ field choice, which gliner2 lists again before the document). None when words are
458+ not counted (no bounded splitter is installed); the prompt's
459+ characters are still checked.
460+
461+ Raises:
462+ InvalidInputError: The prompt takes more than ``max_prompt_tokens``.
463+ """
464+ count = self ._count_subwords
465+ strings = [entry for entry in entries if isinstance (entry , str )]
466+
467+ def tokens () -> int :
468+ if count is None :
469+ return 0
470+ return sum (count (strings )) + per_entry * len (strings ) + extra + _ROW_OVERHEAD_TOKENS
471+
472+ prompt = self ._prompt_limit .check (strings , tokens , (per_entry , extra , * key ))
473+ return prompt if count is not None else None
474+
475+ def _row_tokens (self , windows : list [tuple [str , int | None ]], prompt : int | None ) -> list [int ] | None :
422476 """Estimated tokens of each item's encoder row: the task prompt, then the words it reads.
423477
424478 None when the words were not counted (no bounded splitter is installed).
425479 """
426- count = self ._count_subwords
427- if count is None or any (subwords is None for _ , subwords in windows ):
480+ if prompt is None or any (subwords is None for _ , subwords in windows ):
428481 return None
429- entries = [entry for entry in prompt_entries if isinstance (entry , str )]
430- prompt = sum (count (entries )) + per_entry * len (entries ) + _ROW_OVERHEAD_TOKENS
431482 return [prompt + (subwords or 0 ) for _ , subwords in windows ]
432483
433484 def _run_planned (
0 commit comments