Inteligência artificial, sem ruído.
Tutoriais7 min

Como treinar o Gemma-3 para raciocínio matemático com GRPO, Tunix e LoRA

Tutorial completo: treine o Gemma-3 com GRPO, JAX e LoRA para raciocínio matemático estruturado no dataset GSM8K — tudo em uma única GPU. Com código, hiperparâmetros e troubleshooting.

Como treinar o Gemma-3 para raciocínio matemático com GRPO, Tunix e LoRA

Ensinar um modelo de linguagem a raciocinar matematicamente de forma estruturada é um dos desafios mais interessantes em RL para LLMs. Neste tutorial, vamos construir um pipeline completo de GRPO (Group Relative Policy Optimization) usando Tunix, JAX e LoRA para treinar o Gemma-3 (1B) no dataset GSM8K — e o melhor: tudo roda em um único acelerador.

O GRPO é uma variante de reinforcement learning que não precisa de um modelo de referência separado (como o PPO tradicional). Em vez disso, ele amostra múltiplas respostas do mesmo prompt e usa a média do grupo como baseline — mais eficiente em memória e mais simples de implementar.

Pré-requisitos

Requisitos para rodar o tutorial
ComponenteMínimoRecomendado
GPU1× T4 (Colab free)1× A100 ou TPU v3-8
RAM16 GB32 GB
Disco20 GB50 GB
Python3.10+3.11+
ConhecimentoPython básico + JAXExperiência com RL e transformers
Tempo estimado~30 min (setup + treino com 100 steps)

✅ O que você ganha

  • ✅ Pipeline GRPO funcional com uma GPU consumer
  • ✅ Modelo que gera raciocínio estruturado (tags <reasoning> e <answer>)
  • ✅ Funções de recompensa customizáveis (formato + exatidão + extração numérica)
  • ✅ Adaptadores LoRA (rank 32) para fine-tuning eficiente
  • ✅ Checkpoints e logs para TensorBoard

⚠️ O que você NÃO ganha

  • ⚠️ Performance de modelo 70B+ — é um setup educacional com Gemma-3 1B
  • ⚠️ Convergência garantida em 100 steps — GRPO é sensível a hiperparâmetros
  • ⚠️ Pipeline multi-GPU automático — o exemplo usa uma única GPU

Passo 1: Instalando o ambiente

O ecossistema Tunix + JAX é instalado diretamente do GitHub. O setup completo leva de 5 a 8 minutos e inclui Flax, Qwix, TensorFlow (apenas para datasets), Hugging Face Hub e o próprio Tunix:

# Instalação do Tunix + ecossistema JAX
%pip install -q ipywidgets tensorboardX transformers grain nest_asyncio
%pip install -q datasets huggingface_hub "numpy>2"
%pip install -q tensorflow tensorflow_datasets
%pip install -q git+https://github.com/jax-ml/jax
%pip install -q git+https://github.com/google/tunix
%pip install -q git+https://github.com/google/qwix
%pip uninstall -q flax -y
%pip install -q git+https://github.com/google/flax

Por que JAX e não PyTorch? O Tunix é construído sobre JAX pela performance em TPUs e pela capacidade de sharding automático com jax.make_mesh. Além disso, o ecossistema JAX tem suporte nativo a PRNG determinístico, crucial para reprodutibilidade em RL.

Configure o token do Hugging Face (necessário pois o Gemma-3 requer licença):

import os, getpass
os.environ["WANDB_MODE"] = "disabled"
os.environ["TOKENIZERS_PARALLELISM"] = "false"
HF_TOKEN = os.environ.get("HF_TOKEN") or getpass.getpass("Hugging Face token: ")
os.environ["HF_TOKEN"] = HF_TOKEN

Verifique se o JAX reconhece seu acelerador:

import jax
devices = jax.devices()
print(f"JAX backend: {jax.default_backend()} | {len(devices)} device(s): {devices}")
# Esperado: JAX backend: gpu | 1 device(s): [CudaDevice(id=0)]

