Seq2seq: de secuencia a secuencia

Seq2seq: de secuencia a secuencia

25 min read

El modelo de la lección anterior, a nivel de carácter, sabía prolongar un texto: le dabas un arranque y lo seguía, carácter a carácter, hasta donde quisieras. Pero sólo eso, prolongar. No sabía leer una secuencia entera y responder con otra distinta: acortar un párrafo hasta un titular, corregir una línea llena de faltas, pasar una frase de un idioma a otro. Todas esas tareas comparten una forma que aquel modelo no tenía —entra una secuencia y sale otra, que no es su continuación y ni siquiera tiene por qué medir lo mismo— y esa forma es la que esta lección viene a construir.

La respuesta se llama modelo de secuencia a secuencia (seq2seq), y su idea cabe en una frase: usa dos redes en vez de una. La primera, el encoder, lee la secuencia de entrada de principio a fin y la comprime en un vector; la segunda, el decoder, arranca de ese vector y escribe la secuencia de salida, un token cada vez, como escribía el modelo de la lección anterior. Entre las dos no cruza nada más que ese vector: todo lo que el encoder entendió de la entrada tiene que caber ahí.

Piensa en la tarea más simple que tiene esta forma: invertir. La entrada es roma y la salida amor, la misma tira leída del revés. Para escribir la primera letra de la salida —la a— el decoder necesita saber cuál fue la última de la entrada; para escribir la última —la r— necesita recordar la primera. Es decir: para invertir hay que haber leído la entrada entera antes de empezar a escribir, y eso es justo lo que separa las dos redes. El encoder lee roma hasta el final y deja su resumen; sólo entonces el decoder empieza a producir amor.

Dos filas de celdas unidas por una sola caja verde en el centro. La fila de la izquierda, el encoder, tiene tres celdas etiquetadas h con barra sub uno, sub dos y sub T equis; cada una recibe por debajo un token de entrada, x sub uno, x sub dos, x sub T equis, y las celdas se enlazan de izquierda a derecha. La última celda del encoder entra, por una flecha verde, en la caja verde central etiquetada c, con la nota «d sub h numeros» debajo. De la caja c sale otra flecha verde hacia la fila de la derecha, el decoder, con tres celdas etiquetadas s sub uno, s sub dos y s sub tres, tambien enlazadas de izquierda a derecha. Cada celda del decoder recibe por debajo una entrada —GO en la primera, y con circunflejo sub uno y sub dos en las siguientes— y saca por arriba una salida, y con circunflejo sub uno, sub dos y sub tres. Toda la informacion que cruza del encoder al decoder pasa por la unica caja c.
Todo lo que el encoder leyó de la entrada cruza al decoder por un solo sitio, el vector de contexto c. Sea la entrada de tres tokens o de treinta, el decoder recibe el mismo puñado de números.

Ese resumen es el último estado del encoder, y tiene nombre: el vector de contexto, c\mathbf{c}. Es un vector de tamaño fijo —tantas coordenadas como tenga el estado, ni una más— y es la única cosa que el decoder llega a ver de la entrada. El decoder no vuelve a mirar roma: arranca de c\mathbf{c} y, a partir de ahí, escribe realimentándose lo que va sacando, igual que generaba el modelo de la lección anterior.

El encoder resume, el decoder escribe

Las dos redes son la RNN vanilla —red neuronal recurrente (recurrent neural network, RNN)—, sin novedad en la recurrencia; lo único nuevo es cómo se encadenan. El encoder lee los tokens de la entrada x1,,xTx\mathbf{x}_1, \dots, \mathbf{x}_{T_x} —cada xj\mathbf{x}_j el one-hot del jj-ésimo símbolo, como en la lección anterior— y va actualizando su estado:

hˉj=tanh(Wxhencxj+Whhenchˉj1+bhenc),\bar{\mathbf{h}}_j = \tanh\left(\mathbf{W}^{\text{enc}}_{xh}\mathbf{x}_j + \mathbf{W}^{\text{enc}}_{hh}\bar{\mathbf{h}}_{j-1} + \mathbf{b}^{\text{enc}}_h\right),

con hˉ0=0\bar{\mathbf{h}}_0 = \mathbf{0}. La barra sobre hˉj\bar{\mathbf{h}}_j lo marca como estado del encoder, para no confundirlo con el del decoder; el superíndice enc\text{enc} hace lo mismo con sus pesos, y no es un índice de capa —el encoder no tiene capas que numerar— sino una etiqueta de a qué red pertenecen. Cuando el encoder termina de leer, su último estado es el vector de contexto:

