Skip to content

Commit 305046c

Browse files
cursoragentsrlynch1
andcommitted
Sync # Copied from blocks after check_copies ruff invocation change
Regenerate copied code blocks so check_copies passes with python -m ruff formatting. No intentional behavior changes. Co-authored-by: Simon Lynch <srlynch1@users.noreply.github.com>
1 parent 4947ac8 commit 305046c

58 files changed

Lines changed: 4 additions & 5788 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎src/diffusers/loaders/lora_pipeline.py‎

Lines changed: 4 additions & 663 deletions
Large diffs are not rendered by default.

‎src/diffusers/models/autoencoders/autoencoder_cosmos3_audio.py‎

Lines changed: 0 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -237,40 +237,6 @@ def forward(self, audio: torch.Tensor) -> torch.Tensor:
237237

238238
# Copied from diffusers.models.autoencoders.autoencoder_oobleck.OobleckResidualUnit with Oobleck->Cosmos3Audio
239239
class Cosmos3AudioResidualUnit(nn.Module):
240-
"""
241-
A residual unit composed of Snake1d and weight-normalized Conv1d layers with dilations.
242-
"""
243-
244-
def __init__(self, dimension: int = 16, dilation: int = 1):
245-
super().__init__()
246-
pad = ((7 - 1) * dilation) // 2
247-
248-
self.snake1 = Snake1d(dimension)
249-
self.conv1 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=7, dilation=dilation, padding=pad))
250-
self.snake2 = Snake1d(dimension)
251-
self.conv2 = weight_norm(nn.Conv1d(dimension, dimension, kernel_size=1))
252-
253-
def forward(self, hidden_state):
254-
"""
255-
Forward pass through the residual unit.
256-
257-
Args:
258-
hidden_state (`torch.Tensor` of shape `(batch_size, channels, time_steps)`):
259-
Input tensor .
260-
261-
Returns:
262-
output_tensor (`torch.Tensor` of shape `(batch_size, channels, time_steps)`)
263-
Input tensor after passing through the residual unit.
264-
"""
265-
output_tensor = hidden_state
266-
output_tensor = self.conv1(self.snake1(output_tensor))
267-
output_tensor = self.conv2(self.snake2(output_tensor))
268-
269-
padding = (hidden_state.shape[-1] - output_tensor.shape[-1]) // 2
270-
if padding > 0:
271-
hidden_state = hidden_state[..., padding:-padding]
272-
output_tensor = hidden_state + output_tensor
273-
return output_tensor
274240

275241

276242
"""

‎src/diffusers/models/controlnets/controlnet_sd3.py‎

Lines changed: 0 additions & 244 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)