Team Ai
Apppublic

ceyda/ExplaiNER

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
load.py102 linesDownload Raw Back to src
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