root39058 commited on
Commit
fa0920b
·
verified ·
1 Parent(s): 02eeb75

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +68 -233
app.py CHANGED
@@ -7,7 +7,7 @@ import re
7
  import gradio as gr
8
  from datetime import datetime
9
  import math
10
- import numpy as np
11
 
12
  # ============ НАСТРОЙКИ ПУТЕЙ ============
13
  DATA_DIR = '/data'
@@ -39,7 +39,6 @@ WORDS = [
39
  'работа', 'отдых', 'путешествие', 'еда', 'вода', 'спорт', 'люди', 'мир', 'знание',
40
  'будущее', 'прошлое', 'настоящее', 'интерес', 'радость', 'успех', 'дружба', 'любовь',
41
  'семья', 'здоровье', 'счастье', 'удача', 'смех', 'солнце', 'звезды', 'мечта',
42
- # Добавляем маркеры ролей для памяти
43
  'пользователь', 'андрей', 'говорит'
44
  ]
45
 
@@ -53,7 +52,7 @@ vocab_size = len(WORDS) + 3
53
  PAD = 0
54
  UNK = 1
55
  START = 2
56
- MAX_LEN = 30 # Увеличиваем длину, чтобы влезала история
57
 
58
  def tokenize(text):
59
  return [word_to_idx.get(w, UNK) for w in text.lower().split()]
@@ -61,20 +60,15 @@ def tokenize(text):
61
  def detokenize(tokens):
62
  words = []
63
  for t in tokens:
64
- if t == START:
65
- continue
66
- if t == PAD:
67
- break
68
- if t == UNK:
69
- continue
70
  w = idx_to_word.get(t)
71
- if w:
72
- words.append(w)
73
  return ' '.join(words)
74
 
75
  def pad_sequence(seq, max_len=MAX_LEN):
76
- if len(seq) >= max_len:
77
- return seq[:max_len]
78
  return seq + [PAD] * (max_len - len(seq))
79
 
80
  # ============ TRANSFORMER АРХИТЕКТУРА ============
@@ -83,11 +77,9 @@ class PositionalEncoding(nn.Module):
83
  def __init__(self, d_model, dropout=0.1, max_len=5000):
84
  super(PositionalEncoding, self).__init__()
85
  self.dropout = nn.Dropout(p=dropout)
86
-
87
  pe = torch.zeros(max_len, d_model)
88
  position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
89
  div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
90
-
91
  pe[:, 0::2] = torch.sin(position * div_term)
92
  pe[:, 1::2] = torch.cos(position * div_term)
93
  pe = pe.unsqueeze(0)
@@ -102,24 +94,19 @@ class AndreyTransformer(nn.Module):
102
  num_encoder_layers=2, num_decoder_layers=2,
103
  dim_feedforward=256, dropout=0.1, max_len=MAX_LEN):
104
  super().__init__()
105
-
106
  self.d_model = d_model
107
  self.vocab_size = vocab_size
108
-
109
  self.embedding = nn.Embedding(vocab_size, d_model)
110
  self.pos_encoder = PositionalEncoding(d_model, dropout, max_len)
111
  self.pos_decoder = PositionalEncoding(d_model, dropout, max_len)
112
 
113
  self.transformer = nn.Transformer(
114
- d_model=d_model,
115
- nhead=nhead,
116
  num_encoder_layers=num_encoder_layers,
117
  num_decoder_layers=num_decoder_layers,
118
  dim_feedforward=dim_feedforward,
119
- dropout=dropout,
120
- batch_first=True
121
  )
122
-
123
  self.fc_out = nn.Linear(d_model, vocab_size)
124
  self._init_weights()
125
 
@@ -130,8 +117,7 @@ class AndreyTransformer(nn.Module):
130
  self.fc_out.weight.data.uniform_(-initrange, initrange)
131
 
132
  def generate_mask(self, tgt_len):
133
- mask = torch.triu(torch.ones(tgt_len, tgt_len), diagonal=1).bool()
134
- return mask.to(DEVICE)
135
 
136
  def create_pad_mask(self, seq, pad_idx=PAD):
137
  return (seq == pad_idx).to(DEVICE)
