Base de Conocimiento ANE
EN ES

Capítulos · 06

Caso de estudio: WhisperKit — ASR en producción sobre la ANE (Argmax)

Parte de la ANE Knowledge Base. Fuentes: argmaxinc/whisperkittools (clonado en references/whisperkittools) y su dependencia central argmaxtools 0.1.23 (snapshot de PyPI incluido en references/argmaxtools-0.1.23-pypi-snapshot; el repositorio de GitHub ya no es público).

WhisperKit es el stack de reconocimiento de voz en dispositivo de Argmax. whisperkittools es la parte de Python: convierte los checkpoints de OpenAI Whisper (Hugging Face) en modelos Core ML optimizados para la ANE que consume el runtime de WhisperKit en Swift, aplica compresión y hace benchmarking de los resultados (publicados en hf.co/argmaxinc/whisperkit-coreml). Es el mejor ejemplo público de los cuatro principios de Apple aplicados a escala de producción por un tercero — y de los paradigmas que hubo que inventar más allá de ellos: decodificación autorregresiva con KV cache, implementaciones de atención intercambiables, context prefill, verificación programática del despacho a la ANE y compresión mixed-bit (doc 07).

1. Arquitectura: un modelo se convierte en cuatro componentes .mlmodelc#

Un despliegue de Whisper se descompone en modelos Core ML independientes, cada uno trazado/convertido por separado (véase scripts/generate_model.py::rearrange_quantized_variants):

Componente Fuente Rol
MelSpectrogram.mlmodelc whisperkit/audio_encoder.py::WhisperMelSpectrogram audio → características log-mel (torch.stft + banco de filtros mel como un nn.Module exportable)
AudioEncoder.mlmodelc whisperkit/audio_encoder.py::WhisperAudioEncoder características mel → embeddings del encoder (se ejecuta una vez por ventana de 30 s)
TextDecoder.mlmodelc whisperkit/text_decoder.py::WhisperTextDecoder un token por llamada, decodificación autorregresiva con KV cache
TextDecoderContextPrefill.mlmodelc (opcional) text_decoder.py::WhisperTextDecoderContextPrefill tabla de búsqueda de KV cache para prefijos de tarea/idioma

¿Por qué descomponer? Cada componente tiene una cadencia de ejecución diferente (una vez por ventana vs una vez por token), una implementación óptima de SDPA distinta y una sensibilidad a la compresión distinta. La descomposición también permite al runtime planificarlos de forma independiente.

2. Los principios de Apple, al pie de la letra — vía argmaxtools.nn#

argmaxtools es la biblioteca generalizada de transformers para la ANE de Argmax (su equivalente de ane_transformers, ampliada para las necesidades modernas de la era de los LLM). El linaje es explícito — nn.py cita directamente el repositorio de Apple encima de su LayerNorm.

Principio de la KB Dónde aparece en argmaxtools/whisperkittools
P1 — layout BC1S, nn.Conv2d 1×1 Attention.__init__: q/k/v/o_proj = nn.Conv2d(embed_dim, ..., 1); FFN igual (argmaxtools/nn.py). Todas las shapes de I/O son (batch, embed_dim, 1, seq_len)
P1 — LayerNorm en la dimensión de canal argmaxtools.nn.LayerNorm normaliza dim=1, con un clamp clip_mag opcional — copiado de la referencia de Apple (cita en nn.py:498-499)
P2 — chunking / split softmax _sdpa.SplitHeadsQ: splits por cabeza + softmax por cabeza dim=1 más chunking de la secuencia de consulta (chunk_size=256) — una extensión del P2 de Apple
P3 — copias mínimas, einsum El mismo par de einsum (bchq,bkhc->bkhq / bkhq,bchk->bchq) en SplitHeadsQ; un único transpose de key (patrón "transposeThenSplit is more efficient")
P4 — límite por ancho de banda Palettization de pesos (doc 07); la decodificación de un solo token está inherentemente limitada por ancho de banda, de ahí que la compresión sea también una optimización de latencia
Higiene FP16 Las máscaras usan -1e4 (máscara causal, decoder_key_padding_mask); denominadores del softmax perezoso (lazy) con clamp (.clamp(min=1e-6))
Hooks de compatibilidad de checkpoints argmaxtools/utils.py::linear_to_conv2d_map_{attention,ffn,mlp} — generalizados con tablas de alias de nombres (q_proj/query_proj/linear_q/…) para que un solo hook se adapte a muchas convenciones de nomenclatura upstream

