Qwen3-4B の完全な微調整

単一の H100 GPU で Qwen3-4B 大規模言語モデルを完全に微調整します。 このチュートリアルでは、次の方法を示します。

  • 完全な微調整を実行します。すべてのモデル パラメーターが更新され、データが最大限に適応されます
  • 追加のライブラリをインストールせずに Databricks AI v5 環境 を使用する
  • TRL (トランスフォーマー強化学習) を活用して監視対象の微調整を行う
  • ガバナンスとデプロイのために微調整されたモデルを Unity カタログに登録する

主な概念:

  • 完全な微調整: すべてのモデルの重みを更新し、パラメーター効率の高いメソッドよりも高いメモリとコンピューティングを犠牲にして、データセットから学習する最大の容量をモデルに与えます
  • TRL: 言語モデルを強化学習と教師あり微調整でトレーニングするためのライブラリです
  • メモリ効率の高いトレーニング: BF16 の混合精度と勾配チェックポイント処理を使用して、単一の H100 GPU で 4B パラメーターの完全な微調整を行います

完全な微調整と LoRA デシジョン マトリックス

このノートブックでは、すべてのモデル パラメーターを更新する 完全な微調整を使用します。 もう一つの方法であるLoRA (Low-Rank Adaptation)は、ベースモデルを凍結し、小さなアダプター層のみを学習させます。

シナリオ レコメンデーション 理由
モデルの主要な動作の変更 完全な微調整 モデルの動作に対する基本的な変更のすべてのパラメーターを更新します
1 つのタスクで可能な限り高い品質 完全な微調整 低ランクの近似がないため、モデルには適応する完全な容量があります
制限付き GPU メモリ ローラ 最大 1% のパラメーターのみをトレーニングすることで、メモリ内の大規模なモデルに適合します
複数のタスク固有のアダプター ローラ 同じ基本モデルで異なるアダプターをスワップする

4B パラメーター モデルを完全に微調整するには、すべてのパラメーターでオプティマイザーの状態とグラデーションが維持されるため、LoRA よりも GPU メモリが大幅に多く必要です。 このノートブックでは 、BF16 の混合精度勾配チェックポイント処理が 使用されるため、トレーニングは 1 つの H100 (80 GB) GPU に適合します。

サーバーレス GPU コンピューティングに接続する

サーバーレス GPU コンピューティングに接続するには:

  1. ノートブックの [ 接続 ] ドロップダウン メニューをクリックし、[ サーバーレス GPU] を選択します。
  2. アクセラレータとして 1x H100 GPU を選択します。
  3. [環境] パネルを開き、基本環境として AI v5 を選択します。
  4. [適用] をクリックします。

詳細については、 GPU コンピューティングのドキュメントを参照してください

ライブラリをインポートする

Databricks AI v5 環境には、この例に必要なすべてのライブラリ ( trltransformersdatasetsmlflowなど) が既に含まれているため、追加のインストールは必要ありません。

次のセルは、モデルトレーニング、データセット処理、および MLflow 追跡に必要なライブラリをインポートします。

from datasets import load_dataset
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import (
    SFTConfig,
    SFTTrainer,
    setup_chat_format
)
import torch
import mlflow

構成設定

Unity Catalog の統合

次のセルは、微調整されたモデルを格納して登録する場所を構成します。

  • カタログとスキーマ: Unity カタログ名前空間内でモデルを整理する (既定値: main.default)
  • モデル名: ガバナンスとデプロイ用の Unity カタログに登録されているモデル名
  • ボリューム: トレーニング中にモデル チェックポイントを格納するための Unity カタログ ボリューム

これらのウィジェットを使用すると、コードを編集せずにストレージの場所をカスタマイズできます。 このモデルは、簡単にアクセスしてバージョン管理できるように、 {catalog}.{schema}.{model_name} として登録されます。

ハイパーパラメーターのトレーニング

このセルでは、次の主要なトレーニング パラメーターも定義されます。

  • モデルとデータセット: Qwen3-4B と Capybara の会話型データセット
  • バッチサイズ (1): 各トレーニングステップで各GPUが処理する例数。完全なファインチューニングをメモリに収めるため、小さく保ちます
  • グラデーションの累積 (8):8 バッチを超えるグラデーションを累積し、有効なバッチ サイズを 8 に設定します。
  • 学習率 (2e-5): 完全な微調整に適した保守的なレート
  • 最大ステップ数 (50): 高速なデモ実行のため、学習を50ステップに制限します
  • ログ記録とチェックポイント処理: 25 ステップごとに進行状況を保存し、メトリックを 10 ステップごとにログに記録します
dbutils.widgets.text("uc_catalog", "main")
dbutils.widgets.text("uc_schema", "default")
dbutils.widgets.text("uc_model_name", "qwen3_4b_assistant")
dbutils.widgets.text("uc_volume", "checkpoints")

UC_CATALOG = dbutils.widgets.get("uc_catalog")
UC_SCHEMA = dbutils.widgets.get("uc_schema")
UC_MODEL_NAME = dbutils.widgets.get("uc_model_name")
UC_VOLUME = dbutils.widgets.get("uc_volume")

