Zero-Shot Classification
Core ML
GLiNER
GLiNER2
English
coremltools
deberta-v3
apple-silicon
fp16
multifunction
Instructions to use augustoFranke/GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- GLiNER
How to use augustoFranke/GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512 with GLiNER:
from gliner import GLiNER model = GLiNER.from_pretrained("augustoFranke/GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512") text = "Cristiano Ronaldo dos Santos Aveiro was born on 5 February 1985 in Funchal, Madeira, Portugal." labels = ["person", "date", "location"] entities = model.predict_entities(text, labels) for entity in entities: print(entity["text"], "=>", entity["label"]) - GLiNER2
How to use augustoFranke/GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512 with GLiNER2:
from gliner2 import AutoExtractor extractor = AutoExtractor.from_pretrained("augustoFranke/GLiNER2.5-Decide-CoreML-FP16-MultiFn-L64-512") # Extract entities text = "Apple CEO Tim Cook announced iPhone 15 in Cupertino yesterday." result = extractor.extract_entities(text, ["company", "person", "product", "location"]) print(result) - Notebooks
- Google Colab
- Kaggle
GLiNER2.5-Decide Core ML multifunction package (L64-L512) with bucket router
Browse files- .gitignore +13 -0
- LICENSE +202 -0
- NOTICE +15 -0
- README.md +176 -0
- bench/benchbuckets.swift +59 -0
- bench/plan.swift +71 -0
- bench/router_latency.py +33 -0
- bench/time_prep.py +25 -0
- models/GLiNER2.5-Decide-MultiFn-fp16.mlpackage/Data/com.apple.CoreML/model.mlmodel +3 -0
- models/GLiNER2.5-Decide-MultiFn-fp16.mlpackage/Data/com.apple.CoreML/weights/weight.bin +3 -0
- models/GLiNER2.5-Decide-MultiFn-fp16.mlpackage/Manifest.json +18 -0
- models/tokenizer/config.json +24 -0
- models/tokenizer/special_tokens_map.json +123 -0
- models/tokenizer/tokenizer.json +0 -0
- models/tokenizer/tokenizer_config.json +156 -0
- pyproject.toml +34 -0
- research/ane-ceiling/models.py +250 -0
- research/ane-ceiling/summarize.py +12 -0
- scripts/build/build_multifunction.py +40 -0
- scripts/build/convert_bucket.py +189 -0
- scripts/build/convert_names.py +3 -0
- scripts/build/export_model.py +54 -0
- scripts/build/preprocessing.py +67 -0
- scripts/build/runtime.py +71 -0
- scripts/make_golden.py +43 -0
- src/gliner_decide_coreml/__init__.py +4 -0
- src/gliner_decide_coreml/decoding.py +63 -0
- src/gliner_decide_coreml/encoding.py +95 -0
- src/gliner_decide_coreml/router.py +133 -0
- tests/cases.py +104 -0
- tests/golden.json +108 -0
- tests/test_router.py +57 -0
- uv.lock +0 -0
.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
|
|
|