3 puntos por GN⁺ 3 시간 전 | 1 comentarios | Compartir por WhatsApp
  • Partiendo de la atención softmax, se deriva paso a paso la atención lineal que usa un estado de tamaño fijo, DeltaNet que registra solo el error, Gated DeltaNet que atenúa todo el estado, y Kimi Delta Attention (KDA) que atenúa por canal
  • La atención lineal básica guarda en el estado (S_t) la suma de los productos externos key-value del pasado y opera linealmente con la longitud de la secuencia, pero como no asigna el valor nuevo sino que lo suma a la asociación existente, aparece la interferencia de escritura aditiva
  • DeltaNet registra la diferencia entre el valor predicho desde la key actual y el value objetivo, multiplicada por (\beta_t); las tres interpretaciones de condición de reconstrucción inmediata, descenso de gradiente en línea y actualización de estado de rango 1 llevan a la misma fórmula
  • Gated DeltaNet primero atenúa todo el estado con un escalar (\alpha_t), y KDA lo extiende a una matriz diagonal (D_t=\operatorname{Diag}(\alpha_t)) para conservar o borrar información con proporciones distintas según cada canal de key
  • La misma recurrencia de KDA se ejecuta con un kernel Triton recurrente fusionado para decodificación y con un esquema por chunks para entrenamiento y prefilling largo; el método por chunks restaura las dependencias internas entre tokens con una resolución triangular y las recompone como multiplicaciones de matrices

Notación y orden del desarrollo

  • En la notación bra-ket, (\lvert q\rangle) es un vector columna, (\langle k\rvert) es un vector fila, (\langle k\vert q\rangle) es un escalar, y (\lvert v\rangle\langle k\rvert) es una matriz
  • Se asume una sola cabeza de atención causal y vectores reales; las keys de DeltaNet están normalizadas y el estado mapea del espacio de keys al espacio de values
  • El orden del desarrollo es atención softmax → atención lineal → DeltaNetGated DeltaNetKDA, y al final se conecta con implementaciones Triton recurrentes y por chunks
  • Entre las variantes de la familia DeltaNet, dos se usan en las familias de modelos Qwen y Kimi más recientes

De atención con complejidad cuadrática a estado lineal

  • La atención softmax causal típica calcula la similitud entre key y query, normaliza como distribución las puntuaciones sobre todas las keys pasadas, y luego produce una suma ponderada de los vectores value
  • En una secuencia de longitud (T) hay (T^2) pares key-query
    • En inferencia autorregresiva se pueden cachear keys y values, pero el tamaño de la caché crece con la secuencia
    • Cada query nueva también debe revisar todo el pasado
  • Como el denominador del softmax depende conjuntamente de la query actual y de todas las keys previas, no es fácil reordenar el cálculo de forma simple
  • Si se elimina el softmax, la salida puede agruparse como la suma de productos externos key-value del pasado
    • (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)
  • La identidad clave es ((\lvert v\rangle\langle k\rvert)\lvert q\rangle=\langle k\vert q\rangle\lvert v\rangle), y en vez de guardar todas las keys y values pasadas, se almacena el producto externo acumulado en un estado de tamaño fijo (d_v\times d_k)
  • Como recorre los tokens una sola vez, opera linealmente con la longitud de la secuencia, pero a cambio pierde la normalización y selectividad del softmax
    • Las atenciones lineales más sofisticadas usan mapas de características y términos de normalización

El problema de escritura aditiva en la atención lineal

  • Si justo después de registrar (\lvert v_t\rangle\langle k_t\rvert) para la key actual normalizada se lee con esa misma key, se obtiene (S_t\lvert k_t\rangle=S_{t-1}\lvert k_t\rangle+\lvert v_t\rangle)
  • La escritura nueva no hace que la memoria devuelva (v_t) por asignación, sino que suma (v_t) al valor que ya devolvía, al estilo de +=
  • Si el estado anterior ya devolvía el valor correcto, ese mismo value se duplica; además, como las keys no son ortogonales entre sí, cada escritura puede interferir con las anteriores
  • La atención lineal ofrece una memoria asociativa comprimida, pero realiza actualizaciones aditivas en lugar de una actualización más cercana al = que realmente se necesitaría

