ArGrigorov's picture
Upload folder using huggingface_hub
056d296 verified
Raw
History Blame Contribute Delete
28.6 kB
"""base — unified quantization architecture (refactored).
4 entities (symmetric, no hacks):
1. QuantizedWeight — dataclass, persistent (state_dict). Packed buffers + meta.
2. QuantizedActivation — dataclass, ephemeral (per-forward) / long-lifetime (KV-cache).
3. Quantizer — ONE parameterized class. All 25+ formats as configurations.
4. QuantizedModule — nn.Module wrapper. Chunked dequant + dual-path + QAT learnable.
Minimal VRAM: weights packed (not fp32 master), chunked dequant in forward,
learnable parameters (latent weights, scale, boundaries, codebook) via STE.
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from typing import Any, Protocol, runtime_checkable
import torch
import torch.nn as nn
import torch.nn.functional as F
logger = logging.getLogger(__name__)
_COMPUTE_DTYPES = {"fp32": torch.float32, "fp16": torch.float16, "bf16": torch.bfloat16}
def _compute_dtype_to_torch(compute_dtype: str) -> torch.dtype:
return _COMPUTE_DTYPES.get(compute_dtype, torch.float32)
# ---------------------------------------------------------------------------
# QuantizedWeight — persistent container (saved in state_dict)
# ---------------------------------------------------------------------------
@dataclass
class QuantizedWeight:
"""Container for quantized weight data.
weight_buffers — dict of packed tensors (int4 codes, fp4 nibbles, codebook
indices, scales, rotation matrix, outlier indices, etc.).
Stored in packed form, NOT dequantized fp32.
weight_meta — dict of scalars (value_bits, scale_mode, group_size, ...).
"""
weight_buffers: dict[str, torch.Tensor] = field(default_factory=dict)
weight_meta: dict[str, Any] = field(default_factory=dict)
# ---------------------------------------------------------------------------
# QuantizedActivation — ephemeral container (per-forward or KV-cache)
# ---------------------------------------------------------------------------
@dataclass
class QuantizedActivation:
"""Container for quantized activation data.
activation_buffers — dict of packed activation tensors (int4/int8 codes
for W4A4/W8A8, or just scale for dynamic quant).
Empty for weight-only formats (passthrough).
activation_meta — dict of scalars (scale_mode, group_size, smoothing_s,
input_scale, ...).
"""
activation_buffers: dict[str, torch.Tensor] = field(default_factory=dict)
activation_meta: dict[str, Any] = field(default_factory=dict)
# ---------------------------------------------------------------------------
# Op registry — maps source module type to forward op function
# ---------------------------------------------------------------------------
def _linear_op(x, W, bias, op_kwargs):
return F.linear(x, W, bias)
def _conv1d_op(x, W, bias, op_kwargs):
return F.conv1d(x, W, bias, **op_kwargs)
def _conv2d_op(x, W, bias, op_kwargs):
return F.conv2d(x, W, bias, **op_kwargs)
def _conv3d_op(x, W, bias, op_kwargs):
return F.conv3d(x, W, bias, **op_kwargs)
def _conv_transpose1d_op(x, W, bias, op_kwargs):
return F.conv_transpose1d(x, W, bias, **op_kwargs)
def _conv_transpose2d_op(x, W, bias, op_kwargs):
return F.conv_transpose2d(x, W, bias, **op_kwargs)
def _conv_transpose3d_op(x, W, bias, op_kwargs):
return F.conv_transpose3d(x, W, bias, **op_kwargs)
def _embedding_op(x, W, bias, op_kwargs):
return F.embedding(x, W, **op_kwargs)
def _layernorm_op(x, W, bias, op_kwargs):
return F.layer_norm(x, op_kwargs["normalized_shape"], W, bias, op_kwargs["eps"])
_CONV_TYPES = (nn.Conv1d, nn.Conv2d, nn.Conv3d)
_CONV_TRANSPOSE_TYPES = (
nn.ConvTranspose1d, nn.ConvTranspose2d, nn.ConvTranspose3d,
)
def _extract_op_kwargs(module: nn.Module) -> dict[str, Any]:
"""Extract forward kwargs from the source module."""
if isinstance(module, (nn.Linear, nn.Bilinear)):
return {}
if isinstance(module, _CONV_TYPES):
return {
"stride": module.stride,
"padding": module.padding,
"dilation": module.dilation,
"groups": module.groups,
}
if isinstance(module, _CONV_TRANSPOSE_TYPES):
return {
"stride": module.stride,
"padding": module.padding,
"dilation": module.dilation,
"groups": module.groups,
"output_padding": module.output_padding,
}
if isinstance(module, nn.Embedding):
return {
"padding_idx": module.padding_idx,
"scale_grad_by_freq": module.scale_grad_by_freq,
"sparse": module.sparse,
}
if isinstance(module, nn.LayerNorm):
return {
"normalized_shape": tuple(module.normalized_shape),
"eps": module.eps,
}
return {}
def _select_op(module: nn.Module):
"""Pick the op function for a source module type."""
if isinstance(module, (nn.Linear, nn.Bilinear)):
return _linear_op
if isinstance(module, nn.Conv1d):
return _conv1d_op
if isinstance(module, nn.Conv2d):
return _conv2d_op
if isinstance(module, nn.Conv3d):
return _conv3d_op
if isinstance(module, nn.ConvTranspose1d):
return _conv_transpose1d_op
if isinstance(module, nn.ConvTranspose2d):
return _conv_transpose2d_op
if isinstance(module, nn.ConvTranspose3d):
return _conv_transpose3d_op
if isinstance(module, nn.Embedding):
return _embedding_op
if isinstance(module, nn.LayerNorm):
return _layernorm_op
raise TypeError(f"Unsupported module type for quantization: {type(module).__name__}")
def _extract_bias(module: nn.Module) -> torch.Tensor | None:
"""Extract the bias tensor (or None) from a source module."""
b = getattr(module, "bias", None)
if b is None:
return None
return b.detach().float()
SUPPORTED_MODULE_TYPES = (
nn.Linear,
nn.Bilinear,
nn.Conv1d,
nn.Conv2d,
nn.Conv3d,
nn.ConvTranspose1d,
nn.ConvTranspose2d,
nn.ConvTranspose3d,
nn.Embedding,
nn.LayerNorm,
)
_OP_REGISTRY: dict[str, Any] = {
"Linear": _linear_op,
"Bilinear": _linear_op,
"Conv1d": _conv1d_op,
"Conv2d": _conv2d_op,
"Conv3d": _conv3d_op,
"ConvTranspose1d": _conv_transpose1d_op,
"ConvTranspose2d": _conv_transpose2d_op,
"ConvTranspose3d": _conv_transpose3d_op,
"Embedding": _embedding_op,
"LayerNorm": _layernorm_op,
}
# ---------------------------------------------------------------------------
# AdaptiveChunkSize — VRAM-aware runtime tuning (not gradient)
# ---------------------------------------------------------------------------
class AdaptiveChunkSize:
"""Runtime-adaptive chunk size based on VRAM headroom.
After each forward, measures torch.cuda.memory_allocated(). If headroom
is low → decrease chunk_size (avoid OOM). If high → increase (faster).
Not gradient-learnable — engineering optimization.
"""
def __init__(
self,
initial: int = 1024,
min_size: int = 64,
max_size: int = 8192,
vram_headroom_target: float = 0.3,
adjustment_factor: float = 0.5,
):
self.current = initial
self.min_size = min_size
self.max_size = max_size
self.target = vram_headroom_target
self.factor = adjustment_factor
self._adjustments = 0
def get(self) -> int:
return self.current
def maybe_adjust(self):
"""Check VRAM and adjust chunk_size. Call after forward."""
if not torch.cuda.is_available():
return
try:
allocated = torch.cuda.memory_allocated()
reserved = torch.cuda.max_memory_reserved()
total = torch.cuda.get_device_properties(0).total_memory
headroom_ratio = max(0.0, (total - allocated) / total)
if headroom_ratio < 0.10:
new_size = max(self.min_size, int(self.current * self.factor))
if new_size != self.current:
logger.debug(f"AdaptiveChunkSize: OOM risk, {self.current}{new_size}")
self.current = new_size
self._adjustments += 1
elif headroom_ratio > 0.50 and self.current < self.max_size:
new_size = min(self.max_size, int(self.current / self.factor))
if new_size != self.current:
logger.debug(f"AdaptiveChunkSize: headroom {headroom_ratio:.0%}, {self.current}{new_size}")
self.current = new_size
self._adjustments += 1
except Exception:
pass # VRAM check is best-effort.
# ---------------------------------------------------------------------------
# QuantizedModule — generic nn.Module wrapper with chunked dequant + dual-path
# ---------------------------------------------------------------------------
class QuantizedModule(nn.Module):
"""Generic wrapper around any nn.Module with quantized weights.
Features:
- Chunked dequant: forward processes output dim in chunks of chunk_size
to minimize peak VRAM (packed weights → dequant slice → matmul → free).
- Dual-path: optional teacher QuantizedModule for cross-quantization
distillation. forward(x, path="student"|"teacher"|"both").
- QAT learnable: if quantizer.learnable, latent weight is an nn.Parameter
(trainable via STE); scale is a frozen buffer exposed as nn.Parameter
for API uniformity (STE.backward returns grad=None for scale).
- Adaptive chunk: AdaptiveChunkSize monitor adjusts chunk_size at runtime.
Args:
qw: QuantizedWeight (packed buffers + meta).
quantizer: Quantizer instance (format logic).
op_type: source module type name (for op registry).
op_kwargs: forward kwargs (stride, padding, normalized_shape, ...).
bias: dequantized bias (fp32) or None.
weight_shape: original weight shape tuple.
compute_dtype: "fp32"|"fp16"|"bf16" for dequantized matmul.
chunk_size: output-dim chunk for dequant (default 1024). None = no chunking.
adaptive: if True, AdaptiveChunkSize monitor adjusts chunk_size.
dual_path: if True, teacher_path is active.
teacher: optional QuantizedModule for dual-path (frozen, different format).
"""
def __init__(
self,
qw: QuantizedWeight,
quantizer: Any,
op_type: str,
op_kwargs: dict[str, Any],
bias: torch.Tensor | None,
weight_shape: tuple[int, ...],
compute_dtype: str = "fp32",
chunk_size: int | None = 1024,
adaptive: bool = False,
dual_path: bool = False,
teacher: "QuantizedModule | None" = None,
):
super().__init__()
self.op_type = op_type
self.op_kwargs = op_kwargs
self.weight_shape = weight_shape
self.compute_dtype = compute_dtype
self._quantizer = quantizer
self._quantizer_info = quantizer.info()
self._dual_path = dual_path
self._teacher = teacher
self._ternary_active = False # dual-path switch flag
# Chunked dequant config.
self._chunk_size = chunk_size
self._adaptive = AdaptiveChunkSize() if adaptive else None
# Register weight buffers from container (packed, not fp32).
for name, buf in qw.weight_buffers.items():
if buf is None:
continue
self.register_buffer(name, buf)
# Store meta as plain attributes.
for k, v in qw.weight_meta.items():
setattr(self, "_wmeta_" + k, v)
# Learnable (QAT) path: if the quantizer is learnable, create trainable
# latent parameters from the frozen buffers. Forward will use
# fake_quantize (STE) instead of the frozen dequant.
self._learnable = bool(getattr(quantizer, "learnable", False))
if self._learnable:
# latent_weight: full-precision master weight (nn.Parameter).
W_fp = quantizer.dequantize_weight(qw, "fp32").clone()
self.latent_weight = nn.Parameter(W_fp)
# latent_scale: the scale buffer as a trainable parameter.
scale_buf = qw.weight_buffers.get("scale")
if scale_buf is not None:
self.latent_scale = nn.Parameter(scale_buf.to(torch.float32).clone())
else:
# Formats without a scale buffer (none/prune) — single scalar.
self.latent_scale = nn.Parameter(torch.tensor(1.0, dtype=torch.float32))
# Codebook learnable: store codebook as nn.Parameter for STE training.
if qw.weight_meta.get("value_repr") == "codebook":
cb = qw.weight_buffers.get("codebook")
if cb is not None:
self.latent_codebook = nn.Parameter(cb.to(torch.float32).clone())
# Bias.
if bias is not None:
self.register_buffer("bias", bias.to(torch.float32))
else:
self.register_buffer("bias", None)
# -- reconstruction of QuantizedWeight from registered buffers ----------
def _collect_weight_buffers(self) -> dict[str, torch.Tensor]:
"""Collect registered weight buffers (exclude bias, input_*)."""
out = {}
for name, buf in self._buffers.items():
if buf is None or name == "bias" or name.startswith("input_"):
continue
out[name] = buf
return out
def _collect_weight_meta(self) -> dict[str, Any]:
out = {}
for k, v in self.__dict__.items():
if k.startswith("_wmeta_"):
out[k[len("_wmeta_"):]] = v
return out
def _qw(self) -> QuantizedWeight:
return QuantizedWeight(
weight_buffers=self._collect_weight_buffers(),
weight_meta=self._collect_weight_meta(),
)
# -- public API ---------------------------------------------------------
@classmethod
def from_module(
cls,
module: nn.Module,
quantizer: Any,
compute_dtype: str = "fp32",
chunk_size: int | None = 1024,
adaptive: bool = False,
teacher: "QuantizedModule | None" = None,
) -> "QuantizedModule":
"""Build QuantizedModule from an arbitrary nn.Module with a weight."""
if not isinstance(module, SUPPORTED_MODULE_TYPES):
raise TypeError(
f"QuantizedModule.from_module: unsupported type {type(module).__name__}"
)
W = module.weight.detach().float()
bias = _extract_bias(module)
op_kwargs = _extract_op_kwargs(module)
op_type = type(module).__name__
qw = quantizer.quantize_weight(W)
dual_path = teacher is not None
return cls(
qw=qw,
quantizer=quantizer,
op_type=op_type,
op_kwargs=op_kwargs,
bias=bias,
weight_shape=tuple(W.shape),
compute_dtype=compute_dtype,
chunk_size=chunk_size,
adaptive=adaptive,
dual_path=dual_path,
teacher=teacher,
)
def dequantize_weight(self) -> torch.Tensor:
"""Reconstruct the full float weight in compute_dtype."""
return self._quantizer.dequantize_weight(self._qw(), self.compute_dtype)
def dequantize_weight_slice(self, start: int, end: int) -> torch.Tensor:
"""Reconstruct a slice [start:end] of the output dim (for chunked)."""
return self._quantizer.dequantize_weight(self._qw(), self.compute_dtype, slice=(start, end))
@property
def weight(self) -> torch.Tensor:
"""Read-only dequantized weight — for compatibility."""
return self.dequantize_weight()
def _get_chunk_size(self) -> int:
if self._adaptive is not None:
return self._adaptive.get()
return self._chunk_size if self._chunk_size is not None else self.weight_shape[0]
# -- learnable weight (QAT path) -----------------------------------------
def _learnable_weight(self, slice: tuple[int, int] | None = None) -> torch.Tensor:
"""Fake-quantized weight from latent parameters (STE backward).
Used when quantizer.learnable=True. Returns a differentiable tensor
connected to latent_weight / latent_scale via the STE. The fake-quant
uses the SAME scale granularity as the frozen quantizer (per-channel
or per-group), so strip_latent() (re-quantize via the frozen quantizer)
produces matching outputs.
For codebook formats: uses STECodebook (argmin + STE), codebook is
learnable via latent_codebook (nn.Parameter).
"""
from agiws_neural_quant.training.ste import fake_quantize, fake_codebook_quantize
W = self.latent_weight
meta = self._collect_weight_meta()
repr_ = meta.get("value_repr", "int")
# Codebook learnable path.
if repr_ == "codebook" and hasattr(self, "latent_codebook"):
cb = self.latent_codebook
scale = self.latent_scale
if slice is not None:
start, end = slice
W = W[start:end]
if scale.dim() == 1 and scale.shape[0] == self.weight_shape[0]:
scale = scale[start:end]
# Recompute argmin indices (non-differentiable, STE bypasses).
diff = W.unsqueeze(-1) - cb.unsqueeze(0).unsqueeze(0)
indices = diff.abs().argmin(dim=-1)
w_norm = fake_codebook_quantize(W, cb, indices)
if scale.dim() == 1 and scale.numel() > 1:
return w_norm * scale.unsqueeze(1)
return w_norm * scale
# Int / FP path.
scale = self.latent_scale
scale_mode = meta.get("scale_mode", "per-channel")
gs = meta.get("group_size", 0) or 0
n_levels = self._quantizer.n_levels
symmetric = self._quantizer.symmetric
if slice is not None:
start, end = slice
W = W[start:end]
if scale.dim() >= 1 and scale.shape[0] == self.weight_shape[0]:
scale = scale[start:end]
if scale_mode in ("per-group", "per-block") and gs > 0 and scale.dim() == 2 and W.dim() > 1:
# Per-group: pad latent to multiple of group_size, fake-quant per group.
out_f, in_f = W.shape
pad = (gs - (in_f % gs)) % gs
if pad > 0:
W = torch.nn.functional.pad(W, (0, pad))
num_groups = W.shape[1] // gs
W_grouped = W.reshape(out_f, num_groups, gs)
scale_exp = scale.unsqueeze(2).expand_as(W_grouped)
W_fq = fake_quantize(W_grouped.float(), scale_exp.float(), n_levels, symmetric)
W_fq = W_fq.reshape(out_f, -1)[:, :in_f]
return W_fq
# Per-channel / per-tensor.
if scale.dim() == 1 and W.dim() > 1:
scale = scale.unsqueeze(1)
return fake_quantize(W.float(), scale.float(), n_levels, symmetric)
def _student_forward(self, x: torch.Tensor) -> torch.Tensor:
"""Quantized (student) forward with chunked dequant."""
t = _compute_dtype_to_torch(self.compute_dtype)
qw = self._qw()
qa = self._quantizer.quantize_input(x, qw)
x_deq = self._quantizer.dequantize_input(qa, self.compute_dtype)
if x_deq is None:
x_deq = x # passthrough (weight-only formats)
op = _OP_REGISTRY.get(self.op_type)
if op is None:
raise RuntimeError(f"Unknown op_type={self.op_type!r}")
# Learnable (QAT): use fake-quant latent weight with STE gradient.
if self._learnable:
W = self._learnable_weight()
if self.op_type == "LayerNorm":
b = self.bias.to(t) if self.bias is not None else None
return op(x_deq.to(t), W.to(t), b, self.op_kwargs)
out_features = self.weight_shape[0]
chunk = self._get_chunk_size()
if chunk is None or chunk >= out_features:
b = self.bias.to(t) if self.bias is not None else None
return op(x_deq.to(t), W.to(t), b, self.op_kwargs)
results = []
for start in range(0, out_features, chunk):
end = min(start + chunk, out_features)
W_slice = self._learnable_weight(slice=(start, end))
b_slice = self.bias[start:end].to(t) if self.bias is not None else None
r = op(x_deq.to(t), W_slice.to(t), b_slice, self.op_kwargs)
results.append(r)
del W_slice
dim = -1 if self.op_type in ("Linear", "Bilinear") else 1
return torch.cat(results, dim=dim)
# Frozen path: chunked dequant from buffers.
# Embedding: input is indices (Long), no chunking needed.
if self.op_type == "Embedding":
W = self.dequantize_weight()
return op(x_deq, W, self.bias, self.op_kwargs)
# LayerNorm: weight is 1D, no output-dim chunking.
if self.op_type == "LayerNorm":
W = self.dequantize_weight()
b = self.bias.to(t) if self.bias is not None else None
return op(x_deq.to(t), W, b, self.op_kwargs)
# Linear / Conv: chunked dequant over output dim.
out_features = self.weight_shape[0]
chunk = self._get_chunk_size()
if chunk is None or chunk >= out_features:
# No chunking — dequant all at once.
W = self.dequantize_weight()
b = self.bias.to(t) if self.bias is not None else None
result = op(x_deq.to(t), W, b, self.op_kwargs)
else:
# Chunked: dequant slice → op → collect.
results = []
for start in range(0, out_features, chunk):
end = min(start + chunk, out_features)
W_slice = self.dequantize_weight_slice(start, end)
b_slice = None
if self.bias is not None:
b_slice = self.bias[start:end].to(t)
r = op(x_deq.to(t), W_slice, b_slice, self.op_kwargs)
results.append(r)
del W_slice
if self.op_type in ("Linear", "Bilinear"):
# Concatenate along output dim (last for Linear).
result = torch.cat(results, dim=-1)
else:
# Conv: concat along channel dim (dim 1).
result = torch.cat(results, dim=1)
if self._adaptive is not None:
self._adaptive.maybe_adjust()
return result
def _teacher_forward(self, x: torch.Tensor) -> torch.Tensor:
"""Teacher (frozen) forward — different quantization or fp16."""
if self._teacher is None:
raise RuntimeError("teacher_forward called but no teacher set")
with torch.no_grad():
return self._teacher.forward(x, path="student")
def forward(self, x: torch.Tensor, path: str = "student") -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""Forward pass.
Args:
path: "student" → quantized forward (default).
"teacher" → frozen teacher forward (requires dual_path).
"both" → (student_out, teacher_out) for distillation loss.
"""
if path == "teacher":
return self._teacher_forward(x)
if path == "both":
student_out = self._student_forward(x)
teacher_out = self._teacher_forward(x)
return student_out, teacher_out
# Default: student.
return self._student_forward(x)
def memory_bytes(self) -> int:
"""Total stored bytes: quantized weights + scales + bias."""
qw = self._qw()
w = self._quantizer.storage_bytes(qw)
b = self.bias.numel() * 4 if self.bias is not None else 0
return w + b
# -- serialization -------------------------------------------------------
def to_dict(self) -> dict[str, Any]:
"""Serialize this QuantizedModule to a plain dict (JSON-safe + tensors).
Tensors are moved to CPU. Use torch.save/load or json+torch serialization
for persistence. Reconstruct via QuantizedModule.from_dict(d).
Includes: buffers, meta, quantizer config, op_type, op_kwargs, bias,
weight_shape, compute_dtype, chunk_size, adaptive, dual_path, teacher,
and QAT learnable state (latent_weight, latent_scale, latent_codebook).
"""
buffers: dict[str, Any] = {}
for name, buf in self._buffers.items():
if buf is None:
buffers[name] = None
else:
buffers[name] = buf.detach().cpu().clone()
meta = self._collect_weight_meta()
teacher_dict = None
if self._teacher is not None:
teacher_dict = self._teacher.to_dict()
result = {
"buffers": buffers,
"meta": meta,
"quantizer_config": self._quantizer.to_config(),
"op_type": self.op_type,
"op_kwargs": self.op_kwargs,
"weight_shape": list(self.weight_shape),
"compute_dtype": self.compute_dtype,
"chunk_size": self._chunk_size,
"adaptive": self._adaptive is not None,
"dual_path": self._dual_path,
"teacher": teacher_dict,
"learnable": self._learnable,
}
# QAT learnable parameters (latent_weight, latent_scale, latent_codebook).
if self._learnable:
if hasattr(self, "latent_weight"):
result["latent_weight"] = self.latent_weight.detach().cpu().clone()
if hasattr(self, "latent_scale"):
result["latent_scale"] = self.latent_scale.detach().cpu().clone()
if hasattr(self, "latent_codebook"):
result["latent_codebook"] = self.latent_codebook.detach().cpu().clone()
return result
@classmethod
def from_dict(cls, d: dict[str, Any]) -> "QuantizedModule":
"""Reconstruct a QuantizedModule from a to_dict() dict.
If the dict contains learnable state (latent_weight, latent_scale),
the QuantizedModule is created with _learnable=True and those
parameters restored. This allows QAT save/load: trained latent weights
are preserved across save/load cycles.
"""
from agiws_neural_quant.quantizer import Quantizer
quantizer = Quantizer.from_config(d["quantizer_config"])
qw = QuantizedWeight(
weight_buffers={k: v for k, v in d["buffers"].items()
if v is not None and k != "bias"},
weight_meta=d["meta"],
)
teacher = None
if d.get("teacher") is not None:
teacher = cls.from_dict(d["teacher"])
bias = d["buffers"].get("bias", None)
qm = cls(
qw=qw,
quantizer=quantizer,
op_type=d["op_type"],
op_kwargs=d["op_kwargs"],
bias=bias,
weight_shape=tuple(d["weight_shape"]),
compute_dtype=d["compute_dtype"],
chunk_size=d["chunk_size"],
adaptive=d["adaptive"],
dual_path=d["dual_path"],
teacher=teacher,
)
# Restore QAT learnable parameters if present.
if d.get("learnable", False) and "latent_weight" in d:
import torch.nn as nn_module
qm._learnable = True
qm.latent_weight = nn_module.Parameter(d["latent_weight"].clone())
if "latent_scale" in d:
qm.latent_scale = nn_module.Parameter(d["latent_scale"].clone())
if "latent_codebook" in d:
qm.latent_codebook = nn_module.Parameter(d["latent_codebook"].clone())
return qm
def extra_repr(self) -> str:
parts = [f"op={self.op_type}", f"shape={tuple(self.weight_shape)}"]
parts.append(f"bias={self.bias is not None}")
parts.append(f"compute={self.compute_dtype}")
if self._chunk_size is not None:
parts.append(f"chunk={self._chunk_size}")
if self._dual_path:
parts.append("dual_path")
for k, v in self._quantizer_info.items():
parts.append(f"{k}={v}")
return ", ".join(parts)