Team Ai
Apppublic

sunwaee/Perceiver-Multiclass-Emotion-Classification

sourceHugging Faceupdated 3y agoView on Hugging Face
10likes
app.py119 linesDownload Raw Back to root
1import os2 3import gdown as gdown4import nltk5import streamlit as st6from nltk.tokenize import sent_tokenize7 8from source.pipeline import MultiLabelPipeline, inputs_to_dataset9 10 11def download_models(ids):12    """13    Download all models.14 15    :param ids: name and links of models16    :return:17    """18 19    # Download sentence tokenizer20    nltk.download('punkt')21 22    # Download model from drive if not stored locally23    for key in ids:24        if not os.path.isfile(f"model/{key}.pt"):25            url = f"https://drive.google.com/uc?id={ids[key]}"26            gdown.download(url=url, output=f"model/{key}.pt")27 28 29@st.cache30def load_labels():31    """32    Load model labels.33 34    :return:35    """36 37    return [38        "admiration",39        "amusement",40        "anger",41        "annoyance",42        "approval",43        "caring",44        "confusion",45        "curiosity",46        "desire",47        "disappointment",48        "disapproval",49        "disgust",50        "embarrassment",51        "excitement",52        "fear",53        "gratitude",54        "grief",55        "joy",56        "love",57        "nervousness",58        "optimism",59        "pride",60        "realization",61        "relief",62        "remorse",63        "sadness",64        "surprise",65        "neutral"66    ]67 68 69@st.cache(allow_output_mutation=True)70def load_model(model_path):71    """72    Load model and cache it.73 74    :param model_path: path to model75    :return:76    """77 78    model = MultiLabelPipeline(model_path=model_path)79 80    return model81 82 83# Page config84st.set_page_config(layout="centered")85st.title("Multiclass Emotion Classification")86st.write("DeepMind Language Perceiver for Multiclass Emotion Classification (Eng). ")87 88maintenance = False89if maintenance:90    st.write("Unavailable for now (file downloads limit). ")91else:92    # Variables93    ids = {'perceiver-go-emotions': st.secrets['model']}94    labels = load_labels()95 96    # Download all models from drive97    download_models(ids)98 99    # Display labels100    st.markdown(f"__Labels:__ {', '.join(labels)}")101 102    # Model selection103    left, right = st.columns([4, 2])104    inputs = left.text_area('', max_chars=4096, value='This is a space about multiclass emotion classification. Write '105                                                      'something here to see what happens!')106    model_path = right.selectbox('', options=[k for k in ids], index=0, help='Model to use. ')107    split = right.checkbox('Split into sentences', value=True)108    model = load_model(model_path=f"model/{model_path}.pt")109    right.write(model.device)110 111    if split:112        if not inputs.isspace() and inputs != "":113            with st.spinner('Processing text... This may take a while.'):114                left.write(model(inputs_to_dataset(sent_tokenize(inputs)), batch_size=1))115    else:116        if not inputs.isspace() and inputs != "":117            with st.spinner('Processing text... This may take a while.'):118                left.write(model(inputs_to_dataset([inputs]), batch_size=1))119