Tessera-135M-Gate / quantizer.py
yava-code's picture
Publish Tessera-1B-Nano checkpoint
7709bf3 verified
Raw History Blame Contribute Delete
3.79 kB
from __future__ import annotations
from dataclasses import dataclass
import torch
import torch.nn.functional as F
from torch import Tensor, nn
@dataclass
class QuantizerOutput:
codes: Tensor
indices: Tensor
loss: Tensor
class CodebookTransform(nn.Module):
def __init__(self, segments: int, dim: int) -> None:
super().__init__()
self.layers = nn.ModuleList(
[
nn.Sequential(
nn.Linear(dim, 2 * dim),
nn.ReLU(),
nn.Linear(2 * dim, dim),
)
for _ in range(segments)
]
)
def forward(self, codebook: Tensor) -> Tensor:
return torch.stack(
[layer(codes) for layer, codes in zip(self.layers, codebook, strict=True)]
)
class ProductVectorQuantizer(nn.Module):
def __init__(self, hidden_size: int, segments: int, codebook_size: int) -> None:
super().__init__()
if hidden_size % segments:
raise ValueError("hidden_size must be divisible by segments")
self.hidden_size = hidden_size
self.segments = segments
self.codebook_size = codebook_size
self.segment_dim = hidden_size // segments
self.register_buffer(
"codebook",
torch.empty(segments, codebook_size, self.segment_dim),
)
self.transform = CodebookTransform(segments, self.segment_dim)
nn.init.normal_(self.codebook, std=0.02)
def transformed_codes(self) -> Tensor:
return self.transform(self.codebook)
def forward(self, concepts: Tensor) -> QuantizerOutput:
if concepts.shape[-1] != self.hidden_size:
raise ValueError(f"expected hidden size {self.hidden_size}, got {concepts.shape[-1]}")
shape = concepts.shape[:-1]
x = concepts.reshape(*shape, self.segments, self.segment_dim)
target = x.detach()
codes = self.transformed_codes()
distances = (
target.square().sum(dim=-1, keepdim=True)
+ codes.square()
.sum(dim=-1)
.view(*([1] * len(shape)), self.segments, self.codebook_size)
- 2 * torch.einsum("...sd,snd->...sn", target, codes)
)
indices = distances.argmin(dim=-1)
gather_index = indices.unsqueeze(-1).expand(*indices.shape, self.segment_dim)
expanded = codes.view(*([1] * len(shape)), *codes.shape).expand(*shape, *codes.shape)
quantized = expanded.gather(-2, gather_index.unsqueeze(-2)).squeeze(-2)
loss = F.mse_loss(quantized, target)
return QuantizerOutput(quantized.reshape(*shape, self.hidden_size), indices, loss)
def expected(self, logits: Tensor) -> Tensor:
if logits.shape[-2:] != (self.segments, self.codebook_size):
raise ValueError("logits do not match the product codebook")
probs = logits.float().softmax(dim=-1).to(logits.dtype)
predicted = torch.einsum("...sn,snd->...sd", probs, self.transformed_codes())
return predicted.flatten(-2)
@torch.no_grad()
def usage(self, indices: Tensor) -> dict[str, float]:
flat = indices.reshape(-1, self.segments)
perplexities = []
active = []
for segment in range(self.segments):
counts = torch.bincount(flat[:, segment], minlength=self.codebook_size).float()
probs = counts / counts.sum().clamp_min(1)
entropy = -(probs * probs.clamp_min(1e-12).log()).sum()
perplexities.append(entropy.exp())
active.append((counts > 0).float().mean())
return {
"codebook_perplexity": torch.stack(perplexities).mean().item(),
"codebook_usage": torch.stack(active).mean().item(),
}