| import os |
| import numpy as np |
| import onnxruntime as ort |
| from transformers import AutoTokenizer |
|
|
| class EndpointHandler: |
| def __init__(self, path=""): |
| |
| model_path = os.path.join(path, "joint_command_parser.onnx") |
| |
| |
| self.session = ort.InferenceSession(model_path, providers=["CPUExecutionProvider"]) |
| |
| |
| self.tokenizer = AutoTokenizer.from_pretrained("distilbert-base-uncased") |
| |
| |
| self.intents = [ |
| |
| "add_text", |
| "search_delete", |
| "prev_sentence_delete", |
| "bold", |
| "italic", |
| "underline", |
| |
| "navigate_email", |
| "navigate_dm", |
| |
| "email_compose", |
| "email_send", |
| "email_reply", |
| "email_forward", |
| "email_delete", |
| "email_search", |
| |
| "dm_send", |
| "dm_reply", |
| "dm_delete", |
| "dm_search", |
| |
| "unknown", |
| ] |
| self.max_len = 64 |
|
|
| def __call__(self, data): |
| |
| inputs_payload = data.get("inputs", "") |
| if not inputs_payload: |
| return {"error": "Missing 'inputs' string parameter inside payload."} |
| |
| |
| 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) |
| } |
| |
| |
| intent_logits, ner_logits = self.session.run(None, onnx_feeds) |
| |
| |
| logits_exp = np.exp(intent_logits[0]) |
| probs = logits_exp / np.sum(logits_exp) |
| intent_idx = np.argmax(probs) |
| |
| |
| confidence = float(probs[intent_idx]) |
| final_intent = self.intents[intent_idx] if confidence >= 0.65 else "unknown" |
| |
| |
| 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: |
| 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 |
| } |