CoolFace
Apppublic

softwareweaver/MusicGen

sourceHugging Facecc-by-nc-4.0updated 11mo agoView on Hugging Face
0likes
core_vq.py401 linesDownload Raw Back to quantization
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the license found in the5# LICENSE file in the root directory of this source tree.6 7import typing as tp8 9from einops import rearrange, repeat10import flashy11import torch12from torch import nn, einsum13import torch.nn.functional as F14 15 16def exists(val: tp.Optional[tp.Any]) -> bool:17    return val is not None18 19 20def default(val: tp.Any, d: tp.Any) -> tp.Any:21    return val if exists(val) else d22 23 24def l2norm(t):25    return F.normalize(t, p=2, dim=-1)26 27 28def ema_inplace(moving_avg, new, decay: float):29    moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay))30 31 32def laplace_smoothing(x, n_categories: int, epsilon: float = 1e-5):33    return (x + epsilon) / (x.sum() + n_categories * epsilon)34 35 36def uniform_init(*shape: int):37    t = torch.empty(shape)38    nn.init.kaiming_uniform_(t)39    return t40 41 42def sample_vectors(samples, num: int):43    num_samples, device = samples.shape[0], samples.device44 45    if num_samples >= num:46        indices = torch.randperm(num_samples, device=device)[:num]47    else:48        indices = torch.randint(0, num_samples, (num,), device=device)49 50    return samples[indices]51 52 53def kmeans(samples, num_clusters: int, num_iters: int = 10):54    dim, dtype = samples.shape[-1], samples.dtype55 56    means = sample_vectors(samples, num_clusters)57 58    for _ in range(num_iters):59        diffs = rearrange(samples, "n d -> n () d") - rearrange(60            means, "c d -> () c d"61        )62        dists = -(diffs ** 2).sum(dim=-1)63 64        buckets = dists.max(dim=-1).indices65        bins = torch.bincount(buckets, minlength=num_clusters)66        zero_mask = bins == 067        bins_min_clamped = bins.masked_fill(zero_mask, 1)68 69        new_means = buckets.new_zeros(num_clusters, dim, dtype=dtype)70        new_means.scatter_add_(0, repeat(buckets, "n -> n d", d=dim), samples)71        new_means = new_means / bins_min_clamped[..., None]72 73        means = torch.where(zero_mask[..., None], means, new_means)74 75    return means, bins76 77 78def orthogonal_loss_fn(t):79    # eq (2) from https://arxiv.org/abs/2112.0038480    n = t.shape[0]81    normed_codes = l2norm(t)82    identity = torch.eye(n, device=t.device)83    cosine_sim = einsum("i d, j d -> i j", normed_codes, normed_codes)84    return ((cosine_sim - identity) ** 2).sum() / (n ** 2)85 86 87class EuclideanCodebook(nn.Module):88    """Codebook with Euclidean distance.89 90    Args:91        dim (int): Dimension.92        codebook_size (int): Codebook size.93        kmeans_init (bool): Whether to use k-means to initialize the codebooks.94            If set to true, run the k-means algorithm on the first training batch and use95            the learned centroids as initialization.96        kmeans_iters (int): Number of iterations used for k-means algorithm at initialization.97        decay (float): Decay for exponential moving average over the codebooks.98        epsilon (float): Epsilon value for numerical stability.99        threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes100            that have an exponential moving average cluster size less than the specified threshold with101            randomly selected vector from the current batch.102    """103    def __init__(104        self,105        dim: int,106        codebook_size: int,107        kmeans_init: int = False,108        kmeans_iters: int = 10,109        decay: float = 0.8,110        epsilon: float = 1e-5,111        threshold_ema_dead_code: int = 2,112    ):113        super().__init__()114        self.decay = decay115        init_fn: tp.Union[tp.Callable[..., torch.Tensor], tp.Any] = uniform_init if not kmeans_init else torch.zeros116        embed = init_fn(codebook_size, dim)117 118        self.codebook_size = codebook_size119 120        self.kmeans_iters = kmeans_iters121        self.epsilon = epsilon122        self.threshold_ema_dead_code = threshold_ema_dead_code123 124        self.register_buffer("inited", torch.Tensor([not kmeans_init]))125        self.register_buffer("cluster_size", torch.zeros(codebook_size))126        self.register_buffer("embed", embed)127        self.register_buffer("embed_avg", embed.clone())128 129    @torch.jit.ignore130    def init_embed_(self, data):131        if self.inited:132            return133 134        embed, cluster_size = kmeans(data, self.codebook_size, self.kmeans_iters)135        self.embed.data.copy_(embed)136        self.embed_avg.data.copy_(embed.clone())137        self.cluster_size.data.copy_(cluster_size)138        self.inited.data.copy_(torch.Tensor([True]))139        # Make sure all buffers across workers are in sync after initialization140        flashy.distrib.broadcast_tensors(self.buffers())141 142    def replace_(self, samples, mask):143        modified_codebook = torch.where(144            mask[..., None], sample_vectors(samples, self.codebook_size), self.embed145        )146        self.embed.data.copy_(modified_codebook)147 148    def expire_codes_(self, batch_samples):149        if self.threshold_ema_dead_code == 0:150            return151 152        expired_codes = self.cluster_size < self.threshold_ema_dead_code153        if not torch.any(expired_codes):154            return155 156        batch_samples = rearrange(batch_samples, "... d -> (...) d")157        self.replace_(batch_samples, mask=expired_codes)158        flashy.distrib.broadcast_tensors(self.buffers())159 160    def preprocess(self, x):161        x = rearrange(x, "... d -> (...) d")162        return x163 164    def quantize(self, x):165        embed = self.embed.t()166        dist = -(167            x.pow(2).sum(1, keepdim=True)168            - 2 * x @ embed169            + embed.pow(2).sum(0, keepdim=True)170        )171        embed_ind = dist.max(dim=-1).indices172        return embed_ind173 174    def postprocess_emb(self, embed_ind, shape):175        return embed_ind.view(*shape[:-1])176 177    def dequantize(self, embed_ind):178        quantize = F.embedding(embed_ind, self.embed)179        return quantize180 181    def encode(self, x):182        shape = x.shape183        # pre-process184        x = self.preprocess(x)185        # quantize186        embed_ind = self.quantize(x)187        # post-process188        embed_ind = self.postprocess_emb(embed_ind, shape)189        return embed_ind190 191    def decode(self, embed_ind):192        quantize = self.dequantize(embed_ind)193        return quantize194 195    def forward(self, x):196        shape, dtype = x.shape, x.dtype197        x = self.preprocess(x)198        self.init_embed_(x)199 200        embed_ind = self.quantize(x)201        embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)202        embed_ind = self.postprocess_emb(embed_ind, shape)203        quantize = self.dequantize(embed_ind)204 205        if self.training:206            # We do the expiry of code at that point as buffers are in sync207            # and all the workers will take the same decision.208            self.expire_codes_(x)209            ema_inplace(self.cluster_size, embed_onehot.sum(0), self.decay)210            embed_sum = x.t() @ embed_onehot211            ema_inplace(self.embed_avg, embed_sum.t(), self.decay)212            cluster_size = (213                laplace_smoothing(self.cluster_size, self.codebook_size, self.epsilon)214                * self.cluster_size.sum()215            )216            embed_normalized = self.embed_avg / cluster_size.unsqueeze(1)217            self.embed.data.copy_(embed_normalized)218 219        return quantize, embed_ind220 221 222class VectorQuantization(nn.Module):223    """Vector quantization implementation.224    Currently supports only euclidean distance.225 226    Args:227        dim (int): Dimension228        codebook_size (int): Codebook size229        codebook_dim (int): Codebook dimension. If not defined, uses the specified dimension in dim.230        decay (float): Decay for exponential moving average over the codebooks.231        epsilon (float): Epsilon value for numerical stability.232        kmeans_init (bool): Whether to use kmeans to initialize the codebooks.233        kmeans_iters (int): Number of iterations used for kmeans initialization.234        threshold_ema_dead_code (int):235        channels_last (bool): Channels are the last dimension in the input tensors.236        commitment_weight (float): Weight for commitment loss.237        orthogonal_reg_weight (float): Orthogonal regularization weights.238        orthogonal_reg_active_codes_only (bool): Apply orthogonal regularization only on active codes.239        orthogonal_reg_max_codes (optional int): Maximum number of codes to consider240            for orthogonal regularization.241        threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes242            that have an exponential moving average cluster size less than the specified threshold with243            randomly selected vector from the current batch.244    """245    def __init__(246        self,247        dim: int,248        codebook_size: int,249        codebook_dim: tp.Optional[int] = None,250        decay: float = 0.8,251        epsilon: float = 1e-5,252        kmeans_init: bool = False,253        kmeans_iters: int = 10,254        threshold_ema_dead_code: int = 2,255        channels_last: bool = False,256        commitment_weight: float = 1.,257        orthogonal_reg_weight: float = 0.0,258        orthogonal_reg_active_codes_only: bool = False,259        orthogonal_reg_max_codes: tp.Optional[int] = None,260    ):261        super().__init__()262        _codebook_dim: int = default(codebook_dim, dim)263 264        requires_projection = _codebook_dim != dim265        self.project_in = (nn.Linear(dim, _codebook_dim) if requires_projection else nn.Identity())266        self.project_out = (nn.Linear(_codebook_dim, dim) if requires_projection else nn.Identity())267 268        self.epsilon = epsilon269        self.commitment_weight = commitment_weight270 271        self.orthogonal_reg_weight = orthogonal_reg_weight272        self.orthogonal_reg_active_codes_only = orthogonal_reg_active_codes_only273        self.orthogonal_reg_max_codes = orthogonal_reg_max_codes274 275        self._codebook = EuclideanCodebook(dim=_codebook_dim, codebook_size=codebook_size,276                                           kmeans_init=kmeans_init, kmeans_iters=kmeans_iters,277                                           decay=decay, epsilon=epsilon,278                                           threshold_ema_dead_code=threshold_ema_dead_code)279        self.codebook_size = codebook_size280 281        self.channels_last = channels_last282 283    @property284    def codebook(self):285        return self._codebook.embed286 287    @property288    def inited(self):289        return self._codebook.inited290 291    def _preprocess(self, x):292        if not self.channels_last:293            x = rearrange(x, "b d n -> b n d")294        return x295 296    def _postprocess(self, quantize):297        if not self.channels_last:298            quantize = rearrange(quantize, "b n d -> b d n")299        return quantize300 301    def encode(self, x):302        x = self._preprocess(x)303        x = self.project_in(x)304        embed_in = self._codebook.encode(x)305        return embed_in306 307    def decode(self, embed_ind):308        quantize = self._codebook.decode(embed_ind)309        quantize = self.project_out(quantize)310        quantize = self._postprocess(quantize)311        return quantize312 313    def forward(self, x):314        device = x.device315        x = self._preprocess(x)316 317        x = self.project_in(x)318        quantize, embed_ind = self._codebook(x)319 320        if self.training:321            quantize = x + (quantize - x).detach()322 323        loss = torch.tensor([0.0], device=device, requires_grad=self.training)324 325        if self.training:326            if self.commitment_weight > 0:327                commit_loss = F.mse_loss(quantize.detach(), x)328                loss = loss + commit_loss * self.commitment_weight329 330            if self.orthogonal_reg_weight > 0:331                codebook = self.codebook332 333                if self.orthogonal_reg_active_codes_only:334                    # only calculate orthogonal loss for the activated codes for this batch335                    unique_code_ids = torch.unique(embed_ind)336                    codebook = codebook[unique_code_ids]337 338                num_codes = codebook.shape[0]339                if exists(self.orthogonal_reg_max_codes) and num_codes > self.orthogonal_reg_max_codes:340                    rand_ids = torch.randperm(num_codes, device=device)[:self.orthogonal_reg_max_codes]341                    codebook = codebook[rand_ids]342 343                orthogonal_reg_loss = orthogonal_loss_fn(codebook)344                loss = loss + orthogonal_reg_loss * self.orthogonal_reg_weight345 346        quantize = self.project_out(quantize)347        quantize = self._postprocess(quantize)348 349        return quantize, embed_ind, loss350 351 352class ResidualVectorQuantization(nn.Module):353    """Residual vector quantization implementation.354 355    Follows Algorithm 1. in https://arxiv.org/pdf/2107.03312.pdf356    """357    def __init__(self, *, num_quantizers, **kwargs):358        super().__init__()359        self.layers = nn.ModuleList(360            [VectorQuantization(**kwargs) for _ in range(num_quantizers)]361        )362 363    def forward(self, x, n_q: tp.Optional[int] = None):364        quantized_out = 0.0365        residual = x366 367        all_losses = []368        all_indices = []369 370        n_q = n_q or len(self.layers)371 372        for i, layer in enumerate(self.layers[:n_q]):373            quantized, indices, loss = layer(residual)374            residual = residual - quantized375            quantized_out = quantized_out + quantized376            all_indices.append(indices)377            all_losses.append(loss)378 379        out_losses, out_indices = map(torch.stack, (all_losses, all_indices))380        return quantized_out, out_indices, out_losses381 382    def encode(self, x: torch.Tensor, n_q: tp.Optional[int] = None) -> torch.Tensor:383        residual = x384        all_indices = []385        n_q = n_q or len(self.layers)386        for layer in self.layers[:n_q]:387            indices = layer.encode(residual)388            quantized = layer.decode(indices)389            residual = residual - quantized390            all_indices.append(indices)391        out_indices = torch.stack(all_indices)392        return out_indices393 394    def decode(self, q_indices: torch.Tensor) -> torch.Tensor:395        quantized_out = torch.tensor(0.0, device=q_indices.device)396        for i, indices in enumerate(q_indices):397            layer = self.layers[i]398            quantized = layer.decode(indices)399            quantized_out = quantized_out + quantized400        return quantized_out401