Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a8a1d39de3 | ||
|
|
b4edcceb90 | ||
|
|
464f6abff5 | ||
|
|
def36d9973 | ||
|
|
5939883d31 |
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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")))
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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")
|
||||
|
||||
-11
@@ -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
|
||||
Reference in New Issue
Block a user