{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# La atención\n",
    "\n",
    "Cada palabra mira a las demás y se queda con lo que le sirve. Escrita en seis líneas de numpy.\n",
    "\n",
    "Cuaderno de soluciones del capítulo 16 de **Deep learning desde cero**, de Miss Yera.\n",
    "\n",
    "Corre de arriba abajo. Si lo abres en Google Colab no necesitas instalar nada.\n",
    "\n",
    "Capítulo completo: https://missyera.com/guias/deep-learning-desde-cero/atencion/\n",
    "\n",
    "Este es el cuaderno de **soluciones**. Trae el código de cada ejercicio, la\n",
    "explicación de la trampa y la respuesta del quiz. Si vienes del cuaderno de\n",
    "práctica sin haberlo intentado, vuelve 🙂"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "El capítulo 13 terminó con tres cosas mal: los antónimos salían iguales, se\n",
    "perdía el orden, y cada palabra tenía un vector fijo pasara lo que pasara 🤔\n",
    "\n",
    "Antes de entrar, piénsalo tú: **¿cómo sabes que \"banco\" significa una cosa en \"me senté en el banco\" y otra en \"fui al banco\"?** Lo sabes por lo que hay alrededor, y eso es exactamente lo que vamos a construir 🎯\n",
    "\n",
    "Las tres las arregla la misma idea, publicada en 2017 en un paper que se llama\n",
    "[Attention Is All You Need](https://arxiv.org/abs/1706.03762). Y cabe en seis\n",
    "líneas de numpy."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## La idea, con una frase"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Toma \"la bodega de lima no pago\". Para entender qué significa\n",
    "*pago* en esa frase hace falta mirar el *no*. Y para saber de qué\n",
    "bodega hablamos, hay que mirar *lima*.\n",
    "\n",
    "Eso es la atención: **cada palabra mira a todas las demás y se queda con\n",
    "una mezcla de lo que le sirve**. No un vector fijo, sino uno armado para\n",
    "esa frase.\n",
    "\n",
    "Y cómo decide a quién mirar es lo bonito. Cada palabra genera tres cosas:\n",
    "\n",
    "- 🔎 Una **consulta**: qué estoy buscando.\n",
    "\n",
    "- 🏷️ Una **clave**: qué ofrezco yo.\n",
    "\n",
    "- 📦 Un **valor**: qué me llevo si me eligen.\n",
    "\n",
    "Se comparan todas las consultas con todas las claves, y eso decide los pesos\n",
    "de la mezcla. Es literalmente una búsqueda, con la diferencia de que en vez de\n",
    "elegir un resultado se lleva un poquito de todos."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Las seis líneas"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Atención(Q,K,V)=softmax(QK⊤dk)V\n",
    "\n",
    "cada palabra pregunta a todas las demás cuánto le importan, esas notas se vuelven porcentajes y con ellas se hace una mezcla ponderada de la información"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "\n",
    "def softmax(z):\n",
    "    z = z - z.max(axis=-1, keepdims=True)      # el truco del capítulo 3\n",
    "    e = np.exp(z)\n",
    "    return e / e.sum(axis=-1, keepdims=True)\n",
    "\n",
    "frase = ['la', 'bodega', 'de', 'lima', 'no', 'pago']\n",
    "rng = np.random.default_rng(7)\n",
    "d = 8                                          # tamaño de cada vector\n",
    "\n",
    "E = rng.normal(0, 1, (len(frase), d))          # los embeddings del capítulo 13\n",
    "Wq = rng.normal(0, 0.5, (d, d))\n",
    "Wk = rng.normal(0, 0.5, (d, d))\n",
    "Wv = rng.normal(0, 0.5, (d, d))\n",
    "\n",
    "Q = E @ Wq                                     # consultas\n",
    "K = E @ Wk                                     # claves\n",
    "V = E @ Wv                                     # valores\n",
    "\n",
    "puntajes = Q @ K.T / np.sqrt(d)                # quién le interesa a quién\n",
    "A = softmax(puntajes)                          # convertidos en pesos que suman 1\n",
    "salida = A @ V                                 # la mezcla\n",
    "\n",
    "print('formas:', Q.shape, K.shape, V.shape)\n",
    "print('matriz de atención:', A.shape, ' salida:', salida.shape)\n",
    "print('cada fila suma:', np.round(A.sum(axis=1), 4))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Eso es todo. Tres multiplicaciones para sacar Q, K y V, una para compararlos,\n",
    "un softmax y una mezcla 🎯\n",
    "\n",
    "Y ojo con un detalle que parece decorativo: **dividir por la raíz de\n",
    "d**. Sin eso, con vectores largos los puntajes salen enormes, el softmax\n",
    "se satura (capítulo 3) y los pesos quedan en 1 y 0. La raíz los mantiene en un\n",
    "rango donde el softmax todavía tiene pendiente."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Mirar la matriz de atención"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "print('        ' + '  '.join(f'{p:>8}' for p in frase))\n",
    "for p, fila in zip(frase, A):\n",
    "    print(f'{p:8}' + '  '.join(f'{x:8.3f}' for x in fila))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Cada fila es una palabra y dice cuánto mira a cada una de las seis. Las filas\n",
    "suman 1, o sea que cada palabra reparte un total fijo de atención.\n",
    "\n",
    "Y ahora la parte honesta: **estos números no significan nada**.\n",
    "Las matrices Wq, Wk y Wv las saqué al azar, así que la atención que ves es\n",
    "aleatoria. Que \"bodega\" mire a \"no\" con 0,506 es casualidad.\n",
    "\n",
    "Lo que sí es real y es lo que importa aquí es la *maquinaria*: las\n",
    "formas, que las filas sumen 1, y que cada palabra reciba una mezcla distinta. En\n",
    "un modelo entrenado, esas tres matrices se aprenden con la retropropagación del\n",
    "capítulo 6 igual que cualquier otro peso, y ahí los patrones sí significan 🔧"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## La máscara, que es lo que hace que ChatGPT escriba"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Si el modelo tiene que **predecir la palabra siguiente**, no\n",
    "puede dejar que cada palabra mire a las que vienen después: sería copiarse la\n",
    "respuesta.\n",
    "\n",
    "Se arregla poniendo menos infinito en los puntajes de las palabras futuras\n",
    "antes del softmax:"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "mascara = np.triu(np.ones((len(frase), len(frase))), 1) * -1e9\n",
    "A_causal = softmax(puntajes + mascara)\n",
    "\n",
    "print('        ' + '  '.join(f'{p:>8}' for p in frase))\n",
    "for p, fila in zip(frase, A_causal):\n",
    "    print(f'{p:8}' + '  '.join(f'{x:8.3f}' for x in fila))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Mira la forma de triángulo 🔺\n",
    "\n",
    "La primera palabra solo puede mirarse a sí misma, así que su atención es 1,000\n",
    "y el resto ceros. La segunda mira a dos, la tercera a tres. Y la última las ve\n",
    "todas.\n",
    "\n",
    "Ese triángulo es la diferencia entre un modelo que *lee* y uno que\n",
    "*escribe*. Los que escriben (ChatGPT, Claude) llevan esta máscara, y por\n",
    "eso van palabra por palabra sin poder ver lo que todavía no escribieron.\n",
    "\n",
    "Y el `-1e9` no es capricho: `exp(-1e9)` da cero, así que\n",
    "el softmax les asigna peso cero. Poner directamente `-inf` también\n",
    "funciona, pero da `nan` si una fila entera queda enmascarada 🧯"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Lo que arregla, medido"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Aquí está el pago de lo que quedó pendiente en el capítulo 13. Tomamos la\n",
    "misma palabra, *llego*, en dos frases que dicen lo contrario:"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "vocabulario = {'el': 0, 'pedido': 1, 'llego': 2, 'completo': 3, 'no': 4}\n",
    "tabla = rng.normal(0, 1, (5, d))               # un vector fijo por palabra\n",
    "\n",
    "def atiende(palabras):\n",
    "    Emb = tabla[[vocabulario[p] for p in palabras]]\n",
    "    Q, K, V = Emb @ Wq, Emb @ Wk, Emb @ Wv\n",
    "    return softmax(Q @ K.T / np.sqrt(d)) @ V\n",
    "\n",
    "f1 = ['el', 'pedido', 'llego', 'completo']\n",
    "f2 = ['el', 'pedido', 'no', 'llego']\n",
    "\n",
    "def coseno(a, b):\n",
    "    return float(a @ b / (np.linalg.norm(a) * np.linalg.norm(b)))\n",
    "\n",
    "fijo = coseno(tabla[vocabulario['llego']], tabla[vocabulario['llego']])\n",
    "tras = coseno(atiende(f1)[f1.index('llego')], atiende(f2)[f2.index('llego')])\n",
    "print('el vector fijo de \"llego\" en las dos frases:', round(fijo, 4))\n",
    "print('después de la atención                     :', round(tras, 4))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Ahí está 🎉\n",
    "\n",
    "Con el embedding del capítulo 13, \"llego\" es el mismo vector en las dos\n",
    "frases: parecido 1,0, siempre, pase lo que pase.\n",
    "\n",
    "Después de la atención, los dos \"llego\" tienen un parecido de\n",
    "**0,3808**. El de \"el pedido llego completo\" y el de \"el pedido no\n",
    "llego\" pasaron a ser vectores distintos, porque cada uno se mezcló con las\n",
    "palabras que lo rodean.\n",
    "\n",
    "Eso es lo que quiere decir **representación contextual**, y es\n",
    "literalmente el motivo de que los modelos de lenguaje de hoy funcionen y los de\n",
    "hace diez años no.\n",
    "\n",
    "Y de paso arregla el orden: como la mezcla depende de con quién está cada\n",
    "palabra, \"no llego\" y \"llego no\" ya no dan lo mismo."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Ejercicios"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 1. Qué pasa sin la raíz de d\n",
    "\n",
    "Quita la división y mira los pesos."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "sin_raiz = softmax(Q @ K.T)\n",
    "print('con raíz  :', np.round(A[1], 3))\n",
    "print('sin raíz  :', np.round(sin_raiz[1], 3))\n",
    "print()\n",
    "print('el peso más grande, con raíz :', round(float(A.max()), 4))\n",
    "print('el peso más grande, sin raíz :', round(float(sin_raiz.max()), 4))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "con raíz  : [0.103 0.194 0.048 0.078 0.506 0.071]\n",
    "sin raíz  : [0.01  0.061 0.001 0.005 0.92  0.004]\n",
    "\n",
    "el peso más grande, con raíz : 0.7268\n",
    "el peso más grande, sin raíz : 0.9682\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Sin la raíz, el peso más grande sube de 0,7268 a 0,9682, y mira la fila de\n",
    "\"bodega\": pasa de repartir (0,103 0,194 0,048 0,078 0,506 0,071) a llevárselo casi\n",
    "todo un solo sitio (0,920).\n",
    "\n",
    "Y no es solo que quede feo: un softmax saturado tiene pendiente cero\n",
    "(capítulo 3), así que **por ahí ya no pasa gradiente y esa parte deja de\n",
    "aprender**. Una división que parece cosmética y sostiene el\n",
    "entrenamiento 🔑"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 2. Varias cabezas mirando cosas distintas\n",
    "\n",
    "La atención de verdad se hace varias veces en paralelo. Haz\n",
    "cuatro cabezas de 2 dimensiones cada una."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "CABEZAS, d_cabeza = 4, 2\n",
    "salidas = []\n",
    "for c in range(CABEZAS):\n",
    "    wq = rng.normal(0, 0.5, (d, d_cabeza))\n",
    "    wk = rng.normal(0, 0.5, (d, d_cabeza))\n",
    "    wv = rng.normal(0, 0.5, (d, d_cabeza))\n",
    "    a = softmax((E @ wq) @ (E @ wk).T / np.sqrt(d_cabeza))\n",
    "    salidas.append(a @ (E @ wv))\n",
    "    print(f'cabeza {c}: la palabra \"pago\" mira más a \"{frase[int(a[5].argmax())]}\"')\n",
    "\n",
    "junta = np.concatenate(salidas, axis=1)\n",
    "print()\n",
    "print('cada cabeza da', salidas[0].shape, 'y juntas dan', junta.shape)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "cabeza 0: la palabra \"pago\" mira más a \"de\"\n",
    "cabeza 1: la palabra \"pago\" mira más a \"no\"\n",
    "cabeza 2: la palabra \"pago\" mira más a \"no\"\n",
    "cabeza 3: la palabra \"pago\" mira más a \"de\"\n",
    "\n",
    "cada cabeza da (6, 2) y juntas dan (6, 8)\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Cuatro cabezas y cada una mira a un sitio distinto. Al final se pegan una al\n",
    "lado de la otra y vuelve a salir el mismo tamaño de siempre.\n",
    "\n",
    "En un modelo entrenado, cada cabeza se especializa: unas siguen la sintaxis,\n",
    "otras enlazan un pronombre con su sujeto. Aquí son cuatro al azar, así que lo\n",
    "único real es **que se pueden mirar varias cosas a la vez sin gastar más\n",
    "tamaño** 👀"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 3. La atención no sabe de orden\n",
    "\n",
    "Baraja la frase y mira si la salida de una palabra\n",
    "cambia."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "orden = [3, 1, 0, 5, 2, 4]\n",
    "E_barajado = E[orden]\n",
    "Qb, Kb, Vb = E_barajado @ Wq, E_barajado @ Wk, E_barajado @ Wv\n",
    "salida_b = softmax(Qb @ Kb.T / np.sqrt(d)) @ Vb\n",
    "\n",
    "donde = orden.index(1)              # dónde quedó \"bodega\"\n",
    "print('salida de \"bodega\" en la frase normal  :', np.round(salida[1][:4], 4))\n",
    "print('salida de \"bodega\" en la frase barajada:', np.round(salida_b[donde][:4], 4))\n",
    "print('¿son iguales?', np.allclose(salida[1], salida_b[donde]))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "salida de \"bodega\" en la frase normal  : [ 0.0826 -0.1186 -0.0884 -0.3718]\n",
    "salida de \"bodega\" en la frase barajada: [ 0.0826 -0.1186 -0.0884 -0.3718]\n",
    "¿son iguales? True\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Iguales. La atención mira *quién* está en la frase, no **en qué\n",
    "orden**.\n",
    "\n",
    "Y eso es un problema serio, porque \"la bodega no pago\" y \"no la bodega pago\"\n",
    "darían lo mismo. Se arregla sumándole a cada embedding un vector que depende de\n",
    "su posición, y a eso se le llama **codificación posicional**. Es el\n",
    "ejercicio 4 🔢"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 4. Meterle la posición\n",
    "\n",
    "Suma senos y cosenos de distinta frecuencia según la\n",
    "posición, que es como se hizo en el paper original."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def posiciones(n, dim):\n",
    "    pos = np.arange(n)[:, None]\n",
    "    i = np.arange(dim)[None, :]\n",
    "    angulo = pos / (10000 ** (2 * (i // 2) / dim))\n",
    "    P = np.zeros((n, dim))\n",
    "    P[:, 0::2] = np.sin(angulo[:, 0::2])\n",
    "    P[:, 1::2] = np.cos(angulo[:, 1::2])\n",
    "    return P\n",
    "\n",
    "P = posiciones(len(frase), d)\n",
    "E_con_pos = E + P\n",
    "Qp, Kp, Vp = E_con_pos @ Wq, E_con_pos @ Wk, E_con_pos @ Wv\n",
    "salida_p = softmax(Qp @ Kp.T / np.sqrt(d)) @ Vp\n",
    "\n",
    "E_bar_pos = E[orden] + P\n",
    "Qbp, Kbp, Vbp = E_bar_pos @ Wq, E_bar_pos @ Wk, E_bar_pos @ Wv\n",
    "salida_bp = softmax(Qbp @ Kbp.T / np.sqrt(d)) @ Vbp\n",
    "\n",
    "print('con posición, ¿siguen siendo iguales?',\n",
    "      np.allclose(salida_p[1], salida_bp[orden.index(1)]))\n",
    "print('parecido entre las dos:', round(coseno(salida_p[1], salida_bp[orden.index(1)]), 4))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "con posición, ¿siguen siendo iguales? False\n",
    "parecido entre las dos: 0.9672\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Ya no son iguales, que era lo que buscábamos. Pero el parecido sigue siendo\n",
    "0,9672, o sea que **la posición movió muy poco**.\n",
    "\n",
    "Y eso también hay que decirlo: aquí los embeddings salen de una normal con\n",
    "desviación 1 y la codificación posicional va entre -1 y 1, así que es un\n",
    "empujoncito al lado de un vector que ya era grande. En un modelo entrenado los\n",
    "embeddings aprenden a dejarle sitio a esa señal, porque el entrenamiento castiga\n",
    "confundir el orden. Aquí no hay entrenamiento, así que solo se ve el\n",
    "mecanismo 🌊\n",
    "\n",
    "Los senos y cosenos parecen una rareza y tienen su razón: dan un patrón\n",
    "distinto para cada posición y permiten que el modelo calcule distancias entre\n",
    "posiciones con sumas y restas. Hoy se usan otras variantes y la idea es la\n",
    "misma."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 5. Cuánto cuesta la atención\n",
    "\n",
    "Cuenta las comparaciones según lo larga que sea la\n",
    "entrada."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for palabras_n in [6, 100, 1_000, 100_000]:\n",
    "    print(f'{palabras_n:7,} palabras: {palabras_n ** 2:15,} comparaciones')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "      6 palabras:              36 comparaciones\n",
    "    100 palabras:          10,000 comparaciones\n",
    "  1,000 palabras:       1,000,000 comparaciones\n",
    "100,000 palabras:  10,000,000,000 comparaciones\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Cada palabra mira a todas, así que el costo crece **al cuadrado**.\n",
    "Con 100.000 palabras son diez mil millones de comparaciones para una sola capa,\n",
    "de una sola cabeza.\n",
    "\n",
    "Ahí tienes por qué los modelos tienen un límite de contexto y por qué cuesta\n",
    "tanto ampliarlo. Es el problema abierto más caro del campo, y hay una industria\n",
    "entera buscándole la vuelta 💸"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 6. La máscara sobre una frase más larga\n",
    "\n",
    "Comprueba que el triángulo escala."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "larga = 5\n",
    "p_larga = rng.normal(0, 1, (larga, larga))\n",
    "m_larga = np.triu(np.ones((larga, larga)), 1) * -1e9\n",
    "A_larga = softmax(p_larga + m_larga)\n",
    "\n",
    "print(np.round(A_larga, 3))\n",
    "print()\n",
    "print('ceros por encima de la diagonal:',\n",
    "      int((A_larga[np.triu_indices(larga, 1)] == 0).sum()),\n",
    "      'de', len(np.triu_indices(larga, 1)[0]))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "[[1.    0.    0.    0.    0.   ]\n",
    " [0.024 0.976 0.    0.    0.   ]\n",
    " [0.131 0.117 0.752 0.    0.   ]\n",
    " [0.073 0.053 0.421 0.453 0.   ]\n",
    " [0.037 0.104 0.598 0.009 0.252]]\n",
    "\n",
    "ceros por encima de la diagonal: 10 de 10\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Los diez huecos de arriba a la derecha son cero exacto, siempre.\n",
    "\n",
    "Esa es la garantía de que el modelo no se copia del futuro. Y es una garantía\n",
    "de verdad, no una tendencia: matemáticamente no puede 🔒"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 7. El error de la máscara del tamaño equivocado\n",
    "\n",
    "Aplica una máscara de 5 por 5 a una frase de 6\n",
    "palabras."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "softmax(puntajes + np.triu(np.ones((5, 5)), 1) * -1e9)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "ValueError: operands could not be broadcast together with shapes (6,6) (5,5) \n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Los puntajes son 6 por 6 y la máscara 5 por 5, así que numpy no sabe cómo\n",
    "alinearlas.\n",
    "\n",
    "Este error es de los buenos, porque el equivalente silencioso existe y es\n",
    "peor: si la máscara fuera de 1 por 6 o de 6 por 1, numpy la estiraría sin\n",
    "quejarse y estarías enmascarando lo que no toca. El modelo entrenaría, la pérdida\n",
    "bajaría, y estaría viendo el futuro sin que nadie se entere 😬\n",
    "\n",
    "Cuando montes atención a mano, **imprime la forma de la máscara y la de\n",
    "los puntajes antes de sumarlas**."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Comprueba que lo tienes\n",
    "\n",
    "En el capítulo la palabra llego pasa de un parecido de 1,0 consigo misma a repartir 0,3808 hacia otra palabra. ¿Qué hizo la atención?\n",
    "\n",
    "a) Mezclar información de las demás palabras en la representación de esa\n",
    "\n",
    "b) Cambiar el significado de la palabra\n",
    "\n",
    "c) Ordenar la frase\n",
    "\n",
    "d) Elegir qué palabra es más importante en la frase\n",
    "\n",
    "---\n",
    "\n",
    "**La correcta es la a.**\n",
    "\n",
    "*b)* La palabra es la misma. Lo que cambia es lo que su vector arrastra de las demás.\n",
    "\n",
    "*c)* El orden lo aporta la posición, que es otra pieza.\n",
    "\n",
    "*d)* No elige una: reparte porcentajes entre todas.\n",
    "\n",
    "Cada palabra pregunta a las demás cuánto le importan, y con esas notas hace una mezcla."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Y la pérdida que baja demasiado rápido"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### La trampa\n",
    "\n",
    "Entrenas un modelo que escribe texto y la pérdida baja hasta casi cero en unas pocas vueltas. Nunca habías visto entrenar algo tan rápido.\n",
    "\n",
    "```\n",
    "puntajes = Q @ K.T / np.sqrt(d)\n",
    "pesos = softmax(puntajes)\n",
    "salida = pesos @ V\n",
    "\n",
    "# perdida tras 200 vueltas: 0.0004\n",
    "```\n",
    "\n",
    "**Qué está mal**\n",
    "\n",
    "Falta la máscara 🎭 Sin ella, al predecir la palabra número cinco el modelo puede mirar las palabras seis, siete y ocho, que son justo las que tiene que adivinar. Está copiando la respuesta del renglón de al lado.\n",
    "\n",
    "Por eso la pérdida baja tan rápido, y esa velocidad es la señal: si tu modelo de lenguaje aprende sospechosamente rápido, casi siempre es que ve el futuro.\n",
    "\n",
    "La máscara pone menos infinito en el triángulo de arriba antes del softmax, para que esos pesos salgan cero. Y no es un detalle de implementación: **es lo que hace que un modelo pueda escribir**, porque lo obliga a aprender a continuar en vez de a rellenar huecos. El día que generes texto con el modelo sin máscara, saldrá un desastre y la pérdida de entrenamiento seguirá diciendo 0,0004."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Lo que te llevas"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "- 🔎 Cada palabra genera consulta, clave y valor, y se lleva una mezcla de las\n",
    "demás según qué tanto le interesan.\n",
    "\n",
    "- 🧮 Son seis líneas de numpy y las filas de la matriz de atención suman 1.\n",
    "\n",
    "- 🔑 Dividir por la raíz de d evita que el softmax se sature: sin eso el peso\n",
    "mayor pasa de 0,7268 a 0,9682, y por ahí deja de pasar gradiente.\n",
    "\n",
    "- 🔺 La máscara causal es un triángulo de ceros, y es lo que separa un modelo\n",
    "que lee de uno que escribe.\n",
    "\n",
    "- 🎉 Arregla lo del capítulo 13: \"llego\" pasa de tener parecido 1,0 consigo\n",
    "mismo en cualquier frase a 0,3808 entre dos frases distintas.\n",
    "\n",
    "- 🔢 La atención sola no sabe de orden: barajar la frase da exactamente la\n",
    "misma salida. Hay que sumarle la posición.\n",
    "\n",
    "- 💸 Cuesta al cuadrado: 100.000 palabras son diez mil millones de\n",
    "comparaciones. Por eso hay límite de contexto.\n",
    "\n",
    "Y si de todo el capítulo te llevas una sola frase, que sea esta:"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Toda la atención cabe en seis líneas. Lo difícil no fue la idea, fue tardar treinta años en tenerla."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Y si quieres el vocabulario suelto de transformers, atención y tokens, lo tengo en dos líneas por término en el [glosario de IA](https://missyera.com/glosario-ia/) 📖\n",
    "\n",
    "En el capítulo 17 juntamos esto en un transformer y vemos qué está pasando\n",
    "exactamente cuando le escribes a un modelo de lenguaje.\n",
    "\n",
    "Que tengas lindo día! 🌸"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "\n",
    "Ese era el capítulo 16 de **Deep learning desde cero**. El texto completo, con las salidas de cada bloque, está en https://missyera.com/guias/deep-learning-desde-cero/atencion/\n",
    "\n",
    "Que tengas lindo día! 🌸"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3.11"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
