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 , possui scores por head. Escrever essa matriz na high-bandwidth memory (HBM), lê-la para softmax, escrever probabilidades e relê-las para multiplicar 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 e soma exponencial . Quando um novo bloco tem máximo , defina . Somas e saídas acumuladas são reescaladas por ; o novo bloco, por . 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 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
Conclua o teach-back e acerte o quiz para finalizar a aula.
Fontes
- Tri Dao et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.