diffusers-internal-dev/chronoedit-modular
02
1# Copyright (c) 2025 Wan and Hugging Face Teams. 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.configuration_utils import FrozenDict17from diffusers.modular_pipelines import ModularPipelineBlocks, ComponentSpec, InputParam, OutputParam, PipelineState18from typing import List, Union19from diffusers import AutoencoderKLWan20from diffusers.video_processor import VideoProcessor21import torch22import PIL23import numpy as np24 25 26class ChronoEditDecodeStep(ModularPipelineBlocks):27 model_name = "chronoedit"28 29 @property30 def expected_components(self) -> List[ComponentSpec]:31 return [32 ComponentSpec("vae", AutoencoderKLWan),33 ComponentSpec(34 "video_processor",35 VideoProcessor,36 config=FrozenDict({"vae_scale_factor": 8}),37 default_creation_method="from_config",38 ),39 ]40 41 @property42 def description(self) -> str:43 return "Step that decodes the denoised latents into images"44 45 @property46 def inputs(self) -> List[InputParam]:47 return [48 InputParam(49 "latents",50 required=True,51 type_hint=torch.Tensor,52 description="The denoised latents from the denoising step",53 ),54 InputParam("output_type", default="pil"),55 ]56 57 @property58 def intermediate_outputs(self) -> List[str]:59 return [60 OutputParam(61 "videos",62 type_hint=Union[List[List[PIL.Image.Image]], List[torch.Tensor], List[np.ndarray]],63 description="The generated videos, can be a PIL.Image.Image, torch.Tensor or a numpy array",64 )65 ]66 67 @torch.no_grad()68 def __call__(self, components, state: PipelineState) -> PipelineState:69 block_state = self.get_block_state(state)70 vae_dtype = components.vae.dtype71 72 if not block_state.output_type == "latent":73 latents = block_state.latents74 latents_mean = (75 torch.tensor(components.vae.config.latents_mean)76 .view(1, components.vae.config.z_dim, 1, 1, 1)77 .to(latents.device, latents.dtype)78 )79 latents_std = 1.0 / torch.tensor(components.vae.config.latents_std).view(80 1, components.vae.config.z_dim, 1, 1, 181 ).to(latents.device, latents.dtype)82 latents = latents / latents_std + latents_mean83 latents = latents.to(vae_dtype)84 block_state.videos = components.vae.decode(latents, return_dict=False)[0]85 else:86 block_state.videos = block_state.latents87 88 block_state.videos = components.video_processor.postprocess_video(89 block_state.videos, output_type=block_state.output_type90 )91 92 self.set_block_state(state, block_state)93 94 return components, state95 