Team Ai
Modelpublic

diffusers/matrix-game-2-modular

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes14downloads
before_denoise.py605 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 15import inspect16from typing import List, Optional, Union, Dict17 18import torch19 20from diffusers import AutoencoderKLWan21from diffusers.configuration_utils import FrozenDict22from diffusers.schedulers import UniPCMultistepScheduler23from diffusers.utils import logging24from diffusers.utils.torch_utils import randn_tensor25from diffusers.video_processor import VideoProcessor26from diffusers.modular_pipelines import ModularPipeline, ModularPipelineBlocks, PipelineState27from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, InputParam, OutputParam28 29logger = logging.get_logger(__name__)  # pylint: disable=invalid-name30 31# Constants32FRAME_MULTIPLE = 433DEFAULT_SAMPLES_PER_ACTION = 434DEFAULT_FRAMES_PER_ACTION = 1235 36DEFAULT_MOUSE_DIM = 237DEFAULT_KEYBOARD_DIM = 438 39# Camera movement configuration40CAMERA_MOVEMENT_VALUE = 0.141CAMERA_VALUE_MAP = {42    "camera_up": [CAMERA_MOVEMENT_VALUE, 0],43    "camera_down": [-CAMERA_MOVEMENT_VALUE, 0],44    "camera_l": [0, -CAMERA_MOVEMENT_VALUE],45    "camera_r": [0, CAMERA_MOVEMENT_VALUE],46    "camera_ur": [CAMERA_MOVEMENT_VALUE, CAMERA_MOVEMENT_VALUE],47    "camera_ul": [CAMERA_MOVEMENT_VALUE, -CAMERA_MOVEMENT_VALUE],48    "camera_dr": [-CAMERA_MOVEMENT_VALUE, CAMERA_MOVEMENT_VALUE],49    "camera_dl": [-CAMERA_MOVEMENT_VALUE, -CAMERA_MOVEMENT_VALUE],50}51 52# Define available actions53MOVEMENT_ACTIONS = ["forward", "left", "right"]54COMPOUND_MOVEMENTS = ["forward_left", "forward_right"]55CAMERA_ACTIONS = list(CAMERA_VALUE_MAP.keys())56 57# Keyboard action indices58KEYBOARD_ACTION_INDICES = {"forward": 0, "back": 1, "left": 2, "right": 3}59 60 61def sync_actions_to_frames(62    actions: List[str],63    num_frames: int,64    min_duration: int = 1265) -> List[Dict[str, Union[str, int]]]:66    """67    Synchronize a list of actions to fit exactly within the given number of frames68    using equal distribution strategy.69 70    Args:71        actions: List of action names to perform72        num_frames: Total frames to fill73        min_duration: Minimum frames per action (should be multiple of frame_multiple)74        frame_multiple: Actions must be multiples of this value75 76    Returns:77        List of action dictionaries with 'type', 'start_frame', and 'duration'78    """79 80    if not actions:81        raise ValueError("No actions provided")82 83    max_possible_actions = num_frames // DEFAULT_FRAMES_PER_ACTION84    if len(actions) > max_possible_actions:85        actions = actions[:max_possible_actions]86 87    num_actions = len(actions)88 89    frames_per_action = num_frames // num_actions90    frames_per_action = (frames_per_action // FRAME_MULTIPLE) * FRAME_MULTIPLE91    frames_per_action = max(DEFAULT_FRAMES_PER_ACTION, frames_per_action)92 93    remaining_frames = num_frames - (frames_per_action * num_actions)94    output = []95    current_frame = 096 97    for i, action in enumerate(actions):98        duration = frames_per_action if i != num_actions - 1 else num_frames - current_frame99 100        output.append({101            "action_type": action,102            "start_frame": current_frame,103            "duration": duration104        })105 106        current_frame += duration107 108    return output109 110 111def actions_to_condition_tensors(actions, num_frames):112    keyboard_conditions = torch.zeros((num_frames, DEFAULT_KEYBOARD_DIM))113    mouse_conditions = torch.zeros((num_frames, DEFAULT_MOUSE_DIM))114 115    for action in actions:116        action_type = action['action_type']117        start_frame = action['start_frame']118        end_frame = start_frame + action['duration']119 120        action_components = action_type.split("_")121        for component in action_components:122            if component in KEYBOARD_ACTION_INDICES:123                action_idx = KEYBOARD_ACTION_INDICES[component]124                keyboard_conditions[start_frame:end_frame, action_idx] = 1.0125 126        if not "camera" in action_type:127            continue128 129        mouse_x = mouse_y = 0.0130        for idx, component in enumerate(action_components):131            if not action_components[idx] == "camera":132                continue133 134            camera_action = f"camera_{action_components[idx+1]}"135            if camera_action not in CAMERA_VALUE_MAP:136                continue137 138            camera_values = CAMERA_VALUE_MAP[camera_action]139            mouse_x += camera_values[0]140            mouse_y += camera_values[1]141 142        mouse_conditions[start_frame:end_frame, 0] = mouse_x143        mouse_conditions[start_frame:end_frame, 1] = mouse_y144 145    return keyboard_conditions, mouse_conditions146 147 148def _build_test_actions(149    movement_actions: List[str],150    compound_movements: List[str],151    camera_actions: List[str],152) -> List[str]:153    """Build comprehensive list of test action combinations.154 155    Args:156        movement_actions: List of basic movement actions157        compound_movements: List of compound movement actions158        camera_actions: List of camera control actions159 160    Returns:161        List of all action combinations to test162    """163    # Create base test actions with repetition for variety164    test_actions = compound_movements * 5 + camera_actions * 5 + movement_actions * 5165 166    # Add combined movement + camera actions167    for movement in movement_actions + compound_movements:168        for camera in camera_actions:169            combined_action = f"{movement}_{camera}"170            test_actions.append(combined_action)171 172    return test_actions173 174 175def generate_random_condition_tensors(num_frames: int) -> Dict[str, torch.Tensor]:176    """Generate benchmark action sequences for testing.177 178    Args:179        num_frames: Total number of frames to generate180        num_samples_per_action: Number of samples per action type181 182    Returns:183        Dictionary containing keyboard and mouse conditions for benchmark actions184    """185    # Build test action combinations186    actions = _build_test_actions(187        MOVEMENT_ACTIONS, COMPOUND_MOVEMENTS, CAMERA_ACTIONS188    )189    actions = sync_actions_to_frames(actions, num_frames)190    return actions_to_condition_tensors(actions, num_frames)191 192 193# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps194def retrieve_timesteps(195    scheduler,196    num_inference_steps: Optional[int] = None,197    device: Optional[Union[str, torch.device]] = None,198    timesteps: Optional[List[int]] = None,199    sigmas: Optional[List[float]] = None,200    **kwargs,201):202    r"""203    Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles204    custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.205 206    Args:207        scheduler (`SchedulerMixin`):208            The scheduler to get timesteps from.209        num_inference_steps (`int`):210            The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`211            must be `None`.212        device (`str` or `torch.device`, *optional*):213            The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.214        timesteps (`List[int]`, *optional*):215            Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,216            `num_inference_steps` and `sigmas` must be `None`.217        sigmas (`List[float]`, *optional*):218            Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,219            `num_inference_steps` and `timesteps` must be `None`.220 221    Returns:222        `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the223        second element is the number of inference steps.224    """225    if timesteps is not None and sigmas is not None:226        raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")227    if timesteps is not None:228        accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())229        if not accepts_timesteps:230            raise ValueError(231                f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"232                f" timestep schedules. Please check whether you are using the correct scheduler."233            )234        scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)235        timesteps = scheduler.timesteps236        num_inference_steps = len(timesteps)237    elif sigmas is not None:238        accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())239        if not accept_sigmas:240            raise ValueError(241                f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"242                f" sigmas schedules. Please check whether you are using the correct scheduler."243            )244        scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)245        timesteps = scheduler.timesteps246        num_inference_steps = len(timesteps)247    else:248        scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)249        timesteps = scheduler.timesteps250    return timesteps, num_inference_steps251 252 253def retrieve_latents(254    encoder_output: torch.Tensor, generator: Optional[torch.Generator] = None, sample_mode: str = "sample"255):256    if hasattr(encoder_output, "latent_dist") and sample_mode == "sample":257        return encoder_output.latent_dist.sample(generator)258    elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax":259        return encoder_output.latent_dist.mode()260    elif hasattr(encoder_output, "latents"):261        return encoder_output.latents262    else:263        raise AttributeError("Could not access latents of provided encoder_output")264 265 266class MatrixGameWanActionInputStep(ModularPipelineBlocks):267    model_name = "MatrixGameWan"268 269    @property270    def description(self) -> str:271        return "Action Input step"272 273    @property274    def expected_components(self) -> List[ComponentSpec]:275        return []276 277    @property278    def inputs(self) -> List[InputParam]:279        return [InputParam("num_frames", type_hint=int, required=True), InputParam("actions", type_hint=List[str])]280 281    @property282    def intermediate_outputs(self) -> List[OutputParam]:283        return [284            OutputParam(285                "keyboard_conditions",286                type_hint=torch.Tensor,287                description="image embeddings used to guide the image generation",288            ),289            OutputParam(290                "mouse_conditions",291                type_hint=torch.Tensor,292                description="image embeddings used to guide the image generation",293            )294        ]295 296    @torch.no_grad()297    def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:298        # Get inputs and intermediates299        block_state = self.get_block_state(state)300        block_state.device = components._execution_device301        actions = block_state.actions302 303        if actions is not None:304            actions = sync_actions_to_frames(actions, block_state.num_frames)305            keyboard_conditions, mouse_conditions = actions_to_condition_tensors(actions, block_state.num_frames)306        else:307            keyboard_conditions, mouse_conditions = generate_random_condition_tensors(block_state.num_frames)308 309        block_state.keyboard_conditions = keyboard_conditions.to(block_state.device)310        block_state.mouse_conditions = mouse_conditions.to(block_state.device)311 312        # Add outputs313        self.set_block_state(state, block_state)314        return components, state315 316 317class MatrixGameWanSetTimestepsStep(ModularPipelineBlocks):318    model_name = "MatrixGameWan"319 320    @property321    def expected_components(self) -> List[ComponentSpec]:322        return [323            ComponentSpec("scheduler", UniPCMultistepScheduler),324        ]325 326    @property327    def description(self) -> str:328        return "Step that sets the scheduler's timesteps for inference"329 330    @property331    def inputs(self) -> List[InputParam]:332        return [333            InputParam("num_inference_steps", default=4),334            InputParam("timesteps"),335            InputParam("sigmas"),336        ]337 338    @property339    def intermediate_outputs(self) -> List[OutputParam]:340        return [341            OutputParam("timesteps", type_hint=torch.Tensor, description="The timesteps to use for inference"),342            OutputParam(343                "num_inference_steps",344                type_hint=int,345                description="The number of denoising steps to perform at inference time",346            ),347        ]348 349    @torch.no_grad()350    def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:351        block_state = self.get_block_state(state)352        block_state.device = components._execution_device353 354        block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps(355            components.scheduler,356            block_state.num_inference_steps,357            block_state.device,358            block_state.timesteps,359            block_state.sigmas,360        )361 362        self.set_block_state(state, block_state)363        return components, state364 365 366class MatrixGameWanPrepareLatentsStep(ModularPipelineBlocks):367    model_name = "MatrixGameWan"368 369    @property370    def expected_components(self) -> List[ComponentSpec]:371        return [ComponentSpec("vae", AutoencoderKLWan),]372 373    @property374    def description(self) -> str:375        return "Prepare latents step that prepares the latents for the text-to-video generation process"376 377    @property378    def inputs(self) -> List[InputParam]:379        return [380            InputParam("height", type_hint=int),381            InputParam("width", type_hint=int),382            InputParam("num_frames", type_hint=int),383            InputParam("latents", type_hint=Optional[torch.Tensor]),384            InputParam("num_videos_per_prompt", type_hint=int, default=1),385            InputParam("generator"),386            InputParam("dtype", type_hint=torch.dtype, description="The dtype of the model inputs"),387        ]388 389    @property390    def intermediate_outputs(self) -> List[OutputParam]:391        return [392            OutputParam(393                "latents", type_hint=torch.Tensor, description="The initial latents to use for the denoising process"394            )395        ]396 397    @staticmethod398    def check_inputs(components, block_state):399        if (block_state.height is not None and block_state.height % components.vae_scale_factor_spatial != 0) or (400            block_state.width is not None and block_state.width % components.vae_scale_factor_spatial != 0401        ):402            raise ValueError(403                f"`height` and `width` have to be divisible by {components.vae_scale_factor_spatial} but are {block_state.height} and {block_state.width}."404            )405        if block_state.num_frames is not None and (406            block_state.num_frames < 1 or (block_state.num_frames - 1) % components.vae_scale_factor_temporal != 0407        ):408            raise ValueError(409                f"`num_frames` has to be greater than 0, and (num_frames - 1) must be divisible by {components.vae_scale_factor_temporal}, but got {block_state.num_frames}."410            )411 412    @staticmethod413    def prepare_latents(414        components,415        batch_size: int,416        num_channels_latents: int = 16,417        height: int = 352,418        width: int = 640,419        num_frames: int = 81,420        dtype: Optional[torch.dtype] = None,421        device: Optional[torch.device] = None,422        generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,423        latents: Optional[torch.Tensor] = None,424    ) -> torch.Tensor:425        if latents is not None:426            return latents.to(device=device, dtype=dtype)427 428        num_latent_frames = (num_frames - 1) // components.vae_scale_factor_temporal + 1429        shape = (430            batch_size,431            num_channels_latents,432            num_latent_frames,433            int(height) // components.vae_scale_factor_spatial,434            int(width) // components.vae_scale_factor_spatial,435        )436        if isinstance(generator, list) and len(generator) != batch_size:437            raise ValueError(438                f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"439                f" size of {batch_size}. Make sure the batch size matches the length of the generators."440            )441 442        latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)443        return latents444 445    @torch.no_grad()446    def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:447        block_state = self.get_block_state(state)448 449        block_state.height = block_state.height or components.default_height450        block_state.width = block_state.width or components.default_width451        block_state.num_frames = block_state.num_frames or components.default_num_frames452        block_state.device = components._execution_device453        block_state.dtype = torch.float32  # Wan latents should be torch.float32 for best quality454        block_state.num_channels_latents = components.num_channels_latents455 456        self.check_inputs(components, block_state)457 458        block_state.latents = self.prepare_latents(459            components,460            1,461            block_state.num_channels_latents,462            block_state.height,463            block_state.width,464            block_state.num_frames,465            block_state.dtype,466            block_state.device,467            block_state.generator,468            block_state.latents,469        )470 471        self.set_block_state(state, block_state)472 473        return components, state474 475 476class MatrixGameWanPrepareImageMaskLatentsStep(ModularPipelineBlocks):477    model_name = "MatrixGameWan"478 479    @property480    def expected_components(self) -> List[ComponentSpec]:481        return [482            ComponentSpec("vae", AutoencoderKLWan),483            ComponentSpec("video_processor", VideoProcessor, config=FrozenDict({"vae_scale_factor": 8}))484        ]485 486    @property487    def description(self) -> str:488        return "Prepare latents step that prepares the latents for the text-to-video generation process"489 490    @property491    def inputs(self) -> List[InputParam]:492        return [493            InputParam("image"),494            InputParam("height", type_hint=int),495            InputParam("width", type_hint=int),496            InputParam("num_frames", type_hint=int),497            InputParam("image_mask_latents", type_hint=Optional[torch.Tensor]),498            InputParam("num_videos_per_prompt", type_hint=int, default=1),499            InputParam("generator"),500            InputParam("dtype", type_hint=torch.dtype, description="The dtype of the model inputs"),501        ]502 503    @property504    def intermediate_outputs(self) -> List[OutputParam]:505        return [506            OutputParam(507                "image_mask_latents", type_hint=torch.Tensor, description="The initial latents to use for the denoising process"508            )509        ]510 511    @staticmethod512    def check_inputs(components, block_state):513        if (block_state.height is not None and block_state.height % components.vae_scale_factor_spatial != 0) or (514            block_state.width is not None and block_state.width % components.vae_scale_factor_spatial != 0515        ):516            raise ValueError(517                f"`height` and `width` have to be divisible by {components.vae_scale_factor_spatial} but are {block_state.height} and {block_state.width}."518            )519 520    @staticmethod521    @torch.no_grad()522    def prepare_latents(523        components,524        image,525        batch_size: int,526        num_channels_latents: int = 16,527        height: int = 352,528        width: int = 640,529        num_frames: int = 81,530        dtype: Optional[torch.dtype] = None,531        device: Optional[torch.device] = None,532        generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,533        latents: Optional[torch.Tensor] = None,534    ) -> torch.Tensor:535        if latents is not None:536            return latents.to(device=device, dtype=dtype)537 538        image = components.video_processor.preprocess(image, height, width).to(device, torch.float32)539        image = image.unsqueeze(2)  # [batch_size, channels, 1, height, width]540 541        video_condition = torch.cat(542            [image, image.new_zeros(image.shape[0], image.shape[1], num_frames - 1, height, width)], dim=2543        )544        video_condition = video_condition.to(device=device, dtype=components.vae.dtype)545        latent_condition = retrieve_latents(components.vae.encode(video_condition), sample_mode="argmax")546        latent_condition = latent_condition.repeat(batch_size, 1, 1, 1, 1)547 548        latents_mean = (549            torch.tensor(components.vae.config.latents_mean)550            .view(1, components.vae.config.z_dim, 1, 1, 1)551            .to(device, dtype)552        )553        latents_std = 1.0 / torch.tensor(components.vae.config.latents_std).view(1, components.vae.config.z_dim, 1, 1, 1).to(554            device, dtype555        )556        latent_condition = latent_condition.to(dtype)557        latent_condition = (latent_condition - latents_mean) * latents_std558 559        latent_height = height // components.vae_scale_factor_spatial560        latent_width = width // components.vae_scale_factor_spatial561 562        mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height, latent_width)563        mask_lat_size[:, :, list(range(1, num_frames))] = 0564 565        first_frame_mask = mask_lat_size[:, :, 0:1]566        first_frame_mask = torch.repeat_interleave(first_frame_mask, dim=2, repeats=components.vae_scale_factor_temporal)567 568        mask_lat_size = torch.concat([first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)569        mask_lat_size = mask_lat_size.view(batch_size, -1, components.vae_scale_factor_temporal, latent_height, latent_width)570        mask_lat_size = mask_lat_size.transpose(1, 2).to(latent_condition.device)571 572        image_mask_latents = torch.concat([mask_lat_size, latent_condition], dim=1)573        return image_mask_latents574 575    @torch.no_grad()576    def __call__(self, components: ModularPipeline, state: PipelineState) -> PipelineState:577        block_state = self.get_block_state(state)578 579        block_state.height = block_state.height or components.default_height580        block_state.width = block_state.width or components.default_width581        block_state.num_frames = block_state.num_frames or components.default_num_frames582        block_state.device = components._execution_device583        block_state.dtype = torch.float32  # Wan latents should be torch.float32 for best quality584        block_state.num_channels_latents = components.num_channels_latents585 586        self.check_inputs(components, block_state)587        block_state.image_mask_latents = self.prepare_latents(588            components,589            block_state.image,590            1,591            block_state.num_channels_latents,592            block_state.height,593            block_state.width,594            block_state.num_frames,595            block_state.dtype,596            block_state.device,597            block_state.generator,598            block_state.image_mask_latents,599        )600 601        self.set_block_state(state, block_state)602 603        return components, state604 605