Criar IA de produção em Cloud TPUs com JAX

A pilha de IA JAX expande o núcleo numérico JAX com uma coleção de bibliotecas compostas apoiadas pela Google, transformando-o numa plataforma de código aberto robusta, ponto a ponto, para aprendizagem automática em escalas extremas. Como tal, a pilha de IA do JAX consiste num ecossistema abrangente e robusto que aborda todo o ciclo de vida da AA:

  • Base à escala industrial: a pilha de IA JAX foi arquitetada para uma escala massiva, tirando partido dos ML Pathways para orquestrar a preparação em dezenas de milhares de chips e do Orbax para a criação de pontos de verificação assíncronos resilientes e de elevado débito, o que permite a preparação de modelos de última geração de nível de produção.

  • Conjunto de ferramentas completo e pronto para produção: a pilha de IA JAX oferece um conjunto abrangente de bibliotecas para todo o processo de desenvolvimento: Flax para a criação flexível de modelos, Optax para estratégias de otimização compostas e Grain para os pipelines de dados determinísticos essenciais para execuções reproduzíveis em grande escala.

  • Desempenho especializado de pico: para alcançar a utilização máxima do hardware, a pilha de IA JAX oferece bibliotecas especializadas, incluindo Tokamax para kernels personalizados de vanguarda, Qwix para quantização não intrusiva que aumenta a velocidade de preparação e inferência, e XProf para criação de perfis de desempenho profundos e integrados no hardware.

  • Caminho completo para a produção: a pilha de IA JAX oferece uma transição perfeita da investigação à implementação. Isto inclui o MaxText como referência escalável para a preparação de modelos de base, o Tunix para a aprendizagem por reforço (AR) e o alinhamento de vanguarda, e uma solução de inferência unificada com a integração de TPU vLLM e o tempo de execução de serviço JAX.

A filosofia da pilha de IA do JAX é a de componentes fracamente acoplados, cada um dos quais faz uma coisa bem. Em vez de ser uma framework de ML monolítica, o JAX em si tem um âmbito restrito e foca-se em operações de matriz eficientes e transformações de programas. O ecossistema baseia-se nesta estrutura essencial para oferecer uma vasta gama de funcionalidades relacionadas com a preparação de modelos de ML e outros tipos de cargas de trabalho, como a computação científica.

Este sistema de componentes pouco acoplados permite-lhe selecionar e combinar bibliotecas da melhor forma para se adequar aos seus requisitos. Do ponto de vista da engenharia de software, esta arquitetura também permite atualizar a funcionalidade que seria tradicionalmente considerada componentes essenciais da framework (por exemplo, pipelines de dados e checkpointing) de forma iterativa sem o risco de desestabilizar a framework essencial ou ficar presa em ciclos de lançamento. Uma vez que a maioria das funcionalidades é implementada em bibliotecas em vez de alterações a uma estrutura monolítica, isto torna a biblioteca numérica principal mais duradoura e adaptável a mudanças futuras no panorama tecnológico.

As secções seguintes oferecem uma vista geral técnica da pilha de IA JAX, das respetivas principais funcionalidades, das decisões de design subjacentes e da forma como se combinam para criar uma plataforma duradoura para cargas de trabalho de ML modernas.

A pilha de IA JAX e outros componentes do ecossistema

Componente Função / descrição
Núcleo e componentes da pilha de IA JAX1
JAX Cálculo de matrizes orientado por aceleradores e transformação de programas (JIT, grad, vmap, pmap).
Flax Biblioteca de criação de redes neurais flexível para a criação e modificação intuitivas de modelos.
Optax Uma biblioteca de transformações de processamento e otimização de gradientes compostas.
Orbax Biblioteca de pontos de verificação distribuídos "any-scale" para resiliência de treino em grande escala.
Grão Uma biblioteca de data pipelines de entrada escalável, determinística e com pontos de verificação.
JAX AI stack - Infrastructure
XLA Compilador de aprendizagem automática de código aberto para TPUs, CPUs e GPUs.
Pathways Tempo de execução distribuído para orquestrar a computação em dezenas de milhares de chips.
Coleção de IA JAX - Adv. Programação
Pallas Uma extensão JAX para escrever kernels personalizados de baixo nível e alto desempenho implementados em Python.
Tokamax Uma biblioteca organizada de kernels personalizados de alto desempenho e de última geração (por exemplo, Attention).
Qwix Uma biblioteca abrangente e não intrusiva para a quantização (PTQ, QAT e QLoRA).
JAX AI stack – Aplicação
MaxText / MaxDiffusion Estruturas de referência emblemáticas e escaláveis para preparar modelos de base (por exemplo, LLM e Diffusion).
Tunix Uma estrutura para o alinhamento e o pós-treino de vanguarda (ARFH e ODP).
vLLM Uma solução de inferência de LLM de alto desempenho que usa a integração incorporada da framework vLLM.
XProf Um perfilador profundo integrado no hardware para análise do desempenho ao nível do sistema.