@@ -140,36 +126,27 @@ class AndreyTransformer(nn.Module):
140
  src_key_padding_mask=None, tgt_key_padding_mask=None):
141
  src_emb = self.pos_encoder(self.embedding(src) * math.sqrt(self.d_model))
142
  tgt_emb = self.pos_decoder(self.embedding(tgt) * math.sqrt(self.d_model))
143
-
144
  output = self.transformer(
145
- src_emb, tgt_emb,
146
- src_mask=src_mask,
147
- tgt_mask=tgt_mask,
148
  src_key_padding_mask=src_key_padding_mask,
149
  tgt_key_padding_mask=tgt_key_padding_mask
150
  )
151
-
152
- output = self.fc_out(output)
153
- return output
154
 
155
  def encode(self, src):
156
  src_emb = self.pos_encoder(self.embedding(src) * math.sqrt(self.d_model))
157
  src_key_padding_mask = self.create_pad_mask(src)
158
- memory = self.transformer.encoder(src_emb, src_key_padding_mask=src_key_padding_mask)
159
- return memory
160
 
161
  def decode_step(self, tgt, memory, tgt_mask=None, tgt_key_padding_mask=None):
162
  tgt_emb = self.pos_decoder(self.embedding(tgt) * math.sqrt(self.d_model))
163
-
164
  output = self.transformer.decoder(
165
- tgt_emb, memory,
166
- tgt_mask=tgt_mask,
167
  tgt_key_padding_mask=tgt_key_padding_mask
168
  )
169
-
170
  return self.fc_out(output)
171
 
