¿Un contexto de día completo supera a uno de diez minutos? Atención flash y la pregunta sobre la longitud de la secuencia
Esta es la pregunta que este artículo pretende responder: si un transformador pudiera atender un día de negociación completo con una resolución de un segundo en lugar de una ventana de diez minutos, ¿predeciría mejor?
Hasta hace poco ni siquiera se podía preguntar. Necesidades de atención estándar memoria, por lo que un día de 23.400 pasos con 12 cabezas en float16 necesita aproximadamente 12,9 GB solo para la matriz de puntuación, más que los parámetros del modelo y más de lo que le darán la mayoría de las tarjetas. La pregunta se cerró mediante aritmética antes de que alguien pudiera probarla.
Flash Attention (Dao et al., 2022) lo abre. No aproximando la atención (calcula exactamente el mismo resultado), sino reestructurando el cálculo para que sea consciente de IO, minimizando el tráfico entre los niveles de memoria de la GPU. Ese es el contenido realmente interesante aquí, y la mayor parte de este artículo se dedica a cómo funciona: mosaico, recurrencia de softmax en línea, Límite de IO y recálculo de paso hacia atrás.
Pero el mecanismo es el facilitador, no el reclamo. "Un contexto más largo es mejor" es una afirmación empírica sobre los mercados y la posición de este blog, de Temporal Fusion Transformers, que encontró que los transformadores básicos en modelos recurrentes sobreajustados de series financieras y retrospectivos siguen siendo competitivos en alta frecuencia, va en sentido contrario. Entonces el artículo se cierra con la medición, no con el mecanismo.
Por qué la atención está ligada a la memoria
La atención calcula — la primitiva en sí, en un contexto comercial, se trata en Transformadores de fusión temporal para pronósticos de carteras multihorizonte. Todo el problema es una línea de eso: la matriz de puntuación intermedia es , se escribe en la memoria, se vuelve a leer para softmax, se vuelve a escribir y se lee de nuevo para el matmul final, y debe conservarse para la retropropagación.
La intensidad aritmética de la atención es , entonces aproximadamente 64 FLOP/byte en — bastante a la izquierda del punto de la cresta A100. Se asienta en el techo inclinado del ancho de banda, no en el techo de cómputo plano: la GPU pasa más tiempo moviéndose alrededor que multiplicar cualquier cosa. El marco de la línea del techo que utiliza (punto de cumbrera, techo inclinado versus plano, y por qué el mismo razonamiento decide si vale la pena comprar una GPU) se construye con números medidos en Cuando la GPU da sus frutos.
La jerarquía de memoria que explota el algoritmo
| Nivel de memoria | Tamaño | Ancho de banda | Latencia |
|---|---|---|---|
| HBM (Memoria de gran ancho de banda) | 40-80GB | 2,0 TB/s | ~400 ns |
| SRAM (memoria compartida en chip) | 20MB | 19 TB/s | ~4 segundos |
La SRAM tiene aproximadamente 10 veces el ancho de banda y 100 veces menos latencia, a una milésima parte de la capacidad. Todo lo que hace Flash Attention se deriva de ese intercambio: renunciar a capacidad, comprar ancho de banda y latencia. La misma medida de "reestructurar el algoritmo en lugar de comprar hardware", medida en una prueba retrospectiva de la CPU, es la escalera de velocidad de la prueba retrospectiva.
El algoritmo de atención flash
Flash Attention procesa la atención en mosaicos de tamaño adecuado para SRAM y nunca materializa la atención completa. matriz en HBM.
Dividir en bloques de filas y en bloques de columnas, con elegido para que una ficha y sus acumuladores encajen en el chip. Para cada bloque de consulta, repita todos los bloques clave-valor:
For each query block Q_i:
Initialize: O_i = 0, m_i = -inf, l_i = 0 # output, running max, running sum
For each KV block (K_j, V_j):
1. Load Q_i, K_j, V_j from HBM to SRAM
2. Compute S_ij = Q_i @ K_j^T / sqrt(d) # in SRAM
3. Compute local max: m_ij = rowmax(S_ij)
4. Compute P_ij = exp(S_ij - m_ij) # in SRAM
5. Compute local sum: l_ij = rowsum(P_ij)
6. Update running statistics:
m_new = max(m_i, m_ij)
l_new = l_i * exp(m_i - m_new) + l_ij * exp(m_ij - m_new)
O_i = O_i * (l_i * exp(m_i - m_new) / l_new)
+ P_ij @ V_j * (exp(m_ij - m_new) / l_new)
m_i = m_new, l_i = l_new
Write O_i to HBM
La recurrencia de softmax en línea
El truco que hace posible el mosaico es softmax en línea. Un softmax ingenuo necesita dos pasadas sobre la fila: una para encontrar el máximo (para estabilidad numérica), otra para exponenciar y normalizar. Dos pasadas sobre una fila que se niega a almacenar es una contradicción, por lo que Flash Attention sigue ejecutando estadísticas y reescalando a medida que avanza.
Después de los bloques :
y el acumulador de salida se corrige por el mismo factor:
Cada vez que un nuevo bloque aumenta el máximo de ejecución, la salida acumulada previamente se reescala retroactivamente por – como si el nuevo máximo se hubiera conocido desde el principio. El resultado es algebraicamente idéntico al softmax de dos pasadas. En aritmética exacta esto no es una aproximación; es una reasociación. (En precisión finita, es una ruta de redondeo diferente, lo que importa; consulte la verificación de exactitud a continuación).
Complejidad IO
Esta es la declaración formal de la victoria. Atención Flash realiza
Accesos de HBM, donde es el tamaño de SRAM, contra para la implementación estándar. Tenga en cuenta que aparece en el denominador: cuanto más grande es el scratchpad en el chip, menos viajes de ida y vuelta, razón por la cual el algoritmo se expresa en términos de jerarquía de memoria en lugar de recuento de FLOP. Para tipico y KB, la proporción favorece la Atención Flash con aproximadamente 5 a 10 veces menos accesos.
Pase hacia atrás: volver a calcular en lugar de almacenar
La retropropagación a través de la atención normalmente necesita la matriz que el pase adelantado simplemente se negó a mantener. Flash Atención recalcula los mosaicos de durante el pase hacia atrás, almacenando sólo la salida y las estadísticas de softmax - ambos , no . Cambia una modesta cantidad de aritmética redundante por el término de memoria que fue todo el problema. Esta es la misma ganga que el punto de control de gradiente, aplicado a la granularidad del mosaico dentro de un solo operador.
FA2: paralelismo
Flash Attention 2 (Dao, 2023) mantuvo el algoritmo y arregló la programación:
- Menos FLOP que no son matmul. FA1 dedicó tiempo real a reescalado, búsqueda máxima y exponenciación: operaciones que se ejecutan en núcleos CUDA, no en núcleos tensoriales. FA2 difiere el cambio de escala hasta el final del bucle interno.
- Paralelismo sobre la longitud de la secuencia. FA1 realiza paralelismo sobre lotes y cabezales únicamente. FA2 también realiza paralelismo sobre bloques de consulta. Esto es importante específicamente para el caso comercial, donde a menudo tienes una secuencia muy larga por activo y un tamaño de lote de 1 a 4, exactamente el régimen en el que el paralelismo entre lotes y cabezales priva a la GPU.
- Partición del trabajo de deformación. Cada deformación toma un subconjunto diferente de bloques de consulta en lugar de dividir un cálculo de puntuación y reducir entre deformaciones, eliminando una reducción entre deformaciones.
Resultado informado: ~70 % de los FLOP máximos teóricos en A100 frente a ~35 % para FA1.
FA3: Mecánica de la tolva
Flash Attention 3 (Dao, Shah, 2024) es una arquitectura específica de H100:
- Especialización en deformación asíncrona. El acelerador de memoria tensor (TMA) de Hopper mueve HBM→SRAM de forma asíncrona. FA3 divide los warps en productores que emiten cargas TMA para el siguiente bloque KV y consumidores que calculan el actual, por lo que el movimiento de datos se esconde detrás de la aritmética.
- Matmul y softmax intercalados. para un bloque se ejecuta en núcleos tensoriales, mientras que el softmax para el bloque anterior se ejecuta en núcleos CUDA: dos unidades de hardware diferentes, genuinamente concurrentes en lugar de divididas en tiempo.
- FP8 con procesamiento incoherente. H100 realiza FP8 con un rendimiento 2x FP16. La atención ingenua del 8PM se ve arruinada por los valores atípicos; FA3 rota aleatoriamente los vectores antes de la cuantificación en bloques para distribuir la magnitud atípica entre las coordenadas, lo que se reporta con un error numérico 2,6 veces menor que el ingenuo FP8.
| Versión | GPU | Utilización | Aceleración frente a estándar |
|---|---|---|---|
| FA1 | A100 | ~35% | 2-4x |
| FA2 | A100 | ~70% | 5-7x |
| FA3 (FP16) | H100 | ~75% | 3-5x frente a FA2 |
| FA3 (FP8) | H100 | ~75% | 1,6x frente a FA3 FP16 |
El enmascaramiento causal es donde el comercio obtiene un descuento
El enmascaramiento causal es obligatorio para las series de tiempo (el modelo no debe atender al futuro) y el mosaico no es un costo adicional sino un costo ahorrado. Cualquier mosaico cuyas claves estén completamente en el futuro en relación con sus consultas se omiten por completo, nunca se cargan ni se calculan, lo que reduce aproximadamente la mitad del trabajo. En PyTorch esto es is_causal=True; no se requiere nada más.
La integración es de ocho líneas.
Casi nada del código que necesita trata sobre Flash Attention. Cambie la ruta explícita de la matriz de puntuación para el núcleo fusionado:
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_head)
scores = scores.masked_fill(causal_mask, float("-inf"))
out = torch.matmul(torch.softmax(scores, dim=-1), v)
out = torch.nn.functional.scaled_dot_product_attention(
q, k, v,
dropout_p=self.dropout if self.training else 0.0,
is_causal=True,
)
Ese es todo el cambio. q, k, v tienen forma (batch, heads, seq, head_dim); la máscara causal desaparece porque el núcleo la construye. Para obtener un modelo comercial de PyTorch anotado completo en el que colocar esto (proyección de entrada, bloques y un cabezal arriba/plano/abajo de 3 clases), use DeepLOB, y para ver un programa de capacitación completo, consulte Transformadores de fusión temporal. Construir una cuarta copia de ese andamiaje aquí no enseñaría nada.
Requisitos: capacidad de cálculo >= 8.0 (A100, H100, RTX 3090+), entradas de media precisión, PyTorch >= 2.0. Verifique la ruta rápida realmente comprometida con torch.backends.cuda.sdp_kernel diagnóstico y torch.cuda.max_memory_allocated() — SDPA recurre silenciosamente al núcleo matemático si falla alguna condición previa, y un recurso silencioso se parece exactamente a un modelo funcional que es simplemente lento.
La medición: ¿merece la pena el contexto más largo?
Todo lo anterior dice que un contexto de 32K o 128K ahora es asequible. No dice nada sobre si es útil. El experimento honesto:
Entrena la misma arquitectura en con SDPA-Flash en la serie BTC utilizado en otras partes de esta serie, manteniendo los parámetros, el optimizador y el objetivo fijos, por lo que la longitud de la secuencia es la única variable. Informa dos cosas:
- Costo. Reloj de pared medido por época y
torch.cuda.max_memory_allocated()en cada . - Beneficio. Rendimiento predictivo fuera de muestra versus , en una división hacia adelante.
Un borrador anterior de este artículo incluía una tabla de cifras de memoria por longitud de secuencia derivada analíticamente del Fórmula de activación. Esas filas se eliminan: nunca se midieron y no estaban de acuerdo con la aritmética del presupuesto de memoria del propio artículo. Un número derivado presentado en una tabla de resultados es un resultado fabricado y este blog no los envía.
La propiedad interesante de este experimento es que se puede publicar en cualquier dirección. Si el rendimiento fuera de la muestra aumenta monótonamente con , eso justifica todo el programa de largo contexto. Si se estabiliza en unos pocos miles de pasos, o se degrada, se trata de una pieza más fuerte, un compañero de lo negativo honesto - y significaría que el muro de la memoria nunca fue la limitación vinculante para el comercio de transformadores.
Más contexto es más capacidad, por lo tanto más superficie de sobreajuste
Hay una razón específica para esperar un resultado plano o negativo. Transformadores de fusión temporal ya documenta que los transformadores básicos se aplicaron ingenuamente al sobreajuste de series financieras: carecen de sesgos inductivos temporales y los modelos recurrentes retrospectivos siguen siendo competitivos a alta frecuencia. Ampliar el contexto de 512 a 32.768 pasos no agrega información proporcional a la longitud; el rezago marginal número 32.000 de una serie de precios casi eficiente conlleva muy poco. Lo que agrega de manera confiable es el valor de los parámetros de las cosas que se ajustan.
Así que el barrido debe ser tratado como lo que es: una búsqueda de selección de modelo, con la misma maquinaria que este blog aplica a cualquier otra búsqueda. Tres longitudes de secuencia multiplicadas por cualquier otra variación es un conteo de prueba, y el ganador debe superar una Proporción de Sharpe deflactada calculada contra ese conteo de prueba y una puerta PBO, no simplemente vencer a sus vecinos. De lo contrario, "el contexto largo gana" es indistinguible de elegir lo mejor de tres carreras ruidosas.
Una verificación de exactitud, porque "exacto" implica mucho trabajo
La atención flash es exacta en aritmética exacta. La recomendación adjunta: ejecutar en fp16 o bf16, y en H100 considerar FP8, no lo es. Esas son afirmaciones separadas y la segunda domina en la práctica: reasociar una suma y reducir la precisión a la mitad son ambas perturbaciones, y el artículo que introdujo la garantía de pedido no debería entonces agitar la garantía de precisión.
El blog ya tiene el instrumento adecuado. La trampa de precisión de la GPU establece el estándar: la baja precisión no le advierte, devuelve basura plausible y usted demuestra la exactitud con un oráculo de paridad en una cantidad discreta posterior (recuentos comerciales) no observando curvas. Aplicado aquí:
- Computar atención con SDPA-Flash en bf16 y con una implementación de referencia fp64 en entradas idénticas; informar error relativo máximo en el tensor de salida.
- Empújelo hasta la decisión: para un modelo que emite una etiqueta arriba/plana/abajo, informe cuántas etiquetas se invierten entre las dos rutas, como una fracción del total de decisiones.
Los desacuerdos pequeños, limitados y explicables son la firma de un camino rápido y correcto. Un valor ilimitado significa que la recomendación del 8PM nunca fue segura para este modelo. Ninguno de los números se conoce hasta que se ejecuta.
Cuándo alcanzarlo
Comprimido a la decisión, que tiene la misma forma que la guía de decisión de GPU:
- Por encima de ~2K pasos de tiempo en una GPU CUDA: sí, incondicionalmente. Es un cambio de una línea que produce un resultado exacto y la ganancia crece con . No hay ningún escenario donde quieras que se materialice. camino en su lugar.
- Por debajo de ~512 pasos de tiempo, en CPU o con arquitecturas sin atención (CNN, SSM como Mamba): irrelevante. A la izquierda de la cresta, los gastos generales fijos son el costo total y la atención nunca fue su cuello de botella.
- Los umbrales anteriores son folklore, no medidas: provienen de la literatura general y el cruce en su propio modelo y tarjeta es un punto de referencia de diez líneas. Ejecútelo en lugar de confiar en los números redondos.
Conclusión
Flash Attention es un resultado limpio y genuinamente importante: al respetar la jerarquía de la memoria y reasociar el softmax, calcula la atención exacta con memoria en lugar de , y el IObound explica precisamente por qué. Adoptarlo en un transformador comercial es un cambio de una línea sin costo de precisión en aritmética exacta y una gran ganancia de memoria.
Lo que no hace es responder la pregunta de arriba. Convierte "un contexto de día completo es imposible" en "un contexto de día completo es barato", lo que supone un cambio en el coste del experimento, no en su resultado. El muro de la memoria que se derrumba es una invitación a medir, y la medición es lo que convierte esto de un resumen en papel en un hallazgo.
Referencias
- Dao, T., Fu, D.Y., Ermon, S., Rudra, A., Re, C. "FlashAttention: atención exacta rápida y eficiente en memoria con IO-Awareness". NeurIPS (2022). arXiv:2205.14135
- Dao, T. "FlashAttention-2: Atención más rápida con mejor paralelismo y partición del trabajo". ICLR (2024). arXiv:2307.08691
- Dao, T., Shah, J. "FlashAttention-3: Atención rápida y precisa con asincronía y baja precisión". NeurIPS (2024). arXiv:2407.08608
- Vaswani, A., et al. "Atención es todo lo que necesitas". NeurIPS (2017).
- Milakov, M., Gimelshein, N. "Cálculo del normalizador en línea para softmax". arXiv:1805.02867 (2018).
Authors
Trading-systems engineer
Trading-systems engineer building bots since 2017: cross-exchange arbitrage (connected up to 30 venues), cointegration-based pairs arbitrage across spot and futures, scalping, news and sentiment-driven strategies, trend algorithms, and portfolio management and balancing algorithms. Also builds sub-millisecond order execution, big-data warehouses, backtesting engines, AI agents, and trading interfaces (incl. open-source profitmaker.cc). Stack: JS/TS, Python, Rust/Zig/Go, DevOps, backend, frontend, architecture.