Nueva micro-optimización que no está en el trabajo de Apple: el sesgo de la proyección de key se elimina por completok_proj = nn.Conv2d(..., bias=False) con el comentario "key bias is redundant due to softmax invariance" (nn.py:88-89); el hook del state-dict descarta el k_proj.bias entrante. (Añadir una constante a cada columna de logits de atención no cambia el softmax — el softmax sobre keys se aplica por query contra todas las keys desplazadas por igual... con precisión: un sesgo en K aporta q·b, constante entre keys para una query dada, y el softmax es invariante a desplazamientos.)

También a diferencia del port de Apple: argmaxtools.LayerNorm mantiene el orden estándar w*x + b, de modo que no hace falta un hook de corrección de bias/scale.

3. Nuevo paradigma n.º 1: SDPA como estrategia intercambiable por componente#

El artículo de Apple de 2022 prescribía una forma de atención. La idea clave de Argmax es que la implementación óptima de SDPA depende de la carga de trabajo, por lo que Attention recibe un sdpa_implementation configurable en tiempo de ejecución (argmaxtools/_sdpa.py) con una interfaz común:

Implementación Técnica Mejor para
Cat (por defecto) Sin splits: view 4D a (B, h, c, x), un einsum bhcq,bhck->bhqk, un softmax dim=3. Las views son gratuitas; número mínimo de ops Secuencias de consulta cortas — usado en TextDecoder (q_seq_len = 1)
SplitHeadsQ Split por cabeza al estilo de Apple + split softmax (dim=1), más chunking de la secuencia de consulta cada 256 posiciones Secuencias largas — usado en AudioEncoder (1500 tokens)
SplitKV Atención eficiente en memoria con softmax perezoso (online) (Rabe & Staats 2021): divide la secuencia de key/value en chunks, mantiene el máximo/suma corrientes, y los fusiona Secuencias de KV muy largas donde la matriz de atención completa no cabe en caché
SharedSplitKVCached Softmax perezoso de dos chunks sobre (key_cache, current_key) — atiende a una caché compartida de batch-1 desde una query de batch-N Decodificación en batch con KV cache (p. ej., patrones de token-tree / especulativos)

Los valores por defecto de producción en scripts/generate_model.py son la conclusión:

--audio-encoder-sdpa-implementation  default: SplitHeadsQ   # long seq → split softmax wins
--text-decoder-sdpa-implementation   default: Cat           # seq_len=1 → compact graph wins

Lección para la KB: la receta de split-softmax de Apple no es universalmente óptima — para la decodificación de un solo token, la sobrecarga del chunking supera los beneficios de residencia en caché. Elige según la carga de trabajo; mantén estable la interfaz del modelo para que la elección sea un cambio de una línea.

Attention también generaliza más allá de Whisper: AttentionHeadType soporta MHA / GQA / MQA (tiling de KV-heads vía repeat_kv), además de RoPE opcional (_positional_encoding.py) y RMSNorm — es decir, la misma biblioteca está lista para Llama.

4. Nuevo paradigma n.º 2: decodificación autorregresiva con KV cache en la ANE#

La ANE necesita shapes estáticos, por lo que la clásica "KV cache creciente" debe rediseñarse. WhisperTextDecoder (con AttentionType.KVCachedSelfAttention) muestra todo el patrón:

La caché como I/O explícito del modelo, en BC1S, fusionada entre capas#

key_cache / value_cache inputs: (batch, embed_dim * n_layers, 1, max_seq_len)

Las cachés de todas las capas viajan como un solo tensor concatenado en el eje de canal, con split(d_model, dim=1) dentro del grafo (text_decoder.py:207). Las salidas devuelven solo las actualizaciones de un token (key_cache_updates, value_cache_updates, shape (..., 1)); el runtime de Swift las escribe en su búfer persistente. Un par de I/O en lugar de 2×n_layers tensores.

