3 pontos por GN⁺ 3 시간 전 | 1 comentários | Compartilhar no WhatsApp
  • Parte da atenção softmax e deriva, passo a passo, a atenção linear com estado de tamanho fixo, a DeltaNet, que registra apenas erros, a Gated DeltaNet, que atenua o estado inteiro, e a Kimi Delta Attention (KDA), que atenua por canal
  • A atenção linear básica armazena no estado (S_t) a soma dos produtos externos key-value passados e opera de forma linear no comprimento da sequência, mas sofre interferência de escrita aditiva, pois soma às associações existentes em vez de atribuir novos valores
  • DeltaNet registra a diferença entre o valor previsto a partir da key atual e o value alvo, multiplicada por (\beta_t); três interpretações — condição de reconstrução imediata, descida de gradiente online e atualização de estado de posto 1 — levam à mesma fórmula
  • Gated DeltaNet primeiro atenua o estado inteiro com um escalar (\alpha_t), enquanto KDA expande isso para uma matriz diagonal (D_t=\operatorname{Diag}(\alpha_t)), mantendo ou apagando informações em proporções diferentes para cada canal de key
  • A mesma recorrência da KDA é executada com um kernel Triton recorrente fundido para decode e com um método por chunks para treinamento e prefill longo; o método por chunks restaura dependências internas dos tokens por resolução triangular e as reformula como multiplicações de matrizes

Notação e ordem da derivação

  • Na notação bra-ket, (\lvert q\rangle) é um vetor coluna, (\langle k\rvert) é um vetor linha, (\langle k\vert q\rangle) é um escalar, e (\lvert v\rangle\langle k\rvert) é uma matriz
  • Usamos uma única cabeça de atenção causal e vetores reais, e assumimos que as keys da DeltaNet são normalizadas e que o estado mapeia do espaço das keys para o espaço dos values
  • A ordem da derivação é atenção softmax → atenção linear → DeltaNetGated DeltaNetKDA, conectando no fim às implementações Triton recorrente e em chunks
  • Duas variantes da família DeltaNet são usadas nas famílias de modelos Qwen e Kimi mais recentes

Da atenção de complexidade quadrática ao estado linear

  • A atenção softmax causal comum calcula a similaridade entre keys e queries, normaliza as pontuações para todas as keys passadas como uma distribuição e então produz uma soma ponderada dos vetores value
  • Em uma sequência de comprimento (T), há (T^2) pares key-query
    • Na inferência autorregressiva, keys e values podem ser armazenados em cache, mas o tamanho do cache cresce com a sequência
    • Uma nova query também precisa consultar todo o passado
  • Como o denominador do softmax depende conjuntamente da query atual e de todas as keys anteriores, é difícil simplesmente reorganizar a ordem de cálculo
  • Ao remover o softmax, a saída pode ser agrupada como a soma dos produtos externos key-value passados
    • (S_t=\sum_{i\le t}\lvert v_i\rangle\langle k_i\rvert)
    • (S_t=S_{t-1}+\lvert v_t\rangle\langle k_t\rvert)
    • (\lvert o_t\rangle=S_t\lvert q_t\rangle)
  • A identidade central é ((\lvert v\rangle\langle k\rvert)\lvert q\rangle=\langle k\vert q\rangle\lvert v\rangle), e, em vez de todas as keys e values passados, armazena-se o produto externo somado em um estado de tamanho fixo (d_v\times d_k)
  • Como percorre os tokens uma única vez, opera linearmente no comprimento da sequência, mas ao custo de perder a normalização e a seletividade do softmax
    • Atenções lineares mais sofisticadas usam mapas de características e termos de normalização

O problema da escrita aditiva na atenção linear

  • Logo após registrar (\lvert v_t\rangle\langle k_t\rvert) na key atual normalizada, ler com a mesma key resulta em (S_t\lvert k_t\rangle=S_{t-1}\lvert k_t\rangle+\lvert v_t\rangle)
  • A nova escrita não atribui o valor para que a memória retorne (v_t); ela soma no estilo += ao valor retornado existente
  • Se o estado anterior já retornava o valor correto, o mesmo value é duplicado; como as keys não são mutuamente ortogonais, cada escrita pode interferir nas escritas existentes
  • A atenção linear oferece uma memória associativa comprimida, mas realiza atualizações aditivas em vez de atualizações próximas do = necessário

