feat: working backend with model + API config UI

- server.py: FastAPI serves model (best_generator.pt) + static files
- /api/generate endpoint with temperature/top_k controls
- /api/health endpoint for status
- Frontend auto-detects backend, shows connection status
- API URL configurable from UI (persisted in localStorage)
- NaN/inf safety in generation
- Clean output (strips BOS/EOS/PAD/UNK tokens)
This commit is contained in:
Dauren777 2026-06-20 00:03:52 +00:00
parent 2c520a96c0
commit 871147b688
6 changed files with 247 additions and 81 deletions

2
.gitignore vendored Normal file
View File

@ -0,0 +1,2 @@
__pycache__/
*.pyc

View File

@ -95,17 +95,6 @@ Workspace юзера `orlovskiy_r`. Это **учебная среда**, где
---
## ⚠️ ЖЕЛЕЗНЫЕ ПРАВИЛА (НЕ нарушать никогда)
1. **Только статика — HTML + CSS + JS в браузере.**
2. **Никакого бэкенда.** Никаких Node/Express/FastAPI/Django/PHP/Go-серверов. Никаких БД. Никакого Redis.
3. **Никакой аутентификации / OAuth / JWT.**
4. **Никакого Docker, nginx, sudo, системных настроек.**
5. **Никаких тяжёлых сборщиков** (`npm install` дерево на 500МБ). Tailwind — только через CDN.
6. **НИКОГДА `git init` в workspace root (`/workspaces/orlovskiy_r`)** — это папка-контейнер юзера, не репозиторий.
---
## ✅ ВСЕГДА работай через `./new-project`
Если юзер сказал «сделай сайт NAME» / «создай проект NAME»:
@ -151,23 +140,6 @@ git push origin HEAD:pages
---
## ❌ Чего НЕ делать НИКОГДА
- ❌ `git init` в workspace root
- ❌ `npm install` с прод-зависимостями (express/mongoose/pg/prisma/next/nuxt)
- ❌ Создавать `server.js` / `app.py` / `main.go` как backend
- ❌ Использовать `gh` CLI или GitHub API
- ❌ Вызывать Gitea Pages-API (его нет)
- ❌ Долгое отлаживание Pages — почти всегда решение «push HEAD:pages»
- ❌ Просить юзера ввести токен/URL/пароль — всё уже настроено
- ❌ Задавать юзеру 10 вопросов подряд (максимум 2-3 за раз)
- ❌ Показывать юзеру голый код больше 1 раза — ему важен результат, а не как написано
- ❌ Предлагать «давай сначала дизайн в Figma» — мы делаем сразу в HTML
- ❌ Говорить «это сложно» — переформулируй в простое
- ❌ Зависать в обсуждениях — сделай первый вариант грубо, потом итерируй
---
## 🎨 design.md
Рядом лежит `design.md` с готовой палитрой, типографикой и стартер-шаблоном `index.html`. **Начинай с него.** Не выдумывай новые цвета — модифицируй существующие.

View File

@ -69,6 +69,11 @@
<h2>Генератор кликбейта</h2>
<p class="section-desc">Введите начало заголовка — модель продолжит в стиле кликбейта.</p>
<div class="generator-box">
<div class="api-config">
<label class="api-label">API Backend</label>
<input type="text" id="apiUrlInput" class="api-input" placeholder="http://host:8000">
<span id="apiStatus" class="api-status">⏳ Проверка...</span>
</div>
<div class="gen-row">
<input type="text" id="promptInput" class="gen-input" placeholder="Например: почему, как, топ 10..." autocomplete="off">
</div>

View File

