Capítulo 11 de 15 7 secciones 15 min

La atención, que es la idea que cambió todo

Cada palabra mira a las demás y se queda con lo que le sirve. Escrita en seis líneas de numpy.

La atención hace que cada palabra mire a todas las demás y se quede con una mezcla de lo que le sirve. Se escribe en seis líneas: tres proyecciones (consulta, clave y valor), un producto entre consultas y claves dividido por la raíz de la dimensión, un softmax y una mezcla. Y arregla lo que quedó mal en el capítulo 10: el vector de "llego" deja de ser fijo y pasa a valer 0,3808 de parecido entre dos frases distintas.

El capítulo 10 terminó con tres cosas mal: los antónimos salían iguales, se perdía el orden, y cada palabra tenía un vector fijo pasara lo que pasara 🤔

Las tres las arregla la misma idea, publicada en 2017 en un paper que se llama Attention Is All You Need. Y cabe en seis líneas de numpy.

La idea, con una frase

Toma "la bodega de lima no pago". Para entender qué significa pago en esa frase hace falta mirar el no. Y para saber de qué bodega hablamos, hay que mirar lima.

Eso es la atención: cada palabra mira a todas las demás y se queda con una mezcla de lo que le sirve. No un vector fijo, sino uno armado para esa frase.

Y cómo decide a quién mirar es lo bonito. Cada palabra genera tres cosas:

  • 🔎 Una consulta: qué estoy buscando.
  • 🏷️ Una clave: qué ofrezco yo.
  • 📦 Un valor: qué me llevo si me eligen.

Se comparan todas las consultas con todas las claves, y eso decide los pesos de la mezcla. Es literalmente una búsqueda, con la diferencia de que en vez de elegir un resultado se lleva un poquito de todos.

En la frase el banco me cobro, la palabra banco compara su consulta con todas las demas palabras: da peso alto a cobro y peso bajo a me y a el, y con eso resuelve que se trata de una entidad financiera.
Toda la atencion es esto: cada palabra mira a las demas a la vez y decide a cuales hacer caso para entenderse a si misma. No recorre la frase, la mira entera de golpe, y por eso se puede entrenar en paralelo y la distancia entre palabras deja de importar.

Las seis líneas

Atención(Q,K,V)=softmax(QKdk)V

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

import numpy as np

def softmax(z):
    z = z - z.max(axis=-1, keepdims=True)      # el truco del capítulo 3
    e = np.exp(z)
    return e / e.sum(axis=-1, keepdims=True)

frase = ['la', 'bodega', 'de', 'lima', 'no', 'pago']
rng = np.random.default_rng(7)
d = 8                                          # tamaño de cada vector

E = rng.normal(0, 1, (len(frase), d))          # los embeddings del capítulo 10
Wq = rng.normal(0, 0.5, (d, d))
Wk = rng.normal(0, 0.5, (d, d))
Wv = rng.normal(0, 0.5, (d, d))

Q = E @ Wq                                     # consultas
K = E @ Wk                                     # claves
V = E @ Wv                                     # valores

puntajes = Q @ K.T / np.sqrt(d)                # quién le interesa a quién
A = softmax(puntajes)                          # convertidos en pesos que suman 1
salida = A @ V                                 # la mezcla

print('formas:', Q.shape, K.shape, V.shape)
print('matriz de atención:', A.shape, ' salida:', salida.shape)
print('cada fila suma:', np.round(A.sum(axis=1), 4))
formas: (6, 8) (6, 8) (6, 8)
matriz de atención: (6, 6)  salida: (6, 8)
cada fila suma: [1. 1. 1. 1. 1. 1.]

Eso es todo. Tres multiplicaciones para sacar Q, K y V, una para compararlos, un softmax y una mezcla 🎯

Y ojo con un detalle que parece decorativo: dividir por la raíz de d. Sin eso, con vectores largos los puntajes salen enormes, el softmax se satura (capítulo 3) y los pesos quedan en 1 y 0. La raíz los mantiene en un rango donde el softmax todavía tiene pendiente.

