neuron / app.py
X
Update app.py
b0468e2 verified
Raw
History Blame
37.7 kB
"""
AI PLATFORMER + NEURAL CHATBOT (FROM SCRATCH)
Fixed jump physics, larger player, AABB collision, rich graphics
Custom PyTorch NLP model - no external NLP libraries
"""
import os, json, random, threading, logging, time, re, math
from collections import deque, Counter
from dataclasses import dataclass
from typing import Dict, List, Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from flask import Flask, jsonify, request, render_template_string
logging.basicConfig(level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s')
logger = logging.getLogger(__name__)
@dataclass
class Cfg:
W: int = 80; H: int = 20; GROUND: int = 17; CHUNK: int = 30
SAFE: int = 15; VIEW: int = 40
GRAV: float = 0.45 # Fixed: was 0.35
JUMP: float = -3.8 # Fixed: was -6.5 (~3.6 tiles high now)
SPEED: float = 0.35
STATE: int = 40 * 40; ACTS: int = 4; MEM: int = 10000
BATCH: int = 64; GAMMA: float = 0.99; LR: float = 5e-4
EPS_DEC: float = 0.995; PORT: int = 7860
MODEL: str = "dqn_model.pth"; CHAT: str = "chat_data.json"
VOCAB_MAX: int = 2000; EMB_DIM: int = 64; HIDDEN: int = 128
CHAT_LR: float = 1e-3; CHAT_EPOCHS: int = 80; SIM_THRESH: float = 0.5
C = Cfg()
# ============================================================================
# CUSTOM TOKENIZER (From Scratch)
# ============================================================================
class Tokenizer:
PAD = "<PAD>"; UNK = "<UNK>"
def __init__(self, max_vocab: int = 2000):
self.max_vocab = max_vocab
self.word2idx: Dict[str, int] = {self.PAD: 0, self.UNK: 1}
self.idx2word: Dict[int, str] = {0: self.PAD, 1: self.UNK}
self.frozen = False
def _tokenize(self, text: str) -> List[str]:
text = text.lower().strip()
text = re.sub(r'[^\w\sа-яё]', ' ', text)
words = text.split()
tokens = []
for w in words:
tokens.append(w)
if len(w) > 2:
tokens.extend([w[i:i+2] for i in range(len(w)-1)])
return tokens
def build_vocab(self, texts: List[str]):
counter = Counter()
for t in texts:
counter.update(self._tokenize(t))
most_common = counter.most_common(self.max_vocab - 2)
for word, _ in most_common:
idx = len(self.word2idx)
self.word2idx[word] = idx
self.idx2word[idx] = word
self.frozen = True
logger.info(f"📝 Vocab built: {len(self.word2idx)} tokens")
def encode(self, text: str, max_len: int = 32) -> List[int]:
tokens = self._tokenize(text)[:max_len]
ids = [self.word2idx.get(t, 1) for t in tokens]
ids += [0] * (max_len - len(ids))
return ids
@property
def vocab_size(self) -> int:
return len(self.word2idx)
# ============================================================================
# NEURAL CHAT MODEL (From Scratch)
# ============================================================================
class ChatEncoder(nn.Module):
def __init__(self, vocab_size: int, emb_dim: int, hidden: int, out_dim: int):
super().__init__()
self.embedding = nn.Embedding(vocab_size, emb_dim, padding_idx=0)
self.fc1 = nn.Linear(emb_dim, hidden)
self.fc2 = nn.Linear(hidden, hidden)
self.fc3 = nn.Linear(hidden, out_dim)
self.relu = nn.ReLU()
self.dropout = nn.Dropout(0.2)
self.ln1 = nn.LayerNorm(hidden)
self.ln2 = nn.LayerNorm(hidden)
def forward(self, x):
emb = self.embedding(x)
mask = (x != 0).unsqueeze(-1).float()
pooled = (emb * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1)
h = self.ln1(self.fc1(pooled))
h = self.relu(h)
h = self.dropout(h)
h = self.ln2(self.fc2(h))
h = self.relu(h)
h = self.dropout(h)
out = self.fc3(h)
return nn.functional.normalize(out, p=2, dim=-1)
class NeuralChat:
SEQ_LEN = 32
OUT_DIM = 64
def __init__(self):
self.tokenizer = Tokenizer(C.VOCAB_MAX)
self.data: Dict[str, str] = {}
self.questions: List[str] = []
self.q_embeddings: Optional[torch.Tensor] = None
self.model: Optional[ChatEncoder] = None
self.device = torch.device('cpu')
self._load_data()
self._build_and_train()
logger.info(f"✅ Neural chat ready. {len(self.data)} entries.")
def _load_data(self):
if os.path.exists(C.CHAT):
try:
with open(C.CHAT, 'r', encoding='utf-8') as f:
self.data = json.load(f)
except Exception as e:
logger.warning(f"⚠️ Chat load failed: {e}")
if not self.data:
self.data = {
"как играть": "Стрелки ⬅️➡️ для движения, ⬆️/Пробел для прыжка",
"что делает нейросеть": "DQN учится играть методом проб и ошибок, получая награду за монеты",
"как обучить бота": "Напиши /data вопрос|ответ чтобы добавить знание",
"какой алгоритм": "Deep Q-Network с replay buffer и target network",
"зачем эпсилон": "Epsilon-greedy: чем выше ε, тем больше случайных действий для исследования",
"как сбросить уровень": "Нажми 🔄 Новый уровень",
"что такое dqn": "Deep Q-Network предсказывает ценность каждого действия в состоянии",
"сколько нейронов": "1600→256→256→128→4 (вход=40x40 сетка)",
"привет": "Привет! Я нейросетевой чат-бот платформера 🧠",
"как работает чат": "Я использую кастомную нейросеть на PyTorch с эмбеддингами и косинусным сходством",
"кто тебя создал": "Я написан с нуля на PyTorch без внешних NLP библиотек",
}
self._save_data()
def _save_data(self):
with open(C.CHAT, 'w', encoding='utf-8') as f:
json.dump(self.data, f, ensure_ascii=False, indent=2)
def _build_and_train(self):
self.questions = list(self.data.keys())
if not self.questions:
return
self.tokenizer.build_vocab(self.questions)
self.model = ChatEncoder(
vocab_size=self.tokenizer.vocab_size,
emb_dim=C.EMB_DIM, hidden=C.HIDDEN, out_dim=self.OUT_DIM
).to(self.device)
self._train_model()
self._update_index()
def _train_model(self):
if len(self.questions) < 2:
logger.warning("⚠️ Too few entries to train")
return
optimizer = optim.Adam(self.model.parameters(), lr=C.CHAT_LR)
q_ids = torch.LongTensor([
self.tokenizer.encode(q, self.SEQ_LEN) for q in self.questions
]).to(self.device)
n = len(self.questions)
best_loss = float('inf')
logger.info(f"🧠 Training chat NN: {n} samples, {C.CHAT_EPOCHS} epochs...")
for epoch in range(C.CHAT_EPOCHS):
total_loss = 0.0; num_pairs = 0
indices = list(range(n)); random.shuffle(indices)
for i in indices:
anchor = q_ids[i:i+1]
positive = q_ids[i:i+1]
neg_idx = random.choice([j for j in range(n) if j != i])
negative = q_ids[neg_idx:neg_idx+1]
emb_a = self.model(anchor)
emb_p = self.model(positive)
emb_n = self.model(negative)
pos_sim = nn.functional.cosine_similarity(emb_a, emb_p)
neg_sim = nn.functional.cosine_similarity(emb_a, emb_n)
loss = torch.relu(0.3 - pos_sim + neg_sim).mean()
optimizer.zero_grad(); loss.backward(); optimizer.step()
total_loss += loss.item(); num_pairs += 1
avg_loss = total_loss / max(num_pairs, 1)
if avg_loss < best_loss: best_loss = avg_loss
if (epoch + 1) % 20 == 0:
logger.info(f" Epoch {epoch+1}/{C.CHAT_EPOCHS}, loss={avg_loss:.4f}")
logger.info(f"✅ Chat training complete. Best loss: {best_loss:.4f}")
def _update_index(self):
if not self.questions or self.model is None:
self.q_embeddings = None; return
self.model.eval()
with torch.no_grad():
ids = torch.LongTensor([
self.tokenizer.encode(q, self.SEQ_LEN) for q in self.questions
]).to(self.device)
self.q_embeddings = self.model(ids)
self.model.train()
def ask(self, query: str) -> str:
if not self.questions or self.q_embeddings is None:
return "🤖 База пуста. Обучи меня: /data вопрос|ответ"
self.model.eval()
with torch.no_grad():
q_id = torch.LongTensor([self.tokenizer.encode(query, self.SEQ_LEN)]).to(self.device)
q_emb = self.model(q_id)
sims = nn.functional.cosine_similarity(q_emb, self.q_embeddings)
best_idx = torch.argmax(sims).item()
best_score = sims[best_idx].item()
self.model.train()
if best_score >= C.SIM_THRESH:
return f"{self.data[self.questions[best_idx]]} (🧠 {int(best_score*100)}%)"
return f"🤖 Не знаю «{query}». Научи: /data {query}|ответ"
def teach(self, question: str, answer: str) -> str:
question = question.strip().lower(); answer = answer.strip()
if not question or not answer: return "❌ Формат: /data вопрос|ответ"
is_new = question not in self.data
self.data[question] = answer; self._save_data()
self.questions = list(self.data.keys())
self.tokenizer = Tokenizer(C.VOCAB_MAX)
self.tokenizer.build_vocab(self.questions)
self.model = ChatEncoder(
vocab_size=self.tokenizer.vocab_size,
emb_dim=C.EMB_DIM, hidden=C.HIDDEN, out_dim=self.OUT_DIM
).to(self.device)
self._train_model(); self._update_index()
action = "Добавлено" if is_new else "Обновлено"
return f"✅ {action}: «{question}» → «{answer}». Модель переобучена."
@property
def stats(self) -> dict:
params = sum(p.numel() for p in self.model.parameters()) if self.model else 0
return {'entries': len(self.data), 'vocab': self.tokenizer.vocab_size,
'params': params, 'emb_dim': self.OUT_DIM}
# ============================================================================
# GAME ENGINE (Fixed Physics + Larger Player)
# ============================================================================
class Engine:
PW = 0.8 # Player width — FIXED: was 0.6
PH = 0.95 # Player height — FIXED: was 0.9
def __init__(self, seed=None):
self.seed = seed or random.randint(0, 999999)
self.reset()
def reset(self):
self.px, self.py = 5.0, float(C.GROUND)
self.vx, self.vy = 0.0, 0.0
self.grounded = True; self.alive = True
self.score = 0; self.coins = 0; self.step_n = 0
self.chunks: Dict[int, dict] = {}
self.obs: List[dict] = []; self.enemies: List[dict] = []; self.coin_list: List[dict] = []
self._load_chunks()
return self.get_state()
def _gen_chunk(self, cid: int) -> dict:
rng = random.Random((cid * 1337 + self.seed) % 999999)
bx = cid * C.CHUNK; obs, ens, cns = [], [], []
diff = max(1.0, abs(cid) * 0.1); safe = bx < C.SAFE
if not safe:
for _ in range(rng.randint(3, 6) + int(diff)):
x = bx + rng.randint(5, 25); h = rng.randint(1, 3 + int(diff * 0.5)); w = rng.randint(1, 3)
obs.append({'x': x, 'y': C.GROUND - h, 'w': w, 'h': h, 'pit': False})
for _ in range(rng.randint(1, 2)):
x = bx + rng.randint(10, 20)
obs.append({'x': x, 'y': C.GROUND + 1, 'w': rng.randint(2, 4), 'h': 1, 'pit': True})
for _ in range(rng.randint(1, 2)):
x = bx + rng.randint(10, 20)
ens.append({'x': x, 'y': C.GROUND - 1, 'type': rng.choice(['walker', 'jumper']),
'dir': rng.choice([-1, 1]), 'spd': 0.3 + rng.random() * 0.3,
'rng': rng.randint(3, 8), 'ox': x})
for _ in range(rng.randint(5, 10) + int(diff)):
cns.append({'x': bx + rng.randint(2, 28), 'y': rng.randint(5, C.GROUND - 2), 'collected': False})
return {'obs': obs, 'ens': ens, 'cns': cns}
def _load_chunks(self):
cc = int(self.px // C.CHUNK)
for i in range(cc - 1, cc + 3):
if i not in self.chunks: self.chunks[i] = self._gen_chunk(i)
vl, vr = self.px - C.W / 2, self.px + C.W / 2
self.obs, self.enemies, self.coin_list = [], [], []
for i in range(cc - 1, cc + 3):
ch = self.chunks.get(i, {})
self.obs.extend([o for o in ch.get('obs', []) if vl <= o['x'] <= vr])
self.enemies.extend([e for e in ch.get('ens', []) if vl <= e['x'] <= vr])
self.coin_list.extend([c for c in ch.get('cns', []) if not c['collected'] and vl <= c['x'] <= vr])
def _aabb(self, ax, ay, aw, ah, bx, by, bw, bh):
return ax < bx + bw and ax + aw > bx and ay < by + bh and ay + ah > by
def get_state(self):
s = np.zeros((C.VIEW, C.VIEW), dtype=np.float32); h = C.VIEW // 2
px, py = int(round(self.px)), int(round(self.py)); s[h, h] = 1.0
for o in self.obs:
dx, dy = int(round(o['x'])) - px, int(round(o['y'])) - py
v = -1.0 if o.get('pit') else 0.8
for ww in range(o.get('w', 1)):
for hh in range(o.get('h', 1)):
sx, sy = h + dx + ww, h + dy + hh
if 0 <= sx < C.VIEW and 0 <= sy < C.VIEW: s[sy, sx] = v
for e in self.enemies:
dx, dy = int(round(e['x'])) - px, int(round(e['y'])) - py
if 0 <= h + dx < C.VIEW and 0 <= h + dy < C.VIEW: s[h + dy, h + dx] = 0.7
for c in self.coin_list:
dx, dy = int(round(c['x'])) - px, int(round(c['y'])) - py
if 0 <= h + dx < C.VIEW and 0 <= h + dy < C.VIEW: s[h + dy, h + dx] = 0.3
return s.flatten()
def step(self, action: int):
sound = None
# Input
self.vx = 0.0
if action == 1: self.vx = -C.SPEED
elif action == 2: self.vx = C.SPEED
if action == 3 and self.grounded:
self.vy = C.JUMP; self.grounded = False; sound = 'jump'
# X axis movement + collision
self.px += self.vx
for o in self.obs:
if o.get('pit'): continue
if self._aabb(self.px, self.py, self.PW, self.PH, o['x'], o['y'], o['w'], o['h']):
if self.vx > 0: self.px = o['x'] - self.PW
elif self.vx < 0: self.px = o['x'] + o['w']
self.vx = 0
# Y axis movement + collision
self.vy += C.GRAV; self.py += self.vy; self.grounded = False
if self.py >= C.GROUND:
self.py = C.GROUND; self.vy = 0.0; self.grounded = True
for o in self.obs:
if o.get('pit'): continue
if self._aabb(self.px, self.py, self.PW, self.PH, o['x'], o['y'], o['w'], o['h']):
if self.vy > 0:
self.py = o['y'] - self.PH; self.vy = 0.0; self.grounded = True
elif self.vy < 0:
self.py = o['y'] + o['h']; self.vy = 0.0
# Death checks
if self.py > C.H + 2:
self.alive = False; return self.get_state(), -50.0, True, 'die'
for o in self.obs:
if o.get('pit') and o['x'] <= self.px + self.PW / 2 <= o['x'] + o['w'] and self.py >= C.GROUND:
self.alive = False; return self.get_state(), -50.0, True, 'die'
for e in self.enemies:
if self._aabb(self.px, self.py, self.PW, self.PH, e['x'] - 0.3, e['y'] - 0.3, 0.6, 0.6):
self.alive = False; return self.get_state(), -50.0, True, 'die'
# Coins
got = 0
for c in self.coin_list:
if not c['collected'] and self._aabb(self.px, self.py, self.PW, self.PH, c['x'] - 0.3, c['y'] - 0.3, 0.6, 0.6):
c['collected'] = True; got += 1
if got: self.coins += got; self.score += got * 10; sound = 'coin'
# Update enemies
t = time.time()
for e in self.enemies:
if e['type'] == 'walker':
e['x'] += e['spd'] * e['dir']
if abs(e['x'] - e['ox']) > e['rng']: e['dir'] *= -1
else:
e['y'] = (C.GROUND - 1) + np.sin(t * e['spd'] * 3) * 0.5
self.score += 1; self.step_n += 1; self._load_chunks()
done = self.step_n > 3000
return self.get_state(), 1.0 + got * 5.0, done, sound
def world_data(self):
return {
'player': [round(self.px, 2), round(self.py, 2)],
'obstacles': self.obs, 'entities': self.enemies,
'coins': [c for c in self.coin_list if not c['collected']],
'ground': C.GROUND, 'score': self.score,
'coins_collected': self.coins, 'alive': self.alive
}
# ============================================================================
# DQN AGENT
# ============================================================================
class Net(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Linear(C.STATE, 256), nn.ReLU(),
nn.Linear(256, 256), nn.ReLU(),
nn.Linear(256, 128), nn.ReLU(),
nn.Linear(128, C.ACTS)
)
def forward(self, x): return self.net(x)
class Agent:
def __init__(self):
self.dev = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
self.model = Net().to(self.dev)
self.target = Net().to(self.dev)
self.target.load_state_dict(self.model.state_dict())
self.opt = optim.Adam(self.model.parameters(), lr=C.LR)
self.crit = nn.MSELoss()
self.mem = deque(maxlen=C.MEM)
self.eps = 1.0; self.steps = 0; self.best = 0; self.training = False
if os.path.exists(C.MODEL):
try:
self.model.load_state_dict(torch.load(C.MODEL, map_location=self.dev))
self.target.load_state_dict(self.model.state_dict())
logger.info("✅ DQN loaded")
except Exception as e: logger.warning(f"⚠️ DQN load failed: {e}")
def act(self, s):
if random.random() <= self.eps: return random.randrange(C.ACTS)
with torch.no_grad():
return torch.argmax(self.model(torch.FloatTensor(s).unsqueeze(0).to(self.dev))).item()
def remember(self, s, a, r, ns, d): self.mem.append((s, a, r, ns, d))
def replay(self):
if len(self.mem) < C.BATCH: return
b = random.sample(self.mem, C.BATCH)
st = torch.FloatTensor([x[0] for x in b]).to(self.dev)
ac = torch.LongTensor([x[1] for x in b]).to(self.dev)
rw = torch.FloatTensor([x[2] for x in b]).to(self.dev)
ns = torch.FloatTensor([x[3] for x in b]).to(self.dev)
dn = torch.FloatTensor([x[4] for x in b]).to(self.dev)
q = self.model(st).gather(1, ac.unsqueeze(1)).squeeze()
nq = self.target(ns).max(1)[0].detach()
tgt = rw + C.GAMMA * nq * (1 - dn)
loss = self.crit(q, tgt)
self.opt.zero_grad(); loss.backward(); self.opt.step()
if self.eps > 0.01: self.eps *= C.EPS_DEC
self.steps += 1
if self.steps % 100 == 0: self.target.load_state_dict(self.model.state_dict())
def train_ep(self):
env = Engine(); s = env.reset(); tr = 0.0; d = False; n = 0
while not d and n < 500:
a = self.act(s); ns, r, d, _ = env.step(a)
self.remember(s, a, r, ns, d); self.replay()
s = ns; tr += r; n += 1
if tr > self.best: self.best = tr; self.save()
return tr
def save(self): torch.save(self.model.state_dict(), C.MODEL)
# ============================================================================
# GLOBAL STATE
# ============================================================================
agent = Agent()
chat = NeuralChat()
seed = random.randint(0, 999999)
ai_env = Engine(seed)
pl_env = Engine(seed)
is_training = False
# ============================================================================
# HTML (Rich Graphics + Larger Player Rendering)
# ============================================================================
HTML = """
<!DOCTYPE html>
<html lang="ru">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>🧠 AI Platformer + Neural Chat</title>
<style>
*{margin:0;padding:0;box-sizing:border-box}
body{background:#0d1117;color:#eee;font-family:'Segoe UI',sans-serif;display:flex;justify-content:center;padding:20px;min-height:100vh}
.wrap{max-width:1100px;width:100%}
h1{text-align:center;padding:15px 0;background:linear-gradient(135deg,#ff6b6b,#4ecdc4);-webkit-background-clip:text;-webkit-text-fill-color:transparent;font-size:2.2em}
.sub{text-align:center;color:#888;margin-bottom:15px}
.row{display:flex;gap:20px;flex-wrap:wrap}
.box{flex:1;min-width:320px;background:#161b22;border-radius:16px;padding:15px;box-shadow:0 8px 32px rgba(0,0,0,.5);border:1px solid #30363d}
.box h3{text-align:center;margin-bottom:10px;color:#c9d1d9}
canvas{width:100%;aspect-ratio:4/1;border-radius:8px;display:block;image-rendering:pixelated;background:#0d1117}
.ctrl{display:flex;justify-content:center;gap:12px;margin:15px 0;flex-wrap:wrap}
.ctrl button{padding:12px 30px;font-size:1.1em;border:none;border-radius:10px;cursor:pointer;font-weight:bold;transition:all .15s;color:#fff;text-shadow:0 1px 2px rgba(0,0,0,.5)}
.ctrl button:hover{transform:scale(1.05);filter:brightness(1.2)}
.ctrl button:active{transform:scale(.93)}
.bl,.br{background:linear-gradient(135deg,#ff6b6b,#ee5a24)}
.bj{background:linear-gradient(135deg,#4ecdc4,#2ecc71);padding:12px 45px}
.brs{background:linear-gradient(135deg,#a29bfe,#6c5ce7)}
.stats{background:#161b22;border-radius:12px;padding:12px 20px;margin:10px 0;display:flex;justify-content:space-around;flex-wrap:wrap;gap:10px;font-size:1.1em;border:1px solid #30363d}
.stats span{color:#ff6b6b;font-weight:bold}
.tabs{display:flex;gap:10px;margin:15px 0;flex-wrap:wrap}
.tab{padding:10px 22px;background:#161b22;border-radius:10px;cursor:pointer;border:2px solid #30363d;transition:all .3s;color:#c9d1d9}
.tab:hover{border-color:#ff6b6b}
.tab.active{border-color:#ff6b6b;background:#1c2333}
.tc{background:#161b22;border-radius:12px;padding:20px;min-height:200px;border:1px solid #30363d}
.ca{display:flex;gap:10px;margin-top:10px}
.ca input{flex:1;padding:10px;border-radius:8px;border:1px solid #30363d;background:#0d1117;color:#eee;font-size:1em}
.ca button{padding:10px 25px;background:#ff6b6b;color:#fff;border:none;border-radius:8px;cursor:pointer;font-weight:bold}
.cm{max-height:200px;overflow-y:auto;padding:5px}
.cm div{padding:6px 12px;margin:3px 0;border-radius:6px;background:#0d1117}
.cm .u{border-left:3px solid #ff6b6b}
.cm .b{border-left:3px solid #4ecdc4}
.hidden{display:none}
.nn-info{font-size:0.85em;color:#888;margin-top:8px;padding:8px;background:#0d1117;border-radius:6px}
</style>
</head>
<body>
<div class="wrap">
<h1>🧠 AI vs Player Platformer</h1>
<p class="sub">🤖 Нейросеть слева 🎮 Ты справа (⬅️ ➡️ ⬆️) | 💬 Чат = нейросеть с нуля</p>
<div class="row">
<div class="box"><h3>🤖 Нейросеть</h3><canvas id="ac"></canvas></div>
<div class="box"><h3>🎮 Ты</h3><canvas id="pc"></canvas></div>
</div>
<div class="stats">
<div>🤖 ИИ: <span id="as">0</span></div>
<div>🎮 Ты: <span id="ps">0</span></div>
<div>🪙 Монет: <span id="cc">0</span></div>
<div>🧠 ε: <span id="ep">1.00</span></div>
<div>🏆 Рекорд: <span id="bs">0</span></div>
</div>
<div class="ctrl">
<button class="bl" id="bL">⬅️ Влево</button>
<button class="bj" id="bJ">⬆️ ПРЫЖОК</button>
<button class="br" id="bR">➡️ Вправо</button>
<button class="brs" id="bReset">🔄 Новый уровень</button>
</div>
<div class="tabs">
<div class="tab active" data-tab="chat">💬 Нейро-чат</div>
<div class="tab" data-tab="train">🧠 Тренировка DQN</div>
<div class="tab" data-tab="stats">📊 Статистика</div>
</div>
<div class="tc">
<div id="chatTab">
<div class="cm" id="msgs">
<div class="b">🧠 Привет! Я нейросетевой чат, написанный с нуля на PyTorch.</div>
<div class="b">Команды: /data вопрос|ответ, /stats, /train</div>
<div class="b">Спрашивай что угодно — я ищу по смыслу, не по словам!</div>
</div>
<div class="ca"><input id="ci" placeholder="Задай вопрос или /data вопрос|ответ..." onkeydown="if(event.key==='Enter')sendChat()"><button onclick="sendChat()">➤</button></div>
<div class="nn-info" id="nnInfo">Загрузка модели...</div>
</div>
<div id="trainTab" class="hidden">
<h3>🧠 Тренировка DQN</h3><p>Deep Q-Network (256→256→128)</p>
<button onclick="startTrain()" style="padding:12px 35px;background:linear-gradient(135deg,#ff6b6b,#ee5a24);color:#fff;border:none;border-radius:10px;font-size:1.1em;cursor:pointer;margin-top:10px">🚀 Запустить</button>
<div id="ts" style="margin-top:10px;color:#888">⏸ Остановлена</div>
</div>
<div id="statsTab" class="hidden"><h3>📊 Статистика</h3><div id="sc">Загрузка...</div></div>
</div>
</div>
<script>
function initC(id){const c=document.getElementById(id);c.width=800;c.height=200;return c.getContext('2d')}
const aC=initC('ac'),pC=initC('pc');
let pA=0;
function draw(ctx,d,show){
const W=ctx.canvas.width,H=ctx.canvas.height,cW=W/80,cH=H/20;
ctx.clearRect(0,0,W,H);
// Sky
const sg=ctx.createLinearGradient(0,0,0,H);
sg.addColorStop(0,'#0f0c29');sg.addColorStop(0.5,'#302b63');sg.addColorStop(1,'#24243e');
ctx.fillStyle=sg;ctx.fillRect(0,0,W,H);
// Stars
ctx.fillStyle='rgba(255,255,255,0.3)';
for(let i=0;i<30;i++){const sx=(i*137+d.player[0]*0.1)%W,sy=(i*97)%((d.ground-2)*cH);ctx.fillRect(sx,sy,2,2)}
const cam=Math.max(0,d.player[0]-40);
function toS(wx,wy){return[(wx-cam)*cW,wy*cH]}
// Ground
const gy=d.ground*cH;
const gg=ctx.createLinearGradient(0,gy,0,H);
gg.addColorStop(0,'#4a7c59');gg.addColorStop(0.15,'#3d6b4e');gg.addColorStop(0.5,'#5c4033');gg.addColorStop(1,'#3e2723');
ctx.fillStyle=gg;ctx.fillRect(0,gy,W,H-gy);
ctx.fillStyle='#6abf69';ctx.fillRect(0,gy,W,cH*0.3);
ctx.fillStyle='#81c784';for(let gx=0;gx<W;gx+=8)ctx.fillRect(gx,gy-cH*0.1,4,cH*0.15);
// Obstacles
for(const o of d.obstacles){
const[x,y]=toS(o.x,o.y);
if(o.pit){
const pg=ctx.createLinearGradient(0,y-cH,0,y+cH);
pg.addColorStop(0,'#1a1a2e');pg.addColorStop(1,'#000');
ctx.fillStyle=pg;ctx.fillRect(x,y-cH,o.w*cW,cH*2);
ctx.fillStyle='#ff4444';ctx.fillRect(x,y-cH*0.5,o.w*cW,2);
}else{
const bg=ctx.createLinearGradient(x,y,x,y+o.h*cH);
bg.addColorStop(0,'#8d6e63');bg.addColorStop(1,'#6d4c41');
ctx.fillStyle=bg;ctx.fillRect(x,y,o.w*cW,o.h*cH);
ctx.strokeStyle='rgba(0,0,0,0.3)';ctx.lineWidth=1;
for(let by=0;by<o.h;by++){
const yy=y+by*cH;ctx.beginPath();ctx.moveTo(x,yy);ctx.lineTo(x+o.w*cW,yy);ctx.stroke();
const off=(by%2)*cW*0.5;
for(let bx=off;bx<o.w*cW;bx+=cW){ctx.beginPath();ctx.moveTo(x+bx,yy);ctx.lineTo(x+bx,yy+cH);ctx.stroke()}
}
ctx.fillStyle='rgba(255,255,255,0.15)';ctx.fillRect(x,y,o.w*cW,cH*0.15);
ctx.fillStyle='rgba(0,0,0,0.3)';ctx.fillRect(x+o.w*cW,y,3,o.h*cH);
}
}
// Enemies
const t=Date.now()/200;
for(const e of d.entities){
const[x,y]=toS(e.x,e.y);const bounce=Math.sin(t+e.x)*2;
ctx.save();ctx.translate(x+cW/2,y+cH/2+bounce);
const eg=ctx.createRadialGradient(0,0,2,0,0,cH/2);
eg.addColorStop(0,'#ff6b6b');eg.addColorStop(1,'#c0392b');
ctx.fillStyle=eg;ctx.beginPath();ctx.arc(0,0,cH/2.5,0,Math.PI*2);ctx.fill();
ctx.fillStyle='#fff';ctx.beginPath();ctx.arc(-4,-3,3,0,Math.PI*2);ctx.arc(4,-3,3,0,Math.PI*2);ctx.fill();
ctx.fillStyle='#000';const ex=e.dir*2;
ctx.beginPath();ctx.arc(-4+ex,-3,1.5,0,Math.PI*2);ctx.arc(4+ex,-3,1.5,0,Math.PI*2);ctx.fill();
ctx.shadowColor='#ff6b6b';ctx.shadowBlur=10;
ctx.strokeStyle='#ff6b6b';ctx.lineWidth=1;ctx.beginPath();ctx.arc(0,0,cH/2.2,0,Math.PI*2);ctx.stroke();
ctx.restore();
}
// Coins
for(const c of d.coins){
const[x,y]=toS(c.x,c.y);const pulse=1+Math.sin(t*2+c.x)*0.15;
ctx.save();ctx.translate(x+cW/2,y+cH/2);ctx.scale(pulse,pulse);
const cg=ctx.createRadialGradient(-2,-2,1,0,0,cH/3);
cg.addColorStop(0,'#fff9c4');cg.addColorStop(0.5,'#ffd700');cg.addColorStop(1,'#f9a825');
ctx.fillStyle=cg;ctx.beginPath();ctx.arc(0,0,cH/3,0,Math.PI*2);ctx.fill();
ctx.shadowColor='#ffd700';ctx.shadowBlur=12;
ctx.strokeStyle='#ffeb3b';ctx.lineWidth=1.5;ctx.beginPath();ctx.arc(0,0,cH/3,0,Math.PI*2);ctx.stroke();
ctx.fillStyle='rgba(255,255,255,0.8)';ctx.beginPath();ctx.arc(-3,-3,2,0,Math.PI*2);ctx.fill();
ctx.restore();
}
// PLAYER — FIXED: larger, more visible
if(show&&d.alive){
const[px,py]=toS(d.player[0],d.player[1]);
const pw=cW*0.85, ph=cH*0.95;
const ox=px+(cW-pw)/2, oy=py+(cH-ph)/2;
ctx.save();
ctx.shadowColor='#00ff88';ctx.shadowBlur=25;
const pg=ctx.createLinearGradient(ox,oy,ox+pw,oy+ph);
pg.addColorStop(0,'#00ff88');pg.addColorStop(0.5,'#00e676');pg.addColorStop(1,'#00b894');
ctx.fillStyle=pg;ctx.fillRect(ox,oy,pw,ph);
ctx.shadowBlur=0;
ctx.strokeStyle='#b9f6ca';ctx.lineWidth=2;ctx.strokeRect(ox,oy,pw,ph);
// Eyes
ctx.fillStyle='#fff';
ctx.fillRect(ox+pw*0.15,oy+ph*0.2,pw*0.25,ph*0.22);
ctx.fillRect(ox+pw*0.6,oy+ph*0.2,pw*0.25,ph*0.22);
ctx.fillStyle='#0d1117';
ctx.fillRect(ox+pw*0.22,oy+ph*0.27,pw*0.12,ph*0.1);
ctx.fillRect(ox+pw*0.67,oy+ph*0.27,pw*0.12,ph*0.1);
// Mouth
ctx.fillStyle='#0d1117';
ctx.fillRect(ox+pw*0.3,oy+ph*0.6,pw*0.4,ph*0.08);
ctx.restore();
}
}
async function update(){
try{
const r=await fetch('/step',{method:'POST',headers:{'Content-Type':'application/json'},body:JSON.stringify({action:pA})});
const d=await r.json();
draw(aC,d.ai,true);draw(pC,d.player,d.player.alive);
document.getElementById('as').textContent=d.ai.score;
document.getElementById('ps').textContent=d.player.score;
document.getElementById('cc').textContent=d.player.coins_collected;
document.getElementById('ep').textContent=d.epsilon.toFixed(3);
document.getElementById('bs').textContent=d.best_score;
}catch(e){}
}
const sA=v=>{pA=v};
document.getElementById('bL').onmousedown=()=>sA(1);document.getElementById('bL').onmouseup=()=>sA(0);
document.getElementById('bR').onmousedown=()=>sA(2);document.getElementById('bR').onmouseup=()=>sA(0);
document.getElementById('bJ').onmousedown=()=>sA(3);document.getElementById('bJ').onmouseup=()=>sA(0);
document.addEventListener('keydown',e=>{
if(e.key==='ArrowLeft'){e.preventDefault();sA(1)}
else if(e.key==='ArrowRight'){e.preventDefault();sA(2)}
else if(e.key==='ArrowUp'||e.key===' '){e.preventDefault();sA(3)}
});
document.addEventListener('keyup',e=>{if(['ArrowLeft','ArrowRight','ArrowUp',' '].includes(e.key)){e.preventDefault();sA(0)}});
document.getElementById('bReset').onclick=async()=>{
const r=await fetch('/reset',{method:'POST'});const d=await r.json();
draw(aC,d.ai,true);draw(pC,d.player,true);
document.getElementById('as').textContent=d.ai.score;
document.getElementById('ps').textContent=d.player.score;
};
async function sendChat(){
const inp=document.getElementById('ci');const msg=inp.value.trim();if(!msg)return;inp.value='';
const m=document.getElementById('msgs');
m.innerHTML+=`<div class="u">👤 ${msg}</div>`;m.scrollTop=m.scrollHeight;
const r=await fetch('/chat',{method:'POST',headers:{'Content-Type':'application/json'},body:JSON.stringify({message:msg})});
const d=await r.json();
m.innerHTML+=`<div class="b">🧠 ${d.response}</div>`;m.scrollTop=m.scrollHeight;
}
async function startTrain(){
document.getElementById('ts').textContent='⏳ Запуск...';
const r=await fetch('/train',{method:'POST'});const d=await r.json();
document.getElementById('ts').textContent=d.message;
}
document.querySelectorAll('.tab').forEach(t=>t.onclick=function(){
document.querySelectorAll('.tab').forEach(x=>x.classList.remove('active'));
this.classList.add('active');const n=this.dataset.tab;
document.querySelectorAll('.tc>div').forEach(d=>d.classList.add('hidden'));
document.getElementById(n+'Tab').classList.remove('hidden');
if(n==='stats')fetch('/stats').then(r=>r.json()).then(d=>{
document.getElementById('sc').innerHTML=`
<p>🧠 DQN шагов: ${d.steps}</p>
<p>📉 ε: ${d.epsilon}</p>
<p>🏆 Рекорд: ${d.best_score}</p>
<p>⚡ Тренируется: ${d.training?'✅':'❌'}</p>
<hr style="border-color:#30363d;margin:8px 0">
<p>💬 Чат записей: ${d.chat.entries}</p>
<p>📝 Словарь: ${d.chat.vocab}</p>
<p>🔢 Параметров чата: ${d.chat.params.toLocaleString()}</p>
<p>📐 Эмбеддинг: ${d.chat.emb_dim}D</p>`;
});
});
fetch('/stats').then(r=>r.json()).then(d=>{
document.getElementById('nnInfo').textContent=
`🧠 Чат-нейросеть: ${d.chat.params.toLocaleString()} параметров | Словарь: ${d.chat.vocab} | Эмбеддинг: ${d.chat.emb_dim}D | Записей: ${d.chat.entries}`;
});
setInterval(update,100);update();
</script>
</body>
</html>
"""
# ============================================================================
# FLASK ROUTES
# ============================================================================
app = Flask(__name__)
@app.route('/')
def index(): return render_template_string(HTML)
@app.route('/step', methods=['POST'])
def step():
global ai_env, pl_env
action = request.json.get('action', 0)
if ai_env.alive: ai_env.step(agent.act(ai_env.get_state()))
else: ai_env.reset()
if pl_env.alive: pl_env.step(action)
else: pl_env.reset()
return jsonify({
'ai': ai_env.world_data(), 'player': pl_env.world_data(),
'epsilon': agent.eps, 'best_score': agent.best
})
@app.route('/reset', methods=['POST'])
def reset():
global seed, ai_env, pl_env
seed = random.randint(0, 999999)
ai_env = Engine(seed); pl_env = Engine(seed)
return jsonify({'ai': ai_env.world_data(), 'player': pl_env.world_data()})
@app.route('/chat', methods=['POST'])
def chat_route():
msg = request.json.get('message', '').strip()
if msg.startswith('/data '):
p = msg[6:].split('|')
if len(p) != 2: return jsonify({'response': "❌ Формат: /data вопрос|ответ"})
return jsonify({'response': chat.teach(p[0], p[1])})
elif msg == '/stats':
s = chat.stats
return jsonify({'response': f"🧠 Чат: {s['params']} парам., {s['vocab']} слов, {s['entries']} записей, {s['emb_dim']}D"})
elif msg == '/train':
return jsonify({'response': start_training()})
else:
return jsonify({'response': chat.ask(msg)})
@app.route('/train', methods=['POST'])
def train_route(): return jsonify({'message': start_training()})
@app.route('/stats')
def stats():
return jsonify({
'steps': agent.steps, 'epsilon': round(agent.eps, 3),
'best_score': agent.best, 'training': is_training,
'chat': chat.stats
})
def start_training():
global is_training
if is_training: return "⏳ Уже тренируется!"
is_training = True
def _t():
global is_training
try:
for ep in range(100):
if not is_training: break
sc = agent.train_ep()
if ep % 10 == 0: logger.info(f"DQN Ep {ep}: score={sc:.1f}, ε={agent.eps:.3f}")
except Exception as e: logger.error(f"DQN error: {e}")
finally: is_training = False
threading.Thread(target=_t, daemon=True).start()
return "🚀 Тренировка DQN запущена!"
if __name__ == '__main__':
app.run(host='0.0.0.0', port=C.PORT, debug=False)