1
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

医療LLMのRL post-trainingで CHORD を ms-swift に載せて回した

1
Last updated at Posted at 2026-08-27

TL;DR

  • CHORD (L = (1-μ)·L_GRPO + μ·L_SFT) を ms-swift release/3.11 に実装し、Qwen3-Next-80B-A3B-Instruct を医療 MCQ で学習した
  • 実装は ms-swift 内の 2 ファイル (grpo_trainer.py / megatron_args.py) の置き換えのみで完結
  • ハイパラは μ_peak=0.1, μ_valley=0.01, warmup=0, decay=500, φ=off
  • φ は論文・公式実装が on 推奨だが、予備走で精度が伸びず off に戻した (SFT データが評価フォーマットと乖離していたことが原因と推定)
  • 結果: 医師国試 89.5 → 90.4 (+0.9)、専門医試験 69.0 → 70.0 (+1.0, McNemar p=0.119)、フォーマット崩壊 0 件

はじめに

RAMENチームとして医師国家試験のデータで Qwen3 系列に post-training をかけた一本になります。今回のテーマは、CHORD というアルゴリズムを ms-swift (Megatron-LM backend) に実装して 80B モデルで回してみた話です。

対象タスクは医師国試の多肢選択問題 (医療 MCQ)、ベースモデルは Qwen3-Next-80B-A3B-Instruct。学習は 8 ノード × 8 GPU (計 64 GPU) の SLURM ジョブで実行しました。

参照論文:

Xie et al. (2025) "On-Policy RL Meets Off-Policy Experts: Harmonizing Supervised Fine-Tuning and Reinforcement Learning via Dynamic Weighting" arXiv:2508.11408。公式実装は Trinity-RFT/examples/mix_chord

CHORD の loss は以下です:

L_CHORD = (1 - μ_t) · L_GRPO + μ_t · L_SFT

μ_t はステップに応じて動くスケジューラで、token 単位の重み φ もオプションで載せられます。今回のプロジェクトでは、最初 φ=on で予備走したのち、SFT データ側の要因で φ=off に戻して本走を行いました (詳細は「3.3 φ 関数」の節)。


目次

  1. CHORD は何をしているのか
  2. ms-swift のどこを変えたか
  3. アルゴリズム本体 (loss / μ / φ)
  4. SFT データの作り方
  5. 実験結果と考察
  6. まとめ

1. CHORD は何をしているのか

CHORD framework overview (Figure 3, arXiv:2508.11408)

Figure 3: An overview of the proposed CHORD framework that unifies SFT and RL, featuring a global coefficient μ and a token-wise weighting function ϕ(·).

1.1 loss 形式

loss は以下の凸結合です:

L_CHORD = (1 - μ_t) · L_GRPO + μ_t · L_SFT

μ_t は 0〜1 の実数で、学習ステップに応じて動くスケジューラです。学習序盤で μ を大きく (SFT 側の重みを大きく) 設定し、後半に μ を減衰させて RL 側の重みを大きくする、という使い方が想定されています。

1.2 論文の 2 バリアント: Chord-μ と Chord-φ

CHORD 論文には 2 種類のバリアントが提案されています。

バリアント 特徴 φ 設定
Chord-μ μ を peak (0.9) から valley (0.05) へ decay off
Chord-φ μ を 0.1 に固定、token 単位の φ で制御 on

論文 Table 1 では Chord-φ が Chord-μ を上回っており、公式実装 mix_chord.yaml でも enable_phi_function: true がデフォルトになっています。

CHORD Table 1 (arXiv:2508.11408)

Table 1: Performance comparisons on reasoning problems and tool-use tasks.

本プロジェクトの設定は φ=off (μ を decay) の Chord-μ 系ですが、μ_peak は 0.1 と Chord-φ 側の値に近い形になっています。この設定に至った経緯は「3.3 φ 関数」の節で述べます。


2. ms-swift のどこを変えたか

2.1 前提のバージョン

コンポーネント バージョン
ms-swift modelscope/ms-swiftrelease/3.11 ブランチ
Singularity image swift3.9.3
Megatron-LM core_r0.14.0

