From b8569af9114bc49aecf1dd331cea25385d0a57fa Mon Sep 17 00:00:00 2001 From: fujisawa-yoshihiko Date: Thu, 23 Jul 2026 21:14:51 +0900 Subject: [PATCH 1/3] feat(global_ptq): add MDBF differentiable adapter 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 --- .../global_ptq/_core/mdbf_adapter.py | 319 ++++++++++++++++++ 1 file changed, 319 insertions(+) create mode 100644 global_ptq/onecomp_globalptq/global_ptq/_core/mdbf_adapter.py diff --git a/global_ptq/onecomp_globalptq/global_ptq/_core/mdbf_adapter.py b/global_ptq/onecomp_globalptq/global_ptq/_core/mdbf_adapter.py new file mode 100644 index 0000000..47ef14d --- /dev/null +++ b/global_ptq/onecomp_globalptq/global_ptq/_core/mdbf_adapter.py @@ -0,0 +1,319 @@ +"""MDBF differentiable parameter management for global PTQ. + +Makes MultipathMDBFLinear layers trainable by exposing amplitude parameters +(A_amp, B_amp, Q_U_amp, Q_V_amp) per path as optimisable tensors, and +optionally enabling differentiable binary-sign optimisation via smooth sign STE. + +Architecture recap (per MDBFLinear path): + F = A_sign * (A_amp @ Q_U_amp^T) shape: (n, r) + G = B_sign * (Q_V_amp @ B_amp^T) shape: (r, m) + y = x @ G^T @ F^T + +Trainable (continuous): + A_amp, B_amp, Q_U_amp, Q_V_amp — amplitude/scale factors per path. +Trainable (discrete, opt-in): + A_sign, B_sign — ±1 binary factor matrices per path, via sign STE. + +Copyright 2025-2026 Fujitsu Ltd. + +Authors: Keiji Kimura + +""" + +from types import MethodType +from typing import Dict, List, Tuple + +import torch +import torch.nn as nn + +from .helpers import smooth_sign_ste + +_AMP_ATTRS = ("A_amp", "B_amp", "Q_U_amp", "Q_V_amp") +_BINARY_SIGN_NAMES = ("A", "B") + +# Sharpness for sign-STE: same rationale as DBF adapter (values near ±1, +# tanh saturation avoided by using k=2 instead of the GPTQ default k=100). +_BINARY_STE_K = 2.0 + + +# --------------------------------------------------------------------------- +# Finding MDBF modules +# --------------------------------------------------------------------------- + + +def find_mdbf_modules(model: nn.Module) -> List[Tuple[str, nn.Module]]: + """Return all ``MultipathMDBFLinear`` modules as ``(name, module)`` pairs.""" + from onecomp.quantizer.mdbf.mdbf_layer import MultipathMDBFLinear + + from .helpers import find_target_modules + + return find_target_modules(model, MultipathMDBFLinear) + + +# --------------------------------------------------------------------------- +# Differentiable forward (per MDBFLinear path) +# --------------------------------------------------------------------------- + + +def _make_mdbf_differentiable_forward(): + """Build a differentiable ``forward`` for a ``MultipathMDBFLinear``. + + Each path's computation graph is reconstructed from the (possibly + optimisable) amplitude parameters and, when ``_opt_A_sign_{p}`` / + ``_opt_B_sign_{p}`` exist, through :func:`smooth_sign_ste` so that + gradients flow to the float sign-weight tensors as well. + """ + from onecomp.quantizer.mdbf.mdbf_layer import unpack_binary + + def differentiable_forward(self, x: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + k = getattr(self, "_binary_ste_k", _BINARY_STE_K) + + y = None + for path in self.paths: + # ---- amplitude parameters (always trainable) ---- + A_amp = getattr(path, "_opt_A_amp", path.A_amp).to(dtype) + B_amp = getattr(path, "_opt_B_amp", path.B_amp).to(dtype) + Q_U_amp = getattr(path, "_opt_Q_U_amp", path.Q_U_amp).to(dtype) + Q_V_amp = getattr(path, "_opt_Q_V_amp", path.Q_V_amp).to(dtype) + + # ---- sign matrices (STE or packed buffer) ---- + if hasattr(path, "_opt_A_sign"): + A_sign = smooth_sign_ste(path._opt_A_sign, k=k).to(dtype) + else: + A_sign = unpack_binary(path._packed_sign("A", x.device), (path.n, path.r)).to( + dtype + ) + + if hasattr(path, "_opt_B_sign"): + B_sign = smooth_sign_ste(path._opt_B_sign, k=k).to(dtype) + else: + B_sign = unpack_binary(path._packed_sign("B", x.device), (path.r, path.m)).to( + dtype + ) + + # F = A_sign * (A_amp @ Q_U_amp^T) shape: (n, r) + F = A_sign * (A_amp @ Q_U_amp.T) + # G = B_sign * (Q_V_amp @ B_amp^T) shape: (r, m) + G = B_sign * (Q_V_amp @ B_amp.T) + + # y += x @ G^T @ F^T + path_out = x @ G.T @ F.T + + y = path_out if y is None else y + path_out + + if self.bias is not None: + y = y + self.bias.to(dtype) + return y + + return differentiable_forward + + +# --------------------------------------------------------------------------- +# Parameter setup +# --------------------------------------------------------------------------- + + +def setup_mdbf_differentiable( + mdbf_modules: List[Tuple[str, nn.Module]], + optimize_binary: bool = False, + ste_k: float = _BINARY_STE_K, +) -> Tuple[Dict[str, object], List[torch.Tensor], List[torch.Tensor]]: + """Make MDBF amplitude (and optionally sign) parameters trainable. + + For each ``MultipathMDBFLinear`` module, the original ``forward`` is + replaced with a differentiable version. Amplitude parameters + (A_amp, B_amp, Q_U_amp, Q_V_amp) of every path are promoted to + float32 ``nn.Parameter`` objects stored as ``_opt_*`` attributes on the + individual ``MDBFLinear`` path modules. + + Args: + mdbf_modules: List of ``(name, module)`` pairs from + :func:`find_mdbf_modules`. + optimize_binary: When True, also expose unpacked ±1 sign matrices + as float tensors with ``requires_grad=True`` so that gradients + flow through :func:`smooth_sign_ste`. + ste_k: Sharpness for binary sign STE (``tanh(k*x)`` backward). + Default is :data:`_BINARY_STE_K` (2.0). + + Returns: + ``(original_forwards, amp_params, binary_params)`` + + *original_forwards* maps module name → original forward (for restore). + *amp_params* is a flat list of float32 ``nn.Parameter`` objects. + *binary_params* is a flat list of float tensors (empty when + *optimize_binary* is ``False``). + """ + from onecomp.quantizer.mdbf.mdbf_layer import unpack_binary + + original_forwards: Dict[str, object] = {} + amp_params: List[torch.Tensor] = [] + binary_params: List[torch.Tensor] = [] + + for name, mod in mdbf_modules: + for path in mod.paths: + # ---- continuous amplitude parameters ---- + for attr in _AMP_ATTRS: + buf = getattr(path, attr) + fp32 = buf.data.detach().clone().float() + new_param = nn.Parameter(fp32, requires_grad=True) + setattr(path, f"_opt_{attr}", new_param) + amp_params.append(new_param) + + # ---- discrete sign parameters (optional) ---- + if optimize_binary: + for which in _BINARY_SIGN_NAMES: + shape = (path.n, path.r) if which == "A" else (path.r, path.m) + packed_key = f"{which}_sign_packed" + packed = path._buffers.get(packed_key) + if packed is None: + # GemLite mode: stashed on CPU + packed = path._packed_cpu.get(which) + if packed is None: + continue + unpacked = unpack_binary(packed, shape).float().detach().clone() + new_param = nn.Parameter(unpacked, requires_grad=True) + setattr(path, f"_opt_{which}_sign", new_param) + binary_params.append(new_param) + + original_forwards[name] = mod.forward + mod._binary_ste_k = float(ste_k) + mod.forward = MethodType(_make_mdbf_differentiable_forward(), mod) + + return original_forwards, amp_params, binary_params + + +# --------------------------------------------------------------------------- +# Forward restore / re-install +# --------------------------------------------------------------------------- + + +def restore_mdbf_original( + mdbf_modules: List[Tuple[str, nn.Module]], + original_forwards: Dict[str, object], + cleanup: bool = False, +) -> None: + """Restore every module's original ``forward`` method.""" + for name, mod in mdbf_modules: + if name in original_forwards: + mod.__dict__.pop("forward", None) + if not hasattr(mod, "forward") or mod.forward != original_forwards[name]: + mod.forward = original_forwards[name] + + if cleanup: + if hasattr(mod, "_binary_ste_k"): + delattr(mod, "_binary_ste_k") + for path in mod.paths: + for attr in _AMP_ATTRS: + opt_attr = f"_opt_{attr}" + if hasattr(path, opt_attr): + delattr(path, opt_attr) + for which in _BINARY_SIGN_NAMES: + opt_attr = f"_opt_{which}_sign" + if hasattr(path, opt_attr): + delattr(path, opt_attr) + + +def setup_mdbf_forwards_only( + mdbf_modules: List[Tuple[str, nn.Module]], + original_forwards: Dict[str, object], +) -> None: + """Re-install differentiable forwards for continued training after eval.""" + for name, mod in mdbf_modules: + if name not in original_forwards: + original_forwards[name] = mod.forward + mod.forward = MethodType(_make_mdbf_differentiable_forward(), mod) + + +# --------------------------------------------------------------------------- +# Write-back +# --------------------------------------------------------------------------- + + +def write_back_mdbf_binary(mdbf_modules: List[Tuple[str, nn.Module]]) -> None: + """Write optimised float sign tensors back to packed uint8 buffers.""" + from onecomp.quantizer.mdbf.mdbf_layer import pack_binary + + with torch.no_grad(): + for _name, mod in mdbf_modules: + for path in mod.paths: + for which in _BINARY_SIGN_NAMES: + opt_attr = f"_opt_{which}_sign" + if not hasattr(path, opt_attr): + continue + w = getattr(path, opt_attr) + q = w.sign() + q[q == 0] = 1 + shape = (path.n, path.r) if which == "A" else (path.r, path.m) + packed, _ = pack_binary(q.to(torch.int8).reshape(shape)) + buf_key = f"{which}_sign_packed" + if buf_key in path._buffers: + path._buffers[buf_key].copy_(packed.to(path._buffers[buf_key].device)) + elif which in path._packed_cpu: + path._packed_cpu[which].copy_(packed.cpu()) + + +def write_back_mdbf_amp(mdbf_modules: List[Tuple[str, nn.Module]]) -> None: + """Copy float32 optimised amp params back to fp16 buffers for inference.""" + with torch.no_grad(): + for _name, mod in mdbf_modules: + for path in mod.paths: + for attr in _AMP_ATTRS: + opt_attr = f"_opt_{attr}" + if not hasattr(path, opt_attr): + continue + opt_param = getattr(path, opt_attr) + buf = getattr(path, attr) + buf.copy_(opt_param.data.half()) + + +# --------------------------------------------------------------------------- +# State save / load (for rollback) +# --------------------------------------------------------------------------- + + +def save_mdbf_state(mdbf_modules: List[Tuple[str, nn.Module]]) -> Dict: + """Snapshot amplitude buffers and packed sign buffers.""" + state: Dict[str, dict] = {} + for name, mod in mdbf_modules: + paths_state = {} + for p, path in enumerate(mod.paths): + d: dict = {} + for attr in _AMP_ATTRS: + d[attr] = getattr(path, attr).data.clone() + for which in _BINARY_SIGN_NAMES: + buf_key = f"{which}_sign_packed" + if buf_key in path._buffers: + d[buf_key] = path._buffers[buf_key].clone() + elif which in path._packed_cpu: + d[buf_key] = path._packed_cpu[which].clone() + paths_state[p] = d + state[name] = paths_state + return state + + +def load_mdbf_state( + mdbf_modules: List[Tuple[str, nn.Module]], + state: Dict, +) -> None: + """Restore a previously saved snapshot.""" + with torch.no_grad(): + for name, mod in mdbf_modules: + if name not in state: + continue + paths_state = state[name] + for p, path in enumerate(mod.paths): + if p not in paths_state: + continue + d = paths_state[p] + for attr in _AMP_ATTRS: + if attr in d: + getattr(path, attr).copy_(d[attr]) + for which in _BINARY_SIGN_NAMES: + buf_key = f"{which}_sign_packed" + if buf_key not in d: + continue + if buf_key in path._buffers: + path._buffers[buf_key].copy_(d[buf_key]) + elif which in path._packed_cpu: + path._packed_cpu[which].copy_(d[buf_key].cpu()) From b9940cc97a598a6dce43659c060cbb538f01fc48 Mon Sep 17 00:00:00 2001 From: fujisawa-yoshihiko Date: Thu, 23 Jul 2026 21:15:02 +0900 Subject: [PATCH 2/3] feat(global_ptq): wire MDBF QAT into the KL distillation loop - 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 --- .../global_ptq/_core/core.py | 110 +++++++++++++++--- .../global_ptq/_core/helpers.py | 22 ++-- .../global_ptq/_core/trainer.py | 34 +++++- .../global_ptq/global_ptq.py | 13 ++- .../global_ptq/global_ptq_distributed.py | 79 +++++++++++-- 5 files changed, 223 insertions(+), 35 deletions(-) diff --git a/global_ptq/onecomp_globalptq/global_ptq/_core/core.py b/global_ptq/onecomp_globalptq/global_ptq/_core/core.py index e5a5e3f..9fa2cb9 100644 --- a/global_ptq/onecomp_globalptq/global_ptq/_core/core.py +++ b/global_ptq/onecomp_globalptq/global_ptq/_core/core.py @@ -54,6 +54,15 @@ write_back_dbf_binary, write_back_dbf_scaling, ) +from .mdbf_adapter import ( + load_mdbf_state, + restore_mdbf_original, + save_mdbf_state, + setup_mdbf_differentiable, + setup_mdbf_forwards_only, + write_back_mdbf_amp, + write_back_mdbf_binary, +) logger = getLogger(__name__) @@ -418,6 +427,20 @@ def cosine_warmup_lr_lambda( # --------------------------------------------------------------------------- +@torch.no_grad() +def _teacher_logits( + teacher_model: nn.Module, + input_ids: torch.Tensor, + teacher_dev: torch.device, + student_dev: torch.device, +) -> torch.Tensor: + """Run teacher forward; move logits to *student_dev* if devices differ.""" + if teacher_dev == student_dev: + return get_logits(teacher_model(input_ids)) + logits_t = get_logits(teacher_model(input_ids.to(teacher_dev))) + return logits_t.to(student_dev) + + @torch.no_grad() def eval_kl( model: nn.Module, @@ -425,10 +448,12 @@ def eval_kl( dataloader: List[Dict[str, torch.Tensor]], dev: torch.device, temperature: float = 1.0, + teacher_dev: Optional[torch.device] = None, ) -> float: """Mean KL divergence over *dataloader* batches.""" was_training = model.training model.eval() + teacher_dev = teacher_dev or dev total, n = 0.0, 0 for batch in dataloader: input_ids = batch["input_ids"].to(dev) @@ -437,7 +462,7 @@ def eval_kl( attention_mask = attention_mask.to(dev) logits_s = get_logits(model(input_ids)) - logits_t = get_logits(teacher_model(input_ids)) + logits_t = _teacher_logits(teacher_model, input_ids, teacher_dev, dev) total += compute_kl_loss( logits_t, logits_s, temperature, attention_mask=attention_mask, ).item() @@ -587,6 +612,7 @@ def run_kl_distillation( gptq_intweight_lr: float = 1e-4, optimize_binary: bool = False, ste_k: float = 100.0, + mdbf_ste_k: float = 2.0, calibration_dataset=None, num_calibration_samples: int = 128, max_length: int = 2048, @@ -616,12 +642,17 @@ def run_kl_distillation( early_stopping_patience: int = 0, use_mixed_precision: bool = False, grad_accum_steps: int = 1, + student_device: Optional[str] = None, + teacher_device: Optional[str] = None, ) -> Dict: """Run KL-distillation global PTQ on a GPTQ or DBF quantized model. The model is modified **in-place**. Returns a results dict. """ - dev = torch.device("cuda" if torch.cuda.is_available() else "cpu") + dev = torch.device( + student_device or ("cuda" if torch.cuda.is_available() else "cpu") + ) + teacher_dev = torch.device(teacher_device) if teacher_device else dev # ------------------------------------------------------------------ # 1. Detect method @@ -631,7 +662,7 @@ def run_kl_distillation( logger.warning("No quantized layers detected — skipping global PTQ.") return {"global_executed": False, "reason": "not_quantized"} - if method not in ("gptq", "dbf"): + if method not in ("gptq", "dbf", "mdbf"): logger.info("Method '%s' detected — not supported.", method) return {"global_executed": False, "reason": f"unsupported_method_{method}"} @@ -665,7 +696,8 @@ def run_kl_distillation( teacher_model.eval() for p in teacher_model.parameters(): p.requires_grad = False - teacher_model.to(dev) + if teacher_dev.type != "cpu": + teacher_model.to(teacher_dev) # ------------------------------------------------------------------ # 4. Move student to GPU and set up differentiable parameters @@ -675,6 +707,7 @@ def run_kl_distillation( gptq_modules: list = [] dbf_modules: list = [] + mdbf_modules: list = [] original_forwards: Dict[str, object] = {} param_groups: list = [] binary_params: list = [] @@ -710,6 +743,24 @@ def run_kl_distillation( f", {len(binary_params)} binary" if binary_params else "", ) + elif method == "mdbf": + mdbf_modules = detected_modules + original_forwards, scaling_params, binary_params = setup_mdbf_differentiable( + mdbf_modules, optimize_binary, ste_k=mdbf_ste_k, + ) + logger.info("MDBF binary STE sharpness mdbf_ste_k=%.4g", mdbf_ste_k) + all_mdbf_params = list(scaling_params) + if binary_params: + all_mdbf_params += binary_params + param_groups = [{"params": all_mdbf_params, "lr": dbf_lr}] + + logger.info( + "Trainable: %d amp params%s across %d MDBF modules", + len(scaling_params), + f", {len(binary_params)} binary" if binary_params else "", + len(mdbf_modules), + ) + total_trainable = sum(len(pg["params"]) for pg in param_groups) if total_trainable == 0: logger.warning("No trainable parameters — skipping.") @@ -717,6 +768,8 @@ def run_kl_distillation( restore_gptq_original(gptq_modules, original_forwards) elif method == "dbf": restore_dbf_original(dbf_modules, original_forwards) + elif method == "mdbf": + restore_mdbf_original(mdbf_modules, original_forwards) quantized_model.cpu() del teacher_model gc.collect() @@ -871,17 +924,24 @@ def run_kl_distillation( if method == "gptq": initial_state = save_gptq_state(gptq_modules) restore_gptq_original(gptq_modules, original_forwards) - else: + elif method == "dbf": initial_state = save_dbf_state(dbf_modules) restore_dbf_original(dbf_modules, original_forwards) + else: # mdbf + initial_state = save_mdbf_state(mdbf_modules) + restore_mdbf_original(mdbf_modules, original_forwards) - initial_kl = eval_kl(quantized_model, teacher_model, dataloader, dev, temperature) + initial_kl = eval_kl( + quantized_model, teacher_model, dataloader, dev, temperature, teacher_dev, + ) logger.info("Initial KL = %.6f", initial_kl) if method == "gptq": setup_gptq_forwards_only(gptq_modules, original_forwards, gptq_optimize_intweight) elif method == "dbf": setup_dbf_forwards_only(dbf_modules, original_forwards) + elif method == "mdbf": + setup_mdbf_forwards_only(mdbf_modules, original_forwards) # ------------------------------------------------------------------ # 7. Training loop @@ -928,8 +988,9 @@ def _forward_and_loss() -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: with amp_ctx: logits_s = get_logits(quantized_model(input_ids)) - with torch.no_grad(): - logits_t = get_logits(teacher_model(input_ids)) + logits_t = _teacher_logits( + teacher_model, input_ids, teacher_dev, dev, + ) kl = compute_kl_loss( logits_t, logits_s, temperature, attention_mask=attention_mask, @@ -1052,16 +1113,24 @@ def _forward_and_loss() -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: elif method == "dbf": write_back_dbf_binary(dbf_modules) restore_dbf_original(dbf_modules, original_forwards) + elif method == "mdbf": + write_back_mdbf_binary(mdbf_modules) + write_back_mdbf_amp(mdbf_modules) + restore_mdbf_original(mdbf_modules, original_forwards) - current_kl = eval_kl(quantized_model, teacher_model, dataloader, dev, temperature) + current_kl = eval_kl( + quantized_model, teacher_model, dataloader, dev, temperature, teacher_dev, + ) if current_kl < best_kl: best_kl = current_kl patience_counter = 0 if method == "gptq": best_state = save_gptq_state(gptq_modules) - else: + elif method == "dbf": best_state = save_dbf_state(dbf_modules) + else: # mdbf + best_state = save_mdbf_state(mdbf_modules) else: patience_counter += 1 @@ -1069,6 +1138,8 @@ def _forward_and_loss() -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: setup_gptq_forwards_only(gptq_modules, original_forwards, gptq_optimize_intweight) elif method == "dbf": setup_dbf_forwards_only(dbf_modules, original_forwards) + elif method == "mdbf": + setup_mdbf_forwards_only(mdbf_modules, original_forwards) # Restore non-EMA params for continued training if ema_tracker is not None: @@ -1103,15 +1174,19 @@ def _forward_and_loss() -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: if best_state is not None and best_kl < initial_kl: if method == "gptq": load_gptq_state(gptq_modules, best_state) - else: + elif method == "dbf": load_dbf_state(dbf_modules, best_state) + else: # mdbf + load_mdbf_state(mdbf_modules, best_state) logger.info("Loaded best state (KL=%.6f)", best_kl) elif best_kl >= initial_kl: logger.info("No improvement — rolling back to initial state.") if method == "gptq": load_gptq_state(gptq_modules, initial_state) - else: + elif method == "dbf": load_dbf_state(dbf_modules, initial_state) + else: # mdbf + load_mdbf_state(mdbf_modules, initial_state) best_kl = initial_kl else: if method == "gptq": @@ -1119,11 +1194,16 @@ def _forward_and_loss() -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: elif method == "dbf": write_back_dbf_binary(dbf_modules) write_back_dbf_scaling(dbf_modules) + elif method == "mdbf": + write_back_mdbf_binary(mdbf_modules) + write_back_mdbf_amp(mdbf_modules) if method == "gptq": restore_gptq_original(gptq_modules, original_forwards, cleanup=False) elif method == "dbf": restore_dbf_original(dbf_modules, original_forwards, cleanup=False) + elif method == "mdbf": + restore_mdbf_original(mdbf_modules, original_forwards, cleanup=False) # Cleanup hooks if use_inter_loss: @@ -1140,7 +1220,9 @@ def _forward_and_loss() -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: # Final evaluation quantized_model.eval() - final_kl = eval_kl(quantized_model, teacher_model, dataloader, dev, temperature) + final_kl = eval_kl( + quantized_model, teacher_model, dataloader, dev, temperature, teacher_dev, + ) # Cleanup if method == "gptq": @@ -1148,6 +1230,8 @@ def _forward_and_loss() -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: restore_gptq_original(gptq_modules, original_forwards, cleanup=True) elif method == "dbf": restore_dbf_original(dbf_modules, original_forwards, cleanup=True) + elif method == "mdbf": + restore_mdbf_original(mdbf_modules, original_forwards, cleanup=True) del teacher_model gc.collect() diff --git a/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py b/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py index af32b36..b9b65b3 100644 --- a/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py +++ b/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py @@ -76,30 +76,34 @@ def detect_quantization_method( """Auto-detect the quantization method applied to *model*. Returns: - (method, modules) where *method* is ``"gptq"``, ``"dbf"``, or - ``None``, and *modules* is the list of ``(name, module)`` pairs - for the detected quantized layers. + (method, modules) where *method* is ``"gptq"``, ``"dbf"``, + ``"mdbf"``, or ``None``, and *modules* is the list of + ``(name, module)`` pairs for the detected quantized layers. - When both GPTQ and DBF layers are present (mixed quantization), - a warning is emitted and only GPTQ layers are returned. + Priority: GPTQ > DBF > MDBF. When multiple types coexist a warning is + emitted and only the highest-priority layers are returned. """ from onecomp.quantizer.gptq.gptq_layer import GPTQLinear from onecomp.quantizer.dbf.dbf_layer import DoubleBinaryLinear + from onecomp.quantizer.mdbf.mdbf_layer import MultipathMDBFLinear gptq_modules = find_target_modules(model, GPTQLinear) dbf_modules = find_target_modules(model, DoubleBinaryLinear) + mdbf_modules = find_target_modules(model, MultipathMDBFLinear) - if gptq_modules and dbf_modules: + if gptq_modules and (dbf_modules or mdbf_modules): logger.warning( - "Mixed GPTQ + DBF model detected (gptq=%d, dbf=%d). " + "Mixed GPTQ + DBF/MDBF model detected (gptq=%d, dbf=%d, mdbf=%d). " "Global PTQ currently optimises GPTQ layers only; " - "DBF layers will be skipped.", - len(gptq_modules), len(dbf_modules), + "other layers will be skipped.", + len(gptq_modules), len(dbf_modules), len(mdbf_modules), ) if gptq_modules: return "gptq", gptq_modules if dbf_modules: return "dbf", dbf_modules + if mdbf_modules: + return "mdbf", mdbf_modules return None, [] diff --git a/global_ptq/onecomp_globalptq/global_ptq/_core/trainer.py b/global_ptq/onecomp_globalptq/global_ptq/_core/trainer.py index 6a71253..33b5486 100644 --- a/global_ptq/onecomp_globalptq/global_ptq/_core/trainer.py +++ b/global_ptq/onecomp_globalptq/global_ptq/_core/trainer.py @@ -31,6 +31,12 @@ setup_dbf_forwards_only, write_back_dbf_binary, ) +from .mdbf_adapter import ( + restore_mdbf_original, + setup_mdbf_forwards_only, + write_back_mdbf_binary, + write_back_mdbf_amp, +) logger = getLogger(__name__) @@ -91,9 +97,11 @@ def __init__( self, *, teacher_model: nn.Module, + teacher_device=None, method: str, gptq_modules: list, dbf_modules: list, + mdbf_modules: list = None, original_forwards: dict, optimize_intweight: bool, optimize_binary: bool, @@ -105,9 +113,11 @@ def __init__( ): super().__init__(**kwargs) self.teacher_model = teacher_model + self.teacher_device = teacher_device self.method = method self.gptq_modules = gptq_modules self.dbf_modules = dbf_modules + self.mdbf_modules = mdbf_modules or [] self.original_forwards = original_forwards self.optimize_intweight = optimize_intweight self.optimize_binary = optimize_binary @@ -157,9 +167,19 @@ def compute_loss(self, model, inputs, return_outputs=False, **kwargs): loss = torch.tensor(0.0, device=logits_s.device) if self.w_distill > 0 and self.teacher_model is not None: with torch.no_grad(): - # Teacher also gets all available inputs - teacher_outputs = self.teacher_model(**inputs) - logits_t = get_logits(teacher_outputs) + if ( + self.teacher_device is not None + and self.teacher_device != logits_s.device + ): + teacher_inputs = { + k: v.to(self.teacher_device) + if isinstance(v, torch.Tensor) else v + for k, v in inputs.items() + } + teacher_outputs = self.teacher_model(**teacher_inputs) + else: + teacher_outputs = self.teacher_model(**inputs) + logits_t = get_logits(teacher_outputs).to(logits_s.device) loss = loss + self.w_distill * compute_kl_loss( logits_t, logits_s, self.temperature, @@ -191,6 +211,10 @@ def evaluate(self, eval_dataset=None, ignore_keys=None, elif self.method == "dbf": write_back_dbf_binary(self.dbf_modules) restore_dbf_original(self.dbf_modules, self.original_forwards) + elif self.method == "mdbf": + write_back_mdbf_binary(self.mdbf_modules) + write_back_mdbf_amp(self.mdbf_modules) + restore_mdbf_original(self.mdbf_modules, self.original_forwards) result = super().evaluate(eval_dataset, ignore_keys, metric_key_prefix) @@ -203,5 +227,9 @@ def evaluate(self, eval_dataset=None, ignore_keys=None, setup_dbf_forwards_only( self.dbf_modules, self.original_forwards, ) + elif self.method == "mdbf": + setup_mdbf_forwards_only( + self.mdbf_modules, self.original_forwards, + ) return result diff --git a/global_ptq/onecomp_globalptq/global_ptq/global_ptq.py b/global_ptq/onecomp_globalptq/global_ptq/global_ptq.py index 3f875b2..348debc 100644 --- a/global_ptq/onecomp_globalptq/global_ptq/global_ptq.py +++ b/global_ptq/onecomp_globalptq/global_ptq/global_ptq.py @@ -76,8 +76,10 @@ class GlobalPTQ(PostQuantizationProcess): ste_k (float): Smoothness parameter for GPTQ integer-weight Smooth STE rounding. Only used when ``gptq_optimize_intweight=True``. - DBF binary STE uses a fixed internal sharpness (k=2). Default is 100.0. + mdbf_ste_k (float): + Sharpness for MDBF binary sign STE (``tanh(k*x)`` backward). + Default is 2.0. calibration_dataset (list or None): List of text strings to use as calibration data. If ``None`` (default), the AllenAI C4 dataset is @@ -159,7 +161,6 @@ class GlobalPTQ(PostQuantizationProcess): optimiser update. Default is 1 (no accumulation). Incompatible with ``use_sam=True``; when both are set, this value is silently forced to 1. - Examples: >>> from onecomp import Runner, ModelConfig, GPTQ >>> from onecomp_globalptq import GlobalPTQ @@ -184,6 +185,7 @@ class GlobalPTQ(PostQuantizationProcess): dbf_lr: float = 5e-5 optimize_binary: bool = False ste_k: float = 100.0 + mdbf_ste_k: float = 2.0 calibration_dataset: Optional[List[str]] = None num_calibration_samples: int = 128 max_length: int = 2048 @@ -236,6 +238,10 @@ class GlobalPTQ(PostQuantizationProcess): # --- Gradient Accumulation --- grad_accum_steps: int = 1 + # --- Device placement (multi-GPU / CPU teacher) --- + student_device: Optional[str] = None + teacher_device: Optional[str] = None + def __post_init__(self): super().__post_init__() if self.epochs < 1: @@ -286,6 +292,7 @@ def run( gptq_intweight_lr=self.gptq_intweight_lr, optimize_binary=self.optimize_binary, ste_k=self.ste_k, + mdbf_ste_k=self.mdbf_ste_k, calibration_dataset=self.calibration_dataset, num_calibration_samples=self.num_calibration_samples, max_length=self.max_length, @@ -315,6 +322,8 @@ def run( early_stopping_patience=self.early_stopping_patience, use_mixed_precision=self.use_mixed_precision, grad_accum_steps=self.grad_accum_steps, + student_device=self.student_device, + teacher_device=self.teacher_device, ) except Exception: diff --git a/global_ptq/onecomp_globalptq/global_ptq/global_ptq_distributed.py b/global_ptq/onecomp_globalptq/global_ptq/global_ptq_distributed.py index 7ca9dc5..20d2bd2 100644 --- a/global_ptq/onecomp_globalptq/global_ptq/global_ptq_distributed.py +++ b/global_ptq/onecomp_globalptq/global_ptq/global_ptq_distributed.py @@ -68,8 +68,10 @@ class GlobalPTQDistributed(PostQuantizationProcess): ste_k (float): Smoothness parameter for GPTQ integer-weight Smooth STE rounding. Only used when ``gptq_optimize_intweight=True``. - DBF binary STE uses a fixed internal sharpness (k=2). Default is 100.0. + mdbf_ste_k (float): + Sharpness for MDBF binary sign STE (``tanh(k*x)`` backward). + Default is 2.0. dbf_lr (float): Learning rate for DBF scaling parameters. Default is 5e-5. @@ -169,9 +171,10 @@ class GlobalPTQDistributed(PostQuantizationProcess): gptq_intweight_lr: float = 1e-4 ste_k: float = 100.0 - # --- DBF --- + # --- DBF / MDBF --- dbf_lr: float = 5e-5 optimize_binary: bool = False + mdbf_ste_k: float = 2.0 # --- Calibration --- calibration_dataset: Optional[List[str]] = None @@ -192,6 +195,7 @@ class GlobalPTQDistributed(PostQuantizationProcess): # --- Distributed --- deepspeed_config: Optional[str] = None + teacher_device: Optional[str] = None # --- Output / Logging / Checkpointing --- output_dir: Optional[str] = None @@ -274,6 +278,14 @@ def run( write_back_dbf_scaling, restore_dbf_original, ) + from ._core.mdbf_adapter import ( + load_mdbf_state, + save_mdbf_state, + setup_mdbf_differentiable, + write_back_mdbf_binary, + write_back_mdbf_amp, + restore_mdbf_original, + ) from onecomp import CalibrationConfig from onecomp.calibration import prepare_calibration_dataset from transformers import TrainingArguments, default_data_collator @@ -292,7 +304,7 @@ def run( if method is None: logger.warning("No quantized layers detected — skipping.") return - if method not in ("gptq", "dbf"): + if method not in ("gptq", "dbf", "mdbf"): logger.info("Method '%s' not supported — skipping.", method) return @@ -333,6 +345,7 @@ def run( gptq_modules = [] dbf_modules = [] + mdbf_modules = [] original_forwards = {} param_groups = [] @@ -375,6 +388,31 @@ def run( f", {len(binary_params)} binary" if binary_params else "", ) + elif method == "mdbf": + mdbf_modules = detected_modules + original_forwards, amp_params, binary_params = ( + setup_mdbf_differentiable( + mdbf_modules, + self.optimize_binary, + ste_k=self.mdbf_ste_k, + ) + ) + logger.info("MDBF binary STE sharpness mdbf_ste_k=%.4g", self.mdbf_ste_k) + all_mdbf_params = list(amp_params) + if binary_params: + all_mdbf_params += binary_params + param_groups = [{ + "params": all_mdbf_params, + "lr": self.dbf_lr, + "weight_decay": 0.0, + }] + logger.info( + "Trainable: %d amp params%s across %d MDBF modules", + len(amp_params), + f", {len(binary_params)} binary" if binary_params else "", + len(mdbf_modules), + ) + # DeepSpeed ZeRO requires contiguous tensors for all-reduce. for pg in param_groups: for p in pg["params"]: @@ -388,6 +426,8 @@ def run( restore_gptq_original(gptq_modules, original_forwards, cleanup=True) elif method == "dbf": restore_dbf_original(dbf_modules, original_forwards) + elif method == "mdbf": + restore_mdbf_original(mdbf_modules, original_forwards, cleanup=True) quantized_model.cpu() return @@ -403,13 +443,26 @@ def run( # 5. Teacher model # ------------------------------------------------------------------ need_teacher = self.w_distill > 0 + teacher_dev = dev if need_teacher: - logger.info("Loading FP16 teacher model...") + world_size = int(os.environ.get("WORLD_SIZE", "1")) + resolved_teacher = self.teacher_device + if resolved_teacher is None and ( + self.deepspeed_config or world_size > 1 + ): + resolved_teacher = "cpu" + if resolved_teacher is not None: + teacher_dev = torch.device(resolved_teacher) + logger.info( + "Loading FP16 teacher model (teacher_device=%s)...", + teacher_dev, + ) teacher_model = model_config.load_model(device_map="cpu") teacher_model.eval() for p in teacher_model.parameters(): p.requires_grad = False - teacher_model.to(dev) + if teacher_dev.type != "cpu": + teacher_model.to(teacher_dev) else: logger.info( "w_distill=0 — skipping teacher model load (pure QAT mode)." @@ -417,7 +470,7 @@ def run( # ------------------------------------------------------------------ # 6. TrainingArguments # ------------------------------------------------------------------ - lr = self.gptq_lr if method == "gptq" else self.dbf_lr + lr = self.gptq_lr if method == "gptq" else self.dbf_lr # dbf_lr used for both dbf and mdbf resolved_output_dir = ( self.output_dir if self.output_dir is not None @@ -459,8 +512,10 @@ def run( # Save initial state for rollback if training degrades quality if method == "gptq": _initial_state = save_gptq_state(gptq_modules) - else: + elif method == "dbf": _initial_state = save_dbf_state(dbf_modules) + else: # mdbf + _initial_state = save_mdbf_state(mdbf_modules) # ------------------------------------------------------------------ # 7. Train @@ -468,9 +523,11 @@ def run( trainer = _GlobalPTQTrainer( model=quantized_model, teacher_model=teacher_model, + teacher_device=teacher_dev if need_teacher else None, method=method, gptq_modules=gptq_modules, dbf_modules=dbf_modules, + mdbf_modules=mdbf_modules, original_forwards=original_forwards, optimize_intweight=self.gptq_optimize_intweight, optimize_binary=self.optimize_binary, @@ -523,8 +580,10 @@ def run( ) if method == "gptq": load_gptq_state(gptq_modules, _initial_state) - else: + elif method == "dbf": load_dbf_state(dbf_modules, _initial_state) + else: # mdbf + load_mdbf_state(mdbf_modules, _initial_state) else: logger.info( "Best eval_loss: %.6f (initial=%.6f) at step %d.", @@ -539,6 +598,10 @@ def run( write_back_dbf_binary(dbf_modules) write_back_dbf_scaling(dbf_modules) restore_dbf_original(dbf_modules, original_forwards) + elif method == "mdbf": + write_back_mdbf_binary(mdbf_modules) + write_back_mdbf_amp(mdbf_modules) + restore_mdbf_original(mdbf_modules, original_forwards, cleanup=True) finally: if original_use_cache is not None: quantized_model.config.use_cache = original_use_cache From 85896e244f44a6802942a8d0196b3741f30fb3e5 Mon Sep 17 00:00:00 2001 From: fujisawa-yoshihiko Date: Wed, 29 Jul 2026 22:42:35 +0900 Subject: [PATCH 3/3] fix(global_ptq): make MDBF detection tolerate a missing MDBF quantizer 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. --- .../onecomp_globalptq/global_ptq/_core/helpers.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py b/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py index b9b65b3..5a5b2bc 100644 --- a/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py +++ b/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py @@ -85,11 +85,20 @@ def detect_quantization_method( """ from onecomp.quantizer.gptq.gptq_layer import GPTQLinear from onecomp.quantizer.dbf.dbf_layer import DoubleBinaryLinear - from onecomp.quantizer.mdbf.mdbf_layer import MultipathMDBFLinear + + try: + from onecomp.quantizer.mdbf.mdbf_layer import MultipathMDBFLinear + except ImportError: + # MDBF quantizer is optional; without it only GPTQ/DBF are detectable. + MultipathMDBFLinear = None gptq_modules = find_target_modules(model, GPTQLinear) dbf_modules = find_target_modules(model, DoubleBinaryLinear) - mdbf_modules = find_target_modules(model, MultipathMDBFLinear) + mdbf_modules = ( + find_target_modules(model, MultipathMDBFLinear) + if MultipathMDBFLinear is not None + else [] + ) if gptq_modules and (dbf_modules or mdbf_modules): logger.warning(