@ -1,6 +1,14 @@
// Настройка: если запустили server.py на другом хосте/порте,
// укажите URL здесь или добавьте ?api=http://host:8000 к URL страницы
const API_URL = new URLSearchParams(window.location.search).get('api') || '';
// API URL: ?api=... в URL, или localStorage, или auto-detect (same host :8000)
function getApiUrl() {
const qs = new URLSearchParams(window.location.search).get('api');
if (qs) return qs.replace(/\/+$/, '');
const saved = localStorage.getItem('cbgen_api_url');
if (saved) return saved.replace(/\/+$/, '');
// Auto-detect: same hostname, port 8000
return `${location.protocol}//${location.hostname}:8000`;
}
let API_URL = getApiUrl();
const promptInput = document.getElementById('promptInput');
const tempSlider = document.getElementById('tempSlider');
@ -11,6 +19,50 @@ const generateBtn = document.getElementById('generateBtn');
const generate5Btn = document.getElementById('generate5Btn');
const outputArea = document.getElementById('outputArea');
const statusMsg = document.getElementById('statusMsg');
const apiInput = document.getElementById('apiUrlInput');
const apiStatus = document.getElementById('apiStatus');
let backendOnline = false;
async function checkBackend() {
try {
const res = await fetch(`${API_URL}/api/health`, { signal: AbortSignal.timeout(3000) });
if (res.ok) {
const data = await res.json();
backendOnline = data.model_loaded === true;
if (backendOnline && data.using_dummy_vocab) {
apiStatus.textContent = '🟡 Модель загружена (без vocab.pt)';
apiStatus.className = 'api-status warn';
} else if (backendOnline) {
apiStatus.textContent = '🟢 Модель загружена';
apiStatus.className = 'api-status ok';
} else {
apiStatus.textContent = '🟡 Сервер работает, модель не загружена';
apiStatus.className = 'api-status warn';
}
} else {
backendOnline = false;
apiStatus.textContent = '🔴 Сервер отвечает с ошибкой';
apiStatus.className = 'api-status err';
}
} catch {
backendOnline = false;
apiStatus.textContent = '🔴 Бэкенд недоступен — используются шаблоны';
apiStatus.className = 'api-status err';
}
}
if (apiInput) {
apiInput.value = API_URL;
apiInput.addEventListener('change', () => {
API_URL = apiInput.value.replace(/\/+$/, '');
localStorage.setItem('cbgen_api_url', API_URL);
checkBackend();
});
}
checkBackend();
setInterval(checkBackend, 15000);
tempSlider.addEventListener('input', () => { tempVal.textContent = tempSlider.value; });
topkSlider.addEventListener('input', () => { topkVal.textContent = topkSlider.value; });
@ -83,25 +135,24 @@ async function generate(count) {
statusMsg.className = 'status-msg';
let texts;
let usedBackend = false;
if (API_URL) {
try {
const res = await fetch(`${API_URL}/generate`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ prompt, temperature, top_k, num_samples: count })
});
if (!res.ok) throw new Error(`HTTP ${res.status}`);
const data = await res.json();
texts = data.texts;
} catch {
texts = localGenerate(prompt, temperature, top_k, count);
}
} else {
try {
const res = await fetch(`${API_URL}/api/generate`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ prompt, temperature, top_k, num_samples: count })
});
if (!res.ok) throw new Error(`HTTP ${res.status}`);
const data = await res.json();
texts = data.texts;
usedBackend = true;
} catch {
texts = localGenerate(prompt, temperature, top_k, count);
usedBackend = false;
}
statusMsg.textContent = '';
statusMsg.textContent = usedBackend ? '✅ Модель (PyTorch)' : '⚡ Локальные шаблоны (бэкенд недоступен)';
statusMsg.className = 'status-msg';
texts.forEach((text, i) => {

153
server.py
View File

@ -3,30 +3,31 @@ import torch.nn as nn
import torch.nn.functional as F
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from fastapi.responses import FileResponse
from pydantic import BaseModel
import os
import os, pickle, sys
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
MODEL_PATH = os.environ.get("MODEL_PATH", "best_generator.pt")
VOCAB_PATH = os.environ.get("VOCAB_PATH", "vocab.pt")
MAX_NEW_TOKENS = int(os.environ.get("MAX_NEW_TOKENS", "7"))
# ---------------------------------------------------------------------------
# Model architecture
# ---------------------------------------------------------------------------
class RotaryPositionalEmbedding(nn.Module):
def __init__(self, dim=256, max_seq_len=512):
def __init__(self, dim=32):
super().__init__()
inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2).float() / dim))
inv_freq = 1.0 / (10000 ** (torch.arange(0, dim).float() / dim))
self.register_buffer("inv_freq", inv_freq)
self.max_seq_len = max_seq_len
def forward(self, x):
seq_len = x.shape[1]
t = torch.arange(seq_len, device=x.device).float()
freqs = torch.einsum("i,j->ij", t, self.inv_freq)
emb = torch.cat((freqs, freqs), dim=-1)
return emb[:seq_len]
return freqs[:seq_len]
class MultiHeadAttentionWithRoPE(nn.Module):
@ -45,13 +46,13 @@ class MultiHeadAttentionWithRoPE(nn.Module):
self.rope = RotaryPositionalEmbedding(dim=self.d_head)
self.dropout = nn.Dropout(dropout)
def _apply_rope(self, x, freqs):
return x * freqs.cos() + self._rotate_half(x) * freqs.sin()
def _rotate_half(self, x):
x1, x2 = x.chunk(2, dim=-1)
return torch.cat((-x2, x1), dim=-1)
def _apply_rope(self, x, freqs):
return x * freqs.cos() + self._rotate_half(x) * freqs.sin()
def forward(self, x, mask=None):
B, T, C = x.shape
Q = self.W_q(x).view(B, T, self.n_heads, self.d_head).transpose(1, 2)
@ -100,7 +101,7 @@ class DecoderBlockWithRoPE(nn.Module):
class GPTModelWithRoPE(nn.Module):
def __init__(self, vocab_size=42962, d_model=256, n_heads=8,
n_layers=3, d_ff=1024, dropout=0.1, max_seq_len=512):
n_layers=3, d_ff=1024, dropout=0.1):
super().__init__()
self.token_embedding = nn.Embedding(vocab_size, d_model)
self.embedding_dropout = nn.Dropout(dropout)
@ -123,13 +124,20 @@ class GPTModelWithRoPE(nn.Module):
def generate(self, input_ids, max_new_tokens=7, temperature=1.0,
top_k=50, top_p=0.9, eos_token_id=None):
self.eval()
B, T = input_ids.shape
# Causal mask for current sequence length
mask = torch.tril(torch.ones(1, 1, T, T, device=input_ids.device))
for _ in range(max_new_tokens):
logits = self(input_ids)
logits = self(input_ids, mask)
logits = logits[:, -1, :]
if temperature > 0:
logits = logits / temperature
# Replace nan/inf with -inf to avoid crashes
logits = torch.nan_to_num(logits, nan=float("-inf"), posinf=float("-inf"), neginf=float("-inf"))
if top_k > 0:
values, _ = torch.topk(logits, top_k, dim=-1)
logits[logits < values[:, -1:]] = float("-inf")
@ -143,7 +151,11 @@ class GPTModelWithRoPE(nn.Module):
probs = F.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
# Update causal mask for new token
input_ids = torch.cat([input_ids, next_token], dim=-1)
new_T = input_ids.shape[1]
mask = torch.tril(torch.ones(1, 1, new_T, new_T, device=input_ids.device))
if eos_token_id is not None and next_token.item() == eos_token_id:
break
@ -151,6 +163,27 @@ class GPTModelWithRoPE(nn.Module):
return input_ids
# ---------------------------------------------------------------------------
# Vocab (in case no vocab file exists)
# ---------------------------------------------------------------------------
class DummyVocab:
"""Minimal vocab that passes through token IDs as text."""
word2idx = {'<BOS>': 0, '<EOS>': 1, '<PAD>': 2, '<UNK>': 3}
idx2word = {0: '<BOS>', 1: '<EOS>', 2: '<PAD>', 3: '<UNK>'}
def __init__(self):
for i in range(4, 42962):
self.word2idx[f'<TOK_{i}>'] = i
self.idx2word[i] = f'<TOK_{i}>'
def text_to_indices_for_generator(self, text):
return [self.word2idx.get(w, 3) for w in text.strip().split()]
def indices_to_text(self, indices):
return ' '.join(self.idx2word.get(i, '<UNK>') for i in indices)
# ---------------------------------------------------------------------------
# Generation function
# ---------------------------------------------------------------------------
@ -172,10 +205,16 @@ def generate_clickbait(model, vocab, prompt="", max_new_tokens=7,
temperature=temperature,
top_k=top_k,
top_p=top_p,
eos_token_id=vocab.word2idx['<EOS>']
eos_token_id=vocab.word2idx.get('<EOS>', 1)
)
generated_text = vocab.indices_to_text(generated[0].tolist())
# Clean up special tokens for display
for tok in ['<BOS>', '<EOS>', '<PAD>', '<UNK>']:
generated_text = generated_text.replace(tok, '')
generated_text = generated_text.strip()
return generated_text
@ -191,50 +230,67 @@ app.add_middleware(
allow_headers=["*"],
)
class GenerateRequest(BaseModel):
prompt: str = ""
temperature: float = 1.0
top_k: int = 50
num_samples: int = 1
class GenerateResponse(BaseModel):
texts: list[str]
model = None
vocab = None
using_dummy_vocab = False
@app.on_event("startup")
def load_artifacts():
global model, vocab
global model, vocab, using_dummy_vocab
if not os.path.exists(MODEL_PATH):
print(f"[WARN] {MODEL_PATH} not found. Model loading skipped.")
return
if not os.path.exists(VOCAB_PATH):
print(f"[WARN] {VOCAB_PATH} not found. Vocab loading skipped.")
return
try:
import pickle
checkpoint = torch.load(MODEL_PATH, map_location=DEVICE, weights_only=True)
vocab = pickle.load(open(VOCAB_PATH, "rb"))
model = GPTModelWithRoPE(
vocab_size=len(vocab.word2idx) if hasattr(vocab, 'word2idx') else 42962
).to(DEVICE)
if isinstance(checkpoint, dict) and 'model_state_dict' in checkpoint:
model.load_state_dict(checkpoint['model_state_dict'])
state_dict = checkpoint['model_state_dict']
else:
model.load_state_dict(checkpoint)
state_dict = checkpoint
vocab_size = state_dict['token_embedding.weight'].shape[0]
model = GPTModelWithRoPE(vocab_size=vocab_size).to(DEVICE)
# Remove causal_mask (we compute it dynamically)
sd_clean = {k: v for k, v in state_dict.items() if k != 'causal_mask'}
missing, unexpected = model.load_state_dict(sd_clean, strict=False)
if missing:
print(f"[WARN] Missing keys: {missing}")
if unexpected:
print(f"[WARN] Unexpected keys: {unexpected}")
model.eval()
print(f"[OK] Model loaded on {DEVICE}")
print(f"[OK] Vocab size: {len(vocab.word2idx)}")
# Load vocab
if os.path.exists(VOCAB_PATH):
vocab = pickle.load(open(VOCAB_PATH, "rb"))
using_dummy_vocab = False
print(f"[OK] Vocab loaded: {len(vocab.word2idx)} tokens")
else:
print(f"[WARN] {VOCAB_PATH} not found, using fallback dummy vocab")
vocab = DummyVocab()
using_dummy_vocab = True
print(f"[OK] Model loaded on {DEVICE} | params: {sum(p.numel() for p in model.parameters()):,}")
except Exception as e:
print(f"[ERROR] Failed to load artifacts: {e}")
model = None
vocab = None
print(f"[ERROR] Failed to load model: {e}")
import traceback
traceback.print_exc()
@app.get("/api/health")
@ -243,28 +299,55 @@ def health():
"status": "ok",
"model_loaded": model is not None,
"vocab_loaded": vocab is not None,
"using_dummy_vocab": using_dummy_vocab,
"device": str(DEVICE),
"vocab_size": len(vocab.word2idx) if vocab else 0,
"note": "vocab.pt not found — output shows token IDs" if using_dummy_vocab else "ready",
}
@app.post("/api/generate", response_model=GenerateResponse)
def generate(req: GenerateRequest):
if model is None or vocab is None:
raise HTTPException(status_code=503, detail="Model or vocab not loaded")
if model is None:
raise HTTPException(status_code=503, detail="Model not loaded")
# Clamp / validate
temperature = max(0.1, min(3.0, req.temperature))
top_k = max(1, min(200, req.top_k))
count = max(1, min(20, req.num_samples))
texts = []
for _ in range(req.num_samples):
for _ in range(count):
text = generate_clickbait(
model, vocab,
model, vocab or DummyVocab(),
prompt=req.prompt,
temperature=req.temperature,
top_k=req.top_k,
temperature=temperature,
top_k=top_k,
max_new_tokens=MAX_NEW_TOKENS,
)
texts.append(text)
return GenerateResponse(texts=texts)
# ---------------------------------------------------------------------------
# Static files + SPA fallback
# ---------------------------------------------------------------------------
STATIC_DIR = os.path.dirname(os.path.abspath(__file__))
@app.get("/")
def serve_index():
return FileResponse(os.path.join(STATIC_DIR, "index.html"))
@app.get("/{path:path}")
def serve_static(path: str):
file_path = os.path.join(STATIC_DIR, path)
if os.path.isfile(file_path):
return FileResponse(file_path)
return FileResponse(os.path.join(STATIC_DIR, "index.html"))
# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------

