Team Ai
Apppublic

modelscope/DiffSynth-Painter

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
14likes
sd_motion.py200 linesDownload Raw Back to models
1from .sd_unet import SDUNet, Attention, GEGLU2import torch3from einops import rearrange, repeat4 5 6class TemporalTransformerBlock(torch.nn.Module):7 8    def __init__(self, dim, num_attention_heads, attention_head_dim, max_position_embeddings=32):9        super().__init__()10 11        # 1. Self-Attn12        self.pe1 = torch.nn.Parameter(torch.zeros(1, max_position_embeddings, dim))13        self.norm1 = torch.nn.LayerNorm(dim, elementwise_affine=True)14        self.attn1 = Attention(q_dim=dim, num_heads=num_attention_heads, head_dim=attention_head_dim, bias_out=True)15 16        # 2. Cross-Attn17        self.pe2 = torch.nn.Parameter(torch.zeros(1, max_position_embeddings, dim))18        self.norm2 = torch.nn.LayerNorm(dim, elementwise_affine=True)19        self.attn2 = Attention(q_dim=dim, num_heads=num_attention_heads, head_dim=attention_head_dim, bias_out=True)20 21        # 3. Feed-forward22        self.norm3 = torch.nn.LayerNorm(dim, elementwise_affine=True)23        self.act_fn = GEGLU(dim, dim * 4)24        self.ff = torch.nn.Linear(dim * 4, dim)25 26 27    def forward(self, hidden_states, batch_size=1):28 29        # 1. Self-Attention30        norm_hidden_states = self.norm1(hidden_states)31        norm_hidden_states = rearrange(norm_hidden_states, "(b f) h c -> (b h) f c", b=batch_size)32        attn_output = self.attn1(norm_hidden_states + self.pe1[:, :norm_hidden_states.shape[1]])33        attn_output = rearrange(attn_output, "(b h) f c -> (b f) h c", b=batch_size)34        hidden_states = attn_output + hidden_states35 36        # 2. Cross-Attention37        norm_hidden_states = self.norm2(hidden_states)38        norm_hidden_states = rearrange(norm_hidden_states, "(b f) h c -> (b h) f c", b=batch_size)39        attn_output = self.attn2(norm_hidden_states + self.pe2[:, :norm_hidden_states.shape[1]])40        attn_output = rearrange(attn_output, "(b h) f c -> (b f) h c", b=batch_size)41        hidden_states = attn_output + hidden_states42 43        # 3. Feed-forward44        norm_hidden_states = self.norm3(hidden_states)45        ff_output = self.act_fn(norm_hidden_states)46        ff_output = self.ff(ff_output)47        hidden_states = ff_output + hidden_states48 49        return hidden_states50 51 52class TemporalBlock(torch.nn.Module):53    54    def __init__(self, num_attention_heads, attention_head_dim, in_channels, num_layers=1, norm_num_groups=32, eps=1e-5):55        super().__init__()56        inner_dim = num_attention_heads * attention_head_dim57 58        self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=eps, affine=True)59        self.proj_in = torch.nn.Linear(in_channels, inner_dim)60 61        self.transformer_blocks = torch.nn.ModuleList([62            TemporalTransformerBlock(63                inner_dim,64                num_attention_heads,65                attention_head_dim66            )67            for d in range(num_layers)68        ])69 70        self.proj_out = torch.nn.Linear(inner_dim, in_channels)71 72    def forward(self, hidden_states, time_emb, text_emb, res_stack, batch_size=1):73        batch, _, height, width = hidden_states.shape74        residual = hidden_states75 76        hidden_states = self.norm(hidden_states)77        inner_dim = hidden_states.shape[1]78        hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * width, inner_dim)79        hidden_states = self.proj_in(hidden_states)80 81        for block in self.transformer_blocks:82            hidden_states = block(83                hidden_states,84                batch_size=batch_size85            )86 87        hidden_states = self.proj_out(hidden_states)88        hidden_states = hidden_states.reshape(batch, height, width, inner_dim).permute(0, 3, 1, 2).contiguous()89        hidden_states = hidden_states + residual90 91        return hidden_states, time_emb, text_emb, res_stack92 93 94class SDMotionModel(torch.nn.Module):95    def __init__(self):96        super().__init__()97        self.motion_modules = torch.nn.ModuleList([98            TemporalBlock(8, 40, 320, eps=1e-6),99            TemporalBlock(8, 40, 320, eps=1e-6),100            TemporalBlock(8, 80, 640, eps=1e-6),101            TemporalBlock(8, 80, 640, eps=1e-6),102            TemporalBlock(8, 160, 1280, eps=1e-6),103            TemporalBlock(8, 160, 1280, eps=1e-6),104            TemporalBlock(8, 160, 1280, eps=1e-6),105            TemporalBlock(8, 160, 1280, eps=1e-6),106            TemporalBlock(8, 160, 1280, eps=1e-6),107            TemporalBlock(8, 160, 1280, eps=1e-6),108            TemporalBlock(8, 160, 1280, eps=1e-6),109            TemporalBlock(8, 160, 1280, eps=1e-6),110            TemporalBlock(8, 160, 1280, eps=1e-6),111            TemporalBlock(8, 160, 1280, eps=1e-6),112            TemporalBlock(8, 160, 1280, eps=1e-6),113            TemporalBlock(8, 80, 640, eps=1e-6),114            TemporalBlock(8, 80, 640, eps=1e-6),115            TemporalBlock(8, 80, 640, eps=1e-6),116            TemporalBlock(8, 40, 320, eps=1e-6),117            TemporalBlock(8, 40, 320, eps=1e-6),118            TemporalBlock(8, 40, 320, eps=1e-6),119        ])120        self.call_block_id = {121            1: 0,122            4: 1,123            9: 2,124            12: 3,125            17: 4,126            20: 5,127            24: 6,128            26: 7,129            29: 8,130            32: 9,131            34: 10,132            36: 11,133            40: 12,134            43: 13,135            46: 14,136            50: 15,137            53: 16,138            56: 17,139            60: 18,140            63: 19,141            66: 20142        }143        144    def forward(self):145        pass146 147    @staticmethod148    def state_dict_converter():149        return SDMotionModelStateDictConverter()150 151 152class SDMotionModelStateDictConverter:153    def __init__(self):154        pass155 156    def from_diffusers(self, state_dict):157        rename_dict = {158            "norm": "norm",159            "proj_in": "proj_in",160            "transformer_blocks.0.attention_blocks.0.to_q": "transformer_blocks.0.attn1.to_q",161            "transformer_blocks.0.attention_blocks.0.to_k": "transformer_blocks.0.attn1.to_k",162            "transformer_blocks.0.attention_blocks.0.to_v": "transformer_blocks.0.attn1.to_v",163            "transformer_blocks.0.attention_blocks.0.to_out.0": "transformer_blocks.0.attn1.to_out",164            "transformer_blocks.0.attention_blocks.0.pos_encoder": "transformer_blocks.0.pe1",165            "transformer_blocks.0.attention_blocks.1.to_q": "transformer_blocks.0.attn2.to_q",166            "transformer_blocks.0.attention_blocks.1.to_k": "transformer_blocks.0.attn2.to_k",167            "transformer_blocks.0.attention_blocks.1.to_v": "transformer_blocks.0.attn2.to_v",168            "transformer_blocks.0.attention_blocks.1.to_out.0": "transformer_blocks.0.attn2.to_out",169            "transformer_blocks.0.attention_blocks.1.pos_encoder": "transformer_blocks.0.pe2",170            "transformer_blocks.0.norms.0": "transformer_blocks.0.norm1",171            "transformer_blocks.0.norms.1": "transformer_blocks.0.norm2",172            "transformer_blocks.0.ff.net.0.proj": "transformer_blocks.0.act_fn.proj",173            "transformer_blocks.0.ff.net.2": "transformer_blocks.0.ff",174            "transformer_blocks.0.ff_norm": "transformer_blocks.0.norm3",175            "proj_out": "proj_out",176        }177        name_list = sorted([i for i in state_dict if i.startswith("down_blocks.")])178        name_list += sorted([i for i in state_dict if i.startswith("mid_block.")])179        name_list += sorted([i for i in state_dict if i.startswith("up_blocks.")])180        state_dict_ = {}181        last_prefix, module_id = "", -1182        for name in name_list:183            names = name.split(".")184            prefix_index = names.index("temporal_transformer") + 1185            prefix = ".".join(names[:prefix_index])186            if prefix != last_prefix:187                last_prefix = prefix188                module_id += 1189            middle_name = ".".join(names[prefix_index:-1])190            suffix = names[-1]191            if "pos_encoder" in names:192                rename = ".".join(["motion_modules", str(module_id), rename_dict[middle_name]])193            else:194                rename = ".".join(["motion_modules", str(module_id), rename_dict[middle_name], suffix])195            state_dict_[rename] = state_dict[name]196        return state_dict_197    198    def from_civitai(self, state_dict):199        return self.from_diffusers(state_dict)200