lmnindzja commited on
Commit
c4caa93
·
verified ·
1 Parent(s): 17eec13

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +109 -0
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()