Singularity image と Megatron-LM は公式のものをそのまま利用しています。CHORD の実装追加は ms-swift 側の 2 ファイルに閉じています 。

2.2 修正した 2 ファイル (CHORD 用)

ms-swift 内の以下 2 ファイルを丸ごと置き換えています。

ms-swift 内のファイル 修正内容
swift/megatron/trainers/grpo_trainer.py CHORD 本体を追加。GRPOTrainer に対して、SFT データローダ、μ スケジューラ (_get_chord_mu)、φ 重み計算 (_compute_phi_weights)、GRPO バッチへの SFT サンプル連結 (_merge_chord_into_batch)、loss_func への (1-μ)·L_GRPO + μ·L_SFT 合成を追加
swift/megatron/argument/megatron_args.py CLI 引数を追加。chord_sft_dataset / chord_sft_per_device_train_batch_size / chord_mu_peak / chord_mu_valley / chord_mu_warmup_steps / chord_mu_decay_steps / chord_enable_phi_function の 7 引数と validation (_init_chord)

2.3 CLI の設計

CHORD を有効化するときの CLI:

--rlhf_type grpo \
--loss_type grpo \
--chord_sft_dataset ${SFT_DATASET_JSONL} \
--chord_mu_peak 0.1 \
--chord_mu_valley 0.01 \
--chord_mu_warmup_steps 0 \
--chord_mu_decay_steps 500 \
--chord_enable_phi_function false \

--rlhf_type chord のような専用フラグは追加していません。--rlhf_type grpo 実行時に --chord_sft_dataset が渡されていれば CHORD モードに切り替わる形にしています。GRPOTrainer.__init__ の中で以下のように判定します:

# swift/megatron/trainers/grpo_trainer.py の GRPOTrainer.__init__ に追加
self.chord_enabled = args.chord_sft_dataset is not None
self.chord_sft_dataset = args.chord_sft_dataset
self.chord_mu_peak = args.chord_mu_peak
self.chord_mu_valley = args.chord_mu_valley
self.chord_mu_warmup_steps = args.chord_mu_warmup_steps
self.chord_mu_decay_steps = args.chord_mu_decay_steps
self.chord_enable_phi_function = args.chord_enable_phi_function

3. アルゴリズム本体

3.1 loss のコア

swift/megatron/trainers/grpo_trainer.pyloss_func() に追加した CHORD 部分:

# 追加:SFT損失計算
sft_loss = None
if num_sft_samples > 0 and chord_mu > 0 and sft_per_token_logps is not None:
    sft_loss = -(sft_per_token_logps * sft_completion_mask).sum() / sft_completion_mask.sum().clamp(min=1.0)

    # φ関数の適用(オプション)
    if self.chord_enable_phi_function:
        phi_weights = self._compute_phi_weights(sft_per_token_logps, sft_completion_mask)
        sft_loss = -(sft_per_token_logps * phi_weights * sft_completion_mask).sum() / \
                sft_completion_mask.sum().clamp(min=1.0)

# ★追加: 最終損失の計算(CHORD混合)
if sft_loss is not None and chord_mu > 0:
    loss = (1 - chord_mu) * grpo_loss + chord_mu * sft_loss
else:
    loss = grpo_loss

構造:

  • grpo_losssft_loss を独立に集計し、最後に凸結合
  • sft_loss は token-level NLL の平均 (-(logps * mask).sum() / mask.sum())
  • chord_mu > 0 かつ SFT サンプルがある場合のみ SFT loss を混ぜる。それ以外は素の GRPO loss
  • wandb には chord/mu, chord/grpo_loss, chord/sft_loss の 3 メトリクスを流す

3.2 μ スケジューラ

μ_t は 3 相スケジューラで動きます (_get_chord_mu()):

