ファインチューニング

小さいモデルを使用して用途特化の推論器を作成する。

モチベーション

ChromeにPrompt APIというLanguageModelを使用したAI機能が生えている。ただし、現状コンテキストが小さく、賢さも少し足りない。 transformer.jsを使用すれば、コンテキスト上限の調整も効くはず。モデルの小ささから来る賢さの不足も、やりたいことにチューニングしたモデルを作成すればうまく行くんじゃなかろうか。

ベースモデル

Qwen2.5-1.5B-Instruct 容量が小さく、日本語能力が高くあってほしいので、Qwen2.5 1.5Bを採用する。3系はMoEで容量が大きくブラウザに載らないので2.5にした。outputしてほしい形式が決まっているため、忠実性の高いInstruct番を使用する。


RTX 3070 Ti (VRAM 8GB) で Q をファインチューニングして推論するまで

会社にGPUがあったので、古くてスペックは低いけど1.5Bには十分だから利用する。


1. 動作環境


2. 環境構築とライブラリの選定

はじめは爆速・省メモリで話題の Unsloth の導入を試みたが、Windowsネイティブ環境での依存関係やビルド周りのトラブル(flit_core やソースビルドのエラー)が多発したため、公式の標準的な Hugging Face エコシステム(transformers, peft, trl, bitsandbytes, datasets, accelerate を使うことにした。

必要なライブラリのインストール

# 基本ツールのアップデート
python -m pip install --upgrade pip setuptools wheel

# PyTorch (CUDA 11.8対応版) のインストール
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 --upgrade --prefer-binary

# 標準QLoRAに必要なライブラリ一式
pip install transformers peft trl bitsandbytes datasets accelerate

3. 学習用データセットの準備 (train_data.jsonl)

学習データは、より賢い gpt5.6 Lunaを使用して生成し200件ほど用意した。

{"instruction":"以下のスキーマに基づいて、ユーザーの要望を満たすODataの$filter、$orderbyのいずれかもしくは両方を`&`で結合したクエリを出力してください。\nスキーマ:\n型の末尾の ? は null 可、! は null 不可を表します。\n対象エンティティ: SiteFile\nフィールド一覧:\n- FileType (string!)\n- FilePath (string!)\n- ContentType (string?)\n- FileBinary (string?)\nナビゲーションプロパティ:\n(ナビゲーションプロパティなし)","input":"ファイル種別がPDFで、コンテンツ種別がapplication/pdfのものを、ファイルパスの昇順で並べてほしい。","output":"$filter=FileType eq 'PDF' and ContentType eq 'application/pdf'&$orderby=FilePath asc"}
{"instruction":"以下のスキーマに基づいて、ユーザーの要望を満たすODataの$filter、$orderbyのいずれかもしくは両方を`&`で結合したクエリを出力してください。\nスキーマ:\n型の末尾の ? は null 可、! は null 不可を表します。\n対象エンティティ: SiteFile\nフィールド一覧:\n- FileType (string!)\n- FilePath (string!)\n- ContentType (string?)\n- FileBinary (string?)\nナビゲーションプロパティ:\n(ナビゲーションプロパティなし)","input":"PDFのファイルで、コンテンツ種別がapplication/pdfのやつを、ファイルパス順に並べてくれるかな。","output":"$filter=FileType eq 'PDF' and ContentType eq 'application/pdf'&$orderby=FilePath asc"}
...

色々試していたので学習データ自体はAlpaca形式(instruction, input, output)の jsonl データで作ってしまっていたが、Qwen2.5のファインチューニングはChatML形式(messages リスト)がいいらしく、変換して学習に投入できるようにした。

{"instruction": "...", "input": "...", "output": "..."}

4. 学習スクリプトのポイントと最大の難所

VRAM 8GBという制限があるため、QLoRA (4-bit量子化) を採用。

BFloat16と混合精度(AMP)の互換性エラーにハマった。

RuntimeError: "_amp_foreach_non_finite_check_and_unscale_cuda" not implemented for 'BFloat16'

RTX 3070 TiはBFloat16をサポートしているものの、BitsAndBytesなどのオプティマイザとPyTorchのスケーラーの噛み合わせで勾配のアンースケール処理時にパニックを起こしていた。

これを完全に回避するため、混合精度(FP16/BF16)を完全にオフにし、ストレートなFP32計算で回す設定にすることで無事に解決しました。1.5Bモデルであれば、FP32でもVRAM 8GBの範囲内に収まる。

最終的な学習設定 (train.py の抜粋)

from datasets import load_dataset
from peft import LoraConfig, prepare_model_for_kbit_training
import torch
from transformers import (
    AutoModelForCausalLM,
    AutoTokenizer,
    BitsAndBytesConfig,
)
from trl import SFTConfig, SFTTrainer

model_id = "Qwen/Qwen2.5-1.5B-Instruct"

# 4-bit 量子化 (QLoRA) 設定
bnb_config = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=torch.float16,
    bnb_4bit_use_double_quant=True,
)

