{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "# Redes que leen en orden\n",
    "\n",
    "La recurrente y la LSTM escritas a mano, y la medición de por qué las dos se quedaron cortas.\n",
    "\n",
    "Cuaderno de soluciones del capítulo 15 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/recurrentes-y-lstm/\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, y una era\n",
    "que **el orden se pierde**. Este capítulo arregla esa 🔗\n",
    "\n",
    "Antes de empezar, una pregunta: **¿cómo le explicarías a una máquina la\n",
    "diferencia entre \"llegó el pedido pero no la factura\" y \"llegó la factura pero no\n",
    "el pedido\"?** Son las mismas ocho palabras, la misma cuenta de cada una, y\n",
    "significan lo contrario 🤔\n",
    "\n",
    "Y te aviso de una vez cómo termina, porque este libro no le hace publicidad a\n",
    "nada: **las redes de este capítulo perdieron**. Hoy casi nadie las\n",
    "usa. Están aquí porque sin entender qué no pudieron hacer, la atención del\n",
    "capítulo 16 parece una idea bonita en vez de lo que es, que es la\n",
    "solución a un problema concreto y medido."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Una tarea que la bolsa de palabras no puede"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Vamos a fabricar comentarios de clientes de la distribuidora. Dos formas, y\n",
    "en las dos falta algo distinto."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "\n",
    "PALABRAS = ['ayer', 'hoy', 'llego', 'el', 'la', 'pedido', 'factura', 'pero',\n",
    "            'no', 'todavia']\n",
    "IDX = {p: i for i, p in enumerate(PALABRAS)}\n",
    "\n",
    "# Mismas palabras, orden distinto, significado contrario.\n",
    "def comentarios(n, semilla):\n",
    "    rng = np.random.default_rng(semilla)\n",
    "    X, y = [], []\n",
    "    for _ in range(n):\n",
    "        dia = 'ayer' if rng.random() < 0.5 else 'hoy'\n",
    "        cola = ['todavia'] if rng.random() < 0.5 else []\n",
    "        if rng.random() < 0.5:\n",
    "            f = [dia, 'llego', 'el', 'pedido', 'pero', 'no', 'la', 'factura']\n",
    "            etiqueta = 0                       # lo que falta es la factura\n",
    "        else:\n",
    "            f = [dia, 'llego', 'la', 'factura', 'pero', 'no', 'el', 'pedido']\n",
    "            etiqueta = 1                       # lo que falta es el pedido\n",
    "        X.append(f + cola)\n",
    "        y.append(etiqueta)\n",
    "    return X, np.array(y)\n",
    "\n",
    "Xtr, ytr = comentarios(600, 0)\n",
    "Xte, yte = comentarios(200, 1)\n",
    "print(' '.join(Xtr[0]), ' -> falta', ['la factura', 'el pedido'][ytr[0]])\n",
    "for i, f in enumerate(Xtr):\n",
    "    if ytr[i] != ytr[0]:\n",
    "        print(' '.join(f), ' -> falta', ['la factura', 'el pedido'][ytr[i]])\n",
    "        break"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Esto no es un juego de palabras: es el correo que le llega a operaciones un\n",
    "lunes. Y de la respuesta depende a quién se le reclama 📦\n",
    "\n",
    "Ahora contemos las palabras, que es lo que hace una bolsa de palabras, el\n",
    "método más común de todos:"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from sklearn.linear_model import LogisticRegression\n",
    "\n",
    "def bolsa(X):\n",
    "    M = np.zeros((len(X), len(PALABRAS)))\n",
    "    for i, f in enumerate(X):\n",
    "        for p in f:\n",
    "            M[i, IDX[p]] += 1\n",
    "    return M\n",
    "\n",
    "modelo = LogisticRegression(max_iter=1000).fit(bolsa(Xtr), ytr)\n",
    "print('filas distintas en la bolsa:', len(np.unique(bolsa(Xtr), axis=0)),\n",
    "      'de', len(Xtr))\n",
    "print('acierto de la bolsa de palabras:', round(modelo.score(bolsa(Xte), yte), 4))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Mira la primera línea, que es la que lo explica todo: **600 comentarios\n",
    "y solo 4 filas distintas** 😳\n",
    "\n",
    "Las dos variantes cuentan exactamente las mismas palabras, así que al contar\n",
    "se vuelven idénticas. Las cuatro combinaciones salen de *ayer* u\n",
    "*hoy* y de si lleva *todavía*, que no dicen nada.\n",
    "\n",
    "Y el acierto es 0,53, que con dos clases equilibradas es **tirar una\n",
    "moneda**. No es que el modelo esté mal entrenado: es que le estamos dando\n",
    "una entrada donde la respuesta no está. Ningún modelo, por grande que sea,\n",
    "adivina lo que le borraste 🪙"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## La recurrente, escrita"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "La idea es de una línea: **leer las palabras en orden y arrastrar un\n",
    "estado**. Ese estado es lo que la red se acuerda de lo que ya leyó, y en\n",
    "cada palabra se actualiza mezclando lo que venía con lo nuevo."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "V, D, H = len(PALABRAS), 8, 12\n",
    "\n",
    "def sigmoide(z):\n",
    "    return 1 / (1 + np.exp(-z))\n",
    "\n",
    "def arranca(semilla):\n",
    "    r = np.random.default_rng(semilla)\n",
    "    return dict(E=r.normal(0, .5, (V, D)), Wx=r.normal(0, .5, (D, H)),\n",
    "                Wh=r.normal(0, .5, (H, H)) * .5, b=np.zeros(H),\n",
    "                Wo=r.normal(0, .5, (H, 1)), bo=np.zeros(1))\n",
    "\n",
    "def adelante(p, ids):\n",
    "    hs = [np.zeros(H)]\n",
    "    for t in ids:\n",
    "        x = p['E'][t]\n",
    "        hs.append(np.tanh(x @ p['Wx'] + hs[-1] @ p['Wh'] + p['b']))\n",
    "    return hs, float(sigmoide(hs[-1] @ p['Wo'] + p['bo'])[0])\n",
    "\n",
    "p = arranca(1)\n",
    "frase = ['hoy', 'llego', 'el', 'pedido']\n",
    "hs, _ = adelante(p, [IDX[w] for w in frase])\n",
    "for w, h in zip(frase, hs[1:]):\n",
    "    print(f'{w:8} estado -> {np.round(h[:4], 3)} ...')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Fíjate en la línea del medio, que es toda la novedad del capítulo:"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "hs.append(np.tanh(x @ p['Wx'] + hs[-1] @ p['Wh'] + p['b']))\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Es **la misma neurona del capítulo 2** con un\n",
    "sumando de más: `hs[-1] @ Wh`, o sea el estado anterior. Eso es una\n",
    "recurrente entera. No hay más 🔁\n",
    "\n",
    "Y por eso el vector de *pedido* es distinto según lo que venga antes:\n",
    "no depende solo de la palabra, depende del camino. La misma `Wh` se\n",
    "usa en todos los pasos, igual que el filtro del capítulo 11 se\n",
    "usaba en todas las posiciones. Es la misma idea de compartir pesos, aplicada al\n",
    "tiempo en vez de al espacio 🔑"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Entrenarla, hacia atrás y en el tiempo"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "La retropropagación del capítulo 6 sirve igual, con un\n",
    "detalle: hay que recorrer la frase al revés, acumulando en las mismas matrices\n",
    "porque son las mismas en todos los pasos. Eso tiene nombre propio:\n",
    "**retropropagación en el tiempo**."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def gradiente(p, ids, y):\n",
    "    hs, prob = adelante(p, ids)\n",
    "    g = {k: np.zeros_like(v) for k, v in p.items()}\n",
    "    d = np.array([prob - y])\n",
    "    g['Wo'] = np.outer(hs[-1], d); g['bo'] = d\n",
    "    dh = (p['Wo'] @ d).ravel()\n",
    "    for t in range(len(ids) - 1, -1, -1):\n",
    "        dz = dh * (1 - hs[t + 1] ** 2)\n",
    "        g['Wx'] += np.outer(p['E'][ids[t]], dz)\n",
    "        g['Wh'] += np.outer(hs[t], dz)\n",
    "        g['b'] += dz\n",
    "        g['E'][ids[t]] += p['Wx'] @ dz\n",
    "        dh = p['Wh'] @ dz\n",
    "    return g\n",
    "\n",
    "def entrena(p, X, y, vueltas=10, paso=0.1):\n",
    "    I = [[IDX[w] for w in f] for f in X]\n",
    "    for v in range(vueltas):\n",
    "        for i in np.random.default_rng(v).permutation(len(I)):\n",
    "            g = gradiente(p, I[i], y[i])\n",
    "            for k in p:\n",
    "                p[k] -= paso * g[k]\n",
    "    return p\n",
    "\n",
    "def acierta(p, X, y):\n",
    "    return float(np.mean([(adelante(p, [IDX[w] for w in f])[1] >= .5) == yy\n",
    "                          for f, yy in zip(X, y)]))\n",
    "\n",
    "p = entrena(arranca(1), Xtr, ytr)\n",
    "print('acierto de la recurrente:', round(acierta(p, Xte, yte), 4))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "De 0,53 a 1,0 🎉\n",
    "\n",
    "Y el modelo no es más grande ni más listo: es que **recibe la\n",
    "información que la bolsa de palabras tiraba**. Casi siempre que un modelo\n",
    "no puede con algo, la pregunta buena no es qué modelo poner, sino qué le estamos\n",
    "dando de entrada."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## El problema, medido paso a paso"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Ahora la parte que explica por qué estas redes ya no se usan. Mira otra vez\n",
    "la última línea del bucle de arriba:"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "dh = p['Wh'] @ dz\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Cada paso hacia atrás multiplica por `Wh`. Otra vez. Y otra. En una\n",
    "frase de treinta palabras, lo que le llega al primer paso pasó por treinta\n",
    "multiplicaciones seguidas.\n",
    "\n",
    "Vamos a medir cuánto queda."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def cuanto_gradiente_llega(p, ids):\n",
    "    hs, prob = adelante(p, ids)\n",
    "    dh = (p['Wo'] @ np.array([prob - 1])).ravel()\n",
    "    normas = []\n",
    "    for t in range(len(ids) - 1, -1, -1):\n",
    "        dz = dh * (1 - hs[t + 1] ** 2)\n",
    "        normas.append(np.linalg.norm(dz))\n",
    "        dh = p['Wh'] @ dz\n",
    "    return normas[::-1]\n",
    "\n",
    "largo = 30\n",
    "ids = [IDX['pedido']] + list(np.random.default_rng(0).integers(0, V, size=largo - 1))\n",
    "n = cuanto_gradiente_llega(arranca(1), ids)\n",
    "print(f'{\"paso\":>5} {\"gradiente que llega\":>21}')\n",
    "for t in [0, 5, 10, 15, 20, 25, 29]:\n",
    "    print(f'{t+1:5} {n[t]:21.3e}')\n",
    "print(f'\\ndel paso 30 al paso 1 se divide entre {n[-1]/n[0]:.0f}')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Ahí está, y es de las cosas que se entienden mejor viéndolas que\n",
    "leyéndolas 📉\n",
    "\n",
    "La señal que corrige el último paso vale 0,46. La que le llega al primero vale\n",
    "0,0000007, o sea **seiscientas setenta y dos mil veces menos**. El\n",
    "primer paso prácticamente no se entera de que hubo un error.\n",
    "\n",
    "Eso se llama **gradiente que se desvanece**, y es el mismo\n",
    "fenómeno del capítulo 7 con las capas profundas, con una\n",
    "diferencia que lo empeora: allá el número de capas lo eliges tú, y aquí lo\n",
    "elige la longitud de la frase. Un comentario de cincuenta palabras son cincuenta\n",
    "capas, quieras o no 😰\n",
    "\n",
    "La consecuencia práctica es directa: **la red aprende lo cercano y no\n",
    "aprende lo lejano**. Con nuestras frases de ocho palabras funciona\n",
    "perfecto. Con un párrafo, no."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## El primer error que te va a salir"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Todo lo de arriba funciona porque nuestras frases miden ocho o nueve palabras\n",
    "y las procesamos de una en una. En cuanto quieras hacer varias a la vez, que es\n",
    "lo que se hace de verdad, aparece esto 👇"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "**Esto revienta a propósito.** Se ejecuta dentro de un `try` para que puedas seguir con \"ejecutar todo\" y aun así ver la queja."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "try:\n",
    "    lote = [[IDX[w] for w in f] for f in Xtr[:3]]\n",
    "    for f in lote:\n",
    "        print(len(f), f)\n",
    "\n",
    "    np.array(lote)\n",
    "except Exception as e:\n",
    "    print(f'{type(e).__name__}: {e}')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Y la queja que tiene que salir es esta:\n",
    "\n",
    "```\n",
    "ValueError: setting an array element with a sequence. The requested array has an inhomogeneous shape after 1 dimensions. The detected shape was (3,) + inhomogeneous part.\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "*inhomogeneous shape*. Una matriz de numpy quiere filas del mismo\n",
    "largo, y los comentarios no lo son: unos llevan *todavía* y otros no.\n",
    "\n",
    "Esto no es un detalle de numpy, es **el problema central de trabajar con\n",
    "secuencias**. Todo lo que sabes de matrices asume filas iguales, y el\n",
    "lenguaje no viene así. La solución estándar es rellenar las cortas hasta el largo\n",
    "de la más larga, y esa solución trae su propia trampa, que es la de este\n",
    "capítulo. Léela con calma 👀"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## La LSTM, o darle un carril propio a la memoria"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "En 1997, Hochreiter y Schmidhuber publicaron el arreglo, y es de esas ideas\n",
    "que una vez que la ves ya no se te olvida.\n",
    "\n",
    "Si el problema es que el recuerdo tiene que atravesar una multiplicación en\n",
    "cada paso, entonces **hay que darle un camino donde no lo multipliquen**.\n",
    "La LSTM añade una segunda línea, la *memoria*, que solo se suma y se\n",
    "borra, y tres puertas que deciden qué pasa con ella:\n",
    "\n",
    "- 🗑️ La puerta de **olvido** decide qué se borra de lo guardado.\n",
    "\n",
    "- 📥 La de **entrada** decide qué de lo nuevo se guarda.\n",
    "\n",
    "- 👁️ La de **salida** decide qué de lo guardado se deja ver.\n",
    "\n",
    "Cada puerta es una sigmoide, o sea un número entre 0 y 1 por cada casilla de\n",
    "la memoria: 0 es cerrado y 1 es abierto. Y se aprenden, como todo lo demás."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def arranca_lstm(semilla):\n",
    "    r = np.random.default_rng(semilla)\n",
    "    p = dict(E=r.normal(0, .5, (V, D)),\n",
    "             W=r.normal(0, .3, (D + H, 4 * H)), b=np.zeros(4 * H),\n",
    "             Wo=r.normal(0, .5, (H, 1)), bo=np.zeros(1))\n",
    "    p['b'][:H] = 1.0            # la puerta de olvido arranca abierta\n",
    "    return p\n",
    "\n",
    "def adelante_lstm(p, ids):\n",
    "    h = np.zeros(H); c = np.zeros(H); pasos = []\n",
    "    for t in ids:\n",
    "        x = p['E'][t]\n",
    "        z = np.concatenate([x, h])\n",
    "        a = z @ p['W'] + p['b']\n",
    "        olvido = sigmoide(a[:H])              # que borro de la memoria\n",
    "        entrada = sigmoide(a[H:2*H])          # que dejo entrar\n",
    "        salida = sigmoide(a[2*H:3*H])         # que dejo ver\n",
    "        nuevo = np.tanh(a[3*H:])              # lo que propongo guardar\n",
    "        c_ant = c\n",
    "        c = olvido * c_ant + entrada * nuevo\n",
    "        tc = np.tanh(c)\n",
    "        h = salida * tc\n",
    "        pasos.append((t, z, olvido, entrada, salida, nuevo, c_ant, c, tc, h))\n",
    "    return pasos, float(sigmoide(h @ p['Wo'] + p['bo'])[0])\n",
    "\n",
    "pasos, _ = adelante_lstm(arranca_lstm(1), [IDX[w] for w in frase])\n",
    "_, _, olvido, entrada, salida, _, _, _, _, _ = pasos[-1]\n",
    "print('en el ultimo paso las puertas valen (primeras 4 casillas):')\n",
    "print('  olvido ', np.round(olvido[:4], 3))\n",
    "print('  entrada', np.round(entrada[:4], 3))\n",
    "print('  salida ', np.round(salida[:4], 3))"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Las tres puertas son tres números por casilla y ninguno es 0 ni 1: la red\n",
    "decide *cuánto*, no sí o no. Y fíjate en la de olvido, que sale por encima\n",
    "de 0,7 en las cuatro: sin entrenar todavía, ya está conservando la mayor parte\n",
    "de lo que guardó, que es exactamente para lo que se puso ese sesgo 🚪\n",
    "\n",
    "La línea que importa de todo el bloque es esta:"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "c = olvido * c_ant + entrada * nuevo\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "La memoria nueva es la vieja, con una parte borrada, **más** lo\n",
    "que entra. Es una suma, no una transformación. Si la puerta de olvido está\n",
    "abierta (vale 1), la memoria pasa entera al siguiente paso sin que nadie la\n",
    "multiplique por una matriz.\n",
    "\n",
    "Y por eso `p['b'][:H] = 1.0` no es un detalle: arrancar con la\n",
    "puerta de olvido abierta es lo que hace que al principio del entrenamiento la\n",
    "memoria se conserve en vez de irse a cero. Es una de esas decisiones que no\n",
    "aparecen en los diagramas y deciden si la cosa entrena o no 🔧"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## La misma medición, sobre la LSTM"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Repetimos exactamente el experimento de antes, con la misma frase de treinta\n",
    "pasos. La cuenta hacia atrás cambia porque la memoria tiene su propio camino."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def cuanto_gradiente_llega_lstm(p, ids):\n",
    "    pasos, prob = adelante_lstm(p, ids)\n",
    "    dh = (p['Wo'] @ np.array([prob - 1])).ravel()\n",
    "    dc = np.zeros(H); normas = []\n",
    "    for t in range(len(ids) - 1, -1, -1):\n",
    "        _, z, olvido, entrada, salida, nuevo, c_ant, c, tc, h = pasos[t]\n",
    "        d_salida = dh * tc\n",
    "        dc = dc + dh * salida * (1 - tc ** 2)\n",
    "        normas.append(np.linalg.norm(dc))\n",
    "        da = np.concatenate([\n",
    "            dc * c_ant * olvido * (1 - olvido),\n",
    "            dc * nuevo * entrada * (1 - entrada),\n",
    "            d_salida * salida * (1 - salida),\n",
    "            dc * entrada * (1 - nuevo ** 2)])\n",
    "        dz = p['W'] @ da\n",
    "        dh = dz[D:]\n",
    "        dc = dc * olvido\n",
    "    return normas[::-1]\n",
    "\n",
    "m = cuanto_gradiente_llega_lstm(arranca_lstm(1), ids)\n",
    "print(f'{\"paso\":>5} {\"recurrente\":>13} {\"LSTM\":>13}')\n",
    "for t in [0, 10, 20, 29]:\n",
    "    print(f'{t+1:5} {n[t]:13.3e} {m[t]:13.3e}')\n",
    "print(f'\\nla recurrente se divide entre {n[-1]/n[0]:.0f}')\n",
    "print(f'la LSTM        se divide entre {m[-1]/m[0]:.0f}')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "**672.620 contra 63** 🎯\n",
    "\n",
    "Cuatro órdenes de magnitud, y ninguna de las dos redes está entrenada: es\n",
    "puramente la forma de la arquitectura. Mira la última línea del bucle, que es la\n",
    "que lo consigue:"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "dc = dc * olvido\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Hacia atrás, la memoria se multiplica por la puerta de olvido, que es un\n",
    "número entre 0 y 1 y arranca cerca de 1. La recurrente se multiplicaba por una\n",
    "matriz entera. Multiplicar por algo cercano a 1 treinta veces deja casi todo;\n",
    "multiplicar por una matriz treinta veces no deja nada 🚪"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Entonces por qué no ganaron"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Por una razón que no tiene nada que ver con la memoria, y que se ve en cuatro\n",
    "líneas."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import time\n",
    "\n",
    "X = np.random.default_rng(0).normal(0, 1, (200, H))\n",
    "W = np.random.default_rng(1).normal(0, 1, (H, H))\n",
    "\n",
    "inicio = time.time()\n",
    "for _ in range(50):\n",
    "    h = np.zeros(H)\n",
    "    for fila in X:                       # en cadena: cada paso espera al anterior\n",
    "        h = np.tanh(fila @ W + h @ W)\n",
    "uno_a_uno = time.time() - inicio\n",
    "\n",
    "inicio = time.time()\n",
    "for _ in range(50):\n",
    "    _ = np.tanh(X @ W)                   # de golpe: las 200 a la vez\n",
    "de_golpe = time.time() - inicio\n",
    "\n",
    "print(f'200 pasos en cadena   : {uno_a_uno:.4f} s')\n",
    "print(f'las 200 filas de golpe: {de_golpe:.4f} s')\n",
    "print(f'la cadena tarda {uno_a_uno/de_golpe:.0f} veces mas')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "El número exacto cambia en tu máquina; lo que no cambia es de qué lado está.\n",
    "En la mía la cadena tardó decenas de veces más 🐌\n",
    "\n",
    "Y la culpa no es de numpy: es que **el paso 2 necesita el resultado del\n",
    "paso 1**. No hay forma de hacerlos a la vez, ni con mil tarjetas gráficas.\n",
    "Una recurrente es secuencial por definición.\n",
    "\n",
    "Eso, que con frases de ocho palabras da igual, es lo que decidió la historia.\n",
    "Cuando en 2017 apareció una arquitectura que mira todas las palabras a la vez\n",
    "*y* resuelve las dependencias largas, se acabó la discusión: las\n",
    "recurrentes tardaban semanas donde la otra tardaba días, y encima entendía\n",
    "mejor 🏁\n",
    "\n",
    "Esa arquitectura es la del capítulo 16, y ahora ya sabes contra\n",
    "qué compite."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Una recurrente recuerda multiplicando, y multiplicar muchas veces borra. Una LSTM recuerda sumando, y por eso aguanta."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Dónde siguen vivas"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Que perdieran en lenguaje no quiere decir que no sirvan. Se siguen usando\n",
    "donde la secuencia es corta y el equipo es chico, porque una LSTM entrena en un\n",
    "portátil y un transformer no 💻"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "| Dónde | Por qué ahí sí |\n",
    "|---|---|\n",
    "| Sensores y series cortas | Pocas variables, muchos datos, secuencias de decenas de pasos y no de miles |\n",
    "| Dispositivos con poca memoria | Una LSTM chica cabe donde un transformer no |\n",
    "| Equipos sin GPU | Entrenan con lo que tengas, aunque tarden |\n",
    "| Texto largo o cualquier cosa con lenguaje | Aquí ya no. Es territorio de transformers desde 2017 |"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Y si lo que tienes es una serie de tiempo de negocio, antes de una LSTM prueba\n",
    "lo simple: para pronosticar ventas por mes, un modelo de estadística clásica le\n",
    "gana a una red casi siempre, por la misma razón del capítulo\n",
    "1. Eso está en el\n",
    "[libro de estadística desde cero](https://missyera.com/guias/estadistica-desde-cero/) 📊"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### La trampa\n",
    "\n",
    "Un equipo entrena una LSTM para clasificar comentarios de clientes. Cada comentario es una lista de palabras de largo distinto, así que para meterlos todos en una matriz los rellenan con ceros hasta el más largo. Esto es lo que corre.\n",
    "\n",
    "```\n",
    "# comentarios: lista de listas de indices, de largos distintos\n",
    "largo_max = max(len(c) for c in comentarios)\n",
    "X = np.zeros((len(comentarios), largo_max), dtype=int)\n",
    "for i, c in enumerate(comentarios):\n",
    "    X[i, :len(c)] = c            # el resto se queda en 0\n",
    "\n",
    "h = np.zeros((len(comentarios), H))\n",
    "for t in range(largo_max):\n",
    "    h = paso_lstm(X[:, t], h)    # todos los comentarios avanzan a la vez\n",
    "\n",
    "prediccion = clasifica(h)        # se usa el estado del ultimo paso\n",
    "```\n",
    "\n",
    "**Qué está mal**\n",
    "\n",
    "El estado que se usa al final es el del paso `largo_max`, y para casi todos los comentarios ese paso es relleno. Un comentario de seis palabras en un lote donde el más largo tiene ochenta pasa setenta y cuatro pasos procesando ceros, y en cada uno de esos pasos la puerta de olvido sigue borrando lo que había guardado. Cuando llega el final, lo que la red \"recuerda\" del comentario ya se diluyó, y encima el 0 es un índice válido del vocabulario, así que la red está leyendo setenta y cuatro veces la palabra número 0 como si fuera una palabra de verdad. Se arregla de dos maneras y hay que hacer las dos: guardar el largo real de cada comentario y tomar el estado de ESE paso, y reservar el índice 0 para un relleno que no signifique nada. Es el error más común de quien empieza con secuencias, y no da error: da un modelo que acierta el 50%."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Ejercicios"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Seis, y el 4 es el que de verdad enseña. Intenta antes de abrir 💛"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 1. Cambia el largo de la memoria\n",
    "\n",
    "Prueba `H = 2` y `H = 40` en la tarea de los comentarios\n",
    "y compara el acierto."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for tamano in [2, 4, 12, 40]:\n",
    "    H = tamano\n",
    "    p = entrena(arranca(1), Xtr, ytr)\n",
    "    print(f'H={tamano:3}  acierto {acierta(p, Xte, yte):.4f}')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Ojo con una cosa al correrlo: `H` es una variable global que usan\n",
    "`arranca` y `adelante`, así que reasignarla cambia la red\n",
    "entera. Si te sale un error de formas que no cuadran, es que quedó un\n",
    "`p` viejo de otro tamaño dando vueltas."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 2. Rómpelo a propósito\n",
    "\n",
    "Quita el sumando del estado anterior, o sea deja\n",
    "`np.tanh(x @ p['Wx'] + p['b'])`, y vuelve a entrenar."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Tiene que caerse a la zona del 0,5, porque sin ese sumando ya no es una\n",
    "recurrente: es una neurona normal aplicada a la última palabra. Es la forma más\n",
    "rápida de comprobar que entendiste qué línea hace el trabajo."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 3. Alarga la frase\n",
    "\n",
    "Cambia `largo = 30` por 60 y por 100 en la medición del gradiente,\n",
    "y mira cómo crece la división."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Vas a ver que no crece: se dispara. Cada paso multiplica, así que la caída es\n",
    "geométrica, y por eso el problema no se arregla \"entrenando más\"."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 4. La prueba de memoria de verdad\n",
    "\n",
    "El experimento que cierra el capítulo, y va aparte porque tarda un par de\n",
    "minutos. La tarea: la palabra clave va al principio, después vienen treinta de\n",
    "relleno, y hay que decir cuál era la clave."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def memoria(n, relleno, semilla):\n",
    "    rng = np.random.default_rng(semilla)\n",
    "    X, y = [], []\n",
    "    for _ in range(n):\n",
    "        clave = int(rng.random() < 0.5)\n",
    "        X.append([clave] + list(rng.integers(2, 6, size=relleno)))\n",
    "        y.append(clave)\n",
    "    return X, np.array(y)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Entrena las dos redes sobre eso, con 200 filas, 60 vueltas y paso 0,1, y\n",
    "prueba tres semillas de arranque. Esto es lo que me salió a mí:"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "```\n",
    "recurrente  0.475  0.490  0.500   aciertan 0/3\n",
    "LSTM        1.000  0.475  1.000   aciertan 2/3\n",
    "```"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "La recurrente **no aprende nunca**, y lo comprobé subiendo a 150\n",
    "vueltas: se queda igual. No es que le falte entrenamiento, es que el gradiente\n",
    "no llega. La LSTM lo resuelve perfecto en dos de tres arranques, y en el tercero\n",
    "se queda plantada, que es exactamente lo mismo que nos pasó con el XOR en el\n",
    "capítulo 4: poder resolverlo y llegar a encontrarlo no son la\n",
    "misma cosa."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 5. Rellena el lote y mira qué se rompe\n",
    "\n",
    "Arregla el error de *inhomogeneous shape* de la forma ingenua, la que\n",
    "salta primero: rellenar con ceros hasta el largo máximo."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "largo_max = max(len(f) for f in lote)\n",
    "X = np.zeros((len(lote), largo_max), dtype=int)\n",
    "for i, f in enumerate(lote):\n",
    "    X[i, :len(f)] = f\n",
    "print(X)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Ahora mira la matriz que sale y contesta dos cosas antes de seguir: qué\n",
    "palabra del vocabulario es el índice 0, y qué pasa si la red la lee como si fuera\n",
    "una palabra más. Si te cuesta, la respuesta entera está en la trampa de este\n",
    "capítulo."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### 6. Cierra la puerta de olvido\n",
    "\n",
    "Cambia `p['b'][:H] = 1.0` por `-1.0` y repite la\n",
    "medición del gradiente de la LSTM."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Con el sesgo negativo la puerta arranca casi cerrada, la memoria se borra en\n",
    "cada paso y la LSTM se parece mucho más a la recurrente. Ese `1.0`\n",
    "es de las líneas que más deciden y menos se explican."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Comprueba que lo tienes\n",
    "\n",
    "Tienes comentarios de clientes de unas 200 palabras y un servidor sin tarjeta gráfica. ¿Qué haces?\n",
    "\n",
    "a) Empezar por una bolsa de palabras y medirla\n",
    "\n",
    "b) Una LSTM, que aguanta las dependencias largas\n",
    "\n",
    "c) Un transformer, que es lo que se usa hoy\n",
    "\n",
    "d) Una recurrente simple, que es más chica\n",
    "\n",
    "---\n",
    "\n",
    "**La correcta es la a.**\n",
    "\n",
    "*b)* Puede ser el paso dos. Pero antes de escribir una red hay algo que contestar: ¿el orden importa para tu pregunta, o no?\n",
    "\n",
    "*c)* En un servidor sin GPU y con datos tuyos, entrenarlo desde cero no es realista. Y sigue faltando el paso previo.\n",
    "\n",
    "*d)* Con 200 palabras es justo donde el gradiente se desvanece. Si eliges red, la LSTM antes que esta.\n",
    "\n",
    "Si la bolsa de palabras te da 0,85 y te alcanza, ahí se acabó el proyecto y ese es el mejor final posible."
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Lo que hay que llevarse"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "- 🔁 Una recurrente es la neurona de siempre más un sumando: el estado\n",
    "anterior.\n",
    "\n",
    "- 📉 Hacia atrás, cada paso multiplica, y multiplicar treinta veces divide\n",
    "entre 672.620. Por eso no aprende lo lejano.\n",
    "\n",
    "- 🚪 La LSTM le da a la memoria un carril donde solo se suma y se borra, y la\n",
    "misma cuenta baja a 63.\n",
    "\n",
    "- 🐌 Y aun así perdieron, porque leen de a una palabra y eso no se\n",
    "paraleliza.\n",
    "\n",
    "- 🎯 Que es justo lo que arregla el capítulo 16.\n",
    "\n",
    "Si el numpy de este capítulo te costó más que la idea, eso es de Python y se\n",
    "arregla en el [libro de Python desde cero](https://missyera.com/guias/python-desde-cero/) 🐍\n",
    "\n",
    "Y si quieres el vocabulario suelto, está definido en dos líneas por término en\n",
    "el [glosario de IA](https://missyera.com/glosario-ia/) 📖\n",
    "\n",
    "Que tengas lindo día! 🌸"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "---\n",
    "\n",
    "Ese era el capítulo 15 de **Deep learning desde cero**. El texto completo, con las salidas de cada bloque, está en https://missyera.com/guias/deep-learning-desde-cero/recurrentes-y-lstm/\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
}