Matriz de atención de la frase el pedido de la bodega llegó tarde, con cuadros más oscuros donde una palabra mira a otra. El triángulo superior está vacío por la máscara.
Cada fila es una palabra preguntando y cada columna una palabra respondiendo. El triángulo de arriba está vacío porque ninguna palabra puede mirar a las que todavía no se escribieron, y eso es lo que permite generar texto.
Matriz de atención de la frase el banco me cobró una comisión consigo misma. Cada fila suma 1. La fila de banco tiene sus valores altos en cobró y en comisión, y la fila de cobró lo tiene en banco.
Los números de la atención, que casi nunca se enseñan. Lee la fila de banco: reparte su atención y le da la mayor parte a cobró y a comisión. Ese reparto no se lo puso nadie, salió de entrenar, y es literalmente lo que hace que el modelo entienda de qué banco hablas.

Mirar la matriz de atención

print('        ' + '  '.join(f'{p:>8}' for p in frase))
for p, fila in zip(frase, A):
    print(f'{p:8}' + '  '.join(f'{x:8.3f}' for x in fila))
              la    bodega        de      lima        no      pago
la         0.129     0.236     0.052     0.146     0.271     0.168
bodega     0.103     0.194     0.048     0.078     0.506     0.071
de         0.010     0.008     0.625     0.339     0.002     0.016
lima       0.045     0.026     0.360     0.261     0.002     0.306
no         0.011     0.011     0.727     0.217     0.033     0.002
pago       0.167     0.242     0.001     0.007     0.103     0.480

Cada fila es una palabra y dice cuánto mira a cada una de las seis. Las filas suman 1, o sea que cada palabra reparte un total fijo de atención.

Y ahora la parte honesta: estos números no significan nada. Las matrices Wq, Wk y Wv las saqué al azar, así que la atención que ves es aleatoria. Que "bodega" mire a "no" con 0,506 es casualidad.

Lo que sí es real y es lo que importa aquí es la maquinaria: las formas, que las filas sumen 1, y que cada palabra reciba una mezcla distinta. En un modelo entrenado, esas tres matrices se aprenden con la retropropagación del capítulo 6 igual que cualquier otro peso, y ahí los patrones sí significan 🔧

La máscara, que es lo que hace que ChatGPT escriba

Si el modelo tiene que predecir la palabra siguiente, no puede dejar que cada palabra mire a las que vienen después: sería copiarse la respuesta.

Se arregla poniendo menos infinito en los puntajes de las palabras futuras antes del softmax:

mascara = np.triu(np.ones((len(frase), len(frase))), 1) * -1e9
A_causal = softmax(puntajes + mascara)

print('        ' + '  '.join(f'{p:>8}' for p in frase))
for p, fila in zip(frase, A_causal):
    print(f'{p:8}' + '  '.join(f'{x:8.3f}' for x in fila))
              la    bodega        de      lima        no      pago
la         1.000     0.000     0.000     0.000     0.000     0.000
bodega     0.347     0.653     0.000     0.000     0.000     0.000
de         0.016     0.013     0.971     0.000     0.000     0.000
lima       0.065     0.037     0.519     0.378     0.000     0.000
no         0.011     0.011     0.728     0.217     0.033     0.000
pago       0.167     0.242     0.001     0.007     0.103     0.480

Mira la forma de triángulo 🔺

La primera palabra solo puede mirarse a sí misma, así que su atención es 1,000 y el resto ceros. La segunda mira a dos, la tercera a tres. Y la última las ve todas.

Ese triángulo es la diferencia entre un modelo que lee y uno que escribe. Los que escriben (ChatGPT, Claude) llevan esta máscara, y por eso van palabra por palabra sin poder ver lo que todavía no escribieron.

Y el -1e9 no es capricho: exp(-1e9) da cero, así que el softmax les asigna peso cero. Poner directamente -inf también funciona, pero da nan si una fila entera queda enmascarada 🧯

Cuatro ondas de distinta frecuencia recorriendo las posiciones de una secuencia. Las de frecuencia alta cambian rápido entre posiciones vecinas y las de frecuencia baja cambian despacio a lo largo de toda la secuencia.
La codificación posicional es esto: en cada posición se toma el valor de todas las ondas a la vez, y esa combinación no se repite. Las ondas rápidas distinguen palabras vecinas y las lentas ubican en qué zona del texto estás.