172
- # ============ ДИАЛОГИ ДЛЯ ОБУЧЕНИЯ ============
173
  DIALOGUES = [
174
  ("привет", "привет как дела"), ("здравствуй", "здравствуй рад тебя видеть"),
175
  ("доброе утро", "доброе утро хорошего дня"), ("добрый день", "добрый день чем помочь"),
@@ -275,60 +252,35 @@ DIALOGUES = [
275
 
276
  FALLBACK_DICT = {q.lower(): a for q, a in DIALOGUES}
277
 
278
- # ============ ПОДГОТОВКА ДАННЫХ ============
279
  def prepare_data():
280
- X_questions = []
281
- Y_answers_input = []
282
- Y_answers_target = []
283
-
284
  for q, a in DIALOGUES:
285
- q_toks = tokenize(q)
286
- a_toks = tokenize(a)
287
-
288
  if q_toks and a_toks:
289
- q_padded = pad_sequence(q_toks, MAX_LEN)
290
-
291
- a_input = [START] + a_toks
292
- a_input_padded = pad_sequence(a_input, MAX_LEN)
293
-
294
- a_target = a_toks + [PAD]
295
- a_target_padded = pad_sequence(a_target, MAX_LEN)
296
-
297
- X_questions.append(q_padded)
298
- Y_answers_input.append(a_input_padded)
299
- Y_answers_target.append(a_target_padded)
300
-
301
- X_tensor = torch.tensor(X_questions, dtype=torch.long)
302
- Y_input_tensor = torch.tensor(Y_answers_input, dtype=torch.long)
303
- Y_target_tensor = torch.tensor(Y_answers_target, dtype=torch.long)
304
-
305
- print(f"📚 Всего примеров: {len(X_tensor)}")
306
- return X_tensor, Y_input_tensor, Y_target_tensor
307
 
308
- # ============ АНДРЕЙ TRANSFORMER ============
309
  class AndreyAI:
310
  def __init__(self, bin_file=MODEL_PATH):
311
  self.bin_file = bin_file
312
- self.memory = {
313
- 'chat_history': [], # Здесь хранится реальная история
314
- 'epochs_trained': 0
315
- }
316
  self.model = None
317
  self.load()
318
 
319
  def _get_state(self):
320
  return {
321
  'model_state': self.model.state_dict() if self.model else None,
322
- 'vocab_size': vocab_size,
323
- 'd_model': 128,
324
- 'nhead': 4,
325
- 'num_encoder_layers': 2,
326
- 'num_decoder_layers': 2,
327
- 'dim_feedforward': 256,
328
  'word_to_idx': word_to_idx,
329
  'idx_to_word': {str(k): v for k, v in idx_to_word.items()},
330
- 'memory': self.memory,
331
- 'version': '6.0-Memory-Transformer',
332
  'created': datetime.now().strftime("%Y-%m-%d %H:%M:%S")
333
  }
334
 
@@ -337,84 +289,59 @@ class AndreyAI:
337
  if data.get('model_state'):
338
  self.model = AndreyTransformer(
339
  vocab_size=data.get('vocab_size', vocab_size),
340
- d_model=data.get('d_model', 128),
341
- nhead=data.get('nhead', 4),
342
  num_encoder_layers=data.get('num_encoder_layers', 2),
343
  num_decoder_layers=data.get('num_decoder_layers', 2),
344
  dim_feedforward=data.get('dim_feedforward', 256)
345
  )
346
  self.model.load_state_dict(data['model_state'])
347
- self.model.to(DEVICE)
348
- self.model.eval()
349
  return True
350
  return False
351
 
352
  def load(self):
353
  if os.path.exists(self.bin_file):
354
  try:
355
- data = torch.load(self.bin_file, map_location=DEVICE)
356
- if self._restore(data):
357
  print(f"✅ Андрей загружен из {self.bin_file}")
358
- print(f"🧠 Обучен: {self.memory.get('epochs_trained', 0)} эпох")
359
- print(f"💾 История чатов: {len(self.memory.get('chat_history', []))} сообщений")
360
  return True
361
- except Exception as e:
362
- print(f"⚠️ Ошибка загрузки: {e}")
363
-
364
  print("📝 Создаю нового Андрея...")
365
- self.model = AndreyTransformer().to(DEVICE)
366
- self.model.eval()
367
  return False
368
 
369
  def save(self):
370
  os.makedirs(os.path.dirname(self.bin_file), exist_ok=True)
371
  torch.save(self._get_state(), self.bin_file)
372
- size = os.path.getsize(self.bin_file) / 1024
373
- print(f"✅ Сохранён в /data: {size:.1f} КБ")
374
 
 
375
  def train(self, epochs=150):
376
  print("="*60)
377
  print(f"🧠 ОБУЧЕНИЕ АНДРЕЯ — {epochs} ЭПОХ")
378
- print(f"📂 Модель: {self.bin_file}")
379
- print("="*60 + "\n")
380
-
381
  X_data, Y_input, Y_target = prepare_data()
382
-
383
  criterion = nn.CrossEntropyLoss(ignore_index=PAD)
384
- optimizer = optim.Adam(self.model.parameters(), lr=0.0005)
385
- scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=50, gamma=0.5)
386
-
387
  self.model.train()
388
 
389
  for epoch in range(epochs):
390
- total_loss = 0
391
- n_batches = 0
392
-
393
  indices = list(range(len(X_data)))
394
  random.shuffle(indices)
395
 
396
  for idx in indices:
397
  src = X_data[idx].unsqueeze(0).to(DEVICE)
398
- tgt_input = Y_input[idx].unsqueeze(0).to(DEVICE)
399
- tgt_target = Y_target[idx].unsqueeze(0).to(DEVICE)
400
 
401
- tgt_len = tgt_input.size(1)
402
- tgt_mask = self.model.generate_mask(tgt_len)
403
-
404
- src_key_padding_mask = self.model.create_pad_mask(src)
405
- tgt_key_padding_mask = self.model.create_pad_mask(tgt_input)
406
 
407
  optimizer.zero_grad()
408
-
409
- output = self.model(
410
- src, tgt_input,
411
- tgt_mask=tgt_mask,
412
- src_key_padding_mask=src_key_padding_mask,
413
- tgt_key_padding_mask=tgt_key_padding_mask
414
- )
415
-
416
- loss = criterion(output.view(-1, vocab_size), tgt_target.view(-1))
417
-
418
  loss.backward()
419
  torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
420
  optimizer.step()
@@ -422,175 +349,83 @@ class AndreyAI:
422
  total_loss += loss.item()
423
  n_batches += 1
424
 
425
- scheduler.step()
426
-
427
- if epoch % 10 == 0:
428
- avg_loss = total_loss / n_batches
429
- print(f"Эпоха {epoch:3d}/{epochs} | Потери: {avg_loss:.4f}")
430
 
431
  self.memory['epochs_trained'] = epochs
432
  print("\n✅ Обучение готово!")
433
  self.save()
434
  self.model.eval()
435
-
436
  def get_fallback_answer(self, question):
437
  q_clean = question.lower().strip()
438
- if q_clean in FALLBACK_DICT:
439
- return FALLBACK_DICT[q_clean]
440
  for q, a in DIALOGUES:
441
- if q in q_clean or q_clean in q:
442
- return a
443
  return "Интересный вопрос! Я еще учусь."
444
 
445
  def generate(self, question, history=None, temperature=0.6, max_length=15):
446
- """
447
- Генерирует ответ с учетом истории (Real Memory).
448
- history: список кортежей [(user_msg, bot_msg), ...]
449
- """
450
  q = question.lower().strip()
451
-
452
- # 1. Формируем контекст из истории
453
  context_parts = []
454
  if history:
455
- # Берем последние 3 пары сообщений, чтобы не превысить MAX_LEN
456
- recent_history = history[-3:]
457
- for user_msg, bot_msg in recent_history:
458
  context_parts.append(f"пользователь говорит {user_msg}")
459
  context_parts.append(f"андрей говорит {bot_msg}")
460
-
461
- # Добавляем текущий вопрос
462
  context_parts.append(f"пользователь говорит {question}")
463
-
464
- # Склеиваем в одну строку
465
  full_context = " ".join(context_parts)
466
 
467
- # 2. Токенизация контекста
468
  ctx_tokens = tokenize(full_context)
469
- if not ctx_tokens:
470
- return self.get_fallback_answer(q)
471
-
472
- # Обрезаем до MAX_LEN, оставляя конец (самое важное - текущий вопрос)
473
- if len(ctx_tokens) > MAX_LEN:
474
- ctx_tokens = ctx_tokens[-MAX_LEN:]
475
-
476
- ctx_padded = pad_sequence(ctx_tokens, MAX_LEN)
477
- src = torch.tensor([ctx_padded], dtype=torch.long).to(DEVICE)
478
 
 
479
  generated_text = ""
480
 
481
  try:
482
  with torch.no_grad():
483
  memory = self.model.encode(src)
484
-
485
  response_tokens = []
486
  decoder_input = torch.tensor([[START]], dtype=torch.long).to(DEVICE)
487
 
488
  for i in range(max_length):
489
- tgt_len = decoder_input.size(1)
490
- tgt_mask = self.model.generate_mask(tgt_len).to(DEVICE)
491
-
492
  output = self.model.decode_step(decoder_input, memory, tgt_mask=tgt_mask)
493
-
494
  logits = output[:, -1, :] / temperature
495
  probs = torch.softmax(logits, dim=-1)
496
-
497
- top_prob, next_token = torch.max(probs, dim=-1)
498
  next_token = next_token.item()
499
- confidence = top_prob.item()
500
-
501
- if next_token == PAD or next_token == UNK or confidence < 0.05:
502
- break
503
 
 
504
  response_tokens.append(next_token)
505
- next_token_tensor = torch.tensor([[next_token]], dtype=torch.long).to(DEVICE)
506
- decoder_input = torch.cat([decoder_input, next_token_tensor], dim=1)
507
 
508
  generated_text = detokenize(response_tokens)
509
-
510
- except Exception as e:
511
- print(f"Ошибка генерации: {e}")
512
- generated_text = ""
513
 
514
- # 3. Fallback если модель промолчала
515
- if not generated_text or len(generated_text.split()) < 1:
516
- return self.get_fallback_answer(q)
517
-
518
- return generated_text
519
-
520
- def chat(self):
521
- print("="*60)
522
- print("🤖 АНДРЕЙ v6.0 (Transformer + Real Memory)")
523
- print("💾 Память сохраняется в /data")
524
- print("="*60 + "\n")
525
-
526
- history = self.memory.get('chat_history', [])
527
-
528
- while True:
529
- user = input("👤 Вы: ").strip()
530
-
531
- if user.lower() in ['пока', 'выход', 'exit']:
532
- print("🤖 Андрей: Пока! 👋")
533
- self.save()
534
- break
535
-
536
- if not user:
537
- continue
538
-
539
- answer = self.generate(user, history=history)
540
- print(f"🤖 Андрей: {answer}\n")
541
-
542
- # Обновляем историю
543
- history.append((user, answer))
544
- self.memory['chat_history'] = history
545
 
546
  # ============ GRADIO ИНТЕРФЕЙС ============
547
  def gradio_chat(question, history):
548
- if not question:
549
- return "", history
550
-
551
- # Передаем текущую историю в модель
552
  answer = andrey.generate(question, history=history)
553
-
554
- # Обновляем историю
555
  new_history = history + [(question, answer)]
556
-
557
- # Сохраняем обновленную историю в память объекта
558
  andrey.memory['chat_history'] = new_history
559
-
560
  return "", new_history
561
 
562
  def launch_gradio():
563
- with gr.Blocks(title="Андрей AI", theme=gr.themes.Soft()) as demo:
564
- gr.Markdown("""
565
- # 🤖 Андрей AI v6.0
566
- ### Transformer с реальной памятью
567
- **Память:** Хранится в `/data`
568
- **Контекст:** Помнит последние сообщения
569
- """)
570
-
571
- chatbot = gr.Chatbot(height=400, label="Диалог с Андреем")
572
- msg = gr.Textbox(label="Ваше сообщение", placeholder="Напишите что-нибудь...")
573
- clear = gr.Button("🧹 Очистить историю")
574
 
575
  msg.submit(gradio_chat, [msg, chatbot], [msg, chatbot])
576
  clear.click(lambda: [], None, chatbot)
577
-
578
- gr.Markdown("""
579
- ### ❓ Примеры:
580
- - Привет
581
- - Меня зовут Евгений
582
- - Как меня зовут? (должен вспомнить)
583
- - 2+2
584
- """)
585
 
586
- demo.launch(share=True)
587
 
588
- # ============ ЗАПУСК ============
589
  if __name__ == "__main__":
590
  andrey = AndreyAI(MODEL_PATH)
591
-
592
  if andrey.model is None or andrey.memory.get('epochs_trained', 0) == 0:
593
  andrey.train(150)
594
-
595
- print("\n🚀 Запуск Gradio интерфейса...")
596
  launch_gradio()
 
7
  import gradio as gr
8
  from datetime import datetime
9
  import math
10
+ import spaces # <--- ВАЖНО: Импортируем библиотеку spaces
11
 
12
  # ============ НАСТРОЙКИ ПУТЕЙ ============
13
  DATA_DIR = '/data'
 
39
  'работа', 'отдых', 'путешествие', 'еда', 'вода', 'спорт', 'люди', 'мир', 'знание',
40
  'будущее', 'прошлое', 'настоящее', 'интерес', 'радость', 'успех', 'дружба', 'любовь',
41
  'семья', 'здоровье', 'счастье', 'удача', 'смех', 'солнце', 'звезды', 'мечта',
 
42
  'пользователь', 'андрей', 'говорит'
43
  ]
