hugging-apps/echo-memory
0
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 