Self-Forcing / pipeline /causal_inference.py
Cccccz's picture
Add files using upload-large-folder tool
d1d122e verified
Raw
History Blame Contribute Delete
23.7 kB
from typing import List, Optional
import torch
from model.predictor_v4 import SelfForcingPredictorV4
from predictor_training.rollout_cache import (
build_predictor_workspace,
reset_predictor_workspace,
)
from utils.wan_wrapper import WanDiffusionWrapper, WanTextEncoder, WanVAEWrapper
from demo_utils.memory import gpu, get_cuda_free_memory_gb, DynamicSwapInstaller, move_model_to_device_with_memory_preservation
def velocity_reuse_steps(
schedule: str,
temporal_block_index: int,
) -> tuple[int, ...]:
"""Return denoising steps that reuse step-0 velocity for one chunk."""
schedule = str(schedule).lower()
if schedule not in {"frrf", "frrr"}:
raise ValueError(f"Unknown velocity-reuse schedule {schedule!r}")
if int(temporal_block_index) == 0:
return ()
return (1, 2, 3) if schedule == "frrr" else (1, 2)
class CausalInferencePipeline(torch.nn.Module):
def __init__(
self,
args,
device,
generator=None,
text_encoder=None,
vae=None
):
super().__init__()
# Step 1: Initialize all models
self.generator = WanDiffusionWrapper(
**getattr(args, "model_kwargs", {}), is_causal=True) if generator is None else generator
self.text_encoder = WanTextEncoder() if text_encoder is None else text_encoder
self.vae = WanVAEWrapper() if vae is None else vae
# Step 2: Initialize all causal hyperparmeters
self.scheduler = self.generator.get_scheduler()
self.denoising_step_list = torch.tensor(
args.denoising_step_list, dtype=torch.long)
if args.warp_denoising_step:
timesteps = torch.cat((self.scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
self.denoising_step_list = timesteps[1000 - self.denoising_step_list]
self.num_transformer_blocks = 30
self.frame_seq_length = 1560
self.kv_cache1 = None
self.predictor_v4: SelfForcingPredictorV4 | None = None
self.predictor_schedule = "fppf"
self.args = args
self.num_frame_per_block = getattr(args, "num_frame_per_block", 1)
self.independent_first_frame = args.independent_first_frame
self.local_attn_size = self.generator.model.local_attn_size
self.reuse_first_step_velocity = bool(
getattr(args, "reuse_first_step_velocity", False)
)
self.reuse_first_step_velocity_schedule = str(
getattr(args, "reuse_first_step_velocity_schedule", "frrf")
).lower()
if self.reuse_first_step_velocity and len(self.denoising_step_list) != 4:
raise ValueError(
"reuse_first_step_velocity requires exactly four denoising steps"
)
if self.reuse_first_step_velocity_schedule not in {"frrf", "frrr"}:
raise ValueError(
"reuse_first_step_velocity_schedule must be frrf or frrr, got "
f"{self.reuse_first_step_velocity_schedule!r}"
)
print(f"KV inference with {self.num_frame_per_block} frames per block")
if self.reuse_first_step_velocity:
print(
"Denoising schedule: chunk 0 F-F-F-F; later chunks "
f"{self.reuse_first_step_velocity_schedule.upper()} "
"(R reuses step 0 velocity)"
)
if self.num_frame_per_block > 1:
self.generator.model.num_frame_per_block = self.num_frame_per_block
def enable_predictor_v4(
self,
predictor: SelfForcingPredictorV4,
*,
schedule: str = "fppf",
) -> None:
"""Enable F-P-P-F or F-P-P-P after the first all-Full chunk."""
if self.reuse_first_step_velocity:
raise ValueError("Predictor F-P-P-F and velocity F-R-R-F are mutually exclusive")
if len(self.denoising_step_list) != 4:
raise ValueError("Predictor F-P-P-F requires exactly four denoising steps")
if self.independent_first_frame:
raise NotImplementedError(
"Predictor-v4 currently supports the regular three-frame T2V chunks"
)
if self.num_frame_per_block != 3:
raise ValueError("Predictor-v4 was trained for three latent frames per chunk")
schedule = str(schedule).lower()
if schedule not in {"fppf", "fppp"}:
raise ValueError(f"Unknown Predictor schedule {schedule!r}")
self.predictor_v4 = predictor.eval()
self.predictor_schedule = schedule
print(
"Denoising schedule: chunk 0 F-F-F-F; later chunks "
f"{schedule.upper()} "
f"(Predictor blocks {predictor.source_block_ids})"
)
def _full_step_with_optional_hidden(
self,
*,
noisy_input: torch.Tensor,
conditional_dict: dict,
timestep: torch.Tensor,
current_start: int,
capture_hidden: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
"""Run the loaded generator and optionally capture its pre-head hidden."""
captured: list[torch.Tensor] = []
handle = None
if capture_hidden:
def capture_head_input(module, args):
if not args or not torch.is_tensor(args[0]):
raise RuntimeError("Wan head hook did not receive final hidden states")
captured.append(args[0])
handle = self.generator.model.head.register_forward_pre_hook(
capture_head_input
)
try:
flow_pred, denoised_pred = self.generator(
noisy_image_or_video=noisy_input,
conditional_dict=conditional_dict,
timestep=timestep,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start,
)
finally:
if handle is not None:
handle.remove()
if capture_hidden and len(captured) != 1:
raise RuntimeError(
f"Expected one Wan final hidden capture, received {len(captured)}"
)
return flow_pred, denoised_pred, captured[0] if captured else None
def inference(
self,
noise: torch.Tensor,
text_prompts: List[str],
initial_latent: Optional[torch.Tensor] = None,
return_latents: bool = False,
profile: bool = False,
low_memory: bool = False,
) -> torch.Tensor:
"""
Perform inference on the given noise and text prompts.
Inputs:
noise (torch.Tensor): The input noise tensor of shape
(batch_size, num_output_frames, num_channels, height, width).
text_prompts (List[str]): The list of text prompts.
initial_latent (torch.Tensor): The initial latent tensor of shape
(batch_size, num_input_frames, num_channels, height, width).
If num_input_frames is 1, perform image to video.
If num_input_frames is greater than 1, perform video extension.
return_latents (bool): Whether to return the latents.
Outputs:
video (torch.Tensor): The generated video tensor of shape
(batch_size, num_output_frames, num_channels, height, width).
It is normalized to be in the range [0, 1].
"""
batch_size, num_frames, num_channels, height, width = noise.shape
if not self.independent_first_frame or (self.independent_first_frame and initial_latent is not None):
# If the first frame is independent and the first frame is provided, then the number of frames in the
# noise should still be a multiple of num_frame_per_block
assert num_frames % self.num_frame_per_block == 0
num_blocks = num_frames // self.num_frame_per_block
else:
# Using a [1, 4, 4, 4, 4, 4, ...] model to generate a video without image conditioning
assert (num_frames - 1) % self.num_frame_per_block == 0
num_blocks = (num_frames - 1) // self.num_frame_per_block
num_input_frames = initial_latent.shape[1] if initial_latent is not None else 0
num_output_frames = num_frames + num_input_frames # add the initial latent frames
conditional_dict = self.text_encoder(
text_prompts=text_prompts
)
if low_memory:
gpu_memory_preservation = get_cuda_free_memory_gb(gpu) + 5
move_model_to_device_with_memory_preservation(self.text_encoder, target_device=gpu, preserved_memory_gb=gpu_memory_preservation)
output = torch.zeros(
[batch_size, num_output_frames, num_channels, height, width],
device=noise.device,
dtype=noise.dtype
)
# Set up profiling if requested
if profile:
init_start = torch.cuda.Event(enable_timing=True)
init_end = torch.cuda.Event(enable_timing=True)
diffusion_start = torch.cuda.Event(enable_timing=True)
diffusion_end = torch.cuda.Event(enable_timing=True)
vae_start = torch.cuda.Event(enable_timing=True)
vae_end = torch.cuda.Event(enable_timing=True)
block_times = []
block_start = torch.cuda.Event(enable_timing=True)
block_end = torch.cuda.Event(enable_timing=True)
init_start.record()
# Step 1: Initialize KV cache to all zeros
if self.kv_cache1 is None:
self._initialize_kv_cache(
batch_size=batch_size,
dtype=noise.dtype,
device=noise.device
)
self._initialize_crossattn_cache(
batch_size=batch_size,
dtype=noise.dtype,
device=noise.device
)
else:
# reset cross attn cache
for block_index in range(self.num_transformer_blocks):
self.crossattn_cache[block_index]["is_init"] = False
# reset kv cache
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index]["global_end_index"] = torch.tensor(
[0], dtype=torch.long, device=noise.device)
self.kv_cache1[block_index]["local_end_index"] = torch.tensor(
[0], dtype=torch.long, device=noise.device)
# Step 2: Cache context feature
current_start_frame = 0
if initial_latent is not None:
timestep = torch.ones([batch_size, 1], device=noise.device, dtype=torch.int64) * 0
if self.independent_first_frame:
# Assume num_input_frames is 1 + self.num_frame_per_block * num_input_blocks
assert (num_input_frames - 1) % self.num_frame_per_block == 0
num_input_blocks = (num_input_frames - 1) // self.num_frame_per_block
output[:, :1] = initial_latent[:, :1]
self.generator(
noisy_image_or_video=initial_latent[:, :1],
conditional_dict=conditional_dict,
timestep=timestep * 0,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame * self.frame_seq_length,
)
current_start_frame += 1
else:
# Assume num_input_frames is self.num_frame_per_block * num_input_blocks
assert num_input_frames % self.num_frame_per_block == 0
num_input_blocks = num_input_frames // self.num_frame_per_block
for _ in range(num_input_blocks):
current_ref_latents = \
initial_latent[:, current_start_frame:current_start_frame + self.num_frame_per_block]
output[:, current_start_frame:current_start_frame + self.num_frame_per_block] = current_ref_latents
self.generator(
noisy_image_or_video=current_ref_latents,
conditional_dict=conditional_dict,
timestep=timestep * 0,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame * self.frame_seq_length,
)
current_start_frame += self.num_frame_per_block
if profile:
init_end.record()
torch.cuda.synchronize()
diffusion_start.record()
# Step 3: Temporal denoising loop
all_num_frames = [self.num_frame_per_block] * num_blocks
if self.independent_first_frame and initial_latent is None:
all_num_frames = [1] + all_num_frames
previous_chunk_hidden: dict[int, torch.Tensor] = {}
for temporal_block_index, current_num_frames in enumerate(all_num_frames):
if profile:
block_start.record()
noisy_input = noise[
:, current_start_frame - num_input_frames:current_start_frame + current_num_frames - num_input_frames]
# Step 3.1: Spatial denoising loop
first_step_flow_pred = None
same_chunk_anchor = None
chunk_step_hidden: dict[int, torch.Tensor] = {}
predictor_kv_cache = None
predictor_steps = (
(1, 2, 3) if self.predictor_schedule == "fppp" else (1, 2)
)
if self.predictor_v4 is not None and temporal_block_index > 0:
history_tokens = current_start_frame * self.frame_seq_length
predictor_kv_cache = build_predictor_workspace(
self.kv_cache1,
source_block_ids=tuple(self.predictor_v4.source_block_ids),
history_tokens=history_tokens,
current_tokens=current_num_frames * self.frame_seq_length,
)
reuse_step_indices = velocity_reuse_steps(
self.reuse_first_step_velocity_schedule,
temporal_block_index,
)
for index, current_timestep in enumerate(self.denoising_step_list):
reuse_velocity = (
self.reuse_first_step_velocity
and index in reuse_step_indices
)
use_predictor = (
self.predictor_v4 is not None
and temporal_block_index > 0
and index in predictor_steps
)
if use_predictor:
step_mode = "predictor"
elif reuse_velocity:
step_mode = "reuse"
else:
step_mode = "full"
print(f"current_timestep: {current_timestep} ({step_mode})")
# set current timestep
timestep = torch.ones(
[batch_size, current_num_frames],
device=noise.device,
dtype=torch.int64) * current_timestep
if use_predictor:
reset_predictor_workspace(
predictor_kv_cache,
history_tokens=current_start_frame * self.frame_seq_length,
)
if same_chunk_anchor is None:
raise RuntimeError(
f"Predictor step {index} has no same-chunk anchor hidden"
)
if index not in previous_chunk_hidden:
raise RuntimeError(
f"Predictor step {index} has no previous-chunk hidden"
)
predictor_output = self.predictor_v4(
target_latent=noisy_input,
target_timestep=timestep,
anchor_hidden=same_chunk_anchor,
previous_chunk_hidden=previous_chunk_hidden[index],
kv_cache=predictor_kv_cache,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame * self.frame_seq_length,
)
flow_pred = predictor_output["pred_flow"]
denoised_pred = self.generator._convert_flow_pred_to_x0(
flow_pred=flow_pred.flatten(0, 1),
xt=noisy_input.flatten(0, 1),
timestep=timestep.flatten(0, 1),
).unflatten(0, flow_pred.shape[:2])
same_chunk_anchor = predictor_output["pred_hidden"]
chunk_step_hidden[index] = same_chunk_anchor
elif reuse_velocity:
if first_step_flow_pred is None:
raise RuntimeError(
"Missing step-0 velocity for F-R-R-F denoising"
)
denoised_pred = self.generator._convert_flow_pred_to_x0(
flow_pred=first_step_flow_pred.flatten(0, 1),
xt=noisy_input.flatten(0, 1),
timestep=timestep.flatten(0, 1),
).unflatten(0, first_step_flow_pred.shape[:2])
else:
capture_hidden = (
self.predictor_v4 is not None
and (
index == 0
or (
temporal_block_index == 0
and index in predictor_steps
)
)
)
flow_pred, denoised_pred, final_hidden = (
self._full_step_with_optional_hidden(
noisy_input=noisy_input,
conditional_dict=conditional_dict,
timestep=timestep,
current_start=current_start_frame
* self.frame_seq_length,
capture_hidden=capture_hidden,
)
)
if final_hidden is not None:
chunk_step_hidden[index] = final_hidden
if index == 0:
same_chunk_anchor = final_hidden
if index == 0 and self.reuse_first_step_velocity:
first_step_flow_pred = flow_pred.detach()
if index < len(self.denoising_step_list) - 1:
next_timestep = self.denoising_step_list[index + 1]
noisy_input = self.scheduler.add_noise(
denoised_pred.flatten(0, 1),
torch.randn_like(denoised_pred.flatten(0, 1)),
next_timestep * torch.ones(
[batch_size * current_num_frames], device=noise.device, dtype=torch.long)
).unflatten(0, denoised_pred.shape[:2])
# Step 3.2: record the model's output
output[:, current_start_frame:current_start_frame + current_num_frames] = denoised_pred
if self.predictor_v4 is not None:
missing = set(predictor_steps).difference(chunk_step_hidden)
if missing:
raise RuntimeError(
f"Temporal chunk {temporal_block_index} lacks hidden steps "
f"{sorted(missing)}"
)
previous_chunk_hidden = {
step: chunk_step_hidden[step].detach()
for step in predictor_steps
}
# Step 3.3: rerun with timestep zero to update KV cache using clean context
context_timestep = torch.ones_like(timestep) * self.args.context_noise
self.generator(
noisy_image_or_video=denoised_pred,
conditional_dict=conditional_dict,
timestep=context_timestep,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame * self.frame_seq_length,
)
if profile:
block_end.record()
torch.cuda.synchronize()
block_time = block_start.elapsed_time(block_end)
block_times.append(block_time)
# Step 3.4: update the start and end frame indices
current_start_frame += current_num_frames
if profile:
# End diffusion timing and synchronize CUDA
diffusion_end.record()
torch.cuda.synchronize()
diffusion_time = diffusion_start.elapsed_time(diffusion_end)
init_time = init_start.elapsed_time(init_end)
vae_start.record()
# Step 4: Decode the output
video = self.vae.decode_to_pixel(output, use_cache=False)
video = (video * 0.5 + 0.5).clamp(0, 1)
if profile:
# End VAE timing and synchronize CUDA
vae_end.record()
torch.cuda.synchronize()
vae_time = vae_start.elapsed_time(vae_end)
total_time = init_time + diffusion_time + vae_time
print("Profiling results:")
print(f" - Initialization/caching time: {init_time:.2f} ms ({100 * init_time / total_time:.2f}%)")
print(f" - Diffusion generation time: {diffusion_time:.2f} ms ({100 * diffusion_time / total_time:.2f}%)")
for i, block_time in enumerate(block_times):
print(f" - Block {i} generation time: {block_time:.2f} ms ({100 * block_time / diffusion_time:.2f}% of diffusion)")
print(f" - VAE decoding time: {vae_time:.2f} ms ({100 * vae_time / total_time:.2f}%)")
print(f" - Total time: {total_time:.2f} ms")
if return_latents:
return video, output
else:
return video
def _initialize_kv_cache(self, batch_size, dtype, device):
"""
Initialize a Per-GPU KV cache for the Wan model.
"""
kv_cache1 = []
if self.local_attn_size != -1:
# Use the local attention size to compute the KV cache size
kv_cache_size = self.local_attn_size * self.frame_seq_length
else:
# Use the default KV cache size
kv_cache_size = 32760
for _ in range(self.num_transformer_blocks):
kv_cache1.append({
"k": torch.zeros([batch_size, kv_cache_size, 12, 128], dtype=dtype, device=device),
"v": torch.zeros([batch_size, kv_cache_size, 12, 128], 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)
})
self.kv_cache1 = kv_cache1 # always store the clean cache
def _initialize_crossattn_cache(self, batch_size, dtype, device):
"""
Initialize a Per-GPU cross-attention cache for the Wan model.
"""
crossattn_cache = []
for _ in range(self.num_transformer_blocks):
crossattn_cache.append({
"k": torch.zeros([batch_size, 512, 12, 128], dtype=dtype, device=device),
"v": torch.zeros([batch_size, 512, 12, 128], dtype=dtype, device=device),
"is_init": False
})
self.crossattn_cache = crossattn_cache