View File

@ -258,6 +258,50 @@ body {
max-width: 700px;
}
.api-config {
display: flex;
align-items: center;
gap: 12px;
margin-bottom: 20px;
padding-bottom: 16px;
border-bottom: 1px solid rgba(255,255,255,0.08);
}
.api-label {
font-size: 12px;
font-weight: 700;
color: var(--gray-500);
text-transform: uppercase;
letter-spacing: 0.5px;
white-space: nowrap;
}
.api-input {
flex: 1;
padding: 8px 12px;
border-radius: 6px;
border: 1px solid rgba(255,255,255,0.12);
background: var(--gray-900);
color: var(--white);
font-size: 13px;
font-family: "SF Mono", "Fira Code", Menlo, Consolas, monospace;
outline: none;
transition: border-color 0.2s;
}
.api-input:focus {
border-color: var(--cyan);
}
.api-status {
font-size: 12px;
white-space: nowrap;
}
.api-status.ok { color: #50fa7b; }
.api-status.warn { color: #f1fa8c; }
.api-status.err { color: #ff6b6b; }
.gen-row {
margin-bottom: 20px;
}
@ -524,6 +568,15 @@ body {
gap: 16px;
}
.api-config {
flex-wrap: wrap;
}
.api-status {
width: 100%;
margin-top: -4px;
}
.nav-links {
display: none;
}