Skip to content
118 changes: 104 additions & 14 deletions global_ptq/onecomp_globalptq/global_ptq/_core/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -418,17 +427,33 @@ 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,
teacher_model: nn.Module,
dataloader: List[Dict[str, torch.Tensor]],
dev: torch.device,
temperature: float = 1.0,
teacher_dev: Optional[torch.device] = None,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

devを別途受け渡しているのですが、それを使用する形では難しいでしょうか。

) -> 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)
Expand All @@ -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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

好みの世界で恐縮ですが、 _get_teacher_logits と動詞を含めた方が可読性が高まるかと思いました。

total += compute_kl_loss(
logits_t, logits_s, temperature, attention_mask=attention_mask,
).item()
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -616,12 +642,23 @@ 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,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Noneの場合の条件分岐が何か所かで入っていることを考慮するとOptionalではなく必須の変数にするのも良いかと思いましたがいかがでしょうか。

) -> Dict:
"""Run KL-distillation global PTQ on a GPTQ or DBF quantized model.
"""Run KL-distillation global PTQ on a GPTQ, DBF or MDBF quantized model.

The quantization method is auto-detected from the layer types present in
*quantized_model* (see :func:`detect_quantization_method`). For MDBF the
per-path amplitude factors are trained with ``dbf_lr``; with
``optimize_binary=True`` the +/-1 sign matrices are trained as well through
a smooth sign STE whose sharpness is controlled by ``mdbf_ste_k``.

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
Expand All @@ -631,7 +668,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}"}

Expand Down Expand Up @@ -665,7 +702,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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

内部的にdeviceを変更する場合はloggerでメッセージを上げていただきたいです。


# ------------------------------------------------------------------
# 4. Move student to GPU and set up differentiable parameters
Expand All @@ -675,6 +713,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 = []
Expand Down Expand Up @@ -710,13 +749,33 @@ 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.")
if method == "gptq":
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()
Expand Down Expand Up @@ -871,17 +930,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
Expand Down Expand Up @@ -928,8 +994,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,
Expand Down Expand Up @@ -1052,23 +1119,33 @@ 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

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)

# Restore non-EMA params for continued training
if ema_tracker is not None:
Expand Down Expand Up @@ -1103,27 +1180,36 @@ 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":
write_back_gptq_params(gptq_modules, gptq_optimize_intweight)
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:
Expand All @@ -1140,14 +1226,18 @@ 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":
# Final cleanup of differentiable parameters
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()
Expand Down
31 changes: 22 additions & 9 deletions global_ptq/onecomp_globalptq/global_ptq/_core/helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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, []


Expand Down
Loading