Residual Vector Quantizer. Args: dimension (int): Dimension of the codebooks. n_q (int): Number of residual vector quantizers used. bins (int): Codebook size. decay (float): Decay for exponential moving average over the codebooks. kmeans_init (bool): Wheth
| 26 | |
| 27 | |
| 28 | class ResidualVectorQuantizer(nn.Module): |
| 29 | """Residual Vector Quantizer. |
| 30 | Args: |
| 31 | dimension (int): Dimension of the codebooks. |
| 32 | n_q (int): Number of residual vector quantizers used. |
| 33 | bins (int): Codebook size. |
| 34 | decay (float): Decay for exponential moving average over the codebooks. |
| 35 | kmeans_init (bool): Whether to use kmeans to initialize the codebooks. |
| 36 | kmeans_iters (int): Number of iterations used for kmeans initialization. |
| 37 | threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes |
| 38 | that have an exponential moving average cluster size less than the specified threshold with |
| 39 | randomly selected vector from the current batch. |
| 40 | """ |
| 41 | |
| 42 | def __init__( |
| 43 | self, |
| 44 | dimension: int = 256, |
| 45 | n_q: int = 8, |
| 46 | bins: int = 1024, |
| 47 | decay: float = 0.99, |
| 48 | kmeans_init: bool = True, |
| 49 | kmeans_iters: int = 50, |
| 50 | threshold_ema_dead_code: int = 2, |
| 51 | ): |
| 52 | super().__init__() |
| 53 | self.n_q = n_q |
| 54 | self.dimension = dimension |
| 55 | self.bins = bins |
| 56 | self.decay = decay |
| 57 | self.kmeans_init = kmeans_init |
| 58 | self.kmeans_iters = kmeans_iters |
| 59 | self.threshold_ema_dead_code = threshold_ema_dead_code |
| 60 | self.vq = ResidualVectorQuantization( |
| 61 | dim=self.dimension, |
| 62 | codebook_size=self.bins, |
| 63 | num_quantizers=self.n_q, |
| 64 | decay=self.decay, |
| 65 | kmeans_init=self.kmeans_init, |
| 66 | kmeans_iters=self.kmeans_iters, |
| 67 | threshold_ema_dead_code=self.threshold_ema_dead_code, |
| 68 | ) |
| 69 | |
| 70 | def forward( |
| 71 | self, |
| 72 | x: torch.Tensor, |
| 73 | n_q: tp.Optional[int] = None, |
| 74 | layers: tp.Optional[list] = None, |
| 75 | ) -> QuantizedResult: |
| 76 | """Residual vector quantization on the given input tensor. |
| 77 | Args: |
| 78 | x (torch.Tensor): Input tensor. |
| 79 | n_q (int): Number of quantizer used to quantize. Default: All quantizers. |
| 80 | layers (list): Layer that need to return quantized. Defalt: None. |
| 81 | Returns: |
| 82 | QuantizedResult: |
| 83 | The quantized (or approximately quantized) representation with |
| 84 | the associated numbert quantizers and layer quantized required to return. |
| 85 | """ |