Documentação

Investigação de engenharia: este port MLX pode ficar mais 10× mais rápido?

Data: 2026-09-19. Máquina: Apple M3 Max, 40 núcleos de GPU, 128 GiB de memória unificada, macOS 27.2, MLX / MLX Metal 0.32.2, inferência FP16. Este relatório contém experimentos locais reais, incluindo um kernel Metal escrito à mão. Ele não altera o runtime de produção nem publica pesos quantizados.

As mudanças de engenharia testadas não entregam 10×. Medições intercaladas sustentam melhorias modestas e dependentes de forma vindas da compilação e da poda de saídas não usadas da camada final da cabeça de decisão. Casos selecionados melhoraram aproximadamente 3–8% usando medianas pareadas por rodada. Alguns intervalos de lotes maiores não incluem melhoria. Um kernel personalizado de GELU/gate com erf exato foi numericamente bem-sucedido, mas não forneceu um benefício ponta a ponta adicional consistente sobre a compilação do MLX. A quantização ingênua do backbone em 8 e 4 bits reduziu o armazenamento, falhou em acelerar as cargas de trabalho piloto maiores e alterou previsões ou probabilidades calibradas.

Os limites matemáticos e os trade-offs de aproximação são examinados separadamente em MATH_10X_RESEARCH.md. A revisão original da implementação está em PERFORMANCE_RESEARCH.md; o benchmark dos checkpoints lançados permanece em BENCHMARKS.md.

Controles e limites experimentais

Todo o trabalho de GPU da pesquisa rodou serialmente. Outro trabalho do agent usou apenas CPU/sistema de arquivos/rede. A máquina estava na tomada, sem aviso térmico/de desempenho do pmset registrado e sem uso de swap relatado durante o experimento. A atividade normal da área de trabalho continuou. Esta não é uma câmara térmica controlada nem uma máquina dedicada de benchmark ociosa.

As primeiras execuções de triagem executaram cada candidato em um processo novo, com 4–5 warmups e 12–16 amostras. Elas revelaram deriva substancial entre execuções. Por exemplo, o piloto inglês de uma pergunta sugeriu uma melhoria de compilação de 1.24×, enquanto o experimento intercalado seguinte encontrou apenas cerca de 1.03×. As latências dos pilotos sequenciais são, portanto, evidência de triagem, não a alegação causal principal de ganho de velocidade.

O script de confirmação paired.py:

  • Rotaciona a ordem dos candidatos dentro de cada rodada e usa as mesmas entradas para cada candidato nessa rodada.
  • Muda o texto real do state entre rodadas. Gera até 16 variantes de state e retém as variantes com a mesma forma de tensor; os casos curtos multilíngues têm 10 dessas variantes, enquanto os outros casos relatados têm 16.
  • Usa perguntas distintas em linguagem natural, incluindo 50 instruções diferentes para a maior carga de trabalho curta. Verifica hashes de entrada e não faz cache de respostas, não deduplica perguntas nem reutiliza estados contextuais do encoder.
  • Avalia os resultados e sincroniza a GPU antes de parar cada cronômetro. Mede tanto chamadas preparadas de forward quanto o caminho público de predição, incluindo tokenização e formatação de saída. O carregamento do modelo fica de fora.
  • Roda 32 rodadas medidas para o experimento inglês de cabeça/compilação e 16 para os experimentos multilíngue e de Metal personalizado, após o warmup. Cada candidato vê a mesma contagem de rodadas e a mesma sequência de entrada.

Estas entradas diferem dos fixtures de linha de base publicados. As comparações abaixo são dentro do experimento de pesquisa, não comparações antes/depois obtidas por divisão de tabelas não relacionadas. O limite de lote da pesquisa é 64, enquanto a API lançada usa 16 por padrão. Variantes de state repetidas são medições de repetição intencionais; não há cache de resultados.

analyze.py calcula razões por rodada de eager_time / candidate_time e intervalos exploratórios de bootstrap de percentil para sua mediana, usando 2,000 reamostragens de índices de rodada. Esses intervalos não levam em conta todas as fontes de ruído do sistema operacional ou correlação serial e não substituem a replicação em múltiplas sessões. Uma razão de valores p50 calculados independentemente pode diferir da razão pareada mediana.

