Views
No views yet
1# Loading the label mappings
2import json
3def load_label_mappings():
4 with open("./label_mapping.json", encoding="utf-8") as f:
5 data = json.load(f)
6 return data['labels']
7
8label_mappings = load_label_mappings()
9
10# Loading the model
11import onnxruntime as ort
12from transformers import DistilBertTokenizer
13tokenizer = DistilBertTokenizer.from_pretrained('distilbert-base-uncased')
14ort_session = ort.InferenceSession("./toxic-or-neutral-text-labelled.onnx")
15
16# Predicting label for given text
17def predict_via_onnx(text, ort_session, tokenizer, label_mappings):
18 model_expected_input_shape = ort_session.get_inputs()[0].shape
19 print("Model expects input shape:", model_expected_input_shape)
20 inputs = tokenizer(text, return_tensors="np", padding="max_length", truncation=True, max_length=model_expected_input_shape[1])
21 print("input shape", inputs['input_ids'].shape)
22
23 input_ids = inputs['input_ids']
24 if input_ids.ndim == 1:
25 input_ids = input_ids[np.newaxis, :]
26 ort_inputs = {ort_session.get_inputs()[0].name: input_ids}
27
28 ort_inputs['input_ids'] = ort_inputs['input_ids'].astype(np.int64)
29
30 ort_outputs = ort_session.run(None, ort_inputs)
31 predictions = np.argmax(ort_outputs, axis=-1)
32
33 predicted_label = label_mappings[predictions.item()]
34 return predicted_label
35
36predicted_label = predict_via_onnx("How do I get to the beach?", ort_session, tokenizer, label_mappings)
37print(predicted_label)