Download handler.py from crystas/distilbert-command-data-tagger: direct link, hf CLI and curl.
- Browser
- Download file 3.53 kB
-
https://huggingface.co/crystas/distilbert-command-data-tagger/resolve/main/handler.py
- Command line
-
hf download hf://crystas/distilbert-command-data-tagger/handler.py
-
curl -L -o handler.py https://huggingface.co/crystas/distilbert-command-data-tagger/resolve/main/handler.py
3.53 kB
| 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 | |
| } |