Skip to content

MDBF QAT - #42

Open
fujisawa-yoshihiko wants to merge 3 commits into
FujitsuResearch:export/global_ptqfrom
fujisawa-yoshihiko:feature/mdbf-qat-globalptq
Open

MDBF QAT#42
fujisawa-yoshihiko wants to merge 3 commits into
FujitsuResearch:export/global_ptqfrom
fujisawa-yoshihiko:feature/mdbf-qat-globalptq

Conversation

@fujisawa-yoshihiko

Copy link
Copy Markdown

Summary

onecomp-globalptq の KL 蒸留パイプラインに MDBF 量子化モデルの QAT 経路を追加します。MDBF の振幅パラメータと ±1 符号行列を smooth sign STE で同時に学習できるようになります。

ライブラリ本体のみの変更で、既存の GPTQ / DBF 経路の挙動は変えていません。

1. MDBF 微分可能アダプタ (b8569af)

背景・動機

global_ptq の KL 蒸留ループは、量子化方式ごとのアダプタが「学習対象パラメータの取り出し → 微分可能 forward への差し替え → 学習後の書き戻し」を担う設計になっています。MDBF は重みを

W ≈ Σ_p  (A_sign ⊙ A_amp·Q_U_amp^T) · (B_sign ⊙ Q_V_amp·B_amp^T)

±1 符号行列と多段振幅の積で表すため、既存の GPTQ / DBF 用アダプタでは学習対象を取り出せません。MDBF 専用のアダプタが必要でした。

変更内容(新規 _core/mdbf_adapter.py、+319)

  • MultipathMDBFLinear の forward を微分可能な実装に差し替え、パスごとの振幅パラメータ(A_amp, B_amp, Q_U_amp, Q_V_amp)を fp32 の nn.Parameter として最適化対象に昇格
  • optimize_binary=True のとき、符号行列 A_sign / B_sign を float として露出し、smooth sign STE(mdbf_ste_k、既定 2.0)で学習
  • setup_mdbf_differentiable() / restore_mdbf_original() で元の forward へ復帰可能
  • 学習結果の書き戻しは write_back_mdbf_amp() / write_back_mdbf_binary()(packed / fp16 バッファへ)、ロールバック用の状態保存・復元は save_mdbf_state() / load_mdbf_state()

2. QATへの接続 (b9940cc)

背景・動機

3 点あります。

  1. 蒸留ループは detect_quantization_method() で量子化方式を自動判別しますが、MDBF が候補に無いため、MDBF 量子化済みモデルを渡しても「量子化層なし」として素通りしていました
  2. 1bpw の Llama2-7B / 13B を QAT するには、生徒を DeepSpeed ZeRO-2 + CPU optimizer offload に載せたうえで、FP16 教師を別 GPU に逃がす必要があります(実測で GPU0 ≈52.5GB / GPU1 ≈28GB)。従来は生徒と教師が同一デバイス前提で、この構成が組めませんでした
  3. 符号 STE の鋭さ k は論文値に近づけるための主要な探索軸のため、設定として露出させる必要がありました(今回の実測は k=2.0、次段で k=100 を試す予定)

変更内容

  • _core/helpers.py: detect_quantization_method()MultipathMDBFLinear"mdbf" として検出。優先度は GPTQ > DBF > MDBF で、混在時は警告を出して最優先のもののみ返します(従来の GPTQ / DBF の挙動は不変)
  • _core/core.py: run_kl_distillation() に mdbf 分岐を追加(setup / teardown は上記アダプタ経由、振幅・符号のパラメータ群は dbf_lr を使用)。eval_kl()teacher_device を生徒と分離して受け取れるように変更
  • _core/trainer.py: _GlobalPTQTrainer.compute_loss() も同様に教師デバイスの分離に対応
  • global_ptq.py / global_ptq_distributed.py: mdbf_ste_k / student_device / teacher_device を設定として追加し、run_kl_distillation および Trainer ベースの分散経路へ引き渡し

依存関係についてのご注意

本 PR のコードは onecomp/quantizer/mdbf/(MDBF 量子化器本体)に依存しますが、このブランチには MDBF 本体を含めていません。MDBF は既に #28feature/mdbf にマージ済みのため、そちらとの統合をお願いしたく思います。

パッケージが MDBF 非搭載でも壊れないよう、global_ptq 内の MDBF 参照はすべて関数内の遅延 import にしてあります。

  • onecomp_globalptq の import、および既存の gptq / dbf 経路は MDBF 非搭載でもそのまま動作します(動作確認済み)
  • MDBF 経路を実行したときのみ ImportError になります
feature/mdbf へマージする場合の解決手順(コンフリクト 3 ファイル)

