diff --git a/README.md b/README.md index 1eec36bd2..f506d95c2 100644 --- a/README.md +++ b/README.md @@ -47,7 +47,8 @@ One SIE cluster runs the inference behind a whole agent. Each task is a handful | **Document to markdown** | PDFs, Office files, and scans become clean markdown. | [`lightonocr`](packages/sie_server/models/lightonai__LightOnOCR-2-1B.yaml), [`glm-ocr`](packages/sie_server/models/zai-org__GLM-OCR.yaml), [`mineru`](packages/sie_server/models/opendatalab__MinerU2.5-Pro-2604-1.2B.yaml), [`paddleocr-vl`](packages/sie_server/models/PaddlePaddle__PaddleOCR-VL-1.5.yaml), [`docling`](packages/sie_server/models/docling.yaml) | | **Structured output** | Schema-valid JSON, extracted or generated. | [`gliner2`](packages/sie_server/models/fastino__gliner2-large-v1.yaml), [`nuner-zero`](packages/sie_server/models/numind__NuNER_Zero.yaml), [`qwen3.8-27b`](packages/sie_server/models/Qwen__Qwen3.8-27B-FP8.yaml), [`qwen3.6-27b`](packages/sie_server/models/Qwen__Qwen3.6-27B.yaml) | | **Decide** | Choice, yes/no, and score answers with probabilities to typed questions about a text or JSON state. | [`laya`](packages/sie_server/models/convaiinnovations__laya.yaml), [`laya-multilingual`](packages/sie_server/models/convaiinnovations__laya-multilingual.yaml), [`laya-typed-decisions`](packages/sie_server/models/convaiinnovations__laya-typed-decisions.yaml) | -| **Guard content** | A Yes/No safety verdict, with the decision threshold tunable in the model config. | [`granite-guardian-2b`](packages/sie_server/models/ibm-granite__granite-guardian-3.0-2b.yaml) | +| **Classify** | Zero-shot labels, with several label groups answered in one call. The instruct models also follow a task instruction and few-shot examples. | [`gliclass-large-v3`](packages/sie_server/models/knowledgator__gliclass-large-v3.0.yaml), [`gliclass-instruct-large`](packages/sie_server/models/knowledgator__gliclass-instruct-large-v1.0.yaml), [`gliclass-multilang-mini`](packages/sie_server/models/knowledgator__gliclass-multilang-mini.yaml) | +| **Guard content** | A safety verdict: Yes/No with the threshold set in the model config, or safe/unsafe and policy-label scores with the threshold chosen per request. | [`granite-guardian-2b`](packages/sie_server/models/ibm-granite__granite-guardian-3.0-2b.yaml), [`opir-multitask-large`](packages/sie_server/models/knowledgator__opir-multitask-large-v1.0.yaml), [`opir-edge`](packages/sie_server/models/knowledgator__opir-edge-v1.0.yaml) | | **Run the agent loop** | Plan steps and call tools with an open LLM, streaming included. | [`qwen3.8-27b`](packages/sie_server/models/Qwen__Qwen3.8-27B-FP8.yaml), [`qwen3.6-27b`](packages/sie_server/models/Qwen__Qwen3.6-27B.yaml) | | **Translate** | Text between 400+ languages. | [`madlad400-3b-mt`](packages/sie_server/models/google__madlad400-3b-mt.yaml) | | **See images** | Caption, detect objects, and answer questions about images. | [`florence-2`](packages/sie_server/models/microsoft__Florence-2-large.yaml), [`owlv2`](packages/sie_server/models/google__owlv2-base-patch16-ensemble.yaml), [`grounding-dino`](packages/sie_server/models/IDEA-Research__grounding-dino-base.yaml) | diff --git a/packages/sie_gateway/src/handlers/proxy.rs b/packages/sie_gateway/src/handlers/proxy.rs index 541a1e00d..82db30833 100644 --- a/packages/sie_gateway/src/handlers/proxy.rs +++ b/packages/sie_gateway/src/handlers/proxy.rs @@ -374,6 +374,10 @@ const RESOURCE_EXHAUSTED_RETRY_AFTER: &str = RetryAfter::DEFAULT.resource_exhaus const LORA_LOADING_ERROR_CODE: &str = "LORA_LOADING"; const LORA_LOADING_RETRY_AFTER: &str = RetryAfter::DEFAULT.lora_loading; const INVALID_INPUT_ERROR_CODE: &str = "INVALID_INPUT"; +/// Worker-side input exceeds the model's context window (for example a label +/// set that does not fit). Caller-fixable, so it maps to 400 like +/// ``INVALID_INPUT``; see ``sie_server.adapters.errors.InputTooLongError``. +const INPUT_TOO_LONG_ERROR_CODE: &str = "INPUT_TOO_LONG"; const PAYLOAD_TOO_LARGE_ERROR_CODE: &str = err_code::PAYLOAD_TOO_LARGE; /// Fallback `max_tokens` applied to a chat-completions request that @@ -9002,6 +9006,7 @@ fn unanimous_terminal_client_error( let first = errors.first()?.error_code.as_deref()?; let (status, canonical) = match first { INVALID_INPUT_ERROR_CODE => (StatusCode::BAD_REQUEST, INVALID_INPUT_ERROR_CODE), + INPUT_TOO_LONG_ERROR_CODE => (StatusCode::BAD_REQUEST, INPUT_TOO_LONG_ERROR_CODE), PAYLOAD_TOO_LARGE_ERROR_CODE => { (StatusCode::PAYLOAD_TOO_LARGE, PAYLOAD_TOO_LARGE_ERROR_CODE) } @@ -18468,6 +18473,7 @@ mod tests { fn test_unanimous_terminal_client_errors_map_to_400_and_413() { for (code, status) in [ (INVALID_INPUT_ERROR_CODE, StatusCode::BAD_REQUEST), + (INPUT_TOO_LONG_ERROR_CODE, StatusCode::BAD_REQUEST), (PAYLOAD_TOO_LARGE_ERROR_CODE, StatusCode::PAYLOAD_TOO_LARGE), ] { let first = _err_result(Some(code), "rejected 1"); @@ -18485,6 +18491,27 @@ mod tests { unanimous_terminal_client_error(&[&invalid, &oversized]), None ); + let too_long = _err_result(Some(INPUT_TOO_LONG_ERROR_CODE), "labels do not fit"); + let failed = _err_result(Some("inference_error"), "backend failure"); + assert_eq!(unanimous_terminal_client_error(&[&too_long, &failed]), None); + } + + #[tokio::test] + async fn test_input_too_long_is_a_native_400_with_its_code() { + let too_long = _err_result(Some(INPUT_TOO_LONG_ERROR_CODE), "labels do not fit"); + let (status, code) = unanimous_terminal_client_error(&[&too_long]).unwrap(); + let response = build_terminal_client_error_response(status, code, "labels do not fit"); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!( + response.headers().get("x-sie-error-code").unwrap(), + INPUT_TOO_LONG_ERROR_CODE + ); + let body = axum::body::to_bytes(response.into_body(), 16 * 1024) + .await + .unwrap(); + let value: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!(value["detail"]["code"], INPUT_TOO_LONG_ERROR_CODE); + assert_eq!(value["detail"]["message"], "labels do not fit"); } #[test] diff --git a/packages/sie_sdk/README.md b/packages/sie_sdk/README.md index f640daa9f..6fb16504f 100644 --- a/packages/sie_sdk/README.md +++ b/packages/sie_sdk/README.md @@ -34,6 +34,73 @@ for entry in scores["scores"]: print(entry["item_id"], entry["score"]) ``` +## Zero-shot classification + +GLiClass models (`knowledgator/gliclass-*` and the `knowledgator/opir-*` +guardrail models) score text against labels passed with each request. Every +label comes back in `classifications`, sorted by score. By default the scores +form one distribution that sums to 1; `options={"classification_type": +"multi-label"}` scores each label independently instead. + +```python +ticket = Item(text="I was charged twice and support has not replied in three days.") + +result = client.extract( + "knowledgator/gliclass-instruct-large-v1.0", + ticket, + labels=["billing", "bug report", "feature request"], + instruction="Classify the support ticket by its main topic.", +) +print(result["classifications"][0]["label"]) # billing +``` + +`instruction` is the model's task prompt; the `gliclass-instruct-*` and +`opir-*` models are trained to follow one. Few-shot examples go in +`options={"examples": [{"text": ..., "labels": [...]}]}`, with labels taken +from the request's label set. They can move scores a long way, so check them +on your own data. The model reads examples after the document, so a long +document can push them out of the model's window; send +`options={"overflow_policy": "truncate_text"}` to shorten the document instead. + +To answer several questions in one call, pass named `label_groups` instead of +`labels`. Scores are normalized within each group, and `data` holds one answer +per group: + +```python +result = client.extract( + "knowledgator/gliclass-instruct-large-v1.0", + ticket, + instruction="Triage the support ticket.", + options={ + "label_groups": { + "topic": ["billing", "bug report", "feature request"], + "urgency": ["low", "medium", "high"], + "needs_human": ["yes", "no"], + } + }, +) +urgency = result["data"]["urgency"] +print(urgency["choice"], urgency["confidence"]) # e.g. medium 0.78 +print(urgency["probabilities"]) # {"low": ..., "medium": ..., "high": ...} +``` + +Each group answers `{"type": "choice", "choice", "probabilities", +"confidence"}`, where `confidence` is `1 - entropy / log(number of labels)`. +With `"classification_type": "multi-label"` a group answers `{"labels", +"probabilities"}`: every label scored independently, and `labels` lists those +at or above `options.threshold` (0.5 when no threshold is set). +`classifications` lists the same scores under `group.label` names. In a +grouped request, an example's labels may be written as `"urgency.high"` or as +`{"urgency": "high"}`. + +Usage counts each item's document tokens plus the tokens of the instruction +and example texts sent with it, since the model encodes them for every item. +Label names are not counted. With an instruction or examples, each item's +count is capped at the model window minus the label prompt, unless the +document count alone is already higher. An item whose document pushes the +labels out of the window comes back with an `INPUT_TOO_LONG` error in its +`error` field, and the other items still succeed. + ## Generation prompts and guard verdicts `generate` and `stream_generate` treat text-only prompts as raw continuation diff --git a/packages/sie_sdk/src/sie_sdk/types.py b/packages/sie_sdk/src/sie_sdk/types.py index 6b1bd2d88..f07a67650 100644 --- a/packages/sie_sdk/src/sie_sdk/types.py +++ b/packages/sie_sdk/src/sie_sdk/types.py @@ -511,7 +511,9 @@ class ExtractResult(TypedDict, total=False): relations: List of extracted relation triples. classifications: List of classification results. objects: List of detected objects with bounding boxes. - data: Additional structured extraction data (if output_schema was provided). + data: Structured extraction data: schema-driven results when + output_schema was provided, document parses, or one answer per + group when GLiClass options.label_groups was provided. error: Stable per-item failure when extraction did not complete. request: Request-scoped id, metered usage, and settled debit when supplied by the gateway. diff --git a/packages/sie_server/README.md b/packages/sie_server/README.md index acce396a2..1b0e4e6c3 100644 --- a/packages/sie_server/README.md +++ b/packages/sie_server/README.md @@ -48,6 +48,25 @@ removes matched stop sequences consistently from returned text, completion-token usage, and logprobs. Long completions delay the first visible text; existing request timeouts still apply. Other adapters are unaffected. +### GLiClass usage + +For GLiClass classification, `usage.input_tokens` counts each item's document +tokens plus the instruction and few-shot example texts sent with the request, +because the model encodes that text again for every item. Label names, +including the labels attached to examples, are not counted. With an +instruction or examples, an item's count is capped at the model's +`max_sequence_length` minus the label prompt, unless the document count alone +is already higher. An item refused because its document pushes the labels out +of the window returns a per-item `INPUT_TOO_LONG` error and counts nothing. A +request that sends no instruction or examples is counted exactly as before. + +The instruction and each example text may be at most 2,048 characters, and +together with the example labels at most 8,192 characters. Up to 32 examples +are accepted, and they must leave room for the document in the model window. +Label names are refused when their total length exceeds 16 characters per +token of the window (8,192 characters for a 512-token model), more than any +label prompt can fit. + ## Configuration `sie-server` reads its config from `SIE_*` environment variables (Pydantic diff --git a/packages/sie_server/bundles/default.yaml b/packages/sie_server/bundles/default.yaml index 6962c1d24..10ce4bd1d 100644 --- a/packages/sie_server/bundles/default.yaml +++ b/packages/sie_server/bundles/default.yaml @@ -79,8 +79,9 @@ deps: gliner2: '>=1.3.1,<2' # glirel glirel: '>=1.0,<2' - # gliclass - gliclass: '>=0.1,<1' + # gliclass (0.1.17: cross-attention scorer for the multilingual checkpoints; + # 0.1.18+ needs transformers 5) + gliclass: '>=0.1.17,<1' # gliner/glirel/gliclass shared dep loguru: '>=0.7,<1' # donut, florence2 diff --git a/packages/sie_server/models/knowledgator__gliclass-base-v3.0.yaml b/packages/sie_server/models/knowledgator__gliclass-base-v3.0.yaml new file mode 100644 index 000000000..eebb79d85 --- /dev/null +++ b/packages/sie_server/models/knowledgator__gliclass-base-v3.0.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/gliclass-base-v3.0 +hf_id: knowledgator/gliclass-base-v3.0 +hf_revision: 77a70e6cd52e602ed18184ef37d18bdd3741e3d5 +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 512 +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__gliclass-edge-v3.0.yaml b/packages/sie_server/models/knowledgator__gliclass-edge-v3.0.yaml new file mode 100644 index 000000000..ee11e3645 --- /dev/null +++ b/packages/sie_server/models/knowledgator__gliclass-edge-v3.0.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/gliclass-edge-v3.0 +hf_id: knowledgator/gliclass-edge-v3.0 +hf_revision: df03993a2ed98e5e4a0d2dd7efbbd105abe874cf +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 512 +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__gliclass-instruct-base-v1.0.yaml b/packages/sie_server/models/knowledgator__gliclass-instruct-base-v1.0.yaml new file mode 100644 index 000000000..2f22fc351 --- /dev/null +++ b/packages/sie_server/models/knowledgator__gliclass-instruct-base-v1.0.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/gliclass-instruct-base-v1.0 +hf_id: knowledgator/gliclass-instruct-base-v1.0 +hf_revision: 4f6a108b08a5537f395521d19b5073e197923dd3 +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 512 +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__gliclass-instruct-edge-v1.0.yaml b/packages/sie_server/models/knowledgator__gliclass-instruct-edge-v1.0.yaml new file mode 100644 index 000000000..cd17cd793 --- /dev/null +++ b/packages/sie_server/models/knowledgator__gliclass-instruct-edge-v1.0.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/gliclass-instruct-edge-v1.0 +hf_id: knowledgator/gliclass-instruct-edge-v1.0 +hf_revision: 727be8a417f6a7718e591b025e07054c146d8139 +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 512 +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__gliclass-instruct-large-v1.0.yaml b/packages/sie_server/models/knowledgator__gliclass-instruct-large-v1.0.yaml new file mode 100644 index 000000000..4abb4c0ec --- /dev/null +++ b/packages/sie_server/models/knowledgator__gliclass-instruct-large-v1.0.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/gliclass-instruct-large-v1.0 +hf_id: knowledgator/gliclass-instruct-large-v1.0 +hf_revision: 825e5478c1bf4bffbf297690517097ccbdb2e006 +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 512 +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__gliclass-multilang-edge.yaml b/packages/sie_server/models/knowledgator__gliclass-multilang-edge.yaml new file mode 100644 index 000000000..921138df3 --- /dev/null +++ b/packages/sie_server/models/knowledgator__gliclass-multilang-edge.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/gliclass-multilang-edge +hf_id: knowledgator/gliclass-multilang-edge +hf_revision: d16c08ef70547514081952104e6fc3d190d9ee39 +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 512 +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__gliclass-multilang-mini.yaml b/packages/sie_server/models/knowledgator__gliclass-multilang-mini.yaml new file mode 100644 index 000000000..7ce1624da --- /dev/null +++ b/packages/sie_server/models/knowledgator__gliclass-multilang-mini.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/gliclass-multilang-mini +hf_id: knowledgator/gliclass-multilang-mini +hf_revision: 0bd888b6c3ef9fca5f0a9d407bddfbbc7623486b +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 512 +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__opir-edge-multilang-v1.0.yaml b/packages/sie_server/models/knowledgator__opir-edge-multilang-v1.0.yaml new file mode 100644 index 000000000..f3bcb6ce8 --- /dev/null +++ b/packages/sie_server/models/knowledgator__opir-edge-multilang-v1.0.yaml @@ -0,0 +1,23 @@ +sie_id: knowledgator/opir-edge-multilang-v1.0 +hf_id: knowledgator/opir-edge-multilang-v1.0 +hf_revision: 969a05f901f8de234b0204f5309cb7e6450e6371 +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 1024 # trained with 1024-token inputs +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__opir-edge-v1.0.yaml b/packages/sie_server/models/knowledgator__opir-edge-v1.0.yaml new file mode 100644 index 000000000..5430ba3b5 --- /dev/null +++ b/packages/sie_server/models/knowledgator__opir-edge-v1.0.yaml @@ -0,0 +1,23 @@ +sie_id: knowledgator/opir-edge-v1.0 +hf_id: knowledgator/opir-edge-v1.0 +hf_revision: 467a0431f744522f92f4e37d7756309b4c12332a +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 1024 # trained with 1024-token inputs +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__opir-multitask-large-v1.0.yaml b/packages/sie_server/models/knowledgator__opir-multitask-large-v1.0.yaml new file mode 100644 index 000000000..4cc3dcaa8 --- /dev/null +++ b/packages/sie_server/models/knowledgator__opir-multitask-large-v1.0.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/opir-multitask-large-v1.0 +hf_id: knowledgator/opir-multitask-large-v1.0 +hf_revision: 69bb27407d66eab4797049d36ba75eef3579335b +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 1024 # trained with 1024-token inputs +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/models/knowledgator__opir-multitask-multilang-v1.0.yaml b/packages/sie_server/models/knowledgator__opir-multitask-multilang-v1.0.yaml new file mode 100644 index 000000000..1764f4076 --- /dev/null +++ b/packages/sie_server/models/knowledgator__opir-multitask-multilang-v1.0.yaml @@ -0,0 +1,29 @@ +sie_id: knowledgator/opir-multitask-multilang-v1.0 +hf_id: knowledgator/opir-multitask-multilang-v1.0 +hf_revision: 1c1d66bd22e0ba93a435c710a9d9499be292cd65 +inputs: + text: true + image: false + audio: false + video: false +tasks: + encode: null + score: null + extract: {} +max_sequence_length: 1024 # trained with 1024-token inputs +profiles: + default: + max_batch_tokens: 16384 + compute_precision: null + adapter_path: sie_server.adapters.gliclass:GLiClassAdapter + adapter_options: + loadtime: + classification_type: single-label + runtime: + threshold: 0.0 + multi-label: + extends: default + adapter_options: + runtime: + classification_type: multi-label + threshold: 0.0 diff --git a/packages/sie_server/pyproject.toml b/packages/sie_server/pyproject.toml index f3a25add0..2e1aed445 100644 --- a/packages/sie_server/pyproject.toml +++ b/packages/sie_server/pyproject.toml @@ -47,7 +47,9 @@ dependencies = [ "gliner>=0.2,<1", "gliner2>=1.3.1,<2", "glirel>=1.0,<2", - "gliclass>=0.1,<1", + # 0.1.17 adds the cross-attention scorer and pass-through pooling used by + # the multilingual GLiClass checkpoints; 0.1.18+ requires transformers 5. + "gliclass>=0.1.17,<1", # Docling, composite-document parser (PDF/DOCX/HTML) for extract(). # PINNED TO THE ARTIFACT REVISION, NOT FLOATED (#2872). Docling resolves # its RapidOCR det/cls/rec files by a version-specific path table, so the diff --git a/packages/sie_server/src/sie_server/adapters/gliclass/__init__.py b/packages/sie_server/src/sie_server/adapters/gliclass/__init__.py index 18e1fde31..794b323d5 100644 --- a/packages/sie_server/src/sie_server/adapters/gliclass/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/gliclass/__init__.py @@ -9,25 +9,43 @@ (100 texts, 5 labels). GLiClass is a single-pass architecture (not N×M expansion like NLI cross-encoders), so the gliclass library pipeline has minimal overhead. No separate "GLiClassFlashAdapter" is needed - the library is already efficient. + +Request surface (all optional; a request that sets none of them runs exactly the +pipeline call used before these fields existed): + +- ``instruction``: task description passed to the pipeline as ``prompt``. +- ``options.examples``: few-shot examples, ``[{"text": ..., "labels": [...]}]``. +- ``options.classification_type``: ``"single-label"`` or ``"multi-label"`` for + this request, overriding the load-time default. +- ``options.label_groups``: named label groups, ``{"urgency": ["low", "high"], + ...}``, used instead of ``labels``. Single-label scores are normalized within + each group, and ``data`` holds one answer per group. + +Usage counts each item's document tokens plus the instruction and example +texts encoded with it; label names are not counted. """ from __future__ import annotations import math import re +from dataclasses import dataclass from numbers import Real from pathlib import Path -from typing import TYPE_CHECKING, Any, ClassVar, Literal +from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast import torch +from transformers import AutoTokenizer, PreTrainedTokenizerFast from sie_server.adapters._base_adapter import BaseAdapter from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_NOT_LOADED, ComputePrecision from sie_server.adapters.errors import InputTooLongError -from sie_server.core.inference_output import ExtractOutput +from sie_server.core.extract_cost import MAX_EXTRACT_LABELS +from sie_server.core.inference_output import ExtractItemError, ExtractOutput +from sie_server.types.inputs import InvalidInputError from sie_server.types.overflow_policy import DEFAULT_OVERFLOW_POLICY, OverflowPolicy -from sie_server.types.responses import Classification +from sie_server.types.responses import Classification, ErrorCode if TYPE_CHECKING: from gliclass import ZeroShotClassificationPipeline # ty:ignore[unresolved-import] @@ -49,6 +67,136 @@ # never matches at all. _INDEX_OOB_RE = re.compile(r"index (\d+) is out of bounds for dimension \d+ with size (\d+)") ClassificationType = Literal["single-label", "multi-label"] +_CLASSIFICATION_TYPES: tuple[ClassificationType, ...] = ("single-label", "multi-label") +# gliclass joins a label group and its labels with this separator (the +# pipeline's default ``label_separator``); the model sees "urgency.high". +_LABEL_GROUP_SEPARATOR = "." +# ``ZeroShotClassificationPipeline.__call__`` default. The grouped path runs the +# model directly and mirrors the pipeline's sub-batching so padding matches. +_PIPELINE_BATCH_SIZE = 8 +_MAX_EXAMPLES = 32 +_EXAMPLE_KEYS = frozenset({"text", "labels"}) +_DEFAULT_MULTI_LABEL_THRESHOLD = 0.5 +# Free-form request text is fused into every item's input, so it is capped +# before any tokenization. The model window (512-1024 tokens) holds far less. +_MAX_INSTRUCTION_CHARS = 2048 +_MAX_EXAMPLE_TEXT_CHARS = 2048 +_MAX_CONTEXT_CHARS = 8192 +# Used when the tokenizer's vocabulary cannot be inspected. +_DEFAULT_MAX_TOKEN_CHARS = 32 +# Label names are natural language. Some vocabularies hold long whitespace or +# symbol-run tokens (hundreds of characters), but no label prompt averages more +# than this many characters per token. +_MAX_LABEL_CHARS_PER_TOKEN = 16 +# Items whose estimated fit is this close to the window edge get an exact check +# on the fused encoding: tokenizing the document alone can differ by a token +# or two from tokenizing it next to the label prompt. +_FIT_MARGIN_TOKENS = 8 +_ERR_ITEM_LABELS_TRUNCATED = ( + "The document pushes the labels out of the gliclass model's max sequence length. " + "Shorten the document, or send options.overflow_policy='truncate_text'." +) + + +@dataclass(frozen=True) +class _RequestLayout: + """Token layout shared by every item of a request, tokenized once. + + ``window`` is the number of non-special tokens the model keeps. The label + prompt (markers, label names, separator) is ``label_tokens`` long and its + last marker sits at ``last_marker``. ``context_tokens`` counts the + billable instruction and example texts. + """ + + prompt_first: bool + window: int + label_tokens: int + last_marker: int + context_tokens: int + + def labels_fit(self, document_tokens: int) -> bool: + """Whether every label marker survives truncation next to this document.""" + if self.prompt_first: + return self.last_marker < self.window + return document_tokens + self.last_marker < self.window + + +def _choice_confidence(probabilities: list[float]) -> float: + """``1 - H(p) / log(k)``: 1.0 for a certain answer, 0.0 for a uniform one.""" + if len(probabilities) < 2: + return 1.0 + entropy = -sum(p * math.log(p) for p in probabilities if p > 0.0) + return min(1.0, max(0.0, 1.0 - entropy / math.log(len(probabilities)))) + + +def _pipeline_context(prompt: str | None, examples: list[dict[str, Any]] | None) -> dict[str, Any]: + """Pipeline keyword arguments for the fields a request actually set. + + Leaving unset fields out keeps a request without them on exactly the + pipeline call the adapter has always made. + """ + context: dict[str, Any] = {} + if prompt is not None: + context["prompt"] = prompt + if examples is not None: + context["examples"] = examples + return context + + +def _longest_token_chars(tokenizer: Any) -> int: + """Length in characters of the longest entry in the tokenizer's vocabulary.""" + try: + return max(len(token) for token in tokenizer.get_vocab()) + except Exception: # noqa: BLE001 -- fall back to a generous default + return _DEFAULT_MAX_TOKEN_CHARS + + +def _string_list(value: Any) -> list[str] | None: + """Return ``value`` as a list of strings, or None when it is not one.""" + if isinstance(value, list) and all(isinstance(item, str) for item in value): + return cast("list[str]", value) + return None + + +def _restore_modernbert_rope_fields(encoder_config: Any) -> bool: + """Carry a transformers-5 ModernBERT RoPE config over to transformers 4. + + transformers 5 saves ModernBERT RoPE bases per layer type in + ``rope_parameters`` and no longer writes ``global_rope_theta`` or + ``local_rope_theta``. transformers 4 reads only the latter, so a checkpoint + saved with 5 would silently run its sliding-window layers with the 4.x + default base of 10000. Returns True when the config was changed. + """ + rope_parameters = getattr(encoder_config, "rope_parameters", None) + if ( + getattr(encoder_config, "model_type", None) != "modernbert" + or not isinstance(rope_parameters, dict) + or not hasattr(encoder_config, "local_rope_theta") # transformers 5 has no legacy fields + ): + return False + fields = {"full_attention": "global_rope_theta", "sliding_attention": "local_rope_theta"} + if not set(rope_parameters) <= set(fields): + raise ValueError( + f"ModernBERT rope_parameters keys {sorted(rope_parameters)} are not supported by transformers 4" + ) + for layer_type, field in fields.items(): + params = rope_parameters.get(layer_type) + if params is None: + continue + if params.get("rope_type", "default") != "default" or "rope_theta" not in params: + raise ValueError(f"ModernBERT {layer_type} RoPE parameters {params!r} are not supported by transformers 4") + setattr(encoder_config, field, float(params["rope_theta"])) + layer_types = getattr(encoder_config, "layer_types", None) + if layer_types is not None: + every = encoder_config.global_attn_every_n_layers + expected = [ + "full_attention" if index % every == 0 else "sliding_attention" for index in range(len(layer_types)) + ] + if list(layer_types) != expected: + raise ValueError( + "ModernBERT layer_types do not follow global_attn_every_n_layers; transformers 4 cannot load them" + ) + return True class GLiClassAdapter(BaseAdapter): @@ -64,7 +212,7 @@ class GLiClassAdapter(BaseAdapter): spec: ClassVar[AdapterSpec] = AdapterSpec( inputs=("text",), outputs=("json",), - unload_fields=("_pipeline", "_tokenizer"), + unload_fields=("_pipeline", "_pipelines", "_tokenizer"), ) def __init__( @@ -83,7 +231,9 @@ def __init__( Args: model_name_or_path: HuggingFace model ID or local path. classification_type: "single-label" for mutually exclusive classes, - "multi-label" for multiple classes per text. + "multi-label" for multiple classes per text. Callers can + override it per request via + ``options={"classification_type": ...}``. threshold: Default server-side post-filter threshold (0-1). Defaults to 0.0 so all requested labels are returned with their scores. Callers can override per-request via ``options={"threshold": ...}``. @@ -96,15 +246,19 @@ def __init__( **kwargs: Additional arguments (ignored for compatibility). """ self._model_name_or_path = str(model_name_or_path) - self._classification_type = classification_type + self._classification_type = self._validate_classification_type(classification_type) self._threshold = threshold self._max_seq_length = max_seq_length self._compute_precision = compute_precision self._revision = revision self._pipeline: ZeroShotClassificationPipeline | None = None + # One pipeline per classification type, sharing the model and + # tokenizer; ``_pipeline`` is the load-time type's entry. + self._pipelines: dict[ClassificationType, ZeroShotClassificationPipeline] | None = None self._tokenizer: PreTrainedTokenizerBase | None = None self._special_count: int = 0 + self._max_token_chars: int = _DEFAULT_MAX_TOKEN_CHARS self._device: str | None = None def load(self, device: str) -> None: @@ -113,8 +267,11 @@ def load(self, device: str) -> None: Args: device: Target device (cuda:0, cuda:1, cpu, mps). """ - from gliclass import GLiClassModel, ZeroShotClassificationPipeline # ty:ignore[unresolved-import] - from transformers import AutoTokenizer + from gliclass import ( # ty:ignore[unresolved-import] + GLiClassModel, + GLiClassModelConfig, + ZeroShotClassificationPipeline, + ) self._device = device @@ -132,9 +289,13 @@ def load(self, device: str) -> None: shared_kwargs: dict[str, Any] = {} if self._revision is not None: shared_kwargs["revision"] = self._revision - model = GLiClassModel.from_pretrained(self._model_name_or_path, **shared_kwargs) + config = GLiClassModelConfig.from_pretrained(self._model_name_or_path, **shared_kwargs) + if _restore_modernbert_rope_fields(config.encoder_config): + model = GLiClassModel.from_pretrained(self._model_name_or_path, config=config, **shared_kwargs) + else: + model = GLiClassModel.from_pretrained(self._model_name_or_path, **shared_kwargs) model = model.to(device, dtype=torch_dtype) - self._tokenizer = AutoTokenizer.from_pretrained(self._model_name_or_path, **shared_kwargs) + self._tokenizer = self._load_tokenizer(shared_kwargs) # Bound the tokenizer's max length so any internal tokenization in the # gliclass library auto-truncates to the model's actual capacity. @@ -150,19 +311,36 @@ def load(self, device: str) -> None: pipeline_kwargs: dict[str, Any] = { "model": model, "tokenizer": self._tokenizer, - "classification_type": self._classification_type, "device": device, } if self._max_seq_length is not None: pipeline_kwargs["max_length"] = self._max_seq_length self._special_count = int(self._tokenizer.num_special_tokens_to_add(pair=False)) - self._pipeline = ZeroShotClassificationPipeline(**pipeline_kwargs) + self._max_token_chars = _longest_token_chars(self._tokenizer) + self._pipelines = { + classification_type: ZeroShotClassificationPipeline( + **pipeline_kwargs, classification_type=classification_type + ) + for classification_type in _CLASSIFICATION_TYPES + } + self._pipeline = self._pipelines[self._classification_type] + + def _load_tokenizer(self, shared_kwargs: dict[str, Any]) -> PreTrainedTokenizerBase: + try: + return AutoTokenizer.from_pretrained(self._model_name_or_path, **shared_kwargs) + except ValueError as exc: + # Checkpoints saved with transformers 5 record the generic fast + # tokenizer as "TokenizersBackend", a class transformers 4 lacks. + # Their tokenizer.json loads unchanged as a PreTrainedTokenizerFast. + if "TokenizersBackend" not in str(exc): + raise + return PreTrainedTokenizerFast.from_pretrained(self._model_name_or_path, **shared_kwargs) def _extract_text(self, item: Item) -> str: if not item.text: msg = "Item must have text for classification" - raise ValueError(msg) + raise InvalidInputError(msg) return item.text def _apply_overflow_policy( @@ -170,6 +348,9 @@ def _apply_overflow_policy( texts: list[str], labels: list[str], policy: OverflowPolicy = DEFAULT_OVERFLOW_POLICY, + *, + prompt: str | None = None, + examples: list[dict[str, Any]] | None = None, ) -> list[str]: """Enforce overflow_policy by pre-tokenizing text and label_prompt separately. @@ -178,7 +359,9 @@ def _apply_overflow_policy( ``observed = text_tokens + label_prompt_tokens + special_count``, where ``special_count`` is the BERT-style ``[CLS]``/``[SEP]`` wrap (2 for all current gliclass models). We recover the same total without running the - model by tokenizing each part with ``add_special_tokens=False``. + model by tokenizing each part with ``add_special_tokens=False``. The + label prompt includes the task prompt and few-shot examples when the + request sets them, so ``truncate_text`` shortens only the document. On overflow: - ``default`` returns texts unchanged (upstream as-is — may crash inside @@ -202,7 +385,8 @@ def _apply_overflow_policy( if self._max_seq_length is None: raise RuntimeError(ERR_NOT_LOADED) - label_prompt = self._pipeline.pipe.prepare_input(text="", labels=labels) # ty:ignore[unresolved-attribute] + context = _pipeline_context(prompt, examples) + label_prompt = self._pipeline.pipe.prepare_input(text="", labels=labels, **context) # ty:ignore[unresolved-attribute] label_prompt_tokens = len(self._tokenizer(label_prompt, add_special_tokens=False)["input_ids"]) overhead = label_prompt_tokens + self._special_count budget = self._max_seq_length - overhead @@ -248,26 +432,51 @@ def extract( Args: items: List of items to classify (must have text). labels: Classification labels (e.g., ["positive", "negative", "neutral"]). - Required for zero-shot classification. + Required unless ``options["label_groups"]`` is set. output_schema: Unused (included for interface compatibility). - instruction: Unused (included for interface compatibility). + instruction: Optional task description, passed to the gliclass + pipeline as its ``prompt``. options: Adapter options to override model config defaults. - Supported: threshold (float), classification_type (str). + Supported: ``threshold`` (float), ``classification_type`` + ("single-label" or "multi-label"), ``examples`` (few-shot + ``[{"text": str, "labels": [str, ...]}]``), ``label_groups`` + (``{group: [label, ...]}``, used instead of ``labels``), and + ``overflow_policy``. Returns: ExtractOutput where ``classifications[i]`` is the list of ``Classification(label, score)`` for ``items[i]``, sorted by score - descending. + descending. With ``label_groups``, labels read ``"group.label"`` + and ``data[i]`` maps each group to its answer: ``{"type": + "choice", "choice", "probabilities", "confidence"}`` for + single-label groups, ``{"labels", "probabilities"}`` for + multi-label ones. Probabilities are never threshold-filtered. Raises: RuntimeError: If model not loaded. - ValueError: If labels not provided or items lack text, or if the - input produced an empty tensor inside the gliclass pipeline. + InvalidInputError: If labels are missing or options are malformed. + InputTooLongError: If the labels do not fit in the model window. + ValueError: If items lack text or the pipeline returns malformed + scores. """ if self._pipeline is None: raise RuntimeError(ERR_NOT_LOADED) - normalized_labels = self._validate_labels(labels) + opts = options or {} + label_groups = self._validate_label_groups(opts.get("label_groups")) + classification_type = self._validate_classification_type( + opts.get("classification_type", self._classification_type) + ) + if label_groups is None: + normalized_labels = self._validate_labels(labels) + elif labels: + raise InvalidInputError("GLiClass accepts either labels or options.label_groups, not both") + else: + normalized_labels = self._flatten_label_groups(label_groups, classification_type) + self._check_label_size(normalized_labels) + prompt = self._validate_instruction(instruction) + examples = self._validate_examples(opts.get("examples"), normalized_labels, label_groups) + self._check_context_size(prompt, examples) # Extract texts from all items (batch processing) texts = [self._extract_text(item) for item in items] @@ -275,12 +484,41 @@ def extract( # Get options with fallback to model defaults. The threshold is applied # server-side as a post-filter so we always get all label scores from the # underlying pipeline regardless of caller preferences. - opts = options or {} effective_threshold = self._validate_threshold(opts.get("threshold", self._threshold)) overflow_policy = opts.get("overflow_policy", DEFAULT_OVERFLOW_POLICY) - texts = self._apply_overflow_policy(texts, normalized_labels, overflow_policy) - input_token_counts = self._doc_input_token_counts(texts) + texts = self._apply_overflow_policy(texts, normalized_labels, overflow_policy, prompt=prompt, examples=examples) + layout = self._request_layout(normalized_labels, prompt, examples) + fits = self._items_fit(texts, layout, normalized_labels, prompt, examples) + input_token_counts = self._input_token_counts(texts, prompt, examples, layout, fits) + errors = self._item_errors(fits) + kept = [index for index, ok in enumerate(fits) if ok] + kept_texts = [texts[index] for index in kept] + + if label_groups is not None: + return self._extract_grouped( + kept_texts, + label_groups, + normalized_labels, + classification_type=classification_type, + prompt=prompt, + examples=examples, + threshold=effective_threshold, + input_token_counts=input_token_counts, + kept=kept, + errors=errors, + ) + + if not kept_texts: + return ExtractOutput( + entities=[[] for _ in items], + classifications=[[] for _ in items], + errors=errors, + input_token_counts=input_token_counts, + ) + + context = _pipeline_context(prompt, examples) + pipeline = self._select_pipeline(classification_type) # Run batch classification. # - threshold=0.0: never let the gliclass library drop labels for us @@ -290,11 +528,12 @@ def extract( # of ``{label: score}`` dicts with every requested label present. try: with torch.inference_mode(): - batch_results = self._pipeline( - texts, + batch_results = pipeline( + kept_texts, normalized_labels, threshold=0.0, return_hierarchical=True, + **context, ) except (RuntimeError, IndexError) as exc: # The gliclass library crashes inside the pipeline when inputs exceed @@ -327,8 +566,8 @@ def extract( raise InputTooLongError(_ERR_INPUT_TOO_LONG) from exc raise - all_classifications: list[list[Classification]] = [] - for item_results in batch_results: + all_classifications: list[list[Classification]] = [[] for _ in items] + for index, item_results in zip(kept, batch_results, strict=True): # With return_hierarchical=True and a flat label list the library # returns a dict {label: score}. Anything else (e.g. None for an # empty input) yields no classifications rather than crashing. @@ -348,35 +587,292 @@ def extract( # Sort by score descending classifications.sort(key=lambda x: x["score"], reverse=True) - all_classifications.append(classifications) + all_classifications[index] = classifications return ExtractOutput( entities=[[] for _ in items], classifications=all_classifications, + errors=errors, + input_token_counts=input_token_counts, + ) + + def _select_pipeline(self, classification_type: ClassificationType) -> ZeroShotClassificationPipeline: + if classification_type == self._classification_type: + if self._pipeline is None: + raise RuntimeError(ERR_NOT_LOADED) + return self._pipeline + if self._pipelines is None: + raise RuntimeError(ERR_NOT_LOADED) + return self._pipelines[classification_type] + + def _extract_grouped( + self, + texts: list[str], + label_groups: list[tuple[str, list[str]]], + flat_labels: list[str], + *, + classification_type: ClassificationType, + prompt: str | None, + examples: list[dict[str, Any]] | None, + threshold: float, + input_token_counts: list[int] | None, + kept: list[int], + errors: list[ExtractItemError | None] | None, + ) -> ExtractOutput: + """Score grouped labels and return one answer per group. + + All groups share one forward pass per text, exactly as the gliclass + pipeline encodes a dict of labels ("group.label" after flattening). + The pipeline's own single-label mode applies one softmax across every + flattened label, which makes a group's scores depend on the other + groups. Here each group gets its own softmax over the raw logits + instead; multi-label scores are independent sigmoids either way. + + A single-label group answers like a choice question: + ``{"type": "choice", "choice", "probabilities", "confidence"}`` with + ``confidence = 1 - H(p) / log(k)``. A multi-label group answers + ``{"labels", "probabilities"}``. + """ + rows = ( + self._grouped_scores(texts, label_groups, flat_labels, classification_type, prompt, examples) + if texts + else [] + ) + item_count = len(errors) if errors is not None else len(texts) + # Multi-label groups list the labels at or above the request threshold, + # or at or above 0.5 (an even sigmoid) when no threshold is set. + selection_threshold = threshold if threshold > 0.0 else _DEFAULT_MULTI_LABEL_THRESHOLD + + all_classifications: list[list[Classification]] = [[] for _ in range(item_count)] + all_data: list[dict[str, Any]] = [{} for _ in range(item_count)] + for index, row in zip(kept, rows, strict=True): + position = 0 + answers: dict[str, dict[str, Any]] = {} + classifications: list[Classification] = [] + for group, group_labels in label_groups: + probabilities: dict[str, float] = {} + for label in group_labels: + score = self._validate_score(row[position]) + position += 1 + probabilities[label] = score + classifications.append(Classification(label=f"{group}{_LABEL_GROUP_SEPARATOR}{label}", score=score)) + if classification_type == "single-label": + answers[group] = { + "type": "choice", + "choice": max(probabilities, key=probabilities.__getitem__), + "probabilities": probabilities, + "confidence": _choice_confidence(list(probabilities.values())), + } + else: + answers[group] = { + "labels": [label for label, score in probabilities.items() if score >= selection_threshold], + "probabilities": probabilities, + } + if threshold > 0.0: + classifications = [c for c in classifications if c["score"] >= threshold] + classifications.sort(key=lambda x: x["score"], reverse=True) + all_classifications[index] = classifications + all_data[index] = answers + + return ExtractOutput( + entities=[[] for _ in range(item_count)], + classifications=all_classifications, + data=all_data, + errors=errors, input_token_counts=input_token_counts, ) + def _grouped_scores( + self, + texts: list[str], + label_groups: list[tuple[str, list[str]]], + flat_labels: list[str], + classification_type: ClassificationType, + prompt: str | None, + examples: list[dict[str, Any]] | None, + ) -> list[list[float]]: + """Run the model on flattened group labels and normalize per group.""" + if self._pipeline is None: + raise RuntimeError(ERR_NOT_LOADED) + pipe = self._pipeline.pipe # ty:ignore[unresolved-attribute] + model = pipe.model + num_labels = len(flat_labels) + forward_kwargs: dict[str, Any] = {} + resolve_max_num_classes = getattr(pipe, "_resolve_max_num_classes", None) + if resolve_max_num_classes is not None: + forward_kwargs["max_num_classes"] = resolve_max_num_classes(flat_labels, True) + + chunks: list[torch.Tensor] = [] + with torch.inference_mode(): + for start in range(0, len(texts), _PIPELINE_BATCH_SIZE): + inputs = pipe.prepare_inputs( + texts[start : start + _PIPELINE_BATCH_SIZE], + flat_labels, + same_labels=True, + examples=examples, + prompt=prompt, + ) + logits = model(**inputs, **forward_kwargs).logits + if logits.shape[-1] < num_labels: + raise InputTooLongError(_ERR_INPUT_TOO_LONG) + chunks.append(logits[:, :num_labels].float()) + + logits = torch.cat(chunks) + if classification_type == "multi-label": + scores = torch.sigmoid(logits) + else: + scores = torch.empty_like(logits) + start = 0 + for _, group_labels in label_groups: + end = start + len(group_labels) + scores[:, start:end] = torch.softmax(logits[:, start:end], dim=-1) + start = end + return scores.cpu().tolist() + + @staticmethod + def _validate_classification_type(value: object) -> ClassificationType: + if value == "single-label": + return "single-label" + if value == "multi-label": + return "multi-label" + raise InvalidInputError("GLiClass classification_type must be 'single-label' or 'multi-label'") + + @staticmethod + def _validate_instruction(instruction: object) -> str | None: + if instruction is None: + return None + if not isinstance(instruction, str): + raise InvalidInputError("GLiClass instruction must be a string") + if len(instruction) > _MAX_INSTRUCTION_CHARS: + raise InvalidInputError(f"GLiClass instruction must be at most {_MAX_INSTRUCTION_CHARS} characters") + return instruction if instruction.strip() else None + + @classmethod + def _validate_label_groups(cls, value: Any) -> list[tuple[str, list[str]]] | None: + if value is None: + return None + if not isinstance(value, dict) or not value: + raise InvalidInputError( + "GLiClass label_groups must be a non-empty object mapping group names to label lists" + ) + groups: list[tuple[str, list[str]]] = [] + for name, group_labels in value.items(): + if not isinstance(name, str) or not name.strip(): + raise InvalidInputError("GLiClass label_groups names must be non-empty strings") + if not isinstance(group_labels, list) or not group_labels: + raise InvalidInputError(f"GLiClass label_groups[{name!r}] must be a non-empty list of labels") + groups.append((name.strip(), cls._validate_labels(group_labels))) + if len({name for name, _ in groups}) != len(groups): + raise InvalidInputError("GLiClass label_groups names must be unique") + return groups + + @staticmethod + def _flatten_label_groups( + label_groups: list[tuple[str, list[str]]], + classification_type: ClassificationType, + ) -> list[str]: + flat: list[str] = [] + for name, group_labels in label_groups: + if classification_type == "single-label" and len(group_labels) < 2: + raise InvalidInputError(f"GLiClass single-label group {name!r} needs at least two labels") + flat.extend(f"{name}{_LABEL_GROUP_SEPARATOR}{label}" for label in group_labels) + if len(flat) > MAX_EXTRACT_LABELS: + raise InvalidInputError(f"GLiClass label_groups must contain at most {MAX_EXTRACT_LABELS} labels in total") + if len(set(flat)) != len(flat): + raise InvalidInputError( + f"GLiClass label_groups repeat a label once group and label names are joined with " + f"{_LABEL_GROUP_SEPARATOR!r}" + ) + return flat + + @staticmethod + def _validate_examples( + value: Any, + labels: list[str], + label_groups: list[tuple[str, list[str]]] | None, + ) -> list[dict[str, Any]] | None: + """Normalize few-shot examples to the pipeline's ``{"text", "labels"}`` form. + + Example labels must come from the request's own label set. With + ``label_groups`` they may be written as ``"group.label"`` strings or + as an object mapping a group to one label or a list of labels. + """ + if value is None: + return None + if not isinstance(value, list): + raise InvalidInputError("GLiClass examples must be a list of {text, labels} objects") + if not value: + return None + if len(value) > _MAX_EXAMPLES: + raise InvalidInputError(f"GLiClass examples must contain at most {_MAX_EXAMPLES} entries") + allowed = set(labels) + groups = dict(label_groups) if label_groups is not None else None + normalized: list[dict[str, Any]] = [] + for index, entry in enumerate(value): + where = f"GLiClass examples[{index}]" + if not isinstance(entry, dict) or set(entry) != _EXAMPLE_KEYS: + raise InvalidInputError(f"{where} must be an object with exactly 'text' and 'labels'") + example = cast("dict[str, Any]", entry) + text = example["text"] + if not isinstance(text, str) or not text.strip(): + raise InvalidInputError(f"{where}.text must be a non-empty string") + if len(text) > _MAX_EXAMPLE_TEXT_CHARS: + raise InvalidInputError(f"{where}.text must be at most {_MAX_EXAMPLE_TEXT_CHARS} characters") + raw_labels = example["labels"] + example_labels: list[str] = [] + # An example can name each requested label at most once; a longer + # list is refused before it is walked. + if isinstance(raw_labels, (dict, list)) and len(raw_labels) > len(labels): + raise InvalidInputError(f"{where}.labels names more labels than the request has") + if isinstance(raw_labels, dict): + if groups is None: + raise InvalidInputError(f"{where}.labels can be an object only when label_groups is set") + for group, chosen in cast("dict[Any, Any]", raw_labels).items(): + if not isinstance(group, str) or group.strip() not in groups: + raise InvalidInputError(f"{where}.labels names an unknown group {group!r}") + chosen_labels = _string_list([chosen] if isinstance(chosen, str) else chosen) + if chosen_labels is None: + raise InvalidInputError(f"{where}.labels[{group!r}] must be a label or a list of labels") + if len(chosen_labels) > len(groups[group.strip()]): + raise InvalidInputError(f"{where}.labels[{group!r}] names more labels than the group has") + for label in chosen_labels: + if label.strip() not in groups[group.strip()]: + raise InvalidInputError(f"{where} uses {label!r}, which is not a label of group {group!r}") + example_labels.append(f"{group.strip()}{_LABEL_GROUP_SEPARATOR}{label.strip()}") + elif isinstance(raw_labels, list): + flat_labels = _string_list(raw_labels) + if flat_labels is None: + raise InvalidInputError(f"{where}.labels must be a list of strings") + example_labels = [label.strip() for label in flat_labels] + unknown = [label for label in example_labels if label not in allowed] + if unknown: + raise InvalidInputError(f"{where} uses labels outside the requested label set: {unknown}") + else: + raise InvalidInputError(f"{where}.labels must be a list of labels") + normalized.append({"text": text, "labels": list(dict.fromkeys(example_labels))}) + return normalized + @staticmethod def _validate_labels(labels: list[str] | None) -> list[str]: if not labels: - raise ValueError(_ERR_REQUIRES_LABELS) + raise InvalidInputError(_ERR_REQUIRES_LABELS) if any(not isinstance(label, str) or not label.strip() for label in labels): - raise ValueError("GLiClass labels must be non-empty strings") + raise InvalidInputError("GLiClass labels must be non-empty strings") normalized = [label.strip() for label in labels] if len(set(normalized)) != len(normalized): - raise ValueError("GLiClass labels must be unique") + raise InvalidInputError("GLiClass labels must be unique") return normalized @staticmethod def _validate_threshold(value: object) -> float: if isinstance(value, bool) or not isinstance(value, Real): - raise ValueError("GLiClass threshold must be a finite number between 0 and 1") + raise InvalidInputError("GLiClass threshold must be a finite number between 0 and 1") try: threshold = float(value) except OverflowError as exc: - raise ValueError("GLiClass threshold must be a finite number between 0 and 1") from exc + raise InvalidInputError("GLiClass threshold must be a finite number between 0 and 1") from exc if not math.isfinite(threshold) or not 0.0 <= threshold <= 1.0: - raise ValueError("GLiClass threshold must be a finite number between 0 and 1") + raise InvalidInputError("GLiClass threshold must be a finite number between 0 and 1") return threshold @staticmethod @@ -391,6 +887,200 @@ def _validate_score(value: object) -> float: raise ValueError("GLiClass returned an invalid classification score") return score + def _request_layout( + self, + labels: list[str], + prompt: str | None, + examples: list[dict[str, Any]] | None, + ) -> _RequestLayout | None: + """Tokenize the parts every item shares, once per request. + + Returns None when the model does not put label markers in its input + (only uni-encoder GLiClass models do), so there is nothing to check. + + Raises: + InputTooLongError: If the label prompt alone overflows the window. + InvalidInputError: If the instruction and examples leave no room + for the document. + """ + pipe = getattr(self._pipeline, "pipe", None) + config = getattr(getattr(pipe, "model", None), "config", None) + tokenizer = self._tokenizer + if ( + pipe is None + or config is None + or tokenizer is None + or getattr(config, "architecture_type", None) != "uni-encoder" + ): + return None + + def count(text: str) -> int: + return len(tokenizer(text, add_special_tokens=False)["input_ids"]) + + label_prompt = "".join(f"{pipe.label_token}{label}" for label in labels) + pipe.sep_token + label_ids = tokenizer(label_prompt, add_special_tokens=False)["input_ids"] + markers = [position for position, token in enumerate(label_ids) if token == config.class_token_index] + if len(markers) < len(labels): + return None + window = pipe.max_length - self._special_count + if markers[-1] >= window: + raise InputTooLongError(_ERR_INPUT_TOO_LONG) + prompt_tokens = count(prompt) if prompt else 0 + example_text_tokens = sum(count(example["text"]) for example in examples or []) + if prompt or examples: + examples_tokens = count(pipe._format_examples_for_input(examples)) if examples else 0 + if len(label_ids) + prompt_tokens + examples_tokens >= window: + raise InvalidInputError( + "GLiClass instruction, examples and labels leave no room for the document in the " + f"model's {window}-token window; shorten the instruction or send fewer examples" + ) + return _RequestLayout( + prompt_first=bool(getattr(config, "prompt_first", False)), + window=window, + label_tokens=len(label_ids), + last_marker=markers[-1], + context_tokens=prompt_tokens + example_text_tokens, + ) + + def _items_fit( + self, + texts: list[str], + layout: _RequestLayout | None, + labels: list[str], + prompt: str | None, + examples: list[dict[str, Any]] | None, + ) -> list[bool]: + """Whether each item's label markers survive truncation. + + Documents are tokenized once, on their own. Tokenizing a document next + to the label prompt can differ from that by a token or two (a trailing + space before a marker, for example), so items within a few tokens of + the window edge are checked exactly on their fused encoding. + """ + if layout is None or layout.prompt_first or self._tokenizer is None: + return [True] * len(texts) + encoded = self._tokenizer(texts, add_special_tokens=False, truncation=True, max_length=layout.window) + fits: list[bool] = [] + for text, ids in zip(texts, encoded["input_ids"], strict=True): + estimate = len(ids) + layout.last_marker + if estimate < layout.window - _FIT_MARGIN_TOKENS: + fits.append(True) + elif estimate >= layout.window + _FIT_MARGIN_TOKENS: + fits.append(False) + else: + fits.append(self._labels_survive(text, list(ids), labels, prompt, examples)) + return fits + + def _labels_survive( + self, + text: str, + text_ids: list[int], + labels: list[str], + prompt: str | None, + examples: list[dict[str, Any]] | None, + ) -> bool: + """Exact check: count the label markers left in the truncated fused input.""" + pipe = getattr(self._pipeline, "pipe", None) + tokenizer = self._tokenizer + if pipe is None or tokenizer is None: + return False + marker = pipe.model.config.class_token_index + fused = pipe.prepare_input(text, labels, examples, prompt) + ids = tokenizer(fused, truncation=True, max_length=pipe.max_length)["input_ids"] + # Markers written in the document itself come first and are not labels. + return ids.count(marker) >= text_ids.count(marker) + len(labels) + + @staticmethod + def _item_errors(fits: list[bool]) -> list[ExtractItemError | None] | None: + """Per-item INPUT_TOO_LONG for items whose labels would be cut off. + + Reporting these per item keeps one oversized document from failing + every request batched with it. + """ + if all(fits): + return None + return [ + None if ok else ExtractItemError(code=ErrorCode.INPUT_TOO_LONG.value, message=_ERR_ITEM_LABELS_TRUNCATED) + for ok in fits + ] + + def _input_token_counts( + self, + texts: list[str], + prompt: str | None, + examples: list[dict[str, Any]] | None, + layout: _RequestLayout | None, + fits: list[bool], + ) -> list[int] | None: + """Billable input tokens per item. + + Each item is billed for its document and for the free-form request + text encoded with it: the instruction and the few-shot example texts. + Label names (including the labels attached to examples) and the + pipeline's marker tokens are not billed. With an instruction or + examples, an item's total is capped at the model window minus the + label prompt, the most free-form text it can encode, unless the + document count alone is already higher. Items refused because their + labels would be cut off are billed nothing. + """ + counts = self._doc_input_token_counts(texts) + if counts is None: + return None + if prompt or examples: + if layout is not None: + context_tokens = layout.context_tokens + cap = layout.window + self._special_count - layout.label_tokens + else: + context_tokens = self._context_token_count(prompt, examples) + if context_tokens is None: + return None + cap = self._max_seq_length + counts = [count if cap is None else max(count, min(count + context_tokens, cap)) for count in counts] + if cap is None: + counts = [count + context_tokens for count in counts] + return [count if ok else 0 for count, ok in zip(counts, fits, strict=True)] + + def _context_token_count(self, prompt: str | None, examples: list[dict[str, Any]] | None) -> int | None: + if self._tokenizer is None: + return None + parts = ([prompt] if prompt else []) + [example["text"] for example in examples or []] + try: + return sum(len(ids) for ids in self._tokenizer(parts, add_special_tokens=False)["input_ids"]) + except Exception: # noqa: BLE001 -- metering must not fail classification + return None + + def _check_label_size(self, labels: list[str]) -> None: + """Refuse label sets whose text cannot fit the model window, before tokenizing. + + A token covers at most ``_max_token_chars`` characters (the longest + vocabulary entry), and label text averages far fewer: at most + ``_MAX_LABEL_CHARS_PER_TOKEN``. A label prompt longer than that many + characters per token of the window needs more tokens than the window + holds, so it can never be scored correctly, and only tokenizing it + would take time proportional to its size. + """ + window = self._max_seq_length or getattr(getattr(self._pipeline, "pipe", None), "max_length", None) + if not isinstance(window, int): + return + limit = min(self._max_token_chars, _MAX_LABEL_CHARS_PER_TOKEN) * window + total = sum(len(label) for label in labels) + if total > limit: + raise InvalidInputError( + f"GLiClass labels total {total} characters; at most {limit} can fit in the model's " + f"{window}-token window" + ) + + @staticmethod + def _check_context_size(prompt: str | None, examples: list[dict[str, Any]] | None) -> None: + """Bound the free-form text fused into every item, before any tokenization.""" + size = len(prompt or "") + for example in examples or []: + size += len(example["text"]) + sum(len(label) for label in example["labels"]) + if size > _MAX_CONTEXT_CHARS: + raise InvalidInputError( + f"GLiClass instruction and examples must total at most {_MAX_CONTEXT_CHARS} characters" + ) + def _doc_input_token_counts(self, texts: list[str]) -> list[int] | None: """Count document-only model-tokenizer input units for billing. diff --git a/packages/sie_server/src/sie_server/api/extract.py b/packages/sie_server/src/sie_server/api/extract.py index 64d3181e2..77b4c559f 100644 --- a/packages/sie_server/src/sie_server/api/extract.py +++ b/packages/sie_server/src/sie_server/api/extract.py @@ -383,6 +383,12 @@ async def extract( # Request-level instruction takes precedence; fall back to profile instruction if instruction is None: instruction = options.get("instruction") + if instruction is not None and not isinstance(instruction, str): + span.set_attribute("error", "invalid_instruction") + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"code": ErrorCode.INVALID_INPUT.value, "message": "instruction must be a string"}, + ) items = request.items diff --git a/packages/sie_server/src/sie_server/core/worker/handlers/extract.py b/packages/sie_server/src/sie_server/core/worker/handlers/extract.py index 64749db3a..93d446d4f 100644 --- a/packages/sie_server/src/sie_server/core/worker/handlers/extract.py +++ b/packages/sie_server/src/sie_server/core/worker/handlers/extract.py @@ -7,6 +7,8 @@ from typing import TYPE_CHECKING, Any +from sie_sdk._msgpack import packb as pack_msgpack + from sie_server.core.inference_output import ExtractOutput from sie_server.core.worker.handlers.base import OperationHandler, make_hashable @@ -18,6 +20,22 @@ from sie_server.types.responses import Classification, DetectedObject, Relation +def _ordered_key(value: dict[str, Any]) -> Any: + """Batching key that keeps object key order. + + Requests in one group all run with the first request's options and + schema, and key order can be part of their meaning (named label groups, + schema properties). Requests that differ only in order must therefore + land in different groups. The msgpack encoding is the key the queue + executor already uses for the same purpose. + """ + try: + return pack_msgpack(value, use_bin_type=True) + except (TypeError, ValueError): + # Wire requests always encode; keep the old key for anything else. + return make_hashable(value) + + class ExtractHandler(OperationHandler[ExtractOutput]): """Handler for extract (NER, RE) operations. @@ -40,8 +58,8 @@ def make_config_key(self, metadata: RequestMetadata) -> tuple[Any, ...]: """ # Label order is part of the model input for label-conditioned extractors. labels_key = tuple(metadata.labels) if metadata.labels else None - schema_key = make_hashable(metadata.output_schema) if metadata.output_schema is not None else None - options_key = make_hashable(metadata.options) if metadata.options else None + schema_key = _ordered_key(metadata.output_schema) if metadata.output_schema is not None else None + options_key = _ordered_key(metadata.options) if metadata.options else None return ( labels_key, schema_key, @@ -69,9 +87,12 @@ def run_inference( Returns: ExtractOutput with entities. """ - labels_tuple, _schema_key, instruction, options_tuple = config_key + labels_tuple, _schema_key, instruction, _options_key = config_key labels = list(labels_tuple) if labels_tuple else None - options = dict(options_tuple) if options_tuple else None + # Read original options from metadata like encode and score do: the + # config key turns nested lists and objects into tuples, which would + # hand adapters a lossy copy of structured options. + options = metadata_list[0].options or None output_schema = metadata_list[0].output_schema return adapter.extract( diff --git a/packages/sie_server/src/sie_server/core/worker/model_worker.py b/packages/sie_server/src/sie_server/core/worker/model_worker.py index 4d802e9a6..b92d29c4c 100644 --- a/packages/sie_server/src/sie_server/core/worker/model_worker.py +++ b/packages/sie_server/src/sie_server/core/worker/model_worker.py @@ -44,6 +44,7 @@ WorkerStats, ) from sie_server.observability.worker_telemetry import worker_telemetry, worker_telemetry_enabled +from sie_server.types.inputs import InvalidInputError if TYPE_CHECKING: from sie_server.adapters.base import ModelAdapter @@ -1484,6 +1485,24 @@ async def _dispatch( self._stats.batches_processed += 1 self._stats.total_tokens_processed += batch.total_tokens + def _config_key_or_fail(self, metadata: RequestMetadata) -> tuple[Any, ...] | None: + """Batching key for one request, or None after failing that request alone. + + The key is built from caller-supplied fields. A value that cannot be + hashed must fail only the request that sent it, never the requests + batched with it. + """ + try: + handler = self._handlers[metadata.operation] + config_key = (metadata.operation, *handler.make_config_key(metadata)) + hash(config_key) + except Exception as exc: # noqa: BLE001 -- isolate one malformed request from its batch + logger.warning("Rejecting request with an unbatchable configuration: %s", exc) + if not metadata.future.done(): + metadata.future.set_exception(InvalidInputError(f"Request options cannot be processed: {exc}")) + return None + return config_key + def _group_by_inference_config( self, batch: FormattedBatch[HasCost, RequestMetadata], @@ -1509,9 +1528,9 @@ def _group_by_inference_config( metadata_list = batch.metadata if len(metadata_list) > 1 and all(m is metadata_list[0] for m in metadata_list): first_meta = metadata_list[0] - handler = self._handlers[first_meta.operation] - handler_key = handler.make_config_key(first_meta) - config_key = (first_meta.operation, *handler_key) + config_key = self._config_key_or_fail(first_meta) + if config_key is None: + return {} items_list: list[Item] = [] indices_list: list[int] = [] @@ -1527,11 +1546,14 @@ def _group_by_inference_config( tuple[list[Item], list[RequestMetadata], list[int], list[HasCost]], ] = {} + failed: set[int] = set() for prepared_item, metadata in zip(batch.items, metadata_list, strict=True): - # Get handler and create config key - handler = self._handlers[metadata.operation] - handler_key = handler.make_config_key(metadata) - config_key = (metadata.operation, *handler_key) + if id(metadata) in failed: + continue + config_key = self._config_key_or_fail(metadata) + if config_key is None: + failed.add(id(metadata)) + continue if config_key not in groups: groups[config_key] = ([], [], [], []) diff --git a/packages/sie_server/src/sie_server/queue_executor.py b/packages/sie_server/src/sie_server/queue_executor.py index a734c19f3..65c13a06a 100644 --- a/packages/sie_server/src/sie_server/queue_executor.py +++ b/packages/sie_server/src/sie_server/queue_executor.py @@ -12,6 +12,7 @@ import yaml from sie_sdk._msgpack import packb as pack_msgpack +from sie_server.adapters.errors import InputTooLongError from sie_server.api.ws import compute_bundle_config_hash_cached from sie_server.config.model import ModelConfig from sie_server.core.encode_pipeline import EncodePipeline, resolve_encode_output_types @@ -1423,6 +1424,12 @@ async def process_extract_batch(self, req: ProcessExtractBatchRequest) -> BatchO for bi in req.items: try: options = merge_runtime_options(config, bi.options) + # Same precedence as the HTTP extract path: the request's own + # instruction, else one from the options (profile defaults + # included). + instruction = bi.instruction if bi.instruction is not None else options.get("instruction") + if instruction is not None and not isinstance(instruction, str): + raise InvalidInputError("instruction must be a string") server_item = decode_item(bi.item) timing = RequestTiming() timing.start_tokenization() @@ -1464,7 +1471,7 @@ async def process_extract_batch(self, req: ProcessExtractBatchRequest) -> BatchO model_id, [server_item], config, - instruction=bi.instruction, + instruction=instruction, task=task, ) prepared_items = prepared_batch.items @@ -1477,7 +1484,7 @@ async def process_extract_batch(self, req: ProcessExtractBatchRequest) -> BatchO [server_item], labels=bi.labels, output_schema=bi.output_schema, - instruction=bi.instruction, + instruction=instruction, options=options, ) prepared_items = build_extract_prepared_items([server_item], item_costs=item_costs) @@ -1490,7 +1497,7 @@ async def process_extract_batch(self, req: ProcessExtractBatchRequest) -> BatchO items=[server_item], labels=bi.labels, output_schema=bi.output_schema, - instruction=bi.instruction, + instruction=instruction, options=options, request_id=bi.request_id, timing=timing, @@ -2009,6 +2016,10 @@ def _inference_exception_outcome( # park items in a batcher today, so this arm is a contract guard # against a future caller that submits through the queueing path. return _nak_outcome(bi) + if isinstance(exc, InputTooLongError): + # The input exceeds the model's window: INPUT_TOO_LONG (HTTP 400), as + # the HTTP path reports it, not a server-side inference failure. + return _error_outcome(bi, ErrorCode.INPUT_TOO_LONG.value, str(exc)) if isinstance(exc, (InvalidInputError, msgspec.ValidationError)): # A typed-decode failure (decode_item) or a media contract violation; # both surface as INVALID_INPUT (HTTP 400), matching the HTTP path. diff --git a/packages/sie_server/src/sie_server/types/requests.py b/packages/sie_server/src/sie_server/types/requests.py index 3195e26f4..8c2c14a66 100644 --- a/packages/sie_server/src/sie_server/types/requests.py +++ b/packages/sie_server/src/sie_server/types/requests.py @@ -58,6 +58,11 @@ class ExtractParams(msgspec.Struct): def __post_init__(self) -> None: if self.labels is not None and len(self.labels) > MAX_EXTRACT_LABELS: raise msgspec.ValidationError(f"Field 'labels' must contain at most {MAX_EXTRACT_LABELS} labels") + if self.options is not None and "instruction" in self.options: + try: + msgspec.convert(self.options["instruction"], type=str | None, strict=True) + except msgspec.ValidationError as exc: + raise msgspec.ValidationError(f"{exc} - at `$.params.options.instruction`") from exc class ExtractRequest(msgspec.Struct): diff --git a/packages/sie_server/tests/adapters/test_gliclass_contracts.py b/packages/sie_server/tests/adapters/test_gliclass_contracts.py index 52d721ec6..2a3b91ff6 100644 --- a/packages/sie_server/tests/adapters/test_gliclass_contracts.py +++ b/packages/sie_server/tests/adapters/test_gliclass_contracts.py @@ -95,3 +95,97 @@ def test_pipeline_label_set_must_match_requested_labels(result: dict[str, float] [Item(text="hello")], labels=["positive", "negative"], ) + + +class TestTokenizerLoading: + def test_transformers5_tokenizer_class_falls_back_to_fast_tokenizer(self, monkeypatch: pytest.MonkeyPatch) -> None: + import transformers + + calls: list[tuple[str, dict[str, object]]] = [] + + def auto(name: str, **kwargs: object) -> object: + raise ValueError("Tokenizer class TokenizersBackend does not exist or is not currently imported.") + + def fast(name: str, **kwargs: object) -> str: + calls.append((name, kwargs)) + return "fast-tokenizer" + + monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", auto) + monkeypatch.setattr(transformers.PreTrainedTokenizerFast, "from_pretrained", fast) + adapter = GLiClassAdapter("org/model", revision="abc123") + + assert adapter._load_tokenizer({"revision": "abc123"}) == "fast-tokenizer" + assert calls == [("org/model", {"revision": "abc123"})] + + def test_other_tokenizer_errors_propagate(self, monkeypatch: pytest.MonkeyPatch) -> None: + import transformers + + def auto(name: str, **kwargs: object) -> object: + raise ValueError("Unrecognized model identifier") + + monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", auto) + + with pytest.raises(ValueError, match="Unrecognized model identifier"): + GLiClassAdapter("org/model")._load_tokenizer({}) + + +class TestModernBertRopeCompatibility: + """Checkpoints saved with transformers 5 keep ModernBERT RoPE bases in rope_parameters.""" + + @staticmethod + def _config(**overrides: object) -> object: + from transformers import ModernBertConfig + + fields: dict[str, object] = {"num_hidden_layers": 4, "global_attn_every_n_layers": 3} + fields.update(overrides) + return ModernBertConfig(**fields) + + def test_rope_parameters_fill_the_transformers4_fields(self) -> None: + from sie_server.adapters.gliclass import _restore_modernbert_rope_fields + + config = self._config( + rope_parameters={ + "full_attention": {"rope_theta": 160000.0, "rope_type": "default"}, + "sliding_attention": {"rope_theta": 160000.0, "rope_type": "default"}, + }, + layer_types=["full_attention", "sliding_attention", "sliding_attention", "full_attention"], + ) + assert config.local_rope_theta == 10000.0 # the transformers 4 default a v5 checkpoint would get + + assert _restore_modernbert_rope_fields(config) is True + assert config.local_rope_theta == 160000.0 + assert config.global_rope_theta == 160000.0 + + def test_transformers4_checkpoints_are_left_alone(self) -> None: + from sie_server.adapters.gliclass import _restore_modernbert_rope_fields + + config = self._config(local_rope_theta=10000.0, global_rope_theta=160000.0) + + assert _restore_modernbert_rope_fields(config) is False + assert config.local_rope_theta == 10000.0 + + def test_non_modernbert_encoders_are_left_alone(self) -> None: + from sie_server.adapters.gliclass import _restore_modernbert_rope_fields + from transformers import DebertaV2Config + + assert _restore_modernbert_rope_fields(DebertaV2Config()) is False + + @pytest.mark.parametrize( + ("overrides", "match"), + [ + ({"rope_parameters": {"sliding_attention": {"rope_theta": 1.0, "rope_type": "yarn"}}}, "not supported"), + ({"rope_parameters": {"rope_theta": 1.0, "rope_type": "default"}}, "keys"), + ( + { + "rope_parameters": {"full_attention": {"rope_theta": 1.0}}, + "layer_types": ["full_attention", "full_attention", "sliding_attention", "full_attention"], + }, + "layer_types", + ), + ], + ) + def test_layouts_transformers4_cannot_express_are_rejected(self, overrides: dict[str, object], match: str) -> None: + from sie_server.adapters.gliclass import _restore_modernbert_rope_fields + + with pytest.raises(ValueError, match=match): + _restore_modernbert_rope_fields(self._config(**overrides)) diff --git a/packages/sie_server/tests/adapters/test_gliclass_request_options.py b/packages/sie_server/tests/adapters/test_gliclass_request_options.py new file mode 100644 index 000000000..3af62679a --- /dev/null +++ b/packages/sie_server/tests/adapters/test_gliclass_request_options.py @@ -0,0 +1,794 @@ +"""GLiClass request fields: instruction, examples, classification_type, label_groups.""" + +from __future__ import annotations + +import math +import re +from types import SimpleNamespace +from typing import Any +from unittest.mock import MagicMock + +import pytest +import torch +from sie_server.adapters.errors import InputTooLongError +from sie_server.adapters.gliclass import GLiClassAdapter +from sie_server.types.inputs import InvalidInputError, Item + +_CLASS_TOKEN = 7 +_SPECIAL_IDS = {"<