DeltaNet: escribir el error de predicción en lugar del valor

  • DeltaNet primero lee la predicción existente para la nueva key, (\widehat v_t=S_{t-1}k_t), y registra solo la diferencia en vez del value completo
    • (e_t=\beta_t(v_t-S_{t-1}k_t))
    • (S_t=S_{t-1}+e_tk_t^\mathsf T)
    • La intensidad de escritura aprendida (\beta_t) está en el rango ([0,1])
  • Si se vuelve a leer de inmediato con esa misma key, se obtiene ((1-\beta_t)S_{t-1}k_t+\beta_tv_t)
    • Si (\beta_t=1), devuelve exactamente (v_t)
    • Valores menores solo mueven parcialmente la predicción existente hacia el objetivo
  • La actualización es local en el espacio de keys
    • En direcciones de query ortogonales a la key actual, la actualización por producto externo es 0, así que la respuesta no cambia
    • Solo reemplaza selectivamente la asociación en la dirección de la key actual
  • Derivación desde la pérdida de reconstrucción

    • Si se ve el estado (S) como una transformación lineal y se define la pérdida del par key-value actual como (\frac12\lVert Sk_t-v_t\rVert_2^2), el gradiente es ((Sk_t-v_t)k_t^\mathsf T)
    • Si desde (S_{t-1}) se hace un paso de descenso de gradiente con tamaño (\beta_t), se obtiene exactamente la ecuación de actualización de DeltaNet
    • La misma actualización puede interpretarse de tres formas
      • En operaciones de memoria, (\beta_t) es la fuerza de reemplazo de la asociación existente
      • En aprendizaje en línea, (\beta_t) es la tasa de aprendizaje
      • En álgebra lineal, es el producto externo de rango 1 entre el error de predicción y la key
  • Transición de estado estructurada

    • Al desarrollar la actualización se obtiene (S_t=S_{t-1}(I-\beta_tk_tk_t^\mathsf T)+\beta_tv_tk_t^\mathsf T)
    • Para una key unitaria, (I-\beta_tk_tk_t^\mathsf T) tiene autovalor (1-\beta_t) en la dirección de la key actual y autovalor 1 en todas las direcciones ortogonales
    • Primero elimina la asociación existente en la dirección de la key y luego agrega la nueva, pero todavía no resuelve la gestión de vida útil de todo el estado

Gated DeltaNet: primero olvidar todo el estado

  • Si todo el pasado se comprime en una sola matriz, ya no es posible omitir selectivamente solo tokens individuales que ya fueron fusionados en el estado
  • DeltaNet corrige alrededor de la key actual, pero la información vieja en otras direcciones permanece y puede seguir contribuyendo a lecturas futuras
  • Gated DeltaNet aplica una compuerta escalar aprendida de retención (\alpha_t\in[0,1])
    1. Olvida con (\widetilde S_t=\alpha_tS_{t-1})
    2. Predice con (\widehat v_t=\widetilde S_tk_t)
    3. Corrige con (e_t=\beta_t(v_t-\widehat v_t))
    4. Registra con (S_t=\widetilde S_t+e_tk_t^\mathsf T)
  • El orden olvidar → predecir → corregir → escribir es importante
    • Si se predice antes de la atenuación, la memoria usada para calcular el error no coincide con la memoria que realmente se actualiza
  • La regla delta se encarga del reemplazo sobre la key objetivo y la compuerta escalar de la eliminación global, así que resuelven problemas distintos
  • Aun así, como un único (\alpha_t) se aplica a toda la matriz, todos los canales de key deben conservarse u olvidarse en la misma proporción

