import time import torch import torch.nn.functional as F import streamlit as st from parser import extract_titles from generator import ( get_model_and_tokenizer, generate_prompt, generate_comment, extract_comment_from_prompt ) @st.cache_resource def preprocess(model_name, model_type="gpt2"): tokenizer, model = get_model_and_tokenizer("cpu", "./chkp", model_name, model_type=model_type) return tokenizer, model # Greetings st.title("Generate YouTube Comments") st.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.") # Pre-process tokenizer, model = preprocess("GPT2_02_Ep0_St300000.pt", model_type="gpt2") model.eval() filter_value = -float("Inf") entry_length = 20 # Settings input_url = st.text_input(label="Drop URL below", placeholder="https://youtu.be/mCV44C5rQ2M") col1, col2 = st.columns(2) with col1: temperature = st.slider('Temperature of sampling', 0.0, 1.0, value=0.7, key=7) with col2: top_p = st.slider('Top-p parameter', 0.0, 1.0, value=0.8, key=8) col1, _ = st.columns(2) with col1: num_of_coms = st.slider('Number of comments', 1, 5, value=1, key=9) # Action if st.button("Generate text", type="primary"): if len(input_url): try: # generate & print channel, title = extract_titles(input_url) prompt = generate_prompt(title) with torch.no_grad(): raw_generated = torch.tensor(tokenizer.encode(prompt)).unsqueeze(0) st.markdown("""---""") place_holders = [] texts = [] for index in range(num_of_coms): with st.spinner('Please wait while your comment is being generated...'): place_holders.append(st.empty()) generated = raw_generated.clone() with torch.no_grad(): for i in range(entry_length): outputs = model(generated, labels=generated) loss, logits = outputs[:2] logits = logits[:, -1, :] / (temperature if temperature > 0 else 1.0) sorted_logits, sorted_indices = torch.sort(logits, descending=True) cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) sorted_indices_to_remove = cumulative_probs > top_p sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone() sorted_indices_to_remove[..., 0] = 0 indices_to_remove = sorted_indices[sorted_indices_to_remove] logits[:, indices_to_remove] = filter_value next_token = torch.multinomial(F.softmax(logits, dim=-1), num_samples=1) generated = torch.cat((generated, next_token), dim=1) output_list = list(generated.squeeze().numpy()) output_text = f"{tokenizer.decode(output_list)} " output_text = extract_comment_from_prompt(output_text) place_holders[index].text_area(label=f"Comment #{index + 1}", value=output_text, disabled=True, key=(index + 1)*1000+i) if next_token in tokenizer.encode(""): break output_list = list(generated.squeeze().numpy()) output_text = f"{tokenizer.decode(output_list)} " texts.append(extract_comment_from_prompt(output_text)) for index, text in enumerate(texts): place_holders[index].text_area(label=f"Comment #{index + 1}", value=text, disabled=False, key=index+100) except RuntimeError as e: st.error(e.args[0], icon="🚨") else: st.error('Please enter a url', icon="🚨")