Caché de longitud fija + máscaras en lugar de shapes dinámicas#

  • La caché siempre tiene longitud max_seq_len (p. ej., 448). Los slots no usados se desactivan con decoder_key_padding_mask (-1e4 aditivo, la convención segura para FP16 de la KB).
  • kv_cache_update_mask es un vector one-hot que marca la posición del token actual; dentro de la atención, la actualización de la caché es una mezcla enmascarada (masked blend) — ops puramente elementwise, aptas para la ANE:
# argmaxtools/nn.py::_finalize_kv
key_cache   = key_cache   * (1. - kv_cache_update_mask) + current_key   * kv_cache_update_mask
value_cache = value_cache * (1. - kv_cache_update_mask) + current_value * kv_cache_update_mask
  • cache_length (una entrada int) selecciona el position embedding para el paso actual: embed_positions(cache_length).

Caché de cross-attention: computar una vez, eliminar las proyecciones#

El decoder de Whisper hace cross-attention sobre una salida del encoder fija, por lo que K/V se computan una vez por ventana de audio. StatefulKVCachedEncoderDecoderCrossAttention literalmente hace delattr(self, "k_proj")/"v_proj" — esas proyecciones pertenecen a la pasada del lado del encoder, y el decoder solo consume tensores cacheados (nn.py:452-463).

Variante stateful: Core ML MLState (iOS 18 / macOS 15)#

StatefulKVCachedAttention registra las cachés como register_buffers; el helper de conversión las mapea a entradas ct.StateType (argmaxtools/test_utils.py::_create_coreml_model), y la inferencia usa model.make_state() + predict(..., state=...). La caché entonces vive dentro del modelo Core ML — sin tráfico de I/O de caché por token en absoluto. Condicionado al target de despliegue: los states requieren iOS18/macOS15+.

5. Nuevo paradigma n.º 3: context prefill como modelo de tabla de búsqueda#

Cada ventana de transcripción de Whisper empieza con el mismo patrón de prefijo de 3 tokens: <|startoftranscript|> <|language|> <|task|>. En lugar de ejecutar el decoder 3 veces, WhisperTextDecoderContextPrefill:

  1. Enumera todos los prefijos válidos (idioma × tarea) (~99 idiomas × 2 tareas),
  2. Precomputa las KV caches del decoder para cada prefijo y las almacena aplanadas en dos tablas de búsqueda nn.Embedding,
  3. Se exporta como un modelo Core ML diminuto: (task, language) → (key_cache_prefill, value_cache_prefill).

La observación sutil (documentada en el código, text_decoder.py:342-345): las cachés del prefijo deberían depender de la salida del encoder, pero debido al enmascaramiento causal durante el entrenamiento, los embeddings KV del decoder para los tokens forzados del prefijo no están correlacionados con el audio — así que pueden estimarse una vez con una media por batch sobre salidas aleatorias del encoder. En tiempo de ejecución esto ahorra 3 de ~N pasadas forward del decoder por ventana y empieza a decodificar en cache_length=3.

6. Marcas de tiempo por palabra vía pesos de atención de las alignment heads#

Para las marcas de tiempo a nivel de token, configure_for_token_timestamps() activa _return_w en la cross-attention de alignment heads específicas (del generation_config del modelo), recoge sus pesos de atención con un register_forward_hook y devuelve su media como una 4.ª salida del modelo, alignment_heads_weights (text_decoder.py:143-170). La alineación DTW ocurre después en el runtime de Swift. Esto refleja el propio patrón de Apple de recoger pesos de atención mediante hooks en lugar de canalizarlos a través de valores de retorno.