O JSON bruto inclui todos os tempos, hashes de entrada, metadados de ambiente, métricas de paridade e a impressão digital da origem registrada no momento da medição. Os scripts de experimento foram posteriormente formatados e estendidos com candidatos opcionais disjuntos; impressões digitais anteriores descrevem aquelas versões anteriores dos scripts.

Compilação e poda exata da cabeça final

Quatro caminhos foram comparados:

  1. Eager: o DecisionModel FP16 lançado.
  2. Compilado: mx.compile em torno do modelo carregado, avaliado e congelado, usando especialização de forma normal.
  3. Q selecionado + compilado: na última camada da cabeça, preserve a projeção QKV de comprimento total e K/V, mas emita apenas queries de atenção de CLS/marcadores de opção. Rode a projeção de saída e a FFN apenas nessas saídas selecionadas.
  4. Atenção completa + saídas selecionadas + compilado: preserve o QKV de comprimento total original e a chamada SDPA, depois reúna as saídas de CLS/opção antes da projeção de saída e da FFN. Isso mantém a forma original do kernel de atenção enquanto remove a maior parte do trabalho denso não usado da cabeça final.

Ambos os protótipos de poda preservam as dependências matemáticas do modelo. Eles ainda calculam todas as projeções QKV; não realizam a economia adicional da projeção apenas de Q no cálculo de limite superior matemático. Mudar as formas de GEMM e SDPA pode alterar o arredondamento de ponto flutuante. Nenhum dos protótipos é um cache de decoder, uma saída antecipada ou uma aproximação que descarta camadas anteriores do transformer.

Latência p50 ponta a ponta, em milissegundos:

Modelo / requisição B × L Eager Compilado Q selecionado + compilado Atenção completa + saídas selecionadas + compilado
Inglês curto 1 1 × 78 16.628 16.185 15.925 15.636
Inglês curto 16 16 × 82 116.920 113.700 112.009 110.217
Inglês longo 1 1 × 512 53.921 53.078 52.301 52.121
Inglês longo 8 8 × 512 531.166 518.428 504.135 488.980
Inglês curto 50 50 × 82 456.333 439.013 445.223 438.293
Multilíngue curto 1 1 × 80 8.050 7.570 7.438 7.388
Multilíngue curto 16 16 × 83 44.351 43.830 42.281 42.968
Multilíngue longo 1 1 × 1024 41.964 42.017 40.492 41.120
Multilíngue longo 8 8 × 1024 326.327 323.053 327.842 319.010

Fontes: dados pareados em inglês e dados pareados multilíngues.

Para o caminho de atenção completa/saídas selecionadas, o ganho de velocidade pareado mediano e os intervalos exploratórios de 95% incluem:

Requisição Ganho de velocidade pareado mediano Intervalo de bootstrap
Inglês curto 1 1.049× 1.043–1.056×
Inglês curto 16 1.059× 1.033–1.077×
Inglês longo 1 1.039× 1.027–1.052×
Inglês longo 8 1.061× 1.020–1.095×
Inglês curto 50 1.022× 0.977–1.050×
Multilíngue curto 1 1.077× 1.046–1.140×
Multilíngue curto 16 1.042× 1.017–1.067×
Multilíngue longo 1 1.027× 1.012–1.054×
Multilíngue longo 8 1.067× 0.958–1.082×

Os intervalos inglês de 50 perguntas e multilíngue de lote longo incluem 1. Eles não estabelecem uma melhoria repetível. O caminho de Q selecionado é um pouco melhor para os casos multilíngues curto-16 e longo-1, mas nenhum caminho de poda domina todas as formas. Todos os intervalos candidatos, medições de forward e razões brutas por rodada estão em paired_analysis.json.

A compilação correspondeu exatamente aos logits eager, aos logits de ação e às probabilidades calibradas nas 1,530 comparações de perguntas com entradas alteradas nas duas famílias de modelo neste experimento de cabeça/compilação. Ambos os caminhos de poda concordaram em todas as 1,530 decisões de argmax, com diferença máxima de probabilidade calibrada de 0.0001883. O caminho de poda com atenção completa também passou na suíte separada de 63 perguntas para cada modelo: 126/126 de concordância, com diferenças máximas de probabilidade de 4.31e-5 para o inglês e 6.48e-6 para o multilíngue. São verificações de regressão, não uma alegação de precisão na tarefa em 1,530 exemplos rotulados independentemente.

