GPT Diffusion

MALA: la atención que asigna su propio cómputo — qué gana de verdad tu entrenamiento (y qué no)

2026-09-29 · Devs #modelos#arquitectura#benchmark#optimizacion#llm

MALA: la atención que asigna su propio cómputo, y lo que no te va a ahorrar

TL;DR

  • Verificado verbatim contra el HTML v1 del paper (29/09/2026): MassAlloc Attention (MALA) es una primitiva de atención fusionada que conserva acceso de scores a toda interacción causal legal y usa la contribución normalizada para asignar el cómputo post-score, con una tolerancia común (τ=1) que gobierna forward, backward, prefill y decoding.
  • Todo lo demás es paper-reported (HKUST Guangzhou + BAAI + Université Paris Cité, v1 del 26/09/2026): 2.2× en forward y 3.0× en backward durante entrenamiento y 1.6× en decoding de inferencia frente a FullAttn, medidos sobre el operador de atención a 128K en 8×H100 con TP=8. No existe reproducción independiente a fecha de hoy.
  • Lo que no es: el cálculo de scores QK sigue siendo completo — la complejidad sigue siendo cuadrática en longitud de secuencia. Las ganancias son reducciones de factor constante que dependen de la distribución real de atención. Esto no es contexto subcuadrático, aunque algún titular lo venda así.
  • El matiz que el −23.1% esconde: ese ahorro de FLOPs de entrenamiento a 14B es en entrenamiento long-context a 32K. En pretraining a 4K —donde vive la mayoría de los fine-tunings— el mismo paper reporta un −2.5%. Casi ruido.
  • Código abierto hoy: HKUSTDial/flash-sparse-attention (BSD-3-Clause, 763 stars a 29/09/2026), pip install flash-sparse-attn, con la tolerancia del paper expuesta como softmax_threshold=1.0.

Contexto: qué es MALA exactamente

El problema que ataca no es nuevo: los kernels densos de atención ejecutan el camino post-score completo tras formar cada tile de QK, aunque la masa softmax normalizada de gran parte del espacio causal sea despreciable. Pagas la exponenciación, la carga de V y la acumulación PV por cada tile, y muchos de esos tiles aportan prácticamente nada al resultado.

MALA propone usar precisamente esa contribución normalizada como señal de asignación de cómputo. Durante el forward, acota las contribuciones con el normalizador online-softmax evolutivo que el kernel ya mantiene; durante el backward, reutiliza el normalizador finalizado que el forward guardó para reevaluar cada tile recomputado. En ambos pasos, los tiles cuya contribución cae por debajo de una tolerancia normalizada por longitud omiten su trabajo post-score. La tolerancia es la misma para todo: τ=1 en forward, backward, prefill y decoding, y en el límite τ→0 las condiciones de salto desaparecen y ambos pasos recuperan FullAttn.

Detrás hay 9 autores liderados por Jingze Shi, de la HKUST de Guangzhou con la Beijing Academy of Artificial Intelligence y la Université Paris Cité. El v1 entró en arXiv el 26/09/2026, y el código está publicado desde el primer día en HKUSTDial/flash-sparse-attention bajo BSD-3-Clause — no es un “código próximamente”.

Los números del paper, con fecha

Fuente: HTML v1 del paper, consultado el 29/09/2026. Cifras citadas verbatim; todas reportadas por los autores, sin auditoría externa todavía:

Claim verbatimQué significaContexto del experimento
”mean omitted mass of 0.0188% versus 0.0182%“Con exactamente el mismo trabajo post-score total, MALA se acerca a un oráculo de referencia por instancia (que omite un 0.0182% de masa). El techo de calidad de esta idea está muy cerca.Estudio matched-work a 8K
”89.67% accuracy at 8K compared with 89.97% for FullAttn”En recall asociativo controlado, MALA casi empata con FullAttn; bajo el mismo presupuesto de slots, DSA se queda en 52.61%, MoBA en 47.25% y NSA en 22.61%.Recall asociativo, 1K–8K
”reduces forward and backward latency during training by 2.2x and 3.0x and decoding latency during inference by 1.6x”Latencia del operador de atención completo, incluyendo el descubrimiento de scores.Benchmark del operador a 128K, 8×H100 con TP=8
”reduces total training FLOPs by 2.5% during 4K pre-training and 23.1% during 32K long-context training”Ahorro de FLOPs del modelo entero a 14B, con configuraciones matched de 4K y 32K en cinco tamaños de 0.6B a 14B.Scaling laws, 128 GPUs H100, Megatron-LM
”achieve comparable knowledge, reasoning, and long-context retrieval scores to FullAttn”Los 14B del scaling study y unos 32B de continued training (partiendo del mismo checkpoint de Qwen3-32B) quedan comparables, no superiores. Por ejemplo en RULER 128K a 14B: 65.75 vs 65.84 de FullAttn.Eval con LM Evaluation Harness