def _get_chord_mu(self) -> float:
    current_step = self._step
    warmup_steps = self.chord_mu_warmup_steps
    decay_steps = self.chord_mu_decay_steps or args.train_iters or 1000
    mu_peak = self.chord_mu_peak
    mu_valley = self.chord_mu_valley

    # Warmup phase: 0 → mu_peak
    if current_step < warmup_steps:
        return mu_peak * (current_step / warmup_steps) if warmup_steps > 0 else mu_peak

    # Decay phase: mu_peak → mu_valley
    decay_start = warmup_steps
    decay_end = warmup_steps + decay_steps
    if current_step >= decay_end:
        return mu_valley
    decay_progress = (current_step - decay_start) / max(decay_steps, 1)
    return mu_peak - (mu_peak - mu_valley) * decay_progress

今回の設定:

パラメータ
μ_peak 0.1
μ_valley 0.01
warmup_steps 0
decay_steps 500

3.3 φ 関数: 最初は on で回したが、精度が伸びず off に戻した

φ の定義

CHORD-φ は SFT loss を token ごとに φ(p_t) = p_t (1 - p_t) で重み付けする仕組みです (_compute_phi_weights()):

def _compute_phi_weights(self, per_token_logps, completion_mask):
    if not self.chord_enable_phi_function:
        return torch.ones_like(per_token_logps)
    # 公式のφ定義: φ = p_t * (1 - p_t)
    probs = torch.exp(per_token_logps.clamp(max=0))
    phi_weights = probs * (1 - probs)
    phi_weights = phi_weights / (phi_weights.mean() + 1e-8)  # 正規化
    return phi_weights * completion_mask

p が 0.5 付近の token に SFT の勾配が大きく効き、p が 0 または 1 に近い token では小さくなります。

論文と公式実装

  • 論文 Table 1: Chord-φ (φ=on) が Chord-μ (φ=off) を数学タスク・Tool-use タスクの両方で上回る
  • 公式 mix_chord.yaml: enable_phi_function: true がデフォルト

試走の結果

最初の試走では --chord_enable_phi_function true で回しました。結果として、φ=on の run は φ=off の run より evaluation accuracy が伸びませんでした。

原因の切り分け

φ 自体の動作 (p の値に応じて SFT の勾配に重み付け) は仕様通りに動作しています。

一方、今回の CHORD 用 SFT データは後述の通り「MCQ を自然言語質問に変換 → 5 部構成の医療文で回答」という形式で、評価タスク (MCQ + [ans] タグ) とフォーマットが異なります。

φ=on の場合、SFT loss で焼き付けられる token はモデルがまだ迷っている箇所に集中します。SFT データが評価と乖離している状態では、この選択的焼き付けが評価精度を下げる方向に働いた可能性があります。

φ=off の場合、SFT の勾配は全 token に一様にかかるため、乖離の影響は薄く広がります。

本走の設定

以上の切り分けから、本走 (train_chord.sh 系) はすべて --chord_enable_phi_function false で固定しました。


4. SFT データの作り方

4.1 パイプライン

SFT データは 2 段パイプラインで合成しました。

Step 入力 出力 温度 目的
Step 1 MCQ 問題文 自然な質問文 0.0 「〜について教えてください」形式にリライト
Step 2 質問文 5 部構成の医療文 0.3 要点 / 原因 / 対応 / 受診目安 / 追加質問 の順で回答

Step 1 のシステムプロンプト (抜粋):

医学の選択式問題を、ユーザーが日常的にしそうな自然な質問に変換してください。
【ルール】
- 選択肢の列挙形式(A/B/C/D等)は使わない
- 正答を質問文に含めない
- 医学的テーマは保持する

Step 2 のシステムプロンプト (抜粋):

【回答構造】
1) 要点(3行以内)
2) 考えられる原因(3〜5つ)
3) 自分でできる対応
4) 受診の目安
5) 追加確認質問(最大3つ)
最後に「医療機関での相談が必要な場合があります」と添える。

出力形式:

{"messages": [{"role": "user", "content": 質問}, {"role": "assistant", "content": 構造化回答}]}

4.2 validation

「100 字以上」かつ「要点/原因/対応/受診/医療機関」の 5 マーカーのうち 2 個以上出現、をチェックしています。

4.3 SFT データと評価データのフォーマット差

項目 SFT データ 評価データ
入力 自然言語の質問 MCQ (選択肢 a〜e 列挙)
出力 5 部構成の構造化医療文 (数百字) 考察 + [ans]選択肢[/ans] タグ