Lo que arregla, medido

Aquí está el pago de lo que quedó pendiente en el capítulo 10. Tomamos la misma palabra, llego, en dos frases que dicen lo contrario:

vocabulario = {'el': 0, 'pedido': 1, 'llego': 2, 'completo': 3, 'no': 4}
tabla = rng.normal(0, 1, (5, d))               # un vector fijo por palabra

def atiende(palabras):
    Emb = tabla[[vocabulario[p] for p in palabras]]
    Q, K, V = Emb @ Wq, Emb @ Wk, Emb @ Wv
    return softmax(Q @ K.T / np.sqrt(d)) @ V

f1 = ['el', 'pedido', 'llego', 'completo']
f2 = ['el', 'pedido', 'no', 'llego']

def coseno(a, b):
    return float(a @ b / (np.linalg.norm(a) * np.linalg.norm(b)))

fijo = coseno(tabla[vocabulario['llego']], tabla[vocabulario['llego']])
tras = coseno(atiende(f1)[f1.index('llego')], atiende(f2)[f2.index('llego')])
print('el vector fijo de "llego" en las dos frases:', round(fijo, 4))
print('después de la atención                     :', round(tras, 4))
el vector fijo de "llego" en las dos frases: 1.0
después de la atención                     : 0.3808

Ahí está 🎉

Con el embedding del capítulo 10, "llego" es el mismo vector en las dos frases: parecido 1,0, siempre, pase lo que pase.

Después de la atención, los dos "llego" tienen un parecido de 0,3808. El de "el pedido llego completo" y el de "el pedido no llego" pasaron a ser vectores distintos, porque cada uno se mezcló con las palabras que lo rodean.

Eso es lo que quiere decir representación contextual, y es literalmente el motivo de que los modelos de lenguaje de hoy funcionen y los de hace diez años no.

Y de paso arregla el orden: como la mezcla depende de con quién está cada palabra, "no llego" y "llego no" ya no dan lo mismo.

Ejercicios

1. Qué pasa sin la raíz de d

Quita la división y mira los pesos.

sin_raiz = softmax(Q @ K.T)
print('con raíz  :', np.round(A[1], 3))
print('sin raíz  :', np.round(sin_raiz[1], 3))
print()
print('el peso más grande, con raíz :', round(float(A.max()), 4))
print('el peso más grande, sin raíz :', round(float(sin_raiz.max()), 4))
con raíz  : [0.103 0.194 0.048 0.078 0.506 0.071]
sin raíz  : [0.01  0.061 0.001 0.005 0.92  0.004]

el peso más grande, con raíz : 0.7268
el peso más grande, sin raíz : 0.9682

Sin la raíz, el peso más grande sube de 0,7268 a 0,9682, y mira la fila de "bodega": pasa de repartir (0,103 0,194 0,048 0,078 0,506 0,071) a llevárselo casi todo un solo sitio (0,920).

Y no es solo que quede feo: un softmax saturado tiene pendiente cero (capítulo 3), así que por ahí ya no pasa gradiente y esa parte deja de aprender. Una división que parece cosmética y sostiene el entrenamiento 🔑

2. Varias cabezas mirando cosas distintas

La atención de verdad se hace varias veces en paralelo. Haz cuatro cabezas de 2 dimensiones cada una.

CABEZAS, d_cabeza = 4, 2
salidas = []
for c in range(CABEZAS):
    wq = rng.normal(0, 0.5, (d, d_cabeza))
    wk = rng.normal(0, 0.5, (d, d_cabeza))
    wv = rng.normal(0, 0.5, (d, d_cabeza))
    a = softmax((E @ wq) @ (E @ wk).T / np.sqrt(d_cabeza))
    salidas.append(a @ (E @ wv))
    print(f'cabeza {c}: la palabra "pago" mira más a "{frase[int(a[5].argmax())]}"')