c=hˉTxRdh.\mathbf{c} = \bar{\mathbf{h}}_{T_x} \in \mathbb{R}^{d_h}.

Fíjate en el tamaño de c\mathbf{c}: dhd_h coordenadas, las del estado, y ese número no depende de TxT_x. Una entrada de tres tokens y una de treinta producen un c\mathbf{c} del mismo tamaño.

El decoder es otra RNN, con sus propios pesos, que empieza donde el encoder terminó: su estado inicial es s0=c\mathbf{s}_0 = \mathbf{c}. A partir de ahí produce la salida paso a paso. En el paso ii lee un token —el símbolo de arranque GO en el primer paso, y después el token anterior— y actualiza su estado:

si=tanh(Wxhdecoui1+Whhdecsi1+bhdec),\mathbf{s}_i = \tanh\left(\mathbf{W}^{\text{dec}}_{xh}\mathbf{o}_{u_{i-1}} + \mathbf{W}^{\text{dec}}_{hh}\mathbf{s}_{i-1} + \mathbf{b}^{\text{dec}}_h\right),

donde oui1\mathbf{o}_{u_{i-1}} es el one-hot del token ui1u_{i-1} que entra en la posición ii. De cada estado cuelga una distribución sobre el próximo símbolo, como en el modelo de lenguaje de la lección anterior:

y^i=softmax(Whysi+by),\hat{\mathbf{y}}_i = \text{softmax}\left(\mathbf{W}_{hy}\mathbf{s}_i + \mathbf{b}_y\right),

y la pérdida suma una entropía cruzada por posición de salida, promediada por token:

=1Tyi=1Tylogy^i,yi,\ell = \frac{1}{T_y}\sum_{i=1}^{T_y} -\log \hat{y}_{i,\,y_i},

con yiy_i el símbolo que de verdad tocaba en la posición ii y TyT_y el largo de la salida, que no tiene por qué igualar a TxT_x. Lo que hay que retener es dónde vive esa pérdida: sólo en el decoder. El encoder no predice nada en ninguna de sus posiciones, así que no aporta ni un sumando a \ell. En el modelo de la lección anterior cada posición tenía su salida y su término de pérdida; aquí, la mitad de la red no tiene ninguno.

Una nota sobre el token ui1u_{i-1} que el decoder lee en cada paso. Al entrenar se le da el símbolo correcto —el de la salida real, aunque el decoder hubiera fallado el suyo—; a esto se le llama teacher forcing, y hace el entrenamiento más estable. Al generar no hay respuesta correcta que darle, así que se le realimenta su propia salida. Es la misma distinción entre entrenar y generar de la lección anterior, con GO en el arranque en lugar de un espacio.

Cómo aprende el encoder sin salida propia

Si el encoder no tiene pérdida, ¿de dónde saca su gradiente? De un solo sitio —y por eso este modelo entrena de punta a punta pese a tener dos redes—. La vuelta empieza en la única pérdida que hay, la del decoder, y paso a paso es la de la lección sobre el modelo de lenguaje a nivel de carácter, sin cambiar una línea: cada posición arranca su error desde su salida y lo transporta al paso anterior. Lo nuevo pasa al final de esa vuelta. El primer estado del decoder es s0=c\mathbf{s}_0 = \mathbf{c}, así que cuando el error termina de recorrer el decoder hacia atrás, deja un gradiente sobre el vector de contexto:

c=(Whhdec)δ1,\nabla_{\mathbf{c}}\ell = \left(\mathbf{W}^{\text{dec}}_{hh}\right)^{\top}\boldsymbol{\delta}_1,

donde δ1\boldsymbol{\delta}_1 es el error del primer paso del decoder, el mismo de la lección sobre el modelo de lenguaje a nivel de carácter. Ese c\nabla_{\mathbf{c}}\ell es el único gradiente que el encoder recibe, y le basta: como c=hˉTx\mathbf{c} = \bar{\mathbf{h}}_{T_x}, es exactamente el error del último estado del encoder, y desde él el encoder corre su propia BPTT —la vuelta por el tiempo (backpropagation through time, BPTT) de la lección que la derivó— como si c\nabla_{\mathbf{c}}\ell fuese la señal que allí entraba por la salida. Repartido hacia atrás por sus TxT_x pasos, ese único gradiente ajusta Whhenc\mathbf{W}^{\text{enc}}_{hh}, Wxhenc\mathbf{W}^{\text{enc}}_{xh} y bhenc\mathbf{b}^{\text{enc}}_h.