44
 
 
52
  PAD = 0
53
  UNK = 1
54
  START = 2
55
+ MAX_LEN = 30
56
 
57
  def tokenize(text):
58
  return [word_to_idx.get(w, UNK) for w in text.lower().split()]
 
60
  def detokenize(tokens):
61
  words = []
62
  for t in tokens:
63
+ if t == START: continue
64
+ if t == PAD: break
65
+ if t == UNK: continue
 
 
 
66
  w = idx_to_word.get(t)
67
+ if w: words.append(w)
 
68
  return ' '.join(words)
69
 
70
  def pad_sequence(seq, max_len=MAX_LEN):
71
+ if len(seq) >= max_len: return seq[:max_len]
 
72
  return seq + [PAD] * (max_len - len(seq))
73
 
74
  # ============ TRANSFORMER АРХИТЕКТУРА ============
 
77
  def __init__(self, d_model, dropout=0.1, max_len=5000):
78
  super(PositionalEncoding, self).__init__()
79
  self.dropout = nn.Dropout(p=dropout)
 
80
  pe = torch.zeros(max_len, d_model)
81
  position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
82
  div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
 
83
  pe[:, 0::2] = torch.sin(position * div_term)
84
  pe[:, 1::2] = torch.cos(position * div_term)
85
  pe = pe.unsqueeze(0)
 
