ceyda/ExplaiNER
1
1import streamlit as st2from sentence_transformers import SentenceTransformer3from transformers import AutoModelForTokenClassification # type: ignore4from transformers import AutoTokenizer # type: ignore5 6 7@st.experimental_singleton()8def get_model(model_name: str, labels=None):9 if labels is None:10 return AutoModelForTokenClassification.from_pretrained(11 model_name,12 output_attentions=True,13 ) # type: ignore14 else:15 id2label = {idx: tag for idx, tag in enumerate(labels)}16 label2id = {tag: idx for idx, tag in enumerate(labels)}17 return AutoModelForTokenClassification.from_pretrained(18 model_name,19 output_attentions=True,20 num_labels=len(labels),21 id2label=id2label,22 label2id=label2id,23 ) # type: ignore24 25 26@st.experimental_singleton()27def get_encoder(model_name: str, device: str = "cpu"):28 return SentenceTransformer(model_name, device=device)29 30 31@st.experimental_singleton()32def get_tokenizer(tokenizer_name: str):33 return AutoTokenizer.from_pretrained(tokenizer_name)34 