RustyMark/dots.tts
0
1import math2 3import torch4import torch.nn as nn5 6from dots_tts.modules.backbone.layers import Mlp, MultiHeadAttention7 8 9def modulate(x, shift, scale, **_kwargs):10 return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)11 12 13class TimestepEmbedder(nn.Module):14 def __init__(self, hidden_size, frequency_embedding_size=256):15 super().__init__()16 self.mlp = nn.Sequential(17 nn.Linear(frequency_embedding_size, hidden_size, bias=True),18 nn.SiLU(),19 nn.Linear(hidden_size, hidden_size, bias=True),20 )21 self.frequency_embedding_size = frequency_embedding_size22 23 @staticmethod24 def timestep_embedding(t, dim, max_period=10000):25 half = dim // 226 freqs = torch.exp(27 -math.log(max_period)28 * torch.arange(start=0, end=half, dtype=torch.float32)29 / half30 ).to(device=t.device)31 args = t[:, None].float() * freqs[None]32 embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)33 if dim % 2:34 embedding = torch.cat(35 [embedding, torch.zeros_like(embedding[:, :1])], dim=-136 )37 return embedding38 39 def forward(self, t):40 t_freq = self.timestep_embedding(t, self.frequency_embedding_size)41 return self.mlp(t_freq)42 43 44class FinalLayer(nn.Module):45 def __init__(self, hidden_size, output_size):46 super().__init__()47 self.adaLN_modulation = nn.Sequential(48 nn.SiLU(),49 nn.Linear(hidden_size, 2 * hidden_size, bias=True),50 )51 self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-5)52 self.linear = nn.Linear(hidden_size, output_size, bias=True)53 54 def forward(self, x, c, **_kwargs):55 shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)56 x = modulate(self.norm(x), shift, scale)57 return self.linear(x)58 59 60class DiTBlock(nn.Module):61 def __init__(62 self,63 attention: nn.Module,64 ffn: nn.Module,65 hidden_size: int = 1024,66 modulation: bool = False,67 eps: float = 1e-5,68 **_kwargs,69 ):70 super().__init__()71 self.norm1 = nn.LayerNorm(72 hidden_size, elementwise_affine=not modulation, eps=eps73 )74 self.norm2 = nn.LayerNorm(75 hidden_size, elementwise_affine=not modulation, eps=eps76 )77 self.attn = attention78 self.ffn = ffn79 self.modulation = modulation80 if modulation:81 self.adaLN_modulation = nn.Sequential(82 nn.SiLU(),83 nn.Linear(hidden_size, 6 * hidden_size, bias=True),84 )85 86 def forward(self, x, condition=None, mask=None, **kwargs):87 if condition is None:88 assert not self.modulation, (89 "Without global condition, must set modulation to False"90 )91 else:92 assert self.modulation, "With global condition, must set modulation to True"93 shift_attn, scale_attn, gate_attn, shift_ffn, scale_ffn, gate_ffn = (94 self.adaLN_modulation(condition).chunk(6, dim=1)95 )96 97 if condition is not None:98 pack_indices = kwargs.get("pack_indices")99 if pack_indices is not None:100 gate_attn = gate_attn[pack_indices]101 gate_ffn = gate_ffn[pack_indices]102 else:103 gate_attn = gate_attn.unsqueeze(1)104 gate_ffn = gate_ffn.unsqueeze(1)105 106 if condition is not None:107 x = x + gate_attn * self.attn(108 modulate(self.norm1(x), shift_attn, scale_attn, **kwargs),109 mask=mask,110 **kwargs,111 )112 else:113 x = x + self.attn(self.norm1(x), mask=mask, **kwargs)114 115 if condition is not None:116 x = x + gate_ffn * self.ffn(117 modulate(self.norm2(x), shift_ffn, scale_ffn, **kwargs)118 )119 else:120 x = x + self.ffn(self.norm2(x), mask=mask)121 return x122 123 124class DiT(nn.Module):125 def __init__(126 self,127 in_dim,128 out_dim,129 transformer_config,130 *,131 mode: str = "flow_matching",132 ):133 super().__init__()134 if mode not in {"flow_matching", "meanflow"}:135 raise ValueError(136 f"DiT mode must be 'flow_matching' or 'meanflow', got {mode!r}."137 )138 139 transformer_kwargs = transformer_config.to_dict()140 model_dim = transformer_config.hidden_size141 self.mode = mode142 self.num_layers = transformer_config.num_layers143 144 self.input_layer = nn.Linear(in_dim, model_dim)145 self.time_embedder = TimestepEmbedder(model_dim)146 if mode == "meanflow":147 self.duration_embedder = TimestepEmbedder(model_dim)148 149 self.blocks = nn.ModuleList()150 for i in range(self.num_layers):151 attn_block = MultiHeadAttention(**transformer_kwargs, name=f"layer_{i}")152 ffn_block = Mlp(153 act_layer=lambda: nn.GELU(approximate="tanh"), **transformer_kwargs154 )155 self.blocks.append(156 DiTBlock(attention=attn_block, ffn=ffn_block, **transformer_kwargs)157 )158 159 self.output_layer = FinalLayer(model_dim, out_dim)160 self.initialize_weights()161 162 def initialize_weights(self):163 def _basic_init(module):164 if isinstance(module, nn.Linear):165 torch.nn.init.xavier_uniform_(module.weight)166 if module.bias is not None:167 nn.init.constant_(module.bias, 0)168 169 self.apply(_basic_init)170 171 nn.init.normal_(self.time_embedder.mlp[0].weight, std=0.02)172 nn.init.normal_(self.time_embedder.mlp[2].weight, std=0.02)173 174 for block in self.blocks:175 if hasattr(block, "adaLN_modulation"):176 nn.init.constant_(block.adaLN_modulation[-1].weight, 0)177 nn.init.constant_(block.adaLN_modulation[-1].bias, 0)178 179 nn.init.constant_(self.output_layer.adaLN_modulation[-1].weight, 0)180 nn.init.constant_(self.output_layer.adaLN_modulation[-1].bias, 0)181 nn.init.constant_(self.output_layer.linear.weight, 0)182 nn.init.constant_(self.output_layer.linear.bias, 0)183 184 def forward(185 self,186 x,187 timesteps,188 duration: torch.Tensor | None = None,189 mask=None,190 attn_mask=None,191 g_cond: torch.Tensor | None = None,192 **kwargs,193 ):194 t = self.time_embedder(timesteps)195 c = t196 duration_embedder = getattr(self, "duration_embedder", None)197 if duration_embedder is not None and duration is not None:198 c = c + duration_embedder(duration)199 if g_cond is not None:200 c = c + g_cond201 202 x = self.input_layer(x)203 for block in self.blocks:204 x = block(x, c, mask=attn_mask, **kwargs)205 return self.output_layer(x, c, **kwargs)206 