Compare commits

..

2 Commits

Author SHA1 Message Date
871147b688 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)
2026-06-20 00:03:52 +00:00
2c520a96c0 fix: local generation, regenerate button, API via ?api= 2026-06-19 06:28:35 +00:00
6 changed files with 326 additions and 128 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` ## ✅ ВСЕГДА работай через `./new-project`
Если юзер сказал «сделай сайт NAME» / «создай проект NAME»: Если юзер сказал «сделай сайт 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
Рядом лежит `design.md` с готовой палитрой, типографикой и стартер-шаблоном `index.html`. **Начинай с него.** Не выдумывай новые цвета — модифицируй существующие. Рядом лежит `design.md` с готовой палитрой, типографикой и стартер-шаблоном `index.html`. **Начинай с него.** Не выдумывай новые цвета — модифицируй существующие.

View File

@ -69,6 +69,11 @@
<h2>Генератор кликбейта</h2> <h2>Генератор кликбейта</h2>
<p class="section-desc">Введите начало заголовка — модель продолжит в стиле кликбейта.</p> <p class="section-desc">Введите начало заголовка — модель продолжит в стиле кликбейта.</p>
<div class="generator-box"> <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"> <div class="gen-row">
<input type="text" id="promptInput" class="gen-input" placeholder="Например: почему, как, топ 10..." autocomplete="off"> <input type="text" id="promptInput" class="gen-input" placeholder="Например: почему, как, топ 10..." autocomplete="off">
</div> </div>

202
script.js
View File

@ -1,4 +1,14 @@
const API_URL = window.location.origin + '/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 promptInput = document.getElementById('promptInput');
const tempSlider = document.getElementById('tempSlider'); const tempSlider = document.getElementById('tempSlider');
@ -9,95 +19,157 @@ const generateBtn = document.getElementById('generateBtn');
const generate5Btn = document.getElementById('generate5Btn'); const generate5Btn = document.getElementById('generate5Btn');
const outputArea = document.getElementById('outputArea'); const outputArea = document.getElementById('outputArea');
const statusMsg = document.getElementById('statusMsg'); const statusMsg = document.getElementById('statusMsg');
const apiInput = document.getElementById('apiUrlInput');
const apiStatus = document.getElementById('apiStatus');
tempSlider.addEventListener('input', () => { let backendOnline = false;
tempVal.textContent = tempSlider.value;
});
topkSlider.addEventListener('input', () => { async function checkBackend() {
topkVal.textContent = topkSlider.value; 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; });
generateBtn.addEventListener('click', () => generate(1)); generateBtn.addEventListener('click', () => generate(1));
generate5Btn.addEventListener('click', () => generate(5)); generate5Btn.addEventListener('click', () => generate(5));
// ---------------------------------------------------------------------------
// Client-side clickbait generator (fallback / default)
// ---------------------------------------------------------------------------
const TEMPLATES = [
'этот секрет знают единицы',
'врачи в шоке от этого открытия',
'ты не поверишь что произошло дальше',
'всего один ингредиент меняет всё',
'вот почему никто не говорит об этом',
'результат превзошёл все ожидания',
'это должен знать каждый',
'никто не ожидал такого поворота',
'учёные озадачены этим феноменом',
'это изменит вашу жизнь навсегда',
'вы будете в шоке от правды',
'самое важное о чём молчат СМИ',
'это открытие перевернуло науку',
'простые вещи которые творят чудеса',
'теперь это знает весь мир',
'о чём вам не расскажут в школе',
'главный секрет успешных людей',
'это работает безотказно',
'вы упускаете это каждый день',
'пора узнать правду'
];
const PREFIXES = ['', 'невероятно но ', 'шокирующе но ', 'оказывается ', 'представьте себе '];
function pick(arr, seed) {
const idx = Math.abs(seed * 7 + seed * seed * 31) % arr.length;
return arr[idx];
}
function localGenerate(prompt, temperature, top_k, count) {
const now = Date.now();
const results = [];
for (let i = 0; i < count; i++) {
const seed = now + i * 137 + Math.floor(temperature * 100) + top_k;
const prefix = pick(PREFIXES, seed + 3);
const template = prompt
? `${prompt}${pick(TEMPLATES, seed + 7)}`
: pick(TEMPLATES, seed + 11);
const extra = temperature > 1.3
? (pick(['невероятно', 'шокирующе', 'фантастически'], seed + 19) + ' ')
: '';
results.push(prefix + extra + template);
}
return results;
}
// ---------------------------------------------------------------------------
// Main generation
// ---------------------------------------------------------------------------
async function generate(count) { async function generate(count) {
const prompt = promptInput.value.trim(); const prompt = promptInput.value.trim();
const temperature = parseFloat(tempSlider.value); const temperature = parseFloat(tempSlider.value);
const top_k = parseInt(topkSlider.value); const top_k = parseInt(topkSlider.value);
outputArea.innerHTML = ''; outputArea.innerHTML = '';
statusMsg.textContent = 'Генерация...'; statusMsg.textContent = '⚡ Генерируем...';
statusMsg.className = 'status-msg'; statusMsg.className = 'status-msg';
let texts;
let usedBackend = false;
try { try {
const res = await fetch(`${API_URL}/generate`, { const res = await fetch(`${API_URL}/api/generate`, {
method: 'POST', method: 'POST',
headers: { 'Content-Type': 'application/json' }, headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ body: JSON.stringify({ prompt, temperature, top_k, num_samples: count })
prompt,
temperature,
top_k,
num_samples: count
})
}); });
if (!res.ok) throw new Error(`HTTP ${res.status}`);
if (!res.ok) {
const err = await res.json().catch(() => ({}));
throw new Error(err.detail || `Ошибка ${res.status}`);
}
const data = await res.json(); const data = await res.json();
statusMsg.textContent = ''; texts = data.texts;
statusMsg.className = 'status-msg'; usedBackend = true;
} catch {
data.texts.forEach((text, i) => { texts = localGenerate(prompt, temperature, top_k, count);
const div = document.createElement('div'); usedBackend = false;
div.className = 'gen-output-item';
if (count > 1) {
div.innerHTML = `<span class="gen-badge">#${i + 1}</span>${text}...😱`;
} else {
div.textContent = `${text}...😱`;
}
outputArea.appendChild(div);
});
} catch (err) {
statusMsg.textContent = '❌ Сервер недоступен. Запустите API: python server.py';
statusMsg.className = 'status-msg error';
fallbackGenerate(count, prompt, temperature, top_k);
} }
}
function fallbackGenerate(count, prompt, temperature, top_k) { statusMsg.textContent = usedBackend ? '✅ Модель (PyTorch)' : '⚡ Локальные шаблоны (бэкенд недоступен)';
const templates = [ statusMsg.className = 'status-msg';
prompt
? `${prompt} — этот секрет знают единицы`
: 'Этот секрет знают единицы',
prompt
? `${prompt} — врачи в шоке от открытия`
: 'Врачи в шоке от этого открытия',
prompt
? `${prompt} — ты не поверишь что произошло`
: 'Ты не поверишь что произошло дальше',
prompt
? `${prompt} — всего один ингредиент меняет всё`
: 'Всего один ингредиент меняет всё',
prompt
? `${prompt} — вот почему никто не говорит об этом`
: 'Вот почему никто не говорит об этом'
];
outputArea.innerHTML = ''; texts.forEach((text, i) => {
for (let i = 0; i < count && i < templates.length; i++) {
const div = document.createElement('div'); const div = document.createElement('div');
div.className = 'gen-output-item'; div.className = 'gen-output-item';
if (count > 1) { div.style.animationDelay = `${i * 0.08}s`;
div.innerHTML = `<span class="gen-badge">#${i + 1}</span>${templates[i]}...😱`; const label = count > 1 ? `<span class="gen-badge">#${i + 1}</span>` : '';
} else { div.innerHTML = `${label}${text}...😱`;
div.textContent = `${templates[i]}...😱`;
}
outputArea.appendChild(div); outputArea.appendChild(div);
} });
// Regenerate button
const reBtn = document.createElement('button');
reBtn.className = 'btn btn-secondary btn-regenerate';
reBtn.textContent = '🔄 Сгенерировать ещё';
reBtn.addEventListener('click', () => generate(count));
outputArea.appendChild(reBtn);
} }
// Example tabs // Example tabs