94
  num_encoder_layers=2, num_decoder_layers=2,
95
  dim_feedforward=256, dropout=0.1, max_len=MAX_LEN):
96
  super().__init__()
 
97
  self.d_model = d_model
98
  self.vocab_size = vocab_size
 
99
  self.embedding = nn.Embedding(vocab_size, d_model)
100
  self.pos_encoder = PositionalEncoding(d_model, dropout, max_len)
101
  self.pos_decoder = PositionalEncoding(d_model, dropout, max_len)
102
 
103
  self.transformer = nn.Transformer(
104
+ d_model=d_model, nhead=nhead,
 
105
  num_encoder_layers=num_encoder_layers,
106
  num_decoder_layers=num_decoder_layers,
107
  dim_feedforward=dim_feedforward,
108
+ dropout=dropout, batch_first=True
 
109
  )
 
110
  self.fc_out = nn.Linear(d_model, vocab_size)
111
  self._init_weights()
112
 
 
117
  self.fc_out.weight.data.uniform_(-initrange, initrange)
118
 
119
  def generate_mask(self, tgt_len):
120
+ return torch.triu(torch.ones(tgt_len, tgt_len), diagonal=1).bool().to(DEVICE)
 
121
 
122
  def create_pad_mask(self, seq, pad_idx=PAD):
