Mecanismo de atención — Determinación contextual del microambiente tumoral
Al finalizar este capítulo
En el capítulo #5 dejamos marcados los espacios y pospusimos la exploración del corazón del transformador, la atención (attention). Analizaremos por qué es necesaria, qué hace y cómo se calcula el contexto de 128K tokens que manejan hoy en día los LLM.
Este capítulo representa el clímax de la narrativa. Las piezas de los capítulos #1 a #5 — predicción de la siguiente palabra, 33,62 millones de perillas, descenso de gradiente, retropropagación, incrustaciones (embeddings), residuos y normalización — se unen con esta atención para completar el transformador. La ingeniería de prompts, RAG y agentes que trataremos desde el capítulo #7 también tienen como raíz las propiedades de esta atención.
Volviendo al problema del "hígado" del capítulo #1
Recordemos el ejemplo de polisemia del capítulo #1.
- "Se me ha hinchado el hígado" — órgano corporal (肝)
- "Le queda justo" — grado de salinidad
- "Después de que pasen 3 años" — transcurso del tiempo
La incrustación (embedding) de la palabra "hígado/pasa", como vimos en el capítulo #5, es un único vector fijo. Sin embargo, este vector se utiliza con significados diferentes en las tres oraciones. ¿Cómo distingue y procesa la red neuronal entrenada estos tres casos?
El método es el siguiente. Las otras palabras presentes en el contexto fluyen hacia la representación de esta "palabra" y la transforman en un significado adecuado al contexto. Si hay "hinchado", se inclina hacia el órgano corporal; si hay "queda justo", hacia la salinidad; si hay "3 años", hacia el flujo del tiempo. El componente responsable de este flujo es la atención.
Es decir, la atención es el proceso dentro del transformador mediante el cual la representación de una palabra se contextualiza (contextualize) según el contexto. Si el vector de incrustación es una expresión estática de diccionario, el vector posterior a la atención es un significado vivo dentro de la oración.
Ahora aprenderemos cómo logra esto la atención.
El fracaso del método ingenuo — El problema del promedio uniforme
El método más simple para transmitir información entre palabras. Simplemente promediar los vectores de todos los tokens de la oración.
contextualizado(hígado) = (vec(hígado) + vec(está) + vec(hinchado) + vec(ahora)) / 4Problema: Todas las palabras se reflejan de manera uniforme. Determinar el significado de "hígado" depende decisivamente de "está hinchado", mientras que "el" tiene poca importancia; sin embargo, el promedio uniforme no logra distinguir esto.
Lo que necesitamos: Un método para ajustar dinámicamente cuánto refleja cada palabra según el contexto. Por ejemplo, al determinar el significado de "hígado", asignar un 80% a "está hinchado", un 5% a "tres años" y un 5% a "el".
El cálculo de "cuánto reflejar" es lo que hace la atención. Calcula los pesos de atención para cada par de palabras (consulta-clave) y utiliza estos pesos para realizar un promedio ponderado de la información.
Patólogo que identifica células T en una lámina de patología
Una analogía biológica. Continuación con la lámina de patología vista en el episodio #2.
Supongamos que, como patólogo, estás observando una célula T en una lámina de tejido. Debes determinar su estado actual: ¿está activada, agotada (exhaustion) o es una célula T reguladora?
Esta determinación no puede basarse únicamente en la morfología de la propia célula T; es necesario observar el entorno celular circundante (el microambiente tumoral).
- Si hay células tumorales numerosas justo al lado de esta célula T → Está atacando el tumor o está inactivada.
- Si hay muchas otras células T y células dendríticas cerca → Es un sitio donde la respuesta inmune está activada.
- Si el entorno está engrosado por fibroblastos y una matriz de colágeno → Está atrapado en una barría estromal.
- Si está cerca de células T reguladoras (Treg) → Existe la posibilidad de que esté suprimida.
La importancia de cada célula circundante para esta determinación varía. La densidad de las células tumorales puede ser decisiva, mientras que el grosor del estroma es relativamente menos decisivo. El patólogo asigna inconscientemente diferentes pesos a las células circundantes para realizar la determinación.
Cuantificar esto revela la estructura exacta de la atención. Se necesitan tres componentes:
- Consulta (Query): Lo que la célula T que se está evaluando "quiere saber". Por ejemplo, "soy una célula T CD8+ y quiero saber si estoy atacando el tumor o si me he agotado".
- Clave (Key): La información que cada célula circundante "puede proporcionar". Por ejemplo, las células tumorales proporcionan "estado de expresión de PD-L1", las células dendríticas proporcionan "estado de presentación de antígeno" y las células estromales proporcionan "grado de barrera física".
- Valor (Value): El contenido real de la información transmitida. Por ejemplo, el nivel de expresión de PD-L1 en las células tumorales, el estado de activación de las células dendríticas, etc.
Se compara la Consulta de la célula T actual con la Clave de cada célula circundante para puntuar "qué tan bien coinciden". El Valor de las células que coinciden bien se refleja significativamente en la determinación de esta célula T, mientras que las que no coinciden se ignoran.
Esta es la base de la atención.
Query, Key, Value — Las tres caras de una palabra
El mecanismo de atención del Transformer sigue exactamente esta estructura triangular.
Cada vector de palabra en la oración se transforma en tres vectores derivados:
- Query (Q): "¿Qué quiere saber esta palabra en el contexto?"
- Key (K): "¿Cómo anuncia esta palabra su identidad a otras palabras?"
- Value (V): "El contenido real que transmite esta palabra"
Cada uno se obtiene multiplicando el vector de incrustación (embedding) de la palabra por tres matrices de pesos aprendidas W_Q, W_K, W_V.
Q_i = W_Q · x_i
K_i = W_K · x_i
V_i = W_V · x_ix_i es el vector de incrustación (embedding) del i-ésimo token. Estos tres vectores suelen tener una dimensión menor que el vector original (por ejemplo, de 4096 dimensiones en el vector original a 128 dimensiones para Q, K y V por separado).
Intuición. Se puede pensar en un único vector de token como si tuviera tres "rostros". El rostro Query pregunta a otros tokens: "¿Eres útil para mí?". El rostro Key responde a las preguntas de otros tokens: "Tengo esta información disponible". Y el rostro Value transmite la información real.
Es importante que los tres rostros (Q, K y V) sean diferentes. La forma en que un token se anuncia (Key) puede diferir de la forma en que investiga a otros tokens (Query), y la información que realmente transmite (Value) también puede ser distinta a ambas. Esta separación permite que la atención exprese un flujo de información muy flexible.
Fórmula de la Atención de Producto Punto Escalado
Ahora, el cálculo real. La puntuación de atención (attention score) que calcula el token i con respecto al token j:
score(i, j) = ⟨Q_i, K_j⟩ / sqrt(d_k)Dos vectores están siendo multiplicados escalarmente (producto punto). Si los dos vectores apuntan en direcciones similares, la puntuación es alta; si apuntan en direcciones diferentes, la puntuación es baja.
La razón para dividir por sqrt(d_k). Q y K son vectores de dimensión d_k, y si la dimensión de los vectores aleatorios es d_k, entonces el producto escalar se vuelve aproximadamente sqrt(d_k) más grande. Sin esta división, cuando la dimensión del vector aumenta, la puntuación se dispara, lo que hace que el softmax en el siguiente paso se desvíe hacia los extremos. Esta escala es el héroe silencioso de la estabilidad del entrenamiento del transformador. La razón exacta se deriva en el Apéndice A.2.
Ahora, desde la perspectiva de la palabra i, hemos calculado las puntuaciones para todas las palabras en la oración. Convertiremos estas puntuaciones en una distribución de probabilidad usando softmax.
α(i, j) = exp(score(i, j)) / Σ_k exp(score(i, k))Σ_j α(i, j) = 1. Cada palabra asigna el 100 % de su atención a las demás palabras.
Softmax es una función que ya se encontró en la sección A.2 del capítulo 1 y en la sección A.4 del capítulo 3. Aquí, su función consiste en convertir la competencia entre palabras en una distribución de probabilidad.
Finalmente, se realiza un promedio ponderado de los vectores Value utilizando estos pesos:
output_i = Σ_j α(i, j) · V_jEsta es la expresión posterior a la atención de la palabra i.
Escribir en forma de matriz por oración completa sería mucho más conciso.
Attention(Q, K, V) = softmax(Q · K^T / sqrt(d_k)) · VEsta línea es el corazón del transformador. El título del artículo original "Attention Is All You Need" se refiere a esta ecuación.
Atención de múltiples cabezas — decisiones paralelas desde múltiples perspectivas
Aunque con una sola atención es posible gestionar el flujo de información, en la práctica siempre se utilizan varias atenciones en paralelo. Esto se denomina atención de múltiples cabezas (multi-head attention).
Cada cabeza posee sus propias W_Q, W_K y W_V y calcula la atención de forma independiente. Los resultados se concatenan y luego se combinan nuevamente mediante una única matriz de pesos de salida W_O.
head_h = Attention(Q · W_Q^h, K · W_K^h, V · W_V^h)
MultiHead = Concat(head_1, ..., head_H) · W_OH es el número de cabezas. En GPT-2 medium, son 16; en GPT-3, 96; y en LLaMA-70B, 64.
¿Por qué necesitamos múltiples cabezas? Cada cabeza está especializada en un único "punto de vista". Por ejemplo, una cabeza puede estar centrada en "la palabra anterior", otra en "los sustantivos al principio de la frase" y otra en "hacer coincidir los pares de paréntesis". Tener múltiples cabezas permite procesar simultáneamente diversas relaciones gramaticales y semánticas.
Analogía biológica: evaluación paralela de múltiples marcadores. Cuando un patólogo evalúa una célula T, no se basa en una sola característica. Examina simultáneamente marcadores como CD8/CD4, marcadores de activación (CD69, HLA-DR), marcadores de agotamiento (PD-1, TIM-3, LAG-3) y marcadores reguladores (FoxP3), y luego integra los resultados para llegar a una conclusión. Cada prueba de marcador corresponde a una cabeza de atención, y la integración final corresponde a la combinación W_O. Esta analogía de puntos de vista paralelos es la clave de las múltiples cabezas.
Interpretabilidad. Los estudios que han analizado las cabezas de atención de los transformadores entrenados han demostrado que diferentes cabezas realmente se encargan de diferentes patrones lingüísticos. Por ejemplo, una de las cabezas de GPT-2 se centra estrictamente en "la palabra anterior", mientras que otra rastrea "la aparición anterior del mismo sustantivo". Las investigaciones recientes sobre la interpretabilidad mecánica están catalogando estas funciones específicas de cada cabeza.
Enmascaramiento causal: el problema de no poder ver el futuro
Los LLM se entrenan para predecir la siguiente palabra (episodio n.º 1). En este proceso, existe una restricción importante: no se debe hacer referencia a las palabras futuras que aún no se han predicho.
Por ejemplo, si ya hemos visto "el gato le gusta el pescado" y queremos predecir la siguiente palabra, sería inútil hacer referencia a la palabra "come" que aparecerá a continuación. El entrenamiento siempre debe basarse en "lo que se ha visto hasta ahora" para predecir lo siguiente.
La atención es, en esencia, una estructura que puede hacer referencia a todas las palabras de una frase. Por lo tanto, para evitar que "vea" el futuro, es necesario bloquearlo explícitamente. Esto se conoce como enmascaramiento causal.
Método: en la matriz de puntuación, la parte superior triangular (la posición en la que la palabra i es posterior a la palabra j > i) se establece en -∞. Después de pasar por la función Softmax, -∞ tendrá una probabilidad de 0, lo que hará que las palabras futuras se ignoren.
score(i, j) = ⟨Q_i, K_j⟩ / sqrt(d_k) if j ≤ i
score(i, j) = -∞ if j > iEl atención con enmascaramiento se denomina atención causal o atención del decodificador. Las series GPT la utilizan. La atención sin enmascaramiento es atención bidireccional. Las series BERT la utilizan. La mayoría de los LLM son causales.
Correspondencia biológica — Sutil pero análoga. Cuando una célula toma una decisión de diferenciación durante el desarrollo, esa célula no puede utilizar información de las células downstream que aún no se han desarrollado. Solo refleja las señales de las células upstream en orden temporal. Lo que hace el enmascaramiento causal es imponer exactamente esta causalidad temporal.
Self-Attention vs Cross-Attention
La atención que hemos discutido hasta ahora implica que Q, K y V provienen todos del mismo conjunto de tokens dentro de una oración. Esto se denomina atención propia (self-attention). Es una estructura en la que los tokens reciben información entre sí dentro de la propia oración.
La atención cruzada (cross-attention) ocurre cuando Q proviene de un conjunto, mientras que K y V provienen de otro conjunto. Por ejemplo, cuando el token de destino en un modelo de traducción consulta los tokens de origen. O cuando el token de texto consulta los parches de imagen en la generación de descripciones de imágenes. Es la base de los modelos multimodales.
- Atención propia: Q, K y V pertenecen a la misma secuencia. La mayoría de los LLM.
- Atención cruzada: Q pertenece a una secuencia, K y V a otra. Multimodalidad y traducción.
La arquitectura codificador-decodificador del artículo original sobre Transformers utiliza tanto atención propia como atención cruzada. Los LLM puros basados en decodificador (series GPT) utilizan solo atención propia. Recientemente, la atención cruzada ha vuelto a destacar en modelos de visión-lenguaje recientes (GPT-4V, Claude Vision, etc.).
KV Cache — Acelerar la inferencia 100 veces
Al pasar de la fase de entrenamiento a la de inferencia, aparece una optimización interesante.
Cuando un LLM genera texto, el proceso es el siguiente:
- Se introduce el prompt "El gato".
- El modelo predice el siguiente token: "el pez".
- El prompt se convierte en "El gato el pez" y se reintroduce en el modelo.
- La siguiente predicción es "gusta".
- El prompt se convierte en "El gato el pez gusta" y se reintroduce.
- ...
Si se recalcula todo el prompt en cada paso, la cantidad de cálculos en el paso t será O(t^2). Para generar 100 veces, se requieren 10000 unidades de cálculo.
Observación: Las K y V de los tokens calculados en pasos anteriores son idénticas en el siguiente paso. Gracias al enmascaramiento causal, las Q, K y V de los tokens pasados no cambian incluso cuando se añade un nuevo token. Por lo tanto, podemos almacenarlas en caché.
Esto es la caché KV (KV cache). Se almacenan los vectores K y V de cada capa en la memoria de la GPU, y en cada paso solo se calcula la Q para una nueva palabra, atendiéndola con las K y V almacenadas en caché. Esto reduce el costo computacional por paso a O(t). La generación de 100 pasos se acelera 100 veces, pasando de 10000 a 100 unidades.
El precio. Se consume mucha memoria de la GPU. Capas × 96 × cabezas × 96 × dimensiones × 12288 × contexto × 128K × tamaño del lote... al calcularlo, se requieren cientos de GB. Esta es la razón por la que la inferencia con contexto largo en modelos grandes está limitada por la memoria de la GPU.
Esta optimización de la caché KV es la razón decisiva por la que ChatGPT y Claude pueden mantener conversaciones en tiempo real hoy en día. Sin la caché, la respuesta tardaría decenas de segundos.
GQA y MQA: La evolución para reducir la caché KV
El tamaño de la caché KV era demasiado grande, lo que llevó a nuevas ideas. Varias cabezas Query comparten una sola cabeza Key/Value.
- Multi-Head Attention (MHA): Q, K y V tienen cada una H cabezas. La forma básica.
- Multi-Query Attention (MQA): Q tiene H cabezas, mientras que K y V tienen solo 1 cabeza cada una. Extremo. Ahorra H veces la caché KV.
- Grouped-Query Attention (GQA): Punto intermedio. Q tiene H cabezas, mientras que K y V tienen G cabezas cada una (donde G < H). Por ejemplo: si H=32 y G=8, 4 cabezas Q comparten 1 cabeza K/V.
MQA reduce la caché al máximo, pero conlleva una pérdida de calidad. GQA es un compromiso que mantiene casi toda la calidad mientras ahorra significativamente en caché. Modelos recientes como LLaMA-2, LLaMA-3, Qwen y Mistral utilizan GQA. Es el estándar actual.
Flash Attention: Rompiendo la barrera de memoria
El cuello de botella del cálculo de atención no es realmente la cantidad de operaciones, sino el ancho de banda de la memoria de la GPU. Si Q · K^T es una matriz de tamaño n × n (n = longitud de la secuencia), se necesita una matriz 128000 × 128000 para un contexto de 128K. El tiempo que tarda en ir y volver esta gran matriz desde HBM (memoria de la GPU) a SRAM (caché dentro del chip) es mucho más largo que el cálculo real.
Flash Attention (Dao et al., 2022) es una implementación que minimiza estos viajes ida y vuelta. Divide la gran matriz de atención en bloques, procesándolos dentro de la SRAM sin materializar la gran matriz en HBM. El resultado matemático es idéntico, pero la velocidad es de 2 a 10 veces mayor y el uso de memoria se reduce de O(n²) a O(n).
Esta es la razón decisiva por la que contextos de 128K, 200K y 1M tokens se han vuelto prácticos. Hoy en día, casi todos los entrenamientos e inferencias de modelos grandes utilizan Flash Attention. Sigue mejorándose con Flash Attention 2 y 3.
Punto clave. No es la teoría de algoritmos, sino la optimización de implementación adaptada al hardware lo que respalda la expansión actual de los LLM. La importancia de la ingeniería de sistemas de aprendizaje profundo.
Después del mecanismo de atención — Hasta la predicción de la siguiente palabra
Resumen de los componentes anteriores. El bloque de atención contextualiza los vectores de palabras; luego, el FFN (capa de alimentación hacia adelante) descrito en el capítulo #5 procesa estas representaciones contextualizadas; tras pasar por múltiples bloques, finalmente, la cabeza de modelado del lenguaje (una matriz de tamaño d × V similar a la incrustación) genera la distribución de probabilidad de la siguiente palabra.
Durante el entrenamiento, la entropía cruzada entre esta distribución y la palabra real siguiente constituye la función de pérdida. La retropropagación descrita en los capítulos #3 y #4 propaga este error hasta los parámetros de atención para el aprendizaje. Como resultado, la atención se asienta automáticamente como "la forma más útil de contextualizar el lenguaje".
Interpretación de las matrices de atención. Al visualizar las puntuaciones de atención α(i, j) de un modelo entrenado, se revelan ciertos patrones lingüísticos. Por ejemplo:
- Cabezas de token anterior inmediato: Puntuaciones concentradas justo debajo de la diagonal. Adyacencia gramatical.
- Cabezas de referencia a sustantivos anteriores: Puntuaciones que van de los pronombres a los sustantivos previos. Resolución de coreferencia.
- Cabezas de final de oración: Puntuaciones que van desde el token final hasta el inicio de la misma oración. Resumen de la estructura de la oración.
Esta visualización es el punto de partida para la investigación en interpretabilidad mecánica. Anthropic, OpenAI y Google están actualmente realizando ingeniería inversa de los circuitos internos de los LLM mediante el análisis de patrones de atención.
Escenarios de aplicación biológica
Escenario 1 — Predicción de contactos de aminoácidos con ESM
La atención del modelo de lenguaje proteico ESM, mencionado en los capítulos #4 y #5, muestra propiedades notables. Al investigar la matriz de atención de ESM una vez entrenado, se observa que la atención se concentra naturalmente y con fuerza en pares de aminoácidos que están físicamente en contacto en la estructura tridimensional.
Es decir, aunque no se le entrenó explícitamente para "predecir la estructura", sino solo para predecir el siguiente aminoácido, la atención aprende automáticamente los contactos tridimensionales. Este mapa de atención se utilizó posteriormente como herramienta de predicción de contactos en las etapas iniciales de AlphaFold.
Este es un ejemplo que demuestra el poder general de la atención: el principio de prestar atención a la información relevante se aplica tanto al lenguaje como a las secuencias biológicas.
Escenario 2 — Conexión entre potenciadores y promotores con Enformer
La atención del modelo Enformer (transformador genómico), mencionado en el capítulo #5, conecta promotores (sitios de inicio génico) con potenciadores (elementos reguladores) situados a decenas de kilobases de distancia, dentro de un contexto de 100 kb. Una vez finalizado el entrenamiento, algunos cabezales de atención aprenden automáticamente los pares potenciador-promotor.
Tradicionalmente, la coincidencia de potenciadores y promotores requería experimentos de contacto genómico 3D (Hi-C), pero Enformer predice esta relación únicamente a partir de la secuencia. Esto se debe a que el mecanismo de atención tiene la "capacidad de aprender interacciones a larga distancia".
Escenario 3 — Determinación del contexto celular en imágenes patológicas
Este es un caso real de implementación del escenario descrito en los episodios #2 y #6. Un modelo de patología basado en Vision Transformer (ViT) calcula la atención para cada parche (una pequeña región cuadrada) de una diapositiva, con el fin de determinar el microentorno tumoral. Cada parche de célula T establece una relación de atención con los parches de células tumorales y estromales circundantes; al sintetizar esta información, se predice el estado de la célula T.
Modelos de esta familia (como HIPT, CTransPath, etc.) están comenzando a mostrar un rendimiento comparable al de patólogos en el diagnóstico preciso de tumores y la predicción del pronóstico.
Resumen clave
- La atención es el proceso mediante el cual una expresión léxica se contextualiza adecuadamente según su contexto. El tratamiento de la palabra polisémica "hígado" en el episodio #1 ocurre aquí.
- Cada token posee tres caras: Q, K y V. Se calcula el "grado de atención" entre tokens mediante el producto interno de Q y K, se convierte en probabilidad mediante softmax y se recibe información mediante un promedio ponderado de V.
- Una sola línea de fórmula:
Attention(Q, K, V) = softmax(Q · K^T / sqrt(d)) · V. - El multicabeza permite el procesamiento paralelo desde múltiples perspectivas, análogo a la evaluación de múltiples marcadores y la coherencia conceptual por parte de un patólogo.
- El enmascaramiento causal bloquea la información futura. Es esencial para el entrenamiento de LLM.
- La caché KV acelera la inferencia 100 veces. GQA reduce el tamaño de la caché.
- Flash Attention hace práctico el contexto de 128K+. Las implementaciones conscientes del hardware respaldan la expansión actual de los LLM.
- La atención es un principio general que trasciende no solo el lenguaje, sino también las proteínas, el genoma y las imágenes. La idea de "prestar atención a la información relevante" se reutiliza a través de dominios.
Próximos conceptos
- Episodio #7
prompt-engineering— Diseño de prompts que guían eficazmente la atención de un transformador entrenado. - Episodio #8
rag-and-context— RAG para ampliar la información que la atención puede manejar. - Episodio #11
hallucination-and-alignment— Cómo las fallas en la atención se conectan con las alucinaciones.
📐 Apéndice — Fórmulas matemáticas para expertos
Dificultad: Muy difícil (Very Hard) Dirigido a: Lectores con conocimientos de álgebra lineal, probabilidad y análisis numérico a nivel de posgrado
A.1 Fórmula completa de la Atención de Producto Punto Escalado
Longitud de la secuencia n, dimensión del embedding d, dimensión Q/K d_k, dimensión V d_v.
Entrada:
Q ∈ ℝ^{n × d_k}(consulta)K ∈ ℝ^{n × d_k}(clave)V ∈ ℝ^{n × d_v}(valor)
Salida:
Attention(Q, K, V) = softmax(Q K^T / sqrt(d_k)) VAplicación de Softmax en la dirección de las filas:
softmax(M)_{ij} = exp(M_{ij}) / Σ_k exp(M_{ik})Cada fila es una distribución de probabilidad (suma igual a 1).
Formato de salida final: ℝ^{n × d_v}. Cada fila representa la expresión contextualizada de la palabra correspondiente.
A.2 Derivación de la razón para dividir por sqrt(d_k)
Asumiendo que cada elemento de Q y K es un valor aleatorio con media 0 y varianza 1.
Producto interno de Q_i y K_j:
⟨Q_i, K_j⟩ = Σ_{k=1}^{d_k} Q_{ik} K_{jk}Si cada término Q_{ik} K_{jk} es independiente, la media es 0 y la varianza es 1. La varianza de la suma es:
Var(⟨Q_i, K_j⟩) = d_k
Std(⟨Q_i, K_j⟩) = sqrt(d_k)Es decir, si d_k es grande, la magnitud del producto interno se vuelve sqrt(d_k). Si d_k = 128, la desviación estándar es aproximadamente 11.3.
Si esto se introduce directamente en softmax, los valores grandes se agrupan extremadamente, haciendo que la salida de softmax sea cercana a un one-hot. El gradiente de softmax se vuelve prácticamente cero, lo que impide el entrenamiento.
Al dividir por sqrt(d_k), la desviación estándar se mantiene cerca de 1, permitiendo que softmax produzca una distribución suave. El gradiente fluye adecuadamente, haciendo posible el entrenamiento.
A.3 Derivadas parciales de softmax
Cuando p = softmax(z):
∂p_i/∂z_j = p_i · (δ_{ij} - p_j)δ_{ij} es la delta de Kronecker (1 si i=j, 0 en caso contrario).
Combinación con la pérdida de entropía cruzada:
Etiqueta verdadera y (one-hot), pérdida L = -Σ y_i log p_i.
Regla de la cadena:
∂L/∂z_j = p_j - y_jEl elegante resultado de la Sección A.4, Parte 3. El gradiente del entrenamiento de atención también aprovecha repetidamente este principio.
Ecuaciones de Multi-Head Attention (A.4)
Hcabezas. Cada cabeza h:
Q_h = X · W_Q^h, K_h = X · W_K^h, V_h = X · W_V^h
head_h = Attention(Q_h, K_h, V_h)W_Q^h, W_K^h ∈ ℝ^{d × d_k}, W_V^h ∈ ℝ^{d × d_v}. Por lo general, d_k = d_v = d/H.
Unión:
MultiHead(X) = Concat(head_1, ..., head_H) · W_OW_O ∈ ℝ^{H · d_v × d}.
Número de parámetros:
- Cada proyección Q, K y V tiene
d × d. Al haber tres, el total es3 d^2. - La proyección de salida tiene
d^2. - Total por bloque:
4 d^2.
Conforme al cálculo en la sección A.10 del apéndice 5.
A.5 Implementación de la máscara causal
Matriz de máscara:
M_{ij} = 0 if j ≤ i
M_{ij} = -∞ if j > iAl calcular la atención:
Attention_masked(Q, K, V) = softmax((Q K^T / sqrt(d_k)) + M) V-∞ se convierte en exp(-∞) = 0 en softmax, asignando una probabilidad de 0 a la posición futura.
En la práctica, utilice un número negativo grande (por ejemplo, -1e9) en lugar de -∞ para garantizar la estabilidad numérica.
A.6 Algoritmo de KV Cache
Durante la inferencia:
- Paso 1: Calcular
Q_1, K_1, V_1. Guardar[K_1], [V_1]en el caché. - Paso 2: Calcular
Q_2, K_2, V_2para la nueva palabra. Añadir al caché:[K_1, K_2], [V_1, V_2]. - Paso
t: Calcular solo la nueva palabraQ_t. Atención con K y V en caché.
Memoria:
- Número de cabezas por capa
H, dimensión de la cabezad_h, contexton, loteB - KV cache para una capa:
2 · B · n · H · d_h · bytes_per_value
Ejemplo: LLaMA-70B, H=64, d_h=128, n=128K, B=1, fp16:
2 · 1 · 128000 · 64 · 128 · 2 = 4.19 GB per layerSi hay 80 capas, 335 GB. La razón por la que el contexto largo de los modelos grandes es un límite de memoria.
A.7 GQA (Atención con Consulta Agrupada)
- Número de cabezales Q:
H_Q - Número de cabezales K, V:
H_KV - Tamaño del grupo
g = H_Q / H_KV
Los cabezales Q del mismo grupo comparten un solo cabezal K y V:
head_i = Attention(Q_i, K_{i // g}, V_{i // g})Reducción del tamaño del caché KV en g. Si g = H_Q (todos los Q comparten un único KV), es MQA.
LLaMA-2-70B: H_Q = 64, H_KV = 8, g = 8. Ahorro de caché de 8 veces.
A.8 Algoritmo de bloque Flash Attention
Idea clave: No materializar una gran matriz de atención n × n en SRAM, sino calcular mediante tiling por bloques.
Online Softmax: Algoritmo para calcular Softmax de forma fluida (stream) por bloques. Mantiene el máximo parcial y la suma exponencial parcial de cada bloque para su posterior combinación.
Tamaño de bloque B_r(filas) × B_c(columnas):
- Carga de Q en fragmentos de
B_r - Recorrido de K, V en fragmentos de
B_c - Cálculo de puntuación parcial y suma exponencial parcial para cada (Q_block, K_block)
- Combinación con los resultados del bloque anterior (fórmula de online softmax)
Resultado:
- Tiempo: FLOPs teóricos idénticos, pero el tiempo real (wallclock) es 2~4 veces más rápido (elimina el cuello de botella de memoria)
- Memoria: O(n²) → O(n)
- Precisión: Exactamente la misma (error numérico mínimo)
A.9 Familia de Sparse Attention
Técnicas de aproximación para manejar contextos largos.
- Sliding Window Attention: Cada token solo referencia a
wtokens cercanos. Cálculo deO(nw). Adoptado por Mistral·Longformer. - Longformer: Ventana deslizante + algunos tokens globales.
- BigBird: Ventana deslizante + global + atención aleatoria.
- Sparse Attention (GPT-3): Atención con patrones específicos (strided, factorized).
Compromiso (Trade-off): Pérdida de calidad debido a la aproximación. Tras el éxito de Flash Attention, se ha vuelto relativamente menos necesario. Sigue siendo útil en contextos extremadamente largos (1M+).
A.10 Complejidad de cómputo y memoria de la atención
Longitud de secuencia n, dimensión d:
- Full Attention: Tiempo
O(n²d), memoriaO(n² + nd) - Flash Attention: Tiempo
O(n²d), memoriaO(n · d)(solo el tamaño del bloque en SRAM) - Sliding Window: Tiempo
O(nwd), memoriaO(nd)
Cuando n = 128K, n² es de 1,6 mil millones. El hecho de que el cálculo de matrices de esta escala sea práctico se debe a Flash Attention. La capacidad de contexto largo de los LLM actuales es producto de la colaboración entre este algoritmo y el sistema.
Referencias
Todo el contenido, escenarios y analogías de esta sección han sido desarrollados por BioPlayground. A continuación, se presentan referencias externas que pueden ayudar al aprendizaje de los conceptos.
- Artículo original del transformador (definición de atención): Vaswani et al., "Attention Is All You Need" (NeurIPS 2017)
- Flash Attention: Dao et al., "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness" (NeurIPS 2022)
- Flash Attention 2: Dao, "FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning" (2023)
- Multi-Query Attention: Shazeer, "Fast Transformer Decoding: One Write-Head is All You Need" (2019)
- Grouped-Query Attention: Ainslie et al., "GQA: Training Generalized Multi-Query Transformer Models" (EMNLP 2023)
- Interpretación de cabezales de atención: Elhage et al., "A Mathematical Framework for Transformer Circuits" (Anthropic 2021)
- Catálogo de cabezales de atención: Olsson et al., "In-context Learning and Induction Heads" (Anthropic 2022)
- Predicción de contactos ESM: Rao et al., "Transformer protein language models are unsupervised structure learners" (ICLR 2021)
- Vision Transformer: Dosovitskiy et al., "An Image is Worth 16x16 Words" (ICLR 2021)
- Educación en visualización de aprendizaje profundo: 3Blue1Brown "Deep Learning" Capítulos 6·7 (YouTube) — Referencia pedagógica
El número #6 marca el final de la sección de principios (Fase 1). A partir del número #7, se abordará cómo utilizar realmente este transformador entrenado: ingeniería de prompts, RAG y agentes.