DeltaNet: escrever o erro de previsão em vez do valor

  • DeltaNet primeiro lê a previsão existente (\widehat v_t=S_{t-1}k_t) para a nova key e registra apenas a diferença, não o value inteiro
    • (e_t=\beta_t(v_t-S_{t-1}k_t))
    • (S_t=S_{t-1}+e_tk_t^\mathsf T)
    • A intensidade de escrita aprendida (\beta_t) fica no intervalo ([0,1])
  • Se a leitura for feita imediatamente de novo com a mesma key, obtém-se ((1-\beta_t)S_{t-1}k_t+\beta_tv_t)
    • Se (\beta_t=1), retorna exatamente (v_t)
    • Valores menores apenas movem parcialmente a previsão existente na direção do alvo
  • A atualização é local no espaço das keys
    • Em direções de query ortogonais à key atual, a atualização por produto externo é 0, portanto a resposta não muda
    • Apenas a associação na direção da key atual é substituída seletivamente
  • Derivando a partir da perda de reconstrução

    • Vendo o estado (S) como um mapeamento linear e definindo a perda do par key-value atual como (\frac12\lVert Sk_t-v_t\rVert_2^2), o gradiente é ((Sk_t-v_t)k_t^\mathsf T)
    • Ao executar um passo de descida de gradiente a partir de (S_{t-1}) com tamanho (\beta_t), obtém-se exatamente a regra de atualização da DeltaNet
    • A mesma atualização pode ser interpretada de três formas
      • Em operações de memória, (\beta_t) é a intensidade de substituição da associação existente
      • Em aprendizado online, (\beta_t) é a taxa de aprendizado
      • Em álgebra linear, é o produto externo de posto 1 entre o erro de previsão e a key
  • Transição de estado estruturada

    • Expandindo a atualização, (S_t=S_{t-1}(I-\beta_tk_tk_t^\mathsf T)+\beta_tv_tk_t^\mathsf T)
    • Para uma key unitária, (I-\beta_tk_tk_t^\mathsf T) tem autovalor (1-\beta_t) na direção da key atual e autovalor 1 em todas as direções ortogonais
    • Primeiro remove a associação na direção da key existente e depois adiciona a nova associação, mas o gerenciamento da vida útil do estado inteiro ainda não é resolvido

Gated DeltaNet: esquecer primeiro o estado inteiro

  • Ao comprimir todo o passado em uma única matriz, não é possível pular seletivamente tokens individuais que já foram incorporados ao estado
  • A DeltaNet corrige em torno da key atual, mas informações antigas em outras direções permanecem e podem continuar contribuindo para leituras futuras
  • Gated DeltaNet aplica um gate escalar de retenção aprendido (\alpha_t\in[0,1])
    1. Esquece com (\widetilde S_t=\alpha_tS_{t-1})
    2. Prevê com (\widehat v_t=\widetilde S_tk_t)
    3. Corrige com (e_t=\beta_t(v_t-\widehat v_t))
    4. Registra com (S_t=\widetilde S_t+e_tk_t^\mathsf T)
  • A ordem esquecer → prever → corrigir → escrever é importante
    • Se a previsão for feita antes da atenuação, a memória usada para calcular o erro e a memória efetivamente atualizada serão diferentes
  • A regra delta cuida da substituição na key alvo, enquanto o gate escalar cuida do apagamento global; elas resolvem problemas diferentes
  • No entanto, como um único (\alpha_t) é aplicado à matriz inteira, todos os canais de key precisam ser mantidos ou esquecidos na mesma proporção

