File size: 11,603 Bytes
8d26565
 
 
 
 
 
 
 
 
 
81d11e3
d888fc6
 
 
8d26565
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d888fc6
 
 
 
 
 
 
 
 
 
 
 
 
8d26565
d888fc6
8d26565
 
 
 
 
 
 
 
d888fc6
 
 
 
 
8d26565
 
 
 
d888fc6
 
 
 
 
 
8d26565
 
 
d888fc6
8d26565
 
 
 
 
 
6dc26fe
8d26565
 
6dc26fe
 
d888fc6
 
 
56ed8a9
81d11e3
 
 
 
 
 
 
56ed8a9
 
 
 
 
 
 
 
 
 
 
81d11e3
 
 
 
 
 
 
 
 
 
 
56ed8a9
 
d888fc6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
56ed8a9
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
# ================================================================
# MTP - app.py para Hugging Face Space (Gradio, CPU)
# Carga el checkpoint MTP_MODEL.pt desde el repo TeszenAI/MTP-1
# ================================================================
import os
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import gradio as gr
from starlette.middleware import Middleware
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from typing import Optional
from huggingface_hub import hf_hub_download

# ---------------- Optimización para CPU ----------------
# Limita hilos a los núcleos disponibles (evita overhead en Spaces pequeños)
torch.set_num_threads(max(1, os.cpu_count() or 1))
torch.set_grad_enabled(False)  # solo inferencia, nunca necesitamos gradientes

DEVICE = "cpu"

REPO_ID = "TeszenAI/MTP-1"
FILENAME = "MTP_MODEL.pt"

# ---------------- Arquitectura (idéntica a la de entrenamiento) ----------------
class CausalSelfAttention(nn.Module):
    def __init__(self, n_embd, n_head, block_size, dropout):
        super().__init__()
        self.n_head = n_head
        self.head_dim = n_embd // n_head
        self.qkv = nn.Linear(n_embd, 3 * n_embd)
        self.proj = nn.Linear(n_embd, n_embd)
        self.attn_dropout = nn.Dropout(dropout)
        self.resid_dropout = nn.Dropout(dropout)
        mask = torch.tril(torch.ones(block_size, block_size)).view(1, 1, block_size, block_size)
        self.register_buffer("mask", mask)

    def forward(self, x):
        B, T, C = x.shape
        qkv = self.qkv(x)
        q, k, v = qkv.split(C, dim=2)
        q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
        k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
        v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
        att = (q @ k.transpose(-2, -1)) / math.sqrt(self.head_dim)
        att = att.masked_fill(self.mask[:, :, :T, :T] == 0, float("-inf"))
        att = F.softmax(att, dim=-1)
        att = self.attn_dropout(att)
        out = (att @ v).transpose(1, 2).contiguous().view(B, T, C)
        return self.resid_dropout(self.proj(out))


class FeedForward(nn.Module):
    def __init__(self, n_embd, dropout):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(n_embd, 4 * n_embd), nn.GELU(),
            nn.Linear(4 * n_embd, n_embd), nn.Dropout(dropout),
        )

    def forward(self, x):
        return self.net(x)


class Block(nn.Module):
    def __init__(self, n_embd, n_head, block_size, dropout):
        super().__init__()
        self.ln1 = nn.LayerNorm(n_embd)
        self.attn = CausalSelfAttention(n_embd, n_head, block_size, dropout)
        self.ln2 = nn.LayerNorm(n_embd)
        self.ff = FeedForward(n_embd, dropout)

    def forward(self, x):
        x = x + self.attn(self.ln1(x))
        x = x + self.ff(self.ln2(x))
        return x


class MTP(nn.Module):
    def __init__(self, vocab_size, block_size, n_layer, n_head, n_embd, dropout):
        super().__init__()
        self.block_size = block_size
        self.tok_emb = nn.Embedding(vocab_size, n_embd)
        self.pos_emb = nn.Embedding(block_size, n_embd)
        self.drop = nn.Dropout(dropout)
        self.blocks = nn.ModuleList([Block(n_embd, n_head, block_size, dropout) for _ in range(n_layer)])
        self.ln_f = nn.LayerNorm(n_embd)
        self.lm_head = nn.Linear(n_embd, vocab_size, bias=False)
        self.lm_head.weight = self.tok_emb.weight

    def forward(self, idx):
        B, T = idx.shape
        pos = torch.arange(T, device=idx.device)
        x = self.tok_emb(idx) + self.pos_emb(pos)
        x = self.drop(x)
        for block in self.blocks:
            x = block(x)
        x = self.ln_f(x)
        return self.lm_head(x)


