Avançado

FlashAttention e IO-awareness

FlashAttention calcula attention exata em tiles, sem materializar a matriz completa de scores na memória lenta do dispositivo.

Atualizada em

1

Conceito

A fórmula de attention é compacta, mas uma implementação direta cria intermediários grandes. Para sequência de comprimento NN, QKQK^\top possui N2N^2 scores por head. Escrever essa matriz na high-bandwidth memory (HBM), lê-la para softmax, escrever probabilidades e relê-las para multiplicar VV move muito mais dados do que a fórmula sugere.

Aceleradores têm uma hierarquia. SRAM on-chip é pequena e rápida; HBM é maior e mais lenta. Unidades aritméticas esperam enquanto tensores viajam. IO-awareness significa projetar o algoritmo em torno desse movimento, não apenas contar operações. FlashAttention observa que attention pode ser calculada em tiles que cabem na memória rápida sem guardar toda a matriz.

Divida queries em blocos de linhas e keys/values em blocos de colunas. Carregue um tile de query e outro de key/value na SRAM, calcule scores locais e atualize uma saída parcial. Depois avance. O obstáculo é softmax: normalizar uma linha parece exigir todos os scores simultaneamente.

Softmax online resolve isso. Para cada linha de query, mantenha máximo corrente mm e soma exponencial \ell. Quando um novo bloco tem máximo mm', defina mnew=max(m,m)m_{new}=\max(m,m'). Somas e saídas acumuladas são reescaladas por emmnewe^{m-m_{new}}; o novo bloco, por esmnewe^{s-m_{new}}. No fim, a soma ponderada é dividida pela normalização.

É a mesma identidade estável de softmax aplicada por partes. A ordem muda e pode alterar últimos bits, mas o algoritmo não aproxima, esparsifica ou trunca attention de propósito. Causal mask, dropout, comprimentos variáveis e gradients exigem bookkeeping adicional nos kernels.

Recalcular pode ser mais barato que armazenar. No backward pass, FlashAttention refaz scores locais em vez de ler uma matriz N2N^2 salva. Há mais aritmética e menos tráfego e pico de memória. A troca faz sentido porque operações matriciais são baratas diante de transferências repetidas de HBM no hardware suportado.

O benefício depende do shape. Sequências curtas sofrem launch overhead. Head dimension, dtype, máscara, dropout, geração do hardware e biblioteca definem os kernels disponíveis. O framework pode fazer fallback silencioso quando uma feature não tem suporte. Meça o kernel selecionado, não deduza por um flag.

Na inferência, prefill se beneficia por processar posições numerosas. Decode de um token tem outro shape e costuma ser dominado por pesos e KV cache; não se deve assumir o mesmo ganho. Layouts paginados adicionam endereçamento que o kernel precisa entender.

Compiladores e frameworks podem fundir operações vizinhas, mudar o layout ou escolher versões distintas do kernel conforme o shape. Por isso, um profiler de memória e timeline é mais informativo do que o nome da função na aplicação. O teste deve incluir a mesma máscara, dtype e política de padding usada em produção.

Testes de corretude comparam saídas e gradients contra referência em máscaras, tamanhos ímpares, dtypes e logits extremos. Testes de performance medem fase e memória ponta a ponta, não só um microbenchmark favorável. A lição é maior: quando compute fica barato, o algoritmo precisa contar bytes movidos. FlashAttention é attention exata reorganizada em torno dessa realidade física.

2

Como explicar para uma criança de cinco anos

Uma confeiteira combina cada item de duas listas longas, mas a bancada só comporta poucas tigelas. O método ingênuo registra toda mistura em bandejas, leva-as a um depósito distante e busca tudo para normalizar. Uma confeiteira IO-aware trabalha em tiles que cabem na bancada, mantém totais correntes e envia apenas os bolos finais. FlashAttention recalcula quantidades locais baratas para não transportar uma matriz gigantesca.

3

Ensine de volta

Explique por que materializar attention é caro em IO e como softmax online preserva attention exata sem guardar a matriz completa.

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

Salvo somente neste dispositivo.

Ver uma resposta-modelo

A implementação ingênua escreve a matriz N por N de scores e resultados intermediários em HBM e os lê novamente, fazendo tráfego dominar aritmética. FlashAttention divide Q, K e V em tiles na memória on-chip. Para cada tile de queries, mantém máximo e soma de normalização correntes, reescalando saídas anteriores quando novas keys chegam. Esse softmax online produz o mesmo resultado matemático, salvo efeitos normais de precisão, sem materializar a matriz inteira.

4

Teste seu entendimento

1. Qual objeto FlashAttention evita materializar em HBM?
Resposta e explicação

A matriz completa de scores ou probabilidades de attention — Tiling percorre blocos de scores e mantém só estatísticas necessárias.

2. FlashAttention aproxima softmax attention comum?
Resposta e explicação

Não; reorganiza o cálculo exato, sujeito a efeitos de ponto flutuante — A inovação é tiling IO-aware e normalização online, não aproximação esparsa ou linear.

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

Fontes

  1. Tri Dao et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.