MCPcopy Create free account
hub / github.com/YesianRohn/TextSSR / convert_adapter

Function convert_adapter

diffusers/scripts/convert_original_t2i_adapter.py:26–100  ·  view source on GitHub ↗
(src_state, in_channels)

Source from the content-addressed store, hash-verified

24
25
26def convert_adapter(src_state, in_channels):
27 original_body_length = max([int(x.split(".")[1]) for x in src_state.keys() if "body." in x]) + 1
28
29 assert original_body_length == 8
30
31 # (0, 1) -> channels 1
32 assert src_state["body.0.block1.weight"].shape == (320, 320, 3, 3)
33
34 # (2, 3) -> channels 2
35 assert src_state["body.2.in_conv.weight"].shape == (640, 320, 1, 1)
36
37 # (4, 5) -> channels 3
38 assert src_state["body.4.in_conv.weight"].shape == (1280, 640, 1, 1)
39
40 # (6, 7) -> channels 4
41 assert src_state["body.6.block1.weight"].shape == (1280, 1280, 3, 3)
42
43 res_state = {
44 "adapter.conv_in.weight": src_state.pop("conv_in.weight"),
45 "adapter.conv_in.bias": src_state.pop("conv_in.bias"),
46 # 0.resnets.0
47 "adapter.body.0.resnets.0.block1.weight": src_state.pop("body.0.block1.weight"),
48 "adapter.body.0.resnets.0.block1.bias": src_state.pop("body.0.block1.bias"),
49 "adapter.body.0.resnets.0.block2.weight": src_state.pop("body.0.block2.weight"),
50 "adapter.body.0.resnets.0.block2.bias": src_state.pop("body.0.block2.bias"),
51 # 0.resnets.1
52 "adapter.body.0.resnets.1.block1.weight": src_state.pop("body.1.block1.weight"),
53 "adapter.body.0.resnets.1.block1.bias": src_state.pop("body.1.block1.bias"),
54 "adapter.body.0.resnets.1.block2.weight": src_state.pop("body.1.block2.weight"),
55 "adapter.body.0.resnets.1.block2.bias": src_state.pop("body.1.block2.bias"),
56 # 1
57 "adapter.body.1.in_conv.weight": src_state.pop("body.2.in_conv.weight"),
58 "adapter.body.1.in_conv.bias": src_state.pop("body.2.in_conv.bias"),
59 # 1.resnets.0
60 "adapter.body.1.resnets.0.block1.weight": src_state.pop("body.2.block1.weight"),
61 "adapter.body.1.resnets.0.block1.bias": src_state.pop("body.2.block1.bias"),
62 "adapter.body.1.resnets.0.block2.weight": src_state.pop("body.2.block2.weight"),
63 "adapter.body.1.resnets.0.block2.bias": src_state.pop("body.2.block2.bias"),
64 # 1.resnets.1
65 "adapter.body.1.resnets.1.block1.weight": src_state.pop("body.3.block1.weight"),
66 "adapter.body.1.resnets.1.block1.bias": src_state.pop("body.3.block1.bias"),
67 "adapter.body.1.resnets.1.block2.weight": src_state.pop("body.3.block2.weight"),
68 "adapter.body.1.resnets.1.block2.bias": src_state.pop("body.3.block2.bias"),
69 # 2
70 "adapter.body.2.in_conv.weight": src_state.pop("body.4.in_conv.weight"),
71 "adapter.body.2.in_conv.bias": src_state.pop("body.4.in_conv.bias"),
72 # 2.resnets.0
73 "adapter.body.2.resnets.0.block1.weight": src_state.pop("body.4.block1.weight"),
74 "adapter.body.2.resnets.0.block1.bias": src_state.pop("body.4.block1.bias"),
75 "adapter.body.2.resnets.0.block2.weight": src_state.pop("body.4.block2.weight"),
76 "adapter.body.2.resnets.0.block2.bias": src_state.pop("body.4.block2.bias"),
77 # 2.resnets.1
78 "adapter.body.2.resnets.1.block1.weight": src_state.pop("body.5.block1.weight"),
79 "adapter.body.2.resnets.1.block1.bias": src_state.pop("body.5.block1.bias"),
80 "adapter.body.2.resnets.1.block2.weight": src_state.pop("body.5.block2.weight"),
81 "adapter.body.2.resnets.1.block2.bias": src_state.pop("body.5.block2.bias"),
82 # 3.resnets.0
83 "adapter.body.3.resnets.0.block1.weight": src_state.pop("body.6.block1.weight"),

Callers 1

Calls 3

T2IAdapterClass · 0.90
load_state_dictMethod · 0.80
popMethod · 0.45

Tested by

no test coverage detected