Documentação

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

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 experiências locais reais, incluindo um kernel Metal escrito à mão. 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 suportam melhorias modestas e dependentes da forma a partir 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 emparelhadas por ronda. Alguns intervalos de lotes maiores não incluem qualquer melhoria. Um kernel personalizado de GELU/gate com erf exato teve sucesso numérico, mas não proporcionou um benefício adicional consistente ponta a ponta em relação à compilação do MLX. A quantização ingénua da espinha dorsal em 8 bits 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 compromissos 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 continua a ser BENCHMARKS.md.

Controlos e limites experimentais

Todo o trabalho de investigação com GPU correu em série. Outro trabalho do agent usou apenas CPU/sistema de ficheiros/rede. A máquina estava ligada à corrente, sem qualquer aviso térmico/de desempenho do pmset registado e sem uso de swap reportado durante a experiência. A atividade normal do ambiente de trabalho continuou. Isto não é uma câmara térmica controlada nem uma máquina de benchmark dedicada e de outro modo inativa.

As primeiras execuções de triagem executaram cada candidato num processo novo, com 4–5 warmups e 12–16 amostras. Revelaram uma deriva substancial entre execuções. Por exemplo, o piloto inglês de pergunta única sugeriu uma melhoria de 1.24× com a compilação, enquanto a experiência intercalada subsequente encontrou apenas cerca de 1.03×. As latências do piloto sequencial são, por isso, evidência de triagem, não a principal reivindicação causal de aumento de velocidade.

O script de confirmação paired.py:

  • Roda a ordem dos candidatos dentro de cada ronda e usa as mesmas entradas para todos os candidatos nessa ronda.
  • Muda o texto real do estado entre rondas. Gera até 16 variantes de estado e retém as variantes com a mesma forma de tensor; os casos curtos multilingues têm 10 dessas variantes, enquanto os outros casos reportados têm 16.
  • Usa perguntas distintas em linguagem natural, incluindo 50 instruções diferentes para a maior carga de trabalho curta. Verifica os hashes de entrada e não faz cache de respostas, não deduplica perguntas nem reutiliza estados contextuais do codificador.
  • Avalia os resultados e sincroniza a GPU antes de parar cada cronómetro. Mede tanto chamadas de passagem direta preparadas como o caminho de predição público, incluindo tokenização e formatação da saída. O carregamento do modelo está excluído.
  • Corre 32 rondas medidas para a experiência inglesa de cabeça/compilação e 16 para as experiências multilingue e de Metal personalizado, após o warmup. Cada candidato vê o mesmo número de rondas e a mesma sequência de entradas.

Estas entradas diferem dos fixtures de referência publicados. As comparações abaixo são dentro da experiência de investigação, não comparações antes/depois obtidas pela divisão de tabelas não relacionadas. O limite de lote da investigação é 64, enquanto a API lançada usa 16 por predefinição. As variantes de estado repetidas são medições repetidas intencionais; não há cache de resultados.

analyze.py calcula as razões eager_time / candidate_time por ronda e intervalos bootstrap exploratórios de percentis para a sua mediana, usando 2,000 reamostragens de índices de ronda. Esses intervalos não têm em conta todas as fontes de ruído do sistema operativo ou de correlação serial e não substituem a replicação em várias sessões. A razão de valores de p50 calculados independentemente pode diferir da razão mediana emparelhada.

O JSON em bruto inclui todos os tempos, hashes de entrada, metadados de ambiente, métricas de paridade e a impressão digital da origem registada no momento da medição. Os scripts das experiências foram subsequentemente formatados e ampliados com candidatos opcionais disjuntos; as impressões digitais anteriores descrevem essas versões anteriores dos scripts.

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