Aquí asoma ya el límite. El decoder entero cuelga de c\mathbf{c}, y el encoder entero se corrige a través de c\mathbf{c}: si algo de la entrada no quedó en esas dhd_h coordenadas, ni el decoder puede usarlo ni el gradiente puede pedir que se guarde. Cuanto más larga es la entrada, más ha de apretujarse en el mismo vector.

Ver la vuelta completa, del decoder al encoder

La vuelta del decoder es la de la lección sobre el modelo de lenguaje a nivel de carácter, con s\mathbf{s} en lugar de h\mathbf{h}: el error de la preactivación en cada posición es

δi=(1sisi)(Why(y^iyi)+(Whhdec)δi+1),\boldsymbol{\delta}_i = \left(\mathbf{1} - \mathbf{s}_i \odot \mathbf{s}_i\right) \odot \left(\mathbf{W}_{hy}^{\top}\left(\hat{\mathbf{y}}_i - \mathbf{y}_i\right) + \left(\mathbf{W}^{\text{dec}}_{hh}\right)^{\top}\boldsymbol{\delta}_{i+1}\right),

con δTy+1=0\boldsymbol{\delta}_{T_y+1} = \mathbf{0}. Recorridos los TyT_y pasos, el transporte llega a s0=c\mathbf{s}_0 = \mathbf{c} y se detiene ahí: c=(Whhdec)δ1\nabla_{\mathbf{c}}\ell = (\mathbf{W}^{\text{dec}}_{hh})^{\top}\boldsymbol{\delta}_1. Ese vector arranca la vuelta del encoder, que es la de la lección sobre BPTT pero sin término de salida —el encoder no predice en ningún paso, así que el paréntesis con y^\hat{\mathbf{y}} desaparece y sólo queda el transporte—:

δˉj=(1hˉjhˉj)(Whhenc)δˉj+1,\bar{\boldsymbol{\delta}}_j = \left(\mathbf{1} - \bar{\mathbf{h}}_j \odot \bar{\mathbf{h}}_j\right) \odot \left(\mathbf{W}^{\text{enc}}_{hh}\right)^{\top}\bar{\boldsymbol{\delta}}_{j+1},

arrancando en el último paso con δˉTx=(1cc)c\bar{\boldsymbol{\delta}}_{T_x} = \left(\mathbf{1} - \mathbf{c} \odot \mathbf{c}\right) \odot \nabla_{\mathbf{c}}\ell. Ese único punto de entrada, c\mathbf{c}, es todo el contacto del encoder con la pérdida.

Invertir secuencias en NumPy

Vamos a entrenar un seq2seq en la tarea de juguete: invertir tiras de símbolos. De juguete a propósito —una traducción de verdad no cabe en el navegador—, pero con la forma entera del problema: dos redes, un vector de contexto en medio, y una salida que no es la continuación de la entrada. El alfabeto son seis letras, más un símbolo GO que sólo entra al decoder para arrancar. Cada ejemplo es una tira aleatoria de entre 3 y 7 letras y su reverso. La celda monta las dos redes, encadena las dos vueltas —la del decoder y, colgada de c\mathbf{c}, la del encoder— y entrena. Mira la pérdida por carácter que imprime cada pocos cientos de pasos.

import numpy as np

# Tarea de juguete: leer una tira de letras y escribirla al reves.
V = list("abcdef") # las 6 letras: entrada del encoder y salida del decoder
GO = 6 # simbolo de arranque; solo entra al decoder, nunca se predice
n_v, n_sim = 7, 6 # 7 entradas (letras + GO); 6 clases de salida
d_h, eta, theta, iters = 64, 0.10, 5.0, 1200

rng = np.random.default_rng(0) # semilla fija: veras estos numeros
sc = 0.1
Wxh_e = rng.normal(size=(d_h, n_v)) * sc; Whh_e = rng.normal(size=(d_h, d_h)) * sc; bh_e = np.zeros(d_h)
Wxh_d = rng.normal(size=(d_h, n_v)) * sc; Whh_d = rng.normal(size=(d_h, d_h)) * sc; bh_d = np.zeros(d_h)
Why = rng.normal(size=(n_sim, d_h)) * sc; by = np.zeros(n_sim)
pesos = [Wxh_e, Whh_e, bh_e, Wxh_d, Whh_d, bh_d, Why, by]


