diffusers-internal-dev/chronoedit-modular
02
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 