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 enreferences/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 completo — k_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 condecoder_key_padding_mask(-1e4aditivo, la convención segura para FP16 de la KB). kv_cache_update_maskes 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:
- Enumera todos los prefijos válidos (idioma × tarea) (~99 idiomas × 2 tareas),
- Precomputa las KV caches del decoder para cada prefijo y las almacena aplanadas en dos tablas de búsqueda
nn.Embedding, - 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.conv2dcon 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 unnn.Moduleconregister_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.weightvíaF.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_NEvsCPU_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.mlmodelccompilado conct.models.compute_plan.MLComputePlany registra, por op, el dispositivo de despacho, los dispositivos soportados y la cuota de coste estimada — resumidos comoANE support coverage %yANE 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.mlpackagecon 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/AppleSiliconContextMixincapturan 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-modelcalcula el WER sobrelibrispeech/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)#
- Descompón el pipeline en modelos Core ML por cadencia (extractor de características / encoder / decoder / LUT de prefill).
- Haz el SDPA intercambiable y elige por componente: split softmax para secuencias largas,
Catcompacto para la decodificación de un solo token. - 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
MLStateen iOS18+. - K/V de cross-attention: computar una vez por contexto, eliminar las proyecciones del grafo del decoder.
- Precomputa todo lo enumerable en un modelo de embedding-LUT (context prefill).
- Elimina los parámetros matemáticamente redundantes (sesgo de key).
- Verifica con umbrales de PSNR + planes de cómputo programáticos en CI, no con inspección manual en Xcode.
- Comprime para regímenes limitados por ancho de banda — véase doc 07.
Siguiente: 07 — Model Compression for the ANE.