153
server.py
View File

@ -3,30 +3,31 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from fastapi import FastAPI, HTTPException from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from fastapi.responses import FileResponse
from pydantic import BaseModel from pydantic import BaseModel
import os import os, pickle, sys
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
MODEL_PATH = os.environ.get("MODEL_PATH", "best_generator.pt") MODEL_PATH = os.environ.get("MODEL_PATH", "best_generator.pt")
VOCAB_PATH = os.environ.get("VOCAB_PATH", "vocab.pt") VOCAB_PATH = os.environ.get("VOCAB_PATH", "vocab.pt")
MAX_NEW_TOKENS = int(os.environ.get("MAX_NEW_TOKENS", "7"))
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Model architecture # Model architecture
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
class RotaryPositionalEmbedding(nn.Module): class RotaryPositionalEmbedding(nn.Module):
def __init__(self, dim=256, max_seq_len=512): def __init__(self, dim=32):
super().__init__() 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.register_buffer("inv_freq", inv_freq)
self.max_seq_len = max_seq_len
def forward(self, x): def forward(self, x):
seq_len = x.shape[1] seq_len = x.shape[1]
t = torch.arange(seq_len, device=x.device).float() t = torch.arange(seq_len, device=x.device).float()
freqs = torch.einsum("i,j->ij", t, self.inv_freq) freqs = torch.einsum("i,j->ij", t, self.inv_freq)
emb = torch.cat((freqs, freqs), dim=-1) return freqs[:seq_len]
return emb[:seq_len]
class MultiHeadAttentionWithRoPE(nn.Module): class MultiHeadAttentionWithRoPE(nn.Module):
@ -45,13 +46,13 @@ class MultiHeadAttentionWithRoPE(nn.Module):
self.rope = RotaryPositionalEmbedding(dim=self.d_head) self.rope = RotaryPositionalEmbedding(dim=self.d_head)
self.dropout = nn.Dropout(dropout) 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): def _rotate_half(self, x):
x1, x2 = x.chunk(2, dim=-1) x1, x2 = x.chunk(2, dim=-1)
return torch.cat((-x2, x1), 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): def forward(self, x, mask=None):
B, T, C = x.shape B, T, C = x.shape
Q = self.W_q(x).view(B, T, self.n_heads, self.d_head).transpose(1, 2) 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): class GPTModelWithRoPE(nn.Module):
def __init__(self, vocab_size=42962, d_model=256, n_heads=8, 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__() super().__init__()
self.token_embedding = nn.Embedding(vocab_size, d_model) self.token_embedding = nn.Embedding(vocab_size, d_model)
self.embedding_dropout = nn.Dropout(dropout) 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, def generate(self, input_ids, max_new_tokens=7, temperature=1.0,
top_k=50, top_p=0.9, eos_token_id=None): top_k=50, top_p=0.9, eos_token_id=None):
self.eval() 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): for _ in range(max_new_tokens):
logits = self(input_ids) logits = self(input_ids, mask)
logits = logits[:, -1, :] logits = logits[:, -1, :]
if temperature > 0: if temperature > 0:
logits = logits / temperature 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: if top_k > 0:
values, _ = torch.topk(logits, top_k, dim=-1) values, _ = torch.topk(logits, top_k, dim=-1)
logits[logits < values[:, -1:]] = float("-inf") logits[logits < values[:, -1:]] = float("-inf")
@ -143,7 +151,11 @@ class GPTModelWithRoPE(nn.Module):
probs = F.softmax(logits, dim=-1) probs = F.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=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) 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: if eos_token_id is not None and next_token.item() == eos_token_id:
break break
@ -151,6 +163,27 @@ class GPTModelWithRoPE(nn.Module):
return input_ids 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 # Generation function
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@ -172,10 +205,16 @@ def generate_clickbait(model, vocab, prompt="", max_new_tokens=7,
temperature=temperature, temperature=temperature,
top_k=top_k, top_k=top_k,
top_p=top_p, 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()) 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 return generated_text
@ -191,50 +230,67 @@ app.add_middleware(
allow_headers=["*"], allow_headers=["*"],
) )
class GenerateRequest(BaseModel): class GenerateRequest(BaseModel):
prompt: str = "" prompt: str = ""
temperature: float = 1.0 temperature: float = 1.0
top_k: int = 50 top_k: int = 50
num_samples: int = 1 num_samples: int = 1
class GenerateResponse(BaseModel): class GenerateResponse(BaseModel):
texts: list[str] texts: list[str]
model = None model = None
vocab = None vocab = None
using_dummy_vocab = False
@app.on_event("startup") @app.on_event("startup")
def load_artifacts(): def load_artifacts():
global model, vocab global model, vocab, using_dummy_vocab
if not os.path.exists(MODEL_PATH): if not os.path.exists(MODEL_PATH):
print(f"[WARN] {MODEL_PATH} not found. Model loading skipped.") print(f"[WARN] {MODEL_PATH} not found. Model loading skipped.")
return return
if not os.path.exists(VOCAB_PATH):
print(f"[WARN] {VOCAB_PATH} not found. Vocab loading skipped.")
return
try: try:
import pickle
checkpoint = torch.load(MODEL_PATH, map_location=DEVICE, weights_only=True) 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: 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: 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() 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: except Exception as e:
print(f"[ERROR] Failed to load artifacts: {e}") print(f"[ERROR] Failed to load model: {e}")
model = None import traceback
vocab = None traceback.print_exc()
@app.get("/api/health") @app.get("/api/health")
@ -243,28 +299,55 @@ def health():
"status": "ok", "status": "ok",
"model_loaded": model is not None, "model_loaded": model is not None,
"vocab_loaded": vocab is not None, "vocab_loaded": vocab is not None,
"using_dummy_vocab": using_dummy_vocab,
"device": str(DEVICE), "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) @app.post("/api/generate", response_model=GenerateResponse)
def generate(req: GenerateRequest): def generate(req: GenerateRequest):
if model is None or vocab is None: if model is None:
raise HTTPException(status_code=503, detail="Model or vocab not loaded") 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 = [] texts = []
for _ in range(req.num_samples): for _ in range(count):
text = generate_clickbait( text = generate_clickbait(
model, vocab, model, vocab or DummyVocab(),
prompt=req.prompt, prompt=req.prompt,
temperature=req.temperature, temperature=temperature,
top_k=req.top_k, top_k=top_k,
max_new_tokens=MAX_NEW_TOKENS,
) )
texts.append(text) texts.append(text)
return GenerateResponse(texts=texts) 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 # Entry point
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------

View File

@ -127,6 +127,17 @@ body {
color: var(--cyan); color: var(--cyan);
} }
.btn-regenerate {
margin-top: 16px;
font-size: 14px;
padding: 10px 20px;
align-self: flex-start;
}
.gen-output .btn-regenerate {
align-self: flex-start;
}
/* Sections */ /* Sections */
.section { .section {
padding: 80px 0; padding: 80px 0;
@ -247,6 +258,50 @@ body {
max-width: 700px; 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 { .gen-row {
margin-bottom: 20px; margin-bottom: 20px;
} }
@ -513,6 +568,15 @@ body {
gap: 16px; gap: 16px;
} }
.api-config {
flex-wrap: wrap;
}
.api-status {
width: 100%;
margin-top: -4px;
}
.nav-links { .nav-links {
display: none; display: none;
} }