Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0d1fe20fdb | ||
|
|
7f7431779e |
@@ -0,0 +1,40 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_lingbotworld_fast"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LingBot-World-Fast-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"The video presents a soaring journey through a fantasy jungle. The "
|
||||
"wind whips past the rider's blue hands gripping the reins, causing "
|
||||
"the leather straps to vibrate. The ancient gothic castle approaches "
|
||||
"steadily, its stone details becoming clearer against the backdrop of "
|
||||
"floating islands and distant waterfalls.")
|
||||
image_path = ("https://raw.githubusercontent.com/Robbyant/lingbot-world/"
|
||||
"main/examples/00/image.jpg")
|
||||
action_path = "examples/inference/basic/lingbotworld_examples/00"
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
image_path=image_path,
|
||||
action_path=action_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
num_frames=81,
|
||||
height=480,
|
||||
width=832,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -18,12 +18,13 @@ from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
|
||||
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
from fastvideo.configs.models.dits.lingbotworld2 import LingBotWorld2CausalFastVideoConfig
|
||||
from fastvideo.configs.models.dits.lingbotworld_fast import LingBotWorldFastVideoConfig
|
||||
from fastvideo.configs.models.dits.lingbot_video import LingBotVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
|
||||
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
|
||||
"MiniMaxH3Config", "ZImageDiTConfig", "MMAudioArchConfig", "MMAudioTransformerConfig"
|
||||
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotWorldFastVideoConfig",
|
||||
"LingBotVideoConfig", "MiniMaxH3Config", "ZImageDiTConfig", "MMAudioArchConfig", "MMAudioTransformerConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.configs.models.dits.lingbotworld2 import (
|
||||
LingBotWorld2CausalFastArchConfig, )
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorldFastArchConfig(LingBotWorld2CausalFastArchConfig):
|
||||
"""Arch config for the released LingBot-World-Fast checkpoint.
|
||||
|
||||
The tensor layout is identical to the LingBot World 2 causal-fast model, so
|
||||
the parent's shapes and ``param_names_mapping`` are reused verbatim. Only
|
||||
the sampling/attention-window values released with this checkpoint differ.
|
||||
"""
|
||||
|
||||
# The released `generate_fast.py` leaves `--local_attn_size` at -1, so
|
||||
# self-attention stays global and the KV cache never evicts. That also makes
|
||||
# `sink_size` unreachable, so it is deliberately not pinned here (the loader
|
||||
# overwrites it from the checkpoint config either way).
|
||||
local_attn_size: int = -1
|
||||
# `chunk_size` and `timesteps_index` are absent from the checkpoint config,
|
||||
# so these values are what the sampling loop actually runs with.
|
||||
chunk_size: int = 3
|
||||
timesteps_index: tuple[int, int, int, int] = (0, 179, 358, 679)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorldFastVideoConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=LingBotWorldFastArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
@@ -8,6 +8,7 @@ from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelin
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5DMDConfig, Kandinsky5I2VConfig, Kandinsky5T2VConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld2 import LingBotWorld2CausalFastI2V480PConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld_fast import LingBotWorldFastI2V480PConfig
|
||||
from fastvideo.configs.pipelines.lingbot_video import LingBotVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
|
||||
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
|
||||
@@ -22,6 +23,7 @@ __all__ = [
|
||||
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
|
||||
"Kandinsky5I2VConfig", "Kandinsky5DMDConfig", "LingBotWorld2CausalFastI2V480PConfig", "LingBotVideoT2VConfig",
|
||||
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "MMAudioV2AConfig", "get_pipeline_config_cls_from_name"
|
||||
"Kandinsky5I2VConfig", "Kandinsky5DMDConfig", "LingBotWorld2CausalFastI2V480PConfig",
|
||||
"LingBotWorldFastI2V480PConfig", "LingBotVideoT2VConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig",
|
||||
"MMAudioV2AConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig
|
||||
from fastvideo.configs.models.dits.lingbotworld_fast import (
|
||||
LingBotWorldFastVideoConfig, )
|
||||
from fastvideo.configs.models.encoders import T5Config
|
||||
from fastvideo.configs.pipelines.lingbotworld2 import (
|
||||
LingBotWorld2CausalFastI2V480PConfig, )
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorldFastI2V480PConfig(LingBotWorld2CausalFastI2V480PConfig):
|
||||
"""Pipeline config for LingBot-World-Fast 480P image-to-video.
|
||||
|
||||
Shares the LingBot World 2 causal-fast sampling loop and Wan VAE, but this
|
||||
checkpoint ships the stock ``UMT5EncoderModel`` text encoder (``d_model``
|
||||
fields) rather than LingBot World 2's custom one (``dim`` fields), so the
|
||||
standard ``T5Config`` is restored here alongside this DiT's arch config.
|
||||
"""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=LingBotWorldFastVideoConfig)
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
@@ -389,7 +389,11 @@ class WanCrossAttention(CausalWanSelfAttention):
|
||||
else:
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
x = flash_attention(q, k, v, k_lens=context_lens)
|
||||
# Dispatch through `attention` so the TORCH_SDPA backend this model
|
||||
# advertises actually works; self-attention already does the same. The
|
||||
# SDPA path ignores `k_lens`, which is safe here because the caller
|
||||
# always passes `context_lens=None`.
|
||||
x = attention(q, k, v, k_lens=context_lens)
|
||||
return self.o(x.flatten(2))
|
||||
|
||||
|
||||
|
||||
@@ -61,6 +61,13 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
"lingbotworld2",
|
||||
"LingBotWorld2CausalFastTransformer3DModel",
|
||||
),
|
||||
# LingBot-World-Fast ships this class name; its tensor layout is identical
|
||||
# to the LingBot World 2 causal-fast DiT.
|
||||
"CausalLingBotWorldTransformer3DModel": (
|
||||
"dits",
|
||||
"lingbotworld2",
|
||||
"LingBotWorld2CausalFastTransformer3DModel",
|
||||
),
|
||||
"MatrixGame2WanModel": ("dits", "matrixgame2", "MatrixGame2WanModel"),
|
||||
"CausalMatrixGame2WanModel": ("dits", "matrixgame2", "CausalMatrixGame2WanModel"),
|
||||
# Legacy aliases for older HF model_index.json files
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from .fast_pipeline import LingBotWorldFastPipeline
|
||||
|
||||
__all__ = ["LingBotWorldFastPipeline"]
|
||||
|
||||
EntryClass = LingBotWorldFastPipeline
|
||||
@@ -0,0 +1,38 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LingBot-World-Fast causal image-to-video pipeline."""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler, )
|
||||
from fastvideo.pipelines.basic.lingbotworld2.causal_fast_pipeline import (
|
||||
LingBotWorld2CausalFastPipeline, )
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LingBotWorldFastPipeline(LingBotWorld2CausalFastPipeline):
|
||||
"""LingBot-World-Fast I2V generation.
|
||||
|
||||
The released checkpoint uses the same chunked causal sampling loop as
|
||||
LingBot World 2 causal-fast; the chunk size, timestep indices, and
|
||||
attention window come from ``LingBotWorldFastArchConfig``.
|
||||
"""
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Install the flow-matching scheduler the released model samples with.
|
||||
|
||||
This checkpoint ships a stock ``UniPCMultistepScheduler``, whose
|
||||
``set_timesteps`` takes no ``shift``. The reference ``generate_fast.py``
|
||||
builds a ``FlowUniPCMultistepScheduler`` instead, so replace the loaded
|
||||
one unconditionally rather than only when absent.
|
||||
"""
|
||||
arch_config = fastvideo_args.pipeline_config.dit_config.arch_config
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
num_train_timesteps=arch_config.num_train_timesteps,
|
||||
shift=1,
|
||||
use_dynamic_shifting=False,
|
||||
)
|
||||
|
||||
|
||||
EntryClass = LingBotWorldFastPipeline
|
||||
@@ -0,0 +1,35 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LingBot-World-Fast pipeline preset."""
|
||||
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Causal-fast denoising pass",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
LINGBOTWORLD_FAST_I2V = InferencePreset(
|
||||
name="lingbotworld_fast_i2v",
|
||||
version=1,
|
||||
model_family="lingbotworld_fast",
|
||||
description="LingBot-World-Fast 14B causal I2V",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 4,
|
||||
"fps": 16,
|
||||
"seed": 42,
|
||||
"num_frames": 81,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"negative_prompt": "",
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (LINGBOTWORLD_FAST_I2V, )
|
||||
+28
-2
@@ -32,6 +32,7 @@ from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsk
|
||||
from fastvideo.configs.pipelines.lingbot_video import LingBotVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld2 import LingBotWorld2CausalFastI2V480PConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld_fast import LingBotWorldFastI2V480PConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.flux_2 import (
|
||||
@@ -526,6 +527,27 @@ def _register_configs() -> None:
|
||||
default_preset="lingbotworld2_causal_fast_i2v",
|
||||
)
|
||||
|
||||
# LingBotWorld-Fast — registered BEFORE the LingBotWorld base entry so its
|
||||
# detector wins when both fire. The base detector only excludes the
|
||||
# "causal-fast" spelling, so it also matches this checkpoint's path and its
|
||||
# `lingbotworldcausaldmdpipeline` model_index name; the detector loop keeps
|
||||
# the first match, which must be this more specific entry.
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=LingBotWorldFastI2V480PConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/LingBot-World-Fast-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: ("lingbot-world-fast" in path.lower() or "lingbotworldfast" in path.lower() or
|
||||
"lingbotworldcausaldmdpipeline" in path.lower())
|
||||
],
|
||||
model_family="lingbotworld_fast",
|
||||
default_preset="lingbotworld_fast_i2v",
|
||||
pipeline_cls_name="LingBotWorldFastPipeline",
|
||||
)
|
||||
|
||||
# LingBotWorld
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
@@ -535,8 +557,9 @@ def _register_configs() -> None:
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: (("lingbotworld" in path.lower() or "lingbot-world" in path.lower()) and "causal-fast" not in
|
||||
path.lower() and "causalfast" not in path.lower())
|
||||
lambda path:
|
||||
(("lingbotworld" in path.lower() or "lingbot-world" in path.lower()) and "causal-fast" not in path.lower()
|
||||
and "causalfast" not in path.lower() and "-fast" not in path.lower() and "causaldmd" not in path.lower())
|
||||
],
|
||||
model_family="lingbotworld",
|
||||
default_preset="lingbotworld_i2v",
|
||||
@@ -1307,6 +1330,8 @@ def _register_presets() -> None:
|
||||
ALL_PRESETS as LINGBOTWORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.lingbotworld2.presets import (
|
||||
ALL_PRESETS as LINGBOTWORLD2_PRESETS, )
|
||||
from fastvideo.pipelines.basic.lingbotworld_fast.presets import (
|
||||
ALL_PRESETS as LINGBOTWORLD_FAST_PRESETS, )
|
||||
from fastvideo.pipelines.basic.lingbot_video.presets import (
|
||||
ALL_PRESETS as LINGBOT_VIDEO_PRESETS, )
|
||||
from fastvideo.pipelines.basic.longcat.presets import (
|
||||
@@ -1347,6 +1372,7 @@ def _register_presets() -> None:
|
||||
LINGBOT_VIDEO_PRESETS,
|
||||
LINGBOTWORLD_PRESETS,
|
||||
LINGBOTWORLD2_PRESETS,
|
||||
LINGBOTWORLD_FAST_PRESETS,
|
||||
LONGCAT_PRESETS,
|
||||
LTX2_PRESETS,
|
||||
MATRIXGAME2_PRESETS,
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""SSIM-based similarity test for LingBot-World-Fast causal I2V.
|
||||
|
||||
The camera trajectory is read straight from the LingBot example npy files
|
||||
(poses.npy/intrinsics.npy) by the pipeline itself, matching the official
|
||||
`generate_fast.py` workflow.
|
||||
|
||||
Note: this checkpoint is 4-step distilled, so the step count is fixed by the
|
||||
model rather than reduced for CI.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
build_generated_output_dir,
|
||||
build_reference_folder_path,
|
||||
get_cuda_device_name,
|
||||
resolve_device_reference_folder,
|
||||
select_ssim_params,
|
||||
)
|
||||
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 2
|
||||
|
||||
# The released checkpoint is 4-step distilled; see LingBotWorldFastArchConfig.
|
||||
NUM_DISTILLED_STEPS = 4
|
||||
|
||||
|
||||
def _find_lingbotworld_examples_root() -> str | None:
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
repo_root = os.path.abspath(os.path.join(script_dir, "..", "..", ".."))
|
||||
candidate = os.path.join(repo_root, "examples", "inference", "basic",
|
||||
"lingbotworld_examples")
|
||||
if (os.path.exists(os.path.join(candidate, "00", "poses.npy"))
|
||||
and os.path.exists(os.path.join(candidate, "00",
|
||||
"intrinsics.npy"))):
|
||||
return os.path.abspath(candidate)
|
||||
return None
|
||||
|
||||
|
||||
device_name = get_cuda_device_name()
|
||||
device_reference_folder = resolve_device_reference_folder(
|
||||
(
|
||||
("A40", "A40"),
|
||||
("L40S", "L40S"),
|
||||
("H100", "H100"),
|
||||
("H200", "H200"),
|
||||
),
|
||||
device_name=device_name,
|
||||
logger=logger,
|
||||
)
|
||||
|
||||
LINGBOT_FAST_PARAMS = {
|
||||
"model_path": "FastVideo/LingBot-World-Fast-Diffusers",
|
||||
"num_gpus": 2,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 25, # must be 4k+1; trimmed to a whole number of chunks
|
||||
"seed": 42,
|
||||
"fps": 16,
|
||||
"example_case": "00",
|
||||
"image_path": ("https://raw.githubusercontent.com/Robbyant/lingbot-world/"
|
||||
"main/examples/00/image.jpg"),
|
||||
}
|
||||
LINGBOT_FAST_FULL_QUALITY_PARAMS = {
|
||||
**LINGBOT_FAST_PARAMS,
|
||||
"num_frames": 81,
|
||||
}
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"The video presents a soaring journey through a fantasy jungle. The wind "
|
||||
"whips past the rider's blue hands gripping the reins, causing the leather "
|
||||
"straps to vibrate. The ancient gothic castle approaches steadily, its stone "
|
||||
"details becoming clearer against the backdrop of floating islands and "
|
||||
"distant waterfalls.",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
|
||||
def test_lingbot_fast_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
|
||||
|
||||
params = select_ssim_params(LINGBOT_FAST_PARAMS,
|
||||
LINGBOT_FAST_FULL_QUALITY_PARAMS)
|
||||
|
||||
if device_reference_folder is None:
|
||||
pytest.skip(
|
||||
f"Unsupported device for LingBot-World-Fast SSIM test: {device_name}"
|
||||
)
|
||||
if torch.cuda.device_count() < params["num_gpus"]:
|
||||
pytest.skip(
|
||||
f"LingBot-World-Fast SSIM test requires {params['num_gpus']} GPUs, "
|
||||
f"but only {torch.cuda.device_count()} detected.")
|
||||
|
||||
examples_root = _find_lingbotworld_examples_root()
|
||||
if examples_root is None:
|
||||
pytest.skip(
|
||||
"lingbotworld_examples not found under examples/inference/basic.")
|
||||
|
||||
action_path = os.path.join(examples_root, params["example_case"])
|
||||
|
||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
model_id = "LingBot-World-Fast-Diffusers"
|
||||
output_dir = build_generated_output_dir(
|
||||
script_dir,
|
||||
device_reference_folder,
|
||||
model_id,
|
||||
ATTENTION_BACKEND,
|
||||
)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": params["num_gpus"],
|
||||
"use_fsdp_inference": True,
|
||||
"dit_cpu_offload": True,
|
||||
"dit_layerwise_offload": False,
|
||||
"text_encoder_cpu_offload": True,
|
||||
"vae_cpu_offload": False,
|
||||
"pin_cpu_memory": True,
|
||||
}
|
||||
generation_kwargs = {
|
||||
"output_path": output_dir,
|
||||
"image_path": params["image_path"],
|
||||
"action_path": action_path,
|
||||
"height": params["height"],
|
||||
"width": params["width"],
|
||||
"num_frames": params["num_frames"],
|
||||
"seed": params["seed"],
|
||||
"fps": params["fps"],
|
||||
}
|
||||
|
||||
generator: VideoGenerator | None = None
|
||||
try:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=params["model_path"], **init_kwargs)
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
finally:
|
||||
if generator is not None:
|
||||
generator.shutdown()
|
||||
|
||||
generated_video_path = os.path.join(output_dir, output_video_name)
|
||||
assert os.path.exists(generated_video_path), (
|
||||
f"Output video was not generated at {generated_video_path}")
|
||||
|
||||
reference_folder = build_reference_folder_path(
|
||||
script_dir,
|
||||
device_reference_folder,
|
||||
model_id,
|
||||
ATTENTION_BACKEND,
|
||||
)
|
||||
if not os.path.exists(reference_folder):
|
||||
raise FileNotFoundError(
|
||||
f"Reference video folder does not exist: {reference_folder}")
|
||||
|
||||
reference_video_name = None
|
||||
for filename in os.listdir(reference_folder):
|
||||
if filename.endswith(".mp4") and prompt[:100].strip() in filename:
|
||||
reference_video_name = filename
|
||||
break
|
||||
if not reference_video_name:
|
||||
raise FileNotFoundError(
|
||||
f"Reference video missing for prompt/backend under {reference_folder}"
|
||||
)
|
||||
|
||||
reference_video_path = os.path.join(reference_folder, reference_video_name)
|
||||
logger.info("Computing SSIM between %s and %s", reference_video_path,
|
||||
generated_video_path)
|
||||
ssim_values = compute_video_ssim_torchvision(reference_video_path,
|
||||
generated_video_path,
|
||||
use_ms_ssim=True)
|
||||
mean_ssim = ssim_values[0]
|
||||
logger.info("SSIM mean value: %s", mean_ssim)
|
||||
|
||||
write_ssim_results(output_dir, ssim_values, reference_video_path,
|
||||
generated_video_path, NUM_DISTILLED_STEPS, prompt)
|
||||
|
||||
min_acceptable_ssim = 0.70
|
||||
assert mean_ssim >= min_acceptable_ssim, (
|
||||
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
|
||||
f"for {model_id} with backend {ATTENTION_BACKEND}")
|
||||
Reference in New Issue
Block a user