GungnirAP/Youtube-Comments
0
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="๐จ")