123
  return (seq == pad_idx).to(DEVICE)
 
126
  src_key_padding_mask=None, tgt_key_padding_mask=None):
127
  src_emb = self.pos_encoder(self.embedding(src) * math.sqrt(self.d_model))
128
  tgt_emb = self.pos_decoder(self.embedding(tgt) * math.sqrt(self.d_model))
 
129
  output = self.transformer(
130
+ src_emb, tgt_emb, src_mask=src_mask, tgt_mask=tgt_mask,
 
 
131
  src_key_padding_mask=src_key_padding_mask,
132
  tgt_key_padding_mask=tgt_key_padding_mask
133
  )
134
+ return self.fc_out(output)
 
 
135
 
136
  def encode(self, src):
137
  src_emb = self.pos_encoder(self.embedding(src) * math.sqrt(self.d_model))
138
  src_key_padding_mask = self.create_pad_mask(src)
139
+ return self.transformer.encoder(src_emb, src_key_padding_mask=src_key_padding_mask)
 
140
 
141
  def decode_step(self, tgt, memory, tgt_mask=None, tgt_key_padding_mask=None):
142
  tgt_emb = self.pos_decoder(self.embedding(tgt) * math.sqrt(self.d_model))
 
143
  output = self.transformer.decoder(
144
+ tgt_emb, memory, tgt_mask=tgt_mask,
 
145
  tgt_key_padding_mask=tgt_key_padding_mask
146
  )
 
147
  return self.fc_out(output)
148
 
