CoolFace
Apppublic

RustyMark/dots.tts

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
0likes
dit.py206 linesDownload Raw Back to backbone
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