augustoFranke commited on
Commit
cb96101
·
verified ·
1 Parent(s): 283836c

GLiNER2.5-Decide Core ML multifunction package (L64-L512) with bucket router

Browse files
.gitignore ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ .venv/
2
+ __pycache__/
3
+ .pytest_cache/
4
+ *.egg-info/
5
+ bench/benchbuckets
6
+ bench/plan
7
+
8
+ # Compiled Core ML model: rebuilt from the .mlpackage on first load.
9
+ models/*.mlmodelc/
10
+
11
+ # The .mlpackage is ~900 MB. Keep it out of plain git; use Git LFS (or rebuild it with
12
+ # scripts/build) if this folder ever becomes a repository.
13
+ models/*.mlpackage/
LICENSE ADDED
@@ -0,0 +1,202 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ Apache License
3
+ Version 2.0, January 2004
4
+ http://www.apache.org/licenses/
5
+
6
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
7
+
8
+ 1. Definitions.
9
+
10
+ "License" shall mean the terms and conditions for use, reproduction,
11
+ and distribution as defined by Sections 1 through 9 of this document.
12
+
13
+ "Licensor" shall mean the copyright owner or entity authorized by
14
+ the copyright owner that is granting the License.
15
+
16
+ "Legal Entity" shall mean the union of the acting entity and all
17
+ other entities that control, are controlled by, or are under common
18
+ control with that entity. For the purposes of this definition,
19
+ "control" means (i) the power, direct or indirect, to cause the
20
+ direction or management of such entity, whether by contract or
21
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
22
+ outstanding shares, or (iii) beneficial ownership of such entity.
23
+
24
+ "You" (or "Your") shall mean an individual or Legal Entity
25
+ exercising permissions granted by this License.
26
+
27
+ "Source" form shall mean the preferred form for making modifications,
28
+ including but not limited to software source code, documentation
29
+ source, and configuration files.
30
+
31
+ "Object" form shall mean any form resulting from mechanical
32
+ transformation or translation of a Source form, including but
33
+ not limited to compiled object code, generated documentation,
34
+ and conversions to other media types.
35
+
36
+ "Work" shall mean the work of authorship, whether in Source or
37
+ Object form, made available under the License, as indicated by a
38
+ copyright notice that is included in or attached to the work
39
+ (an example is provided in the Appendix below).
40
+
41
+ "Derivative Works" shall mean any work, whether in Source or Object
42
+ form, that is based on (or derived from) the Work and for which the
43
+ editorial revisions, annotations, elaborations, or other modifications
44
+ represent, as a whole, an original work of authorship. For the purposes
45
+ of this License, Derivative Works shall not include works that remain
46
+ separable from, or merely link (or bind by name) to the interfaces of,
47
+ the Work and Derivative Works thereof.
48
+
49
+ "Contribution" shall mean any work of authorship, including
50
+ the original version of the Work and any modifications or additions
51
+ to that Work or Derivative Works thereof, that is intentionally
52
+ submitted to Licensor for inclusion in the Work by the copyright owner
53
+ or by an individual or Legal Entity authorized to submit on behalf of
54
+ the copyright owner. For the purposes of this definition, "submitted"
55
+ means any form of electronic, verbal, or written communication sent
56
+ to the Licensor or its representatives, including but not limited to
57
+ communication on electronic mailing lists, source code control systems,
58
+ and issue tracking systems that are managed by, or on behalf of, the
59
+ Licensor for the purpose of discussing and improving the Work, but
60
+ excluding communication that is conspicuously marked or otherwise
61
+ designated in writing by the copyright owner as "Not a Contribution."
62
+
63
+ "Contributor" shall mean Licensor and any individual or Legal Entity
64
+ on behalf of whom a Contribution has been received by Licensor and
65
+ subsequently incorporated within the Work.
66
+
67
+ 2. Grant of Copyright License. Subject to the terms and conditions of
68
+ this License, each Contributor hereby grants to You a perpetual,
69
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
70
+ copyright license to reproduce, prepare Derivative Works of,
71
+ publicly display, publicly perform, sublicense, and distribute the
72
+ Work and such Derivative Works in Source or Object form.
73
+
74
+ 3. Grant of Patent License. Subject to the terms and conditions of
75
+ this License, each Contributor hereby grants to You a perpetual,
76
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
77
+ (except as stated in this section) patent license to make, have made,
78
+ use, offer to sell, sell, import, and otherwise transfer the Work,
79
+ where such license applies only to those patent claims licensable
80
+ by such Contributor that are necessarily infringed by their
81
+ Contribution(s) alone or by combination of their Contribution(s)
82
+ with the Work to which such Contribution(s) was submitted. If You
83
+ institute patent litigation against any entity (including a
84
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
85
+ or a Contribution incorporated within the Work constitutes direct
86
+ or contributory patent infringement, then any patent licenses
87
+ granted to You under this License for that Work shall terminate
88
+ as of the date such litigation is filed.
89
+
90
+ 4. Redistribution. You may reproduce and distribute copies of the
91
+ Work or Derivative Works thereof in any medium, with or without
92
+ modifications, and in Source or Object form, provided that You
93
+ meet the following conditions:
94
+
95
+ (a) You must give any other recipients of the Work or
96
+ Derivative Works a copy of this License; and
97
+
98
+ (b) You must cause any modified files to carry prominent notices
99
+ stating that You changed the files; and
100
+
101
+ (c) You must retain, in the Source form of any Derivative Works
102
+ that You distribute, all copyright, patent, trademark, and
103
+ attribution notices from the Source form of the Work,
104
+ excluding those notices that do not pertain to any part of
105
+ the Derivative Works; and
106
+
107
+ (d) If the Work includes a "NOTICE" text file as part of its
108
+ distribution, then any Derivative Works that You distribute must
109
+ include a readable copy of the attribution notices contained
110
+ within such NOTICE file, excluding those notices that do not
111
+ pertain to any part of the Derivative Works, in at least one
112
+ of the following places: within a NOTICE text file distributed
113
+ as part of the Derivative Works; within the Source form or
114
+ documentation, if provided along with the Derivative Works; or,
115
+ within a display generated by the Derivative Works, if and
116
+ wherever such third-party notices normally appear. The contents
117
+ of the NOTICE file are for informational purposes only and
118
+ do not modify the License. You may add Your own attribution
119
+ notices within Derivative Works that You distribute, alongside
120
+ or as an addendum to the NOTICE text from the Work, provided
121
+ that such additional attribution notices cannot be construed
122
+ as modifying the License.
123
+
124
+ You may add Your own copyright statement to Your modifications and
125
+ may provide additional or different license terms and conditions
126
+ for use, reproduction, or distribution of Your modifications, or
127
+ for any such Derivative Works as a whole, provided Your use,
128
+ reproduction, and distribution of the Work otherwise complies with
129
+ the conditions stated in this License.
130
+
131
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
132
+ any Contribution intentionally submitted for inclusion in the Work
133
+ by You to the Licensor shall be under the terms and conditions of
134
+ this License, without any additional terms or conditions.
135
+ Notwithstanding the above, nothing herein shall supersede or modify
136
+ the terms of any separate license agreement you may have executed
137
+ with Licensor regarding such Contributions.
138
+
139
+ 6. Trademarks. This License does not grant permission to use the trade
140
+ names, trademarks, service marks, or product names of the Licensor,
141
+ except as required for reasonable and customary use in describing the
142
+ origin of the Work and reproducing the content of the NOTICE file.
143
+
144
+ 7. Disclaimer of Warranty. Unless required by applicable law or
145
+ agreed to in writing, Licensor provides the Work (and each
146
+ Contributor provides its Contributions) on an "AS IS" BASIS,
147
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
148
+ implied, including, without limitation, any warranties or conditions
149
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
150
+ PARTICULAR PURPOSE. You are solely responsible for determining the
151
+ appropriateness of using or redistributing the Work and assume any
152
+ risks associated with Your exercise of permissions under this License.
153
+
154
+ 8. Limitation of Liability. In no event and under no legal theory,
155
+ whether in tort (including negligence), contract, or otherwise,
156
+ unless required by applicable law (such as deliberate and grossly
157
+ negligent acts) or agreed to in writing, shall any Contributor be
158
+ liable to You for damages, including any direct, indirect, special,
159
+ incidental, or consequential damages of any character arising as a
160
+ result of this License or out of the use or inability to use the
161
+ Work (including but not limited to damages for loss of goodwill,
162
+ work stoppage, computer failure or malfunction, or any and all
163
+ other commercial damages or losses), even if such Contributor
164
+ has been advised of the possibility of such damages.
165
+
166
+ 9. Accepting Warranty or Additional Liability. While redistributing
167
+ the Work or Derivative Works thereof, You may choose to offer,
168
+ and charge a fee for, acceptance of support, warranty, indemnity,
169
+ or other liability obligations and/or rights consistent with this
170
+ License. However, in accepting such obligations, You may act only
171
+ on Your own behalf and on Your sole responsibility, not on behalf
172
+ of any other Contributor, and only if You agree to indemnify,
173
+ defend, and hold each Contributor harmless for any liability
174
+ incurred by, or claims asserted against, such Contributor by reason
175
+ of your accepting any such warranty or additional liability.
176
+
177
+ END OF TERMS AND CONDITIONS
178
+
179
+ APPENDIX: How to apply the Apache License to your work.
180
+
181
+ To apply the Apache License to your work, attach the following
182
+ boilerplate notice, with the fields enclosed by brackets "[]"
183
+ replaced with your own identifying information. (Don't include
184
+ the brackets!) The text should be enclosed in the appropriate
185
+ comment syntax for the file format. We also recommend that a
186
+ file or class name and description of purpose be included on the
187
+ same "printed page" as the copyright notice for easier
188
+ identification within third-party archives.
189
+
190
+ Copyright [yyyy] [name of copyright owner]
191
+
192
+ Licensed under the Apache License, Version 2.0 (the "License");
193
+ you may not use this file except in compliance with the License.
194
+ You may obtain a copy of the License at
195
+
196
+ http://www.apache.org/licenses/LICENSE-2.0
197
+
198
+ Unless required by applicable law or agreed to in writing, software
199
+ distributed under the License is distributed on an "AS IS" BASIS,
200
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
201
+ See the License for the specific language governing permissions and
202
+ limitations under the License.
NOTICE ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512
2
+
3
+ This product includes a Core ML conversion of GLiNER2.5-Decide
4
+ (https://huggingface.co/fastino/GLiNER2.5-Decide), created by Fastino and
5
+ licensed under the Apache License, Version 2.0.
6
+
7
+ It includes files derived from gliner2-5-decide-coreml
8
+ (https://huggingface.co/FluidInference/gliner2-5-decide-coreml) by Fluid
9
+ Inference, licensed under the Apache License, Version 2.0:
10
+ scripts/build/export_model.py, convert_names.py, preprocessing.py and runtime.py (unmodified)
11
+ scripts/build/convert_bucket.py (modified from convert-coreml.py)
12
+ src/gliner_decide_coreml/encoding.py and decoding.py (adapted from preprocessing.py and runtime.py)
13
+
14
+ The model packages in models/ were converted from the Fastino checkpoint at
15
+ revision 65624f1a0265b3f612bae66a2685a06b94a68a9d.
README.md ADDED
@@ -0,0 +1,176 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: coremltools
4
+ pipeline_tag: zero-shot-classification
5
+ base_model: fastino/GLiNER2.5-Decide
6
+ language:
7
+ - en
8
+ tags:
9
+ - coreml
10
+ - gliner
11
+ - gliner2
12
+ - deberta-v3
13
+ - apple-silicon
14
+ - fp16
15
+ - multifunction
16
+ - zero-shot-classification
17
+ ---
18
+
19
+ # GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512
20
+
21
+ [fastino/GLiNER2.5-Decide](https://huggingface.co/fastino/GLiNER2.5-Decide) (340M, DeBERTa-v3-large, zero-shot
22
+ classification) packaged for fast local inference on Apple Silicon:
23
+
24
+ - **Core ML, FP16**: the classification path (encoder + label classifier), converted with the tooling from
25
+ [FluidInference/gliner2-5-decide-coreml](https://huggingface.co/FluidInference/gliner2-5-decide-coreml).
26
+ - **MultiFn**: one multifunction package with four fixed-length graphs (`L64`, `L128`, `L256`, `L512`) that
27
+ share a single copy of the weights: **0.9 GB instead of 3.5 GB** for four separate packages.
28
+ - **Routing**: every request is tokenized and sent to the smallest bucket that fits it, so short requests
29
+ pay short-request latency. Text longer than 512 tokens is split into chunks automatically.
30
+
31
+ ## Results (Apple M1 Pro, macOS 27.2, `CPU_AND_GPU`)
32
+
33
+ | Bucket | Model only (Swift) | End to end (Python router) | `ALL` compute units |
34
+ |---:|---:|---:|---:|
35
+ | 64 | 19.3 ms | 20.0 ms | 70.7 ms |
36
+ | 128 | 31.6 ms | 32.2 ms | 197.6 ms |
37
+ | 256 | 63.3 ms | 60.9 ms | 79.6 ms |
38
+ | 512 | 133.9 ms | 135.3 ms | 208.9 ms |
39
+
40
+ p50 latency. "End to end" includes tokenization, routing, prediction and decoding; Python preprocessing costs
41
+ only 0.2–0.7 ms, so a native Swift runtime would not be meaningfully faster. **Use the GPU**: on this machine
42
+ `ALL` (GPU + Neural Engine) is slower at every size, and the Neural Engine alone is ~10x slower (see
43
+ [Research notes](#research-notes)).
44
+
45
+ **Fidelity.** Every bucket was verified against the native PyTorch model at conversion time (logit error
46
+ ≤ 4.8e-7 for the export wrapper; identical labels on every example that fits). The test suite checks the
47
+ router against native answers on six requests covering all four buckets (1–4 heads, multi-label,
48
+ label descriptions, prompts): identical labels, largest confidence difference **0.0007**.
49
+
50
+ ## Quick start
51
+
52
+ Requires Apple Silicon, **macOS 15+** (multifunction models), Python 3.10–3.12 and [uv](https://docs.astral.sh/uv/).
53
+
54
+ ```bash
55
+ hf download augustoFranke/GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512 --local-dir GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512
56
+ cd GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512
57
+ uv sync
58
+ ```
59
+
60
+ ```python
61
+ from gliner_decide_coreml import DecideRouter
62
+
63
+ router = DecideRouter() # loads all four buckets once; do this at startup
64
+
65
+ router.classify(
66
+ "My subscription renewed after the service was already down. Can I get that charge refunded?",
67
+ {
68
+ "intent": ["order_status", "refund_request", "cancel_subscription", "other"],
69
+ "urgency": ["low", "normal", "high"],
70
+ "topics": {"labels": ["billing", "outage", "account"], "multi_label": True, "cls_threshold": 0.4},
71
+ },
72
+ )
73
+ # {"intent": {"label": "refund_request", "confidence": 0.99}, "urgency": {...}, "topics": [{...}, ...]}
74
+
75
+ result, route = router.classify_with_route(text, tasks)
76
+ # Route(tokens=63, bucket=64, chunks=1, calls=1)
77
+ ```
78
+
79
+ The task format is the same as native `classify_text`: a list of labels, or a dict with `labels` (list, or
80
+ `{label: description}`), `multi_label`, `cls_threshold` and `prompt`.
81
+
82
+ The first load of each function compiles it for the GPU (can take a minute or more on a new machine);
83
+ macOS caches the result, and later loads take about 10 s for all four.
84
+
85
+ ## How it works
86
+
87
+ ```
88
+ text + tasks ──► GLiNER2 schema tokenizer ──► n tokens ──► smallest bucket ≥ n ──► Core ML function ──► decode
89
+ "( [P] intent ( [L] refund …" 64 / 128 / 256 / 512 (GPU, fp16) softmax or
90
+ [SEP_TEXT] my subscription …" └─ n > 512: word chunks, sigmoid +
91
+ merged per head threshold
92
+ ```
93
+
94
+ - **Input budget.** The schema (questions and labels) shares the sequence with the text. With 2 questions and
95
+ 8 labels (38 tokens) the buckets leave room for about 26 / 90 / 218 / 474 text tokens. Label descriptions
96
+ and prompts cost tokens too.
97
+ - **Heads.** Each call scores up to 4 heads with up to 32 labels each. More than 4 heads are split across calls.
98
+ Heads in one call see each other's labels, so results can shift slightly from a single native pass.
99
+ - **Chunking.** For text over 512 tokens: overlapping word chunks (32 words of overlap), each classified,
100
+ then merged: single-label heads take the most confident chunk, multi-label heads take the per-label
101
+ maximum. This is a heuristic; no chunk sees the whole text.
102
+
103
+ ## Layout
104
+
105
+ ```
106
+ models/
107
+ GLiNER2.5-Decide-MultiFn-fp16.mlpackage functions L64 / L128 / L256 / L512 (default L128)
108
+ tokenizer/ DeBERTa-v3 tokenizer + GLiNER2 special tokens
109
+ src/gliner_decide_coreml/
110
+ router.py DecideRouter: bucket selection, head grouping, chunking
111
+ encoding.py GLiNER2 schema preprocessing and bucket padding
112
+ decoding.py activations, thresholds, chunk merging
113
+ tests/ router vs native answers (golden.json), chunking, head splitting
114
+ scripts/
115
+ make_golden.py record native PyTorch answers for the tests
116
+ build/convert_bucket.py convert + verify one bucket (from the Fluid Inference tooling)
117
+ build/build_multifunction.py
118
+ bench/ Swift latency and compute-plan tools, Python router benchmark
119
+ research/ane-ceiling/ Neural Engine feasibility experiment (see below)
120
+ ```
121
+
122
+ ## Tests and benchmarks
123
+
124
+ ```bash
125
+ uv run pytest # router vs native answers, all buckets
126
+ uv run python bench/router_latency.py # end-to-end latency per bucket
127
+
128
+ cd bench
129
+ swiftc -O benchbuckets.swift -o benchbuckets # model-only latency of compiled packages
130
+ swiftc -O plan.swift -o plan # per-op device placement (MLComputePlan)
131
+ ```
132
+
133
+ `scripts/make_golden.py` regenerates `tests/golden.json` from the native model (downloads the pinned
134
+ checkpoint, ~1.7 GB).
135
+
136
+ ## Rebuilding the package
137
+
138
+ ```bash
139
+ cd scripts/build
140
+ for L in 64 128 256 512; do
141
+ uv run python convert_bucket.py --length $L --max-heads 4 --max-options 32 --output-dir build
142
+ done
143
+ uv run python build_multifunction.py --build-dir build
144
+ ```
145
+
146
+ Source checkpoint: `fastino/GLiNER2.5-Decide` at revision `65624f1a0265b3f612bae66a2685a06b94a68a9d`.
147
+ Pinned toolchain: coremltools 9.0, torch 2.7.0, transformers 4.57.6, gliner2 2.0.0.
148
+
149
+ ## Research notes
150
+
151
+ Measured on an M1 Pro while building this package.
152
+
153
+ - **Compute units.** Under `ALL`, Core ML splits every layer between the GPU and the Neural Engine (48
154
+ switches per request), which is slower than the GPU alone. Under `CPU_AND_NEURAL_ENGINE` the model takes
155
+ ~600 ms: DeBERTa's disentangled attention uses `gather_along_axis` (49 ops: the c2p and p2c
156
+ relative-position lookups in all 24 layers, plus the label-marker gather), which the Neural Engine cannot run.
157
+ - **Neural Engine ceiling** (`research/ane-ceiling`, random-weight encoders with this model's shape). A plain
158
+ encoder runs in **32–36 ms** on the Neural Engine, and adding DeBERTa's extra position matmuls costs only
159
+ +6 ms. Replacing the gathers exactly with gather-free "relative shift" code keeps every op on the Neural
160
+ Engine but is slow: reshape-based shift **807 ms**, barrel shifter (static slices + blends) **187 ms**.
161
+ An exact ANE rewrite would only beat the GPU if the shift cost less than ~1 ms per layer.
162
+ - **GPU headroom.** The model needs roughly 175 GFLOP per 256-token request; at the M1 Pro GPU's peak that is a
163
+ ~35 ms floor, versus 61–63 ms achieved. Bigger wins come from shorter inputs (this package) or a smaller,
164
+ distilled model, not from kernel tuning.
165
+
166
+ ## Limitations
167
+
168
+ - Classification only: GLiNER2's entity, relation and structure extraction heads are not included.
169
+ - Preprocessing uses the GLiNER2 Python package, which imports PyTorch (no model weights are loaded).
170
+ - Latency figures are from one machine (M1 Pro); newer chips are faster and may favour different compute units.
171
+ - Multifunction packages need macOS 15 / iOS 18 or later.
172
+
173
+ ## License and attribution
174
+
175
+ Apache-2.0 (see `LICENSE` and `NOTICE`). Model by [Fastino](https://huggingface.co/fastino/GLiNER2.5-Decide);
176
+ original Core ML conversion tooling by [Fluid Inference](https://huggingface.co/FluidInference/gliner2-5-decide-coreml).
bench/benchbuckets.swift ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // Latency of each GLiNER2.5-Decide bucket package with valid, full-length inputs.
2
+ // Usage: benchbuckets <units,...> <iterations> <pkg.mlmodelc>...
3
+ import CoreML
4
+ import Foundation
5
+
6
+ let args = CommandLine.arguments
7
+ let unitsByName: [String: MLComputeUnits] = [
8
+ "all": .all, "cpuAndGPU": .cpuAndGPU, "cpuOnly": .cpuOnly, "cpuAndNeuralEngine": .cpuAndNeuralEngine,
9
+ ]
10
+ let unitNames = args[1].split(separator: ",").map(String.init)
11
+ let iters = Int(args[2])!
12
+ let packages = args.dropFirst(3).map { URL(fileURLWithPath: $0) }
13
+
14
+ func filled(_ shape: [NSNumber], _ type: MLMultiArrayDataType, _ value: (Int) -> Double) throws -> MLMultiArray {
15
+ let a = try MLMultiArray(shape: shape, dataType: type)
16
+ for i in 0..<a.count { a[i] = NSNumber(value: value(i)) }
17
+ return a
18
+ }
19
+
20
+ // Every position is a real (unmasked) token: the worst case for a bucket. Padding does not
21
+ // change the compute, so this is also the cost of any request routed to this bucket.
22
+ func inputs(for desc: MLModelDescription) throws -> MLFeatureProvider {
23
+ let ids = desc.inputDescriptionsByName["input_ids"]!.multiArrayConstraint!.shape
24
+ let grid = desc.inputDescriptionsByName["marker_indices"]!.multiArrayConstraint!.shape
25
+ let L = ids[1].intValue, K = grid[2].intValue
26
+ let used = [5, 3, 0, 0] // two questions with 5 and 3 labels, like the example request
27
+ return try MLDictionaryFeatureProvider(dictionary: [
28
+ "input_ids": try filled(ids, .int32) { _ in Double(Int.random(in: 1000..<100_000)) },
29
+ "attention_mask": try filled(ids, .int32) { _ in 1 },
30
+ "marker_indices": try filled(grid, .int32) { i in
31
+ let h = i / K, k = i % K
32
+ return k < used[h] ? Double(min(L - 1, 2 + h * 12 + k * 2)) : 0
33
+ },
34
+ "marker_mask": try filled(grid, .float32) { i in i % K < used[i / K] ? 1 : 0 },
35
+ ])
36
+ }
37
+
38
+ print("bucket units load p50 p90")
39
+ for url in packages {
40
+ for label in unitNames {
41
+ let config = MLModelConfiguration()
42
+ config.computeUnits = unitsByName[label]!
43
+ let t0 = Date()
44
+ let model = try MLModel(contentsOf: url, configuration: config)
45
+ let load = Date().timeIntervalSince(t0)
46
+ let input = try inputs(for: model.modelDescription)
47
+ let L = model.modelDescription.inputDescriptionsByName["input_ids"]!.multiArrayConstraint!.shape[1]
48
+ for _ in 0..<5 { _ = try model.prediction(from: input) }
49
+ var t: [Double] = []
50
+ for _ in 0..<iters {
51
+ let s = Date()
52
+ _ = try model.prediction(from: input)
53
+ t.append(Date().timeIntervalSince(s) * 1000)
54
+ }
55
+ t.sort()
56
+ print(String(format: "L%-6@ %-10@ %5.1fs %6.1f ms %6.1f ms",
57
+ L.stringValue as NSString, label as NSString, load, t[t.count / 2], t[Int(Double(t.count) * 0.9)]))
58
+ }
59
+ }
bench/plan.swift ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // Step 0 diagnostic: where does CoreML place each op of the GLiNER2.5-Decide package?
2
+ // Usage: plan <path.mlpackage|.mlmodelc> <out.json>
3
+ import CoreML
4
+ import Foundation
5
+
6
+ func deviceName(_ d: MLComputeDevice) -> String {
7
+ switch d {
8
+ case .cpu: return "CPU"
9
+ case .gpu: return "GPU"
10
+ case .neuralEngine: return "ANE"
11
+ @unknown default: return "?"
12
+ }
13
+ }
14
+
15
+ struct OpRecord: Codable {
16
+ let index: Int
17
+ let op: String
18
+ let outputs: [String]
19
+ let preferred: String
20
+ let supported: [String]
21
+ let cost: Double
22
+ }
23
+
24
+ func walk(_ block: MLModelStructure.Program.Block, plan: MLComputePlan, into records: inout [OpRecord]) {
25
+ for op in block.operations {
26
+ let usage = plan.deviceUsage(for: op)
27
+ let cost = plan.estimatedCost(of: op)?.weight ?? 0
28
+ records.append(OpRecord(
29
+ index: records.count,
30
+ op: op.operatorName,
31
+ outputs: op.outputs.map { $0.name },
32
+ preferred: usage.map { deviceName($0.preferred) } ?? "none",
33
+ supported: usage.map { $0.supported.map(deviceName) } ?? [],
34
+ cost: cost))
35
+ for inner in op.blocks { walk(inner, plan: plan, into: &records) }
36
+ }
37
+ }
38
+
39
+ let args = CommandLine.arguments
40
+ let modelURL = URL(fileURLWithPath: args[1])
41
+ let outURL = URL(fileURLWithPath: args[2])
42
+
43
+ let compiledURL: URL
44
+ if modelURL.pathExtension == "mlmodelc" {
45
+ compiledURL = modelURL
46
+ } else {
47
+ let t0 = Date()
48
+ let tmp = try await MLModel.compileModel(at: modelURL)
49
+ compiledURL = modelURL.deletingPathExtension().appendingPathExtension("mlmodelc")
50
+ try? FileManager.default.removeItem(at: compiledURL)
51
+ try FileManager.default.moveItem(at: tmp, to: compiledURL)
52
+ print("compiled in \(String(format: "%.1f", Date().timeIntervalSince(t0)))s -> \(compiledURL.lastPathComponent)")
53
+ }
54
+
55
+ var result: [String: [OpRecord]] = [:]
56
+ for (label, units) in [("cpuAndNeuralEngine", MLComputeUnits.cpuAndNeuralEngine), ("all", MLComputeUnits.all)] {
57
+ let config = MLModelConfiguration()
58
+ config.computeUnits = units
59
+ let plan = try await MLComputePlan.load(contentsOf: compiledURL, configuration: config)
60
+ guard case let .program(program) = plan.modelStructure, let main = program.functions["main"] else {
61
+ fatalError("not an ML program")
62
+ }
63
+ var records: [OpRecord] = []
64
+ walk(main.block, plan: plan, into: &records)
65
+ result[label] = records
66
+ print("\(label): \(records.count) ops")
67
+ }
68
+ let enc = JSONEncoder()
69
+ enc.outputFormatting = [.prettyPrinted, .sortedKeys]
70
+ try enc.encode(result).write(to: outURL)
71
+ print("wrote \(outURL.path)")
bench/router_latency.py ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """End-to-end router latency (tokenize + route + Core ML predict + decode) per bucket, from Python.
2
+
3
+ uv run python bench/router_latency.py
4
+ """
5
+ import statistics
6
+ import sys
7
+ import time
8
+ import warnings
9
+ from pathlib import Path
10
+
11
+ warnings.filterwarnings("ignore")
12
+ sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "tests"))
13
+
14
+ from cases import CASES # noqa: E402
15
+ from gliner_decide_coreml import DecideRouter # noqa: E402
16
+
17
+ router = DecideRouter()
18
+ seen = set()
19
+ print(f"{'case':<28}{'tokens':>7}{'bucket':>8}{'p50':>10}{'p90':>10}")
20
+ for case in CASES:
21
+ if case["bucket"] in seen:
22
+ continue
23
+ seen.add(case["bucket"])
24
+ for _ in range(5):
25
+ router.classify(case["text"], case["tasks"])
26
+ times = []
27
+ for _ in range(30):
28
+ start = time.perf_counter()
29
+ _, route = router.classify_with_route(case["text"], case["tasks"])
30
+ times.append((time.perf_counter() - start) * 1000)
31
+ times.sort()
32
+ print(f"{case['id']:<28}{route.tokens:>7}{route.bucket:>8}{statistics.median(times):>8.1f}ms"
33
+ f"{times[int(len(times) * 0.9)]:>8.1f}ms")
bench/time_prep.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Time the Python preprocessing (tokenize + build schema + pad) for requests of different sizes."""
2
+ import statistics
3
+ import time
4
+ import warnings
5
+
6
+ warnings.filterwarnings("ignore")
7
+ from preprocessing import load_processor, prepare_decision # noqa: E402
8
+
9
+ proc = load_processor(".")
10
+ tasks = {"intent": ["order_status", "refund_request", "cancel_subscription", "update_payment", "other"],
11
+ "urgency": ["low", "normal", "high"]}
12
+ sentence = ("My subscription renewed on April 15 for 5,400 yen after the service was already down. "
13
+ "Can I get that charge refunded? ")
14
+
15
+ for bucket, repeats in [(64, 1), (128, 3), (256, 8), (512, 18)]:
16
+ text = (sentence * repeats).strip()
17
+ for _ in range(5):
18
+ arrays = prepare_decision(proc, text, tasks, bucket, 4, 32)
19
+ times = []
20
+ for _ in range(50):
21
+ t = time.perf_counter()
22
+ arrays = prepare_decision(proc, text, tasks, bucket, 4, 32)
23
+ times.append((time.perf_counter() - t) * 1000)
24
+ n = int(arrays["attention_mask"].sum())
25
+ print(f"L{bucket:<4} {n:4d} real tokens prep p50 {statistics.median(times):5.2f} ms")
models/GLiNER2.5-Decide-MultiFn-fp16.mlpackage/Data/com.apple.CoreML/model.mlmodel ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2af62288520359e51612dcde770266a181de4922dd637b9e33bafebbf2ed734e
3
+ size 1585697
models/GLiNER2.5-Decide-MultiFn-fp16.mlpackage/Data/com.apple.CoreML/weights/weight.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1d9ea7dadc22a0e666a5b27c89b0d087a89aeae778dbb0a45e1abe32acea0d88
3
+ size 943634368
models/GLiNER2.5-Decide-MultiFn-fp16.mlpackage/Manifest.json ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "fileFormatVersion": "1.0.0",
3
+ "itemInfoEntries": {
4
+ "1F0D1536-6E76-4AF5-A08C-54426203B6AD": {
5
+ "author": "com.apple.CoreML",
6
+ "description": "CoreML Model Weights",
7
+ "name": "weights",
8
+ "path": "com.apple.CoreML/weights"
9
+ },
10
+ "D12A5651-2419-4D92-A578-C314C16AE33B": {
11
+ "author": "com.apple.CoreML",
12
+ "description": "CoreML Model Specification",
13
+ "name": "model.mlmodel",
14
+ "path": "com.apple.CoreML/model.mlmodel"
15
+ }
16
+ },
17
+ "rootModelIdentifier": "D12A5651-2419-4D92-A578-C314C16AE33B"
18
+ }
models/tokenizer/config.json ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_attn_implementation_autoset": true,
3
+ "_name_or_path": "/home/urchadezaratiana/checkpoints/checkpoint-45500",
4
+ "architecture": "span",
5
+ "architecture_version": 1,
6
+ "architectures": [
7
+ "SpanExtractor"
8
+ ],
9
+ "attn_implementation": "sdpa",
10
+ "config_version": 3,
11
+ "counting_layer": "count_lstm",
12
+ "max_len": null,
13
+ "max_width": 8,
14
+ "model_name": "microsoft/deberta-v3-large",
15
+ "model_type": "extractor",
16
+ "span_head": {
17
+ "dropout": 0.1,
18
+ "max_width": 8,
19
+ "span_mode": "markerV0"
20
+ },
21
+ "token_pooling": "first",
22
+ "transformers_version": "4.48.1",
23
+ "use_moe": false
24
+ }
models/tokenizer/special_tokens_map.json ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "additional_special_tokens": [
3
+ {
4
+ "content": "[SEP_STRUCT]",
5
+ "lstrip": false,
6
+ "normalized": false,
7
+ "rstrip": false,
8
+ "single_word": false
9
+ },
10
+ {
11
+ "content": "[SEP_TEXT]",
12
+ "lstrip": false,
13
+ "normalized": false,
14
+ "rstrip": false,
15
+ "single_word": false
16
+ },
17
+ {
18
+ "content": "[P]",
19
+ "lstrip": false,
20
+ "normalized": false,
21
+ "rstrip": false,
22
+ "single_word": false
23
+ },
24
+ {
25
+ "content": "[C]",
26
+ "lstrip": false,
27
+ "normalized": false,
28
+ "rstrip": false,
29
+ "single_word": false
30
+ },
31
+ {
32
+ "content": "[E]",
33
+ "lstrip": false,
34
+ "normalized": false,
35
+ "rstrip": false,
36
+ "single_word": false
37
+ },
38
+ {
39
+ "content": "[R]",
40
+ "lstrip": false,
41
+ "normalized": false,
42
+ "rstrip": false,
43
+ "single_word": false
44
+ },
45
+ {
46
+ "content": "[L]",
47
+ "lstrip": false,
48
+ "normalized": false,
49
+ "rstrip": false,
50
+ "single_word": false
51
+ },
52
+ {
53
+ "content": "[EXAMPLE]",
54
+ "lstrip": false,
55
+ "normalized": false,
56
+ "rstrip": false,
57
+ "single_word": false
58
+ },
59
+ {
60
+ "content": "[OUTPUT]",
61
+ "lstrip": false,
62
+ "normalized": false,
63
+ "rstrip": false,
64
+ "single_word": false
65
+ },
66
+ {
67
+ "content": "[DESCRIPTION]",
68
+ "lstrip": false,
69
+ "normalized": false,
70
+ "rstrip": false,
71
+ "single_word": false
72
+ }
73
+ ],
74
+ "bos_token": {
75
+ "content": "[CLS]",
76
+ "lstrip": false,
77
+ "normalized": false,
78
+ "rstrip": false,
79
+ "single_word": false
80
+ },
81
+ "cls_token": {
82
+ "content": "[CLS]",
83
+ "lstrip": false,
84
+ "normalized": false,
85
+ "rstrip": false,
86
+ "single_word": false
87
+ },
88
+ "eos_token": {
89
+ "content": "[SEP]",
90
+ "lstrip": false,
91
+ "normalized": false,
92
+ "rstrip": false,
93
+ "single_word": false
94
+ },
95
+ "mask_token": {
96
+ "content": "[MASK]",
97
+ "lstrip": false,
98
+ "normalized": false,
99
+ "rstrip": false,
100
+ "single_word": false
101
+ },
102
+ "pad_token": {
103
+ "content": "[PAD]",
104
+ "lstrip": false,
105
+ "normalized": false,
106
+ "rstrip": false,
107
+ "single_word": false
108
+ },
109
+ "sep_token": {
110
+ "content": "[SEP]",
111
+ "lstrip": false,
112
+ "normalized": false,
113
+ "rstrip": false,
114
+ "single_word": false
115
+ },
116
+ "unk_token": {
117
+ "content": "[UNK]",
118
+ "lstrip": false,
119
+ "normalized": true,
120
+ "rstrip": false,
121
+ "single_word": false
122
+ }
123
+ }
models/tokenizer/tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
models/tokenizer/tokenizer_config.json ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": true,
3
+ "added_tokens_decoder": {
4
+ "0": {
5
+ "content": "[PAD]",
6
+ "lstrip": false,
7
+ "normalized": false,
8
+ "rstrip": false,
9
+ "single_word": false,
10
+ "special": true
11
+ },
12
+ "1": {
13
+ "content": "[CLS]",
14
+ "lstrip": false,
15
+ "normalized": false,
16
+ "rstrip": false,
17
+ "single_word": false,
18
+ "special": true
19
+ },
20
+ "2": {
21
+ "content": "[SEP]",
22
+ "lstrip": false,
23
+ "normalized": false,
24
+ "rstrip": false,
25
+ "single_word": false,
26
+ "special": true
27
+ },
28
+ "3": {
29
+ "content": "[UNK]",
30
+ "lstrip": false,
31
+ "normalized": true,
32
+ "rstrip": false,
33
+ "single_word": false,
34
+ "special": true
35
+ },
36
+ "128000": {
37
+ "content": "[MASK]",
38
+ "lstrip": false,
39
+ "normalized": false,
40
+ "rstrip": false,
41
+ "single_word": false,
42
+ "special": true
43
+ },
44
+ "128001": {
45
+ "content": "[SEP_STRUCT]",
46
+ "lstrip": false,
47
+ "normalized": false,
48
+ "rstrip": false,
49
+ "single_word": false,
50
+ "special": true
51
+ },
52
+ "128002": {
53
+ "content": "[SEP_TEXT]",
54
+ "lstrip": false,
55
+ "normalized": false,
56
+ "rstrip": false,
57
+ "single_word": false,
58
+ "special": true
59
+ },
60
+ "128003": {
61
+ "content": "[P]",
62
+ "lstrip": false,
63
+ "normalized": false,
64
+ "rstrip": false,
65
+ "single_word": false,
66
+ "special": true
67
+ },
68
+ "128004": {
69
+ "content": "[C]",
70
+ "lstrip": false,
71
+ "normalized": false,
72
+ "rstrip": false,
73
+ "single_word": false,
74
+ "special": true
75
+ },
76
+ "128005": {
77
+ "content": "[E]",
78
+ "lstrip": false,
79
+ "normalized": false,
80
+ "rstrip": false,
81
+ "single_word": false,
82
+ "special": true
83
+ },
84
+ "128006": {
85
+ "content": "[R]",
86
+ "lstrip": false,
87
+ "normalized": false,
88
+ "rstrip": false,
89
+ "single_word": false,
90
+ "special": true
91
+ },
92
+ "128007": {
93
+ "content": "[L]",
94
+ "lstrip": false,
95
+ "normalized": false,
96
+ "rstrip": false,
97
+ "single_word": false,
98
+ "special": true
99
+ },
100
+ "128008": {
101
+ "content": "[EXAMPLE]",
102
+ "lstrip": false,
103
+ "normalized": false,
104
+ "rstrip": false,
105
+ "single_word": false,
106
+ "special": true
107
+ },
108
+ "128009": {
109
+ "content": "[OUTPUT]",
110
+ "lstrip": false,
111
+ "normalized": false,
112
+ "rstrip": false,
113
+ "single_word": false,
114
+ "special": true
115
+ },
116
+ "128010": {
117
+ "content": "[DESCRIPTION]",
118
+ "lstrip": false,
119
+ "normalized": false,
120
+ "rstrip": false,
121
+ "single_word": false,
122
+ "special": true
123
+ }
124
+ },
125
+ "additional_special_tokens": [
126
+ "[SEP_STRUCT]",
127
+ "[SEP_TEXT]",
128
+ "[P]",
129
+ "[C]",
130
+ "[E]",
131
+ "[R]",
132
+ "[L]",
133
+ "[EXAMPLE]",
134
+ "[OUTPUT]",
135
+ "[DESCRIPTION]"
136
+ ],
137
+ "backend": "tokenizers",
138
+ "bos_token": "[CLS]",
139
+ "clean_up_tokenization_spaces": false,
140
+ "cls_token": "[CLS]",
141
+ "do_lower_case": false,
142
+ "eos_token": "[SEP]",
143
+ "extra_special_tokens": {},
144
+ "is_local": false,
145
+ "local_files_only": false,
146
+ "mask_token": "[MASK]",
147
+ "model_max_length": 1000000000000000019884624838656,
148
+ "pad_token": "[PAD]",
149
+ "sep_token": "[SEP]",
150
+ "sp_model_kwargs": {},
151
+ "split_by_punct": false,
152
+ "tokenizer_class": "DebertaV2Tokenizer",
153
+ "unk_id": 3,
154
+ "unk_token": "[UNK]",
155
+ "vocab_type": "spm"
156
+ }
pyproject.toml ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [project]
2
+ name = "gliner-decide-coreml"
3
+ version = "0.1.0"
4
+ description = "GLiNER2.5-Decide as one Core ML multifunction package (64/128/256/512-token buckets) with automatic bucket routing, tuned for Apple Silicon GPUs."
5
+ readme = "README.md"
6
+ license = "Apache-2.0"
7
+ requires-python = ">=3.10,<3.13"
8
+ # Pinned to the versions the Core ML buckets were converted and verified with.
9
+ # torch/transformers are needed only because GLiNER2's schema preprocessing imports them.
10
+ dependencies = [
11
+ "coremltools==9.0",
12
+ "gliner2[local]==2.0.0",
13
+ "huggingface-hub>=0.34,<1",
14
+ "numpy<2.3",
15
+ "sentencepiece>=0.2,<0.3",
16
+ "torch==2.7.0",
17
+ "transformers==4.57.6",
18
+ ]
19
+
20
+ [dependency-groups]
21
+ dev = [
22
+ "pytest>=8",
23
+ ]
24
+
25
+ [build-system]
26
+ requires = ["hatchling"]
27
+ build-backend = "hatchling.build"
28
+
29
+ [tool.hatch.build.targets.wheel]
30
+ packages = ["src/gliner_decide_coreml"]
31
+
32
+ [tool.pytest.ini_options]
33
+ testpaths = ["tests"]
34
+ filterwarnings = ["ignore::DeprecationWarning", "ignore::UserWarning", "ignore::FutureWarning"]
research/ane-ceiling/models.py ADDED
@@ -0,0 +1,250 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """ANE ceiling test: DeBERTa-v3-large-shaped encoders with random weights.
2
+
3
+ Variants
4
+ a-std : plain post-LN encoder, standard (B, S, C) layout with nn.Linear
5
+ a-ane : same math, Apple's ANE layout (B, C, 1, S), 1x1 Conv2d, per-head attention
6
+ b-ane : a-ane + DeBERTa's c2p / p2c terms, computed gather-free:
7
+ constant per-layer position tables (one row per possible distance)
8
+ + the "relative shift" (reshape / slice) instead of gather_along_axis
9
+ Speed does not depend on weight values, so random weights give real timings.
10
+ """
11
+ import argparse
12
+ import math
13
+
14
+ import numpy as np
15
+ import torch
16
+ import torch.nn as nn
17
+ import torch.nn.functional as F
18
+
19
+ C, H, D, FF = 1024, 16, 64, 4096 # DeBERTa-v3-large dims (from the port's encoder_config)
20
+
21
+
22
+ def skew(t: torch.Tensor, S: int) -> torch.Tensor:
23
+ """(B, S, 1, 2S-1) scores against every distance -> (B, S, 1, S).
24
+
25
+ out[b, i, 0, j] = t[b, i, 0, j - i + S - 1]; done with reshape/slice only.
26
+ Row i of the flattened tensor starts at i*(2S-1); the wanted entry is at
27
+ i*(2S-2) + (S-1) + j, i.e. rows of stride 2S-2 once we drop the first S-1.
28
+ """
29
+ B = t.shape[0]
30
+ f = t.reshape(B, S * (2 * S - 1))
31
+ f = f[:, S - 1 : S - 1 + S * (2 * S - 2)]
32
+ f = f.reshape(B, S, 2 * S - 2)[:, :, :S]
33
+ return f.reshape(B, S, 1, S)
34
+
35
+
36
+ def shift_last(x: torch.Tensor, S: int, N: int) -> torch.Tensor:
37
+ """Barrel shifter along the last axis, no reshapes.
38
+
39
+ x: (B, N*S, 1, 2S-1), rows grouped per head. out[.., n*S+i, 0, j] = x[.., n*S+i, 0, j + S-1-i].
40
+ The per-row shift (0..S-1) is applied in log2(S) steps of 1, 2, 4, ...; each step is two
41
+ static slices and a select with a constant row mask.
42
+ """
43
+ w = 2 * S - 1 # static ints only: traced shapes would become dynamic ops
44
+ for b, (m, inv) in enumerate(barrel_masks(S, N, "rows")):
45
+ step = 1 << b
46
+ w -= step
47
+ x = x[..., step:step + w] * m + x[..., :w] * inv
48
+ return x
49
+
50
+
51
+ def shift_chan(x: torch.Tensor, S: int) -> torch.Tensor:
52
+ """Barrel shifter along the channel axis (lands p2c in query-major order, no transpose).
53
+
54
+ x: (N, 2S-1, 1, S). out[n, i, 0, j] = x[n, i + S-1-j, 0, j].
55
+ """
56
+ c = 2 * S - 1
57
+ for b, (m, inv) in enumerate(barrel_masks(S, 1, "cols")):
58
+ step = 1 << b
59
+ c -= step
60
+ x = x[:, step:step + c] * m + x[:, :c] * inv
61
+ return x
62
+
63
+
64
+ _MASKS = {}
65
+
66
+
67
+ def barrel_masks(S: int, N: int, kind: str):
68
+ """Constant 0/1 masks (and their complements) for each barrel step.
69
+
70
+ Built once outside tracing so they enter the graph as plain constants; blending with
71
+ m * hi + (1 - m) * lo is exact because one of the two products is always zero.
72
+ """
73
+ key = (S, N, kind)
74
+ if key not in _MASKS:
75
+ shift = S - 1 - torch.arange(S)
76
+ shape = (1, N * S, 1, 1) if kind == "rows" else (1, 1, 1, S)
77
+ steps = []
78
+ b = 0
79
+ while (1 << b) < S:
80
+ m = ((shift >> b) & 1).repeat(N).float().view(shape)
81
+ steps.append((m, 1.0 - m))
82
+ b += 1
83
+ _MASKS[key] = steps
84
+ return _MASKS[key]
85
+
86
+
87
+ class LayerNormANE(nn.Module):
88
+ """LayerNorm over the channel axis of a (B, C, 1, S) tensor."""
89
+
90
+ def __init__(self, c, eps=1e-5):
91
+ super().__init__()
92
+ self.eps = eps
93
+ self.weight = nn.Parameter(torch.ones(c))
94
+ self.bias = nn.Parameter(torch.zeros(c))
95
+
96
+ def forward(self, x):
97
+ zc = x - x.mean(dim=1, keepdim=True)
98
+ var = (zc * zc).mean(dim=1, keepdim=True)
99
+ out = zc * torch.rsqrt(var + self.eps)
100
+ return out * self.weight.view(1, -1, 1, 1) + self.bias.view(1, -1, 1, 1)
101
+
102
+
103
+ class LayerANE(nn.Module):
104
+ """rel: None (plain), "skew" (reshape shift), "barrel" (slice/select shift),
105
+ "noshift" (extra matmuls only, shift skipped: wrong math, timing control)."""
106
+
107
+ def __init__(self, S: int, rel=None):
108
+ super().__init__()
109
+ self.S, self.rel = S, rel
110
+ self.q, self.k, self.v, self.o = (nn.Conv2d(C, C, 1) for _ in range(4))
111
+ self.ln1, self.ln2 = LayerNormANE(C), LayerNormANE(C)
112
+ self.ff1, self.ff2 = nn.Conv2d(C, FF, 1), nn.Conv2d(FF, C, 1)
113
+ # DeBERTa scales by sqrt(d * (1 + number of position terms)).
114
+ self.scale = 1.0 / math.sqrt(D * (3 if rel else 1))
115
+ if rel:
116
+ # Precomputed LN(rel_embeddings) @ W_k / W_q, expanded to every distance.
117
+ self.register_buffer("pk", torch.randn(1, C, 1, 2 * S - 1) * 0.02)
118
+ self.register_buffer("pq", torch.randn(1, C, 1, 2 * S - 1) * 0.02)
119
+ self.register_buffer("pqT", torch.randn(1, 2 * S - 1, 1, C) * 0.02)
120
+ if rel == "barrel":
121
+ # Build the masks now, so tracing captures them as constants rather than ops.
122
+ barrel_masks(S, H, "rows")
123
+ barrel_masks(S, 1, "cols")
124
+
125
+ def rel_terms(self, qT, ks, kT):
126
+ """Per-head (c2p, p2c) score tables, each (B, Sq, 1, Sk)."""
127
+ S = self.S
128
+ pks = self.pk.split(D, dim=1)
129
+ if self.rel == "skew":
130
+ pqs = self.pq.split(D, dim=1)
131
+ c2p = [skew(torch.einsum("bchr,bqhc->bqhr", pks[h], qT[h]), S) for h in range(H)]
132
+ p2c = [skew(torch.einsum("bchr,bkhc->bkhr", pqs[h], kT[h]), S).transpose(1, 3)
133
+ for h in range(H)]
134
+ return c2p, p2c
135
+ # Heads batched so each shift runs once per layer, not once per head.
136
+ a = torch.cat([torch.einsum("bchr,bqhc->bqhr", pks[h], qT[h]) for h in range(H)], dim=1)
137
+ pqT = self.pqT.split(D, dim=3)
138
+ m = torch.cat([torch.einsum("bchk,brhc->brhk", ks[h], pqT[h]) for h in range(H)], dim=0)
139
+ if self.rel == "barrel":
140
+ a, m = shift_last(a, S, H), shift_chan(m, S)
141
+ else: # noshift
142
+ a, m = a[..., :S], m[:, :S]
143
+ return a.split(S, dim=1), m.split(1, dim=0)
144
+
145
+ def forward(self, x, mask): # x: (B, C, 1, S); mask: (B, 1, 1, S) additive
146
+ q, k, v = self.q(x), self.k(x), self.v(x)
147
+ qT = q.transpose(1, 3).split(D, dim=3) # H x (B, S, 1, D)
148
+ ks, vs = k.split(D, dim=1), v.split(D, dim=1) # H x (B, D, 1, S)
149
+ if self.rel:
150
+ kT = k.transpose(1, 3).split(D, dim=3) if self.rel == "skew" else None
151
+ c2p, p2c = self.rel_terms(qT, ks, kT)
152
+ heads = []
153
+ for h in range(H):
154
+ s = torch.einsum("bchk,bqhc->bqhk", ks[h], qT[h]) # (B, Sq, 1, Sk)
155
+ if self.rel:
156
+ s = s + c2p[h] + p2c[h]
157
+ w = (s * self.scale + mask).softmax(dim=3)
158
+ heads.append(torch.einsum("bqhk,bchk->bchq", w, vs[h])) # (B, D, 1, Sq)
159
+ x = self.ln1(x + self.o(torch.cat(heads, dim=1)))
160
+ return self.ln2(x + self.ff2(F.gelu(self.ff1(x))))
161
+
162
+
163
+ class LayerStd(nn.Module):
164
+ def __init__(self, S: int):
165
+ super().__init__()
166
+ self.q, self.k, self.v, self.o = (nn.Linear(C, C) for _ in range(4))
167
+ self.ln1, self.ln2 = nn.LayerNorm(C), nn.LayerNorm(C)
168
+ self.ff1, self.ff2 = nn.Linear(C, FF), nn.Linear(FF, C)
169
+ self.scale = 1.0 / math.sqrt(D)
170
+
171
+ def forward(self, x, mask): # x: (B, S, C); mask: (B, 1, 1, S)
172
+ B, S, _ = x.shape
173
+ q = self.q(x).view(B, S, H, D).transpose(1, 2)
174
+ k = self.k(x).view(B, S, H, D).transpose(1, 2)
175
+ v = self.v(x).view(B, S, H, D).transpose(1, 2)
176
+ w = (q @ k.transpose(-1, -2) * self.scale + mask).softmax(-1)
177
+ a = (w @ v).transpose(1, 2).reshape(B, S, C)
178
+ x = self.ln1(x + self.o(a))
179
+ return self.ln2(x + self.ff2(F.gelu(self.ff1(x))))
180
+
181
+
182
+ class Encoder(nn.Module):
183
+ def __init__(self, variant: str, S: int, layers: int):
184
+ super().__init__()
185
+ make = {"a-std": lambda: LayerStd(S),
186
+ "a-ane": lambda: LayerANE(S),
187
+ "b-ane": lambda: LayerANE(S, "skew"),
188
+ "c-ane": lambda: LayerANE(S, "barrel"),
189
+ "n-ane": lambda: LayerANE(S, "noshift")}[variant]
190
+ self.layers = nn.ModuleList(make() for _ in range(layers))
191
+
192
+ def forward(self, x, mask):
193
+ for layer in self.layers:
194
+ x = layer(x, mask)
195
+ return x
196
+
197
+
198
+ def check_skew():
199
+ """The reshape/slice shift must equal the gather it replaces."""
200
+ S = 7
201
+ t = torch.randn(2, S, 1, 2 * S - 1)
202
+ i = torch.arange(S).view(S, 1)
203
+ j = torch.arange(S).view(1, S)
204
+ idx = (j - i + S - 1).expand(2, S, S).unsqueeze(2) # (B, S, 1, S)
205
+ ref = torch.gather(t, 3, idx)
206
+ assert torch.equal(skew(t, S), ref), "skew != gather"
207
+
208
+ S, N = 8, 3 # barrel shifters need S to be a power of two
209
+ t = torch.randn(1, N * S, 1, 2 * S - 1)
210
+ i = torch.arange(S).view(S, 1)
211
+ j = torch.arange(S).view(1, S)
212
+ idx = (j - i + S - 1).repeat(N, 1).view(1, N * S, 1, S)
213
+ assert torch.equal(shift_last(t, S, N), torch.gather(t, 3, idx)), "shift_last != gather"
214
+ t = torch.randn(N, 2 * S - 1, 1, S)
215
+ idx = (i + S - 1 - j).view(1, S, 1, S).expand(N, S, 1, S)
216
+ assert torch.equal(shift_chan(t, S), torch.gather(t, 1, idx)), "shift_chan != gather"
217
+ print("skew, shift_last, shift_chan == gather: OK")
218
+
219
+
220
+ if __name__ == "__main__":
221
+ import coremltools as ct
222
+
223
+ p = argparse.ArgumentParser()
224
+ p.add_argument("variant", choices=["a-std", "a-ane", "b-ane", "c-ane", "n-ane"])
225
+ p.add_argument("--seq", type=int, default=256)
226
+ p.add_argument("--layers", type=int, default=24)
227
+ p.add_argument("--out", required=True)
228
+ a = p.parse_args()
229
+
230
+ check_skew()
231
+ torch.manual_seed(0)
232
+ model = Encoder(a.variant, a.seq, a.layers).eval()
233
+ n = sum(p.numel() for p in model.parameters()) + sum(b.numel() for b in model.buffers())
234
+ print(f"{a.variant}: {a.layers} layers, seq {a.seq}, {n/1e6:.1f}M params+constants")
235
+
236
+ x = torch.randn(1, a.seq, C) if a.variant == "a-std" else torch.randn(1, C, 1, a.seq)
237
+ mask = torch.zeros(1, 1, 1, a.seq)
238
+ with torch.no_grad():
239
+ traced = torch.jit.trace(model, (x, mask))
240
+ ml = ct.convert(
241
+ traced,
242
+ inputs=[ct.TensorType(name="hidden", shape=x.shape, dtype=np.float16),
243
+ ct.TensorType(name="mask", shape=mask.shape, dtype=np.float16)],
244
+ outputs=[ct.TensorType(name="out", dtype=np.float16)],
245
+ convert_to="mlprogram",
246
+ compute_precision=ct.precision.FLOAT16,
247
+ minimum_deployment_target=ct.target.macOS14,
248
+ )
249
+ ml.save(a.out)
250
+ print("saved", a.out)
research/ane-ceiling/summarize.py ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json, sys, collections
2
+ d = json.load(open(sys.argv[1]))
3
+ for mode in ("cpuAndNeuralEngine", "all"):
4
+ ops = [o for o in d[mode] if o["op"] != "const"]
5
+ dev = collections.Counter(o["preferred"] for o in ops)
6
+ seq = [o["preferred"] for o in ops]
7
+ sw = sum(1 for a, b in zip(seq, seq[1:]) if a != b)
8
+ off = collections.Counter(o["op"] for o in ops if o["preferred"] != "ANE")
9
+ nosup = collections.Counter(o["op"] for o in ops if "ANE" not in o["supported"])
10
+ print(f" {mode:19} ops={len(ops):5} {dict(dev)} switches={sw}")
11
+ if off: print(f" not on ANE: {dict(off)}")
12
+ if nosup: print(f" ANE can't run: {dict(nosup)}")
scripts/build/build_multifunction.py ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Merge the four bucket packages into one multifunction package with shared weights.
2
+
3
+ Run convert_bucket.py for each length first, e.g.
4
+
5
+ cd scripts/build
6
+ for L in 64 128 256 512; do
7
+ uv run python convert_bucket.py --length $L --max-heads 4 --max-options 32 --output-dir build
8
+ done
9
+ uv run python build_multifunction.py --build-dir build
10
+
11
+ Each function (L64, L128, L256, L512) keeps its own fixed-shape graph; identical weight
12
+ tensors are stored once (~0.9 GB instead of ~3.5 GB for four separate packages).
13
+ """
14
+ import argparse
15
+ from pathlib import Path
16
+
17
+ import coremltools as ct
18
+
19
+ from runtime import package_name
20
+
21
+ BUCKETS = (64, 128, 256, 512)
22
+ ROOT = Path(__file__).resolve().parents[2]
23
+
24
+
25
+ def main():
26
+ parser = argparse.ArgumentParser()
27
+ parser.add_argument("--build-dir", default="build")
28
+ parser.add_argument("--output", default=str(ROOT / "models" / "GLiNER2.5-Decide-MultiFn-fp16.mlpackage"))
29
+ args = parser.parse_args()
30
+ desc = ct.utils.MultiFunctionDescriptor()
31
+ for length in BUCKETS:
32
+ source = Path(args.build_dir) / package_name("fp16", length, 4, 32)
33
+ desc.add_function(str(source), src_function_name="main", target_function_name=f"L{length}")
34
+ desc.default_function_name = "L128"
35
+ ct.utils.save_multifunction(desc, args.output)
36
+ print("saved", args.output)
37
+
38
+
39
+ if __name__ == "__main__":
40
+ main()
scripts/build/convert_bucket.py ADDED
@@ -0,0 +1,189 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Convert and verify the pinned GLiNER2.5-Decide multi-head classification path.
2
+
3
+ Modified from Fluid Inference's gliner2-5-decide-coreml `convert-coreml.py` (Apache-2.0):
4
+ traces with an example that fits the requested length and verifies every example that fits,
5
+ so buckets shorter than the original examples (L64) can be built and checked.
6
+ """
7
+ import argparse
8
+ import json
9
+ import math
10
+ from pathlib import Path
11
+
12
+ import coremltools as ct
13
+ import numpy as np
14
+ import torch
15
+ from gliner2 import AutoExtractor
16
+ from huggingface_hub import snapshot_download
17
+ from transformers.models.deberta_v2 import modeling_deberta_v2
18
+
19
+ from convert_names import MODEL_ID, MODEL_REVISION
20
+ from export_model import GLiNER2DecideExport, coreml_safe_attention_forward
21
+ from preprocessing import classification_schema, prepare_decision
22
+ from runtime import decode, package_name
23
+
24
+ EXAMPLES = [
25
+ ("My subscription renewed on April 15 for ¥5,400 after the service was already down. "
26
+ "Can I get that charge refunded?",
27
+ {"intent": ["order_status", "refund_request", "cancel_subscription", "update_payment", "login_problem",
28
+ "shipping_delay", "bug_report", "speak_to_human"]}),
29
+ ("Battery dies before lunch, but the keyboard and the screen are the best I have used on a laptop.",
30
+ {"sentiment": ["positive", "negative", "mixed", "neutral"],
31
+ "aspects": {"labels": ["battery", "keyboard", "screen", "camera", "price", "support"],
32
+ "multi_label": True, "cls_threshold": 0.4}}),
33
+ ("From: compliance@group.example\nSubject: Protocol update — action required today\n\n"
34
+ "Please confirm the new retention rule is applied before Friday's audit.",
35
+ {"intent": ["fyi", "request", "approval", "complaint", "newsletter", "security_alert"],
36
+ "urgency": ["low", "normal", "high", "critical"],
37
+ "route": ["support", "billing", "legal", "security", "finance", "archive"]}),
38
+ ("My subscription renewed on April 15 for 5,400 yen after the service was already down. "
39
+ "Can I get that charge refunded?",
40
+ {"intent": ["order_status", "refund_request", "cancel_subscription", "update_payment", "other"],
41
+ "urgency": ["low", "normal", "high"]}),
42
+ ]
43
+
44
+
45
+ def fits(processor, text, tasks, bucket):
46
+ try:
47
+ prepare_decision(processor, text, tasks, *bucket)
48
+ return True
49
+ except ValueError:
50
+ return False
51
+
52
+
53
+ def native_logits(native, text, tasks):
54
+ """Per-head native logits via the upstream span collator, encoder and shared classifier."""
55
+ from gliner2.training.trainer import ExtractorCollator
56
+
57
+ collator = ExtractorCollator(native.processor, is_training=False, max_len=None, architecture="span")
58
+ batch = collator([(text, classification_schema(tasks).build())])
59
+ _, schema_embs = native._encode_batch(batch)
60
+ return [native.classifier(torch.stack(embs[1:])).squeeze(-1) for embs in schema_embs[0]]
61
+
62
+
63
+ def main():
64
+ parser = argparse.ArgumentParser()
65
+ parser.add_argument("--output-dir", default="build")
66
+ parser.add_argument("--length", type=int, default=128)
67
+ parser.add_argument("--max-heads", type=int, default=4)
68
+ parser.add_argument("--max-options", type=int, default=8)
69
+ parser.add_argument("--precision", choices=["fp16", "fp32"], default="fp16")
70
+ args = parser.parse_args()
71
+ torch.set_num_threads(4)
72
+ source = snapshot_download(
73
+ MODEL_ID, revision=MODEL_REVISION,
74
+ allow_patterns=[
75
+ "config.json", "encoder_config/*", "model.safetensors", "tokenizer.json", "tokenizer_config.json",
76
+ "special_tokens_map.json",
77
+ ],
78
+ )
79
+ native = AutoExtractor.from_pretrained(source, map_location="cpu").eval()
80
+ wrapper = GLiNER2DecideExport(native).eval()
81
+ bucket = (args.length, args.max_heads, args.max_options)
82
+ fitting = [e for e in EXAMPLES if fits(native.processor, *e, bucket)]
83
+ if not fitting:
84
+ raise RuntimeError(f"no example fits bucket {bucket}")
85
+ print(f"examples fitting L{args.length}: {len(fitting)}/{len(EXAMPLES)}")
86
+ text, tasks = fitting[-1] if args.length < 128 else (EXAMPLES[2] if EXAMPLES[2] in fitting else fitting[-1])
87
+ arrays = prepare_decision(native.processor, text, tasks, *bucket)
88
+ tensors = tuple(torch.from_numpy(value) for value in arrays.values())
89
+
90
+ def wrapper_error():
91
+ with torch.no_grad():
92
+ expected = native_logits(native, text, tasks)
93
+ actual = wrapper(*tensors)[0][0]
94
+ return max(float((row - actual[h, : len(row)]).abs().max()) for h, row in enumerate(expected))
95
+
96
+ error = wrapper_error()
97
+ if error > 1e-4:
98
+ raise RuntimeError(f"Wrapper/native logit mismatch: {error}")
99
+ # The upstream scale is a constant for a fixed DeBERTa attention head width.
100
+ # Its traced int32 sqrt is rejected by Core ML; freeze the identical float32
101
+ # value while tracing, and restore the upstream implementation immediately.
102
+ original_scale = modeling_deberta_v2.scaled_size_sqrt
103
+ original_rpos = modeling_deberta_v2.build_rpos
104
+ original_attention = modeling_deberta_v2.DisentangledSelfAttention.forward
105
+
106
+ def static_scale(query_layer, scale_factor):
107
+ value = math.sqrt(float(query_layer.shape[-1] * scale_factor))
108
+ return torch.tensor(value, dtype=torch.float32, device=query_layer.device)
109
+
110
+ modeling_deberta_v2.scaled_size_sqrt = static_scale
111
+ # The encoder only uses self-attention: query and key sequence lengths are
112
+ # identical, so the scripted build_rpos returns relative_pos unchanged.
113
+ modeling_deberta_v2.build_rpos = lambda query, key, relative_pos, buckets, max_pos: relative_pos
114
+ modeling_deberta_v2.DisentangledSelfAttention.forward = coreml_safe_attention_forward
115
+ try:
116
+ frozen_error = wrapper_error()
117
+ if frozen_error > 1e-4:
118
+ raise RuntimeError(f"Frozen attention scale changed native logits: {frozen_error}")
119
+ with torch.no_grad():
120
+ traced = torch.jit.trace(wrapper, tensors, check_trace=False)
121
+ finally:
122
+ modeling_deberta_v2.scaled_size_sqrt = original_scale
123
+ modeling_deberta_v2.build_rpos = original_rpos
124
+ modeling_deberta_v2.DisentangledSelfAttention.forward = original_attention
125
+ grid = (1, args.max_heads, args.max_options)
126
+ converted = ct.convert(
127
+ traced, convert_to="mlprogram", minimum_deployment_target=ct.target.iOS17,
128
+ compute_precision=ct.precision.FLOAT16 if args.precision == "fp16" else ct.precision.FLOAT32,
129
+ compute_units=ct.ComputeUnit.CPU_ONLY,
130
+ inputs=[
131
+ ct.TensorType(name="input_ids", shape=(1, args.length), dtype=np.int32),
132
+ ct.TensorType(name="attention_mask", shape=(1, args.length), dtype=np.int32),
133
+ ct.TensorType(name="marker_indices", shape=grid, dtype=np.int32),
134
+ ct.TensorType(name="marker_mask", shape=grid, dtype=np.float32),
135
+ ],
136
+ outputs=[ct.TensorType(name="logits", dtype=np.float32), ct.TensorType(name="probabilities", dtype=np.float32)],
137
+ )
138
+ converted.short_description = "GLiNER2.5-Decide multi-head schema classification path"
139
+ converted.author = "Fastino (original); Fluid Inference (Core ML conversion)"
140
+ converted.license = "Apache-2.0"
141
+ converted.user_defined_metadata.update({
142
+ "source_model": MODEL_ID, "source_revision": MODEL_REVISION,
143
+ "scope": "classification only; span and count heads not exported",
144
+ "length": str(args.length), "max_heads": str(args.max_heads), "max_options": str(args.max_options),
145
+ })
146
+ out = Path(args.output_dir)
147
+ out.mkdir(parents=True, exist_ok=True)
148
+ package = out / package_name(args.precision, *bucket)
149
+ converted.save(str(package))
150
+ runtime = ct.models.MLModel(str(package), compute_units=ct.ComputeUnit.ALL)
151
+ cases = []
152
+ for text, tasks in fitting:
153
+ native_output = native.classify_text(text, tasks, include_confidence=True)
154
+ logits = np.asarray(runtime.predict(prepare_decision(native.processor, text, tasks, *bucket))["logits"])[0]
155
+ coreml_output = decode(tasks, logits)
156
+ if json.dumps(labels_only(native_output)) != json.dumps(labels_only(coreml_output)):
157
+ raise RuntimeError(f"Core ML/native mismatch: {coreml_output} != {native_output}")
158
+ cases.append({"text": text, "native": native_output, "coreml": coreml_output,
159
+ "max_confidence_error": confidence_error(native_output, coreml_output)})
160
+ report = {
161
+ "source_model": MODEL_ID, "source_revision": MODEL_REVISION, "package": str(package),
162
+ "package_bytes": sum(f.stat().st_size for f in package.rglob("*") if f.is_file()),
163
+ "native_total_parameters": sum(p.numel() for p in native.parameters()),
164
+ "exported_parameters": sum(p.numel() for p in wrapper.parameters()),
165
+ "wrapper_max_logit_error": error, "coremltools": ct.__version__, "torch": torch.__version__,
166
+ "cases": cases,
167
+ }
168
+ (out / f"conversion-{package.stem}.json").write_text(json.dumps(report, indent=2, ensure_ascii=False) + "\n")
169
+ print(json.dumps(report, indent=2, ensure_ascii=False))
170
+
171
+
172
+ def _entries(value):
173
+ return value if isinstance(value, list) else [value]
174
+
175
+
176
+ def labels_only(result: dict) -> dict:
177
+ return {task: sorted(entry["label"] for entry in _entries(value)) for task, value in result.items()}
178
+
179
+
180
+ def confidence_error(native: dict, coreml: dict) -> float:
181
+ errors = []
182
+ for task, value in native.items():
183
+ predicted = {entry["label"]: entry["confidence"] for entry in _entries(coreml[task])}
184
+ errors += [abs(entry["confidence"] - predicted[entry["label"]]) for entry in _entries(value)]
185
+ return max(errors)
186
+
187
+
188
+ if __name__ == "__main__":
189
+ main()
scripts/build/convert_names.py ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ """Pinned source checkpoint shared by the conversion, verification and scoring scripts."""
2
+ MODEL_ID = "fastino/GLiNER2.5-Decide"
3
+ MODEL_REVISION = "65624f1a0265b3f612bae66a2685a06b94a68a9d"
scripts/build/export_model.py ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Native GLiNER2.5-Decide multi-head classification path with explicit marker routing."""
2
+ import torch
3
+ from torch import nn
4
+ from transformers.models.deberta_v2 import modeling_deberta_v2
5
+
6
+
7
+ def coreml_safe_attention_forward(
8
+ self, hidden_states, attention_mask, output_attentions=False,
9
+ query_states=None, relative_pos=None, rel_embeddings=None,
10
+ ):
11
+ """Native DeBERTa attention with a finite mask sentinel for FP16 Core ML."""
12
+ if query_states is None:
13
+ query_states = hidden_states
14
+ query = self.transpose_for_scores(self.query_proj(query_states), self.num_attention_heads)
15
+ key = self.transpose_for_scores(self.key_proj(hidden_states), self.num_attention_heads)
16
+ value = self.transpose_for_scores(self.value_proj(hidden_states), self.num_attention_heads)
17
+ factor = 1 + int("c2p" in self.pos_att_type) + int("p2c" in self.pos_att_type)
18
+ scale = modeling_deberta_v2.scaled_size_sqrt(query, factor)
19
+ scores = torch.bmm(query, key.transpose(-1, -2) / scale.to(dtype=query.dtype))
20
+ if self.relative_attention:
21
+ relative = self.disentangled_attention_bias(
22
+ query, key, relative_pos, self.pos_dropout(rel_embeddings), factor
23
+ )
24
+ scores = scores + relative
25
+ scores = scores.view(-1, self.num_attention_heads, scores.size(-2), scores.size(-1))
26
+ scores = scores.masked_fill(~attention_mask.bool(), -1e4)
27
+ probabilities = self.dropout(torch.softmax(scores, dim=-1))
28
+ context = torch.bmm(probabilities.view(-1, probabilities.size(-2), probabilities.size(-1)), value)
29
+ context = context.view(-1, self.num_attention_heads, context.size(-2), context.size(-1))
30
+ context = context.permute(0, 2, 1, 3).contiguous()
31
+ context = context.view(context.size()[:-2] + (-1,))
32
+ return (context, probabilities) if output_attentions else (context, None)
33
+
34
+
35
+ class GLiNER2DecideExport(nn.Module):
36
+ """Encoder plus the shared label classifier, scored for every (head, label) marker at once.
37
+
38
+ Logits are returned per head so the host can apply softmax (single-label) or sigmoid
39
+ (multi-label) exactly as the native runtime does. Probabilities are the per-head softmax.
40
+ """
41
+
42
+ def __init__(self, native: nn.Module):
43
+ super().__init__()
44
+ self.encoder = native.encoder
45
+ self.classifier = native.classifier
46
+
47
+ def forward(self, input_ids, attention_mask, marker_indices, marker_mask):
48
+ hidden = self.encoder(input_ids=input_ids.long(), attention_mask=attention_mask.long()).last_hidden_state
49
+ heads, options = marker_indices.shape[1], marker_indices.shape[2]
50
+ flat = marker_indices.long().reshape(1, heads * options, 1).expand(-1, -1, hidden.shape[-1])
51
+ states = hidden.gather(1, flat)
52
+ logits = self.classifier(states).reshape(1, heads, options)
53
+ logits = torch.where(marker_mask > 0.5, logits, torch.full_like(logits, -1e4))
54
+ return logits, torch.softmax(logits, dim=-1)
scripts/build/preprocessing.py ADDED
@@ -0,0 +1,67 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Native GLiNER2 span-architecture schema preprocessing for a fixed Core ML bucket."""
2
+ import numpy as np
3
+ from gliner2 import Schema
4
+ from gliner2.models.base import load_extractor_tokenizer
5
+ from gliner2.processor import SchemaTransformer
6
+ from gliner2.training.trainer import ExtractorCollator
7
+
8
+
9
+ def load_processor(tokenizer_dir: str):
10
+ """Load only the tokenizer and schema formatter needed by the Core ML model."""
11
+ return SchemaTransformer(tokenizer=load_extractor_tokenizer(tokenizer_dir), token_pooling="first")
12
+
13
+
14
+ def classification_schema(tasks: dict) -> Schema:
15
+ """Same task-dict handling as native ``classify_text``."""
16
+ schema = Schema()
17
+ for name, config in tasks.items():
18
+ if isinstance(config, dict) and "labels" in config:
19
+ cfg = config.copy()
20
+ labels = cfg.pop("labels")
21
+ schema.classification(name, labels, **cfg)
22
+ else:
23
+ schema.classification(name, config)
24
+ return schema
25
+
26
+
27
+ def task_labels(tasks: dict) -> dict[str, list[str]]:
28
+ result = {}
29
+ for name, config in tasks.items():
30
+ labels = config["labels"] if isinstance(config, dict) and "labels" in config else config
31
+ result[name] = list(labels.keys()) if isinstance(labels, dict) else list(labels)
32
+ return result
33
+
34
+
35
+ def prepare_decision(processor, text: str, tasks: dict, length: int, max_heads: int, max_options: int):
36
+ """Tokenize ``tasks`` exactly as the native span collator does and pad into the bucket."""
37
+ if not 1 <= len(tasks) <= max_heads:
38
+ raise ValueError(f"Expected 1..{max_heads} decision heads, got {len(tasks)}")
39
+ labels = task_labels(tasks)
40
+ for name, values in labels.items():
41
+ if not 1 <= len(values) <= max_options:
42
+ raise ValueError(f"Head {name!r} has {len(values)} labels; bucket holds 1..{max_options}")
43
+ collator = ExtractorCollator(processor, is_training=False, max_len=None, architecture="span")
44
+ batch = collator([(text, classification_schema(tasks).build())])
45
+ ids = batch.input_ids.numpy()
46
+ if ids.shape[1] > length:
47
+ raise ValueError(f"Schema and text require {ids.shape[1]} subwords; bucket holds {length}")
48
+ groups = batch.schema_special_indices[0]
49
+ if len(groups) != len(labels):
50
+ raise ValueError("Schema head count does not match the requested tasks")
51
+ indices = np.zeros((1, max_heads, max_options), dtype=np.int32)
52
+ mask = np.zeros((1, max_heads, max_options), dtype=np.float32)
53
+ for head, (positions, values) in enumerate(zip(groups, labels.values())):
54
+ # positions[0] is the [P] prompt marker; the rest are one [L] marker per label.
55
+ markers = list(positions[1:])
56
+ if len(markers) != len(values):
57
+ raise ValueError("Label markers were truncated or merged")
58
+ indices[0, head, : len(markers)] = markers
59
+ mask[0, head, : len(markers)] = 1.0
60
+ attention = batch.attention_mask.numpy()
61
+ pad = processor.tokenizer.pad_token_id
62
+ return {
63
+ "input_ids": np.pad(ids, ((0, 0), (0, length - ids.shape[1])), constant_values=pad).astype(np.int32),
64
+ "attention_mask": np.pad(attention, ((0, 0), (0, length - attention.shape[1]))).astype(np.int32),
65
+ "marker_indices": indices,
66
+ "marker_mask": mask,
67
+ }
scripts/build/runtime.py ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Run the GLiNER2.5-Decide Core ML classifier without loading the original model weights."""
2
+ import argparse
3
+ import json
4
+ from pathlib import Path
5
+
6
+ import coremltools as ct
7
+ import numpy as np
8
+
9
+ from preprocessing import load_processor, prepare_decision, task_labels
10
+
11
+
12
+ def package_name(precision: str, length: int, max_heads: int, max_options: int) -> str:
13
+ return f"gliner2_decide_classification_{precision}_L{length}_H{max_heads}_K{max_options}.mlpackage"
14
+
15
+
16
+ def decode(tasks: dict, logits: np.ndarray) -> dict:
17
+ """Apply the native activation and label rules to per-head logits of shape (heads, options)."""
18
+ results = {}
19
+ for head, (name, labels) in enumerate(task_labels(tasks).items()):
20
+ config = tasks[name] if isinstance(tasks[name], dict) else {}
21
+ values = logits[head, : len(labels)].astype(np.float64)
22
+ activation = config.get("class_act", "auto")
23
+ multi = config.get("multi_label", False)
24
+ if activation == "sigmoid" or (activation == "auto" and multi):
25
+ probs = 1.0 / (1.0 + np.exp(-values))
26
+ else:
27
+ probs = np.exp(values - values.max())
28
+ probs /= probs.sum()
29
+ if multi:
30
+ threshold = config.get("cls_threshold", 0.5)
31
+ chosen = [{"label": labels[j], "confidence": float(probs[j])} for j in range(len(labels))
32
+ if probs[j] >= threshold]
33
+ best = int(probs.argmax())
34
+ results[name] = chosen or [{"label": labels[best], "confidence": float(probs[best])}]
35
+ else:
36
+ best = int(probs.argmax())
37
+ results[name] = {"label": labels[best], "confidence": float(probs[best])}
38
+ return results
39
+
40
+
41
+ class CoreMLDecide:
42
+ def __init__(self, model_dir: str, precision: str = "fp16", length: int = 128, max_heads: int = 4,
43
+ max_options: int | None = None, compute_units=ct.ComputeUnit.ALL):
44
+ model_dir = Path(model_dir)
45
+ # Published buckets: L128 holds 8 labels per head, L256 and L512 hold 32.
46
+ max_options = max_options or (8 if length == 128 else 32)
47
+ self.length, self.max_heads, self.max_options = length, max_heads, max_options
48
+ self.processor = load_processor(str(model_dir))
49
+ package = model_dir / package_name(precision, length, max_heads, max_options)
50
+ self.model = ct.models.MLModel(str(package), compute_units=compute_units)
51
+
52
+ def classify(self, text: str, tasks: dict) -> dict:
53
+ arrays = prepare_decision(self.processor, text, tasks, self.length, self.max_heads, self.max_options)
54
+ logits = np.asarray(self.model.predict(arrays)["logits"])[0]
55
+ return decode(tasks, logits)
56
+
57
+
58
+ def main():
59
+ parser = argparse.ArgumentParser()
60
+ parser.add_argument("--model-dir", required=True)
61
+ parser.add_argument("--text", required=True)
62
+ parser.add_argument("--tasks", required=True, help='JSON object, e.g. {"intent": ["a", "b"]}')
63
+ parser.add_argument("--precision", default="fp16")
64
+ parser.add_argument("--length", type=int, default=128)
65
+ args = parser.parse_args()
66
+ model = CoreMLDecide(args.model_dir, args.precision, args.length)
67
+ print(json.dumps(model.classify(args.text, json.loads(args.tasks)), indent=2))
68
+
69
+
70
+ if __name__ == "__main__":
71
+ main()
scripts/make_golden.py ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Record native (PyTorch) GLiNER2.5-Decide answers for tests/cases.py into tests/golden.json.
2
+
3
+ Downloads the pinned source checkpoint (~1.7 GB) on first run. The tests compare the Core ML
4
+ router against this file, so they need neither PyTorch weights nor a network connection.
5
+
6
+ uv run python scripts/make_golden.py
7
+ """
8
+ import json
9
+ import sys
10
+ import warnings
11
+ from pathlib import Path
12
+
13
+ warnings.filterwarnings("ignore")
14
+ ROOT = Path(__file__).resolve().parents[1]
15
+ sys.path.insert(0, str(ROOT / "tests"))
16
+ sys.path.insert(0, str(ROOT / "scripts" / "build"))
17
+
18
+ from gliner2 import AutoExtractor # noqa: E402
19
+ from huggingface_hub import snapshot_download # noqa: E402
20
+
21
+ from cases import CASES # noqa: E402
22
+ from convert_names import MODEL_ID, MODEL_REVISION # noqa: E402
23
+
24
+
25
+ def main():
26
+ source = snapshot_download(
27
+ MODEL_ID, revision=MODEL_REVISION,
28
+ allow_patterns=["config.json", "encoder_config/*", "model.safetensors", "tokenizer.json",
29
+ "tokenizer_config.json", "special_tokens_map.json"],
30
+ )
31
+ native = AutoExtractor.from_pretrained(source, map_location="cpu").eval()
32
+ golden = {}
33
+ for case in CASES:
34
+ golden[case["id"]] = native.classify_text(case["text"], case["tasks"], include_confidence=True)
35
+ print(case["id"], json.dumps(golden[case["id"]]))
36
+ out = ROOT / "tests" / "golden.json"
37
+ out.write_text(json.dumps({"source_model": MODEL_ID, "source_revision": MODEL_REVISION, "cases": golden},
38
+ indent=2) + "\n")
39
+ print("wrote", out)
40
+
41
+
42
+ if __name__ == "__main__":
43
+ main()
src/gliner_decide_coreml/__init__.py ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ """GLiNER2.5-Decide on Core ML: one multifunction package, automatic 64/128/256/512 bucket routing."""
2
+ from .router import BUCKETS, MAX_HEADS, MAX_OPTIONS, DecideRouter, Route
3
+
4
+ __all__ = ["BUCKETS", "MAX_HEADS", "MAX_OPTIONS", "DecideRouter", "Route"]
src/gliner_decide_coreml/decoding.py ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Per-head logits -> labels, with the native ``classify_text`` activation and threshold rules.
2
+
3
+ Single-label heads use softmax and return one label; multi-label heads use sigmoid and return
4
+ every label at or above ``cls_threshold`` (or the single best one if none pass), as in
5
+ Fluid Inference's gliner2-5-decide-coreml `runtime.decode` (Apache-2.0).
6
+ """
7
+ import numpy as np
8
+
9
+ from .encoding import task_labels
10
+
11
+
12
+ def head_config(tasks: dict, name: str) -> dict:
13
+ return tasks[name] if isinstance(tasks[name], dict) else {}
14
+
15
+
16
+ def probabilities(config: dict, logits: np.ndarray) -> np.ndarray:
17
+ values = logits.astype(np.float64)
18
+ activation = config.get("class_act", "auto")
19
+ if activation == "sigmoid" or (activation == "auto" and config.get("multi_label", False)):
20
+ return 1.0 / (1.0 + np.exp(-values))
21
+ probs = np.exp(values - values.max())
22
+ return probs / probs.sum()
23
+
24
+
25
+ def format_head(config: dict, labels: list[str], probs: np.ndarray):
26
+ if config.get("multi_label", False):
27
+ threshold = config.get("cls_threshold", 0.5)
28
+ chosen = [{"label": labels[j], "confidence": float(probs[j])}
29
+ for j in range(len(labels)) if probs[j] >= threshold]
30
+ best = int(probs.argmax())
31
+ return chosen or [{"label": labels[best], "confidence": float(probs[best])}]
32
+ best = int(probs.argmax())
33
+ return {"label": labels[best], "confidence": float(probs[best])}
34
+
35
+
36
+ def head_probabilities(tasks: dict, logits: np.ndarray) -> dict[str, np.ndarray]:
37
+ """logits: (heads, options) from the model -> probabilities per head name."""
38
+ return {
39
+ name: probabilities(head_config(tasks, name), logits[h, : len(labels)])
40
+ for h, (name, labels) in enumerate(task_labels(tasks).items())
41
+ }
42
+
43
+
44
+ def decode(tasks: dict, probs: dict[str, np.ndarray]) -> dict:
45
+ labels = task_labels(tasks)
46
+ return {name: format_head(head_config(tasks, name), labels[name], p) for name, p in probs.items()}
47
+
48
+
49
+ def merge_chunks(tasks: dict, per_chunk: list[dict[str, np.ndarray]]) -> dict[str, np.ndarray]:
50
+ """Combine per-chunk probabilities for text longer than the largest bucket.
51
+
52
+ Heuristic, since no single chunk saw the whole text:
53
+ single-label -> the distribution from the chunk that is most confident for that head
54
+ multi-label -> per-label maximum across chunks (a label is on if any chunk says so)
55
+ """
56
+ merged = {}
57
+ for name in task_labels(tasks):
58
+ rows = [chunk[name] for chunk in per_chunk]
59
+ if head_config(tasks, name).get("multi_label", False):
60
+ merged[name] = np.max(np.stack(rows), axis=0)
61
+ else:
62
+ merged[name] = max(rows, key=lambda p: float(p.max()))
63
+ return merged
src/gliner_decide_coreml/encoding.py ADDED
@@ -0,0 +1,95 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Text + task schema -> model inputs, using GLiNER2's own processor.
2
+
3
+ This is the exact preprocessing the Core ML buckets were converted and verified with.
4
+ Adapted from Fluid Inference's gliner2-5-decide-coreml `preprocessing.py` (Apache-2.0).
5
+ """
6
+ from dataclasses import dataclass
7
+ from pathlib import Path
8
+
9
+ import numpy as np
10
+ from gliner2 import Schema
11
+ from gliner2.models.base import load_extractor_tokenizer
12
+ from gliner2.processor import SchemaTransformer
13
+ from gliner2.training.trainer import ExtractorCollator
14
+
15
+
16
+ def load_collator(tokenizer_dir: Path) -> ExtractorCollator:
17
+ """Tokenizer + schema formatter + collator; no PyTorch weights are loaded."""
18
+ processor = SchemaTransformer(tokenizer=load_extractor_tokenizer(str(tokenizer_dir)), token_pooling="first")
19
+ return ExtractorCollator(processor, is_training=False, max_len=None, architecture="span")
20
+
21
+
22
+ def task_labels(tasks: dict) -> dict[str, list[str]]:
23
+ """Label names per head, in declaration order."""
24
+ result = {}
25
+ for name, config in tasks.items():
26
+ labels = config["labels"] if isinstance(config, dict) and "labels" in config else config
27
+ result[name] = list(labels.keys()) if isinstance(labels, dict) else list(labels)
28
+ return result
29
+
30
+
31
+ def classification_schema(tasks: dict) -> Schema:
32
+ """Same task-dict handling as native ``classify_text``."""
33
+ schema = Schema()
34
+ for name, config in tasks.items():
35
+ if isinstance(config, dict) and "labels" in config:
36
+ cfg = config.copy()
37
+ labels = cfg.pop("labels")
38
+ schema.classification(name, labels, **cfg)
39
+ else:
40
+ schema.classification(name, config)
41
+ return schema
42
+
43
+
44
+ @dataclass(frozen=True)
45
+ class Encoded:
46
+ """One request as the encoder sees it, before padding."""
47
+
48
+ input_ids: np.ndarray # (n,) int32: schema tokens, [SEP_TEXT], text tokens
49
+ attention_mask: np.ndarray # (n,) int32
50
+ markers: list[list[int]] # per head: token position of each [L] label marker
51
+
52
+ @property
53
+ def length(self) -> int:
54
+ return len(self.input_ids)
55
+
56
+
57
+ def encode(collator: ExtractorCollator, text: str, tasks: dict) -> Encoded:
58
+ batch = collator([(text, classification_schema(tasks).build())])
59
+ labels = task_labels(tasks)
60
+ groups = batch.schema_special_indices[0]
61
+ if len(groups) != len(labels):
62
+ raise ValueError("Schema head count does not match the requested tasks")
63
+ markers = []
64
+ for positions, values in zip(groups, labels.values()):
65
+ # positions[0] is the [P] prompt marker; the rest are one [L] marker per label.
66
+ found = [int(p) for p in positions[1:]]
67
+ if len(found) != len(values):
68
+ raise ValueError("Label markers were truncated or merged")
69
+ markers.append(found)
70
+ return Encoded(
71
+ input_ids=batch.input_ids[0].numpy().astype(np.int32),
72
+ attention_mask=batch.attention_mask[0].numpy().astype(np.int32),
73
+ markers=markers,
74
+ )
75
+
76
+
77
+ def to_bucket(enc: Encoded, length: int, max_heads: int, max_options: int, pad_id: int) -> dict[str, np.ndarray]:
78
+ """Pad one request into a fixed (length, heads, options) bucket."""
79
+ n = enc.length
80
+ if n > length:
81
+ raise ValueError(f"Request needs {n} tokens; bucket holds {length}")
82
+ if len(enc.markers) > max_heads:
83
+ raise ValueError(f"{len(enc.markers)} heads; bucket holds {max_heads}")
84
+ ids = np.full((1, length), pad_id, dtype=np.int32)
85
+ ids[0, :n] = enc.input_ids
86
+ attention = np.zeros((1, length), dtype=np.int32)
87
+ attention[0, :n] = enc.attention_mask
88
+ indices = np.zeros((1, max_heads, max_options), dtype=np.int32)
89
+ mask = np.zeros((1, max_heads, max_options), dtype=np.float32)
90
+ for head, markers in enumerate(enc.markers):
91
+ if len(markers) > max_options:
92
+ raise ValueError(f"Head has {len(markers)} labels; bucket holds {max_options}")
93
+ indices[0, head, : len(markers)] = markers
94
+ mask[0, head, : len(markers)] = 1.0
95
+ return {"input_ids": ids, "attention_mask": attention, "marker_indices": indices, "marker_mask": mask}
src/gliner_decide_coreml/router.py ADDED
@@ -0,0 +1,133 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Route each request to the smallest Core ML bucket that fits it.
2
+
3
+ One multifunction package holds four fixed-shape graphs (functions ``L64``, ``L128``, ``L256``,
4
+ ``L512``) that share a single copy of the weights. A request is tokenized once, sent to the
5
+ smallest function whose length fits, and decoded on the host. Text that does not fit even
6
+ ``L512`` is split into overlapping word chunks whose results are merged.
7
+ """
8
+ from dataclasses import dataclass
9
+ from pathlib import Path
10
+
11
+ import coremltools as ct
12
+
13
+ from .decoding import decode, head_probabilities, merge_chunks
14
+ from .encoding import Encoded, encode, load_collator, task_labels, to_bucket
15
+
16
+ BUCKETS = (64, 128, 256, 512)
17
+ MAX_HEADS = 4 # heads scored per model call; more are split across calls
18
+ MAX_OPTIONS = 32 # labels per head
19
+ PACKAGE = "GLiNER2.5-Decide-MultiFn-fp16"
20
+ DEFAULT_MODEL_DIR = Path(__file__).resolve().parents[2] / "models"
21
+
22
+
23
+ @dataclass(frozen=True)
24
+ class Route:
25
+ """How a request was served."""
26
+
27
+ tokens: int # schema + text tokens (largest head group)
28
+ bucket: int | None # function used; None when the text had to be chunked
29
+ chunks: int # 1 unless the text was longer than the largest bucket
30
+ calls: int # model predictions made
31
+
32
+
33
+ def _ensure_compiled(package: Path) -> Path:
34
+ compiled = package.with_suffix(".mlmodelc")
35
+ if not compiled.exists() or compiled.stat().st_mtime < package.stat().st_mtime:
36
+ ct.models.utils.compile_model(str(package), destination_path=str(compiled))
37
+ return compiled
38
+
39
+
40
+ class DecideRouter:
41
+ """GLiNER2.5-Decide classification on Core ML with automatic bucket selection.
42
+
43
+ Loads every bucket at construction, since the first load of a function compiles it for
44
+ the device, so do this once at startup, not per request.
45
+ ``compute_units`` defaults to the GPU: on an M1 Pro it beats ``ALL`` at every bucket size.
46
+ """
47
+
48
+ def __init__(self, model_dir: Path | str = DEFAULT_MODEL_DIR,
49
+ compute_units: ct.ComputeUnit = ct.ComputeUnit.CPU_AND_GPU,
50
+ chunk_overlap_words: int = 32):
51
+ model_dir = Path(model_dir)
52
+ self.collator = load_collator(model_dir / "tokenizer")
53
+ self.pad_id = self.collator.processor.tokenizer.pad_token_id
54
+ compiled = _ensure_compiled(model_dir / f"{PACKAGE}.mlpackage")
55
+ self.models = {
56
+ length: ct.models.CompiledMLModel(str(compiled), compute_units=compute_units, function_name=f"L{length}")
57
+ for length in BUCKETS
58
+ }
59
+ self.chunk_overlap_words = chunk_overlap_words
60
+
61
+ def count_tokens(self, text: str, tasks: dict) -> int:
62
+ return encode(self.collator, text, tasks).length
63
+
64
+ def classify(self, text: str, tasks: dict) -> dict:
65
+ return self.classify_with_route(text, tasks)[0]
66
+
67
+ def classify_with_route(self, text: str, tasks: dict) -> tuple[dict, Route]:
68
+ if not tasks:
69
+ raise ValueError("At least one classification head is required")
70
+ for name, labels in task_labels(tasks).items():
71
+ if not 1 <= len(labels) <= MAX_OPTIONS:
72
+ raise ValueError(f"Head {name!r} has {len(labels)} labels; supported: 1..{MAX_OPTIONS}")
73
+ # Heads beyond MAX_HEADS go to separate calls. Heads in one call see each other's labels,
74
+ # so splitting can shift results slightly compared with a single native pass.
75
+ names = list(tasks)
76
+ groups = [{n: tasks[n] for n in names[i:i + MAX_HEADS]} for i in range(0, len(names), MAX_HEADS)]
77
+ result, routes = {}, []
78
+ for group in groups:
79
+ group_result, route = self._classify_group(text, group)
80
+ result.update(group_result)
81
+ routes.append(route)
82
+ largest = max(routes, key=lambda r: r.tokens)
83
+ return {name: result[name] for name in names}, Route(
84
+ tokens=largest.tokens,
85
+ bucket=None if any(r.bucket is None for r in routes) else max(r.bucket for r in routes),
86
+ chunks=max(r.chunks for r in routes),
87
+ calls=sum(r.calls for r in routes),
88
+ )
89
+
90
+ def _classify_group(self, text: str, tasks: dict) -> tuple[dict, Route]:
91
+ enc = encode(self.collator, text, tasks)
92
+ bucket = self._bucket_for(enc)
93
+ if bucket is not None:
94
+ return decode(tasks, self._probabilities(enc, tasks, bucket)), Route(enc.length, bucket, 1, 1)
95
+ chunks = self._split(text, tasks)
96
+ per_chunk = []
97
+ for chunk in chunks:
98
+ chunk_enc = encode(self.collator, chunk, tasks)
99
+ per_chunk.append(self._probabilities(chunk_enc, tasks, self._bucket_for(chunk_enc)))
100
+ return decode(tasks, merge_chunks(tasks, per_chunk)), Route(enc.length, None, len(chunks), len(chunks))
101
+
102
+ @staticmethod
103
+ def _bucket_for(enc: Encoded) -> int | None:
104
+ return next((length for length in BUCKETS if enc.length <= length), None)
105
+
106
+ def _probabilities(self, enc: Encoded, tasks: dict, bucket: int):
107
+ arrays = to_bucket(enc, bucket, MAX_HEADS, MAX_OPTIONS, self.pad_id)
108
+ logits = self.models[bucket].predict(arrays)["logits"][0]
109
+ return head_probabilities(tasks, logits)
110
+
111
+ def _split(self, text: str, tasks: dict) -> list[str]:
112
+ """Greedy word chunks that each fit the largest bucket, overlapping by a few words."""
113
+ words = text.split()
114
+ limit = BUCKETS[-1]
115
+
116
+ def fits(start: int, end: int) -> bool:
117
+ return encode(self.collator, " ".join(words[start:end]), tasks).length <= limit
118
+
119
+ if not fits(0, 1):
120
+ raise ValueError(f"The task schema alone does not fit in {limit} tokens")
121
+ chunks, start = [], 0
122
+ while True:
123
+ lo, hi = start + 1, len(words) # invariant: words[start:lo] fits
124
+ while lo < hi:
125
+ mid = (lo + hi + 1) // 2
126
+ if fits(start, mid):
127
+ lo = mid
128
+ else:
129
+ hi = mid - 1
130
+ chunks.append(" ".join(words[start:lo]))
131
+ if lo == len(words):
132
+ return chunks
133
+ start = max(lo - self.chunk_overlap_words, start + 1)
tests/cases.py ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Test requests, sized to land in each bucket. Shared by scripts/make_golden.py and the tests."""
2
+
3
+ SUPPORT_TASKS = {
4
+ "intent": ["order_status", "refund_request", "cancel_subscription", "update_payment", "other"],
5
+ "urgency": ["low", "normal", "high"],
6
+ }
7
+
8
+ CASES = [
9
+ {
10
+ "id": "short-refund",
11
+ "bucket": 64,
12
+ "text": "My subscription renewed on April 15 for 5,400 yen after the service was already down. "
13
+ "Can I get that charge refunded?",
14
+ "tasks": SUPPORT_TASKS,
15
+ },
16
+ {
17
+ "id": "laptop-review-aspects",
18
+ "bucket": 64,
19
+ "text": "Battery dies before lunch, but the keyboard and the screen are the best I have used on a laptop.",
20
+ "tasks": {
21
+ "sentiment": ["positive", "negative", "mixed", "neutral"],
22
+ "aspects": {"labels": ["battery", "keyboard", "screen", "camera", "price", "support"],
23
+ "multi_label": True, "cls_threshold": 0.4},
24
+ },
25
+ },
26
+ {
27
+ "id": "compliance-email-3-heads",
28
+ "bucket": 128,
29
+ "text": "From: compliance@group.example\nSubject: Protocol update, action required today\n\n"
30
+ "Please confirm the new retention rule is applied before Friday's audit.",
31
+ "tasks": {
32
+ "intent": ["fyi", "request", "approval", "complaint", "newsletter", "security_alert"],
33
+ "urgency": ["low", "normal", "high", "critical"],
34
+ "route": ["support", "billing", "legal", "security", "finance", "archive"],
35
+ },
36
+ },
37
+ {
38
+ "id": "descriptions-and-prompt",
39
+ "bucket": 128,
40
+ "text": "My subscription renewed on April 15 for 5,400 yen after the service was already down. "
41
+ "Can I get that charge refunded?",
42
+ "tasks": {
43
+ "intent": {"labels": {
44
+ "order_status": "The customer asks where an order is",
45
+ "refund_request": "The customer wants money back",
46
+ "cancel_subscription": "The customer wants to stop a recurring plan",
47
+ "update_payment": "The customer wants to change a card or billing method",
48
+ "other": "None of the above",
49
+ }},
50
+ "urgency": {"labels": ["low", "normal", "high"], "prompt": "How urgently should support respond?"},
51
+ },
52
+ },
53
+ {
54
+ "id": "hotel-complaint-4-heads",
55
+ "bucket": 256,
56
+ "text": "Hello, I'm writing from room 1408. The air conditioning stopped working yesterday afternoon "
57
+ "and the room is now above 30 degrees, so neither of us could sleep last night. We called the "
58
+ "front desk twice and were told a technician would come, but nobody showed up. We have a "
59
+ "conference presentation tomorrow morning and cannot spend another night like this. Please "
60
+ "either fix it within the next two hours or move us to another room tonight. If neither is "
61
+ "possible we will check out early and expect the remaining nights to be refunded. We would "
62
+ "also like the incidentals hold on our card to be released, since we were told at check-in "
63
+ "that it would only last one day. Thank you.",
64
+ "tasks": {
65
+ "intent": ["maintenance", "room_change", "checkout", "billing", "complaint", "amenity_request"],
66
+ "priority": ["low", "normal", "high", "urgent"],
67
+ "needs_human": ["yes", "no"],
68
+ "topics": {"labels": ["hvac", "billing", "housekeeping", "noise", "safety"],
69
+ "multi_label": True, "cls_threshold": 0.4},
70
+ },
71
+ },
72
+ {
73
+ "id": "incident-report-long",
74
+ "bucket": 512,
75
+ "text": "Incident summary for the payments platform, prepared for the weekly operations review. "
76
+ "On Tuesday at 09:12 UTC the card authorization service began returning timeouts for roughly "
77
+ "one in five requests. The on-call engineer was paged four minutes later by the latency alert "
78
+ "and started investigating the database tier, which showed normal load. At 09:31 the team noticed "
79
+ "that a configuration change deployed the previous evening had lowered the connection pool size "
80
+ "for the authorization service from two hundred to twenty connections. Under the morning traffic "
81
+ "peak the pool was exhausted, requests queued, and the upstream gateway gave up after its "
82
+ "three second deadline. The change had been reviewed, but the reviewer assumed the value was "
83
+ "per worker rather than per service, and the staging environment does not receive enough "
84
+ "traffic to reveal the problem. The configuration was rolled back at 09:44 and error rates "
85
+ "returned to normal within two minutes. In total about eleven thousand authorizations failed; "
86
+ "most customers retried successfully, but around four hundred orders were abandoned. No card "
87
+ "data was exposed and there is no indication of any security issue. Follow-up actions: add a "
88
+ "load test that replays production traffic volumes against staging before configuration "
89
+ "changes to the authorization service are approved; make pool sizes explicit about their "
90
+ "scope in the configuration schema; add an alert on connection pool saturation so the "
91
+ "on-call engineer is pointed at the right component immediately; and review whether the "
92
+ "gateway deadline should be longer for authorization calls. Finance has asked for an estimate "
93
+ "of lost revenue from abandoned orders by the end of the week, and customer support has "
94
+ "prepared a short statement in case merchants ask about the failed transactions. The service "
95
+ "owner will present these actions at the architecture forum next Thursday and track them to "
96
+ "completion in the reliability backlog.",
97
+ "tasks": {
98
+ "doc_type": ["incident_report", "feature_request", "meeting_notes", "security_advisory", "marketing"],
99
+ "severity": ["sev1", "sev2", "sev3", "sev4"],
100
+ "teams": {"labels": ["payments", "security", "finance", "support", "marketing", "legal"],
101
+ "multi_label": True, "cls_threshold": 0.4},
102
+ },
103
+ },
104
+ ]
tests/golden.json ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "source_model": "fastino/GLiNER2.5-Decide",
3
+ "source_revision": "65624f1a0265b3f612bae66a2685a06b94a68a9d",
4
+ "cases": {
5
+ "short-refund": {
6
+ "intent": {
7
+ "label": "refund_request",
8
+ "confidence": 0.9987769722938538
9
+ },
10
+ "urgency": {
11
+ "label": "low",
12
+ "confidence": 0.3725019693374634
13
+ }
14
+ },
15
+ "laptop-review-aspects": {
16
+ "sentiment": {
17
+ "label": "positive",
18
+ "confidence": 0.9994937181472778
19
+ },
20
+ "aspects": [
21
+ {
22
+ "label": "battery",
23
+ "confidence": 0.7610353231430054
24
+ },
25
+ {
26
+ "label": "keyboard",
27
+ "confidence": 0.9939987659454346
28
+ },
29
+ {
30
+ "label": "screen",
31
+ "confidence": 0.9846210479736328
32
+ }
33
+ ]
34
+ },
35
+ "compliance-email-3-heads": {
36
+ "intent": {
37
+ "label": "request",
38
+ "confidence": 0.44433778524398804
39
+ },
40
+ "urgency": {
41
+ "label": "critical",
42
+ "confidence": 0.411828875541687
43
+ },
44
+ "route": {
45
+ "label": "legal",
46
+ "confidence": 0.5800722241401672
47
+ }
48
+ },
49
+ "descriptions-and-prompt": {
50
+ "intent": {
51
+ "label": "refund_request",
52
+ "confidence": 0.914971113204956
53
+ },
54
+ "urgency": {
55
+ "label": "normal",
56
+ "confidence": 0.38015663623809814
57
+ }
58
+ },
59
+ "hotel-complaint-4-heads": {
60
+ "intent": {
61
+ "label": "checkout",
62
+ "confidence": 0.9280232191085815
63
+ },
64
+ "priority": {
65
+ "label": "urgent",
66
+ "confidence": 0.3841314911842346
67
+ },
68
+ "needs_human": {
69
+ "label": "yes",
70
+ "confidence": 0.6280083060264587
71
+ },
72
+ "topics": [
73
+ {
74
+ "label": "hvac",
75
+ "confidence": 0.8771321773529053
76
+ },
77
+ {
78
+ "label": "billing",
79
+ "confidence": 0.6899283528327942
80
+ }
81
+ ]
82
+ },
83
+ "incident-report-long": {
84
+ "doc_type": {
85
+ "label": "incident_report",
86
+ "confidence": 0.5693653225898743
87
+ },
88
+ "severity": {
89
+ "label": "sev3",
90
+ "confidence": 0.33189529180526733
91
+ },
92
+ "teams": [
93
+ {
94
+ "label": "payments",
95
+ "confidence": 0.9709568619728088
96
+ },
97
+ {
98
+ "label": "finance",
99
+ "confidence": 0.8886100053787231
100
+ },
101
+ {
102
+ "label": "support",
103
+ "confidence": 0.8042415976524353
104
+ }
105
+ ]
106
+ }
107
+ }
108
+ }
tests/test_router.py ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The Core ML router must reproduce the native PyTorch model (tests/golden.json)."""
2
+ import json
3
+ from pathlib import Path
4
+
5
+ import pytest
6
+
7
+ from cases import CASES, SUPPORT_TASKS
8
+ from gliner_decide_coreml import BUCKETS, DecideRouter
9
+
10
+ GOLDEN = json.loads((Path(__file__).parent / "golden.json").read_text())["cases"]
11
+ CONFIDENCE_TOLERANCE = 0.02 # fp16 Core ML vs fp32 PyTorch
12
+
13
+
14
+ @pytest.fixture(scope="module")
15
+ def router():
16
+ return DecideRouter()
17
+
18
+
19
+ def entries(value):
20
+ return value if isinstance(value, list) else [value]
21
+
22
+
23
+ @pytest.mark.parametrize("case", CASES, ids=[c["id"] for c in CASES])
24
+ def test_matches_native(router, case):
25
+ result, route = router.classify_with_route(case["text"], case["tasks"])
26
+ assert route.bucket == case["bucket"] and route.calls == 1
27
+ native = GOLDEN[case["id"]]
28
+ assert list(result) == list(native)
29
+ for head, expected in native.items():
30
+ got = {e["label"]: e["confidence"] for e in entries(result[head])}
31
+ want = {e["label"]: e["confidence"] for e in entries(expected)}
32
+ assert set(got) == set(want), f"{head}: {sorted(got)} != {sorted(want)}"
33
+ for label, confidence in want.items():
34
+ assert got[label] == pytest.approx(confidence, abs=CONFIDENCE_TOLERANCE), f"{head}/{label}"
35
+
36
+
37
+ def test_every_bucket_is_exercised():
38
+ assert {c["bucket"] for c in CASES} == set(BUCKETS)
39
+
40
+
41
+ def test_long_text_is_chunked(router):
42
+ text = " ".join([CASES[0]["text"]] * 40) # ~1,000 tokens: larger than the 512 bucket
43
+ result, route = router.classify_with_route(text, SUPPORT_TASKS)
44
+ assert route.bucket is None and route.chunks >= 2 and route.calls == route.chunks
45
+ assert result["intent"]["label"] == "refund_request"
46
+
47
+
48
+ def test_more_than_four_heads_are_split_across_calls(router):
49
+ tasks = {**CASES[4]["tasks"], "language": ["english", "spanish", "german"], "sentiment": ["positive", "negative"]}
50
+ result, route = router.classify_with_route(CASES[4]["text"], tasks)
51
+ assert list(result) == list(tasks) and route.calls == 2
52
+ assert result["language"]["label"] == "english"
53
+
54
+
55
+ def test_too_many_labels_is_rejected(router):
56
+ with pytest.raises(ValueError, match="labels"):
57
+ router.classify("hello", {"topic": [f"label_{i}" for i in range(33)]})
uv.lock ADDED
The diff for this file is too large to render. See raw diff