Kimi Delta Attention: atenuação por canal

  • Kimi Delta Attention transforma o escalar (\alpha_t) em um vetor de dimensão (d_k) e constrói (D_t=\operatorname{Diag}(\alpha_t))
  • Como o estado mapeia do espaço das keys para o espaço dos values, os canais de key correspondem às colunas de (S), e a multiplicação à direita (S_{t-1}D_t) aplica uma taxa de retenção diferente a cada coluna
  • A KDA opera na seguinte ordem
    1. Atenua por canal de key com (\widetilde S_t=S_{t-1}D_t)
    2. Prevê com (\widehat v_t=\widetilde S_tk_t)
    3. Corrige com (e_t=\beta_t(v_t-\widehat v_t))
    4. Registra com (S_t=\widetilde S_t+e_tk_t^\mathsf T)
    5. Lê com (o_t=S_t(d_k^{-1/2}q_t))
  • A mudança conceitual da Gated DeltaNet para a KDA é apenas elevar (\alpha_t) a (D_t), mas isso permite apagar um canal enquanto mantém outro
  • Transição diagonal-mais-baixo-posto

    • Expandindo a KDA, (S_t=S_{t-1}A_t+\beta_tv_tk_t^\mathsf T), onde (A_t=D_t(I-\beta_tk_tk_t^\mathsf T))
    • Pode ser escrito como (A_t=D_t-b_ta_t^\mathsf T), (b_t=D_tk_t), (a_t^\mathsf T=\beta_tk_t^\mathsf T), tornando-se uma transição diagonal-mais-baixo-posto (DPLR)
    • DPLR significa uma transição (d_k\times d_k) que atua no espaço das keys; o próprio estado de memória continua sendo uma matriz (d_v\times d_k)
    • Cada variante acrescenta as seguintes capacidades
      • Atenção linear: memória recorrente de tamanho fixo
      • DeltaNet: substituição seletiva na direção alvo
      • Gated DeltaNet: atenuação do estado inteiro
      • KDA: atenuação por canal de key
    • A implementação normalmente armazena (g_t=\log\alpha_t\le0) e então obtém a taxa de retenção com (\exp(g_t))
    • Uma implementação de referência em 5 etapas com layout transposto (d_k\times d_v) pode ser vista em naive_recurrent_kda

Kernel Triton recorrente fundido para decode

  • A KDA tem duas principais formas de execução
    • Forma recorrente fundida: adequada para decode, sequências curtas e serving com estado persistente
    • Forma por chunks: adequada para treinamento e prefill longo
  • fused_recurrent_kda_fwd executa um programa Triton por sequência, cabeça de value e tile de value de largura 32
    • BK cobre a dimensão de key nas configurações com suporte típicas
    • Cada programa possui um tile [BK, BV] do estado transposto e percorre os tokens em ordem
    • Tiles de value, heads e sequências diferentes são executados independentemente
  • O kernel realiza exatamente a recorrência: atenuação do estado, redução da previsão para a key, cálculo do residual, escrita por produto externo e redução de leitura da query
  • É adequado para decode, em que apenas um novo token chega por vez, mas é desfavorável para treinamento e prefill longo porque não consegue transformar operações vetoriais em grandes multiplicações de matrizes eficientes em Tensor Cores

Chunkwise KDA: reorganizando a recorrência como multiplicações de matrizes

  • A Chunkwise KDA precisa processar (C) tokens juntos e ainda produzir exatamente o mesmo estado e as mesmas saídas da execução recorrente token a token
  • Cada chunk calcula dois resultados
    • (S_{c+1}), o estado após processar todo o chunk a partir do estado de entrada (S_c)
    • As saídas causais de todos os tokens dentro do chunk
  • A principal dificuldade é que o erro delta de cada token depende das escritas anteriores dentro do mesmo chunk
  • Atenuação acumulada e erros provisórios

    • Denote a atenuação diagonal do token (i) como (D_i) e a atenuação acumulada da fronteira do chunk até o token (i) como (D_{0:i}=D_0D_1\cdots D_i)
    • Quando a escrita do token (j) se propaga até o token (i), aplica-se (D_{j+1:i}); como as matrizes diagonais comutam, as matrizes de atenuação podem ser trocadas de ordem
    • Primeiro calcula-se em paralelo um erro provisório ignorando outras escritas dentro do chunk
      • (\bar e_i=\beta_i(v_i-S_cD_{0:i}k_i))
    • Exceto para o primeiro token, os erros provisórios omitem o efeito das escritas anteriores dentro do chunk e, portanto, não podem ser usados diretamente
  • Restauração das dependências causais

    • Defina o coeficiente pelo qual o token anterior (j) afeta o erro do token atual (i) como (\rho_{ij}=\beta_i k_j^\mathsf TD_{j+1:i}k_i)
    • O erro real tem dependência sequencial na forma (e_i=\bar e_i-\sum_{j<i}\rho_{ij}e_j)
    • Colocando (\rho_{ij}) em uma matriz estritamente triangular inferior (R_c), a matriz de erros empilhados é calculada como (E_c=\bar E_c(A_c^{kk})^\mathsf T), (A_c^{kk}=(I+R_c)^{-1})
    • Não é necessária uma inversa densa genérica
      • (I+R_c) é uma matriz triangular com diagonal unitária
      • Basta executar uma resolução triangular causal para cada canal de value
  • Cálculo do estado ao fim do chunk

    • O estado de entrada atravessa todas as atenuações do chunk, e cada escrita interna do chunk atravessa apenas as atenuações posteriores a ela
    • Empilhando em linhas, em (K_c^{\mathrm{end}}), as keys atenuadas até o fim do chunk, o estado pode ser organizado como a seguinte multiplicação de matrizes
      • (S_{c+1}=S_cD_{0:C-1}+E_cK_c^{\mathrm{end}})
    • Várias escritas por produto externo de posto 1 são combinadas em uma única multiplicação de matrizes para avançar o estado do chunk inteiro de uma só vez
  • Cálculo de todas as saídas dentro do chunk

    • Como a KDA lê depois de escrever o token atual, a saída do token (i) também inclui sua própria escrita
    • Defina o coeficiente pelo qual uma escrita anterior (j) afeta a query (i) como (\chi_{ij}=s,k_j^\mathsf TD_{j+1:i}q_i), com (j\le i)
    • Coloque os coeficientes na matriz triangular inferior de leitura (A_c^{qk})
      • Os zeros na triangular superior bloqueiam a contribuição de tokens futuros
      • A diagonal reflete a leitura do token atual depois de sua própria escrita
    • Empilhando em (Q_c^{\mathrm{boundary}}) os vetores atenuados da fronteira até cada query, a saída total é
      • (O_c=sS_cQ_c^{\mathrm{boundary}}+E_c(A_c^{qk})^\mathsf T)
    • A primeira multiplicação de matrizes lê o estado de entrada do chunk atenuado, e a segunda soma a contribuição das escritas causais internas do chunk

