Views
No views yet
1
2%%capture
3# Skip restarting message in Colab
4import sys; modules = list(sys.modules.keys())
5for x in modules: sys.modules.pop(x) if "PIL" in x or "google" in x else None
6
7!pip install unsloth vllm
8!pip install --upgrade pillow
9# If you are running this notebook on local, you need to install `diffusers` too
10# !pip install diffusers
11# Temporarily install a specific TRL nightly version
12!pip install git+https://github.com/huggingface/trl.git@e95f9fb74a3c3647b86f251b7e230ec51c64b72b1from unsloth import FastLanguageModel, PatchFastRL
2PatchFastRL("GRPO", FastLanguageModel)
3
4import re
5import torch
6from datasets import load_dataset, Dataset
7from transformers import AutoTokenizer, AutoModelForCausalLM
8from peft import LoraConfig
9from trl import GRPOConfig, GRPOTrainer
10from unsloth import is_bfloat16_supported1model_id="llm-jp/llm-jp-3-13b-instruct2"
2adpter_id="morizon/llm-jp-3-13b-instruct2-grpo-MATH-lighteval_step1000_lora"
3
4# --- モデルの読み込みと LoRA 適用 ---
5max_seq_length = 1024 # 推論トレースの最大長
6lora_rank = 64 # LoRA のランク(推奨値:64)
7
8# FastLanguageModel 経由でモデルとトークナイザーを読み込み
9# ※ モデル名は使用するものに合わせてください
10model, tokenizer = FastLanguageModel.from_pretrained(
11 model_name=model_id,
12 max_seq_length=max_seq_length,
13 load_in_4bit=True, # 4bit量子化(LoRAファインチューニング時は設定に注意)
14 fast_inference=True, # vLLM 高速推論を有効化
15 max_lora_rank=lora_rank,
16 gpu_memory_utilization=0.7,
17)
18
19# LoRA (PEFT) を適用
20model = FastLanguageModel.get_peft_model(
21 model,
22 r=lora_rank,
23 target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
24 lora_alpha=lora_rank,
25 use_gradient_checkpointing="unsloth",
26 random_state=3407,
27)1# --- プロンプトとデータセットの準備 ---
2# 推奨:システムプロンプトを排除し、ユーザープロンプトに全指示を統合
3USER_INSTRUCTION = (
4 "Please ensure your response begins with \"<reasoning>\n\". "
5 "Please reason step by step, and put your final answer within \\boxed{}. "
6)
7
8# テストデータの例(リスト形式)
9test_data = [
10 {"id": 0, "text": "$x^{-1}>x$を満たす正の整数$x$の個数を求めなさい。", "gold": "0", "response": "", "type": "Algebra", "level": "Level 2"},
11 ##評価したいテストデータを入力してください
12]
13
14def extract_boxed_answer_rev(text: str) -> str:
15 """
16 テキスト中から最初の \boxed{...} の中身(ネストを考慮)を抽出する。
17 例: r"\boxed{\frac{\pi}{6}}" -> "\frac{\pi}{6}"
18 """
19 key = r"\boxed{"
20 start_idx = text.find(key)
21 if start_idx == -1:
22 return ""
23 # \boxed{ の直後の位置を開始位置とする
24 start_idx += len(key)
25 brace_count = 1 # 最初の { を既にカウント
26 i = start_idx
27 while i < len(text) and brace_count > 0:
28 if text[i] == "{":
29 brace_count += 1
30 elif text[i] == "}":
31 brace_count -= 1
32 i += 1
33 # i-1 が閉じ括弧に対応する位置
34 return text[start_idx:i-1].strip()
35
36from vllm import SamplingParams
37
38correct = 0
39total = len(test_data)
40
41# 正解ケースと誤答ケースを記録するリスト
42correct_cases = []
43incorrect_cases = []
44
45for item in test_data:
46 # プロンプト生成(USER_INSTRUCTION を先頭に追加)
47 prompt = USER_INSTRUCTION + item["text"]
48 text = tokenizer.apply_chat_template([
49 {"role": "user", "content": prompt},
50 ], tokenize=False, add_generation_prompt=True)
51
52 # 推論実行
53 sampling_params = SamplingParams(
54 temperature=0.6,
55 max_tokens=2048,
56 )
57 output = model.fast_generate(
58 text,
59 sampling_params=sampling_params,
60 lora_request = model.load_lora(adpter_id),
61 # lora_request = model.load_lora("grpo_saved_lora"),
62 )[0].outputs[0].text
63
64 # \boxed{...} の中身を抽出する関数で回答を取得
65 boxed_answer = extract_boxed_answer_rev(output)
66
67 # 結果の表示用
68 print("\n----------Test ID:", item["id"], "----------")
69 print("Prompt:")
70 print(prompt)
71 print("\nLLM Output:")
72 print(output)
73 print("\nExtracted Answer:")
74 print(boxed_answer)
75 print("Gold Answer:", item["gold"])
76
77 # 抽出回答と gold の一致で正解判定
78 if boxed_answer == item["gold"]:
79 correct += 1
80 correct_cases.append({
81 "id": item["id"],
82 "prompt": prompt,
83 "LLM_output": output,
84 "extracted_answer": boxed_answer,
85 "gold": item["gold"]
86 })
87 else:
88 incorrect_cases.append({
89 "id": item["id"],
90 "prompt": prompt,
91 "LLM_output": output,
92 "extracted_answer": boxed_answer,
93 "gold": item["gold"]
94 })
95
96# 正解ケースの表示
97print("\n========== 正解ケース ==========")
98for case in correct_cases:
99 print("\nTest ID:", case["id"])
100 print("Prompt:")
101 print(case["prompt"])
102 print("LLM Output:")
103 print(case["LLM_output"])
104 print("Extracted Answer:", case["extracted_answer"])
105 print("Gold Answer:", case["gold"])
106 print("-" * 40)
107
108# 誤答ケースの表示
109print("\n========== 誤答ケース ==========")
110for case in incorrect_cases:
111 print("\nTest ID:", case["id"])
112 print("Prompt:")
113 print(case["prompt"])
114 print("LLM Output:")
115 print(case["LLM_output"])
116 print("Extracted Answer:", case["extracted_answer"])
117 print("Gold Answer:", case["gold"])
118 print("-" * 40)
119
120accuracy = correct / total * 100
121print("\nOverall Accuracy: {}/{} ({:.2f}%)".format(correct, total, accuracy))
122