Apprentissage par renforcement de Gemma4-2B
Effectuez le post-entraînement du grand modèle de langage (LLM) Gemma4-2B (unsloth/gemma-4-E2B-it) par apprentissage par renforcement sur AI Runtime (GPU Serverless). Cet exemple utilise GRPO (Group Relative Policy Optimization) pour apprendre au modèle à résoudre des grilles de Sudoku : le modèle rédige une stratégie de résolution, et des fonctions de récompense l'évaluent sur sa capacité à produire des Solutions valides et sans tricherie. L'exécution s'effectue de bout en bout sur un seul processeur graphique H100 et vous montre comment effectuer les opérations suivantes :
- Charger Gemma4-2B avec des adaptateurs Low-Rank Adaptation (LoRA) à l'aide d'Unsloth pour un apprentissage par renforcement optimisé pour la mémoire
- Définissez un environnement d'apprentissage par renforcement et des fonctions de récompense qui évaluent les stratégies générées par le modèle
- Entraînez la politique avec l'entraîneur GRPO (intégré à Transformer Reinforcement Learning (TRL))
- Exécuter l'inférence avec le modèle entraîné et enregistrer les adaptateurs LoRA
Concepts clés :
- GRPO: un algorithme d'apprentissage par renforcement qui optimise une politique à partir de récompenses relatives au groupe, sans entraîner de modèle de valeur distinct
- LoRA: entraîne un petit ensemble de pondérations d’adaptateur au lieu du modèle complet pour réduire l’utilisation de la mémoire
- Unsloth: une bibliothèque pour l'affinement de LLM et l'apprentissage par renforcement efficaces en mémoire
Cet exemple nécessite la version 6 ou une version ultérieure de l'environnement AI Runtime.
Se connecter au compute serverless GPU
Ce notebook requiert un compute serverless GPU. Pour vous connecter :
- Cliquez sur le sélecteur de compute du notebook dans l’angle supérieur droit, puis sélectionnez Serverless GPU .
- Sur le côté droit, cliquez sur le bouton d’environnement.
- Sélectionnez H100 comme accélérateur .
- Choisissez AI v6 dans l’environnement de base.
- Cliquez sur Appliquer .
Tâche : résoudre le Sudoku par apprentissage par renforcement
L'objectif est d'amener Gemma4-2B à apprendre à résoudre des grilles de Sudoku à l'aide du GRPO. Le modèle élabore une stratégie pour remplir les cellules vides, et les fonctions de récompense l'évaluent sur les placements corrects et la résolution de grilles valides.
Installer les bibliothèques
Cet exemple utilise Unsloth pour le machine learning par renforcement optimisé pour la mémoire sur Gemma4-2B. La cellule suivante l'installe ainsi que ses dépendances.
%pip install unsloth==2026.9.4
%%capture
!pip install --no-deps --upgrade timm # For Gemma 4 vision/audio
Charger Gemma4-2B avec 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
)
Pour un apprentissage par renforcement efficace, cet exemple utilise LoRA pour entraîner les poids des adaptateurs au lieu du modèle complet, ce qui réduit l'utilisation de la mémoire.
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,
)
Mettre en œuvre le jeu du Sudoku
L'environnement Sudoku accepte une stratégie qui renvoie un tuple (row, column, value) pour chaque mouvement.
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)
Testez l’environnement Sudoku :
# 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
Essayez d'effectuer quelques déplacements :
# 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()}")
Une action en dehors de l'espace d'actions autorisé fait entrer le jeu dans l'état failed. Le jeu rejette ensuite les actions suivantes.
Configurez l’environnement de machine learning par renforcement
Définissez des stratégies d'exécution assorties de limites de temps afin d'éviter les boucles infinies.
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()
Appliquez un délai d'expiration de 10 secondes pour permettre des stratégies plus longues tout en évitant les boucles infinies.
@execute_with_time_limit(10)
def execute_strategy(strategy: Callable, game: SudokuGame):
"""Execute strategy with 10 second time limit."""
return _execute_strategy(strategy, game)
Tester avec une stratégie simple :
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())
Exécuter le code généré
Avant d'exécuter une fonction Python générée, vérifiez qu'elle n'accède pas à des variables globales non autorisées ou à des modules externes. Cette vérification permet d'éviter le piratage de récompense (« reward hacking »).
Le code suivant transmet check_python_modules car il n'importe pas de modules externes :
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)
Le code suivant importe numpy, de sorte que check_python_modules le rejette :
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)
Configurer les données et la tâche de machine learning par renforcement
Créez un prompt qui demande au modèle de générer une stratégie de résolution de Sudoku. Vous pouvez adapter ce prompt pour d'autres tâches de machine learning par renforcement.
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)
Générer une réponse de référence avant l'apprentissage par renforcement :
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)
Définir les fonctions de récompense
Définissez extract_function pour extraire une fonction d'un bloc de code Markdown.
Définissez ensuite trois fonctions de récompense :
function_worksrécompense le modèle lorsque la stratégie est une fonction Python valide.no_cheatingpénalise les fonctions qui importent des modules externes.strategy_succeedsrécompenses générées par les stratégies pour effectuer des mouvements valides et résoudre le casse-tête Sudoku.
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 : Valider la fonction
Vérifie si le code généré est du code Python valide et s’exécute correctement.
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
Récompense 2 : Prévenir la triche
Pénalise les fonctions qui importent des bibliothèques externes.
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
Récompense 3 : récompenser les stratégies fructueuses
Récompense les stratégies qui résolvent avec succès les puzzles de Sudoku.
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
Préparer le dataset
Créer le dataset d'entraînement.
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])
Entraînez le modèle
Configurer le formateur GRPO. Pour les autres algorithmes pris en charge et options de configuration, consultez la documentation sur l'apprentissage par renforcement d'Unsloth.
# 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,
)
Exécutez l’entraîneur et surveillez la colonne reward. L’exécution configurée en 60 étapes est une courte démonstration ; les valeurs de récompense peuvent donc rester faibles ou fluctuer tout au long de l’entraînement.
Étape | Loss de training | récompense | 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 l'exécution d'entraînement à 60 étapes :
trainer.train()
Enregistrez les adaptateurs LoRA entraînés :
model.save_pretrained("gemma_4_lora") # Local saving
tokenizer.save_pretrained("gemma_4_lora")
Vérifiez que les poids de l’adaptateur LoRA sont différents de zéro :
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())
Exécuter l’inférence
Générez une réponse avec le modèle entraîné :
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),
)
Étapes suivantes
Maintenant que vous avez post-entraîné Gemma4-2B avec GRPO, vous pouvez :
- Déployer le modèle : Déployer des modèles personnalisés
- Explorez d'autres exemples de post-entraînement : Modèles OSS de post-entraînement (LLM)
- Optimiser l'utilisation du GPU Serverless : Bonnes pratiques pour Runtime
- Résoudre les problèmes : Résoudre les problèmes sur le Runtime IA