A compilação do modelo inteiro e por bloco foram ambas triadas. O experimento por bloco também preservou todas as 63 saídas dos fixtures em inglês, mas não estabeleceu uma vantagem material sobre a compilação do modelo inteiro. A especialização de forma deve ser limitada em um serviço. O modelo usa reshapes e máscaras dependentes de forma em Python, então aplicar shapeless=True indiscriminadamente é inseguro. O guia oficial de compilação documenta especialização de forma e captura de estado.

A primeira chamada do candidato inglês de modelo inteiro levou 2,166.7 ms, seguida de cerca de 12.75 ms de p50 de forward a quente nesse piloto; a primeira chamada de uma nova forma B16 levou 272.4 ms. O campo JSON se chama cold_forward, mas significa a primeira chamada do candidato após a inferência de referência eager, não uma aplicação totalmente a frio nem um driver Metal recém-inicializado. Candidatos subsequentes reutilizaram kernels Metal já compilados, então seus tempos de primeira chamada não são um ranking controlado do custo de partida a frio. A memória MLX ativa/pico do inglês compilado de curto-1 foi de cerca de 803.6/918.6 MiB no piloto; a do multilíngue foi de cerca de 614.1/676.9 MiB. Essas medições do alocador não incluem todas as alocações do compilador no host e não estabelecem limites de memória sob churn ilimitado de formas. Veja piloto de compilação em inglês e piloto de compilação multilíngue.

Quantização seletiva: economia útil de armazenamento, inadequada como alegação de velocidade

O protótipo chama nn.quantize depois de carregar o modelo denso FP16. Ele seleciona apenas os módulos lineares encoder.layers.*, com tamanho de grupo afim 64, e então compila o modelo resultante. Embeddings, norms, a cabeça de decisão, o scorer e a cabeça de ação permanecem em FP16. Isso evita converter pesos inteiros empacotados pelo loader denso atual e evita a largura de entrada não divisível 1028/772 da cabeça de ação. Nenhum formato de checkpoint quantizado nem contrato de carregamento está sendo lançado. A implementação oficial de camadas quantizadas do MLX fornece este mecanismo de seleção.

Modelo / precisão do encoder Armazenamento total de tensores Concordância no fixture Maior mudança de probabilidade no fixture Concordância em carga de trabalho distinta Maior mudança de probabilidade em carga de trabalho distinta
Inglês FP16 803.55 MiB Referência — Referência —
Inglês 8 bits 496.76 MiB 62/63 0.0401 18/18 0.0312
Inglês 4 bits 333.13 MiB 50/63 0.3256 18/18 0.2224
Multilíngue FP16 613.99 MiB Referência — Referência —
Multilíngue 8 bits 515.38 MiB 63/63 0.0133 26/26 0.0358
Multilíngue 4 bits 462.79 MiB 63/63 0.1268 19/26 0.8008

O resultado multilíngue de 4 bits ilustra por que a pequena suíte de fixtures sozinha é insuficiente: seus 63 argmaxes de fixture permaneceram iguais, mas 7 de 26 decisões distintas de carga de trabalho mudaram. Estas são medições de concordância em relação ao FP16, não medições de precisão contra a verdade de referência. Uma mudança absoluta de probabilidade de 0.8008 é 80.08 pontos percentuais.

Em entradas piloto inglesas de curto-16, o p50 ponta a ponta FP16 eager/compilado foi 91.26/87.94 ms; o 8 bits/4 bits compilado foi 96.66/93.20 ms. A quantização de curto-1 pareceu um pouco mais rápida nessa execução de triagem, enquanto formas maiores não. A triagem multilíngue de formas grandes também não mostrou ganho de velocidade, mas suas execuções sequenciais tiveram deriva substancial. Essas observações justificam rejeitar uma alegação de ganho de velocidade ou de lançamento sem qualificação, e não atribuir fatores precisos de lentidão sem replicação quantizada intercalada. Trabalho adicional de quantização precisa de calibração ciente de ativação ou ajuste fino e uma suíte representativa de qualidade rotulada.