export/global_ptqfeature/mdbf の分岐に由来するもので、本 PR の 2 コミットが原因ではありません。衝突するのは CHANGELOG.md / onecomp/quantizer/__init__.py / onecomp/quantized_model_loader.py の 3 ファイルで、いずれも「MDBF 側の追加」と「OneBit 側の追加」が並列に入っただけなので、両方を残すのが正しい解決です。

  • CHANGELOG.md: 両方のエントリを残す
  • onecomp/quantizer/__init__.py: feature/mdbf 側(HEAD)が上位集合なので HEAD を採用

onecomp/quantized_model_loader.py は 6 ハンクありますが、最後の 2 ハンクは elif の分岐が共通の末尾(from_saved_state(...) の引数列)を共有しているため、機械的に両側を連結すると SyntaxError になります。両分岐をそれぞれ完全な形で書き出してください。

            elif effective_method == "mdbf":
                layer_target_bits = resolve_mdbf_layer_bits(name, quant_config)
                quantized_module = MultipathMDBFLinear.from_saved_state(
                    layer_sd,
                    in_features=in_features,
                    out_features=out_features,
                    empty=True,
                    target_bits=layer_target_bits,
                )
            elif effective_method == "onebit":
                quantized_module = OneBitLinear.from_saved_state(
                    layer_sd,
                    in_features=in_features,
                    out_features=out_features,
                    empty=True,
                )

残りの 3 ハンクは以下のとおりです。

# import(両方残す)
from .quantizer.mdbf.config import resolve_mdbf_layer_bits
from .quantizer.mdbf.mdbf_layer import MultipathMDBFLinear
from .quantizer.onebit.onebit_layer import OneBitLinear
from .utils.device import get_default_device

# layers_cls の分岐(両方残す)
            elif effective_method == "mdbf":
                layers_cls = [MultipathMDBFLinear]
            elif effective_method == "onebit":
                layers_cls = [OneBitLinear]

# skip_types(4 クラスすべて)
        skip_types = (
            GPTQLinear,
            DoubleBinaryLinear,
            MultipathMDBFLinear,
            OneBitLinear,
        )

マージ後の確認:

python -c "
from onecomp.quantizer.mdbf import MDBF
from onecomp_globalptq import GlobalPTQ, GlobalPTQDistributed
from onecomp_globalptq.global_ptq._core import mdbf_adapter
print('OK')
"

変更ファイル

ファイル 変更内容
global_ptq/…/_core/mdbf_adapter.py 新規 (+319)。MDBF 微分可能アダプタ
global_ptq/…/_core/core.py KL 蒸留への mdbf 分岐、teacher_device の分離
global_ptq/…/_core/helpers.py detect_quantization_method の MDBF 対応
global_ptq/…/_core/trainer.py 教師デバイス分離への対応
global_ptq/…/global_ptq.py mdbf_ste_k / student_device / teacher_device の追加
global_ptq/…/global_ptq_distributed.py mdbf_ste_k / teacher_device の追加と分散経路への引き渡し

検証

本 PR の QAT 経路を用いて、Llama2-7B / 13B で MDBF 論文プロトコル(1.00 BPW)を実行し、PTQ・QAT とも完走することを確認しています。

実行環境: TSUBAME4 node_h, H100 80GB × 2(QAT 時 GPU0 ≈52.5GB / GPU1 ≈28GB(FP16 教師))、DeepSpeed ZeRO-2

設定: target_bits=1.0, P=1, ADMM iters=1000 / inner=3 / reg=0.03, gradient refine iters=1500 / lr=0.01, activation_aware=True, act_init="osvd", QEP percdamp=0.01 / perccorr=0.5, 先頭 4 + 末尾 4 ブロックを FP16 除外, calibration = C4 512 サンプル × 2048 トークン。Table 1 は l=8 の PTQ のみ、Table 2 は l=3 + QAT(KL 蒸留のみ w_distill=1.0 / w_ntp=0.0, optimize_binary=True, epochs=5, dbf_lr=5e-5, mdbf_ste_k=2.0, calibration_strategy="drop_rand")。

7B PTQ (T1) 13B PTQ (T1) 7B QAT (T2) 13B QAT (T2)
WikiText-2 PPL 47.23 38.65 16.03 13.85
C4 PPL 33.80 29.83 17.39 15.28
PTB PPL 180.97 767.59 88.72 384.57
PPL 平均 87.33 278.69 40.71 137.90
ACC 平均(6 タスク) 0.4298 0.4315 0.4685 0.4929
論文 PPL 参照 99.26 144.89 — (WT2: 9.36)
  • PTQ は論文値と同等以上(7B: 87.33 vs 99.26)
  • QAT により WikiText-2 PPL が 7B で 47.23 → 16.03、13B で 38.65 → 13.85 に改善。ACC も +3.9pt / +6.1pt
  • 13B は WikiText-2 / C4 で 7B を上回るが、1bpw ではパラメータ約 2 倍に対する改善は限定的