Leído en frío: el estudio matched-work es el diseño más honesto de la pila — aísla el beneficio de asignar por distribución frente a un oráculo que sabe la masa real de cada instancia, con el mismo presupuesto de trabajo. MALA se queda a un pelo de ese oráculo (0.0188% de masa omitida frente a 0.0182%). El problema es que, tres días después del v1, todo lo que tenemos son las palabras del autor.

Cómo funciona: presupuestos no, distribuciones

Los enfoques sparse habituales fijan un presupuesto — top-k slots, ventanas, un router aprendido — y esperan que el modelo se apañe. MALA invierte la lógica: no hay presupuesto, hay una tolerancia. El kernel online-softmax ya mantiene por cada fila el máximo acumulado y la suma de normalización; MALA usa esa suma para acotar cuánto puede aportar un tile al softmax normalizado. Si la cota cae bajo τ, el tile se descarta y te ahorras su exponenciación, su carga de V y su acumulación. Una fila de atención nítida rechaza muchos tiles; una fila difusa retiene más. El trabajo retenido lo decide la distribución de tu modelo, capa a capa, head a head, query a query.

La fidelidad reportada es notable: en su suite operador a lo largo de 1K–32K, la masa de probabilidad omitida media es como máximo 0.0062% (P95: 0.032%) y el error L2 relativo de salida como máximo 0.021%, con errores de gradiente por debajo del 0.4%. Son cifras del paper, pero son las que importarían verificar primero en una reproducción independiente.

Y ahora la parte que muchos titulares se van a saltar. El propio paper lo dice claro: MALA calcula los scores QK de todos los tiles causales legales, así que su coste sigue siendo cuadrático en longitud de secuencia; el ahorro son “reducciones de factor constante dependientes de los datos” en aritmética post-score y tráfico de memoria. Ojo al detalle del benchmark de latencia: mide el operador completo incluyendo score discovery — es decir, los 2.2×/3.0× ya pagan el camino cuadrático íntegro y aun así ganan eso. Es un resultado serio. Solo que no es el que dice “subcuadrático”.

Qué NO es MALA (y dónde se va a exagerar)

  • No es contexto subcuadrático end-to-end. La cobertura QK sigue completa; lo que se salta es el trabajo post-score de baja contribución. Si necesitas subcuadraticidad real, esto no lo es.
  • Las latencias son del operador, no de tu run de entrenamiento. En un paso de entrenamiento real, la atención es una fracción del tiempo: MLPs, comunicaciones de tensor paralelismo, optimizer. La cifra honesta para tu presupuesto de cómputo es la del modelo completo: −2.5% de FLOPs a 4K, −23.1% a 32K, en 14B.
  • La ganancia depende de tu distribución de atención. Si tu modelo atiende de forma muy difusa, retiene más tiles y ahorra menos. El paper entrena modelos propios de 0.6B a 14B; no hay garantía de traslación a tu fine-tuning.
  • Es un preprint sin peer review, con tres días de vida y cero reproducciones independientes. El estudio de 32B parte de un checkpoint de Qwen3-32B con continued training de 64B tokens a 32K; los resultados son razonables, pero son suyos.

En la práctica: repo, requisitos y esa tolerancia como parámetro

El código está en HKUSTDial/flash-sparse-attention, con documentación en hkustdial.github.io/flash-sparse-attention. Lo verificado en el README a 29/09/2026:

  • Instalación: pip install flash-sparse-attn (paquete publicado en PyPI) o build local con pip install .
  • Requisitos: Linux (Ubuntu 22.04 o posterior), Python ≥3.9, PyTorch ≥2.5.1, Triton ≥3.6.0. Soporta backends GPU/XPU/NPU/PPU. En macOS no hay nada declarado — no cuentes con entrenar esto en tu portátil.
  • La tolerancia del paper está expuesta tal cual: softmax_threshold=1.0 en flash_sparse_attn_func y en flash_sparse_attn_with_kvcache_func. El τ=1 del paper es un parámetro que puedes tocar, no un valor enterrado en C++.
  • Backend CuTe es el más rápido hoy; el backend Gluon está en progreso. Los scripts de benchmark del operador están en tests/ (benchmark_forward.py, backward y decode).
  • Baselines que compara el paper: FlashAttention, MoBA (de Moonshot AI) y DSA (el indexer de DeepSeek, sobre DeepGEMM).

