O que acontece quando o cache do JIT enche demais
Você roda um modelo em JAX por horas, talvez dias, e de repente começa a ver alertas de memória subindo sem motivo óbvio. O processo não está processando dados novos demais, não está carregando datasets gigantes — simplesmente o cache da compilação just-in-time foi acumulando traces até ocupar gigabytes. Isso é o problema central que quem trabalha com JAX encontra na prática, e é algo que a documentação oficial nem sempre deixa claro desde o início.
O que é jit fade cacheado
O termo que eu costumo usar internamente, jit fade cacheado, se refere à situação em que o cache do JAX se torna tão grande que passa a causar efeitos negativos: ou o processo consome memória demais, ou novas funções são compiladas quando não deveria ser necessário porque o cache antigo "sufocou" o mecanismo de lookup. Na minha experiência, isso aparece principalmente em pipelines de treinamento onde os shapes dos inputs variam levemente entre batches, gerando entradas novas no cache a cada época. O JAX funciona de maneira diferente do TensorFlow ou PyTorch padrão. Quando você decora uma função com @jax.jit, a primeira execução compila um trace da função baseado nos shapes e tipos dos argumentos. Esse trace compilado fica armazenado em um dicionário interno. Chamadas subsequentes com o mesmo signature recuperam o trace compilado, evitando recompilação. O problema é que esse dicionário não tem tamanho limitado por padrão.
Por que o cache cresce sem controle
Existem dois cenários principais. O primeiro é usar argumentos dinâmicos para shapes que mudam constantemente. Se você passa arrays com shapes diferentes repetidamente — digamos, batches de tamanhos ligeiramente variados no final de uma época — o JAX trata cada shape único como uma nova entrada no cache. Em um treinamento de 100 épocas com 10.000 batches, isso pode significar centenas de milhares de entradas no cache. O segundo cenário é menos óbvio. Você pode estar usando static_argnums ou static_argnames de forma incorreta, passando valores que mudam a cada chamada como argumentos estáticos quando na verdade deveriam ser dinâmicos, ou vice-versa. O JAX diferencia estritamente entre argumentos estáticos (que afetam o trace) e dinâmicos (que não afetam). Misturar isso gera entradas extras no cache que parecem redundantes mas na verdade criam traces distintos.
Um detalhe que muita gente perde: o cache do JAX é baseado em uma hash dos shapes e tipos. Mesmo mudanças minúsculas — um array que era float32 e passa a ser float64 por causa de uma operação intermediária — gera uma entrada nova. Eu já perdi duas tardesando um leak de memória só porque uma operação de normalização estava promovendo tipos silenciosamente.
Como diagnosticar o problema
A ferramenta mais direta é jax.jit.clear_cache(). Você pode chamar isso em um ponto de verificação para ver quantos traces estão ativos:
import jax
Ver tamanho do cache
print(len(jax._src.xla.xla._xla_computations))
Limpar cache
jax.jit.clear_cache()
Existe também o flag JAX_TRACE_CACHE_LOGGING=1 que habilita logging detalhado de todas as entradas e saídas do cache. Isso gera bastante output, mas é útil para entender exatamente quais chamadas estão gerando traces novos. Eu uso isso em scripts de debug antes de ir. Outra abordagem é monitorar o uso de memória do processo com ferramentas como memory_profiler ou simplesmente psutil em intervalos regulares durante o treinamento. Se a memória sobe de forma monotônica sem relação direta com o tamanho dos dados processados, o cache é o suspeito principal.
👉 Clique no botão abaixo para saber mais sobre o assunto!
Soluções práticas
Limpeza periódica do cache: A solução mais simples é chamar jax.jit.clear_cache() em intervalos regulares. Em pipelines longos, eu costumo limpar a cada N passos ou épocas. Isso remove traces que já não são mais necessários. O custo é que a próxima vez que aquela função for chamada com o mesmo shape, ela precisará ser recompilada. Para funções que são chamadas frequentemente com shapes fixos, esse custo de recompilação é mínimo comparado ao ganho de memória. Uso correto de argumentos estáticos: A maioria dos problemas de cache explodindo vem de argumentos que deveriam ser estáticos mas não são, ou vice-versa. Shapes que mudam frequentemente devem ser tratados como dinâmicos, enquanto configurações que definem a estrutura do trace (como o número de camadas) devem ser estáticas. A regra prática é: se mudar esse argumento muda a estrutura computacional da função, ele é estático. Se muda apenas os dados mas não a estrutura, é dinâmico.
import jax
Errado: shape variável tratado como dinâmico gera muitas entradas
@jax.jit
def process_data(x, config):
return x @ config['weight']
Certo: config é estático porque define a estrutura
@jax.jit
def process_data(x, config):
return x @ config['weight']
process_data_jitted = jax.jit(process_data, static_argnames=['config'])
Buffer donation: Quando você não precisa mais de um array de entrada após a computação, pode indicá-lo para doação de buffer usando donate_argnums. Isso permite que o JAX reutilize a memória do buffer de entrada para o resultado, reduzindo o pico de memória sem afetar o cache diretamente, mas ajudando no contexto geral de gestão de memória.
@jax.jit(donate_argnums=(0,))
def train_step(state, batch):
'state' será liberado após esta operação
new_state, loss = update(state, batch)
return new_state, loss
Shape padding ou truncamento: Em casos onde os batches têm tamanhos variáveis no último dimensão, você pode padronizar todos para o mesmo tamanho máximo usando jnp.pad ou truncar. Isso reduz o número de shapes únicos e portanto o número de entradas no cache. O overhead computacional é geralmente pequeno comparado ao custo de memória economizado.
Limitações e tradeoffs honestos
Limpar o cache regularmente tem um custo que precisa ser considerado. A recompilação pode levar de alguns segundos a minutos dependendo da complexidade da função. Em benchmarks meus, funções JIT simples recompilam em cerca de 2-5 segundos, enquanto modelos maiores podem levar 30 segundos ou mais. Se seu pipeline é sensível a latência de inicialização, essa abordagem pode não ser ideal. Outro ponto: o JAX não implementa um LRU cache ou qualquer política de evicção inteligente por padrão. A limpeza é tudo ou nada. Existem propostas na comunidade para implementar cache com size limit, mas até o momento a solução permanece sendo gerenciamento manual. A versão mais recente do JAX trouxe melhorias no gerenciamento de memória, mas o problema fundamental do cache ilimitado persiste.
Para casos extremos onde o cache continua sendo um problema mesmo após todas as otimizações acima, considere usar jax.profiler para identificar exatamente quais funções estão consumindo mais memória de cache, e então aplicar limpeza seletiva apenas naquelas funções críticas. Muitas vezes, um subconjunto pequeno de funções é responsável pela maioria do crescimento do cache.
Quando evitar JIT completamente
Nem todo mundo precisa de JIT. Se seu modelo é pequeno o suficiente para rodar em eager mode sem problemas de performance, ou se você está em fase de desenvolvimento e depuração, pular a compilação evita todo esse problema. O JAX funciona perfeitamente em modo eager — é apenas mais lento para funções chamadas repetidamente. Uma boa regra prática: use JIT para o loop de treinamento principal, mantenha o resto em eager mode durante desenvolvimento, e ative JIT globalmente apenas quando estiver pronto para production. A lição que eu levanto depois de meses lidando com isso: o cache do JAX é uma faca de dois gumes. Ele dá performance extraordinária quando bem gerenciado, mas pode silently drenar memória do seu processo. Monitorar ativamente, usar argumentos estáticos corretamente, e limpar periodicamente são as três práticas que fazem a diferença entre um pipeline que roda por dias e um que crasha na terceira época.