lmnindzja commited on
Commit
c25ca80
·
verified ·
1 Parent(s): e67ebe7

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +48 -42
app.py CHANGED
@@ -7,7 +7,6 @@ from transformers import AutoModelForCausalLM, AutoTokenizer
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
@@ -25,6 +24,17 @@ def load_model():
25
  st.error(f"Ошибка при загрузке модели: {e}")
26
  st.stop()
27
 
 
 
 
 
 
 
 
 
 
 
 
28
  # ================ Форматирование текста ================
29
  def format_text(text):
30
  """Форматирует текст: заглавные буквы и корректные точки."""
@@ -40,48 +50,29 @@ def format_text(text):
40
  formatted_sentences.append(sentence)
41
  return " ".join(formatted_sentences)
42
 
43
- # ================ Оптимизированная генерация ================
44
- def generate_text(model, tokenizer, prompt, params):
45
- """Генерирует текст блоками для ускорения."""
46
- input_ids = tokenizer.encode(prompt, return_tensors="pt")
47
- output_ids = input_ids
48
-
49
- for _ in range(params["max_length"] // 10): # Генерация по 10 токенов
50
- with torch.no_grad():
51
- outputs = model.generate(
52
- input_ids=output_ids,
53
- max_new_tokens=10,
54
- temperature=params["temperature"],
55
- top_p=params["top_p"],
56
- top_k=params["top_k"],
57
- do_sample=params["do_sample"],
58
- no_repeat_ngram_size=params["no_repeat_ngram_size"],
59
- pad_token_id=tokenizer.eos_token_id,
60
- )
61
- new_tokens = outputs[0][output_ids.shape[1]:]
62
- output_ids = outputs
63
- new_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
64
- yield new_text.strip()
65
-
66
  # ================ Основное приложение ================
67
  def main():
68
  model, tokenizer = load_model()
69
 
70
- # Параметры генерации
71
- params = {
72
- "max_length": st.sidebar.slider("Максимальная длина", 50, 300, 150, help="Максимальная длина текста."),
73
- "temperature": st.sidebar.slider("Температура", 0.01, 2.0, 0.7, 0.05, help="Степень случайности."),
74
- "top_p": st.sidebar.slider("Top-p", 0.01, 1.0, 0.9, 0.05, help="Фильтрация по вероятности."),
75
- "top_k": st.sidebar.number_input("Top-k", 1, 100, 50, help="Ограничение по количеству кандидатов."),
76
- "do_sample": st.sidebar.checkbox("Случайная выборка", value=True, help="Включает случайность генерации."),
77
- "no_repeat_ngram_size": st.sidebar.number_input("Размер n-грамм", 1, 10, 2, help="Предотвращает повторение фраз."),
78
- }
79
-
80
- # Поля ввода
81
  category = st.selectbox("Категория:", ["Кафе", "Ресторан", "Парк", "Музей"], help="Выберите категорию.")
82
  rating = st.slider("Рейтинг", 1, 5, 5)
83
  key_words = st.text_input("Ключевые слова", "вкусно, уютно, быстро", help="Введите ключевые слова для отзыва.")
84
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
85
  if st.button("Сгенерировать"):
86
  if not key_words:
87
  st.warning("Введите ключевые слова.")
@@ -90,16 +81,31 @@ def main():
90
  input_prompt = f"Напиши отзыв для заведения. Категория: {category}. Рейтинг: {rating}. Ключевые слова: {key_words}."
91
  st.info("Генерация текста...")
92
 
93
- placeholder = st.empty()
94
- full_text = ""
95
 
96
- for chunk in generate_text(model, tokenizer, input_prompt, params):
97
- full_text += " " + chunk
98
- formatted_text = format_text(full_text)
99
- placeholder.write(formatted_text)
 
 
 
 
 
 
 
 
 
 
 
 
100
 
101
  st.success("Генерация завершена!")
102
- st.download_button("Скачать отзыв", formatted_text, "generated_review.txt")
 
 
 
103
 
104
  if __name__ == "__main__":
105
  main()
 
7
  st.set_page_config(page_title="Генератор отзывов", initial_sidebar_state="expanded")
8
  st.title("Генератор отзывов на основе ИИ")
9
  st.write("Создавайте текстовые отзывы с помощью нейросети на основе категорий, рейтинга и ключевых слов.")
 
10
 
11
  # ================ Загрузка модели ================
12
  @st.cache_resource
 
24
  st.error(f"Ошибка при загрузке модели: {e}")
25
  st.stop()
26
 
27
+ # ================ Предустановленные стили ================
28
+ def get_preset_params(style):
29
+ """Возвращает параметры генерации для выбранного стиля."""
30
+ if style == "Строгий":
31
+ return {"temperature": 0.3, "top_p": 0.8, "top_k": 30, "no_repeat_ngram_size": 3}
32
+ elif style == "Умеренный":
33
+ return {"temperature": 0.7, "top_p": 0.9, "top_k": 50, "no_repeat_ngram_size": 2}
34
+ elif style == "Безумный":
35
+ return {"temperature": 1.5, "top_p": 1.0, "top_k": 100, "no_repeat_ngram_size": 1}
36
+ return {}
37
+
38
  # ================ Форматирование текста ================
39
  def format_text(text):
40
  """Форматирует текст: заглавные буквы и корректные точки."""
 
50
  formatted_sentences.append(sentence)
51
  return " ".join(formatted_sentences)
52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
53
  # ================ Основное приложение ================
54
  def main():
55
  model, tokenizer = load_model()
56
 
57
+ # Основные параметры
 
 
 
 
 
 
 
 
 
 
58
  category = st.selectbox("Категория:", ["Кафе", "Ресторан", "Парк", "Музей"], help="Выберите категорию.")
59
  rating = st.slider("Рейтинг", 1, 5, 5)
60
  key_words = st.text_input("Ключевые слова", "вкусно, уютно, быстро", help="Введите ключевые слова для отзыва.")
61
 
62
+ # Выбор стиля генерации
63
+ style = st.selectbox("Стиль генерации:", ["Строгий", "Умеренный", "Безумный"], help="Выберите стиль текста.")
64
+ params = get_preset_params(style)
65
+
66
+ # Расширенные настройки в сворачивающемся блоке
67
+ with st.expander("Дополнительные настройки"):
68
+ params["temperature"] = st.slider("Температура", 0.01, 2.0, params["temperature"], 0.05, help="Степень случайности.")
69
+ params["top_p"] = st.slider("Top-p", 0.01, 1.0, params["top_p"], 0.05, help="Фильтрация по вероятности.")
70
+ params["top_k"] = st.number_input("Top-k", 1, 100, params["top_k"], help="Ограничение по количеству кандидатов.")
71
+ params["no_repeat_ngram_size"] = st.number_input("Размер n-грамм", 1, 10, params["no_repeat_ngram_size"], help="Предотвращает повторение фраз.")
72
+ max_length = st.slider("Максимальная длина", 50, 300, 150, help="Максимальная длина текста.")
73
+ num_return_sequences = st.number_input("Количество вариантов", 1, 5, 1, help="Сколько вариантов текста нужно сгенерировать.")
74
+
75
+ # Генерация текста
76
  if st.button("Сгенерировать"):
77
  if not key_words:
78
  st.warning("Введите ключевые слова.")
 
81
  input_prompt = f"Напиши отзыв для заведения. Категория: {category}. Рейтинг: {rating}. Ключевые слова: {key_words}."
82
  st.info("Генерация текста...")
83
 
84
+ placeholders = [st.empty() for _ in range(num_return_sequences)]
85
+ texts = []
86
 
87
+ for i in range(num_return_sequences):
88
+ with torch.no_grad():
89
+ output = model.generate(
90
+ tokenizer.encode(input_prompt, return_tensors="pt"),
91
+ max_new_tokens=max_length,
92
+ temperature=params["temperature"],
93
+ top_p=params["top_p"],
94
+ top_k=params["top_k"],
95
+ do_sample=True,
96
+ no_repeat_ngram_size=params["no_repeat_ngram_size"],
97
+ pad_token_id=tokenizer.eos_token_id,
98
+ )
99
+ generated_text = tokenizer.decode(output[0], skip_special_tokens=True)
100
+ formatted_text = format_text(generated_text)
101
+ texts.append(formatted_text)
102
+ placeholders[i].write(f"**Вариант {i + 1}:**\n{formatted_text}")
103
 
104
  st.success("Генерация завершена!")
105
+
106
+ # Скачивание всех вариантов
107
+ all_texts = "\n\n".join([f"Вариант {i + 1}:\n{text}" for i, text in enumerate(texts)])
108
+ st.download_button("Скачать все отзывы", all_texts, "generated_reviews.txt")
109
 
110
  if __name__ == "__main__":
111
  main()