powered by TechFeed
表示モード
ハウツー

Llama 3をカスタムAPI呼び出しに特化させる — プロンプトだけでは崩れやすいJSONを、ファインチューニングで安定させる方法

10月9日、Machine Learning Masteryが「How to Fine-Tune Llama 3 for Custom Tool Calling with Unsloth in Python」と題した記事を公開した。UnslothとQLoRAを使ってLlama 3 8BをカスタムのツールコールAPIに対応させるファインチューニングの具体的な手順を解説している。なぜプロンプトエンジニアリングでは足りないのかFunction Calling / Tool Callingとは、LLMが自然言語のクエリを受け取り、外部APIやツールを呼び出すための構造化されたJSONを出力する仕組みだ。OpenAI APIやGeminiなど主要なLLMサービスでは既に広く普及しており、エージェント型アプリケーションの基盤技術となっている。問題は、汎用のオープンソースモデルをそのまま使う場合にある。LLMベースのエージェントを実運用に乗せようとすると、必ずぶつかる壁がある。「テスト時は正しいJSONを返すのに、複雑なクエリや長い会話になると形式が崩れやすい」という問題だ。prose(散文)が混入...

10月9日、Machine Learning Masteryが「How to Fine-Tune Llama 3 for Custom Tool Calling with Unsloth in Python」と題した記事を公開した。UnslothとQLoRAを使ってLlama 3 8BをカスタムのツールコールAPIに対応させるファインチューニングの具体的な手順を解説している。

なぜプロンプトエンジニアリングでは足りないのか

Function Calling / Tool Callingとは、LLMが自然言語のクエリを受け取り、外部APIやツールを呼び出すための構造化されたJSONを出力する仕組みだ。OpenAI APIやGeminiなど主要なLLMサービスでは既に広く普及しており、エージェント型アプリケーションの基盤技術となっている。

問題は、汎用のオープンソースモデルをそのまま使う場合にある。LLMベースのエージェントを実運用に乗せようとすると、必ずぶつかる壁がある。「テスト時は正しいJSONを返すのに、複雑なクエリや長い会話になると形式が崩れやすい」という問題だ。prose(散文)が混入したり、クォートスタイルが変わったり、必須フィールドが欠落したりする。下流のパーサーはその瞬間に壊れる。

記事はこの問題を「知識の問題ではなく、振る舞いの問題」と整理する。外部の事実を与えたいならRAG(Retrieval-Augmented Generation)が適しているが、一貫した出力フォーマットを強制したいならファインチューニングが必要だ。重みレベルで行動パターンを書き換えることが、プロンプトエンジニアリングとの本質的な違いである。

なお、Llama 3はMetaが2024年にリリースしたオープンソースLLMで、8Bと70Bのパラメータ規模が公開されている。本記事では8Bモデルを対象とする。

**QLoRA(Quantized Low-Rank Adaptation)はファインチューニングを現実的なコストで実現する手法だ。ベースモデルの重みを凍結した上で、Attention層に小さな訓練可能な行列(LoRAアダプター)を注入する。更新するパラメータは全体の約1%**に留まり、無料のGoogle Colab T4 GPUで動作する。

環境構築とモデルのロード

Google ColabでT4 GPUランタイムを選択し、以下でインストールする。

!pip install "unsloth[colab-new] @ git+https://github.com/unslothai/unsloth.git"
!pip install --no-deps xformers trl peft accelerate bitsandbytes

UnslothはCUDAの設定を自動処理するため、追加セットアップは不要だ。モデルのロードには、Unslothが配布する4bit事前量子化済みのLlama 3 8Bを使う。Hugging Faceのゲート認証が不要で、ダウンロードも通常の約4倍速い。

model, tokenizer = FastLanguageModel.from_pretrained(
    model_name="unsloth/llama-3-8b-Instruct-bnb-4bit",
    max_seq_length=1024,
    dtype=None,
    load_in_4bit=True,
)

続いてLoRAアダプターをアタッチする。r=8(ランク)は小規模データセットでの過学習リスクを抑える値だ。lora_alphaはランクの2倍の16を指定する。**lora_dropout=0は必須**で、非ゼロにするとUnslothの最適化カーネルが無効化されて低速なコードパスにフォールバックする。

model = FastLanguageModel.get_peft_model(
    model,
    r=8,
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
                    "gate_proj", "up_proj", "down_proj"],
    lora_alpha=16,
    lora_dropout=0,
    bias="none",
    use_gradient_checkpointing="unsloth",
    random_state=42,
)

データセット構築が最重要工程

記事が「最も重要なステップ」と位置づけるのがデータセットの設計だ。各トレーニング例には3つの要素が必要になる。

  1. システムプロンプト:利用可能なツールとそのJSONスキーマの定義
  2. ユーザークエリ:自然言語での問い
  3. アシスタント出力:モデルが返すべき正確なJSON
tool_calling_data = [
    {
        "system": """You have access to these tools:
get_weather(location: str) -> dict
fetch_stock_price(ticker: str) -> dict
Respond ONLY with a valid JSON tool call. No other text.""",
        "user": "What is the weather like in Tokyo right now?",
        "output": '{"name": "get_weather", "arguments": {"location": "Tokyo"}}'
    },
    {
        "system": """...""",
        "user": "Get me the current stock price for Apple.",
        "output": '{"name": "fetch_stock_price", "arguments": {"ticker": "AAPL"}}'
    },
]

記事が強調する重要な指摘がある。量より質だ。ノイズの多い2,000件より、丁寧に構築された200件の方が一貫して良い結果を出す。本番環境では、曖昧なクエリ、複数引数のツール、似たツールの区別など、エッジケースを網羅した例を追加すべきである。

フォーマットにはtokenizer.apply_chat_template()を使う。Llama 3の特殊トークンを正確に処理し、将来のモデルバージョンへの移行も容易になる。

学習と推論テスト

TRLのSFTTrainerで学習を実行する。max_steps=60でColab T4上の学習時間を10分以内に収める設定だ。

trainer = SFTTrainer(
    model=model,
    tokenizer=tokenizer,
    train_dataset=dataset,
    dataset_text_field="text",
    max_seq_length=1024,
    args=TrainingArguments(
        per_device_train_batch_size=2,
        gradient_accumulation_steps=4,
        learning_rate=2e-4,
        max_steps=60,
        warmup_steps=5,
        lr_scheduler_type="linear",
        optim="adamw_8bit",
        weight_decay=0.01,
        output_dir="tool_caller_output",
        report_to="none",
    ),
)

学習ロスは0.1〜0.3への漸減が健全な範囲だ。小規模データセットでロスが急激にゼロに近づく場合は、汎化ではなく記憶(メモ化)が起きている可能性があり、より多様な例を追加する必要がある。

学習後の推論テストでは、学習時に含まれていないクエリ「What's the stock price of Tesla?」を投げる。成功したファインチューニングの出力は以下のようになる。

{"name": "fetch_stock_price", "arguments": {"ticker": "TSLA"}}

ベースモデルは同じシステムプロンプトを与えても「Sure! Here's the tool call:」のような会話的なフィラーをJSONの前後に付け加えることが多い。下流のパーサーはこれで壊れる。ファインチューニング後にまだ余分なテキストが出る場合は、学習データのフォーマットが不均一であるか、「no other text」制約を強化する例が不足している。

最後にLoRAアダプターを保存して完了だ。ベースモデルの重みは変わっておらず、アダプターは独立して管理・配布できる。

詳細はHow to Fine-Tune Llama 3 for Custom Tool Calling with Unsloth in Pythonを参照していただきたい。