# ---------------- Carga del checkpoint (una sola vez, al iniciar el Space) ----------------
print("Descargando checkpoint desde el Hub...")
ckpt_path = hf_hub_download(repo_id=REPO_ID, filename=FILENAME)
checkpoint = torch.load(ckpt_path, map_location=DEVICE)

cfg = checkpoint["config"]
stoi = checkpoint["stoi"]
itos = {int(k): v for k, v in checkpoint["itos"].items()}
special = checkpoint["special_tokens"]
gen_defaults = checkpoint["generation_defaults"]

PAD_ID, BOS_ID, EOS_ID, UNK_ID = special["pad_id"], special["bos_id"], special["eos_id"], special["unk_id"]

model = MTP(
    vocab_size=cfg["vocab_size"], block_size=cfg["block_size"],
    n_layer=cfg["n_layer"], n_head=cfg["n_head"],
    n_embd=cfg["n_embd"], dropout=cfg["dropout"],
).to(DEVICE)
model.load_state_dict(checkpoint["model_state_dict"])
model.eval()

# fusiona LayerNorm/Linear estáticamente no aplica aquí, pero fija modo eval
# y evita cualquier dropout durante inferencia.
BLOCK_SIZE = cfg["block_size"]

print(f"MTP cargado ({checkpoint['meta']['model_name']}, "
      f"entrenado con {checkpoint['meta']['trained_examples']} ejemplos)")


def encode_text(s):
    return [stoi.get(ch, UNK_ID) for ch in s]


def decode_ids(ids):
    return "".join(itos.get(i, "") for i in ids if i not in (PAD_ID, BOS_ID, EOS_ID))


# ---------------- Generación ----------------
@torch.inference_mode()
def generate(idx, max_new_tokens, temperature, top_k, top_p, repetition_penalty):
    for _ in range(max_new_tokens):
        idx_cond = idx[:, -BLOCK_SIZE:]
        logits = model(idx_cond)
        logits = logits[:, -1, :] / max(temperature, 1e-5)

        if repetition_penalty and repetition_penalty != 1.0:
            for token_id in set(idx[0].tolist()):
                logits[0, token_id] /= repetition_penalty

        if top_k is not None and top_k > 0:
            v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
            logits[logits < v[:, [-1]]] = float("-inf")

        probs = F.softmax(logits, dim=-1)

        if top_p is not None and 0 < top_p < 1:
            sorted_probs, sorted_idx = torch.sort(probs, descending=True)
            cum_probs = torch.cumsum(sorted_probs, dim=-1)
            cutoff = cum_probs > top_p
            cutoff[:, 1:] = cutoff[:, :-1].clone()
            cutoff[:, 0] = False
            sorted_probs[cutoff] = 0.0
            sorted_probs = sorted_probs / sorted_probs.sum(dim=-1, keepdim=True)
            next_id = sorted_idx.gather(-1, torch.multinomial(sorted_probs, 1))
        else:
            next_id = torch.multinomial(probs, num_samples=1)

        idx = torch.cat([idx, next_id], dim=1)
        if next_id.item() == EOS_ID:
            break
    return idx


def run_inference(text, max_new_tokens=None, temperature=None, top_k=None, top_p=None, repetition_penalty=None):
    """Núcleo de generación, reutilizado por la UI de Gradio y por la API /generate.
    No reduce calidad por estar en CPU: usa exactamente el mismo muestreo
    (top_k + top_p + repetition_penalty) que en la Celda 2 de entrenamiento,
    solo que tarda más en devolver el resultado."""
    max_new_tokens = int(max_new_tokens) if max_new_tokens else gen_defaults["max_new_tokens"]
    temperature = float(temperature) if temperature is not None else gen_defaults["temperature"]
    top_k = int(top_k) if top_k is not None else gen_defaults["top_k"]
    top_p = float(top_p) if top_p is not None else gen_defaults["top_p"]
    repetition_penalty = float(repetition_penalty) if repetition_penalty is not None else gen_defaults["repetition_penalty"]

    # límite de seguridad para no colgar el Space con peticiones abusivas
    max_new_tokens = max(1, min(max_new_tokens, 500))

    prefix = f"Usuario: {text}\nMTP: "
    ids = [BOS_ID] + encode_text(prefix)
    idx = torch.tensor([ids], dtype=torch.long, device=DEVICE)

    out = generate(idx, max_new_tokens, temperature, top_k, top_p, repetition_penalty)
    new_ids = out[0].tolist()[len(ids):]
    return decode_ids(new_ids).strip()