1Incluído no pacote Python.jax-ai-stack

Figura 1: a pilha de IA JAX e os componentes do ecossistema

Coleção de IA JAX

O imperativo arquitetónico: desempenho além das estruturas

À medida que as arquiteturas de modelos convergem, por exemplo, em transformadores multimodais de mistura de especialistas (MoE), a procura do desempenho máximo está a levar à emergência de megakernels. Um megakernel é efetivamente a passagem direta completa (ou uma grande parte) de um modelo específico, codificado manualmente através de uma API de nível inferior, como o CUDA SDK em GPUs NVIDIA. Esta abordagem alcança a máxima utilização do hardware através da sobreposição agressiva de computação, memória e comunicação. O trabalho recente da comunidade de investigação demonstrou que esta abordagem pode gerar ganhos significativos de débito, mais de 22% em alguns casos, para a inferência em GPUs. Esta tendência não se limita à inferência. Os dados sugerem que alguns esforços de preparação em grande escala envolveram o controlo de hardware de baixo nível para alcançar ganhos de eficiência substanciais.

Se esta tendência se acelerar, todas as frameworks de nível superior, tal como existem atualmente, correm o risco de se tornarem menos relevantes, uma vez que o acesso de baixo nível ao hardware é o que, em última análise, importa para o desempenho em arquiteturas estáveis e maduras. Isto representa um desafio para todas as stacks de ML modernas: como fornecer controlo de hardware ao nível de especialista sem sacrificar a produtividade e a flexibilidade de uma estrutura de alto nível.

Para que as TPUs ofereçam um caminho claro para este nível de desempenho, o ecossistema tem de expor uma camada de API mais próxima do hardware, o que permite o desenvolvimento destes núcleos altamente especializados. A pilha JAX foi concebida para resolver este problema, oferecendo um continuum de abstração (consulte a Figura 2), desde as otimizações automatizadas de alto nível do compilador XLA ao controlo manual detalhado da biblioteca de criação de kernels Pallas.

Figura 2: o continuum de abstração do JAX

O continuum de abstração do JAX

A coleção de IA JAX principal

A base da pilha de IA JAX consiste em cinco bibliotecas principais que fornecem a base para o desenvolvimento de modelos:

JAX: uma base para transformação de programas de alto desempenho e compósitos

O JAX é uma biblioteca Python para computação de matrizes orientada para aceleradores e transformação de programas, concebida para computação numérica de elevado desempenho e aprendizagem automática em grande escala. Com o seu modelo de programação funcional e API semelhante ao NumPy, o JAX oferece uma base sólida para bibliotecas de nível superior.

Com o seu design baseado no compilador, o JAX promove inerentemente a escalabilidade através da utilização do XLA (consulte a secção XLA) para uma análise, otimização e segmentação de hardware agressivas de todo o programa. A ênfase do JAX na programação funcional (por exemplo, funções puras) torna as transformações de programas essenciais mais tratáveis e, crucialmente, compostas.

Estas transformações essenciais podem ser combinadas para alcançar um elevado desempenho e escalabilidade das cargas de trabalho em função do tamanho do modelo, do tamanho do cluster e dos tipos de hardware:

  • jit: compilação just-in-time de funções Python em executáveis XLA otimizados e fundidos.
  • grad: diferenciação automática, compatível com o modo direto e inverso, bem como derivadas de ordem superior.
  • vmap: vetorização automática, que permite o processamento em lote e o paralelismo de dados sem problemas, sem modificar a lógica da função.
  • pmap / shard_map: paralelização automática em vários dispositivos (por exemplo, núcleos de TPU), que formam a base para a preparação distribuída.

A integração perfeita com o modelo GSPMD (SPMD de uso geral) do XLA permite que o JAX paralelize automaticamente os cálculos em grandes TPU Pods com alterações mínimas ao código. Na maioria dos casos, a escalabilidade só requer anotações de divisão em fragmentos de alto nível.

Flax: criação flexível de redes neurais

O Flax simplifica a criação, a depuração e a análise de redes neurais no JAX, oferecendo uma abordagem intuitiva e orientada para objetos à criação de modelos. Embora a API funcional do JAX seja poderosa, oferece uma abstração baseada em camadas mais familiar para os programadores habituados a frameworks como o PyTorch, sem qualquer penalização de desempenho.