Fontes brutas: inglês 8 bits, inglês 4 bits, multilíngue 8 bits, multilíngue 4 bits.

Metal escrito à mão: a fusão exata GELU/gate foi implementada e testada

kernels.py implementa um kernel Metal personalizado real que lê os dois ramos concatenados do MLP, calcula o mesmo GELU baseado em erf, multiplica pelo gate e escreve uma única saída. Ele não substitui o tanh-GELU nem uma aproximação por sigmoid. O kernel usa os próprios helpers erf e expm1 do MLX v0.32.2, preservando suas licenças e avisos em vendor/README.md. Ele suporta explicitamente apenas FP16 e usa o modo matemático seguro do Metal. O guia oficial de kernels personalizados descreve esta API e seus controles de modo matemático.

Em oito formas de ativação representativas, 27,958,016 elementos de saída FP16 gerados aleatoriamente tiveram valores exatamente iguais aos da operação original. O microbenchmark compara igualdade numérica, não o bit de sinal do zero. Testes de modelo inteiro com entradas alteradas também corresponderam exatamente: 474/474 comparações de pergunta nas duas famílias de modelo, mais as duas suítes de fixtures de 63 perguntas, com zero diferença de logit, logit de ação ou probabilidade calibrada.

Este resultado de correção não se traduziu em uma vantagem de velocidade consistente sobre a expressão compilada e fundida do MLX. Por exemplo, em 1,312 tokens e largura intermediária 2,624, o tempo de ativação sincronizado por chamada foi 0.378 ms para GELU-e-gate eager, 0.268 ms para mx.compile e 0.280 ms para o kernel personalizado. Em 8,192 tokens e largura 1,152, os valores correspondentes foram 0.846/0.764/0.714 ms. Esses microbenchmarks incluem sobrecarga de despacho e sincronização e são sondas de triagem; não são medições do tempo de execução isolado do dispositivo. Entradas completas, tempos brutos e verificações de igualdade estão em microbench.json.

O kernel personalizado foi então instalado em cada MLP do encoder e medido no modelo completo com ordem rotativa de candidatos e entradas variáveis:

Modelo / requisição p50 compilado original p50 Metal + compilado
Inglês curto 1 23.795 ms 23.837 ms
Inglês curto 16 142.716 ms 139.355 ms
Inglês longo 1 68.241 ms 68.982 ms
Multilíngue curto 1 7.557 ms 7.437 ms
Multilíngue curto 16 49.683 ms 50.301 ms
Multilíngue longo 1 48.906 ms 51.032 ms

As execuções pareadas completas do kernel personalizado usam uma segunda instância de modelo com pesos idênticos, para que as implementações não modificada e personalizada coexistam sem mutação nem capturas compiladas desatualizadas. Seus tempos absolutos não devem ser comparados com a execução anterior de poda da cabeça. Os resultados mistos e modestos não sustentam publicar o kernel personalizado como uma melhoria geral de desempenho. Fontes: dados pareados de Metal em inglês e dados pareados de Metal multilíngue.

Onde a engenharia personalizada valeria mais investigação

O modelo já chama mx.fast.scaled_dot_product_attention, mx.fast.rope, e normalização de camada otimizada. Seu caminho SDPA de máscara booleana D64 é fundido; não há um switch ausente de Flash Attention que explique uma lacuna de 10×. A atenção local ainda percorre tiles densos de key/value. Um kernel real de janela bidirecional poderia pular esses tiles preservando a distância inclusiva <=64 e a semântica de padding, mas sua oportunidade aritmética no modelo inteiro é pequena em entradas curtas e é limitada nas formas longas publicadas. A revisão de código existente e o relatório matemático quantificam essa distinção.