junta = np.concatenate(salidas, axis=1)
print()
print('cada cabeza da', salidas[0].shape, 'y juntas dan', junta.shape)
cabeza 0: la palabra "pago" mira más a "de"
cabeza 1: la palabra "pago" mira más a "no"
cabeza 2: la palabra "pago" mira más a "no"
cabeza 3: la palabra "pago" mira más a "de"

cada cabeza da (6, 2) y juntas dan (6, 8)

Cuatro cabezas y cada una mira a un sitio distinto. Al final se pegan una al lado de la otra y vuelve a salir el mismo tamaño de siempre.

En un modelo entrenado, cada cabeza se especializa: unas siguen la sintaxis, otras enlazan un pronombre con su sujeto. Aquí son cuatro al azar, así que lo único real es que se pueden mirar varias cosas a la vez sin gastar más tamaño 👀

3. La atención no sabe de orden

Baraja la frase y mira si la salida de una palabra cambia.

orden = [3, 1, 0, 5, 2, 4]
E_barajado = E[orden]
Qb, Kb, Vb = E_barajado @ Wq, E_barajado @ Wk, E_barajado @ Wv
salida_b = softmax(Qb @ Kb.T / np.sqrt(d)) @ Vb

donde = orden.index(1)              # dónde quedó "bodega"
print('salida de "bodega" en la frase normal  :', np.round(salida[1][:4], 4))
print('salida de "bodega" en la frase barajada:', np.round(salida_b[donde][:4], 4))
print('¿son iguales?', np.allclose(salida[1], salida_b[donde]))
salida de "bodega" en la frase normal  : [ 0.0826 -0.1186 -0.0884 -0.3718]
salida de "bodega" en la frase barajada: [ 0.0826 -0.1186 -0.0884 -0.3718]
¿son iguales? True

Iguales. La atención mira quién está en la frase, no en qué orden.

Y eso es un problema serio, porque "la bodega no pago" y "no la bodega pago" darían lo mismo. Se arregla sumándole a cada embedding un vector que depende de su posición, y a eso se le llama codificación posicional. Es el ejercicio 4 🔢

4. Meterle la posición

Suma senos y cosenos de distinta frecuencia según la posición, que es como se hizo en el paper original.

