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..5a5b2bc 100644 --- a/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py +++ b/global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py @@ -76,30 +76,43 @@ 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 + 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) + if MultipathMDBFLinear is not None + else [] + ) - 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/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()) 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