5. 実験結果と考察

5.1 数値

Base 80B に対する CHORD の差分:

データセット Base CHORD
医師国試 (igakuqa, N=1122) 89.5 90.4 (+0.9)
専門医試験 (specialist, N=3757) 69.0 70.0 (+1.0)

フォーマット遵守率: [ans]...[/ans] パース失敗率 0.0%。

McNemar 検定 (Base vs CHORD, specialist): p = 0.119 (n.s.)

診療科別で CHORD がリードした科: 整形外科 (+2.8), 心臓外科 (+2.4), 救急 (+1.7)

5.2 考察: SFT データのフォーマット乖離

CHORD の SFT パートは、SFT データのフォーマット (自然言語 QA) を通じて学習に寄与します。今回のケースではこれが評価タスク (MCQ + [ans]) と乖離しているため、SFT loss を混ぜたことによる正の効果が制限された可能性があります。

構造:

  1. SFT loss は「自然言語 QA スタイルで回答するクセ」を学習させる
  2. 評価タスクと GRPO 側は「MCQ で [ans]a[/ans] を返す」ことを要求する (RL 側でフォーマットは固定済み)
  3. 学習方向と評価方向が別のフォーマットに向いている

「3.3 φ 関数」の節で述べた φ=on 予備走で精度が伸びなかったのも、同じ根 (SFT データの評価との乖離) が原因と考えています。

SFT の目的とフォーマット固定の関係

今回の SFT データは「医療知識の補強」を目的に設計しており、フォーマット自体は自然言語 QA で意図的に開いた形にしていました。一方 GRPO 側では出力フォーマットを [ans]a[/ans] に固定していました。「SFT で知識を入れる、フォーマットは RL 側で固定する」という役割分担の想定です。

この構成で精度が伸びなかった事実から、以下が読み取れます:

  • CHORD の loss 混合下では、SFT データが「知識のみを注入する」意図で作られていても、そのフォーマット自体が学習信号として同時に流れ込む
  • RL 側でフォーマットが固定されていても、SFT 側でフォーマットが別方向に開いていると、両者の学習方向が食い違って相殺する
  • 従って CHORD においては、RL 側でフォーマットを固定するのであれば、SFT データも同じフォーマットに揃えて作るべきだった

6. まとめ

  • Qwen3-Next-80B-A3B-Instruct に対して、CHORD (L_CHORD = (1-μ)·L_GRPO + μ·L_SFT) を ms-swift の release/3.11 に実装して 64 GPU で回した
  • 実装は ms-swift 内の 2 ファイル (grpo_trainer.py / megatron_args.py) の置き換えで完結
  • ハイパラは μ_peak=0.1, μ_valley=0.01, warmup=0, decay=500, φ=off
  • φ は最初 on で予備走したが精度が伸びず、SFT データが評価フォーマット (MCQ + [ans]) と乖離していたため off に戻して本走
  • 結果: 医師国試 89.5 → 90.4 (+0.9)、専門医試験 69.0 → 70.0 (+1.0, McNemar p=0.119)、フォーマット崩壊なし

謝辞

この成果は、NEDO(国立研究開発法人新エネルギー・産業技術総合開発機構)の 委託業務(JPNP25006)の結果得られたものです。

参考文献

  1. Xie et al. (2025) "On-Policy RL Meets Off-Policy Experts: Harmonizing Supervised Fine-Tuning and Reinforcement Learning via Dynamic Weighting" arXiv:2508.11408
  2. Shao et al. (2024) "DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models" arXiv:2402.03300 — GRPO の原論文
  3. 公式実装: Trinity-RFT/examples/mix_chord
  4. ms-swift: modelscope/ms-swift (今回は release/3.11 ブランチを使用)
  5. Megatron-LM: NVIDIA/Megatron-LM (今回は core_r0.14.0 タグを使用)

関連記事

リポジトリ
https://github.com/weblab-llm-m/singularity-post-training-medical

学習済みモデル (HuggingFace)

1
0
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
1
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?