Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0c606b7cfa | ||
|
|
0a0cdebaba | ||
|
|
cd9e7a8cac | ||
|
|
87e19dad5a | ||
|
|
79195ad18e |
@@ -40,6 +40,9 @@ dist/
|
||||
*.egg
|
||||
eggs/
|
||||
.eggs/
|
||||
node_modules/
|
||||
.vite
|
||||
vite.config.js.timestamp-*.mjs
|
||||
|
||||
# MkDocs documentation
|
||||
site/
|
||||
|
||||
@@ -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"
|
||||
@@ -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":
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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]
|
||||
@@ -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",
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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>
|
||||
Generated
+1901
File diff suppressed because it is too large
Load Diff
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -0,0 +1,8 @@
|
||||
import App from './App.svelte'
|
||||
import './style.css'
|
||||
|
||||
const app = new App({
|
||||
target: document.getElementById('app')
|
||||
})
|
||||
|
||||
export default app
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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',
|
||||
},
|
||||
},
|
||||
}
|
||||
})
|
||||
@@ -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
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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())
|
||||
Reference in New Issue
Block a user