Spaces:
Paused
Paused
| # Copyright 2026 The HuggingFace Team. All rights reserved. | |
| # | |
| # Licensed under the Apache License, Version 2.0 (the "License"); | |
| # you may not use this file except in compliance with the License. | |
| # You may obtain a copy of the License at | |
| # | |
| # http://www.apache.org/licenses/LICENSE-2.0 | |
| # | |
| # Unless required by applicable law or agreed to in writing, software | |
| # distributed under the License is distributed on an "AS IS" BASIS, | |
| # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | |
| # See the License for the specific language governing permissions and | |
| # limitations under the License. | |
| import torch | |
| from ...configuration_utils import FrozenDict | |
| from ...guiders import ClassifierFreeGuidance | |
| from ...models.transformers import SD3Transformer2DModel | |
| from ...schedulers import FlowMatchEulerDiscreteScheduler | |
| from ...utils import logging | |
| from ..modular_pipeline import ( | |
| BlockState, | |
| LoopSequentialPipelineBlocks, | |
| ModularPipelineBlocks, | |
| PipelineState, | |
| ) | |
| from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam | |
| from .modular_pipeline import StableDiffusion3ModularPipeline | |
| logger = logging.get_logger(__name__) | |
| class StableDiffusion3LoopDenoiser(ModularPipelineBlocks): | |
| model_name = "stable-diffusion-3" | |
| def expected_components(self) -> list[ComponentSpec]: | |
| return [ | |
| ComponentSpec( | |
| "guider", | |
| ClassifierFreeGuidance, | |
| config=FrozenDict({"guidance_scale": 7.0}), | |
| default_creation_method="from_config", | |
| ), | |
| ComponentSpec("transformer", SD3Transformer2DModel), | |
| ] | |
| def description(self) -> str: | |
| return "Step within the denoising loop that denoises the latents." | |
| def inputs(self) -> list[InputParam]: | |
| return [ | |
| InputParam( | |
| "joint_attention_kwargs", | |
| type_hint=dict, | |
| description="A kwargs dictionary passed along to the AttentionProcessor.", | |
| ), | |
| InputParam( | |
| "latents", | |
| required=True, | |
| type_hint=torch.Tensor, | |
| description="The initial latents to use for the denoising process.", | |
| ), | |
| InputParam( | |
| "prompt_embeds", | |
| required=True, | |
| type_hint=torch.Tensor, | |
| description="Text embeddings for guidance.", | |
| ), | |
| InputParam( | |
| "pooled_prompt_embeds", | |
| required=True, | |
| type_hint=torch.Tensor, | |
| description="Pooled text embeddings for guidance.", | |
| ), | |
| InputParam( | |
| "negative_prompt_embeds", | |
| type_hint=torch.Tensor, | |
| description="Negative text embeddings for guidance.", | |
| ), | |
| InputParam( | |
| "negative_pooled_prompt_embeds", | |
| type_hint=torch.Tensor, | |
| description="Negative pooled text embeddings for guidance.", | |
| ), | |
| InputParam( | |
| "num_inference_steps", | |
| type_hint=int, | |
| description="The number of denoising steps.", | |
| ), | |
| ] | |
| def __call__( | |
| self, | |
| components: StableDiffusion3ModularPipeline, | |
| block_state: BlockState, | |
| i: int, | |
| t: torch.Tensor, | |
| ) -> PipelineState: | |
| do_cfg = block_state.negative_prompt_embeds is not None | |
| guider_inputs = { | |
| "hidden_states": (block_state.latents, block_state.latents) if do_cfg else block_state.latents, | |
| "encoder_hidden_states": ( | |
| block_state.prompt_embeds, | |
| block_state.negative_prompt_embeds, | |
| ) | |
| if do_cfg | |
| else block_state.prompt_embeds, | |
| "text_embeds": ( | |
| block_state.pooled_prompt_embeds, | |
| block_state.negative_pooled_prompt_embeds, | |
| ) | |
| if do_cfg | |
| else block_state.pooled_prompt_embeds, | |
| } | |
| components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) | |
| guider_state = components.guider.prepare_inputs(guider_inputs) | |
| for guider_state_batch in guider_state: | |
| components.guider.prepare_models(components.transformer) | |
| latent_model_input = guider_state_batch.hidden_states | |
| prompt_embeds = guider_state_batch.encoder_hidden_states | |
| pooled_projections = getattr(guider_state_batch, "text_embeds", None) | |
| timestep = t.expand(latent_model_input.shape[0]) | |
| guider_state_batch.noise_pred = components.transformer( | |
| hidden_states=latent_model_input, | |
| timestep=timestep, | |
| encoder_hidden_states=prompt_embeds, | |
| pooled_projections=pooled_projections, | |
| joint_attention_kwargs=block_state.joint_attention_kwargs, | |
| return_dict=False, | |
| )[0] | |
| components.guider.cleanup_models(components.transformer) | |
| guider_output = components.guider(guider_state) | |
| block_state.noise_pred = guider_output.pred | |
| return components, block_state | |
| class StableDiffusion3LoopAfterDenoiser(ModularPipelineBlocks): | |
| model_name = "stable-diffusion-3" | |
| def expected_components(self) -> list[ComponentSpec]: | |
| return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] | |
| def intermediate_outputs(self) -> list[OutputParam]: | |
| return [ | |
| OutputParam( | |
| "latents", | |
| type_hint=torch.Tensor, | |
| description="The denoised latent tensors.", | |
| ) | |
| ] | |
| def __call__( | |
| self, | |
| components: StableDiffusion3ModularPipeline, | |
| block_state: BlockState, | |
| i: int, | |
| t: torch.Tensor, | |
| ): | |
| latents_dtype = block_state.latents.dtype | |
| block_state.latents = components.scheduler.step( | |
| block_state.noise_pred, | |
| t, | |
| block_state.latents, | |
| return_dict=False, | |
| )[0] | |
| if block_state.latents.dtype != latents_dtype: | |
| block_state.latents = block_state.latents.to(latents_dtype) | |
| return components, block_state | |
| class StableDiffusion3DenoiseLoopWrapper(LoopSequentialPipelineBlocks): | |
| model_name = "stable-diffusion-3" | |
| def loop_expected_components(self) -> list[ComponentSpec]: | |
| return [ | |
| ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), | |
| ComponentSpec("transformer", SD3Transformer2DModel), | |
| ] | |
| def loop_inputs(self) -> list[InputParam]: | |
| return [ | |
| InputParam("timesteps", required=True, type_hint=torch.Tensor), | |
| InputParam("num_inference_steps", required=True, type_hint=int), | |
| ] | |
| def __call__(self, components: StableDiffusion3ModularPipeline, state: PipelineState) -> PipelineState: | |
| block_state = self.get_block_state(state) | |
| block_state.num_warmup_steps = max( | |
| len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, | |
| 0, | |
| ) | |
| with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: | |
| for i, t in enumerate(block_state.timesteps): | |
| components, block_state = self.loop_step(components, block_state, i=i, t=t) | |
| if i == len(block_state.timesteps) - 1 or ( | |
| (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 | |
| ): | |
| progress_bar.update() | |
| self.set_block_state(state, block_state) | |
| return components, state | |
| class StableDiffusion3DenoiseStep(StableDiffusion3DenoiseLoopWrapper): | |
| block_classes = [StableDiffusion3LoopDenoiser, StableDiffusion3LoopAfterDenoiser] | |
| block_names = ["denoiser", "after_denoiser"] | |