Kimi Delta Attention: atenuación por canal

  • Kimi Delta Attention reemplaza el escalar (\alpha_t) por un vector de dimensión (d_k) y construye (D_t=\operatorname{Diag}(\alpha_t))
  • Como el estado mapea del espacio de keys al de values, los canales de key corresponden a las columnas de (S), y la multiplicación por la derecha (S_{t-1}D_t) aplica una tasa de retención distinta a cada columna
  • KDA funciona en el siguiente orden
    1. Atenuación por canal de key con (\widetilde S_t=S_{t-1}D_t)
    2. Predicción con (\widehat v_t=\widetilde S_tk_t)
    3. Corrección con (e_t=\beta_t(v_t-\widehat v_t))
    4. Escritura con (S_t=\widetilde S_t+e_tk_t^\mathsf T)
    5. Lectura con (o_t=S_t(d_k^{-1/2}q_t))
  • El cambio conceptual de Gated DeltaNet a KDA es solo elevar (\alpha_t) a (D_t), pero ahora puede borrar un canal mientras conserva otros
  • Transición diagonal-bajo rango

    • Si se desarrolla KDA, se obtiene (S_t=S_{t-1}A_t+\beta_tv_tk_t^\mathsf T), donde (A_t=D_t(I-\beta_tk_tk_t^\mathsf T))
    • Puede escribirse como (A_t=D_t-b_ta_t^\mathsf T), con (b_t=D_tk_t) y (a_t^\mathsf T=\beta_tk_t^\mathsf T), lo que da una transición diagonal-plus-low-rank (DPLR)
    • DPLR indica una transición (d_k\times d_k) que actúa en el espacio de keys; el estado de memoria en sí sigue siendo una matriz (d_v\times d_k)
    • Cada variante agrega la siguiente capacidad
      • Atención lineal: memoria recurrente de tamaño fijo
      • DeltaNet: reemplazo selectivo en la dirección objetivo
      • Gated DeltaNet: atenuación de todo el estado
      • KDA: atenuación por canal de key
    • En implementaciones suele guardarse (g_t=\log\alpha_t\le0) y luego obtener la retención con (\exp(g_t))
    • La implementación de referencia en 5 pasos con layout transpuesto (d_k\times d_v) puede verse en naive_recurrent_kda

Kernel Triton recurrente fusionado para decodificación

  • KDA tiene dos formas principales de ejecución
    • Modo recurrente fusionado: adecuado para decode, secuencias cortas y serving con mantenimiento de estado
    • Modo por chunks: adecuado para entrenamiento y prefilling largo
  • fused_recurrent_kda_fwd ejecuta un programa Triton por secuencia, cabeza de value y tile de value de ancho 32
    • BK cubre la dimensión de key en configuraciones de soporte habituales
    • Cada programa posee un tile [BK, BV] del estado transpuesto y recorre los tokens en orden
    • Distintos tiles de value, heads y secuencias se ejecutan de forma independiente
  • El kernel realiza exactamente la recurrencia: atenuación del estado, reducción de predicción sobre la key, cálculo del residual, escritura por producto externo y reducción de lectura con la query
  • Es adecuado para decode, donde solo entra un token nuevo a la vez, pero como no puede convertir eficientemente las operaciones vectoriales en grandes multiplicaciones de matrices favorables para Tensor Cores, resulta desfavorable para entrenamiento y prefilling largo

KDA por chunks: reordenar la recurrencia como multiplicaciones de matrices

  • KDA por chunks debe procesar juntos (C) tokens y aun así producir exactamente el mismo estado y las mismas salidas que el modo recurrente por token
  • Cada chunk calcula dos resultados
    • El estado (S_{c+1}) después de procesar todo el chunk a partir del estado entrante (S_c)
    • Las salidas causales de todos los tokens dentro del chunk
  • La dificultad clave es que el error delta de cada token depende de escrituras previas dentro del mismo chunk
  • Atenuación acumulada y error provisional

    • Sea (D_i) la atenuación diagonal del token (i), y (D_{0:i}=D_0D_1\cdots D_i) la atenuación acumulada desde el borde del chunk hasta el token (i)
    • Cuando la escritura del token (j) se propaga hasta el token (i), se aplica (D_{j+1:i}); como son matrices diagonales, estas atenuaciones conmutan entre sí
    • Primero se calcula en paralelo un error provisional que ignora las demás escrituras dentro del chunk
      • (\bar e_i=\beta_i(v_i-S_cD_{0:i}k_i))
    • El error provisional de todos salvo el primer token omite el efecto de escrituras previas del mismo chunk, así que no puede usarse tal cual
  • Restaurar la dependencia causal

    • Se define el coeficiente con que el token previo (j) afecta el error del token actual (i) como (\rho_{ij}=\beta_i k_j^\mathsf TD_{j+1:i}k_i)
    • El error real sigue la dependencia secuencial (e_i=\bar e_i-\sum_{j<i}\rho_{ij}e_j)
    • Si (\rho_{ij}) se coloca en una matriz estrictamente triangular inferior (R_c), la matriz apilada de errores se calcula como (E_c=\bar E_c(A_c^{kk})^\mathsf T), con (A_c^{kk}=(I+R_c)^{-1})
    • No hace falta una inversión densa general
      • (I+R_c) es triangular con 1 en la diagonal
      • Basta resolver un sistema triangular causal para cada canal de value
  • Calcular el estado al final del chunk

    • El estado entrante atraviesa toda la atenuación del chunk, y cada escritura interna del chunk atraviesa solo las atenuaciones posteriores a ella misma
    • Si las keys atenuadas hasta el final del chunk se apilan por filas en (K_c^{\mathrm{end}}), el estado puede resumirse con la siguiente multiplicación matricial
      • (S_{c+1}=S_cD_{0:C-1}+E_cK_c^{\mathrm{end}})
    • Múltiples escrituras de producto externo de rango 1 se combinan en una sola multiplicación matricial para avanzar todo el estado del chunk de una vez
  • Calcular todas las salidas dentro del chunk

    • KDA escribe el token actual y luego lee, por lo que la salida del token (i) incluye también su propia escritura
    • Se define el coeficiente con que la escritura previa (j) afecta la query (i) como (\chi_{ij}=s,k_j^\mathsf TD_{j+1:i}q_i), con (j\le i)
    • Los coeficientes se colocan en una matriz de lectura triangular inferior (A_c^{qk})
      • Los 0 por encima de la diagonal bloquean la contribución de tokens futuros
      • Los elementos diagonales reflejan que el token actual lee después de su propia escritura
    • Si los vectores atenuados desde el borde hasta cada query se apilan en (Q_c^{\mathrm{boundary}}), la salida total es la siguiente
      • (O_c=sS_cQ_c^{\mathrm{boundary}}+E_c(A_c^{qk})^\mathsf T)
    • La primera multiplicación matricial lee el estado de entrada del chunk ya atenuado, y la segunda suma la contribución causal de las escrituras internas del chunk

