multimodalart's picture
multimodalart HF Staff
Sync the split MiniMax-H3 Spaces (part 2)
cd458ae verified
Raw
History Blame Contribute Delete
8.13 kB
# 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"
@property
def expected_components(self) -> list[ComponentSpec]:
return [
ComponentSpec(
"guider",
ClassifierFreeGuidance,
config=FrozenDict({"guidance_scale": 7.0}),
default_creation_method="from_config",
),
ComponentSpec("transformer", SD3Transformer2DModel),
]
@property
def description(self) -> str:
return "Step within the denoising loop that denoises the latents."
@property
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.",
),
]
@torch.no_grad()
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"
@property
def expected_components(self) -> list[ComponentSpec]:
return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)]
@property
def intermediate_outputs(self) -> list[OutputParam]:
return [
OutputParam(
"latents",
type_hint=torch.Tensor,
description="The denoised latent tensors.",
)
]
@torch.no_grad()
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"
@property
def loop_expected_components(self) -> list[ComponentSpec]:
return [
ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler),
ComponentSpec("transformer", SD3Transformer2DModel),
]
@property
def loop_inputs(self) -> list[InputParam]:
return [
InputParam("timesteps", required=True, type_hint=torch.Tensor),
InputParam("num_inference_steps", required=True, type_hint=int),
]
@torch.no_grad()
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"]