RustyMark/dots.tts
0
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 