Base de Conocimiento ANE
EN ES

Capítulos · 02

Los cuatro principios para optimizar Transformers en el ANE

Parte de la ANE Knowledge Base. Fuente: Apple ML Research, "Deploying Transformers on the Apple Neural Engine" (junio de 2022), y su código de referencia en references/ml-ane-transformers.

El transformer estándar de PyTorch (nn.Linear sobre tensores (B, S, C), atención multi-head fusionada) se mapea mal sobre el ANE. El paquete ane_transformers de Apple reexpresa exactamente la misma matemática en una forma nativa del ANE. Resultado sobre DistilBERT de Hugging Face: hasta 10× menos latencia y 14× menos memoria pico frente a la línea base.

Los cuatro principios siguientes derivan directamente de las restricciones de hardware descritas en el doc 01.


Principio 1: Elegir el formato de datos correcto — (B, C, 1, S)#

Problema. Los transformers de PyTorch usan tensores 3D channels-last, (B, S, C) o (S, B, C). El ANE quiere 4D channels-first, y sus búferes requieren que el último eje sea contiguo y esté alineado a 64 bytes (sin empaquetar — un último eje pequeño recibe padding hasta 64 bytes).

Solución. Migrar todo a (B, C, 1, S) — Apple lo llama BC1S:

  • Batch, Channels (dimensión de embedding), una altura ficticia de 1, y la longitud de secuencia al final. S es el eje que crece, así que amortiza la alineación a 64 bytes.
  • Cambiar cada nn.Linear por un nn.Conv2d con tamaño de kernel 1. Una conv 1×1 sobre (B, C, 1, S) es matemáticamente idéntica a una capa lineal sobre (B, S, C), y las convoluciones son lo que el ANE ejecuta mejor.
# ane_transformers/reference/ffn.py — the FFN is just two 1x1 convs
self.layers = nn.ModuleList([
    nn.Conv2d(embed_dim, ffn_dim, 1),
    nn.ReLU(),
    nn.Dropout(dropout) if dropout > 0. else nn.Identity(),
    nn.Conv2d(ffn_dim, embed_dim, 1),
])

LayerNorm debe seguir la disposición. torch.nn.LayerNorm normaliza la última dimensión; en BC1S el embedding ahora vive en la dim 1. Apple incluye LayerNormANE (ane_transformers/reference/layer_norm.py), que normaliza sobre el eje de canales (dim 1) y está construido a partir de primitivas simples y amigables para el ANE:

# ane_transformers/reference/layer_norm.py (forward, simplified)
channels_mean = inputs.mean(dim=1, keepdims=True)
zero_mean = inputs - channels_mean
zero_mean_sq = zero_mean * zero_mean
denom = (zero_mean_sq.mean(dim=1, keepdims=True) + self.eps).rsqrt()
out = zero_mean * denom
if self.elementwise_affine:
    out = (out + self.bias.view(1, C, 1, 1)) * self.weight.view(1, C, 1, 1)

Dos trampas codificadas en el repositorio:

  1. Inversión del orden de scale/bias: LayerNormANE aplica (x + bias) * weight, mientras que nn.LayerNorm calcula x * weight + bias. Al restaurar un checkpoint preentrenado, el bias debe dividirse primero por el weight (ver correct_for_bias_scale_order_inversion en ane_transformers/huggingface/distilbert.py).
  2. clip_mag opcional hace clamping de las entradas antes de la normalización para reducir el riesgo de overflow en FP16.

La compatibilidad de pesos es mecánica. Los pesos de nn.Linear son (out, in); los pesos de nn.Conv2d 1×1 son (out, in, 1, 1). Un pre-hook de load_state_dict les aplica unsqueeze dos veces, de modo que los checkpoints preentrenados se cargan sin cambios:

# ane_transformers/huggingface/distilbert.py
def linear_to_conv2d_map(state_dict, ...):
    for k in state_dict:
        if is_linear_weight(k) and len(state_dict[k].shape) == 2:
            state_dict[k] = state_dict[k][:, :, None, None]

Principio 2: Fragmentar los tensores intermedios grandes#

Problema. La atención multi-head fusionada crea intermedios muy grandes (proyecciones QKV completas, la matriz de atención completa). Los tensores grandes se salen de la caché L2 del ANE y no pueden repartirse entre los cores del ANE.

Solución. Dividir Q, K, V en fragmentos por cabeza y computar una lista explícita de funciones de atención de una sola cabeza. "Smaller chunks increase the chance of L2 cache residency as well as increasing multicore utilization during compilation."

# ane_transformers/reference/multihead_attention.py (_attention_fn)
mh_q = q.split(self.d_qk // self.n_head, dim=1)   # n_head × (B, d/h, 1, tgt_len)
mh_k = k.transpose(1, 3).split(self.d_qk // self.n_head, dim=3)
mh_v = v.split(self.d_v // self.n_head, dim=1)

attn_weights = [torch.einsum('bchq,bkhc->bkhq', [qi, ki]) * self.q_normalize_fact
                for qi, ki in zip(mh_q, mh_k)]
attn_weights = [aw.softmax(dim=1) for aw in attn_weights]   # ← "split softmax"
attn = [torch.einsum('bkhq,bchk->bchq', wi, vi) for wi, vi in zip(mh_w, mh_v)]
attn = torch.cat(attn, dim=1)                     # (B, d_v, 1, tgt_len)

Notas:

  • El softmax por cabeza sobre la dim 1 (el eje de clave/secuencia de origen en esta disposición) es el split softmax que el trabajo posterior sobre vision transformers destaca como una de las mayores ganancias de latencia — el softmax es cuadrático en la longitud de tokens y, de lo contrario, serializa.
  • La fragmentación multiplica el número de ops (DistilBERT llega a 606 ops), lo que eleva el coste puntual de carga/compilación pero reduce la latencia en régimen estacionario. Carga los modelos de forma asíncrona.

Principio 3: Minimizar las copias de memoria#

Problema. En el ANE, "reshape and transpose operations are likely to trigger memory copies unless specifically handled."

Solución. La atención de referencia evita todos los reshapes e incurre en exactamente un transpose — sobre el tensor de claves, justo antes del matmul QK (k.transpose(1, 3) arriba). Todo lo demás se expresa con fórmulas de einsum cuyas disposiciones de operandos se mapean directamente sobre los batched matmuls del hardware, sin transposes ni reshapes intermedios:

Paso einsum Formas
Pesos de atención bchq,bkhc->bkhq q (B, C/h, 1, T) × kᵀ (B, S, 1, C/h)(B, S, 1, T)
Valores ponderados bkhq,bchk->bchq w (B, S, 1, T) × v (B, C/h, 1, S)(B, C/h, 1, T)

Nótese también lo que no ocurre: no hay ningún barajado con view/permute para formar un tensor de atención (B·h, T, S) como en las implementaciones estándar — las cabezas permanecen como una lista de Python de pequeños tensores 4D desde el split hasta el concat.

Enmascaramiento en esta disposición. Las máscaras son floats aditivos aplicados antes del softmax:

  • qk_mask (como attn_mask): forma (B, S, 1, T) — por ejemplo, máscaras causales.
  • k_mask (como key_padding_mask): forma (B, S, 1, 1) — por ejemplo, tokens de padding.
  • Usar -1e4 para bloquear la atención (seguro en FP16, componible mediante suma), 0 para conservarla.

Principio 4: Manejar la limitación por ancho de banda#

Problema. "Many Transformer configurations become bandwidth-bound on the ANE when the sequence length is relatively short": los pesos se transmiten desde memoria solo para tocar unas pocas activaciones antes de que se traiga el siguiente tensor de pesos. Evidencia: la latencia de DistilBERT es ~constante para longitudes de secuencia 32/64/128 en batch 1, a pesar de que el cómputo se cuadruplica.

Soluciones.

  1. Aumentar el tamaño del batch en cargas de trabajo de inferencia por lotes — más aritmética útil por cada fetch de pesos (Apple reporta ganancias de hasta 10×/14× en el terreno de seq 512 / batch 8, frente a 2.84× en seq 128 / batch 1).
  2. Reducir los pesos con cuantización o poda — Apple señala que el rendimiento pico del ANE estaba "far from saturated" para DistilBERT, así que ahí quedan ganancias adicionales disponibles.

Resumen: transformer de línea base vs. optimizado para el ANE#

Aspecto PyTorch de línea base Optimizado para el ANE (ane_transformers)
Disposición de tensores (B, S, C) 3D channels-last (B, C, 1, S) 4D channels-first (BC1S)
Proyecciones nn.Linear nn.Conv2d kernel 1
LayerNorm nn.LayerNorm (última dim) LayerNormANE (dim de canales), orden (x+b)*w
Atención multi-head Grandes matmuls fusionados + reshape/permute Lista por cabeza vía split, matmuls con einsum
Softmax Un gran softmax Split softmax por cabeza (dim=1)
Transposes Muchos implícitos Exactamente uno (tensor de claves)
Máscaras bool/-inf float aditivo -1e4
eps de LayerNorm 1e-12 (DistilBERT) 1e-7 (seguro en FP16)

El encoder/decoder genérico completo construido a partir de estos bloques vive en references/ml-ane-transformers/ane_transformers/reference/ (transformer.py, encoder.py, decoder.py, multihead_attention.py, ffn.py, layer_norm.py) — refleja la configuración base original de "Attention Is All You Need" pero en forma nativa del ANE.

Siguiente: 03 — Caso de estudio: DistilBERT de Hugging Face.

Generado desde el markdown de la base de conocimiento — cada afirmación traza a una fuente citada.