DeepSeek-R1からQwenへの知識蒸留:TextBrewerを用いた実装ガイド

大規模言語モデルの性能を軽量なモデルに継承する知識蒸留は、実用的なAIデプロイメントにおいて重要な技術です。本記事では、DeepSeek-R1-Distill-Qwen-1.5Bを教師モデル、Qwen-2.5-1.5Bを生徒モデルとして、TextBrewerライブラリを使用した蒸留プロセスを解説します。

蒸留の目的と概要

本実装の主な目標は、DeepSeek-R1が持つ推論能力をQwen-2.5-1.5Bに転移させることです。蒸留により、軽量でありながら高度な推論機能を備えたモデルを構築し、リソース制約のある環境でも効率的に運用可能にします。

環境構成

必要なライブラリ

実装には以下のパッケージが必要です:

pip install torch transformers datasets accelerate textbrewer
pip install evaluate tensorboard

ハードウェア要件

  • GPU:NVIDIA V100/A100相当(VRAM 24GB以上推奨)
  • CUDA:PyTorchと互換性のあるバージョン

モデルの準備

Hugging Faceからモデルをダウンロードします。オフライン環境の場合は、ミラーサイトを利用可能です:

import os
os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"

# 教師モデル:DeepSeek-R1-Distill-Qwen-1.5B
# 生徒モデル:Qwen/Qwen2.5-1.5B

実装の詳細

モデルとトークナイザーの読み込み

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

def initialize_models():
    """教師モデルと生徒モデルを初期化"""
    # 教師モデルの設定
    mentor_path = "./models/DeepSeek-R1-Distill-Qwen-1.5B"
    mentor_tokenizer = AutoTokenizer.from_pretrained(mentor_path, local_files_only=True)
    mentor_model = AutoModelForCausalLM.from_pretrained(
        mentor_path,
        local_files_only=True,
        torch_dtype=torch.float16
    )
    
    # 生徒モデルの設定
    learner_path = "./models/qwen2.5-1.5B"
    learner_tokenizer = AutoTokenizer.from_pretrained(learner_path, local_files_only=True)
    learner_model = AutoModelForCausalLM.from_pretrained(
        learner_path,
        local_files_only=True,
        torch_dtype=torch.float16
    )
    
    return mentor_model, learner_model, mentor_tokenizer, learner_tokenizer

データセットの前処理

Wikitextデータセットを使用してトレーニングデータを準備します:

from datasets import load_dataset
from transformers import DataCollatorForLanguageModeling

def setup_data_pipeline(tokenizer, max_seq_length=512):
    """データパイプラインを構築"""
    # データセットの読み込み
    corpus = load_dataset("wikitext", "wikitext-2-raw-v1")
    
    def tokenize_batch(samples):
        return tokenizer(
            samples["text"],
            truncation=True,
            padding="max_length",
            max_length=max_seq_length
        )
    
    # トレーニング用データの処理
    train_data = corpus["train"].map(tokenize_batch, batched=True)
    val_data = corpus["validation"].map(tokenize_batch, batched=True)
    
    # 因果言語モデル用のコレーター(MLM無効)
    collator = DataCollatorForLanguageModeling(
        tokenizer=tokenizer,
        mlm=False
    )
    
    return train_data, val_data, collator

蒸留設定の定義

TextBrewerのDistillationConfigを使用して、蒸留パラメータを設定します:

from textbrewer import DistillationConfig, TrainingConfig, GeneralDistiller

def configure_distillation():
    """蒸留設定を構成"""
    distill_cfg = DistillationConfig(
        temperature=2.0,  # ソフトラベルの平滑化
        hard_label_weight=0.5,  # 正解ラベルの損失重み
        kd_loss_type="ce",  # 交差エントロピー損失
        intermediate_matches=[
            {
                "layer_T": 6,  # 教師モデルの層
                "layer_S": 6,  # 生徒モデルの層
                "feature": "hidden",  # 隠れ状態のマッチング
                "weight": 1.0,
                "loss": "mse"  # 平均二乗誤差
            }
        ]
    )
    
    train_cfg = TrainingConfig(
        device="cuda" if torch.cuda.is_available() else "cpu",
        log_dir="./logs",
        output_dir="./checkpoints",
        fp16=True
    )
    
    return distill_cfg, train_cfg

トレーニングパラメータ

from transformers import TrainingArguments

def get_training_args():
    return TrainingArguments(
        output_dir="./distillation_output",
        evaluation_strategy="epoch",
        learning_rate=5e-5,
        per_device_train_batch_size=2,
        per_device_eval_batch_size=2,
        num_train_epochs=3,
        weight_decay=0.01,
        logging_dir="./training_logs",
        logging_steps=100,
        gradient_accumulation_steps=4,
        report_to="tensorboard",
        save_strategy="epoch"
    )

完全な実装コード

以下に、主要なコンポーネントを統合した完全な蒸留スクリプトを示します:

import os
import torch
import logging
from transformers import (
    AutoModelForCausalLM,
    AutoTokenizer,
    DataCollatorForLanguageModeling,
    Trainer,
    TrainingArguments,
    get_linear_schedule_with_warmup
)
from textbrewer import GeneralDistiller, TrainingConfig, DistillationConfig
from datasets import load_dataset
from torch.optim import AdamW

# ロギング設定
logging.basicConfig(
    level=logging.INFO,
    format="[%(asctime)s] %(levelname)s: %(message)s"
)
log = logging.getLogger(__name__)