Este design simplifica a modificação ou a combinação de componentes do modelo preparado. As técnicas como LoRA e quantização requerem definições de modelos manipuláveis, que a API NNX do Flax fornece através de uma interface Pythonic. NNX encapsula o estado do modelo, reduzindo a carga cognitiva do utilizador e permitindo a travessia programática e a modificação da hierarquia do modelo.

Principais pontos fortes:

  • API intuitiva orientada por objetos: simplifica a criação de modelos e permite exemplos de utilização avançados, como a substituição de submódulos e a inicialização parcial.
  • Consistente com o JAX principal: o Flax oferece transformações elevadas totalmente compatíveis com o paradigma funcional do JAX, oferecendo o desempenho total do JAX com maior facilidade de utilização para programadores.

Optax: estratégias de otimização e processamento de gradientes compostas

O Optax é uma biblioteca de processamento e otimização de gradientes para o JAX. Foi concebida para oferecer aos criadores de modelos bases que podem ser recombinadas de formas personalizadas para formar modelos de aprendizagem profunda, entre outras aplicações. Baseia-se nas capacidades da biblioteca JAX principal para fornecer uma biblioteca de funções de perda e otimização de alto desempenho bem testada e técnicas associadas que podem ser usadas para preparar modelos de ML.

Motivação

O cálculo e a minimização das perdas estão no centro do que permite o treino de modelos de ML. Com o respetivo suporte para diferenciação automática, a biblioteca JAX principal oferece as capacidades numéricas para formar modelos, mas não oferece implementações padrão de otimizadores populares (por exemplo, RMSProp ou Adam) nem perdas (por exemplo, CrossEntropy ou MSE). Embora possa implementar estas funções (e alguns programadores avançados optem por fazê-lo), um erro numa implementação do otimizador introduziria problemas de qualidade do modelo difíceis de diagnosticar. Em vez de o utilizador implementar estas partes críticas, a Optax fornece implementações destes algoritmos que são testadas quanto à correção e ao desempenho.

O campo da teoria da otimização situa-se claramente no domínio da investigação. No entanto, o seu papel central na preparação também a torna uma parte indispensável da preparação de modelos de ML de produção. Uma biblioteca que desempenhe esta função tem de ser suficientemente flexível para se adaptar a iterações de investigação rápidas e suficientemente robusta e com bom desempenho para ser fiável para a preparação de modelos de produção. Também deve fornecer implementações bem testadas de algoritmos de vanguarda que correspondam às equações padrão. A biblioteca Optax, através da sua arquitetura modular componível e ênfase no código legível correto, foi concebida para alcançar este objetivo.

Design

O Optax foi concebido para melhorar a velocidade da investigação e a transição da investigação para a produção, fornecendo implementações legíveis, bem testadas e eficientes de algoritmos essenciais. O Optax tem utilizações além do contexto da aprendizagem profunda. No entanto, neste contexto, pode ser visto como uma coleção de funções de perda, algoritmos de otimização e transformações de gradientes bem conhecidas implementadas de forma puramente funcional, em conformidade com a filosofia do JAX. A coleção de perdas conhecidas e otimizadores permite que os utilizadores comecem a usar a API com facilidade e confiança.

A abordagem modular adotada pela Optax permite encadear vários otimizadores juntamente com outras transformações comuns (por exemplo, restrição de gradiente) e envolvê-los usando técnicas comuns, como MultiStep ou Lookahead, para alcançar estratégias de otimização eficazes com algumas linhas de código. A interface flexível permite-lhe pesquisar novos algoritmos de otimização e usar técnicas de otimização de segunda ordem avançadas, como shampoo ou muon.

# Optax implementation of a RMSProp optimizer with a custom learning rate
#  schedule, gradient clipping and gradient accumulation.
optimizer = optax.chain(
  optax.clip_by_global_norm(GRADIENT_CLIP_VALUE),
  optax.rmsprop(learning_rate=optax.cosine_decay_schedule(init_value=lr,decay_steps=decay)),
  optax.apply_every(k=ACCUMULATION_STEPS)
)

# The same thing, in PyTorch
optimizer = optim.RMSprop(model_params, lr=LEARNING_RATE)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=TOTAL_STEPS)
for i, (inputs, targets) in enumerate(data_loader):
    # ... Training loop body ...
    if (i + 1) % ACCUMULATION_STEPS == 0:
        torch.nn.utils.clip_grad_norm_(model.parameters