Team Ai
Modelpublic

diffusers/matrix-game-2-modular

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes14downloads
decoders.py101 linesDownload Raw Back to root
1# Copyright 2025 The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7#     http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14 15from typing import Any, List, Tuple, Union16 17import numpy as np18import PIL19import torch20 21from diffusers.configuration_utils import FrozenDict22from diffusers.models import AutoencoderKLWan23from diffusers.utils import logging24from diffusers.video_processor import VideoProcessor25from diffusers.modular_pipelines import ModularPipelineBlocks, PipelineState26from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, InputParam, OutputParam27 28 29logger = logging.get_logger(__name__)  # pylint: disable=invalid-name30 31 32class MatrixGameWanDecodeStep(ModularPipelineBlocks):33    model_name = "MatrixGameWan"34 35    @property36    def expected_components(self) -> List[ComponentSpec]:37        return [38            ComponentSpec("vae", AutoencoderKLWan, repo="Wan-AI/Wan2.1-T2V-1.3B-Diffusers", subfolder="vae"),39            ComponentSpec(40                "video_processor",41                VideoProcessor,42                config=FrozenDict({"vae_scale_factor": 8}),43                default_creation_method="from_config",44            ),45        ]46 47    @property48    def description(self) -> str:49        return "Step that decodes the denoised latents into images"50 51    @property52    def inputs(self) -> List[Tuple[str, Any]]:53        return [54            InputParam("output_type", default="pil"),55            InputParam(56                "latents",57                required=True,58                type_hint=torch.Tensor,59                description="The denoised latents from the denoising step",60            )61        ]62 63    @property64    def intermediate_outputs(self) -> List[str]:65        return [66            OutputParam(67                "videos",68                type_hint=Union[List[List[PIL.Image.Image]], List[torch.Tensor], List[np.ndarray]],69                description="The generated videos, can be a PIL.Image.Image, torch.Tensor or a numpy array",70            )71        ]72 73    @torch.no_grad()74    def __call__(self, components, state: PipelineState) -> PipelineState:75        block_state = self.get_block_state(state)76        vae_dtype = components.vae.dtype77 78        if not block_state.output_type == "latent":79            latents = block_state.latents80            latents_mean = (81                torch.tensor(components.vae.config.latents_mean)82                .view(1, components.vae.config.z_dim, 1, 1, 1)83                .to(latents.device, latents.dtype)84            )85            latents_std = 1.0 / torch.tensor(components.vae.config.latents_std).view(86                1, components.vae.config.z_dim, 1, 1, 187            ).to(latents.device, latents.dtype)88            latents = latents / latents_std + latents_mean89            latents = latents.to(vae_dtype)90            block_state.videos = components.vae.decode(latents, return_dict=False)[0]91        else:92            block_state.videos = block_state.latents93 94        block_state.videos = components.video_processor.postprocess_video(95            block_state.videos, output_type=block_state.output_type96        )97 98        self.set_block_state(state, block_state)99 100        return components, state101