Passo 2: Formatando prompts e definindo recompensas

O GSM8K exige um formato específico: o modelo deve gerar raciocínio entre tags <reasoning> e a resposta numérica entre <answer>. Isso é crucial porque as funções de recompensa usam regex para extrair a resposta e comparar com o gabarito.

reasoning_start, reasoning_end = "<reasoning>", "</reasoning>"
solution_start,  solution_end  = "<answer>",   "</answer>"

SYSTEM_PROMPT = (
    f"You are given a problem. First, think about the problem and provide your "
    f"reasoning between {reasoning_start} and {reasoning_end}. Then give the final "
    f"answer (just one number) between {solution_start} and {solution_end}."
)

As funções de recompensa são o coração do GRPO. Quatro sinais diferentes guiam o treinamento:

  1. match_format_exactly: recompensa 3.0 se o formato exato (tags reasoning + answer) estiver correto
  2. match_format_approximately: conta tags individuais (0.5 por tag correta, -0.5 por erro)
  3. check_answer: compara a resposta exata (3.0 para match exato, 1.5 para whitespace diferente)
  4. check_numbers: fallback que extrai qualquer número na tag answer (1.5 para match, 0.0 caso contrário)

A combinação de múltiplos sinais de recompensa é o que faz o GRPO convergir mesmo com poucas iterações. O modelo aprende simultaneamente o formato esperado e a exatidão matemática.

Passo 3: Carregando o Gemma-3 com adaptadores LoRA

O download do modelo usa snapshot_download do Hugging Face Hub. Em seguida, aplicamos LoRA com rank 32 e alpha 32.0 usando o Qwix — a biblioteca de adaptadores do Google para JAX:

MODEL_ID = "google/gemma-3-1b-it"
RANK, ALPHA = 32, 32.0

# Download do modelo
local_model_path = snapshot_download(repo_id=MODEL_ID, ignore_patterns=["*.pth"])

# Criar modelo base
model_config = gemma_lib.ModelConfig.gemma3_1b_it()
with mesh:
    base_model = params_safetensors_lib.create_model_from_safe_tensors(
        local_model_path, model_config, mesh)

# Aplicar LoRA nos módulos de atenção e projeção
def apply_lora(base, mesh):
    provider = qwix.LoraProvider(
        module_path=".*q_einsum|.*kv_einsum|.*gate_proj|.*down_proj|.*up_proj|.*attn_vec_einsum",
        rank=RANK, alpha=ALPHA)
    return qwix.apply_lora_to_model(base, provider, **base.get_model_input())

Por que rank 32? Para um modelo de 1B de parâmetros, rank 32 oferece um bom equilíbrio entre capacidade de adaptação (~0.5% de parâmetros treináveis) e eficiência de memória. Ranks mais altos (64-128) consomem mais VRAM sem ganho proporcional em tasks de raciocínio matemático.

Passo 4: Executando o treinamento GRPO

Agora configuramos o GRPO com hiperparâmetros conservadores para uma GPU consumer:

grpo_config = GRPOConfig(
    beta=0.08,          # Penalidade KL (menor = mais exploração)
    epsilon=0.2,        # Clip do ratio (padrão PPO)
    temperature=0.9,    # Diversidade nas gerações
    top_p=1.0,
    top_k=50,
    num_generations=2,  # Amostras por prompt para baseline do grupo
    num_iterations=1,   # Iterações GRPO por batch
    max_steps=100,      # Passos totais de treino
    learning_rate=3e-6,
    max_grad_norm=0.1,
    warmup_steps=10
)

Interpretando os hiperparâmetros:num_generations=2 significa que para cada prompt, o modelo gera 2 respostas e usa a média como baseline — típico de GRPO. beta=0.08 é uma penalidade KL baixa, permitindo que o modelo se afaste mais da política original. temperature=0.9 mantém diversidade suficiente para exploração sem gerar nonsense.

