Compare commits

...
Author SHA1 Message Date
Will Lin a8a1d39de3 update 2026-01-15 15:21:04 -08:00
Will Lin b4edcceb90 update 2026-01-15 14:02:46 -08:00
Will Lin 464f6abff5 update 2026-01-15 12:39:27 -08:00
Will Lin def36d9973 update 2026-01-15 12:22:23 -08:00
Will Lin 5939883d31 update 2026-01-15 12:03:34 -08:00
39 changed files with 4218 additions and 2337 deletions
+23 -1
View File
@@ -168,6 +168,10 @@ Pipelines are composed of stages, each handling a specific part of the diffusion
- **DenoisingStage**: Performs denoising diffusion
- **DecodingStage**: Converts latents to pixels
Note: `DenoisingStage` uses the unified denoising engine under the hood. You
can inject a custom strategy via `strategy_cls` if a pipeline needs specialized
denoising behavior.
### Creating Your Pipeline
```python
@@ -230,7 +234,7 @@ class MyCustomPipeline(ComposedPipelineBase):
stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")
scheduler=self.get_module("scheduler"),
)
)
@@ -245,6 +249,24 @@ class MyCustomPipeline(ComposedPipelineBase):
EntryClass = MyCustomPipeline
```
### Customizing Denoising Strategies
If your model requires a custom denoising loop, pass a strategy class:
```python
from fastvideo.pipelines.stages.denoising_cosmos_strategy import (
CosmosStrategy)
self.add_stage(
stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
strategy_cls=CosmosStrategy,
)
)
```
### Creating Custom Stages (Optional)
If existing stages don't meet your needs, create custom ones:
+10
View File
@@ -38,8 +38,10 @@ if TYPE_CHECKING:
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
FASTVIDEO_SERVER_DEV_MODE: bool = False
FASTVIDEO_STAGE_LOGGING: bool = False
FASTVIDEO_DENOISING_PERF_LOGGING: bool = False
FASTVIDEO_HOST_IP: str = ""
FASTVIDEO_LOOPBACK_IP: str = ""
FASTVIDEO_DISABLE_PIN_MEMORY: str | None = None
def get_default_cache_root() -> str:
@@ -136,6 +138,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_LOOPBACK_IP":
lambda: os.getenv("FASTVIDEO_LOOPBACK_IP", ""),
# Disable pinned memory (e.g., on platforms that do not support it)
"FASTVIDEO_DISABLE_PIN_MEMORY":
lambda: os.getenv("FASTVIDEO_DISABLE_PIN_MEMORY", None),
# Number of GPUs per worker in Ray, if it is set to be a fraction,
# it allows ray to schedule multiple actors on a single GPU,
# so that users can colocate other actors on the same GPUs as FastVideo.
@@ -275,6 +281,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
# taken for each stage
"FASTVIDEO_STAGE_LOGGING":
lambda: bool(int(os.getenv("FASTVIDEO_STAGE_LOGGING", "0"))),
# Enable per-step denoising perf logging hooks
"FASTVIDEO_DENOISING_PERF_LOGGING":
lambda: bool(int(os.getenv("FASTVIDEO_DENOISING_PERF_LOGGING", "0"))),
}
# end-env-vars-definition
+6
View File
@@ -608,6 +608,7 @@ class FastVideoArgs:
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
from fastvideo.platforms import current_platform
from fastvideo.pin_memory import is_pin_memory_available
if current_platform.is_mps():
self.use_fsdp_inference = False
@@ -688,6 +689,11 @@ class FastVideoArgs:
self.pipeline_config.vae_config.load_encoder = True
self.preprocess_config.check_preprocess_config()
if self.pin_cpu_memory and not is_pin_memory_available():
logger.warning("Pinned memory is unavailable on this system; "
"disabling pin_cpu_memory.")
self.pin_cpu_memory = False
_current_fastvideo_args = None
+6 -2
View File
@@ -18,6 +18,7 @@ from fastvideo.layers.linear import (ColumnParallelLinear, LinearBase,
QKVParallelLinear, ReplicatedLinear,
RowParallelLinear)
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.pin_memory import is_pin_memory_available
from fastvideo.utils import get_mixed_precision_state
torch._dynamo.config.recompile_limit = 16
@@ -155,8 +156,11 @@ class BaseLayerWithLoRA(nn.Module):
get_local_torch_device(),
non_blocking=True).full_tensor().to(current_device))
offload_policy = CPUOffloadPolicy() if "cpu" in str(
current_device) else OffloadPolicy()
if "cpu" in str(current_device):
offload_policy = CPUOffloadPolicy(
pin_memory=is_pin_memory_available())
else:
offload_policy = OffloadPolicy()
mp_policy = get_mixed_precision_state().mp_policy
self.base_layer = fully_shard(unsharded_base_layer,
+4
View File
@@ -11,6 +11,8 @@ from typing import Dict, Set, Optional, Tuple
import torch
from fastvideo.pin_memory import is_pin_memory_available
class LayerwiseOffloadManager:
"""A lightweight layerwise CPU offload manager.
@@ -33,6 +35,8 @@ class LayerwiseOffloadManager:
self.module_list_attr = module_list_attr
self.num_layers = int(num_layers)
self.pin_cpu_memory = bool(pin_cpu_memory)
if self.pin_cpu_memory and not is_pin_memory_available():
self.pin_cpu_memory = False
self.enabled = bool(enabled and torch.cuda.is_available())
self.device = (
+6
View File
@@ -20,6 +20,7 @@ from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule,
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.logger import init_logger
from fastvideo.pin_memory import is_pin_memory_available
from fastvideo.models.loader.utils import (get_param_names_mapping,
hf_to_custom_state_dict)
from fastvideo.models.loader.weight_utils import safetensors_weights_iterator
@@ -212,6 +213,11 @@ def shard_model(
"mp_policy": mp_policy,
}
if cpu_offload:
if pin_cpu_memory and not is_pin_memory_available():
logger.warning(
"Pinned memory is unavailable; disabling pin_cpu_memory for "
"FSDP offload.")
pin_cpu_memory = False
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(
pin_memory=pin_cpu_memory)
+49
View File
@@ -0,0 +1,49 @@
# SPDX-License-Identifier: Apache-2.0
"""
Scheduler adapter interfaces for unified denoising.
"""
from __future__ import annotations
from typing import Any, Protocol
import torch
class SchedulerAdapter(Protocol):
def scale_model_input(self, latents: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
...
def step(self, noise_pred: torch.Tensor, t: torch.Tensor,
latents: torch.Tensor, **kwargs: Any) -> Any:
...
def add_noise(self, latents: torch.Tensor, noise: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
...
def set_timesteps(self, num_steps: int, device: torch.device | None = None,
**kwargs: Any) -> Any:
...
class DefaultSchedulerAdapter:
def __init__(self, scheduler: Any) -> None:
self.scheduler = scheduler
def scale_model_input(self, latents: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
return self.scheduler.scale_model_input(latents, t)
def step(self, noise_pred: torch.Tensor, t: torch.Tensor,
latents: torch.Tensor, **kwargs: Any) -> Any:
return self.scheduler.step(noise_pred, t, latents, **kwargs)
def add_noise(self, latents: torch.Tensor, noise: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
return self.scheduler.add_noise(latents, noise, t)
def set_timesteps(self, num_steps: int, device: torch.device | None = None,
**kwargs: Any) -> Any:
return self.scheduler.set_timesteps(num_steps, device=device, **kwargs)
+42
View File
@@ -0,0 +1,42 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import os
from functools import cache
import torch
from fastvideo.logger import init_logger
from fastvideo.platforms import current_platform
logger = init_logger(__name__)
def _probe_pin_memory() -> bool:
if os.getenv("FASTVIDEO_DISABLE_PIN_MEMORY",
"") not in ("", "0", "false", "False"):
return False
if current_platform.is_cpu() or current_platform.is_mps():
return False
try:
if torch.cuda.is_available():
torch.cuda.current_device()
_ = torch.empty(1024, device="cpu").pin_memory()
_ = torch.empty(1024, device="cpu", pin_memory=True)
except Exception as exc:
logger.warning("Pinned memory is unavailable: %s", exc)
return False
return True
@cache
def _cached_pin_memory_available(pid: int) -> bool:
return _probe_pin_memory()
def is_pin_memory_available() -> bool:
return _cached_pin_memory_available(os.getpid())
@@ -11,11 +11,12 @@ from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage, CosmosDenoisingStage,
from fastvideo.pipelines.stages import (ConditioningStage, DenoisingStage,
CosmosLatentPreparationStage,
DecodingStage, InputValidationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages.denoising_cosmos_strategy import CosmosStrategy
logger = init_logger(__name__)
@@ -73,9 +74,10 @@ class Cosmos2VideoToWorldPipeline(ComposedPipelineBase):
vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=CosmosDenoisingStage(
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
scheduler=self.get_module("scheduler"),
strategy_cls=CosmosStrategy))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
@@ -16,13 +16,15 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (
DecodingStage,
DenoisingStage,
InputValidationStage,
TextEncodingStage,
TimestepPreparationStage,
)
from fastvideo.pipelines.stages.denoising_longcat_strategy import (
LongCatI2VStrategy)
from fastvideo.pipelines.stages.longcat_image_vae_encoding import LongCatImageVAEEncodingStage
from fastvideo.pipelines.stages.longcat_i2v_latent_preparation import LongCatI2VLatentPreparationStage
from fastvideo.pipelines.stages.longcat_i2v_denoising import LongCatI2VDenoisingStage
from fastvideo.pipelines.stages.longcat_refine_init import LongCatRefineInitStage
from fastvideo.pipelines.stages.longcat_refine_timestep import LongCatRefineTimestepStage
@@ -132,12 +134,13 @@ class LongCatImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
# 8. Denoising with I2V support
self.add_stage(stage_name="denoising_stage",
stage=LongCatI2VDenoisingStage(
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self))
pipeline=self,
strategy_cls=LongCatI2VStrategy))
# 9. Decoding
self.add_stage(stage_name="decoding_stage",
@@ -11,12 +11,14 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (
DecodingStage,
DenoisingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
)
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
from fastvideo.pipelines.stages.denoising_longcat_strategy import (
LongCatStrategy)
from fastvideo.pipelines.stages.longcat_refine_init import LongCatRefineInitStage
from fastvideo.pipelines.stages.longcat_refine_timestep import LongCatRefineTimestepStage
@@ -127,12 +129,13 @@ class LongCatPipeline(LoRAPipeline, ComposedPipelineBase):
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=LongCatDenoisingStage(
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self))
pipeline=self,
strategy_cls=LongCatStrategy))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae"),
@@ -16,14 +16,16 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (
DecodingStage,
DenoisingStage,
InputValidationStage,
TextEncodingStage,
TimestepPreparationStage,
)
from fastvideo.pipelines.stages.denoising_longcat_strategy import (
LongCatVCStrategy)
from fastvideo.pipelines.stages.longcat_video_vae_encoding import LongCatVideoVAEEncodingStage
from fastvideo.pipelines.stages.longcat_i2v_latent_preparation import LongCatI2VLatentPreparationStage
from fastvideo.pipelines.stages.longcat_kv_cache_init import LongCatKVCacheInitStage
from fastvideo.pipelines.stages.longcat_vc_denoising import LongCatVCDenoisingStage
logger = init_logger(__name__)
@@ -131,12 +133,13 @@ class LongCatVideoContinuationPipeline(LoRAPipeline, ComposedPipelineBase):
# 7. Denoising with VC and KV cache support
self.add_stage(stage_name="denoising_stage",
stage=LongCatVCDenoisingStage(
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self))
pipeline=self,
strategy_cls=LongCatVCStrategy))
# 8. Decoding
self.add_stage(stage_name="decoding_stage",
@@ -11,10 +11,11 @@ from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
CausalDMDDenosingStage,
InputValidationStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage)
from fastvideo.pipelines.stages.denoising_causal_strategy import (
CausalBlockStrategy)
# isort: on
logger = init_logger(__name__)
@@ -47,11 +48,12 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=CausalDMDDenosingStage(
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
vae=self.get_module("vae"),
strategy_cls=CausalBlockStrategy))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
@@ -14,10 +14,11 @@ from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DmdDenoisingStage, InputValidationStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages.denoising_dmd_strategy import DmdStrategy
# isort: on
logger = init_logger(__name__)
@@ -63,9 +64,10 @@ class WanDMDPipeline(LoRAPipeline, ComposedPipelineBase):
use_btchw_layout=True))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
scheduler=FlowMatchEulerDiscreteScheduler(shift=8.0),
strategy_cls=DmdStrategy))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
@@ -13,12 +13,13 @@ from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DmdDenoisingStage,
ImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.stages.denoising_dmd_strategy import DmdStrategy
logger = init_logger(__name__)
@@ -69,9 +70,10 @@ class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
scheduler=FlowMatchEulerDiscreteScheduler(shift=8.0),
strategy_cls=DmdStrategy))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
+1 -9
View File
@@ -7,12 +7,9 @@ complete diffusion pipelines.
"""
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.pipelines.stages.conditioning import ConditioningStage
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.denoising import (CosmosDenoisingStage,
DenoisingStage,
DmdDenoisingStage)
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.image_encoding import (
ImageEncodingStage, MatrixGameImageEncodingStage, RefImageEncodingStage,
@@ -31,7 +28,6 @@ from fastvideo.pipelines.stages.timestep_preparation import (
# LongCat stages
from fastvideo.pipelines.stages.longcat_video_vae_encoding import LongCatVideoVAEEncodingStage
from fastvideo.pipelines.stages.longcat_kv_cache_init import LongCatKVCacheInitStage
from fastvideo.pipelines.stages.longcat_vc_denoising import LongCatVCDenoisingStage
__all__ = [
"PipelineStage",
@@ -41,10 +37,7 @@ __all__ = [
"CosmosLatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenosingStage",
"MatrixGameCausalDenoisingStage",
"CosmosDenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
@@ -58,5 +51,4 @@ __all__ = [
# LongCat stages
"LongCatVideoVAEEncodingStage",
"LongCatKVCacheInitStage",
"LongCatVCDenoisingStage",
]
@@ -1,497 +0,0 @@
import torch # type: ignore
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
logger = init_logger(__name__)
class CausalDMDDenosingStage(DenoisingStage):
"""
Denoising stage for causal diffusion.
"""
def __init__(self,
transformer,
scheduler,
transformer_2=None,
vae=None) -> None:
super().__init__(transformer, scheduler, transformer_2)
# KV and cross-attention cache state (initialized on first forward)
self.transformer = transformer
self.transformer_2 = transformer_2
self.vae = vae
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
try:
self.local_attn_size = getattr(self.transformer.model,
"local_attn_size",
-1) # type: ignore
except Exception:
self.local_attn_size = -1
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
patch_ratio = self.transformer.config.arch_config.patch_size[
-1] * self.transformer.config.arch_config.patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
# TODO(will): make this a parameter once we add i2v support
independent_first_frame = self.transformer.independent_first_frame if hasattr(
self.transformer, 'independent_first_frame') else False
# Timesteps for DMD
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
pos_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
# "encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
# STA
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
# Latents and prompts
assert batch.latents is not None, "latents must be provided"
latents = batch.latents # [B, C, T, H, W]
b, c, t, h, w = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset caches
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
# Initialize the low noise kv cache
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
def _get_kv_cache(timestep: float) -> list[dict]:
if boundary_timestep is not None:
if timestep >= boundary_timestep:
return kv_cache1
else:
assert kv_cache2 is not None, "kv_cache2 is not initialized"
return kv_cache2
return kv_cache1
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.text_encoder_configs[0].
arch_config.text_len,
dtype=target_dtype,
device=latents.device)
pos_start_base = 0
# Determine block sizes
if t % self.num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
)
num_blocks = t // self.num_frames_per_block
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
# For now hardcode the first block to be 1 frame assuming the model is Wan2.2-MoE
if boundary_timestep is not None:
block_sizes[0] = 1
first_frame_latent = None
if batch.pil_image is not None:
# Causal video gen directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert self.vae is not None, "VAE is not provided for causal video gen task"
self.vae = self.vae.to(get_local_torch_device())
first_frame_latent = self.vae.encode(batch.pil_image).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
first_frame_latent -= self.vae.shift_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
first_frame_latent = first_frame_latent * self.vae.scaling_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent = first_frame_latent * self.vae.scaling_factor
if fastvideo_args.vae_cpu_offload:
self.vae = self.vae.to("cpu")
# Fill the low noise and high noise kv cache with first_frame_latent and timestep 0
t_zero = torch.zeros([latents.shape[0], 1],
device=latents.device,
dtype=torch.long)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
self.transformer(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
if boundary_timestep is not None:
self.transformer_2(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += 1
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for current_num_frames in block_sizes:
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
# use BTCHW for DMD conversion routines
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
video_raw_latent_shape = noise_latents_btchw.shape
for i, t_cur in enumerate(timesteps):
if boundary_timestep is not None and t_cur < boundary_timestep:
current_model = self.transformer_2
else:
current_model = self.transformer
# Copy for pred conversion
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(target_dtype)
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(target_dtype)
],
dim=2)
# Prepare inputs
t_expand = t_cur.repeat(latent_model_input.shape[0])
# Attention metadata if needed
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(
)
attn_metadata = self.attn_metadata_builder.build( # type: ignore
current_timestep=i, # type: ignore
raw_latent_shape=(current_num_frames, h,
w), # type: ignore
patch_size=fastvideo_args.pipeline_config.
dit_config.patch_size, # type: ignore
STA_param=batch.STA_param, # type: ignore
VSA_sparsity=fastvideo_args.
VSA_sparsity, # type: ignore
device=get_local_torch_device(), # type: ignore
) # type: ignore
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=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
# Run transformer; follow DMD stage pattern
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], 1),
device=latent_model_input.device,
dtype=torch.long)
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=_get_kv_cache(t_cur),
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
# Convert pred noise to pred video with FM Euler scheduler utilities
if boundary_timestep is not None and t_cur >= 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) *
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)
noise = torch.randn(
video_raw_latent_shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else
batch.generator)).to(self.device)
noise_btchw = noise
if boundary_timestep is not None and i < len(
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) *
boundary_timestep).unflatten(
0, pred_video_btchw.shape[:2])
elif boundary_timestep is not None and i == len(
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)
if progress_bar is not None:
progress_bar.update()
# Write back and advance
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Re-run with context timestep to update KV cache using clean context
context_noise = getattr(fastvideo_args.pipeline_config,
"context_noise", 0)
t_context = torch.ones([latents.shape[0]],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(target_dtype)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
if boundary_timestep is not None:
self.transformer_2(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
self.transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += current_num_frames
if boundary_timestep is not None:
num_frames_to_remove = self.num_frames_per_block - 1
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
def _initialize_kv_cache(self, batch_size, dtype, device) -> list[dict]:
"""
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
kv_cache1 = []
num_attention_heads = self.transformer.num_attention_heads
attention_head_dim = self.transformer.attention_head_dim
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
for _ in range(self.num_transformer_blocks):
kv_cache1.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_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> list[dict]:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
crossattn_cache = []
num_attention_heads = self.transformer.num_attention_heads
attention_head_dim = self.transformer.attention_head_dim
for _ in range(self.num_transformer_blocks):
crossattn_cache.append({
"k":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"is_init":
False,
})
return crossattn_cache
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify denoising stage inputs."""
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
result.add_check("image_embeds", batch.image_embeds, V.is_list)
result.add_check("image_latent", batch.image_latent,
V.none_or_tensor_with_dims(5))
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
result.add_check("guidance_scale", batch.guidance_scale,
V.positive_float)
result.add_check("eta", batch.eta, V.non_negative_float)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
+21 -910
View File
@@ -13,41 +13,14 @@ from tqdm.auto import tqdm
from fastvideo.attention import get_attn_backend
from fastvideo.configs.pipelines.base import STA_Mode
from fastvideo.distributed import (get_local_torch_device, get_world_group)
from fastvideo.distributed import get_world_group
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import dict_to_3d_list, masks_like
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
try:
from fastvideo.attention.backends.vmoba import VMOBAAttentionBackend
from fastvideo.utils import is_vmoba_available
vmoba_attn_available = is_vmoba_available()
except ImportError:
vmoba_attn_available = False
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
logger = init_logger(__name__)
@@ -65,13 +38,15 @@ class DenoisingStage(PipelineStage):
scheduler,
pipeline=None,
transformer_2=None,
vae=None) -> None:
vae=None,
strategy_cls=None) -> None:
super().__init__()
self.transformer = transformer
self.transformer_2 = transformer_2
self.scheduler = scheduler
self.vae = vae
self.pipeline = weakref.ref(pipeline) if pipeline else None
self.strategy_cls = strategy_cls
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
@@ -92,416 +67,23 @@ class DenoisingStage(PipelineStage):
) -> ForwardBatch:
"""
Run the denoising loop.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
if pipeline:
pipeline.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
from fastvideo.pipelines.stages.denoising_standard_strategy import (
StandardStrategy)
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
strategy_cls = self.strategy_cls or StandardStrategy
engine = DenoisingEngine(strategy_cls(self),
hooks=self._build_engine_hooks())
return engine.run(batch, fastvideo_args)
# Setup precision and autocast settings
# TODO(will): make the precision configurable for inference
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
# TODO(will): remove this once we add input/output validation for stages
if timesteps is None:
raise ValueError("Timesteps must be provided")
num_inference_steps = batch.num_inference_steps
num_warmup_steps = len(
timesteps) - num_inference_steps * self.scheduler.order
# Prepare image latents and embeddings for I2V generation
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert not torch.isnan(
image_embeds[0]).any(), "image_embeds contains nan"
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_image": image_embeds,
"mask_strategy": dict_to_3d_list(
None, t_max=50, l_max=60, h_max=24)
},
)
pos_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
neg_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_neg,
"encoder_attention_mask": batch.negative_attention_mask,
},
)
action_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"mouse_cond": batch.mouse_cond,
"keyboard_cond": batch.keyboard_cond,
},
)
# Prepare STA parameters
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds
assert neg_prompt_embeds is not None
assert not torch.isnan(
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
if batch.boundary_ratio is not None:
logger.info("Overriding boundary ratio from %s to %s",
boundary_ratio, batch.boundary_ratio)
boundary_ratio = batch.boundary_ratio
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
else:
boundary_timestep = None
latent_model_input = latents.to(target_dtype)
assert latent_model_input.shape[0] == 1, "only support batch size 1"
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
# TI2V directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert batch.image_latent is None, "TI2V task should not have image latents"
assert self.vae is not None, "VAE is not provided for TI2V task"
z = self.vae.encode(batch.pil_image).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
z -= self.vae.shift_factor.to(z.device, z.dtype)
else:
z -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
z = z * self.vae.scaling_factor.to(z.device, z.dtype)
else:
z = z * self.vae.scaling_factor
latent_model_input = latent_model_input.squeeze(0)
_, mask2 = masks_like([latent_model_input], zero=True)
latent_model_input = (1. -
mask2[0]) * z + mask2[0] * latent_model_input
# latent_model_input = latent_model_input.unsqueeze(0)
latent_model_input = latent_model_input.to(get_local_torch_device())
latents = latent_model_input
F = batch.num_frames
temporal_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_temporal
spatial_scale = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
seq_len = ((F - 1) // temporal_scale +
1) * (batch.height // spatial_scale) * (
batch.width // spatial_scale) // (patch_size[1] *
patch_size[2])
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
# Skip if interrupted
if hasattr(self, 'interrupt') and self.interrupt:
continue
if boundary_timestep is None or t >= boundary_timestep:
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and self.transformer_2 is not None and next(
self.transformer_2.parameters()).device.type
== 'cuda'):
self.transformer_2.to('cpu')
current_model = self.transformer
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and not fastvideo_args.use_fsdp_inference
and current_model is not None):
transformer_device = next(
current_model.parameters()).device.type
if transformer_device == 'cpu':
current_model.to(get_local_torch_device())
current_guidance_scale = batch.guidance_scale
else:
# low-noise stage in wan2.2
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and next(self.transformer.parameters()).device.type
== 'cuda'):
self.transformer.to('cpu')
current_model = self.transformer_2
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and not fastvideo_args.use_fsdp_inference
and current_model is not None):
transformer_2_device = next(
current_model.parameters()).device.type
if transformer_2_device == 'cpu':
current_model.to(get_local_torch_device())
current_guidance_scale = batch.guidance_scale_2
assert current_model is not None, "current_model is None"
# Expand latents for V2V/I2V
latent_model_input = latents.to(target_dtype)
if batch.video_latent is not None:
latent_model_input = torch.cat([
latent_model_input, batch.video_latent,
torch.zeros_like(latents)
],
dim=1).to(target_dtype)
elif batch.image_latent is not None:
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
dim=1).to(target_dtype)
assert not torch.isnan(
latent_model_input).any(), "latent_model_input contains nan"
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
timestep = torch.stack([t]).to(get_local_torch_device())
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep).flatten()
temp_ts = torch.cat([
temp_ts,
temp_ts.new_ones(seq_len - temp_ts.size(0)) * timestep
])
timestep = temp_ts.unsqueeze(0)
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
else:
t_expand = t.repeat(latent_model_input.shape[0])
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t)
# Prepare inputs for transformer
guidance_expand = (
torch.tensor(
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_local_torch_device(),
).to(target_dtype) *
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
is not None else None)
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
if (st_attn_available
and self.attn_backend == SlidingTileAttentionBackend
) or (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(
)
# TODO(will): clean this up
attn_metadata = self.attn_metadata_builder.build( # type: ignore
current_timestep=i, # type: ignore
raw_latent_shape=batch.
raw_latent_shape[2:5], # type: ignore
patch_size=fastvideo_args.
pipeline_config. # type: ignore
dit_config.patch_size, # type: ignore
STA_param=batch.STA_param, # type: ignore
VSA_sparsity=fastvideo_args.
VSA_sparsity, # type: ignore
device=get_local_torch_device(),
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
elif (vmoba_attn_available
and self.attn_backend == VMOBAAttentionBackend):
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(
)
# Prepare V-MoBA parameters from config
moba_params = fastvideo_args.moba_config.copy()
moba_params.update({
"current_timestep":
i,
"raw_latent_shape":
batch.raw_latent_shape[2:5],
"patch_size":
fastvideo_args.pipeline_config.dit_config.
patch_size,
"device":
get_local_torch_device(),
})
attn_metadata = self.attn_metadata_builder.build(
**moba_params)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
# TODO(will): finalize the interface. vLLM uses this to
# support torch dynamo compilation. They pass in
# attn_metadata, vllm_config, and num_tokens. We can pass in
# fastvideo_args or training_args, and attn_metadata.
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch,
# fastvideo_args=fastvideo_args
):
# Run transformer
noise_pred = current_model(
latent_model_input,
prompt_embeds,
t_expand,
guidance=guidance_expand,
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
)
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch,
):
noise_pred_uncond = current_model(
latent_model_input,
neg_prompt_embeds,
t_expand,
guidance=guidance_expand,
**image_kwargs,
**neg_cond_kwargs,
**action_kwargs,
)
noise_pred_text = noise_pred
noise_pred = noise_pred_uncond + current_guidance_scale * (
noise_pred_text - noise_pred_uncond)
# Apply guidance rescale if needed
if batch.guidance_rescale > 0.0:
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
noise_pred = self.rescale_noise_cfg(
noise_pred,
noise_pred_text,
guidance_rescale=batch.guidance_rescale,
)
# Compute the previous noisy sample
latents = self.scheduler.step(noise_pred,
t,
latents,
**extra_step_kwargs,
return_dict=False)[0]
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
latents = latents.squeeze(0)
latents = (1. - mask2[0]) * z + mask2[0] * latents
# latents = latents.unsqueeze(0)
# save trajectory latents if needed
if batch.return_trajectory_latents:
trajectory_timesteps.append(t)
trajectory_latents.append(latents)
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0
and progress_bar is not None):
progress_bar.update()
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
dim=0)
else:
trajectory_tensor = None
trajectory_timesteps_tensor = None
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
batch.trajectory_latents = trajectory_tensor.cpu()
# Update batch with final latents
batch.latents = latents
if fastvideo_args.dit_layerwise_offload:
mgr = getattr(self.transformer, "_layerwise_offload_manager", None)
if mgr is not None and getattr(mgr, "enabled", False):
mgr.release_all()
if self.transformer_2 is not None:
mgr2 = getattr(self.transformer_2, "_layerwise_offload_manager",
None)
if mgr2 is not None and getattr(mgr2, "enabled", False):
mgr2.release_all()
# Save STA mask search results if needed
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
self.save_sta_search_results(batch)
# deallocate transformer if on mps
if torch.backends.mps.is_available():
logger.info("Memory before deallocating transformer: %s",
torch.mps.current_allocated_memory())
del self.transformer
if pipeline is not None and "transformer" in pipeline.modules:
del pipeline.modules["transformer"]
fastvideo_args.model_loaded["transformer"] = False
logger.info("Memory after deallocating transformer: %s",
torch.mps.current_allocated_memory())
return batch
def _build_engine_hooks(self):
import fastvideo.envs as envs
if not envs.FASTVIDEO_DENOISING_PERF_LOGGING:
return []
from fastvideo.pipelines.stages.denoising_engine_hooks import (
PerfLoggingHook)
return [PerfLoggingHook()]
def prepare_extra_func_kwargs(self, func, kwargs) -> dict[str, Any]:
"""
@@ -712,8 +294,9 @@ class DenoisingStage(PipelineStage):
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify denoising stage inputs."""
result = VerificationResult()
result.add_check("timesteps", batch.timesteps,
[V.is_tensor, V.min_dims(1)])
if batch.timesteps is not None:
result.add_check("timesteps", batch.timesteps,
[V.is_tensor, V.min_dims(1)])
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
@@ -741,475 +324,3 @@ class DenoisingStage(PipelineStage):
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
return result
class CosmosDenoisingStage(DenoisingStage):
"""
Denoising stage for Cosmos models using FlowMatchEulerDiscreteScheduler.
"""
def __init__(self, transformer, scheduler, pipeline=None) -> None:
super().__init__(transformer, scheduler, pipeline)
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
if pipeline:
pipeline.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
if hasattr(self.transformer, 'module'):
transformer_dtype = next(self.transformer.module.parameters()).dtype
else:
transformer_dtype = next(self.transformer.parameters()).dtype
target_dtype = transformer_dtype
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latents = batch.latents
num_inference_steps = batch.num_inference_steps
guidance_scale = batch.guidance_scale
sigma_max = 80.0
sigma_min = 0.002
sigma_data = 1.0
final_sigmas_type = "sigma_min"
if self.scheduler is not None:
self.scheduler.register_to_config(
sigma_max=sigma_max,
sigma_min=sigma_min,
sigma_data=sigma_data,
final_sigmas_type=final_sigmas_type,
)
self.scheduler.set_timesteps(num_inference_steps, device=latents.device)
timesteps = self.scheduler.timesteps
if (hasattr(self.scheduler.config, 'final_sigmas_type')
and self.scheduler.config.final_sigmas_type == "sigma_min"
and len(self.scheduler.sigmas) > 1):
self.scheduler.sigmas[-1] = self.scheduler.sigmas[-2]
conditioning_latents = getattr(batch, 'conditioning_latents', None)
unconditioning_latents = conditioning_latents
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if hasattr(self, 'interrupt') and self.interrupt:
continue
current_sigma = self.scheduler.sigmas[i]
current_t = current_sigma / (current_sigma + 1)
c_in = 1 - current_t
c_skip = 1 - current_t
c_out = -current_t
timestep = current_t.view(1, 1, 1, 1,
1).expand(latents.size(0), -1,
latents.size(2), -1,
-1) # [B, 1, T, 1, 1]
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
cond_latent = latents * c_in
if hasattr(
batch, 'cond_indicator'
) and batch.cond_indicator is not None and conditioning_latents is not None:
cond_latent = batch.cond_indicator * conditioning_latents + (
1 - batch.cond_indicator) * cond_latent
else:
logger.warning(
"Step %s: Missing conditioning data - cond_indicator: %s, conditioning_latents: %s",
i, hasattr(batch, 'cond_indicator'),
conditioning_latents is not None)
cond_latent = cond_latent.to(target_dtype)
cond_timestep = timestep
if hasattr(batch, 'cond_indicator'
) and batch.cond_indicator is not None:
sigma_conditioning = 0.0001
t_conditioning = sigma_conditioning / (
sigma_conditioning + 1)
cond_timestep = batch.cond_indicator * t_conditioning + (
1 - batch.cond_indicator) * timestep
cond_timestep = cond_timestep.to(target_dtype)
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
# Use conditioning masks from CosmosLatentPreparationStage
condition_mask = batch.cond_mask.to(
target_dtype) if hasattr(batch,
'cond_mask') else None
padding_mask = torch.zeros(1,
1,
batch.height,
batch.width,
device=cond_latent.device,
dtype=target_dtype)
# Fallback if masks not available
if condition_mask is None:
batch_size, num_channels, num_frames, height, width = cond_latent.shape
condition_mask = torch.zeros(
batch_size,
1,
num_frames,
height,
width,
device=cond_latent.device,
dtype=target_dtype)
noise_pred = self.transformer(
hidden_states=cond_latent,
timestep=cond_timestep.to(target_dtype),
encoder_hidden_states=batch.prompt_embeds[0].to(
target_dtype),
fps=24, # TODO: get fps from batch or config
condition_mask=condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
cond_pred = (c_skip * latents +
c_out * noise_pred.float()).to(target_dtype)
if hasattr(
batch, 'cond_indicator'
) and batch.cond_indicator is not None and conditioning_latents is not None:
cond_pred = batch.cond_indicator * conditioning_latents + (
1 - batch.cond_indicator) * cond_pred
if batch.do_classifier_free_guidance and batch.negative_prompt_embeds is not None:
uncond_latent = latents * c_in
if hasattr(
batch, 'uncond_indicator'
) and batch.uncond_indicator is not None and unconditioning_latents is not None:
uncond_latent = batch.uncond_indicator * unconditioning_latents + (
1 - batch.uncond_indicator) * uncond_latent
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
uncond_condition_mask = batch.uncond_mask.to(
target_dtype
) if hasattr(
batch, 'uncond_mask'
) and batch.uncond_mask is not None else condition_mask
uncond_timestep = timestep
if hasattr(batch, 'uncond_indicator'
) and batch.uncond_indicator is not None:
sigma_conditioning = 0.0001
t_conditioning = sigma_conditioning / (
sigma_conditioning + 1)
uncond_timestep = batch.uncond_indicator * t_conditioning + (
1 - batch.uncond_indicator) * timestep
uncond_timestep = uncond_timestep.to(
target_dtype)
noise_pred_uncond = self.transformer(
hidden_states=uncond_latent.to(target_dtype),
timestep=uncond_timestep.to(target_dtype),
encoder_hidden_states=batch.
negative_prompt_embeds[0].to(target_dtype),
fps=24, # TODO: get fps from batch or config
condition_mask=uncond_condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
uncond_pred = (
c_skip * latents +
c_out * noise_pred_uncond.float()).to(target_dtype)
if hasattr(
batch, 'uncond_indicator'
) and batch.uncond_indicator is not None and unconditioning_latents is not None:
uncond_pred = batch.uncond_indicator * unconditioning_latents + (
1 - batch.uncond_indicator) * uncond_pred
guidance_diff = cond_pred - uncond_pred
final_pred = cond_pred + guidance_scale * guidance_diff
else:
final_pred = cond_pred
# Convert to noise for scheduler step
if current_sigma > 1e-8:
noise_for_scheduler = (latents - final_pred) / current_sigma
else:
logger.warning(
"Step %s: current_sigma too small (%s), using final_pred directly",
i, current_sigma)
noise_for_scheduler = final_pred
if torch.isnan(noise_for_scheduler).sum() > 0:
logger.error(
"Step %s: NaN detected in noise_for_scheduler, sum: %s",
i,
noise_for_scheduler.float().sum().item())
logger.error(
"Step %s: latents sum: %s, final_pred sum: %s, current_sigma: %s",
i,
latents.float().sum().item(),
final_pred.float().sum().item(), current_sigma)
latents = self.scheduler.step(noise_for_scheduler,
t,
latents,
**extra_step_kwargs,
return_dict=False)[0]
progress_bar.update()
batch.latents = latents
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify Cosmos denoising stage inputs."""
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
result.add_check("guidance_scale", batch.guidance_scale,
V.positive_float)
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify Cosmos denoising stage outputs."""
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
return result
class DmdDenoisingStage(DenoisingStage):
"""
Denoising stage for DMD.
"""
def __init__(self, transformer, scheduler) -> None:
super().__init__(transformer, scheduler)
self.scheduler = FlowMatchEulerDiscreteScheduler(shift=8.0)
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Run the denoising loop.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
# Setup precision and autocast settings
# TODO(will): make the precision configurable for inference
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
# TODO(will): remove this once we add input/output validation for stages
if timesteps is None:
raise ValueError("Timesteps must be provided")
num_inference_steps = batch.num_inference_steps
num_warmup_steps = len(
timesteps) - num_inference_steps * self.scheduler.order
# Prepare image latents and embeddings for I2V generation
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_image": image_embeds,
"mask_strategy": dict_to_3d_list(
None, t_max=50, l_max=60, h_max=24)
},
)
pos_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
# Prepare STA parameters
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
# Get latents and embeddings
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
video_raw_latent_shape = latents.shape
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
# Run denoising loop
with self.progress_bar(total=len(timesteps)) as progress_bar:
for i, t in enumerate(timesteps):
# Skip if interrupted
if hasattr(self, 'interrupt') and self.interrupt:
continue
# Expand latents for I2V
noise_latents = latents.clone()
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.permute(0, 2, 1, 3, 4)
],
dim=2).to(target_dtype)
assert not torch.isnan(
latent_model_input).any(), "latent_model_input contains nan"
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (
torch.tensor(
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_local_torch_device(),
).to(target_dtype) *
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
is not None else None)
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
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(
)
# TODO(will): clean this up
attn_metadata = self.attn_metadata_builder.build( # type: ignore
current_timestep=i, # type: ignore
raw_latent_shape=batch.
raw_latent_shape[2:5], # type: ignore
patch_size=fastvideo_args.
pipeline_config. # type: ignore
dit_config.patch_size, # type: ignore
STA_param=batch.STA_param, # type: ignore
VSA_sparsity=fastvideo_args.
VSA_sparsity, # type: ignore
device=get_local_torch_device(), # type: ignore
) # type: ignore
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch,
# fastvideo_args=fastvideo_args
):
# Run transformer
pred_noise = self.transformer(
latent_model_input.permute(0, 2, 1, 3, 4),
prompt_embeds,
t_expand,
guidance=guidance_expand,
**image_kwargs,
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
[1], dtype=torch.long, device=pred_video.device)
noise = torch.randn(video_raw_latent_shape,
dtype=pred_video.dtype,
generator=batch.generator[0]).to(
self.device)
latents = self.scheduler.add_noise(
pred_video.flatten(0, 1), noise.flatten(0, 1),
next_timestep).unflatten(0, pred_video.shape[:2])
else:
latents = pred_video
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0
and progress_bar is not None):
progress_bar.update()
# Gather results if using sequence parallelism
latents = latents.permute(0, 2, 1, 3, 4)
# Update batch with final latents
batch.latents = latents
return batch
@@ -0,0 +1,623 @@
# SPDX-License-Identifier: Apache-2.0
"""
Causal block denoising strategy (DMD).
"""
from __future__ import annotations
from typing import Any
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_strategies import (
BlockContext,
BlockDenoisingStrategy,
BlockPlan,
BlockPlanItem,
ModelInputs,
StrategyState,
)
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
class CausalBlockStrategy(BlockDenoisingStrategy):
def __init__(self, stage: Any) -> None:
self.stage = stage
self.num_transformer_blocks = 0
self.num_frames_per_block = 0
self.sliding_window_num_frames = 0
self.local_attn_size = -1
self.frame_seq_length = 0
def _ensure_model_constants(self) -> None:
transformer = self.stage.transformer
self.num_transformer_blocks = len(transformer.blocks)
arch_config = transformer.config.arch_config
self.num_frames_per_block = arch_config.num_frames_per_block
self.sliding_window_num_frames = arch_config.sliding_window_num_frames
try:
self.local_attn_size = getattr(transformer.model, "local_attn_size",
-1)
except Exception:
self.local_attn_size = -1
def _initialize_kv_cache(self, batch_size, dtype, device) -> list[dict]:
kv_cache1 = []
num_attention_heads = self.stage.transformer.num_attention_heads
attention_head_dim = self.stage.transformer.attention_head_dim
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
for _ in range(self.num_transformer_blocks):
kv_cache1.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_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> list[dict]:
crossattn_cache = []
num_attention_heads = self.stage.transformer.num_attention_heads
attention_head_dim = self.stage.transformer.attention_head_dim
for _ in range(self.num_transformer_blocks):
crossattn_cache.append({
"k":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"is_init":
False,
})
return crossattn_cache
def prepare(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> StrategyState:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
self._ensure_model_constants()
latents = batch.latents
if latents is None:
raise ValueError("latents must be provided")
latent_seq_length = latents.shape[-1] * latents.shape[-2]
patch_ratio = (
self.stage.transformer.config.arch_config.patch_size[-1] *
self.stage.transformer.config.arch_config.patch_size[-2])
self.frame_seq_length = latent_seq_length // patch_ratio
independent_first_frame = getattr(self.stage.transformer,
"independent_first_frame", False)
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if getattr(fastvideo_args.pipeline_config, "warp_denoising_step",
False):
scheduler_timesteps = torch.cat((
self.stage.scheduler.timesteps.cpu(),
torch.tensor([0], dtype=torch.float32),
))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
if boundary_ratio is not None:
boundary_timestep = (boundary_ratio *
self.stage.scheduler.num_train_timesteps)
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
image_kwargs: dict[str, Any] = {}
pos_cond_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.transformer.forward,
{
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
if (st_attn_available
and self.stage.attn_backend == SlidingTileAttentionBackend):
self.stage.prepare_sta_param(batch, fastvideo_args)
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
kv_cache2 = self._initialize_kv_cache(
batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device,
)
text_len = None
if fastvideo_args.pipeline_config.text_encoder_configs:
text_len = getattr(
fastvideo_args.pipeline_config.text_encoder_configs[0].
arch_config, "text_len", None)
if not text_len:
if batch.prompt_attention_mask:
text_len = batch.prompt_attention_mask[0].shape[-1]
elif batch.prompt_embeds:
text_len = batch.prompt_embeds[0].shape[1]
else:
text_len = 0
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=text_len,
dtype=target_dtype,
device=latents.device,
)
num_frames = latents.shape[2]
if num_frames % self.num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for "
"causal DMD denoising")
num_blocks = num_frames // self.num_frames_per_block
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
if boundary_timestep is not None:
block_sizes[0] = 1
pos_start_base = 0
if batch.pil_image is not None:
assert self.stage.vae is not None, (
"VAE is not provided for causal video gen task")
self.stage.vae = self.stage.vae.to(get_local_torch_device())
first_frame_latent = self.stage.vae.encode(
batch.pil_image).mean.float()
if (hasattr(self.stage.vae, "shift_factor")
and self.stage.vae.shift_factor is not None):
if isinstance(self.stage.vae.shift_factor, torch.Tensor):
first_frame_latent -= self.stage.vae.shift_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent -= self.stage.vae.shift_factor
if isinstance(self.stage.vae.scaling_factor, torch.Tensor):
first_frame_latent = (
first_frame_latent * self.stage.vae.scaling_factor.to(
first_frame_latent.device, first_frame_latent.dtype))
else:
first_frame_latent = (first_frame_latent *
self.stage.vae.scaling_factor)
if fastvideo_args.vae_cpu_offload:
self.stage.vae = self.stage.vae.to("cpu")
t_zero = torch.zeros([latents.shape[0], 1],
device=latents.device,
dtype=torch.long)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
self.stage.transformer(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
if boundary_timestep is not None:
self.stage.transformer_2(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += 1
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
progress_bar = self.stage.progress_bar(total=len(block_sizes) *
len(timesteps))
extra: dict[str, Any] = {
"batch":
batch,
"fastvideo_args":
fastvideo_args,
"target_dtype":
target_dtype,
"autocast_enabled":
autocast_enabled,
"boundary_timestep":
boundary_timestep,
"high_noise_timesteps":
high_noise_timesteps,
"image_kwargs":
image_kwargs,
"pos_cond_kwargs":
pos_cond_kwargs,
"kv_cache1":
kv_cache1,
"kv_cache2":
kv_cache2,
"crossattn_cache":
crossattn_cache,
"block_sizes":
block_sizes,
"start_index":
start_index,
"progress_bar":
progress_bar,
"independent_first_frame":
independent_first_frame,
"context_noise":
getattr(fastvideo_args.pipeline_config, "context_noise", 0),
"pos_start_base":
pos_start_base,
}
return StrategyState(
latents=latents,
timesteps=timesteps,
num_inference_steps=len(timesteps),
prompt_embeds=batch.prompt_embeds,
negative_prompt_embeds=batch.negative_prompt_embeds,
prompt_attention_mask=batch.prompt_attention_mask,
negative_attention_mask=batch.negative_attention_mask,
image_embeds=batch.image_embeds,
guidance_scale=batch.guidance_scale,
guidance_scale_2=batch.guidance_scale_2,
guidance_rescale=batch.guidance_rescale,
do_cfg=batch.do_classifier_free_guidance,
extra=extra,
)
def block_plan(self, state: StrategyState) -> BlockPlan:
block_sizes = state.extra["block_sizes"]
start_index = state.extra["start_index"]
items: list[BlockPlanItem] = []
for block_size in block_sizes:
items.append(
BlockPlanItem(
start_index=start_index,
num_frames=block_size,
use_kv_cache=True,
model_selector="default",
))
start_index += block_size
return BlockPlan(items=items)
def init_block_context(self, state: StrategyState,
block_item: BlockPlanItem,
block_idx: int) -> BlockContext:
return BlockContext(
kv_cache=state.extra["kv_cache1"],
kv_cache_2=state.extra["kv_cache2"],
crossattn_cache=state.extra["crossattn_cache"],
action_cache=None,
extra={
"block_idx": block_idx,
"start_index": block_item.start_index,
"num_frames": block_item.num_frames,
},
)
def process_block(self, state: StrategyState, block_ctx: BlockContext,
block_item: BlockPlanItem) -> None:
batch = state.extra["batch"]
fastvideo_args = state.extra["fastvideo_args"]
target_dtype = state.extra["target_dtype"]
autocast_enabled = state.extra["autocast_enabled"]
boundary_timestep = state.extra["boundary_timestep"]
high_noise_timesteps = state.extra["high_noise_timesteps"]
image_kwargs = state.extra["image_kwargs"]
pos_cond_kwargs = state.extra["pos_cond_kwargs"]
progress_bar = state.extra["progress_bar"]
independent_first_frame = state.extra["independent_first_frame"]
start_index = block_item.start_index
current_num_frames = block_item.num_frames
kv_cache1 = block_ctx.kv_cache
kv_cache2 = block_ctx.kv_cache_2
crossattn_cache = block_ctx.crossattn_cache
def _get_kv_cache(timestep_val: float) -> list[dict]:
if boundary_timestep is not None:
if timestep_val >= boundary_timestep:
return kv_cache1
if kv_cache2 is None:
raise ValueError("kv_cache2 is not initialized")
return kv_cache2
return kv_cache1
current_latents = state.latents[:, :, start_index:start_index +
current_num_frames, :, :]
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
video_raw_latent_shape = noise_latents_btchw.shape
h, w = current_latents.shape[-2:]
attn_metadata = None
for i, t_cur in enumerate(state.timesteps):
if boundary_timestep is not None and t_cur < boundary_timestep:
current_model = self.stage.transformer_2
else:
current_model = self.stage.transformer
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(target_dtype)
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(target_dtype)],
dim=2)
t_expand = t_cur.repeat(latent_model_input.shape[0])
if (vsa_available
and self.stage.attn_backend == VideoSparseAttentionBackend):
self.attn_metadata_builder_cls = (
self.stage.attn_backend.get_builder_cls())
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = (
self.attn_metadata_builder_cls())
attn_metadata = self.attn_metadata_builder.build( # type: ignore
current_timestep=i, # type: ignore
raw_latent_shape=(current_num_frames, h,
w), # type: ignore
patch_size=fastvideo_args.pipeline_config.dit_config.
patch_size, # type: ignore
STA_param=batch.STA_param, # type: ignore
VSA_sparsity=fastvideo_args.
VSA_sparsity, # type: ignore
device=get_local_torch_device(), # type: ignore
)
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=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], 1),
device=latent_model_input.device,
dtype=torch.long)
pred_noise_btchw = current_model(
latent_model_input,
batch.prompt_embeds,
t_expanded_noise,
kv_cache=_get_kv_cache(t_cur),
crossattn_cache=crossattn_cache,
current_start=start_index * self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
if boundary_timestep is not None and t_cur >= 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) *
boundary_timestep,
scheduler=self.stage.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.stage.scheduler,
).unflatten(0, pred_noise_btchw.shape[:2])
if i < len(state.timesteps) - 1:
next_timestep = state.timesteps[i + 1] * torch.ones(
[1],
dtype=torch.long,
device=pred_video_btchw.device,
)
noise = torch.randn(
video_raw_latent_shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else batch.generator),
).to(self.stage.device)
noise_btchw = noise
if (boundary_timestep is not None
and high_noise_timesteps is not None
and i < len(high_noise_timesteps) - 1):
noise_latents_btchw = self.stage.scheduler.add_noise_high(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep,
torch.ones_like(next_timestep) * boundary_timestep,
).unflatten(0, pred_video_btchw.shape[:2])
elif (boundary_timestep is not None
and high_noise_timesteps is not None
and i == len(high_noise_timesteps) - 1):
noise_latents_btchw = pred_video_btchw
else:
noise_latents_btchw = self.stage.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)
if progress_bar is not None:
progress_bar.update()
block_ctx.extra["attn_metadata"] = attn_metadata
state.latents[:, :, start_index:start_index +
current_num_frames, :, :] = (current_latents)
def update_context(self, state: StrategyState, block_ctx: BlockContext,
block_item: BlockPlanItem) -> None:
batch = state.extra["batch"]
target_dtype = state.extra["target_dtype"]
autocast_enabled = state.extra["autocast_enabled"]
boundary_timestep = state.extra["boundary_timestep"]
kv_cache1 = state.extra["kv_cache1"]
kv_cache2 = state.extra["kv_cache2"]
crossattn_cache = state.extra["crossattn_cache"]
image_kwargs = state.extra["image_kwargs"]
pos_cond_kwargs = state.extra["pos_cond_kwargs"]
context_noise = state.extra["context_noise"]
start_index = block_item.start_index
current_num_frames = block_item.num_frames
current_latents = state.latents[:, :, start_index:start_index +
current_num_frames, :, :]
latents_device = current_latents.device
t_context = torch.ones([current_latents.shape[0]],
device=latents_device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(target_dtype)
attn_metadata = block_ctx.extra.get("attn_metadata")
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
if boundary_timestep is not None:
self.stage.transformer_2(
context_bcthw,
batch.prompt_embeds,
t_expanded_context,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=start_index * self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
self.stage.transformer(
context_bcthw,
batch.prompt_embeds,
t_expanded_context,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=start_index * self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
def postprocess(self, state: StrategyState) -> ForwardBatch:
progress_bar = state.extra.get("progress_bar")
if progress_bar is not None:
progress_bar.close()
batch = state.extra["batch"]
boundary_timestep = state.extra["boundary_timestep"]
latents = state.latents
if boundary_timestep is not None:
num_frames_to_remove = self.num_frames_per_block - 1
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
raise NotImplementedError
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
raise NotImplementedError
def cfg_combine(self, state: StrategyState,
noise_pred: torch.Tensor) -> torch.Tensor:
raise NotImplementedError
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
raise NotImplementedError
@@ -0,0 +1,325 @@
# SPDX-License-Identifier: Apache-2.0
"""
Cosmos denoising strategy using FlowMatchEulerDiscreteScheduler.
"""
from __future__ import annotations
from typing import Any
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_strategies import (
DenoisingStrategy,
ModelInputs,
StrategyState,
)
logger = init_logger(__name__)
class CosmosStrategy(DenoisingStrategy):
def __init__(self, stage: Any) -> None:
self.stage = stage
def prepare(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> StrategyState:
pipeline = self.stage.pipeline() if self.stage.pipeline else None
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.stage.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
if pipeline:
pipeline.add_module("transformer", self.stage.transformer)
fastvideo_args.model_loaded["transformer"] = True
extra_step_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
if hasattr(self.stage.transformer, "module"):
transformer_dtype = next(
self.stage.transformer.module.parameters()).dtype
else:
transformer_dtype = next(self.stage.transformer.parameters()).dtype
target_dtype = transformer_dtype
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latents = batch.latents
num_inference_steps = batch.num_inference_steps
sigma_max = 80.0
sigma_min = 0.002
sigma_data = 1.0
final_sigmas_type = "sigma_min"
if self.stage.scheduler is not None:
self.stage.scheduler.register_to_config(
sigma_max=sigma_max,
sigma_min=sigma_min,
sigma_data=sigma_data,
final_sigmas_type=final_sigmas_type,
)
self.stage.scheduler.set_timesteps(num_inference_steps,
device=latents.device)
timesteps = self.stage.scheduler.timesteps
if (hasattr(self.stage.scheduler.config, "final_sigmas_type")
and self.stage.scheduler.config.final_sigmas_type == "sigma_min"
and len(self.stage.scheduler.sigmas) > 1):
self.stage.scheduler.sigmas[-1] = self.stage.scheduler.sigmas[-2]
conditioning_latents = getattr(batch, "conditioning_latents", None)
unconditioning_latents = conditioning_latents
progress_bar = self.stage.progress_bar(total=num_inference_steps)
extra: dict[str, Any] = {
"batch": batch,
"fastvideo_args": fastvideo_args,
"extra_step_kwargs": extra_step_kwargs,
"target_dtype": target_dtype,
"autocast_enabled": autocast_enabled,
"progress_bar": progress_bar,
"conditioning_latents": conditioning_latents,
"unconditioning_latents": unconditioning_latents,
}
return StrategyState(
latents=latents,
timesteps=timesteps,
num_inference_steps=num_inference_steps,
prompt_embeds=batch.prompt_embeds,
negative_prompt_embeds=batch.negative_prompt_embeds,
prompt_attention_mask=batch.prompt_attention_mask,
negative_attention_mask=batch.negative_attention_mask,
image_embeds=batch.image_embeds,
guidance_scale=batch.guidance_scale,
guidance_scale_2=batch.guidance_scale_2,
guidance_rescale=batch.guidance_rescale,
do_cfg=batch.do_classifier_free_guidance,
extra=extra,
)
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
state.extra["step_idx"] = step_idx
return ModelInputs(
latent_model_input=state.latents,
timestep=t,
prompt_embeds=state.prompt_embeds,
prompt_attention_mask=state.prompt_attention_mask,
)
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
batch = state.extra["batch"]
target_dtype = state.extra["target_dtype"]
autocast_enabled = state.extra["autocast_enabled"]
conditioning_latents = state.extra["conditioning_latents"]
unconditioning_latents = state.extra["unconditioning_latents"]
step_idx = state.extra["step_idx"]
if getattr(self.stage, "interrupt", False):
return state.latents
current_sigma = self.stage.scheduler.sigmas[step_idx]
current_t = current_sigma / (current_sigma + 1)
c_in = 1 - current_t
c_skip = 1 - current_t
c_out = -current_t
timestep = current_t.view(1, 1, 1, 1,
1).expand(state.latents.size(0), -1,
state.latents.size(2), -1, -1)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
cond_latent = state.latents * c_in
if (hasattr(batch, "cond_indicator")
and batch.cond_indicator is not None
and conditioning_latents is not None):
cond_latent = (batch.cond_indicator * conditioning_latents +
(1 - batch.cond_indicator) * cond_latent)
else:
logger.warning(
"Step %s: Missing conditioning data - "
"cond_indicator: %s, conditioning_latents: %s", step_idx,
hasattr(batch, "cond_indicator"), conditioning_latents
is not None)
cond_latent = cond_latent.to(target_dtype)
cond_timestep = timestep
if hasattr(batch,
"cond_indicator") and batch.cond_indicator is not None:
sigma_conditioning = 0.0001
t_conditioning = sigma_conditioning / (sigma_conditioning + 1)
cond_timestep = (batch.cond_indicator * t_conditioning +
(1 - batch.cond_indicator) * timestep)
cond_timestep = cond_timestep.to(target_dtype)
with set_forward_context(
current_timestep=step_idx,
attn_metadata=None,
forward_batch=batch,
):
condition_mask = (batch.cond_mask.to(target_dtype) if hasattr(
batch, "cond_mask") else None)
padding_mask = torch.zeros(1,
1,
batch.height,
batch.width,
device=cond_latent.device,
dtype=target_dtype)
if condition_mask is None:
batch_size, _, num_frames, height, width = cond_latent.shape
condition_mask = torch.zeros(batch_size,
1,
num_frames,
height,
width,
device=cond_latent.device,
dtype=target_dtype)
noise_pred = self.stage.transformer(
hidden_states=cond_latent,
timestep=cond_timestep.to(target_dtype),
encoder_hidden_states=batch.prompt_embeds[0].to(
target_dtype),
fps=24,
condition_mask=condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
cond_pred = (c_skip * state.latents +
c_out * noise_pred.float()).to(target_dtype)
if (hasattr(batch, "cond_indicator")
and batch.cond_indicator is not None
and conditioning_latents is not None):
cond_pred = (batch.cond_indicator * conditioning_latents +
(1 - batch.cond_indicator) * cond_pred)
if (state.do_cfg and batch.negative_prompt_embeds is not None):
uncond_latent = state.latents * c_in
if (hasattr(batch, "uncond_indicator")
and batch.uncond_indicator is not None
and unconditioning_latents is not None):
uncond_latent = (
batch.uncond_indicator * unconditioning_latents +
(1 - batch.uncond_indicator) * uncond_latent)
with set_forward_context(
current_timestep=step_idx,
attn_metadata=None,
forward_batch=batch,
):
uncond_condition_mask = (
batch.uncond_mask.to(target_dtype) if
(hasattr(batch, "uncond_mask")
and batch.uncond_mask is not None) else condition_mask)
uncond_timestep = timestep
if (hasattr(batch, "uncond_indicator")
and batch.uncond_indicator is not None):
sigma_conditioning = 0.0001
t_conditioning = sigma_conditioning / (
sigma_conditioning + 1)
uncond_timestep = (
batch.uncond_indicator * t_conditioning +
(1 - batch.uncond_indicator) * timestep)
uncond_timestep = uncond_timestep.to(target_dtype)
noise_pred_uncond = self.stage.transformer(
hidden_states=uncond_latent.to(target_dtype),
timestep=uncond_timestep.to(target_dtype),
encoder_hidden_states=batch.negative_prompt_embeds[0].
to(target_dtype),
fps=24,
condition_mask=uncond_condition_mask,
padding_mask=padding_mask,
return_dict=False,
)[0]
uncond_pred = (
c_skip * state.latents +
c_out * noise_pred_uncond.float()).to(target_dtype)
if (hasattr(batch, "uncond_indicator")
and batch.uncond_indicator is not None
and unconditioning_latents is not None):
uncond_pred = (
batch.uncond_indicator * unconditioning_latents +
(1 - batch.uncond_indicator) * uncond_pred)
guidance_diff = cond_pred - uncond_pred
final_pred = cond_pred + state.guidance_scale * guidance_diff
else:
final_pred = cond_pred
if current_sigma > 1e-8:
noise_for_scheduler = (state.latents - final_pred) / current_sigma
else:
logger.warning(
"Step %s: current_sigma too small (%s), using final_pred directly",
step_idx, current_sigma)
noise_for_scheduler = final_pred
if torch.isnan(noise_for_scheduler).sum() > 0:
logger.error(
"Step %s: NaN detected in noise_for_scheduler, sum: %s",
step_idx,
noise_for_scheduler.float().sum().item())
logger.error(
"Step %s: latents sum: %s, final_pred sum: %s, current_sigma: %s",
step_idx,
state.latents.float().sum().item(),
final_pred.float().sum().item(), current_sigma)
return noise_for_scheduler
def cfg_combine(self, state: StrategyState,
noise_pred: torch.Tensor) -> torch.Tensor:
return noise_pred
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
latents = self.stage.scheduler.step(
noise_pred,
t,
state.latents,
**state.extra["extra_step_kwargs"],
return_dict=False,
)[0]
progress_bar = state.extra["progress_bar"]
if progress_bar is not None:
progress_bar.update()
return latents
def postprocess(self, state: StrategyState) -> ForwardBatch:
progress_bar = state.extra.get("progress_bar")
if progress_bar is not None:
progress_bar.close()
batch = state.extra["batch"]
batch.latents = state.latents
return batch
@@ -0,0 +1,269 @@
# SPDX-License-Identifier: Apache-2.0
"""
DMD denoising strategy (FlowMatch-based).
"""
from __future__ import annotations
from typing import Any
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_strategies import (
DenoisingStrategy,
ModelInputs,
StrategyState,
)
from fastvideo.utils import dict_to_3d_list
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
logger = init_logger(__name__)
class DmdStrategy(DenoisingStrategy):
def __init__(self, stage: Any) -> None:
self.stage = stage
def prepare(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> StrategyState:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
timesteps = batch.timesteps
if timesteps is None:
raise ValueError("Timesteps must be provided")
num_inference_steps = batch.num_inference_steps
num_warmup_steps = len(
timesteps) - num_inference_steps * self.stage.scheduler.order
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.transformer.forward,
{
"encoder_hidden_states_image": image_embeds,
"mask_strategy": dict_to_3d_list(
None, t_max=50, l_max=60, h_max=24)
},
)
pos_cond_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
if (st_attn_available
and self.stage.attn_backend == SlidingTileAttentionBackend):
self.stage.prepare_sta_param(batch, fastvideo_args)
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
progress_bar = self.stage.progress_bar(total=len(timesteps))
extra: dict[str, Any] = {
"batch": batch,
"fastvideo_args": fastvideo_args,
"target_dtype": target_dtype,
"autocast_enabled": autocast_enabled,
"num_warmup_steps": num_warmup_steps,
"image_kwargs": image_kwargs,
"pos_cond_kwargs": pos_cond_kwargs,
"progress_bar": progress_bar,
"video_raw_latent_shape": latents.shape,
}
return StrategyState(
latents=latents,
timesteps=timesteps,
num_inference_steps=num_inference_steps,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=batch.negative_prompt_embeds,
prompt_attention_mask=batch.prompt_attention_mask,
negative_attention_mask=batch.negative_attention_mask,
image_embeds=image_embeds,
guidance_scale=batch.guidance_scale,
guidance_scale_2=batch.guidance_scale_2,
guidance_rescale=batch.guidance_rescale,
do_cfg=batch.do_classifier_free_guidance,
extra=extra,
)
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
state.extra["step_idx"] = step_idx
return ModelInputs(
latent_model_input=state.latents,
timestep=t,
prompt_embeds=state.prompt_embeds,
prompt_attention_mask=state.prompt_attention_mask,
)
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
batch = state.extra["batch"]
fastvideo_args = state.extra["fastvideo_args"]
target_dtype = state.extra["target_dtype"]
autocast_enabled = state.extra["autocast_enabled"]
step_idx = state.extra["step_idx"]
if getattr(self.stage, "interrupt", False):
state.extra["pred_latents"] = state.latents
return state.latents
latents = state.latents
noise_latents = latents.clone()
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
latent_model_input = torch.cat(
[latent_model_input,
batch.image_latent.permute(0, 2, 1, 3, 4)],
dim=2).to(target_dtype)
t_expand = model_inputs.timestep.repeat(latent_model_input.shape[0])
guidance_expand = None
if fastvideo_args.pipeline_config.embedded_cfg_scale is not None:
guidance_expand = (torch.tensor(
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_local_torch_device(),
).to(target_dtype) * 1000.0)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
if (vsa_available
and self.stage.attn_backend == VideoSparseAttentionBackend):
self.attn_metadata_builder_cls = (
self.stage.attn_backend.get_builder_cls())
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = (
self.attn_metadata_builder_cls())
attn_metadata = self.attn_metadata_builder.build( # type: ignore
current_timestep=step_idx, # type: ignore
raw_latent_shape=batch.
raw_latent_shape[2:5], # type: ignore
patch_size=fastvideo_args.pipeline_config.dit_config.
patch_size, # type: ignore
STA_param=batch.STA_param, # type: ignore
VSA_sparsity=fastvideo_args.
VSA_sparsity, # type: ignore
device=get_local_torch_device(), # type: ignore
)
assert attn_metadata is not None, (
"attn_metadata cannot be None")
else:
attn_metadata = None
else:
attn_metadata = None
with set_forward_context(
current_timestep=step_idx,
attn_metadata=attn_metadata,
forward_batch=batch,
):
pred_noise = self.stage.transformer(
latent_model_input.permute(0, 2, 1, 3, 4),
state.prompt_embeds,
t_expand,
guidance=guidance_expand,
**state.extra["image_kwargs"],
**state.extra["pos_cond_kwargs"],
).permute(0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.stage.scheduler).unflatten(0, pred_noise.shape[:2])
if step_idx < len(state.timesteps) - 1:
next_timestep = state.timesteps[step_idx + 1] * torch.ones(
[1], dtype=torch.long, device=pred_video.device)
generator = batch.generator
if isinstance(generator, list):
generator = generator[0] if generator else None
noise = torch.randn(state.extra["video_raw_latent_shape"],
dtype=pred_video.dtype,
generator=generator).to(self.stage.device)
latents = self.stage.scheduler.add_noise(pred_video.flatten(0, 1),
noise.flatten(0, 1),
next_timestep).unflatten(
0,
pred_video.shape[:2])
else:
latents = pred_video
state.extra["pred_latents"] = latents
return latents
def cfg_combine(self, state: StrategyState,
noise_pred: torch.Tensor) -> torch.Tensor:
return noise_pred
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
progress_bar = state.extra["progress_bar"]
step_idx = state.extra["step_idx"]
num_warmup_steps = state.extra["num_warmup_steps"]
if step_idx == len(state.timesteps) - 1 or (
(step_idx + 1) > num_warmup_steps and
(step_idx + 1) % self.stage.scheduler.order == 0
and progress_bar is not None):
progress_bar.update()
return state.extra.get("pred_latents", state.latents)
def postprocess(self, state: StrategyState) -> ForwardBatch:
progress_bar = state.extra.get("progress_bar")
if progress_bar is not None:
progress_bar.close()
batch = state.extra["batch"]
latents = state.extra["pred_latents"].permute(0, 2, 1, 3, 4)
batch.latents = latents
return batch
@@ -0,0 +1,177 @@
# SPDX-License-Identifier: Apache-2.0
"""
Scaffolding for a unified denoising engine with hook support.
"""
from __future__ import annotations
from functools import lru_cache
from typing import Protocol, TYPE_CHECKING
from collections.abc import Sequence
import torch
from fastvideo.models.schedulers.adapter import SchedulerAdapter
from fastvideo.pipelines.stages.denoising_strategies import (
BlockDenoisingStrategy,
DenoisingStrategy,
StrategyState,
)
if TYPE_CHECKING:
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
class EngineHook(Protocol):
def on_init(self, engine: DenoisingEngine, batch: ForwardBatch,
args: FastVideoArgs) -> None:
...
def pre_run(self, state: StrategyState) -> None:
...
def pre_step(self, state: StrategyState, step_idx: int,
t: torch.Tensor) -> None:
...
def post_step(self, state: StrategyState, step_idx: int,
t: torch.Tensor) -> None:
...
def post_run(self, state: StrategyState, batch: ForwardBatch) -> None:
...
class BaseEngineHook:
def on_init(self, engine: DenoisingEngine, batch: ForwardBatch,
args: FastVideoArgs) -> None:
return None
def pre_run(self, state: StrategyState) -> None:
return None
def pre_step(self, state: StrategyState, step_idx: int,
t: torch.Tensor) -> None:
return None
def post_step(self, state: StrategyState, step_idx: int,
t: torch.Tensor) -> None:
return None
def post_run(self, state: StrategyState, batch: ForwardBatch) -> None:
return None
class GuidanceCache:
def __init__(self, maxsize: int = 8) -> None:
self._build = lru_cache(maxsize=maxsize)(self._build_impl)
def _build_impl(self, batch_size: int, dtype: torch.dtype,
device: torch.device, guidance_val: float) -> torch.Tensor:
return (torch.full(
(batch_size, ),
guidance_val,
dtype=torch.float32,
device=device,
).to(dtype) * 1000.0)
def get(self, batch_size: int, dtype: torch.dtype, device: torch.device,
guidance_val: float | None) -> torch.Tensor | None:
if guidance_val is None:
return None
return self._build(batch_size, dtype, device, guidance_val)
class DenoisingEngine:
def __init__(
self,
strategy: DenoisingStrategy,
*,
scheduler_adapter: SchedulerAdapter | None = None,
hooks: Sequence[EngineHook] | None = None,
) -> None:
self.strategy = strategy
self.scheduler_adapter = scheduler_adapter
self.hooks = list(hooks) if hooks is not None else []
def run(self, batch: ForwardBatch, args: FastVideoArgs) -> ForwardBatch:
for hook in self.hooks:
hook.on_init(self, batch, args)
state = self.strategy.prepare(batch, args)
if self.scheduler_adapter is not None:
state.extra.setdefault("scheduler_adapter", self.scheduler_adapter)
for hook in self.hooks:
hook.pre_run(state)
if isinstance(self.strategy, BlockDenoisingStrategy):
self.run_blocks(state)
else:
timesteps = state.timesteps
for i, t in enumerate(timesteps):
for hook in self.hooks:
hook.pre_step(state, i, t)
model_inputs = self.strategy.make_model_inputs(state, t, i)
noise_pred = self.strategy.forward(state, model_inputs)
noise_pred = self.strategy.cfg_combine(state, noise_pred)
state.latents = self.strategy.scheduler_step(
state,
noise_pred,
t,
)
for hook in self.hooks:
hook.post_step(state, i, t)
for hook in self.hooks:
hook.post_run(state, batch)
return self.strategy.postprocess(state)
def run_blocks(
self,
state: StrategyState,
*,
block_plan=None,
start_block: int = 0,
num_blocks: int | None = None,
) -> None:
if not isinstance(self.strategy, BlockDenoisingStrategy):
raise TypeError("run_blocks requires a BlockDenoisingStrategy")
strategy = self.strategy
if block_plan is None:
block_plan = strategy.block_plan(state)
items = block_plan.items
end_block = len(items)
if num_blocks is not None:
end_block = min(end_block, start_block + num_blocks)
for block_idx in range(start_block, end_block):
block_item = items[block_idx]
t_hook = self._block_hook_t(state, block_idx)
for hook in self.hooks:
hook.pre_step(state, block_idx, t_hook)
block_ctx = strategy.init_block_context(
state,
block_item,
block_idx,
)
strategy.process_block(state, block_ctx, block_item)
strategy.update_context(state, block_ctx, block_item)
for hook in self.hooks:
hook.post_step(state, block_idx, t_hook)
def _block_hook_t(self, state: StrategyState,
block_idx: int) -> torch.Tensor:
timesteps = state.timesteps
if timesteps.numel() > 0:
return timesteps[0]
return torch.tensor(block_idx, device=state.latents.device)
@@ -0,0 +1,77 @@
# SPDX-License-Identifier: Apache-2.0
"""
Engine hook implementations for denoising runs.
"""
from __future__ import annotations
import time
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
from fastvideo.pipelines.stages.denoising_strategies import StrategyState
logger = init_logger(__name__)
class PerfLoggingHook:
"""
Record per-step and total denoising times into batch logging info.
"""
def __init__(self, stage_name: str = "DenoisingEngine") -> None:
self.stage_name = stage_name
self._batch: ForwardBatch | None = None
self._run_start = 0.0
self._step_starts: dict[int, float] = {}
self._step_times_ms: list[float] = []
def on_init(self, engine: DenoisingEngine, batch: ForwardBatch, args):
self._batch = batch
def pre_run(self, state: StrategyState) -> None:
self._run_start = time.perf_counter()
self._step_starts.clear()
self._step_times_ms.clear()
def pre_step(self, state: StrategyState, step_idx: int, t) -> None:
self._step_starts[step_idx] = time.perf_counter()
def post_step(self, state: StrategyState, step_idx: int, t) -> None:
start = self._step_starts.pop(step_idx, None)
if start is None:
return
self._step_times_ms.append((time.perf_counter() - start) * 1000.0)
def post_run(self, state: StrategyState, batch: ForwardBatch) -> None:
total_ms = (time.perf_counter() - self._run_start) * 1000.0
target_batch = self._batch or batch
if target_batch is None:
return
target_batch.logging_info.add_stage_metric(self.stage_name,
"denoise_step_times_ms",
list(self._step_times_ms))
target_batch.logging_info.add_stage_metric(self.stage_name,
"denoise_total_ms", total_ms)
if not self._step_times_ms:
return
if not self._is_primary_rank():
return
mean_ms = sum(self._step_times_ms) / len(self._step_times_ms)
logger.info(
"[%s] denoise steps=%d total_ms=%.2f mean_step_ms=%.2f min=%.2f max=%.2f",
self.stage_name,
len(self._step_times_ms),
total_ms,
mean_ms,
min(self._step_times_ms),
max(self._step_times_ms),
)
def _is_primary_rank(self) -> bool:
try:
from fastvideo.distributed import get_world_group
return get_world_group().local_rank == 0
except Exception:
return True
@@ -0,0 +1,551 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat denoising strategies (base, I2V, VC).
"""
from __future__ import annotations
import time
from typing import Any
import torch
from tqdm import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_strategies import (
DenoisingStrategy,
ModelInputs,
StrategyState,
)
logger = init_logger(__name__)
class _BaseLongCatStrategy(DenoisingStrategy):
def __init__(self, stage: Any) -> None:
self.stage = stage
def _load_transformer(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
attach_pipeline: bool) -> None:
if fastvideo_args.model_loaded["transformer"]:
return
loader = TransformerLoader()
self.stage.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
if attach_pipeline:
pipeline = self.stage.pipeline() if self.stage.pipeline else None
if pipeline:
pipeline.add_module("transformer", self.stage.transformer)
fastvideo_args.model_loaded["transformer"] = True
def _build_prompt_inputs(self, batch: ForwardBatch):
prompt_embeds = batch.prompt_embeds[0]
prompt_attention_mask = (batch.prompt_attention_mask[0]
if batch.prompt_attention_mask else None)
do_cfg = batch.do_classifier_free_guidance
if do_cfg:
negative_prompt_embeds = batch.negative_prompt_embeds[0]
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
if batch.negative_attention_mask
else None)
prompt_embeds_combined = torch.cat(
[negative_prompt_embeds, prompt_embeds], dim=0)
if prompt_attention_mask is not None:
prompt_attention_mask_combined = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask],
dim=0)
else:
prompt_attention_mask_combined = None
else:
prompt_embeds_combined = prompt_embeds
prompt_attention_mask_combined = prompt_attention_mask
return prompt_embeds, prompt_attention_mask, prompt_embeds_combined, \
prompt_attention_mask_combined
def optimized_scale(self, positive_flat: torch.Tensor,
negative_flat: torch.Tensor) -> torch.Tensor:
"""
Calculate optimized scale from CFG-zero paper.
st_star = (v_cond^T * v_uncond) / ||v_uncond||^2
"""
dot_product = torch.sum(positive_flat * negative_flat,
dim=1,
keepdim=True)
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
return dot_product / squared_norm
def cfg_combine(self, state: StrategyState,
noise_pred: torch.Tensor) -> torch.Tensor:
return noise_pred
class LongCatStrategy(_BaseLongCatStrategy):
def prepare(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> StrategyState:
self._load_transformer(batch, fastvideo_args, attach_pipeline=True)
if hasattr(self.stage.transformer, "module"):
transformer_dtype = next(
self.stage.transformer.module.parameters()).dtype
else:
transformer_dtype = next(self.stage.transformer.parameters()).dtype
target_dtype = transformer_dtype
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
prompt_embeds, prompt_attention_mask, prompt_embeds_combined, \
prompt_attention_mask_combined = self._build_prompt_inputs(batch)
timesteps = batch.timesteps
num_inference_steps = len(timesteps)
progress_bar = tqdm(total=num_inference_steps, desc="LongCat Denoising")
extra: dict[str, Any] = {
"batch": batch,
"target_dtype": target_dtype,
"autocast_enabled": autocast_enabled,
"prompt_embeds_combined": prompt_embeds_combined,
"prompt_attention_mask_combined": prompt_attention_mask_combined,
"progress_bar": progress_bar,
}
return StrategyState(
latents=batch.latents,
timesteps=timesteps,
num_inference_steps=num_inference_steps,
prompt_embeds=[prompt_embeds],
negative_prompt_embeds=batch.negative_prompt_embeds,
prompt_attention_mask=[prompt_attention_mask]
if prompt_attention_mask is not None else None,
negative_attention_mask=batch.negative_attention_mask,
image_embeds=batch.image_embeds,
guidance_scale=batch.guidance_scale,
guidance_scale_2=batch.guidance_scale_2,
guidance_rescale=batch.guidance_rescale,
do_cfg=batch.do_classifier_free_guidance,
extra=extra,
)
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
target_dtype = state.extra["target_dtype"]
latents = state.latents
if state.do_cfg:
latent_model_input = torch.cat([latents] * 2)
else:
latent_model_input = latents
latent_model_input = latent_model_input.to(target_dtype)
timestep = t.expand(latent_model_input.shape[0]).to(target_dtype)
state.extra["step_idx"] = step_idx
return ModelInputs(
latent_model_input=latent_model_input,
timestep=timestep,
prompt_embeds=state.prompt_embeds,
prompt_attention_mask=state.prompt_attention_mask,
)
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
batch = state.extra["batch"]
target_dtype = state.extra["target_dtype"]
autocast_enabled = state.extra["autocast_enabled"]
prompt_embeds_combined = state.extra["prompt_embeds_combined"]
prompt_attention_mask_combined = state.extra[
"prompt_attention_mask_combined"]
step_idx = state.extra["step_idx"]
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=step_idx,
attn_metadata=None,
forward_batch=batch,
), torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
noise_pred = self.stage.transformer(
hidden_states=model_inputs.latent_model_input,
encoder_hidden_states=prompt_embeds_combined,
timestep=model_inputs.timestep,
encoder_attention_mask=prompt_attention_mask_combined,
)
if state.do_cfg:
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
B = noise_pred_cond.shape[0]
positive = noise_pred_cond.reshape(B, -1)
negative = noise_pred_uncond.reshape(B, -1)
st_star = self.optimized_scale(positive, negative)
st_star = st_star.view(B, 1, 1, 1, 1)
noise_pred = (noise_pred_uncond * st_star + state.guidance_scale *
(noise_pred_cond - noise_pred_uncond * st_star))
noise_pred = -noise_pred
return noise_pred
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
latents = self.stage.scheduler.step(noise_pred,
t,
state.latents,
return_dict=False)[0]
progress_bar = state.extra["progress_bar"]
if progress_bar is not None:
progress_bar.update()
return latents
def postprocess(self, state: StrategyState) -> ForwardBatch:
progress_bar = state.extra.get("progress_bar")
if progress_bar is not None:
progress_bar.close()
batch = state.extra["batch"]
batch.latents = state.latents
return batch
class LongCatI2VStrategy(_BaseLongCatStrategy):
def prepare(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> StrategyState:
self._load_transformer(batch, fastvideo_args, attach_pipeline=False)
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
prompt_embeds, prompt_attention_mask, prompt_embeds_combined, \
prompt_attention_mask_combined = self._build_prompt_inputs(batch)
timesteps = batch.timesteps
num_inference_steps = len(timesteps)
progress_bar = tqdm(total=num_inference_steps, desc="I2V Denoising")
num_cond_latents = getattr(batch, "num_cond_latents", 0)
if num_cond_latents > 0:
logger.info("I2V Denoising: num_cond_latents=%s, latent_shape=%s",
num_cond_latents, batch.latents.shape)
extra: dict[str, Any] = {
"batch": batch,
"target_dtype": target_dtype,
"autocast_enabled": autocast_enabled,
"prompt_embeds_combined": prompt_embeds_combined,
"prompt_attention_mask_combined": prompt_attention_mask_combined,
"num_cond_latents": num_cond_latents,
"progress_bar": progress_bar,
}
return StrategyState(
latents=batch.latents,
timesteps=timesteps,
num_inference_steps=num_inference_steps,
prompt_embeds=[prompt_embeds],
negative_prompt_embeds=batch.negative_prompt_embeds,
prompt_attention_mask=[prompt_attention_mask]
if prompt_attention_mask is not None else None,
negative_attention_mask=batch.negative_attention_mask,
image_embeds=batch.image_embeds,
guidance_scale=batch.guidance_scale,
guidance_scale_2=batch.guidance_scale_2,
guidance_rescale=batch.guidance_rescale,
do_cfg=batch.do_classifier_free_guidance,
extra=extra,
)
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
target_dtype = state.extra["target_dtype"]
num_cond_latents = state.extra["num_cond_latents"]
latents = state.latents
if state.do_cfg:
latent_model_input = torch.cat([latents] * 2)
else:
latent_model_input = latents
latent_model_input = latent_model_input.to(target_dtype)
timestep = t.expand(latent_model_input.shape[0]).to(target_dtype)
timestep = timestep.unsqueeze(-1).repeat(1, latent_model_input.shape[2])
if num_cond_latents > 0:
timestep[:, :num_cond_latents] = 0
state.extra["step_idx"] = step_idx
return ModelInputs(
latent_model_input=latent_model_input,
timestep=timestep,
prompt_embeds=state.prompt_embeds,
prompt_attention_mask=state.prompt_attention_mask,
extra_kwargs={"num_cond_latents": num_cond_latents},
)
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
batch = state.extra["batch"]
target_dtype = state.extra["target_dtype"]
autocast_enabled = state.extra["autocast_enabled"]
prompt_embeds_combined = state.extra["prompt_embeds_combined"]
prompt_attention_mask_combined = state.extra[
"prompt_attention_mask_combined"]
step_idx = state.extra["step_idx"]
num_cond_latents = state.extra["num_cond_latents"]
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=step_idx,
attn_metadata=None,
forward_batch=batch,
), torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
noise_pred = self.stage.transformer(
hidden_states=model_inputs.latent_model_input,
encoder_hidden_states=prompt_embeds_combined,
timestep=model_inputs.timestep,
encoder_attention_mask=prompt_attention_mask_combined,
num_cond_latents=num_cond_latents,
)
if state.do_cfg:
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
B = noise_pred_cond.shape[0]
positive = noise_pred_cond.reshape(B, -1)
negative = noise_pred_uncond.reshape(B, -1)
st_star = self.optimized_scale(positive, negative)
st_star = st_star.view(B, 1, 1, 1, 1)
noise_pred = (noise_pred_uncond * st_star + state.guidance_scale *
(noise_pred_cond - noise_pred_uncond * st_star))
noise_pred = -noise_pred
return noise_pred
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
num_cond_latents = state.extra["num_cond_latents"]
latents = state.latents
if num_cond_latents > 0:
latents[:, :, num_cond_latents:] = self.stage.scheduler.step(
noise_pred[:, :, num_cond_latents:],
t,
latents[:, :, num_cond_latents:],
return_dict=False)[0]
else:
latents = self.stage.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
progress_bar = state.extra["progress_bar"]
if progress_bar is not None:
progress_bar.update()
return latents
def postprocess(self, state: StrategyState) -> ForwardBatch:
progress_bar = state.extra.get("progress_bar")
if progress_bar is not None:
progress_bar.close()
batch = state.extra["batch"]
batch.latents = state.latents
return batch
class LongCatVCStrategy(_BaseLongCatStrategy):
def prepare(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> StrategyState:
self._load_transformer(batch, fastvideo_args, attach_pipeline=False)
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
prompt_embeds, prompt_attention_mask, prompt_embeds_combined, \
prompt_attention_mask_combined = self._build_prompt_inputs(batch)
timesteps = batch.timesteps
num_inference_steps = len(timesteps)
progress_bar = tqdm(total=num_inference_steps, desc="VC Denoising")
num_cond_latents = getattr(batch, "num_cond_latents", 0)
use_kv_cache = getattr(batch, "use_kv_cache", False)
kv_cache_dict = getattr(batch, "kv_cache_dict", {})
logger.info(
"VC Denoising: num_cond_latents=%d, use_kv_cache=%s, latent_shape=%s",
num_cond_latents, use_kv_cache, batch.latents.shape)
extra: dict[str, Any] = {
"batch": batch,
"target_dtype": target_dtype,
"autocast_enabled": autocast_enabled,
"prompt_embeds_combined": prompt_embeds_combined,
"prompt_attention_mask_combined": prompt_attention_mask_combined,
"num_cond_latents": num_cond_latents,
"use_kv_cache": use_kv_cache,
"kv_cache_dict": kv_cache_dict,
"progress_bar": progress_bar,
"step_times": [],
}
return StrategyState(
latents=batch.latents,
timesteps=timesteps,
num_inference_steps=num_inference_steps,
prompt_embeds=[prompt_embeds],
negative_prompt_embeds=batch.negative_prompt_embeds,
prompt_attention_mask=[prompt_attention_mask]
if prompt_attention_mask is not None else None,
negative_attention_mask=batch.negative_attention_mask,
image_embeds=batch.image_embeds,
guidance_scale=batch.guidance_scale,
guidance_scale_2=batch.guidance_scale_2,
guidance_rescale=batch.guidance_rescale,
do_cfg=batch.do_classifier_free_guidance,
extra=extra,
)
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
target_dtype = state.extra["target_dtype"]
num_cond_latents = state.extra["num_cond_latents"]
use_kv_cache = state.extra["use_kv_cache"]
latents = state.latents
if state.do_cfg:
latent_model_input = torch.cat([latents] * 2)
else:
latent_model_input = latents
latent_model_input = latent_model_input.to(target_dtype)
timestep = t.expand(latent_model_input.shape[0]).to(target_dtype)
timestep = timestep.unsqueeze(-1).repeat(1, latent_model_input.shape[2])
if not use_kv_cache and num_cond_latents > 0:
timestep[:, :num_cond_latents] = 0
state.extra["step_idx"] = step_idx
state.extra["step_start"] = time.time()
extra_kwargs = {"num_cond_latents": num_cond_latents}
if use_kv_cache:
extra_kwargs["kv_cache_dict"] = state.extra["kv_cache_dict"]
return ModelInputs(
latent_model_input=latent_model_input,
timestep=timestep,
prompt_embeds=state.prompt_embeds,
prompt_attention_mask=state.prompt_attention_mask,
extra_kwargs=extra_kwargs,
)
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
batch = state.extra["batch"]
target_dtype = state.extra["target_dtype"]
autocast_enabled = state.extra["autocast_enabled"]
prompt_embeds_combined = state.extra["prompt_embeds_combined"]
prompt_attention_mask_combined = state.extra[
"prompt_attention_mask_combined"]
step_idx = state.extra["step_idx"]
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=step_idx,
attn_metadata=None,
forward_batch=batch,
), torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
noise_pred = self.stage.transformer(
hidden_states=model_inputs.latent_model_input,
encoder_hidden_states=prompt_embeds_combined,
timestep=model_inputs.timestep,
encoder_attention_mask=prompt_attention_mask_combined,
**model_inputs.extra_kwargs,
)
if state.do_cfg:
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
B = noise_pred_cond.shape[0]
positive = noise_pred_cond.reshape(B, -1)
negative = noise_pred_uncond.reshape(B, -1)
st_star = self.optimized_scale(positive, negative)
st_star = st_star.view(B, 1, 1, 1, 1)
noise_pred = (noise_pred_uncond * st_star + state.guidance_scale *
(noise_pred_cond - noise_pred_uncond * st_star))
noise_pred = -noise_pred
return noise_pred
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
num_cond_latents = state.extra["num_cond_latents"]
use_kv_cache = state.extra["use_kv_cache"]
latents = state.latents
if use_kv_cache:
latents = self.stage.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
else:
if num_cond_latents > 0:
latents[:, :, num_cond_latents:] = self.stage.scheduler.step(
noise_pred[:, :, num_cond_latents:],
t,
latents[:, :, num_cond_latents:],
return_dict=False,
)[0]
else:
latents = self.stage.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
step_time = time.time() - state.extra["step_start"]
state.extra["step_times"].append(step_time)
if state.extra["step_idx"] < 3:
logger.info("Step %d: %.2fs", state.extra["step_idx"], step_time)
progress_bar = state.extra["progress_bar"]
if progress_bar is not None:
progress_bar.update()
return latents
def postprocess(self, state: StrategyState) -> ForwardBatch:
progress_bar = state.extra.get("progress_bar")
if progress_bar is not None:
progress_bar.close()
batch = state.extra["batch"]
use_kv_cache = state.extra["use_kv_cache"]
step_times = state.extra["step_times"]
latents = state.latents
if use_kv_cache and hasattr(
batch, "cond_latents") and batch.cond_latents is not None:
latents = torch.cat([batch.cond_latents, latents], dim=2)
logger.info(
"Concatenated conditioning latents back, final shape: %s",
latents.shape)
if step_times:
avg_time = sum(step_times) / len(step_times)
logger.info("Average step time: %.2fs (total: %.1fs)", avg_time,
sum(step_times))
batch.latents = latents
return batch
@@ -0,0 +1,324 @@
# SPDX-License-Identifier: Apache-2.0
"""
MatrixGame causal block denoising strategy.
"""
from __future__ import annotations
from typing import Any
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_strategies import (
BlockContext,
BlockDenoisingStrategy,
BlockPlan,
BlockPlanItem,
ModelInputs,
StrategyState,
)
from fastvideo.pipelines.stages.matrixgame_denoising import BlockProcessingContext
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
class MatrixGameBlockStrategy(BlockDenoisingStrategy):
def __init__(self, stage: Any) -> None:
self.stage = stage
def prepare(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> StrategyState:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latents = batch.latents
if latents is None:
raise ValueError("latents must be provided")
latent_seq_length = latents.shape[-1] * latents.shape[-2]
patch_size = self.stage.transformer.patch_size
patch_ratio = patch_size[-1] * patch_size[-2]
self.stage.frame_seq_length = latent_seq_length // patch_ratio
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if getattr(fastvideo_args.pipeline_config, "warp_denoising_step",
False):
scheduler_timesteps = torch.cat((
self.stage.scheduler.timesteps.cpu(),
torch.tensor([0], dtype=torch.float32),
))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
"boundary_ratio", None)
if boundary_ratio is not None:
boundary_timestep = (boundary_ratio *
self.stage.scheduler.num_train_timesteps)
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = {"encoder_hidden_states_image": image_embeds}
pos_cond_kwargs: dict[str, Any] = {}
if (st_attn_available
and self.stage.attn_backend == SlidingTileAttentionBackend):
self.stage.prepare_sta_param(batch, fastvideo_args)
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
kv_cache1 = self.stage._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
kv_cache2 = self.stage._initialize_kv_cache(
batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device,
)
kv_cache_mouse = None
kv_cache_keyboard = None
if self.stage.use_action_module:
kv_cache_mouse, kv_cache_keyboard = (
self.stage._initialize_action_kv_cache(
batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device,
))
crossattn_cache = self.stage._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=257,
dtype=target_dtype,
device=latents.device,
)
num_frames = latents.shape[2]
if num_frames % self.stage.num_frame_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frame_per_block for "
"causal denoising")
num_blocks = num_frames // self.stage.num_frame_per_block
block_sizes = [self.stage.num_frame_per_block] * num_blocks
start_index = 0
if boundary_timestep is not None:
block_sizes[0] = 1
ctx = BlockProcessingContext(
batch=batch,
block_idx=0,
start_index=0,
kv_cache1=kv_cache1,
kv_cache2=kv_cache2,
kv_cache_mouse=kv_cache_mouse,
kv_cache_keyboard=kv_cache_keyboard,
crossattn_cache=crossattn_cache,
timesteps=timesteps,
block_sizes=block_sizes,
noise_pool=None,
fastvideo_args=fastvideo_args,
target_dtype=target_dtype,
autocast_enabled=autocast_enabled,
boundary_timestep=boundary_timestep,
high_noise_timesteps=high_noise_timesteps,
context_noise=getattr(fastvideo_args.pipeline_config,
"context_noise", 0),
image_kwargs=image_kwargs,
pos_cond_kwargs=pos_cond_kwargs,
)
progress_bar = self.stage.progress_bar(total=len(block_sizes) *
len(timesteps))
extra: dict[str, Any] = {
"batch": batch,
"fastvideo_args": fastvideo_args,
"ctx": ctx,
"block_sizes": block_sizes,
"start_index": start_index,
"progress_bar": progress_bar,
"boundary_timestep": boundary_timestep,
}
return StrategyState(
latents=latents,
timesteps=timesteps,
num_inference_steps=len(timesteps),
prompt_embeds=batch.prompt_embeds,
negative_prompt_embeds=batch.negative_prompt_embeds,
prompt_attention_mask=batch.prompt_attention_mask,
negative_attention_mask=batch.negative_attention_mask,
image_embeds=image_embeds,
guidance_scale=batch.guidance_scale,
guidance_scale_2=batch.guidance_scale_2,
guidance_rescale=batch.guidance_rescale,
do_cfg=batch.do_classifier_free_guidance,
extra=extra,
)
def block_plan(self, state: StrategyState) -> BlockPlan:
block_sizes = state.extra["block_sizes"]
start_index = state.extra["start_index"]
items: list[BlockPlanItem] = []
for block_size in block_sizes:
items.append(
BlockPlanItem(
start_index=start_index,
num_frames=block_size,
use_kv_cache=True,
model_selector="default",
))
start_index += block_size
return BlockPlan(items=items)
def init_block_context(self, state: StrategyState,
block_item: BlockPlanItem,
block_idx: int) -> BlockContext:
ctx = state.extra["ctx"]
ctx.block_idx = block_idx
ctx.start_index = block_item.start_index
action_kwargs = self.stage._prepare_action_kwargs(
state.extra["batch"],
block_item.start_index,
block_item.num_frames,
)
return BlockContext(
kv_cache=ctx.kv_cache1,
kv_cache_2=ctx.kv_cache2,
crossattn_cache=ctx.crossattn_cache,
action_cache=None,
extra={
"ctx": ctx,
"action_kwargs": action_kwargs,
"start_index": block_item.start_index,
"num_frames": block_item.num_frames,
},
)
def process_block(self, state: StrategyState, block_ctx: BlockContext,
block_item: BlockPlanItem) -> None:
ctx = block_ctx.extra["ctx"]
batch = state.extra["batch"]
progress_bar = state.extra["progress_bar"]
start_index = block_item.start_index
current_num_frames = block_item.num_frames
action_kwargs = block_ctx.extra["action_kwargs"]
current_latents = state.latents[:, :, start_index:start_index +
current_num_frames, :, :]
noise_generator = None
if ctx.noise_pool is not None:
latents_device = state.latents.device
def noise_generator(shape: tuple, dtype: torch.dtype,
step_idx: int) -> torch.Tensor:
if step_idx < len(ctx.noise_pool):
noise = ctx.noise_pool[step_idx]
if noise.shape != shape:
noise = noise[:, :shape[1], :, :, :]
return noise.to(device=latents_device, dtype=dtype)
generator = batch.generator
if isinstance(generator, list):
generator = generator[0] if generator else None
return torch.randn(shape, dtype=dtype,
generator=generator).to(latents_device)
current_latents = self.stage._process_single_block(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
timesteps=state.timesteps,
ctx=ctx,
action_kwargs=action_kwargs,
progress_bar=progress_bar,
noise_generator=noise_generator,
)
state.latents[:, :, start_index:start_index +
current_num_frames, :, :] = (current_latents)
def update_context(self, state: StrategyState, block_ctx: BlockContext,
block_item: BlockPlanItem) -> None:
ctx = block_ctx.extra["ctx"]
batch = state.extra["batch"]
action_kwargs = block_ctx.extra["action_kwargs"]
start_index = block_item.start_index
current_num_frames = block_item.num_frames
current_latents = state.latents[:, :, start_index:start_index +
current_num_frames, :, :]
self.stage._update_context_cache(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=action_kwargs,
context_noise=ctx.context_noise,
)
def postprocess(self, state: StrategyState) -> ForwardBatch:
progress_bar = state.extra.get("progress_bar")
if progress_bar is not None:
progress_bar.close()
batch = state.extra["batch"]
boundary_timestep = state.extra["boundary_timestep"]
latents = state.latents
if boundary_timestep is not None:
num_frames_to_remove = self.stage.num_frame_per_block - 1
if num_frames_to_remove > 0:
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
raise NotImplementedError
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
raise NotImplementedError
def cfg_combine(self, state: StrategyState,
noise_pred: torch.Tensor) -> torch.Tensor:
raise NotImplementedError
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
raise NotImplementedError
@@ -0,0 +1,552 @@
# SPDX-License-Identifier: Apache-2.0
"""
Standard denoising strategy backed by the legacy DenoisingStage utilities.
"""
from __future__ import annotations
from typing import Any
import torch
from fastvideo.configs.pipelines.base import STA_Mode
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_strategies import (
DenoisingStrategy,
ModelInputs,
StrategyState,
)
from fastvideo.utils import dict_to_3d_list, masks_like
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.vmoba import VMOBAAttentionBackend
from fastvideo.utils import is_vmoba_available
vmoba_attn_available = is_vmoba_available()
except ImportError:
vmoba_attn_available = False
VMOBAAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
logger = init_logger(__name__)
class StandardStrategy(DenoisingStrategy):
def __init__(self, stage: Any) -> None:
self.stage = stage
def prepare(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> StrategyState:
pipeline = self.stage.pipeline() if self.stage.pipeline else None
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.stage.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
if pipeline:
pipeline.add_module("transformer", self.stage.transformer)
fastvideo_args.model_loaded["transformer"] = True
extra_step_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
timesteps = batch.timesteps
if timesteps is None:
raise ValueError("Timesteps must be provided")
num_inference_steps = batch.num_inference_steps
num_warmup_steps = len(
timesteps) - num_inference_steps * self.stage.scheduler.order
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert not torch.isnan(
image_embeds[0]).any(), "image_embeds contains nan"
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.transformer.forward,
{
"encoder_hidden_states_image": image_embeds,
"mask_strategy": dict_to_3d_list(
None, t_max=50, l_max=60, h_max=24)
},
)
pos_cond_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
neg_cond_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_neg,
"encoder_attention_mask": batch.negative_attention_mask,
},
)
action_kwargs = self.stage.prepare_extra_func_kwargs(
self.stage.transformer.forward,
{
"mouse_cond": batch.mouse_cond,
"keyboard_cond": batch.keyboard_cond,
},
)
if (st_attn_available
and self.stage.attn_backend == SlidingTileAttentionBackend):
self.stage.prepare_sta_param(batch, fastvideo_args)
latents = batch.latents
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
neg_prompt_embeds = None
if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds
assert neg_prompt_embeds is not None
assert not torch.isnan(
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
boundary_ratio = (
fastvideo_args.pipeline_config.dit_config.boundary_ratio)
if batch.boundary_ratio is not None:
logger.info("Overriding boundary ratio from %s to %s",
boundary_ratio, batch.boundary_ratio)
boundary_ratio = batch.boundary_ratio
if boundary_ratio is not None:
boundary_timestep = (boundary_ratio *
self.stage.scheduler.num_train_timesteps)
else:
boundary_timestep = None
latent_model_input = latents.to(target_dtype)
if latent_model_input.ndim == 5:
assert latent_model_input.shape[0] == 1, (
"only support batch size 1")
ti2v_mask = None
ti2v_z = None
ti2v_seq_len = None
if (fastvideo_args.pipeline_config.ti2v_task
and batch.pil_image is not None):
assert batch.image_latent is None, (
"TI2V task should not have image latents")
assert self.stage.vae is not None, (
"VAE is not provided for TI2V task")
z = self.stage.vae.encode(batch.pil_image).mean.float()
if (hasattr(self.stage.vae, "shift_factor")
and self.stage.vae.shift_factor is not None):
if isinstance(self.stage.vae.shift_factor, torch.Tensor):
z -= self.stage.vae.shift_factor.to(z.device, z.dtype)
else:
z -= self.stage.vae.shift_factor
if isinstance(self.stage.vae.scaling_factor, torch.Tensor):
z = z * self.stage.vae.scaling_factor.to(z.device, z.dtype)
else:
z = z * self.stage.vae.scaling_factor
latent_model_input = latents.to(target_dtype).squeeze(0)
_, mask2 = masks_like([latent_model_input], zero=True)
latent_model_input = ((1. - mask2[0]) * z +
mask2[0] * latent_model_input)
latent_model_input = latent_model_input.to(get_local_torch_device())
latents = latent_model_input
F = batch.num_frames
temporal_scale = (fastvideo_args.pipeline_config.vae_config.
arch_config.scale_factor_temporal)
spatial_scale = (fastvideo_args.pipeline_config.vae_config.
arch_config.scale_factor_spatial)
patch_size = (fastvideo_args.pipeline_config.dit_config.arch_config.
patch_size)
seq_len = ((F - 1) // temporal_scale +
1) * (batch.height // spatial_scale) * (
batch.width // spatial_scale) // (patch_size[1] *
patch_size[2])
ti2v_mask = mask2[0]
ti2v_z = z
ti2v_seq_len = seq_len
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_latents: list[torch.Tensor] | None = None
if batch.return_trajectory_latents:
trajectory_timesteps = []
trajectory_latents = []
progress_bar = self.stage.progress_bar(total=num_inference_steps)
extra: dict[str, Any] = {
"batch": batch,
"fastvideo_args": fastvideo_args,
"extra_step_kwargs": extra_step_kwargs,
"target_dtype": target_dtype,
"autocast_enabled": autocast_enabled,
"num_warmup_steps": num_warmup_steps,
"image_kwargs": image_kwargs,
"pos_cond_kwargs": pos_cond_kwargs,
"neg_cond_kwargs": neg_cond_kwargs,
"action_kwargs": action_kwargs,
"boundary_timestep": boundary_timestep,
"progress_bar": progress_bar,
"trajectory_timesteps": trajectory_timesteps,
"trajectory_latents": trajectory_latents,
"ti2v_mask": ti2v_mask,
"ti2v_z": ti2v_z,
"ti2v_seq_len": ti2v_seq_len,
}
return StrategyState(
latents=latents,
timesteps=timesteps,
num_inference_steps=num_inference_steps,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=neg_prompt_embeds,
prompt_attention_mask=batch.prompt_attention_mask,
negative_attention_mask=batch.negative_attention_mask,
image_embeds=image_embeds,
guidance_scale=batch.guidance_scale,
guidance_scale_2=batch.guidance_scale_2,
guidance_rescale=batch.guidance_rescale,
do_cfg=batch.do_classifier_free_guidance,
extra=extra,
)
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
batch = state.extra["batch"]
fastvideo_args = state.extra["fastvideo_args"]
target_dtype = state.extra["target_dtype"]
boundary_timestep = state.extra["boundary_timestep"]
if getattr(self.stage, "interrupt", False):
state.extra["skip_step"] = True
else:
state.extra["skip_step"] = False
if boundary_timestep is None or t >= boundary_timestep:
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and self.stage.transformer_2 is not None
and next(self.stage.transformer_2.parameters()).device.type
== 'cuda'):
self.stage.transformer_2.to('cpu')
current_model = self.stage.transformer
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and not fastvideo_args.use_fsdp_inference
and current_model is not None):
transformer_device = next(
current_model.parameters()).device.type
if transformer_device == 'cpu':
current_model.to(get_local_torch_device())
current_guidance_scale = batch.guidance_scale
else:
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and next(self.stage.transformer.parameters()).device.type
== 'cuda'):
self.stage.transformer.to('cpu')
current_model = self.stage.transformer_2
if (fastvideo_args.dit_cpu_offload
and not fastvideo_args.dit_layerwise_offload
and not fastvideo_args.use_fsdp_inference
and current_model is not None):
transformer_2_device = next(
current_model.parameters()).device.type
if transformer_2_device == 'cpu':
current_model.to(get_local_torch_device())
current_guidance_scale = batch.guidance_scale_2
assert current_model is not None, "current_model is None"
state.extra["current_model"] = current_model
state.extra["current_guidance_scale"] = current_guidance_scale
state.extra["step_idx"] = step_idx
latent_model_input = state.latents.to(target_dtype)
if batch.video_latent is not None:
latent_model_input = torch.cat([
latent_model_input, batch.video_latent,
torch.zeros_like(state.latents)
],
dim=1).to(target_dtype)
elif batch.image_latent is not None:
assert not fastvideo_args.pipeline_config.ti2v_task, (
"image latents should not be provided for TI2V task")
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
dim=1).to(target_dtype)
if (fastvideo_args.pipeline_config.ti2v_task
and batch.pil_image is not None):
timestep = torch.stack([t]).to(get_local_torch_device())
mask2 = state.extra["ti2v_mask"]
seq_len = state.extra["ti2v_seq_len"]
temp_ts = (mask2[0][:, ::2, ::2] * timestep).flatten()
temp_ts = torch.cat([
temp_ts,
temp_ts.new_ones(seq_len - temp_ts.size(0)) * timestep
])
timestep = temp_ts.unsqueeze(0)
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
else:
t_expand = t.repeat(latent_model_input.shape[0])
latent_model_input = self.stage.scheduler.scale_model_input(
latent_model_input, t)
guidance_expand = None
if fastvideo_args.pipeline_config.embedded_cfg_scale is not None:
guidance_expand = (torch.tensor(
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_local_torch_device(),
).to(target_dtype) * 1000.0)
state.extra["guidance_expand"] = guidance_expand
return ModelInputs(
latent_model_input=latent_model_input,
timestep=t_expand,
prompt_embeds=state.prompt_embeds,
prompt_attention_mask=state.prompt_attention_mask,
)
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
if state.extra.get("skip_step", False):
return state.latents
batch = state.extra["batch"]
fastvideo_args = state.extra["fastvideo_args"]
target_dtype = state.extra["target_dtype"]
autocast_enabled = state.extra["autocast_enabled"]
current_model = state.extra["current_model"]
current_guidance_scale = state.extra["current_guidance_scale"]
step_idx = state.extra["step_idx"]
guidance_expand = state.extra["guidance_expand"]
image_kwargs = state.extra["image_kwargs"]
pos_cond_kwargs = state.extra["pos_cond_kwargs"]
neg_cond_kwargs = state.extra["neg_cond_kwargs"]
action_kwargs = state.extra["action_kwargs"]
if ((st_attn_available
and self.stage.attn_backend == SlidingTileAttentionBackend) or
(vsa_available
and self.stage.attn_backend == VideoSparseAttentionBackend)):
self.attn_metadata_builder_cls = (
self.stage.attn_backend.get_builder_cls())
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls()
attn_metadata = self.attn_metadata_builder.build( # type: ignore
current_timestep=step_idx, # type: ignore
raw_latent_shape=batch.
raw_latent_shape[2:5], # type: ignore
patch_size=fastvideo_args.pipeline_config.dit_config.
patch_size, # type: ignore
STA_param=batch.STA_param, # type: ignore
VSA_sparsity=fastvideo_args.VSA_sparsity, # type: ignore
device=get_local_torch_device(),
)
assert attn_metadata is not None, (
"attn_metadata cannot be None")
else:
attn_metadata = None
elif (vmoba_attn_available
and self.stage.attn_backend == VMOBAAttentionBackend):
self.attn_metadata_builder_cls = (
self.stage.attn_backend.get_builder_cls())
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls()
moba_params = fastvideo_args.moba_config.copy()
moba_params.update({
"current_timestep":
step_idx,
"raw_latent_shape":
batch.raw_latent_shape[2:5],
"patch_size":
fastvideo_args.pipeline_config.dit_config.patch_size,
"device":
get_local_torch_device(),
})
attn_metadata = self.attn_metadata_builder.build(**moba_params)
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=target_dtype,
enabled=autocast_enabled):
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=step_idx,
attn_metadata=attn_metadata,
forward_batch=batch,
):
noise_pred = current_model(
model_inputs.latent_model_input,
state.prompt_embeds,
model_inputs.timestep,
guidance=guidance_expand,
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
)
if state.do_cfg:
batch.is_cfg_negative = True
with set_forward_context(
current_timestep=step_idx,
attn_metadata=attn_metadata,
forward_batch=batch,
):
noise_pred_uncond = current_model(
model_inputs.latent_model_input,
state.negative_prompt_embeds,
model_inputs.timestep,
guidance=guidance_expand,
**image_kwargs,
**neg_cond_kwargs,
**action_kwargs,
)
noise_pred_text = noise_pred
noise_pred = noise_pred_uncond + current_guidance_scale * (
noise_pred_text - noise_pred_uncond)
if state.guidance_rescale > 0.0:
noise_pred = self.stage.rescale_noise_cfg(
noise_pred,
noise_pred_text,
guidance_rescale=state.guidance_rescale,
)
return noise_pred
def cfg_combine(self, state: StrategyState,
noise_pred: torch.Tensor) -> torch.Tensor:
return noise_pred
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
if state.extra.get("skip_step", False):
return state.latents
batch = state.extra["batch"]
extra_step_kwargs = state.extra["extra_step_kwargs"]
latents = self.stage.scheduler.step(
noise_pred,
t,
state.latents,
**extra_step_kwargs,
return_dict=False,
)[0]
if (state.extra["ti2v_mask"] is not None
and batch.pil_image is not None):
mask2 = state.extra["ti2v_mask"]
z = state.extra["ti2v_z"]
latents = latents.squeeze(0)
latents = (1. - mask2) * z + mask2 * latents
if state.extra["trajectory_latents"] is not None:
state.extra["trajectory_timesteps"].append(t)
state.extra["trajectory_latents"].append(latents)
progress_bar = state.extra["progress_bar"]
step_idx = state.extra["step_idx"]
num_warmup_steps = state.extra["num_warmup_steps"]
timesteps = state.timesteps
if step_idx == len(timesteps) - 1 or (
(step_idx + 1) > num_warmup_steps and
(step_idx + 1) % self.stage.scheduler.order == 0
and progress_bar is not None):
progress_bar.update()
return latents
def postprocess(self, state: StrategyState) -> ForwardBatch:
batch = state.extra["batch"]
fastvideo_args = state.extra["fastvideo_args"]
progress_bar = state.extra.get("progress_bar")
if state.extra["trajectory_latents"]:
trajectory_tensor = torch.stack(state.extra["trajectory_latents"],
dim=1)
trajectory_timesteps_tensor = torch.stack(
state.extra["trajectory_timesteps"], dim=0)
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
batch.trajectory_latents = trajectory_tensor.cpu()
batch.latents = state.latents
if fastvideo_args.dit_layerwise_offload:
mgr = getattr(self.stage.transformer, "_layerwise_offload_manager",
None)
if mgr is not None and getattr(mgr, "enabled", False):
mgr.release_all()
if self.stage.transformer_2 is not None:
mgr2 = getattr(self.stage.transformer_2,
"_layerwise_offload_manager", None)
if mgr2 is not None and getattr(mgr2, "enabled", False):
mgr2.release_all()
if (st_attn_available
and self.stage.attn_backend == SlidingTileAttentionBackend
and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING):
self.stage.save_sta_search_results(batch)
pipeline = self.stage.pipeline() if self.stage.pipeline else None
if torch.backends.mps.is_available():
logger.info("Memory before deallocating transformer: %s",
torch.mps.current_allocated_memory())
del self.stage.transformer
if pipeline is not None and "transformer" in pipeline.modules:
del pipeline.modules["transformer"]
fastvideo_args.model_loaded["transformer"] = False
logger.info("Memory after deallocating transformer: %s",
torch.mps.current_allocated_memory())
if progress_bar is not None:
progress_bar.close()
return batch
@@ -0,0 +1,106 @@
# SPDX-License-Identifier: Apache-2.0
"""
Strategy interfaces and shared types for unified denoising.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Protocol, runtime_checkable
import torch
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
@dataclass
class StrategyState:
latents: torch.Tensor
timesteps: torch.Tensor
num_inference_steps: int
prompt_embeds: list[torch.Tensor]
negative_prompt_embeds: list[torch.Tensor] | None
prompt_attention_mask: list[torch.Tensor] | None
negative_attention_mask: list[torch.Tensor] | None
image_embeds: list[torch.Tensor]
guidance_scale: float
guidance_scale_2: float | None
guidance_rescale: float
do_cfg: bool
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
class ModelInputs:
latent_model_input: torch.Tensor
timestep: torch.Tensor
prompt_embeds: torch.Tensor | list[torch.Tensor]
prompt_attention_mask: torch.Tensor | list[torch.Tensor] | None
extra_kwargs: dict[str, Any] = field(default_factory=dict)
@dataclass
class BlockPlanItem:
start_index: int
num_frames: int
use_kv_cache: bool
model_selector: str
@dataclass
class BlockPlan:
items: list[BlockPlanItem] = field(default_factory=list)
@dataclass
class BlockContext:
kv_cache: list[dict] | None
kv_cache_2: list[dict] | None
crossattn_cache: list[dict] | None
action_cache: dict[str, list[dict]] | None
extra: dict[str, Any] = field(default_factory=dict)
class DenoisingStrategy(Protocol):
def prepare(self, batch: ForwardBatch, args) -> StrategyState:
...
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
...
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
...
def cfg_combine(self, state: StrategyState,
noise_pred: torch.Tensor) -> torch.Tensor:
...
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
...
def postprocess(self, state: StrategyState) -> ForwardBatch:
...
@runtime_checkable
class BlockDenoisingStrategy(DenoisingStrategy, Protocol):
def block_plan(self, state: StrategyState) -> BlockPlan:
...
def init_block_context(self, state: StrategyState,
block_item: BlockPlanItem,
block_idx: int) -> BlockContext:
...
def process_block(self, state: StrategyState, block_ctx: BlockContext,
block_item: BlockPlanItem) -> None:
...
def update_context(self, state: StrategyState, block_ctx: BlockContext,
block_item: BlockPlanItem) -> None:
...
@@ -1,179 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat-specific denoising stage implementing CFG-zero optimized guidance.
"""
import torch
from tqdm import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.forward_context import set_forward_context
logger = init_logger(__name__)
class LongCatDenoisingStage(DenoisingStage):
"""
LongCat denoising stage with CFG-zero optimized guidance scale.
Implements:
1. Optimized CFG scale from CFG-zero paper
2. Negation of noise prediction before scheduler step (flow matching convention)
3. Batched CFG computation (unlike standard FastVideo separate passes)
"""
def optimized_scale(self, positive_flat, negative_flat) -> torch.Tensor:
"""
Calculate optimized scale from CFG-zero paper.
st_star = (v_cond^T * v_uncond) / ||v_uncond||^2
Args:
positive_flat: Conditional prediction, flattened [B, -1]
negative_flat: Unconditional prediction, flattened [B, -1]
Returns:
st_star: Optimized scale [B, 1]
"""
# Calculate dot product
dot_product = torch.sum(positive_flat * negative_flat,
dim=1,
keepdim=True)
# Squared norm of uncondition
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
# st_star = v_cond^T * v_uncond / ||v_uncond||^2
st_star = dot_product / squared_norm
return st_star
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Run LongCat denoising loop with optimized CFG.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
if not fastvideo_args.model_loaded["transformer"]:
from fastvideo.models.loader.component_loader import TransformerLoader
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
pipeline = self.pipeline() if self.pipeline else None
if pipeline:
pipeline.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
# Get transformer dtype
if hasattr(self.transformer, 'module'):
transformer_dtype = next(self.transformer.module.parameters()).dtype
else:
transformer_dtype = next(self.transformer.parameters()).dtype
target_dtype = transformer_dtype
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Extract batch parameters
latents = batch.latents
timesteps = batch.timesteps
prompt_embeds = batch.prompt_embeds[0] # LongCat uses single encoder
prompt_attention_mask = batch.prompt_attention_mask[
0] if batch.prompt_attention_mask else None
guidance_scale = batch.guidance_scale
do_classifier_free_guidance = batch.do_classifier_free_guidance
# Get negative prompts if doing CFG
if do_classifier_free_guidance:
negative_prompt_embeds = batch.negative_prompt_embeds[0]
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
if batch.negative_attention_mask
else None)
# Concatenate for batched processing
prompt_embeds_combined = torch.cat(
[negative_prompt_embeds, prompt_embeds], dim=0)
if prompt_attention_mask is not None:
prompt_attention_mask_combined = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask],
dim=0)
else:
prompt_attention_mask_combined = None
else:
prompt_embeds_combined = prompt_embeds
prompt_attention_mask_combined = prompt_attention_mask
# Denoising loop
num_inference_steps = len(timesteps)
with tqdm(total=num_inference_steps,
desc="LongCat Denoising") as progress_bar:
for i, t in enumerate(timesteps):
# Expand latents for CFG
if do_classifier_free_guidance:
latent_model_input = torch.cat([latents] * 2)
else:
latent_model_input = latents
latent_model_input = latent_model_input.to(target_dtype)
# Expand timestep to match batch size
timestep = t.expand(
latent_model_input.shape[0]).to(target_dtype)
# Run transformer with context
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
), torch.autocast(device_type='cuda',
dtype=target_dtype,
enabled=autocast_enabled):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds_combined,
timestep=timestep,
encoder_attention_mask=prompt_attention_mask_combined,
)
# Apply CFG with optimized scale
if do_classifier_free_guidance:
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
B = noise_pred_cond.shape[0]
positive = noise_pred_cond.reshape(B, -1)
negative = noise_pred_uncond.reshape(B, -1)
# Calculate optimized scale (CFG-zero)
st_star = self.optimized_scale(positive, negative)
# Reshape for broadcasting
st_star = st_star.view(B, 1, 1, 1, 1)
# Apply optimized CFG formula
noise_pred = (
noise_pred_uncond * st_star + guidance_scale *
(noise_pred_cond - noise_pred_uncond * st_star))
# CRITICAL: Negate noise prediction for flow matching scheduler
noise_pred = -noise_pred
# Compute previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
progress_bar.update()
# Update batch with denoised latents
batch.latents = latents
return batch
@@ -1,171 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat I2V Denoising Stage with conditioning support.
This stage implements Tier 3 I2V denoising:
1. Per-frame timestep masking (timestep[:, :num_cond_latents] = 0)
2. Passes num_cond_latents to transformer (for RoPE skipping)
3. Selective denoising (only updates non-conditioned frames)
4. CFG-zero optimized guidance
"""
import torch
from tqdm import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
logger = init_logger(__name__)
class LongCatI2VDenoisingStage(LongCatDenoisingStage):
"""
LongCat denoising with I2V conditioning support.
Key modifications from base LongCat denoising:
1. Sets timestep=0 for conditioning frames
2. Passes num_cond_latents to transformer
3. Only applies scheduler step to non-conditioned frames
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""Run denoising loop with I2V conditioning."""
# Load transformer if needed
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
fastvideo_args.model_loaded["transformer"] = True
# Setup
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latents = batch.latents
timesteps = batch.timesteps
prompt_embeds = batch.prompt_embeds[0]
prompt_attention_mask = (batch.prompt_attention_mask[0]
if batch.prompt_attention_mask else None)
guidance_scale = batch.guidance_scale
do_classifier_free_guidance = batch.do_classifier_free_guidance
# Get num_cond_latents from batch
num_cond_latents = getattr(batch, 'num_cond_latents', 0)
if num_cond_latents > 0:
logger.info("I2V Denoising: num_cond_latents=%s, latent_shape=%s",
num_cond_latents, latents.shape)
# Prepare negative prompts for CFG
if do_classifier_free_guidance:
negative_prompt_embeds = batch.negative_prompt_embeds[0]
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
if batch.negative_attention_mask
else None)
prompt_embeds_combined = torch.cat(
[negative_prompt_embeds, prompt_embeds], dim=0)
if prompt_attention_mask is not None:
prompt_attention_mask_combined = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask],
dim=0)
else:
prompt_attention_mask_combined = None
else:
prompt_embeds_combined = prompt_embeds
prompt_attention_mask_combined = prompt_attention_mask
# Denoising loop
num_inference_steps = len(timesteps)
with tqdm(total=num_inference_steps,
desc="I2V Denoising") as progress_bar:
for i, t in enumerate(timesteps):
# 1. Expand latents for CFG
if do_classifier_free_guidance:
latent_model_input = torch.cat([latents] * 2)
else:
latent_model_input = latents
latent_model_input = latent_model_input.to(target_dtype)
# 2. Expand timestep to match batch size
timestep = t.expand(
latent_model_input.shape[0]).to(target_dtype)
# 3. CRITICAL: Expand timestep to temporal dimension
# and set conditioning frames to timestep=0
timestep = timestep.unsqueeze(-1).repeat(
1, latent_model_input.shape[2])
# Mark conditioning frames as clean (timestep=0)
if num_cond_latents > 0:
timestep[:, :num_cond_latents] = 0
# 4. Run transformer with num_cond_latents
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
), torch.autocast(device_type='cuda',
dtype=target_dtype,
enabled=autocast_enabled):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds_combined,
timestep=timestep,
encoder_attention_mask=prompt_attention_mask_combined,
num_cond_latents=num_cond_latents,
)
# 5. Apply CFG with optimized scale
if do_classifier_free_guidance:
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
B = noise_pred_cond.shape[0]
positive = noise_pred_cond.reshape(B, -1)
negative = noise_pred_uncond.reshape(B, -1)
# CFG-zero optimized scale
st_star = self.optimized_scale(positive, negative)
st_star = st_star.view(B, 1, 1, 1, 1)
noise_pred = (
noise_pred_uncond * st_star + guidance_scale *
(noise_pred_cond - noise_pred_uncond * st_star))
# 6. CRITICAL: Negate for flow matching scheduler
noise_pred = -noise_pred
# 7. CRITICAL: Only update non-conditioned frames
# The conditioning frames stay FIXED throughout denoising
if num_cond_latents > 0:
latents[:, :, num_cond_latents:] = self.scheduler.step(
noise_pred[:, :, num_cond_latents:],
t,
latents[:, :, num_cond_latents:],
return_dict=False)[0]
else:
# No conditioning, update all frames
latents = self.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
progress_bar.update()
# Update batch with denoised latents
batch.latents = latents
return batch
@@ -1,217 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat VC Denoising Stage with KV cache support.
This stage extends the I2V denoising stage to support:
1. KV cache for conditioning frames
2. Video continuation with multiple conditioning frames
"""
import time
import torch
from tqdm import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.longcat_denoising import LongCatDenoisingStage
logger = init_logger(__name__)
class LongCatVCDenoisingStage(LongCatDenoisingStage):
"""
LongCat denoising with Video Continuation and KV cache support.
Key differences from I2V denoising:
- Supports KV cache (reuses cached K/V from conditioning frames)
- Handles larger num_cond_latents
- Concatenates conditioning latents back after denoising
When use_kv_cache=True:
- batch.latents contains ONLY noise frames (cond removed by KV cache init)
- batch.kv_cache_dict contains cached K/V
- batch.cond_latents contains conditioning latents for post-concat
When use_kv_cache=False:
- batch.latents contains ALL frames (cond + noise)
- Timestep masking: timestep[:, :num_cond_latents] = 0
- Selective denoising: only update noise frames
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""Run denoising loop with VC conditioning and optional KV cache."""
# Load transformer if needed
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
fastvideo_args.model_loaded["transformer"] = True
# Setup
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latents = batch.latents
timesteps = batch.timesteps
prompt_embeds = batch.prompt_embeds[0]
prompt_attention_mask = (batch.prompt_attention_mask[0]
if batch.prompt_attention_mask else None)
guidance_scale = batch.guidance_scale
do_classifier_free_guidance = batch.do_classifier_free_guidance
# Get VC-specific parameters
num_cond_latents = getattr(batch, 'num_cond_latents', 0)
use_kv_cache = getattr(batch, 'use_kv_cache', False)
kv_cache_dict = getattr(batch, 'kv_cache_dict', {})
logger.info(
"VC Denoising: num_cond_latents=%d, use_kv_cache=%s, latent_shape=%s",
num_cond_latents, use_kv_cache, latents.shape)
# Prepare negative prompts for CFG
if do_classifier_free_guidance:
negative_prompt_embeds = batch.negative_prompt_embeds[0]
negative_prompt_attention_mask = (batch.negative_attention_mask[0]
if batch.negative_attention_mask
else None)
prompt_embeds_combined = torch.cat(
[negative_prompt_embeds, prompt_embeds], dim=0)
if prompt_attention_mask is not None:
prompt_attention_mask_combined = torch.cat(
[negative_prompt_attention_mask, prompt_attention_mask],
dim=0)
else:
prompt_attention_mask_combined = None
else:
prompt_embeds_combined = prompt_embeds
prompt_attention_mask_combined = prompt_attention_mask
# Denoising loop
num_inference_steps = len(timesteps)
step_times = []
with tqdm(total=num_inference_steps,
desc="VC Denoising") as progress_bar:
for i, t in enumerate(timesteps):
step_start = time.time()
# 1. Expand latents for CFG
if do_classifier_free_guidance:
latent_model_input = torch.cat([latents] * 2)
else:
latent_model_input = latents
latent_model_input = latent_model_input.to(target_dtype)
# 2. Expand timestep to match batch size
timestep = t.expand(
latent_model_input.shape[0]).to(target_dtype)
# 3. Expand timestep to temporal dimension
timestep = timestep.unsqueeze(-1).repeat(
1, latent_model_input.shape[2])
# 4. Timestep masking (only when NOT using KV cache)
if not use_kv_cache and num_cond_latents > 0:
timestep[:, :num_cond_latents] = 0
# 5. Prepare transformer kwargs
# IMPORTANT: num_cond_latents is ALWAYS passed - needed for RoPE position offset
transformer_kwargs = {
'num_cond_latents': num_cond_latents,
}
if use_kv_cache:
transformer_kwargs['kv_cache_dict'] = kv_cache_dict
# 6. Run transformer
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
), torch.autocast(device_type='cuda',
dtype=target_dtype,
enabled=autocast_enabled):
noise_pred = self.transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds_combined,
timestep=timestep,
encoder_attention_mask=prompt_attention_mask_combined,
**transformer_kwargs,
)
# 7. Apply CFG with optimized scale (CFG-zero)
if do_classifier_free_guidance:
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
B = noise_pred_cond.shape[0]
positive = noise_pred_cond.reshape(B, -1)
negative = noise_pred_uncond.reshape(B, -1)
st_star = self.optimized_scale(positive, negative)
st_star = st_star.view(B, 1, 1, 1, 1)
noise_pred = (
noise_pred_uncond * st_star + guidance_scale *
(noise_pred_cond - noise_pred_uncond * st_star))
# 8. Negate for flow matching scheduler
noise_pred = -noise_pred
# 9. Scheduler step
if use_kv_cache:
# All latents are noise frames (conditioning is in cache)
latents = self.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
else:
# Only update noise frames (skip conditioning)
if num_cond_latents > 0:
latents[:, :, num_cond_latents:] = self.scheduler.step(
noise_pred[:, :, num_cond_latents:],
t,
latents[:, :, num_cond_latents:],
return_dict=False,
)[0]
else:
latents = self.scheduler.step(noise_pred,
t,
latents,
return_dict=False)[0]
step_time = time.time() - step_start
step_times.append(step_time)
# Log timing for first few steps
if i < 3:
logger.info("Step %d: %.2fs", i, step_time)
progress_bar.update()
# 10. If using KV cache, concatenate conditioning latents back
if use_kv_cache and hasattr(
batch, 'cond_latents') and batch.cond_latents is not None:
latents = torch.cat([batch.cond_latents, latents], dim=2)
logger.info(
"Concatenated conditioning latents back, final shape: %s",
latents.shape)
# Log average timing
avg_time = sum(step_times) / len(step_times)
logger.info("Average step time: %.2fs (total: %.1fs)", avg_time,
sum(step_times))
# Update batch with denoised latents
batch.latents = latents
return batch
@@ -15,14 +15,6 @@ from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
@@ -123,169 +115,22 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
self._streaming_initialized: bool = False
self._streaming_ctx: BlockProcessingContext | None = None
self._streaming_engine = None
self._streaming_state = None
self._streaming_block_plan = None
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
from fastvideo.pipelines.stages.denoising_matrixgame_strategy import (
MatrixGameBlockStrategy)
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
patch_size = self.transformer.patch_size
patch_ratio = patch_size[-1] * patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
'boundary_ratio', None)
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
# directly set the kwarg.
image_kwargs = {"encoder_hidden_states_image": image_embeds}
pos_cond_kwargs: dict[str, Any] = {}
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
b, c, t, h, w = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache_mouse = None
kv_cache_keyboard = None
if self.use_action_module:
kv_cache_mouse, kv_cache_keyboard = self._initialize_action_kv_cache(
batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=257, # 1 CLS + 256 patch tokens
dtype=target_dtype,
device=latents.device)
if t % self.num_frame_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frame_per_block for causal denoising"
)
num_blocks = t // self.num_frame_per_block
block_sizes = [self.num_frame_per_block] * num_blocks
start_index = 0
if boundary_timestep is not None:
block_sizes[0] = 1
# NOTE: MatrixGame does NOT process the first frame separately.
# The first frame information is already encoded in batch.image_latent (cond_concat)
# and will be used by the model via channel concatenation: torch.cat([x, cond_concat], dim=1)
ctx = BlockProcessingContext(
batch=batch,
block_idx=0,
start_index=0,
kv_cache1=kv_cache1,
kv_cache2=kv_cache2,
kv_cache_mouse=kv_cache_mouse,
kv_cache_keyboard=kv_cache_keyboard,
crossattn_cache=crossattn_cache,
timesteps=timesteps,
block_sizes=block_sizes,
noise_pool=None,
fastvideo_args=fastvideo_args,
target_dtype=target_dtype,
autocast_enabled=autocast_enabled,
boundary_timestep=boundary_timestep,
high_noise_timesteps=high_noise_timesteps,
context_noise=getattr(fastvideo_args.pipeline_config,
"context_noise", 0),
image_kwargs=image_kwargs,
pos_cond_kwargs=pos_cond_kwargs,
)
context_noise = getattr(fastvideo_args.pipeline_config, "context_noise",
0)
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for block_idx, current_num_frames in enumerate(block_sizes):
ctx.block_idx = block_idx
ctx.start_index = start_index
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
action_kwargs = self._prepare_action_kwargs(
batch, start_index, current_num_frames)
current_latents = self._process_single_block(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
timesteps=timesteps,
ctx=ctx,
action_kwargs=action_kwargs,
progress_bar=progress_bar,
)
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Update KV caches with clean context
self._update_context_cache(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=action_kwargs,
context_noise=context_noise,
)
start_index += current_num_frames
if boundary_timestep is not None:
num_frames_to_remove = self.num_frame_per_block - 1
if num_frames_to_remove > 0:
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
engine = DenoisingEngine(MatrixGameBlockStrategy(self),
hooks=self._build_engine_hooks())
return engine.run(batch, fastvideo_args)
def _prepare_action_kwargs(self, batch: ForwardBatch, start_index: int,
num_frames: int) -> dict[str, Any]:
@@ -652,123 +497,44 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
def streaming_reset(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
from fastvideo.pipelines.stages.denoising_matrixgame_strategy import (
MatrixGameBlockStrategy)
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
patch_size = self.transformer.patch_size
patch_ratio = patch_size[-1] * patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
strategy = MatrixGameBlockStrategy(self)
engine = DenoisingEngine(strategy, hooks=self._build_engine_hooks())
state = strategy.prepare(batch, fastvideo_args)
for hook in engine.hooks:
hook.on_init(engine, batch, fastvideo_args)
for hook in engine.hooks:
hook.pre_run(state)
block_plan = strategy.block_plan(state)
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
progress_bar = state.extra.get("progress_bar")
if progress_bar is not None:
progress_bar.close()
state.extra["progress_bar"] = None
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
'boundary_ratio', None)
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
# directly set the kwarg.
image_kwargs = {"encoder_hidden_states_image": image_embeds}
pos_cond_kwargs: dict[str, Any] = {}
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
ctx = state.extra["ctx"]
latents = state.latents
assert latents is not None, "latents must be provided"
b, c, t, h, w = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize caches
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache_mouse = None
kv_cache_keyboard = None
if self.use_action_module:
kv_cache_mouse, kv_cache_keyboard = self._initialize_action_kv_cache(
batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=257, # 1 CLS + 256 patch tokens
dtype=target_dtype,
device=latents.device)
# Calculate block sizes
if t % self.num_frame_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frame_per_block for causal denoising"
)
num_blocks = t // self.num_frame_per_block
block_sizes = [self.num_frame_per_block] * num_blocks
if boundary_timestep is not None:
block_sizes[0] = 1
# Pre-allocate noise pool
num_denoising_steps = len(timesteps)
num_denoising_steps = len(state.timesteps)
noise_shape = (b, self.num_frame_per_block, c, h, w)
noise_pool = [
torch.randn(
noise_shape,
dtype=target_dtype,
dtype=ctx.target_dtype,
device=latents.device,
) for _ in range(num_denoising_steps - 1)
) for _ in range(max(num_denoising_steps - 1, 0))
]
ctx.noise_pool = noise_pool
# Create and store context
self._streaming_ctx = BlockProcessingContext(
batch=batch,
block_idx=0,
start_index=0,
kv_cache1=kv_cache1,
kv_cache2=kv_cache2,
kv_cache_mouse=kv_cache_mouse,
kv_cache_keyboard=kv_cache_keyboard,
crossattn_cache=crossattn_cache,
timesteps=timesteps,
block_sizes=block_sizes,
noise_pool=noise_pool,
fastvideo_args=fastvideo_args,
target_dtype=target_dtype,
autocast_enabled=autocast_enabled,
boundary_timestep=boundary_timestep,
high_noise_timesteps=high_noise_timesteps,
context_noise=getattr(fastvideo_args.pipeline_config,
"context_noise", 0),
image_kwargs=image_kwargs,
pos_cond_kwargs=pos_cond_kwargs,
)
self._streaming_engine = engine
self._streaming_state = state
self._streaming_block_plan = block_plan
self._streaming_ctx = ctx
self._streaming_initialized = True
return batch
@@ -776,23 +542,25 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
self,
keyboard_action: torch.Tensor | None = None,
mouse_action: torch.Tensor | None = None) -> ForwardBatch:
if not self._streaming_initialized or self._streaming_ctx is None:
if (not self._streaming_initialized or self._streaming_ctx is None
or self._streaming_engine is None
or self._streaming_state is None
or self._streaming_block_plan is None):
raise RuntimeError(
"Streaming not initialized! Call streaming_reset first.")
ctx = self._streaming_ctx
if ctx.block_idx >= len(ctx.block_sizes):
block_plan = self._streaming_block_plan
if ctx.block_idx >= len(block_plan.items):
return ctx.batch
batch = ctx.batch
latents = batch.latents
assert latents is not None, "latents must be set in batch"
current_num_frames = ctx.block_sizes[ctx.block_idx]
start_index = ctx.start_index
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
block_item = block_plan.items[ctx.block_idx]
start_index = block_item.start_index
current_num_frames = block_item.num_frames
# Update batch with new actions for this block
if keyboard_action is not None or mouse_action is not None:
@@ -810,58 +578,32 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
batch.mouse_cond[:, start_frame:start_frame +
n] = mouse_action.to(batch.mouse_cond.device)
action_kwargs = self._prepare_action_kwargs(batch, start_index,
current_num_frames)
# Create noise generator that uses pre-allocated noise pool
def streaming_noise_generator(shape: tuple, dtype: torch.dtype,
step_idx: int) -> torch.Tensor:
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:
# Fallback to dynamic allocation if pool not available
return torch.randn(
shape,
dtype=dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else batch.generator)).to(
latents.device)
current_latents = self._process_single_block(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
timesteps=ctx.timesteps,
ctx=ctx,
action_kwargs=action_kwargs,
noise_generator=streaming_noise_generator,
)
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Update KV caches with clean context
self._update_context_cache(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=action_kwargs,
context_noise=ctx.context_noise,
self._streaming_engine.run_blocks(
self._streaming_state,
block_plan=block_plan,
start_block=ctx.block_idx,
num_blocks=1,
)
# Advance streaming state
ctx.start_index += current_num_frames
ctx.start_index = start_index + current_num_frames
ctx.block_idx += 1
return batch
def streaming_clear(self) -> None:
if (self._streaming_engine is not None
and self._streaming_state is not None):
batch = (self._streaming_ctx.batch if self._streaming_ctx
is not None else self._streaming_state.extra.get("batch"))
if batch is not None:
for hook in self._streaming_engine.hooks:
hook.post_run(self._streaming_state, batch)
self._streaming_initialized = False
self._streaming_ctx = None
self._streaming_engine = None
self._streaming_state = None
self._streaming_block_plan = None
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
+5
View File
@@ -21,3 +21,8 @@ def distributed_setup():
yield
cleanup_dist_env_and_memory()
def pytest_configure(config):
config.addinivalue_line("markers",
"gpu: requires CUDA with BF16 support")
@@ -1,11 +0,0 @@
{
"mean_ssim": 0.7606927804004999,
"min_ssim": 0.7035917639732361,
"max_ssim": 0.7920367121696472,
"reference_video": "/mnt/fast-disks/hao_lab/loay/FastVideo/fastvideo/tests/ssim/L40S_reference_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
"generated_video": "/mnt/fast-disks/hao_lab/loay/FastVideo/fastvideo/tests/ssim/generated_videos/TurboWan2.1-T2V-1.3B-Diffusers/SLA_ATTN/Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of.mp4",
"parameters": {
"num_inference_steps": 4,
"prompt": "Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
}
}
@@ -19,6 +19,8 @@ if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "A100" in device_name:
device_reference_folder = "A100" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}")
@@ -22,6 +22,8 @@ if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "A100" in device_name:
device_reference_folder = "A100" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}")
@@ -24,6 +24,8 @@ if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "A100" in device_name:
device_reference_folder = "A100" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}, using L40S references")
@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
from fastvideo.pipelines.stages.denoising_engine_hooks import PerfLoggingHook
from fastvideo.pipelines.stages.denoising_strategies import (ModelInputs,
StrategyState)
class DummyStrategy:
def prepare(self, batch: ForwardBatch,
args: FastVideoArgs) -> StrategyState:
latents = torch.zeros(1, 1, 1, 1, 1)
timesteps = torch.tensor([2, 1], dtype=torch.long)
return StrategyState(
latents=latents,
timesteps=timesteps,
num_inference_steps=len(timesteps),
prompt_embeds=[torch.zeros(1, 1, 1)],
negative_prompt_embeds=None,
prompt_attention_mask=None,
negative_attention_mask=None,
image_embeds=[],
guidance_scale=1.0,
guidance_scale_2=None,
guidance_rescale=0.0,
do_cfg=False,
extra={"batch": batch},
)
def make_model_inputs(self, state: StrategyState, t: torch.Tensor,
step_idx: int) -> ModelInputs:
return ModelInputs(
latent_model_input=state.latents,
timestep=t,
prompt_embeds=state.prompt_embeds,
prompt_attention_mask=state.prompt_attention_mask,
)
def forward(self, state: StrategyState,
model_inputs: ModelInputs) -> torch.Tensor:
return state.latents
def cfg_combine(self, state: StrategyState,
noise_pred: torch.Tensor) -> torch.Tensor:
return noise_pred
def scheduler_step(self, state: StrategyState, noise_pred: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
return state.latents
def postprocess(self, state: StrategyState) -> ForwardBatch:
return state.extra["batch"]
def test_perf_logging_hook_records_steps():
batch = ForwardBatch(data_type="test")
args = FastVideoArgs(model_path="dummy")
hook = PerfLoggingHook()
engine = DenoisingEngine(DummyStrategy(), hooks=[hook])
engine.run(batch, args)
info = batch.logging_info.get_stage_info("DenoisingEngine")
step_times = info.get("denoise_step_times_ms")
total_ms = info.get("denoise_total_ms")
assert isinstance(step_times, list)
assert len(step_times) == 2
assert total_ms is not None
assert total_ms >= 0.0
@@ -0,0 +1,811 @@
# SPDX-License-Identifier: Apache-2.0
import contextlib
from types import SimpleNamespace
import torch
import pytest
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising_engine import DenoisingEngine
from fastvideo.pipelines.stages.denoising_strategies import (
BlockDenoisingStrategy)
from fastvideo.pipelines.stages.denoising_standard_strategy import (
StandardStrategy)
from fastvideo.pipelines.stages.denoising_cosmos_strategy import (
CosmosStrategy)
from fastvideo.pipelines.stages.denoising_dmd_strategy import DmdStrategy
from fastvideo.pipelines.stages.denoising_longcat_strategy import (
LongCatStrategy,
LongCatI2VStrategy,
LongCatVCStrategy,
)
from fastvideo.pipelines.stages.denoising_causal_strategy import (
CausalBlockStrategy)
from fastvideo.pipelines.stages.denoising_matrixgame_strategy import (
MatrixGameBlockStrategy)
class DummyProgressBar:
def __init__(self):
self.updates = 0
def update(self):
self.updates += 1
def close(self):
return None
class DummyScheduler:
def __init__(self, timesteps: torch.Tensor):
self.order = 1
self.timesteps = timesteps
self.sigmas = self._make_sigmas(timesteps)
self.num_train_timesteps = int(timesteps.max().item()
) if timesteps.numel() else 1
self.init_noise_sigma = 1.0
self.config = SimpleNamespace(final_sigmas_type=None)
def _make_sigmas(self, timesteps: torch.Tensor) -> torch.Tensor:
if timesteps.numel() == 0:
return timesteps
max_val = float(timesteps.max().item()) or 1.0
return timesteps.float() / max_val
def register_to_config(self, **kwargs):
for key, value in kwargs.items():
setattr(self.config, key, value)
def set_timesteps(self, num_steps: int, device=None):
self.timesteps = torch.linspace(1.0, 0.0, num_steps, device=device)
self.sigmas = self._make_sigmas(self.timesteps)
return self.timesteps
def scale_model_input(self, latents: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
return latents
def step(self, noise_pred: torch.Tensor, t: torch.Tensor,
latents: torch.Tensor, **kwargs):
return (latents - noise_pred * 0.1, )
def add_noise(self, latents: torch.Tensor, noise: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
return latents + noise * 0.01
def add_noise_high(self, latents: torch.Tensor, noise: torch.Tensor,
t: torch.Tensor, boundary_timestep: torch.Tensor
) -> torch.Tensor:
return latents + noise * 0.01
class DummyTransformer(torch.nn.Module):
def __init__(self,
hidden_size: int = 8,
num_attention_heads: int = 2,
patch_size: tuple[int, int, int] = (1, 1, 1),
num_frames_per_block: int = 1,
sliding_window_num_frames: int = 1,
local_attn_size: int = 4) -> None:
super().__init__()
self.dummy = torch.nn.Parameter(torch.zeros(1))
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
self.attention_head_dim = hidden_size // num_attention_heads
self.blocks = [object(), object()]
self.config = SimpleNamespace(
arch_config=SimpleNamespace(
patch_size=patch_size,
num_frames_per_block=num_frames_per_block,
sliding_window_num_frames=sliding_window_num_frames,
))
self.patch_size = patch_size
self.local_attn_size = local_attn_size
self.model = SimpleNamespace(local_attn_size=local_attn_size)
self.independent_first_frame = False
def forward(
self,
hidden_states=None,
encoder_hidden_states=None,
timestep=None,
guidance=None,
encoder_hidden_states_image=None,
encoder_hidden_states_2=None,
encoder_attention_mask=None,
mask_strategy=None,
mouse_cond=None,
keyboard_cond=None,
kv_cache=None,
crossattn_cache=None,
current_start=None,
start_frame=None,
num_cond_latents=None,
*args,
**kwargs,
):
latent = hidden_states
if latent is None and args:
latent = args[0]
if latent is None:
raise ValueError("Missing latent input")
t_val = timestep
if t_val is None and len(args) > 2:
t_val = args[2]
scale = float(torch.as_tensor(t_val).float().mean().item()
) if t_val is not None else 0.0
out = latent * 0.01 + scale
out_channels = getattr(self, "out_channels", None)
if out_channels is not None and out.shape[1] != out_channels:
out = out[:, :out_channels, ...]
if kwargs.get("return_dict") is False:
return (out, )
return out
class DummyStandardStage:
def __init__(self, transformer, scheduler):
self.transformer = transformer
self.transformer_2 = None
self.scheduler = scheduler
self.vae = None
self.pipeline = None
self.attn_backend = object()
def prepare_extra_func_kwargs(self, func, kwargs):
extra_step_kwargs = {}
for key, value in kwargs.items():
if key in set(func.__code__.co_varnames):
extra_step_kwargs[key] = value
return extra_step_kwargs
def progress_bar(self, iterable=None, total=None):
return DummyProgressBar()
def rescale_noise_cfg(self, noise_cfg, noise_pred_text,
guidance_rescale=0.0) -> torch.Tensor:
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)),
keepdim=True)
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)),
keepdim=True)
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
return (guidance_rescale * noise_pred_rescaled +
(1 - guidance_rescale) * noise_cfg)
def prepare_sta_param(self, batch, fastvideo_args):
return None
def save_sta_search_results(self, batch):
return None
class DummyLongCatStage:
def __init__(self, transformer, scheduler):
self.transformer = transformer
self.scheduler = scheduler
def optimized_scale(self, positive_flat, negative_flat) -> torch.Tensor:
dot_product = torch.sum(positive_flat * negative_flat,
dim=1,
keepdim=True)
squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8
return dot_product / squared_norm
class DummyDmdStage(DummyStandardStage):
@property
def device(self) -> torch.device:
return torch.device("cpu")
class DummyCausalStage:
def __init__(self, transformer, scheduler):
self.transformer = transformer
self.transformer_2 = None
self.scheduler = scheduler
self.vae = None
self.num_transformer_blocks = len(transformer.blocks)
self.num_frames_per_block = transformer.config.arch_config.num_frames_per_block
self.sliding_window_num_frames = (
transformer.config.arch_config.sliding_window_num_frames)
self.local_attn_size = transformer.model.local_attn_size
self.attn_backend = object()
self.frame_seq_length = 1
@property
def device(self) -> torch.device:
return next(self.transformer.parameters()).device
def prepare_extra_func_kwargs(self, func, kwargs):
extra_step_kwargs = {}
for key, value in kwargs.items():
if key in set(func.__code__.co_varnames):
extra_step_kwargs[key] = value
return extra_step_kwargs
def prepare_sta_param(self, batch, fastvideo_args):
return None
def progress_bar(self, iterable=None, total=None):
return DummyProgressBar()
def _initialize_kv_cache(self, batch_size, dtype, device):
return [{} for _ in range(self.num_transformer_blocks)]
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device):
return [{} for _ in range(self.num_transformer_blocks)]
class DummyMatrixGameStage:
def __init__(self, transformer, scheduler):
self.transformer = transformer
self.transformer_2 = None
self.scheduler = scheduler
self.vae = None
self.num_transformer_blocks = len(transformer.blocks)
self.num_frame_per_block = transformer.config.arch_config.num_frames_per_block
self.sliding_window_num_frames = (
transformer.config.arch_config.sliding_window_num_frames)
self.local_attn_size = transformer.local_attn_size
self.use_action_module = False
self.attn_backend = object()
self.frame_seq_length = 1
@property
def device(self) -> torch.device:
return torch.device("cpu")
def progress_bar(self, iterable=None, total=None):
return DummyProgressBar()
def prepare_sta_param(self, batch, fastvideo_args):
return None
def _initialize_kv_cache(self, batch_size, dtype, device):
return [{} for _ in range(self.num_transformer_blocks)]
def _initialize_action_kv_cache(self, batch_size, dtype, device):
return None, None
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device):
return [{} for _ in range(self.num_transformer_blocks)]
def _prepare_action_kwargs(self, batch: ForwardBatch, start_index: int,
num_frames: int) -> dict:
return {}
def _process_single_block(self, current_latents: torch.Tensor,
batch: ForwardBatch, start_index: int,
current_num_frames: int, timesteps: torch.Tensor,
ctx, action_kwargs, noise_generator=None,
progress_bar=None):
latents = current_latents
for t_cur in timesteps:
noise_pred = self.transformer(latents, batch.prompt_embeds, t_cur)
latents = self.scheduler.step(noise_pred,
t_cur,
latents,
return_dict=False)[0]
if progress_bar is not None:
progress_bar.update()
return latents
def _update_context_cache(self, current_latents: torch.Tensor,
batch: ForwardBatch, start_index: int,
current_num_frames: int, ctx, action_kwargs,
context_noise: float):
return None
@pytest.fixture(autouse=True)
def _patch_autocast(monkeypatch, request):
if request.node.get_closest_marker("gpu"):
return
@contextlib.contextmanager
def _autocast(*args, **kwargs):
yield
monkeypatch.setattr(torch, "autocast", _autocast)
@pytest.fixture(autouse=True)
def _patch_local_device(monkeypatch, request):
if request.node.get_closest_marker("gpu"):
return
def _cpu_device():
return torch.device("cpu")
import fastvideo.distributed as dist
import fastvideo.pipelines.stages.denoising_standard_strategy as standard
import fastvideo.pipelines.stages.denoising_dmd_strategy as dmd
import fastvideo.pipelines.stages.denoising_causal_strategy as causal
import fastvideo.pipelines.stages.denoising_matrixgame_strategy as matrixgame
monkeypatch.setattr(dist, "get_local_torch_device", _cpu_device)
monkeypatch.setattr(standard, "get_local_torch_device", _cpu_device)
monkeypatch.setattr(dmd, "get_local_torch_device", _cpu_device)
monkeypatch.setattr(causal, "get_local_torch_device", _cpu_device)
monkeypatch.setattr(matrixgame, "get_local_torch_device", _cpu_device)
def _make_args() -> FastVideoArgs:
args = FastVideoArgs(model_path="dummy")
args.pipeline_config.dit_config.arch_config.patch_size = (1, 1, 1)
args.pipeline_config.dit_config.boundary_ratio = None
args.pipeline_config.ti2v_task = False
args.pipeline_config.embedded_cfg_scale = None
return args
def _make_batch(latents: torch.Tensor,
timesteps: torch.Tensor,
guidance_scale: float,
seed: int,
device: torch.device | str | None = None,
generator_device: torch.device | str = "cpu") -> ForwardBatch:
if device is None:
device = latents.device
embed_generator = torch.Generator(device=device).manual_seed(seed)
batch_generator = torch.Generator(device=generator_device).manual_seed(seed)
batch = ForwardBatch(data_type="test")
batch.latents = latents.clone()
batch.timesteps = timesteps.clone()
batch.num_inference_steps = len(timesteps)
batch.num_frames = latents.shape[2]
batch.height = latents.shape[-2]
batch.width = latents.shape[-1]
batch.raw_latent_shape = latents.shape
batch.prompt_embeds = [
torch.randn(1, 4, 8, generator=embed_generator, device=device)
]
batch.negative_prompt_embeds = [
torch.randn(1, 4, 8, generator=embed_generator, device=device)
]
batch.prompt_attention_mask = [
torch.ones(1, 4, dtype=torch.long, device=device)
]
batch.negative_attention_mask = [
torch.ones(1, 4, dtype=torch.long, device=device)
]
batch.guidance_scale = guidance_scale
batch.guidance_scale_2 = guidance_scale
batch.do_classifier_free_guidance = guidance_scale > 1.0
batch.guidance_rescale = 0.0
batch.generator = batch_generator
batch.image_embeds = []
batch.clip_embedding_pos = [
torch.randn(1, 4, 8, generator=embed_generator, device=device)
]
batch.clip_embedding_neg = [
torch.randn(1, 4, 8, generator=embed_generator, device=device)
]
return batch
def _run_manual(strategy, batch, args):
state = strategy.prepare(batch, args)
if isinstance(strategy, BlockDenoisingStrategy):
block_plan = strategy.block_plan(state)
for block_idx, block_item in enumerate(block_plan.items):
block_ctx = strategy.init_block_context(state, block_item, block_idx)
strategy.process_block(state, block_ctx, block_item)
strategy.update_context(state, block_ctx, block_item)
else:
for i, t in enumerate(state.timesteps):
model_inputs = strategy.make_model_inputs(state, t, i)
noise_pred = strategy.forward(state, model_inputs)
noise_pred = strategy.cfg_combine(state, noise_pred)
state.latents = strategy.scheduler_step(state, noise_pred, t)
return strategy.postprocess(state)
def _cuda_bf16_available() -> bool:
if not torch.cuda.is_available():
return False
is_bf16_supported = getattr(torch.cuda, "is_bf16_supported", None)
if is_bf16_supported is None:
return False
return torch.cuda.is_bf16_supported()
def test_standard_strategy_parity_cpu():
timesteps = torch.tensor([2, 1], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
scheduler = DummyScheduler(timesteps)
stage_a = DummyStandardStage(DummyTransformer(), scheduler)
stage_b = DummyStandardStage(DummyTransformer(), DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=0)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=0)
engine_out = DenoisingEngine(StandardStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(StandardStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_standard_strategy_parity_i2v_cpu():
timesteps = torch.tensor([2, 1], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
scheduler = DummyScheduler(timesteps)
stage_a = DummyStandardStage(DummyTransformer(), scheduler)
stage_b = DummyStandardStage(DummyTransformer(), DummyScheduler(timesteps))
stage_a.transformer.out_channels = latents.shape[1]
stage_b.transformer.out_channels = latents.shape[1]
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=5)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=5)
cond_gen = torch.Generator().manual_seed(123)
image_latent = torch.randn(latents.shape,
generator=cond_gen,
device=latents.device,
dtype=latents.dtype)
image_embeds = [
torch.randn(1,
4,
8,
generator=cond_gen,
device=latents.device,
dtype=latents.dtype)
]
batch_a.image_latent = image_latent
batch_b.image_latent = image_latent.clone()
batch_a.image_embeds = image_embeds
batch_b.image_embeds = [image_embeds[0].clone()]
engine_out = DenoisingEngine(StandardStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(StandardStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_standard_strategy_parity_v2v_cpu():
timesteps = torch.tensor([2, 1], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
scheduler = DummyScheduler(timesteps)
stage_a = DummyStandardStage(DummyTransformer(), scheduler)
stage_b = DummyStandardStage(DummyTransformer(), DummyScheduler(timesteps))
stage_a.transformer.out_channels = latents.shape[1]
stage_b.transformer.out_channels = latents.shape[1]
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=6)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=6)
cond_gen = torch.Generator().manual_seed(456)
video_latent = torch.randn(latents.shape,
generator=cond_gen,
device=latents.device,
dtype=latents.dtype)
batch_a.video_latent = video_latent
batch_b.video_latent = video_latent.clone()
engine_out = DenoisingEngine(StandardStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(StandardStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_standard_strategy_parity_boundary_cpu():
timesteps = torch.tensor([2, 1], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
args.pipeline_config.dit_config.boundary_ratio = 0.75
scheduler = DummyScheduler(timesteps)
stage_a = DummyStandardStage(DummyTransformer(), scheduler)
stage_b = DummyStandardStage(DummyTransformer(), DummyScheduler(timesteps))
stage_a.transformer_2 = DummyTransformer()
stage_b.transformer_2 = DummyTransformer()
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=7)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=7)
batch_a.guidance_scale_2 = 1.5
batch_b.guidance_scale_2 = 1.5
engine_out = DenoisingEngine(StandardStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(StandardStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_cosmos_strategy_parity_cpu():
timesteps = torch.tensor([2, 1], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
scheduler = DummyScheduler(timesteps)
stage_a = DummyStandardStage(DummyTransformer(), scheduler)
stage_b = DummyStandardStage(DummyTransformer(), DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=1)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=1)
engine_out = DenoisingEngine(CosmosStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(CosmosStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_dmd_strategy_parity_cpu():
timesteps = torch.tensor([4, 2], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
args.pipeline_config.dmd_denoising_steps = [4, 2]
scheduler = DummyScheduler(timesteps)
scheduler.timesteps = timesteps
scheduler.sigmas = scheduler._make_sigmas(timesteps)
stage_a = DummyDmdStage(DummyTransformer(), scheduler)
stage_b = DummyDmdStage(DummyTransformer(),
DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=1.0, seed=2)
batch_b = _make_batch(latents, timesteps, guidance_scale=1.0, seed=2)
engine_out = DenoisingEngine(DmdStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(DmdStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_dmd_strategy_parity_i2v_cpu():
timesteps = torch.tensor([4, 2], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
args.pipeline_config.dmd_denoising_steps = [4, 2]
scheduler = DummyScheduler(timesteps)
scheduler.timesteps = timesteps
scheduler.sigmas = scheduler._make_sigmas(timesteps)
stage_a = DummyDmdStage(DummyTransformer(), scheduler)
stage_b = DummyDmdStage(DummyTransformer(), DummyScheduler(timesteps))
stage_a.transformer.out_channels = latents.shape[1]
stage_b.transformer.out_channels = latents.shape[1]
batch_a = _make_batch(latents, timesteps, guidance_scale=1.0, seed=8)
batch_b = _make_batch(latents, timesteps, guidance_scale=1.0, seed=8)
cond_gen = torch.Generator().manual_seed(789)
image_latent = torch.randn(latents.shape,
generator=cond_gen,
device=latents.device,
dtype=latents.dtype)
image_embeds = [
torch.randn(1,
4,
8,
generator=cond_gen,
device=latents.device,
dtype=latents.dtype)
]
batch_a.image_latent = image_latent
batch_b.image_latent = image_latent.clone()
batch_a.image_embeds = image_embeds
batch_b.image_embeds = [image_embeds[0].clone()]
engine_out = DenoisingEngine(DmdStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(DmdStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_longcat_strategy_parity_cpu():
timesteps = torch.tensor([2, 1], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
scheduler = DummyScheduler(timesteps)
stage_a = DummyLongCatStage(DummyTransformer(), scheduler)
stage_b = DummyLongCatStage(DummyTransformer(), DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=3)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=3)
engine_out = DenoisingEngine(LongCatStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(LongCatStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_longcat_i2v_strategy_parity_cpu():
timesteps = torch.tensor([2, 1], dtype=torch.long)
latents = torch.randn(1, 2, 3, 2, 2)
args = _make_args()
scheduler = DummyScheduler(timesteps)
stage_a = DummyLongCatStage(DummyTransformer(), scheduler)
stage_b = DummyLongCatStage(DummyTransformer(), DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=4)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=4)
batch_a.num_cond_latents = 1
batch_b.num_cond_latents = 1
engine_out = DenoisingEngine(LongCatI2VStrategy(stage_a)).run(batch_a,
args)
manual_out = _run_manual(LongCatI2VStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_longcat_vc_strategy_parity_cpu():
timesteps = torch.tensor([2, 1], dtype=torch.long)
latents = torch.randn(1, 2, 3, 2, 2)
args = _make_args()
scheduler = DummyScheduler(timesteps)
stage_a = DummyLongCatStage(DummyTransformer(), scheduler)
stage_b = DummyLongCatStage(DummyTransformer(), DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=5)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=5)
batch_a.num_cond_latents = 1
batch_b.num_cond_latents = 1
batch_a.use_kv_cache = True
batch_b.use_kv_cache = True
batch_a.kv_cache_dict = {}
batch_b.kv_cache_dict = {}
batch_a.cond_latents = torch.randn(1, 2, 1, 2, 2)
batch_b.cond_latents = batch_a.cond_latents.clone()
engine_out = DenoisingEngine(LongCatVCStrategy(stage_a)).run(batch_a,
args)
manual_out = _run_manual(LongCatVCStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_causal_block_strategy_parity_cpu():
timesteps = torch.tensor([4, 2], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
args.pipeline_config.dmd_denoising_steps = [4, 2]
scheduler = DummyScheduler(timesteps)
scheduler.timesteps = timesteps
scheduler.sigmas = scheduler._make_sigmas(timesteps)
stage_a = DummyCausalStage(DummyTransformer(num_frames_per_block=1),
scheduler)
stage_b = DummyCausalStage(DummyTransformer(num_frames_per_block=1),
DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=1.0, seed=6)
batch_b = _make_batch(latents, timesteps, guidance_scale=1.0, seed=6)
engine_out = DenoisingEngine(CausalBlockStrategy(stage_a)).run(batch_a,
args)
manual_out = _run_manual(CausalBlockStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_matrixgame_block_strategy_parity_cpu():
timesteps = torch.tensor([3, 1], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
args.pipeline_config.dmd_denoising_steps = [3, 1]
scheduler = DummyScheduler(timesteps)
scheduler.timesteps = timesteps
scheduler.sigmas = scheduler._make_sigmas(timesteps)
stage_a = DummyMatrixGameStage(DummyTransformer(num_frames_per_block=1),
scheduler)
stage_b = DummyMatrixGameStage(DummyTransformer(num_frames_per_block=1),
DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=1.0, seed=7)
batch_b = _make_batch(latents, timesteps, guidance_scale=1.0, seed=7)
engine_out = DenoisingEngine(MatrixGameBlockStrategy(stage_a)).run(
batch_a, args)
manual_out = _run_manual(MatrixGameBlockStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
def test_matrixgame_block_strategy_streaming_parity_cpu():
timesteps = torch.tensor([3, 1], dtype=torch.long)
latents = torch.randn(1, 2, 2, 2, 2)
args = _make_args()
args.pipeline_config.dmd_denoising_steps = [3, 1]
scheduler = DummyScheduler(timesteps)
scheduler.timesteps = timesteps
scheduler.sigmas = scheduler._make_sigmas(timesteps)
stage_full = DummyMatrixGameStage(DummyTransformer(num_frames_per_block=1),
scheduler)
stage_stream = DummyMatrixGameStage(
DummyTransformer(num_frames_per_block=1), DummyScheduler(timesteps))
batch_full = _make_batch(latents, timesteps, guidance_scale=1.0, seed=10)
batch_stream = _make_batch(latents, timesteps, guidance_scale=1.0, seed=10)
batch_full.generator = None
batch_stream.generator = None
torch.manual_seed(1234)
noise_shape = (latents.shape[0], stage_full.num_frame_per_block,
latents.shape[1], latents.shape[3], latents.shape[4])
noise_pool = [
torch.randn(noise_shape, dtype=latents.dtype)
for _ in range(max(len(timesteps) - 1, 0))
]
strategy_full = MatrixGameBlockStrategy(stage_full)
state_full = strategy_full.prepare(batch_full, args)
state_full.extra["ctx"].noise_pool = [t.clone() for t in noise_pool]
block_plan_full = strategy_full.block_plan(state_full)
strategy_stream = MatrixGameBlockStrategy(stage_stream)
state_stream = strategy_stream.prepare(batch_stream, args)
state_stream.extra["ctx"].noise_pool = [t.clone() for t in noise_pool]
block_plan_stream = strategy_stream.block_plan(state_stream)
engine_full = DenoisingEngine(strategy_full)
engine_full.run_blocks(state_full, block_plan=block_plan_full)
engine_stream = DenoisingEngine(strategy_stream)
for block_idx in range(len(block_plan_stream.items)):
engine_stream.run_blocks(
state_stream,
block_plan=block_plan_stream,
start_block=block_idx,
num_blocks=1,
)
batch_full = strategy_full.postprocess(state_full)
batch_stream = strategy_stream.postprocess(state_stream)
assert torch.allclose(batch_full.latents, batch_stream.latents)
@pytest.mark.gpu
@pytest.mark.skipif(not _cuda_bf16_available(),
reason="CUDA BF16 not available")
def test_standard_strategy_parity_cuda():
device = torch.device("cuda")
timesteps = torch.tensor([2, 1], dtype=torch.long, device=device)
latents = torch.randn(1, 2, 2, 2, 2, device=device)
args = _make_args()
scheduler = DummyScheduler(timesteps)
stage_a = DummyStandardStage(DummyTransformer().to(device), scheduler)
stage_b = DummyStandardStage(DummyTransformer().to(device),
DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=2.0, seed=8)
batch_b = _make_batch(latents, timesteps, guidance_scale=2.0, seed=8)
engine_out = DenoisingEngine(StandardStrategy(stage_a)).run(batch_a, args)
manual_out = _run_manual(StandardStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
@pytest.mark.gpu
@pytest.mark.skipif(not _cuda_bf16_available(),
reason="CUDA BF16 not available")
def test_causal_block_strategy_parity_cuda():
device = torch.device("cuda")
timesteps = torch.tensor([4, 2], dtype=torch.long, device=device)
latents = torch.randn(1, 2, 2, 2, 2, device=device)
args = _make_args()
args.pipeline_config.dmd_denoising_steps = [4, 2]
scheduler = DummyScheduler(timesteps)
scheduler.timesteps = timesteps
scheduler.sigmas = scheduler._make_sigmas(timesteps)
stage_a = DummyCausalStage(
DummyTransformer(num_frames_per_block=1).to(device), scheduler)
stage_b = DummyCausalStage(
DummyTransformer(num_frames_per_block=1).to(device),
DummyScheduler(timesteps))
batch_a = _make_batch(latents, timesteps, guidance_scale=1.0, seed=9)
batch_b = _make_batch(latents, timesteps, guidance_scale=1.0, seed=9)
engine_out = DenoisingEngine(CausalBlockStrategy(stage_a)).run(batch_a,
args)
manual_out = _run_manual(CausalBlockStrategy(stage_b), batch_b, args)
assert torch.allclose(engine_out.latents, manual_out.latents)
@@ -0,0 +1,52 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.models.schedulers.adapter import DefaultSchedulerAdapter
class DummyScheduler:
def __init__(self) -> None:
self.calls: dict[str, tuple] = {}
def scale_model_input(self, latents: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
self.calls["scale_model_input"] = (latents, t)
return latents + 1
def step(self, noise_pred: torch.Tensor, t: torch.Tensor,
latents: torch.Tensor, **kwargs):
self.calls["step"] = (noise_pred, t, latents, kwargs)
return "step-result"
def add_noise(self, latents: torch.Tensor, noise: torch.Tensor,
t: torch.Tensor) -> torch.Tensor:
self.calls["add_noise"] = (latents, noise, t)
return latents + noise
def set_timesteps(self, num_steps: int, device=None, **kwargs):
self.calls["set_timesteps"] = (num_steps, device, kwargs)
return ["ok"]
def test_default_scheduler_adapter_delegates():
scheduler = DummyScheduler()
adapter = DefaultSchedulerAdapter(scheduler)
latents = torch.zeros(2, 2)
noise = torch.ones(2, 2)
t = torch.tensor(1)
scaled = adapter.scale_model_input(latents, t)
assert torch.allclose(scaled, latents + 1)
assert scheduler.calls["scale_model_input"] == (latents, t)
step_out = adapter.step(torch.ones(1), t, latents, foo="bar")
assert step_out == "step-result"
assert scheduler.calls["step"][3]["foo"] == "bar"
noised = adapter.add_noise(latents, noise, t)
assert torch.allclose(noised, latents + noise)
assert scheduler.calls["add_noise"] == (latents, noise, t)
set_out = adapter.set_timesteps(4, device=torch.device("cpu"), baz=3)
assert set_out == ["ok"]
assert scheduler.calls["set_timesteps"][0] == 4