Compararam-se quatro caminhos:

  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, preserva a projeção QKV de comprimento total e os K/V, mas emite apenas queries de atenção de CLS/marcadores de opção. Corre a projeção de saída e a FFN apenas nessas saídas selecionadas.
  4. Atenção completa + saídas selecionadas + compilado: preserva o QKV de comprimento total e a chamada SDPA originais, e depois reúne as saídas de CLS/opção antes da projeção de saída e da FFN. Isto mantém a forma original do kernel de atenção removendo 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. Continuam a calcular todas as projeções QKV; não realizam a poupança adicional da projeção apenas Q do cálculo matemático do limite superior. Alterar as formas de GEMM e SDPA pode alterar o arredondamento de vírgula flutuante. Nenhum dos protótipos é uma cache de descodificador, uma saída antecipada ou uma aproximação que descarta camadas Transformer anteriores.

Latência p50 ponta a ponta, milissegundos:

Modelo / pedido 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
Multilingual curto 1 1 × 80 8.050 7.570 7.438 7.388
Multilingual curto 16 16 × 83 44.351 43.830 42.281 42.968
Multilingual longo 1 1 × 1024 41.964 42.017 40.492 41.120
Multilingual longo 8 8 × 1024 326.327 323.053 327.842 319.010

Fontes: dados emparelhados em inglês e dados emparelhados multilingues.

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

Pedido Aumento de velocidade mediano emparelhado Intervalo 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×
Multilingual curto 1 1.077× 1.046–1.140×
Multilingual curto 16 1.042× 1.017–1.067×
Multilingual longo 1 1.027× 1.012–1.054×
Multilingual longo 8 1.067× 0.958–1.082×

Os intervalos de 50 perguntas em inglês e de lote longo multilingue incluem 1. Não estabelecem uma melhoria repetível. O caminho de Q selecionado é um pouco melhor para os casos multilingues curto-16 e longo-1, mas nenhum caminho de poda domina todas as formas. Todos os intervalos dos candidatos, medições de forward e razões em bruto por ronda 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 entre as duas famílias de modelos nesta experiência de cabeça/compilação. Ambos os caminhos de poda concordaram nas 1,530 decisões de argmax, com uma diferença máxima de probabilidade calibrada de 0.0001883. O caminho de poda de atenção completa também passou a suite separada de 63 perguntas para cada modelo: 126/126 de concordância, com diferenças máximas de probabilidade de 4.31e-5 para inglês e 6.48e-6 para multilingue. Estes são testes de regressão, não uma afirmação de exatidão de tarefa em 1,530 exemplos rotulados independentemente.

Foram triadas tanto a compilação do modelo completo como a compilação por bloco. A experiência 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 completo. A especialização de forma tem de ser limitada num serviço. O modelo usa reshapes e máscaras dependentes da forma em Python, por isso aplicar shapeless=True indiscriminadamente é inseguro. O guia oficial de compilação documenta a especialização de forma e a captura de estado.

A primeira chamada do candidato inglês de modelo completo levou 2,166.7 ms, seguida de cerca de 12.75 ms de forward a quente p50 nesse piloto; uma primeira chamada de uma nova forma B16 levou 272.4 ms. O campo JSON chama-se 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. Os candidatos subsequentes reutilizaram kernels Metal previamente compilados, por isso os seus tempos de primeira chamada não são uma classificação controlada do custo de arranque a frio. A memória ativa/pico do MLX compilado em inglês curto-1 foi cerca de 803.6/918.6 MiB no piloto; a multilingue foi cerca de 614.1/676.9 MiB. Estas medições do alocador não incluem todas as alocações do compilador do lado do anfitrião e não estabelecem limites de memória sob mudança de forma sem limites. Vê o piloto de compilação em inglês e o piloto de compilação multilingue.

Quantização seletiva: poupanças de armazenamento úteis, inadequadas como reivindicação de velocidade

O protótipo chama nn.quantize depois de carregar o modelo denso FP16. Seleciona apenas os módulos lineares encoder.layers.*, com tamanho de grupo afim 64, e depois compila o modelo resultante. Embeddings, norms, a cabeça de decisão, o scorer e a cabeça de ação permanecem em FP16. Isto evita converter pesos inteiros empacotados através do carregador denso atual e evita a largura de entrada não divisível 1028/772 da cabeça de ação. Não está a ser disponibilizado nenhum formato de checkpoint quantizado nem contrato de carregamento. A implementação oficial de camadas quantizadas do MLX fornece este mecanismo de seleção.

