¡Hola Devs! Vamos a hablar de servir MoE en serio
Si alguna vez intentaste servir un modelo Mixture-of-Experts (MoE) de frontera en aceleradores especializados, ya conoces el muro: no puedes resolverlo a punta de prueba y error. Qwen 3.5-397B-A17B es el caso perfecto. Tiene 397B parámetros totales (unos 400 GB de pesos), pero solo activa 17B por token — una tasa de activación del 4.3%. Esa esparsidad es el punto: obtienes la inteligencia de un modelo de 400B al costo de servir uno de ~20B.
Pero esa promesa solo se cumple si tu ingeniería de sistemas está impecable. El sharding tensor-parallel ingenuo se rompe de inmediato cuando el modelo tiene exactamente 2 KV heads en sus capas GQA y 512 experts en el stack MoE. No puedes partir 2 heads en 8 devices. Replicarlas desperdicia HBM. Y en un chip con 192 GB de HBM (contra 288 GB en el Blackwell GB300 — ~50% menos), desperdiciar HBM es fatal.
La solución no es un kernel mágico. Es una metodología modular y agnóstica al modelo — un playbook de bloques reutilizables (Batched RPA, Grouped GEMMs, unpermutation en SparseCore) que portas a nuevas arquitecturas con casi cero fricción. Vamos a ver cómo se aplicó al Qwen 3.5 en Ironwood (TPU v7x), logrando ~3.1x en decode y ~4.7x en prefill en el tier de 512 de concurrencia.
El reporte técnico completo está aquí.

Las restricciones arquitectónicas que rompen el sharding ingenuo
El layout híbrido de Qwen 3.5 no es un stack Transformer uniforme. Son 60 capas organizadas en 15 bloques repetidos, cada uno con una proporción 3:1 de atención lineal Gated DeltaNet (GDN) a atención GQA estándar:
- GDN linear attention (75% de las capas): 64 V-heads, 16 QK-heads, head dim 128. Mantiene una matriz de estado recurrente de tamaño constante por head, actualizada con la regla delta. Escala O(S) en vez de O(S²).
- GQA (25% de las capas): 32 query heads, exactamente 2 KV heads, head dim 256, RoPE dim 64. Comprime el KV cache, pero impone restricciones brutales de sharding.
- MoE FFN: 512 experts pequeños, dim intermedia 1024, ruteo
top_k=10más 1 shared expert siempre activo.
Por qué TP=8 falla
Con solo 2 KV heads, un tensor-parallel de tamaño 8 fuerza sharding fraccional (2/8 = 0.25 heads por device) — físicamente imposible. Replicar las heads en 8 cores duplica el KV cache en cada device, limitando la concurrencia real a ~200 en vez de los 512 planeados.
La solución: DP de atención + EP de experts
El equipo co-diseñó un esquema de 8-way Attention Data Parallelism (DP=8) + 8-way Expert Parallelism (EP=8):
# Topología conceptual de sharding para Qwen 3.5 en mesh TPU de 8 devices
# Comentario: atención en DP, MoE en EP sobre el mesh de 8 chips
sharding = {
"gdn_layers": "replicated", # Las 2 KV heads completas por device
"gqa_layers": "replicated", # Consistencia local del KV cache
"moe_experts": "expert_parallel", # 512 experts / 8 devices = 64 por chip
}
# Ruteo cross-device: All-Gather de metadatos, Reduce-Scatter de salidas
# Comentario: metadatos de ruteo por All-Gather, salida de experts por Reduce-Scatter
Los pesos de atención se replican en los 8 devices (cada core ve las 2 KV heads completas), eliminando comunicación intra-atención. Los 512 experts ruteados se distribuyen equitativamente (64 por device), evitando duplicar los 400 GB de pesos.
Tres colectivas se vuelven dos
En EP ingenuo, preparar el MoE local requería tres All-Gathers: índices de expert, pesos topk y hidden states. Como índices (int) y pesos (float) comparten shape idéntico [1024, 10], se apilaron, bitcast y empaquetaron en un blob de 32-bit — un All-Gather en vez de dos. Latencia de metadatos de ruteo cortada a la mitad.

Kernels Pallas, co-diseño con SparseCore y el roofline
Con la topología correcta, las ganancias restantes vinieron de kernels hand-scheduled que se saltan el camino estándar de XLA.
Tres optimizaciones de alto impacto
| Optimización | Técnica | Impacto medido |
|---|---|---|
| RPA con indexación gruesa | KV page size 16 → 256 (--block-size=256) | Latencia de decode en C=512: 428µs → 283µs (33.8% más rápido) |
| Unpermutation en SparseCore | Offload de unpermutation + reducción local al SparseCore | Lecturas HBM 20→10, escrituras 15→5 |
| Fusión GDN + conv causal 1D | Sliding window en registros VPU | Eliminó 6 round-trips redundantes a HBM |
Fusión a nivel de registro en el camino GDN
La actualización recurrente de GDN se compilaba como op independiente, forzando las salidas de la convolución a ir y volver de HBM. El equipo fusionó la conv causal 1D (K=4) y la actualización de estado GDN en un solo bloque, cacheando estados históricos directo en los registros de la VPU. También migraron las variables de estado del SSM de Float32 a BFloat16, doblando el throughput vectorial de la VPU sin comprometer la convergencia numérica.
Validación por roofline
En concurrencia 64 con layout 8K/1K, el throughput empírico quedó muy cerca de los límites teóricos del roofline — o sea, los kernels están empujando el hardware cerca de sus límites físicos. Bajo alta concurrencia, las matrices de gating y ruteo son sensibles a errores de acumulación en baja precisión, así que una Capa de Verificación Numérica audita continuamente los bloques de escala FP8, confirmando cero desviación del camino de referencia en Float32.

Limitaciones, advertencias y siguientes pasos
Lo que este playbook no resuelve
- HBM es el techo duro. Los 192 GB por chip del TPU v7 son ~50% menores que los 288 GB del Blackwell GB300. Ninguna fusión de kernel recupera esa diferencia.
- Las novedades específicas del modelo todavía cuestan tiempo. El playbook reduce fricción, pero las capas GDN y el
top_k=10de Qwen 3.5 exigieron kernels Pallas a medida. - La corrección numérica debe re-verificarse por modelo. La acumulación FP8 es sensible al modelo. Lo que funciona aquí puede no transferirse limpio.
Roadmap restante
Dos pistas abiertas: (1) fusionar el kernel de selección top_k directo en la VPU, eliminando el cuello de botella de serialización TensorCore→VPU; y (2) reducir más el overhead de colectivas cross-device con pipelining a nivel de chunk.
Si vas a construir sobre esto
Empieza por la topología de sharding antes de tocar kernels. Acertar DP+EP vale más que cualquier micro-optimización. Después haz profiling agresivo — los cuellos de botella se esconden en colectivas, no solo en matmuls.