alinikkhah/SpoilerDetectionClassification
0
1import torch2from transformers import BertForSequenceClassification3import gradio as gr4from transformers import BertTokenizer5import torch6from transformers import BertForSequenceClassification, BertTokenizer7import gradio as gr8 9import torch10from transformers import BertForSequenceClassification11 12# Load the model architecture with the number of labels13model = BertForSequenceClassification.from_pretrained('bert-base-uncased', num_labels=2)14 15# Load the state dict while mapping to CPU16try:17 model.load_state_dict(torch.load('bert_model_complete.pth', map_location=torch.device('cpu')), strict=False)18except Exception as e:19 print(f"Error loading state dict: {e}")20 21 22model.eval() # Set the model to evaluation mode23 24 25# Load the tokenizer26tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')27 28def predict(text):29 inputs = tokenizer(text, return_tensors="pt", padding=True, truncation=True)30 with torch.no_grad():31 outputs = model(**inputs)32 logits = outputs.logits33 predicted_class = logits.argmax().item()34 return predicted_class35 36# Set up the Gradio interface37interface = gr.Interface(fn=predict, inputs="text", outputs="label", title="BERT Text Classification")38 39# Load model and tokenizer40model = BertForSequenceClassification.from_pretrained('bert-base-uncased')41model.load_state_dict(torch.load('bert_model_complete.pth'))42model.eval()43 44tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')45 46# Define prediction function47def predict(text):48 inputs = tokenizer(text, return_tensors="pt", padding=True, truncation=True)49 with torch.no_grad():50 outputs = model(**inputs)51 logits = outputs.logits52 predicted_class = logits.argmax().item()53 return predicted_class54 55# Set up Gradio interface56interface = gr.Interface(fn=predict, inputs="text", outputs="label", title="BERT Text Classification")57 58# Launch the interface59if __name__ == "__main__":60 interface.launch()61 