print(f"UC_CATALOG: {UC_CATALOG}")
print(f"UC_SCHEMA: {UC_SCHEMA}")
print(f"UC_MODEL_NAME: {UC_MODEL_NAME}")
print(f"UC_VOLUME: {UC_VOLUME}")

# MLflow and Unity Catalog configuration

# Model selection
MODEL_NAME = "Qwen/Qwen3-4B"
DATASET_NAME = "trl-lib/Capybara"
OUTPUT_DIR = f"/Volumes/{UC_CATALOG}/{UC_SCHEMA}/{UC_VOLUME}/{UC_MODEL_NAME}"

# Training hyperparameters
BATCH_SIZE = 1
GRADIENT_ACCUMULATION_STEPS = 8
LEARNING_RATE = 2e-5
MAX_STEPS = 50
EVAL_STEPS = 25
LOGGING_STEPS = 10
SAVE_STEPS = 25

データセットの読み込みと準備

次のセルは、トレーニング データセットを読み込み、微調整用に準備します。

  • データセット: trl-lib/Capybara - 命令に従って最適化された高品質の会話データ
  • トレーニング/検証分割: テスト セットが存在しない場合は 90/10 分割を作成します
  • データ検証: 会話の微調整に適した書式を確保する
dataset = load_dataset(DATASET_NAME)
print(f"✓ Dataset loaded: {dataset}")

if "test" not in dataset:
    print("Creating validation split from training data...")
    dataset = dataset["train"].train_test_split(test_size=0.1, seed=42)
    print("✓ Data split: 90% train, 10% validation")

モデルとトークナイザーを初期化する

次のセルは、基本モデルとトークナイザーを読み込み、会話の微調整用に構成します。

  • モデルの読み込み: BF16精度でHugging FaceからQwen3-4Bをダウンロード
  • トークナイザーのセットアップ: 適切なパディングを使用して高速トークナイザーを構成する
  • チャットの書式設定: トークナイザーでまだ定義されていない場合は、構造化された会話にチャット テンプレートを適用します
  • トークンの構成: 適切なシーケンス処理のために埋め込みトークンを EOS トークンに設定します
model = AutoModelForCausalLM.from_pretrained(
    MODEL_NAME,
    torch_dtype=torch.bfloat16,
    trust_remote_code=True,
)

tokenizer = AutoTokenizer.from_pretrained(
    MODEL_NAME,
    trust_remote_code=True,
    use_fast=True
)

# Chat template formatting for conversational fine-tuning
if tokenizer.chat_template is None:
    print("Adding chat template for proper conversation formatting...")
    model, tokenizer = setup_chat_format(model, tokenizer, format="chatml")
    print("✓ ChatML format applied for structured conversations")

if tokenizer.pad_token is None:
    tokenizer.pad_token = tokenizer.eos_token
    print("✓ Padding token set to EOS token")

print("✓ Model and tokenizer loaded successfully")

モデルをトレーニングする

次のセルは、完全な微調整プロセスを構成して実行します。

トレーニング構成

  • バッチ構成: 8 つの勾配蓄積ステップを持つデバイスあたり 1 サンプル (有効なバッチ サイズ: 8)
  • 最適化: ウォームアップステップ、重量減衰、および評価損失に基づく最適なモデル選択
  • ログ記録: 実験追跡のためにメトリックを MLflow に報告する

有効になっているキーの最適化

  • BF16混合精度: メモリ使用量を抑えつつ高速な計算を実現し、H100 GPUに適している
  • 勾配チェックポイント: 追加の計算コストと引き換えに活性化メモリを大幅に削減する手法で、これにより4Bモデルのフルファインチューニングを単一のH100に収めることができます
  • 勾配累積: より大きなバッチ サイズをシミュレートして安定したトレーニングを行います
  • チェックポイント処理: 25 ステップごとにモデルを保存し、チェックポイント数を 2 に制限します。

トレーニング ループでは、10 ステップごとに進行状況が記録され、25 ステップごとに評価されます。

