Team Ai
Apppublic

Kafke/Code-Realize-TTS

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
hubert_model.py222 linesDownload Raw Back to root
1import copy2from typing import Optional, Tuple3import random4 5import torch6import torch.nn as nn7import torch.nn.functional as F8from torch.nn.modules.utils import consume_prefix_in_state_dict_if_present9 10class Hubert(nn.Module):11    def __init__(self, num_label_embeddings: int = 100, mask: bool = True):12        super().__init__()13        self._mask = mask14        self.feature_extractor = FeatureExtractor()15        self.feature_projection = FeatureProjection()16        self.positional_embedding = PositionalConvEmbedding()17        self.norm = nn.LayerNorm(768)18        self.dropout = nn.Dropout(0.1)19        self.encoder = TransformerEncoder(20            nn.TransformerEncoderLayer(21                768, 12, 3072, activation="gelu", batch_first=True22            ),23            12,24        )25        self.proj = nn.Linear(768, 256)26 27        self.masked_spec_embed = nn.Parameter(torch.FloatTensor(768).uniform_())28        self.label_embedding = nn.Embedding(num_label_embeddings, 256)29 30    def mask(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:31        mask = None32        if self.training and self._mask:33            mask = _compute_mask((x.size(0), x.size(1)), 0.8, 10, x.device, 2)34            x[mask] = self.masked_spec_embed.to(x.dtype)35        return x, mask36 37    def encode(38        self, x: torch.Tensor, layer: Optional[int] = None39    ) -> Tuple[torch.Tensor, torch.Tensor]:40        x = self.feature_extractor(x)41        x = self.feature_projection(x.transpose(1, 2))42        x, mask = self.mask(x)43        x = x + self.positional_embedding(x)44        x = self.dropout(self.norm(x))45        x = self.encoder(x, output_layer=layer)46        return x, mask47 48    def logits(self, x: torch.Tensor) -> torch.Tensor:49        logits = torch.cosine_similarity(50            x.unsqueeze(2),51            self.label_embedding.weight.unsqueeze(0).unsqueeze(0),52            dim=-1,53        )54        return logits / 0.155 56    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:57        x, mask = self.encode(x)58        x = self.proj(x)59        logits = self.logits(x)60        return logits, mask61 62 63class HubertSoft(Hubert):64    def __init__(self):65        super().__init__()66 67    @torch.inference_mode()68    def units(self, wav: torch.Tensor) -> torch.Tensor:69        wav = F.pad(wav, ((400 - 320) // 2, (400 - 320) // 2))70        x, _ = self.encode(wav)71        return self.proj(x)72 73 74class FeatureExtractor(nn.Module):75    def __init__(self):76        super().__init__()77        self.conv0 = nn.Conv1d(1, 512, 10, 5, bias=False)78        self.norm0 = nn.GroupNorm(512, 512)79        self.conv1 = nn.Conv1d(512, 512, 3, 2, bias=False)80        self.conv2 = nn.Conv1d(512, 512, 3, 2, bias=False)81        self.conv3 = nn.Conv1d(512, 512, 3, 2, bias=False)82        self.conv4 = nn.Conv1d(512, 512, 3, 2, bias=False)83        self.conv5 = nn.Conv1d(512, 512, 2, 2, bias=False)84        self.conv6 = nn.Conv1d(512, 512, 2, 2, bias=False)85 86    def forward(self, x: torch.Tensor) -> torch.Tensor:87        x = F.gelu(self.norm0(self.conv0(x)))88        x = F.gelu(self.conv1(x))89        x = F.gelu(self.conv2(x))90        x = F.gelu(self.conv3(x))91        x = F.gelu(self.conv4(x))92        x = F.gelu(self.conv5(x))93        x = F.gelu(self.conv6(x))94        return x95 96 97class FeatureProjection(nn.Module):98    def __init__(self):99        super().__init__()100        self.norm = nn.LayerNorm(512)101        self.projection = nn.Linear(512, 768)102        self.dropout = nn.Dropout(0.1)103 104    def forward(self, x: torch.Tensor) -> torch.Tensor:105        x = self.norm(x)106        x = self.projection(x)107        x = self.dropout(x)108        return x109 110 111class PositionalConvEmbedding(nn.Module):112    def __init__(self):113        super().__init__()114        self.conv = nn.Conv1d(115            768,116            768,117            kernel_size=128,118            padding=128 // 2,119            groups=16,120        )121        self.conv = nn.utils.weight_norm(self.conv, name="weight", dim=2)122 123    def forward(self, x: torch.Tensor) -> torch.Tensor:124        x = self.conv(x.transpose(1, 2))125        x = F.gelu(x[:, :, :-1])126        return x.transpose(1, 2)127 128 129class TransformerEncoder(nn.Module):130    def __init__(131        self, encoder_layer: nn.TransformerEncoderLayer, num_layers: int132    ) -> None:133        super(TransformerEncoder, self).__init__()134        self.layers = nn.ModuleList(135            [copy.deepcopy(encoder_layer) for _ in range(num_layers)]136        )137        self.num_layers = num_layers138 139    def forward(140        self,141        src: torch.Tensor,142        mask: torch.Tensor = None,143        src_key_padding_mask: torch.Tensor = None,144        output_layer: Optional[int] = None,145    ) -> torch.Tensor:146        output = src147        for layer in self.layers[:output_layer]:148            output = layer(149                output, src_mask=mask, src_key_padding_mask=src_key_padding_mask150            )151        return output152 153 154def _compute_mask(155    shape: Tuple[int, int],156    mask_prob: float,157    mask_length: int,158    device: torch.device,159    min_masks: int = 0,160) -> torch.Tensor:161    batch_size, sequence_length = shape162 163    if mask_length < 1:164        raise ValueError("`mask_length` has to be bigger than 0.")165 166    if mask_length > sequence_length:167        raise ValueError(168            f"`mask_length` has to be smaller than `sequence_length`, but got `mask_length`: {mask_length} and `sequence_length`: {sequence_length}`"169        )170 171    # compute number of masked spans in batch172    num_masked_spans = int(mask_prob * sequence_length / mask_length + random.random())173    num_masked_spans = max(num_masked_spans, min_masks)174 175    # make sure num masked indices <= sequence_length176    if num_masked_spans * mask_length > sequence_length:177        num_masked_spans = sequence_length // mask_length178 179    # SpecAugment mask to fill180    mask = torch.zeros((batch_size, sequence_length), device=device, dtype=torch.bool)181 182    # uniform distribution to sample from, make sure that offset samples are < sequence_length183    uniform_dist = torch.ones(184        (batch_size, sequence_length - (mask_length - 1)), device=device185    )186 187    # get random indices to mask188    mask_indices = torch.multinomial(uniform_dist, num_masked_spans)189 190    # expand masked indices to masked spans191    mask_indices = (192        mask_indices.unsqueeze(dim=-1)193        .expand((batch_size, num_masked_spans, mask_length))194        .reshape(batch_size, num_masked_spans * mask_length)195    )196    offsets = (197        torch.arange(mask_length, device=device)[None, None, :]198        .expand((batch_size, num_masked_spans, mask_length))199        .reshape(batch_size, num_masked_spans * mask_length)200    )201    mask_idxs = mask_indices + offsets202 203    # scatter indices to mask204    mask = mask.scatter(1, mask_idxs, True)205 206    return mask207 208 209def hubert_soft(210    path: str211) -> HubertSoft:212    r"""HuBERT-Soft from `"A Comparison of Discrete and Soft Speech Units for Improved Voice Conversion"`.213    Args:214        path (str): path of a pretrained model215    """216    hubert = HubertSoft()217    checkpoint = torch.load(path)218    consume_prefix_in_state_dict_if_present(checkpoint, "module.")219    hubert.load_state_dict(checkpoint)220    hubert.eval()221    return hubert222