diffusers/matrix-game-2-modular
014
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 