with mlflow.start_run(run_name=f"{MODEL_NAME}_full-fine-tuning", log_system_metrics=True):
    try:
        print(f"Learning rate: {LEARNING_RATE}")

        training_args_dict = {
            "output_dir": OUTPUT_DIR,
            "per_device_train_batch_size": BATCH_SIZE,
            "per_device_eval_batch_size": BATCH_SIZE,
            "gradient_accumulation_steps": GRADIENT_ACCUMULATION_STEPS,
            "learning_rate": LEARNING_RATE,
            "max_steps": MAX_STEPS,
            "eval_steps": EVAL_STEPS,
            "logging_steps": LOGGING_STEPS,
            "save_steps": SAVE_STEPS,
            "save_total_limit": 2,
            "report_to": "mlflow",  # Log to MLflow
            "warmup_steps": 10,
            "weight_decay": 0.01,
            "metric_for_best_model": "eval_loss",
            "greater_is_better": False,
            "eval_strategy": "steps",  # Run evaluation every eval_steps
            "save_strategy": "steps",  # Checkpoint on the same cadence as eval
            "load_best_model_at_end": True,  # Register the best-eval checkpoint, not the last
            "dataloader_pin_memory": False,
            "remove_unused_columns": False,
            "bf16": True,  # Mixed precision training
            "gradient_checkpointing": True,  # Reduce activation memory for full fine-tuning
            "gradient_checkpointing_kwargs": {"use_reentrant": False},
        }

        training_args = SFTConfig(**training_args_dict)

        trainer = SFTTrainer(
            model=model,
            args=training_args,
            train_dataset=dataset["train"],
            eval_dataset=dataset["test"],
            processing_class=tokenizer,
        )

        print("\n" + "="*50)
        print("STARTING TRAINING")
        print("="*50)

        print("🚀 Full fine-tuning Qwen3-4B on a single H100 GPU")

        trainer.train()
        print("\n✓ Training completed successfully!")

    except Exception as e:
        print(f"✗ Training failed: {e}")
        raise

モデル成果物を保存する

次のセルは、トレーニング済みのモデルとトークナイザーを Unity カタログ ボリュームに保存します。

  • モデル全体の重み: 微調整済みモデル全体を保存し、推論用にそのまま直接読み込めるようにします
  • トークナイザー: 推論用のトークナイザー構成を保存します
  • 保存場所: 保存先 /Volumes/{catalog}/{schema}/{volume}/{model_name}
try:
    print("\nSaving trained model...")

    trainer.save_model(training_args.output_dir)
    print("✓ Full model weights saved")

    tokenizer.save_pretrained(training_args.output_dir)
    print("✓ Tokenizer saved with model")
    print(f"\n🎉 All artifacts saved to: {training_args.output_dir}")

except Exception as e:
    print(f"✗ Model saving failed: {e}")
    raise

Unity カタログにモデルを登録する

次のセルは、ガバナンスとデプロイのために Unity カタログに微調整されたモデルを登録します。

モデル登録ワークフロー

  1. トレーニング済みモデルの読み込み: 保存済みのフルウェイト モデルとトークナイザーを読み込みます
  2. ログ記録の準備: モデルとトークナイザーを使用してトランスフォーマー モデル ディクショナリを作成します
  3. Unity カタログへの登録: MLflow へのログと Unity カタログへの登録
  4. メタデータの追加: タスクの種類、モデル ファミリ、サイズ情報が含まれます

Unity カタログ登録の利点

  • ガバナンス: アクセス制御と系列追跡を使用した一元化されたモデル レジストリ
  • バージョン管理: モデル ライフサイクルの自動バージョン管理
  • デプロイ: サービス エンドポイントをモデル化するための簡単なデプロイ
  • 検出可能性: モデルは検索可能であり、Unity カタログに記載されています
mlflow_run_id = mlflow.last_active_run().info.run_id
print("\nRegistering model with MLflow and Unity Catalog...")

with mlflow.start_run(run_id=mlflow_run_id):
    try:
        # Load the trained full-weight model for registration
        print("Loading fine-tuned model for registration...")
        trained_model = AutoModelForCausalLM.from_pretrained(
            training_args.output_dir,
            torch_dtype=torch.bfloat16,
            trust_remote_code=True
        )
        tokenizer = AutoTokenizer.from_pretrained(training_args.output_dir)
        model_type = "Full fine-tuning"
        size_params = "4b"

        # Prepare transformers model dictionary
        transformers_model = {
            "model": trained_model,
            "tokenizer": tokenizer
        }

        # Create Unity Catalog model name
        full_model_name = f"{UC_CATALOG}.{UC_SCHEMA}.{UC_MODEL_NAME}"

        print(f"Registering model as: {full_model_name}")

        # Start MLflow run and log model
        task = "llm/v1/chat"
        model_info = mlflow.transformers.log_model(
            transformers_model=transformers_model,
            task=task,
            registered_model_name=full_model_name,
            metadata={
                "task": task,
                "pretrained_model_name": MODEL_NAME,
                "databricks_model_family": "Qwen3ForCausalLM",
                "databricks_model_size_parameters": size_params,
            },
            repo_type="local",  # Fix: specify repo_type for local path
        )

        print(f"✓ Model successfully registered in Unity Catalog: {full_model_name}")
        print(f"✓ MLflow model URI: {model_info.model_uri}")

        # Print deployment information
        print(f"\n📦 Model Registration Complete!")
        print(f"Unity Catalog Path: {full_model_name}")
        print(f"Model Type: {model_type}")

    except Exception as e:
        print(f"✗ Model registration failed: {e}")
        print("Model is still saved locally and can be registered manually")
        print(f"Local model path: {training_args.output_dir}")
        raise

次のステップ

お使いのQwen3-4Bモデルは、フルウェイトファインチューニングによって正常にファインチューニングされ、Unity Catalog に登録されました。 次に、次のことができます。

ノートブックの例

Qwen3-4B の完全な微調整

ノートブックを入手