def posiciones(n, dim):
    pos = np.arange(n)[:, None]
    i = np.arange(dim)[None, :]
    angulo = pos / (10000 ** (2 * (i // 2) / dim))
    P = np.zeros((n, dim))
    P[:, 0::2] = np.sin(angulo[:, 0::2])
    P[:, 1::2] = np.cos(angulo[:, 1::2])
    return P

P = posiciones(len(frase), d)
E_con_pos = E + P
Qp, Kp, Vp = E_con_pos @ Wq, E_con_pos @ Wk, E_con_pos @ Wv
salida_p = softmax(Qp @ Kp.T / np.sqrt(d)) @ Vp

E_bar_pos = E[orden] + P
Qbp, Kbp, Vbp = E_bar_pos @ Wq, E_bar_pos @ Wk, E_bar_pos @ Wv
salida_bp = softmax(Qbp @ Kbp.T / np.sqrt(d)) @ Vbp

print('con posición, ¿siguen siendo iguales?',
      np.allclose(salida_p[1], salida_bp[orden.index(1)]))
print('parecido entre las dos:', round(coseno(salida_p[1], salida_bp[orden.index(1)]), 4))
con posición, ¿siguen siendo iguales? False
parecido entre las dos: 0.9672

Ya no son iguales, que era lo que buscábamos. Pero el parecido sigue siendo 0,9672, o sea que la posición movió muy poco.

Y eso también hay que decirlo: aquí los embeddings salen de una normal con desviación 1 y la codificación posicional va entre -1 y 1, así que es un empujoncito al lado de un vector que ya era grande. En un modelo entrenado los embeddings aprenden a dejarle sitio a esa señal, porque el entrenamiento castiga confundir el orden. Aquí no hay entrenamiento, así que solo se ve el mecanismo 🌊

Los senos y cosenos parecen una rareza y tienen su razón: dan un patrón distinto para cada posición y permiten que el modelo calcule distancias entre posiciones con sumas y restas. Hoy se usan otras variantes y la idea es la misma.

5. Cuánto cuesta la atención

Cuenta las comparaciones según lo larga que sea la entrada.

for palabras_n in [6, 100, 1_000, 100_000]:
    print(f'{palabras_n:7,} palabras: {palabras_n ** 2:15,} comparaciones')
      6 palabras:              36 comparaciones
    100 palabras:          10,000 comparaciones
  1,000 palabras:       1,000,000 comparaciones
100,000 palabras:  10,000,000,000 comparaciones

Cada palabra mira a todas, así que el costo crece al cuadrado. Con 100.000 palabras son diez mil millones de comparaciones para una sola capa, de una sola cabeza.

Ahí tienes por qué los modelos tienen un límite de contexto y por qué cuesta tanto ampliarlo. Es el problema abierto más caro del campo, y hay una industria entera buscándole la vuelta 💸

6. La máscara sobre una frase más larga

Comprueba que el triángulo escala.

larga = 5
p_larga = rng.normal(0, 1, (larga, larga))
m_larga = np.triu(np.ones((larga, larga)), 1) * -1e9
A_larga = softmax(p_larga + m_larga)

print(np.round(A_larga, 3))
print()
print('ceros por encima de la diagonal:',
      int((A_larga[np.triu_indices(larga, 1)] == 0).sum()),
      'de', len(np.triu_indices(larga, 1)[0]))
[[1.    0.    0.    0.    0.   ]
 [0.024 0.976 0.    0.    0.   ]
 [0.131 0.117 0.752 0.    0.   ]
 [0.073 0.053 0.421 0.453 0.   ]
 [0.037 0.104 0.598 0.009 0.252]]

ceros por encima de la diagonal: 10 de 10

Los diez huecos de arriba a la derecha son cero exacto, siempre.

Esa es la garantía de que el modelo no se copia del futuro. Y es una garantía de verdad, no una tendencia: matemáticamente no puede 🔒

7. El error de la máscara del tamaño equivocado

Aplica una máscara de 5 por 5 a una frase de 6 palabras.

softmax(puntajes + np.triu(np.ones((5, 5)), 1) * -1e9)
ValueError: operands could not be broadcast together with shapes (6,6) (5,5) 

Los puntajes son 6 por 6 y la máscara 5 por 5, así que numpy no sabe cómo alinearlas.

Este error es de los buenos, porque el equivalente silencioso existe y es peor: si la máscara fuera de 1 por 6 o de 6 por 1, numpy la estiraría sin quejarse y estarías enmascarando lo que no toca. El modelo entrenaría, la pérdida bajaría, y estaría viendo el futuro sin que nadie se entere 😬

Cuando montes atención a mano, imprime la forma de la máscara y la de los puntajes antes de sumarlas.

Comprueba que lo tienes

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?

  • Mezclar información de las demás palabras en la representación de esa
  • Cambiar el significado de la palabra
  • Ordenar la frase
  • Elegir qué palabra es más importante en la frase

Lo que te llevas

  • 🔎 Cada palabra genera consulta, clave y valor, y se lleva una mezcla de las demás según qué tanto le interesan.
  • 🧮 Son seis líneas de numpy y las filas de la matriz de atención suman 1.
  • 🔑 Dividir por la raíz de d evita que el softmax se sature: sin eso el peso mayor pasa de 0,7268 a 0,9682, y por ahí deja de pasar gradiente.
  • 🔺 La máscara causal es un triángulo de ceros, y es lo que separa un modelo que lee de uno que escribe.
  • 🎉 Arregla lo del capítulo 10: "llego" pasa de tener parecido 1,0 consigo mismo en cualquier frase a 0,3808 entre dos frases distintas.
  • 🔢 La atención sola no sabe de orden: barajar la frase da exactamente la misma salida. Hay que sumarle la posición.
  • 💸 Cuesta al cuadrado: 100.000 palabras son diez mil millones de comparaciones. Por eso hay límite de contexto.

En el capítulo 12 juntamos esto en un transformer y vemos qué está pasando exactamente cuando le escribes a un modelo de lenguaje.

Que tengas lindo día! 🌸

¿Tienes alguna duda o consulta?