| """training_unified — QAT wrapper + dual-path distillation for the unified |
| QuantizedModule architecture. |
| |
| STE (Straight-Through Estimator) is reused from training/ste.py. |
| |
| Dual-path distillation: MSE(student_out, teacher_out.detach()) — local block |
| distillation, gradients flow only through the student path. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import torch |
| import torch.nn as nn |
|
|
|
|
| class UnifiedQATWrapper(nn.Module): |
| """Wrap a unified QuantizedModule for QAT (learnable latent weights + STE). |
| |
| The QuantizedModule must have been created with a Quantizer(learnable=True). |
| This wrapper provides the training-time interface: |
| - forward: fake-quant (quantize then dequantize) via STE |
| - distillation_loss: dual-path MSE(student, teacher) |
| After training, call strip_latent() (module-level helper) to freeze. |
| """ |
|
|
| def __init__(self, quantized_module: nn.Module): |
| super().__init__() |
| self.qm = quantized_module |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return self.qm(x) |
|
|
| def distillation_loss( |
| self, |
| x: torch.Tensor, |
| reduction: str = "mean", |
| ) -> torch.Tensor: |
| """Dual-path distillation loss: MSE(student, teacher). |
| |
| Requires the QuantizedModule to have a teacher (dual_path=True). |
| """ |
| student_out, teacher_out = self.qm(x, path="both") |
| if reduction == "mean": |
| return torch.nn.functional.mse_loss(student_out, teacher_out.detach()) |
| elif reduction == "sum": |
| return torch.nn.functional.mse_loss(student_out, teacher_out.detach(), reduction="sum") |
| return torch.nn.functional.mse_loss(student_out, teacher_out.detach(), reduction="none") |
|
|
|
|
| def dual_path_loss( |
| student_out: torch.Tensor, |
| teacher_out: torch.Tensor, |
| loss_type: str = "mse", |
| ) -> torch.Tensor: |
| """Compute distillation loss between student and teacher outputs. |
| |
| Args: |
| student_out: quantized (student) forward output. |
| teacher_out: frozen (teacher) forward output. |
| loss_type: "mse" (default), "cosine", "kl". |
| """ |
| teacher_detached = teacher_out.detach() |
| if loss_type == "mse": |
| return torch.nn.functional.mse_loss(student_out, teacher_detached) |
| elif loss_type == "cosine": |
| s = student_out.flatten() |
| t = teacher_detached.flatten() |
| cos = torch.nn.functional.cosine_similarity(s.unsqueeze(0), t.unsqueeze(0)) |
| return 1.0 - cos |
| elif loss_type == "kl": |
| |
| return torch.nn.functional.kl_div( |
| torch.nn.functional.log_softmax(student_out, dim=-1), |
| torch.nn.functional.softmax(teacher_detached, dim=-1), |
| reduction="batchmean", |
| ) |
| raise ValueError(f"Unknown loss_type: {loss_type}") |
|
|
|
|
| def strip_latent(quantized_module: nn.Module) -> nn.Module: |
| """Convert a learnable (QAT) QuantizedModule to inference-only. |
| |
| After training, the latent weight + learnable scale are "baked" into |
| frozen QuantizedWeight buffers (no gradient, no STE). The QuantizedModule |
| re-quantizes latent_weight with the final latent_scale and replaces its |
| buffers, then drops the latent parameters. |
| """ |
| from agiws_neural_quant.base import QuantizedModule |
| if not isinstance(quantized_module, QuantizedModule): |
| return quantized_module |
| if not getattr(quantized_module, "_learnable", False): |
| return quantized_module |
| |
| quantizer = quantized_module._quantizer |
| W_final = quantized_module.latent_weight.detach() |
| qw_final = quantizer.quantize_weight(W_final) |
| |
| for name in list(quantized_module._buffers.keys()): |
| if name in qw_final.weight_buffers: |
| quantized_module._buffers[name] = qw_final.weight_buffers[name] |
| |
| if hasattr(quantized_module, "latent_codebook"): |
| trained_cb = quantized_module.latent_codebook.detach() |
| if "codebook" in quantized_module._buffers: |
| quantized_module._buffers["codebook"] = trained_cb |
| |
| meta = quantized_module._collect_weight_meta() |
| if meta.get("codebook_source") != "vq": |
| W_norm = W_final / quantized_module.latent_scale.detach().unsqueeze(1) \ |
| if quantized_module.latent_scale.dim() == 1 and quantized_module.latent_scale.numel() > 1 \ |
| else W_final / quantized_module.latent_scale.detach() |
| diff = W_norm.unsqueeze(2) - trained_cb.unsqueeze(0).unsqueeze(0) |
| quantized_module._buffers["indices"] = diff.abs().argmin(dim=2).to(torch.int32) |
| |
| del quantized_module.latent_weight |
| del quantized_module.latent_scale |
| |
| if hasattr(quantized_module, "latent_codebook"): |
| del quantized_module.latent_codebook |
| quantized_module._learnable = False |
| for p in quantized_module.parameters(): |
| p.requires_grad = False |
| return quantized_module |
|
|
|
|
| __all__ = [ |
| "UnifiedQATWrapper", |
| "dual_path_loss", |
| "strip_latent", |
| ] |