Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
sdxl_motion.py105 linesDownload Raw Back to models
1from .sd_motion import TemporalBlock2import torch3 4 5 6class SDXLMotionModel(torch.nn.Module):7    def __init__(self):8        super().__init__()9        self.motion_modules = torch.nn.ModuleList([10            TemporalBlock(8, 320//8, 320, eps=1e-6),11            TemporalBlock(8, 320//8, 320, eps=1e-6),12 13            TemporalBlock(8, 640//8, 640, eps=1e-6),14            TemporalBlock(8, 640//8, 640, eps=1e-6),15 16            TemporalBlock(8, 1280//8, 1280, eps=1e-6),17            TemporalBlock(8, 1280//8, 1280, eps=1e-6),18 19            TemporalBlock(8, 1280//8, 1280, eps=1e-6),20            TemporalBlock(8, 1280//8, 1280, eps=1e-6),21            TemporalBlock(8, 1280//8, 1280, eps=1e-6),22 23            TemporalBlock(8, 640//8, 640, eps=1e-6),24            TemporalBlock(8, 640//8, 640, eps=1e-6),25            TemporalBlock(8, 640//8, 640, eps=1e-6),26 27            TemporalBlock(8, 320//8, 320, eps=1e-6),28            TemporalBlock(8, 320//8, 320, eps=1e-6),29            TemporalBlock(8, 320//8, 320, eps=1e-6),30        ])31        self.call_block_id = {32            0: 0,33            2: 1,34            7: 2,35            10: 3,36            15: 4,37            18: 5,38            25: 6,39            28: 7,40            31: 8,41            35: 9,42            38: 10,43            41: 11,44            44: 12,45            46: 13,46            48: 14,47        }48        49    def forward(self):50        pass51 52    @staticmethod53    def state_dict_converter():54        return SDMotionModelStateDictConverter()55 56 57class SDMotionModelStateDictConverter:58    def __init__(self):59        pass60 61    def from_diffusers(self, state_dict):62        rename_dict = {63            "norm": "norm",64            "proj_in": "proj_in",65            "transformer_blocks.0.attention_blocks.0.to_q": "transformer_blocks.0.attn1.to_q",66            "transformer_blocks.0.attention_blocks.0.to_k": "transformer_blocks.0.attn1.to_k",67            "transformer_blocks.0.attention_blocks.0.to_v": "transformer_blocks.0.attn1.to_v",68            "transformer_blocks.0.attention_blocks.0.to_out.0": "transformer_blocks.0.attn1.to_out",69            "transformer_blocks.0.attention_blocks.0.pos_encoder": "transformer_blocks.0.pe1",70            "transformer_blocks.0.attention_blocks.1.to_q": "transformer_blocks.0.attn2.to_q",71            "transformer_blocks.0.attention_blocks.1.to_k": "transformer_blocks.0.attn2.to_k",72            "transformer_blocks.0.attention_blocks.1.to_v": "transformer_blocks.0.attn2.to_v",73            "transformer_blocks.0.attention_blocks.1.to_out.0": "transformer_blocks.0.attn2.to_out",74            "transformer_blocks.0.attention_blocks.1.pos_encoder": "transformer_blocks.0.pe2",75            "transformer_blocks.0.norms.0": "transformer_blocks.0.norm1",76            "transformer_blocks.0.norms.1": "transformer_blocks.0.norm2",77            "transformer_blocks.0.ff.net.0.proj": "transformer_blocks.0.act_fn.proj",78            "transformer_blocks.0.ff.net.2": "transformer_blocks.0.ff",79            "transformer_blocks.0.ff_norm": "transformer_blocks.0.norm3",80            "proj_out": "proj_out",81        }82        name_list = sorted([i for i in state_dict if i.startswith("down_blocks.")])83        name_list += sorted([i for i in state_dict if i.startswith("mid_block.")])84        name_list += sorted([i for i in state_dict if i.startswith("up_blocks.")])85        state_dict_ = {}86        last_prefix, module_id = "", -187        for name in name_list:88            names = name.split(".")89            prefix_index = names.index("temporal_transformer") + 190            prefix = ".".join(names[:prefix_index])91            if prefix != last_prefix:92                last_prefix = prefix93                module_id += 194            middle_name = ".".join(names[prefix_index:-1])95            suffix = names[-1]96            if "pos_encoder" in names:97                rename = ".".join(["motion_modules", str(module_id), rename_dict[middle_name]])98            else:99                rename = ".".join(["motion_modules", str(module_id), rename_dict[middle_name], suffix])100            state_dict_[rename] = state_dict[name]101        return state_dict_102    103    def from_civitai(self, state_dict):104        return self.from_diffusers(state_dict)105