MCPcopy Create free account
hub / github.com/Shakker-Labs/RepText / from_transformer

Method from_transformer

controlnet_flux.py:183–214  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected