adrianrossv commited on
Commit
75432cc
·
verified ·
1 Parent(s): 33a8e75

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +122 -25
app.py CHANGED
@@ -1,6 +1,21 @@
1
  # ================================================================
2
  # MTP - app.py para Hugging Face Space (Gradio, CPU)
3
  # Carga el checkpoint MTP_MODEL.pt desde el repo TeszenAI/MTP-1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4
  # ================================================================
5
  import os
6
  import math
@@ -17,10 +32,22 @@ from huggingface_hub import hf_hub_download
17
  # ---------------- Optimización para CPU ----------------
18
  # Limita hilos a los núcleos disponibles (evita overhead en Spaces pequeños)
19
  torch.set_num_threads(max(1, os.cpu_count() or 1))
 
 
 
 
 
 
 
 
20
  torch.set_grad_enabled(False) # solo inferencia, nunca necesitamos gradientes
21
 
22
  DEVICE = "cpu"
23
 
 
 
 
 
24
  REPO_ID = "TeszenAI/MTP-1.2"
25
  FILENAME = "MTP_MODEL.pt"
26
 
@@ -35,21 +62,52 @@ class CausalSelfAttention(nn.Module):
35
  self.attn_dropout = nn.Dropout(dropout)
36
  self.resid_dropout = nn.Dropout(dropout)
37
  mask = torch.tril(torch.ones(block_size, block_size)).view(1, 1, block_size, block_size)
 
 
 
 
38
  self.register_buffer("mask", mask)
39
 
40
- def forward(self, x):
41
  B, T, C = x.shape
42
  qkv = self.qkv(x)
43
  q, k, v = qkv.split(C, dim=2)
44
  q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
45
  k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
46
  v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
47
- att = (q @ k.transpose(-2, -1)) / math.sqrt(self.head_dim)
48
- att = att.masked_fill(self.mask[:, :, :T, :T] == 0, float("-inf"))
49
- att = F.softmax(att, dim=-1)
50
- att = self.attn_dropout(att)
51
- out = (att @ v).transpose(1, 2).contiguous().view(B, T, C)
52
- return self.resid_dropout(self.proj(out))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
53
 
54
 
55
  class FeedForward(nn.Module):
@@ -72,10 +130,11 @@ class Block(nn.Module):
72
  self.ln2 = nn.LayerNorm(n_embd)
73
  self.ff = FeedForward(n_embd, dropout)
74
 
75
- def forward(self, x):
76
- x = x + self.attn(self.ln1(x))
 
77
  x = x + self.ff(self.ln2(x))
78
- return x
79
 
80
 
81
  class MTP(nn.Module):
@@ -90,15 +149,22 @@ class MTP(nn.Module):
90
  self.lm_head = nn.Linear(n_embd, vocab_size, bias=False)
91
  self.lm_head.weight = self.tok_emb.weight
92
 
93
- def forward(self, idx):
94
  B, T = idx.shape
95
- pos = torch.arange(T, device=idx.device)
96
  x = self.tok_emb(idx) + self.pos_emb(pos)
97
  x = self.drop(x)
98
- for block in self.blocks:
99
- x = block(x)
 
 
 
 
 
 
100
  x = self.ln_f(x)
101
- return self.lm_head(x)
 
102
 
103
 
104
  # ---------------- Carga del checkpoint (una sola vez, al iniciar el Space) ----------------
@@ -122,12 +188,11 @@ model = MTP(
122
  model.load_state_dict(checkpoint["model_state_dict"])
123
  model.eval()
124
 
125
- # fusiona LayerNorm/Linear estáticamente no aplica aquí, pero fija modo eval
126
- # y evita cualquier dropout durante inferencia.
127
  BLOCK_SIZE = cfg["block_size"]
128
 
129
  print(f"MTP cargado ({checkpoint['meta']['model_name']}, "
130
- f"entrenado con {checkpoint['meta']['trained_examples']} ejemplos)")
 
131
 
132
 
133
  def encode_text(s):
@@ -138,17 +203,48 @@ def decode_ids(ids):
138
  return "".join(itos.get(i, "") for i in ids if i not in (PAD_ID, BOS_ID, EOS_ID))
139
 
140
 
141
- # ---------------- Generación ----------------
142
  @torch.inference_mode()
143
  def generate(idx, max_new_tokens, temperature, top_k, top_p, repetition_penalty):
 
 
 
144
  for _ in range(max_new_tokens):
145
- idx_cond = idx[:, -BLOCK_SIZE:]
146
- logits = model(idx_cond)
147
- logits = logits[:, -1, :] / max(temperature, 1e-5)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
148
 
149
  if repetition_penalty and repetition_penalty != 1.0:
150
- for token_id in set(idx[0].tolist()):
151
- logits[0, token_id] /= repetition_penalty
 
 
152
 
153
  if top_k is not None and top_k > 0:
154
  v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
@@ -171,6 +267,7 @@ def generate(idx, max_new_tokens, temperature, top_k, top_p, repetition_penalty)
171
  idx = torch.cat([idx, next_id], dim=1)
172
  if next_id.item() == EOS_ID:
173
  break
 
174
  return idx
175
 
176
 
@@ -178,7 +275,7 @@ def run_inference(text, max_new_tokens=None, temperature=None, top_k=None, top_p
178
  """Núcleo de generación, reutilizado por la UI de Gradio y por la API /generate.
179
  No reduce calidad por estar en CPU: usa exactamente el mismo muestreo
180
  (top_k + top_p + repetition_penalty) que en la Celda 2 de entrenamiento,
181
- solo que tarda más en devolver el resultado."""
182
  max_new_tokens = int(max_new_tokens) if max_new_tokens else gen_defaults["max_new_tokens"]
183
  temperature = float(temperature) if temperature is not None else gen_defaults["temperature"]
184
  top_k = int(top_k) if top_k is not None else gen_defaults["top_k"]
 
1
  # ================================================================
2
  # MTP - app.py para Hugging Face Space (Gradio, CPU)
3
  # Carga el checkpoint MTP_MODEL.pt desde el repo TeszenAI/MTP-1
4
+ #
5
+ # OPTIMIZACIÓN DE VELOCIDAD (sin tocar arquitectura ni pesos):
6
+ # - KV-cache en la atención: en generación autoregresiva, cada paso
7
+ # antes recomputaba TODO el contexto desde cero (O(n^2) en total).
8
+ # Ahora se reutiliza lo ya calculado y solo se procesa el token
9
+ # nuevo (O(n) en total). Es el mismo cálculo matemático, solo que
10
+ # no se repite trabajo ya hecho.
11
+ # - F.scaled_dot_product_attention: kernel fusionado de PyTorch,
12
+ # mismo resultado que el softmax manual pero más rápido en CPU.
13
+ # Si la versión de PyTorch no lo trae, cae automáticamente al
14
+ # cálculo manual (fallback), así que no se rompe en ningún entorno.
15
+ # - repetition_penalty vectorizado (sin bucle Python + set() por token).
16
+ #
17
+ # El modelo, los pesos, el muestreo (top_k/top_p/temperature/repetition)
18
+ # y las respuestas de la API/Gradio son EXACTAMENTE los mismos que antes.
19
  # ================================================================
20
  import os
21
  import math
 
32
  # ---------------- Optimización para CPU ----------------
33
  # Limita hilos a los núcleos disponibles (evita overhead en Spaces pequeños)
34
  torch.set_num_threads(max(1, os.cpu_count() or 1))
35
+
36
+ # set_num_interop_threads solo puede llamarse una vez y antes de cualquier
37
+ # operación paralela; lo protegemos por si el entorno ya lo fijó.
38
+ try:
39
+ torch.set_num_interop_threads(1)
40
+ except RuntimeError:
41
+ pass
42
+
43
  torch.set_grad_enabled(False) # solo inferencia, nunca necesitamos gradientes
44
 
45
  DEVICE = "cpu"
46
 
47
+ # Disponibilidad de scaled_dot_product_attention (PyTorch >= 2.0).
48
+ # Si no está disponible, usamos el softmax manual original como fallback.
49
+ _HAS_SDPA = hasattr(F, "scaled_dot_product_attention")
50
+
51
  REPO_ID = "TeszenAI/MTP-1.2"
52
  FILENAME = "MTP_MODEL.pt"
53
 
 
62
  self.attn_dropout = nn.Dropout(dropout)
63
  self.resid_dropout = nn.Dropout(dropout)
64
  mask = torch.tril(torch.ones(block_size, block_size)).view(1, 1, block_size, block_size)
65
+ # Se mantiene el buffer para que el state_dict del checkpoint cargue
66
+ # igual que antes (la clave "attn.mask" existe en el checkpoint).
67
+ # Ya no se usa en el forward optimizado con SDPA; solo lo usa el
68
+ # fallback manual si SDPA no está disponible.
69
  self.register_buffer("mask", mask)
70
 
71
+ def forward(self, x, past_kv=None, use_cache=False):
72
  B, T, C = x.shape
73
  qkv = self.qkv(x)
74
  q, k, v = qkv.split(C, dim=2)
75
  q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
76
  k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
77
  v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
78
+
79
+ if past_kv is not None:
80
+ past_k, past_v = past_kv
81
+ k = torch.cat([past_k, k], dim=2)
82
+ v = torch.cat([past_v, v], dim=2)
83
+
84
+ present_kv = (k, v) if use_cache else None
85
+
86
+ # Causal solo hace falta cuando hay varias queries nuevas sin pasado
87
+ # (prefill del prompt). En un paso de decodificación (T=1 con caché)
88
+ # el único token nuevo ya puede ver todo el pasado sin máscara.
89
+ is_causal = (past_kv is None) and (T > 1)
90
+
91
+ if _HAS_SDPA:
92
+ out = F.scaled_dot_product_attention(
93
+ q, k, v,
94
+ attn_mask=None,
95
+ dropout_p=0.0, # en eval() el dropout original no hace nada
96
+ is_causal=is_causal,
97
+ )
98
+ else:
99
+ att = (q @ k.transpose(-2, -1)) / math.sqrt(self.head_dim)
100
+ if is_causal:
101
+ Tk = k.size(-2)
102
+ causal_mask = torch.tril(torch.ones(T, Tk, device=x.device, dtype=torch.bool))
103
+ att = att.masked_fill(~causal_mask, float("-inf"))
104
+ att = F.softmax(att, dim=-1)
105
+ att = self.attn_dropout(att)
106
+ out = att @ v
107
+
108
+ out = out.transpose(1, 2).contiguous().view(B, T, C)
109
+ out = self.resid_dropout(self.proj(out))
110
+ return out, present_kv
111
 
112
 
113
  class FeedForward(nn.Module):
 
130
  self.ln2 = nn.LayerNorm(n_embd)
131
  self.ff = FeedForward(n_embd, dropout)
132
 
133
+ def forward(self, x, past_kv=None, use_cache=False):
134
+ attn_out, present_kv = self.attn(self.ln1(x), past_kv=past_kv, use_cache=use_cache)
135
+ x = x + attn_out
136
  x = x + self.ff(self.ln2(x))
137
+ return x, present_kv
138
 
139
 
140
  class MTP(nn.Module):
 
149
  self.lm_head = nn.Linear(n_embd, vocab_size, bias=False)
150
  self.lm_head.weight = self.tok_emb.weight
151
 
152
+ def forward(self, idx, past_key_values=None, use_cache=False, pos_offset=0):
153
  B, T = idx.shape
154
+ pos = torch.arange(pos_offset, pos_offset + T, device=idx.device)
155
  x = self.tok_emb(idx) + self.pos_emb(pos)
156
  x = self.drop(x)
157
+
158
+ new_past = [] if use_cache else None
159
+ for i, block in enumerate(self.blocks):
160
+ past_kv = past_key_values[i] if past_key_values is not None else None
161
+ x, present_kv = block(x, past_kv=past_kv, use_cache=use_cache)
162
+ if use_cache:
163
+ new_past.append(present_kv)
164
+
165
  x = self.ln_f(x)
166
+ logits = self.lm_head(x)
167
+ return logits, new_past
168
 
169
 
170
  # ---------------- Carga del checkpoint (una sola vez, al iniciar el Space) ----------------
 
188
  model.load_state_dict(checkpoint["model_state_dict"])
189
  model.eval()
190
 
 
 
191
  BLOCK_SIZE = cfg["block_size"]
192
 
193
  print(f"MTP cargado ({checkpoint['meta']['model_name']}, "
194
+ f"entrenado con {checkpoint['meta']['trained_examples']} ejemplos)"
195
+ f" | SDPA={'sí' if _HAS_SDPA else 'no (fallback manual)'}")
196
 
197
 
198
  def encode_text(s):
 
203
  return "".join(itos.get(i, "") for i in ids if i not in (PAD_ID, BOS_ID, EOS_ID))
204
 
205
 
206
+ # ---------------- Generación (con KV-cache) ----------------
207
  @torch.inference_mode()
208
  def generate(idx, max_new_tokens, temperature, top_k, top_p, repetition_penalty):
209
+ past_key_values = None
210
+ cache_len = 0 # cuántos tokens del extremo derecho de `idx` ya están en la caché
211
+
212
  for _ in range(max_new_tokens):
213
+ total_len = idx.shape[1]
214
+
215
+ if total_len <= BLOCK_SIZE:
216
+ if past_key_values is None:
217
+ # Primer paso: una sola pasada ("prefill") sobre todo el prompt.
218
+ logits, past_key_values = model(idx, use_cache=True)
219
+ cache_len = total_len
220
+ else:
221
+ # Pasos siguientes: solo se procesa el último token generado,
222
+ # reutilizando la caché de todo lo anterior.
223
+ last_token = idx[:, -1:]
224
+ logits, past_key_values = model(
225
+ last_token,
226
+ past_key_values=past_key_values,
227
+ use_cache=True,
228
+ pos_offset=cache_len,
229
+ )
230
+ cache_len += 1
231
+ logits = logits[:, -1, :]
232
+ else:
233
+ # Se superó block_size: mismo comportamiento que el modelo original
234
+ # (ventana deslizante recalculada por completo). Solo ocurre en
235
+ # respuestas muy largas; la caché se reinicia para esa ventana.
236
+ idx_cond = idx[:, -BLOCK_SIZE:]
237
+ logits, past_key_values = model(idx_cond, use_cache=True)
238
+ cache_len = BLOCK_SIZE
239
+ logits = logits[:, -1, :]
240
+
241
+ logits = logits / max(temperature, 1e-5)
242
 
243
  if repetition_penalty and repetition_penalty != 1.0:
244
+ # Vectorizado: antes era `for token_id in set(idx[0].tolist())`,
245
+ # un bucle Python nuevo por cada token generado.
246
+ unique_ids = torch.unique(idx[0])
247
+ logits[0, unique_ids] /= repetition_penalty
248
 
249
  if top_k is not None and top_k > 0:
250
  v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
 
267
  idx = torch.cat([idx, next_id], dim=1)
268
  if next_id.item() == EOS_ID:
269
  break
270
+
271
  return idx
272
 
273
 
 
275
  """Núcleo de generación, reutilizado por la UI de Gradio y por la API /generate.
276
  No reduce calidad por estar en CPU: usa exactamente el mismo muestreo
277
  (top_k + top_p + repetition_penalty) que en la Celda 2 de entrenamiento,
278
+ solo que ahora con KV-cache es notablemente más rápido en respuestas largas."""
279
  max_new_tokens = int(max_new_tokens) if max_new_tokens else gen_defaults["max_new_tokens"]
280
  temperature = float(temperature) if temperature is not None else gen_defaults["temperature"]
281
  top_k = int(top_k) if top_k is not None else gen_defaults["top_k"]