besstie-bert-gru-sentiment / modeling_bert_gru.py
milesrack's picture
Upload folder using huggingface_hub
551bb09 verified
Raw
History Blame Contribute Delete
3.21 kB
import torch
import torch.nn as nn
from transformers import BertModel
class BERTGRUSentiment(nn.Module):
def __init__(
self,
bert_model_name="bert-base-uncased",
hidden_dim=256,
output_dim=1,
n_layers=2,
bidirectional=True,
dropout=0.25,
):
super().__init__()
self.bert = BertModel.from_pretrained(bert_model_name)
self.embedding_dim = self.bert.config.hidden_size
self.rnn = nn.GRU(
input_size=self.embedding_dim,
hidden_size=hidden_dim,
num_layers=n_layers,
bidirectional=bidirectional,
dropout=dropout if n_layers > 1 else 0,
batch_first=True,
)
self.out = nn.Linear(
hidden_dim * 2 if bidirectional else hidden_dim,
output_dim,
)
self.dropout = nn.Dropout(dropout)
def forward(self, input_ids):
# Matches the teammate's training code:
# embedded = self.bert(text)[0]
with torch.no_grad():
embedded = self.bert(input_ids)[0]
_, hidden = self.rnn(embedded)
if self.rnn.bidirectional:
hidden = self.dropout(
torch.cat((hidden[-2, :, :], hidden[-1, :, :]), dim=1)
)
else:
hidden = self.dropout(hidden[-1, :, :])
return self.out(hidden)
def build_model(device="cpu"):
model = BERTGRUSentiment(
bert_model_name="bert-base-uncased",
hidden_dim=256,
output_dim=1,
n_layers=2,
bidirectional=True,
dropout=0.25,
)
model.to(device)
model.eval()
return model
def load_model(checkpoint_path, device="cpu"):
model = build_model(device=device)
state_dict = torch.load(checkpoint_path, map_location=device)
if isinstance(state_dict, dict) and "state_dict" in state_dict:
state_dict = state_dict["state_dict"]
if hasattr(state_dict, "state_dict"):
state_dict = state_dict.state_dict()
missing, unexpected = model.load_state_dict(state_dict, strict=True)
if missing or unexpected:
raise RuntimeError(
f"Checkpoint did not load cleanly. Missing={missing}, Unexpected={unexpected}"
)
model.eval()
return model
def predict_sentiment(model, tokenizer, sentence, device="cpu"):
init_token_idx = tokenizer.convert_tokens_to_ids(tokenizer.cls_token)
eos_token_idx = tokenizer.convert_tokens_to_ids(tokenizer.sep_token)
max_input_length = tokenizer.model_max_length
tokens = tokenizer.tokenize(sentence)
tokens = tokens[: max_input_length - 2]
indexed = (
[init_token_idx]
+ tokenizer.convert_tokens_to_ids(tokens)
+ [eos_token_idx]
)
tensor = torch.LongTensor(indexed).to(device)
tensor = tensor.unsqueeze(0)
with torch.no_grad():
probability = torch.sigmoid(model(tensor)).item()
label = "Positive" if probability >= 0.5 else "Negative"
confidence = probability if label == "Positive" else 1 - probability
return {
"label": label,
"positive_probability": probability,
"confidence": confidence,
}