Compare commits
No commits in common. "871147b688893d8fc8fd2b94713a00ca3ba1d335" and "43d99a94687b322982abeec8226d34a0583e7729" have entirely different histories.
871147b688
...
43d99a9468
2
.gitignore
vendored
2
.gitignore
vendored
@ -1,2 +0,0 @@
|
|||||||
__pycache__/
|
|
||||||
*.pyc
|
|
||||||
28
AGENTS.md
28
AGENTS.md
@ -95,6 +95,17 @@ 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»:
|
||||||
@ -140,6 +151,23 @@ 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`. **Начинай с него.** Не выдумывай новые цвета — модифицируй существующие.
|
||||||
|
|||||||
@ -69,11 +69,6 @@
|
|||||||
<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>
|
||||||
|
|||||||
196
script.js
196
script.js
@ -1,14 +1,4 @@
|
|||||||
// API URL: ?api=... в URL, или localStorage, или auto-detect (same host :8000)
|
const API_URL = window.location.origin + '/api';
|
||||||
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');
|
||||||
@ -19,157 +9,95 @@ 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');
|
|
||||||
|
|
||||||
let backendOnline = false;
|
tempSlider.addEventListener('input', () => {
|
||||||
|
tempVal.textContent = tempSlider.value;
|
||||||
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();
|
topkSlider.addEventListener('input', () => {
|
||||||
setInterval(checkBackend, 15000);
|
topkVal.textContent = topkSlider.value;
|
||||||
|
});
|
||||||
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}/api/generate`, {
|
const res = await fetch(`${API_URL}/generate`, {
|
||||||
method: 'POST',
|
method: 'POST',
|
||||||
headers: { 'Content-Type': 'application/json' },
|
headers: { 'Content-Type': 'application/json' },
|
||||||
body: JSON.stringify({ prompt, temperature, top_k, num_samples: count })
|
body: JSON.stringify({
|
||||||
|
prompt,
|
||||||
|
temperature,
|
||||||
|
top_k,
|
||||||
|
num_samples: count
|
||||||
|
})
|
||||||
});
|
});
|
||||||
if (!res.ok) throw new Error(`HTTP ${res.status}`);
|
|
||||||
const data = await res.json();
|
if (!res.ok) {
|
||||||
texts = data.texts;
|
const err = await res.json().catch(() => ({}));
|
||||||
usedBackend = true;
|
throw new Error(err.detail || `Ошибка ${res.status}`);
|
||||||
} catch {
|
|
||||||
texts = localGenerate(prompt, temperature, top_k, count);
|
|
||||||
usedBackend = false;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
statusMsg.textContent = usedBackend ? '✅ Модель (PyTorch)' : '⚡ Локальные шаблоны (бэкенд недоступен)';
|
const data = await res.json();
|
||||||
|
statusMsg.textContent = '';
|
||||||
statusMsg.className = 'status-msg';
|
statusMsg.className = 'status-msg';
|
||||||
|
|
||||||
texts.forEach((text, i) => {
|
data.texts.forEach((text, i) => {
|
||||||
const div = document.createElement('div');
|
const div = document.createElement('div');
|
||||||
div.className = 'gen-output-item';
|
div.className = 'gen-output-item';
|
||||||
div.style.animationDelay = `${i * 0.08}s`;
|
if (count > 1) {
|
||||||
const label = count > 1 ? `<span class="gen-badge">#${i + 1}</span>` : '';
|
div.innerHTML = `<span class="gen-badge">#${i + 1}</span>${text}...😱`;
|
||||||
div.innerHTML = `${label}${text}...😱`;
|
} else {
|
||||||
|
div.textContent = `${text}...😱`;
|
||||||
|
}
|
||||||
outputArea.appendChild(div);
|
outputArea.appendChild(div);
|
||||||
});
|
});
|
||||||
|
} catch (err) {
|
||||||
|
statusMsg.textContent = '❌ Сервер недоступен. Запустите API: python server.py';
|
||||||
|
statusMsg.className = 'status-msg error';
|
||||||
|
fallbackGenerate(count, prompt, temperature, top_k);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Regenerate button
|
function fallbackGenerate(count, prompt, temperature, top_k) {
|
||||||
const reBtn = document.createElement('button');
|
const templates = [
|
||||||
reBtn.className = 'btn btn-secondary btn-regenerate';
|
prompt
|
||||||
reBtn.textContent = '🔄 Сгенерировать ещё';
|
? `${prompt} — этот секрет знают единицы`
|
||||||
reBtn.addEventListener('click', () => generate(count));
|
: 'Этот секрет знают единицы',
|
||||||
outputArea.appendChild(reBtn);
|
prompt
|
||||||
|
? `${prompt} — врачи в шоке от открытия`
|
||||||
|
: 'Врачи в шоке от этого открытия',
|
||||||
|
prompt
|
||||||
|
? `${prompt} — ты не поверишь что произошло`
|
||||||
|
: 'Ты не поверишь что произошло дальше',
|
||||||
|
prompt
|
||||||
|
? `${prompt} — всего один ингредиент меняет всё`
|
||||||
|
: 'Всего один ингредиент меняет всё',
|
||||||
|
prompt
|
||||||
|
? `${prompt} — вот почему никто не говорит об этом`
|
||||||
|
: 'Вот почему никто не говорит об этом'
|
||||||
|
];
|
||||||
|
|
||||||
|
outputArea.innerHTML = '';
|
||||||
|
for (let i = 0; i < count && i < templates.length; i++) {
|
||||||
|
const div = document.createElement('div');
|
||||||
|
div.className = 'gen-output-item';
|
||||||
|
if (count > 1) {
|
||||||
|
div.innerHTML = `<span class="gen-badge">#${i + 1}</span>${templates[i]}...😱`;
|
||||||
|
} else {
|
||||||
|
div.textContent = `${templates[i]}...😱`;
|
||||||
|
}
|
||||||
|
outputArea.appendChild(div);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Example tabs
|
// Example tabs
|
||||||
|
|||||||
153
server.py
153
server.py
@ -3,31 +3,30 @@ 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, pickle, sys
|
import os
|
||||||
|
|
||||||
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=32):
|
def __init__(self, dim=256, max_seq_len=512):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
inv_freq = 1.0 / (10000 ** (torch.arange(0, dim).float() / dim))
|
inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2).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)
|
||||||
return freqs[:seq_len]
|
emb = torch.cat((freqs, freqs), dim=-1)
|
||||||
|
return emb[:seq_len]
|
||||||
|
|
||||||
|
|
||||||
class MultiHeadAttentionWithRoPE(nn.Module):
|
class MultiHeadAttentionWithRoPE(nn.Module):
|
||||||
@ -46,13 +45,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)
|
||||||
@ -101,7 +100,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):
|
n_layers=3, d_ff=1024, dropout=0.1, max_seq_len=512):
|
||||||
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)
|
||||||
@ -124,20 +123,13 @@ 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, mask)
|
logits = self(input_ids)
|
||||||
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")
|
||||||
@ -151,11 +143,7 @@ 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
|
||||||
@ -163,27 +151,6 @@ 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
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@ -205,16 +172,10 @@ 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.get('<EOS>', 1)
|
eos_token_id=vocab.word2idx['<EOS>']
|
||||||
)
|
)
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
|
|
||||||
@ -230,67 +191,50 @@ 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, using_dummy_vocab
|
global model, 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:
|
||||||
state_dict = checkpoint['model_state_dict']
|
model.load_state_dict(checkpoint['model_state_dict'])
|
||||||
else:
|
else:
|
||||||
state_dict = checkpoint
|
model.load_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}")
|
||||||
# Load vocab
|
print(f"[OK] Vocab size: {len(vocab.word2idx)}")
|
||||||
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 model: {e}")
|
print(f"[ERROR] Failed to load artifacts: {e}")
|
||||||
import traceback
|
model = None
|
||||||
traceback.print_exc()
|
vocab = None
|
||||||
|
|
||||||
|
|
||||||
@app.get("/api/health")
|
@app.get("/api/health")
|
||||||
@ -299,55 +243,28 @@ 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:
|
if model is None or vocab is None:
|
||||||
raise HTTPException(status_code=503, detail="Model not loaded")
|
raise HTTPException(status_code=503, detail="Model or vocab 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(count):
|
for _ in range(req.num_samples):
|
||||||
text = generate_clickbait(
|
text = generate_clickbait(
|
||||||
model, vocab or DummyVocab(),
|
model, vocab,
|
||||||
prompt=req.prompt,
|
prompt=req.prompt,
|
||||||
temperature=temperature,
|
temperature=req.temperature,
|
||||||
top_k=top_k,
|
top_k=req.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
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
64
style.css
64
style.css
@ -127,17 +127,6 @@ 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;
|
||||||
@ -258,50 +247,6 @@ 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;
|
||||||
}
|
}
|
||||||
@ -568,15 +513,6 @@ 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;
|
||||||
}
|
}
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user