Compare commits

...
Author SHA1 Message Date
SolitaryThinker c4a0f789da wip 2025-10-31 20:54:49 +00:00
SolitaryThinker 5ac14938d2 add example 2025-10-31 20:54:00 +00:00
SolitaryThinker 9d239e9f8b prepare for wan2.2 self-forcing 2025-10-31 20:52:31 +00:00
9 changed files with 409 additions and 33 deletions
@@ -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()
+2 -1
View File
@@ -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,
+10
View File
@@ -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
+36 -16
View File
@@ -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
}
+15 -2
View File
@@ -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
+1 -1
View File
@@ -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"
+70 -12
View File
@@ -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):
+24 -1
View File
@@ -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
+208
View File
@@ -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)