lmnindzja commited on
Commit
94aa6e6
·
verified ·
1 Parent(s): a75fee2

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +81 -50
app.py CHANGED
@@ -11,6 +11,7 @@ st.write("Создавайте текстовые отзывы с помощью
11
  # ================ Загрузка модели ================
12
  @st.cache_resource
13
  def load_model():
 
14
  model_path = "./model"
15
  if not os.path.exists(model_path):
16
  st.error(f"Модель не найдена по пути {model_path}. Проверьте путь.")
@@ -26,86 +27,116 @@ def load_model():
26
 
27
  # ================ Рейтинг звёздочками ================
28
  def star_rating():
 
29
  stars = ["⭐", "⭐⭐", "⭐⭐⭐", "⭐⭐⭐⭐", "⭐⭐⭐⭐⭐"]
30
- return st.radio("Рейтинг:", range(1, 6), format_func=lambda x: stars[x - 1])
31
-
32
- # ================ Стиль генерации ================
33
- def get_style_params(style):
34
- styles = {
35
- "Строгий": {"temperature": 0.3, "top_p": 0.8, "top_k": 30},
36
- "Умеренный": {"temperature": 0.7, "top_p": 0.9, "top_k": 50},
37
- "Безумный": {"temperature": 1.2, "top_p": 1.0, "top_k": 100},
38
- }
39
- return styles.get(style, {"temperature": 0.7, "top_p": 0.9, "top_k": 50})
40
 
41
  # ================ Форматирование текста ================
42
  def format_text(text):
 
43
  sentences = text.split(". ")
44
  formatted = []
45
  for sentence in sentences:
46
  sentence = sentence.strip()
47
  if sentence:
48
- sentence = sentence.capitalize() if not sentence[0].isupper() else sentence
49
- sentence += "." if not sentence.endswith((".", "!", "?")) else ""
 
 
50
  formatted.append(sentence)
51
  return " ".join(formatted)
52
 
53
- # ================ Генерация текста ================
54
- def generate_text(model, tokenizer, prompt, params, max_length):
 
55
  input_ids = tokenizer.encode(prompt, return_tensors="pt")
56
- with torch.no_grad():
57
- outputs = model.generate(
58
- input_ids=input_ids,
59
- max_length=max_length,
60
- temperature=params["temperature"],
61
- top_p=params["top_p"],
62
- top_k=params["top_k"],
63
- do_sample=True,
64
- pad_token_id=tokenizer.eos_token_id
65
- )
66
- return tokenizer.decode(outputs[0], skip_special_tokens=True)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
67
 
68
  # ================ Основное приложение ================
69
  def main():
70
- model, tokenizer = load_model()
 
 
71
 
72
- # Инициализация состояния
73
- if "is_generating" not in st.session_state:
74
- st.session_state.is_generating = False # По умолчанию генерация выключена
75
 
76
  # Основные параметры
77
  category = st.selectbox("Категория:", ["Кафе", "Ресторан", "Парк", "Музей"])
78
  rating = star_rating()
79
- key_words = st.text_input("Ключевые слова", "вкусно, уютно, быстро")
80
 
81
- # Стиль генерации
82
- style = st.selectbox("Стиль генерации:", ["Строгий", "Умеренный", "Безумный"])
83
- style_params = get_style_params(style)
 
 
84
 
85
  # Дополнительные настройки
86
  with st.expander("Дополнительные настройки"):
87
  max_length = st.slider("Максимальная длина", 50, 300, 150)
 
 
 
 
 
 
88
 
89
- # Управление кнопкой
90
- button_label = "🔴 Остановить" if st.session_state.is_generating else "🟢 Сгенерировать"
91
-
92
- if st.button(button_label):
93
- st.session_state.is_generating = not st.session_state.is_generating # Переключение состояния
94
 
95
  # Генерация текста
96
- if st.session_state.is_generating:
 
 
 
 
 
 
97
  st.info("Генерация текстов...")
98
- try:
99
- prompt = f"Категория: {category}. Рейтинг: {rating}. Ключевые слова: {key_words}. Стиль: {style}."
100
- generated_text = generate_text(
101
- model, tokenizer, prompt, style_params, max_length
102
- )
103
- st.success("Генерация завершена!")
104
- st.write(format_text(generated_text))
105
- except Exception as e:
106
- st.error(f"Ошибка во время генерации: {e}")
107
- finally:
108
- st.session_state.is_generating = False # Возврат в исходное состояние
 
 
 
 
 
 
 
109
 
110
  if __name__ == "__main__":
111
  main()
 
11
  # ================ Загрузка модели ================
12
  @st.cache_resource
13
  def load_model():
14
+ """Загрузка модели и токенизатора (кэшируется для предотвращения перезагрузки)."""
15
  model_path = "./model"
16
  if not os.path.exists(model_path):
17
  st.error(f"Модель не найдена по пути {model_path}. Проверьте путь.")
 
27
 
28
  # ================ Рейтинг звёздочками ================
29
  def star_rating():
30
+ """Выбор рейтинга с отображением звёздочек."""
31
  stars = ["⭐", "⭐⭐", "⭐⭐⭐", "⭐⭐⭐⭐", "⭐⭐⭐⭐⭐"]
32
+ selected_star = st.radio("Рейтинг:", range(1, 6), format_func=lambda x: stars[x - 1], help="Выберите оценку от 1 до 5 звёзд.")
33
+ return selected_star
 
 
 
 
 
 
 
 
34
 
35
  # ================ Форматирование текста ================
36
  def format_text(text):
37
+ """Форматирует текст: заглавные буквы и точки."""
38
  sentences = text.split(". ")
39
  formatted = []
40
  for sentence in sentences:
41
  sentence = sentence.strip()
42
  if sentence:
43
+ if not sentence[0].isupper():
44
+ sentence = sentence.capitalize()
45
+ if not sentence.endswith((".", "!", "?")):
46
+ sentence += "."
47
  formatted.append(sentence)
48
  return " ".join(formatted)
49
 
50
+ # ================ Потоковая генерация текста ================
51
+ def stream_generate(model, tokenizer, prompt, params, max_length, batch_size):
52
+ """Генерирует текст блоками с обновлением контекста."""
53
  input_ids = tokenizer.encode(prompt, return_tensors="pt")
54
+ output_ids = input_ids
55
+
56
+ for _ in range(max_length // batch_size):
57
+ if st.session_state.get("stop_generation", False):
58
+ st.warning("Генерация остановлена!")
59
+ break
60
+
61
+ with torch.no_grad():
62
+ outputs = model.generate(
63
+ input_ids=output_ids,
64
+ max_new_tokens=batch_size,
65
+ temperature=params["temperature"],
66
+ top_p=params["top_p"],
67
+ top_k=params["top_k"],
68
+ do_sample=True,
69
+ pad_token_id=tokenizer.eos_token_id,
70
+ no_repeat_ngram_size=params["no_repeat_ngram_size"],
71
+ )
72
+ new_tokens = outputs[0, output_ids.shape[1]:]
73
+ output_ids = torch.cat([output_ids, new_tokens.unsqueeze(0)], dim=-1) # Обновляем контекст
74
+ new_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
75
+ yield new_text.strip()
76
+
77
+ if tokenizer.eos_token_id in new_tokens:
78
+ break
79
 
80
  # ================ Основное приложение ================
81
  def main():
82
+ # Инициализация состояния для кнопки остановки
83
+ if "stop_generation" not in st.session_state:
84
+ st.session_state.stop_generation = False
85
 
86
+ model, tokenizer = load_model()
 
 
87
 
88
  # Основные параметры
89
  category = st.selectbox("Категория:", ["Кафе", "Ресторан", "Парк", "Музей"])
90
  rating = star_rating()
91
+ key_words = st.text_input("Ключевые слова", "вкусно, уютно, быстро", help="Введите ключевые слова для использования в отзыве.")
92
 
93
+ # Стили генерации
94
+ style = st.selectbox("Стиль генерации:", ["Строгий", "Умеренный", "Безумный"],
95
+ help="Выберите стиль текста: Строгий – формальный и точный; Безумный – творческий и креативный.")
96
+ style_params = {"Строгий": (0.3, 30), "Умеренный": (0.7, 50), "Безумный": (1.5, 100)}
97
+ temperature, top_k = style_params[style]
98
 
99
  # Дополнительные настройки
100
  with st.expander("Дополнительные настройки"):
101
  max_length = st.slider("Максимальная длина", 50, 300, 150)
102
+ num_variants = st.number_input("Количество вариантов", 1, 5, 1)
103
+ batch_size = st.slider("Размер батча токенов", 1, 20, 5, help="Количество токенов за шаг.")
104
+ temperature = st.slider("Температура", 0.01, 2.0, temperature, 0.05)
105
+ top_p = st.slider("Top-p", 0.01, 1.0, 0.9, 0.05)
106
+ top_k = st.number_input("Top-k", 1, 100, top_k)
107
+ no_repeat_ngram_size = st.slider("Размер n-грамм для предотвращения повторов", 1, 10, 3)
108
 
109
+ # Кнопка для остановки генерации
110
+ if st.button("Остановить генерацию"):
111
+ st.session_state.stop_generation = True
 
 
112
 
113
  # Генерация текста
114
+ if st.button("Сгенерировать"):
115
+ st.session_state.stop_generation = False # Сбрасываем состояние остановки
116
+ if not key_words:
117
+ st.warning("Введите ключевые слова.")
118
+ return
119
+
120
+ input_prompt = f"Напиши отзыв для заведения. Категория: {category}. Рейтинг: {rating} из 5. Ключевые слова: {key_words}."
121
  st.info("Генерация текстов...")
122
+
123
+ placeholders = [st.empty() for _ in range(num_variants)]
124
+ texts = [""] * num_variants
125
+
126
+ for i in range(num_variants):
127
+ for chunk in stream_generate(
128
+ model, tokenizer, input_prompt,
129
+ {"temperature": temperature, "top_p": top_p, "top_k": top_k, "no_repeat_ngram_size": no_repeat_ngram_size},
130
+ max_length, batch_size
131
+ ):
132
+ texts[i] += chunk + " "
133
+ placeholders[i].write(f"**Вариант {i + 1}:**\n{format_text(texts[i])}")
134
+
135
+ st.success("Генерация завершена!")
136
+
137
+ # Скачивание всех вариантов
138
+ all_texts = "\n\n".join([f"Вариант {i + 1}:\n{format_text(text)}" for i, text in enumerate(texts)])
139
+ st.download_button("Скачать все отзывы", all_texts, "generated_reviews.txt")
140
 
141
  if __name__ == "__main__":
142
  main()