Modelo / precisão do codificador Armazenamento total de tensores Concordância nos fixtures Maior mudança de probabilidade nos fixtures Concordância em cargas de trabalho distintas Maior mudança de probabilidade em cargas de trabalho distintas
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
Multilingual FP16 613.99 MiB Referência — Referência —
Multilingual 8 bits 515.38 MiB 63/63 0.0133 26/26 0.0358
Multilingual 4 bits 462.79 MiB 63/63 0.1268 19/26 0.8008

O resultado multilingue de 4 bits ilustra porque a pequena suite de fixtures por si só é insuficiente: os seus 63 argmaxes de fixtures mantiveram-se iguais, mas 7 de 26 decisões de cargas de trabalho distintas mudaram. Estas são medições de concordância em relação a FP16, não medições de exatidão de referência (ground truth). Uma mudança absoluta de probabilidade de 0.8008 equivale a 80.08 pontos percentuais.

Em entradas piloto inglesas curto-16, o p50 ponta a ponta eager/compilado FP16 foi 91.26/87.94 ms; o compilado 8 bits/4 bits foi 96.66/93.20 ms. A quantização curto-1 pareceu algo mais rápida nessa execução de triagem, enquanto as formas maiores não. A triagem multilingue de formas grandes também não mostrou uma vantagem de velocidade, mas as suas execuções sequenciais tiveram uma deriva substancial. Estas observações justificam rejeitar uma reivindicação não qualificada de aumento de velocidade ou de lançamento, não atribuir fatores precisos de abrandamento sem replicação quantizada intercalada. Mais trabalho de quantização precisa de calibração ou ajuste fino ciente das ativações e de uma suite de qualidade rotulada representativa.

Fontes em bruto: inglês 8 bits, inglês 4 bits, multilingue 8 bits, multilingue 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 da MLP, calcula o mesmo GELU baseado em erf, multiplica pelo gate e escreve uma única saída. Não substitui por tanh-GELU nem por uma aproximação sigmoid. O kernel usa os próprios auxiliares erf e expm1 do MLX v0.32.2, preservando as suas licenças e avisos em vendor/README.md. Suporta explicitamente apenas FP16 e usa o modo de matemática segura do Metal. O guia oficial de kernels personalizados descreve esta API e os seus controlos de modo de matemática.

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. Os testes de modelo completo com entradas alteradas também corresponderam exatamente: 474/474 comparações de perguntas entre as duas famílias de modelos, mais ambas as suites 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 numa vantagem de velocidade consistente sobre a expressão compilada fundida do MLX. Por exemplo, com 1,312 tokens e largura intermédia 2,624, o tempo de ativação sincronizado por chamada foi 0.378 ms para o eager GELU-depois-gate, 0.268 ms para mx.compile e 0.280 ms para o kernel personalizado. Com 8,192 tokens e largura 1,152, os valores correspondentes foram 0.846/0.764/0.714 ms. Estes microbenchmarks incluem sobrecarga de despacho e de sincronização e são sondas de triagem; não são medições do tempo de execução isolado do dispositivo. As entradas completas, os tempos em bruto e as verificações de igualdade estão em microbench.json.

O kernel personalizado foi depois instalado em todas as MLPs do codificador e medido no modelo completo com ordem de candidatos rotativa e entradas em mudança:

Modelo / pedido 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
Multilingual curto 1 7.557 ms 7.437 ms
Multilingual curto 16 49.683 ms 50.301 ms
Multilingual longo 1 48.906 ms 51.032 ms

As execuções emparelhadas completas do kernel personalizado usam uma segunda instância do modelo com pesos idênticos, para que as implementações não modificada e personalizada coexistam sem mutação nem capturas compiladas desatualizadas. Os 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 apoiam a publicação do kernel personalizado como uma melhoria de desempenho geral. Fontes: dados emparelhados Metal em inglês e dados emparelhados Metal multilingues.