model = AutoModelForCausalLM.from_pretrained(
    model_id,
    quantization_config=bnb_config,
    device_map="auto",
    trust_remote_code=True,
)
model = prepare_model_for_kbit_training(model)

peft_config = LoraConfig(
    r=16,
    lora_alpha=32,
    target_modules=[
        "q_proj",
        "k_proj",
        "v_proj",
        "o_proj",
        "gate_proj",
        "up_proj",
        "down_proj",
    ],
    lora_dropout=0.05,
    bias="none",
    task_type="CAUSAL_LM",
)

# 混合精度をオフ(FP32)にすることでエラーを完全に回避
training_args = SFTConfig(
    output_dir="./results",
    num_train_epochs=3,
    per_device_train_batch_size=1,
    gradient_accumulation_steps=4,
    learning_rate=2e-4,
    logging_steps=1,
    save_strategy="epoch",
    fp16=False,
    bf16=False,
    optim="adamw_torch",
    gradient_checkpointing=True,
    max_length=1024,
    dataset_text_field="text",
)

trainer = SFTTrainer(
    model=model,
    train_dataset=dataset,
    peft_config=peft_config,
    tokenizer=tokenizer,
    args=training_args,
)

trainer.train()
trainer.model.save_pretrained("./qwen2.5-1.5b-finetuned")

5. 推論の実行 (test.py)

学習が完了したら、保存されたLoRAアダプター(最後のチェックポイントフォルダ等)をベースモデルに適用して推論を行う。学習はFP32で行ったが、推論時もFP32である必要は無いらしく torch.float16 を指定してFP16で計算できた。

from peft import PeftModel
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

base_model_id = "Qwen/Qwen2.5-1.5B-Instruct"
adapter_path = "./results/checkpoint-162"  # または保存先フォルダ

tokenizer = AutoTokenizer.from_pretrained(base_model_id, trust_remote_code=True)

base_model = AutoModelForCausalLM.from_pretrained(
    base_model_id,
    torch_dtype=torch.float16,
    device_map="auto",
    trust_remote_code=True,
)

model = PeftModel.from_pretrained(base_model, adapter_path)