def chat_fn(message, history, max_new_tokens, temperature, top_k, top_p, repetition_penalty):
    return run_inference(message, max_new_tokens, temperature, top_k, top_p, repetition_penalty)


# ---------------- Interfaz Gradio (para probar el modelo desde el navegador) ----------------
with gr.Blocks(title="MTP Chat") as demo:
    gr.Markdown("# MTP\nModelo GPT entrenado desde cero (char-level). Ejecutándose en CPU.")

    with gr.Accordion("Parámetros de generación", open=False):
        max_new_tokens_ui = gr.Slider(16, 400, value=gen_defaults["max_new_tokens"], step=1, label="max_new_tokens")
        temperature_ui = gr.Slider(0.1, 2.0, value=gen_defaults["temperature"], step=0.05, label="temperature")
        top_k_ui = gr.Slider(0, 100, value=gen_defaults["top_k"], step=1, label="top_k")
        top_p_ui = gr.Slider(0.1, 1.0, value=gen_defaults["top_p"], step=0.05, label="top_p")
        repetition_penalty_ui = gr.Slider(1.0, 2.0, value=gen_defaults["repetition_penalty"], step=0.05,
                                           label="repetition_penalty")

    chatbot = gr.ChatInterface(
        fn=chat_fn,
        additional_inputs=[max_new_tokens_ui, temperature_ui, top_k_ui, top_p_ui, repetition_penalty_ui],
        title=None,
        examples=[
            ["Hola, ¿cómo estás?"],
            ["¿Cuánto es 8 + 5?"],
            ["Explícame qué es un algoritmo."],
        ],
        cache_examples=False,
    )

demo.queue(max_size=16)

# ---------------- API REST /generate (la que consume el PHP) ----------------
# El PHP hace: fetch(url, { method:'POST', body: JSON.stringify({text, max_tokens, temperature}) })
# y espera de vuelta: { "reply": "..." }
#
# IMPORTANTE:
# - ssr_mode=False: Gradio 6 usa un servidor Node.js aparte para SSR, que
#   intentaba levantarse en el puerto 7861 y chocaba. Lo desactivamos porque
#   no lo necesitamos para servir la API.
# - El middleware CORS se pasa vía app_kwargs ANTES de llamar a launch(),
#   porque una vez que la app arranca, Starlette ya no permite añadir
#   middleware (por eso fallaba con app.add_middleware() después).

class GenerateRequest(BaseModel):
    text: str
    max_tokens: Optional[int] = None
    temperature: Optional[float] = None
    top_k: Optional[int] = None
    top_p: Optional[float] = None
    repetition_penalty: Optional[float] = None


PORT = int(os.environ.get("PORT", 7860))
demo.launch(
    server_name="0.0.0.0",
    server_port=PORT,
    prevent_thread_lock=True,
    ssr_mode=False,
    app_kwargs={
        "middleware": [
            Middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"]),
        ]
    },
)

app = demo.app


@app.post("/generate")
def generate_endpoint(req: GenerateRequest):
    if not req.text or not req.text.strip():
        return {"reply": "Escribe algo para que pueda responder."}
    try:
        reply = run_inference(
            req.text,
            max_new_tokens=req.max_tokens,
            temperature=req.temperature,
            top_k=req.top_k,
            top_p=req.top_p,
            repetition_penalty=req.repetition_penalty,
        )
        if not reply:
            reply = "No pude generar una respuesta."
        return {"reply": reply}
    except Exception as e:
        return {"reply": f"Error del modelo: {e}"}


@app.get("/generate")
def generate_health():
    # Solo para poder comprobar en el navegador que la ruta existe (GET no genera texto)
    return {"status": "ok", "info": "Usa POST con JSON {text, max_tokens, temperature}"}


# demo.launch(prevent_thread_lock=True) ya dejó el servidor corriendo en un
# hilo en segundo plano (un solo proceso, un solo puerto). Mantenemos vivo
# el hilo principal para que el contenedor del Space no termine.
demo.block_thread()