def paso(x, y): # x: la fuente; y: la fuente al reves
L = len(x)
He = np.zeros((d_h, L)); h = np.zeros(d_h) # --- encoder: lee y resume ---
for j in range(L):
h = np.tanh(Wxh_e[:, x[j]] + Whh_e @ h + bh_e); He[:, j] = h
c = h # vector de contexto: el ultimo estado
ent = np.concatenate(([GO], y[:-1])) # decoder con teacher forcing: GO, y1, ..., y_{L-1}
S = np.zeros((d_h, L)); P = np.zeros((n_sim, L)); s = c.copy(); loss = 0.0
for i in range(L): # --- decoder: escribe desde c ---
s = np.tanh(Wxh_d[:, ent[i]] + Whh_d @ s + bh_d); S[:, i] = s
o = Why @ s + by; o -= o.max(); p = np.exp(o); p /= p.sum(); P[:, i] = p
loss -= np.log(p[y[i]] + 1e-12)
gWhy = np.zeros_like(Why); gby = np.zeros_like(by) # --- vuelta del decoder ---
gWxh_d = np.zeros_like(Wxh_d); gWhh_d = np.zeros_like(Whh_d); gbh_d = np.zeros_like(bh_d)
ds = np.zeros(d_h)
for i in reversed(range(L)):
dO = P[:, i].copy(); dO[y[i]] -= 1.0
gWhy += np.outer(dO, S[:, i]); gby += dO
dp = (1 - S[:, i] ** 2) * (Why.T @ dO + ds)
s_ant = S[:, i - 1] if i > 0 else c
gWhh_d += np.outer(dp, s_ant); gWxh_d[:, ent[i]] += dp; gbh_d += dp
ds = Whh_d.T @ dp
gWxh_e = np.zeros_like(Wxh_e); gWhh_e = np.zeros_like(Whh_e); gbh_e = np.zeros_like(bh_e)
dh = ds # unico enlace: el gradiente en c
for j in reversed(range(L)): # --- vuelta del encoder ---
dp = (1 - He[:, j] ** 2) * dh
h_ant = He[:, j - 1] if j > 0 else np.zeros(d_h)
gWhh_e += np.outer(dp, h_ant); gWxh_e[:, x[j]] += dp; gbh_e += dp
dh = Whh_e.T @ dp
return loss / L, [g / L for g in (gWxh_e, gWhh_e, gbh_e, gWxh_d, gWhh_d, gbh_d, gWhy, gby)]


suave = None
for it in range(iters):
L = int(rng.integers(3, 8)) # longitudes 3..7
x = rng.integers(0, n_sim, size=L); y = x[::-1].copy()
loss, grads = paso(x, y)
for g in grads:
nrm = np.linalg.norm(g)
if nrm > theta: g *= theta / nrm # recorte del gradiente
for Pm, g in zip(pesos, grads):
Pm -= eta * g
suave = loss if suave is None else 0.99 * suave + 0.01 * loss
if it % 300 == 0 or it == iters - 1:
print("iter %4d perdida/car %.3f" % (it, suave))
numpy

La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.

La pérdida arranca cerca de ln61.79\ln 6 \approx 1.79 —lo que vale repartir la probabilidad por igual entre las seis letras, el eco del ln32\ln 32 del modelo de la lección anterior con seis caras en vez de treinta y dos— y baja hasta rondar 0.750.75. No llega a cero, y no debería: el bucle se corta pronto para caber en los pocos segundos que el navegador da a una celda, y aun entrenado del todo este modelo tiene un techo que la siguiente celda deja ver. Con los pesos entrenados, le pedimos que invierta tiras que no vio en el entrenamiento, y medimos cuántos caracteres acierta según la longitud.

import numpy as np

# Necesita la celda anterior: los pesos entrenados, V, n_sim.
def escribe(x): # genera sin teacher forcing: se realimenta
h = np.zeros(d_h)
for j in range(len(x)):
h = np.tanh(Wxh_e[:, x[j]] + Whh_e @ h + bh_e)
s = h; ix = GO; out = []
for _ in range(len(x)):
s = np.tanh(Wxh_d[:, ix] + Whh_d @ s + bh_d)
ix = int((Why @ s + by).argmax()); out.append(ix) # el mas probable (voraz)
return out


ej = np.random.default_rng(7)
print("Unos ejemplos, de corto a largo:")
for L in [3, 4, 8, 12]:
x = ej.integers(0, n_sim, size=L)
got = escribe(x)
print(" %-9s -> quiero %-9s y escribe %-9s %s" % (
"".join(V[i] for i in x), "".join(V[i] for i in x[::-1]),
"".join(V[i] for i in got), "OK" if list(got) == list(x[::-1]) else ""))

