crystas commited on
Commit
7e7b191
·
verified ·
1 Parent(s): 08a77b4

Create README.md

Browse files

# 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
}

Files changed (1) hide show
  1. README.md +19 -0
README.md ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ language:
3
+ - en
4
+ license: mit
5
+ tags:
6
+ - distilbert
7
+ - onnx
8
+ - intent-classification
9
+ - token-classification
10
+ - ner
11
+ - custom-handler
12
+ pipeline_tag: text-classification
13
+ widget:
14
+ - text: bold this sentence
15
+ - text: send email to john@example.com
16
+ - text: search for previous sentence delete
17
+ base_model:
18
+ - distilbert/distilbert-base-uncased
19
+ ---