class KnowledgeDistiller:
    def __init__(self, mentor_dir, learner_dir):
        self.mentor_dir = mentor_dir
        self.learner_dir = learner_dir
        self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
        
    def load_models(self):
        """モデルとトークナイザーを読み込み"""
        log.info("Loading mentor model from %s", self.mentor_dir)
        self.mentor_tokenizer = AutoTokenizer.from_pretrained(self.mentor_dir)
        self.mentor_model = AutoModelForCausalLM.from_pretrained(
            self.mentor_dir,
            torch_dtype=torch.float16
        ).to(self.device)
        
        log.info("Loading learner model from %s", self.learner_dir)
        self.learner_tokenizer = AutoTokenizer.from_pretrained(self.learner_dir)
        self.learner_model = AutoModelForCausalLM.from_pretrained(
            self.learner_dir,
            torch_dtype=torch.float16
        ).to(self.device)
        
        # トークナイザーの整合性確認
        if self.mentor_tokenizer.vocab != self.learner_tokenizer.vocab:
            log.warning("Vocabulary mismatch detected, adjusting...")
            self.learner_tokenizer.add_special_tokens({'pad_token': '[PAD]'})
            self.learner_model.resize_token_embeddings(len(self.learner_tokenizer))
    
    def prepare_datasets(self, dataset_name="wikitext", subset="wikitext-2-raw-v1"):
        """データセットを準備"""
        log.info("Loading dataset: %s", dataset_name)
        raw_data = load_dataset(dataset_name, subset)
        
        def tokenize(samples):
            return self.learner_tokenizer(
                samples["text"],
                truncation=True,
                padding="max_length",
                max_length=512
            )
        
        self.train_set = raw_data["train"].map(tokenize, batched=True)
        self.val_set = raw_data["validation"].map(tokenize, batched=True)
        
        self.collator = DataCollatorForLanguageModeling(
            tokenizer=self.learner_tokenizer,
            mlm=False
        )
    
    def setup_distiller(self):
        """蒸留器を設定"""
        distill_config = DistillationConfig(
            temperature=2.0,
            hard_label_weight=0.4,
            kd_loss_weight=0.6,
            kd_loss_type="ce",
            intermediate_matches=[
                {
                    "layer_T": [8, 16, 24],
                    "layer_S": [4, 8, 12],
                    "feature": "hidden",
                    "loss": "cosine",
                    "weight": 0.5
                }
            ]
        )
        
        training_config = TrainingConfig(
            device=self.device,
            output_dir="./distilled_model",
            log_dir="./distill_logs",
            fp16=torch.cuda.is_available(),
            gradient_accumulation_steps=4,
            max_grad_norm=1.0
        )
        
        self.distiller = GeneralDistiller(
            train_config=training_config,
            distill_config=distill_config,
            model_T=self.mentor_model,
            model_S=self.learner_model,
            adaptor_T=None,
            adaptor_S=None
        )
    
    def run_distillation(self, epochs=3, batch_size=2):
        """蒸留を実行"""
        optimizer = AdamW(
            self.learner_model.parameters(),
            lr=5e-5,
            weight_decay=0.01
        )
        
        scheduler = get_linear_schedule_with_warmup(
            optimizer,
            num_warmup_steps=500,
            num_training_steps=len(self.train_set) // batch_size * epochs
        )
        
        log.info("Starting knowledge distillation...")
        with self.distiller:
            self.distiller.train(
                optimizer=optimizer,
                scheduler=scheduler,
                train_dataset=self.train_set,
                eval_dataset=self.val_set,
                batch_size=batch_size,
                num_epochs=epochs,
                data_collator=self.collator
            )
        
        self.save_model()
    
    def save_model(self, output_dir="./final_distilled_model"):
        """最終モデルを保存"""
        self.learner_model.save_pretrained(output_dir)
        self.learner_tokenizer.save_pretrained(output_dir)
        log.info("Model saved to %s", output_dir)


def main():
    distiller = KnowledgeDistiller(
        mentor_dir="./models/DeepSeek-R1-Distill-Qwen-1.5B",
        learner_dir="./models/Qwen2.5-1.5B"
    )
    
    distiller.load_models()
    distiller.prepare_datasets()
    distiller.setup_distiller()
    distiller.run_distillation(epochs=3, batch_size=2)


if __name__ == "__main__":
    main()

中間層マッチングの戦略

蒸留において、教師モデルと生徒モデルの中間層をマッチングさせることで、より深い知識の転移が可能になります。以下の設定は、隠れ状態と注意機構の両方を考慮した多層マッチングの例です:

intermediate_config = [
    {
        "layer_T": [6, 12, 18],
        "layer_S": [3, 6, 9],
        "feature": "hidden",
        "loss": "cosine",  # 余弦類似度による類似性評価
        "weight": 0.5,
        "proj": ["linear", 1536, 1024]  # 次元削減プロジェクション
    },
    {
        "layer_T": [8, 16],
        "layer_S": [4, 8],
        "feature": "attention",
        "loss": "mse",
        "weight": 0.3
    }
]

この設定では、隠れ状態に対して余弦類似度損失、注意機構に対してMSE損失を適用し、それぞれに異なる重みを割り当てています。次元が異なる場合は、線形プロジェクション層を用いて調整します。

タグ: Knowledge Distillation DeepSeek-R1 Qwen TextBrewer PyTorch

9月12日 02:19 投稿