| |
| |
| |
| |
| 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 |
|
|
| |
| |
| torch.set_num_threads(max(1, os.cpu_count() or 1)) |
| torch.set_grad_enabled(False) |
|
|
| DEVICE = "cpu" |
|
|
| REPO_ID = "TeszenAI/MTP-1.2" |
| FILENAME = "MTP_MODEL.pt" |
|
|
| |
| 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) |
|
|
|
|
| |
| 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() |
|
|
| |
| |
| 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)) |
|
|
|
|
| |
| @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"] |
|
|
| |
| |
| |
| |
| MAX_TOKENS_HARD_LIMIT = 4000 |
| max_new_tokens = max(1, min(max_new_tokens, MAX_TOKENS_HARD_LIMIT)) |
|
|
| 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) |
|
|
|
|
| |
| 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, 4000, value=gen_defaults["max_new_tokens"], step=10, 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) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| 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(): |
| |
| return {"status": "ok", "info": "Usa POST con JSON {text, max_tokens, temperature}"} |
|
|
|
|
| |
| |
| |
| demo.block_thread() |