Views
No views yet
1from transformers import AutoTokenizer
2import onnxruntime
3import numpy as np
4import torch.nn.functional as F
5
6def decode_sentiment(idx: int) -> str:
7 sentiment_map = {0: 'positive', 1: 'neutral', 2: 'negative'}
8 return sentiment_map[idx]
9
10def decode_market(idx: int) -> str:
11 market_map = {
12 0: 'strong bullish',
13 1: 'bullish',
14 2: 'neutral',
15 3: 'bearish',
16 4: 'strong bearish'
17 }
18 return market_map[idx]
19
20def softmax(x, axis=1):
21 exp_x = np.exp(x - np.max(x, axis=axis, keepdims=True))
22 return exp_x / np.sum(exp_x, axis=axis, keepdims=True)
23
24tokenizer = AutoTokenizer.from_pretrained("microsoft/deberta-v3-small")
25
26text = "input-text-goes-here"
27
28inputs = tokenizer(
29 text,
30 return_tensors="pt",
31 padding="max_length",
32 truncation=True,
33 max_length=512
34)
35input_ids = inputs["input_ids"]
36attention_mask = inputs["attention_mask"]
37
38ort_inputs = {
39 "input_ids": input_ids.cpu().numpy(),
40 "attention_mask": attention_mask.cpu().numpy()
41}
42
43session = onnxruntime.InferenceSession("a1-debertav3.onnx")
44
45sentiment_logits, market_logits = session.run(None, ort_inputs)
46
47sentiment_probs = softmax(sentiment_logits, axis=1)
48market_probs = softmax(market_logits, axis=1)
49
50sentiment_pred = np.argmax(sentiment_probs, axis=1)
51market_pred = np.argmax(market_probs, axis=1)
52
53decoded_sentiment = decode_sentiment(sentiment_pred.item())
54decoded_market = decode_market(market_pred.item())
55print(f"Sentiment: {decoded_sentiment}")
56print(f"Market: {decoded_market}")
57