File size: 3,785 Bytes
7709bf3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
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(),
        }