Team Ai
Modelpublic

diffusers-internal-dev/chronoedit-modular

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes2downloads
before_denoise.py213 linesDownload Raw Back to root
1# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.2# SPDX-License-Identifier: Apache-2.03#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15 16from diffusers.modular_pipelines import (17    ModularPipelineBlocks,18    ComponentSpec,19    PipelineState,20    ModularPipeline,21    OutputParam,22    InputParam,23)24from diffusers.modular_pipelines.wan.before_denoise import retrieve_timesteps25from typing import Optional, List, Union, Tuple26from diffusers.image_processor import PipelineImageInput27from diffusers.utils.torch_utils import randn_tensor28import torch29from diffusers import AutoencoderKLWan, UniPCMultistepScheduler30 31# One needs Wan anyway to run ChronoEdit (`AutoencoderKLWan`).32from diffusers.pipelines.wan.pipeline_wan_i2v import retrieve_latents33 34 35class ChronoEditSetTimestepsStep(ModularPipelineBlocks):36    model_name = "chronoedit"37 38    @property39    def expected_components(self) -> List[ComponentSpec]:40        return [ComponentSpec("scheduler", UniPCMultistepScheduler)]41 42    @property43    def inputs(self) -> List[InputParam]:44        return [InputParam("num_inference_steps", default=50), InputParam("timesteps"), InputParam("sigmas")]45 46    @property47    def intermediate_outputs(self) -> List[OutputParam]:48        return [49            OutputParam("timesteps", type_hint=torch.Tensor, description="The timesteps to use for inference"),50            OutputParam(51                "num_inference_steps",52                type_hint=int,53                description="The number of denoising steps to perform at inference time",54            ),55        ]56 57    @torch.no_grad()58    def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:59        block_state = self.get_block_state(state)60        block_state.device = components._execution_device61 62        block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps(63            components.scheduler,64            block_state.num_inference_steps,65            block_state.device,66            block_state.timesteps,67            block_state.sigmas,68        )69 70        self.set_block_state(state, block_state)71        return components, state72 73 74class ChronoEditPrepareLatentStep(ModularPipelineBlocks):75    model_name = "chronoedit"76 77    @property78    def expected_components(self) -> List[ComponentSpec]:79        return [ComponentSpec("vae", AutoencoderKLWan)]80 81    @property82    def inputs(self) -> List[InputParam]:83        return [84            InputParam("processed_image", type_hint=PipelineImageInput),85            InputParam("image_embeds", type_hint=torch.Tensor),86            InputParam("height", type_hint=int, default=480),87            InputParam("width", type_hint=int, default=832),88            InputParam("num_frames", type_hint=int, default=81),89            InputParam("batch_size"),90            InputParam("num_videos_per_prompt", type_hint=int, default=1),91            InputParam("latents", type_hint=Optional[torch.Tensor]),92            InputParam("generator"),93        ]94 95    @property96    def intermediate_outputs(self) -> List[OutputParam]:97        return [98            OutputParam(99                "latents",100                type_hint=torch.Tensor,101                description="The initial latents to use for the denoising process.",102            ),103            OutputParam(104                "condition",105                type_hint=torch.Tensor,106                description="Conditioning latents for the denoising process.",107            ),108        ]109 110    @staticmethod111    def check_inputs(height, width):112        if height % 16 != 0 or width % 16 != 0:113            raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.")114 115    @staticmethod116    def prepare_latents(117        components,118        image: PipelineImageInput,119        batch_size: int,120        num_channels_latents: int = 16,121        height: int = 480,122        width: int = 832,123        num_frames: int = 81,124        dtype: Optional[torch.dtype] = None,125        device: Optional[torch.device] = None,126        generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,127        latents: Optional[torch.Tensor] = None,128    ) -> Tuple[torch.Tensor, torch.Tensor]:129        num_latent_frames = (num_frames - 1) // components.vae_scale_factor_temporal + 1130        latent_height = height // components.vae_scale_factor_spatial131        latent_width = width // components.vae_scale_factor_spatial132 133        shape = (batch_size, num_channels_latents, num_latent_frames, latent_height, latent_width)134        if isinstance(generator, list) and len(generator) != batch_size:135            raise ValueError(136                f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"137                f" size of {batch_size}. Make sure the batch size matches the length of the generators."138            )139 140        if latents is None:141            latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)142        else:143            latents = latents.to(device=device, dtype=dtype)144 145        image = image.unsqueeze(2)146        video_condition = torch.cat(147            [image, image.new_zeros(image.shape[0], image.shape[1], num_frames - 1, height, width)], dim=2148        )149        video_condition = video_condition.to(device=device, dtype=dtype)150 151        latents_mean = (152            torch.tensor(components.vae.config.latents_mean)153            .view(1, components.vae.config.z_dim, 1, 1, 1)154            .to(latents.device, latents.dtype)155        )156        latents_std = 1.0 / torch.tensor(components.vae.config.latents_std).view(157            1, components.vae.config.z_dim, 1, 1, 1158        ).to(latents.device, latents.dtype)159 160        if isinstance(generator, list):161            latent_condition = [162                retrieve_latents(components.vae.encode(video_condition), sample_mode="argmax") for _ in generator163            ]164            latent_condition = torch.cat(latent_condition)165        else:166            latent_condition = retrieve_latents(components.vae.encode(video_condition), sample_mode="argmax")167            latent_condition = latent_condition.repeat(batch_size, 1, 1, 1, 1)168 169        latent_condition = (latent_condition - latents_mean) * latents_std170 171        mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height, latent_width)172        mask_lat_size[:, :, list(range(1, num_frames))] = 0173        first_frame_mask = mask_lat_size[:, :, 0:1]174        first_frame_mask = torch.repeat_interleave(175            first_frame_mask, dim=2, repeats=components.vae_scale_factor_temporal176        )177        mask_lat_size = torch.concat([first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)178        mask_lat_size = mask_lat_size.view(179            batch_size, -1, components.vae_scale_factor_temporal, latent_height, latent_width180        )181        mask_lat_size = mask_lat_size.transpose(1, 2)182        mask_lat_size = mask_lat_size.to(latent_condition.device)183 184        return latents, torch.concat([mask_lat_size, latent_condition], dim=1)185 186    @torch.no_grad()187    def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:188        block_state = self.get_block_state(state)189 190        self.check_inputs(block_state.height, block_state.width)191 192        block_state.device = components._execution_device193        block_state.num_channels_latents = components.num_channels_latents194 195        batch_size = block_state.batch_size * block_state.num_videos_per_prompt196        block_state.latents, block_state.condition = self.prepare_latents(197            components,198            block_state.processed_image,199            batch_size,200            block_state.num_channels_latents,201            block_state.height,202            block_state.width,203            block_state.num_frames,204            torch.bfloat16,205            block_state.device,206            block_state.generator,207            block_state.latents,208        )209 210        self.set_block_state(state, block_state)211 212        return components, state213