注記: これらを実行した論文プロトコル再現スクリプトと、そのために必要だった評価用データセット読み込みの修正(onecomp/utils/perplexity.py ほか)は、本 PR のスコープ外としています。必要でしたら別 PR で提出します。

Test plan

  • MDBF 非搭載環境で onecomp_globalptq が import でき、既存の GPTQ / DBF 経路が従来どおり動作すること
  • detect_quantization_method() が MDBF モデルで "mdbf" を返し、GPTQ 混在時は "gptq" を優先すること
  • feature/mdbf とマージした状態で MDBF QAT(optimize_binary=True)が完走すること
  • teacher_device を CPU / 別 GPU に指定して KL 蒸留が動作すること
  • write_back_mdbf_amp() / write_back_mdbf_binary() 後の量子化モデルが保存・再読み込みできること

既知の制限・今後

  • 論文の QAT 値(WikiText-2 PPL 9.36)には未到達(現状 13.85)。主因は QAT 学習データ量と見ており、データ量拡大 → SmoothSign の k を 2.0 → 100 → 中間層蒸留の順で追試予定です
  • 中間層蒸留(use_inter_loss / lambda_inter)は GlobalPTQ にのみ実装されており、GlobalPTQDistributed 側には未接続です
  • global_ptq/ 配下の既存ファイルは base の時点で black / isort 未適用のため、本 PR でも整形はかけていません(無関係な差分を増やさないため)。新規追加した mdbf_adapter.py は black(--line-length=99)/ isort 適用済みです

スコープ外(別途提出を検討)

  • MDBF 量子化器本体(onecomp/quantizer/mdbf/)— feature/mdbf: gemlite対応とHessian/scale_bits修正 #28feature/mdbf にマージ済み
  • Llama2 論文プロトコル再現スクリプトと、評価用 HF データセット読み込みの修正
  • calibration_dataset に事前ロード済みテキスト(list[str])を渡せるようにする拡張
  • 符号反転を損失で検証する座標降下(CD)ステップ

fujisawa-yoshihiko and others added 3 commits July 29, 2026 21:38
Introduce the piece needed to make MultipathMDBFLinear layers trainable
during GlobalPTQ-style QAT:

- mdbf_adapter.py: promotes per-path amplitude parameters (A_amp, B_amp,
  Q_U_amp, Q_V_amp) to fp32 nn.Parameter and, when optimize_binary=True,
  exposes +/-1 sign matrices as float shadow tensors trained through a
  smooth sign STE. Provides differentiable forward reconstruction,
  write-back to packed/fp16 buffers for inference, and state
  snapshot/restore for rollback.

This module is not wired into the training loop yet; see the following
commit for integration into GlobalPTQ/GlobalPTQDistributed.

Co-authored-by: Cursor <cursoragent@cursor.com>
- helpers.detect_quantization_method: recognize MultipathMDBFLinear
  layers as method "mdbf" (priority gptq > dbf > mdbf). Without this a
  MDBF-quantized model was reported as having no quantized layers and
  the distillation loop skipped it entirely.
- core.run_kl_distillation: add the "mdbf" branch (setup/teardown via
  mdbf_adapter, param groups for amplitude/binary params using dbf_lr)
  and expose mdbf_ste_k for the sign STE sharpness.
- core.eval_kl / trainer._GlobalPTQTrainer.compute_loss: support a
  separate teacher_device so the FP16 teacher can be kept on CPU or a
  second GPU while the student trains under DeepSpeed ZeRO-2 with CPU
  optimizer offload. Required to fit 1bpw Llama2-7B/13B QAT on 80GB
  cards.
- global_ptq.py / global_ptq_distributed.py: expose the corresponding
  GlobalPTQConfig(Distributed) fields (mdbf_ste_k, student_device,
  teacher_device) and thread them through to run_kl_distillation and
  the Trainer-based distributed path.

Co-authored-by: Cursor <cursoragent@cursor.com>
detect_quantization_method() imported MultipathMDBFLinear
unconditionally, so on an installation without onecomp.quantizer.mdbf
*every* call raised ModuleNotFoundError -- including for plain GPTQ and
DBF models.  On this PR's base branch that broke the existing
TestDetectQuantizationMethod tests (3 failures).

Guard the import and treat MDBF as absent when it cannot be imported,
so GPTQ/DBF detection keeps working without the MDBF quantizer.

Verified both ways: without MDBF the global_ptq unit tests pass
(69 passed, integration tests deselected), and with MDBF present
(merged with feature/mdbf) detection still resolves MultipathMDBFLinear.
@fujisawa-yoshihiko
fujisawa-yoshihiko force-pushed the feature/mdbf-qat-globalptq branch from 75fec39 to 85896e2 Compare July 29, 2026 13:57
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant