Skip to content

Repository files navigation

Teaching Gemma to Reason: A Full Tunix Post-Training Pipeline Under Real TPU Constraints

Abstract

Most open-weight language models can generate correct answers, but they rarely expose the causal chain of logic behind those answers. This lack of "transparent reasoning" limits trust, interpretability, and downstream usefulness. In this project, I built a complete, reproducible post-training pipeline using Google’s Tunix (JAX-native) framework to teach Gemma 2-2B to produce explicit, structured reasoning traces before providing final answers.

The work was performed entirely within the hard constraints of a single Kaggle TPU v3-8 session (limited time, 16GB HBM per core, KV-cache limits, and no background processes). The pipeline consists of three carefully designed tiers:

  1. Tier-1 (SFT): Supervised Fine-Tuning to learn strict reasoning syntax.
  2. Tier-2 (RL): Reinforcement Learning via GRPO (Generalized Reward Policy Optimization) to refine reasoning quality without a memory-heavy critic.
  3. Tier-3 (Self-Distillation): A "Best-of-N" rejection sampling loop to densify high-quality reasoning data.

Every design decision- from dynamic token bounding to iterative reward shaping- was driven by practical TPU limitations rather than idealized assumptions. The final model reliably outputs a <reasoning> trace followed by the <answer> across math, science, coding, and logic tasks.


Motivation

Reasoning quality does not emerge automatically from 2B-parameter models; it requires explicit formatting constraints, rewards that penalize hallucination, and strict length control to fit in edge-device contexts.

This project demonstrates that Tunix + Gemma allows us to solve this in practice, specifically tackling the JAX memory management and TPU throughput challenges that often block complex post-training workflows.


Model & Infrastructure

All training stages are executed within a single Kaggle notebook session.

  • Base Model: google/gemma-2-2b (Base)
  • Framework: Tunix (JAX-native). Chosen for XLA compilation benefits.
  • Hardware: Single Kaggle TPU v3-8 (8 cores).
  • Precision: bfloat16 (native TPU support).
  • Training Style: LoRA (Low-Rank Adaptation) targeting q_proj, v_proj, k_proj, o_proj.
  • Constraint: The entire pipeline - SFT, RL, and Distillation - had to fit within the 9-hour session limit.

Data Engineering: The 50k "GlassBox" Corpus

The backbone of this project is a meticulously curated, 50,000-sample dataset, built from scratch over a 10–15 day intensive data engineering sprint. We realized early on that "more data" isn't enough; we needed balanced data that forces the model to reason across different modalities.

Tier-1 Dataset Construction (The Foundation) We aggregated and sanitized multiple public high-quality sources- specifically leveraging the Sagi reasoning datasets, Nvidia’s help/quality datasets, and distilled traces from DeepSeek R1. We didn't just dump these into a loader; we balanced them into a perfectly rounded 50k corpus comprising exactly 10k examples for each of five critical domains:

  1. Coding: Algorithmic logic and syntax correction.
  2. Science: Multi-step deduction and fact retrieval.
  3. Math: Numeric reasoning (GSM8K-style) with strict step validation.
  4. General Reasoning: Logic puzzles and commonsense inference.
  5. Long-Context ("Big Para"): Reading comprehension and summarization tasks designed to test memory retention over long sequences.

Every single row was reformatted into our strict <reasoning> internal chain-of-thought XML schema.

Tier-2 & Tier-3 Strategy (Avoiding Leakage) For the subsequent stages, we were hyper-careful about data hygiene to prevent the model from memorizing answers.

  • For Tier-2 (RL): We selected a pristine, 5,000-row subset from the Sagi/Public datasets, cleaning it manually to ensure the ground truth was unambiguous. This stability was crucial for the GRPO reward signal.
  • For Tier-3 (Self-Distillation): We returned to the Sagi source but strictly selected two different, non-overlapping subsets of rows. By using completely unseen prompts for the "Best-of-N" generation loop, we ensured the model was optimizing its general reasoning ability, not just overfitting to the Tier-2 training set.

Methodology

Tier-1: Supervised Fine-Tuning (SFT)

Goal: Teach the model the syntax, structure, and discipline of reasoning traces.

Dataset & Formatting: I utilized a small, high-quality supervised dataset (a few thousand examples) derived from CoT sources. Each example was formatted to strictly follow this XML schema:

<reasoning>
... step-by-step logic ...
</reasoning>
<answer>
... final result ...
</answer>