Pipeline Triton por chunks

  • La implementación por chunks no es un solo kernel gigante, sino un pipeline de varias llamadas a kernel
  • Primero calcula la atenuación logarítmica acumulada dentro del chunk
    • Con la diferencia de dos prefix sums se expresa (D_{j+1:i}) sin multiplicar directamente largos vectores de retención
  • Luego construye las matrices de interacción causal (A^{qk}) y (A^{kk}), y con (A^{kk}) forma una representación tipo WY para las escrituras corregidas del chunk
  • El kernel de estado realiza el único recorrido entre chunks
    • Genera el estado que entra a cada chunk
    • Resuelve los errores delta del chunk
  • Una vez calculado el estado de entrada, el kernel de salida puede procesar en paralelo los tokens de distintos chunks y tiles
  • La implementación real primero calcula bloques diagonales de interacción de 16 tokens y luego ejecuta kernels fusionados para las partes no diagonales y la resolución triangular
  • chunk_kda_fwd coordina los pasos, y los puntos de entrada principales son chunk_kda_fwd_intra, chunk_gated_delta_rule_fwd_h, chunk_gla_fwd_o_gk
    • v_new en el código es el error ya resuelto
    • h es el estado de entrada al chunk
    • kg es la key atenuada hasta el final del chunk
  • El modo recurrente y el modo por chunks no son dos atenciones distintas, sino dos schedules de ejecución de la misma recurrencia KDA
    • El modo recurrente son operaciones vectoriales seriales para decode de baja latencia
    • El modo por chunks son operaciones matriciales orientadas a Tensor Cores para entrenamiento y prefilling

