File size: 5,389 Bytes
7e7b191 2a3dce9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | ---
language:
- en
license: mit
tags:
- distilbert
- onnx
- intent-classification
- token-classification
- ner
- custom-handler
pipeline_tag: text-classification
widget:
- text: bold this sentence
- text: send email to john@example.com
- text: search for previous sentence delete
base_model:
- distilbert/distilbert-base-uncased
---
# distilbert-command-data-tagger
`distilbert-command-data-tagger` is a lightweight, optimized ONNX model trained for joint **Intent Classification** and **Named Entity Recognition (NER) / Data Tagging** from natural language user commands.
It is designed for lightweight deployment on low-resource runtimes and Hugging Face Inference Endpoints using ONNX Runtime.
---
## Model Architecture & Files
- **Base Model**: `distilbert-base-uncased`
- **Model Format**: ONNX (`joint_command_parser.onnx` + `joint_command_parser.onnx.data`)
- **Max Sequence Length**: 64 tokens
- **Output Heads**:
1. **Intent Logits**: Decodes top-level command intent (with a strict `0.65` confidence threshold fallback to `unknown`).
2. **NER / Data Tag Logits**: Tagging token sequences (e.g., target text to edit, recipient email address, or search queries).
---
## Supported Intents
The model categorizes inputs into one of the following structured categories:
- **Document Editing**: `add_text`, `search_delete`, `prev_sentence_delete`, `bold`, `italic`, `underline`
- **Navigation**: `navigate_email`, `navigate_dm`
- **Email Page Operations**: `email_compose`, `email_send`, `email_reply`, `email_forward`, `email_delete`, `email_search`
- **Direct Message (DM) Operations**: `dm_send`, `dm_reply`, `dm_delete`, `dm_search`
- **Fallback**: `unknown`
---
## Inference with Custom Endpoint Handler (`handler.py`)
Below is the standard Hugging Face `EndpointHandler` used to load and run inference on the ONNX graph:
```python
import os
import numpy as np
import onnxruntime as ort
from transformers import AutoTokenizer
class EndpointHandler:
def __init__(self, path=""):
# 1. Resolve paths for BOTH structural graph and matrix weights
model_path = os.path.join(path, "joint_command_parser.onnx")
# Initialize the ONNX Runtime execution engine thread pool
self.session = ort.InferenceSession(model_path, providers=["CPUExecutionProvider"])
# 2. Match the exact tokenizer backbone used during training
self.tokenizer = AutoTokenizer.from_pretrained("distilbert-base-uncased")
# 3. Structural target layouts
self.intents = [
# document editing
"add_text",
"search_delete",
"prev_sentence_delete",
"bold",
"italic",
"underline",
# navigation
"navigate_email",
"navigate_dm",
# email page
"email_compose",
"email_send",
"email_reply",
"email_forward",
"email_delete",
"email_search",
# dm page
"dm_send",
"dm_reply",
"dm_delete",
"dm_search",
# fallback
"unknown",
]
self.max_len = 64
def __call__(self, data):
# Extract input text payload sent via HTTP POST Request
inputs_payload = data.get("inputs", "")
if not inputs_payload:
return {"error": "Missing 'inputs' string parameter inside payload."}
# Tokenize incoming sequence matching matrix bounds
tokenized = self.tokenizer(
inputs_payload,
max_length=self.max_len,
padding="max_length",
truncation=True,
return_tensors="np"
)
onnx_feeds = {
"input_ids": tokenized["input_ids"].astype(np.int64),
"attention_mask": tokenized["attention_mask"].astype(np.int64)
}
# Run execution graph matrix calculations
intent_logits, ner_logits = self.session.run(None, onnx_feeds)
# Decode Intent array using Softmax calculation
logits_exp = np.exp(intent_logits[0])
probs = logits_exp / np.sum(logits_exp)
intent_idx = np.argmax(probs)
# Apply strict fallback thresholds
confidence = float(probs[intent_idx])
final_intent = self.intents[intent_idx] if confidence >= 0.65 else "unknown"
# Process and decode non-zero NER spans
ner_tags = np.argmax(ner_logits[0], axis=-1)
tokens = self.tokenizer.convert_ids_to_tokens(tokenized["input_ids"][0])
extracted_tokens = []
for token, tag in zip(tokens, ner_tags):
if token in [self.tokenizer.cls_token, self.tokenizer.sep_token, self.tokenizer.pad_token]:
continue
if tag > 0: # Valid B-DATA or I-DATA sequences
cleaned = token.replace("##", "")
if token.startswith("##") and extracted_tokens:
extracted_tokens[-1] += cleaned
else:
extracted_tokens.append(cleaned)
extracted_data = " ".join(extracted_tokens).strip()
return {
"intent": final_intent,
"confidence": confidence,
"extracted_data": extracted_data
}
|