Team Ai
Modelpublic

diffusers-internal-dev/chronoedit-modular

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes2downloads
decoders.py95 linesDownload Raw Back to root
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