Avançado

FlashAttention e IO-awareness

FlashAttention calcula attention exata em tiles dimensionados para a memória on-chip, e é isso que torna fisicamente possível o contexto de 262.144 tokens do Qwen3.8-27B.

Atualizada em

01 · Conceito

Conceito

O Qwen3.8-27B aceita um contexto nativo de 262.144 tokens, e 16 de suas 64 camadas executam full softmax attention sobre todo esse intervalo. Alimente o modelo com um prompt máximo e cada uma dessas camadas precisa, em princípio, comparar cada query com cada key. A fórmula de attention é uma linha; a questão é se os objetos intermediários que ela implica podem existir fisicamente.

Faça a aritmética. Com N=262,144=218N = 262{,}144 = 2^{18}, a matriz de scores QKQK^\top tem 218×218=2366.9×10102^{18} \times 2^{18} = 2^{36} \approx 6.9 \times 10^{10} entradas por head. A dois bytes cada em bf16, isso dá 2372^{37} bytes, ou 128 GiB, para um head de uma camada. O Qwen3.8-27B tem 24 query heads por camada de full attention, então materializar os scores de uma camada tomaria cerca de 3 TiB, e suas 16 camadas de full attention juntas, aproximadamente 48 TiB — escritos na high-bandwidth memory (HBM), lidos de volta para o softmax, escritos novamente como probabilidades e lidos mais uma vez para multiplicar VV. Um acelerador da classe de 80 GB não consegue guardar nem a matriz de um único head. Attention ingênua nesse comprimento de contexto não é lenta; é impossível sem mudar o plano de memória.

Aceleradores têm uma hierarquia de memória: um pequeno pool de SRAM on-chip, medido em dezenas de megabytes, que é muito rápido, e a HBM, que é ordens de magnitude maior mas bem mais lenta de alcançar. As unidades aritméticas ficam ociosas enquanto tensores viajam entre as duas. IO-awareness significa projetar o algoritmo em torno dos bytes movidos por essa hierarquia, e não apenas contar operações de ponto flutuante. A observação de FlashAttention é que attention pode ser calculada em tiles que cabem na SRAM, de modo que a matriz de scores nunca precisa existir por inteiro.

Divida as queries em blocos de linhas e as keys e values em blocos de colunas, e dimensione os tiles pela head dimension. Os heads de full attention do Qwen usam head dimension 256, então um tile de key de 128 linhas ocupa 128×256×2=65,536128 \times 256 \times 2 = 65{,}536 bytes, 64 KiB, e seu gêmeo de value outros 64 KiB. Isso é o dobro da pegada dos heads de 128 dimensões para os quais muitos kernels antigos foram ajustados, o que força blocos de linhas menores ou mais passadas dentro do mesmo orçamento de SRAM — e significa que o suporte do kernel a d=256d = 256 precisa ser verificado para a versão exata da biblioteca, não presumido a partir de um flag de API.

O obstáculo ao tiling é o softmax: normalizar uma linha de scores parece exigir ver a linha inteira de uma vez. O equívoco clássico é aplicar softmax localmente a cada tile e concatenar os resultados. Experimente numa linha de query cujos scores chegam como um bloco [1.0, 3.0][1.0,\ 3.0] seguido de um bloco [5.0][5.0]. Um softmax local do primeiro bloco dá ao score 3,0 um peso de e3/(e1+e3)0.881e^{3}/(e^{1}+e^{3}) \approx 0.881 — mas o peso verdadeiro, quando o 5,0 dominante chega, é só cerca de 0,117. A normalização local congela um denominador que os tiles seguintes invalidam.

O softmax online corrige isso carregando duas estatísticas correntes por linha: o máximo mm e a soma exponencial \ell. Após o primeiro bloco, m1=3m_1 = 3 e 1=e13+e33=0.135+1=1.135\ell_1 = e^{1-3} + e^{3-3} = 0.135 + 1 = 1.135. O segundo bloco eleva o máximo para mnew=5m_{new} = 5, então a soma antiga é reescalada por em1mnew=e2e^{m_1 - m_{new}} = e^{-2}: 1.135×0.135=0.154\ell \leftarrow 1.135 \times 0.135 = 0.154, e o novo termo soma e55=1e^{5-5} = 1, resultando em =1.154\ell = 1.154. Calculando diretamente, e4+e2+e0=0.018+0.135+1=1.154e^{-4} + e^{-2} + e^{0} = 0.018 + 0.135 + 1 = 1.154 — idêntico. Os acumuladores de saída parcial são reescalados pelo mesmo fator, então após o último tile a divisão por \ell produz attention exata. A ordem das operações muda, então os últimos bits de ponto flutuante podem diferir, mas nada é aproximado, esparsificado ou truncado.