# テスト推論の実行
messages = [
    {"role": "system", "content":
"""
以下のスキーマに基づいて、ユーザーの要望を満たすODataの$filter、$orderbyのいずれかもしくは両方を`&`で結合したクエリを出力してください。
スキーマ:
型の末尾の ? は null 可、! は null 不可を表します。
対象エンティティ: PurchaseOrder
フィールド一覧:
- Id (number!)
- OrderType (number?)
- OrderStatus (number!)
- OrderDate (date!)
- ShipRequestDate (date?)
- ShipDate (date?)
- DeliveryCompleteDate (date?)
- CustomerId (number!)
- PaymentMethod (number!)
- PaymentStatus (number!)
- TaxRate (number!)
- DeliveryCharge (number!)
- PointPaymentForDeliveryCharge (number!)
- TotalPayment (number!)
- TotalUsagePoint (number!)
- ReturnPrice (number!)
- CashOnDeliveryCharge (number!)
- CardPayment (number!)
- TotalAmount (number!)
- CancelBeforeStatus (number?)
- ReturnDate (date?)
- ExtraPoint (number?)
- ExtraPointSummary (number?)
- SiteId (number?)
- MemoId (number?)
- BookingEnable (boolean?)
- ChargePointSummary (number?)
- OriginalOrderId (number?)
- EstimateShipDate (date?)
- AllocationCompleteDate (date?)
- DeliveryReportDate (date?)
- ShipBookAt (date?)
- DeliveryBookAt (date?)
- ChangeStamp (string?)
- DiscountPrice (number!)
- PointPaymentForPaymentCharge (number!)
ナビゲーションプロパティ:
- DeliveryOrder (単一, → DeliveryOrder): "DeliveryOrder/<フィールド名>"で使用
  DeliveryOrderフィールド一覧:
  - OrderId (number!)
  - DeliveryDate (date?)
  - HourRange (string?)
  - WrappingTypeId (number?)
  - SenderName (string?)
  - AddressName (string!)
  - DeliveryNo (string?)
  - AddressId (number?)
  - ShipSourceId (number?)
  - MailAddress (string?)
- OrderCustomer (単一, → OrderCustomer): "OrderCustomer/<フィールド名>"で使用
  OrderCustomerフィールド一覧:
  - Id (number!)
  - IsGuest (boolean!)
  - UserNo (number?)
  - MemberRank (number?)
  - MemberStatus (number?)
  - OrdererName (string?)
  - FirstName (string!)
  - LastName (string!)
  - FirstNameKana (string!)
  - LastNameKana (string!)
  - MailAddress (string!)
  - CreditCardNumber (string?)
  - CreditCardExpire (string?)
  - CreditSecurity (string?)
  - AuthorizeNo (string?)
  - AuthorizeAt (date?)
  - CreditResponseStatus (string?)
  - TransactionNo (string?)
  - NumberOfPayments (number?)
  - AuthorizeError (boolean?)
  - OrderedAddressId (number?)
  - InvoiceAddressId (number?)
  - Birthday (date?)
  - Sex (number?)
  - OriginalOrderId (string?)
  - PaymentDetail (string?)
  - PaymentSlipNumber (string?)
  - PaymentSlipUrl (string?)
  - CompletePayment (boolean?)
  - PaymentMailAddress (string?)
  - GuestAccessKey (string?)
  - AutoCancelDate (date?)
- OrderLines (コレクション, → OrderLine): any|allなど用いる
  OrderLineフィールド一覧:
  - Id (number!)
  - OrderLineType (number!)
  - OrderId (number!)
  - ReserveId (number?)
  - ArriveNoticeId (number?)
  - ProductId (number!)
  - ParentId (number?)
  - OrderAmount (number!)
  - AllocateAmount (number!)
  - Allocating (boolean!)
  - UnitPrice (number!)
  - ReportPrice (number?)
  - LinePrice (number!)
  - Tax (number!)
  - PointUsage (number!)
  - PointUsageTax (number!)
  - PointUsagePrice (number!)
  - PointCharge (number!)
  - PointChargeRate (number!)
  - AllocateCompleteDate (date?)
  - ExtraPoint (number?)
  - MemoId (number?)
  - ProductName (string?)
  - DiscountPrice (number!)
  - Description (string?)
- RelatePurchaseOrders (コレクション, → PurchaseOrder): any|allなど用いる
  PurchaseOrderフィールド一覧:
  - Id (number!)
  - OrderType (number?)
  - OrderStatus (number!)
  - OrderDate (date!)
  - ShipRequestDate (date?)
  - ShipDate (date?)
  - DeliveryCompleteDate (date?)
  - CustomerId (number!)
  - PaymentMethod (number!)
  - PaymentStatus (number!)
  - TaxRate (number!)
  - DeliveryCharge (number!)
  - PointPaymentForDeliveryCharge (number!)
  - TotalPayment (number!)
  - TotalUsagePoint (number!)
  - ReturnPrice (number!)
  - CashOnDeliveryCharge (number!)
  - CardPayment (number!)
  - TotalAmount (number!)
  - CancelBeforeStatus (number?)
  - ReturnDate (date?)
  - ExtraPoint (number?)
  - ExtraPointSummary (number?)
  - SiteId (number?)
  - MemoId (number?)
  - BookingEnable (boolean?)
  - ChargePointSummary (number?)
  - OriginalOrderId (number?)
  - EstimateShipDate (date?)
  - AllocationCompleteDate (date?)
  - DeliveryReportDate (date?)
  - ShipBookAt (date?)
  - DeliveryBookAt (date?)
  - ChangeStamp (string?)
  - DiscountPrice (number!)
  - PointPaymentForPaymentCharge (number!)
- OriginalPurchaseOrder (単一, → PurchaseOrder): "OriginalPurchaseOrder/<フィールド名>"で使用
  PurchaseOrderフィールド一覧:
  - Id (number!)
  - OrderType (number?)
  - OrderStatus (number!)
  - OrderDate (date!)
  - ShipRequestDate (date?)
  - ShipDate (date?)
  - DeliveryCompleteDate (date?)
  - CustomerId (number!)
  - PaymentMethod (number!)
  - PaymentStatus (number!)
  - TaxRate (number!)
  - DeliveryCharge (number!)
  - PointPaymentForDeliveryCharge (number!)
  - TotalPayment (number!)
  - TotalUsagePoint (number!)
  - ReturnPrice (number!)
  - CashOnDeliveryCharge (number!)
  - CardPayment (number!)
  - TotalAmount (number!)
  - CancelBeforeStatus (number?)
  - ReturnDate (date?)
  - ExtraPoint (number?)
  - ExtraPointSummary (number?)
  - SiteId (number?)
  - MemoId (number?)
  - BookingEnable (boolean?)
  - ChargePointSummary (number?)
  - OriginalOrderId (number?)
  - EstimateShipDate (date?)
  - AllocationCompleteDate (date?)
  - DeliveryReportDate (date?)
  - ShipBookAt (date?)
  - DeliveryBookAt (date?)
  - ChangeStamp (string?)
  - DiscountPrice (number!)
  - PointPaymentForPaymentCharge (number!)
- ReserveStocks (コレクション, → ReserveStock): any|allなど用いる
  ReserveStockフィールド一覧:
  - ReserveId (number!)
  - UserNo (number?)
  - OrderId (number?)
  - ProductId (number?)
  - Amount (number?)
  - ExpireAt (date?)
  - Purchased (boolean?)
  - ExternalSourceId (number?)
- ReserveRequests (コレクション, → ReserveRequest): any|allなど用いる
  ReserveRequestフィールド一覧:
  - Id (number!)
  - RequestAt (date!)
  - UserNo (number!)
  - ProductId (number!)
  - Amount (number!)
  - Status (number!)
  - OrderId (number?)
  - ReserveStockId (number?)
  - MailAddress (string?)
  - SiteId (number!)
- ReturnOrders (コレクション, → ReturnOrder): any|allなど用いる
  ReturnOrderフィールド一覧:
  - Id (number!)
  - OrderType (number?)
  - OrderStatus (number!)
  - OrderDate (date!)
  - ShipRequestDate (date?)
  - ShipDate (date?)
  - DeliveryCompleteDate (date?)
  - CustomerId (number!)
  - PaymentMethod (number!)
  - PaymentStatus (number!)
  - TaxRate (number!)
  - DeliveryCharge (number!)
  - PointPaymentForDeliveryCharge (number!)
  - TotalPayment (number!)
  - TotalUsagePoint (number!)
  - ReturnPrice (number!)
  - CashOnDeliveryCharge (number!)
  - CardPayment (number!)
  - TotalAmount (number!)
  - CancelBeforeStatus (number?)
  - ReturnDate (date?)
  - ExtraPoint (number?)
  - ExtraPointSummary (number?)
  - SiteId (number?)
  - MemoId (number?)
  - BookingEnable (boolean?)
  - ChargePointSummary (number?)
  - OriginalOrderId (number?)
  - EstimateShipDate (date?)
  - AllocationCompleteDate (date?)
  - DeliveryReportDate (date?)
  - ShipBookAt (date?)
  - DeliveryBookAt (date?)
  - ReturnReasonId (number!)
  - PurchaseOrderId (number!)
  - DiscountPrice (number!)
  - PointPaymentForPaymentCharge (number!)
- ServiceValues (コレクション, → PurchaseOrderService): any|allなど用いる
  PurchaseOrderServiceフィールド一覧:
  - OrderId (number!)
  - Code (string!)
  - Value (string?)
- OrderLineServices (コレクション, → OrderLineService): any|allなど用いる
  OrderLineServiceフィールド一覧:
  - OrderLineId (number!)
  - Code (string!)
  - OrderId (number!)
  - Value (string?)
- OrdersLastUpdate (単一, → OrdersLastUpdate): "OrdersLastUpdate/<フィールド名>"で使用
  OrdersLastUpdateフィールド一覧:
  - Id (number!)
  - LastUpdate (date?)
- PurchaseOrderImport (単一, → PurchaseOrderImport): "PurchaseOrderImport/<フィールド名>"で使用
  PurchaseOrderImportフィールド一覧:
  - OrderId (number!)
  - ImportOrderId (string?)
- CampaignRelations (コレクション, → PurchaseOrderApplyCampaign): any|allなど用いる
  PurchaseOrderApplyCampaignフィールド一覧:
  - OrderId (number!)
  - CampaignId (number!)
- Campaigns (コレクション, → Campaign): any|allなど用いる
  Campaignフィールド一覧:
  - Id (number!)
  - Name (string!)
  - StartDate (date!)
  - EndDate (date!)
  - CampaignType (number!)
  - Count (string?)
  - MemoId (number?)
  - ApplyPrice (number?)
定数値情報:
- HumanSexes (対象: OrderCustomer.Sex)
  - 0: NotKnown / 不明
  - 1: Male / 男性
  - 2: Female / 女性
  - 9: NotApplicable / 適用不能
- OrderLineType (対象: OrderLine.OrderType, ReturnOrderLine.OrderType)
  - 0: NormalOrder / 通常
  - 1: Canceled / キャンセル
  - 2: Returned / 返品
  - 3: Benefit / 特典
  - 4: SetProductParent / セット親
  - 5: SetProductChild / セット子
  - 6: Unavailable / 無効
  - 7: Additional / 追加商品
- OrderType (対象: PurchaseOrder.OrderType, ReturnOrder.OrderType)
  - 0: Normal / 旧譜
  - 1: Reserve / 新譜
  - 2: Pending / 取り置き
  - 3: Bto / 受注生産品
  - 4: ReSend / 返品再送
- OrderStatus (対象: PurchaseOrder.OrderStatus, ReturnOrder.OrderStatus)
  - 0: Accept / 受注
  - 1: UnAllocate / 未引当
  - 2: PartAllocate / 中途引当
  - 3: Allocated / 引当済
  - 6: ShipIndicate / 出荷指示済
  - 7: Ship / 出荷済
  - 8: Arrival / 着荷済
  - 9: Booked / 計上済
  - 10: Cancel / キャンセル
  - 11: ShipSuspend / 出荷保留
  - 12: WorkerProcessing / バッチ処理中
- PaymentMethod (対象: PurchaseOrder.PaymentMethod, ReturnOrder.PaymentMethod)
  - 0: CashOnDelivery / 代引き
  - 1: CreditCard / クレジットカード / カード番号をサーバーに通過させる場合(要PCIDSS認証取得)
  - 2: PointOnly / 全額ポイント
  - 3: Cvs / コンビニ決済
  - 4: Offsite / オフサイト / 銀行振込など、サイト内で決済処理を行わない決済
  - 5: Offline / オフライン / 実店舗での決済など、オンラインでない決済
  - 6: Token / トークン / トークン方式のカード決済等
  - 100: External / 外部
  - -1: None / None
- PaymentStatus (対象: PurchaseOrder.PaymentStatus, ReturnOrder.PaymentStatus)
  - 0: Ready / 未決済
  - 1: Success / 決済成功
  - 2: Fail / 決済失敗
  - 3: Processing / 処理中
  - 4: Cancel / キャンセル
- MemberStatus (対象: OrderCustomers.MemberStatus)
  - 0: Default / 通常
  - 1: Attention / 注意
  - 2: Warning / 警告
- ReserveRequestStatus (対象: ReserveRequests.Status, ReserveRequestStatusCounts.Status)
  - 0: Accept / 入荷待ち / 標準動作では、販売可能数が更新された際に自動で、「引当済み」に変更されます。
  - 1: Allocated / 引当済み
"""},
    {"role": "user", "content": "未出荷の受注を古い順"},
]

text = tokenizer.apply_chat_template(
    messages, tokenize=False, add_generation_prompt=True
)
inputs = tokenizer([text], return_tensors="pt").to("cuda")

outputs = model.generate(
    **inputs, max_new_tokens=128, temperature=0.7, do_sample=True
)
response = tokenizer.decode(outputs[0], skip_special_tokens=True)

print(response)

出力結果はいい感じ。

$filter=OrderStatus eq 7 and ShipDate ge null&$orderby=OrderDate desc

まとめ

ファインチューニングはできた。出力も、余計な説明や装飾がなく指定したものだけが出るようにできているので追学習の効果が出てると思われる。

ただし、出力結果のODataクエリは構文に間違いが無いが意味に間違いがある。 これは、用意した学習データにはこのPurchaseOrderに対する教師データが1件もなく、近しい学習データがなかったために思う。未出荷が指し示すものがOrderStatusが7未満であるというのは訓練しないとわからないのだろう。学習データを増やせば良くなりそうだ。ただし、未出荷なのに7=出荷済を選択していることには少し潜在的な知能不足を感じるので学習データだけではなんともならないのかもしれない。