Update app.py
Browse files
app.py
CHANGED
|
@@ -3,7 +3,6 @@ import torch.nn as nn
|
|
| 3 |
import torch.optim as optim
|
| 4 |
import random
|
| 5 |
import os
|
| 6 |
-
import re
|
| 7 |
import gradio as gr
|
| 8 |
from datetime import datetime
|
| 9 |
import math
|
|
@@ -12,12 +11,7 @@ import spaces # Обязательно для ZeroGPU
|
|
| 12 |
# ============ НАСТРОЙКИ ============
|
| 13 |
DATA_DIR = '/data'
|
| 14 |
os.makedirs(DATA_DIR, exist_ok=True)
|
| 15 |
-
MODEL_PATH = os.path.join(DATA_DIR, '
|
| 16 |
-
|
| 17 |
-
# На ZeroGPU мы не можем полагаться на глобальный DEVICE при инициализации,
|
| 18 |
-
# поэтому будем определять его внутри функций, обернутых в @spaces.GPU
|
| 19 |
-
# или использовать CPU для легких операций, если нужно.
|
| 20 |
-
# Но для модели нужен CUDA.
|
| 21 |
|
| 22 |
# ============ СЛОВАРЬ ============
|
| 23 |
WORDS = [
|
|
@@ -74,7 +68,7 @@ def pad_sequence(seq, max_len=MAX_LEN):
|
|
| 74 |
if len(seq) >= max_len: return seq[:max_len]
|
| 75 |
return seq + [PAD] * (max_len - len(seq))
|
| 76 |
|
| 77 |
-
# ============ МОДЕЛЬ ============
|
| 78 |
class PositionalEncoding(nn.Module):
|
| 79 |
def __init__(self, d_model, dropout=0.1, max_len=5000):
|
| 80 |
super(PositionalEncoding, self).__init__()
|
|
@@ -142,7 +136,7 @@ class AndreyTransformer(nn.Module):
|
|
| 142 |
)
|
| 143 |
return self.fc_out(output)
|
| 144 |
|
| 145 |
-
# ============ ДАННЫЕ ============
|
| 146 |
DIALOGUES = [
|
| 147 |
("привет", "привет как дела"), ("здравствуй", "здравствуй рад тебя видеть"),
|
| 148 |
("доброе утро", "доброе утро хорошего дня"), ("добрый день", "добрый день чем помочь"),
|
|
@@ -262,12 +256,11 @@ def prepare_data():
|
|
| 262 |
torch.tensor(Y_answers_input, dtype=torch.long),
|
| 263 |
torch.tensor(Y_answers_target, dtype=torch.long))
|
| 264 |
|
|
|
|
| 265 |
class AndreyAI:
|
| 266 |
def __init__(self, bin_file=MODEL_PATH):
|
| 267 |
self.bin_file = bin_file
|
| 268 |
self.memory = {'chat_history': [], 'epochs_trained': 0}
|
| 269 |
-
self.model = None
|
| 270 |
-
# Загружаем структуру, но веса загрузим позже
|
| 271 |
self.model = AndreyTransformer()
|
| 272 |
|
| 273 |
def load_weights(self):
|
|
@@ -290,14 +283,14 @@ class AndreyAI:
|
|
| 290 |
'num_encoder_layers': 2, 'num_decoder_layers': 2, 'dim_feedforward': 256,
|
| 291 |
'word_to_idx': word_to_idx,
|
| 292 |
'idx_to_word': {str(k): v for k, v in idx_to_word.items()},
|
| 293 |
-
'memory': self.memory, 'version': '8.
|
| 294 |
'created': datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
| 295 |
}
|
| 296 |
torch.save(state, self.bin_file)
|
| 297 |
print(f"✅ Сохранено в /data: {os.path.getsize(self.bin_file)/1024:.1f} КБ")
|
| 298 |
|
| 299 |
@spaces.GPU(duration=120)
|
| 300 |
-
def train(self, epochs=50):
|
| 301 |
print("🚀 Начало обучения на ZeroGPU...")
|
| 302 |
device = torch.device('cuda')
|
| 303 |
self.model.to(device)
|
|
@@ -335,7 +328,7 @@ class AndreyAI:
|
|
| 335 |
print(f"Эпоха {epoch}/{epochs} | Loss: {total_loss/n_batches:.4f}")
|
| 336 |
|
| 337 |
self.memory['epochs_trained'] = epochs
|
| 338 |
-
self.model.cpu()
|
| 339 |
self.save_weights()
|
| 340 |
print("✅ Обучение завершено!")
|
| 341 |
|
|
@@ -346,7 +339,7 @@ class AndreyAI:
|
|
| 346 |
if q in q_clean or q_clean in q: return a
|
| 347 |
return "Интересный вопрос! Я еще учусь."
|
| 348 |
|
| 349 |
-
@spaces.GPU(duration=10)
|
| 350 |
def generate(self, question, history=None, temperature=0.6, max_length=15):
|
| 351 |
device = torch.device('cuda')
|
| 352 |
self.model.to(device)
|
|
@@ -368,7 +361,6 @@ class AndreyAI:
|
|
| 368 |
|
| 369 |
if len(ctx_tokens) > MAX_LEN: ctx_tokens = ctx_tokens[-MAX_LEN:]
|
| 370 |
|
| 371 |
-
# Создаем тензор уже на устройстве
|
| 372 |
src = torch.tensor([pad_sequence(ctx_tokens, MAX_LEN)], dtype=torch.long).to(device)
|
| 373 |
generated_text = ""
|
| 374 |
|
|
@@ -394,22 +386,31 @@ class AndreyAI:
|
|
| 394 |
except Exception as e:
|
| 395 |
print(f"Ошибка генерации: {e}")
|
| 396 |
finally:
|
| 397 |
-
self.model.cpu()
|
| 398 |
|
| 399 |
return generated_text if generated_text else self.get_fallback_answer(q)
|
| 400 |
|
| 401 |
-
# ============ GRADIO ============
|
| 402 |
def gradio_chat(message, history):
|
| 403 |
-
if not message:
|
|
|
|
|
|
|
| 404 |
answer = andrey.generate(message, history=history)
|
| 405 |
new_history = history + [(message, answer)]
|
| 406 |
return "", new_history
|
| 407 |
|
| 408 |
with gr.Blocks(title="Андрей AI") as demo:
|
| 409 |
-
gr.Markdown("# 🤖 Андрей AI (ZeroGPU)\n### Transformer с памятью")
|
| 410 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 411 |
msg = gr.Textbox(label="Сообщение", placeholder="Напишите что-нибудь...")
|
| 412 |
-
clear = gr.Button("🧹 Очистить")
|
| 413 |
|
| 414 |
msg.submit(gradio_chat, [msg, chatbot], [msg, chatbot])
|
| 415 |
clear.click(lambda: [], None, chatbot)
|
|
@@ -419,6 +420,6 @@ if __name__ == "__main__":
|
|
| 419 |
loaded = andrey.load_weights()
|
| 420 |
|
| 421 |
if not loaded or andrey.memory.get('epochs_trained', 0) == 0:
|
| 422 |
-
andrey.train(150)
|
| 423 |
|
| 424 |
demo.launch()
|
|
|
|
| 3 |
import torch.optim as optim
|
| 4 |
import random
|
| 5 |
import os
|
|
|
|
| 6 |
import gradio as gr
|
| 7 |
from datetime import datetime
|
| 8 |
import math
|
|
|
|
| 11 |
# ============ НАСТРОЙКИ ============
|
| 12 |
DATA_DIR = '/data'
|
| 13 |
os.makedirs(DATA_DIR, exist_ok=True)
|
| 14 |
+
MODEL_PATH = os.path.join(DATA_DIR, 'andrey_zerogpu_final.bin')
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 15 |
|
| 16 |
# ============ СЛОВАРЬ ============
|
| 17 |
WORDS = [
|
|
|
|
| 68 |
if len(seq) >= max_len: return seq[:max_len]
|
| 69 |
return seq + [PAD] * (max_len - len(seq))
|
| 70 |
|
| 71 |
+
# ============ МОДЕЛЬ TRANSFORMER ============
|
| 72 |
class PositionalEncoding(nn.Module):
|
| 73 |
def __init__(self, d_model, dropout=0.1, max_len=5000):
|
| 74 |
super(PositionalEncoding, self).__init__()
|
|
|
|
| 136 |
)
|
| 137 |
return self.fc_out(output)
|
| 138 |
|
| 139 |
+
# ============ ДАННЫЕ ДЛЯ ОБУЧЕНИЯ ============
|
| 140 |
DIALOGUES = [
|
| 141 |
("привет", "привет как дела"), ("здравствуй", "здравствуй рад тебя видеть"),
|
| 142 |
("доброе утро", "доброе утро хорошего дня"), ("добрый день", "добрый день чем помочь"),
|
|
|
|
| 256 |
torch.tensor(Y_answers_input, dtype=torch.long),
|
| 257 |
torch.tensor(Y_answers_target, dtype=torch.long))
|
| 258 |
|
| 259 |
+
# ============ КЛАСС ANDREY AI ============
|
| 260 |
class AndreyAI:
|
| 261 |
def __init__(self, bin_file=MODEL_PATH):
|
| 262 |
self.bin_file = bin_file
|
| 263 |
self.memory = {'chat_history': [], 'epochs_trained': 0}
|
|
|
|
|
|
|
| 264 |
self.model = AndreyTransformer()
|
| 265 |
|
| 266 |
def load_weights(self):
|
|
|
|
| 283 |
'num_encoder_layers': 2, 'num_decoder_layers': 2, 'dim_feedforward': 256,
|
| 284 |
'word_to_idx': word_to_idx,
|
| 285 |
'idx_to_word': {str(k): v for k, v in idx_to_word.items()},
|
| 286 |
+
'memory': self.memory, 'version': '8.1-ZeroGPU-Fixed',
|
| 287 |
'created': datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
| 288 |
}
|
| 289 |
torch.save(state, self.bin_file)
|
| 290 |
print(f"✅ Сохранено в /data: {os.path.getsize(self.bin_file)/1024:.1f} КБ")
|
| 291 |
|
| 292 |
@spaces.GPU(duration=120)
|
| 293 |
+
def train(self, epochs=50):
|
| 294 |
print("🚀 Начало обучения на ZeroGPU...")
|
| 295 |
device = torch.device('cuda')
|
| 296 |
self.model.to(device)
|
|
|
|
| 328 |
print(f"Эпоха {epoch}/{epochs} | Loss: {total_loss/n_batches:.4f}")
|
| 329 |
|
| 330 |
self.memory['epochs_trained'] = epochs
|
| 331 |
+
self.model.cpu()
|
| 332 |
self.save_weights()
|
| 333 |
print("✅ Обучение завершено!")
|
| 334 |
|
|
|
|
| 339 |
if q in q_clean or q_clean in q: return a
|
| 340 |
return "Интересный вопрос! Я еще учусь."
|
| 341 |
|
| 342 |
+
@spaces.GPU(duration=10)
|
| 343 |
def generate(self, question, history=None, temperature=0.6, max_length=15):
|
| 344 |
device = torch.device('cuda')
|
| 345 |
self.model.to(device)
|
|
|
|
| 361 |
|
| 362 |
if len(ctx_tokens) > MAX_LEN: ctx_tokens = ctx_tokens[-MAX_LEN:]
|
| 363 |
|
|
|
|
| 364 |
src = torch.tensor([pad_sequence(ctx_tokens, MAX_LEN)], dtype=torch.long).to(device)
|
| 365 |
generated_text = ""
|
| 366 |
|
|
|
|
| 386 |
except Exception as e:
|
| 387 |
print(f"Ошибка генерации: {e}")
|
| 388 |
finally:
|
| 389 |
+
self.model.cpu()
|
| 390 |
|
| 391 |
return generated_text if generated_text else self.get_fallback_answer(q)
|
| 392 |
|
| 393 |
+
# ============ GRADIO ИНТЕРФЕЙС ============
|
| 394 |
def gradio_chat(message, history):
|
| 395 |
+
if not message:
|
| 396 |
+
return "", history
|
| 397 |
+
|
| 398 |
answer = andrey.generate(message, history=history)
|
| 399 |
new_history = history + [(message, answer)]
|
| 400 |
return "", new_history
|
| 401 |
|
| 402 |
with gr.Blocks(title="Андрей AI") as demo:
|
| 403 |
+
gr.Markdown("# 🤖 Андрей AI (ZeroGPU)\n### Transformer с реальной памятью")
|
| 404 |
+
|
| 405 |
+
# ВАЖНО: type="tuples" исправляет ошибку формата
|
| 406 |
+
chatbot = gr.Chatbot(
|
| 407 |
+
height=400,
|
| 408 |
+
label="Диалог",
|
| 409 |
+
type="tuples"
|
| 410 |
+
)
|
| 411 |
+
|
| 412 |
msg = gr.Textbox(label="Сообщение", placeholder="Напишите что-нибудь...")
|
| 413 |
+
clear = gr.Button("🧹 Очистить историю")
|
| 414 |
|
| 415 |
msg.submit(gradio_chat, [msg, chatbot], [msg, chatbot])
|
| 416 |
clear.click(lambda: [], None, chatbot)
|
|
|
|
| 420 |
loaded = andrey.load_weights()
|
| 421 |
|
| 422 |
if not loaded or andrey.memory.get('epochs_trained', 0) == 0:
|
| 423 |
+
andrey.train(150)
|
| 424 |
|
| 425 |
demo.launch()
|