@@ -206,247 +206,3 @@ def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int
206206
207207 # Copied from diffusers.models.transformers.transformer_sd3.SD3Transformer2DModel.fuse_qkv_projections
208208 def fuse_qkv_projections (self ):
209- """
210- Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query, key, value)
211- are fused. For cross-attention modules, key and value projection matrices are fused.
212-
213- > [!WARNING] > This API is 🧪 experimental.
214- """
215- self .original_attn_processors = None
216-
217- for _ , attn_processor in self .attn_processors .items ():
218- if "Added" in str (attn_processor .__class__ .__name__ ):
219- raise ValueError ("`fuse_qkv_projections()` is not supported for models having added KV projections." )
220-
221- self .original_attn_processors = self .attn_processors
222-
223- for module in self .modules ():
224- if isinstance (module , Attention ):
225- module .fuse_projections (fuse = True )
226-
227- self .set_attn_processor (FusedJointAttnProcessor2_0 ())
228-
229- # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections
230- def unfuse_qkv_projections (self ):
231- """Disables the fused QKV projection if enabled.
232-
233- > [!WARNING] > This API is 🧪 experimental.
234-
235- """
236- if self .original_attn_processors is not None :
237- self .set_attn_processor (self .original_attn_processors )
238-
239- # Notes: This is for SD3.5 8b controlnet, which shares the pos_embed with the transformer
240- # we should have handled this in conversion script
241- def _get_pos_embed_from_transformer (self , transformer ):
242- pos_embed = PatchEmbed (
243- height = transformer .config .sample_size ,
244- width = transformer .config .sample_size ,
245- patch_size = transformer .config .patch_size ,
246- in_channels = transformer .config .in_channels ,
247- embed_dim = transformer .inner_dim ,
248- pos_embed_max_size = transformer .config .pos_embed_max_size ,
249- )
250- pos_embed .load_state_dict (transformer .pos_embed .state_dict (), strict = True )
251- return pos_embed
252-
253- @classmethod
254- def from_transformer (
255- cls , transformer , num_layers = 12 , num_extra_conditioning_channels = 1 , load_weights_from_transformer = True
256- ):
257- config = transformer .config
258- config ["num_layers" ] = num_layers or config .num_layers
259- config ["extra_conditioning_channels" ] = num_extra_conditioning_channels
260- controlnet = cls .from_config (config )
261-
262- if load_weights_from_transformer :
263- controlnet .pos_embed .load_state_dict (transformer .pos_embed .state_dict ())
264- controlnet .time_text_embed .load_state_dict (transformer .time_text_embed .state_dict ())
265- controlnet .context_embedder .load_state_dict (transformer .context_embedder .state_dict ())
266- controlnet .transformer_blocks .load_state_dict (transformer .transformer_blocks .state_dict (), strict = False )
267-
268- controlnet .pos_embed_input = zero_module (controlnet .pos_embed_input )
269-
270- return controlnet
271-
272- @apply_lora_scale ("joint_attention_kwargs" )
273- def forward (
274- self ,
275- hidden_states : torch .Tensor ,
276- controlnet_cond : torch .Tensor ,
277- conditioning_scale : float = 1.0 ,
278- encoder_hidden_states : torch .Tensor = None ,
279- pooled_projections : torch .Tensor = None ,
280- timestep : torch .LongTensor = None ,
281- joint_attention_kwargs : dict [str , Any ] | None = None ,
282- return_dict : bool = True ,
283- ) -> torch .Tensor | Transformer2DModelOutput :
284- """
285- The [`SD3Transformer2DModel`] forward method.
286-
287- Args:
288- hidden_states (`torch.Tensor` of shape `(batch size, channel, height, width)`):
289- Input `hidden_states`.
290- controlnet_cond (`torch.Tensor`):
291- The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`.
292- conditioning_scale (`float`, defaults to `1.0`):
293- The scale factor for ControlNet outputs.
294- encoder_hidden_states (`torch.Tensor` of shape `(batch size, sequence_len, embed_dims)`):
295- Conditional embeddings (embeddings computed from the input conditions such as prompts) to use.
296- pooled_projections (`torch.Tensor` of shape `(batch_size, projection_dim)`): Embeddings projected
297- from the embeddings of input conditions.
298- timestep ( `torch.LongTensor`):
299- Used to indicate denoising step.
300- joint_attention_kwargs (`dict`, *optional*):
301- A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
302- `self.processor` in
303- [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
304- return_dict (`bool`, *optional*, defaults to `True`):
305- Whether or not to return a [`~models.transformer_2d.Transformer2DModelOutput`] instead of a plain
306- tuple.
307-
308- Returns:
309- If `return_dict` is True, an [`~models.transformer_2d.Transformer2DModelOutput`] is returned, otherwise a
310- `tuple` where the first element is the sample tensor.
311- """
312- if self .pos_embed is not None and hidden_states .ndim != 4 :
313- raise ValueError ("hidden_states must be 4D when pos_embed is used" )
314-
315- # SD3.5 8b controlnet does not have a `pos_embed`,
316- # it use the `pos_embed` from the transformer to process input before passing to controlnet
317- elif self .pos_embed is None and hidden_states .ndim != 3 :
318- raise ValueError ("hidden_states must be 3D when pos_embed is not used" )
319-
320- if self .context_embedder is not None and encoder_hidden_states is None :
321- raise ValueError ("encoder_hidden_states must be provided when context_embedder is used" )
322- # SD3.5 8b controlnet does not have a `context_embedder`, it does not use `encoder_hidden_states`
323- elif self .context_embedder is None and encoder_hidden_states is not None :
324- raise ValueError ("encoder_hidden_states should not be provided when context_embedder is not used" )
325-
326- if self .pos_embed is not None :
327- hidden_states = self .pos_embed (hidden_states ) # takes care of adding positional embeddings too.
328-
329- temb = self .time_text_embed (timestep , pooled_projections )
330-
331- if self .context_embedder is not None :
332- encoder_hidden_states = self .context_embedder (encoder_hidden_states )
333-
334- # add
335- hidden_states = hidden_states + self .pos_embed_input (controlnet_cond )
336-
337- block_res_samples = ()
338-
339- for block in self .transformer_blocks :
340- if torch .is_grad_enabled () and self .gradient_checkpointing :
341- if self .context_embedder is not None :
342- encoder_hidden_states , hidden_states = self ._gradient_checkpointing_func (
343- block ,
344- hidden_states ,
345- encoder_hidden_states ,
346- temb ,
347- )
348- else :
349- # SD3.5 8b controlnet use single transformer block, which does not use `encoder_hidden_states`
350- hidden_states = self ._gradient_checkpointing_func (block , hidden_states , temb )
351-
352- else :
353- if self .context_embedder is not None :
354- encoder_hidden_states , hidden_states = block (
355- hidden_states = hidden_states , encoder_hidden_states = encoder_hidden_states , temb = temb
356- )
357- else :
358- # SD3.5 8b controlnet use single transformer block, which does not use `encoder_hidden_states`
359- hidden_states = block (hidden_states , temb )
360-
361- block_res_samples = block_res_samples + (hidden_states ,)
362-
363- controlnet_block_res_samples = ()
364- for block_res_sample , controlnet_block in zip (block_res_samples , self .controlnet_blocks ):
365- block_res_sample = controlnet_block (block_res_sample )
366- controlnet_block_res_samples = controlnet_block_res_samples + (block_res_sample ,)
367-
368- # 6. scaling
369- controlnet_block_res_samples = [sample * conditioning_scale for sample in controlnet_block_res_samples ]
370-
371- if not return_dict :
372- return (controlnet_block_res_samples ,)
373-
374- return SD3ControlNetOutput (controlnet_block_samples = controlnet_block_res_samples )
375-
376-
377- class SD3MultiControlNetModel (ModelMixin ):
378- r"""
379- `SD3ControlNetModel` wrapper class for Multi-SD3ControlNet
380-
381- This module is a wrapper for multiple instances of the `SD3ControlNetModel`. The `forward()` API is designed to be
382- compatible with `SD3ControlNetModel`.
383-
384- Args:
385- controlnets (`list[SD3ControlNetModel]`):
386- Provides additional conditioning to the unet during the denoising process. You must set multiple
387- `SD3ControlNetModel` as a list.
388- """
389-
390- def __init__ (self , controlnets ):
391- super ().__init__ ()
392- self .nets = nn .ModuleList (controlnets )
393-
394- def forward (
395- self ,
396- hidden_states : torch .Tensor ,
397- controlnet_cond : list [torch .tensor ],
398- conditioning_scale : list [float ],
399- pooled_projections : torch .Tensor ,
400- encoder_hidden_states : torch .Tensor = None ,
401- timestep : torch .LongTensor = None ,
402- joint_attention_kwargs : dict [str , Any ] | None = None ,
403- return_dict : bool = True ,
404- ) -> SD3ControlNetOutput | tuple :
405- r"""
406- Args:
407- hidden_states (`torch.Tensor`):
408- Input `hidden_states`.
409- controlnet_cond (`list` of `torch.Tensor`):
410- A list of conditional input tensors, one per ControlNet.
411- conditioning_scale (`list` of `float`):
412- A list of scale factors applied to the ControlNet outputs.
413- pooled_projections (`torch.Tensor`):
414- Embeddings projected from the embeddings of input conditions.
415- encoder_hidden_states (`torch.Tensor`, *optional*):
416- Conditional embeddings (embeddings computed from the input conditions such as prompts) to use.
417- timestep (`torch.LongTensor`, *optional*):
418- Used to indicate denoising step.
419- joint_attention_kwargs (`dict`, *optional*):
420- A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
421- `self.processor` in
422- [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
423- return_dict (`bool`, *optional*, defaults to `True`):
424- Whether or not to return a [`SD3ControlNetOutput`] instead of a plain tuple.
425-
426- Returns:
427- [`SD3ControlNetOutput`] or `tuple`:
428- If `return_dict` is True, a [`SD3ControlNetOutput`] is returned, otherwise a plain `tuple` is returned.
429- """
430- for i , (image , scale , controlnet ) in enumerate (zip (controlnet_cond , conditioning_scale , self .nets )):
431- block_samples = controlnet (
432- hidden_states = hidden_states ,
433- timestep = timestep ,
434- encoder_hidden_states = encoder_hidden_states ,
435- pooled_projections = pooled_projections ,
436- controlnet_cond = image ,
437- conditioning_scale = scale ,
438- joint_attention_kwargs = joint_attention_kwargs ,
439- return_dict = return_dict ,
440- )
441-
442- # merge samples
443- if i == 0 :
444- control_block_samples = block_samples
445- else :
446- control_block_samples = [
447- control_block_sample + block_sample
448- for control_block_sample , block_sample in zip (control_block_samples [0 ], block_samples [0 ])
449- ]
450- control_block_samples = (tuple (control_block_samples ),)
451-
452- return control_block_samples
0 commit comments