Spaces:
Sleeping
Sleeping
Update app.py
Browse files
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 |
-
|
| 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 |
-
|
| 49 |
-
|
|
|
|
|
|
|
| 50 |
formatted.append(sentence)
|
| 51 |
return " ".join(formatted)
|
| 52 |
|
| 53 |
-
# ================
|
| 54 |
-
def
|
|
|
|
| 55 |
input_ids = tokenizer.encode(prompt, return_tensors="pt")
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 67 |
|
| 68 |
# ================ Основное приложение ================
|
| 69 |
def main():
|
| 70 |
-
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 84 |
|
| 85 |
# Дополнительные настройки
|
| 86 |
with st.expander("Дополнительные настройки"):
|
| 87 |
max_length = st.slider("Максимальная длина", 50, 300, 150)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 88 |
|
| 89 |
-
#
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
if st.button(button_label):
|
| 93 |
-
st.session_state.is_generating = not st.session_state.is_generating # Переключение состояния
|
| 94 |
|
| 95 |
# Генерация текста
|
| 96 |
-
if st.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 97 |
st.info("Генерация текстов...")
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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()
|