Pipeline Triton da versão por chunks

  • A implementação por chunks não é um único kernel enorme, mas uma pipeline composta por várias chamadas de kernel
  • Primeiro calcula a atenuação logarítmica acumulada dentro do chunk
    • Expressa (D_{j+1:i}) sem multiplicar longamente vetores de retenção, usando a diferença entre duas prefix sums
  • Em seguida, constrói as matrizes de interação causais (A^{qk}) e (A^{kk}), e usa (A^{kk}) para montar uma forma WY para as escritas corrigidas do chunk
  • O kernel de estado realiza a única varredura entre chunks
    • Gera o estado que entra em cada chunk
    • Resolve os erros delta do chunk
  • Depois que o estado de entrada é calculado, o kernel de saída pode processar em paralelo tokens de chunks e tiles diferentes
  • A implementação real primeiro calcula blocos diagonais de interação de 16 tokens e depois executa um kernel fundido para os blocos não diagonais e a resolução triangular
  • chunk_kda_fwd coordena as etapas, e os principais pontos de entrada são chunk_kda_fwd_intra, chunk_gated_delta_rule_fwd_h, chunk_gla_fwd_o_gk
    • No código, v_new é o erro resolvido
    • h é o estado de entrada do chunk
    • kg é a key atenuada até o fim do chunk
  • A forma recorrente e a forma por chunks não são atenções diferentes, mas dois cronogramas de execução da mesma recorrência KDA
    • A forma recorrente é uma operação vetorial serial para decode de baixa latência
    • A forma por chunks usa operações matriciais voltadas a Tensor Cores para treinamento e prefill