7. Adaptaciones a la ANE específicas de Whisper#

  • Stem Conv1d → Conv2d: las dos capas Conv1d de Whisper se ejecutan como F.conv2d con los pesos Conv1d "unsqueezed" (weight[:, :, None, :]) para que todo el grafo permanezca 4D/BC1S (audio_encoder.py::pre_transformer_proj).
  • Espectrograma mel como modelo: torch.stft + ventana de Hann + banco de filtros mel envueltos en un nn.Module con register_buffers → su propio .mlmodelc; mantiene el DSP fuera de la ruta de código de CPU de la app.
  • Embeddings atados para los logits: la proyección final reutiliza embed_tokens.weight vía F.linear — sin una matriz de proyección de vocabulario separada que almacenar o palettizar.

8. Nuevo paradigma n.º 4: la verificación como CI, no como paso manual en Xcode#

El doc 05 de la KB verifica el despacho con la pestaña Performance de Xcode. argmaxtools lo automatiza todo en mixins de unittest (argmaxtools/test_utils.py):

  • Corrección: el PSNR entre las salidas de PyTorch y de Core ML debe superar los 35 dB (compute_psnr, TEST_PSNR_THR).
  • Speedup: la latencia mediana en CPU_AND_NE vs CPU_ONLY (el mismo .mlpackage, recargado con una compute unit diferente) debe superar un umbral; además de un contador de FLOP para TFlop/s.
  • La compresión no debe regresar: la variante palettizada a 1 bit debe conservar ≥ 0.95× de velocidad.
  • Plan de cómputo programático (_print_compute_plan, coremltools ≥ 8.1): carga el .mlmodelc compilado con ct.models.compute_plan.MLComputePlan y registra, por op, el dispositivo de despacho, los dispositivos soportados y la cuota de coste estimada — resumidos como ANE support coverage % y ANE dispatch %, y volcados a un .mlcomputeplan.json. Esto es el informe de rendimiento de Xcode, pero scriptable en CI.
  • Modelos multifunción (ct.utils.MultiFunctionDescriptor, iOS 18): varias variantes de shape de entrada exportadas como funciones de un .mlpackage con pesos compartidos — la respuesta moderna a "el tracing fija shapes estáticos".
  • Bisección de modelos (ct.models.utils.bisect_model): divide un modelo sobredimensionado en un pipeline fragmentado (chunked).
  • Metadatos de reproducibilidad: InferenceContextSpec/AppleSiliconContextMixin capturan el dispositivo (gpu_core_count, RAM, nombre del chip), el OS, el commit del código y el commit del modelo para cada benchmark; los campos de metadatos de Core ML (whisperkit_version, descripciones por entrada) se estampan en cada modelo exportado (whisperkit/test_utils.py::set_metadata_for_whisper_decoder).
  • Evaluaciones de extremo a extremo: whisperkit-evaluate-model calcula el WER sobre librispeech/earnings22/Common Voice vía datasets de HF; los resultados se publican en un espacio público de HF. La calidad se somete a tests de regresión, no se juzga a ojo.

Detalle pragmático destacable: el pipeline de generación relaja TEST_MIN_SPEEDUP_VS_CPU a 0.3 (generate_model.py:24) — para algunos componentes (un paso de decoder con KV cache es diminuto), superar a la CPU no es el objetivo; no bloquear el pipeline sí lo es.

9. Lecciones transferibles (adiciones al checklist de la KB)#

  1. Descompón el pipeline en modelos Core ML por cadencia (extractor de características / encoder / decoder / LUT de prefill).
  2. Haz el SDPA intercambiable y elige por componente: split softmax para secuencias largas, Cat compacto para la decodificación de un solo token.
  3. KV cache en la ANE = longitud máxima fija + máscara de actualización one-hot + máscara de padding, con la caché fusionada entre capas en el eje de canal; o MLState en iOS18+.
  4. K/V de cross-attention: computar una vez por contexto, eliminar las proyecciones del grafo del decoder.
  5. Precomputa todo lo enumerable en un modelo de embedding-LUT (context prefill).
  6. Elimina los parámetros matemáticamente redundantes (sesgo de key).
  7. Verifica con umbrales de PSNR + planes de cómputo programáticos en CI, no con inspección manual en Xcode.
  8. Comprime para regímenes limitados por ancho de banda — véase doc 07.

Siguiente: 07 — Model Compression for the ANE.

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