Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c4a0f789da | ||
|
|
5ac14938d2 | ||
|
|
9d239e9f8b |
@@ -0,0 +1,43 @@
|
||||
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
|
||||
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
|
||||
num_frame_per_block=7,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -13,7 +13,7 @@ from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig, WANV2VConfig)
|
||||
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -37,6 +37,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
|
||||
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWan2_2_T2V480PConfig,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
|
||||
|
||||
@@ -176,3 +176,13 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
|
||||
is_causal: bool = True
|
||||
flow_shift: float | None = 12.0
|
||||
boundary_ratio: float | None = 0.875
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 850, 700, 550, 350, 275, 200, 125])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
@@ -9,7 +9,7 @@ from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.sample.wan import (
|
||||
FastWanT2V480PConfig,
|
||||
FastWanT2V480P_SamplingParam,
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
@@ -19,7 +19,8 @@ from fastvideo.configs.sample.wan import (
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam,
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
SelfForcingWanT2V480PConfig,
|
||||
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -29,35 +30,50 @@ from fastvideo.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
# Registry maps specific model weights to their config classes
|
||||
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
"FastVideo/FastHunyuan-diffusers":
|
||||
FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo":
|
||||
HunyuanSamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers":
|
||||
StepVideoT2VSamplingParam,
|
||||
|
||||
# Wan2.1
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
|
||||
WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
|
||||
# Wan2.2
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
|
||||
# FastWan2.1
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
|
||||
FastWanT2V480P_SamplingParam,
|
||||
|
||||
# FastWan2.2
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
|
||||
# Causal Self-Forcing Wan2.1
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
|
||||
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers":
|
||||
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -67,6 +83,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -77,7 +95,9 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
"wanpipeline":
|
||||
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
|
||||
"stepvideo": StepVideoT2VSamplingParam
|
||||
"wandmdpipeline": FastWanT2V480P_SamplingParam,
|
||||
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
"stepvideo": StepVideoT2VSamplingParam,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -97,7 +97,7 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
|
||||
# DMD parameters
|
||||
# dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
|
||||
num_inference_steps: int = 3
|
||||
@@ -183,5 +183,18 @@ class Wan2_2_Fun_A14B_Control_SamplingParam(
|
||||
# ============= Causal Self-Forcing =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class SelfForcingWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
|
||||
Wan2_2_T2V_A14B_SamplingParam):
|
||||
guidance_scale: float = 2.0
|
||||
guidance_scale_2: float = 2.0
|
||||
num_inference_steps: int = 8
|
||||
num_frames: int = 81
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
fps: int = 16
|
||||
|
||||
@@ -208,7 +208,7 @@ class VideoGenerator:
|
||||
|
||||
def _sanitize_filename_component(name: str) -> str:
|
||||
# Remove characters invalid on common filesystems, strip spaces/dots
|
||||
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
|
||||
sanitized = re.sub(r'[\/:*?"<>|]', '', name)
|
||||
sanitized = sanitized.strip().strip('.')
|
||||
sanitized = re.sub(r'\s+', ' ', sanitized)
|
||||
return sanitized or "video"
|
||||
|
||||
@@ -20,7 +20,7 @@ import fastvideo.envs as envs
|
||||
from fastvideo.attention import (DistributedAttention,
|
||||
LocalAttention)
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size, get_local_torch_device
|
||||
from fastvideo.forward_context import get_forward_context
|
||||
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
|
||||
RMSNorm, ScaleResidual,
|
||||
@@ -33,9 +33,29 @@ 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 WanT2VCrossAttention, WanTimeTextImageEmbedding
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class CacheAppend(torch.autograd.Function):
|
||||
"""
|
||||
KV cache with shape [batch, seq_len, heads, head_dim].
|
||||
"""
|
||||
@staticmethod
|
||||
def forward(ctx, storage, active_cache, x, start, end):
|
||||
# Ensure storage has the same dtype as x
|
||||
storage.data[:, start:end] = x
|
||||
ctx.save_for_backward(storage.to(x.dtype))
|
||||
ctx.start = start
|
||||
ctx.end = end
|
||||
return storage[:, :end].to(x.dtype) # Ensure returned value has same dtype as input
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
start = ctx.start
|
||||
end = ctx.end
|
||||
return None, grad_output[:, :start], grad_output[:, start:end], None, None
|
||||
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
@@ -67,6 +87,10 @@ class CausalWanSelfAttention(nn.Module):
|
||||
causal=False,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
self.k_cache = None
|
||||
self.v_cache = None
|
||||
self.counter = 0
|
||||
|
||||
def forward(self,
|
||||
q: torch.Tensor,
|
||||
@@ -84,6 +108,19 @@ class CausalWanSelfAttention(nn.Module):
|
||||
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
||||
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
||||
"""
|
||||
if kv_cache is not None:
|
||||
# Then we are in _forward_inference mode
|
||||
if self.k_cache is None:
|
||||
assert self.counter == 0
|
||||
self.counter += 1
|
||||
del self.k_cache
|
||||
self.register_buffer("k_cache", torch.empty(1, self.max_attention_size, self.num_heads, self.head_dim, device=q.device, dtype=q.dtype), persistent=False)
|
||||
self.k_cache.requires_grad_(True)
|
||||
if self.v_cache is None:
|
||||
del self.v_cache
|
||||
self.register_buffer("v_cache", torch.empty(1, self.max_attention_size, self.num_heads, self.head_dim, device=v.device, dtype=v.dtype), persistent=False)
|
||||
self.v_cache.requires_grad_(True)
|
||||
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
@@ -128,6 +165,8 @@ class CausalWanSelfAttention(nn.Module):
|
||||
num_new_tokens = roped_query.shape[1]
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
raise Exception("Not implemented")
|
||||
# @TODO(Wei): This part has not been thoroughly tested yet. Use with caution.
|
||||
# Calculate the number of new tokens added in this step
|
||||
# Shift existing cache content left to discard oldest tokens
|
||||
# Clone the source slice to avoid overlapping memory error
|
||||
@@ -137,26 +176,44 @@ class CausalWanSelfAttention(nn.Module):
|
||||
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# self.k_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
# self.k_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# self.v_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
# self.v_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# Insert the new keys/values at the end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - \
|
||||
kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
# local_k = CacheAppend.apply(self.k_cache, kv_cache["k"], roped_key, local_start_index, local_end_index)
|
||||
# local_v = CacheAppend.apply(self.v_cache, kv_cache["v"], v, local_start_index, local_end_index)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
else:
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
kv_cache["k"] = kv_cache["k"].detach()
|
||||
kv_cache["v"] = kv_cache["v"].detach()
|
||||
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
# kv_cache["k"] = kv_cache["k"].detach()
|
||||
# kv_cache["v"] = kv_cache["v"].detach()
|
||||
# kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
# kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
local_k = CacheAppend.apply(self.k_cache, kv_cache["k"], roped_key, local_start_index, local_end_index)
|
||||
# logger.info("Is local_k meta tensor: %s, %s", local_k.is_meta, local_k.shape)
|
||||
# logger.info("local_start_index: %d, local_end_index: %d, number of zeros in local k: %d", local_start_index, local_end_index, (local_k == 0).sum().item())
|
||||
local_v = CacheAppend.apply(self.v_cache, kv_cache["v"], v, local_start_index, local_end_index)
|
||||
|
||||
x = self.attn(
|
||||
roped_query,
|
||||
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
|
||||
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
|
||||
local_k,
|
||||
local_v
|
||||
)
|
||||
# x = self.attn(
|
||||
# roped_query,
|
||||
# kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
|
||||
# kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
|
||||
# )
|
||||
|
||||
kv_cache["k"] = local_k
|
||||
kv_cache["v"] = local_v
|
||||
kv_cache["global_end_index"].fill_(current_end)
|
||||
kv_cache["local_end_index"].fill_(local_end_index)
|
||||
|
||||
@@ -233,6 +290,9 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
self.null_shift = torch.tensor([0], device=get_local_torch_device())
|
||||
self.null_scale = torch.tensor([0], device=get_local_torch_device())
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -283,9 +343,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
attn_output, _ = self.to_out(attn_output)
|
||||
attn_output = attn_output.squeeze(1)
|
||||
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
hidden_states, attn_output, gate_msa, self.null_shift, self.null_scale)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -452,7 +511,6 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
This function will be run for num_frame times.
|
||||
Process the latent frames one by one (1560 tokens each)
|
||||
"""
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
|
||||
@@ -212,9 +212,18 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
start_index = 0
|
||||
|
||||
# DMD loop in causal blocks
|
||||
# Optional per-block callback for streaming
|
||||
on_block = None
|
||||
try:
|
||||
on_block = getattr(batch, "extra",
|
||||
{}).get("on_block",
|
||||
None) # type: ignore[attr-defined]
|
||||
except Exception:
|
||||
on_block = None
|
||||
|
||||
with self.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps)) as progress_bar:
|
||||
for current_num_frames in block_sizes:
|
||||
for block_idx, current_num_frames in enumerate(block_sizes):
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
# use BTCHW for DMD conversion routines
|
||||
@@ -355,6 +364,20 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
)
|
||||
start_index += current_num_frames
|
||||
|
||||
# Invoke callback with block metadata (no large tensor transfer required)
|
||||
try:
|
||||
if callable(on_block):
|
||||
on_block(
|
||||
block_index=block_idx,
|
||||
total_blocks=len(block_sizes),
|
||||
start_index=start_index - current_num_frames,
|
||||
num_frames=current_num_frames,
|
||||
latents=current_latents,
|
||||
)
|
||||
except Exception as e:
|
||||
# Swallow callback errors so they don't break generation
|
||||
logger.warning("on_block callback failed: %s", str(e))
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ import time
|
||||
from collections.abc import Callable
|
||||
from multiprocessing.process import BaseProcess
|
||||
from typing import Any, cast
|
||||
from collections.abc import Iterator
|
||||
|
||||
import psutil
|
||||
|
||||
@@ -22,6 +23,8 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.utils import decorate_logs, get_distributed_init_method, get_exception_traceback, get_loopback_ip, get_mp_context, get_open_port, kill_itself_when_parent_died, force_spawn
|
||||
from fastvideo.worker.executor import Executor
|
||||
from fastvideo.worker.worker_base import WorkerWrapperBase
|
||||
import torch
|
||||
import torchvision
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -88,6 +91,61 @@ class MultiprocExecutor(Executor):
|
||||
|
||||
return result_batch
|
||||
|
||||
def execute_forward_streaming(
|
||||
self, forward_batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> Iterator[dict[str, Any]]:
|
||||
"""
|
||||
Multiprocess streaming interface.
|
||||
|
||||
Broadcasts a streaming forward request to all workers. Rank 0 streams
|
||||
events (progress/block/complete) back through its pipe; other ranks
|
||||
run the forward pass and return a small completion status.
|
||||
"""
|
||||
# Send request to all workers
|
||||
for worker in self.workers:
|
||||
worker.pipe.send({
|
||||
"method": "execute_forward_streaming",
|
||||
"args": (),
|
||||
"kwargs": {
|
||||
"forward_batch": forward_batch,
|
||||
"fastvideo_args": fastvideo_args,
|
||||
},
|
||||
})
|
||||
|
||||
done_workers: set[int] = set()
|
||||
rank0_complete = False
|
||||
pipes = [w.pipe for w in self.workers]
|
||||
|
||||
# Drain messages until all workers are done
|
||||
while len(done_workers) < len(self.workers):
|
||||
ready = mp.connection.wait(pipes)
|
||||
for pipe in ready:
|
||||
# Identify which worker sent this
|
||||
idx = next(
|
||||
(i for i, w in enumerate(self.workers) if w.pipe is pipe),
|
||||
-1)
|
||||
if idx < 0:
|
||||
continue
|
||||
msg = pipe.recv()
|
||||
|
||||
# Rank 0 streams events
|
||||
if idx == 0 and isinstance(msg, dict):
|
||||
msg_type = msg.get("type")
|
||||
if msg_type in ("progress", "block", "complete"):
|
||||
yield msg
|
||||
if msg_type == "complete":
|
||||
rank0_complete = True
|
||||
done_workers.add(idx)
|
||||
elif msg.get("status") == "done":
|
||||
done_workers.add(idx)
|
||||
else:
|
||||
# Non-zero ranks: expect simple completion status
|
||||
if isinstance(msg, dict) and msg.get("status") == "done":
|
||||
done_workers.add(idx)
|
||||
else:
|
||||
# Any unexpected message from non-zero ranks counts as done
|
||||
done_workers.add(idx)
|
||||
|
||||
def set_lora_adapter(self,
|
||||
lora_nickname: str,
|
||||
lora_path: str | None = None) -> None:
|
||||
@@ -467,6 +525,156 @@ class WorkerMultiprocProc:
|
||||
"output_batch": output_batch.output.cpu(),
|
||||
"logging_info": logging_info
|
||||
})
|
||||
if method == 'execute_forward_streaming':
|
||||
forward_batch = kwargs['forward_batch']
|
||||
fastvideo_args = kwargs['fastvideo_args']
|
||||
# Install a lightweight per-block callback for streaming metadata from rank 0
|
||||
try:
|
||||
extra = getattr(forward_batch, 'extra', None)
|
||||
if extra is None:
|
||||
forward_batch.extra = {}
|
||||
extra = forward_batch.extra
|
||||
except Exception:
|
||||
forward_batch.extra = {}
|
||||
extra = forward_batch.extra
|
||||
|
||||
def _on_block(**evt):
|
||||
# Only rank 0 streams events to the executor pipe
|
||||
if self.rank != 0:
|
||||
return
|
||||
try:
|
||||
block_index = int(evt.get('block_index', 0))
|
||||
total_blocks = int(evt.get('total_blocks', 0))
|
||||
num_frames = int(evt.get('num_frames', 0))
|
||||
latents = evt.get('latents')
|
||||
if latents is None:
|
||||
# Fallback to meta event if no latents provided
|
||||
self.pipe.send({
|
||||
'type': 'block_meta',
|
||||
'block_index': block_index,
|
||||
'total_blocks': total_blocks,
|
||||
'num_frames': num_frames,
|
||||
})
|
||||
return
|
||||
|
||||
# Decode latents to pixels using pipeline VAE
|
||||
vae = None
|
||||
try:
|
||||
vae = self.worker.pipeline.get_module('vae')
|
||||
except Exception:
|
||||
vae = None
|
||||
if vae is None:
|
||||
# If no VAE, send meta only
|
||||
self.pipe.send({
|
||||
'type': 'block_meta',
|
||||
'block_index': block_index,
|
||||
'total_blocks': total_blocks,
|
||||
'num_frames': num_frames,
|
||||
})
|
||||
return
|
||||
|
||||
# Prepare latents for VAE (apply scaling/shift like DecodingStage)
|
||||
z = latents.permute(
|
||||
0, 2, 1, 3, 4
|
||||
) # [B,T,C,H,W] -> [B,C,T,H,W] expected by many VAEs
|
||||
if hasattr(vae, 'scaling_factor'
|
||||
) and vae.scaling_factor is not None:
|
||||
sf = vae.scaling_factor
|
||||
if isinstance(sf, torch.Tensor):
|
||||
z = z / sf.to(z.device, z.dtype)
|
||||
else:
|
||||
z = z / sf
|
||||
if hasattr(vae, 'shift_factor'
|
||||
) and vae.shift_factor is not None:
|
||||
shf = vae.shift_factor
|
||||
if isinstance(shf, torch.Tensor):
|
||||
z = z + shf.to(z.device, z.dtype)
|
||||
else:
|
||||
z = z + shf
|
||||
|
||||
with torch.autocast(device_type='cuda',
|
||||
dtype=torch.bfloat16,
|
||||
enabled=True):
|
||||
pixels = vae.decode(
|
||||
z
|
||||
) # [B,C,T,H,W] in [-1,1] or already normalized depending on VAE
|
||||
|
||||
# Normalize to [0,1]
|
||||
pixels = (pixels / 2 + 0.5).clamp(0, 1)
|
||||
# Convert to uint8 HWC per-frame for first batch element
|
||||
pixels = (pixels[0].permute(1, 2, 3, 0) *
|
||||
255).to(torch.uint8).cpu().numpy()
|
||||
# Now pixels shape: [T,H,W,C]
|
||||
frames = [
|
||||
pixels[i] for i in range(pixels.shape[0])
|
||||
]
|
||||
|
||||
self.pipe.send({
|
||||
'type': 'block',
|
||||
'block_index': block_index,
|
||||
'total_blocks': total_blocks,
|
||||
'frames': frames,
|
||||
})
|
||||
except Exception:
|
||||
# Fail-safe: don't break generation if decode fails
|
||||
try:
|
||||
self.pipe.send({
|
||||
'type':
|
||||
'block_meta',
|
||||
'block_index':
|
||||
int(evt.get('block_index', 0)),
|
||||
'total_blocks':
|
||||
int(evt.get('total_blocks', 0)),
|
||||
'num_frames':
|
||||
int(evt.get('num_frames', 0)),
|
||||
})
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
extra['on_block'] = _on_block
|
||||
except Exception:
|
||||
pass
|
||||
# Run forward on all ranks
|
||||
output_batch = self.worker.execute_forward(
|
||||
forward_batch, fastvideo_args)
|
||||
if self.rank != 0:
|
||||
# Non-zero ranks send a small completion message
|
||||
self.pipe.send({"status": "done"})
|
||||
else:
|
||||
# Rank 0 builds frames and streams a single block then complete
|
||||
samples = output_batch.output # Tensor [b,c,t,h,w]
|
||||
try:
|
||||
videos = samples.permute(2, 0, 1, 3, 4)
|
||||
frames = []
|
||||
for x in videos:
|
||||
grid = torchvision.utils.make_grid(x,
|
||||
nrow=6)
|
||||
grid = grid.transpose(0, 1).transpose(
|
||||
1, 2).squeeze(-1)
|
||||
frame = (grid * 255).to(
|
||||
torch.uint8).cpu().numpy()
|
||||
frames.append(frame)
|
||||
# Emit a single block with all frames
|
||||
self.pipe.send({
|
||||
"type": "block",
|
||||
"block_index": 0,
|
||||
"total_blocks": 1,
|
||||
"frames": frames,
|
||||
})
|
||||
except Exception:
|
||||
# If frame construction fails, still complete
|
||||
pass
|
||||
# Emit a complete event
|
||||
self.pipe.send({
|
||||
"type": "complete",
|
||||
"result": {
|
||||
"num_frames":
|
||||
int(samples.shape[2])
|
||||
if hasattr(samples, 'shape')
|
||||
and len(samples.shape) >= 3 else None
|
||||
},
|
||||
})
|
||||
else:
|
||||
result = self.worker.execute_method(method, *args, **kwargs)
|
||||
self.pipe.send(result)
|
||||
|
||||
Reference in New Issue
Block a user