nlc-explorer / model_loader.py
butterswords's picture
Create model_loader.py
2facc36 verified
Raw
History Blame Contribute Delete
991 Bytes
import os, spacy
from transformers import AutoTokenizer, AutoModelForSequenceClassification, TextClassificationPipeline
import torch.nn.functional as F
import torch
from lime.lime_text import LimeTextExplainer
# Load spaCy once
try:
nlp = spacy.load("en_core_web_lg")
except OSError:
os.system("python -m spacy download en_core_web_lg")
nlp = spacy.load("en_core_web_lg")
# Load the transformer model once
MODEL_NAME = "distilbert-base-uncased-finetuned-sst-2-english"
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
model = AutoModelForSequenceClassification.from_pretrained(MODEL_NAME)
pipe = TextClassificationPipeline(model=model, tokenizer=tokenizer, top_k=None)
# LIME explainer
explainer = LimeTextExplainer(class_names=['negative', 'positive'])
# Predictor function used by LIME
def predictor(texts):
outputs = model(**tokenizer(texts, return_tensors="pt", padding=True))
probas = F.softmax(outputs.logits, dim=1).detach().numpy()
return probas