Automatic Speech Recognition
Transformers
asr
speaker-diarization
timestamps
quantization
low-bit
arm
on-device
Instructions to use yongyizang/TinyMOSS-Diarize with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use yongyizang/TinyMOSS-Diarize with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("automatic-speech-recognition", model="yongyizang/TinyMOSS-Diarize")# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("yongyizang/TinyMOSS-Diarize", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """Embedding / tied lm_head fake-quant variants for the embed PTQ study. | |
| New schemes live here so existing ``rtn_quant`` / ``wrap_model`` stay untouched. | |
| Supports W3+ RTN (including bits not allowed by ``rtn_quantize``), frequency-aware | |
| mixed precision, and product quantization. All wrappers keep an fp32 master and | |
| apply fake-quant only in forward; tied ``lm_head`` shares the same master. | |
| """ | |
| from __future__ import annotations | |
| from collections.abc import Iterable, Sequence | |
| from typing import Any | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from .rtn_quant import rtn_quantize | |
| def embed_storage_bits( | |
| *, | |
| vocab: int, | |
| dim: int, | |
| scheme: str, | |
| hot_tokens: int = 32_768, | |
| group_size: int = 64, | |
| pq_subvector: int = 8, | |
| pq_codebook_size: int = 256, | |
| scale_bits: int = 16, | |
| codebook_elem_bits: int = 16, | |
| ) -> dict[str, float]: | |
| """Report effective bits/param and packed size in MB for one tied matrix.""" | |
| params = float(vocab * dim) | |
| if scheme == "bf16": | |
| total_bits = params * 16.0 | |
| elif scheme == "w4_rtn_per_channel": | |
| total_bits = params * 4.0 + vocab * scale_bits | |
| elif scheme == "w3_rtn_group": | |
| groups = (dim + group_size - 1) // group_size | |
| total_bits = params * 3.0 + vocab * groups * scale_bits | |
| elif scheme == "freq_mixed": | |
| hot = min(hot_tokens, vocab) | |
| cold = vocab - hot | |
| cold_groups = (dim + group_size - 1) // group_size | |
| total_bits = ( | |
| hot * dim * 4.0 | |
| + hot * scale_bits | |
| + cold * dim * 2.0 | |
| + cold * cold_groups * scale_bits | |
| ) | |
| elif scheme == "pq": | |
| if dim % pq_subvector: | |
| raise ValueError(f"dim {dim} must be divisible by pq_subvector {pq_subvector}") | |
| n_sub = dim // pq_subvector | |
| index_bits = vocab * n_sub * 8.0 | |
| codebook_bits = n_sub * pq_codebook_size * pq_subvector * codebook_elem_bits | |
| total_bits = index_bits + codebook_bits | |
| else: | |
| raise ValueError(f"unknown scheme: {scheme}") | |
| return { | |
| "params": params, | |
| "total_bits": total_bits, | |
| "bits": total_bits / params, | |
| "mb": total_bits / 8.0 / 1_000_000.0, | |
| } | |
| def _validate_extended_rtn( | |
| weight: torch.Tensor, | |
| bits: int, | |
| granularity: str, | |
| group_size: int, | |
| ) -> None: | |
| if weight.ndim != 2: | |
| raise ValueError(f"RTN only supports 2-D weights, got {weight.ndim}-D") | |
| if not weight.is_floating_point(): | |
| raise ValueError("RTN fake quantization requires a floating-point weight") | |
| if bits not in {2, 3, 4, 8}: | |
| raise ValueError("bits must be one of {2, 3, 4, 8}") | |
| if granularity not in {"per_channel", "per_group"}: | |
| raise ValueError("granularity must be 'per_channel' or 'per_group'") | |
| if not isinstance(group_size, int) or group_size <= 0: | |
| raise ValueError("group_size must be a positive integer") | |
| def extended_rtn_quantize( | |
| weight: torch.Tensor, | |
| bits: int = 4, | |
| granularity: str = "per_channel", | |
| group_size: int = 64, | |
| ) -> torch.Tensor: | |
| """Symmetric absmax RTN supporting W2/W3 (and W4/W8). | |
| W4/W8 with the same args as ``rtn_quantize`` delegates to that helper so the | |
| study baseline matches the existing quantlib grid exactly. | |
| """ | |
| _validate_extended_rtn(weight, bits, granularity, group_size) | |
| if bits in {4, 8}: | |
| return rtn_quantize( | |
| weight, bits=bits, granularity=granularity, group_size=group_size | |
| ) | |
| rows, columns = weight.shape | |
| actual_group_size = columns if granularity == "per_channel" else group_size | |
| number_of_groups = (columns + actual_group_size - 1) // actual_group_size | |
| padded_columns = number_of_groups * actual_group_size | |
| groups = weight | |
| if padded_columns != columns: | |
| groups = F.pad(groups, (0, padded_columns - columns)) | |
| groups = groups.reshape(rows, number_of_groups, actual_group_size) | |
| qmax = 2 ** (bits - 1) - 1 | |
| scale_dtype = ( | |
| weight.dtype | |
| if weight.dtype in {torch.float16, torch.bfloat16} | |
| else torch.float32 | |
| ) | |
| scales = groups.float().abs().amax(dim=-1, keepdim=True) / qmax | |
| scales = scales.to(scale_dtype).float() | |
| safe_scales = torch.where(scales == 0, torch.ones_like(scales), scales) | |
| integers = torch.round(groups.float() / safe_scales).clamp(-qmax, qmax) | |
| quantized = (integers * scales).reshape(rows, -1)[:, :columns] | |
| return quantized.to(weight.dtype) | |
| class _ExtendedRTNSTE(torch.autograd.Function): | |
| def forward( | |
| ctx: Any, | |
| weight: torch.Tensor, | |
| bits: int, | |
| granularity: str, | |
| group_size: int, | |
| ) -> torch.Tensor: | |
| del ctx | |
| return extended_rtn_quantize(weight, bits, granularity, group_size) | |
| def backward( | |
| ctx: Any, grad_output: torch.Tensor | |
| ) -> tuple[torch.Tensor, None, None, None]: | |
| del ctx | |
| return grad_output, None, None, None | |
| def extended_rtn_quantize_ste( | |
| weight: torch.Tensor, | |
| bits: int = 4, | |
| granularity: str = "per_channel", | |
| group_size: int = 64, | |
| ) -> torch.Tensor: | |
| return _ExtendedRTNSTE.apply(weight, bits, granularity, group_size) | |
| def find_embed_tokens(model: nn.Module) -> tuple[str, nn.Embedding]: | |
| """Locate the LLM ``embed_tokens`` module (prefer language_model path).""" | |
| candidates: list[tuple[str, nn.Embedding]] = [] | |
| for name, module in model.named_modules(): | |
| if isinstance(module, nn.Embedding) and name.endswith("embed_tokens"): | |
| candidates.append((name, module)) | |
| if not candidates: | |
| raise RuntimeError("could not find embed_tokens Embedding on model") | |
| for name, module in candidates: | |
| if "language_model" in name: | |
| return name, module | |
| return candidates[0] | |
| def _set_module_by_name(root: nn.Module, dotted: str, value: nn.Module) -> None: | |
| parent_path, _, leaf = dotted.rpartition(".") | |
| parent = root.get_submodule(parent_path) if parent_path else root | |
| setattr(parent, leaf, value) | |
| class SharedQuantLMHead(nn.Module): | |
| """lm_head that matmuls against a sibling embedding's fake-quant weight. | |
| The embedding is stored as a plain attribute (not a registered submodule) so | |
| the tied module is not duplicated in ``named_modules`` / ``state_dict``. | |
| """ | |
| def __init__(self, embedding: nn.Module, bias: torch.Tensor | None = None) -> None: | |
| super().__init__() | |
| object.__setattr__(self, "_quant_embed", embedding) | |
| if bias is not None: | |
| self.bias = nn.Parameter(bias.detach().float().clone()) | |
| else: | |
| self.register_parameter("bias", None) | |
| def forward(self, hidden: torch.Tensor) -> torch.Tensor: | |
| weight = self._quant_embed.quantized_weight().to(dtype=hidden.dtype) | |
| bias = self.bias.to(hidden.dtype) if self.bias is not None else None | |
| return F.linear(hidden, weight, bias) | |
| def install_tied_embed_quant(model: nn.Module, embedding: nn.Module) -> dict[str, str]: | |
| """Replace ``embed_tokens`` and retie ``lm_head`` to the same fake-quant module.""" | |
| embed_path, _old = find_embed_tokens(model) | |
| _set_module_by_name(model, embed_path, embedding) | |
| if not hasattr(model, "lm_head"): | |
| raise RuntimeError("model has no lm_head to retie") | |
| old_head = model.lm_head | |
| bias = old_head.bias.detach().clone() if getattr(old_head, "bias", None) is not None else None | |
| model.lm_head = SharedQuantLMHead(embedding, bias=bias) | |
| return {"embed_tokens": embed_path, "lm_head": "lm_head"} | |
| def enable_tied_embed_eval_cache(model: nn.Module) -> None: | |
| """Cache fake-quant embed matrix once (required for fast autoregressive eval).""" | |
| for name, module in model.named_modules(): | |
| if not name.endswith("embed_tokens"): | |
| continue | |
| if hasattr(module, "enable_eval_cache"): | |
| module.enable_eval_cache() | |
| return | |
| if hasattr(module, "quantized_weight"): | |
| module.quantized_weight() | |
| return | |
| raise RuntimeError("no embed_tokens module found for eval cache") | |
| class ConfigurableRTNEmbedding(nn.Embedding): | |
| """Embedding with configurable RTN fake-quant (W2/W3/W4/W8).""" | |
| def __init__( | |
| self, | |
| num_embeddings: int, | |
| embedding_dim: int, | |
| padding_idx: int | None = None, | |
| max_norm: float | None = None, | |
| norm_type: float = 2.0, | |
| scale_grad_by_freq: bool = False, | |
| sparse: bool = False, | |
| *, | |
| bits: int = 4, | |
| granularity: str = "per_channel", | |
| group_size: int = 64, | |
| device: torch.device | str | None = None, | |
| ) -> None: | |
| _validate_extended_rtn( | |
| torch.empty(num_embeddings, embedding_dim), bits, granularity, group_size | |
| ) | |
| super().__init__( | |
| num_embeddings, | |
| embedding_dim, | |
| padding_idx=padding_idx, | |
| max_norm=max_norm, | |
| norm_type=norm_type, | |
| scale_grad_by_freq=scale_grad_by_freq, | |
| sparse=sparse, | |
| device=device, | |
| dtype=torch.float32, | |
| ) | |
| self.bits = bits | |
| self.granularity = granularity | |
| self.group_size = group_size | |
| self._eval_cache: torch.Tensor | None = None | |
| def from_embedding( | |
| cls, | |
| embedding: nn.Embedding, | |
| *, | |
| bits: int = 4, | |
| granularity: str = "per_channel", | |
| group_size: int = 64, | |
| ) -> "ConfigurableRTNEmbedding": | |
| converted = cls( | |
| embedding.num_embeddings, | |
| embedding.embedding_dim, | |
| padding_idx=embedding.padding_idx, | |
| max_norm=embedding.max_norm, | |
| norm_type=embedding.norm_type, | |
| scale_grad_by_freq=embedding.scale_grad_by_freq, | |
| sparse=embedding.sparse, | |
| bits=bits, | |
| granularity=granularity, | |
| group_size=group_size, | |
| device=embedding.weight.device, | |
| ) | |
| with torch.no_grad(): | |
| converted.weight.copy_(embedding.weight.detach().float()) | |
| converted.weight.requires_grad_(embedding.weight.requires_grad) | |
| converted.train(embedding.training) | |
| return converted | |
| def enable_eval_cache(self) -> torch.Tensor: | |
| """Materialize fake-quant weights once for PTQ / generation.""" | |
| with torch.no_grad(): | |
| cached = extended_rtn_quantize( | |
| self.weight, | |
| bits=self.bits, | |
| granularity=self.granularity, | |
| group_size=self.group_size, | |
| ).detach() | |
| self._eval_cache = cached | |
| return cached | |
| def clear_eval_cache(self) -> None: | |
| self._eval_cache = None | |
| def quantized_weight(self) -> torch.Tensor: | |
| if self._eval_cache is not None: | |
| return self._eval_cache | |
| return extended_rtn_quantize_ste( | |
| self.weight, | |
| bits=self.bits, | |
| granularity=self.granularity, | |
| group_size=self.group_size, | |
| ) | |
| def forward(self, input: torch.Tensor) -> torch.Tensor: | |
| return F.embedding( | |
| input, | |
| self.quantized_weight(), | |
| self.padding_idx, | |
| self.max_norm, | |
| self.norm_type, | |
| self.scale_grad_by_freq, | |
| self.sparse, | |
| ) | |
| def count_token_frequencies( | |
| texts: Iterable[str], | |
| tokenizer, | |
| *, | |
| vocab_size: int, | |
| ) -> torch.Tensor: | |
| """Return int64 counts ``[vocab_size]`` from raw transcript strings.""" | |
| counts = torch.zeros(vocab_size, dtype=torch.int64) | |
| batch: list[int] = [] | |
| flush_every = 1 << 16 | |
| def flush() -> None: | |
| nonlocal batch | |
| if not batch: | |
| return | |
| flat = torch.tensor(batch, dtype=torch.int64) | |
| flat = flat[(flat >= 0) & (flat < vocab_size)] | |
| if flat.numel(): | |
| counts.scatter_add_(0, flat, torch.ones_like(flat)) | |
| batch = [] | |
| for text in texts: | |
| if not text: | |
| continue | |
| ids = tokenizer(text, add_special_tokens=False)["input_ids"] | |
| if ids: | |
| batch.extend(ids) | |
| if len(batch) >= flush_every: | |
| flush() | |
| flush() | |
| return counts | |
| def topk_token_mask(counts: torch.Tensor, k: int) -> torch.Tensor: | |
| """Boolean mask over vocab rows marked as frequent (True = hot / W4).""" | |
| vocab = counts.numel() | |
| k = min(int(k), vocab) | |
| if k <= 0: | |
| return torch.zeros(vocab, dtype=torch.bool, device=counts.device) | |
| # Stable: break ties by lower token id so tests are deterministic. | |
| order = torch.argsort( | |
| counts.float() | |
| + (vocab - torch.arange(vocab, device=counts.device)).float() * 1e-12, | |
| descending=True, | |
| ) | |
| mask = torch.zeros(vocab, dtype=torch.bool, device=counts.device) | |
| mask[order[:k]] = True | |
| return mask | |
| class FreqMixedEmbedding(nn.Embedding): | |
| """Hot rows W4 per-channel; cold rows W2 group-quantized.""" | |
| def __init__( | |
| self, | |
| num_embeddings: int, | |
| embedding_dim: int, | |
| padding_idx: int | None = None, | |
| max_norm: float | None = None, | |
| norm_type: float = 2.0, | |
| scale_grad_by_freq: bool = False, | |
| sparse: bool = False, | |
| *, | |
| hot_bits: int = 4, | |
| cold_bits: int = 2, | |
| cold_group_size: int = 64, | |
| device: torch.device | str | None = None, | |
| ) -> None: | |
| super().__init__( | |
| num_embeddings, | |
| embedding_dim, | |
| padding_idx=padding_idx, | |
| max_norm=max_norm, | |
| norm_type=norm_type, | |
| scale_grad_by_freq=scale_grad_by_freq, | |
| sparse=sparse, | |
| device=device, | |
| dtype=torch.float32, | |
| ) | |
| self.hot_bits = hot_bits | |
| self.cold_bits = cold_bits | |
| self.cold_group_size = cold_group_size | |
| self._eval_cache: torch.Tensor | None = None | |
| self.register_buffer( | |
| "hot_mask", | |
| torch.zeros(num_embeddings, dtype=torch.bool, device=device), | |
| persistent=True, | |
| ) | |
| def from_embedding( | |
| cls, | |
| embedding: nn.Embedding, | |
| hot_mask: torch.Tensor, | |
| *, | |
| hot_bits: int = 4, | |
| cold_bits: int = 2, | |
| cold_group_size: int = 64, | |
| ) -> "FreqMixedEmbedding": | |
| if hot_mask.shape != (embedding.num_embeddings,): | |
| raise ValueError( | |
| f"hot_mask shape {tuple(hot_mask.shape)} != ({embedding.num_embeddings},)" | |
| ) | |
| converted = cls( | |
| embedding.num_embeddings, | |
| embedding.embedding_dim, | |
| padding_idx=embedding.padding_idx, | |
| max_norm=embedding.max_norm, | |
| norm_type=embedding.norm_type, | |
| scale_grad_by_freq=embedding.scale_grad_by_freq, | |
| sparse=embedding.sparse, | |
| hot_bits=hot_bits, | |
| cold_bits=cold_bits, | |
| cold_group_size=cold_group_size, | |
| device=embedding.weight.device, | |
| ) | |
| with torch.no_grad(): | |
| converted.weight.copy_(embedding.weight.detach().float()) | |
| converted.hot_mask.copy_(hot_mask.to(device=converted.hot_mask.device)) | |
| converted.weight.requires_grad_(embedding.weight.requires_grad) | |
| converted.train(embedding.training) | |
| return converted | |
| def enable_eval_cache(self) -> torch.Tensor: | |
| with torch.no_grad(): | |
| hot = extended_rtn_quantize( | |
| self.weight, | |
| bits=self.hot_bits, | |
| granularity="per_channel", | |
| group_size=self.weight.shape[1], | |
| ) | |
| cold = extended_rtn_quantize( | |
| self.weight, | |
| bits=self.cold_bits, | |
| granularity="per_group", | |
| group_size=self.cold_group_size, | |
| ) | |
| cached = torch.where(self.hot_mask.unsqueeze(-1), hot, cold).detach() | |
| self._eval_cache = cached | |
| return cached | |
| def clear_eval_cache(self) -> None: | |
| self._eval_cache = None | |
| def quantized_weight(self) -> torch.Tensor: | |
| if self._eval_cache is not None: | |
| return self._eval_cache | |
| hot = extended_rtn_quantize_ste( | |
| self.weight, | |
| bits=self.hot_bits, | |
| granularity="per_channel", | |
| group_size=self.weight.shape[1], | |
| ) | |
| cold = extended_rtn_quantize_ste( | |
| self.weight, | |
| bits=self.cold_bits, | |
| granularity="per_group", | |
| group_size=self.cold_group_size, | |
| ) | |
| mask = self.hot_mask.unsqueeze(-1) | |
| return torch.where(mask, hot, cold) | |
| def forward(self, input: torch.Tensor) -> torch.Tensor: | |
| return F.embedding( | |
| input, | |
| self.quantized_weight(), | |
| self.padding_idx, | |
| self.max_norm, | |
| self.norm_type, | |
| self.scale_grad_by_freq, | |
| self.sparse, | |
| ) | |
| def _kmeans_torch( | |
| points: torch.Tensor, | |
| k: int, | |
| *, | |
| iters: int = 15, | |
| seed: int = 0, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| """Lloyd k-means on ``[N, D]``; returns ``(centroids [K,D], labels [N])``.""" | |
| if points.ndim != 2: | |
| raise ValueError("points must be 2-D") | |
| n, dim = points.shape | |
| del dim | |
| if n == 0: | |
| raise ValueError("empty points for k-means") | |
| k = min(k, n) | |
| generator = torch.Generator(device=points.device) | |
| generator.manual_seed(seed) | |
| perm = torch.randperm(n, generator=generator, device=points.device) | |
| centroids = points[perm[:k]].clone() | |
| labels = torch.zeros(n, dtype=torch.int64, device=points.device) | |
| flat = points.float() | |
| for _ in range(iters): | |
| chunk = max(1, min(n, 8192)) | |
| label_chunks: list[torch.Tensor] = [] | |
| for start in range(0, n, chunk): | |
| block = flat[start : start + chunk] | |
| dist = torch.cdist(block, centroids.float(), p=2) | |
| label_chunks.append(dist.argmin(dim=1)) | |
| labels = torch.cat(label_chunks, dim=0) | |
| for center_id in range(k): | |
| members = flat[labels == center_id] | |
| if members.numel() == 0: | |
| idx = int( | |
| torch.randint( | |
| 0, n, (1,), generator=generator, device=points.device | |
| ).item() | |
| ) | |
| centroids[center_id] = flat[idx] | |
| else: | |
| centroids[center_id] = members.mean(dim=0) | |
| return centroids.to(points.dtype), labels | |
| class PQEmbedding(nn.Module): | |
| """Product-quantization fake embedding: M×(K×d) codebooks + uint8 indices.""" | |
| def __init__( | |
| self, | |
| num_embeddings: int, | |
| embedding_dim: int, | |
| *, | |
| subvector_dim: int = 8, | |
| codebook_size: int = 256, | |
| padding_idx: int | None = None, | |
| device: torch.device | str | None = None, | |
| ) -> None: | |
| super().__init__() | |
| if embedding_dim % subvector_dim: | |
| raise ValueError( | |
| f"embedding_dim ({embedding_dim}) must be divisible by " | |
| f"subvector_dim ({subvector_dim})" | |
| ) | |
| self.num_embeddings = num_embeddings | |
| self.embedding_dim = embedding_dim | |
| self.subvector_dim = subvector_dim | |
| self.codebook_size = codebook_size | |
| self.n_subvectors = embedding_dim // subvector_dim | |
| self.padding_idx = padding_idx | |
| self.weight = nn.Parameter( | |
| torch.empty(num_embeddings, embedding_dim, device=device, dtype=torch.float32), | |
| requires_grad=False, | |
| ) | |
| self.register_buffer( | |
| "codebooks", | |
| torch.zeros( | |
| self.n_subvectors, | |
| codebook_size, | |
| subvector_dim, | |
| device=device, | |
| dtype=torch.float32, | |
| ), | |
| persistent=True, | |
| ) | |
| self.register_buffer( | |
| "codes", | |
| torch.zeros(num_embeddings, self.n_subvectors, device=device, dtype=torch.uint8), | |
| persistent=True, | |
| ) | |
| self.register_buffer( | |
| "_reconstructed", | |
| torch.zeros(num_embeddings, embedding_dim, device=device, dtype=torch.float32), | |
| persistent=False, | |
| ) | |
| self._cache_valid = False | |
| def from_embedding( | |
| cls, | |
| embedding: nn.Embedding, | |
| *, | |
| subvector_dim: int = 8, | |
| codebook_size: int = 256, | |
| sample_frac: float = 0.2, | |
| kmeans_iters: int = 15, | |
| seed: int = 0, | |
| device: torch.device | str | None = None, | |
| ) -> "PQEmbedding": | |
| target_device = torch.device(device) if device is not None else embedding.weight.device | |
| converted = cls( | |
| embedding.num_embeddings, | |
| embedding.embedding_dim, | |
| subvector_dim=subvector_dim, | |
| codebook_size=codebook_size, | |
| padding_idx=embedding.padding_idx, | |
| device=target_device, | |
| ) | |
| with torch.no_grad(): | |
| converted.weight.copy_(embedding.weight.detach().float().to(target_device)) | |
| converted.fit( | |
| sample_frac=sample_frac, | |
| kmeans_iters=kmeans_iters, | |
| seed=seed, | |
| ) | |
| return converted | |
| def fit( | |
| self, | |
| *, | |
| sample_frac: float = 0.2, | |
| kmeans_iters: int = 15, | |
| seed: int = 0, | |
| ) -> None: | |
| weight = self.weight.detach() | |
| vocab, _dim = weight.shape | |
| n_sub = self.n_subvectors | |
| d = self.subvector_dim | |
| k = self.codebook_size | |
| sample_n = max(k, min(vocab, int(round(vocab * sample_frac)))) | |
| generator = torch.Generator(device=weight.device) | |
| generator.manual_seed(seed) | |
| sample_idx = torch.randperm(vocab, generator=generator, device=weight.device)[ | |
| :sample_n | |
| ] | |
| sampled = weight[sample_idx].reshape(sample_n, n_sub, d) | |
| codebooks = torch.empty_like(self.codebooks) | |
| for sub in range(n_sub): | |
| points = sampled[:, sub, :].contiguous() | |
| centroids, _ = _kmeans_torch( | |
| points, k, iters=kmeans_iters, seed=seed + sub | |
| ) | |
| if centroids.shape[0] < k: | |
| padded = torch.zeros(k, d, device=weight.device, dtype=weight.dtype) | |
| padded[: centroids.shape[0]] = centroids | |
| centroids = padded | |
| codebooks[sub] = centroids | |
| self.codebooks.copy_(codebooks) | |
| codes = torch.empty(vocab, n_sub, device=weight.device, dtype=torch.int64) | |
| reshaped = weight.reshape(vocab, n_sub, d) | |
| chunk = max(1, min(vocab, 4096)) | |
| for sub in range(n_sub): | |
| centroids = codebooks[sub].float() | |
| label_chunks: list[torch.Tensor] = [] | |
| for start in range(0, vocab, chunk): | |
| block = reshaped[start : start + chunk, sub, :].float() | |
| dist = torch.cdist(block, centroids, p=2) | |
| label_chunks.append(dist.argmin(dim=1)) | |
| codes[:, sub] = torch.cat(label_chunks, dim=0) | |
| self.codes.copy_(codes.to(torch.uint8)) | |
| self._cache_valid = False | |
| self._refresh_cache() | |
| def _refresh_cache(self) -> None: | |
| vocab = self.num_embeddings | |
| n_sub = self.n_subvectors | |
| d = self.subvector_dim | |
| reconstructed = torch.empty( | |
| vocab, n_sub, d, device=self.codebooks.device, dtype=self.codebooks.dtype | |
| ) | |
| codes = self.codes.long() | |
| for sub in range(n_sub): | |
| reconstructed[:, sub, :] = self.codebooks[sub][codes[:, sub]] | |
| self._reconstructed.copy_(reconstructed.reshape(vocab, self.embedding_dim)) | |
| self._cache_valid = True | |
| def enable_eval_cache(self) -> torch.Tensor: | |
| self._refresh_cache() | |
| return self._reconstructed | |
| def quantized_weight(self) -> torch.Tensor: | |
| if not self._cache_valid: | |
| self._refresh_cache() | |
| return self._reconstructed | |
| def forward(self, input: torch.Tensor) -> torch.Tensor: | |
| return F.embedding(input, self.quantized_weight(), self.padding_idx) | |
| def apply_scheme_a(model: nn.Module) -> dict[str, Any]: | |
| """W4 RTN per-channel baseline on the tied embed matrix.""" | |
| _path, embed = find_embed_tokens(model) | |
| quantized = ConfigurableRTNEmbedding.from_embedding( | |
| embed, bits=4, granularity="per_channel", group_size=embed.embedding_dim | |
| ) | |
| paths = install_tied_embed_quant(model, quantized) | |
| stats = embed_storage_bits( | |
| vocab=quantized.num_embeddings, | |
| dim=quantized.embedding_dim, | |
| scheme="w4_rtn_per_channel", | |
| ) | |
| return {"name": "A_w4", "paths": paths, **stats} | |
| def apply_scheme_b(model: nn.Module, *, group_size: int = 64) -> dict[str, Any]: | |
| """W3 RTN group64 on the tied embed matrix.""" | |
| _path, embed = find_embed_tokens(model) | |
| quantized = ConfigurableRTNEmbedding.from_embedding( | |
| embed, bits=3, granularity="per_group", group_size=group_size | |
| ) | |
| paths = install_tied_embed_quant(model, quantized) | |
| stats = embed_storage_bits( | |
| vocab=quantized.num_embeddings, | |
| dim=quantized.embedding_dim, | |
| scheme="w3_rtn_group", | |
| group_size=group_size, | |
| ) | |
| return {"name": "B_w3", "paths": paths, **stats} | |
| def apply_scheme_c( | |
| model: nn.Module, | |
| hot_mask: torch.Tensor, | |
| *, | |
| cold_group_size: int = 64, | |
| ) -> dict[str, Any]: | |
| """Frequency-aware mixed W4/W2 on the tied embed matrix.""" | |
| _path, embed = find_embed_tokens(model) | |
| quantized = FreqMixedEmbedding.from_embedding( | |
| embed, hot_mask=hot_mask.to(embed.weight.device), cold_group_size=cold_group_size | |
| ) | |
| paths = install_tied_embed_quant(model, quantized) | |
| stats = embed_storage_bits( | |
| vocab=quantized.num_embeddings, | |
| dim=quantized.embedding_dim, | |
| scheme="freq_mixed", | |
| hot_tokens=int(hot_mask.sum().item()), | |
| group_size=cold_group_size, | |
| ) | |
| return {"name": "C_freq", "paths": paths, **stats} | |
| def apply_scheme_d( | |
| model: nn.Module, | |
| *, | |
| sample_frac: float = 0.2, | |
| device: torch.device | str | None = None, | |
| ) -> dict[str, Any]: | |
| """PQ 8-dim × 256 codebooks on the tied embed matrix.""" | |
| _path, embed = find_embed_tokens(model) | |
| quantized = PQEmbedding.from_embedding( | |
| embed, sample_frac=sample_frac, device=device or embed.weight.device | |
| ) | |
| paths = install_tied_embed_quant(model, quantized) | |
| stats = embed_storage_bits( | |
| vocab=quantized.num_embeddings, | |
| dim=quantized.embedding_dim, | |
| scheme="pq", | |
| ) | |
| return {"name": "D_pq", "paths": paths, **stats} | |
| NINE_LANG_TRAIN_MANIFESTS: tuple[str, ...] = ( | |
| "fleurs-zh-train", | |
| "fleurs-en-train", | |
| "fleurs-ja-train", | |
| "fleurs-ko-train", | |
| "fleurs-de-train", | |
| "fleurs-es-train", | |
| "fleurs-fr-train", | |
| "fleurs-it-train", | |
| "fleurs-ru-train", | |
| ) | |
| def load_manifest_texts(manifest_paths: Sequence[Any]) -> list[str]: | |
| from datapipe.io import read_jsonl | |
| texts: list[str] = [] | |
| for path in manifest_paths: | |
| for row in read_jsonl(path): | |
| text = row.get("text") or row.get("target") or "" | |
| if text: | |
| texts.append(str(text)) | |
| return texts | |