augustoFranke's picture
GLiNER2.5-Decide Core ML multifunction package (L64-L512) with bucket router
cb96101 verified
Raw History Blame
9.65 kB
"""Convert and verify the pinned GLiNER2.5-Decide multi-head classification path.
Modified from Fluid Inference's gliner2-5-decide-coreml `convert-coreml.py` (Apache-2.0):
traces with an example that fits the requested length and verifies every example that fits,
so buckets shorter than the original examples (L64) can be built and checked.
"""
import argparse
import json
import math
from pathlib import Path
import coremltools as ct
import numpy as np
import torch
from gliner2 import AutoExtractor
from huggingface_hub import snapshot_download
from transformers.models.deberta_v2 import modeling_deberta_v2
from convert_names import MODEL_ID, MODEL_REVISION
from export_model import GLiNER2DecideExport, coreml_safe_attention_forward
from preprocessing import classification_schema, prepare_decision
from runtime import decode, package_name
EXAMPLES = [
("My subscription renewed on April 15 for ¥5,400 after the service was already down. "
"Can I get that charge refunded?",
{"intent": ["order_status", "refund_request", "cancel_subscription", "update_payment", "login_problem",
"shipping_delay", "bug_report", "speak_to_human"]}),
("Battery dies before lunch, but the keyboard and the screen are the best I have used on a laptop.",
{"sentiment": ["positive", "negative", "mixed", "neutral"],
"aspects": {"labels": ["battery", "keyboard", "screen", "camera", "price", "support"],
"multi_label": True, "cls_threshold": 0.4}}),
("From: compliance@group.example\nSubject: Protocol update — action required today\n\n"
"Please confirm the new retention rule is applied before Friday's audit.",
{"intent": ["fyi", "request", "approval", "complaint", "newsletter", "security_alert"],
"urgency": ["low", "normal", "high", "critical"],
"route": ["support", "billing", "legal", "security", "finance", "archive"]}),
("My subscription renewed on April 15 for 5,400 yen after the service was already down. "
"Can I get that charge refunded?",
{"intent": ["order_status", "refund_request", "cancel_subscription", "update_payment", "other"],
"urgency": ["low", "normal", "high"]}),
]
def fits(processor, text, tasks, bucket):
try:
prepare_decision(processor, text, tasks, *bucket)
return True
except ValueError:
return False
def native_logits(native, text, tasks):
"""Per-head native logits via the upstream span collator, encoder and shared classifier."""
from gliner2.training.trainer import ExtractorCollator
collator = ExtractorCollator(native.processor, is_training=False, max_len=None, architecture="span")
batch = collator([(text, classification_schema(tasks).build())])
_, schema_embs = native._encode_batch(batch)
return [native.classifier(torch.stack(embs[1:])).squeeze(-1) for embs in schema_embs[0]]
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--output-dir", default="build")
parser.add_argument("--length", type=int, default=128)
parser.add_argument("--max-heads", type=int, default=4)
parser.add_argument("--max-options", type=int, default=8)
parser.add_argument("--precision", choices=["fp16", "fp32"], default="fp16")
args = parser.parse_args()
torch.set_num_threads(4)
source = snapshot_download(
MODEL_ID, revision=MODEL_REVISION,
allow_patterns=[
"config.json", "encoder_config/*", "model.safetensors", "tokenizer.json", "tokenizer_config.json",
"special_tokens_map.json",
],
)
native = AutoExtractor.from_pretrained(source, map_location="cpu").eval()
wrapper = GLiNER2DecideExport(native).eval()
bucket = (args.length, args.max_heads, args.max_options)
fitting = [e for e in EXAMPLES if fits(native.processor, *e, bucket)]
if not fitting:
raise RuntimeError(f"no example fits bucket {bucket}")
print(f"examples fitting L{args.length}: {len(fitting)}/{len(EXAMPLES)}")
text, tasks = fitting[-1] if args.length < 128 else (EXAMPLES[2] if EXAMPLES[2] in fitting else fitting[-1])
arrays = prepare_decision(native.processor, text, tasks, *bucket)
tensors = tuple(torch.from_numpy(value) for value in arrays.values())
def wrapper_error():
with torch.no_grad():
expected = native_logits(native, text, tasks)
actual = wrapper(*tensors)[0][0]
return max(float((row - actual[h, : len(row)]).abs().max()) for h, row in enumerate(expected))
error = wrapper_error()
if error > 1e-4:
raise RuntimeError(f"Wrapper/native logit mismatch: {error}")
# The upstream scale is a constant for a fixed DeBERTa attention head width.
# Its traced int32 sqrt is rejected by Core ML; freeze the identical float32
# value while tracing, and restore the upstream implementation immediately.
original_scale = modeling_deberta_v2.scaled_size_sqrt
original_rpos = modeling_deberta_v2.build_rpos
original_attention = modeling_deberta_v2.DisentangledSelfAttention.forward
def static_scale(query_layer, scale_factor):
value = math.sqrt(float(query_layer.shape[-1] * scale_factor))
return torch.tensor(value, dtype=torch.float32, device=query_layer.device)
modeling_deberta_v2.scaled_size_sqrt = static_scale
# The encoder only uses self-attention: query and key sequence lengths are
# identical, so the scripted build_rpos returns relative_pos unchanged.
modeling_deberta_v2.build_rpos = lambda query, key, relative_pos, buckets, max_pos: relative_pos
modeling_deberta_v2.DisentangledSelfAttention.forward = coreml_safe_attention_forward
try:
frozen_error = wrapper_error()
if frozen_error > 1e-4:
raise RuntimeError(f"Frozen attention scale changed native logits: {frozen_error}")
with torch.no_grad():
traced = torch.jit.trace(wrapper, tensors, check_trace=False)
finally:
modeling_deberta_v2.scaled_size_sqrt = original_scale
modeling_deberta_v2.build_rpos = original_rpos
modeling_deberta_v2.DisentangledSelfAttention.forward = original_attention
grid = (1, args.max_heads, args.max_options)
converted = ct.convert(
traced, convert_to="mlprogram", minimum_deployment_target=ct.target.iOS17,
compute_precision=ct.precision.FLOAT16 if args.precision == "fp16" else ct.precision.FLOAT32,
compute_units=ct.ComputeUnit.CPU_ONLY,
inputs=[
ct.TensorType(name="input_ids", shape=(1, args.length), dtype=np.int32),
ct.TensorType(name="attention_mask", shape=(1, args.length), dtype=np.int32),
ct.TensorType(name="marker_indices", shape=grid, dtype=np.int32),
ct.TensorType(name="marker_mask", shape=grid, dtype=np.float32),
],
outputs=[ct.TensorType(name="logits", dtype=np.float32), ct.TensorType(name="probabilities", dtype=np.float32)],
)
converted.short_description = "GLiNER2.5-Decide multi-head schema classification path"
converted.author = "Fastino (original); Fluid Inference (Core ML conversion)"
converted.license = "Apache-2.0"
converted.user_defined_metadata.update({
"source_model": MODEL_ID, "source_revision": MODEL_REVISION,
"scope": "classification only; span and count heads not exported",
"length": str(args.length), "max_heads": str(args.max_heads), "max_options": str(args.max_options),
})
out = Path(args.output_dir)
out.mkdir(parents=True, exist_ok=True)
package = out / package_name(args.precision, *bucket)
converted.save(str(package))
runtime = ct.models.MLModel(str(package), compute_units=ct.ComputeUnit.ALL)
cases = []
for text, tasks in fitting:
native_output = native.classify_text(text, tasks, include_confidence=True)
logits = np.asarray(runtime.predict(prepare_decision(native.processor, text, tasks, *bucket))["logits"])[0]
coreml_output = decode(tasks, logits)
if json.dumps(labels_only(native_output)) != json.dumps(labels_only(coreml_output)):
raise RuntimeError(f"Core ML/native mismatch: {coreml_output} != {native_output}")
cases.append({"text": text, "native": native_output, "coreml": coreml_output,
"max_confidence_error": confidence_error(native_output, coreml_output)})
report = {
"source_model": MODEL_ID, "source_revision": MODEL_REVISION, "package": str(package),
"package_bytes": sum(f.stat().st_size for f in package.rglob("*") if f.is_file()),
"native_total_parameters": sum(p.numel() for p in native.parameters()),
"exported_parameters": sum(p.numel() for p in wrapper.parameters()),
"wrapper_max_logit_error": error, "coremltools": ct.__version__, "torch": torch.__version__,
"cases": cases,
}
(out / f"conversion-{package.stem}.json").write_text(json.dumps(report, indent=2, ensure_ascii=False) + "\n")
print(json.dumps(report, indent=2, ensure_ascii=False))
def _entries(value):
return value if isinstance(value, list) else [value]
def labels_only(result: dict) -> dict:
return {task: sorted(entry["label"] for entry in _entries(value)) for task, value in result.items()}
def confidence_error(native: dict, coreml: dict) -> float:
errors = []
for task, value in native.items():
predicted = {entry["label"]: entry["confidence"] for entry in _entries(coreml[task])}
errors += [abs(entry["confidence"] - predicted[entry["label"]]) for entry in _entries(value)]
return max(errors)
if __name__ == "__main__":
main()