Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -0,0 +1,109 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import streamlit as st
|
| 2 |
+
import os
|
| 3 |
+
import torch
|
| 4 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 5 |
+
|
| 6 |
+
# ================ Конфигурация страницы ================
|
| 7 |
+
st.set_page_config(page_title="Генератор отзывов", initial_sidebar_state="expanded")
|
| 8 |
+
st.title("Генератор отзывов на основе ИИ")
|
| 9 |
+
st.write("Создавайте текстовые отзывы с помощью нейросети на основе категорий, рейтинга и ключевых слов.")
|
| 10 |
+
st.sidebar.title("Параметры генерации")
|
| 11 |
+
|
| 12 |
+
# ================ Загрузка модели ================
|
| 13 |
+
@st.cache_resource
|
| 14 |
+
def load_model():
|
| 15 |
+
"""Загружает модель и токенизатор из локальной директории."""
|
| 16 |
+
model_path = "./model"
|
| 17 |
+
if not os.path.exists(model_path):
|
| 18 |
+
st.error(f"Модель не найдена по пути {model_path}. Проверьте путь.")
|
| 19 |
+
st.stop()
|
| 20 |
+
try:
|
| 21 |
+
with st.spinner("Загрузка модели..."):
|
| 22 |
+
model = AutoModelForCausalLM.from_pretrained(model_path)
|
| 23 |
+
tokenizer = AutoTokenizer.from_pretrained(model_path)
|
| 24 |
+
return model, tokenizer
|
| 25 |
+
except Exception as e:
|
| 26 |
+
st.error(f"Ошибка при загрузке модели: {e}")
|
| 27 |
+
st.stop()
|
| 28 |
+
|
| 29 |
+
# ================ Параметры генерации ================
|
| 30 |
+
def get_generation_params():
|
| 31 |
+
"""Получает параметры для генерации текста из боковой панели."""
|
| 32 |
+
params = {
|
| 33 |
+
"max_length": st.sidebar.slider("Максимальная длина", 50, 300, 200),
|
| 34 |
+
"num_return_sequences": st.sidebar.number_input("Количество вариантов", 1, 5, 3),
|
| 35 |
+
"temperature": st.sidebar.slider("Температура", 0.01, 2.0, 0.9, 0.05),
|
| 36 |
+
"top_p": st.sidebar.slider("Top-p", 0.01, 1.0, 0.95, 0.05),
|
| 37 |
+
"top_k": st.sidebar.number_input("Top-k", 1, 100, 60),
|
| 38 |
+
"do_sample": st.sidebar.checkbox("Случайная выборка", value=True),
|
| 39 |
+
"no_repeat_ngram_size": st.sidebar.number_input("Размер n-грамм", 1, 10, 2),
|
| 40 |
+
}
|
| 41 |
+
return params
|
| 42 |
+
|
| 43 |
+
# ================ Генерация текста ================
|
| 44 |
+
def generate_text(model, tokenizer, prompt, params):
|
| 45 |
+
"""Генерирует текст с использованием модели и заданных параметров."""
|
| 46 |
+
input_ids = tokenizer.encode(prompt, return_tensors="pt")
|
| 47 |
+
|
| 48 |
+
if params["num_return_sequences"] > 1 and not params["do_sample"]:
|
| 49 |
+
st.warning("Для генерации нескольких текстов активирована случайная выборка.")
|
| 50 |
+
params["do_sample"] = True
|
| 51 |
+
|
| 52 |
+
with torch.no_grad():
|
| 53 |
+
outputs = model.generate(
|
| 54 |
+
input_ids=input_ids,
|
| 55 |
+
max_length=params["max_length"],
|
| 56 |
+
temperature=params["temperature"],
|
| 57 |
+
top_p=params["top_p"],
|
| 58 |
+
top_k=params["top_k"],
|
| 59 |
+
do_sample=params["do_sample"],
|
| 60 |
+
no_repeat_ngram_size=params["no_repeat_ngram_size"],
|
| 61 |
+
num_return_sequences=params["num_return_sequences"],
|
| 62 |
+
pad_token_id=tokenizer.eos_token_id,
|
| 63 |
+
)
|
| 64 |
+
return [tokenizer.decode(output, skip_special_tokens=True) for output in outputs]
|
| 65 |
+
|
| 66 |
+
# ================ Форматирование текста ================
|
| 67 |
+
def format_text(text):
|
| 68 |
+
"""Добавляет заглавные буквы и точки для улучшения читаемости."""
|
| 69 |
+
sentences = text.split(". ")
|
| 70 |
+
formatted = [s.capitalize().strip() + "." if not s.endswith(".") else s.capitalize().strip() for s in sentences if s]
|
| 71 |
+
return " ".join(formatted)
|
| 72 |
+
|
| 73 |
+
# ================ Основное приложение ================
|
| 74 |
+
def main():
|
| 75 |
+
st.sidebar.subheader("Параметры генерации текста")
|
| 76 |
+
model, tokenizer = load_model()
|
| 77 |
+
params = get_generation_params()
|
| 78 |
+
|
| 79 |
+
# Поля ввода данных
|
| 80 |
+
categories = ["Кафе", "Ресторан", "Кондитерская", "Парк", "Музей", "Отель", "Магазин"]
|
| 81 |
+
category = st.selectbox("Категория:", categories)
|
| 82 |
+
custom_category = st.text_input("Или введите свою категорию:")
|
| 83 |
+
category = custom_category if custom_category else category
|
| 84 |
+
rating = st.slider("Рейтинг", 1, 5, 5)
|
| 85 |
+
key_words = st.text_input("Ключевые слова (через запятую)", "десерт, торт, вкусно")
|
| 86 |
+
|
| 87 |
+
# Генерация текста
|
| 88 |
+
if st.button("Сгенерировать"):
|
| 89 |
+
if not category or not key_words:
|
| 90 |
+
st.warning("Пожалуйста, заполните все поля.")
|
| 91 |
+
return
|
| 92 |
+
|
| 93 |
+
input_prompt = f"Напиши отзыв для заведения. Категория: {category}. Рейтинг: {rating} из 5. Включи ключевые слова: {key_words}. Отзыв:"
|
| 94 |
+
st.info("Генерация текстов...")
|
| 95 |
+
|
| 96 |
+
try:
|
| 97 |
+
generated_texts = generate_text(model, tokenizer, input_prompt, params)
|
| 98 |
+
for i, text in enumerate(generated_texts):
|
| 99 |
+
st.subheader(f"Вариант {i + 1}")
|
| 100 |
+
st.write(format_text(text))
|
| 101 |
+
|
| 102 |
+
# Скачать все результаты
|
| 103 |
+
all_texts = "\n\n".join([f"Вариант {i + 1}:\n{text}" for i, text in enumerate(generated_texts)])
|
| 104 |
+
st.download_button("Скачать отзывы", all_texts, file_name="generated_reviews.txt", mime="text/plain")
|
| 105 |
+
except Exception as e:
|
| 106 |
+
st.error(f"Ошибка генерации текста: {e}")
|
| 107 |
+
|
| 108 |
+
if __name__ == "__main__":
|
| 109 |
+
main()
|