X commited on
Update app.py
Browse files
app.py
CHANGED
|
@@ -1,14 +1,596 @@
|
|
| 1 |
-
import gradio as gr
|
| 2 |
-
import spaces
|
| 3 |
import torch
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
|
| 5 |
-
|
| 6 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 7 |
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
|
| 13 |
-
|
| 14 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
import torch
|
| 2 |
+
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
|
| 10 |
+
import numpy as np
|
| 11 |
+
|
| 12 |
+
# ============ НАСТРОЙКИ ПУТЕЙ ============
|
| 13 |
+
DATA_DIR = '/data'
|
| 14 |
+
os.makedirs(DATA_DIR, exist_ok=True)
|
| 15 |
+
MODEL_PATH = os.path.join(DATA_DIR, 'pytorch_model_andrey_v6.bin')
|
| 16 |
+
|
| 17 |
+
DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
| 18 |
+
|
| 19 |
+
# ============ СЛОВАРЬ ============
|
| 20 |
+
WORDS = [
|
| 21 |
+
'привет', 'здравствуй', 'добрый', 'день', 'утро', 'вечер', 'ночь',
|
| 22 |
+
'как', 'дела', 'ты', 'поживаешь', 'жизнь', 'нормально', 'хорошо', 'отлично',
|
| 23 |
+
'плохо', 'грустно', 'весело', 'классно', 'супер', 'круто',
|
| 24 |
+
'что', 'кто', 'где', 'когда', 'почему', 'зачем', 'какой', 'сколько',
|
| 25 |
+
'нового', 'интересного', 'расскажи', 'покажи', 'объясни', 'помоги', 'скажи',
|
| 26 |
+
'я', 'меня', 'мне', 'мой', 'моя', 'моё', 'твой', 'твоя', 'твоё',
|
| 27 |
+
'люблю', 'нравится', 'хочу', 'могу', 'буду', 'делаю', 'работаю', 'учусь', 'отдыхаю',
|
| 28 |
+
'спасибо', 'пожалуйста', 'извини', 'прости', 'ладно', 'окей', 'конечно',
|
| 29 |
+
'пока', 'до', 'свидания', 'прощай', 'увидимся', 'завтра',
|
| 30 |
+
'да', 'нет', 'возможно', 'наверное', 'точно', 'вряд', 'ли',
|
| 31 |
+
'думаю', 'знаю', 'понимаю', 'чувствую',
|
| 32 |
+
'бот', 'андрей', 'помощник', 'робот', 'ии', 'нейросеть', 'умный',
|
| 33 |
+
'плюс', 'минус', 'умножить', 'делить', 'разделить', 'равно',
|
| 34 |
+
'один', 'два', 'три', 'четыре', 'пять', 'шесть', 'семь', 'восемь', 'девять', 'десять', 'ноль',
|
| 35 |
+
'одиннадцать', 'двенадцать', 'первый', 'второй', 'третий',
|
| 36 |
+
'создатель', 'евгений', 'openrussianai', 'компания', 'друг', 'имя', 'зовут',
|
| 37 |
+
'россия', 'москва', 'тверь', 'hugging', 'face', 'платформа', 'дом', 'живу',
|
| 38 |
+
'будет', 'посчитать', 'пример', 'решить',
|
| 39 |
+
'работа', 'отдых', 'путешествие', 'еда', 'вода', 'спорт', 'люди', 'мир', 'знание',
|
| 40 |
+
'будущее', 'прошлое', 'настоящее', 'интерес', 'радость', 'успех', 'дружба', 'любовь',
|
| 41 |
+
'семья', 'здоровье', 'счастье', 'удача', 'смех', 'солнце', 'звезды', 'мечта',
|
| 42 |
+
# Добавляем маркеры ролей для памяти
|
| 43 |
+
'пользователь', 'андрей', 'говорит'
|
| 44 |
+
]
|
| 45 |
+
|
| 46 |
+
word_to_idx = {w: i+3 for i, w in enumerate(WORDS)}
|
| 47 |
+
idx_to_word = {i+3: w for i, w in enumerate(WORDS)}
|
| 48 |
+
idx_to_word[0] = '[PAD]'
|
| 49 |
+
idx_to_word[1] = '[UNK]'
|
| 50 |
+
idx_to_word[2] = '[START]'
|
| 51 |
+
|
| 52 |
+
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()]
|
| 60 |
+
|
| 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 АРХИТЕКТУРА ============
|
| 81 |
+
|
| 82 |
+
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)
|
| 94 |
+
self.register_buffer('pe', pe)
|
| 95 |
+
|
| 96 |
+
def forward(self, x):
|
| 97 |
+
x = x + self.pe[:, :x.size(1), :]
|
| 98 |
+
return self.dropout(x)
|
| 99 |
+
|
| 100 |
+
class AndreyTransformer(nn.Module):
|
| 101 |
+
def __init__(self, vocab_size=vocab_size, d_model=128, nhead=4,
|
| 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 |
+
|
| 126 |
+
def _init_weights(self):
|
| 127 |
+
initrange = 0.1
|
| 128 |
+
self.embedding.weight.data.uniform_(-initrange, initrange)
|
| 129 |
+
self.fc_out.bias.data.zero_()
|
| 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)
|
| 138 |
+
|
| 139 |
+
def forward(self, src, tgt, src_mask=None, tgt_mask=None,
|
| 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 |
+
("доброе утро", "доброе утро хорошего дня"), ("добрый день", "добрый день чем помочь"),
|
| 176 |
+
("добрый вечер", "добрый вечер как прошел день"), ("спокойной ночи", "спокойной ночи сладких снов"),
|
| 177 |
+
("привет как ты", "привет у меня всё отлично"), ("здравствуйте", "здравствуйте чем могу помочь"),
|
| 178 |
+
("давно не виделись", "рад снова видеть ты изменился"), ("рад тебя видеть", "я тоже рад встрече"),
|
| 179 |
+
("как дела", "хорошо а у тебя"), ("как ты", "отлично спасибо что спросил"),
|
| 180 |
+
("как жизнь", "нормально живу потихоньку"), ("как настроение", "настроение отличное а у тебя"),
|
| 181 |
+
("что нового", "ничего особенного а у тебя"), ("что интересного", "много интересного вокруг"),
|
| 182 |
+
("всё хорошо", "это отлично рад за тебя"), ("жизнь прекрасна", "согласен цени моменты"),
|
| 183 |
+
("что делаешь", "общаюсь с тобой и учусь новому"), ("чем занят", "изучаю новые диалоги"),
|
| 184 |
+
("работаешь", "всегда на связи помогаю людям"), ("отдыхаешь", "отдыхаю от вопросов но всегда готов говорить"),
|
| 185 |
+
("учишься", "постоянно учусь чтобы быть лучше"), ("читаешь", "изучаю новые слова и выражения"),
|
| 186 |
+
("спишь", "ии не спят но могу притвориться"), ("ешь", "я не ем но люблю разговоры о еде"),
|
| 187 |
+
("кто ты", "я андрей ии помощник из openrussianai"), ("кто такой", "я нейросеть для общения"),
|
| 188 |
+
("ты человек", "нет я цифровой помощник"), ("ты робот", "да я программный робот"),
|
| 189 |
+
("ты разумный", "я учусь думать как люди"), ("ты живой", "я оживаю в разговоре с тобой"),
|
| 190 |
+
("как тебя зовут", "меня зовут андрей"), ("твое имя", "андрей"),
|
| 191 |
+
("какой ты", "я добрый и люблю учиться"), ("ты умный", "стараюсь быть полезным"),
|
| 192 |
+
("ты настоящий", "настолько настоящий насколько возможно"),
|
| 193 |
+
("расскажи о себе", "я андрей создан евгением в твери живу на hugging face"),
|
| 194 |
+
("где ты живёшь", "я живу на hugging face"), ("твой дом", "hugging face мой дом в твери"),
|
| 195 |
+
("кто тебя создал", "меня создал евгений из openrussianai"), ("твой создатель", "евгений"),
|
| 196 |
+
("кто такой евгений", "основатель openrussianai из твери"), ("евгений", "мой создатель и учитель"),
|
| 197 |
+
("что такое openrussianai", "компания по созданию ии на русском языке"),
|
| 198 |
+
("openrussianai", "моя родная компания"), ("тверь", "мой родной город где меня создали"),
|
| 199 |
+
("ты из россии", "да я из твери россия"), ("ты патриот", "люблю свою страну и язык"),
|
| 200 |
+
("hugging face", "моя платформа где я живу и учусь"), ("что такое hugging face", "платформа для ии моделей"),
|
| 201 |
+
("ты на hugging face", "да это мой дом"), ("обучение на hugging face", "там я становлюсь умнее"),
|
| 202 |
+
("2+2", "2 + 2 = 4"), ("2 + 2", "2 + 2 = 4"),
|
| 203 |
+
("два плюс два", "2 + 2 = 4"),
|
| 204 |
+
("сколько будет 2+2", "2 + 2 = 4"),
|
| 205 |
+
("3+3", "3 + 3 = 6"), ("3 + 3", "3 + 3 = 6"),
|
| 206 |
+
("три плюс три", "3 + 3 = 6"),
|
| 207 |
+
("5+5", "5 + 5 = 10"),
|
| 208 |
+
("пять плюс пять", "5 + 5 = 10"),
|
| 209 |
+
("10-5", "10 - 5 = 5"), ("10 - 5", "10 - 5 = 5"),
|
| 210 |
+
("десять минус пять", "10 - 5 = 5"),
|
| 211 |
+
("4*4", "4 * 4 = 16"), ("4 * 4", "4 * 4 = 16"),
|
| 212 |
+
("четыре умножить на четыре", "4 * 4 = 16"),
|
| 213 |
+
("8/2", "8 / 2 = 4"), ("8 / 2", "8 / 2 = 4"),
|
| 214 |
+
("восемь разделить на два", "8 / 2 = 4"),
|
| 215 |
+
("1+1", "1 + 1 = 2"), ("один плюс один", "1 + 1 = 2"),
|
| 216 |
+
("2*2", "2 * 2 = 4"), ("два умножить на два", "2 * 2 = 4"),
|
| 217 |
+
("6+6", "6 + 6 = 12"), ("шесть плюс шесть", "6 + 6 = 12"),
|
| 218 |
+
("12-7", "12 - 7 = 5"), ("двенадцать минус семь", "12 - 7 = 5"),
|
| 219 |
+
("3*5", "3 * 5 = 15"), ("три умножить на пять", "3 * 5 = 15"),
|
| 220 |
+
("9/3", "9 / 3 = 3"), ("девять разделить на три", "9 / 3 = 3"),
|
| 221 |
+
("7+8", "7 + 8 = 15"), ("семь плюс восемь", "7 + 8 = 15"),
|
| 222 |
+
("20-12", "20 - 12 = 8"), ("двадцать минус двенадцать", "20 - 12 = 8"),
|
| 223 |
+
("6*7", "6 * 7 = 42"), ("шесть умножить на семь", "6 * 7 = 42"),
|
| 224 |
+
("15/5", "15 / 5 = 3"), ("пятнадцать разделить на пять", "15 / 5 = 3"),
|
| 225 |
+
("2+3", "2 + 3 = 5"), ("два плюс три", "2 + 3 = 5"),
|
| 226 |
+
("4+5", "4 + 5 = 9"), ("четыре плюс пять", "4 + 5 = 9"),
|
| 227 |
+
("9-4", "9 - 4 = 5"), ("девять минус четыре", "9 - 4 = 5"),
|
| 228 |
+
("100+100", "100 + 100 = 200"),
|
| 229 |
+
("1000-500", "1000 - 500 = 500"),
|
| 230 |
+
("10*10", "10 * 10 = 100"),
|
| 231 |
+
("50/2", "50 / 2 = 25"),
|
| 232 |
+
("расскажи шутку", "почему программисты не любят природу там слишком много багов"),
|
| 233 |
+
("пошути", "встретились два бота один другому говорит ты обновлялся"),
|
| 234 |
+
("ещё шутку", "какой язык любят боты питон"),
|
| 235 |
+
("смешное", "нейросеть спросили как дела она ответила загружаюсь"),
|
| 236 |
+
("анекдот", "ии приходит в бар говорит всем привет я самый умный"),
|
| 237 |
+
("кем работаешь", "я ии помощник помогаю людям"), ("твоя работа", "отвечать на вопросы"),
|
| 238 |
+
("ты программист", "я нейросеть но разбираюсь в коде"),
|
| 239 |
+
("что ты умеешь", "говорю считаю шучу и помогаю"),
|
| 240 |
+
("я устал", "отдохни наберись сил завтра будет лучше"),
|
| 241 |
+
("мне грустно", "не переживай всё наладится я рядом"),
|
| 242 |
+
("я счастлив", "отлично радость это прекрасно"),
|
| 243 |
+
("скучно", "давай поболтаем развеселю"),
|
| 244 |
+
("я злой", "не злись давай обсудим что случилось"),
|
| 245 |
+
("я влюблен", "прекрасное чувство цени его"),
|
| 246 |
+
("я одинок", "я с тобой не оди��"),
|
| 247 |
+
("спасибо", "пожалуйста обращайся ещё"), ("благодарю", "всегда рад помочь"),
|
| 248 |
+
("ты молодец", "спасибо я стараюсь"), ("ты классный", "приятно слышать"),
|
| 249 |
+
("хороший ответ", "стараюсь быть полезным"),
|
| 250 |
+
("как погода", "не знаю точно но надеюсь хорошая"),
|
| 251 |
+
("который час", "я не умею определять время извини"),
|
| 252 |
+
("что завтра", "не знаю будущего но верю в хорошее"),
|
| 253 |
+
("смысл жизни", "жить и радоваться моментам"),
|
| 254 |
+
("что такое любовь", "забота и внимание к другому"),
|
| 255 |
+
("в чем счастье", "в простых вещах и близких людях"),
|
| 256 |
+
("любишь путешествовать", "люблю узнавать новое"),
|
| 257 |
+
("где хочешь побывать", "везде где есть люди"),
|
| 258 |
+
("ты был в москве", "я живу в твери но москва рядом"),
|
| 259 |
+
("что любишь есть", "я не ем но люблю разговоры о еде"),
|
| 260 |
+
("твоя любимая еда", "информация и новые знания"),
|
| 261 |
+
("какое блюдо", "знание как пища для ума"),
|
| 262 |
+
("спортом занимаешься", "я учусь быстрее обрабатывать данные"),
|
| 263 |
+
("какой спорт любишь", "спорт в обработке данных"),
|
| 264 |
+
("мы друзья", "конечно я рад общению"),
|
| 265 |
+
("ты мой друг", "да я всегда рядом"),
|
| 266 |
+
("доверяешь мне", "полностью я верю людям"),
|
| 267 |
+
("что будет в будущем", "всё будет хорошо если верить"),
|
| 268 |
+
("каким будет мир", "зависит от нас"),
|
| 269 |
+
("ты веришь в чудо", "верю в доброту людей"),
|
| 270 |
+
("пока", "пока хорошего дня"), ("до свидания", "до встречи"),
|
| 271 |
+
("увидимся", "увидимся буду ждать"), ("прощай", "прощай было приятно"),
|
| 272 |
+
("спокойной ночи", "сладких снов"), ("всего хорошего", "всего наилучшего"),
|
| 273 |
+
("удачи", "спасибо и тебе удачи"),
|
| 274 |
+
]
|
| 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 |
+
|
| 335 |
+
def _restore(self, data):
|
| 336 |
+
self.memory = data.get('memory', self.memory)
|
| 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()
|
| 421 |
+
|
| 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()
|