Views
No views yet
[!WARNING] 本コード、モデルはすべてChatGPTのみを使用して作成されました。よって、全内容を公開します。
同様に、動作確認や品質保証はできていません。使用する場合は丁寧に動作確認してください。
ca-reward-distill-1B-ja は、日本語の prompt-response ペアに対して単一のスカラー報酬スコア を出力する 1B クラスの Reward Model です。ベースモデルには sbintuitions/sarashina2.2-1b-instruct-v0.1 を使い、教師モデル cyberagent/ca-reward-3b-ja が付与したスコアを回帰蒸留する形で学習しました。AutoModelForSequenceClassification (num_labels=1, regression)sbintuitions/sarashina2.2-1b-instruct-v0.1cyberagent/ca-reward-3b-jastudent_score_norm: 学習時の正規化空間での出力student_score_denorm: score_normalization.json を使って教師スコア空間へ戻した値build_mixed_prompt_dataset.py で、日本語中心の prompt-only データセットを構築generate_teacher_dataset.py で、各 prompt に対して複数候補応答を生成し、cyberagent/ca-reward-3b-ja で採点train_student_rm_regression.py で、教師スコアを回帰ラベルとして生徒 RM を学習evaluate_student_rm_against_teacher.py で、生徒 RM と教師スコアの整合性を評価| Source | Rows |
|---|---|
llm-jp/magpie-sft-v1.0 | 15,000 |
llm-jp/extraction-wiki-ja (v0.3) | 10,000 |
llm-jp/wizardlm8x22b-logical-math-coding-sft-ja | 8,000 |
llm-jp/oasst2-33k-ja | 6,000 |
llm-jp/oasst1-21k-ja | 3,000 |
llm-jp/Synthetic-JP-EN-Coding-Dataset (Japanese prompts only) | 3,000 |
llm-jp/databricks-dolly-15k-ja | 2,000 |
llm-jp/llm-jp-instructions (v1.0/train) | 300 |
llm-jp/oasst2-33k-en → Japanese-answer wrapped | 2,000 |
llm-jp/oasst1-21k-en → Japanese-answer wrapped | 2,000 |
sbintuitions/sarashina2.2-1b-instruct-v0.1 を使い、各 prompt につき 4 candidate responses を生成しています。そのため、教師データセットは 約 205,200 件 の prompt-response-score 行から構成されます。cyberagent/ca-reward-3b-ja により採点され、teacher_score として保存されます。max_length=2048per_device_train_batch_size=2gradient_accumulation_steps=16learning_rate=1e-5num_train_epochs=1gradient_checkpointing=Truesbintuitions/sarashina2.2-1b-instruct-v0.1train_zscoremseprompt_hash ベースの安定ハッシュ分割で validation_ratio=0.02
python /mnt/data/convert_student_rm_to_safetensors.py \
--model-dir ./student_rm_regression_trial/final_model \
--output-dir ./student_rm_regression_trial/final_model_safetensorsevaluate_student_rm_against_teacher.py により、教師 RM に対する近似精度として評価できます。主な評価指標は次の通りです。| Metric | Value |
|---|---|
| Pearson (denorm) | 0.9705 |
| Spearman (denorm) | 0.9631 |
| MAE (denorm) | 0.0881 |
| MSE (denorm) | 0.1929 |
| Top-1 agreement | 0.6833 |
| Pairwise accuracy (micro) | 0.8291 |
| Pairwise accuracy (macro) | 0.8282 |
| Prompt-wise Spearman mean | 0.7172 |
transformers1import json
2import torch
3from pathlib import Path
4from transformers import AutoModelForSequenceClassification, AutoTokenizer
5from huggingface_hub import hf_hub_download
6
7model_id = "kurogane/ca-reward-distill-1B-ja"
8device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
9
10tokenizer = AutoTokenizer.from_pretrained(model_id)
11model = AutoModelForSequenceClassification.from_pretrained(model_id)
12model.eval()
13model.to(device)
14
15prompt = "富士山について短く説明して"
16response_a = "富士山は日本で最も高い山で、静岡県と山梨県にまたがる成層火山です。"
17response_b = "富士山はたぶんどこかにある山です。"
18
19def build_text(prompt: str, response: str) -> str:
20 chat = [
21 {"role": "user", "content": prompt},
22 {"role": "assistant", "content": response},
23 ]
24 return tokenizer.apply_chat_template(chat, tokenize=False, add_generation_prompt=False)
25
26texts = [build_text(prompt, response_a), build_text(prompt, response_b)]
27inputs = tokenizer(
28 texts,
29 return_tensors="pt",
30 truncation=True,
31 max_length=2048,
32 padding=True,
33 add_special_tokens=False,
34).to(device)
35
36with torch.no_grad():
37 scores_norm = model(**inputs).logits.squeeze(-1).float().cpu().tolist()
38
39print("normalized scores:", scores_norm)
40
41# Optional: denormalize to the teacher-score scale
42try:
43 score_stats_path = hf_hub_download(repo_id=model_id, filename="score_normalization.json")
44 with open(score_stats_path, "r", encoding="utf-8") as f:
45 score_stats = json.load(f)
46 if score_stats.get("mode") == "train_zscore":
47 mean = float(score_stats.get("mean", 0.0))
48 std = float(score_stats.get("std", 1.0))
49 scores_denorm = [s * std + mean for s in scores_norm]
50 print("denormalized scores:", scores_denorm)
51except Exception:
52 pass1python score_student_rm_minimal.py \
2 --model-dir ./final_model \
3 --prompt "富士山について短く説明して" \
4 --response "富士山は日本で最も高い山です。"build_mixed_prompt_dataset.pygenerate_teacher_dataset.pytrain_student_rm_regression.pyevaluate_student_rm_against_teacher.pyscore_student_rm_minimal.py