Reinforcement learning of Gemma4-2B
Post-train the Gemma4-2B (unsloth/gemma-4-E2B-it) large language model (LLM) with reinforcement learning on AI Runtime (serverless GPU). This example uses Group Relative Policy Optimization (GRPO) to teach the model to solve Sudoku puzzles: the model writes a solver strategy, and reward functions score it for producing valid, non-cheating solutions. It runs end-to-end on a single H100 GPU and shows you how to:
- Load Gemma4-2B with Low-Rank Adaptation (LoRA) adapters using Unsloth for memory-efficient reinforcement learning
- Define a reinforcement learning environment and reward functions that score the model's generated strategies
- Train the policy with the GRPO trainer (built on Transformer Reinforcement Learning (TRL))
- Run inference with the trained model and save the LoRA adapters
Key concepts:
- GRPO: A reinforcement learning algorithm that optimizes a policy from group-relative rewards, without training a separate value model
- LoRA: Trains a small set of adapter weights instead of the full model to reduce memory use
- Unsloth: A library for memory-efficient LLM fine-tuning and reinforcement learning
This example requires the AI Runtime environment version 6 or above.
Connect to serverless GPU compute
This notebook requires serverless GPU compute. To connect:
- Click the notebook's compute selector in the top right and select Serverless GPU.
- On the right side, click the environment button.
- Select H100 as the Accelerator.
- Choose AI v6 from the base environment.
- Click Apply.
Task: solve Sudoku with reinforcement learning
The goal is to make Gemma4-2B learn to solve Sudoku puzzles using GRPO. The model devises a strategy to fill in empty cells, and the reward functions score it for correct placements and for completing valid puzzles.
Install libraries
This example uses Unsloth for memory-efficient reinforcement learning on Gemma4-2B. The next cell installs it along with its supporting dependencies.
%pip install unsloth==2026.9.4
%%capture
!pip install --no-deps --upgrade timm # For Gemma 4 vision/audio
Load Gemma4-2B with Unsloth
from unsloth import FastVisionModel
import torch
max_seq_length = 4096 # Can increase for longer reasoning traces
lora_rank = 32 # Larger rank = smarter, but slower
gemma4_models = [
# Gemma-4 instruct models:
"unsloth/gemma-4-E2B-it",
"unsloth/gemma-4-E4B-it",
"unsloth/gemma-4-31B-it",
"unsloth/gemma-4-26B-A4B-it",
# Gemma-4 base models:
"unsloth/gemma-4-E2B",
"unsloth/gemma-4-E4B",
"unsloth/gemma-4-31B",
"unsloth/gemma-4-26B-A4B",
] # More models at https://huggingface.co/unsloth
model, tokenizer = FastVisionModel.from_pretrained(
model_name = "unsloth/gemma-4-E2B-it",
max_seq_length = max_seq_length,
load_in_4bit = False, # False for LoRA 16bit
fast_inference = False, # Enable vllm fast inference
)
For efficient reinforcement learning, this example uses LoRA to train adapter weights instead of the full model, which reduces memory use.
model = FastVisionModel.get_peft_model(
model,
r = lora_rank, # Suggested values: 8, 16, 32, 64, or 128
target_modules = [
"q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj",
],
lora_alpha = lora_rank*2, # *2 speeds up training
use_gradient_checkpointing = "unsloth", # Reduces memory usage
random_state = 3407,
)
Implement the Sudoku game
The Sudoku environment accepts a strategy that returns a (row, column, value) tuple for each move.
from dataclasses import dataclass, field
from typing import List, Tuple, Optional
import random
import copy
def _is_valid_placement(board: List[List[int]], row: int, col: int, num: int) -> bool:
"""Check if placing num at (row, col) is valid."""
# Check row
if num in board[row]:
return False
# Check column
if num in [board[r][col] for r in range(9)]:
return False
# Check 3x3 box
box_row, box_col = 3 * (row // 3), 3 * (col // 3)
for r in range(box_row, box_row + 3):
for c in range(box_col, box_col + 3):
if board[r][c] == num:
return False
return True
def _solve_sudoku(board: List[List[int]]) -> bool:
"""Solve sudoku using backtracking (for puzzle generation)."""
for row in range(9):
for col in range(9):
if board[row][col] == 0:
for num in range(1, 10):
if _is_valid_placement(board, row, col, num):
board[row][col] = num
if _solve_sudoku(board):
return True
board[row][col] = 0
return False
return True
def _generate_complete_board(rng: random.Random) -> List[List[int]]:
"""Generate a complete valid Sudoku board."""
board = [[0 for _ in range(9)] for _ in range(9)]
# Fill diagonal 3x3 boxes first (they don't affect each other)
for box in range(3):
nums = list(range(1, 10))
rng.shuffle(nums)
for i in range(3):
for j in range(3):
board[box * 3 + i][box * 3 + j] = nums[i * 3 + j]
# Solve the rest
_solve_sudoku(board)
return board
@dataclass
class SudokuGame:
difficulty: int = 40 # Number of cells to remove (20 = easy, 40 = medium, 50 = hard)
seed: Optional[int] = None
_rng: random.Random = field(init = False, repr = False)
_board: List[List[int]] = field(init = False, repr = False)
_solution: List[List[int]] = field(init = False, repr = False)
_initial_board: List[List[int]] = field(init = False, repr = False)
_moves: int = field(default = 0, init = False, repr = False)
_state: str = field(default = "ongoing", init = False, repr = False)
def __post_init__(self):
self._rng = random.Random(self.seed)
# Generate complete board
complete_board = _generate_complete_board(self._rng)
self._solution = copy.deepcopy(complete_board)
# Remove cells to create puzzle
self._board = copy.deepcopy(complete_board)
cells = [(r, c) for r in range(9) for c in range(9)]
self._rng.shuffle(cells)
for r, c in cells[:self.difficulty]:
self._board[r][c] = 0
self._initial_board = copy.deepcopy(self._board)
self._update_state()
def board(self) -> List[List[int]]:
"""Return current board state."""
return [row[:] for row in self._board]
def initial_board(self) -> List[List[int]]:
"""Return initial puzzle state."""
return [row[:] for row in self._initial_board]
def state(self) -> str:
"""Return game state: 'ongoing', 'success', or 'failed'."""
return self._state
def moves(self) -> int:
"""Return number of moves made."""
return self._moves
def place_number(self, row: int, col: int, num: int) -> bool:
"""Place a number on the board. Returns True if valid move."""
# Validate input
if not (0 <= row < 9 and 0 <= col < 9):
self._state = "failed"
return False
if not (1 <= num <= 9):
self._state = "failed"
return False
# Can't modify initial cells
if self._initial_board[row][col] != 0:
self._state = "failed"
return False
if self._board[row][col] != 0:
self._state = "failed"
return False
# Check if placement is valid
if not _is_valid_placement(self._board, row, col, num):
self._state = "failed"
return False
# Place number
self._board[row][col] = num
self._moves += 1
self._update_state()
return True
def _update_state(self) -> None:
"""Update game state based on current board."""
# Check if puzzle is complete
if all(self._board[r][c] != 0 for r in range(9) for c in range(9)):
# Verify solution is correct
if self._board == self._solution:
self._state = "success"
else:
self._state = "failed"
else:
self._state = "ongoing"
def pretty(self, colors: bool = True) -> str:
"""Pretty print the Sudoku board."""
RESET = "\x1b[0m"
INITIAL = "\x1b[38;5;45m" # Cyan for initial numbers
PLACED = "\x1b[38;5;226m" # Yellow for placed numbers
EMPTY = "\x1b[38;5;239m" # Gray for empty cells
lines = []
lines.append("┌───────┬───────┬───────┐")
for row in range(9):
row_str = "│ "
for col in range(9):
num = self._board[row][col]
if colors:
if num == 0:
row_str += f"{EMPTY}.{RESET}"
elif self._initial_board[row][col] != 0:
row_str += f"{INITIAL}{num}{RESET}"
else:
row_str += f"{PLACED}{num}{RESET}"
else:
row_str += str(num) if num != 0 else "."
if col % 3 == 2:
row_str += " │ "
else:
row_str += " "
lines.append(row_str.rstrip())
if row == 8:
lines.append("└───────┴───────┴───────┘")
elif row % 3 == 2:
lines.append("├───────┼───────┼───────┤")
return "\n".join(lines)
Test the Sudoku environment:
# Create an easy puzzle
game = SudokuGame(difficulty = 30, seed = 42)
print("Initial puzzle:")
print(game.pretty())
print(f"\nState: {game.state()}, Moves: {game.moves()}")
game
Try making some moves:
# Make a valid move
game.place_number(0, 1, 7)
print("\nAfter placing 7 at (1,0):")
print(game.pretty())
print(f"State: {game.state()}, Moves: {game.moves()}")
An action outside the allowed action space causes the game to enter the failed state. The game then rejects subsequent actions.
Set up the reinforcement learning environment
Run strategies with time limits to prevent infinite loops.
from typing import Callable
from unsloth import execute_with_time_limit
def _execute_strategy(strategy: Callable, game: SudokuGame):
"""Execute a strategy function on a Sudoku game."""
assert callable(strategy)
max_moves = 100
valid_moves = 0 # Track successful moves
while game.state() == "ongoing" and valid_moves < max_moves:
try:
board = game.board()
initial = game.initial_board()
result = strategy(board, initial)
# Validate result format
if not isinstance(result, (tuple, list)) or len(result) != 3:
# Invalid format = immediate fail, but return valid moves made
return valid_moves, "failed"
row, col, num = result
# Validate types
if not all(isinstance(x, int) for x in [row, col, num]):
return valid_moves, "failed"
# Try to place number
success = game.place_number(row, col, num)
if success:
valid_moves += 1 # Count this valid move
else:
# Invalid move = game fails, but return valid_moves made so far
return valid_moves, "failed"
except Exception:
return valid_moves, "failed"
if valid_moves >= max_moves and game.state() == "ongoing":
return valid_moves, "failed"
return valid_moves, game.state()
Apply a 10-second timeout to allow longer strategies while preventing infinite loops.
@execute_with_time_limit(10)
def execute_strategy(strategy: Callable, game: SudokuGame):
"""Execute strategy with 10 second time limit."""
return _execute_strategy(strategy, game)
Test with a simple strategy:
def simple_strategy(board, initial):
"""Simple strategy: fill first empty cell with 1."""
for r in range(9):
for c in range(9):
if board[r][c] == 0 and initial[r][c] == 0:
return (r, c, 7)
return (0, 0, 7)
game = SudokuGame(difficulty = 30, seed = 42)
try:
moves, state = execute_strategy(simple_strategy, game)
print(f"Moves: {moves}, State: {state}")
except TimeoutError as e:
print(f"Timed out: {e}")
print(game.pretty())
Run generated code
Before running a generated Python function, verify that it does not access disallowed global variables or external modules. This check helps prevent reward hacking.
The following code passes check_python_modules because it does not import external modules:
from unsloth import check_python_modules, create_locked_down_function
# Test safe code
sample = """
def strategy(board, initial):
for r in range(9):
for c in range(9):
if board[r][c] == 0:
return (r, c, 1)
return (0, 0, 1)
"""
ok, info = check_python_modules(sample)
print("Safe Python code?", ok)
print(info)
The following code imports numpy, so check_python_modules rejects it:
sample = """
def strategy(board, initial):
import numpy as np
return (0, 0, 1)
"""
ok, info = check_python_modules(sample)
print("Safe Python code?", ok)
print(info)
Set up the data and reinforcement learning task
Create a prompt that instructs the model to generate a Sudoku-solving strategy. You can adapt this prompt for other reinforcement learning tasks.
prompt = """
Create a Sudoku solving strategy using only native Python built-in functions without any import statements.
You are given two lists of lists (9x9 grids):
- board: current state (0 means empty)
- initial: starting puzzle (0 means was empty, numbers are fixed)
Return a tuple (row, col, number) for the next move.
- row: 0-8 (row index)
- col: 0-8 (column index)
- number: 1-9 (digit to place)
Only place numbers in cells that are BOTH empty in initial AND empty in board (initial[row][col] == 0 AND board[row][col] == 0)
Use Sudoku rules: no duplicates in rows, columns, or 3x3 boxes.
Enclose the function in a Python Markdown code block.
All helper functions must be inside def strategy. Output only the function.
""".strip()
print(prompt)
Generate a baseline response before reinforcement learning:
text = tokenizer.apply_chat_template(
[{"role": "user", "content": prompt.strip()}],
tokenize = False,
add_generation_prompt = True,
)
from transformers import TextStreamer
print("=" * 50)
print("BASE MODEL OUTPUT (before RL training):")
print("=" * 50)
inputs = tokenizer(
text = text,
add_special_tokens = False,
return_tensors = "pt",
).to("cuda")
text_streamer = TextStreamer(tokenizer, skip_prompt = True)
result = model.generate(**inputs, streamer = text_streamer, max_new_tokens = 128,
use_cache = True, temperature = 1.0, top_p = 0.95, top_k = 64)
Define reward functions
Define extract_function to extract a function from a Markdown code block.
Then define three reward functions:
function_worksrewards the model when the strategy is a valid Python function.no_cheatingpenalizes functions that import external modules.strategy_succeedsrewards generated strategies for making valid moves and solving the Sudoku puzzle.
def extract_function(text):
"""Extract Python function from markdown code blocks."""
if text.count("```") >= 2:
first = text.find("```") + 3
second = text.find("```", first)
fx = text[first:second].strip()
fx = fx.removeprefix("python\n")
fx = fx[fx.find("def"):]
if fx.startswith("def strategy(board, initial):"):
return fx
return None
Reward 1: Validate the function
Checks whether the generated code is valid Python and runs successfully.
def function_works(completions, **kwargs):
"""Reward for generating valid executable Python code."""
scores = []
for completion in completions:
score = 0
response = completion[0]["content"]
function = extract_function(response)
if function is not None:
ok, info = check_python_modules(function)
if function is None or "error" in info:
score = -2.0 # Invalid function
else:
try:
new_strategy = create_locked_down_function(function)
score = 1.0 # Valid function
except:
score = -1.0 # Function has errors
scores.append(score)
return scores
Reward 2: Prevent cheating
Penalizes functions that import external libraries.
def no_cheating(completions, **kwargs):
"""Penalize use of external imports."""
scores = []
for completion in completions:
response = completion[0]["content"]
function = extract_function(response)
if function is not None:
ok, info = check_python_modules(function)
scores.append(1.0 if ok else -20.0) # Heavy penalty for cheating
else:
scores.append(-1.0) # Failed to create function
return scores
Reward 3: Reward successful strategies
Rewards strategies that successfully solve Sudoku puzzles.
import numpy as np
global PRINTER
PRINTER = 0
def strategy_succeeds(completions, **kwargs):
"""Reward valid moves even if strategy eventually fails."""
global PRINTER
scores = []
seed = np.random.randint(10000)
difficulty = 40
for completion in completions:
printed = False
response = completion[0]["content"]
function = extract_function(response)
if PRINTER % 5 == 0:
printed = True
print("\n" + "=" * 60)
print(function)
print("=" * 60)
PRINTER += 1
if function is not None:
ok, info = check_python_modules(function)
if function is None or "error" in info:
scores.append(0)
continue
try:
new_strategy = create_locked_down_function(function)
except:
scores.append(0)
continue
try:
game = SudokuGame(difficulty = difficulty, seed = seed)
valid_moves, game_state = execute_strategy(new_strategy, game)
if valid_moves == difficulty:
game_state = "success"
print(f"\n Valid moves: {valid_moves}, Final state: {game_state}")
if not printed:
print("Strategy:")
print(function[:200] + "..." if len(function) > 200 else function)
print("\nFinal board:")
print(game.pretty())
if game_state == "success":
scores.append(30.0) # Solved the puzzle
elif valid_moves > 0:
# Reward based on valid moves made before failure
# Each valid move is worth 0.2 points
reward = valid_moves * 0.2
scores.append(reward)
else:
scores.append(-2.0) # Failed immediately with no valid moves
except TimeoutError:
print("Timeout")
scores.append(-1.0)
except Exception as e:
print(f"Exception: {str(e)[:100]}")
scores.append(-3.0)
return scores
Prepare the dataset
Create the training dataset.
from datasets import Dataset
dataset = Dataset.from_list([
{
"prompt": [{"role": "user", "content": prompt.strip()}],
"answer": 0,
}
] * 1000)
maximum_length = len(tokenizer.apply_chat_template(
[{"role": "user", "content": prompt.strip()}],
add_generation_prompt = True
))
print(f"Maximum prompt length: {maximum_length}")
print("\nDataset sample:")
print(dataset[0])
Train the model
Configure the GRPO trainer. For other supported algorithms and configuration options, see the Unsloth reinforcement learning documentation.
# Leave room for the prompt (plus 1 token safety margin)
max_completion_length = max_seq_length - (maximum_length + 1)
from trl import GRPOConfig, GRPOTrainer
training_args = GRPOConfig(
temperature = 1.0,
learning_rate = 5e-5,
weight_decay = 0.001,
warmup_ratio = 0.1,
lr_scheduler_type = "linear",
optim = "adamw_8bit",
logging_steps = 1,
per_device_train_batch_size = 1,
gradient_accumulation_steps = 2, # Increase to 4 for smoother training
num_generations = 2, # Decrease if out of memory
max_completion_length = max_completion_length,
# num_train_epochs = 1, # Set to 1 for a full training run
max_steps = 60,
save_steps = 100,
report_to = "none", # Can use Weights & Biases, TrackIO
output_dir = "outputs",
epsilon = 0.2,
epsilon_high = 0.28, # one sided
delta = 1.5, # two sided
loss_type = 'bnpo',
mask_truncated_completions = True
# For optional training + evaluation
# fp16_full_eval = True,
# per_device_eval_batch_size = 4,
# eval_accumulation_steps = 1,
# eval_strategy = "steps",
# eval_steps = 1,
)
Run the trainer and monitor the reward column. The configured 60-step run is a short demonstration, so reward values can remain low or fluctuate throughout training.
Step | Training Loss | reward | reward_std | completion_length | kl |
|---|---|---|---|---|---|
1 | 0.000000 | 0.125000 | 0.000000 | 200.000000 | 0.000000 |
2 | 0.000000 | 0.072375 | 0.248112 | 200.000000 | 0.000000 |
3 | 0.000000 | -0.079000 | 0.163776 | 182.500000 | 0.000005 |
# For optional training + evaluation
# new_dataset = dataset.train_test_split(test_size = 0.01)
trainer = GRPOTrainer(
model = model,
processing_class = tokenizer,
reward_funcs = [
function_works,
no_cheating,
strategy_succeeds,
],
args = training_args,
train_dataset = dataset,
# For optional training + evaluation
# train_dataset = new_dataset["train"],
# eval_dataset = new_dataset["test"],
)
Start the 60-step training run:
trainer.train()
Save the trained LoRA adapters:
model.save_pretrained("gemma_4_lora") # Local saving
tokenizer.save_pretrained("gemma_4_lora")
Verify that the LoRA adapter weights are nonzero:
from safetensors import safe_open
tensors = {}
with safe_open("gemma_4_lora/adapter_model.safetensors", framework = "pt") as f:
# Verify both A and B are non zero
for key in f.keys():
tensor = f.get_tensor(key)
n_zeros = (tensor == 0).sum() / tensor.numel()
assert(n_zeros.item() != tensor.numel())
Run inference
Generate a response with the trained model:
text = tokenizer.apply_chat_template(
[{"role": "user", "content": prompt.strip()}],
tokenize = False,
add_generation_prompt = True,
)
from transformers import TextStreamer
_ = model.generate(
**tokenizer(images = None,text = text, return_tensors = "pt").to("cuda"),
temperature = 1.0,
max_new_tokens = 512,
streamer = TextStreamer(tokenizer, skip_prompt = False),
)
Next steps
Now that you've post-trained Gemma4-2B with GRPO, you can:
- Deploy the model: Deploy custom models
- Explore more post-training examples: Post-training OSS models (LLMs)
- Optimize serverless GPU usage: Best practices for AI Runtime
- Troubleshoot issues: Troubleshoot issues on AI Runtime