Base de Conocimiento ANE
EN ES

Capítulos · 03

Caso de estudio: Optimizar DistilBERT de Hugging Face para el ANE

Parte de la ANE Knowledge Base. Fuente: Apple ML Research, "Deploying Transformers on the Apple Neural Engine", código en references/ml-ane-transformers/ane_transformers/huggingface/distilbert.py.

Este caso de estudio muestra cómo los cuatro principios se adaptan (retrofit) sobre un modelo de terceros existentedistilbert-base-uncased-finetuned-sst-2-englishsin reentrenar y manteniendo la compatibilidad de checkpoint.

1. Estrategia: subclasificar e intercambiar submódulos#

En lugar de reescribir DistilBERT, Apple subclasifica cada módulo de Hugging Face e intercambia solo los submódulos problemáticos mediante setattr, conservando el resto de la lógica upstream:

class TransformerBlock(modeling_distilbert.TransformerBlock):
    def __init__(self, config):
        super().__init__(config)
        setattr(self, 'attention', MultiHeadSelfAttention(config))     # ANE version
        setattr(self, 'sa_layer_norm', LayerNormANE(config.dim, eps=EPS))
        setattr(self, 'ffn', FFN(config))                              # 1x1 convs
        setattr(self, 'output_layer_norm', LayerNormANE(config.dim, eps=EPS))

Qué se reemplaza, según los cuatro principios:

Módulo de HF Reemplazo Principio
q_lin, k_lin, v_lin, out_lin (nn.Linear) nn.Conv2d(dim, dim, 1) P1 (disposición BC1S)
FFN.lin1/lin2 (nn.Linear) nn.Conv2d 1×1 P1
nn.LayerNorm LayerNormANE P1 (normalizar dim 1)
Forward de MHA fusionado split por cabeza + einsum + split softmax P2, P3
Cabezas de tarea (pre_classifier, classifier, vocab_projector, qa_outputs) nn.Conv2d 1×1 P1

Se cubren las seis variantes de tarea: DistilBertModel, ForMaskedLM, ForSequenceClassification, ForQuestionAnswering, ForTokenClassification, ForMultipleChoice.

2. Compatibilidad de checkpoint sin reentrenar#

Dos pre-hooks de load_state_dict hacen que los pesos preentrenados se carguen en la nueva arquitectura sin tocarlos:

  1. linear_to_conv2d_map — aplica unsqueeze a cada peso Linear (out, in) relevante para convertirlo en (out, in, 1, 1) para Conv2d.
  2. correct_for_bias_scale_order_inversionLayerNormANE aplica (x + bias) * weight mientras que nn.LayerNorm aplica x * weight + bias, así que el bias almacenado se reescala: bias = bias / weight.

3. Higiene de FP16#

  • Epsilon: el eps original de LayerNorm de DistilBERT de 1e-12 "is not friendly with the float16 precision that ANE uses by default" → EPS = 1e-7.
  • Máscaras: la máscara de atención de HF (bool o int64 (B, S)) se convierte en una máscara float aditiva (B, S, 1, 1) con -1e4 para las posiciones enmascaradas:
if mask.dtype == torch.bool:
    mask = mask.logical_not().float() * -1e4
elif mask.dtype == torch.int64:
    mask = (1 - mask).float() * -1e4
  • Solo inferencia: las clases lanzan una excepción con self.training o labels is not None — este port es para inferencia en dispositivo; entrena con la implementación original de HF.
  • return_dict debe ser False: "coremltools does not support dict outputs."

4. Detalles de disposición que conviene notar#

  • Los estados ocultos (hidden states) fluyen por toda la red como (B, dim, 1, seq_len) (BC1S).
  • El pooling del token CLS se convierte en un slice sobre el último eje: hidden_state[:, :, :, 0:1](B, dim, 1, 1) — sin necesidad de transpose.
  • Los logits de tarea se recuperan justo al final con ops finales baratas (squeeze), después de que todo el cómputo pesado se haya hecho en forma amigable para el ANE.

5. Resultados#

Medido sobre el DistilBERT de SST-2 (iPhone 13, iOS 16):

Métrica Valor
Latencia (seq 128, batch 1) 3.47 ms a 0.454 W (también 9.44 ms a 0.072 W en el punto de bajo consumo)
Aceleración en seq 128 / batch 1 (informe de Xcode) 2.84× vs línea base
Aceleración en cargas de trabajo mayores (por ejemplo, seq 512, batch 8) hasta 10× en latencia, 14× en memoria pico
Comparación de referencia del lado del servidor AWS c6i/inf1: ~5–6 ms en seq 128
Dispositivos validados iPhone 12 (iOS 15/16), iPhone 13 (iOS 16), Mac M1 (macOS 13)

Notas operativas del tutorial:

  • El modelo optimizado tiene 606 ops; el número de ops (por la fragmentación) aumenta el tiempo de carga/compilación — un coste puntual, ocúltalo con carga asíncrona.
  • 4 de 606 ops se ejecutan en CPU: las ops de búsqueda de embeddings, que simplemente son más eficientes en CPU para esta configuración. Unas pocas ops en CPU ≠ fallo.
  • La latencia es ~plana a lo largo de longitudes de secuencia 32/64/128 en batch 1 → el modelo está limitado por ancho de banda ahí (Principio 4); queda margen de cuantización/poda.

6. La receta de despliegue (abreviada)#

Flujo de trabajo completo con explicación en el doc 05; la esencia del README del repositorio:

baseline_model = transformers.AutoModelForSequenceClassification.from_pretrained(
    "distilbert-base-uncased-finetuned-sst-2-english",
    return_dict=False, torchscript=True).eval()

optimized_model = ane_distilbert.DistilBertForSequenceClassification(
    baseline_model.config).eval()
optimized_model.load_state_dict(baseline_model.state_dict())  # hooks do the mapping

tokenized = tokenizer(["..."], return_tensors="pt", max_length=128, padding="max_length")
traced = torch.jit.trace(optimized_model,
                         (tokenized["input_ids"], tokenized["attention_mask"]))

mlpackage = ct.convert(traced, convert_to="mlprogram",
    inputs=[ct.TensorType(f"input_{name}", shape=t.shape, dtype=np.int32)
            for name, t in tokenized.items()],
    compute_units=ct.ComputeUnit.ALL)
mlpackage.save("distilbert_seqLen128_batchSize1.mlpackage")

7. Lecciones transferibles#

  1. Rara vez necesitas un modelo nuevo — una reexpresión que preserva la arquitectura y es matemáticamente equivalente, más hooks de state-dict, convierte los checkpoints existentes.
  2. Subclasificar + setattr es un patrón limpio para portar cualquier familia de modelos de HF.
  3. Corrige las constantes de precisión (eps, valores de máscara) al mismo tiempo que corriges la disposición — de lo contrario, la rotura de FP16 es silenciosa.
  4. Juzga el éxito con los informes de rendimiento de Xcode (despacho por op), no solo con "se ejecuta".
  5. Espera algunas ops (embeddings, squeezes finales) en CPU; optimiza el tronco (trunk) del transformer.

Siguiente: 04 — Vision Transformers en el ANE.

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