ac = np.random.default_rng(123)
print("\nAcierto por longitud (150 tiras nuevas cada una):")
for L in range(3, 9):
bien = tot = 0
for _ in range(150):
x = ac.integers(0, n_sim, size=L); y = x[::-1]
bien += int((np.array(escribe(x)) == y).sum()); tot += L
print(" L=%d %.0f%% de caracteres" % (L, 100 * bien / tot))
numpy

La primera ejecución descarga el intérprete de Python (~15 MB). Después queda en la caché del navegador.

Las tiras cortas salen bien —las de tres o cuatro letras, a menudo exactas—; las largas se tuercen, y se tuercen de una forma que dice de dónde viene el problema. El principio de lo que el decoder escribe —que es el final de lo que el encoder leyó, lo más reciente en c\mathbf{c}— suele estar bien; el final —que le pediría recordar el principio de la entrada, lo más lejano— es donde se desordena. La tabla lo cuantifica: cerca del acierto perfecto en longitud 3, y alrededor de la mitad de los caracteres en las más largas. No es que falte entrenamiento. La tira larga y la corta llegan al decoder como el mismo c\mathbf{c}, los mismos dhd_h números; en la larga hay más que meter en el mismo sitio, y algo se queda fuera.

Comprueba tu intuición

Cuatro preguntas: en qué se distingue este modelo del de la lección anterior, cuántos números lleva el resumen, cómo aprende el encoder sin salida propia, y qué arreglaría —y qué no— entrenar mucho más.

El modelo de la lección anterior, a nivel de carácter, prolongaba un texto en marcha. ¿Qué hace distinto un modelo de secuencia a secuencia?

El encoder tiene un estado de dh=64d_h = 64 coordenadas y lee una frase de Tx=30T_x = 30 tokens. ¿Cuántos números recibe el decoder como resumen de esa frase entera?

A margin of ±0 is accepted.

El encoder no tiene salida propia: no predice nada en ninguna de sus posiciones. Aun así, se entrena. ¿Por qué caminos le llega el gradiente? Marca todo lo que sea cierto.

Select every correct option. This is graded all-or-nothing: there is no partial credit.

El inversor acierta casi todos los caracteres en las tiras cortas y baja hasta rondar la mitad en las largas. Si lo entrenaras mucho más, ¿qué cabría esperar?


Con esto el bloque cierra. La RNN vanilla leía cualquier longitud; la vuelta por el tiempo de la lección sobre BPTT la entrenaba; las compuertas de la LSTM y la GRU le daban una memoria capaz de aguantar la distancia. El modelo de secuencia a secuencia los pone a todos a trabajar juntos en la primera tarea del curso que transforma una secuencia en otra distinta, en vez de sólo continuarla o colgarle un veredicto.

Pero la costura entre las dos redes es un solo vector. Todo lo que el encoder leyó —tres tokens o treinta— tiene que caber en los dhd_h números de c\mathbf{c}, y por eso las tiras largas se deshilachan: una frase entera no cabe donde cabía una palabra. Ese es exactamente el problema con el que arranca el bloque 4, sobre el puente hacia la atención. En lugar de exprimir la entrada en un único resumen y no volver a mirarla, la idea será dejar que el decoder vuelva la vista sobre toda la secuencia leída y elija, en cada paso que escribe, qué parte de ella necesita. Cómo se hace eso —y por qué reordena el curso entero— es lo que viene.

Further reading2 sources · 2 papers

Where this lesson comes from, and where to go next. None of it is needed to carry on with the course.

  • Sequence to Sequence Learning with Neural Networks
    paperSutskever, Vinyals y Le, 2014arXiv:1409.3215EN

    La arquitectura de esta lección: un LSTM comprime la entrada en un vector fijo y otro escribe la salida desde él. Invierten la frase de entrada, y —al revés que tu inversor— informan de que las frases largas no les dieron problema.

  • On the Properties of Neural Machine Translation: Encoder-Decoder Approaches
    paperCho, van Merriënboer, Bahdanau y Bengio, 2014arXiv:1409.1259EN

    Mide el cuello de botella de tu última celda: con el vector de contexto de tamaño fijo, la calidad cae en picado según crece la frase de entrada. El arreglo —dejar que el decoder vuelva a mirar la entrada— es el bloque 4.