(
cls,
transformer,
num_layers: int = 4,
num_single_layers: int = 10,
attention_head_dim: int = 128,
num_attention_heads: int = 24,
extra_condition_channels: int = 0,
load_weights_from_transformer=True,
)
| 181 | |
| 182 | @classmethod |
| 183 | def from_transformer( |
| 184 | cls, |
| 185 | transformer, |
| 186 | num_layers: int = 4, |
| 187 | num_single_layers: int = 10, |
| 188 | attention_head_dim: int = 128, |
| 189 | num_attention_heads: int = 24, |
| 190 | extra_condition_channels: int = 0, |
| 191 | load_weights_from_transformer=True, |
| 192 | ): |
| 193 | config = transformer.config |
| 194 | config["num_layers"] = num_layers |
| 195 | config["num_single_layers"] = num_single_layers |
| 196 | config["attention_head_dim"] = attention_head_dim |
| 197 | config["num_attention_heads"] = num_attention_heads |
| 198 | config["extra_condition_channels"] = extra_condition_channels |
| 199 | |
| 200 | controlnet = cls(**config) |
| 201 | |
| 202 | if load_weights_from_transformer: |
| 203 | controlnet.pos_embed.load_state_dict(transformer.pos_embed.state_dict()) |
| 204 | controlnet.time_text_embed.load_state_dict(transformer.time_text_embed.state_dict()) |
| 205 | controlnet.context_embedder.load_state_dict(transformer.context_embedder.state_dict()) |
| 206 | controlnet.x_embedder.load_state_dict(transformer.x_embedder.state_dict()) |
| 207 | controlnet.transformer_blocks.load_state_dict(transformer.transformer_blocks.state_dict(), strict=False) |
| 208 | controlnet.single_transformer_blocks.load_state_dict( |
| 209 | transformer.single_transformer_blocks.state_dict(), strict=False |
| 210 | ) |
| 211 | |
| 212 | controlnet.controlnet_x_embedder = zero_module(controlnet.controlnet_x_embedder) |
| 213 | |
| 214 | return controlnet |
| 215 | |
| 216 | def forward( |
| 217 | self, |
nothing calls this directly
no outgoing calls
no test coverage detected