Próximos projetos úteis, com seus requisitos de evidência, são:

  • Atenção de janela para entradas longas: especialize limites de tile para D64, a janela bidirecional real e lotes com padding. Compare com o SDPA denso fundido em 512/1024 tokens e depois no modelo completo. Este kernel não foi construído nem submetido a benchmark neste relatório.
  • Epílogos e agendamento de kernels densos: investigue fundir o epílogo do MLP com gate na GEMM ou melhorar o agendamento de matrizes com M pequeno. O MLX já usa implementações especializadas de GEMM em Metal, então substituí-las exige um profile real de despacho/kernel e ganhos medidos para as formas exatas de M/N/K. O resultado isolado de ativação mostra por que outro kernel elementwise sozinho é insuficiente.
  • Batching ciente do comprimento e preparação compartilhada na CPU: preserve IDs de entrada exatos enquanto tokeniza o texto de state compartilhado uma vez antes de construir cada sequência de pergunta, e evite fazer padding de itens pequenos até itens longos não relacionados. O piloto multilíngue longo-8 gastou cerca de 13.1 ms preparando entradas, contra centenas de milissegundos de ponta a ponta. Mesmo eliminando totalmente essa preparação, não se produziria 10× nesta carga de trabalho. Latência de fila e contagem de inferências únicas devem fazer parte de qualquer alegação de batching.
  • Um student menor que responde em conjunto: se 10× for um requisito de produto, destile ou redesenhe o modelo para remover a maior parte do trabalho denso ou responder muitas perguntas fixas com uma única codificação contextual. Isso muda o modelo aprendido e precisa de treinamento/avaliação representativos e rotulados; não é uma otimização exata de port. Reutilizar um estado contextual/KV arbitrário entre perguntas no encoder bidirecional atual é inválido.

Oito sondas autônomas de GEMM de projeção de entrada do encoder em FP16 alcançaram 0.66–11.55 TFLOP/s, incluindo sincronização por chamada. A sonda inglesa grande M=4096, N=5248, K=1024 alcançou 11.55 TFLOP/s; a sonda multilíngue M=8192, N=2304, K=768 alcançou 7.96 TFLOP/s. Estes são valores de vazão observados, não especificações de pico de hardware nem limites superiores de vazão do grafo completo. Medições com M pequeno são particularmente dominadas por custos de submissão e sincronização; um grafo em streaming as amortiza de forma diferente. Elas mostram quais formas merecem profiling, não uma prova de que nenhum kernel melhor possa existir. Os orçamentos de vazão de 10× para o mesmo trabalho no relatório matemático continuam sendo requisitos teóricos, e não capacidades medidas do dispositivo.

Reprodução e decisão de lançamento

Os scripts usam o .venv existente e checkpoints locais fixados. Rode os comandos de GPU em sequência, nunca ao lado do benchmark formal:

# Screening: repeat for eager, compiled, blocks, q8, q4, selected-compiled.
.venv/bin/python -m experiments.engineering.run_variants \
  --model laya --variant compiled --iterations 12 --warmup 4 --quality \
  --output experiments/engineering/reproduced-compiled.json

# Primary confirmation, including 50 genuinely different questions.
.venv/bin/python -m experiments.engineering.paired \
  --model laya --iterations 32 \
  --output experiments/engineering/reproduced-laya-paired.json
.venv/bin/python -m experiments.engineering.paired \
  --model laya-multilingual --iterations 16 --cases short1,short16,long1,long8 \
  --output experiments/engineering/reproduced-multilingual-paired.json

# Hand-written kernel microbench and complete-model comparison.
.venv/bin/python -m experiments.engineering.microbench
.venv/bin/python -m experiments.engineering.paired \
  --model laya --iterations 16 --cases short1,short16,long1 --metal \
  --output experiments/engineering/reproduced-metal-paired.json
.venv/bin/python -m experiments.engineering.run_variants \
  --model laya --variant metal-compiled --iterations 5 --warmup 3 \
  --cases short1 --quality --output experiments/engineering/reproduced-metal-quality.json

# CPU-only paired analysis.
.venv/bin/python -m experiments.engineering.analyze

Todos os arquivos Python experimentais passam nas verificações de formatação e lint do Ruff. O runtime estável, os resultados originais de benchmark e os checkpoints FP16 publicados permanecem sendo os artefatos de lançamento. A compilação e a poda exata da cabeça final são otimizações futuras opcionais críveis, após a política de forma a frio/cache e uma validação de qualidade mais ampla; os ganhos medidos não justificam adicionar silenciosamente latência de compilação ou um kernel personalizado ao caminho padrão. Nenhum ganho de 10×, checkpoint quantizado pronto para produção ou ganho medido de kernel de janela local é alegado.