Implementation Details:

  • Tokenizer Strategy: Rather than modifying the tokenizer vocabulary (which can cause instability), format compliance was enforced purely through supervised examples and later reinforced during RL.
  • Loss Calculation: Loss was primarily focused on the model output region. In practice, minor prompt leakage occasionally occurred, which was handled via inference-time post-processing.
  • Outcome: The model successfully learned to separate reasoning from answers, providing a stable foundation for RL.

Tier-2: Reinforcement Learning (GRPO)

Why GRPO over PPO? Proximal Policy Optimization (PPO) requires a Critic Model (Value Function). In a memory-constrained TPU environment, loading a second copy of Gemma 2-2B alongside the Actor and Optimizer states results in immediate Out-Of-Memory (OOM) errors.

GRPO (Generalized Reward Policy Optimization) eliminates the Critic by estimating the baseline from the group mean of sampled outputs.

The Reward Function (The "Judge"): Since I could not run a heavy neural reward model, I implemented an iterative, rule-based reward function that evolved as I encountered failure modes:

  1. Format Gate: Strict checks for <reasoning> and <answer> tags.
  2. Correctness: Exact-match checks for math/coding answers.
  3. Safety & Quality: Successive safeguards were added to penalize empty reasoning, repetition loops, and numeric inconsistencies.

TPU Optimization: To prevent XLA recompilation (which takes minutes), I implemented bucketing and padding, ensuring the TPU computation graph remained static even as generation lengths varied.

Tier-3: LLM-as-Judge Self-Distillation

Tier-3 creates a feedback loop where the model learns from its own best attempts.

Process:

  1. Prompt Sampling: Loaded ~1,000 distinct prompts filtered by length to prevent cache overflow.
  2. Generation: The Tier-2 model generated 4 candidates per prompt.
  3. Filtering: I applied strict token-length filtering and quality gating.
  4. Selection: Approximately 50% of the generated candidates were retained- prioritizing high-confidence, correctly formatted reasoning traces over volume.

Result: The retained "gold standard" examples were merged with the original Tier-1 dataset for a final SFT pass. This ensures the model "internalizes" its own best strategies without forgetting the original instruction following.


Engineering Challenges & Solutions

1. The KV-Cache Bottleneck

On TPUs, the Key-Value (KV) cache for attention grows linearly with sequence length.

  • Problem: Naive generation logic crashed the 16GB memory when batches were large.
  • Solution: I implemented dynamic bounding. During RL and self-distillation, generation length was dynamically calculated based on KV-cache capacity, prompt length, and a safety margin. This prevented OOMs while maximizing available reasoning space.

2. The "Empty Thinking" & Repetition Loops

Early in Tier-2, the model sometimes learned to output:

  • Empty tags: <reasoning></reasoning> (to get to the answer faster).
  • Repetition: Looping the same phrase to satisfy length requirements.
  • Fix: The reward function was updated to return negative rewards for reasoning traces shorter than 20 characters or those exhibiting high n-gram repetition.

Results & Evaluation

We evaluated the final model on a holdout set of math and logic puzzles.

  • Format Compliance: Near 100% adherence to the XML schema.
  • Reasoning Capability: The model consistently generates reasoning traces before answering.
  • Emergent Behavior: In several cases, the model demonstrated emergent correction behavior, revising intermediate steps in the reasoning trace before producing a final answer.

Known Limitations

While the model reliably produces structured reasoning traces, several limitations remain due to the constrained training environment:

  • Formal Logic Traps: Certain complex quantifier-based syllogisms remain challenging without a symbolic logic reward.
  • Answer Verbosity: The model sometimes prioritizes reasoning quality over answer completeness, occasionally producing minimal final answers.
  • Prompt Leakage: Residual prompt tokens (e.g., closing tags) occasionally appear and are handled via inference-time sanitation.

These limitations reflect conscious trade-offs made to fit the entire pipeline within a single Kaggle TPU session and prioritize reasoning structure over perfect linguistic polish.


Conclusion

This project proves that you do not need H100 clusters to train reasoning models. By leveraging Tunix's JAX optimizations, GRPO's memory efficiency, and Gemma 2-2B's architectural strengths, I fitted a complete post-training pipeline into a free Kaggle instance.

The final model is a transparent thinking engine that shows its work. It represents a practical, reproducible step toward accessible, trustworthy AI.

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages