{
 "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 práctica 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",
    "Los ejercicios están al final y traen una celda vacía debajo de cada uno. Las\n",
    "respuestas viven en el cuaderno de soluciones, y merece la pena pelearse un\n",
    "rato antes de abrirlo 💛"
   ]
  },
  {
   "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?** La respuesta está en el cuaderno de soluciones. Míralo tú primero."
   ]
  },
  {
   "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": [
    "# tu turno"
   ]
  },
  {
   "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": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# tu turno"
   ]
  },
  {
   "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": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# tu turno"
   ]
  },
  {
   "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": [
    "# tu turno"
   ]
  },
  {
   "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": [
    "# tu turno"
   ]
  },
  {
   "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": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# tu turno"
   ]
  },
  {
   "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"
   ]
  },
  {
   "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
}
