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
| Componente | Mínimo | Recomendado |
|---|---|---|
| GPU | 1× T4 (Colab free) | 1× A100 ou TPU v3-8 |
| RAM | 16 GB | 32 GB |
| Disco | 20 GB | 50 GB |
| Python | 3.10+ | 3.11+ |
| Conhecimento | Python básico + JAX | Experiê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/flaxPor 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_TOKENVerifique 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:
- match_format_exactly: recompensa 3.0 se o formato exato (tags reasoning + answer) estiver correto
- match_format_approximately: conta tags individuais (0.5 por tag correta, -0.5 por erro)
- check_answer: compara a resposta exata (3.0 para match exato, 1.5 para whitespace diferente)
- 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.