# Criar learner e executar treinamento
learner = GRPOLearner(
    model=model,
    config=grpo_config,
    reward_fns=REWARD_FNS,
    tokenizer=tokenizer,
    mesh=mesh
)

# Loop de treino
for step in range(grpo_config.max_steps):
    batch = next(train_iter)
    metrics = learner.train_step(batch)
    if step % 10 == 0:
        print(f"Step {step}: reward={metrics['reward']:.3f} loss={metrics['loss']:.4f}")

Casos de uso reais

  • Fine-tuning educacional: adaptar LLMs pequenos para tutoria de matemática com raciocínio passo a passo visível
  • Extração estruturada: ensinar modelos a produzir outputs em formato específico para integração com sistemas downstream
  • Prototipagem de RLHF: testar funções de recompensa antes de escalar para modelos maiores com PPO/DPO
  • Pesquisa em alinhamento: estudar como diferentes sinais de recompensa afetam o comportamento de raciocínio

Troubleshooting

❌ JAX vê apenas CPU mesmo com GPU: reinstale o JAX com suporte CUDA: pip install -U "jax[cuda12]" e reinicie o runtime.

❌ Erro de tokenizer: o Gemma-3 usa tokenizer SentencePiece. Se não encontrar, baixe manualmente de storage.googleapis.com/gemma-data/tokenizers/tokenizer_gemma3.model.

❌ OOM (Out of Memory): reduza MAX_PROMPT_LENGTH para 128 e MAX_STEPS para 50. O Gemma-3 1B com LoRA rank 32 deve caber em ~8 GB VRAM.

❌ Recompensa não melhora após 20 steps: aumente temperature para 1.0 (mais exploração) ou reduza beta para 0.04 (menos restrição KL).

❌ Erro 401 no download do modelo: você precisa aceitar a licença do Gemma em huggingface.co/google/gemma-3-1b-it e configurar o token HF.

FAQ

Preciso de uma A100 para rodar? Não. O tutorial foi projetado para rodar em uma T4 (Colab free) com Gemma-3 1B + LoRA rank 32. O consumo de VRAM fica em torno de 6-8 GB.

Posso usar outro modelo? Sim. O Tunix suporta qualquer modelo compatível com seu sistema de configuração. Basta trocar MODEL_ID e ajustar o ModelConfig.

GRPO é melhor que DPO? Depende. GRPO é mais eficiente em memória (sem modelo de referência) e naturalmente explora mais (múltiplas amostras por prompt). DPO é mais simples de implementar e mais estável em datasets pequenos. Para raciocínio matemático com verificação automática, GRPO tende a performar melhor.

Quanto tempo leva para treinar? Com 100 steps e num_generations=2 em uma T4, aproximadamente 25-30 minutos. Em uma A100, menos de 10 minutos.

Posso usar em produção? O setup é educacional. Para produção, você precisaria de mais steps de treino (1000+), validação em held-out set e possivelmente curriculum learning (começar com problemas fáceis e aumentar a dificuldade).

Perspectiva

O GRPO representa uma tendência importante em RL para LLMs: treinamento mais eficiente sem sacrificar a qualidade do alinhamento. À medida que modelos menores como o Gemma-3 se tornam mais capazes, técnicas como GRPO + LoRA permitem que times com recursos limitados façam fine-tuning de raciocínio que antes exigia clusters de GPU. Em 2027, devemos ver bibliotecas como Tunix evoluírem para suportar multi-GPU e multi-node de forma transparente, democratizando ainda mais o acesso ao fine-tuning avançado.


Descubra mais sobre noticiAI

Assine para receber nossas notícias mais recentes por e-mail.

R
Sobre o autorRedação Noticiai

Equipe editorial dedicada a explicar inteligência artificial com clareza, independência e contexto.