Onde a engenharia personalizada valeria a pena investigar mais

O modelo já chama mx.fast.scaled_dot_product_attention, mx.fast.rope e normalização de camada otimizada. O seu caminho SDPA de máscara booleana D64 está fundido; não há nenhum interruptor de Flash Attention em falta que explique uma diferença de 10×. A atenção local continua a percorrer tiles densos de key/value. Um verdadeiro kernel de janela bidirecional poderia saltar esses tiles preservando a distância inclusiva <=64 e a semântica de padding, mas a sua oportunidade aritmética no modelo completo é pequena em entradas curtas e é limitada nas formas longas publicadas. A revisão de código existente e o relatório matemático quantificam esta distinção.

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

  • Atenção de janela para entradas longas: especializar os limites de tiles para D64, a janela bidirecional real e lotes com padding. Comparar com o SDPA denso fundido em 512/1024 tokens e depois no modelo completo. Este kernel não foi construído nem medido neste relatório.
  • Epílogos e escalonamento de kernels densos: investigar a fusão do epílogo da MLP com gating dentro do GEMM ou melhorar o escalonamento de matrizes de M curto. O MLX já usa implementações Metal de GEMM especializadas, por isso substituí-las exige um perfil real de despacho/kernel e ganhos medidos para as formas exatas de M/N/K. O resultado isolado das ativações mostra porque outro kernel elementwise sozinho é insuficiente.
  • Batching ciente do comprimento e preparação partilhada na CPU: preservar IDs de entrada exatos ao tokenizar o texto do estado partilhado uma vez antes de construir cada sequência de pergunta, e evitar fazer padding de itens pequenos para itens longos não relacionados. O piloto multilingue longo-8 gastou cerca de 13.1 ms a preparar entradas, contra centenas de milissegundos ponta a ponta. Mesmo eliminar totalmente essa preparação não produziria 10× nesta carga de trabalho. A latência de fila de espera e a contagem de inferências únicas têm de fazer parte de qualquer reivindicação de batching.
  • Um aluno mais pequeno que responde em conjunto: se 10× for um requisito de produto, destilar ou redesenhar o modelo para remover a maior parte do trabalho denso ou responder a muitas perguntas fixas com uma única codificação contextual. Isto altera o modelo aprendido e precisa de treino/avaliação rotulados representativos; não é uma otimização exata do port. Reutilizar um estado KV contextual arbitrário entre perguntas no atual codificador bidirecional é inválido.

Oito sondas GEMM autónomas FP16 de projeção de entrada do codificador alcançaram 0.66–11.55 TFLOP/s, incluindo a sincronização por chamada. A grande sonda inglesa M=4096, N=5248, K=1024 alcançou 11.55 TFLOP/s; a sonda multilingue M=8192, N=2304, K=768 alcançou 7.96 TFLOP/s. Estes são valores de débito observados, não especificações de pico de hardware nem limites superiores do débito do grafo completo. As medições de M pequeno são particularmente dominadas por custos de submissão e sincronização; um grafo em streaming amortece-os de forma diferente. Mostram que formas merecem a criação de perfis, não uma prova de que não possa existir um kernel melhor. Os orçamentos de débito de mesmo trabalho de 10× no relatório matemático permanecem requisitos teóricos e não capacidades de dispositivo medidas.

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

Os scripts usam o .venv existente e checkpoints locais fixados. Corre os comandos de GPU sequencialmente, nunca a par 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 ficheiros Python experimentais passam as verificações de formatação e lint do Ruff. O runtime estável, os resultados do benchmark original e os checkpoints FP16 publicados continuam a ser os artefactos de lançamento. A compilação e a poda exata da cabeça final são otimizações futuras opcionais credíveis, após a política de forma/cache a frio e uma validação de qualidade mais ampla; os ganhos medidos não justificam acrescentar silenciosamente uma latência de compilação ou um kernel personalizado ao caminho predefinido. Não se reivindica nenhum aumento de velocidade de 10×, nenhum checkpoint quantizado pronto para produção nem qualquer ganho medido do kernel de janela local.