CoolFace
Apppublic

RustyMark/dots.tts

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
0likes
semantic_encoder.py357 linesDownload Raw Back to backbone
1from __future__ import annotations2 3from dataclasses import dataclass4 5import torch6import torch.nn as nn7import torch.nn.functional as F8from einops import rearrange9 10from dots_tts.modules.backbone.layers import Conv1d, Mlp, MultiHeadAttention11 12 13@dataclass14class SemanticEncoderDecodeState:15    conv_tail: torch.Tensor16    layer_caches: tuple[tuple[torch.Tensor, torch.Tensor], ...]17    seq_len: int18 19 20class TransformerEncoderLayer(nn.Module):21    def __init__(22        self,23        hidden_size,24        num_heads=16,25        ffn_hidden_size=4096,26        attn_dropout=0.0,27        ffn_dropout=0.0,28        norm_layer="LayerNorm",29        **kwargs,30    ):31        super().__init__()32        self.attn = MultiHeadAttention(33            hidden_size,34            num_heads,35            attn_drop=attn_dropout,36            norm_layer=norm_layer,37            **kwargs,38        )39        norm_cls = getattr(nn, norm_layer)40        self.attn_norm = norm_cls(hidden_size)41        self.ffn = Mlp(42            hidden_size, ffn_hidden_size, dropout=ffn_dropout, act_layer=nn.SiLU43        )44        self.ffn_norm = norm_cls(hidden_size)45        self.hidden_size = hidden_size46 47    def _build_causal_mask(self, T: int, device):48        return torch.tril(torch.ones(T, T, dtype=torch.bool, device=device))49 50    def _build_padding_mask(self, x_lens, max_len: int, device):51        B = x_lens.size(0)52        positions = torch.arange(max_len, device=device).unsqueeze(0).expand(B, -1)53        return positions < x_lens.unsqueeze(1)54 55    def _fuse_attn_mask(self, causal_mask, padding_mask):56        if causal_mask is None and padding_mask is None:57            return None58        if causal_mask is None:59            row = padding_mask.unsqueeze(2)60            col = padding_mask.unsqueeze(1)61            return row & col62        if padding_mask is None:63            return causal_mask.unsqueeze(0)64 65        _B, _T = padding_mask.shape66        causal = causal_mask.unsqueeze(0)67        row = padding_mask.unsqueeze(2)68        col = padding_mask.unsqueeze(1)69        pad_2d = row & col70        return causal & pad_2d71 72    def forward(73        self,74        x,75        x_lens=None,76        causal=True,77    ):78        _B, T, C = x.shape79        assert self.hidden_size == C80        device = x.device81 82        causal_mask = self._build_causal_mask(T, device) if causal else None83        if x_lens is not None:84            padding_mask = self._build_padding_mask(x_lens, T, device)85        else:86            padding_mask = None87        fused_mask = self._fuse_attn_mask(causal_mask, padding_mask)88 89        h = self.attn_norm(x)90        h = self.attn(91            q=h,92            mask=fused_mask,93        )94        x = x + h95 96        h = self.ffn_norm(x)97        h = self.ffn(h)98        return x + h99 100    def decode_step(101        self,102        x,103        *,104        cache: tuple[torch.Tensor, torch.Tensor],105        positions: torch.Tensor,106    ):107        if x.size(1) <= 0:108            raise ValueError(109                "TransformerEncoderLayer.decode_step expects a non-empty input."110            )111 112        h = self.attn_norm(x)113        h, cache = self.attn.decode_step(h, cache=cache, positions=positions)114        x = x + h115 116        h = self.ffn_norm(x)117        h = self.ffn(h)118        return x + h, cache119 120 121class SuperviseEncoder(nn.Module):122    def __init__(self, config):123        super().__init__()124        self.hidden_size = config.get("hidden_size", 1024)125        self.layers = nn.ModuleList(126            [127                TransformerEncoderLayer(128                    hidden_size=self.hidden_size,129                    num_heads=config.get("num_heads", 16),130                    ffn_hidden_size=config.get("ffn_hidden_size", 4096),131                    norm_layer=config.get("norm_layer", "LayerNorm"),132                )133                for _ in range(config.get("num_layers", 6))134            ]135        )136        self.causal = config.get("causal", False)137 138    def forward(self, x, x_lens=None):139        batch_size, seq_len, _ = x.shape140        if x_lens is None:141            x_lens = torch.full(142                (batch_size,), seq_len, device=x.device, dtype=torch.long143            )144        for layer in self.layers:145            x = layer(x, x_lens=x_lens, causal=self.causal)146        return x147 148    def init_decode_state(149        self,150        *,151        batch_size: int,152        max_seq_len: int,153        device: torch.device,154        dtype: torch.dtype,155    ):156        layer_caches = []157        for layer in self.layers:158            cache_shape = (159                batch_size,160                layer.attn.num_heads,161                max_seq_len,162                layer.attn.head_dim,163            )164            layer_caches.append(165                (166                    torch.zeros(cache_shape, dtype=dtype, device=device),167                    torch.zeros(cache_shape, dtype=dtype, device=device),168                )169            )170        return tuple(layer_caches)171 172    def reset_decode_state(173        self,174        layer_caches: tuple[tuple[torch.Tensor, torch.Tensor], ...],175    ) -> None:176        if len(layer_caches) != len(self.layers):177            raise ValueError("Layer cache count does not match encoder depth.")178        for key_cache, value_cache in layer_caches:179            key_cache.zero_()180            value_cache.zero_()181 182    def decode_step(self, x, *, layer_caches, positions: torch.Tensor):183        if len(layer_caches) != len(self.layers):184            raise ValueError("Layer cache count does not match encoder depth.")185 186        for layer, cache in zip(self.layers, layer_caches, strict=True):187            x, _ = layer.decode_step(x, cache=cache, positions=positions)188        return x189 190 191class VAESemanticEncoder(nn.Module):192    def __init__(self, in_dim, out_dim, config):193        super().__init__()194        in_ds_rate = 2195        self.patch_size = int(config.patch_size)196        self.in_ds_rate = in_ds_rate197        self.ds_proj = Conv1d(198            in_dim, in_dim, kernel_size=in_ds_rate, stride=in_ds_rate, causal=True199        )200        self.in_proj = nn.Linear(in_dim, config.PatchEncoder.hidden_size)201        self.encoder = SuperviseEncoder(config.PatchEncoder)202        self.out_ds_rate = self.patch_size // in_ds_rate203        self.out_proj = nn.Linear(204            config.PatchEncoder.hidden_size * self.out_ds_rate, out_dim205        )206 207    def forward(self, x, x_lens=None):208        x = self._downsample(x)209        x = self.in_proj(x)210        z = self.encoder(x, x_lens=x_lens)211        return self._project_embeddings(z)212 213    def init_decode_state(214        self,215        *,216        max_audio_patch_count: int,217        batch_size: int,218        device: torch.device,219        dtype: torch.dtype,220    ) -> SemanticEncoderDecodeState:221        return SemanticEncoderDecodeState(222            conv_tail=torch.zeros(223                (batch_size, self.ds_proj.in_channels, self.ds_proj.left_padding),224                dtype=dtype,225                device=device,226            ),227            layer_caches=self.encoder.init_decode_state(228                batch_size=batch_size,229                max_seq_len=max_audio_patch_count * self.out_ds_rate,230                device=device,231                dtype=dtype,232            ),233            seq_len=0,234        )235 236    def reset_decode_state(self, state: SemanticEncoderDecodeState) -> None:237        state.conv_tail.zero_()238        self.encoder.reset_decode_state(state.layer_caches)239        state.seq_len = 0240 241    def prefill(242        self,243        x,244        state: SemanticEncoderDecodeState,245    ) -> tuple[torch.Tensor, SemanticEncoderDecodeState]:246        if x.ndim != 3:247            raise ValueError(248                f"VAESemanticEncoder.prefill expects rank-3 input, got {tuple(x.shape)}."249            )250        if x.size(1) % self.patch_size != 0:251            raise ValueError(252                f"Prompt latent length {x.size(1)} must be divisible by patch_size={self.patch_size}."253            )254 255        if x.size(1) == 0:256            return (257                x.new_zeros((x.size(0), 0, self.out_proj.out_features)),258                state,259            )260        if state.conv_tail.size(0) != x.size(0):261            raise ValueError(262                "VAESemanticEncoder.prefill batch size does not match decode state."263            )264 265        step_inputs = self.in_proj(self._downsample(x))266        expected_token_count = (x.size(1) // self.patch_size) * self.out_ds_rate267        if step_inputs.size(1) != expected_token_count:268            raise RuntimeError(269                "Patch encoder prefill produced an unexpected token count: "270                f"expected={expected_token_count} actual={step_inputs.size(1)}."271            )272 273        current_seq_len = state.seq_len274        next_seq_len = current_seq_len + step_inputs.size(1)275        cache_capacity = state.layer_caches[0][0].size(2)276        if next_seq_len > cache_capacity:277            raise ValueError(278                "Patch encoder prefill exceeds decode-state capacity: "279                f"required={next_seq_len} capacity={cache_capacity}."280            )281 282        positions = (283            torch.arange(step_inputs.size(1), device=x.device, dtype=torch.long)284            + current_seq_len285        )286        encoded = self.encoder.decode_step(287            step_inputs,288            layer_caches=state.layer_caches,289            positions=positions,290        )291        embedding = self._project_embeddings(encoded)292        raw = x.transpose(1, 2)293        state.conv_tail.copy_(raw[..., -self.ds_proj.left_padding :])294        state.seq_len = next_seq_len295        return embedding, state296 297    def decode_patch(298        self,299        latent_patch,300        conv_tail: torch.Tensor,301        layer_caches: tuple[tuple[torch.Tensor, torch.Tensor], ...],302        positions: torch.Tensor,303    ) -> tuple[torch.Tensor, torch.Tensor]:304        if latent_patch.ndim != 3:305            raise ValueError(306                f"VAESemanticEncoder.decode_patch expects rank-3 input, got {tuple(latent_patch.shape)}."307            )308        if latent_patch.size(1) != self.patch_size:309            raise ValueError(310                f"decode_patch expects patch length {self.patch_size}, got {latent_patch.size(1)}."311            )312        if positions.ndim != 1 or positions.size(0) != self.out_ds_rate:313            raise ValueError(314                "decode_patch positions must be a rank-1 tensor matching out_ds_rate."315            )316 317        step_inputs, conv_tail = self._downsample_step(318            latent_patch,319            conv_tail=conv_tail,320        )321        if step_inputs.size(1) != self.out_ds_rate:322            raise RuntimeError(323                f"Downsample step produced {step_inputs.size(1)} tokens, expected {self.out_ds_rate}."324            )325 326        encoded = self.encoder.decode_step(327            step_inputs,328            layer_caches=layer_caches,329            positions=positions,330        )331        embedding = self._project_embeddings(encoded)332        return embedding, conv_tail333 334    def _downsample(self, x):335        return self.ds_proj(x.transpose(1, 2)).transpose(1, 2)336 337    def _project_embeddings(self, z):338        if self.out_ds_rate > 1:339            z = rearrange(z, "b (s d) h -> b s (d h)", d=self.out_ds_rate)340        return self.out_proj(z)341 342    def _downsample_step(self, latent_patch, *, conv_tail):343        raw = latent_patch.transpose(1, 2)344        conv_input = torch.cat([conv_tail, raw], dim=-1)345 346        projected = F.conv1d(347            conv_input,348            self.ds_proj.weight,349            self.ds_proj.bias,350            stride=self.ds_proj.stride[0],351            padding=0,352            dilation=self.ds_proj.dilation[0],353            groups=self.ds_proj.groups,354        ).transpose(1, 2)355        new_conv_tail = raw[..., -self.ds_proj.left_padding :]356        return self.in_proj(projected), new_conv_tail357