149
+ # ============ ДИАЛОГИ ============
150
  DIALOGUES = [
151
  ("привет", "привет как дела"), ("здравствуй", "здравствуй рад тебя видеть"),
152
  ("доброе утро", "доброе утро хорошего дня"), ("добрый день", "добрый день чем помочь"),
 
252
 
253
  FALLBACK_DICT = {q.lower(): a for q, a in DIALOGUES}
254
 
 
255
  def prepare_data():
256
+ X_questions, Y_answers_input, Y_answers_target = [], [], []
 
 
 
257
  for q, a in DIALOGUES:
258
+ q_toks, a_toks = tokenize(q), tokenize(a)
 
 
259
  if q_toks and a_toks:
260
+ X_questions.append(pad_sequence(q_toks, MAX_LEN))
261
+ a_input = pad_sequence([START] + a_toks, MAX_LEN)
262
+ a_target = pad_sequence(a_toks + [PAD], MAX_LEN)
263
+ Y_answers_input.append(a_input)
264
+ Y_answers_target.append(a_target)
265
+ return (torch.tensor(X_questions, dtype=torch.long),
266
+ torch.tensor(Y_answers_input, dtype=torch.long),
267
+ torch.tensor(Y_answers_target, dtype=torch.long))
 
 
 
 
 
 
 
 
 
 
268
 
 
269
  class AndreyAI:
270
  def __init__(self, bin_file=MODEL_PATH):
271
  self.bin_file = bin_file
272
+ self.memory = {'chat_history': [], 'epochs_trained': 0}
 
 
 
273
  self.model = None
274
  self.load()
275
 
276
  def _get_state(self):
277
  return {
278
  'model_state': self.model.state_dict() if self.model else None,
279
+ 'vocab_size': vocab_size, 'd_model': 128, 'nhead': 4,
280
+ 'num_encoder_layers': 2, 'num_decoder_layers': 2, 'dim_feedforward': 256,
 
 
 
 
281
  'word_to_idx': word_to_idx,
282
  'idx_to_word': {str(k): v for k, v in idx_to_word.items()},
283
+ 'memory': self.memory, 'version': '7.0-ZeroGPU',
 
284
  'created': datetime.now().strftime("%Y-%m-%d %H:%M:%S")
285
  }
286
 
 
289
  if data.get('model_state'):
290
  self.model = AndreyTransformer(
291
  vocab_size=data.get('vocab_size', vocab_size),
292
+ d_model=data.get('d_model', 128), nhead=data.get('nhead', 4),
 
293
  num_encoder_layers=data.get('num_encoder_layers', 2),
294
  num_decoder_layers=data.get('num_decoder_layers', 2),
295
  dim_feedforward=data.get('dim_feedforward', 256)
296
  )
297
  self.model.load_state_dict(data['model_state'])
298
+ self.model.to(DEVICE).eval()
 
299
  return True
300
  return False
301
 
302
  def load(self):
303
  if os.path.exists(self.bin_file):
304
  try:
305
+ if self._restore(torch.load(self.bin_file, map_location=DEVICE)):
 
306
  print(f"✅ Андрей загружен из {self.bin_file}")
 
 
307
  return True
308
+ except Exception as e: print(f"⚠️ Ошибка загрузки: {e}")
 
 
309
  print("📝 Создаю нового Андрея...")
310
+ self.model = AndreyTransformer().to(DEVICE).eval()
 
311
  return False
312
 
313
  def save(self):
314
  os.makedirs(os.path.dirname(self.bin_file), exist_ok=True)
315
  torch.save(self._get_state(), self.bin_file)
316
+ print(f"✅ Сохранён в /data: {os.path.getsize(self.bin_file)/1024:.1f} КБ")
 
317
 
318
+ @spaces.GPU(duration=120) # <--- ВАЖНО: Декоратор для ZeroGPU
319
  def train(self, epochs=150):
320
  print("="*60)
321
  print(f"🧠 ОБУЧЕНИЕ АНДРЕЯ — {epochs} ЭПОХ")
322
+ print("="*60)
 
 
323
  X_data, Y_input, Y_target = prepare_data()
 
324
  criterion = nn.CrossEntropyLoss(ignore_index=PAD)
325
+ optimizer = optim.Adam(self.model.parameters(), lr=0.001)
 
 
326
  self.model.train()
327
 
328
  for epoch in range(epochs):
329
+ total_loss, n_batches = 0, 0
 
 
330
  indices = list(range(len(X_data)))
331
  random.shuffle(indices)
332
 
333
  for idx in indices:
334
  src = X_data[idx].unsqueeze(0).to(DEVICE)
335
+ tgt_in = Y_input[idx].unsqueeze(0).to(DEVICE)
336
+ tgt_tar = Y_target[idx].unsqueeze(0).to(DEVICE)
337
 
338
+ tgt_mask = self.model.generate_mask(tgt_in.size(1))
339
+ src_pad = self.model.create_pad_mask(src)
340
+ tgt_pad = self.model.create_pad_mask(tgt_in)
 
 
341
 
342
  optimizer.zero_grad()
343
+ output = self.model(src, tgt_in, tgt_mask=tgt_mask, src_key_padding_mask=src_pad, tgt_key_padding_mask=tgt_pad)
344
+ loss = criterion(output.view(-1, vocab_size), tgt_tar.view(-1))
 
 
 
 
 
 
 
 
345
  loss.backward()
346
  torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
347
  optimizer.step()
 
349
  total_loss += loss.item()
350
  n_batches += 1
351
 
352
+ if epoch % 10 == 0: print(f"Эпоха {epoch:3d}/{epochs} | Потери: {total_loss/n_batches:.4f}")
 
 
 
 
353
 
354
  self.memory['epochs_trained'] = epochs
355
  print("\n✅ Обучение готово!")
356
  self.save()
357
  self.model.eval()
358
+
359
  def get_fallback_answer(self, question):
360
  q_clean = question.lower().strip()
361
+ if q_clean in FALLBACK_DICT: return FALLBACK_DICT[q_clean]
 
362
  for q, a in DIALOGUES:
363
+ if q in q_clean or q_clean in q: return a
 
364
  return "Интересный вопрос! Я еще учусь."
365
 
366
  def generate(self, question, history=None, temperature=0.6, max_length=15):
 
 
 
 
367
  q = question.lower().strip()
 
 
368
  context_parts = []
369
  if history:
370
+ for user_msg, bot_msg in history[-3:]:
 
 
371
  context_parts.append(f"пользователь говорит {user_msg}")
372
  context_parts.append(f"андрей говорит {bot_msg}")
 
 
373
  context_parts.append(f"пользователь говорит {question}")
 
 
374
  full_context = " ".join(context_parts)
375
 
 
376
  ctx_tokens = tokenize(full_context)
377
+ if not ctx_tokens: return self.get_fallback_answer(q)
378
+ if len(ctx_tokens) > MAX_LEN: ctx_tokens = ctx_tokens[-MAX_LEN:]
 
 
 
 
 
 
 
379
 
380
+ src = torch.tensor([pad_sequence(ctx_tokens, MAX_LEN)], dtype=torch.long).to(DEVICE)
381
  generated_text = ""
382
 
383
  try:
384
  with torch.no_grad():
385
  memory = self.model.encode(src)
 
386
  response_tokens = []
387
  decoder_input = torch.tensor([[START]], dtype=torch.long).to(DEVICE)
388
 
389
  for i in range(max_length):
390
+ tgt_mask = self.model.generate_mask(decoder_input.size(1)).to(DEVICE)
 
 
391
  output = self.model.decode_step(decoder_input, memory, tgt_mask=tgt_mask)
 
392
  logits = output[:, -1, :] / temperature
393
  probs = torch.softmax(logits, dim=-1)
394
+ _, next_token = torch.max(probs, dim=-1)
 
395
  next_token = next_token.item()
 
 
 
 
396
 
397
+ if next_token in [PAD, UNK]: break
398
  response_tokens.append(next_token)
399
+ decoder_input = torch.cat([decoder_input, torch.tensor([[next_token]], dtype=torch.long).to(DEVICE)], dim=1)
 
400
 
401
  generated_text = detokenize(response_tokens)
402
+ except Exception: pass
 
 
 
403
 
404
+ return generated_text if generated_text else self.get_fallback_answer(q)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
405
 
406
  # ============ GRADIO ИНТЕРФЕЙС ============
407
  def gradio_chat(question, history):
408
+ if not question: return "", history
 
 
 
409
  answer = andrey.generate(question, history=history)
 
 
410
  new_history = history + [(question, answer)]
 
 
411
  andrey.memory['chat_history'] = new_history
 
412
  return "", new_history
413
 
414
  def launch_gradio():
415
+ # Исправлено для Gradio 6.0: theme передается в launch
416
+ with gr.Blocks(title="Андрей AI") as demo:
417
+ gr.Markdown("# 🤖 Андрей AI v7.0 (ZeroGPU)\n### Transformer с реальной памятью")
418
+ chatbot = gr.Chatbot(height=400, label="Диалог")
419
+ msg = gr.Textbox(label="Сообщение", placeholder="Напишите что-нибудь...")
420
+ clear = gr.Button("🧹 Очистить")
 
 
 
 
 
421
 
422
  msg.submit(gradio_chat, [msg, chatbot], [msg, chatbot])
423
  clear.click(lambda: [], None, chatbot)
 
 
 
 
 
 
 
 
424
 
425
+ demo.launch(theme=gr.themes.Soft()) # <--- Исправлено место передачи темы
426
 
 
427
  if __name__ == "__main__":
428
  andrey = AndreyAI(MODEL_PATH)
 
429
  if andrey.model is None or andrey.memory.get('epochs_trained', 0) == 0:
430
  andrey.train(150)
 
 
431
  launch_gradio()