Spaces:
Runtime error
Runtime error
| 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 | |
| ) | |
| 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)} <EOS>" | |
| 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("<EOS>"): | |
| break | |
| output_list = list(generated.squeeze().numpy()) | |
| output_text = f"{tokenizer.decode(output_list)} <EOS>" | |
| 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="🚨") |