1 comentarios

 
GN⁺ 3 시간 전
Opiniones de Hacker News
  • Durante los últimos 15 años, el aprendizaje automático necesitó una notación matemática unificada, y probablemente la siga necesitando. Antes era peor, porque en los papers de investigadores de todo el mundo aparecían notaciones de lo más extravagantes.
    Cuando la notación cambia de un paper a otro, se genera fricción para entender. Al menos este texto explica explícitamente la notación desde el principio, algo que pocos papers hacen. Al principio ni siquiera me di cuenta de que existía la función para cambiar la notación, pero es muy útil.

    • No entiendo por qué se prefiere la notación matemática tradicional, que usa símbolos como ∣q⟩, en lugar de símbolos de una sola letra o tipos de datos explícitos. Tendrá la ventaja de ser concisa, pero creo que las fórmulas serían mucho más fáciles de entender si se escribieran como pseudocódigo o en un lenguaje de programación real como Python.
    • Este texto solo explica un aspecto de la notación, pero no ofrece las definiciones de las variables que usa. Si estudiaste aprendizaje automático, sabes o puedes inferir qué son k, q y S, pero sin ese contexto gran parte del texto se vuelve opaco.
    • Antes yo también pensaba así, pero como paso mucho más tiempo mirando fórmulas que código, una vez que conoces el significado de los símbolos la notación concisa es mucho más fácil de leer. Al escribir con símbolos también evitas el famoso problema de poner nombres, que es difícil.
  • Aunque diga “algo que uno mismo podría haber pensado…”, crear o combinar algo que no existía es enormemente difícil.
    Cuando alguien finalmente publica el resultado de un trabajo difícil, enseguida aparecen reacciones como “no era tan difícil” o “yo también podría haberlo hecho”, y todo empieza a parecer simple. También es común estar desarrollando algo y pensar que inventaste algo nuevo, para luego descubrir que ya se había creado en los años 70 y se usaba ampliamente. Simplemente no sabías que existía porque nunca se había cruzado con tu camino.

  • Para mí, la notación bra-ket hace que todo sea simple e intuitivo. Con la notación vectorial me confundía sobre qué lado era fila y cuál columna, terminaba siguiendo solo bloques y perdía la concentración, pero con bra-ket todo resultó muy intuitivo.
    Creo que me he perdido muchos buenos textos, así que pienso convertir otros a esta notación. Como referencia, tengo un doctorado en física y una dislexia leve.

  • Cuando veo un estilo como “el producto exterior es una matriz y el producto interior es un número. En lugar de almacenar todas las claves y valores del pasado, se guarda la suma de los productos exteriores en un estado de tamaño fijo S_t”, me convenzo de que es un texto escrito por un LLM.

    • Probablemente empezó pidiéndole un título con alguna palabra de moda.
    • Si le pides a Claude que no use guiones largos (), obtienes este tipo de resultado.
  • También hay un tutorial visualizado: https://snowchord.com/blog/linear-attention-visualized/

  • Cada vez que veo textos y títulos así, siento profunda gratitud y humildad ante la enorme cantidad de personas mucho más inteligentes que yo. En la secundaria y la universidad me consideraban muy inteligente, y soy más listo que el promedio, pero sin duda hay millones de personas que me harían ver como un novato.
    Aquí, por inteligente me refiero a la capacidad de contener en la mente conceptos y sistemas enormes y complejos, y razonar sobre ellos; parece un talento especialmente importante para los matemáticos.

    • Aunque las herramientas de IA aceleren cada vez más el trabajo, creo que la fuente de la mayoría de las nuevas ideas seguirá siendo humana.
      Un experimento mental que hice tomando con un amigo fue criar a niños aislándolos de las pantallas y del contenido masivo que alimentan los algoritmos, en un entorno favorable al aprendizaje donde se controle estrictamente la calidad de los medios y materiales, como cuando se entrenan modelos de vanguardia. Sería como un monasterio para niños, enseñándoles el conocimiento más actualizado sobre la realidad mediante matemáticas, ingeniería, ciencias de la computación, deep learning, etc.
      Al final, para ampliar las fronteras del conocimiento usando herramientas avanzadas de IA, todavía se necesitan personas muy inteligentes y con el pensamiento no demasiado contaminado. La idea de que la IA reemplazará por completo a los humanos va en la dirección equivocada.
  • Como referencia, el nombre notación bra-ket de hecho proviene de bracket, “paréntesis” o “corchete”.
    https://en.wikipedia.org/wiki/Bra-ket_notation

  • Al principio dudé, pero la notación ket me terminó gustando porque hace que las operaciones sean mucho más claras. Eso sí, me habría gustado que hubiera también un repaso breve de algunas variables, como d_k en la atención cuadrática.

  • Al principio me desanimó no haber pensado en esta solución, pero al darme cuenta de que incluso me cuesta escribir una búsqueda binaria en JavaScript por mi cuenta, me tranquilicé de inmediato. No había ninguna posibilidad de que se me ocurriera Kimi Delta Attention.

    • El código de álgebra lineal tiene un lado sorprendentemente fácil de escribir. No se enreda con recursión compleja como el código típico de ciencias de la computación, todas las variables tienen relaciones matemáticas entre sí, y para los conceptos matemáticos comunes ya se pueden usar bibliotecas bien implementadas.
      Los bucles rara vez llegan a tener más de dos o tres niveles de profundidad, y si algo es más complejo que eso, de todos modos conviene pasarlo a una biblioteca.