A recomputação estende a mesma lógica ao treinamento: o backward pass pode recalcular scores locais dentro dos tiles em vez de ler uma matriz N2N^2 armazenada, trocando aritmética barata por tráfego caro de HBM.

O benefício é específico por fase. O prefill processa muitas linhas de query e é exatamente o regime que o tiling mira. O decode de um token tem uma única linha de query; como a lição 7.3 estabeleceu, ele é dominado pela leitura de pesos e do KV cache (a lição 7.2 deriva o custo por token do cache do Qwen), e kernels especializados de decode são primos, e não beneficiários, do mesmo ganho de vitrine.

A lição durável sobrevive a um kernel. Quando aritmética é barata e movimentação de memória não é, o design de algoritmos precisa contar bytes através da hierarquia. FlashAttention é attention exata rearranjada em torno dessa realidade física — e é a razão pela qual um contexto de um quarto de milhão de tokens é um recurso de produto, e não um experimento mental.

02 · Analogia

Analogia

Uma confeiteira precisa combinar cada item de duas listas longas de ingredientes, mas a bancada só comporta poucas tigelas. O método ingênuo registra toda mistura par a par em bandejas, leva todas as bandejas a um depósito distante e depois as busca de volta para normalizar. Uma confeiteira IO-aware trabalha em tiles que cabem na bancada, mantém totais correntes e envia para fora apenas os pães finais. FlashAttention, de modo análogo, recalcula quantidades locais baratas para evitar transportar uma matriz de attention gigantesca.

03 · Explique de volta

Explique de volta

Explique por que materializar os scores de attention no contexto nativo do Qwen3.8-27B é inviável, e como o softmax online em tiles calcula o mesmo resultado sem armazenar a matriz de scores.

Mínimo: 80 caracteres e 15 palavras. Seu texto fica somente neste navegador.

Aguardando sua explicação.

Comparar com uma resposta-modelo

Com 262.144 tokens, a matriz de scores tem dois elevado à trigésima sexta entradas por head, cerca de 128 GiB em bf16 — mais do que a memória inteira de um acelerador para um único head de uma única camada, antes de contar os 24 query heads do Qwen e as 16 camadas de full attention. FlashAttention carrega tiles de query e de key/value na SRAM on-chip rápida, calcula scores locais ali e mantém um máximo corrente por linha e uma soma de normalização, reescalando saídas parciais anteriores pela exponencial do máximo antigo menos o novo sempre que um tile posterior o eleva. A saída final é igual à softmax attention padrão, salvo reordenação de ponto flutuante; só as estatísticas dos tiles chegam a existir, então o tráfego de memória escala com as entradas e saídas, não com o quadrado do comprimento da sequência.

04 · Teste seu entendimento

Teste seu entendimento

01Qual é o principal objeto que FlashAttention evita materializar na memória do dispositivo?
Resposta e explicação

A matriz completa de scores/probabilidades de attention — O tiling percorre blocos de scores em fluxo e mantém apenas o máximo corrente e as estatísticas de normalização necessárias para combiná-los exatamente.

02Usando a lição 7.3, por que o ganho principal de FlashAttention se aplica mais ao prefill do que ao decode de um token?
Resposta e explicação

Prefill processa muitas linhas de query de uma vez e é dominado pelo tráfego de scores, enquanto decode tem uma linha de query e é dominado pela leitura de pesos e do KV cache — Prefill é a fase com um bloco grande de queries e uma matriz de scores potencialmente enorme; a linha única do decode faz da largura de banda de pesos e KV cache a restrição dominante.

Conclua o teach-back e acerte o quiz para finalizar a aula.

◎ · Marcador de evidência

Fontes

  1. Tri Dao et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.
  2. Qwen Team (2026). Qwen3.8-27B Model Card.