CoolFace
Apppublic

GungnirAP/Youtube-Comments

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
app.py92 linesDownload Raw Back to root
1import time2import torch3import torch.nn.functional as F4 5import streamlit as st6from parser import extract_titles7from generator import (8    get_model_and_tokenizer, 9    generate_prompt,10    generate_comment,11    extract_comment_from_prompt12)13 14@st.cache_resource15def preprocess(model_name, model_type="gpt2"):16    tokenizer, model = get_model_and_tokenizer("cpu", "./chkp", model_name, model_type=model_type)17    return tokenizer, model18 19 20 21# Greetings22st.title("Generate YouTube Comments")23 24st.write("This mini-app generates English comments for a YouTube video. It uses a fine-tuned version of GPT-2 by OpenAI. You can find the code in Files section.")25 26 27# Pre-process28tokenizer, model = preprocess("GPT2_02_Ep0_St300000.pt", model_type="gpt2")29model.eval()30filter_value = -float("Inf")31entry_length = 2032 33 34# Settings35input_url = st.text_input(label="Drop URL below", placeholder="https://youtu.be/mCV44C5rQ2M")36col1, col2 = st.columns(2)37with col1:38    temperature = st.slider('Temperature of sampling', 0.0, 1.0, value=0.7, key=7)39with col2:40    top_p = st.slider('Top-p parameter', 0.0, 1.0, value=0.8, key=8)41col1, _ = st.columns(2)42with col1:43    num_of_coms = st.slider('Number of comments', 1, 5, value=1, key=9)44 45# Action46if st.button("Generate text", type="primary"):47    if len(input_url):48        try:49            # generate & print50            channel, title = extract_titles(input_url)51            prompt = generate_prompt(title)52            with torch.no_grad():53                raw_generated = torch.tensor(tokenizer.encode(prompt)).unsqueeze(0)54            55            st.markdown("""---""")56 57            place_holders = []58            texts = []59            for index in range(num_of_coms):60                with st.spinner('Please wait while your comment is being generated...'):61                    place_holders.append(st.empty())62                    generated = raw_generated.clone()63                    with torch.no_grad():64                        for i in range(entry_length):65                            outputs = model(generated, labels=generated)66                            loss, logits = outputs[:2]67                            logits = logits[:, -1, :] / (temperature if temperature > 0 else 1.0)68                            sorted_logits, sorted_indices = torch.sort(logits, descending=True)69                            cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)70                            sorted_indices_to_remove = cumulative_probs > top_p71                            sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()72                            sorted_indices_to_remove[..., 0] = 073                            indices_to_remove = sorted_indices[sorted_indices_to_remove]74                            logits[:, indices_to_remove] = filter_value75                            next_token = torch.multinomial(F.softmax(logits, dim=-1), num_samples=1)76                            generated = torch.cat((generated, next_token), dim=1)77                            output_list = list(generated.squeeze().numpy())78                            output_text = f"{tokenizer.decode(output_list)} <EOS>" 79                            output_text = extract_comment_from_prompt(output_text)80                            place_holders[index].text_area(label=f"Comment #{index + 1}", 81                                                           value=output_text, disabled=True, key=(index + 1)*1000+i)82                            if next_token in tokenizer.encode("<EOS>"):83                                break84                        output_list = list(generated.squeeze().numpy())85                        output_text = f"{tokenizer.decode(output_list)} <EOS>" 86                        texts.append(extract_comment_from_prompt(output_text))87            for index, text in enumerate(texts):88                place_holders[index].text_area(label=f"Comment #{index + 1}", value=text, disabled=False, key=index+100)89        except RuntimeError as e:90            st.error(e.args[0], icon="๐Ÿšจ")91    else:92        st.error('Please enter a url', icon="๐Ÿšจ")