Spaces:
Sleeping
Sleeping
Update app.py
Browse files
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 |
-
|
| 94 |
-
|
| 95 |
|
| 96 |
-
for
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 100 |
|
| 101 |
st.success("Генерация завершена!")
|
| 102 |
-
|
|
|
|
|
|
|
|
|
|
| 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()
|