1 comentários

 
GN⁺ 3 시간 전
Comentários do Hacker News
  • Nos últimos 15 anos, o machine learning precisou de uma notação matemática unificada, e provavelmente continuará precisando. Antes era pior, com artigos de pesquisadores do mundo todo usando notações mirabolantes
    Quando cada artigo usa uma notação diferente, isso cria atrito na compreensão. Pelo menos este texto explica explicitamente a notação desde o início, algo raro em artigos. No começo eu nem percebi que havia uma função para alternar a notação, mas ela é muito útil

    • Não consigo entender por que preferem a notação matemática tradicional com caracteres como ∣q⟩ em vez de símbolos de uma letra ou tipos de dados explícitos. Pode ter a vantagem de ser concisa, mas parece muito mais fácil entender fórmulas escritas como pseudocódigo ou numa linguagem de programação real, como Python
    • Este texto explica só um aspecto da notação, mas não fornece as definições das variáveis usadas. Quem estudou machine learning pode saber ou adivinhar o que são k, q e S, mas sem esse conhecimento de base boa parte do texto fica opaca
    • Eu também pensava assim antes, mas como passo muito mais tempo olhando fórmulas do que código, depois que você entende o significado dos símbolos, a notação concisa fica bem mais fácil de ler. Também evita até o problema de naming, famoso por ser difícil
  • Dizem “eu poderia ter pensado nisso...”, mas criar ou combinar algo que não existia é extremamente difícil
    Assim que alguém finalmente publica um trabalho difícil, logo aparecem reações como “nem era tão difícil”, “eu também conseguiria”, e tudo começa a parecer simples. Também é comum estar desenvolvendo algo e achar que inventou uma coisa nova, para depois descobrir que ela já tinha sido criada nos anos 1970 e era amplamente usada. Só não cruzou com ela no seu caminho, então você nem sabia que existia

  • Para mim, a notação bra-ket torna tudo simples e intuitivo. Com notação vetorial eu sempre me confundia sobre o que era horizontal ou vertical, acabava só seguindo os blocos e perdia a concentração, mas com bra-ket tudo pareceu muito intuitivo
    Acho que devo ter perdido muitos bons textos por causa disso, então estou pensando em converter outros textos para essa notação. Para referência, sou doutor em física e tenho dislexia leve

  • Quando vejo um estilo como “o produto externo é uma matriz e o produto interno é um número. Em vez de armazenar todas as chaves e valores passados, armazena-se a soma dos produtos externos no estado de tamanho fixo S_t”, isso me convence de que foi um texto escrito por LLM

    • Provavelmente começou pedindo um título com palavras da moda
    • Se você colocar no prompt do Claude para não usar travessão (), sai algo assim
  • Também há um tutorial visualizado: https://snowchord.com/blog/linear-attention-visualized/

  • Sempre que vejo textos e títulos assim, sinto profunda gratidão e humildade diante das inúmeras pessoas muito mais inteligentes do que eu. No ensino médio e na graduação eu era visto como alguém muito inteligente, e sou mais esperto que a média, mas certamente existem milhões de pessoas que me fariam parecer um novato
    Aqui, ser inteligente significa a capacidade de manter conceitos e sistemas enormes e complexos na cabeça e raciocinar sobre eles, algo que parece ser um talento especialmente importante para matemáticos

    • Mesmo que ferramentas de IA acelerem cada vez mais o trabalho, acredito que a fonte da maioria das novas ideias continuará sendo humana
      Um experimento mental que fiz bebendo com um amigo foi isolar crianças do conteúdo popular fornecido por telas e algoritmos, e criá-las num ambiente propício ao aprendizado em que a qualidade da mídia e dos materiais fosse rigidamente controlada, como se estivéssemos treinando um modelo de ponta. Seria como um mosteiro para crianças, ensinando o conhecimento mais atualizado sobre a realidade por meio de matemática, engenharia, ciência da computação, deep learning etc.
      No fim, para ampliar as fronteiras do conhecimento com ferramentas avançadas de IA, ainda será preciso gente muito inteligente e com o pensamento pouco contaminado. A ideia de que a IA substituirá totalmente os humanos vai na direção errada
  • Como curiosidade, o nome notação bra-ket realmente vem de bracket, isto é, parêntese
    https://en.wikipedia.org/wiki/Bra-ket_notation

  • No começo eu hesitei, mas gostei da notação ket porque deixou as operações muito mais claras. Ainda assim, teria sido bom ter uma recapitulação rápida de algumas variáveis, como d_k na atenção quadrática

  • No começo fiquei desanimado por não ter pensado nessa solução, mas logo fiquei em paz ao perceber que também tenho dificuldade para escrever busca binária em JavaScript por conta própria. Não existe absolutamente nenhuma chance de eu ter inventado Kimi Delta Attention

    • Código de álgebra linear tem, surpreendentemente, um lado mais fácil de escrever. Não há recursão se entrelaçando de forma complexa como em código geral de ciência da computação, todas as variáveis têm relações matemáticas entre si, e conceitos matemáticos comuns já podem usar bibliotecas bem implementadas
      Também é raro que loops precisem passar de dois ou três níveis de profundidade, e se ficar mais complexo do que isso, provavelmente já é melhor delegar a uma biblioteca