La parte que el pip install no arregla: el paper entrena con Megatron-LM sobre 128 H100. Integrar un kernel de atención nuevo en tu stack de entrenamiento — checkpointing, TP, reproducibilidad — es tu trabajo, no un flag.

Cuándo gana y cuándo pierde

Dónde sí. Entrenamiento o continued training long-context: a 32K, el ahorro de FLOPs del modelo entero (−23.1% a 14B) es de los que cambian decisiones de presupuesto. Decoding largo en inferencia servida con tus propios kernels (1.6× sobre el operador a 128K). Y cualquier carga donde la atención sea muy concentrada — el mecanismo está diseñado exactamente para eso.

Dónde no. Pretraining o fine-tuning a contexto corto: −2.5% de FLOPs a 4K no justifica integrar un kernel nuevo y casarte con un preprint de tres días. Tampoco si tu entorno no es Linux con GPUs modernas, o si necesitas resultados respaldados por terceros antes de tocar el training run que paga tu factura de nube.

Y una aclaración de contexto: nada de esto cambia lo que pagas por una API. Esto es palanca para quien entrena o sirve modelos con kernels propios. Si tu problema es la factura de tokens de agentes o de API, hay palancas más directas: cómo reducir los costes de tokens un 30%, recortar costes de coding agents un 50% sin perder calidad o cómo elegir el LLM correcto para cada tarea.

La otra mitad: CoWA, el gemelo que reparte la cobertura

Dieciséis minutos antes del v1 de MALA —15:09 UTC frente a 15:24 UTC del 26/09—, los mismos 9 autores subieron a arXiv CoWA — “CoWindow Attention: Full Causal Coverage Is a Collective Property”. El diagnóstico es distinto: FullAttn expone la historia causal completa a cada head, con redundancia brutal entre heads. CoWA reparte el acceso: todos los heads comparten ventanas near-diagonales y prefix-sink, y ventanas long-range complementarias particionan el resto de la historia. La unión da cobertura causal completa aunque cada head asista de forma sparse a los tokens lejanos. El patrón lo define la posición, sin router aprendido, y es el mismo en entrenamiento e inferencia — alineado de paso con el tensor paralelismo de KV heads.

Sus números, con el mismo disclaimer de siempre (reportados por los autores): en recall asociativo window-matched a 8K, 89.73% frente a 89.97% de FullAttn; en el benchmark de operador a 128K con TP, 7.4× en forward y 8.6× en backward de entrenamiento, 3.0× en decoding, y memoria de decoding por rank 7.6× menor. La misma tanda de scaling laws de 0.6B a 14B, con perplexity que sigue de cerca a FullAttn.

La diferencia en una línea: MALA decide en runtime mirando la distribución real de atención, por query; CoWA fija el reparto por posición antes de ver ni un token. Uno salta cómputo post-score de baja contribución; el otro recorta qué historia ve cada head. Ojo con el repo: implementa MALA — CoWA, hoy por hoy, no aparece en él —, y el paper de MALA usa a MoBA y DSA — seleccionadores de tokens, primos lejanos de CoWA — como baselines. No es casualidad: es el mismo grupo de investigación cubriendo los dos extremos del mismo problema. Si esto acaba en producción, probablemente sea con una de las dos ideas, o con una híbrida.

Qué seguir

Tres cosas que vigilaría: que alguien replique el estudio matched-work (es el más barato de replicar y el que soporta todo lo demás), que algún stack de entrenamiento abierto integre el kernel (cuando aparezca en un framework que no sea de los autores, el riesgo de adopción baja un orden de magnitud), y qué hace falta realmente para que softmax_threshold sobreviva en un run de fine-tuning tuyo — la tolerancia es un parámetro, pero la distribución de atención de tu dominio manda, y eso solo se mide entrenando.

Si estás decidiendo el stack de un run long-context de aquí a un trimestre, MALA es hoy la opción con mejor relación evidencia/coste-entre-integrarla — con la evidencia todavía toda del mismo sitio. Yo esperaría la primera reproducción antes de ponerla en el run que importa, y mientras tanto, el repo tiene los benchmarks listos para hacerla yo mismo si el hardware aparece.

Fuentes

Cargando comentarios...