Compare commits

...
Author SHA1 Message Date
kevin314 0c606b7cfa Add model selection 2026-02-14 02:06:09 +00:00
kevin314 0a0cdebaba Add gpu selection 2026-02-05 04:02:49 +00:00
kevin314 cd9e7a8cac Add gpu provisioning 2026-02-04 07:39:31 +00:00
kevin314 87e19dad5a Add model resetting 2026-02-02 21:36:52 +00:00
kevin314 79195ad18e Add web app for matrix game 2026-02-01 00:12:07 +00:00
36 changed files with 7330 additions and 170 deletions
+3
View File
@@ -40,6 +40,9 @@ dist/
*.egg
eggs/
.eggs/
node_modules/
.vite
vite.config.js.timestamp-*.mjs
# MkDocs documentation
site/
+3 -1
View File
@@ -7,9 +7,11 @@ from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
from fastvideo.configs.models.dits.wangamevideo import WanGameVideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig",
"WanGameVideoConfig"
]
@@ -0,0 +1,122 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class WanGameVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$":
r"blocks.\1.attn1.to_out.0.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
r"blocks.\1.attn2.to_out.0.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
})
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len = 512
num_attention_heads: int = 40
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
text_dim: int = 4096
freq_dim: int = 256
ffn_dim: int = 13824
num_layers: int = 40
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = 6 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 21
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
@dataclass
class WanGameVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=WanGameVideoArchConfig)
prefix: str = "WanGame"
@dataclass
class WanLingBotVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=WanGameVideoArchConfig)
prefix: str = "WanLingBot"
+5 -1
View File
@@ -23,7 +23,7 @@ from fastvideo.configs.pipelines.wan import (
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig,
MatrixGameI2V480PConfig)
MatrixGameI2V480PConfig, WanGameI2V480PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -42,6 +42,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers": HYWorldConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers": WanGameI2V480PConfig,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
@@ -93,6 +94,8 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "hyworld" in id.lower(),
"matrixgame":
lambda id: "matrix-game" in id.lower() or "matrixgame" in id.lower(),
"wangameactionimagetovideo":
lambda id: "wangameactionimagetovideo" in id.lower(),
"wanpipeline":
lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo":
@@ -124,6 +127,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"matrixgame": MatrixGameI2V480PConfig,
"wangameactionimagetovideo": WanGameI2V480PConfig,
"hunyuan15":
Hunyuan15T2V480PConfig, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"hyworld":
+21
View File
@@ -7,6 +7,7 @@ import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.configs.models.dits.matrixgame import MatrixGameWanVideoConfig
from fastvideo.configs.models.dits.wangamevideo import WanGameVideoConfig
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
CLIPVisionConfig, T5Config,
WAN2_1ControlCLIPVisionConfig)
@@ -217,3 +218,23 @@ class MatrixGameI2V480PConfig(WanI2V480PConfig):
context_noise: int = 0
num_frames_per_block: int = 3
# sliding_window_num_frames: int = 15
# =============================================
# ============= WanGame ======================
# =============================================
@dataclass
class WanGameI2V480PConfig(WanI2V480PConfig):
"""Configuration for WanGame image-to-video pipeline."""
dit_config: DiTConfig = field(default_factory=WanGameVideoConfig)
image_encoder_config: EncoderConfig = field(
default_factory=WAN2_1ControlCLIPVisionConfig)
is_causal: bool = True
flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 666, 333])
warp_denoising_step: bool = True
context_noise: int = 0
num_frames_per_block: int = 3
+6
View File
@@ -92,6 +92,9 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"KyleShao/Cosmos-Predict2.5-2B-Diffusers":
Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
# WanGame models
"weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers": MatrixGame2_SamplingParam,
# MatrixGame2.0 models
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGame2_SamplingParam,
@@ -130,6 +133,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline":
lambda id: "wancausaldmdpipeline" in id.lower(),
"wangame":
lambda id: "wangame" in id.lower(),
"matrixgame":
lambda id: "matrixgame" in id.lower() or "matrix-game" in id.lower(),
"turbodiffusion":
@@ -157,6 +162,7 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
"wandmdpipeline": FastWanT2V480P_SamplingParam,
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam,
"wangame": MatrixGame2_SamplingParam,
"matrixgame": MatrixGame2_SamplingParam,
"turbodiffusion":
TurboDiffusionT2V_1_3B_SamplingParam, # Default to T2V for fallback
+5 -2
View File
@@ -171,7 +171,7 @@ class StreamingVideoGenerator(VideoGenerator):
def step(
self, keyboard_cond: torch.Tensor,
mouse_cond: torch.Tensor) -> tuple[list[np.ndarray], Future | None]:
mouse_cond: torch.Tensor) -> tuple[list[np.ndarray], Future | None, dict | None]:
if self.batch is None:
raise RuntimeError("Call reset() before step()")
@@ -195,7 +195,10 @@ class StreamingVideoGenerator(VideoGenerator):
# Returns Future for block file, or None if no block_dir
block_future = self.writer.add_frames(frames)
return frames, block_future
# Extract stage timings if available
stage_timings = getattr(output_batch, 'stage_timings', None)
return frames, block_future, stage_timings
async def step_async(
self, keyboard_cond: torch.Tensor,
@@ -496,14 +496,14 @@ class ActionModule(nn.Module):
else:
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - kv_cache_keyboard["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
assert k.shape[0] == S # BS == 1 or the cache should not be saved/ load method should be modified
kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k[:1]
kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v[:1]
assert k.shape[0] % S == 0 # BS >= 1
kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k[::S]
kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v[::S]
attn = self.keyboard_attn_layer(
q,
kv_cache_keyboard["k"][:, max(0, local_end_index - max_attention_size):local_end_index].repeat(S, 1, 1, 1),
kv_cache_keyboard["v"][:, max(0, local_end_index - max_attention_size):local_end_index].repeat(S, 1, 1, 1),
kv_cache_keyboard["k"][:, max(0, local_end_index - max_attention_size):local_end_index].repeat_interleave(S, dim=0),
kv_cache_keyboard["v"][:, max(0, local_end_index - max_attention_size):local_end_index].repeat_interleave(S, dim=0),
)
kv_cache_keyboard["global_end_index"].fill_(current_end)
@@ -14,7 +14,7 @@ from torch.nn.attention.flex_attention import BlockMask
from fastvideo.forward_context import get_forward_context
flex_attention = torch.compile(
flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
flex_attention, dynamic=True, mode="max-autotune-no-cudagraphs")
from fastvideo.attention import LocalAttention
from fastvideo.configs.models.dits.matrixgame import MatrixGameWanVideoConfig
@@ -465,6 +465,7 @@ _DEFAULT_MATRIXGAME_CONFIG = MatrixGameWanVideoConfig()
class CausalMatrixGameWanModel(BaseDiT):
supports_action_input = True
_concatenates_image_latent = True # Model handles image_latent internally via forward_context
_fsdp_shard_conditions = _DEFAULT_MATRIXGAME_CONFIG._fsdp_shard_conditions
_compile_conditions = _DEFAULT_MATRIXGAME_CONFIG._compile_conditions
+12
View File
@@ -0,0 +1,12 @@
from .model import WanGameActionTransformer3DModel
from .causal_model import (CausalWanGameTransformer3DModel,
CausalWanTransformer3DModel)
from .hyworld_action_module import WanGameActionTimeImageEmbedding, WanGameActionSelfAttention
__all__ = [
"WanGameActionTransformer3DModel",
"CausalWanTransformer3DModel",
"CausalWanGameTransformer3DModel",
"WanGameActionTimeImageEmbedding",
"WanGameActionSelfAttention",
]
@@ -0,0 +1,98 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Any
import torch
from fastvideo.configs.models.dits.wangamevideo import WanGameVideoConfig
from fastvideo.models.dits.matrixgame.causal_model import CausalMatrixGameWanModel
_DEFAULT_WANGAME_CAUSAL_CONFIG = WanGameVideoConfig()
class CausalWanTransformer3DModel(CausalMatrixGameWanModel):
supports_action_input = False
_fsdp_shard_conditions = _DEFAULT_WANGAME_CAUSAL_CONFIG._fsdp_shard_conditions
_compile_conditions = _DEFAULT_WANGAME_CAUSAL_CONFIG._compile_conditions
_supported_attention_backends = _DEFAULT_WANGAME_CAUSAL_CONFIG._supported_attention_backends
param_names_mapping = _DEFAULT_WANGAME_CAUSAL_CONFIG.param_names_mapping
reverse_param_names_mapping = _DEFAULT_WANGAME_CAUSAL_CONFIG.reverse_param_names_mapping
lora_param_names_mapping = _DEFAULT_WANGAME_CAUSAL_CONFIG.lora_param_names_mapping
def __init__(self,
config: WanGameVideoConfig,
hf_config: dict[str, Any],
**kwargs) -> None:
super().__init__(config=config, hf_config=hf_config, **kwargs)
def _normalize_action_inputs(
self,
mouse_cond: torch.Tensor | None,
keyboard_cond: torch.Tensor | None,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
# drop
if len(getattr(self, "action_config", {})) == 0:
return None, None
return mouse_cond, keyboard_cond
def _forward_inference(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
kv_cache: dict | None = None,
kv_cache_mouse: dict | None = None,
kv_cache_keyboard: dict | None = None,
crossattn_cache: dict | None = None,
current_start: int = 0,
cache_start: int = 0,
start_frame: int = 0,
**kwargs,
) -> torch.Tensor:
mouse_cond, keyboard_cond = self._normalize_action_inputs(
mouse_cond, keyboard_cond)
return super()._forward_inference(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
encoder_hidden_states_image=encoder_hidden_states_image,
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond,
kv_cache=kv_cache,
kv_cache_mouse=kv_cache_mouse,
kv_cache_keyboard=kv_cache_keyboard,
crossattn_cache=crossattn_cache,
current_start=current_start,
cache_start=cache_start,
start_frame=start_frame,
**kwargs,
)
def _forward_train(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
**kwargs,
) -> torch.Tensor:
mouse_cond, keyboard_cond = self._normalize_action_inputs(
mouse_cond, keyboard_cond)
return super()._forward_train(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
encoder_hidden_states_image=encoder_hidden_states_image,
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond,
**kwargs,
)
class CausalWanGameTransformer3DModel(CausalWanTransformer3DModel):
pass
@@ -0,0 +1,289 @@
import math
import torch
import torch.nn as nn
from fastvideo.layers.visual_embedding import TimestepEmbedder, ModulateProjection, timestep_embedding
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.attention import DistributedAttention
from fastvideo.forward_context import set_forward_context
from fastvideo.models.dits.wanvideo import WanImageEmbedding
from fastvideo.models.dits.hyworld.camera_rope import prope_qkv
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
from fastvideo.layers.mlp import MLP
class WanGameActionTimeImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
image_embed_dim: int | None = None,
):
super().__init__()
self.time_freq_dim = time_freq_dim
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
self.action_embedder = MLP(
time_freq_dim,
dim,
dim,
bias=True,
act_type="silu"
)
# Initialize fc_in with kaiming_uniform (same as nn.Linear default)
nn.init.kaiming_uniform_(self.action_embedder.fc_in.weight, a=math.sqrt(5))
# Initialize fc_out with zeros for residual-like behavior
nn.init.zeros_(self.action_embedder.fc_out.weight)
if self.action_embedder.fc_out.bias is not None:
nn.init.zeros_(self.action_embedder.fc_out.bias)
def forward(
self,
timestep: torch.Tensor,
action: torch.Tensor,
encoder_hidden_states: torch.Tensor, # Kept for interface compatibility
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
"""
Args:
timestep: [B] diffusion timesteps (one per batch sample)
action: [B, T] action labels (one per frame per batch sample)
Returns:
temb: [B*T, dim] combined timestep + action embedding
timestep_proj: [B*T, 6*dim] modulation projection
"""
# timestep: [B] -> temb: [B, dim]
temb = self.time_embedder(timestep, timestep_seq_len)
# Handle action embedding for batch > 1
# action shape: [B, T] where B=batch_size, T=num_frames
batch_size = action.shape[0]
num_frames = action.shape[1]
# Compute action embeddings: [B, T] -> [B*T] -> [B*T, dim]
action_flat = action.flatten() # [B*T]
action_emb = timestep_embedding(action_flat, self.time_freq_dim)
action_embedder_dtype = next(iter(self.action_embedder.parameters())).dtype
if (
action_emb.dtype != action_embedder_dtype
and action_embedder_dtype != torch.int8
):
action_emb = action_emb.to(action_embedder_dtype)
action_emb = self.action_embedder(action_emb).type_as(temb) # [B*T, dim]
# Expand temb to match action_emb: [B, dim] -> [B, T, dim] -> [B*T, dim]
# Each batch's temb is repeated for all its frames
temb_expanded = temb.unsqueeze(1).expand(-1, num_frames, -1) # [B, T, dim]
temb_expanded = temb_expanded.reshape(batch_size * num_frames, -1) # [B*T, dim]
# Add action embedding to expanded temb
temb = temb_expanded + action_emb # [B*T, dim]
timestep_proj = self.time_modulation(temb) # [B*T, 6*dim]
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
if encoder_hidden_states_image is not None:
assert self.image_embedder is not None
encoder_hidden_states_image = self.image_embedder(
encoder_hidden_states_image)
encoder_hidden_states = torch.zeros((batch_size, 0, temb.shape[-1]),
device=temb.device,
dtype=temb.dtype)
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
class WanGameActionSelfAttention(nn.Module):
"""
Self-attention module with support for:
- Standard RoPE-based attention
- Camera PRoPE-based attention (when viewmats and Ks are provided)
- KV caching for autoregressive generation
"""
def __init__(self,
dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm=True,
eps=1e-6) -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.max_attention_size = 32760 if local_attn_size == -1 else local_attn_size * 1560
# Scaled dot product attention (using DistributedAttention for SP support)
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=self.head_dim,
softmax_scale=None,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
kv_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
is_cache: bool = False,
attention_mask: torch.Tensor | None = None):
"""
Forward pass with camera PRoPE attention combining standard RoPE and projective positional encoding.
Args:
q, k, v: Query, key, value tensors [B, L, num_heads, head_dim]
freqs_cis: RoPE frequency cos/sin tensors
kv_cache: KV cache dict (may have None values for training)
current_start: Current position for KV cache
cache_start: Cache start position
viewmats: Camera view matrices for PRoPE [B, cameras, 4, 4]
Ks: Camera intrinsics for PRoPE [B, cameras, 3, 3]
is_cache: Whether to store to KV cache (for inference)
attention_mask: Attention mask [B, L] (1 = attend, 0 = mask)
"""
if cache_start is None:
cache_start = current_start
# Apply RoPE manually
cos, sin = freqs_cis
query_rope = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
key_rope = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
value_rope = v
# # DEBUG: Check camera matrices
# if self.training and torch.distributed.get_rank() == 0:
# vm_info = f"viewmats={viewmats.shape if viewmats is not None else None}"
# ks_info = f"Ks={Ks.shape if Ks is not None else None}"
# vm_nonzero = (viewmats != 0).sum().item() if viewmats is not None else 0
# ks_nonzero = (Ks != 0).sum().item() if Ks is not None else 0
# print(f"[DEBUG] PRoPE input: {vm_info} nonzero={vm_nonzero}, {ks_info} nonzero={ks_nonzero}", flush=True)
# Get PRoPE transformed q, k, v
query_prope, key_prope, value_prope, apply_fn_o = prope_qkv(
q.transpose(1, 2), # [B, num_heads, L, head_dim]
k.transpose(1, 2),
v.transpose(1, 2),
viewmats=viewmats,
Ks=Ks,
patches_x=40, # hardcoded for now
patches_y=22,
)
# PRoPE returns [B, num_heads, L, head_dim], convert to [B, L, num_heads, head_dim]
query_prope = query_prope.transpose(1, 2)
key_prope = key_prope.transpose(1, 2)
value_prope = value_prope.transpose(1, 2)
# # DEBUG: Check prope_qkv output
# if self.training and torch.distributed.get_rank() == 0:
# q_nz = (query_prope != 0).sum().item()
# k_nz = (key_prope != 0).sum().item()
# v_nz = (value_prope != 0).sum().item()
# print(f"[DEBUG] prope_qkv output: q_nonzero={q_nz}, k_nonzero={k_nz}, v_nonzero={v_nz}", flush=True)
# KV cache handling
if kv_cache is not None:
cache_key = kv_cache.get("k", None)
cache_value = kv_cache.get("v", None)
if cache_value is not None and not is_cache:
cache_key_rope, cache_key_prope = cache_key.chunk(2, dim=-1)
cache_value_rope, cache_value_prope = cache_value.chunk(2, dim=-1)
key_rope = torch.cat([cache_key_rope, key_rope], dim=1)
value_rope = torch.cat([cache_value_rope, value_rope], dim=1)
key_prope = torch.cat([cache_key_prope, key_prope], dim=1)
value_prope = torch.cat([cache_value_prope, value_prope], dim=1)
if is_cache:
# Store to cache (update input dict directly)
kv_cache["k"] = torch.cat([key_rope, key_prope], dim=-1)
kv_cache["v"] = torch.cat([value_rope, value_prope], dim=-1)
# Concatenate rope and prope paths (matching original)
query_all = torch.cat([query_rope, query_prope], dim=0)
key_all = torch.cat([key_rope, key_prope], dim=0)
value_all = torch.cat([value_rope, value_prope], dim=0)
# Check if Q and KV have different sequence lengths (KV cache mode)
# In this case, use LocalAttention (supports different Q/KV lengths)
if query_all.shape[1] != key_all.shape[1]:
# KV cache mode: Q has new tokens only, KV has cached + new tokens
# Use LocalAttention which supports different Q/KV lengths
# LocalAttention will use the appropriate backend (SageAttn, FlashAttn, etc.)
if not hasattr(self, '_kv_cache_attn'):
from fastvideo.attention import LocalAttention
self._kv_cache_attn = LocalAttention(
num_heads=self.num_heads,
head_size=self.head_dim,
causal=False,
supported_attention_backends=(AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
)
hidden_states_all = self._kv_cache_attn(query_all, key_all, value_all)
else:
# Same sequence length: use DistributedAttention (supports SP)
# Create default attention mask if not provided
# NOTE: query_all has shape [2*B, L, ...] (rope+prope concatenated), so mask needs 2*B
if attention_mask is None:
batch_size, seq_len = q.shape[0], q.shape[1]
attention_mask = torch.ones(batch_size * 2, seq_len, device=q.device, dtype=q.dtype)
from fastvideo.attention.backends.sdpa import SDPAMetadataBuilder
attn_metadata_builder = SDPAMetadataBuilder
if q.dtype != torch.float32:
try:
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata_builder = FlashAttnMetadataBuilder
except ImportError:
pass # fall back to SDPA
attn_metadata = attn_metadata_builder().build(
current_timestep=0,
attn_mask=attention_mask,
)
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
hidden_states_all, _ = self.attn(query_all, key_all, value_all, attention_mask=attention_mask)
hidden_states_rope, hidden_states_prope = hidden_states_all.chunk(2, dim=0)
# # DEBUG: Check attention output and apply_fn_o
# if self.training and torch.distributed.get_rank() == 0:
# attn_all_nz = (hidden_states_all != 0).sum().item()
# rope_nz = (hidden_states_rope != 0).sum().item()
# prope_before = (hidden_states_prope != 0).sum().item()
# print(f"[DEBUG] attn output: all_nonzero={attn_all_nz}, rope_nonzero={rope_nz}, prope_before_apply={prope_before}", flush=True)
hidden_states_prope = apply_fn_o(hidden_states_prope.transpose(1, 2)).transpose(1, 2)
# # DEBUG: Check after apply_fn_o
# if self.training and torch.distributed.get_rank() == 0:
# prope_after = (hidden_states_prope != 0).sum().item()
# print(f"[DEBUG] prope_after_apply_fn_o={prope_after}", flush=True)
return hidden_states_rope, hidden_states_prope
+427
View File
@@ -0,0 +1,427 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
from fastvideo.configs.models.dits.wangamevideo import WanGameVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import PatchEmbed
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanI2VCrossAttention
from fastvideo.platforms import AttentionBackendEnum, current_platform
# Import ActionModule
from fastvideo.models.dits.wangame.hyworld_action_module import WanGameActionTimeImageEmbedding, WanGameActionSelfAttention
logger = init_logger(__name__)
class WanGameCrossAttention(WanI2VCrossAttention):
def forward(self, x, context, context_lens=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
"""
context_img = context
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
# output
x = img_x.flatten(2)
x, _ = self.to_out(x)
return x
class WanGameActionTransformerBlock(nn.Module):
"""
Transformer block for WAN Action model with support for:
- Self-attention with RoPE and camera PRoPE
- Cross-attention with text/image context
- Feed-forward network with AdaLN modulation
"""
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = WanGameActionSelfAttention(
dim,
num_heads,
local_attn_size=local_attn_size,
sink_size=sink_size,
qk_norm=qk_norm,
eps=eps)
self.hidden_dim = dim
self.num_attention_heads = num_heads
self.local_attn_size = local_attn_size
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
raise ValueError(f"QK Norm type {qk_norm} not supported")
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=True,
compute_dtype=torch.float32)
# 2. Cross-attention (I2V only for now)
self.attn2 = WanGameCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
# norm3 for FFN input
self.norm3 = LayerNormScaleShift(dim, norm_type="layer", eps=eps,
elementwise_affine=False)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
# PRoPE output projection (initialized via add_discrete_action_parameters on the model)
self.to_out_prope = ReplicatedLinear(dim, dim, bias=True)
nn.init.zeros_(self.to_out_prope.weight)
if self.to_out_prope.bias is not None:
nn.init.zeros_(self.to_out_prope.bias)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
kv_cache: dict | None = None,
crossattn_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
is_cache: bool = False,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# Cast temb to float32 for scale/shift computation
e = self.scale_shift_table + temb.float()
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(6, dim=2)
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).to(orig_dtype).flatten(1, 2)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Self-attention with camera PRoPE
attn_output_rope, attn_output_prope = self.attn1(
query, key, value, freqs_cis,
kv_cache, current_start, cache_start, viewmats, Ks,
is_cache=is_cache
)
# Combine rope and prope outputs
attn_output_rope = attn_output_rope.flatten(2)
attn_output_rope, _ = self.to_out(attn_output_rope)
attn_output_prope = attn_output_prope.flatten(2)
# # DEBUG: Check if prope input is zero
# if self.training and torch.distributed.get_rank() == 0:
# prope_nonzero = (attn_output_prope != 0).sum().item()
# prope_total = attn_output_prope.numel()
# if prope_nonzero == 0:
# print(f"[DEBUG] to_out_prope INPUT is ALL ZEROS! shape={attn_output_prope.shape}", flush=True)
attn_output_prope, _ = self.to_out_prope(attn_output_prope)
attn_output = attn_output_rope.squeeze(1) + attn_output_prope.squeeze(1)
# Self-attention residual + norm in float32
null_shift = null_scale = torch.zeros(1, device=hidden_states.device, dtype=torch.float32)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states.float(), attn_output.float(), gate_msa, null_shift, null_scale)
hidden_states = hidden_states.type_as(attn_output)
norm_hidden_states = norm_hidden_states.type_as(attn_output)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states.to(orig_dtype),
context=encoder_hidden_states,
context_lens=None)
# Cross-attention residual in bfloat16
hidden_states = hidden_states + attn_output
# norm3 for FFN input in float32
norm_hidden_states = self.norm3(
hidden_states.float(), c_shift_msa, c_scale_msa
).type_as(hidden_states)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states.to(orig_dtype))
hidden_states = self.mlp_residual(hidden_states.float(), ff_output.float(), c_gate_msa)
hidden_states = hidden_states.to(orig_dtype) # Cast back to original dtype
return hidden_states
class WanGameActionTransformer3DModel(BaseDiT):
"""
WAN Action Transformer 3D Model for video generation with action conditioning.
Extends the base WAN video model with:
- Action embedding support for controllable generation
- camera PRoPE attention for 3D-aware generation
- KV caching for autoregressive inference
"""
_uses_discrete_action = True # Uses single action tensor [B, T] instead of keyboard/mouse KV caches
_kv_cache_head_dim_multiplier = 2 # PRoPE stores cat([rope, prope], dim=-1) in KV cache
_fsdp_shard_conditions = WanGameVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanGameVideoConfig()._compile_conditions
_supported_attention_backends = WanGameVideoConfig()._supported_attention_backends
param_names_mapping = WanGameVideoConfig().param_names_mapping
reverse_param_names_mapping = WanGameVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = WanGameVideoConfig().lora_param_names_mapping
def __init__(self, config: WanGameVideoConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.local_attn_size = config.local_attn_size
self.inner_dim = inner_dim
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings (with action support)
self.condition_embedder = WanGameActionTimeImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
image_embed_dim=config.image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList([
WanGameActionTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.local_attn_size,
config.sink_size,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.blocks.{i}")
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
# Causal-specific
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.__post_init__()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor],
guidance=None,
action: torch.Tensor | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
kv_cache: list[dict] | None = None,
crossattn_cache: list[dict] | None = None,
current_start: int = 0,
cache_start: int = 0,
start_frame: int = 0,
is_cache: bool = False,
**kwargs
) -> torch.Tensor:
"""
Forward pass for both training and inference with KV caching.
Args:
hidden_states: Video latents [B, C, T, H, W]
encoder_hidden_states: Text embeddings [B, L, D]
timestep: Timestep tensor
encoder_hidden_states_image: Optional image embeddings
action: Action tensor [B, T] for per-frame conditioning
viewmats: Camera view matrices for PRoPE [B, T, 4, 4]
Ks: Camera intrinsics for PRoPE [B, T, 3, 3]
kv_cache: KV cache for autoregressive inference (list of dicts per layer)
crossattn_cache: Cross-attention cache for inference
current_start: Current position for KV cache
cache_start: Cache start position
start_frame: RoPE offset for new frames in autoregressive mode
is_cache: If True, populate KV cache and return early (cache-only mode)
"""
orig_dtype = hidden_states.dtype
# if not isinstance(encoder_hidden_states, torch.Tensor):
# encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image, list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
# else:
# encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height, post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
start_frame=start_frame
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if timestep.dim() == 2:
# condition_embedder expects [B] timestep, not [B*T].
# All frames share the same timestep, so take first column.
timestep = timestep[:, 0]
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, action, encoder_hidden_states, encoder_hidden_states_image=encoder_hidden_states_image)
# condition_embedder returns:
# - temb: [B*T, dim] where T = post_patch_num_frames
# - timestep_proj: [B*T, 6*dim]
# Reshape to [B, T, 6, dim] for transformer blocks
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)) # [B*T, 6, dim]
timestep_proj = timestep_proj.view(batch_size, post_patch_num_frames, 6, self.hidden_size) # [B, T, 6, dim]
encoder_hidden_states = encoder_hidden_states_image
# Transformer blocks
for block_idx, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj, freqs_cis,
kv_cache[block_idx] if kv_cache else None,
crossattn_cache[block_idx] if crossattn_cache else None,
current_start, cache_start,
viewmats, Ks, is_cache)
else:
hidden_states = block(
hidden_states, encoder_hidden_states, timestep_proj, freqs_cis,
kv_cache[block_idx] if kv_cache else None,
crossattn_cache[block_idx] if crossattn_cache else None,
current_start, cache_start,
viewmats, Ks, is_cache)
# If cache-only mode, return early
if is_cache:
return kv_cache
# Output norm, projection & unpatchify
# temb is [B*T, dim], reshape to [B, T, 1, dim]
temb = temb.view(batch_size, post_patch_num_frames, -1).unsqueeze(2) # [B, T, 1, dim]
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2, dim=2)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
+2
View File
@@ -44,6 +44,8 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"MatrixGameWanModel": ("dits", "matrixgame", "MatrixGameWanModel"),
"CausalMatrixGameWanModel": ("dits", "matrixgame", "CausalMatrixGameWanModel"),
"WanGameActionTransformer3DModel": ("dits", "wangame", "WanGameActionTransformer3DModel"),
"CausalWanGameTransformer3DModel": ("dits", "wangame", "CausalWanGameTransformer3DModel"),
}
_TEXT_ENCODER_MODELS = {
@@ -96,28 +96,47 @@ class MatrixGameCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
self._vae_cache = None
def streaming_step(self, keyboard_action, mouse_action) -> ForwardBatch:
import time
denoiser = self._stage_name_mapping["denoising_stage"]
ctx = denoiser._streaming_ctx
assert ctx is not None, "streaming_ctx must be set"
start_idx = ctx.start_index
# Time DiT forward pass
torch.cuda.synchronize()
t0 = time.perf_counter()
batch = denoiser.streaming_step(keyboard_action, mouse_action)
torch.cuda.synchronize()
dit_time_ms = (time.perf_counter() - t0) * 1000
end_idx = ctx.start_index
# Decode only the new generated block
vae_time_ms = 0.0
if end_idx > start_idx:
current_latents = batch.latents[:, :, start_idx:end_idx, :, :]
args = ctx.fastvideo_args
decoder = self._stage_name_mapping["decoding_stage"]
# Time VAE decode
torch.cuda.synchronize()
t0 = time.perf_counter()
decoded_frames, self._vae_cache = decoder.streaming_decode(
current_latents,
args,
cache=self._vae_cache,
is_first_chunk=(start_idx == 0))
torch.cuda.synchronize()
vae_time_ms = (time.perf_counter() - t0) * 1000
batch.output = decoded_frames
else:
batch.output = None
# Store timings in batch for caller to access
batch.stage_timings = {"dit_ms": dit_time_ms, "vae_ms": vae_time_ms}
return batch
def streaming_clear(self) -> None:
@@ -0,0 +1,122 @@
# SPDX-License-Identifier: Apache-2.0
"""WanGame causal DMD pipeline implementation."""
from fastvideo.fastvideo_args import FastVideoArgs
import torch
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, ForwardBatch, LoRAPipeline
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
MatrixGameCausalDenoisingStage)
from fastvideo.pipelines.stages.image_encoding import (
MatrixGameImageEncodingStage, MatrixGameImageVAEEncodingStage)
logger = init_logger(__name__)
class WanGameCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"vae", "transformer", "scheduler", "image_encoder", "image_processor"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
if (self.get_module("text_encoder", None) is not None
and self.get_module("tokenizer", None) is not None):
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
if (self.get_module("image_encoder", None) is not None
and self.get_module("image_processor", None) is not None):
self.add_stage(
stage_name="image_encoding_stage",
stage=MatrixGameImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(
stage_name="image_latent_preparation_stage",
stage=MatrixGameImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=MatrixGameCausalDenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
pipeline=self,
vae=self.get_module("vae")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
logger.info("WanGameCausalDMDPipeline initialized")
@torch.no_grad()
def streaming_reset(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs):
if not self.post_init_called:
self.post_init()
stages_to_run = [
"input_validation_stage", "prompt_encoding_stage",
"image_encoding_stage", "conditioning_stage",
"latent_preparation_stage", "image_latent_preparation_stage"
]
for stage_name in stages_to_run:
if stage_name in self._stage_name_mapping:
batch = self._stage_name_mapping[stage_name].forward(
batch, fastvideo_args)
denoiser = self._stage_name_mapping["denoising_stage"]
denoiser.streaming_reset(batch, fastvideo_args)
self._vae_cache = None
def streaming_step(self, keyboard_action, mouse_action) -> ForwardBatch:
denoiser = self._stage_name_mapping["denoising_stage"]
ctx = denoiser._streaming_ctx
assert ctx is not None, "streaming_ctx must be set"
start_idx = ctx.start_index
batch = denoiser.streaming_step(keyboard_action, mouse_action)
end_idx = ctx.start_index
if end_idx > start_idx:
current_latents = batch.latents[:, :, start_idx:end_idx, :, :]
args = ctx.fastvideo_args
decoder = self._stage_name_mapping["decoding_stage"]
decoded_frames, self._vae_cache = decoder.streaming_decode(
current_latents,
args,
cache=self._vae_cache,
is_first_chunk=(start_idx == 0))
batch.output = decoded_frames
else:
batch.output = None
return batch
def streaming_clear(self) -> None:
denoiser = self._stage_name_mapping.get("denoising_stage")
if denoiser is not None and hasattr(denoiser, "streaming_clear"):
denoiser.streaming_clear()
self._vae_cache = None
EntryClass = [WanGameCausalDMDPipeline]
@@ -0,0 +1,78 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
logger = init_logger(__name__)
class WanGameActionImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(
stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
class WanLingBotImageToVideoPipeline(WanGameActionImageToVideoPipeline):
pass
EntryClass = [WanGameActionImageToVideoPipeline, WanLingBotImageToVideoPipeline]
+2
View File
@@ -33,6 +33,8 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"Cosmos2_5Pipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
"MatrixGameCausalDMDPipeline": "matrixgame",
"WanGameActionImageToVideoPipeline": "wan",
"WanGameCausalDMDPipeline": "wan",
"LongCatPipeline": "longcat",
"LongCatImageToVideoPipeline": "longcat",
"LongCatVideoContinuationPipeline": "longcat",
+23 -4
View File
@@ -845,7 +845,18 @@ class MatrixGameImageVAEEncodingStage(ImageVAEEncodingStage):
num_frames = batch.num_frames if isinstance(
batch.num_frames, int) else batch.num_frames[0]
def _gpu_mem(label=""):
a = torch.cuda.memory_allocated() / 1024**3
r = torch.cuda.memory_reserved() / 1024**3
logger.info("VAE stage [%s]: alloc=%.2fGiB, reserved=%.2fGiB", label, a, r)
_gpu_mem("before vae.to(GPU)")
self.vae = self.vae.to(get_local_torch_device())
_gpu_mem("after vae.to(GPU)")
logger.info("VAE use_feature_cache=%s, use_tiling=%s, use_temporal_tiling=%s",
getattr(self.vae, 'use_feature_cache', 'N/A'),
getattr(self.vae, 'use_tiling', 'N/A'),
getattr(self.vae, 'use_temporal_tiling', 'N/A'))
# Process single image for I2V (latent dimensions computed but not used directly)
@@ -906,15 +917,23 @@ class MatrixGameImageVAEEncodingStage(ImageVAEEncodingStage):
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
# Encode Image
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
_gpu_mem("before encode")
logger.info("video_condition shape=%s dtype=%s device=%s",
video_condition.shape, video_condition.dtype,
video_condition.device)
# Encode Image (no_grad avoids autograd graph that would retain all
# intermediate activations across the feature-cache loop iterations)
with torch.no_grad(), torch.autocast(
device_type="cuda", dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
if not vae_autocast_enabled:
video_condition = video_condition.to(vae_dtype)
_gpu_mem("right before vae.encode()")
encoder_output = self.vae.encode(video_condition)
_gpu_mem("after vae.encode()")
# MatrixGame uses deterministic VAE encode for the first-frame conditioning.
# Sampling would inject random noise into the cond_concat tensor and destroy the action guidance.
+331 -155
View File
@@ -313,32 +313,43 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
attention_head_dim = getattr(
self.transformer, 'attention_head_dim',
self.transformer.hidden_size // num_attention_heads)
# WanGame PRoPE stores cat([rope, prope], dim=-1) in cache
head_dim_mult = getattr(self.transformer,
'_kv_cache_head_dim_multiplier', 1)
attention_head_dim = attention_head_dim * head_dim_mult
if self.local_attn_size != -1:
kv_cache_size = self.local_attn_size * self.frame_seq_length
else:
kv_cache_size = self.frame_seq_length * self.sliding_window_num_frames
# WanGame uses None-initialized cache (concat-based); MatrixGame uses
# zero-tensor cache with index tracking.
use_none_init = head_dim_mult > 1
for _ in range(self.num_transformer_blocks):
kv_cache.append({
"k":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
})
if use_none_init:
kv_cache.append({"k": None, "v": None})
else:
kv_cache.append({
"k":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
})
return kv_cache
@@ -429,6 +440,227 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
})
return crossattn_cache
def _denoise_one_step(
self,
current_latents: torch.Tensor,
noise_latents_btchw: torch.Tensor,
batch: ForwardBatch,
start_index: int,
current_num_frames: int,
timestep: torch.Tensor,
step_idx: int,
next_timestep: torch.Tensor | None,
ctx: BlockProcessingContext,
action_kwargs: dict[str, Any],
noise_generator: Callable[[tuple, torch.dtype, int], torch.Tensor]
| None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Run one denoising step. Returns (updated_latents, updated_noise_latents_btchw)."""
prompt_embeds = batch.prompt_embeds
t_cur = timestep
if ctx.boundary_timestep is not None and t_cur < ctx.boundary_timestep:
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
else:
current_model = self.transformer
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(ctx.target_dtype)
# For models that don't handle image_latent internally (e.g., WanGame),
# concatenate image conditioning along the channel dimension.
# MatrixGame handles this inside its forward() via forward_context.
if (batch.image_latent is not None
and not getattr(current_model, '_concatenates_image_latent',
False)):
image_cond = batch.image_latent.to(ctx.target_dtype)
end = start_index + current_num_frames
if image_cond.shape[2] >= end:
image_cond = image_cond[:, :, start_index:end]
elif image_cond.shape[2] > start_index:
image_cond = image_cond[:, :, start_index:]
pad_frames = current_num_frames - image_cond.shape[2]
if pad_frames > 0:
pad = torch.zeros(image_cond.shape[0],
image_cond.shape[1],
pad_frames,
image_cond.shape[3],
image_cond.shape[4],
device=image_cond.device,
dtype=image_cond.dtype)
image_cond = torch.cat([image_cond, pad], dim=2)
else:
image_cond = torch.zeros(image_cond.shape[0],
image_cond.shape[1],
current_num_frames,
image_cond.shape[3],
image_cond.shape[4],
device=image_cond.device,
dtype=image_cond.dtype)
latent_model_input = torch.cat([latent_model_input, image_cond],
dim=1)
independent_first_frame = getattr(self.transformer,
'independent_first_frame', False)
if batch.image_latent is not None and independent_first_frame and start_index == 0:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(ctx.target_dtype)
],
dim=2)
# t_expand needs to be [batch * frames] to match flattened pred_noise/noise_latents
t_expand = t_cur.repeat(latent_model_input.shape[0] *
current_num_frames)
# Build attention metadata if VSA is available
if vsa_available and self.attn_backend == VideoSparseAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
h, w = current_latents.shape[-2:]
attn_metadata = self.attn_metadata_builder.build(
current_timestep=step_idx,
raw_latent_shape=(current_num_frames, h, w),
patch_size=ctx.fastvideo_args.pipeline_config.
dit_config.patch_size,
STA_param=batch.STA_param,
VSA_sparsity=ctx.fastvideo_args.VSA_sparsity,
device=get_local_torch_device(),
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
with torch.autocast(device_type="cuda",
dtype=ctx.target_dtype,
enabled=ctx.autocast_enabled), \
set_forward_context(current_timestep=step_idx,
attn_metadata=attn_metadata,
forward_batch=batch):
# Expand timestep to per-frame format [batch, num_frames] for causal model
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], current_num_frames),
device=latent_model_input.device,
dtype=torch.long)
model_kwargs = {
"kv_cache": ctx.get_kv_cache(t_cur),
"crossattn_cache": ctx.crossattn_cache,
"current_start": start_index * self.frame_seq_length,
"start_frame": start_index,
}
if self.use_action_module and current_model == self.transformer:
model_kwargs.update({
"kv_cache_mouse":
ctx.kv_cache_mouse,
"kv_cache_keyboard":
ctx.kv_cache_keyboard,
})
model_kwargs.update(action_kwargs)
# For models with discrete action embedding (WanGame-style):
# convert keyboard one-hot to scalar action = trans_label * 9
# and provide identity camera matrices for PRoPE
if getattr(current_model, '_uses_discrete_action', False):
dev = latent_model_input.device
b = latent_model_input.shape[0]
vae_ratio = 4
start_pixel = (0 if start_index == 0 else 1 + vae_ratio *
(start_index - 1))
if (batch.keyboard_cond is not None
and start_pixel < batch.keyboard_cond.shape[1]):
kbd = batch.keyboard_cond[:, start_pixel, :].to(dev)
has_key = kbd.sum(dim=-1) > 0 # [B]
trans_label = (kbd.argmax(dim=-1) + 1) * has_key.long()
else:
trans_label = torch.zeros(b, device=dev, dtype=torch.long)
action_value = trans_label * 9 # composite: trans * 9 + rot
model_kwargs["action"] = action_value.unsqueeze(1).expand(
-1, current_num_frames).to(dev)
# Identity viewmats and default normalized intrinsics for PRoPE
model_kwargs["viewmats"] = torch.eye(
4, device=dev, dtype=ctx.target_dtype
).unsqueeze(0).unsqueeze(0).expand(
b, current_num_frames, -1, -1).contiguous()
default_K = torch.tensor(
[[0.5052, 0.0, 0.5],
[0.0, 0.8979, 0.5],
[0.0, 0.0, 1.0]],
device=dev, dtype=ctx.target_dtype)
model_kwargs["Ks"] = default_K.unsqueeze(0).unsqueeze(
0).expand(b, current_num_frames, -1, -1).contiguous()
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
t_expanded_noise,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
**model_kwargs,
).permute(0, 2, 1, 3, 4)
if ctx.boundary_timestep is not None and t_cur >= ctx.boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
boundary_timestep=torch.ones_like(t_expand) *
ctx.boundary_timestep,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
else:
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if next_timestep is not None:
next_t = next_timestep * torch.ones(
[1], dtype=torch.long, device=pred_video_btchw.device)
# Use custom noise generator if provided (for streaming), else generate
if noise_generator is not None:
noise = noise_generator(pred_video_btchw.shape,
pred_video_btchw.dtype, step_idx)
else:
noise = torch.randn(
pred_video_btchw.shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else batch.generator)).to(
pred_video_btchw.device)
noise_btchw = noise
if ctx.boundary_timestep is not None and ctx.high_noise_timesteps is not None and step_idx < len(
ctx.high_noise_timesteps) - 1:
noise_latents_btchw = self.scheduler.add_noise_high(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1), next_t,
torch.ones_like(next_t) *
ctx.boundary_timestep).unflatten(
0, pred_video_btchw.shape[:2])
elif ctx.boundary_timestep is not None and ctx.high_noise_timesteps is not None and step_idx == len(
ctx.high_noise_timesteps) - 1:
noise_latents_btchw = pred_video_btchw
else:
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0,
1), noise_btchw.flatten(0, 1),
next_t).unflatten(0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(0, 2, 1, 3, 4)
else:
current_latents = pred_video_btchw.permute(0, 2, 1, 3, 4)
return current_latents, noise_latents_btchw
def _process_single_block(
self,
current_latents: torch.Tensor,
@@ -442,144 +674,23 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
| None = None,
progress_bar: Any | None = None,
) -> torch.Tensor:
prompt_embeds = batch.prompt_embeds
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
for i, t_cur in enumerate(timesteps):
if ctx.boundary_timestep is not None and t_cur < ctx.boundary_timestep:
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
else:
current_model = self.transformer
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(ctx.target_dtype)
independent_first_frame = getattr(self.transformer,
'independent_first_frame', False)
if batch.image_latent is not None and independent_first_frame and start_index == 0:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(ctx.target_dtype)
],
dim=2)
# t_expand needs to be [batch * frames] to match flattened pred_noise/noise_latents
t_expand = t_cur.repeat(latent_model_input.shape[0] *
current_num_frames)
# Build attention metadata if VSA is available
if vsa_available and self.attn_backend == VideoSparseAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
h, w = current_latents.shape[-2:]
attn_metadata = self.attn_metadata_builder.build(
current_timestep=i,
raw_latent_shape=(current_num_frames, h, w),
patch_size=ctx.fastvideo_args.pipeline_config.
dit_config.patch_size,
STA_param=batch.STA_param,
VSA_sparsity=ctx.fastvideo_args.VSA_sparsity,
device=get_local_torch_device(),
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
with torch.autocast(device_type="cuda",
dtype=ctx.target_dtype,
enabled=ctx.autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
# Expand timestep to per-frame format [batch, num_frames] for causal model
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], current_num_frames),
device=latent_model_input.device,
dtype=torch.long)
model_kwargs = {
"kv_cache": ctx.get_kv_cache(t_cur),
"crossattn_cache": ctx.crossattn_cache,
"current_start": start_index * self.frame_seq_length,
"start_frame": start_index,
}
if self.use_action_module and current_model == self.transformer:
model_kwargs.update({
"kv_cache_mouse":
ctx.kv_cache_mouse,
"kv_cache_keyboard":
ctx.kv_cache_keyboard,
})
model_kwargs.update(action_kwargs)
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
t_expanded_noise,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
**model_kwargs,
).permute(0, 2, 1, 3, 4)
if ctx.boundary_timestep is not None and t_cur >= ctx.boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
boundary_timestep=torch.ones_like(t_expand) *
ctx.boundary_timestep,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
else:
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
[1], dtype=torch.long, device=pred_video_btchw.device)
# Use custom noise generator if provided (for streaming), else generate
if noise_generator is not None:
noise = noise_generator(pred_video_btchw.shape,
pred_video_btchw.dtype, i)
else:
noise = torch.randn(
pred_video_btchw.shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else batch.generator)).to(
pred_video_btchw.device)
noise_btchw = noise
if ctx.boundary_timestep is not None and ctx.high_noise_timesteps is not None and i < len(
ctx.high_noise_timesteps) - 1:
noise_latents_btchw = self.scheduler.add_noise_high(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1), next_timestep,
torch.ones_like(next_timestep) *
ctx.boundary_timestep).unflatten(
0, pred_video_btchw.shape[:2])
elif ctx.boundary_timestep is not None and ctx.high_noise_timesteps is not None and i == len(
ctx.high_noise_timesteps) - 1:
noise_latents_btchw = pred_video_btchw
else:
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0,
1), noise_btchw.flatten(0, 1),
next_timestep).unflatten(0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(0, 2, 1, 3, 4)
else:
current_latents = pred_video_btchw.permute(0, 2, 1, 3, 4)
next_timestep = timesteps[i + 1] if i < len(timesteps) - 1 else None
current_latents, noise_latents_btchw = self._denoise_one_step(
current_latents=current_latents,
noise_latents_btchw=noise_latents_btchw,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
timestep=t_cur,
step_idx=i,
next_timestep=next_timestep,
ctx=ctx,
action_kwargs=action_kwargs,
noise_generator=noise_generator,
)
if progress_bar is not None:
progress_bar.update()
@@ -605,6 +716,36 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(ctx.target_dtype)
# Channel concatenation for models that don't handle it internally
if (batch.image_latent is not None
and not getattr(self.transformer, '_concatenates_image_latent',
False)):
image_cond = batch.image_latent.to(ctx.target_dtype)
end = start_index + current_num_frames
if image_cond.shape[2] >= end:
image_cond = image_cond[:, :, start_index:end]
elif image_cond.shape[2] > start_index:
image_cond = image_cond[:, :, start_index:]
pad_frames = current_num_frames - image_cond.shape[2]
if pad_frames > 0:
pad = torch.zeros(image_cond.shape[0],
image_cond.shape[1],
pad_frames,
image_cond.shape[3],
image_cond.shape[4],
device=image_cond.device,
dtype=image_cond.dtype)
image_cond = torch.cat([image_cond, pad], dim=2)
else:
image_cond = torch.zeros(image_cond.shape[0],
image_cond.shape[1],
current_num_frames,
image_cond.shape[3],
image_cond.shape[4],
device=image_cond.device,
dtype=image_cond.dtype)
context_bcthw = torch.cat([context_bcthw, image_cond], dim=1)
with torch.autocast(device_type="cuda",
dtype=ctx.target_dtype,
enabled=ctx.autocast_enabled), \
@@ -628,6 +769,37 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
})
context_model_kwargs.update(action_kwargs)
# Discrete action and camera matrices for WanGame-style models
if getattr(self.transformer, '_uses_discrete_action', False):
dev = latents_device
b = context_bcthw.shape[0]
vae_ratio = 4
start_pixel = (0 if start_index == 0 else 1 + vae_ratio *
(start_index - 1))
if (batch.keyboard_cond is not None
and start_pixel < batch.keyboard_cond.shape[1]):
kbd = batch.keyboard_cond[:, start_pixel, :].to(dev)
has_key = kbd.sum(dim=-1) > 0
trans_label = (kbd.argmax(dim=-1) + 1) * has_key.long()
else:
trans_label = torch.zeros(b, device=dev, dtype=torch.long)
action_value = trans_label * 9
context_model_kwargs["action"] = (
action_value.unsqueeze(1).expand(
-1, current_num_frames).to(dev))
context_model_kwargs["viewmats"] = torch.eye(
4, device=dev, dtype=ctx.target_dtype
).unsqueeze(0).unsqueeze(0).expand(
b, current_num_frames, -1, -1).contiguous()
default_K = torch.tensor(
[[0.5052, 0.0, 0.5],
[0.0, 0.8979, 0.5],
[0.0, 0.0, 1.0]],
device=dev, dtype=ctx.target_dtype)
context_model_kwargs["Ks"] = default_K.unsqueeze(
0).unsqueeze(0).expand(
b, current_num_frames, -1, -1).contiguous()
if ctx.boundary_timestep is not None and self.transformer_2 is not None:
self.transformer_2(
context_bcthw,
@@ -641,6 +813,10 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
**ctx.pos_cond_kwargs,
)
# WanGame uses is_cache=True to store KV without prepending old cache
if getattr(self.transformer, '_uses_discrete_action', False):
context_model_kwargs["is_cache"] = True
self.transformer(
context_bcthw,
prompt_embeds,
@@ -0,0 +1,695 @@
"""Multi-user engine implementing ORCA-style batching for streaming generation.
Groups users at the same (block_idx, denoising_step) and batches their forward
passes together, sharing model weights while maintaining per-user KV caches.
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
from typing import Any, Iterator
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.matrixgame_denoising import (
BlockProcessingContext,
MatrixGameCausalDenoisingStage,
)
logger = init_logger(__name__)
@dataclass
class CompletedResult:
"""Result for a user whose block generation finished."""
user_id: str
output_batch: ForwardBatch
@dataclass
class UserSession:
"""Per-user state for multi-user streaming."""
user_id: str
ctx: BlockProcessingContext
vae_cache: list[torch.Tensor | None] | None
batch: ForwardBatch
fastvideo_args: FastVideoArgs
# Denoising state (set when a step is in progress)
denoising_step: int = -1 # -1 means no pending work
current_latents: torch.Tensor | None = None
noise_latents_btchw: torch.Tensor | None = None
action_kwargs: dict[str, Any] = field(default_factory=dict)
block_start_index: int = 0
block_num_frames: int = 0
dit_elapsed_ms: float = 0.0 # accumulated DiT time across denoising steps
class MultiUserEngine:
"""ORCA-style scheduler that batches multiple users on a single GPU."""
def __init__(self, pipeline, disable_batching: bool = False):
self.pipeline = pipeline
self.disable_batching = disable_batching
if not pipeline.post_init_called:
pipeline.post_init()
self.denoiser: MatrixGameCausalDenoisingStage = (
pipeline._stage_name_mapping["denoising_stage"]
)
self.decoder = pipeline._stage_name_mapping["decoding_stage"]
self.users: dict[str, UserSession] = {}
def add_user(
self,
user_id: str,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> None:
"""Initialize a new user session (runs preprocessing + streaming_reset)."""
def _gpu_mem():
a = torch.cuda.memory_allocated() / 1024**3
r = torch.cuda.memory_reserved() / 1024**3
return f"alloc={a:.2f}GiB, reserved={r:.2f}GiB"
logger.info("add_user(%s) start: %s", user_id, _gpu_mem())
logger.info("add_user(%s) memory summary:\n%s",
user_id, torch.cuda.memory_summary())
if user_id in self.users:
logger.warning("User %s already exists, removing first", user_id)
self.remove_user(user_id)
# Free any existing streaming state FIRST to make room on GPU.
# After warmup, the denoiser holds KV caches + noise pool.
import gc
logger.info("add_user(%s) denoiser._streaming_initialized=%s, ctx=%s",
user_id, self.denoiser._streaming_initialized,
self.denoiser._streaming_ctx is not None)
if self.denoiser._streaming_initialized:
ctx = self.denoiser._streaming_ctx
if ctx is not None:
# Explicitly delete large tensors to break any reference cycles
ctx.kv_cache1 = None
ctx.kv_cache2 = None
ctx.kv_cache_mouse = None
ctx.kv_cache_keyboard = None
ctx.crossattn_cache = None
ctx.noise_pool = None
if ctx.batch is not None:
ctx.batch.latents = None
ctx.batch.image_latent = None
ctx.batch.prompt_embeds = None
ctx.batch = None
self.denoiser._streaming_ctx = None
self.denoiser._streaming_initialized = False
logger.info("add_user(%s) after ctx cleanup: %s", user_id, _gpu_mem())
gc.collect()
torch.cuda.empty_cache()
logger.info("add_user(%s) after gc+empty_cache: %s", user_id, _gpu_mem())
logger.info("add_user(%s) post-cleanup summary:\n%s",
user_id, torch.cuda.memory_summary())
# Run preprocessing stages
if not self.pipeline.post_init_called:
self.pipeline.post_init()
# Ensure VAE temporal tiling is enabled to avoid OOM when encoding
# the full video_condition [1,3,num_frames,H,W] through 3D conv layers.
# Enable on both the pipeline's VAE and the stage's VAE reference.
vae = self.pipeline.get_module("vae", None)
if vae is not None:
vae.enable_tiling(use_temporal_tiling=True)
logger.info("Pipeline VAE tiling: use_tiling=%s, temporal=%s, min_frames=%s, id=%s",
vae.use_tiling, vae.use_temporal_tiling,
getattr(vae, 'tile_sample_min_num_frames', 'N/A'), id(vae))
img_stage = self.pipeline._stage_name_mapping.get("image_latent_preparation_stage")
if img_stage is not None and hasattr(img_stage, 'vae'):
img_stage.vae.enable_tiling(use_temporal_tiling=True)
logger.info("Stage VAE tiling: use_tiling=%s, temporal=%s, min_frames=%s, id=%s",
img_stage.vae.use_tiling, img_stage.vae.use_temporal_tiling,
getattr(img_stage.vae, 'tile_sample_min_num_frames', 'N/A'),
id(img_stage.vae))
stages_to_run = [
"input_validation_stage", "prompt_encoding_stage",
"image_encoding_stage", "conditioning_stage",
"latent_preparation_stage", "image_latent_preparation_stage",
]
for stage_name in stages_to_run:
if stage_name in self.pipeline._stage_name_mapping:
batch = self.pipeline._stage_name_mapping[stage_name].forward(
batch, fastvideo_args)
logger.info("add_user(%s) after %s: %s",
user_id, stage_name, _gpu_mem())
# Initialize denoiser state for this user (creates KV caches etc.)
# We call streaming_reset to set up the context, then steal it
logger.info("add_user(%s) before streaming_reset: %s", user_id, _gpu_mem())
self.denoiser.streaming_reset(batch, fastvideo_args)
ctx = self.denoiser._streaming_ctx
assert ctx is not None
# Detach from denoiser so it doesn't interfere with other users
self.denoiser._streaming_ctx = None
self.denoiser._streaming_initialized = False
self.users[user_id] = UserSession(
user_id=user_id,
ctx=ctx,
vae_cache=None,
batch=batch,
fastvideo_args=fastvideo_args,
)
logger.info("Added user %s (total: %d)", user_id, len(self.users))
def remove_user(self, user_id: str) -> None:
"""Remove a user session and free GPU memory."""
session = self.users.pop(user_id, None)
if session is not None:
# Let GC handle tensor cleanup
session.ctx = None # type: ignore
session.vae_cache = None
session.current_latents = None
session.noise_latents_btchw = None
logger.info("Removed user %s (remaining: %d)",
user_id, len(self.users))
def submit_step(
self,
user_id: str,
keyboard_action: torch.Tensor | None,
mouse_action: torch.Tensor | None,
) -> None:
"""Queue a block generation request for a user."""
session = self.users.get(user_id)
if session is None:
raise KeyError(f"User {user_id} not found")
ctx = session.ctx
if ctx.block_idx >= len(ctx.block_sizes):
logger.warning("User %s has no more blocks to generate", user_id)
return
batch = session.batch
latents = batch.latents
assert latents is not None
current_num_frames = ctx.block_sizes[ctx.block_idx]
start_index = ctx.start_index
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
# Update batch with new actions
if keyboard_action is not None or mouse_action is not None:
vae_ratio = 4
start_frame = 0 if start_index == 0 else 1 + vae_ratio * (
start_index - 1)
if keyboard_action is not None:
n = keyboard_action.shape[1]
batch.keyboard_cond[:, start_frame:start_frame +
n] = keyboard_action.to(
batch.keyboard_cond.device)
if mouse_action is not None:
n = mouse_action.shape[1]
batch.mouse_cond[:, start_frame:start_frame +
n] = mouse_action.to(batch.mouse_cond.device)
action_kwargs = self.denoiser._prepare_action_kwargs(
batch, start_index, current_num_frames)
# Set up denoising state
session.denoising_step = 0
session.dit_elapsed_ms = 0.0
session.current_latents = current_latents
session.noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
session.action_kwargs = action_kwargs
session.block_start_index = start_index
session.block_num_frames = current_num_frames
def has_pending_work(self) -> bool:
"""Check if any user has pending denoising work."""
return any(s.denoising_step >= 0 for s in self.users.values())
def run_iteration(self) -> Iterator[CompletedResult]:
"""Run one denoising step for all ready users (ORCA iteration-level scheduling).
Groups users by (block_idx, denoising_step) and batches compatible groups.
Yields completed results immediately after each user's VAE decode finishes,
so responses can be sent without waiting for all users in the batch.
"""
# Group users with pending work by (block_idx, denoising_step)
groups: dict[tuple[int, int], list[UserSession]] = {}
for session in self.users.values():
if session.denoising_step < 0:
continue
key = (session.ctx.block_idx, session.denoising_step)
groups.setdefault(key, []).append(session)
for (block_idx, step_idx), sessions in groups.items():
timesteps = sessions[0].ctx.timesteps
t_cur = timesteps[step_idx]
next_timestep = timesteps[step_idx + 1] if step_idx < len(timesteps) - 1 else None
torch.cuda.synchronize()
t0 = time.perf_counter()
user_ids = [s.user_id[:8] for s in sessions]
if len(sessions) == 1 or self.disable_batching:
# Process each user individually
for s in sessions:
logger.info("ORCA step: block=%d step=%d single user=%s",
block_idx, step_idx, s.user_id[:8])
self._run_single_user_step(s, t_cur, step_idx, next_timestep)
else:
# Multiple users at same state - batch them
logger.info("ORCA step: block=%d step=%d BATCHED %d users=%s",
block_idx, step_idx, len(sessions), user_ids)
self._run_batched_step(sessions, t_cur, step_idx, next_timestep)
torch.cuda.synchronize()
step_ms = (time.perf_counter() - t0) * 1000
# Advance step counters and finalize completed users
for s in sessions:
s.dit_elapsed_ms += step_ms
s.denoising_step += 1
if s.denoising_step >= len(timesteps):
# Denoising complete — VAE decode and yield immediately
result = self._finalize_block(s)
if result is not None:
yield result
def _run_single_user_step(
self,
session: UserSession,
timestep: torch.Tensor,
step_idx: int,
next_timestep: torch.Tensor | None,
) -> None:
"""Run one denoising step for a single user (no batching)."""
noise_gen = self._make_noise_generator(session)
current_latents, noise_latents_btchw = self.denoiser._denoise_one_step(
current_latents=session.current_latents,
noise_latents_btchw=session.noise_latents_btchw,
batch=session.batch,
start_index=session.block_start_index,
current_num_frames=session.block_num_frames,
timestep=timestep,
step_idx=step_idx,
next_timestep=next_timestep,
ctx=session.ctx,
action_kwargs=session.action_kwargs,
noise_generator=noise_gen,
)
session.current_latents = current_latents
session.noise_latents_btchw = noise_latents_btchw
def _run_batched_step(
self,
sessions: list[UserSession],
timestep: torch.Tensor,
step_idx: int,
next_timestep: torch.Tensor | None,
) -> None:
"""Run one denoising step for multiple users batched together."""
n = len(sessions)
# Concatenate along batch dimension
batched_latents = torch.cat(
[s.current_latents for s in sessions], dim=0)
batched_noise = torch.cat(
[s.noise_latents_btchw for s in sessions], dim=0)
# Build a merged ForwardBatch with concatenated per-user tensors
merged_batch = self._merge_batches(sessions)
# Build merged context with concatenated KV caches
merged_ctx = self._merge_contexts(sessions, merged_batch)
# Merge action kwargs
merged_action_kwargs = self._merge_action_kwargs(sessions)
# Create noise generator that concatenates per-user noise
device = batched_latents.device
def batched_noise_gen(shape, dtype, si):
noises = []
for s in sessions:
gen = self._make_noise_generator(s)
if gen is not None:
per_user_shape = (1,) + shape[1:]
noises.append(gen(per_user_shape, dtype, si))
else:
noises.append(torch.randn(
(1,) + shape[1:], dtype=dtype).to(device))
return torch.cat(noises, dim=0)
current_latents, noise_latents_btchw = self.denoiser._denoise_one_step(
current_latents=batched_latents,
noise_latents_btchw=batched_noise,
batch=merged_batch,
start_index=sessions[0].block_start_index,
current_num_frames=sessions[0].block_num_frames,
timestep=timestep,
step_idx=step_idx,
next_timestep=next_timestep,
ctx=merged_ctx,
action_kwargs=merged_action_kwargs,
noise_generator=batched_noise_gen,
)
# Unbatch results back to per-user tensors
latents_list = current_latents.split(1, dim=0)
noise_list = noise_latents_btchw.split(1, dim=0)
# Restore per-user KV caches from the batched caches
self._split_contexts(merged_ctx, sessions)
for i, s in enumerate(sessions):
s.current_latents = latents_list[i]
s.noise_latents_btchw = noise_list[i]
def _merge_batches(self, sessions: list[UserSession]) -> ForwardBatch:
"""Create a merged ForwardBatch with concatenated user tensors."""
ref = sessions[0].batch
# prompt_embeds and image_embeds are shared (same game)
# latents, keyboard_cond, mouse_cond differ per user
merged = ForwardBatch.__new__(ForwardBatch)
merged.__dict__.update(ref.__dict__)
# image_latent: cat along batch dim if present
if ref.image_latent is not None:
merged.image_latent = torch.cat(
[s.batch.image_latent for s in sessions], dim=0)
return merged
def _merge_contexts(
self, sessions: list[UserSession],
merged_batch: ForwardBatch | None = None,
) -> BlockProcessingContext:
"""Create a merged BlockProcessingContext with concatenated KV caches."""
ref = sessions[0].ctx
n = len(sessions)
# Concatenate KV caches along batch dimension
merged_kv1 = self._cat_kv_caches([s.ctx.kv_cache1 for s in sessions])
merged_kv2 = None
if ref.kv_cache2 is not None:
merged_kv2 = self._cat_kv_caches(
[s.ctx.kv_cache2 for s in sessions])
merged_crossattn = self._cat_crossattn_caches(
[s.ctx.crossattn_cache for s in sessions])
merged_kv_mouse = None
merged_kv_keyboard = None
if ref.kv_cache_mouse is not None:
merged_kv_mouse = self._cat_action_kv_caches(
[s.ctx.kv_cache_mouse for s in sessions],
is_mouse=True)
if ref.kv_cache_keyboard is not None:
merged_kv_keyboard = self._cat_kv_caches(
[s.ctx.kv_cache_keyboard for s in sessions])
# image_kwargs: cat image embeds
merged_image_kwargs = dict(ref.image_kwargs)
if "encoder_hidden_states_image" in ref.image_kwargs:
embeds = ref.image_kwargs["encoder_hidden_states_image"]
if isinstance(embeds, list) and len(embeds) > 0 and torch.is_tensor(embeds[0]):
# Cat each tensor in the list across users
merged_embeds = []
for idx in range(len(embeds)):
merged_embeds.append(torch.cat(
[s.ctx.image_kwargs["encoder_hidden_states_image"][idx]
for s in sessions], dim=0))
merged_image_kwargs["encoder_hidden_states_image"] = merged_embeds
if merged_batch is None:
merged_batch = self._merge_batches(sessions)
return BlockProcessingContext(
batch=merged_batch,
block_idx=ref.block_idx,
start_index=ref.start_index,
kv_cache1=merged_kv1,
kv_cache2=merged_kv2,
kv_cache_mouse=merged_kv_mouse,
kv_cache_keyboard=merged_kv_keyboard,
crossattn_cache=merged_crossattn,
timesteps=ref.timesteps,
block_sizes=ref.block_sizes,
noise_pool=None, # noise handled per-user
fastvideo_args=ref.fastvideo_args,
target_dtype=ref.target_dtype,
autocast_enabled=ref.autocast_enabled,
boundary_timestep=ref.boundary_timestep,
high_noise_timesteps=ref.high_noise_timesteps,
context_noise=ref.context_noise,
image_kwargs=merged_image_kwargs,
pos_cond_kwargs=ref.pos_cond_kwargs,
)
def _cat_kv_caches(
self, caches_list: list[list[dict[str, Any]]]
) -> list[dict[str, Any]]:
"""Concatenate KV caches along batch dimension."""
num_layers = len(caches_list[0])
merged = []
for layer_idx in range(num_layers):
merged.append({
"k": torch.cat(
[c[layer_idx]["k"] for c in caches_list], dim=0),
"v": torch.cat(
[c[layer_idx]["v"] for c in caches_list], dim=0),
"global_end_index": caches_list[0][layer_idx]["global_end_index"],
"local_end_index": caches_list[0][layer_idx]["local_end_index"],
})
return merged
def _cat_action_kv_caches(
self,
caches_list: list[list[dict[str, Any]]],
is_mouse: bool = False,
) -> list[dict[str, Any]]:
"""Concatenate action KV caches.
Mouse caches have shape [B*frame_seq_length, ...] so we need
to concatenate correctly.
"""
num_layers = len(caches_list[0])
merged = []
for layer_idx in range(num_layers):
if is_mouse:
# Mouse: [B*F, cache_size, heads, dim] -> [N*F, cache_size, heads, dim]
# Each user has [1*F, ...], so simple cat along dim 0 works
merged.append({
"k": torch.cat(
[c[layer_idx]["k"] for c in caches_list], dim=0),
"v": torch.cat(
[c[layer_idx]["v"] for c in caches_list], dim=0),
"global_end_index": caches_list[0][layer_idx]["global_end_index"],
"local_end_index": caches_list[0][layer_idx]["local_end_index"],
})
else:
merged.append({
"k": torch.cat(
[c[layer_idx]["k"] for c in caches_list], dim=0),
"v": torch.cat(
[c[layer_idx]["v"] for c in caches_list], dim=0),
"global_end_index": caches_list[0][layer_idx]["global_end_index"],
"local_end_index": caches_list[0][layer_idx]["local_end_index"],
})
return merged
def _cat_crossattn_caches(
self, caches_list: list[list[dict[str, Any]]]
) -> list[dict[str, Any]]:
"""Concatenate cross-attention caches along batch dimension."""
num_layers = len(caches_list[0])
merged = []
for layer_idx in range(num_layers):
merged.append({
"k": torch.cat(
[c[layer_idx]["k"] for c in caches_list], dim=0),
"v": torch.cat(
[c[layer_idx]["v"] for c in caches_list], dim=0),
"is_init": caches_list[0][layer_idx]["is_init"],
})
return merged
def _split_contexts(
self,
merged_ctx: BlockProcessingContext,
sessions: list[UserSession],
) -> None:
"""Split merged KV caches back into per-user caches after forward pass."""
n = len(sessions)
self._split_kv_caches(merged_ctx.kv_cache1,
[s.ctx.kv_cache1 for s in sessions])
if merged_ctx.kv_cache2 is not None:
self._split_kv_caches(merged_ctx.kv_cache2,
[s.ctx.kv_cache2 for s in sessions])
self._split_crossattn_caches(merged_ctx.crossattn_cache,
[s.ctx.crossattn_cache for s in sessions])
if merged_ctx.kv_cache_mouse is not None:
frame_seq_len = self.denoiser.frame_seq_length
self._split_action_kv_caches(
merged_ctx.kv_cache_mouse,
[s.ctx.kv_cache_mouse for s in sessions],
is_mouse=True, frame_seq_length=frame_seq_len)
if merged_ctx.kv_cache_keyboard is not None:
self._split_kv_caches(merged_ctx.kv_cache_keyboard,
[s.ctx.kv_cache_keyboard for s in sessions])
def _split_kv_caches(
self,
merged: list[dict[str, Any]],
per_user: list[list[dict[str, Any]]],
) -> None:
"""Copy data from merged KV cache back to per-user caches."""
n = len(per_user)
for layer_idx in range(len(merged)):
k_split = merged[layer_idx]["k"].split(1, dim=0)
v_split = merged[layer_idx]["v"].split(1, dim=0)
for i in range(n):
per_user[i][layer_idx]["k"].copy_(k_split[i])
per_user[i][layer_idx]["v"].copy_(v_split[i])
per_user[i][layer_idx]["global_end_index"].copy_(
merged[layer_idx]["global_end_index"])
per_user[i][layer_idx]["local_end_index"].copy_(
merged[layer_idx]["local_end_index"])
def _split_action_kv_caches(
self,
merged: list[dict[str, Any]],
per_user: list[list[dict[str, Any]]],
is_mouse: bool = False,
frame_seq_length: int = 1,
) -> None:
"""Split action KV caches back to per-user."""
n = len(per_user)
for layer_idx in range(len(merged)):
if is_mouse:
# [N*F, ...] -> N chunks of [F, ...]
k_split = merged[layer_idx]["k"].split(frame_seq_length, dim=0)
v_split = merged[layer_idx]["v"].split(frame_seq_length, dim=0)
else:
k_split = merged[layer_idx]["k"].split(1, dim=0)
v_split = merged[layer_idx]["v"].split(1, dim=0)
for i in range(n):
per_user[i][layer_idx]["k"].copy_(k_split[i])
per_user[i][layer_idx]["v"].copy_(v_split[i])
per_user[i][layer_idx]["global_end_index"].copy_(
merged[layer_idx]["global_end_index"])
per_user[i][layer_idx]["local_end_index"].copy_(
merged[layer_idx]["local_end_index"])
def _split_crossattn_caches(
self,
merged: list[dict[str, Any]],
per_user: list[list[dict[str, Any]]],
) -> None:
"""Split cross-attention caches back to per-user."""
n = len(per_user)
for layer_idx in range(len(merged)):
k_split = merged[layer_idx]["k"].split(1, dim=0)
v_split = merged[layer_idx]["v"].split(1, dim=0)
for i in range(n):
per_user[i][layer_idx]["k"].copy_(k_split[i])
per_user[i][layer_idx]["v"].copy_(v_split[i])
per_user[i][layer_idx]["is_init"] = merged[layer_idx]["is_init"]
def _merge_action_kwargs(
self, sessions: list[UserSession]
) -> dict[str, Any]:
"""Merge action kwargs by concatenating tensors along batch dim."""
ref = sessions[0].action_kwargs
if not ref:
return {}
merged: dict[str, Any] = {}
for key in ref:
vals = [s.action_kwargs[key] for s in sessions]
if torch.is_tensor(vals[0]):
merged[key] = torch.cat(vals, dim=0)
else:
# Scalars (like num_frame_per_block) - use ref value
merged[key] = vals[0]
return merged
def _make_noise_generator(self, session: UserSession):
"""Create a noise generator using the user's pre-allocated noise pool."""
ctx = session.ctx
latents = session.batch.latents
def noise_gen(shape, dtype, step_idx):
if ctx.noise_pool is not None and step_idx < len(ctx.noise_pool):
return ctx.noise_pool[step_idx][:, :shape[1], :, :, :].to(
latents.device)
else:
return torch.randn(shape, dtype=dtype).to(latents.device)
return noise_gen
def _finalize_block(self, session: UserSession) -> CompletedResult | None:
"""Finalize a completed block: update context cache + VAE decode."""
ctx = session.ctx
batch = session.batch
latents = batch.latents
start_index = session.block_start_index
current_num_frames = session.block_num_frames
# Write denoised latents back
latents[:, :, start_index:start_index +
current_num_frames, :, :] = session.current_latents
# Update KV caches with clean context
self.denoiser._update_context_cache(
current_latents=session.current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=session.action_kwargs,
context_noise=ctx.context_noise,
)
# Advance streaming state
old_start = ctx.start_index
ctx.start_index += current_num_frames
ctx.block_idx += 1
# VAE decode
torch.cuda.synchronize()
t0 = time.perf_counter()
current_latents_for_vae = latents[:, :, start_index:start_index +
current_num_frames, :, :]
decoded_frames, session.vae_cache = self.decoder.streaming_decode(
current_latents_for_vae,
session.fastvideo_args,
cache=session.vae_cache,
is_first_chunk=(start_index == 0),
)
torch.cuda.synchronize()
vae_ms = (time.perf_counter() - t0) * 1000
batch.output = decoded_frames
batch.stage_timings = {"dit_ms": session.dit_elapsed_ms, "vae_ms": vae_ms}
# Reset denoising state
session.denoising_step = -1
session.current_latents = None
session.noise_latents_btchw = None
return CompletedResult(user_id=session.user_id, output_batch=batch)
+4
View File
@@ -75,6 +75,10 @@ class Worker:
self.pipeline = build_pipeline(self.fastvideo_args)
def get_pipeline(self):
"""Return the loaded pipeline instance."""
return self.pipeline
def execute_forward(self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args)
+172 -1
View File
@@ -39,13 +39,17 @@ logger = init_logger(__name__)
class StreamingTaskType(str, Enum):
"""
Enumeration for different streaming task types.
Inherits from str to allow string comparison for backward compatibility.
"""
RESET = "reset"
STEP = "step"
CLEAR = "clear"
EXIT = "exit"
# Multi-user task types
USER_JOIN = "user_join"
USER_STEP = "user_step"
USER_LEAVE = "user_leave"
@dataclass
@@ -58,6 +62,8 @@ class StreamingTask:
# For RESET tasks:
batch: ForwardBatch | None = None
fastvideo_args: FastVideoArgs | None = None
# For multi-user tasks:
user_id: str | None = None
@dataclass
@@ -66,6 +72,8 @@ class StreamingResult:
task_type: StreamingTaskType
output_batch: ForwardBatch | None = None
error: Exception | None = None
# For multi-user results:
user_id: str | None = None
class MultiprocExecutor(Executor):
@@ -183,6 +191,13 @@ class MultiprocExecutor(Executor):
self.collective_rpc("start_streaming_queue_loop")
self._streaming_enabled = True
def enable_multi_user_streaming(self) -> None:
if self._streaming_enabled:
return
self.collective_rpc("start_multi_user_streaming_loop")
self._streaming_enabled = True
def disable_streaming(self) -> None:
if not self._streaming_enabled:
return
@@ -225,6 +240,42 @@ class MultiprocExecutor(Executor):
self._streaming_input_queue.put(
StreamingTask(task_type=StreamingTaskType.CLEAR))
def submit_user_join(self, user_id: str, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> None:
if not self._streaming_enabled:
self.enable_streaming()
self._streaming_input_queue.put(
StreamingTask(
task_type=StreamingTaskType.USER_JOIN,
user_id=user_id,
batch=forward_batch,
fastvideo_args=fastvideo_args,
))
def submit_user_step(self, user_id: str,
keyboard_action: torch.Tensor | None,
mouse_action: torch.Tensor | None) -> None:
if not self._streaming_enabled:
raise RuntimeError(
"Streaming mode not enabled. Call enable_streaming() first.")
self._streaming_input_queue.put(
StreamingTask(
task_type=StreamingTaskType.USER_STEP,
user_id=user_id,
keyboard_action=keyboard_action,
mouse_action=mouse_action,
))
def submit_user_leave(self, user_id: str) -> None:
if self._streaming_enabled and self._streaming_input_queue is not None:
self._streaming_input_queue.put(
StreamingTask(
task_type=StreamingTaskType.USER_LEAVE,
user_id=user_id,
))
def get_result(self,
timeout: float | None = None) -> StreamingResult | None:
if not self._streaming_enabled or self._streaming_output_queue is None:
@@ -642,6 +693,11 @@ class WorkerMultiprocProc:
{"status": "streaming_queue_loop_started"})
self.streaming_queue_loop()
continue
if method == "start_multi_user_streaming_loop":
self.pipe.send(
{"status": "multi_user_streaming_loop_started"})
self.multi_user_streaming_loop()
continue
if method == 'execute_forward':
forward_batch = kwargs['forward_batch']
fastvideo_args = kwargs['fastvideo_args']
@@ -715,11 +771,126 @@ class WorkerMultiprocProc:
self.worker.execute_streaming_clear()
self.streaming_output_queue.put(
StreamingResult(task_type=StreamingTaskType.CLEAR))
elif task.task_type in (StreamingTaskType.USER_JOIN,
StreamingTaskType.USER_STEP,
StreamingTaskType.USER_LEAVE):
# Switch to multi-user mode
self.multi_user_streaming_loop(first_task=task)
break
except Exception as e:
logger.error("Worker %d queue loop error: %s", self.rank, e)
self.streaming_output_queue.put(
StreamingResult(task_type=StreamingTaskType.STEP, error=e))
def multi_user_streaming_loop(
self, first_task: StreamingTask | None = None
) -> None:
"""ORCA-style multi-user streaming loop.
Processes USER_JOIN/USER_STEP/USER_LEAVE tasks and batches
compatible users together for each denoising step.
"""
from fastvideo.pipelines.stages.multi_user_engine import MultiUserEngine
if self.streaming_input_queue is None or self.streaming_output_queue is None:
logger.error("Worker %d: streaming queues not initialized",
self.rank)
return
pipeline = self.worker.get_pipeline()
try:
from ui.world_model.server.config import DISABLE_BATCHING
except ImportError:
DISABLE_BATCHING = False
engine = MultiUserEngine(pipeline, disable_batching=DISABLE_BATCHING)
logger.info("Worker %d: multi-user streaming loop started", self.rank)
def handle_task(task: StreamingTask) -> bool:
"""Handle a single task. Returns False if should exit."""
if task.task_type == StreamingTaskType.EXIT:
return False
elif task.task_type == StreamingTaskType.USER_JOIN:
try:
engine.add_user(
task.user_id, task.batch, task.fastvideo_args)
self.streaming_output_queue.put(
StreamingResult(
task_type=StreamingTaskType.USER_JOIN,
user_id=task.user_id))
except Exception as e:
logger.error("Worker %d user_join error for %s: %s",
self.rank, task.user_id, e)
self.streaming_output_queue.put(
StreamingResult(
task_type=StreamingTaskType.USER_JOIN,
user_id=task.user_id, error=e))
elif task.task_type == StreamingTaskType.USER_STEP:
try:
engine.submit_step(
task.user_id, task.keyboard_action, task.mouse_action)
except Exception as e:
logger.error("Worker %d user_step submit error for %s: %s",
self.rank, task.user_id, e)
self.streaming_output_queue.put(
StreamingResult(
task_type=StreamingTaskType.USER_STEP,
user_id=task.user_id, error=e))
elif task.task_type == StreamingTaskType.USER_LEAVE:
engine.remove_user(task.user_id)
self.streaming_output_queue.put(
StreamingResult(
task_type=StreamingTaskType.USER_LEAVE,
user_id=task.user_id))
return True
# Handle the first task that triggered the switch
if first_task is not None:
if not handle_task(first_task):
return
while True:
# 1. Drain input queue (non-blocking) for new requests
while True:
try:
task = self.streaming_input_queue.get_nowait()
if not handle_task(task):
return
except queue.Empty:
break
# 2. Run one ORCA iteration
if engine.has_pending_work():
try:
completed = engine.run_iteration()
for result in completed:
self.streaming_output_queue.put(
StreamingResult(
task_type=StreamingTaskType.USER_STEP,
user_id=result.user_id,
output_batch=result.output_batch))
except Exception as e:
logger.error("Worker %d ORCA iteration error: %s",
self.rank, e, exc_info=True)
# Reset all pending users to prevent infinite error loop
for uid, sess in engine.users.items():
if sess.denoising_step >= 0:
sess.denoising_step = -1
sess.current_latents = None
sess.noise_latents_btchw = None
self.streaming_output_queue.put(
StreamingResult(
task_type=StreamingTaskType.USER_STEP,
user_id=uid, error=e))
# 3. If no work, block on queue to avoid busy-waiting
if not engine.has_pending_work():
try:
task = self.streaming_input_queue.get(timeout=0.1)
if not handle_task(task):
return
except queue.Empty:
continue
@staticmethod
def setup_proc_title_and_log_prefix() -> None:
dp_size = get_dp_group().world_size
+12
View File
@@ -0,0 +1,12 @@
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>World Model Interface</title>
</head>
<body>
<div id="app"></div>
<script type="module" src="/src/main.js"></script>
</body>
</html>
File diff suppressed because it is too large Load Diff
+16
View File
@@ -0,0 +1,16 @@
{
"name": "wm-interface-frontend",
"private": true,
"version": "0.0.1",
"type": "module",
"scripts": {
"dev": "vite",
"build": "vite build",
"preview": "vite preview"
},
"devDependencies": {
"@sveltejs/vite-plugin-svelte": "^3.0.1",
"svelte": "^4.2.8",
"vite": "^5.0.11"
}
}
+18
View File
@@ -0,0 +1,18 @@
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 5.7 KiB

File diff suppressed because it is too large Load Diff
+8
View File
@@ -0,0 +1,8 @@
import App from './App.svelte'
import './style.css'
const app = new App({
target: document.getElementById('app')
})
export default app
+20
View File
@@ -0,0 +1,20 @@
* {
margin: 0;
padding: 0;
box-sizing: border-box;
}
body {
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, Oxygen, Ubuntu, Cantarell, sans-serif;
background: #1a1a1a;
color: white;
min-height: 100vh;
overflow-y: auto;
}
#app {
width: 100%;
min-height: 100vh;
display: flex;
flex-direction: column;
}
+18
View File
@@ -0,0 +1,18 @@
import { defineConfig } from 'vite'
import { svelte } from '@sveltejs/vite-plugin-svelte'
export default defineConfig({
plugins: [svelte()],
server: {
port: 5173,
proxy: {
'/ws': {
target: 'http://localhost:8000',
ws: true,
},
'/status': {
target: 'http://localhost:8000',
},
},
}
})
+61
View File
@@ -0,0 +1,61 @@
from pathlib import Path
# Repo root: ui/world_model/server/config.py -> go up 3 levels
_REPO_ROOT = Path(__file__).resolve().parents[3]
# Model registry - add new models here
MODEL_REGISTRY = {
"matrix-game-2.0-base": {
"name": "Matrix-Game 2.0 Base",
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
"keyboard_dim": 4,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
},
"wangame-1.3b-mc-w-only-still-7k": {
"name": "WANGame 1.3B MC (W-Only, 7k)",
"model_path": "weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers",
"init_weights_from_safetensors": str(_REPO_ROOT / "checkpoints" / "wangame-1.3b-mc-w-only-still-7k" / "checkpoint-7000" / "transformer"),
"override_transformer_cls_name": "WanGameActionTransformer3DModel",
"override_pipeline_cls_name": "WanGameCausalDMDPipeline",
"keyboard_dim": 4,
#"image_url": "/server-assets/mc.png",
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
# "image_path": str(Path(__file__).resolve().parent / "mc.png"),
},
}
DEFAULT_MODEL_ID = "matrix-game-2.0-base"
# Active model configuration (set by server or user selection)
MODEL_CONFIG = MODEL_REGISTRY[DEFAULT_MODEL_ID]
# Keyboard mappings (WASD)
KEYBOARD_MAP = {
"w": [1, 0, 0, 0],
"s": [0, 1, 0, 0],
"a": [0, 0, 1, 0],
"d": [0, 0, 0, 1],
}
# Camera mappings (Arrow keys)
CAM_VALUE = 0.1
CAMERA_MAP = {
"ArrowUp": [CAM_VALUE, 0],
"ArrowDown": [-CAM_VALUE, 0],
"ArrowLeft": [0, -CAM_VALUE],
"ArrowRight": [0, CAM_VALUE],
}
# Generation limits
MAX_BLOCKS = 50
SESSION_TIMEOUT_SECONDS = 900
MAX_USERS_PER_GPU = 16
DISABLE_BATCHING = False # Process users sequentially (batch_size=1) instead of batching
# Frame settings
NUM_FRAMES = 597
FRAME_HEIGHT = 352
FRAME_WIDTH = 640
NUM_INFERENCE_STEPS = 4
JPEG_QUALITY = 85
BATCH_SIZE = 12
File diff suppressed because it is too large Load Diff
+262
View File
@@ -0,0 +1,262 @@
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from contextlib import asynccontextmanager
import asyncio
import base64
import cv2
import os
import time
import uuid
from config import (
KEYBOARD_MAP, CAMERA_MAP, MAX_BLOCKS, JPEG_QUALITY, BATCH_SIZE,
SESSION_TIMEOUT_SECONDS, MODEL_CONFIG, MODEL_REGISTRY, DEFAULT_MODEL_ID
)
from gpu_pool import GPUPool, GPUSlot, get_available_gpus
# Global GPU pool
gpu_pool: GPUPool = None
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Application lifespan manager."""
global gpu_pool
print("Starting server...")
# Get available GPUs
gpu_ids = get_available_gpus()
print(f"Detected GPUs: {gpu_ids}")
# Initialize GPU pool (spawns subprocess per GPU)
gpu_pool = GPUPool(gpu_ids)
await gpu_pool.initialize()
print("Server ready")
yield
print("Shutting down server...")
await gpu_pool.shutdown()
app = FastAPI(lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.get("/status")
async def get_status():
"""Get the current status of the GPU pool."""
if gpu_pool is None:
return {"error": "GPU pool not initialized"}
return gpu_pool.get_status()
@app.get("/models")
async def get_models():
"""Get available models and the currently active model."""
models = []
for model_id, config in MODEL_REGISTRY.items():
models.append({
"id": model_id,
"name": config["name"],
})
return {
"models": models,
"default_model_id": DEFAULT_MODEL_ID,
}
def encode_frames(frames: list) -> list[str]:
"""Encode frames to base64 JPEG."""
encoded = []
for frame in frames:
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
_, buffer = cv2.imencode('.jpg', frame_rgb, [cv2.IMWRITE_JPEG_QUALITY, JPEG_QUALITY])
encoded.append(base64.b64encode(buffer.tobytes()).decode("utf-8"))
return encoded
async def send_frames(websocket: WebSocket, frames: list):
"""Encode and send frames in batches."""
if not frames:
return
encoded_frames = encode_frames(frames)
for i in range(0, len(encoded_frames), BATCH_SIZE):
batch = encoded_frames[i:i + BATCH_SIZE]
await websocket.send_json({"type": "frame_batch", "frames": batch})
@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
await websocket.accept()
client_id = str(uuid.uuid4())
print(f"Client {client_id[:8]} connected")
# Send initial queue status
status = gpu_pool.get_status()
await websocket.send_json({
"type": "queue_status",
"position": status["queue_size"] + 1 if status["available_gpus"] == 0 else 0,
"total_gpus": status["total_gpus"],
"available_gpus": status["available_gpus"]
})
gpu_id = None
slot: GPUSlot = None
timeout_task: asyncio.Task = None
async def session_timeout():
"""Close the session after timeout."""
await asyncio.sleep(SESSION_TIMEOUT_SECONDS)
print(f"[GPU {gpu_id}] Session timeout for client {client_id[:8]}")
try:
await websocket.send_json({
"type": "session_timeout",
"message": f"Session expired after {SESSION_TIMEOUT_SECONDS} seconds"
})
await websocket.close(code=1000, reason="Session timeout")
except Exception:
pass # WebSocket may already be closed
try:
# Wait for model selection from client
model_id = DEFAULT_MODEL_ID
model_config = MODEL_CONFIG
try:
init_data = await asyncio.wait_for(websocket.receive_json(), timeout=5.0)
if init_data.get("type") == "select_model":
model_id = init_data.get("model_id", DEFAULT_MODEL_ID)
if model_id in MODEL_REGISTRY:
model_config = MODEL_REGISTRY[model_id]
print(f"Client {client_id[:8]} selected model: {model_id}")
except asyncio.TimeoutError:
pass # Use default model config
# Acquire a GPU slot (may wait in queue, may share with other users)
gpu_id, slot = await gpu_pool.acquire(client_id, websocket)
# Start session timeout
timeout_task = asyncio.create_task(session_timeout())
# Join the multi-user engine on this GPU (triggers reload if model differs)
await slot.join_user(client_id, model_id=model_id)
# Notify client they're connected to a GPU
await websocket.send_json({
"type": "gpu_assigned",
"gpu_id": gpu_id,
"session_timeout": SESSION_TIMEOUT_SECONDS,
"image_url": model_config.get("image_url")
})
# Generate and send initial frame
block_count = 0
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
try:
frames, timings = await slot.user_step(
client_id, [0, 0, 0, 0], [0, 0])
block_count = 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
await send_frames(websocket, frames)
if timings:
print(f"[GPU {gpu_id}] Initial ({client_id[:8]}): model={timings.get('model_step_ms', 0):.0f}ms")
except Exception as e:
print(f"[GPU {gpu_id}] Initial frame error for {client_id[:8]}: {e}")
# Main event loop
while True:
data = await websocket.receive_json()
message_type = data.get("type", "key")
if message_type == "reset":
print(f"[GPU {gpu_id}] Reset requested by client {client_id[:8]}")
await websocket.send_json({"type": "reset_started"})
try:
# Leave and rejoin to reset user state
await slot.leave_user(client_id)
await slot.join_user(client_id, model_id=model_id)
frames, timings = await slot.user_step(
client_id, [0, 0, 0, 0], [0, 0])
block_count = 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
await send_frames(websocket, frames)
await websocket.send_json({"type": "reset_complete"})
print(f"[GPU {gpu_id}] Reset complete for {client_id[:8]}")
except Exception as e:
print(f"[GPU {gpu_id}] Reset error for {client_id[:8]}: {e}")
await websocket.send_json({"type": "reset_complete"})
continue
# Handle key press
key = data.get("key")
if key and block_count < MAX_BLOCKS:
if key in CAMERA_MAP:
keyboard_vector = [0, 0, 0, 0]
mouse_vector = CAMERA_MAP[key]
else:
keyboard_vector = KEYBOARD_MAP.get(key, [0, 0, 0, 0])
mouse_vector = [0, 0]
print(f"[GPU {gpu_id}] Key '{key}' from {client_id[:8]}, generating block {block_count + 1}...")
t_start = time.time()
frames, timings = await slot.user_step(
client_id, keyboard_vector, mouse_vector)
t_generation = time.time() - t_start
block_count += 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
if frames:
t_encode_start = time.time()
await send_frames(websocket, frames)
t_encode = time.time() - t_encode_start
dit_ms = timings.get('dit_ms', 0)
vae_ms = timings.get('vae_ms', 0)
print(f"[GPU {gpu_id}] ({client_id[:8]}) DiT: {dit_ms:.0f}ms, VAE: {vae_ms:.0f}ms, Encode+Send: {t_encode*1000:.0f}ms, Total: {t_generation*1000:.0f}ms")
except WebSocketDisconnect:
print(f"Client {client_id[:8]} disconnected")
except Exception as e:
print(f"Client {client_id[:8]} error: {e}")
import traceback
traceback.print_exc()
finally:
# Cancel the timeout task if still running
if timeout_task and not timeout_task.done():
timeout_task.cancel()
try:
await timeout_task
except asyncio.CancelledError:
pass
if client_id:
await gpu_pool.release(client_id)
# Serve server-side assets (images, etc.)
server_dir = os.path.dirname(os.path.abspath(__file__))
app.mount("/server-assets", StaticFiles(directory=server_dir), name="server-assets")
# Serve built frontend (must be after API/WebSocket routes)
static_dir = os.path.join(server_dir, "..", "client", "dist")
if os.path.isdir(static_dir):
app.mount("/", StaticFiles(directory=static_dir, html=True), name="static")
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8001)
+332
View File
@@ -0,0 +1,332 @@
"""
Mock server to test latency perception.
Generates frames with a moving object based on keyboard input.
Configurable latency to find the acceptable threshold.
"""
import asyncio
import base64
import io
import time
import uuid
from contextlib import asynccontextmanager
import numpy as np
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from PIL import Image, ImageDraw
import os
# Configuration
LATENCY_MS = 200 # Simulated model latency - adjust this to test
FRAME_WIDTH = 640
FRAME_HEIGHT = 352
NUM_FRAMES = 12 # Frames per block
FPS = 24
BATCH_SIZE = 4
JPEG_QUALITY = 85
SESSION_TIMEOUT_SECONDS = 90
MAX_BLOCKS = 50
# Mock game state
class GameState:
def __init__(self):
self.reset()
def reset(self):
self.x = FRAME_WIDTH // 2
self.y = FRAME_HEIGHT // 2
self.size = 40
self.color = (100, 200, 255) # Light blue
self.trail = [] # Previous positions for motion blur effect
def update(self, keyboard_vector, mouse_vector):
"""Store the movement direction for this block."""
self.keyboard_vector = keyboard_vector
self.mouse_vector = mouse_vector
# Arrow keys for camera (just change color as visual feedback)
if mouse_vector[0] != 0 or mouse_vector[1] != 0:
r = int(100 + mouse_vector[0] * 50) % 256
g = int(200 + mouse_vector[1] * 50) % 256
self.color = (r, g, 255)
def render_frame(self, frame_idx: int) -> np.ndarray:
"""Render a single frame with incremental movement."""
speed_per_frame = 8
# Apply movement for this frame
if hasattr(self, 'keyboard_vector'):
if self.keyboard_vector[0]: # W - up
self.y -= speed_per_frame
if self.keyboard_vector[1]: # A - left
self.x -= speed_per_frame
if self.keyboard_vector[2]: # S - down
self.y += speed_per_frame
if self.keyboard_vector[3]: # D - right
self.x += speed_per_frame
# Wrap around screen
self.x = self.x % FRAME_WIDTH
self.y = self.y % FRAME_HEIGHT
# Update trail
self.trail.append((self.x, self.y))
if len(self.trail) > 5:
self.trail.pop(0)
"""Render a single frame."""
# Create frame with gradient background
img = Image.new('RGB', (FRAME_WIDTH, FRAME_HEIGHT), (20, 20, 30))
draw = ImageDraw.Draw(img)
# Draw grid for visual reference
for x in range(0, FRAME_WIDTH, 40):
draw.line([(x, 0), (x, FRAME_HEIGHT)], fill=(40, 40, 50), width=1)
for y in range(0, FRAME_HEIGHT, 40):
draw.line([(0, y), (FRAME_WIDTH, y)], fill=(40, 40, 50), width=1)
# Interpolate position for smooth animation within block
t = frame_idx / NUM_FRAMES
# Draw trail (motion blur effect)
for i, (tx, ty) in enumerate(self.trail):
alpha = (i + 1) / len(self.trail) * 0.3
trail_size = int(self.size * (0.5 + alpha * 0.5))
trail_color = tuple(int(c * alpha) for c in self.color)
draw.ellipse(
[tx - trail_size//2, ty - trail_size//2,
tx + trail_size//2, ty + trail_size//2],
fill=trail_color
)
# Draw main object (circle with glow)
# Outer glow
for r in range(3, 0, -1):
glow_size = self.size + r * 8
glow_alpha = 0.2 / r
glow_color = tuple(int(c * glow_alpha) for c in self.color)
draw.ellipse(
[self.x - glow_size//2, self.y - glow_size//2,
self.x + glow_size//2, self.y + glow_size//2],
fill=glow_color
)
# Main circle
draw.ellipse(
[self.x - self.size//2, self.y - self.size//2,
self.x + self.size//2, self.y + self.size//2],
fill=self.color
)
# Inner highlight
highlight_size = self.size // 3
highlight_offset = self.size // 6
draw.ellipse(
[self.x - highlight_offset - highlight_size//2,
self.y - highlight_offset - highlight_size//2,
self.x - highlight_offset + highlight_size//2,
self.y - highlight_offset + highlight_size//2],
fill=(255, 255, 255)
)
# Add latency indicator text
draw.text((10, 10), f"Latency: {LATENCY_MS}ms", fill=(150, 150, 150))
draw.text((10, 30), f"Frame: {frame_idx + 1}/{NUM_FRAMES}", fill=(100, 100, 100))
return np.array(img)
def encode_frame(frame: np.ndarray) -> str:
"""Encode frame to base64 JPEG."""
img = Image.fromarray(frame)
buffer = io.BytesIO()
img.save(buffer, format='JPEG', quality=JPEG_QUALITY)
return base64.b64encode(buffer.getvalue()).decode('utf-8')
def generate_frames(game_state: GameState) -> list[str]:
"""Generate and encode a block of frames."""
encoded = []
for i in range(NUM_FRAMES):
frame = game_state.render_frame(i)
encoded.append(encode_frame(frame))
return encoded
# Keyboard/camera mappings (same as real server)
KEYBOARD_MAP = {
'w': [1, 0, 0, 0],
'a': [0, 1, 0, 0],
's': [0, 0, 1, 0],
'd': [0, 0, 0, 1],
}
CAMERA_MAP = {
'ArrowUp': [0, -1],
'ArrowDown': [0, 1],
'ArrowLeft': [-1, 0],
'ArrowRight': [1, 0],
}
@asynccontextmanager
async def lifespan(app: FastAPI):
print(f"Mock server starting with {LATENCY_MS}ms simulated latency...")
yield
print("Mock server shutting down...")
app = FastAPI(lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
async def send_frames(websocket: WebSocket, frames: list[str]):
"""Send frames in batches."""
for i in range(0, len(frames), BATCH_SIZE):
batch = frames[i:i + BATCH_SIZE]
await websocket.send_json({"type": "frame_batch", "frames": batch})
@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
await websocket.accept()
client_id = str(uuid.uuid4())
print(f"Client {client_id[:8]} connected")
game_state = GameState()
block_count = 0
timeout_task = None
async def session_timeout():
await asyncio.sleep(SESSION_TIMEOUT_SECONDS)
try:
await websocket.send_json({
"type": "session_timeout",
"message": f"Session expired after {SESSION_TIMEOUT_SECONDS} seconds"
})
await websocket.close()
except Exception:
pass
try:
# Send initial status
await websocket.send_json({
"type": "queue_status",
"position": 0,
"total_gpus": 1,
"available_gpus": 1
})
# Simulate GPU assignment delay
await asyncio.sleep(0.1)
# Start timeout
timeout_task = asyncio.create_task(session_timeout())
# Send GPU assigned
await websocket.send_json({
"type": "gpu_assigned",
"gpu_id": 0,
"session_timeout": SESSION_TIMEOUT_SECONDS,
"image_url": None # No initial image for mock
})
# Generate initial frame
await websocket.send_json({"type": "block_count", "count": 0, "max": MAX_BLOCKS})
# Simulate latency
await asyncio.sleep(LATENCY_MS / 1000)
frames = generate_frames(game_state)
block_count = 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
await send_frames(websocket, frames)
print(f"[Mock] Initial frames sent")
# Main loop
while True:
data = await websocket.receive_json()
message_type = data.get("type", "key")
if message_type == "reset":
print(f"[Mock] Reset requested")
await websocket.send_json({"type": "reset_started"})
game_state.reset()
await asyncio.sleep(LATENCY_MS / 1000)
frames = generate_frames(game_state)
block_count = 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
await send_frames(websocket, frames)
await websocket.send_json({"type": "reset_complete"})
continue
key = data.get("key")
if key and block_count < MAX_BLOCKS:
if key in CAMERA_MAP:
keyboard_vector = [0, 0, 0, 0]
mouse_vector = CAMERA_MAP[key]
else:
keyboard_vector = KEYBOARD_MAP.get(key, [0, 0, 0, 0])
mouse_vector = [0, 0]
print(f"[Mock] Key '{key}' pressed, generating block {block_count + 1}...")
t_start = time.time()
# Simulate model latency
await asyncio.sleep(LATENCY_MS / 1000)
# Update game state and generate frames
game_state.update(keyboard_vector, mouse_vector)
frames = generate_frames(game_state)
t_total = (time.time() - t_start) * 1000
block_count += 1
await websocket.send_json({"type": "block_count", "count": block_count, "max": MAX_BLOCKS})
await send_frames(websocket, frames)
print(f"[Mock] Block generated in {t_total:.0f}ms (target: {LATENCY_MS}ms)")
except WebSocketDisconnect:
print(f"Client {client_id[:8]} disconnected")
except Exception as e:
print(f"Client {client_id[:8]} error: {e}")
import traceback
traceback.print_exc()
finally:
if timeout_task and not timeout_task.done():
timeout_task.cancel()
# Serve static files
static_dir = os.path.join(os.path.dirname(__file__), "..", "client", "dist")
if os.path.isdir(static_dir):
app.mount("/", StaticFiles(directory=static_dir, html=True), name="static")
if __name__ == "__main__":
import argparse
import uvicorn
parser = argparse.ArgumentParser(description="Mock server for latency testing")
parser.add_argument("--latency", type=int, default=200, help="Simulated latency in ms")
parser.add_argument("--port", type=int, default=8001, help="Server port")
args = parser.parse_args()
LATENCY_MS = args.latency
print(f"Starting mock server with {LATENCY_MS}ms latency on port {args.port}")
uvicorn.run(app, host="0.0.0.0", port=args.port)
+149
View File
@@ -0,0 +1,149 @@
"""Test ORCA batching: synchronized vs random keypresses."""
import argparse
import asyncio
import json
import random
import time
import websockets
SERVER_URL = "ws://localhost:8000/ws"
KEYS = ["w", "a", "s", "d"]
# Think time range for random mode (ms)
MIN_THINK_MS = 200
MAX_THINK_MS = 1500
async def connect_and_wait(client_id, ws):
"""Drain messages until initial frame_batch is received."""
while True:
msg = json.loads(await ws.recv())
if msg["type"] == "frame_batch":
print(f" Client {client_id}: got initial frame")
return
async def synced_client(client_id: int, barrier: asyncio.Barrier, results: list,
num_rounds: int):
"""Synchronized keypresses — all clients press at the same time."""
async with websockets.connect(
SERVER_URL, max_size=50_000_000, ping_interval=None,
) as ws:
await connect_and_wait(client_id, ws)
for round_idx in range(num_rounds):
await barrier.wait()
t0 = time.perf_counter()
await ws.send(json.dumps({"type": "key", "key": "w"}))
while True:
msg = json.loads(await ws.recv())
if msg["type"] == "frame_batch":
break
elapsed_ms = (time.perf_counter() - t0) * 1000
results.append(elapsed_ms)
print(f" Client {client_id} round {round_idx}: {elapsed_ms:.0f}ms")
async def random_client(client_id: int, start_event: asyncio.Event, results: list,
num_rounds: int):
"""Random keypresses with think time between presses."""
async with websockets.connect(
SERVER_URL, max_size=50_000_000, ping_interval=None,
) as ws:
await connect_and_wait(client_id, ws)
start_event.set()
# Small stagger so clients don't all start at once
await asyncio.sleep(random.uniform(0, 1.0))
for _ in range(num_rounds):
think_ms = random.uniform(MIN_THINK_MS, MAX_THINK_MS)
await asyncio.sleep(think_ms / 1000)
t0 = time.perf_counter()
await ws.send(json.dumps({"type": "key", "key": random.choice(KEYS)}))
while True:
msg = json.loads(await ws.recv())
if msg["type"] == "frame_batch":
break
elapsed_ms = (time.perf_counter() - t0) * 1000
results.append(elapsed_ms)
print(f" Client {client_id}: done")
def print_stats(times: list[float], label: str):
"""Print summary statistics and histogram."""
times.sort()
n = len(times)
print(f"\n--- {label} ({n} keypresses) ---")
print(f" Min: {times[0]:.0f}ms")
print(f" p25: {times[n//4]:.0f}ms")
print(f" Median: {times[n//2]:.0f}ms")
print(f" p75: {times[3*n//4]:.0f}ms")
print(f" p95: {times[int(n*0.95)]:.0f}ms")
print(f" Max: {times[-1]:.0f}ms")
print(f" Avg: {sum(times)/n:.0f}ms")
buckets = [0] * 10
labels = ["<1s", "<1.5s", "<2s", "<2.5s", "<3s",
"<3.5s", "<4s", "<4.5s", "<5s", "5s+"]
for ms in times:
idx = min(int(ms / 500), 9)
buckets[idx] += 1
print(f"\n Distribution:")
for lbl, count in zip(labels, buckets):
if count > 0:
bar = "#" * count
pct = count / n * 100
print(f" {lbl:>6s}: {bar} ({count}, {pct:.0f}%)")
async def run_synced(num_clients: int, num_rounds: int):
print(f"=== SYNCHRONIZED MODE ({num_clients} clients, {num_rounds} rounds) ===")
barrier = asyncio.Barrier(num_clients)
results = []
tasks = [synced_client(i, barrier, results, num_rounds) for i in range(num_clients)]
await asyncio.gather(*tasks)
print_stats(results, "Synchronized (worst case)")
async def run_random(num_clients: int, num_rounds: int):
print(f"=== RANDOM MODE ({num_clients} clients, {num_rounds} presses each) ===")
print(f" Think time: {MIN_THINK_MS}-{MAX_THINK_MS}ms")
results = []
events = [asyncio.Event() for _ in range(num_clients)]
tasks = [random_client(i, events[i], results, num_rounds) for i in range(num_clients)]
await asyncio.gather(*tasks)
print_stats(results, "Random (realistic)")
def parse_args():
parser = argparse.ArgumentParser(description="Test ORCA multi-user batching")
parser.add_argument("mode", nargs="?", default="both",
choices=["synced", "random", "both"])
parser.add_argument("-c", "--clients", type=int, default=4,
help="Number of concurrent clients (default: 4)")
parser.add_argument("-r", "--rounds", type=int, default=10,
help="Number of rounds per client (default: 10)")
return parser.parse_args()
async def main():
args = parse_args()
print(f"Connecting {args.clients} clients to {SERVER_URL}...\n")
if args.mode in ("synced", "both"):
await run_synced(args.clients, args.rounds)
if args.mode in ("random", "both"):
if args.mode == "both":
print("\n" + "=" * 60 + "\n")
await run_random(args.clients, args.rounds)
if __name__ == "__main__":
asyncio.run(main())