Pular para o conteúdo
← Todas as notícias

NVIDIA detalha otimizações no JAX para treinamento MoE do DeepSeek-V3 671B

Guia da NVIDIA explica como kernels especializados e comunicação mais eficiente levaram o DeepSeek-V3 671B de 103 para 1.068 TFLOPS/GPU em JAX.

A NVIDIA publicou um guia que mostra como otimizar o treinamento do DeepSeek-V3 671B em JAX com o Transformer Engine. O foco é modelo com mistura de especialistas, em que o roteamento de tokens e a troca de dados entre GPUs viram gargalos importantes.

NVIDIA detalha otimizações no JAX para treinamento MoE do DeepSeek-V3 671B

No exemplo citado, a base sem ajustes alcançava 103 TFLOPS/GPU, com 84% do tempo acumulado dos kernels gasto em comunicação entre GPUs. Com as otimizações descritas, o número chegou a 1.068 TFLOPS/GPU, um avanço de 10,4 vezes.

O texto destaca três partes centrais: o grouped GEMM, que executa as multiplicações de cada especialista usando o número real de tokens; a combinação entre dispatch e combine em um caminho fundido com NCCL EP; e a deduplicação de tokens, que evita enviar o mesmo token mais de uma vez pela rede quando ele vai para vários especialistas no mesmo nó ou em nós remotos.

O guia também cita offloading de ativações intermediárias para a memória do host e coletivas em múltiplos fluxos no XLA, para sobrepor comunicação entre nós e dentro do nó. Segundo a NVIDIA, o pacote de otimizações sustenta 97% de eficiência em 1.024 GPUs. As melhorias estão no contêiner NGC MaxText com Transformer Engine, e a empresa diz que planeja adicionar NVFP4, fusão de quantização com GEMM e sobreposição A2A.

← Todas as notícias