ceyda/ExplaiNER
1
1from typing import Optional2 3import pandas as pd4import streamlit as st5from datasets import Dataset # type: ignore6 7from src.data import encode_dataset, get_collator, get_data, predict8from src.model import get_encoder, get_model, get_tokenizer9from src.subpages import Context10from src.utils import align_sample, device, explode_df11 12_TOKENIZER_NAME = (13 "xlm-roberta-base",14 "gagan3012/bert-tiny-finetuned-ner",15 "distilbert-base-german-cased",16)[0]17 18 19def _load_models_and_tokenizer(20 encoder_model_name: str,21 model_name: str,22 tokenizer_name: Optional[str],23 device: str = "cpu",24):25 sentence_encoder = get_encoder(encoder_model_name, device=device)26 tokenizer = get_tokenizer(tokenizer_name if tokenizer_name else model_name)27 labels = "O B-COMMA".split() if "comma" in model_name else None28 model = get_model(model_name, labels=labels)29 return sentence_encoder, model, tokenizer30 31 32@st.cache(allow_output_mutation=True)33def load_context(34 encoder_model_name: str,35 model_name: str,36 ds_name: str,37 ds_config_name: str,38 ds_split_name: str,39 split_sample_size: int,40 randomize_sample: bool,41 **kw_args,42) -> Context:43 """Utility method loading (almost) everything we need for the application.44 This exists just because we want to cache the results of this function.45 46 Args:47 encoder_model_name (str): Name of the sentence encoder to load.48 model_name (str): Name of the NER model to load.49 ds_name (str): Dataset name or path.50 ds_config_name (str): Dataset config name.51 ds_split_name (str): Dataset split name.52 split_sample_size (int): Number of examples to load from the split.53 54 Returns:55 Context: An object containing everything we need for the application.56 """57 58 sentence_encoder, model, tokenizer = _load_models_and_tokenizer(59 encoder_model_name=encoder_model_name,60 model_name=model_name,61 tokenizer_name=_TOKENIZER_NAME if "comma" in model_name else None,62 device=str(device),63 )64 collator = get_collator(tokenizer)65 66 # load data related stuff67 split: Dataset = get_data(68 ds_name, ds_config_name, ds_split_name, split_sample_size, randomize_sample69 )70 tags = split.features["ner_tags"].feature71 split_encoded, word_ids, ids = encode_dataset(split, tokenizer)72 73 # transform into dataframe74 df = predict(split_encoded, model, tokenizer, collator, tags)75 df["word_ids"] = word_ids76 df["ids"] = ids77 78 # explode, clean, merge79 df_tokens = explode_df(df)80 df_tokens_cleaned = df_tokens.query("labels != 'IGN'")81 df_merged = pd.DataFrame(df.apply(align_sample, axis=1).tolist())82 df_tokens_merged = explode_df(df_merged)83 84 return Context(85 **{86 "model": model,87 "tokenizer": tokenizer,88 "sentence_encoder": sentence_encoder,89 "df": df,90 "df_tokens": df_tokens,91 "df_tokens_cleaned": df_tokens_cleaned,92 "df_tokens_merged": df_tokens_merged,93 "tags": tags,94 "labels": tags.names,95 "split_sample_size": split_sample_size,96 "ds_name": ds_name,97 "ds_config_name": ds_config_name,98 "ds_split_name": ds_split_name,99 "split": split,100 }101 )102 