在 AI 執行階段(無伺服器 GPU)上,使用強化學習對 Gemma4-2B(unsloth/gemma-4-E2B-it)大型語言模型(LLM)進行後訓練。 此範例使用 群組相對策略優化(Group Relative Policy Optimization,GRPO) 來教導模型解決數獨謎題:模型撰寫解題策略,並獎勵函數對產生有效且非作弊解答給予評分。 它可在單一 H100 GPU 上完成端對端執行,並向您展示如何:
- 使用 Unsloth 載入 Gemma4-2B 並搭配 Low-Rank Adaptation(LoRA)介面卡,進行記憶體效率高的強化學習
- 定義強化學習環境與獎勵 函數 ,以評分模型生成策略
- 使用 GRPO 訓練器(建立在 Transformer Reinforcement Learning(TRL)之上) 來訓練策略。
- 用訓練好的模型執行推理,並儲存 LoRA 的轉接器
關鍵概念:
- GRPO:一種強化學習演算法,從群體相對獎勵中優化政策,無需訓練獨立的價值模型
- LoRA:訓練一小組適配器權重,而不是完整模型,以減少記憶體使用量
- Unsloth:一個用於記憶體效率佳的 LLM 微調與強化學習的函式庫
Note
此範例需要 AI 執行環境版本 6 或以上。
連接到無伺服器的 GPU 運算
這款筆記型電腦需要無伺服器 GPU 運算。 連結方式:
- 點選筆電右上角的運算選擇器,選擇 無伺服器 GPU。
- 在右側,點擊環境按鈕。
- 選擇 H100 作為 加速器。
- 從基礎環境中選擇 AI v6 。
- 點擊套用。
任務:用強化學習解決數獨
目標是讓 Gemma4-2B 學會使用 GRPO 來解決數獨謎題。 模型設計策略填補空格,獎勵函數則根據正確位置及完成有效謎題給予評分。
安裝程式庫
此範例使用 Unsloth 針對 Gemma4-2B 進行具記憶體效率的強化學習。 下一個儲存格會安裝它及其所需的相依套件。
%pip install unsloth==2026.9.4
%%capture
!pip install --no-deps --upgrade timm # For Gemma 4 vision/audio
使用 Unsloth 載入 Gemma4-2B
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
)
為了更有效率地強化學習,本範例使用 LoRA 訓練適配器權重,而非完整模型,從而減少記憶體使用。
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,
)
實作數獨遊戲
數獨環境接受一個策略,該策略會在每一步回傳一個 (row, column, value) 元組。
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)
測試數獨環境:
# 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
試著做一些動作:
# 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()}")
若動作超出允許的動作空間,遊戲會進入 failed 狀態。 遊戲接著拒絕後續的行動。
建立強化學習環境
執行有時間限制的策略,避免無限循環。
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()
設定 10 秒的超時,允許更長的策略,同時避免無限循環。
@execute_with_time_limit(10)
def execute_strategy(strategy: Callable, game: SudokuGame):
"""Execute strategy with 10 second time limit."""
return _execute_strategy(strategy, game)
用一個簡單的策略測試:
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())
執行產生程式碼
在執行產生的 Python 函式前,請確認它不會存取不允許的全域變數或外部模組。 此檢查有助於防止鑽獎勵機制漏洞。
以下程式碼通過 check_python_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)
以下程式碼匯 numpy入 ,因此 check_python_modules 拒絕:
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)
設定資料與強化學習任務
建立一個提示,指示模型產生數獨解法的策略。 你可以將這個提示詞調整為用於其他強化學習任務。
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)
在強化學習前先產生基線反應:
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)
定義獎勵函數
定義 extract_function 從 Markdown 程式碼區塊中擷取一個函式。
接著定義三個獎勵函數:
-
function_works當策略是有效的 Python 函數時,會獎勵模型。 -
no_cheating懲罰匯入外部模組的函式。 -
strategy_succeeds會獎勵所產生的策略,使其能下出有效步驟並解決數獨謎題。
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
獎勵一:驗證功能
檢查產生的程式碼是否有效 Python,並能順利執行。
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
獎勵二:防止作弊
會懲罰匯入外部函式庫的函式。
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
獎勵三:獎勵成功的策略
獎勵成功解決數獨謎題的策略。
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
準備資料集
建立訓練資料集。
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])
訓練模型
設定 GRPO 訓練器。 關於其他支援的演算法與設定選項,請參閱 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,
)
執行訓練器並監控 reward 欄。 配置的60步跑步是短暫的示範,因此獎勵值可能在訓練過程中保持偏低或波動。
| Step | 訓練損失 | 獎勵 | reward_std | 完成長度 | 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"],
)
開始60步訓練跑:
trainer.train()
儲存已訓練好的 LoRA 轉接器:
model.save_pretrained("gemma_4_lora") # Local saving
tokenizer.save_pretrained("gemma_4_lora")
確認 LoRA 轉接器權重是否非零:
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())
執行推論
用訓練好的模型產生回應:
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),
)
下一步
現在你已使用 GRPO 完成 Gemma4-2B 的後訓練,你可以:
- 部署模型: 部署自訂模型
- 探索更多後訓練範例: 後訓練 OSS 模型(LLM)
- 優化無伺服器 GPU 使用: AI 執行時的最佳實務
- 故障排除問題: AI 執行時故障排除問題