Views
No views yet
1from transformers import AutoTokenizer, AutoModel
2import torch
3import torch.nn as nn
4
5# Define model class
6class TinyBERTDualClassifier(nn.Module):
7 def __init__(self, num_module_labels, num_date_labels, dropout_rate=0.1):
8 super(TinyBERTDualClassifier, self).__init__()
9 self.encoder = AutoModel.from_pretrained("JayShah07/tinybert-dual-classifier")
10 self.hidden_size = self.encoder.config.hidden_size
11 self.dropout = nn.Dropout(p=dropout_rate)
12 self.module_classifier = nn.Linear(self.hidden_size, num_module_labels)
13 self.date_classifier = nn.Linear(self.hidden_size, num_date_labels)
14
15 def forward(self, input_ids, attention_mask):
16 outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
17 cls_output = outputs.last_hidden_state[:, 0, :]
18 cls_output = self.dropout(cls_output)
19 module_logits = self.module_classifier(cls_output)
20 date_logits = self.date_classifier(cls_output)
21 return module_logits, date_logits
22
23# Load model
24classifier_config = torch.hub.load_state_dict_from_url(
25 f"https://huggingface.co/JayShah07/tinybert-dual-classifier/resolve/main/classifier_heads.pt"
26)
27
28model = TinyBERTDualClassifier(
29 num_module_labels=6,
30 num_date_labels=7
31)
32
33model.module_classifier.load_state_dict(classifier_config['module_classifier'])
34model.date_classifier.load_state_dict(classifier_config['date_classifier'])
35
36tokenizer = AutoTokenizer.from_pretrained("JayShah07/tinybert-dual-classifier")
37
38# Inference
39model.eval()
40text = "Show my holdings for this month"
41inputs = tokenizer(text, return_tensors='pt', padding='max_length',
42 truncation=True, max_length=128)
43
44with torch.no_grad():
45 module_logits, date_logits = model(inputs['input_ids'], inputs['attention_mask'])
46 module_pred = torch.argmax(module_logits, dim=1).item()
47 date_pred = torch.argmax(date_logits, dim=1).item()
48
49module_labels = ['holdings', 'capital_gains', 'scheme_wise_returns', 'investment_account_wise_returns', 'portfolio_update', 'None_module']
50date_labels = ['Current Year', 'Previous Year', 'Daily', 'Monthly', 'Weekly', 'Yearly', 'None_date']
51
52print(f"Module: {module_labels[module_pred]}")
53print(f"Date: {date_labels[date_pred]}")1@misc{tinybert-dual-classifier,
2 author = {Jay Shah},
3 title = {TinyBERT Dual Classifier for Investment Reporting},
4 year = {2025},
5 publisher = {Hugging Face},
6 howpublished = {\url{https://huggingface.co/JayShah07/tinybert-dual-classifier}}
7}