Compare commits
57
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0a861031c7 | ||
|
|
f08c5ee8af | ||
|
|
755f4a7967 | ||
|
|
d543a67b10 | ||
|
|
741aa8d289 | ||
|
|
ad9cd63122 | ||
|
|
68e6ffca9e | ||
|
|
64cdcf6be4 | ||
|
|
aa0d98a6b8 | ||
|
|
93b03bc14d | ||
|
|
160f0c9ccf | ||
|
|
e5d1110a0f | ||
|
|
622217ff2a | ||
|
|
cbab605eff | ||
|
|
ac98869aa1 | ||
|
|
56d4a6074f | ||
|
|
9713ea1275 | ||
|
|
37aa382cce | ||
|
|
3f00983287 | ||
|
|
2dc57f4070 | ||
|
|
9df19be719 | ||
|
|
089eea3970 | ||
|
|
dd8447ecc5 | ||
|
|
dca423fd31 | ||
|
|
1b43af8e8e | ||
|
|
628591b620 | ||
|
|
0980ca563f | ||
|
|
b158388733 | ||
|
|
c4ad4227c0 | ||
|
|
a63ccce73d | ||
|
|
ac56806aff | ||
|
|
0462e1b0e7 | ||
|
|
907f2100ec | ||
|
|
e0a3db5651 | ||
|
|
fca45bc8e1 | ||
|
|
aa95a4c18e | ||
|
|
942f7db3db | ||
|
|
aadb23f409 | ||
|
|
f56f567042 | ||
|
|
15a164a052 | ||
|
|
aaaa7a14a3 | ||
|
|
86d639c848 | ||
|
|
00338aa9ca | ||
|
|
74b409d7cf | ||
|
|
8537dcd6de | ||
|
|
528cef02c4 | ||
|
|
8208536cd1 | ||
|
|
0653f8f3af | ||
|
|
e0d702decb | ||
|
|
541ef014ee | ||
|
|
ffc1a7a58b | ||
|
|
9028953625 | ||
|
|
6eb95693a1 | ||
|
|
c3567eb468 | ||
|
|
3c3da4d057 | ||
|
|
0399713e7b | ||
|
|
0af2e9e8ef |
@@ -10,9 +10,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Clone a reference repo for FastVideo parity tests."
|
||||
)
|
||||
parser = argparse.ArgumentParser(description="Clone a reference repo for FastVideo parity tests.")
|
||||
parser.add_argument("repo_url", help="Official reference repository URL")
|
||||
parser.add_argument("target_dir", help="Directory to clone into")
|
||||
parser.add_argument("--branch", help="Branch or tag to clone")
|
||||
@@ -62,9 +60,7 @@ def gitignore_entry_for(target: Path) -> str:
|
||||
try:
|
||||
relative = resolved.relative_to(root)
|
||||
except ValueError as exc:
|
||||
raise ValueError(
|
||||
"--update-gitignore requires target_dir to be under the current directory"
|
||||
) from exc
|
||||
raise ValueError("--update-gitignore requires target_dir to be under the current directory") from exc
|
||||
|
||||
text = relative.as_posix().rstrip("/")
|
||||
return "/" + text + "/"
|
||||
|
||||
@@ -8,14 +8,12 @@ import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Download a HF model snapshot or selected files into a local directory."
|
||||
)
|
||||
description="Download a HF model snapshot or selected files into a local directory.")
|
||||
parser.add_argument("repo_id", help="HF repo id, for example Org/Model")
|
||||
parser.add_argument("local_dir", help="Destination directory")
|
||||
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
|
||||
|
||||
@@ -10,7 +10,6 @@ import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
|
||||
RAW_WEIGHT_SUFFIXES = (".safetensors", ".pt", ".pth", ".ckpt", ".bin")
|
||||
KNOWN_COMPONENTS = {
|
||||
@@ -34,8 +33,7 @@ KNOWN_COMPONENTS = {
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown."
|
||||
)
|
||||
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown.")
|
||||
parser.add_argument("source", help="HF repo id or local weights directory")
|
||||
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
|
||||
parser.add_argument("--revision", help="HF revision to inspect")
|
||||
@@ -94,14 +92,12 @@ def load_remote_files(
|
||||
) -> list[str]:
|
||||
from huggingface_hub import list_repo_files
|
||||
|
||||
return sorted(
|
||||
list_repo_files(
|
||||
repo_id,
|
||||
repo_type=repo_type,
|
||||
revision=revision,
|
||||
token=token,
|
||||
)
|
||||
)
|
||||
return sorted(list_repo_files(
|
||||
repo_id,
|
||||
repo_type=repo_type,
|
||||
revision=revision,
|
||||
token=token,
|
||||
))
|
||||
|
||||
|
||||
def load_remote_model_index(
|
||||
@@ -215,24 +211,24 @@ def build_result(args: argparse.Namespace) -> dict[str, Any]:
|
||||
"components_seen": components,
|
||||
"file_count": len(files),
|
||||
"file_scan_truncated": truncated,
|
||||
"files_sample": files[: args.sample_limit],
|
||||
"files_sample": files[:args.sample_limit],
|
||||
}
|
||||
|
||||
|
||||
def print_human(result: dict[str, Any]) -> None:
|
||||
for key in (
|
||||
"source",
|
||||
"source_kind",
|
||||
"repo_type",
|
||||
"revision",
|
||||
"token_env",
|
||||
"source_layout",
|
||||
"needs_conversion",
|
||||
"model_index_class",
|
||||
"model_index_diffusers_version",
|
||||
"model_index_error",
|
||||
"file_count",
|
||||
"file_scan_truncated",
|
||||
"source",
|
||||
"source_kind",
|
||||
"repo_type",
|
||||
"revision",
|
||||
"token_env",
|
||||
"source_layout",
|
||||
"needs_conversion",
|
||||
"model_index_class",
|
||||
"model_index_diffusers_version",
|
||||
"model_index_error",
|
||||
"file_count",
|
||||
"file_scan_truncated",
|
||||
):
|
||||
value = result.get(key)
|
||||
if value is not None:
|
||||
|
||||
@@ -18,7 +18,6 @@ import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
@@ -35,15 +34,10 @@ FASTVIDEO_CONFIG_CLASS = "<FastVideoConfig>" # TODO.
|
||||
FASTVIDEO_MODEL_MODULE = "fastvideo.models.<bucket>.<module>" # TODO.
|
||||
FASTVIDEO_MODEL_CLASS = "<FastVideoModel>" # TODO.
|
||||
|
||||
OFFICIAL_REF_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>")
|
||||
)
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY)
|
||||
)
|
||||
CONVERTED_WEIGHTS_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / FAMILY)
|
||||
)
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>"))
|
||||
LOCAL_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY))
|
||||
CONVERTED_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR",
|
||||
REPO_ROOT / "converted_weights" / FAMILY))
|
||||
|
||||
|
||||
def _resolve_hf_token() -> str | None:
|
||||
@@ -99,18 +93,14 @@ def _load_official_model(device: torch.device, dtype: torch.dtype) -> torch.nn.M
|
||||
model = OfficialClass() # TODO: pass official config kwargs.
|
||||
state_dict = {} # TODO: load official state dict from LOCAL_WEIGHTS_DIR.
|
||||
missing, unexpected = model.load_state_dict(state_dict, strict=True)
|
||||
assert not missing and not unexpected, (
|
||||
f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
|
||||
)
|
||||
assert not missing and not unexpected, (f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
"""Load the FastVideo component with the same tensor content."""
|
||||
if not CONVERTED_WEIGHTS_DIR.exists() and not LOCAL_WEIGHTS_DIR.exists():
|
||||
pytest.skip(
|
||||
f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}"
|
||||
)
|
||||
pytest.skip(f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}")
|
||||
|
||||
# TODO: replace with the bucket-specific FastVideo config/class/loader.
|
||||
# DiT examples:
|
||||
@@ -127,8 +117,7 @@ def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.
|
||||
state_dict = {} # TODO: load converted or directly mapped state dict.
|
||||
missing, unexpected = model.load_state_dict(state_dict, strict=True)
|
||||
assert not missing and not unexpected, (
|
||||
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
|
||||
)
|
||||
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
@@ -187,11 +176,9 @@ def test_component_parity():
|
||||
|
||||
assert official_out.shape == fastvideo_out.shape
|
||||
diff = (official_out - fastvideo_out).abs()
|
||||
print(
|
||||
f"official abs_mean={official_out.abs().mean().item():.6f} "
|
||||
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
|
||||
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}"
|
||||
)
|
||||
print(f"official abs_mean={official_out.abs().mean().item():.6f} "
|
||||
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
|
||||
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
|
||||
|
||||
# TODO: pick tolerance by scope:
|
||||
# - single block / same kernel: 1e-4
|
||||
|
||||
@@ -27,7 +27,6 @@ try:
|
||||
except ImportError: # pragma: no cover - optional local conversion dependency
|
||||
snapshot_download = None
|
||||
|
||||
|
||||
# TODO: fill with authoritative component prefixes for monolithic checkpoints.
|
||||
# Example: {"model.model.": "transformer", "pretransform.model.": "vae"}
|
||||
COMPONENT_PREFIXES: dict[str, str] = {}
|
||||
@@ -47,10 +46,7 @@ SKIP_PATTERNS: tuple[str, ...] = ()
|
||||
|
||||
|
||||
def _hf_token() -> str | None:
|
||||
return (
|
||||
os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN")
|
||||
or os.environ.get("HF_API_KEY")
|
||||
)
|
||||
return (os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN") or os.environ.get("HF_API_KEY"))
|
||||
|
||||
|
||||
def resolve_src(src: str, revision: str | None) -> Path:
|
||||
@@ -95,11 +91,10 @@ def apply_mapping(key: str) -> str | None:
|
||||
return key
|
||||
|
||||
|
||||
def split_monolithic(
|
||||
state: dict[str, torch.Tensor],
|
||||
) -> dict[str, OrderedDict[str, torch.Tensor]]:
|
||||
def split_monolithic(state: dict[str, torch.Tensor], ) -> dict[str, OrderedDict[str, torch.Tensor]]:
|
||||
components: dict[str, OrderedDict[str, torch.Tensor]] = {
|
||||
name: OrderedDict() for name in set(COMPONENT_PREFIXES.values())
|
||||
name: OrderedDict()
|
||||
for name in set(COMPONENT_PREFIXES.values())
|
||||
}
|
||||
intentionally_skipped: list[str] = []
|
||||
unowned: list[str] = []
|
||||
@@ -117,10 +112,8 @@ def split_monolithic(
|
||||
unowned.append(key)
|
||||
if unowned:
|
||||
sample = ", ".join(unowned[:10])
|
||||
raise ValueError(
|
||||
f"Unowned monolithic keys: {len(unowned)}. "
|
||||
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}"
|
||||
)
|
||||
raise ValueError(f"Unowned monolithic keys: {len(unowned)}. "
|
||||
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}")
|
||||
if intentionally_skipped:
|
||||
print(f"Intentionally skipped {len(intentionally_skipped)} keys")
|
||||
return {name: weights for name, weights in components.items() if weights}
|
||||
@@ -143,8 +136,12 @@ def build_component_configs(_src_dir: Path) -> dict[str, dict[str, Any]]:
|
||||
# TODO: emit config content accepted by FastVideo loaders. Most components use
|
||||
# config.json; schedulers use scheduler_config.json.
|
||||
return {
|
||||
"transformer": {"_class_name": "<FastVideoTransformerClass>"},
|
||||
"vae": {"_class_name": "<FastVideoVAEClass>"},
|
||||
"transformer": {
|
||||
"_class_name": "<FastVideoTransformerClass>"
|
||||
},
|
||||
"vae": {
|
||||
"_class_name": "<FastVideoVAEClass>"
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -177,19 +174,13 @@ def build_model_index(
|
||||
}
|
||||
if revision:
|
||||
index["_fastvideo_converted_revision"] = revision
|
||||
return {
|
||||
key: value
|
||||
for key, value in index.items()
|
||||
if key.startswith("_") or key in available_components
|
||||
}
|
||||
return {key: value for key, value in index.items() if key.startswith("_") or key in available_components}
|
||||
|
||||
|
||||
def validate_component_configs(configs: dict[str, dict[str, Any]]) -> None:
|
||||
# TODO: instantiate each FastVideo config and call update_model_arch(...) or
|
||||
# update_model_config(...) with this JSON so unknown emitted keys fail here.
|
||||
placeholder_configs = [
|
||||
name for name, config in configs.items() if "<" in json.dumps(config)
|
||||
]
|
||||
placeholder_configs = [name for name, config in configs.items() if "<" in json.dumps(config)]
|
||||
if placeholder_configs:
|
||||
raise ValueError(f"Replace config placeholders for: {placeholder_configs}")
|
||||
|
||||
@@ -201,9 +192,7 @@ def verify_conversion(
|
||||
del dst_dir, components
|
||||
# TODO: load each emitted stateful component through its production loader and
|
||||
# assert strict load, or document exact allowed missing/unexpected keys.
|
||||
raise NotImplementedError(
|
||||
"Implement production config validation and strict-load checks"
|
||||
)
|
||||
raise NotImplementedError("Implement production config validation and strict-load checks")
|
||||
|
||||
|
||||
def write_component(
|
||||
@@ -216,9 +205,7 @@ def write_component(
|
||||
if component_dir.exists() and any(component_dir.iterdir()):
|
||||
shutil.rmtree(component_dir)
|
||||
component_dir.mkdir(parents=True, exist_ok=True)
|
||||
save_file(
|
||||
dict(state), str(component_dir / "diffusion_pytorch_model.safetensors")
|
||||
)
|
||||
save_file(dict(state), str(component_dir / "diffusion_pytorch_model.safetensors"))
|
||||
if config is not None:
|
||||
config_path = component_dir / config_filename(name)
|
||||
with config_path.open("w", encoding="utf-8") as f:
|
||||
@@ -261,9 +248,7 @@ def convert(
|
||||
|
||||
if layout in {"monolithic", "raw_official"}:
|
||||
# TODO: replace model.safetensors with the official monolithic file name.
|
||||
components = split_monolithic(
|
||||
load_checkpoint(default_monolithic_checkpoint(src_path))
|
||||
)
|
||||
components = split_monolithic(load_checkpoint(default_monolithic_checkpoint(src_path)))
|
||||
elif layout in {"separate_components", "mixed"}:
|
||||
if not src_path.is_dir():
|
||||
raise ValueError(f"{layout} layout requires a source directory: {src_path}")
|
||||
@@ -271,9 +256,7 @@ def convert(
|
||||
else:
|
||||
raise ValueError(f"Unsupported template layout: {layout}")
|
||||
|
||||
copied = (
|
||||
copy_passthrough(src_path, dst_dir) if src_path.is_dir() else []
|
||||
)
|
||||
copied = (copy_passthrough(src_path, dst_dir) if src_path.is_dir() else [])
|
||||
configs = build_component_configs(src_path if src_path.is_dir() else src_path.parent)
|
||||
validate_component_configs(configs)
|
||||
for name, state in components.items():
|
||||
@@ -289,9 +272,7 @@ def convert(
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--src", required=True, help="HF repo id, local dir, or checkpoint path"
|
||||
)
|
||||
parser.add_argument("--src", required=True, help="HF repo id, local dir, or checkpoint path")
|
||||
parser.add_argument("--revision", help="HF branch, tag, or commit for repo sources")
|
||||
parser.add_argument(
|
||||
"--dst",
|
||||
|
||||
@@ -24,8 +24,8 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
|
||||
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
|
||||
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
|
||||
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
|
||||
DRILL_LAYER_ENV: str = "<FAMILY>_DEBUG_DRILL_LAYER"
|
||||
HYPOTHESIS_ENV: str = "<FAMILY>_DEBUG_PATCH_<HYPOTHESIS>"
|
||||
REL_THRESHOLD: float = 0.005 # 0.5% abs_mean drift flags a block as divergent
|
||||
@@ -94,6 +94,7 @@ def _attach_block_hooks(
|
||||
handles: list[Any] = []
|
||||
|
||||
def _hook(name: str):
|
||||
|
||||
def fn(_module, _inputs, outputs):
|
||||
t = outputs[0] if isinstance(outputs, tuple) else outputs
|
||||
if not torch.is_tensor(t):
|
||||
@@ -101,6 +102,7 @@ def _attach_block_hooks(
|
||||
log.append({"side": label, **_stat(name, t)})
|
||||
if tensors is not None:
|
||||
tensors[name] = t.detach().float().cpu()
|
||||
|
||||
return fn
|
||||
|
||||
def _pre_hook(name: str):
|
||||
@@ -114,6 +116,7 @@ def _attach_block_hooks(
|
||||
log.append({"side": label, **_stat(key, t)})
|
||||
if tensors is not None:
|
||||
tensors[key] = t.detach().float().cpu()
|
||||
|
||||
return fn
|
||||
|
||||
# TODO: adapt attribute paths to your model. Remove adapter block if absent.
|
||||
@@ -131,43 +134,21 @@ def _attach_block_hooks(
|
||||
# magi-human uses: attention, mlp.pre_norm, mlp.up_gate_proj,
|
||||
# mlp.down_proj (pre+post), mlp, attn_post_norm, mlp_post_norm.
|
||||
if hasattr(layer, "attention"):
|
||||
handles.append(
|
||||
layer.attention.register_forward_hook(_hook(f"{tag}.attention"))
|
||||
)
|
||||
handles.append(layer.attention.register_forward_hook(_hook(f"{tag}.attention")))
|
||||
if hasattr(layer, "mlp"):
|
||||
mlp = layer.mlp
|
||||
if hasattr(mlp, "pre_norm"):
|
||||
handles.append(
|
||||
mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm"))
|
||||
)
|
||||
handles.append(mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm")))
|
||||
if hasattr(mlp, "up_gate_proj"):
|
||||
handles.append(
|
||||
mlp.up_gate_proj.register_forward_hook(
|
||||
_hook(f"{tag}.mlp.up_gate_proj")
|
||||
)
|
||||
)
|
||||
handles.append(mlp.up_gate_proj.register_forward_hook(_hook(f"{tag}.mlp.up_gate_proj")))
|
||||
if hasattr(mlp, "down_proj"):
|
||||
handles.append(
|
||||
mlp.down_proj.register_forward_pre_hook(
|
||||
_pre_hook(f"{tag}.mlp.down_proj")
|
||||
)
|
||||
)
|
||||
handles.append(
|
||||
mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj"))
|
||||
)
|
||||
handles.append(mlp.down_proj.register_forward_pre_hook(_pre_hook(f"{tag}.mlp.down_proj")))
|
||||
handles.append(mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj")))
|
||||
handles.append(mlp.register_forward_hook(_hook(f"{tag}.mlp")))
|
||||
if hasattr(layer, "attn_post_norm"):
|
||||
handles.append(
|
||||
layer.attn_post_norm.register_forward_hook(
|
||||
_hook(f"{tag}.attn_post_norm")
|
||||
)
|
||||
)
|
||||
handles.append(layer.attn_post_norm.register_forward_hook(_hook(f"{tag}.attn_post_norm")))
|
||||
if hasattr(layer, "mlp_post_norm"):
|
||||
handles.append(
|
||||
layer.mlp_post_norm.register_forward_hook(
|
||||
_hook(f"{tag}.mlp_post_norm")
|
||||
)
|
||||
)
|
||||
handles.append(layer.mlp_post_norm.register_forward_hook(_hook(f"{tag}.mlp_post_norm")))
|
||||
return handles
|
||||
|
||||
|
||||
@@ -193,11 +174,9 @@ def _write_log(entries: list[dict], path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(path, "w") as f:
|
||||
for e in entries:
|
||||
f.write(
|
||||
f"{e['name']} {e['shape']} "
|
||||
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
|
||||
f"{e['min']:.6f} {e['max']:.6f}\n"
|
||||
)
|
||||
f.write(f"{e['name']} {e['shape']} "
|
||||
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
|
||||
f"{e['min']:.6f} {e['max']:.6f}\n")
|
||||
|
||||
|
||||
def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
@@ -205,9 +184,14 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
return (0, "")
|
||||
if name.startswith(f"L{drill_layer:02d}."):
|
||||
sub_order = {
|
||||
"attention": 0, "attn_post_norm": 1, "mlp.pre_norm": 2,
|
||||
"mlp.up_gate_proj": 3, "mlp.down_proj<in>": 4,
|
||||
"mlp.down_proj": 5, "mlp": 6, "mlp_post_norm": 7,
|
||||
"attention": 0,
|
||||
"attn_post_norm": 1,
|
||||
"mlp.pre_norm": 2,
|
||||
"mlp.up_gate_proj": 3,
|
||||
"mlp.down_proj<in>": 4,
|
||||
"mlp.down_proj": 5,
|
||||
"mlp": 6,
|
||||
"mlp_post_norm": 7,
|
||||
}.get(name.split(".", 1)[1], 9)
|
||||
return (1, f"block[{drill_layer:02d}]", sub_order)
|
||||
if name.startswith("block["):
|
||||
@@ -216,10 +200,8 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
|
||||
|
||||
def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
|
||||
hdr = (
|
||||
f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
|
||||
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}"
|
||||
)
|
||||
hdr = (f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
|
||||
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}")
|
||||
print(f"\n{hdr}\n{'-' * len(hdr)}")
|
||||
first_div: int | None = None
|
||||
for name in sorted(by_name.keys(), key=lambda n: _sort_key(n, drill_layer)):
|
||||
@@ -235,11 +217,9 @@ def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
|
||||
flag = " <<< DIVERGE"
|
||||
if first_div is None:
|
||||
first_div = int(name[len("block["):-1])
|
||||
print(
|
||||
f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
|
||||
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
|
||||
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}"
|
||||
)
|
||||
print(f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
|
||||
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
|
||||
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}")
|
||||
return first_div
|
||||
|
||||
|
||||
@@ -255,10 +235,8 @@ def _print_elementwise(up_t: dict[str, torch.Tensor], fv_t: dict[str, torch.Tens
|
||||
continue
|
||||
diff = (a - b).abs()
|
||||
rel = (diff.mean().item() / max(a.abs().mean().item(), 1e-9)) * 100
|
||||
print(
|
||||
f"{name:<30} {str(tuple(a.shape)):<22} "
|
||||
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%"
|
||||
)
|
||||
print(f"{name:<30} {str(tuple(a.shape)):<22} "
|
||||
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
|
||||
@@ -43,12 +43,10 @@ def _add_official_to_path() -> Path:
|
||||
|
||||
def _log_tensor_stats(label: str, tensor: torch.Tensor) -> None:
|
||||
value = tensor.detach().float()
|
||||
print(
|
||||
f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
|
||||
f"dtype={tensor.dtype} device={tensor.device} "
|
||||
f"min={value.min().item():.6f} max={value.max().item():.6f} "
|
||||
f"mean={value.mean().item():.6f} std={value.std().item():.6f}"
|
||||
)
|
||||
print(f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
|
||||
f"dtype={tensor.dtype} device={tensor.device} "
|
||||
f"min={value.min().item():.6f} max={value.max().item():.6f} "
|
||||
f"mean={value.mean().item():.6f} std={value.std().item():.6f}")
|
||||
|
||||
|
||||
def _extract_tensor(output: Any, key: str) -> torch.Tensor:
|
||||
@@ -73,10 +71,8 @@ def _run_official_pipeline(
|
||||
device: torch.device,
|
||||
) -> Any:
|
||||
del official_path, params, device
|
||||
pytest.skip(
|
||||
"TODO: import the official pipeline/factory, load official weights, "
|
||||
"run with params, and return the comparison target."
|
||||
)
|
||||
pytest.skip("TODO: import the official pipeline/factory, load official weights, "
|
||||
"run with params, and return the comparison target.")
|
||||
|
||||
|
||||
def _run_fastvideo_pipeline(model_path: Path, params: dict[str, Any]) -> Any:
|
||||
@@ -146,8 +142,6 @@ def test_todo_model_family_pipeline_official_parity() -> None:
|
||||
assert official_tensor.shape == fastvideo_tensor.shape
|
||||
|
||||
diff = (official_tensor - fastvideo_tensor).abs()
|
||||
print(
|
||||
f"diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
|
||||
)
|
||||
print(f"diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}")
|
||||
assert_close(fastvideo_tensor, official_tensor, atol=1e-2, rtol=1e-2)
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
name: macOS MLX Smoke
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
paths:
|
||||
- ".github/workflows/ci-macos-mlx.yml"
|
||||
- "fastvideo/mlx_runtime/**"
|
||||
- "fastvideo/tests/mlx/**"
|
||||
- "fastvideo/tests/platforms/test_mps_vsa_error.py"
|
||||
- "fastvideo/platforms/mps.py"
|
||||
- "fastvideo/platforms/__init__.py"
|
||||
- "fastvideo/__init__.py"
|
||||
- "examples/inference/basic/mlx_*.py"
|
||||
- "fastvideo/benchmarks/mlx_*.py"
|
||||
- "pyproject.toml"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: macos-mlx-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
mlx-smoke:
|
||||
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
|
||||
runs-on: macos-15
|
||||
timeout-minutes: 25
|
||||
env:
|
||||
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
|
||||
TOKENIZERS_PARALLELISM: "false"
|
||||
MASTER_ADDR: localhost
|
||||
MASTER_PORT: "29513"
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
cache: pip
|
||||
|
||||
- uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install lightweight MLX smoke dependencies
|
||||
run: |
|
||||
uv pip install --system \
|
||||
--index-url https://download.pytorch.org/whl/cpu \
|
||||
torch==2.11.0 torchvision torchaudio
|
||||
uv pip install --system \
|
||||
pytest numpy scipy pillow imageio einops cloudpickle filelock \
|
||||
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru mlx \
|
||||
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
|
||||
|
||||
- name: Show Apple runtime
|
||||
run: |
|
||||
python - <<'PY'
|
||||
import platform
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
print("machine:", platform.machine())
|
||||
print("processor:", platform.processor())
|
||||
print("mlx default device:", mx.default_device())
|
||||
memory_size = mx.metal.device_info().get("memory_size") if mx.metal.is_available() else "metal unavailable"
|
||||
print("mlx memory_size:", memory_size)
|
||||
print("torch:", torch.__version__)
|
||||
print("torch mps available:", torch.backends.mps.is_available())
|
||||
PY
|
||||
|
||||
- name: Run MLX smoke tests
|
||||
run: |
|
||||
python -m pytest \
|
||||
fastvideo/tests/mlx/test_dmd_sampling.py \
|
||||
fastvideo/tests/mlx/test_memory_limits.py \
|
||||
fastvideo/tests/mlx/test_quant_capability.py \
|
||||
fastvideo/tests/mlx/test_mlx_dit_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_compile_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_checkpoint.py \
|
||||
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
|
||||
fastvideo/tests/mlx/test_taehv_decode.py \
|
||||
fastvideo/tests/mlx/test_frame_upsample.py \
|
||||
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
|
||||
fastvideo/tests/mlx/test_mlx_refine.py \
|
||||
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
|
||||
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
|
||||
fastvideo/tests/mlx/test_wan22_sample.py \
|
||||
fastvideo/tests/mlx/test_windowed_attention.py \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
|
||||
fastvideo/tests/platforms/test_mps_vsa_error.py \
|
||||
-q
|
||||
|
||||
# Same tests on MLX's CPU backend. Hosted macOS runners are scarce and
|
||||
# slower to schedule; this Linux job gives fast PR signal on the identical
|
||||
# graph (the parity tests were designed to be backend-agnostic), while the
|
||||
# macOS job above stays the source of truth for Metal behavior.
|
||||
mlx-smoke-linux-cpu:
|
||||
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
env:
|
||||
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
|
||||
TOKENIZERS_PARALLELISM: "false"
|
||||
MASTER_ADDR: localhost
|
||||
MASTER_PORT: "29513"
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
cache: pip
|
||||
|
||||
- uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install lightweight MLX smoke dependencies (CPU backend)
|
||||
run: |
|
||||
uv pip install --system \
|
||||
--index-url https://download.pytorch.org/whl/cpu \
|
||||
torch==2.11.0 torchvision torchaudio
|
||||
uv pip install --system \
|
||||
pytest numpy scipy pillow imageio einops cloudpickle filelock \
|
||||
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru "mlx[cpu]" \
|
||||
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
|
||||
|
||||
- name: Run MLX smoke tests (CPU backend)
|
||||
run: |
|
||||
python -m pytest \
|
||||
fastvideo/tests/mlx/test_dmd_sampling.py \
|
||||
fastvideo/tests/mlx/test_memory_limits.py \
|
||||
fastvideo/tests/mlx/test_quant_capability.py \
|
||||
fastvideo/tests/mlx/test_mlx_dit_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_compile_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_checkpoint.py \
|
||||
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
|
||||
fastvideo/tests/mlx/test_taehv_decode.py \
|
||||
fastvideo/tests/mlx/test_frame_upsample.py \
|
||||
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
|
||||
fastvideo/tests/mlx/test_mlx_refine.py \
|
||||
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
|
||||
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
|
||||
fastvideo/tests/mlx/test_wan22_sample.py \
|
||||
fastvideo/tests/mlx/test_windowed_attention.py \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
|
||||
fastvideo/tests/platforms/test_mps_vsa_error.py \
|
||||
-q
|
||||
@@ -6,6 +6,7 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
@@ -16,6 +17,7 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
|
||||
@@ -23,6 +23,7 @@ Miniconda3-latest-Linux-x86_64.sh
|
||||
*validation/
|
||||
data/
|
||||
outputs/
|
||||
outputs_audio/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
|
||||
@@ -22,6 +22,7 @@ repos:
|
||||
hooks:
|
||||
- id: yapf
|
||||
args: [--in-place, --verbose]
|
||||
language_version: python3.12
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.12
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
|
||||
|
||||
## NEWS
|
||||
- `2026/08/19`: FastVideo now supports MLX on Apple Silicon with [FastMetal-QAD](https://huggingface.co/collections/FastVideo/fastmetal), a family of 1.3B, 5B, and 14B models optimized for Mac—follow the [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/) and read the [Blog](https://haoailab.com/blogs/fastmetal/).
|
||||
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
|
||||
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
|
||||
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
|
||||
@@ -62,6 +63,11 @@ UV_TORCH_BACKEND=cu126 uv pip install fastvideo
|
||||
Use `UV_TORCH_BACKEND=cu130` on CUDA 13. Apple silicon users should follow the
|
||||
[MPS installation guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
|
||||
|
||||
> **On an Apple Silicon Mac?** FastVideo runs FastWan text-to-video natively
|
||||
> through an MLX runtime — a 5-second 480p clip generated locally, no cloud,
|
||||
> no discrete GPU. Install with `uv pip install -e '.[mlx]'` and follow the
|
||||
> [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
|
||||
|
||||
> **On an NVIDIA DGX Spark (GB10 / ARM64 + CUDA 13)?** There's no prebuilt ARM wheel for the FastVideo CUDA kernel, so it's an editable from-source install (`UV_TORCH_BACKEND=cu130 uv pip install -e .`, which compiles that kernel for you) rather than `UV_TORCH_BACKEND=cu130 uv pip install fastvideo`. A compatible prebuilt ARM64 FlashAttention wheel is available separately. Follow the [DGX Spark install guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/spark/).
|
||||
|
||||
@@ -3,7 +3,6 @@ from __future__ import annotations
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
TESTS_DIR = Path(__file__).resolve().parent
|
||||
DREAMVERSE_PACKAGE_DIR = TESTS_DIR.parent
|
||||
DREAMVERSE_APP_DIR = DREAMVERSE_PACKAGE_DIR.parent
|
||||
|
||||
@@ -5,7 +5,6 @@ from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
SERVER_DIR = Path(__file__).resolve().parents[1]
|
||||
|
||||
|
||||
@@ -53,9 +52,7 @@ def test_config_defaults_to_cerebras_with_parallel_groq_fallback_stage(monkeypat
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.PROMPT_PROVIDER == "cerebras"
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
|
||||
("cerebras", "groq"),
|
||||
)
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
|
||||
assert module.PROMPT_PROVIDER_PRIORITY == (
|
||||
"cerebras",
|
||||
"groq",
|
||||
@@ -86,9 +83,7 @@ def test_config_ignores_legacy_groq_primary_override(monkeypatch):
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.PROMPT_PROVIDER == "cerebras"
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
|
||||
("cerebras", "groq"),
|
||||
)
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
|
||||
assert module.PROMPT_PROVIDER_PRIORITY == (
|
||||
"cerebras",
|
||||
"groq",
|
||||
@@ -106,24 +101,17 @@ def test_config_uses_local_overlay_paths_when_devtools_enabled(monkeypatch, tmp_
|
||||
|
||||
assert module.DEVTOOLS_ENABLED is True
|
||||
assert module.FRONTEND_ROOT.as_posix().endswith("apps/dreamverse/web")
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith(
|
||||
"dreamverse/prompts.local/next_segment_system_prompt.md"
|
||||
)
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith("dreamverse/prompts.local/next_segment_system_prompt.md")
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
|
||||
"dreamverse/prompts/next_segment_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts/next_segment_system_prompt.md")
|
||||
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_PATH.endswith(
|
||||
"dreamverse/prompts.local/rewrite_user_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts.local/rewrite_user_system_prompt.md")
|
||||
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
|
||||
"dreamverse/prompts/rewrite_user_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts/rewrite_user_system_prompt.md")
|
||||
assert module.CURATED_PRESETS_FILE_PATH.endswith(
|
||||
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json")
|
||||
assert module.CURATED_PRESETS_FALLBACK_FILE_PATH.endswith(
|
||||
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json")
|
||||
assert module.FRONTEND_STATIC_DIR_CANDIDATES[:2] == (
|
||||
str(module.FRONTEND_ROOT / "out"),
|
||||
str(module.FRONTEND_ROOT / "dist"),
|
||||
|
||||
@@ -9,6 +9,7 @@ from fastapi.testclient import TestClient
|
||||
import fastvideo.entrypoints.streaming as streaming_entrypoints
|
||||
import pytest
|
||||
|
||||
|
||||
def _install_stack03_import_stubs(monkeypatch):
|
||||
"""Keep entrypoint tests focused while later-stack runtime modules are absent."""
|
||||
if not hasattr(streaming_entrypoints, "build_health_router"):
|
||||
@@ -17,6 +18,7 @@ def _install_stack03_import_stubs(monkeypatch):
|
||||
gpu_pool_stub = types.ModuleType("dreamverse.gpu_pool")
|
||||
|
||||
class GPUPool:
|
||||
|
||||
def __init__(self, _gpu_ids):
|
||||
pass
|
||||
|
||||
@@ -49,6 +51,7 @@ def _install_stack03_import_stubs(monkeypatch):
|
||||
controller_stub = types.ModuleType("dreamverse.session.controller")
|
||||
|
||||
class SessionController:
|
||||
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
@@ -76,13 +79,11 @@ def _run_cli(module, monkeypatch, argv: list[str]) -> list[dict[str, object]]:
|
||||
uvicorn_stub = types.ModuleType("uvicorn")
|
||||
|
||||
def run(app, host: str, port: int) -> None:
|
||||
calls.append(
|
||||
{
|
||||
"app": app,
|
||||
"host": host,
|
||||
"port": port,
|
||||
}
|
||||
)
|
||||
calls.append({
|
||||
"app": app,
|
||||
"host": host,
|
||||
"port": port,
|
||||
})
|
||||
|
||||
uvicorn_stub.run = run
|
||||
monkeypatch.setitem(sys.modules, "uvicorn", uvicorn_stub)
|
||||
@@ -99,13 +100,11 @@ def test_server_cli_defaults_to_local_web_port(monkeypatch):
|
||||
server_main = _import_server_main(monkeypatch)
|
||||
calls = _run_cli(server_main, monkeypatch, ["dreamverse-server"])
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": server_main.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": server_main.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}]
|
||||
|
||||
|
||||
def test_server_cli_allows_explicit_host_and_port(monkeypatch):
|
||||
@@ -116,13 +115,11 @@ def test_server_cli_allows_explicit_host_and_port(monkeypatch):
|
||||
["dreamverse-server", "--host", "127.0.0.1", "--port", "8123"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": server_main.app,
|
||||
"host": "127.0.0.1",
|
||||
"port": 8123,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": server_main.app,
|
||||
"host": "127.0.0.1",
|
||||
"port": 8123,
|
||||
}]
|
||||
|
||||
|
||||
def test_server_does_not_expose_backend_source_as_static_assets(monkeypatch):
|
||||
@@ -142,13 +139,11 @@ def test_mock_server_cli_defaults_to_local_web_port(monkeypatch):
|
||||
["dreamverse-mock-server"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}]
|
||||
|
||||
|
||||
def test_mock_server_cli_updates_latency(monkeypatch):
|
||||
@@ -161,13 +156,11 @@ def test_mock_server_cli_updates_latency(monkeypatch):
|
||||
["dreamverse-mock-server", "--latency", "321", "--port", "8111"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8111,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8111,
|
||||
}]
|
||||
assert mock_server.LATENCY_MS == 321
|
||||
finally:
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
|
||||
@@ -7,7 +7,6 @@ from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
import dreamverse.gpu_pool as gpu_pool
|
||||
|
||||
|
||||
@@ -85,9 +84,7 @@ def test_send_command_raises_on_worker_death():
|
||||
cmd_q = ctx.Queue()
|
||||
resp_q = ctx.Queue()
|
||||
|
||||
proc = ctx.Process(
|
||||
target=_child_consume_and_exit, args=(cmd_q, resp_q)
|
||||
)
|
||||
proc = ctx.Process(target=_child_consume_and_exit, args=(cmd_q, resp_q))
|
||||
proc.start()
|
||||
|
||||
# Wait for the spawn child to fully boot. Allow generous time —
|
||||
|
||||
@@ -7,7 +7,7 @@ ALLOWED_PREFIXES = (
|
||||
"fastvideo.entrypoints.video_generator",
|
||||
"fastvideo.configs",
|
||||
)
|
||||
ALLOWED_EXACT = ("fastvideo",)
|
||||
ALLOWED_EXACT = ("fastvideo", )
|
||||
FORBIDDEN_PREFIXES = (
|
||||
"fastvideo.pipelines",
|
||||
"fastvideo.models",
|
||||
@@ -38,19 +38,13 @@ def test_dreamverse_server_imports_only_public_fastvideo_surfaces() -> None:
|
||||
except SyntaxError as task_exc:
|
||||
raise AssertionError(f"Failed to parse {path}") from task_exc
|
||||
for node in ast.walk(tree):
|
||||
names = (
|
||||
[a.name for a in node.names] if isinstance(node, ast.Import)
|
||||
else [node.module] if isinstance(node, ast.ImportFrom) and node.module
|
||||
else []
|
||||
)
|
||||
names = ([a.name for a in node.names] if isinstance(node, ast.Import) else
|
||||
[node.module] if isinstance(node, ast.ImportFrom) and node.module else [])
|
||||
for name in names:
|
||||
if not name:
|
||||
continue
|
||||
rel_path = str(path.relative_to(root))
|
||||
if (
|
||||
name.startswith(FORBIDDEN_PREFIXES)
|
||||
and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS
|
||||
):
|
||||
if (name.startswith(FORBIDDEN_PREFIXES) and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS):
|
||||
bad.append((str(path.relative_to(root)), getattr(node, "lineno", 0), name))
|
||||
|
||||
assert bad == [], f"Forbidden internal imports: {bad}"
|
||||
|
||||
@@ -6,7 +6,6 @@ import os
|
||||
|
||||
from fastapi import WebSocketDisconnect
|
||||
|
||||
|
||||
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
|
||||
os.environ.setdefault("GROQ_API_KEY", "dummy")
|
||||
|
||||
@@ -14,6 +13,7 @@ import dreamverse.mock_server as mock_server
|
||||
|
||||
|
||||
class _FakeWebSocket:
|
||||
|
||||
def __init__(self, messages: list[tuple[float, dict[str, object]]]):
|
||||
self._messages = messages
|
||||
self._index = 0
|
||||
@@ -49,34 +49,34 @@ def test_mock_server_matches_current_single5s_protocol():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "simple_prompt_1",
|
||||
"curated_prompts": ["selected prompt"],
|
||||
"single_clip_mode": True,
|
||||
"enhancement_enabled": False,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.01,
|
||||
{
|
||||
"type": "simple_generate",
|
||||
"preset_id": "simple_custom_prompt",
|
||||
"prompt_id": "simple_custom_prompt",
|
||||
"prompt": "custom prompt",
|
||||
"enhancement_enabled": True,
|
||||
"initial_image": None,
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "simple_prompt_1",
|
||||
"curated_prompts": ["selected prompt"],
|
||||
"single_clip_mode": True,
|
||||
"enhancement_enabled": False,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.01,
|
||||
{
|
||||
"type": "simple_generate",
|
||||
"preset_id": "simple_custom_prompt",
|
||||
"prompt_id": "simple_custom_prompt",
|
||||
"prompt": "custom prompt",
|
||||
"enhancement_enabled": True,
|
||||
"initial_image": None,
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -92,24 +92,14 @@ def test_mock_server_matches_current_single5s_protocol():
|
||||
assert message_types.count("ltx2_stream_complete") == 2
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
gpu_assigned_event = next(
|
||||
payload for payload in ws.sent_json if payload["type"] == "gpu_assigned"
|
||||
)
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
gpu_assigned_event = next(payload for payload in ws.sent_json if payload["type"] == "gpu_assigned")
|
||||
assert gpu_assigned_event["session_timeout"] == mock_server.SESSION_TIMEOUT_SECONDS
|
||||
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
|
||||
assert segment_start_events[0]["prompt"] == "selected prompt"
|
||||
assert segment_start_events[1]["prompt"] == "custom prompt"
|
||||
|
||||
step_complete_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "step_complete"
|
||||
]
|
||||
step_complete_events = [payload for payload in ws.sent_json if payload["type"] == "step_complete"]
|
||||
assert len(step_complete_events) == 2
|
||||
assert step_complete_events[0]["latency_ms"] == {
|
||||
"total": 121.0,
|
||||
@@ -134,29 +124,29 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
|
||||
mock_server.LATENCY_MS = 1
|
||||
mock_server.GENERATION_SEGMENT_CAP = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "start a new rollout",
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "start a new rollout",
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -166,11 +156,7 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
|
||||
assert "generation_cap_reached" not in message_types
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
|
||||
assert segment_start_events[0]["prompt"] == "segment one"
|
||||
assert segment_start_events[1]["prompt"] == "segment one [start a new rollout]"
|
||||
@@ -187,54 +173,40 @@ def test_mock_server_rewrite_during_active_segment_restarts_from_first_rewritten
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 100
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one", "segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "restart from rewrite",
|
||||
},
|
||||
),
|
||||
(0.40, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one", "segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "restart from rewrite",
|
||||
},
|
||||
),
|
||||
(0.40, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
|
||||
"segment one",
|
||||
"segment one [restart from rewrite]",
|
||||
]
|
||||
assert all(
|
||||
payload["prompt"] != "segment two"
|
||||
for payload in segment_start_events[1:]
|
||||
)
|
||||
reset_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload.get("type") == "seed_prompts_reset_applied"
|
||||
]
|
||||
assert any(
|
||||
payload.get("reason") == "rewrite_during_generation"
|
||||
for payload in reset_events
|
||||
)
|
||||
assert all(payload["prompt"] != "segment two" for payload in segment_start_events[1:])
|
||||
reset_events = [payload for payload in ws.sent_json if payload.get("type") == "seed_prompts_reset_applied"]
|
||||
assert any(payload.get("reason") == "rewrite_during_generation" for payload in reset_events)
|
||||
finally:
|
||||
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
@@ -247,24 +219,24 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "custom_editable",
|
||||
"preset_label": "Custom rollout",
|
||||
"curated_prompts": [],
|
||||
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "custom_editable",
|
||||
"preset_label": "Custom rollout",
|
||||
"curated_prompts": [],
|
||||
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -276,15 +248,9 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
|
||||
assert "ltx2_stream_start" in message_types
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert segment_start_events
|
||||
assert segment_start_events[0]["prompt"] == (
|
||||
"A moonbase corridor thriller with flooding [segment 1]"
|
||||
)
|
||||
assert segment_start_events[0]["prompt"] == ("A moonbase corridor thriller with flooding [segment 1]")
|
||||
finally:
|
||||
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
@@ -297,35 +263,37 @@ def test_mock_server_can_start_new_project_without_reconnecting():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 40
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.02, {"type": "end_project_keep_session"}),
|
||||
(
|
||||
0.20,
|
||||
{
|
||||
"type": "project_init_v1",
|
||||
"preset_id": "test_preset_2",
|
||||
"preset_label": "Test Preset 2",
|
||||
"curated_prompts": ["segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.40, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.02, {
|
||||
"type": "end_project_keep_session"
|
||||
}),
|
||||
(
|
||||
0.20,
|
||||
{
|
||||
"type": "project_init_v1",
|
||||
"preset_id": "test_preset_2",
|
||||
"preset_label": "Test Preset 2",
|
||||
"curated_prompts": ["segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.40, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -336,16 +304,11 @@ def test_mock_server_can_start_new_project_without_reconnecting():
|
||||
|
||||
project_idle_index = message_types.index("project_idle")
|
||||
stream_start_indexes = [
|
||||
index for index, message_type in enumerate(message_types)
|
||||
if message_type == "ltx2_stream_start"
|
||||
index for index, message_type in enumerate(message_types) if message_type == "ltx2_stream_start"
|
||||
]
|
||||
assert stream_start_indexes[0] < project_idle_index < stream_start_indexes[1]
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
|
||||
"segment one",
|
||||
"segment two",
|
||||
|
||||
@@ -6,7 +6,6 @@ import os
|
||||
import re
|
||||
import time
|
||||
|
||||
|
||||
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
|
||||
os.environ.setdefault("GROQ_API_KEY", "dummy")
|
||||
|
||||
@@ -22,6 +21,7 @@ from dreamverse.prompt_enhancer import (
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self._payload = payload
|
||||
|
||||
@@ -30,6 +30,7 @@ class _FakeResponse:
|
||||
|
||||
|
||||
class _FakeSyncCompletions:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self._payload = payload
|
||||
|
||||
@@ -38,6 +39,7 @@ class _FakeSyncCompletions:
|
||||
|
||||
|
||||
class _FakeSyncClient:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self.chat = type(
|
||||
"_FakeChat",
|
||||
@@ -47,6 +49,7 @@ class _FakeSyncClient:
|
||||
|
||||
|
||||
class _DelayedSyncCompletions:
|
||||
|
||||
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
|
||||
self._payload = payload
|
||||
self._delay_s = delay_s
|
||||
@@ -61,29 +64,26 @@ class _DelayedSyncCompletions:
|
||||
|
||||
|
||||
class _DelayedSyncClient:
|
||||
|
||||
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
|
||||
self.chat = type(
|
||||
"_FakeChat",
|
||||
(),
|
||||
{
|
||||
"completions": _DelayedSyncCompletions(
|
||||
payload,
|
||||
delay_s=delay_s,
|
||||
exc=exc,
|
||||
)
|
||||
},
|
||||
{"completions": _DelayedSyncCompletions(
|
||||
payload,
|
||||
delay_s=delay_s,
|
||||
exc=exc,
|
||||
)},
|
||||
)()
|
||||
|
||||
|
||||
def _chat_payload_with_content(content: str) -> dict:
|
||||
return {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"content": content,
|
||||
}
|
||||
"choices": [{
|
||||
"message": {
|
||||
"content": content,
|
||||
}
|
||||
]
|
||||
}]
|
||||
}
|
||||
|
||||
|
||||
@@ -172,6 +172,7 @@ def _build_staged_enhancer(
|
||||
|
||||
|
||||
class _FakeOpenAIClient:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
self.chat = type(
|
||||
@@ -182,6 +183,7 @@ class _FakeOpenAIClient:
|
||||
|
||||
|
||||
class _FakeCerebrasClient:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
self.chat = type(
|
||||
@@ -192,16 +194,12 @@ class _FakeCerebrasClient:
|
||||
|
||||
|
||||
def test_parse_json_response_accepts_fenced_json_with_prose():
|
||||
parsed = _parse_json_response(
|
||||
"Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks."
|
||||
)
|
||||
parsed = _parse_json_response("Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks.")
|
||||
assert parsed == {"segment_prompts": ["A", "B"]}
|
||||
|
||||
|
||||
def test_parse_json_response_extracts_first_embedded_object():
|
||||
parsed = _parse_json_response(
|
||||
"Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)"
|
||||
)
|
||||
parsed = _parse_json_response("Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)")
|
||||
assert parsed == {"segment_prompts": ["A", "B"]}
|
||||
|
||||
|
||||
@@ -268,16 +266,12 @@ def test_build_client_supports_groq_provider(monkeypatch):
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "preset_a"
|
||||
@@ -286,15 +280,12 @@ def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"rewritten_prompts":["A","B"]}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"rewritten_prompts":["A","B"]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "current_rollout"
|
||||
@@ -303,19 +294,14 @@ def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_segment_dicts_without_top_level_rollout_metadata():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"segments":[{"prompt":"A"},{"text":"B"}]}'
|
||||
)
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segments":[{"prompt":"A"},{"text":"B"}]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
preset_id="preset_a",
|
||||
preset_label="Preset A",
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "preset_a"
|
||||
@@ -329,14 +315,12 @@ def test_rewrite_prompt_sequence_accepts_numbered_prose_output():
|
||||
"The user is asking for a cinematic rewrite.\n\n"
|
||||
"1. A dog bounds across the moon's dusty surface, kicking up silver regolith as it chases a rabbit beneath the black sky.\n"
|
||||
"2. The rabbit darts around a crater rim while the dog lunges after it, Earth glowing blue in the distance.\n"
|
||||
)
|
||||
)
|
||||
))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "current_rollout"
|
||||
@@ -355,12 +339,10 @@ def test_enhance_prompt_prefers_cerebras_before_groq_fallback():
|
||||
groq_delay_s=0.01,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -381,12 +363,10 @@ def test_enhance_prompt_uses_groq_when_cerebras_fails():
|
||||
groq_delay_s=0.01,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -408,13 +388,11 @@ def test_enhance_prompt_can_use_groq_when_cerebras_times_out():
|
||||
enhancer.http_timeout_ms = 50
|
||||
enhancer.default_timeout_ms = 50
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
timeout_ms=50,
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
timeout_ms=50,
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -434,12 +412,10 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
|
||||
groq_delay_s=0.08,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -453,15 +429,12 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("I cannot comply with JSON right now.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("I cannot comply with JSON right now."))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
assert result.raw_response_text == "I cannot comply with JSON right now."
|
||||
@@ -473,9 +446,7 @@ def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
|
||||
def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'))
|
||||
captured = {
|
||||
"body": None,
|
||||
"timeout_seconds": None,
|
||||
@@ -486,8 +457,7 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
captured["timeout_seconds"] = timeout_seconds
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
|
||||
),
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'),
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}',
|
||||
)
|
||||
|
||||
@@ -502,8 +472,7 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
rewrite_model="gpt-test",
|
||||
rewrite_temperature=0.2,
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert captured["body"]["messages"][0] == {
|
||||
@@ -512,12 +481,12 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
}
|
||||
assert captured["body"]["messages"][1]["role"] == "user"
|
||||
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
|
||||
"mode": "edit_existing_rollout",
|
||||
"request": (
|
||||
"Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."
|
||||
),
|
||||
"user_instruction": "make it cinematic",
|
||||
"mode":
|
||||
"edit_existing_rollout",
|
||||
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."),
|
||||
"user_instruction":
|
||||
"make it cinematic",
|
||||
"current_rollout": {
|
||||
"id": "preset_a",
|
||||
"label": "Preset A",
|
||||
@@ -528,11 +497,8 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
|
||||
def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'))
|
||||
captured = {
|
||||
"body": None,
|
||||
}
|
||||
@@ -541,10 +507,8 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
del timeout_seconds
|
||||
captured["body"] = body
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'
|
||||
),
|
||||
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'),
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}',
|
||||
)
|
||||
@@ -560,30 +524,29 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
rewrite_model="gpt-test",
|
||||
rewrite_temperature=0.2,
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.prompts == ["A", "B", "C", "D", "E", "F"]
|
||||
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
|
||||
"mode": "new_rollout",
|
||||
"request": (
|
||||
"Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."
|
||||
),
|
||||
"user_instruction": "A moonbase corridor thriller with flooding and red alarms",
|
||||
"desired_segment_count": 6,
|
||||
"rollout_id_hint": "custom_editable",
|
||||
"rollout_label_hint": "Custom rollout",
|
||||
"mode":
|
||||
"new_rollout",
|
||||
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."),
|
||||
"user_instruction":
|
||||
"A moonbase corridor thriller with flooding and red alarms",
|
||||
"desired_segment_count":
|
||||
6,
|
||||
"rollout_id_hint":
|
||||
"custom_editable",
|
||||
"rollout_label_hint":
|
||||
"Custom rollout",
|
||||
}
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared system prompt"
|
||||
captured = {
|
||||
"body": None,
|
||||
@@ -593,9 +556,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
del timeout_seconds
|
||||
captured["body"] = body
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
),
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'),
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}',
|
||||
)
|
||||
|
||||
@@ -609,8 +570,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
rewrite_instruction="make it cinematic",
|
||||
rewrite_model="gpt-test",
|
||||
system_prompt_override="session specific system prompt",
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert captured["body"]["messages"][0] == {
|
||||
@@ -621,10 +581,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
|
||||
def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
|
||||
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
|
||||
|
||||
@@ -635,24 +592,17 @@ def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
|
||||
|
||||
def test_resolve_rewrite_new_rollout_system_prompt_prefers_override():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
|
||||
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
|
||||
|
||||
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt(
|
||||
"session specific system prompt"
|
||||
)
|
||||
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt("session specific system prompt")
|
||||
|
||||
assert resolved == "session specific system prompt"
|
||||
|
||||
|
||||
def test_generate_auto_prompt_uses_selected_model():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"next_prompt":"Auto next"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Auto next"}'))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -680,8 +630,7 @@ def test_generate_auto_prompt_uses_selected_model():
|
||||
next_segment_idx=2,
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Auto next"
|
||||
@@ -690,9 +639,7 @@ def test_generate_auto_prompt_uses_selected_model():
|
||||
|
||||
|
||||
def test_enhance_prompt_uses_selected_model():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"next_prompt":"Enhanced next"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Enhanced next"}'))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
@@ -722,8 +669,7 @@ def test_enhance_prompt_uses_selected_model():
|
||||
next_segment_idx=2,
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Enhanced next"
|
||||
@@ -732,9 +678,7 @@ def test_enhance_prompt_uses_selected_model():
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"prompt":"Extended single clip"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"prompt":"Extended single clip"}'))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
@@ -764,14 +708,12 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
|
||||
|
||||
enhancer._request_content = _fake_request_content # type: ignore[attr-defined]
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"short 5s idea",
|
||||
mode="single_clip",
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"short 5s idea",
|
||||
mode="single_clip",
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Extended single clip"
|
||||
@@ -784,17 +726,15 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
|
||||
"single 5-second LTX-2.3 video clip. Respond with "
|
||||
'valid JSON only as {"prompt": "..."}.' # noqa: E501
|
||||
),
|
||||
"user_prompt": "short 5s idea",
|
||||
"user_prompt":
|
||||
"short 5s idea",
|
||||
}
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_rejects_plain_text_response():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
"Medium shot of a woman by a rainy cafe window as she lifts her "
|
||||
"phone, exhales softly, and the camera makes a slow push in."
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content("Medium shot of a woman by a rainy cafe window as she lifts her "
|
||||
"phone, exhales softly, and the camera makes a slow push in."))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
|
||||
result = asyncio.run(
|
||||
@@ -803,17 +743,14 @@ def test_enhance_prompt_single_clip_rejects_plain_text_response():
|
||||
mode="single_clip",
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No JSON object found in assistant response." in result.error
|
||||
assert result.prompt == ""
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"segment_prompts":["A","B"]}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segment_prompts":["A","B"]}'))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -824,17 +761,14 @@ def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
|
||||
mode="single_clip",
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "Missing prompt string." in (result.error or "")
|
||||
|
||||
|
||||
def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("A cinematic continuation with slow dolly movement.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("A cinematic continuation with slow dolly movement."))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -846,17 +780,14 @@ def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
next_segment_idx=2,
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
|
||||
|
||||
def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("A calm, grounded continuation with subtle motion.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("A calm, grounded continuation with subtle motion."))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -867,34 +798,30 @@ def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
next_segment_idx=2,
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
|
||||
enhancer = _build_test_enhancer(
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "length",
|
||||
"message": {
|
||||
"content": [],
|
||||
"refusal": None,
|
||||
},
|
||||
}
|
||||
],
|
||||
"usage": {"completion_tokens": 0},
|
||||
}
|
||||
)
|
||||
enhancer = _build_test_enhancer({
|
||||
"choices": [{
|
||||
"finish_reason": "length",
|
||||
"message": {
|
||||
"content": [],
|
||||
"refusal": None,
|
||||
},
|
||||
}],
|
||||
"usage": {
|
||||
"completion_tokens": 0
|
||||
},
|
||||
})
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No rewrite segment prompts found in assistant response." in (result.error or "")
|
||||
assert isinstance(result.raw_response_text, str)
|
||||
@@ -903,10 +830,7 @@ def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
|
||||
|
||||
def test_get_rewrite_model_config_returns_fixed_defaults():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_default_model = "gpt-oss-120b"
|
||||
enhancer.rewrite_model_options = ["gpt-oss-120b"]
|
||||
|
||||
@@ -918,10 +842,7 @@ def test_get_rewrite_model_config_returns_fixed_defaults():
|
||||
|
||||
def test_get_prompt_config_includes_auto_extension_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.enhance_system_prompt_path = "/tmp/next.md"
|
||||
enhancer.auto_system_prompt_path = "/tmp/auto.md"
|
||||
enhancer.rewrite_all_system_prompt_path = "/tmp/rewrite.md"
|
||||
@@ -948,19 +869,14 @@ def test_get_prompt_config_includes_auto_extension_prompt():
|
||||
|
||||
def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
rewrite_fallback_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
rewrite_fallback_path.write_text("rewrite prompt\n", encoding="utf-8")
|
||||
next_path = tmp_path / "next.md"
|
||||
next_path.write_text("next prompt\n", encoding="utf-8")
|
||||
auto_path = tmp_path / "auto.md"
|
||||
auto_path.write_text("auto prompt\n", encoding="utf-8")
|
||||
enhancer.rewrite_all_system_prompt_path = str(
|
||||
tmp_path / "prompts.local" / "rewrite_window_system_prompt.md"
|
||||
)
|
||||
enhancer.rewrite_all_system_prompt_path = str(tmp_path / "prompts.local" / "rewrite_window_system_prompt.md")
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = str(rewrite_fallback_path)
|
||||
enhancer.enhance_system_prompt_path = str(next_path)
|
||||
enhancer.auto_system_prompt_path = str(auto_path)
|
||||
@@ -973,14 +889,9 @@ def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
|
||||
assert config["rewrite_window_system_prompt_path"] == str(rewrite_fallback_path)
|
||||
|
||||
|
||||
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(
|
||||
tmp_path,
|
||||
):
|
||||
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(tmp_path, ):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
@@ -1007,10 +918,7 @@ def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_emp
|
||||
|
||||
def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1021,9 +929,7 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
enhancer.auto_system_prompt_path = str(auto_path)
|
||||
enhancer.rewrite_all_system_prompt_path = str(rewrite_path)
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
auto_extension_system_prompt="auto updated",
|
||||
)
|
||||
config = enhancer.save_prompt_config(auto_extension_system_prompt="auto updated", )
|
||||
|
||||
assert auto_path.read_text(encoding="utf-8").strip() == "auto updated"
|
||||
assert config["auto_extension_system_prompt"] == "auto updated"
|
||||
@@ -1031,10 +937,7 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1052,9 +955,7 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_user_system_prompt_fallback_path = None
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_user_system_prompt="rewrite user updated",
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_user_system_prompt="rewrite user updated", )
|
||||
|
||||
assert rewrite_user_path.read_text(encoding="utf-8").strip() == "rewrite user updated"
|
||||
assert config["rewrite_user_system_prompt"] == "rewrite user updated"
|
||||
@@ -1062,10 +963,7 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1081,9 +979,7 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_model="gpt-alt",
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_model="gpt-alt", )
|
||||
|
||||
assert enhancer.rewrite_default_model == "gpt-alt"
|
||||
assert config["rewrite_model"] == "gpt-alt"
|
||||
@@ -1092,10 +988,7 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1109,9 +1002,7 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
enhancer.auto_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_temperature=1.3,
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_temperature=1.3, )
|
||||
|
||||
assert enhancer.rewrite_default_temperature == 1.3
|
||||
assert config["rewrite_temperature"] == 1.3
|
||||
@@ -1119,10 +1010,7 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
|
||||
def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
@@ -1136,13 +1024,9 @@ def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_pat
|
||||
enhancer.auto_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
|
||||
enhancer.save_prompt_config(
|
||||
rewrite_window_system_prompt="rewrite updated",
|
||||
)
|
||||
enhancer.save_prompt_config(rewrite_window_system_prompt="rewrite updated", )
|
||||
|
||||
backup_paths = sorted(
|
||||
tmp_path.glob("rewrite_window_system_prompt.*.bak.md")
|
||||
)
|
||||
backup_paths = sorted(tmp_path.glob("rewrite_window_system_prompt.*.bak.md"))
|
||||
|
||||
assert rewrite_path.read_text(encoding="utf-8").strip() == "rewrite updated"
|
||||
assert len(backup_paths) == 1
|
||||
|
||||
@@ -27,13 +27,8 @@ try:
|
||||
except ModuleNotFoundError:
|
||||
websockets = None # type: ignore[assignment]
|
||||
|
||||
|
||||
DEFAULT_PRESET_FILE = (
|
||||
Path(__file__).resolve().parents[2]
|
||||
/ "web"
|
||||
/ "prompts"
|
||||
/ "selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
DEFAULT_PRESET_FILE = (Path(__file__).resolve().parents[2] / "web" / "prompts" /
|
||||
"selected_ltx2_continuation_story_presets.json")
|
||||
|
||||
|
||||
def utc_now_iso() -> str:
|
||||
@@ -65,10 +60,7 @@ def safe_percentile(values: list[float], percentile: float) -> float | None:
|
||||
if lower == upper:
|
||||
return sorted_values[lower]
|
||||
fraction = rank - lower
|
||||
return (
|
||||
sorted_values[lower]
|
||||
+ (sorted_values[upper] - sorted_values[lower]) * fraction
|
||||
)
|
||||
return (sorted_values[lower] + (sorted_values[upper] - sorted_values[lower]) * fraction)
|
||||
|
||||
|
||||
def summarize_series(values: list[float]) -> dict[str, float | int | None]:
|
||||
@@ -145,24 +137,16 @@ def load_curated_prompts(
|
||||
selected_id = str(selected.get("id", "")).strip() or "unknown_preset"
|
||||
raw_prompts = selected.get("segment_prompts", [])
|
||||
if not isinstance(raw_prompts, list):
|
||||
raise ValueError(
|
||||
f"Preset {selected_id} has invalid segment_prompts (must be list)."
|
||||
)
|
||||
raise ValueError(f"Preset {selected_id} has invalid segment_prompts (must be list).")
|
||||
|
||||
prompts = [
|
||||
str(prompt).strip()
|
||||
for prompt in raw_prompts
|
||||
if isinstance(prompt, str) and str(prompt).strip()
|
||||
]
|
||||
prompts = [str(prompt).strip() for prompt in raw_prompts if isinstance(prompt, str) and str(prompt).strip()]
|
||||
if not prompts:
|
||||
raise ValueError(f"Preset {selected_id} has no non-empty prompts.")
|
||||
|
||||
limited = prompts[:curated_limit]
|
||||
if not limited:
|
||||
raise ValueError(
|
||||
f"curated_limit={curated_limit} produced no prompts for preset "
|
||||
f"{selected_id}."
|
||||
)
|
||||
raise ValueError(f"curated_limit={curated_limit} produced no prompts for preset "
|
||||
f"{selected_id}.")
|
||||
return selected_id, limited, len(prompts)
|
||||
|
||||
|
||||
@@ -224,11 +208,11 @@ async def run_single_session(
|
||||
|
||||
try:
|
||||
async with websockets.connect(
|
||||
url,
|
||||
max_size=None,
|
||||
ping_interval=None,
|
||||
open_timeout=connect_timeout_s,
|
||||
close_timeout=2.0,
|
||||
url,
|
||||
max_size=None,
|
||||
ping_interval=None,
|
||||
open_timeout=connect_timeout_s,
|
||||
close_timeout=2.0,
|
||||
) as ws:
|
||||
connect_finish_monotonic = time.monotonic()
|
||||
session_data["connect_finish_ts_utc"] = utc_now_iso()
|
||||
@@ -249,9 +233,7 @@ async def run_single_session(
|
||||
timeout_remaining = session_timeout_s - elapsed_s
|
||||
if timeout_remaining <= 0:
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = (
|
||||
f"Session timed out after {session_timeout_s:.1f}s."
|
||||
)
|
||||
session_data["error"] = (f"Session timed out after {session_timeout_s:.1f}s.")
|
||||
break
|
||||
|
||||
recv_start_epoch = time.time()
|
||||
@@ -265,9 +247,7 @@ async def run_single_session(
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = (
|
||||
"Timed out waiting for websocket message."
|
||||
)
|
||||
session_data["error"] = ("Timed out waiting for websocket message.")
|
||||
break
|
||||
except Exception as exc:
|
||||
session_data["status"] = "failed"
|
||||
@@ -288,20 +268,16 @@ async def run_single_session(
|
||||
|
||||
chunk_gap_ms: float | None = None
|
||||
if last_chunk_finish_monotonic is not None:
|
||||
chunk_gap_ms = (
|
||||
recv_finish_monotonic - last_chunk_finish_monotonic
|
||||
) * 1000.0
|
||||
chunk_gap_ms = (recv_finish_monotonic - last_chunk_finish_monotonic) * 1000.0
|
||||
|
||||
session_data["chunks"].append(
|
||||
{
|
||||
"segment_idx": current_segment_idx,
|
||||
"chunk_idx": session_data["total_chunks"],
|
||||
"size_bytes": len(message),
|
||||
"chunk_start_ts_utc": recv_start_iso,
|
||||
"chunk_finish_ts_utc": recv_finish_iso,
|
||||
"chunk_gap_ms": chunk_gap_ms,
|
||||
}
|
||||
)
|
||||
session_data["chunks"].append({
|
||||
"segment_idx": current_segment_idx,
|
||||
"chunk_idx": session_data["total_chunks"],
|
||||
"size_bytes": len(message),
|
||||
"chunk_start_ts_utc": recv_start_iso,
|
||||
"chunk_finish_ts_utc": recv_finish_iso,
|
||||
"chunk_gap_ms": chunk_gap_ms,
|
||||
})
|
||||
last_chunk_finish_monotonic = recv_finish_monotonic
|
||||
last_chunk_finish_epoch = recv_finish_epoch
|
||||
session_data["last_chunk_finish_ts_utc"] = recv_finish_iso
|
||||
@@ -321,9 +297,7 @@ async def run_single_session(
|
||||
if msg_type == "gpu_assigned":
|
||||
session_data["gpu_assigned_ts_utc"] = recv_finish_iso
|
||||
if connect_finish_monotonic is not None:
|
||||
session_data["queue_wait_ms"] = (
|
||||
recv_finish_monotonic - connect_finish_monotonic
|
||||
) * 1000.0
|
||||
session_data["queue_wait_ms"] = (recv_finish_monotonic - connect_finish_monotonic) * 1000.0
|
||||
elif msg_type == "ltx2_stream_start":
|
||||
if initial_total_segments is None:
|
||||
parsed_total = parse_int(data.get("total_segments"))
|
||||
@@ -338,20 +312,13 @@ async def run_single_session(
|
||||
session_data["media_segments_completed"] += 1
|
||||
if first_media_segment_complete_epoch is None:
|
||||
first_media_segment_complete_epoch = recv_finish_epoch
|
||||
session_data[
|
||||
"first_media_segment_complete_ts_utc"
|
||||
] = recv_finish_iso
|
||||
session_data["first_media_segment_complete_ts_utc"] = recv_finish_iso
|
||||
elif msg_type == "ltx2_segment_complete":
|
||||
session_data["segments_completed"] += 1
|
||||
seg_idx = parse_int(data.get("segment_idx"))
|
||||
if (
|
||||
initial_total_segments is not None
|
||||
and seg_idx is not None
|
||||
and seg_idx >= initial_total_segments
|
||||
):
|
||||
session_data[
|
||||
"target_segment_complete_ts_utc"
|
||||
] = recv_finish_iso
|
||||
if (initial_total_segments is not None and seg_idx is not None
|
||||
and seg_idx >= initial_total_segments):
|
||||
session_data["target_segment_complete_ts_utc"] = recv_finish_iso
|
||||
await asyncio.sleep(post_complete_wait_s)
|
||||
session_data["leave_sent_ts_utc"] = utc_now_iso()
|
||||
try:
|
||||
@@ -362,15 +329,11 @@ async def run_single_session(
|
||||
break
|
||||
elif msg_type == "session_timeout":
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = str(
|
||||
data.get("message") or "Backend session timeout"
|
||||
)
|
||||
session_data["error"] = str(data.get("message") or "Backend session timeout")
|
||||
break
|
||||
elif msg_type == "error":
|
||||
session_data["status"] = "failed"
|
||||
session_data["error"] = str(
|
||||
data.get("message") or "Backend error message"
|
||||
)
|
||||
session_data["error"] = str(data.get("message") or "Backend error message")
|
||||
break
|
||||
|
||||
if session_data["status"] == "failed" and session_data["error"] is None:
|
||||
@@ -379,29 +342,18 @@ async def run_single_session(
|
||||
session_data["status"] = "failed"
|
||||
session_data["error"] = f"WebSocket connect/run failed: {exc}"
|
||||
|
||||
if (
|
||||
first_chunk_finish_epoch is not None
|
||||
and last_chunk_finish_epoch is not None
|
||||
and session_data["total_chunk_bytes"] > 0
|
||||
):
|
||||
if (first_chunk_finish_epoch is not None and last_chunk_finish_epoch is not None
|
||||
and session_data["total_chunk_bytes"] > 0):
|
||||
duration_s = last_chunk_finish_epoch - first_chunk_finish_epoch
|
||||
if duration_s > 0:
|
||||
session_data["session_goodput_mbps"] = (
|
||||
session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0
|
||||
)
|
||||
session_data["session_goodput_mbps"] = (session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0)
|
||||
|
||||
if (
|
||||
first_chunk_finish_epoch is not None
|
||||
and first_media_segment_complete_epoch is not None
|
||||
):
|
||||
session_data["first_chunk_before_first_media_complete"] = (
|
||||
first_chunk_finish_epoch < first_media_segment_complete_epoch
|
||||
)
|
||||
if (first_chunk_finish_epoch is not None and first_media_segment_complete_epoch is not None):
|
||||
session_data["first_chunk_before_first_media_complete"] = (first_chunk_finish_epoch
|
||||
< first_media_segment_complete_epoch)
|
||||
|
||||
session_data["close_ts_utc"] = utc_now_iso()
|
||||
session_data["duration_ms"] = (
|
||||
time.monotonic() - session_start_monotonic
|
||||
) * 1000.0
|
||||
session_data["duration_ms"] = (time.monotonic() - session_start_monotonic) * 1000.0
|
||||
return session_data
|
||||
|
||||
|
||||
@@ -412,14 +364,11 @@ async def run_worker_sessions(
|
||||
config: dict[str, Any],
|
||||
) -> list[dict[str, Any]]:
|
||||
tasks = [
|
||||
asyncio.create_task(
|
||||
run_single_session(
|
||||
worker_id=worker_id,
|
||||
worker_session_idx=idx,
|
||||
config=config,
|
||||
)
|
||||
)
|
||||
for idx in range(session_count)
|
||||
asyncio.create_task(run_single_session(
|
||||
worker_id=worker_id,
|
||||
worker_session_idx=idx,
|
||||
config=config,
|
||||
)) for idx in range(session_count)
|
||||
]
|
||||
if not tasks:
|
||||
return []
|
||||
@@ -437,29 +386,23 @@ def worker_entry(
|
||||
try:
|
||||
ready_queue.put({"worker_id": worker_id, "status": "ready"})
|
||||
start_event.wait()
|
||||
sessions = asyncio.run(
|
||||
run_worker_sessions(
|
||||
worker_id=worker_id,
|
||||
session_count=session_count,
|
||||
config=config,
|
||||
)
|
||||
)
|
||||
result_queue.put(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"status": "ok",
|
||||
"sessions": sessions,
|
||||
}
|
||||
)
|
||||
sessions = asyncio.run(run_worker_sessions(
|
||||
worker_id=worker_id,
|
||||
session_count=session_count,
|
||||
config=config,
|
||||
))
|
||||
result_queue.put({
|
||||
"worker_id": worker_id,
|
||||
"status": "ok",
|
||||
"sessions": sessions,
|
||||
})
|
||||
except Exception as exc:
|
||||
result_queue.put(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"status": "error",
|
||||
"error": str(exc),
|
||||
"traceback": traceback.format_exc(),
|
||||
}
|
||||
)
|
||||
result_queue.put({
|
||||
"worker_id": worker_id,
|
||||
"status": "error",
|
||||
"error": str(exc),
|
||||
"traceback": traceback.format_exc(),
|
||||
})
|
||||
|
||||
|
||||
def build_summary(
|
||||
@@ -517,33 +460,22 @@ def build_summary(
|
||||
if len(all_chunk_finish_epochs) >= 2 and total_chunk_bytes > 0:
|
||||
duration_s = max(all_chunk_finish_epochs) - min(all_chunk_finish_epochs)
|
||||
if duration_s > 0:
|
||||
global_goodput_mbps = (
|
||||
total_chunk_bytes * 8.0 / duration_s / 1_000_000.0
|
||||
)
|
||||
global_goodput_mbps = (total_chunk_bytes * 8.0 / duration_s / 1_000_000.0)
|
||||
|
||||
bucket_throughputs_mbps = [
|
||||
(bytes_count * 8.0) / 1_000_000.0
|
||||
for _, bytes_count in sorted(bucket_bytes.items())
|
||||
]
|
||||
bucket_throughputs_mbps = [(bytes_count * 8.0) / 1_000_000.0 for _, bytes_count in sorted(bucket_bytes.items())]
|
||||
bucket_stats = summarize_series(bucket_throughputs_mbps)
|
||||
|
||||
chunk_gap_threshold_breaches = [
|
||||
value for value in chunk_gaps if value >= chunk_gap_threshold_ms
|
||||
]
|
||||
chunk_gap_threshold_breaches = [value for value in chunk_gaps if value >= chunk_gap_threshold_ms]
|
||||
non_success = len(sessions) - status_counts.get("success", 0)
|
||||
|
||||
fail_reasons: list[str] = []
|
||||
if non_success > 0:
|
||||
fail_reasons.append(
|
||||
f"{non_success} session(s) did not complete successfully."
|
||||
)
|
||||
fail_reasons.append(f"{non_success} session(s) did not complete successfully.")
|
||||
if not chunk_gaps:
|
||||
fail_reasons.append("No chunk gap data collected.")
|
||||
if chunk_gap_threshold_breaches:
|
||||
fail_reasons.append(
|
||||
f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
|
||||
f"{chunk_gap_threshold_ms:.0f}ms."
|
||||
)
|
||||
fail_reasons.append(f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
|
||||
f"{chunk_gap_threshold_ms:.0f}ms.")
|
||||
|
||||
passed = len(fail_reasons) == 0
|
||||
progressive_ratio = None
|
||||
@@ -554,20 +486,18 @@ def build_summary(
|
||||
"passed": passed,
|
||||
"fail_reasons": fail_reasons,
|
||||
"sessions": {
|
||||
"total": len(sessions),
|
||||
"success": status_counts.get("success", 0),
|
||||
"failed": status_counts.get("failed", 0),
|
||||
"timeout": status_counts.get("timeout", 0),
|
||||
"protocol_error": status_counts.get("protocol_error", 0),
|
||||
"other": (
|
||||
len(sessions)
|
||||
- (
|
||||
status_counts.get("success", 0)
|
||||
+ status_counts.get("failed", 0)
|
||||
+ status_counts.get("timeout", 0)
|
||||
+ status_counts.get("protocol_error", 0)
|
||||
)
|
||||
),
|
||||
"total":
|
||||
len(sessions),
|
||||
"success":
|
||||
status_counts.get("success", 0),
|
||||
"failed":
|
||||
status_counts.get("failed", 0),
|
||||
"timeout":
|
||||
status_counts.get("timeout", 0),
|
||||
"protocol_error":
|
||||
status_counts.get("protocol_error", 0),
|
||||
"other": (len(sessions) - (status_counts.get("success", 0) + status_counts.get("failed", 0) +
|
||||
status_counts.get("timeout", 0) + status_counts.get("protocol_error", 0))),
|
||||
},
|
||||
"chunk_gap_ms": {
|
||||
**chunk_gap_stats,
|
||||
@@ -606,51 +536,39 @@ def print_summary(
|
||||
bucket_bw = bandwidth["bucketed_1s"]
|
||||
|
||||
print("=== LTX2 Realtime Stress Test Summary ===")
|
||||
print(
|
||||
"Run: "
|
||||
f"url={run_info['url']} clients={run_info['clients']} "
|
||||
f"processes={run_info['processes']} "
|
||||
f"preset={run_info['preset_id']} "
|
||||
f"curated_limit={run_info['curated_limit']}"
|
||||
)
|
||||
print(
|
||||
"Sessions: "
|
||||
f"total={sessions['total']} success={sessions['success']} "
|
||||
f"failed={sessions['failed']} timeout={sessions['timeout']} "
|
||||
f"protocol_error={sessions['protocol_error']}"
|
||||
)
|
||||
print(
|
||||
"Chunk gap ms: "
|
||||
f"min={format_num(chunk_gap['min'])} "
|
||||
f"p50={format_num(chunk_gap['p50'])} "
|
||||
f"p95={format_num(chunk_gap['p95'])} "
|
||||
f"p99={format_num(chunk_gap['p99'])} "
|
||||
f"max={format_num(chunk_gap['max'])} "
|
||||
f"threshold={format_num(chunk_gap['threshold_ms'])} "
|
||||
f"breaches={chunk_gap['breach_count']}"
|
||||
)
|
||||
print(
|
||||
"Queue wait ms: "
|
||||
f"min={format_num(queue_wait['min'])} "
|
||||
f"p50={format_num(queue_wait['p50'])} "
|
||||
f"p95={format_num(queue_wait['p95'])} "
|
||||
f"max={format_num(queue_wait['max'])}"
|
||||
)
|
||||
print("Run: "
|
||||
f"url={run_info['url']} clients={run_info['clients']} "
|
||||
f"processes={run_info['processes']} "
|
||||
f"preset={run_info['preset_id']} "
|
||||
f"curated_limit={run_info['curated_limit']}")
|
||||
print("Sessions: "
|
||||
f"total={sessions['total']} success={sessions['success']} "
|
||||
f"failed={sessions['failed']} timeout={sessions['timeout']} "
|
||||
f"protocol_error={sessions['protocol_error']}")
|
||||
print("Chunk gap ms: "
|
||||
f"min={format_num(chunk_gap['min'])} "
|
||||
f"p50={format_num(chunk_gap['p50'])} "
|
||||
f"p95={format_num(chunk_gap['p95'])} "
|
||||
f"p99={format_num(chunk_gap['p99'])} "
|
||||
f"max={format_num(chunk_gap['max'])} "
|
||||
f"threshold={format_num(chunk_gap['threshold_ms'])} "
|
||||
f"breaches={chunk_gap['breach_count']}")
|
||||
print("Queue wait ms: "
|
||||
f"min={format_num(queue_wait['min'])} "
|
||||
f"p50={format_num(queue_wait['p50'])} "
|
||||
f"p95={format_num(queue_wait['p95'])} "
|
||||
f"max={format_num(queue_wait['max'])}")
|
||||
ratio = progressive["ratio"]
|
||||
ratio_text = "n/a" if ratio is None else f"{ratio * 100:.2f}%"
|
||||
print(
|
||||
"Progressive streaming: "
|
||||
f"{progressive['success_sessions']}/"
|
||||
f"{progressive['eligible_sessions']} ({ratio_text})"
|
||||
)
|
||||
print(
|
||||
"Bandwidth Mbps: "
|
||||
f"per_session_avg={format_num(per_session_bw['avg'])} "
|
||||
f"per_session_p95={format_num(per_session_bw['p95'])} "
|
||||
f"global={format_num(bandwidth['global_goodput_mbps'])} "
|
||||
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
|
||||
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}"
|
||||
)
|
||||
print("Progressive streaming: "
|
||||
f"{progressive['success_sessions']}/"
|
||||
f"{progressive['eligible_sessions']} ({ratio_text})")
|
||||
print("Bandwidth Mbps: "
|
||||
f"per_session_avg={format_num(per_session_bw['avg'])} "
|
||||
f"per_session_p95={format_num(per_session_bw['p95'])} "
|
||||
f"global={format_num(bandwidth['global_goodput_mbps'])} "
|
||||
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
|
||||
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}")
|
||||
print(f"VERDICT: {'PASS' if summary['passed'] else 'FAIL'}")
|
||||
if summary["fail_reasons"]:
|
||||
print("Fail reasons:")
|
||||
@@ -670,10 +588,8 @@ def distribute_sessions(total_clients: int, process_count: int) -> list[int]:
|
||||
|
||||
def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
if websockets is None:
|
||||
raise RuntimeError(
|
||||
"Missing dependency: websockets. Install it before running this "
|
||||
"stress test."
|
||||
)
|
||||
raise RuntimeError("Missing dependency: websockets. Install it before running this "
|
||||
"stress test.")
|
||||
|
||||
preset_file = Path(args.preset_file).expanduser().resolve()
|
||||
selected_preset_id, curated_prompts, total_prompt_count = load_curated_prompts(
|
||||
@@ -735,13 +651,8 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
start_event.set()
|
||||
|
||||
result_deadline = (
|
||||
time.monotonic()
|
||||
+ args.connect_timeout_s
|
||||
+ args.session_timeout_s
|
||||
+ args.post_complete_wait_s
|
||||
+ 180.0
|
||||
)
|
||||
result_deadline = (time.monotonic() + args.connect_timeout_s + args.session_timeout_s +
|
||||
args.post_complete_wait_s + 180.0)
|
||||
worker_results: list[dict[str, Any]] = []
|
||||
while len(worker_results) < len(processes):
|
||||
timeout_s = max(0.1, result_deadline - time.monotonic())
|
||||
@@ -765,24 +676,20 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
if result.get("status") == "ok":
|
||||
sessions.extend(result.get("sessions", []))
|
||||
else:
|
||||
worker_errors.append(
|
||||
{
|
||||
"worker_id": result.get("worker_id"),
|
||||
"error": result.get("error"),
|
||||
"traceback": result.get("traceback"),
|
||||
}
|
||||
)
|
||||
worker_errors.append({
|
||||
"worker_id": result.get("worker_id"),
|
||||
"error": result.get("error"),
|
||||
"traceback": result.get("traceback"),
|
||||
})
|
||||
|
||||
received_workers = {result.get("worker_id") for result in worker_results}
|
||||
expected_workers = set(range(len(processes)))
|
||||
missing_workers = sorted(expected_workers - received_workers)
|
||||
for worker_id in missing_workers:
|
||||
worker_errors.append(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"error": "No worker result received.",
|
||||
}
|
||||
)
|
||||
worker_errors.append({
|
||||
"worker_id": worker_id,
|
||||
"error": "No worker result received.",
|
||||
})
|
||||
|
||||
run_end_epoch = time.time()
|
||||
run_end_iso = iso_from_epoch(run_end_epoch)
|
||||
@@ -795,9 +702,8 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
if worker_errors:
|
||||
summary["passed"] = False
|
||||
summary["fail_reasons"] = list(summary["fail_reasons"]) + [
|
||||
f"{len(worker_errors)} worker error(s) occurred."
|
||||
]
|
||||
summary["fail_reasons"] = list(
|
||||
summary["fail_reasons"]) + [f"{len(worker_errors)} worker error(s) occurred."]
|
||||
|
||||
output_payload = {
|
||||
"run_info": {
|
||||
@@ -833,9 +739,7 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Multiprocess realtime stress test for LTX2 streaming.",
|
||||
)
|
||||
parser = argparse.ArgumentParser(description="Multiprocess realtime stress test for LTX2 streaming.", )
|
||||
parser.add_argument(
|
||||
"-u",
|
||||
"--url",
|
||||
|
||||
@@ -47,13 +47,11 @@ def test_persist_session_init_image_returns_none_when_missing_data():
|
||||
|
||||
def test_persist_session_init_image_rejects_unsupported_mime():
|
||||
with pytest.raises(ValueError, match="PNG, JPEG, or WebP"):
|
||||
persist_session_init_image(
|
||||
{
|
||||
"name": "frame.gif",
|
||||
"mime_type": "image/gif",
|
||||
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
|
||||
}
|
||||
)
|
||||
persist_session_init_image({
|
||||
"name": "frame.gif",
|
||||
"mime_type": "image/gif",
|
||||
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
|
||||
})
|
||||
|
||||
|
||||
def test_persist_session_init_image_rejects_large_payload(monkeypatch):
|
||||
@@ -66,10 +64,8 @@ def test_persist_session_init_image_rejects_large_payload(monkeypatch):
|
||||
monkeypatch.setattr(base64, "b64decode", fake_b64decode)
|
||||
|
||||
with pytest.raises(ValueError, match="15 MB or smaller"):
|
||||
persist_session_init_image(
|
||||
{
|
||||
"name": "frame.png",
|
||||
"mime_type": "image/png",
|
||||
"data_url": data_url,
|
||||
}
|
||||
)
|
||||
persist_session_init_image({
|
||||
"name": "frame.png",
|
||||
"mime_type": "image/png",
|
||||
"data_url": data_url,
|
||||
})
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -8,12 +8,10 @@ import modal
|
||||
|
||||
IMAGE = os.environ.get("DREAMVERSE_IMAGE")
|
||||
if not IMAGE:
|
||||
raise RuntimeError(
|
||||
"DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
|
||||
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
|
||||
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
|
||||
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag."
|
||||
)
|
||||
raise RuntimeError("DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
|
||||
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
|
||||
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
|
||||
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag.")
|
||||
|
||||
# ``@modal.web_server`` invokes ``serve()`` directly and bypasses the image
|
||||
# ENTRYPOINT (``docker/docker_entrypoint.sh``). That entrypoint normally
|
||||
@@ -65,14 +63,10 @@ def serve():
|
||||
# ``or ""`` collapses ``None`` (unset) into an empty string, ``.strip()``
|
||||
# collapses whitespace-only values (e.g. ``" "``) — both should be
|
||||
# treated as missing.
|
||||
missing = [
|
||||
k for k in _REQUIRED_SECRET_KEYS
|
||||
if not (os.environ.get(k) or "").strip()
|
||||
]
|
||||
missing = [k for k in _REQUIRED_SECRET_KEYS if not (os.environ.get(k) or "").strip()]
|
||||
if missing:
|
||||
raise RuntimeError(
|
||||
"dreamverse-api-keys secret is missing required entries: "
|
||||
f"{', '.join(missing)}. Add them with `modal secret create "
|
||||
"dreamverse-api-keys ... --force` and redeploy "
|
||||
"(see apps/dreamverse/scripts/modal/README.md).")
|
||||
raise RuntimeError("dreamverse-api-keys secret is missing required entries: "
|
||||
f"{', '.join(missing)}. Add them with `modal secret create "
|
||||
"dreamverse-api-keys ... --force` and redeploy "
|
||||
"(see apps/dreamverse/scripts/modal/README.md).")
|
||||
subprocess.Popen(["dreamverse-server", "--host", "0.0.0.0", "--port", "8009"])
|
||||
|
||||
@@ -74,10 +74,8 @@ def test_snapshot_shapes_devices(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
}
|
||||
|
||||
|
||||
def test_snapshot_tolerates_missing_sensors(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setitem(sys.modules, "pynvml",
|
||||
_make_fake_pynvml(broken_sensors=True))
|
||||
def test_snapshot_tolerates_missing_sensors(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setitem(sys.modules, "pynvml", _make_fake_pynvml(broken_sensors=True))
|
||||
snap = gpu_mod.get_gpu_snapshot()
|
||||
assert snap["available"] is True
|
||||
g = snap["gpus"][0]
|
||||
@@ -86,8 +84,7 @@ def test_snapshot_tolerates_missing_sensors(
|
||||
assert g["power_limit_watts"] is None
|
||||
|
||||
|
||||
def test_snapshot_reports_nvml_failure(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_snapshot_reports_nvml_failure(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
fake = _make_fake_pynvml()
|
||||
fake.nvmlInit = lambda: (_ for _ in ()).throw(_NVMLError("driver gone"))
|
||||
monkeypatch.setitem(sys.modules, "pynvml", fake)
|
||||
|
||||
@@ -130,8 +130,7 @@ def test_dmd_builds_three_role_models_and_method_knobs() -> None:
|
||||
|
||||
|
||||
def test_dmd_vsa_maps_to_training_vsa_sparsity() -> None:
|
||||
config = build_training_config(
|
||||
_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
|
||||
config = build_training_config(_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
|
||||
assert config["training"]["vsa"]["sparsity"] == 0.9
|
||||
|
||||
|
||||
@@ -161,31 +160,27 @@ def test_validation_callback_only_when_file_given() -> None:
|
||||
without = build_training_config(_job("full_t2v"), "out")
|
||||
assert "validation" not in without["callbacks"]
|
||||
|
||||
with_file = build_training_config(
|
||||
_job("full_t2v", validation_dataset_file="val.json"), "out")
|
||||
with_file = build_training_config(_job("full_t2v", validation_dataset_file="val.json"), "out")
|
||||
validation = with_file["callbacks"]["validation"]
|
||||
assert validation["dataset_file"] == "val.json"
|
||||
assert validation["pipeline_target"].endswith(".WanPipeline")
|
||||
assert validation["sampling_steps"] == [50]
|
||||
|
||||
dmd = build_training_config(
|
||||
_job("dmd_t2v", validation_dataset_file="val.json"), "out")
|
||||
dmd = build_training_config(_job("dmd_t2v", validation_dataset_file="val.json"), "out")
|
||||
validation = dmd["callbacks"]["validation"]
|
||||
assert validation["pipeline_target"].endswith(".WanDMDPipeline")
|
||||
assert validation["sampling_steps"] == [3]
|
||||
assert validation["sampling_timesteps"] == [1000, 757, 522]
|
||||
|
||||
# KD/ODE-init has no sampling-based validation pipeline.
|
||||
ode = build_training_config(
|
||||
_job("ode_init", validation_dataset_file="val.json"), "out")
|
||||
ode = build_training_config(_job("ode_init", validation_dataset_file="val.json"), "out")
|
||||
assert "validation" not in ode["callbacks"]
|
||||
|
||||
|
||||
def test_ltx2_models_are_rejected() -> None:
|
||||
assert is_ltx2_model("Lightricks/LTX-2-19B")
|
||||
with pytest.raises(ValueError, match="LTX-2 training is not supported"):
|
||||
build_training_config(
|
||||
_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
|
||||
build_training_config(_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
|
||||
|
||||
|
||||
def test_unknown_workload_is_rejected() -> None:
|
||||
@@ -195,8 +190,7 @@ def test_unknown_workload_is_rejected() -> None:
|
||||
|
||||
def test_invalid_denoising_steps_are_rejected() -> None:
|
||||
with pytest.raises(ValueError, match="Invalid DMD denoising steps"):
|
||||
build_training_config(
|
||||
_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
|
||||
build_training_config(_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
|
||||
|
||||
|
||||
def test_training_env_has_no_backend_override() -> None:
|
||||
@@ -211,8 +205,7 @@ def test_workloads_match_frontend_job_config() -> None:
|
||||
(src/lib/jobConfig.ts) — drift means creatable-but-unrunnable jobs."""
|
||||
import re
|
||||
|
||||
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" /
|
||||
"jobConfig.ts").read_text(encoding="utf-8")
|
||||
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" / "jobConfig.ts").read_text(encoding="utf-8")
|
||||
all_types = set(re.findall(r'type:\s*"([^"]+)"', job_config))
|
||||
inference_types = {"t2v", "i2v", "t2i"}
|
||||
assert inference_types <= all_types, "jobConfig.ts parse failed"
|
||||
|
||||
+1
-1
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
|
||||
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
|
||||
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
|
||||
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
|
||||
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
|
||||
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
|
||||
|
||||
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
|
||||
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
{
|
||||
"recipes": [
|
||||
{
|
||||
"id": "fastwan21-t2v",
|
||||
"task": "Text to video",
|
||||
"label": "FastWan2.1 1.3B (distilled + VSA)",
|
||||
"model": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"source": "scripts/inference/inference_wan_VSA_DMD_1_3B.yaml",
|
||||
"command": "FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml"
|
||||
},
|
||||
{
|
||||
"id": "wan22-t2v",
|
||||
"task": "Text to video",
|
||||
"label": "Wan2.2 A14B",
|
||||
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_wan2_2.py",
|
||||
"command": "python examples/inference/basic/basic_wan2_2.py"
|
||||
},
|
||||
{
|
||||
"id": "wan21-i2v",
|
||||
"task": "Image to video",
|
||||
"label": "Wan2.1 14B 480P",
|
||||
"model": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
"source": "scripts/inference/inference_wan_i2v.yaml",
|
||||
"command": "fastvideo generate --config scripts/inference/inference_wan_i2v.yaml"
|
||||
},
|
||||
{
|
||||
"id": "turbowan22-i2v",
|
||||
"task": "Image to video",
|
||||
"label": "TurboWan2.2 A14B",
|
||||
"model": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_turbodiffusion_i2v.py",
|
||||
"command": "python examples/inference/basic/basic_turbodiffusion_i2v.py"
|
||||
},
|
||||
{
|
||||
"id": "wan22-ti2v",
|
||||
"task": "Text or image to video",
|
||||
"label": "Wan2.2 TI2V 5B",
|
||||
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_wan2_2_ti2v.py",
|
||||
"command": "python examples/inference/basic/basic_wan2_2_ti2v.py"
|
||||
},
|
||||
{
|
||||
"id": "matrix-game-2",
|
||||
"task": "Interactive world",
|
||||
"label": "Matrix Game 2.0",
|
||||
"model": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"source": "examples/inference/basic/basic_matrixgame2.py",
|
||||
"command": "python examples/inference/basic/basic_matrixgame2.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
(() => {
|
||||
let recipesPromise;
|
||||
|
||||
const loadRecipes = (url) => {
|
||||
recipesPromise ||= fetch(url).then((response) => {
|
||||
if (!response.ok) throw new Error(`HTTP ${response.status}`);
|
||||
return response.json();
|
||||
});
|
||||
return recipesPromise;
|
||||
};
|
||||
|
||||
const init = () => {
|
||||
document.querySelectorAll("[data-cookbook]").forEach(async (root) => {
|
||||
if (root.dataset.initialized) return;
|
||||
root.dataset.initialized = "true";
|
||||
|
||||
const select = root.querySelector("[data-cookbook-recipe]");
|
||||
const model = root.querySelector("[data-cookbook-model]");
|
||||
const source = root.querySelector("[data-cookbook-source]");
|
||||
const command = root.querySelector("[data-cookbook-command]");
|
||||
const status = root.querySelector("[data-cookbook-status]");
|
||||
|
||||
try {
|
||||
const { recipes } = await loadRecipes(root.dataset.recipes);
|
||||
const byId = new Map(recipes.map((recipe) => [recipe.id, recipe]));
|
||||
const groups = new Map();
|
||||
|
||||
select.replaceChildren();
|
||||
recipes.forEach((recipe) => {
|
||||
if (!groups.has(recipe.task)) {
|
||||
const group = document.createElement("optgroup");
|
||||
group.label = recipe.task;
|
||||
groups.set(recipe.task, group);
|
||||
select.append(group);
|
||||
}
|
||||
groups.get(recipe.task).append(new Option(recipe.label, recipe.id));
|
||||
});
|
||||
|
||||
const render = () => {
|
||||
const recipe = byId.get(select.value);
|
||||
model.textContent = recipe.model;
|
||||
source.textContent = recipe.source;
|
||||
source.href = `https://github.com/hao-ai-lab/FastVideo/blob/main/${recipe.source}`;
|
||||
command.textContent = recipe.command;
|
||||
status.textContent = `${recipe.label} selected.`;
|
||||
};
|
||||
|
||||
select.addEventListener("change", render);
|
||||
select.disabled = false;
|
||||
render();
|
||||
} catch (error) {
|
||||
status.textContent = "Recipes could not be loaded. Use the examples link below.";
|
||||
console.error("Failed to load FastVideo cookbook recipes", error);
|
||||
}
|
||||
});
|
||||
};
|
||||
|
||||
if (window.document$) window.document$.subscribe(init);
|
||||
else document.addEventListener("DOMContentLoaded", init);
|
||||
})();
|
||||
@@ -42,6 +42,46 @@ img {
|
||||
margin: 0 auto;
|
||||
}
|
||||
|
||||
.cookbook-picker {
|
||||
padding: 1rem;
|
||||
border: 0.05rem solid var(--md-default-fg-color--lightest);
|
||||
border-radius: 0.2rem;
|
||||
}
|
||||
|
||||
.cookbook-picker select {
|
||||
width: 100%;
|
||||
padding: 0.6rem;
|
||||
color: var(--md-default-fg-color);
|
||||
background: var(--md-default-bg-color);
|
||||
border: 0.05rem solid var(--md-default-fg-color--lighter);
|
||||
border-radius: 0.2rem;
|
||||
}
|
||||
|
||||
.cookbook-picker dl {
|
||||
display: grid;
|
||||
grid-template-columns: max-content 1fr;
|
||||
gap: 0.25rem 1rem;
|
||||
}
|
||||
|
||||
.cookbook-picker dt {
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.cookbook-picker dd {
|
||||
margin: 0;
|
||||
min-width: 0;
|
||||
overflow-wrap: anywhere;
|
||||
}
|
||||
|
||||
.cookbook-picker__status {
|
||||
position: absolute;
|
||||
width: 1px;
|
||||
height: 1px;
|
||||
overflow: hidden;
|
||||
clip: rect(0, 0, 0, 0);
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.md-typeset .copy-page-button.md-button {
|
||||
float: right;
|
||||
margin: 0 0 1rem 1rem;
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# Inference Cookbook
|
||||
|
||||
Choose a complete recipe maintained in the FastVideo repository. Each command
|
||||
runs its checked-in source directly, so coupled model, GPU, offload, and
|
||||
attention settings do not drift into unsupported combinations.
|
||||
|
||||
The commands expect a local clone:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git
|
||||
cd FastVideo
|
||||
```
|
||||
|
||||
<div class="cookbook-picker" data-cookbook data-recipes="../assets/cookbook-recipes.json">
|
||||
<label for="cookbook-recipe"><strong>Recipe</strong></label>
|
||||
<select id="cookbook-recipe" data-cookbook-recipe disabled>
|
||||
<option>Loading recipes…</option>
|
||||
</select>
|
||||
<dl>
|
||||
<dt>Model</dt>
|
||||
<dd data-cookbook-model>Loading…</dd>
|
||||
<dt>Source</dt>
|
||||
<dd><a data-cookbook-source href="../inference/examples/basic/">Browse maintained examples</a></dd>
|
||||
</dl>
|
||||
<pre><code class="language-bash" data-cookbook-command>Loading…</code></pre>
|
||||
<p class="cookbook-picker__status" role="status" aria-live="polite" data-cookbook-status></p>
|
||||
<noscript>
|
||||
JavaScript is needed for the recipe picker. Browse the
|
||||
<a href="../inference/examples/examples_inference_index/">inference examples</a>
|
||||
instead.
|
||||
</noscript>
|
||||
</div>
|
||||
|
||||
## Customize a recipe
|
||||
|
||||
Start from the checked-in source, then change only the settings your model
|
||||
supports:
|
||||
|
||||
- [Configuration](../inference/configuration.md) covers the Python and CLI
|
||||
config surfaces.
|
||||
- [Optimizations](../inference/optimizations.md) covers attention backends,
|
||||
compilation, and memory tradeoffs.
|
||||
- [Support matrix](../inference/support_matrix.md) lists supported models and
|
||||
optimizations.
|
||||
@@ -0,0 +1,128 @@
|
||||
# Fast mode (RIFE) — Apple Silicon
|
||||
|
||||
`--fast` makes local generation ~2.7× faster by **generating fewer frames and
|
||||
interpolating the rest** with an Apple-Silicon-native RIFE model, instead of
|
||||
denoising every frame. Video-diffusion denoise is dominated by self-attention,
|
||||
which is O(tokens²); halving the frames cuts the token count ~2× and the denoise
|
||||
compute ~3.7×, so the wall-clock drops far more than 2×. RIFE (which estimates
|
||||
its own optical flow — no motion vectors needed) fills the dropped frames back
|
||||
in for ~1.4 s, and a light unsharp pass counters its softening.
|
||||
|
||||
Measured on the 1.3B INT8 QAD model (fox, 480×832×81, M4): generate 41 + RIFE→81
|
||||
runs in ~35 s of denoise vs ~90 s full, at reconstruction MS-SSIM **0.97**.
|
||||
Reproduce with `python -m fastvideo.benchmarks.eval_metalfx_rife --mode int8`.
|
||||
|
||||
> **Note:** Apple's *MetalFX* frame interpolation is **not** usable here — it
|
||||
> requires game-engine motion vectors + depth, which diffusion output lacks. We
|
||||
> use the video-native **`rife-mlx`** model instead (Metal-backed, torch-free).
|
||||
|
||||
## Install
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[mlx]" # RIFE ships vendored; this only needs MLX
|
||||
```
|
||||
|
||||
## Use
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/mlx_wan_prompt_to_video.py \
|
||||
--mlx-checkpoint <FastWan2.1-T2V-1.3B-INT8-QAD> \
|
||||
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
|
||||
--num-frames 81 --fast \
|
||||
--output-path video_samples/fox_fast.mp4
|
||||
```
|
||||
|
||||
`--num-frames` stays the *target* length; fast mode generates the smallest
|
||||
VAE-aligned keyframe count that RIFE can interpolate to that target.
|
||||
|
||||
| Flag | Default | Meaning |
|
||||
|---|---|---|
|
||||
| `--fast` / `--no-fast` | off | enable fast mode |
|
||||
| `--fast-factor` | 2 | generate 1/factor of the frames (2 = half) |
|
||||
| `--fast-sharpen` | 0.6 | light unsharp strength to counter RIFE softness (0 disables) |
|
||||
|
||||
Fast mode composes with everything else (`--mlx-quantization int8`,
|
||||
`--mlx-compile`, TAEHV vs `--decode-backend wan-vae`). Keep `--fast-factor` at 2
|
||||
for quality — larger temporal gaps are where RIFE starts inventing motion.
|
||||
|
||||
## Spatial fast mode (`--fast-spatial`)
|
||||
|
||||
The spatial twin of `--fast`: instead of dropping frames, drop pixels. Denoise
|
||||
*and decode* at `height/width // fast-spatial-scale`, then resample the decoded
|
||||
frames up to the requested size. Self-attention is O(tokens²), so halving each
|
||||
spatial axis cuts the token count 4× and the denoise time far more than that —
|
||||
measured on the 1.3B INT8 QAD model at 480×832×81, M4 Max: **86.1 s → 10.3 s**
|
||||
of denoise. It composes with `--fast`; both together run the same clip in
|
||||
**4.5 s** of denoise.
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/mlx_wan_prompt_to_video.py \
|
||||
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
|
||||
--height 480 --width 832 --num-frames 81 --fast-spatial \
|
||||
--output-path video_samples/fox_fast_spatial.mp4
|
||||
```
|
||||
|
||||
| Flag | Default | Meaning |
|
||||
|---|---|---|
|
||||
| `--fast-spatial` / `--no-fast-spatial` | off | enable spatial fast mode |
|
||||
| `--fast-spatial-scale` | 2 | denoise at 1/scale of each spatial axis |
|
||||
| `--fast-spatial-upsample-mode` | `lanczos` | pixel interpolation kernel (`lanczos`, `cubic`, `bilinear`, `nearest`) |
|
||||
| `--fast-spatial-sharpen` | 0.4 | light unsharp strength to counter resampling softness (0 disables) |
|
||||
|
||||
### The upsample must happen in pixel space
|
||||
|
||||
This is the one thing to get right. The obvious implementation — bilinearly
|
||||
upsample the finished latents and decode at the target size — **does not work**,
|
||||
and produces a distinctive failure: correct composition and silhouette under a
|
||||
smeared, hazy veil, with ringing along strong edges.
|
||||
|
||||
A Wan latent cell is a *learned code* for an 8×8 (Wan2.1) or 16×16 (Wan2.2)
|
||||
pixel block, not a low-pass sample of the image. The average of two adjacent
|
||||
codes is not the code of the averaged blocks; it is a vector the decoder was
|
||||
never trained on. Measured on Wan2.1-1.3B at 480×832, a 2× bilinear latent
|
||||
upsample destroys **62%** of the latent's high-frequency energy while leaving
|
||||
its overall magnitude intact — exactly the signature of that veil. At Wan2.2-5B
|
||||
the same operation degrades to black or noise.
|
||||
|
||||
Decoded RGB frames have no such problem: an image *is* a sampled 2-D signal, so
|
||||
Lanczos interpolation is the operation it was defined for. The result is soft —
|
||||
it carries stage-1's real detail budget and no more — but clean and coherent.
|
||||
|
||||
`--refine` gets away with a latent-space upsample only because a second DMD pass
|
||||
re-denoises the hand-off; spatial fast mode passes the latent straight to the
|
||||
decoder, so it cannot.
|
||||
|
||||
## Refine (`--refine`) stage-2 timesteps
|
||||
|
||||
`--refine` hands stage 1 to stage 2 as `(1 - sigma) * upsampled + sigma * noise`,
|
||||
where `sigma` comes from the *first* stage-2 timestep. FastWan's DMD grid opens
|
||||
at `t=1000`, which is `sigma == 1` exactly — so a stage-2 grid that starts there
|
||||
weights the stage-1 result at zero and refine silently degrades into a plain
|
||||
full-resolution run at twice the cost.
|
||||
|
||||
Left unset, `--refine-dmd-denoising-steps` now derives the stage-2 grid from the
|
||||
stage-1 one with leading full-noise steps dropped (`1000,757,522` → `757,522`).
|
||||
That keeps the pass on timesteps the distilled student was trained on while
|
||||
letting stage-1 structure through: hand-off `sigma = 0.757`, stage-1 weight
|
||||
`0.243`. Passing a grid that starts at full noise is now an error rather than a
|
||||
silently wasted pass.
|
||||
|
||||
The run prints the resolved hand-off so it is visible:
|
||||
|
||||
```
|
||||
[refine] stage-2 hand-off sigma=0.7568 (stage-1 weight 0.2432)
|
||||
```
|
||||
|
||||
There is a trade-off in choosing that grid. Later start = more of the draft
|
||||
survives, but fewer stage-2 steps. On Wan2.1 the default `757,522` gives weight
|
||||
0.243 with two steps; `--refine-dmd-denoising-steps 522` gives weight 0.478 with
|
||||
one. `--refine-sigma` decouples the noise level from the timestep entirely — it
|
||||
logs a warning, because the DiT is then told a timestep that does not match the
|
||||
noise it receives.
|
||||
|
||||
**Wan2.2-5B has a lower ceiling.** Its warped schedule maps `1000,757,522` to
|
||||
sigmas `1.000, 0.940, 0.845`, so the best available stage-1 weight is **0.060**
|
||||
(vs 0.243 at 1.3B). Un-warped (`--no-warp`) the same grid gives `1.000, 0.757,
|
||||
0.522` and a weight of 0.243 — but warping is what matches the FastVideo
|
||||
sampling schedule, so turning it off changes the timesteps the distilled student
|
||||
sees. Which is better at 5B is unresolved and needs a run on real 5B weights.
|
||||
@@ -76,6 +76,10 @@ surfaces:
|
||||
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
|
||||
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
|
||||
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
|
||||
VSA_tile_size: "VSA-H3 tile geometry request; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_decode: "MiniMax-H3 sequence-parallel VAE decode opt-in; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_encode: "MiniMax-H3 sequence-parallel reference VAE encode opt-in; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_decode_strategy: "Chunk-transport collective for vae_parallel_decode; model-specific optimization not yet represented in the typed public schema."
|
||||
attention_backend: "Process-wide default attention-backend request applied per component at load time; kernel-selection knob not yet represented in the typed public schema."
|
||||
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
|
||||
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
|
||||
|
||||
import itertools
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
@@ -19,6 +20,40 @@ GENERATED_DOC_PREFIXES = (
|
||||
"training/examples/",
|
||||
"distillation/examples/",
|
||||
)
|
||||
COOKBOOK_DATA = ROOT_DIR / "docs/assets/cookbook-recipes.json"
|
||||
COOKBOOK_SOURCE_ROOTS = (
|
||||
ROOT_DIR / "examples/inference",
|
||||
ROOT_DIR / "scripts/inference",
|
||||
)
|
||||
|
||||
|
||||
def validate_cookbook() -> None:
|
||||
"""Keep cookbook entries tied to checked-in runnable sources."""
|
||||
recipes = json.loads(COOKBOOK_DATA.read_text(encoding="utf-8")).get("recipes")
|
||||
if not isinstance(recipes, list) or not recipes:
|
||||
raise ValueError(f"{COOKBOOK_DATA}: recipes must be a non-empty list")
|
||||
|
||||
seen: set[str] = set()
|
||||
for recipe in recipes:
|
||||
required = ("id", "task", "label", "model", "source", "command")
|
||||
missing = {key for key in required if not recipe.get(key)}
|
||||
if missing:
|
||||
raise ValueError(f"Cookbook recipe is missing: {', '.join(sorted(missing))}")
|
||||
if recipe["id"] in seen:
|
||||
raise ValueError(f"Duplicate cookbook recipe id: {recipe['id']}")
|
||||
seen.add(recipe["id"])
|
||||
|
||||
source = (ROOT_DIR / recipe["source"]).resolve()
|
||||
if not any(source.is_relative_to(root.resolve()) for root in COOKBOOK_SOURCE_ROOTS):
|
||||
raise ValueError(f"Cookbook source is outside an approved directory: {recipe['source']}")
|
||||
if not source.is_file():
|
||||
raise ValueError(f"Cookbook source does not exist: {recipe['source']}")
|
||||
|
||||
source_text = source.read_text(encoding="utf-8")
|
||||
if recipe["model"] not in source_text:
|
||||
raise ValueError(f"Cookbook model is not present in {recipe['source']}: {recipe['model']}")
|
||||
if recipe["source"] not in recipe["command"]:
|
||||
raise ValueError(f"Cookbook command does not invoke its source: {recipe['id']}")
|
||||
|
||||
|
||||
def fix_case(text: str) -> str:
|
||||
@@ -536,6 +571,7 @@ def on_pre_build(config, **kwargs):
|
||||
MkDocs hook to generate examples before building the documentation.
|
||||
This function is called automatically by MkDocs' native hook system.
|
||||
"""
|
||||
validate_cookbook()
|
||||
print("Generating example documentation...")
|
||||
generate_examples(generate_main_index=True)
|
||||
print("Example documentation generated successfully!")
|
||||
@@ -549,6 +585,7 @@ def on_page_context(context, page, **kwargs):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
validate_cookbook()
|
||||
print("Generating example documentation...")
|
||||
generate_examples(generate_main_index=True)
|
||||
print("Example documentation generated successfully!")
|
||||
|
||||
@@ -65,6 +65,7 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
|
||||
- [Quick Start](quick_start.md) - Generate your first video
|
||||
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
|
||||
|
||||
@@ -49,10 +49,12 @@ brew install ffmpeg
|
||||
|
||||
### Installation
|
||||
|
||||
FastWan's native Apple Silicon runtime requires the `mlx` extra.
|
||||
|
||||
#### With uv (recommended)
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
uv pip install "fastvideo[mlx]"
|
||||
```
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
@@ -60,7 +62,7 @@ uv pip install fastvideo
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
uv pip install "fastvideo[mlx]"
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -76,13 +78,13 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
uv pip install -e ".[mlx]"
|
||||
```
|
||||
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
uv pip install -e ".[mlx]"
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
@@ -23,61 +23,21 @@ Also optionally install flash-attn:
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Basic Usage
|
||||
## Choose a maintained recipe
|
||||
|
||||
### Text-to-Video Generation
|
||||
The cookbook selects complete, checked-in recipes instead of mixing model,
|
||||
parallelism, offload, and attention settings independently.
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
|
||||
|
||||
def main():
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
|
||||
# Define a prompt for your video
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
|
||||
|
||||
# Generate the video
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path="my_videos/", # Controls where videos are saved
|
||||
save_video=True
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
### Image-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
def main():
|
||||
# Create the generator
|
||||
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
|
||||
|
||||
# Set up parameters with an initial image
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
generator.generate_video(prompt, sampling_param=sampling_param,
|
||||
output_path="my_videos/",
|
||||
save_video=True)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
!!! tip "Need more control?"
|
||||
Start from a maintained recipe, then use the
|
||||
[configuration](../inference/configuration.md) and
|
||||
[optimization](../inference/optimizations.md) guides for supported changes.
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
|
||||
- [Installation Guide](installation.md) - Detailed installation instructions
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
|
||||
|
||||
@@ -64,6 +64,7 @@ column links a runnable script in `examples/inference/basic/` where one exists.
|
||||
| ltx2 | `FastVideo/LTX2-Distilled-Diffusers`<br>`FastVideo/LTX2.3-Distilled-Diffusers`<br>`FastVideo/LTX-2.3-Distilled-Diffusers` | T2V | [basic_ltx2_distilled.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2_distilled.py) |
|
||||
| ltx2 | `Lightricks/LTX-2.3`<br>`FastVideo/LTX2.3-base`<br>`FastVideo/LTX2.3-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
|
||||
| ltx2 | `Lightricks/LTX-2`<br>`FastVideo/LTX2-base`<br>`FastVideo/LTX2-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
|
||||
| mmaudio | `FastVideo/MMAudio-large-44k-v2-Diffusers` | V2A, T2A | [basic_mmaudio.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mmaudio.py) |
|
||||
| matrixgame | `FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-Base-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Diffusers`<br>`mignonjia/mg_longtuning_distilled_zelda`<br>`mignonjia/mg_sf_distilled_zelda_1k_steps`<br>`mignonjia/mg_sf_distilled_zelda`<br>`mignonjia/mg_causal_zelda`<br>`mignonjia/mg_bidirectional_zelda` | I2V | [basic_matrixgame2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame2.py) |
|
||||
| matrixgame | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | I2V | [basic_matrixgame3.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame3.py) |
|
||||
| minimax_h3 | `MiniMaxAI/MiniMax-H3` | T2V, I2V | [T2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_t2v.py)<br>[FL2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_fl2va.py)<br>[Ref2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_ref2va.py) |
|
||||
@@ -94,6 +95,10 @@ column links a runnable script in `examples/inference/basic/` where one exists.
|
||||
(`StableAudioT2AConfig` / `StableAudioOpenSmallConfig`); they are registered
|
||||
under the generic T2V workload option in the registry.
|
||||
|
||||
**Note (MMAudio)**: the registered Hugging Face model ID is reserved but not
|
||||
yet public. Follow the [MMAudio inference guide](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/pipelines/basic/mmaudio/README.md)
|
||||
to convert the official weights locally and set `MMAUDIO_MODEL_PATH`.
|
||||
|
||||
**Note (MiniMax H3)**: T2VA, FL2VA, and Ref2VA all generate video with stereo
|
||||
audio. Use the Ref2VA example when passing ordered image, video, or audio
|
||||
references.
|
||||
@@ -173,6 +178,17 @@ optimizations: absence means **untested**, not incompatible.
|
||||
| Matrix Game 3.0 Base Distilled | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | 720x1280 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| GEN3C Cosmos 7B | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | 704px1280p | ❌ | ❌ | ❌ | ⭕ | ⭕ |
|
||||
|
||||
## Apple Silicon native runtime
|
||||
|
||||
| Release path | Model | Mode | Validated hardware | Status |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| MLX FastWan T2V | FastWan-QAD-INT8-1.3B `[release model ID pending]` | 480x832, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 36 GB unified-memory class, MLX 0.31.2 | Release candidate; requires release-owner visual sign-off |
|
||||
|
||||
This is a text-to-video-only source-install release. It is validated on the
|
||||
hardware listed above; MLX allocator caps are not evidence of support for a
|
||||
physical 16 GB Mac. See [Apple Silicon FastWan](../getting_started/installation/mps.md)
|
||||
for the supported command and release gates.
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
***Lucy Edit Dev uses a non-commercial model license. FastVideo support is
|
||||
|
||||
@@ -33,6 +33,14 @@ For the typed config/request path added during the inference API refactor:
|
||||
python examples/inference/basic/basic_dmd_new_api.py
|
||||
```
|
||||
|
||||
For the few-step (4-step, DMD2-distilled) MiniMax-H3 preview, generating synchronized video and audio, optionally with block-sparse VSA attention:
|
||||
```
|
||||
python examples/inference/basic/basic_fasth3.py --prompt "your prompt" [--vsa-sparsity 0.9]
|
||||
```
|
||||
The default checkpoint `FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1` is private on the Hub while its license review completes; until it flips public, pass `--model-path` with a local snapshot of the release.
|
||||
|
||||
On Blackwell (sm_100) GPUs with a `fastvideo-kernel` build that carries the sm_100a block-sparse extension, `--vsa-kernel sm100a` routes the tile-64 attention forwards through the CUDA kernel instead of Triton (it sets `FASTVIDEO_VSA_SM100A=1` before the pipeline boots); if the extension or the arch is missing, the run warns once and falls back to Triton.
|
||||
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -12,11 +14,11 @@ def main():
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
@@ -24,22 +26,19 @@ def main():
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
|
||||
@@ -30,8 +30,7 @@ def main():
|
||||
"and casting reflections onto adjacent vehicles. "
|
||||
"The motion creates space in the lineup, signaling activity within the otherwise quiet station. "
|
||||
"It then comes to a smooth stop, resuming its position in line. "
|
||||
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene."
|
||||
)
|
||||
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene.")
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
@@ -47,4 +46,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
@@ -31,8 +31,7 @@ def main():
|
||||
"The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. "
|
||||
"The metal surface beneath the torch shows ongoing signs of heating and melting. "
|
||||
"The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, "
|
||||
"underscoring the ongoing nature of the welding operation."
|
||||
)
|
||||
"underscoring the ongoing nature of the welding operation.")
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
@@ -46,6 +45,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -50,4 +50,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
@@ -5,6 +5,8 @@ from fastvideo import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2"
|
||||
|
||||
|
||||
def main():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
@@ -14,10 +16,10 @@ def main():
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
@@ -25,7 +27,6 @@ def main():
|
||||
load_end_time = time.perf_counter()
|
||||
load_time = load_end_time - load_start_time
|
||||
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.num_frames = 81
|
||||
|
||||
@@ -39,18 +40,16 @@ def main():
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
start_time = time.perf_counter()
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
|
||||
end_time = time.perf_counter()
|
||||
gen_time2 = end_time - start_time
|
||||
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Time taken to generate video2: {gen_time2} seconds")
|
||||
|
||||
@@ -32,11 +32,9 @@ def main():
|
||||
),
|
||||
# PR 2 still routes a few advanced inference knobs through the
|
||||
# compatibility bridge until they get first-class typed fields.
|
||||
pipeline=PipelineSelection(
|
||||
experimental={
|
||||
"VSA_sparsity": 0.8,
|
||||
},
|
||||
),
|
||||
pipeline=PipelineSelection(experimental={
|
||||
"VSA_sparsity": 0.8,
|
||||
}, ),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
@@ -44,14 +42,12 @@ def main():
|
||||
load_end_time = time.perf_counter()
|
||||
load_time = load_end_time - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
|
||||
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
|
||||
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
|
||||
"LED umbrella. Steam rises from a street food cart, and a cat darts "
|
||||
"across the screen. Raindrops are visible on the camera lens, creating "
|
||||
"a cinematic bokeh effect."
|
||||
)
|
||||
prompt = ("A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
|
||||
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
|
||||
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
|
||||
"LED umbrella. Steam rises from a street food cart, and a cat darts "
|
||||
"across the screen. Raindrops are visible on the camera lens, creating "
|
||||
"a cinematic bokeh effect.")
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
output=OutputConfig(
|
||||
@@ -66,13 +62,11 @@ def main():
|
||||
end_time = time.perf_counter()
|
||||
gen_time = end_time - start_time
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently "
|
||||
"in the breeze, enhancing the lion's commanding presence. The tone is "
|
||||
"vibrant, embodying the raw energy of the wild. Low angle, steady "
|
||||
"tracking shot, cinematic."
|
||||
)
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently "
|
||||
"in the breeze, enhancing the lion's commanding presence. The tone is "
|
||||
"vibrant, embodying the raw energy of the wild. Low angle, steady "
|
||||
"tracking shot, cinematic.")
|
||||
request2 = GenerationRequest(
|
||||
prompt=prompt2,
|
||||
output=OutputConfig(
|
||||
|
||||
@@ -2,7 +2,6 @@ import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
|
||||
|
||||
|
||||
@@ -46,10 +45,8 @@ def main():
|
||||
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
|
||||
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
|
||||
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
|
||||
"action_speed_list": [
|
||||
float(value)
|
||||
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
|
||||
],
|
||||
"action_speed_list":
|
||||
[float(value) for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")],
|
||||
}
|
||||
if image_path:
|
||||
kwargs["image_path"] = image_path
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
|
||||
|
||||
FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1 is a 4-step distillation of
|
||||
MiniMaxAI/MiniMax-H3 (data-free DMD2): it walks a 4-step grid on the release
|
||||
sampler's shift-12 schedule instead of the base model's 50 steps, generating
|
||||
synchronized video and audio in one pipeline call.
|
||||
|
||||
The student was trained with block-sparse video attention (VSA, 64-token
|
||||
tiles) and its checkpoint carries the trained sparse-gate parameters
|
||||
(``attn.to_gate_compress``), so this script always runs the VSA-H3 attention
|
||||
backend. At the default ``--vsa-sparsity 0.0`` the attention math is exactly
|
||||
dense (every tile is selected); raise the sparsity for additional speedup.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
CompileConfig,
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1")
|
||||
# The HF repo is private while the MiniMax H3 Community License review
|
||||
# completes; until it flips public, pass --model-path with a local
|
||||
# snapshot of the release instead (e.g. the team export at
|
||||
# /mnt/lustre/vlm-wlsaidhi/fastvideo/exports/FastVideo-Minimax-FastH3-Preview-v0.1).
|
||||
parser.add_argument("--prompt", required=True)
|
||||
parser.add_argument("--output", default="outputs/fasth3")
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
parser.add_argument("--width", type=int, default=1344)
|
||||
parser.add_argument("--num-frames", type=int, default=124)
|
||||
# num_inference_steps counts sigma-GRID POINTS, matching the base model's
|
||||
# convention ("50 steps" = a 50-point grid = 49 transformer forwards). The
|
||||
# student's distilled 4-step grid is 4 FORWARDS, i.e. a 5-point grid
|
||||
# (t = 1000, 750, 500, 250 -> 0 on the shift-12 schedule) — so the correct
|
||||
# default here is 5. Other grids are off-distribution.
|
||||
parser.add_argument("--steps",
|
||||
type=int,
|
||||
default=5,
|
||||
help="num_inference_steps = sigma-grid points; N points run N-1 denoising "
|
||||
"forwards. 5 (default) is the distilled 4-forward grid")
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--num-gpus", type=int, default=4)
|
||||
parser.add_argument("--vsa-sparsity",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Run-level VSA sparsity in [0, 1). 0.0 (default) selects every tile, which is "
|
||||
"exactly dense attention; the student was trained at 0.9")
|
||||
# 64 is the trained contract: the student was TRAINED with 64-token
|
||||
# (4,4,4) tiles, and its to_gate_compress gates were learned against
|
||||
# pooling at that granularity — keep 64 unless you are ablating.
|
||||
parser.add_argument("--vsa-tile-size",
|
||||
type=int,
|
||||
choices=(64, 256),
|
||||
default=64,
|
||||
help="VSA-H3 tile size in tokens; 64 (default) is what the student was trained "
|
||||
"with and runs the native Triton block-sparse path, 256 is the FA4-CuTe-capable "
|
||||
"geometry for ablations")
|
||||
parser.add_argument("--vsa-kernel",
|
||||
choices=("triton", "sm100a"),
|
||||
default="triton",
|
||||
help="Block-sparse kernel for the tile-64 attention forward: triton (default, "
|
||||
"portable fwd+bwd) or sm100a — the opt-in Blackwell CUDA forward "
|
||||
"(fastvideo_kernel.block_sparse_attn_sm100a). sm100a needs an sm_100 GPU and a "
|
||||
"fastvideo-kernel build that carries the extension; if a precondition fails at "
|
||||
"run time the attention layer logs one warning and falls back to Triton. Only "
|
||||
"meaningful with --vsa-tile-size 64")
|
||||
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
|
||||
parser.add_argument("--compile-mode",
|
||||
default=None,
|
||||
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
|
||||
parser.add_argument("--repeats",
|
||||
type=int,
|
||||
default=1,
|
||||
help="generate N times; with --torch-compile the first run pays "
|
||||
"compilation, so steady-state is the last repeat")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
output_dir = Path(args.output)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if args.vsa_kernel == "sm100a":
|
||||
# The attention backend reads FASTVIDEO_VSA_SM100A per forward; set it
|
||||
# before the pipeline boots so spawned GPU workers inherit it. The
|
||||
# kernel is forward-only and inference runs under no-grad, so every
|
||||
# denoising forward qualifies for the CUDA route.
|
||||
os.environ["FASTVIDEO_VSA_SM100A"] = "1"
|
||||
|
||||
# Boot-time run configuration, folded into FastVideoArgs (the same route
|
||||
# examples/inference/basic/basic_minimax_h3_t2v.py uses for sparsity):
|
||||
# - attention_backend: the checkpoint carries trained to_gate_compress
|
||||
# gates, which only exist under the VSA-H3 backend — a dense-backend
|
||||
# load would reject them as unexpected weights. Layers that do not
|
||||
# support VSA-H3 (e.g. the token refiner) fall back to flash attention.
|
||||
# - VSA_tile_size: forwarded even at sparsity 0.0 because the gate-compress
|
||||
# branch pools per tile, and the gates were trained at 64 tokens/tile.
|
||||
experimental: dict[str, object] = {
|
||||
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
|
||||
"VSA_tile_size": args.vsa_tile_size,
|
||||
}
|
||||
if args.vsa_sparsity > 0.0:
|
||||
experimental["VSA_sparsity"] = args.vsa_sparsity
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
pipeline=PipelineSelection(experimental=experimental),
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=args.num_gpus > 1,
|
||||
parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus),
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
dit_layerwise=False,
|
||||
text_encoder=True,
|
||||
vae=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
compile=CompileConfig(
|
||||
enabled=args.torch_compile,
|
||||
mode=args.compile_mode,
|
||||
),
|
||||
),
|
||||
))
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=args.steps,
|
||||
# the base model is guidance-distilled; the student inherits it
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_dir / "fasth3.mp4"),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
print(f"Output written to: {result.video_path}")
|
||||
if result.generation_time is not None:
|
||||
# machine-readable: benchmark harnesses parse this line to separate
|
||||
# generation from model-load time (last occurrence = steady state)
|
||||
print(f"Generation time: {result.generation_time:.2f}s")
|
||||
for _ in range(args.repeats - 1):
|
||||
result = generator.generate(request)
|
||||
if result.generation_time is not None:
|
||||
print(f"Generation time: {result.generation_time:.2f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -64,12 +64,8 @@ def main() -> None:
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (
|
||||
args.num_gpus if args.num_gpus > 1 else 1
|
||||
)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (
|
||||
1 if args.num_gpus > 1 else args.num_gpus
|
||||
)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
|
||||
@@ -21,7 +21,6 @@ from fastvideo.api import (
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
DEFAULT_PROMPT = "a brushed steel espresso machine on a marble counter, morning window light"
|
||||
|
||||
|
||||
|
||||
@@ -9,11 +9,9 @@ import re
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
(
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
|
||||
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
|
||||
"35mm, bokeh"
|
||||
),
|
||||
("a cinematic photo of a red panda wearing a tiny backpack, standing on a "
|
||||
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
|
||||
"35mm, bokeh"),
|
||||
]
|
||||
|
||||
|
||||
@@ -42,9 +40,7 @@ def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(
|
||||
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
|
||||
)
|
||||
p = argparse.ArgumentParser(description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.", )
|
||||
p.add_argument(
|
||||
"--model-path",
|
||||
default="official_weights/FLUX.1-dev",
|
||||
@@ -108,9 +104,7 @@ def main() -> None:
|
||||
try:
|
||||
for i, prompt in enumerate(prompts):
|
||||
seed = args.seed + i
|
||||
filename_base = (
|
||||
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
|
||||
)
|
||||
filename_base = (f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}")
|
||||
_remove_existing_outputs(args.out_dir, filename_base)
|
||||
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
|
||||
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
|
||||
|
||||
@@ -33,21 +33,20 @@ MODEL_PATH = os.environ.get("GAMECRAFT_MODEL_PATH", "FastVideo/HunyuanGameCraft-
|
||||
|
||||
# Default prompts for demo
|
||||
DEFAULT_PROMPTS = {
|
||||
"village": "A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
|
||||
"temple": "A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
|
||||
"forest": "A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
|
||||
"village":
|
||||
"A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
|
||||
"temple":
|
||||
"A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
|
||||
"forest":
|
||||
"A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
|
||||
"beach": "A tropical beach with crystal clear turquoise water, white sand, and palm trees swaying in the breeze.",
|
||||
}
|
||||
|
||||
# I2V: default reference image (URL). Can override with a local path.
|
||||
DEFAULT_I2V_IMAGE_URL = (
|
||||
"https://huggingface.co/datasets/huggingface/documentation-images/"
|
||||
"resolve/main/diffusers/astronaut.jpg"
|
||||
)
|
||||
DEFAULT_I2V_PROMPT = (
|
||||
"An astronaut hatching from an egg, on the surface of the moon, "
|
||||
"the darkness and depth of space realised in the background."
|
||||
)
|
||||
DEFAULT_I2V_IMAGE_URL = ("https://huggingface.co/datasets/huggingface/documentation-images/"
|
||||
"resolve/main/diffusers/astronaut.jpg")
|
||||
DEFAULT_I2V_PROMPT = ("An astronaut hatching from an egg, on the surface of the moon, "
|
||||
"the darkness and depth of space realised in the background.")
|
||||
|
||||
OUTPUT_PATH = "video_samples_gamecraft"
|
||||
|
||||
|
||||
@@ -26,51 +26,35 @@ from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="GEN3C video generation")
|
||||
parser.add_argument("--model_path",
|
||||
type=str,
|
||||
default="converted_weights/GEN3C-Cosmos-7B")
|
||||
parser.add_argument("--image_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Input image for 3D cache conditioning")
|
||||
parser.add_argument("--prompt",
|
||||
type=str,
|
||||
default="A slow camera pan over a sunlit landscape.")
|
||||
parser.add_argument("--model_path", type=str, default="converted_weights/GEN3C-Cosmos-7B")
|
||||
parser.add_argument("--image_path", type=str, default=None, help="Input image for 3D cache conditioning")
|
||||
parser.add_argument("--prompt", type=str, default="A slow camera pan over a sunlit landscape.")
|
||||
parser.add_argument(
|
||||
"--negative_prompt",
|
||||
type=str,
|
||||
default=(
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
|
||||
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
|
||||
"flickering. Overall, the video is of poor quality."
|
||||
),
|
||||
default=("The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
|
||||
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
|
||||
"flickering. Overall, the video is of poor quality."),
|
||||
)
|
||||
parser.add_argument("--trajectory",
|
||||
type=str,
|
||||
default="left",
|
||||
choices=[
|
||||
"left", "right", "up", "down", "zoom_in",
|
||||
"zoom_out", "clockwise", "counterclockwise", "none"
|
||||
])
|
||||
parser.add_argument(
|
||||
"--trajectory",
|
||||
type=str,
|
||||
default="left",
|
||||
choices=["left", "right", "up", "down", "zoom_in", "zoom_out", "clockwise", "counterclockwise", "none"])
|
||||
parser.add_argument("--movement_distance", type=float, default=0.3)
|
||||
parser.add_argument("--camera_rotation",
|
||||
type=str,
|
||||
default="center_facing",
|
||||
choices=[
|
||||
"center_facing", "no_rotation",
|
||||
"trajectory_aligned"
|
||||
])
|
||||
choices=["center_facing", "no_rotation", "trajectory_aligned"])
|
||||
parser.add_argument("--height", type=int, default=704)
|
||||
parser.add_argument("--width", type=int, default=1280)
|
||||
parser.add_argument("--num_frames", type=int, default=121)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=35)
|
||||
parser.add_argument("--guidance_scale", type=float, default=1.0)
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs_video/gen3c.mp4")
|
||||
parser.add_argument("--output_path", type=str, default="outputs_video/gen3c.mp4")
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@ import json
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -12,31 +14,38 @@ def main():
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
video = generator.generate_video(prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt="",
|
||||
num_frames=81,
|
||||
fps=16)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
video2 = generator.generate_video(prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt="",
|
||||
num_frames=81,
|
||||
fps=16)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -3,38 +3,37 @@ import json
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15_1080p"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
|
||||
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
|
||||
# or "weizhou03/HunyuanVideo-1.5-Diffusers-1080p" # 720p -> 1080p
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
|
||||
|
||||
|
||||
@@ -6,6 +6,8 @@ DEFAULT_PROMPT = 'A paved pathway leads towards a stone arch bridge spanning a c
|
||||
DEFAULT_IMAGE = 'https://raw.githubusercontent.com/Tencent-Hunyuan/HY-WorldPlay/main/assets/img/test.png'
|
||||
|
||||
OUTPUT_PATH = "video_samples_hyworld"
|
||||
|
||||
|
||||
def main():
|
||||
import argparse
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ IMAGE_PATH = "assets/girl.png"
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
|
||||
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
@@ -19,9 +19,7 @@ def main():
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A woman stands up and walks away"
|
||||
)
|
||||
prompt = ("A woman stands up and walks away")
|
||||
_ = generator.generate_video(
|
||||
prompt,
|
||||
image_path=IMAGE_PATH,
|
||||
|
||||
@@ -2,6 +2,7 @@ from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
@@ -17,21 +18,28 @@ def main():
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
_ = generator.generate_video(prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=512,
|
||||
width=768,
|
||||
num_frames=121)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=512,
|
||||
width=768,
|
||||
num_frames=121)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -6,7 +6,6 @@ from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
DATASET_DIR = REPO_ROOT / "examples" / "dataset" / "lingbotworld2"
|
||||
OUTPUT_PATH = REPO_ROOT / "outputs" / "lingbotworld2_causal_fast.mp4"
|
||||
|
||||
@@ -3,17 +3,19 @@ from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embeddin
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
OUTPUT_PATH = "video_samples_lingbotworld"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
|
||||
@@ -21,20 +21,16 @@ import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
"A woman sits at a wooden table by the window in a cozy café. She reaches out "
|
||||
"with her right hand, picks up the white coffee cup from the saucer, and gently "
|
||||
"brings it to her lips to take a sip. After drinking, she places the cup back on "
|
||||
"the table and looks out the window, enjoying the peaceful atmosphere."
|
||||
)
|
||||
PROMPT = ("A woman sits at a wooden table by the window in a cozy café. She reaches out "
|
||||
"with her right hand, picks up the white coffee cup from the saucer, and gently "
|
||||
"brings it to her lips to take a sip. After drinking, she places the cup back on "
|
||||
"the table and looks out the window, enjoying the peaceful atmosphere.")
|
||||
|
||||
NEGATIVE_PROMPT = (
|
||||
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards"
|
||||
)
|
||||
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards")
|
||||
|
||||
# Input image path
|
||||
IMAGE_PATH = "assets/girl.png"
|
||||
@@ -51,20 +47,20 @@ def basic_generation():
|
||||
print("=" * 60)
|
||||
print("LongCat I2V: Basic Generation (50 steps, 480p)")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_i2v_basic"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -79,7 +75,7 @@ def basic_generation():
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -94,11 +90,11 @@ def distill_refine_generation():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat I2V: Distill + Refine Pipeline")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Stage 1: Distilled generation (16 steps at 480p)
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
num_gpus=1,
|
||||
@@ -111,9 +107,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_i2v_distill"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -128,14 +124,14 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
# Stage 2: Refinement (480p -> 768p)
|
||||
print("\n[Stage 2] Refinement (480p -> 768p with BSA)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
# Find the actual saved video file from stage 1
|
||||
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
|
||||
if not video_files:
|
||||
@@ -143,7 +139,7 @@ def distill_refine_generation():
|
||||
# Use the most recently created video file
|
||||
distill_video_path = max(video_files, key=os.path.getmtime)
|
||||
print(f"Using stage 1 video: {distill_video_path}")
|
||||
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
# Note: Refinement uses the T2V model (not I2V) since it's upscaling the generated video
|
||||
# For BSA [4, 4, 8]: latent must be divisible by 8
|
||||
@@ -163,9 +159,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_i2v_refine_720p"
|
||||
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -182,7 +178,7 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
@@ -192,13 +188,13 @@ def main():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat Image-to-Video Example")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
|
||||
# Run basic generation
|
||||
basic_generation()
|
||||
|
||||
|
||||
# Run distill+refine pipeline
|
||||
distill_refine_generation()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All generations complete!")
|
||||
print("=" * 60)
|
||||
@@ -206,5 +202,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -15,22 +15,18 @@ import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
"In a realistic photography style, a white boy around seven or eight years old "
|
||||
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
|
||||
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
|
||||
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
|
||||
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
|
||||
"features a green lawn and several tall trees, creating a warm and loving scene."
|
||||
)
|
||||
PROMPT = ("In a realistic photography style, a white boy around seven or eight years old "
|
||||
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
|
||||
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
|
||||
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
|
||||
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
|
||||
"features a green lawn and several tall trees, creating a warm and loving scene.")
|
||||
|
||||
NEGATIVE_PROMPT = (
|
||||
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards"
|
||||
)
|
||||
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards")
|
||||
|
||||
SEED = 42
|
||||
|
||||
@@ -44,20 +40,20 @@ def basic_generation():
|
||||
print("=" * 60)
|
||||
print("LongCat T2V: Basic Generation (50 steps, 480p)")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_t2v_basic"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -71,7 +67,7 @@ def basic_generation():
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -86,11 +82,11 @@ def distill_refine_generation():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat T2V: Distill + Refine Pipeline")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Stage 1: Distilled generation (16 steps at 480p)
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
@@ -103,9 +99,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_t2v_distill"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -119,14 +115,14 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
# Stage 2: Refinement (480p -> 720p)
|
||||
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
# Find the actual saved video file from stage 1
|
||||
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
|
||||
if not video_files:
|
||||
@@ -134,7 +130,7 @@ def distill_refine_generation():
|
||||
# Use the most recently created video file
|
||||
distill_video_path = max(video_files, key=os.path.getmtime)
|
||||
print(f"Using stage 1 video: {distill_video_path}")
|
||||
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
refine_generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
@@ -151,9 +147,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_t2v_refine_720p"
|
||||
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -170,7 +166,7 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
@@ -180,13 +176,13 @@ def main():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat Text-to-Video Example")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
|
||||
# Run basic generation
|
||||
basic_generation()
|
||||
|
||||
|
||||
# Run distill+refine pipeline
|
||||
distill_refine_generation()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All generations complete!")
|
||||
print("=" * 60)
|
||||
@@ -194,5 +190,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -21,21 +21,17 @@ import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
"A person rides a motorcycle along a long, straight road that stretches between "
|
||||
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
|
||||
"the motorcycle centered between the guardrails, while the scenery passes by on "
|
||||
"both sides. The video captures the journey from the rider's perspective, emphasizing "
|
||||
"the sense of motion and adventure."
|
||||
)
|
||||
PROMPT = ("A person rides a motorcycle along a long, straight road that stretches between "
|
||||
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
|
||||
"the motorcycle centered between the guardrails, while the scenery passes by on "
|
||||
"both sides. The video captures the journey from the rider's perspective, emphasizing "
|
||||
"the sense of motion and adventure.")
|
||||
|
||||
NEGATIVE_PROMPT = (
|
||||
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards"
|
||||
)
|
||||
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards")
|
||||
|
||||
# Input video path
|
||||
VIDEO_PATH = "assets/motorcycle.mp4"
|
||||
@@ -55,27 +51,25 @@ def basic_generation():
|
||||
print("=" * 60)
|
||||
print("LongCat VC: Basic Generation (50 steps, 480p)")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Check if video exists
|
||||
if not os.path.exists(VIDEO_PATH):
|
||||
raise FileNotFoundError(
|
||||
f"Video not found at {VIDEO_PATH}. "
|
||||
"Please provide a valid video path."
|
||||
)
|
||||
|
||||
raise FileNotFoundError(f"Video not found at {VIDEO_PATH}. "
|
||||
"Please provide a valid video path.")
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-VC-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_vc_basic"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -91,7 +85,7 @@ def basic_generation():
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -106,18 +100,16 @@ def distill_refine_generation():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat VC: Distill + Refine Pipeline")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Check if video exists
|
||||
if not os.path.exists(VIDEO_PATH):
|
||||
raise FileNotFoundError(
|
||||
f"Video not found at {VIDEO_PATH}. "
|
||||
"Please provide a valid video path."
|
||||
)
|
||||
|
||||
raise FileNotFoundError(f"Video not found at {VIDEO_PATH}. "
|
||||
"Please provide a valid video path.")
|
||||
|
||||
# Stage 1: Distilled generation (16 steps at 480p)
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-VC-Diffusers",
|
||||
num_gpus=1,
|
||||
@@ -130,9 +122,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_vc_distill"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -148,14 +140,14 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
# Stage 2: Refinement (480p -> 720p)
|
||||
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
# Find the actual saved video file from stage 1
|
||||
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
|
||||
if not video_files:
|
||||
@@ -163,7 +155,7 @@ def distill_refine_generation():
|
||||
# Use the most recently created video file
|
||||
distill_video_path = max(video_files, key=os.path.getmtime)
|
||||
print(f"Using stage 1 video: {distill_video_path}")
|
||||
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
# Note: Refinement uses the T2V model (not VC) since it's upscaling the generated video
|
||||
refine_generator = VideoGenerator.from_pretrained(
|
||||
@@ -181,9 +173,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_vc_refine_720p"
|
||||
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -200,7 +192,7 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
@@ -210,13 +202,13 @@ def main():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat Video Continuation Example")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
|
||||
# Run basic generation
|
||||
basic_generation()
|
||||
|
||||
|
||||
# Run distill+refine pipeline
|
||||
distill_refine_generation()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All generations complete!")
|
||||
print("=" * 60)
|
||||
@@ -224,5 +216,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -1,19 +1,16 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic.")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
@@ -36,4 +33,4 @@ def main() -> None:
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -67,25 +67,19 @@ _inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v")
|
||||
)
|
||||
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX23_MODEL_PATH",
|
||||
"FastVideo/LTX-2.3-Distilled-Diffusers")))
|
||||
OUTPUT_DIR = Path(os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v"))
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
DEFAULT_PROMPT = ("A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel.")
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
# Per-stage timing helpers --------------------------------------------------
|
||||
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
"""Print stage execution times and return the sum, or None if missing."""
|
||||
logging_info = result.get("logging_info")
|
||||
@@ -114,9 +108,7 @@ def _collect_stage_times(
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
stage_times.setdefault(name, []).append(float(metrics.get("execution_time", 0.0)))
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
@@ -125,21 +117,18 @@ def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
raise FileNotFoundError(f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`.")
|
||||
|
||||
|
||||
# Main ---------------------------------------------------------------------
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py"
|
||||
)
|
||||
raise SystemExit("LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py")
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
@@ -201,11 +190,13 @@ def main() -> None:
|
||||
|
||||
common_kwargs = dict(
|
||||
prompt=PROMPT,
|
||||
negative_prompt="", # distilled is CFG-free; no negative needed
|
||||
guidance_scale=1.0, # CFG=1 for distilled
|
||||
height=1280, width=832, # portrait runway aspect
|
||||
num_frames=121, fps=24, # ~5s clip
|
||||
num_inference_steps=8, # distilled denoise steps
|
||||
negative_prompt="", # distilled is CFG-free; no negative needed
|
||||
guidance_scale=1.0, # CFG=1 for distilled
|
||||
height=1280,
|
||||
width=832, # portrait runway aspect
|
||||
num_frames=121,
|
||||
fps=24, # ~5s clip
|
||||
num_inference_steps=8, # distilled denoise steps
|
||||
# i2v: anchor the input image at frame 0 with full strength.
|
||||
# `ltx2_image_crf=0.0` skips an extra JPEG re-encode of an already
|
||||
# JPEG conditioning image.
|
||||
@@ -251,10 +242,7 @@ def main() -> None:
|
||||
**common_kwargs,
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
e2e = (result.get("e2e_latency") if isinstance(result, dict) else None) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
@@ -266,10 +254,8 @@ def main() -> None:
|
||||
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
print(f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s")
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
|
||||
@@ -100,23 +100,14 @@ _inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv(
|
||||
"LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"
|
||||
)
|
||||
)
|
||||
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX23_MODEL_PATH",
|
||||
"FastVideo/LTX-2.3-Distilled-Diffusers")))
|
||||
OUTPUT_DIR = Path(os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"))
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
DEFAULT_PROMPT = ("A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel.")
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
|
||||
@@ -147,9 +138,7 @@ def _collect_stage_times(
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
stage_times.setdefault(name, []).append(float(metrics.get("execution_time", 0.0)))
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
@@ -157,20 +146,16 @@ def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
raise FileNotFoundError(f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`.")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_typed.py"
|
||||
)
|
||||
raise SystemExit("LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_typed.py")
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
@@ -220,10 +205,9 @@ def main() -> None:
|
||||
# model-specific VAE precision / decoder defaults are picked up
|
||||
# the same way the legacy example's
|
||||
# ``PipelineConfig.from_pretrained(model_root)`` did them.
|
||||
components=ComponentConfig(
|
||||
upsampler_weights=str(refine_upsampler_path),
|
||||
# Distilled has no refine LoRA — omit ``lora_path``.
|
||||
),
|
||||
components=ComponentConfig(upsampler_weights=str(refine_upsampler_path),
|
||||
# Distilled has no refine LoRA — omit ``lora_path``.
|
||||
),
|
||||
vae_tiling=False,
|
||||
preset_overrides={
|
||||
"refine": {
|
||||
@@ -278,11 +262,7 @@ def main() -> None:
|
||||
for w in range(warmup_runs):
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
|
||||
t0 = time.perf_counter()
|
||||
generator.generate(
|
||||
build_request(
|
||||
OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7
|
||||
)
|
||||
)
|
||||
generator.generate(build_request(OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7))
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
@@ -291,46 +271,30 @@ def main() -> None:
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
for m in range(measured_runs):
|
||||
out_path = (
|
||||
OUTPUT_DIR
|
||||
/ f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4"
|
||||
)
|
||||
print(
|
||||
f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}"
|
||||
)
|
||||
out_path = (OUTPUT_DIR / f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4")
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate(
|
||||
build_request(out_path, seed=2002 + m)
|
||||
)
|
||||
result = generator.generate(build_request(out_path, seed=2002 + m))
|
||||
wall = time.perf_counter() - t0
|
||||
# ``e2e_latency`` is currently surfaced via ``result.extra``;
|
||||
# ``GenerationResult`` exposes ``generation_time`` as a
|
||||
# first-class field but the LTX-2 pipeline only fills the
|
||||
# legacy ``e2e_latency`` key. Prefer the explicit one, fall
|
||||
# back to wall-clock.
|
||||
e2e = (
|
||||
result.extra.get("e2e_latency")
|
||||
if hasattr(result, "extra") else None
|
||||
) or wall
|
||||
e2e = (result.extra.get("e2e_latency") if hasattr(result, "extra") else None) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(
|
||||
f"[measured {m + 1}/{measured_runs}] "
|
||||
f"e2e={e2e:.2f}s wall={wall:.2f}s"
|
||||
)
|
||||
print(f"[measured {m + 1}/{measured_runs}] "
|
||||
f"e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
print("\n=== summary ===")
|
||||
print(
|
||||
f"warmup wall-times: "
|
||||
f"{[round(x, 1) for x in warmup_secs]}"
|
||||
)
|
||||
print(f"warmup wall-times: "
|
||||
f"{[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
print(f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s")
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
|
||||
@@ -1,21 +1,21 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic.")
|
||||
import os
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
|
||||
@@ -12,15 +12,11 @@ from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.layers.quantization.nvfp4_config import NVFP4Config
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
VALIDATION_JSON = (
|
||||
Path(__file__).resolve().parents[2] / "training" / "finetune" / "ltx2" / "validation.json"
|
||||
)
|
||||
VALIDATION_JSON = (Path(__file__).resolve().parents[2] / "training" / "finetune" / "ltx2" / "validation.json")
|
||||
|
||||
# Override with a local snapshot or converted directory when needed, e.g.
|
||||
# export LTX2_MODEL_PATH=/raid/$USER/hf/FastVideo/LTX2-Distilled-Diffusers
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(os.getenv("LTX2_MODEL_PATH", "FastVideo/LTX2-Distilled-Diffusers"))
|
||||
)
|
||||
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX2_MODEL_PATH", "FastVideo/LTX2-Distilled-Diffusers")))
|
||||
OUTPUT_DIR = Path("outputs_video/ltx2_distilled_fast_profile")
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
@@ -69,9 +65,7 @@ def print_stage_breakdown(
|
||||
return total
|
||||
|
||||
|
||||
def extract_sr_forward_latency(
|
||||
result: dict,
|
||||
) -> tuple[float | None, list[tuple[str, float]], list[str]]:
|
||||
def extract_sr_forward_latency(result: dict, ) -> tuple[float | None, list[tuple[str, float]], list[str]]:
|
||||
logging_info = result.get("logging_info")
|
||||
if logging_info is None:
|
||||
return None, [], []
|
||||
@@ -89,12 +83,8 @@ def extract_sr_forward_latency(
|
||||
if sr_match_substr:
|
||||
is_sr_stage = sr_match_substr in stage_name_l
|
||||
else:
|
||||
is_sr_stage = (
|
||||
"srdenoisingstage" in stage_name_l
|
||||
or "sr_denoising" in stage_name_l
|
||||
or "upsample" in stage_name_l
|
||||
or ("refine" in stage_name_l and "denois" in stage_name_l)
|
||||
)
|
||||
is_sr_stage = ("srdenoisingstage" in stage_name_l or "sr_denoising" in stage_name_l
|
||||
or "upsample" in stage_name_l or ("refine" in stage_name_l and "denois" in stage_name_l))
|
||||
if not is_sr_stage:
|
||||
continue
|
||||
exec_time = float(stage_metrics.get("execution_time", 0.0))
|
||||
@@ -161,11 +151,9 @@ def resolve_refine_upsampler_path(model_root: str) -> Path:
|
||||
return candidate
|
||||
|
||||
checked = "\n".join(f" - {candidate}" for candidate in candidates)
|
||||
raise FileNotFoundError(
|
||||
"Could not find an LTX2 refine upsampler directory.\n"
|
||||
"Checked:\n"
|
||||
f"{checked}"
|
||||
)
|
||||
raise FileNotFoundError("Could not find an LTX2 refine upsampler directory.\n"
|
||||
"Checked:\n"
|
||||
f"{checked}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
@@ -314,19 +302,15 @@ def main() -> None:
|
||||
|
||||
measured_times = run_times[measured_start_idx:]
|
||||
avg_time = sum(measured_times) / len(measured_times)
|
||||
print(
|
||||
f"Average video generation time over {len(measured_times)} runs "
|
||||
f"(runs {measured_start_idx + 1}-{len(run_times)}, skipping first {warmup_runs} warmup runs): "
|
||||
f"{avg_time:.2f}s"
|
||||
)
|
||||
print(f"Average video generation time over {len(measured_times)} runs "
|
||||
f"(runs {measured_start_idx + 1}-{len(run_times)}, skipping first {warmup_runs} warmup runs): "
|
||||
f"{avg_time:.2f}s")
|
||||
|
||||
measured_e2e_times = e2e_times[measured_start_idx:]
|
||||
avg_e2e_time = sum(measured_e2e_times) / len(measured_e2e_times)
|
||||
print(
|
||||
f"Average end-to-end latency over {len(measured_e2e_times)} runs "
|
||||
f"(runs {measured_start_idx + 1}-{len(e2e_times)}, skipping first {warmup_runs} warmup runs): "
|
||||
f"{avg_e2e_time:.2f}s"
|
||||
)
|
||||
print(f"Average end-to-end latency over {len(measured_e2e_times)} runs "
|
||||
f"(runs {measured_start_idx + 1}-{len(e2e_times)}, skipping first {warmup_runs} warmup runs): "
|
||||
f"{avg_e2e_time:.2f}s")
|
||||
|
||||
if sr_forward_times:
|
||||
avg_sr_forward = sum(sr_forward_times) / len(sr_forward_times)
|
||||
@@ -338,10 +322,8 @@ def main() -> None:
|
||||
|
||||
if non_stage_overhead_times:
|
||||
avg_non_stage_overhead = sum(non_stage_overhead_times) / len(non_stage_overhead_times)
|
||||
print(
|
||||
"Average non-stage overhead over "
|
||||
f"{len(non_stage_overhead_times)} measured runs: {avg_non_stage_overhead:.3f}s"
|
||||
)
|
||||
print("Average non-stage overhead over "
|
||||
f"{len(non_stage_overhead_times)} measured runs: {avg_non_stage_overhead:.3f}s")
|
||||
else:
|
||||
print("Average non-stage overhead unavailable (no stage timings).")
|
||||
finally:
|
||||
|
||||
@@ -13,24 +13,34 @@ MODEL_VARIANT = "base_distilled_model"
|
||||
# Variant-specific settings
|
||||
VARIANT_CONFIG = {
|
||||
"base_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
4,
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"gta_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
2,
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"templerun_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
7,
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples_matrixgame2"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -42,8 +52,8 @@ def main():
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
|
||||
@@ -14,27 +14,40 @@ MODEL_VARIANT = "base_distilled_model"
|
||||
# Variant-specific settings
|
||||
VARIANT_CONFIG = {
|
||||
"base_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"mode": "universal",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
4,
|
||||
"mode":
|
||||
"universal",
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"gta_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"mode": "gta_drive",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
2,
|
||||
"mode":
|
||||
"gta_drive",
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"templerun_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"mode": "templerun",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
7,
|
||||
"mode":
|
||||
"templerun",
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples_matrixgame2"
|
||||
|
||||
|
||||
async def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -46,8 +59,8 @@ async def main():
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
@@ -56,11 +69,8 @@ async def main():
|
||||
)
|
||||
|
||||
max_blocks = 50
|
||||
num_frames = 597
|
||||
actions = {
|
||||
"keyboard": torch.zeros((num_frames, config["keyboard_dim"])),
|
||||
"mouse": torch.zeros((num_frames, 2))
|
||||
}
|
||||
num_frames = 597
|
||||
actions = {"keyboard": torch.zeros((num_frames, config["keyboard_dim"])), "mouse": torch.zeros((num_frames, 2))}
|
||||
grid_sizes = torch.tensor([150, 44, 80])
|
||||
mode = config["mode"]
|
||||
|
||||
@@ -81,11 +91,11 @@ async def main():
|
||||
|
||||
for block_id in range(max_blocks):
|
||||
print(f"\n=== Block {block_id + 1}/{max_blocks} ===")
|
||||
|
||||
|
||||
action = await get_current_action_async(mode)
|
||||
keyboard_cond, mouse_cond = expand_action_to_frames(action, 12)
|
||||
await generator.step_async(keyboard_cond, mouse_cond)
|
||||
|
||||
|
||||
if (await asyncio.to_thread(input, "\nContinue? (y/n): ")).lower() == 'n':
|
||||
break
|
||||
|
||||
|
||||
@@ -25,6 +25,13 @@ from fastvideo.pipelines.basic.minimax_h3.packing import resolve_canvas_size
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
parser.add_argument("--image", required=True, help="First-frame image path.")
|
||||
parser.add_argument("--last-image", help="Optional last-frame image path.")
|
||||
parser.add_argument("--output", default="outputs/minimax_h3_fl2va")
|
||||
|
||||
@@ -25,6 +25,13 @@ from fastvideo.pipelines.basic.minimax_h3 import MiniMaxH3Reference
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
parser.add_argument("--reference-video", required=True)
|
||||
parser.add_argument("--reference-audio", help="Optional additional audio reference.")
|
||||
parser.add_argument("--output", default="outputs/minimax_h3_ref2va")
|
||||
|
||||
@@ -22,6 +22,13 @@ from fastvideo.api import (
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
parser.add_argument("--prompt", required=True)
|
||||
parser.add_argument("--output", default="outputs/minimax_h3_t2v")
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
@@ -30,13 +37,15 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--num-gpus", type=int, default=4)
|
||||
parser.add_argument("--torch-compile", action="store_true",
|
||||
help="torch.compile the DiT transformer path")
|
||||
parser.add_argument("--compile-mode", default=None,
|
||||
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
|
||||
parser.add_argument("--compile-mode",
|
||||
default=None,
|
||||
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
|
||||
parser.add_argument("--repeats", type=int, default=1,
|
||||
parser.add_argument("--repeats",
|
||||
type=int,
|
||||
default=1,
|
||||
help="generate N times; with --torch-compile the first run pays "
|
||||
"compilation, so steady-state is the last repeat")
|
||||
"compilation, so steady-state is the last repeat")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
@@ -67,24 +76,24 @@ def main() -> None:
|
||||
))
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_dir / "minimax_h3_t2v.mp4"),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_dir / "minimax_h3_t2v.mp4"),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
print(f"Output written to: {result.video_path}")
|
||||
if result.generation_time is not None:
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MMAudio large-44k-v2 video-to-audio example."""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--video-path", required=True)
|
||||
parser.add_argument("--output-path", default="outputs_audio/mmaudio.wav")
|
||||
parser.add_argument("--duration-seconds", type=float, default=8.0)
|
||||
parser.add_argument("--prompt", default="")
|
||||
parser.add_argument("--negative-prompt", default="music")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
os.environ.get(
|
||||
"MMAUDIO_MODEL_PATH",
|
||||
"converted_weights/mmaudio/large_44k_v2",
|
||||
),
|
||||
workload_type="v2a",
|
||||
num_gpus=1,
|
||||
)
|
||||
result = generator.generate_video(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
video_path=args.video_path,
|
||||
audio_end_in_s=args.duration_seconds,
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
)
|
||||
print(result["video_path"])
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,19 +1,20 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
config.text_encoder_precisions = ["fp16"]
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
pipeline_config=config,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
dit_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
dit_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
# Create sampling parameters with reduced number of frames
|
||||
@@ -23,18 +24,19 @@ def main():
|
||||
sampling_param.width = 256
|
||||
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param)
|
||||
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, sampling_param=sampling_param)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -16,27 +18,24 @@ def main():
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
distributed_executor_backend="ray",
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
|
||||
@@ -7,7 +7,6 @@ import os
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a rainy neon-lit street at night, shallow depth of field, sharp focus, 35mm, bokeh",
|
||||
@@ -49,7 +48,9 @@ def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description="Run SD3.5 Medium text-to-image with FastVideo VideoGenerator.")
|
||||
p.add_argument("--model-path", default="stabilityai/stable-diffusion-3.5-medium", help="Path to local diffusers-format SD3.5 weights directory.")
|
||||
p.add_argument("--model-path",
|
||||
default="stabilityai/stable-diffusion-3.5-medium",
|
||||
help="Path to local diffusers-format SD3.5 weights directory.")
|
||||
p.add_argument(
|
||||
"--out-dir",
|
||||
"--outdir",
|
||||
|
||||
@@ -3,6 +3,8 @@ import time
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_causal"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -13,19 +15,18 @@ def main():
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -5,6 +5,8 @@ import json
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -14,8 +16,8 @@ def main():
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
dit_precision="fp32",
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
@@ -37,7 +39,11 @@ def main():
|
||||
for prompt_image_pair in prompt_image_pairs:
|
||||
prompt = prompt_image_pair["prompt"]
|
||||
image_path = prompt_image_pair["image_path"]
|
||||
_ = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
_ = generator.generate_video(prompt,
|
||||
image_path=image_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
sampling_param=sampling_param)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -5,6 +5,8 @@ from fastvideo import VideoGenerator
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -14,15 +16,17 @@ def main():
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
|
||||
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
|
||||
init_weights_from_safetensors=
|
||||
"/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
|
||||
init_weights_from_safetensors_2=
|
||||
"/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
|
||||
num_frame_per_block=7,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
@@ -31,13 +35,11 @@ def main():
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -54,16 +54,15 @@ PROMPT = "Steady lo-fi hip hop drum loop with vinyl crackle."
|
||||
# ...) you want to extend or repair. The pipeline raises if a mask is
|
||||
# passed without a reference, so this must be a real path.
|
||||
REFERENCE_AUDIO_PATH = "path/to/your/loop.wav"
|
||||
KEEP_SECONDS = 6.0 # first KEEP_SECONDS preserved exactly
|
||||
TOTAL_SECONDS = 12.0 # extend the loop to this duration
|
||||
KEEP_SECONDS = 6.0 # first KEEP_SECONDS preserved exactly
|
||||
TOTAL_SECONDS = 12.0 # extend the loop to this duration
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not os.path.isfile(REFERENCE_AUDIO_PATH):
|
||||
raise FileNotFoundError(
|
||||
f"REFERENCE_AUDIO_PATH={REFERENCE_AUDIO_PATH!r} does not exist. "
|
||||
"Edit this script to point at a real audio file (wav/mp3/mp4/"
|
||||
"m4a/flac) before running.")
|
||||
raise FileNotFoundError(f"REFERENCE_AUDIO_PATH={REFERENCE_AUDIO_PATH!r} does not exist. "
|
||||
"Edit this script to point at a real audio file (wav/mp3/mp4/"
|
||||
"m4a/flac) before running.")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
|
||||
@@ -15,19 +15,17 @@ def main() -> None:
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
|
||||
# set to false if using RTX 4090
|
||||
# set to false if using RTX 4090
|
||||
# pin_cpu_memory=False,
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
# TurboDiffusion defaults: guidance_scale=1.0 and num_inference_steps=4 (from config)
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
@@ -36,13 +34,11 @@ def main() -> None:
|
||||
)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic."
|
||||
)
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(
|
||||
prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
|
||||
@@ -17,11 +17,9 @@ def main() -> None:
|
||||
num_gpus=2,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
@@ -30,13 +28,11 @@ def main() -> None:
|
||||
)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic."
|
||||
)
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(
|
||||
prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
|
||||
@@ -18,12 +18,13 @@ def main() -> None:
|
||||
)
|
||||
|
||||
# Example prompt and image for I2V
|
||||
prompt = ("Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside.")
|
||||
prompt = (
|
||||
"Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
|
||||
)
|
||||
|
||||
# Use an example image path
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
image_path=image_path,
|
||||
|
||||
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_t2v"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -12,8 +14,8 @@ def main():
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
@@ -25,24 +27,31 @@ def main():
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, height=720, width=1280, num_frames=81)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
_ = generator.generate_video(prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=720,
|
||||
width=1280,
|
||||
num_frames=81)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=720, width=1280, num_frames=81)
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=720,
|
||||
width=1280,
|
||||
num_frames=81)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -4,6 +4,8 @@ from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_1_Fun"
|
||||
OUTPUT_NAME = "wan2.1_test"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -14,8 +16,8 @@ def main():
|
||||
# "alibaba-pai/Wan2.2-Fun-A14B-Control",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
@@ -30,7 +32,14 @@ def main():
|
||||
image_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/8.png"
|
||||
control_video_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/pose.mp4"
|
||||
|
||||
video = generator.generate_video(prompt, negative_prompt=negative_prompt, image_path=image_path, video_path=control_video_path, output_path=OUTPUT_PATH, output_video_name=OUTPUT_NAME, save_video=True)
|
||||
video = generator.generate_video(prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
image_path=image_path,
|
||||
video_path=control_video_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
output_video_name=OUTPUT_NAME,
|
||||
save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -12,8 +14,8 @@ def main():
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
@@ -24,7 +26,14 @@ def main():
|
||||
prompt = "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside."
|
||||
image_path = "https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG"
|
||||
|
||||
video = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, height=832, width=480, num_frames=81)
|
||||
video = generator.generate_video(prompt,
|
||||
image_path=image_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=832,
|
||||
width=480,
|
||||
num_frames=81)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_wan2_2_5B_ti2v"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -11,11 +13,11 @@ def main():
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
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, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
@@ -28,14 +30,13 @@ def main():
|
||||
# model!
|
||||
|
||||
# T2V mode
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -20,13 +20,11 @@ from fastvideo.api import (
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
DEFAULT_PROMPT = (
|
||||
"Young Chinese woman in red Hanfu, intricate embroidery. Impeccable makeup, red floral forehead pattern. "
|
||||
"Elaborate high bun, golden phoenix headdress, red flowers, beads. Holds round folding fan with lady, trees, bird. "
|
||||
"Neon lightning-bolt lamp (⚡️), bright yellow glow, above extended left palm. Soft-lit outdoor night background, "
|
||||
"silhouetted tiered pagoda (西安大雁塔), blurred colorful distant lights."
|
||||
)
|
||||
"silhouetted tiered pagoda (西安大雁塔), blurred colorful distant lights.")
|
||||
DEFAULT_REVISION = "f332072aa78be7aecdf3ee76d5c247082da564a6"
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tiny MLX RIFE frame-interpolation smoke test."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.mlx_runtime.rife_interp import interpolate, load_model
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="MLX RIFE 4.25 frame interpolation smoke test."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--self-test",
|
||||
action="store_true",
|
||||
help="Run a tiny two-frame interpolation test.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
if not args.self_test:
|
||||
raise SystemExit("Nothing to do; pass --self-test")
|
||||
|
||||
frame0 = np.zeros((64, 96, 3), dtype=np.uint8)
|
||||
frame1 = np.zeros((64, 96, 3), dtype=np.uint8)
|
||||
frame1[:, :, 0] = 255
|
||||
start = time.perf_counter()
|
||||
model = load_model()
|
||||
load_s = time.perf_counter() - start
|
||||
start = time.perf_counter()
|
||||
frames = interpolate([frame0, frame1], factor=2, model=model)
|
||||
interp_s = time.perf_counter() - start
|
||||
assert len(frames) == 3
|
||||
assert frames[1].shape == frame0.shape
|
||||
assert frames[1].dtype == np.uint8
|
||||
print(
|
||||
"MLX RIFE self-test passed: "
|
||||
f"load_s={load_s:.3f} interp_s={interp_s:.3f} shape={frames[1].shape}"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,510 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""End-to-end Wan2.2-TI2V-5B generation on Apple Silicon (MLX DiT + MLX TAEHV).
|
||||
|
||||
Pipeline: torch/MPS UMT5 encode (shared with 1.3B) → MLXWan22DiT 3-step DMD
|
||||
(warped schedule, flow_shift=5) → MLX TAEHV decode (taew2_2.pth). Fully MLX
|
||||
on the heavy DiT + decode path.
|
||||
|
||||
PYTHONPATH=$PWD python examples/inference/basic/mlx_wan22_generate.py \
|
||||
--prompt "A red fox trotting through a snowy pine forest at golden hour" \
|
||||
--output-path video_samples/demo_5b/fox_5b_mlx.mp4
|
||||
|
||||
Decoder backends: ``taehv`` (default, MLX, ~seconds), ``taehv-torch`` (parity),
|
||||
``wan-vae`` (full AutoencoderKLWan on MPS, slow).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.mlx_runtime.fast_spatial import DEFAULT_FAST_SPATIAL_SHARPEN
|
||||
from fastvideo.mlx_runtime.frame_upsample import DEFAULT_PIXEL_UPSAMPLE_MODE, PIXEL_UPSAMPLE_MODES
|
||||
from fastvideo.mlx_runtime.memory import cleanup_mlx
|
||||
from fastvideo.mlx_runtime.prompt_cache import (
|
||||
fingerprint_digest,
|
||||
load_prompt_cache,
|
||||
save_prompt_cache,
|
||||
text_encoder_fingerprint,
|
||||
)
|
||||
from fastvideo.mlx_runtime.rife_interp import aligned_keyframe_count
|
||||
|
||||
FASTWAN21_MODEL_ID = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
|
||||
FASTWAN22_MODEL_ID = "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
|
||||
DEFAULT_HEIGHT = 448
|
||||
DEFAULT_WIDTH = 832
|
||||
DEFAULT_NUM_FRAMES = 121
|
||||
|
||||
def _resolve_model_paths(
|
||||
*,
|
||||
text_encoder_root: Path | None,
|
||||
dit_checkpoint: Path | None,
|
||||
dit_config: Path | None,
|
||||
vae_root: Path | None,
|
||||
mlx_checkpoint: Path | None,
|
||||
decode_backend: str,
|
||||
) -> tuple[Path, Path | None, Path | None, Path | None]:
|
||||
"""Download only the missing assets required by the selected Wan2.2 path."""
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
if text_encoder_root is None:
|
||||
text_encoder_root = Path(snapshot_download(
|
||||
FASTWAN21_MODEL_ID,
|
||||
allow_patterns=["tokenizer/*", "text_encoder/*"],
|
||||
))
|
||||
if mlx_checkpoint is None and (dit_checkpoint is None or dit_config is None):
|
||||
patterns = []
|
||||
if dit_checkpoint is None:
|
||||
patterns.append("transformer/diffusion_pytorch_model.safetensors")
|
||||
if dit_config is None:
|
||||
patterns.append("transformer/config.json")
|
||||
model_root = Path(snapshot_download(FASTWAN22_MODEL_ID, allow_patterns=patterns))
|
||||
dit_checkpoint = dit_checkpoint or model_root / "transformer/diffusion_pytorch_model.safetensors"
|
||||
dit_config = dit_config or model_root / "transformer/config.json"
|
||||
if decode_backend == "wan-vae" and vae_root is None:
|
||||
model_root = Path(snapshot_download(FASTWAN22_MODEL_ID, allow_patterns=["vae/*"]))
|
||||
vae_root = model_root / "vae"
|
||||
return text_encoder_root, dit_checkpoint, dit_config, vae_root
|
||||
|
||||
|
||||
def _prompt_cache_fingerprint(
|
||||
*,
|
||||
prompt: str,
|
||||
prompt_used: str,
|
||||
enhance_prompt: bool,
|
||||
enhance_prompt_backend: str,
|
||||
text_encoder_root: Path,
|
||||
max_sequence_length: int,
|
||||
dtype: str,
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"prompt": prompt,
|
||||
"prompt_used": prompt_used,
|
||||
"enhance_prompt": enhance_prompt,
|
||||
"enhance_prompt_backend": enhance_prompt_backend,
|
||||
"text_encoder": text_encoder_fingerprint(text_encoder_root),
|
||||
"max_sequence_length": max_sequence_length,
|
||||
"dtype": dtype,
|
||||
}
|
||||
|
||||
|
||||
def _default_prompt_cache_path(fingerprint: dict[str, object]) -> Path:
|
||||
"""Content-addressed default cache file for a prompt fingerprint.
|
||||
|
||||
The Wan2.1 entrypoint caches prompt embeddings by default; this one only
|
||||
did so when handed an explicit ``--prompt-embeds-cache`` path, so every 5B
|
||||
run paid a full UMT5 encode (~45s on an M4 Max) even for a repeat prompt.
|
||||
The fingerprint already covers everything that changes the embedding, so
|
||||
hash it for the filename.
|
||||
"""
|
||||
digest = fingerprint_digest(fingerprint)[:32]
|
||||
return Path.home() / ".cache" / "fastvideo" / "prompt_embeds" / f"wan22_{digest}.npy"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="MLX Wan2.2-5B T2V (encode → DiT DMD → TAEHV/VAE decode)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default="A red fox trotting through a snowy pine forest at golden hour, cinematic",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=Path,
|
||||
default=Path("video_samples/demo_5b/fox_5b_mlx.mp4"),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-root",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Root with text_encoder/ + tokenizer/",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt-embeds-cache",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Explicit .npy UMT5 embedding cache file. Overrides the automatic "
|
||||
"content-addressed cache (--prompt-cache).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt-cache",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Cache prompt embeddings under ~/.cache/fastvideo/prompt_embeds so "
|
||||
"repeat runs skip the text encoder entirely. Default: on.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-device",
|
||||
choices=("auto", "cpu", "mps"),
|
||||
default="cpu",
|
||||
help="Device for UMT5 encoding. CPU is safest beside the 5B MLX DiT.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enhance-prompt",
|
||||
action="store_true",
|
||||
help="Apply deterministic local cinematic prompt enrichment before UMT5.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enhance-prompt-backend",
|
||||
choices=("template",),
|
||||
default="template",
|
||||
help="Prompt enrichment backend.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-checkpoint",
|
||||
type=Path,
|
||||
default=None,
|
||||
)
|
||||
parser.add_argument("--dit-config", type=Path, default=None)
|
||||
parser.add_argument(
|
||||
"--mlx-checkpoint",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Pre-quantized MLX DiT checkpoint directory. Rewrapped with Wan2.2 per-token conditioning.",
|
||||
)
|
||||
parser.add_argument("--vae-root", type=Path, default=None)
|
||||
parser.add_argument("--height", type=int, default=DEFAULT_HEIGHT)
|
||||
parser.add_argument("--width", type=int, default=DEFAULT_WIDTH)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=DEFAULT_NUM_FRAMES,
|
||||
help="Pixel frames (121 at 24fps = 5.04 seconds)",
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=1234)
|
||||
parser.add_argument("--renoise-seed", type=int, default=0)
|
||||
parser.add_argument("--fps", type=int, default=24)
|
||||
parser.add_argument("--flow-shift", type=float, default=5.0)
|
||||
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
|
||||
parser.add_argument(
|
||||
"--no-warp",
|
||||
action="store_true",
|
||||
help="Disable schedule warping (debug only).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fast",
|
||||
action="store_true",
|
||||
help="Generate fewer frames then RIFE-interpolate to --num-frames.",
|
||||
)
|
||||
parser.add_argument("--fast-factor", type=int, default=2)
|
||||
parser.add_argument("--fast-sharpen", type=float, default=0.6)
|
||||
parser.add_argument(
|
||||
"--fast-spatial",
|
||||
action="store_true",
|
||||
help="Denoise and decode at reduced spatial resolution, then resample "
|
||||
"the decoded frames up to the target size.",
|
||||
)
|
||||
parser.add_argument("--fast-spatial-scale", type=int, default=2)
|
||||
parser.add_argument(
|
||||
"--fast-spatial-upsample-mode",
|
||||
choices=PIXEL_UPSAMPLE_MODES,
|
||||
default=DEFAULT_PIXEL_UPSAMPLE_MODE,
|
||||
)
|
||||
parser.add_argument("--fast-spatial-sharpen", type=float, default=DEFAULT_FAST_SPATIAL_SHARPEN)
|
||||
parser.add_argument(
|
||||
"--refine",
|
||||
action="store_true",
|
||||
help="Two-pass DMD: coarse denoise, upsample/re-noise, full-res denoise.",
|
||||
)
|
||||
parser.add_argument("--refine-scale", type=int, default=2)
|
||||
parser.add_argument(
|
||||
"--refine-upsample-mode",
|
||||
choices=("bilinear", "nearest"),
|
||||
default="bilinear",
|
||||
)
|
||||
parser.add_argument("--no-refine-add-noise", action="store_true")
|
||||
parser.add_argument(
|
||||
"--decode-backend",
|
||||
choices=("taehv", "taehv-torch", "wan-vae"),
|
||||
default="taehv",
|
||||
)
|
||||
parser.add_argument("--save-latents", type=Path, default=None)
|
||||
parser.add_argument("--metrics-json", type=Path, default=None,
|
||||
help="Write measured run metadata as JSON for reports or galleries.")
|
||||
parser.add_argument(
|
||||
"--compile",
|
||||
action="store_true",
|
||||
help="Compile the DiT forward with mx.compile; fallback to eager on failure.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.fast_factor < 2:
|
||||
parser.error("--fast-factor must be at least 2")
|
||||
# --fast-spatial used to be rejected here because it upsampled the completed
|
||||
# 48-channel latent, which is out of distribution for the decoder and gave
|
||||
# black or noisy video. The upsample now runs on decoded frames, so the
|
||||
# latent never leaves the grid it was denoised on and the mode is usable.
|
||||
if args.refine and args.fast_spatial:
|
||||
print("[wan22] --refine takes precedence over --fast-spatial")
|
||||
args.text_encoder_root, args.dit_checkpoint, args.dit_config, args.vae_root = _resolve_model_paths(
|
||||
text_encoder_root=args.text_encoder_root,
|
||||
dit_checkpoint=args.dit_checkpoint,
|
||||
dit_config=args.dit_config,
|
||||
vae_root=args.vae_root,
|
||||
mlx_checkpoint=args.mlx_checkpoint,
|
||||
decode_backend=args.decode_backend,
|
||||
)
|
||||
target_frames = args.num_frames
|
||||
if args.fast:
|
||||
args.num_frames = aligned_keyframe_count(target_frames, args.fast_factor)
|
||||
print(
|
||||
f"[wan22 fast] generating {args.num_frames} frames, "
|
||||
f"RIFE {args.fast_factor}x -> {target_frames}"
|
||||
)
|
||||
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
from examples.inference.basic.mlx_wan_prompt_to_video import (
|
||||
_postprocess_video,
|
||||
encode_prompt,
|
||||
make_rotary_embeddings,
|
||||
)
|
||||
from fastvideo.mlx_runtime.fast_spatial import plan_fast_spatial
|
||||
from fastvideo.mlx_runtime.refine import (
|
||||
default_refine_timesteps,
|
||||
plan_refine_resolutions,
|
||||
prepare_refine_latents,
|
||||
)
|
||||
from fastvideo.mlx_runtime.wan22 import (
|
||||
mlx_wan22_dit_from_diffusers_safetensors,
|
||||
mlx_wan22_dit_from_mlx_checkpoint,
|
||||
)
|
||||
from fastvideo.mlx_runtime.wan22_sample import build_wan22_dmd_schedule, sample_wan22_dmd
|
||||
from fastvideo.mlx_runtime.wan_vae import decode_latents_to_video
|
||||
|
||||
if args.mlx_checkpoint is not None:
|
||||
config = json.loads((args.mlx_checkpoint / "mlx_dit.json").read_text())["config"]
|
||||
else:
|
||||
config = json.loads(args.dit_config.read_text())
|
||||
patch_size = tuple(config.get("patch_size", (1, 2, 2)))
|
||||
if args.refine:
|
||||
active_plan = plan_refine_resolutions(
|
||||
height=args.height, width=args.width, num_frames=args.num_frames,
|
||||
spatial_scale=args.refine_scale, vae_spatial_compression=16,
|
||||
vae_temporal_compression=4, patch_size=patch_size, enabled=True,
|
||||
)
|
||||
spatial_mode = "refine"
|
||||
elif args.fast_spatial:
|
||||
fast_spatial_plan = plan_fast_spatial(
|
||||
height=args.height, width=args.width, num_frames=args.num_frames,
|
||||
spatial_scale=args.fast_spatial_scale, vae_spatial_compression=16,
|
||||
vae_temporal_compression=4, patch_size=patch_size,
|
||||
upsample_mode=args.fast_spatial_upsample_mode,
|
||||
sharpen=args.fast_spatial_sharpen, enabled=True,
|
||||
)
|
||||
active_plan = fast_spatial_plan.plan
|
||||
spatial_mode = "fast_spatial"
|
||||
else:
|
||||
active_plan = plan_refine_resolutions(
|
||||
height=args.height, width=args.width, num_frames=args.num_frames,
|
||||
spatial_scale=1, vae_spatial_compression=16, vae_temporal_compression=4,
|
||||
patch_size=patch_size, enabled=False,
|
||||
)
|
||||
spatial_mode = "off"
|
||||
lat_h, lat_w = active_plan.stage1_latent_height, active_plan.stage1_latent_width
|
||||
lat_t = active_plan.latent_frames
|
||||
in_ch = int(config["in_channels"])
|
||||
print(f"[5B] latent {in_ch}x{lat_t}x{lat_h}x{lat_w}", flush=True)
|
||||
|
||||
total_start = time.perf_counter()
|
||||
prompt_for_encode = args.prompt
|
||||
enhance_backend = None
|
||||
enhance_elapsed_s = 0.0
|
||||
if args.enhance_prompt:
|
||||
from fastvideo.mlx_runtime.prompt_enhance import enhance_prompt
|
||||
|
||||
enhancement = enhance_prompt(args.prompt, backend=args.enhance_prompt_backend)
|
||||
prompt_for_encode = enhancement.enhanced
|
||||
enhance_backend = enhancement.backend
|
||||
enhance_elapsed_s = enhancement.elapsed_s
|
||||
print(f"[enhance] backend={enhance_backend} in {enhance_elapsed_s:.2f}s", flush=True)
|
||||
print(f"[enhance] prompt: {prompt_for_encode}", flush=True)
|
||||
|
||||
t0 = time.perf_counter()
|
||||
prompt_cache_fingerprint = _prompt_cache_fingerprint(
|
||||
prompt=args.prompt,
|
||||
prompt_used=prompt_for_encode,
|
||||
enhance_prompt=args.enhance_prompt,
|
||||
enhance_prompt_backend=args.enhance_prompt_backend,
|
||||
text_encoder_root=args.text_encoder_root,
|
||||
max_sequence_length=512,
|
||||
dtype="fp16",
|
||||
)
|
||||
prompt_cache_path = args.prompt_embeds_cache
|
||||
if prompt_cache_path is None and args.prompt_cache:
|
||||
prompt_cache_path = _default_prompt_cache_path(prompt_cache_fingerprint)
|
||||
cached_embeds = load_prompt_cache(
|
||||
prompt_cache_path,
|
||||
prompt_cache_fingerprint,
|
||||
)
|
||||
if cached_embeds is not None:
|
||||
embeds = torch.from_numpy(cached_embeds).contiguous()
|
||||
else:
|
||||
embeds = encode_prompt(
|
||||
model_root=args.text_encoder_root,
|
||||
prompt=prompt_for_encode,
|
||||
max_sequence_length=512,
|
||||
device_arg=args.text_encoder_device,
|
||||
dtype_arg="fp16",
|
||||
)
|
||||
save_prompt_cache(
|
||||
prompt_cache_path,
|
||||
embeds.cpu().numpy(),
|
||||
prompt_cache_fingerprint,
|
||||
)
|
||||
ehs = mx.array(embeds.numpy()).astype(mx.float16)
|
||||
prompt_encode_s = time.perf_counter() - t0
|
||||
print(f"[5B] prompt encoded {tuple(ehs.shape)} in {prompt_encode_s:.1f}s", flush=True)
|
||||
|
||||
t1 = time.perf_counter()
|
||||
if args.mlx_checkpoint is not None:
|
||||
dit = mlx_wan22_dit_from_mlx_checkpoint(
|
||||
args.mlx_checkpoint,
|
||||
compile=args.compile,
|
||||
)
|
||||
else:
|
||||
dit = mlx_wan22_dit_from_diffusers_safetensors(
|
||||
args.dit_checkpoint,
|
||||
args.dit_config,
|
||||
dtype="fp16",
|
||||
compile=args.compile,
|
||||
)
|
||||
dit_load_s = time.perf_counter() - t1
|
||||
print(f"[5B] DiT loaded in {dit_load_s:.1f}s", flush=True)
|
||||
|
||||
freqs = make_rotary_embeddings(config, latent_frames=lat_t, latent_height=lat_h, latent_width=lat_w)
|
||||
gen = torch.Generator().manual_seed(args.seed)
|
||||
noise = mx.array(
|
||||
torch.randn(1, in_ch, lat_t, lat_h, lat_w, generator=gen, dtype=torch.float32).numpy()).astype(mx.float16)
|
||||
|
||||
steps = [int(s) for s in args.dmd_denoising_steps.split(",") if s.strip()]
|
||||
t2 = time.perf_counter()
|
||||
mx.reset_peak_memory()
|
||||
latents = sample_wan22_dmd(
|
||||
dit,
|
||||
ehs,
|
||||
noise,
|
||||
freqs,
|
||||
dmd_denoising_steps=steps,
|
||||
flow_shift=args.flow_shift,
|
||||
warp_denoising_step=not args.no_warp,
|
||||
seed=args.renoise_seed,
|
||||
)
|
||||
if spatial_mode == "refine":
|
||||
schedule, warped_steps = build_wan22_dmd_schedule(
|
||||
steps, flow_shift=args.flow_shift, warp_denoising_step=not args.no_warp,
|
||||
)
|
||||
# The grid opens at sigma == 1, where the hand-off
|
||||
# `(1 - sigma) * upsampled + sigma * noise` weights stage 1 at zero and
|
||||
# refine silently becomes a plain full-res run. Drop the leading
|
||||
# full-noise steps so stage 1 actually reaches stage 2.
|
||||
stage2_warped = default_refine_timesteps(schedule, warped_steps)
|
||||
stage2_steps = steps[len(warped_steps) - len(stage2_warped):]
|
||||
sigma = schedule.sigma_for(stage2_warped[0])
|
||||
print(f"[5B refine] stage-2 steps={stage2_steps} sigma={sigma:.4f} "
|
||||
f"(stage-1 weight {1.0 - sigma:.4f})", flush=True)
|
||||
latents = prepare_refine_latents(
|
||||
latents, scale=args.refine_scale, sigma=sigma,
|
||||
add_noise_flag=not args.no_refine_add_noise,
|
||||
upsample_mode=args.refine_upsample_mode, seed=args.renoise_seed + 1,
|
||||
)
|
||||
freqs_stage2 = make_rotary_embeddings(
|
||||
config, latent_frames=lat_t,
|
||||
latent_height=active_plan.stage2_latent_height,
|
||||
latent_width=active_plan.stage2_latent_width,
|
||||
)
|
||||
latents = sample_wan22_dmd(
|
||||
dit, ehs, latents, freqs_stage2, dmd_denoising_steps=stage2_steps,
|
||||
flow_shift=args.flow_shift, warp_denoising_step=not args.no_warp,
|
||||
seed=args.renoise_seed + 2,
|
||||
)
|
||||
# spatial_mode == "fast_spatial" leaves the latents on the stage-1 grid;
|
||||
# the resample happens after decode, in _postprocess_video.
|
||||
denoise_s = time.perf_counter() - t2
|
||||
peak = mx.get_peak_memory() / (1024**3)
|
||||
print(f"[5B] denoise {len(steps)} steps in {denoise_s:.1f}s, peak {peak:.2f} GiB", flush=True)
|
||||
|
||||
latents_np = np.array(latents.astype(mx.float32))
|
||||
if args.save_latents is not None:
|
||||
args.save_latents.parent.mkdir(parents=True, exist_ok=True)
|
||||
np.savez(args.save_latents, latents=latents_np, prompt=args.prompt, seed=args.seed)
|
||||
print(f"[5B] wrote latents {args.save_latents}", flush=True)
|
||||
|
||||
if spatial_mode == "refine":
|
||||
del freqs_stage2
|
||||
del dit, latents, ehs, noise, freqs
|
||||
cleanup_mlx()
|
||||
|
||||
metrics = decode_latents_to_video(
|
||||
latents_np,
|
||||
args.output_path,
|
||||
fps=args.fps,
|
||||
backend=args.decode_backend,
|
||||
vae_dir=args.vae_root if args.decode_backend == "wan-vae" else None,
|
||||
z_dim=in_ch,
|
||||
)
|
||||
# One h264 round-trip for both post-decode passes (see _postprocess_video).
|
||||
rife_s = 0.0
|
||||
rife_request = ({
|
||||
"factor": args.fast_factor,
|
||||
"target_frames": target_frames,
|
||||
"sharpen": args.fast_sharpen,
|
||||
} if args.fast else None)
|
||||
spatial_request = fast_spatial_plan if spatial_mode == "fast_spatial" else None
|
||||
if rife_request is not None or spatial_request is not None:
|
||||
rife_start = time.perf_counter()
|
||||
_postprocess_video(
|
||||
video_path=args.output_path, fps=args.fps,
|
||||
rife=rife_request, spatial=spatial_request,
|
||||
)
|
||||
rife_s = time.perf_counter() - rife_start
|
||||
print(f"[5B] decoded via {metrics['backend']} in {metrics['decode_s']:.1f}s → {args.output_path}", flush=True)
|
||||
summary = {
|
||||
"output_path": str(args.output_path.resolve()),
|
||||
"prompt": args.prompt,
|
||||
"prompt_used": prompt_for_encode,
|
||||
"enhance_prompt": args.enhance_prompt,
|
||||
"enhance_backend": enhance_backend,
|
||||
"enhance_elapsed_s": round(enhance_elapsed_s, 3),
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"fps": args.fps,
|
||||
"target_frames": target_frames,
|
||||
"generated_frames": args.num_frames,
|
||||
"seed": args.seed,
|
||||
"renoise_seed": args.renoise_seed,
|
||||
"dmd_denoising_steps": steps,
|
||||
"flow_shift": args.flow_shift,
|
||||
"warp": not args.no_warp,
|
||||
"spatial_mode": spatial_mode,
|
||||
"fast": args.fast,
|
||||
"fast_factor": args.fast_factor if args.fast else None,
|
||||
"fast_spatial_scale": args.fast_spatial_scale if args.fast_spatial else None,
|
||||
"refine_scale": args.refine_scale if args.refine else None,
|
||||
"decode_backend": args.decode_backend,
|
||||
"prompt_encode_s": round(prompt_encode_s, 3),
|
||||
"dit_load_s": round(dit_load_s, 3),
|
||||
"denoise_s": round(denoise_s, 3),
|
||||
"decode_s": round(metrics["decode_s"], 3),
|
||||
"rife_s": round(rife_s, 3),
|
||||
"wall_total_s": round(time.perf_counter() - total_start, 3),
|
||||
"peak_gib": round(peak, 3),
|
||||
"latent_shape": [in_ch, lat_t, lat_h, lat_w],
|
||||
"stage2_latent_shape": [in_ch, lat_t, active_plan.stage2_latent_height, active_plan.stage2_latent_width],
|
||||
"mlx_checkpoint": str(args.mlx_checkpoint.resolve()) if args.mlx_checkpoint else None,
|
||||
}
|
||||
if args.metrics_json is not None:
|
||||
args.metrics_json.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.metrics_json.write_text(json.dumps(summary, indent=2) + "\n")
|
||||
print(f"[5B] wrote metrics {args.metrics_json}", flush=True)
|
||||
print(json.dumps(summary, indent=2), flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,130 @@
|
||||
"""Compare Wan VAE and TAEHV decode on saved FastWan latents."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
from examples.inference.basic.mlx_wan_prompt_to_video import DEFAULT_MODEL_ROOT, decode_latents_to_video
|
||||
|
||||
|
||||
def _torch_mps_memory() -> dict[str, int | None]:
|
||||
"""
|
||||
Report current and recommended memory usage for the MPS backend.
|
||||
|
||||
Returns:
|
||||
dict[str, int | None]: Memory metrics in bytes, or `None` values when
|
||||
PyTorch or MPS is unavailable.
|
||||
"""
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
return {
|
||||
"current_allocated_bytes": None,
|
||||
"driver_allocated_bytes": None,
|
||||
"recommended_max_bytes": None,
|
||||
}
|
||||
if not torch.backends.mps.is_available():
|
||||
return {
|
||||
"current_allocated_bytes": None,
|
||||
"driver_allocated_bytes": None,
|
||||
"recommended_max_bytes": None,
|
||||
}
|
||||
return {
|
||||
"current_allocated_bytes": int(torch.mps.current_allocated_memory()),
|
||||
"driver_allocated_bytes": int(torch.mps.driver_allocated_memory()),
|
||||
"recommended_max_bytes": int(torch.mps.recommended_max_memory()),
|
||||
}
|
||||
|
||||
|
||||
def _parse_backends(raw: str) -> list[str]:
|
||||
"""
|
||||
Parse and validate a comma-separated list of decoding backends.
|
||||
|
||||
Parameters:
|
||||
raw (str): Comma-separated backend names.
|
||||
|
||||
Returns:
|
||||
list[str]: Trimmed, supported backend names in input order.
|
||||
|
||||
Raises:
|
||||
ValueError: If the input contains an unsupported backend.
|
||||
"""
|
||||
backends = [backend.strip() for backend in raw.split(",") if backend.strip()]
|
||||
allowed = {"wan-vae", "taehv"}
|
||||
unknown = sorted(set(backends) - allowed)
|
||||
if unknown:
|
||||
raise ValueError(f"Unsupported decode backends: {unknown}")
|
||||
return backends
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""
|
||||
Benchmark selected Wan latent decoding backends and record their performance metrics.
|
||||
|
||||
Loads the specified latent array, decodes it with each selected backend, exports the
|
||||
results as MP4 files, and writes per-backend timing and Torch MPS memory metrics to
|
||||
`metrics.json`.
|
||||
"""
|
||||
parser = argparse.ArgumentParser(description="Benchmark decode backends on saved Wan/FastWan latents.")
|
||||
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
|
||||
parser.add_argument("--latents-path", type=Path, required=True)
|
||||
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_decode_benchmark"))
|
||||
parser.add_argument("--backends", default="wan-vae,taehv")
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--torch-device", default="auto")
|
||||
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
|
||||
parser.add_argument("--taehv-source-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-parallel", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
latents = np.load(args.latents_path)
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
rows = []
|
||||
|
||||
for backend in _parse_backends(args.backends):
|
||||
print(f"=== Decode backend: {backend} ===")
|
||||
before = _torch_mps_memory()
|
||||
start = time.perf_counter()
|
||||
output_path = args.output_dir / f"{args.latents_path.stem}_{backend}.mp4"
|
||||
decode_latents_to_video(
|
||||
model_root=args.model_root,
|
||||
latents_np=latents,
|
||||
output_path=output_path,
|
||||
fps=args.fps,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
backend=backend,
|
||||
taehv_source_path=args.taehv_source_path,
|
||||
taehv_checkpoint_path=args.taehv_checkpoint_path,
|
||||
taehv_parallel=args.taehv_parallel,
|
||||
)
|
||||
elapsed = time.perf_counter() - start
|
||||
after = _torch_mps_memory()
|
||||
metrics = {
|
||||
"backend": backend,
|
||||
"latents_path": str(args.latents_path),
|
||||
"latents_shape": list(latents.shape),
|
||||
"decode_export_s": elapsed,
|
||||
"torch_mps_current_before_bytes": before["current_allocated_bytes"],
|
||||
"torch_mps_current_after_bytes": after["current_allocated_bytes"],
|
||||
"torch_mps_driver_before_bytes": before["driver_allocated_bytes"],
|
||||
"torch_mps_driver_after_bytes": after["driver_allocated_bytes"],
|
||||
"torch_mps_recommended_max_bytes": after["recommended_max_bytes"],
|
||||
"output_path": str(output_path),
|
||||
}
|
||||
rows.append(metrics)
|
||||
print(json.dumps(metrics, indent=2))
|
||||
|
||||
metrics_path = args.output_dir / "metrics.json"
|
||||
metrics_path.write_text(json.dumps(rows, indent=2))
|
||||
print(f"Wrote decode metrics to: {metrics_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,365 @@
|
||||
"""Benchmark MLX FastWan quantization modes with one shared prompt encode."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
|
||||
import numpy as np
|
||||
|
||||
from examples.inference.basic.mlx_wan_prompt_to_video import (
|
||||
DEFAULT_MODEL_ROOT,
|
||||
decode_latents_to_video,
|
||||
encode_prompt,
|
||||
make_rotary_embeddings,
|
||||
)
|
||||
from fastvideo.mlx_runtime.memory import cleanup_mlx
|
||||
|
||||
|
||||
def _parse_modes(raw: str) -> list[str]:
|
||||
"""
|
||||
Parse and validate a comma-separated list of quantization modes.
|
||||
|
||||
Parameters:
|
||||
raw (str): Comma-separated mode names.
|
||||
|
||||
Returns:
|
||||
list[str]: Normalized, whitespace-trimmed mode names.
|
||||
|
||||
Raises:
|
||||
ValueError: If any mode is unsupported.
|
||||
"""
|
||||
modes = [mode.strip() for mode in raw.split(",") if mode.strip()]
|
||||
allowed = {"none", "int8", "int4", "mxfp8", "mxfp4", "nvfp4"}
|
||||
unknown = sorted(set(modes) - allowed)
|
||||
if unknown:
|
||||
raise ValueError(f"Unsupported modes: {unknown}")
|
||||
return modes
|
||||
|
||||
|
||||
def _latent_delta_metrics(candidate: np.ndarray, baseline: np.ndarray) -> dict[str, float]:
|
||||
"""
|
||||
Compare candidate and baseline latent arrays using error and signal-quality metrics.
|
||||
|
||||
Parameters:
|
||||
candidate (np.ndarray): Latent array to evaluate.
|
||||
baseline (np.ndarray): Reference latent array for comparison.
|
||||
|
||||
Returns:
|
||||
dict[str, float]: Mean squared error, mean absolute error, maximum absolute
|
||||
error, and signal-to-noise ratio in decibels between the arrays.
|
||||
"""
|
||||
diff = candidate.astype(np.float32) - baseline.astype(np.float32)
|
||||
mse = float(np.mean(np.square(diff)))
|
||||
mae = float(np.mean(np.abs(diff)))
|
||||
max_abs = float(np.max(np.abs(diff)))
|
||||
signal = float(np.mean(np.square(baseline.astype(np.float32))))
|
||||
return {
|
||||
"latent_mse_vs_fp16": mse,
|
||||
"latent_mae_vs_fp16": mae,
|
||||
"latent_max_abs_vs_fp16": max_abs,
|
||||
"latent_snr_db_vs_fp16": float(10.0 * np.log10(signal / mse)) if mse > 0 else float("inf"),
|
||||
}
|
||||
|
||||
|
||||
def _torch_mps_memory() -> dict[str, int | None]:
|
||||
"""
|
||||
Report PyTorch MPS memory statistics when PyTorch MPS is available.
|
||||
|
||||
Returns:
|
||||
dict[str, int | None]: A mapping of MPS memory metric names to byte counts, or `None` values when PyTorch or MPS is unavailable.
|
||||
"""
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
return {
|
||||
"torch_mps_current_allocated_bytes": None,
|
||||
"torch_mps_driver_allocated_bytes": None,
|
||||
"torch_mps_recommended_max_bytes": None,
|
||||
}
|
||||
if not torch.backends.mps.is_available():
|
||||
return {
|
||||
"torch_mps_current_allocated_bytes": None,
|
||||
"torch_mps_driver_allocated_bytes": None,
|
||||
"torch_mps_recommended_max_bytes": None,
|
||||
}
|
||||
return {
|
||||
"torch_mps_current_allocated_bytes": int(torch.mps.current_allocated_memory()),
|
||||
"torch_mps_driver_allocated_bytes": int(torch.mps.driver_allocated_memory()),
|
||||
"torch_mps_recommended_max_bytes": int(torch.mps.recommended_max_memory()),
|
||||
}
|
||||
|
||||
|
||||
def _decode_with_metrics(*, args, latents: np.ndarray, output_path: Path) -> dict[str, float | int | None | str]:
|
||||
"""
|
||||
Decode latents to a video and collect export timing and PyTorch MPS memory metrics.
|
||||
|
||||
Parameters:
|
||||
args: Configuration values for decoding and video export.
|
||||
latents (np.ndarray): Latent representation to decode.
|
||||
output_path (Path): Destination path for the exported video.
|
||||
|
||||
Returns:
|
||||
dict[str, float | int | None | str]: Video export duration and PyTorch MPS memory measurements.
|
||||
"""
|
||||
before = _torch_mps_memory()
|
||||
decode_start = time.perf_counter()
|
||||
decode_latents_to_video(
|
||||
model_root=args.model_root,
|
||||
latents_np=latents,
|
||||
output_path=output_path,
|
||||
fps=args.fps,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
backend=args.decode_backend,
|
||||
taehv_source_path=args.taehv_source_path,
|
||||
taehv_checkpoint_path=args.taehv_checkpoint_path,
|
||||
taehv_parallel=args.taehv_parallel,
|
||||
)
|
||||
decode_time = time.perf_counter() - decode_start
|
||||
after = _torch_mps_memory()
|
||||
return {
|
||||
"decode_export_s": decode_time,
|
||||
"decode_torch_mps_current_before_bytes": before["torch_mps_current_allocated_bytes"],
|
||||
"decode_torch_mps_current_after_bytes": after["torch_mps_current_allocated_bytes"],
|
||||
"decode_torch_mps_driver_before_bytes": before["torch_mps_driver_allocated_bytes"],
|
||||
"decode_torch_mps_driver_after_bytes": after["torch_mps_driver_allocated_bytes"],
|
||||
"decode_torch_mps_recommended_max_bytes": after["torch_mps_recommended_max_bytes"],
|
||||
}
|
||||
|
||||
|
||||
def _run_one_mode(
|
||||
*,
|
||||
mode: str,
|
||||
args,
|
||||
config: dict,
|
||||
checkpoint_path: Path,
|
||||
config_path: Path,
|
||||
prompt_embeds,
|
||||
freqs_cis,
|
||||
):
|
||||
"""
|
||||
Run denoising for one quantization mode and collect performance and memory metrics.
|
||||
|
||||
Parameters:
|
||||
mode (str): Quantization mode to benchmark.
|
||||
args: Benchmark configuration, including dtype, dimensions, seed, scheduler, and denoising settings.
|
||||
config (dict): Model configuration containing the input channel count.
|
||||
checkpoint_path (Path): Path to the transformer checkpoint.
|
||||
config_path (Path): Path to the transformer configuration.
|
||||
prompt_embeds: Encoded prompt embeddings shared across benchmark modes.
|
||||
freqs_cis: Rotary positional embeddings used during denoising.
|
||||
|
||||
Returns:
|
||||
dict: The mode name, generated latent array, and metrics for model loading,
|
||||
denoising, step timing, and MLX memory usage.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
from fastvideo.benchmarks.mlx_fastwan_bench import denoise_dmd_on_device
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.mlx_runtime.fastwan import mlx_dit_from_diffusers_safetensors
|
||||
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step
|
||||
|
||||
mx_dtype = mx.float16 if args.mlx_dtype == "fp16" else mx.float32
|
||||
quantization = None if mode == "none" else mode
|
||||
latent_frames = (args.num_frames - 1) // 4 + 1
|
||||
latent_height = args.height // 8
|
||||
latent_width = args.width // 8
|
||||
|
||||
load_start = time.perf_counter()
|
||||
mx.clear_cache()
|
||||
mx.reset_peak_memory()
|
||||
dit = mlx_dit_from_diffusers_safetensors(
|
||||
checkpoint_path,
|
||||
config_path,
|
||||
dtype=args.mlx_dtype,
|
||||
quantization=quantization,
|
||||
)
|
||||
load_time = time.perf_counter() - load_start
|
||||
load_peak_memory = mx.get_peak_memory()
|
||||
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=args.flow_shift)
|
||||
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
|
||||
timesteps = [int(step.strip()) for step in args.dmd_denoising_steps.split(",") if step.strip()]
|
||||
# Same torch generator sequence as the original host-round-trip loop
|
||||
# (initial latents first, then one re-noise draw per intermediate step),
|
||||
# so every mode still shares identical stochasticity.
|
||||
generator = torch.Generator(device="cpu").manual_seed(args.seed)
|
||||
latents_seed = torch.randn(
|
||||
(1, int(config["in_channels"]), latent_frames, latent_height, latent_width),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
).numpy()
|
||||
renoise_by_step = [
|
||||
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
|
||||
for _ in range(max(0, len(timesteps) - 1))
|
||||
]
|
||||
latents = mx.array(latents_seed).astype(mx_dtype)
|
||||
encoder_hidden_states = mx.array(prompt_embeds.numpy()).astype(mx_dtype)
|
||||
|
||||
denoise_start = time.perf_counter()
|
||||
mx.reset_peak_memory()
|
||||
latents_np, step_times = denoise_dmd_on_device(
|
||||
mx=mx,
|
||||
dit=dit,
|
||||
latents=latents,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
freqs_cis=freqs_cis,
|
||||
timesteps=timesteps,
|
||||
renoise_by_step=renoise_by_step,
|
||||
schedule=schedule,
|
||||
dmd_step=dmd_step,
|
||||
mx_dtype=mx_dtype,
|
||||
)
|
||||
denoise_time = time.perf_counter() - denoise_start
|
||||
denoise_peak_memory = mx.get_peak_memory()
|
||||
active_memory = mx.get_active_memory()
|
||||
return {
|
||||
"mode": mode,
|
||||
"latents": latents_np,
|
||||
"metrics": {
|
||||
"mlx_dit_load_s": load_time,
|
||||
"mlx_denoise_s": denoise_time,
|
||||
"mlx_denoise_first_step_s": step_times[0] if step_times else None,
|
||||
"mlx_load_peak_bytes": int(load_peak_memory),
|
||||
"mlx_denoise_peak_bytes": int(denoise_peak_memory),
|
||||
"mlx_active_after_denoise_bytes": int(active_memory),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""
|
||||
Run the MLX FastWan quantization benchmark for the selected modes and write latency, memory, output, and latent-difference metrics to the output directory.
|
||||
"""
|
||||
parser = argparse.ArgumentParser(description="Benchmark MLX FastWan quantization modes.")
|
||||
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
|
||||
parser.add_argument("--prompt", default="A snow leopard walks across a windy mountain ridge.")
|
||||
parser.add_argument("--height", type=int, default=192)
|
||||
parser.add_argument("--width", type=int, default=320)
|
||||
parser.add_argument("--num-frames", type=int, default=17)
|
||||
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
|
||||
parser.add_argument("--flow-shift", type=float, default=8.0)
|
||||
parser.add_argument("--max-sequence-length", type=int, default=256)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--torch-device", default="auto")
|
||||
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
|
||||
parser.add_argument("--mlx-dtype", choices=("fp16", "fp32"), default="fp16")
|
||||
parser.add_argument("--modes", default="none,int8,int4,mxfp8,mxfp4,nvfp4")
|
||||
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_quant_benchmark"))
|
||||
parser.add_argument("--decode-backend", choices=("none", "wan-vae", "taehv"), default="taehv")
|
||||
parser.add_argument("--taehv-source-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-parallel", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
mx.random.seed(args.seed)
|
||||
torch.manual_seed(args.seed)
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
config_path = args.model_root / "transformer/config.json"
|
||||
checkpoint_path = args.model_root / "transformer/diffusion_pytorch_model.safetensors"
|
||||
config = json.loads(config_path.read_text())
|
||||
latent_frames = (args.num_frames - 1) // 4 + 1
|
||||
latent_height = args.height // 8
|
||||
latent_width = args.width // 8
|
||||
|
||||
prompt_start = time.perf_counter()
|
||||
prompt_embeds = encode_prompt(
|
||||
model_root=args.model_root,
|
||||
prompt=args.prompt,
|
||||
max_sequence_length=args.max_sequence_length,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
)
|
||||
prompt_time = time.perf_counter() - prompt_start
|
||||
freqs_cis = make_rotary_embeddings(
|
||||
config,
|
||||
latent_frames=latent_frames,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
|
||||
from fastvideo.mlx_runtime.fastwan import UnsupportedMLXQuantizationError
|
||||
|
||||
baseline_latents = None
|
||||
rows = []
|
||||
for mode in _parse_modes(args.modes):
|
||||
print(f"=== MLX quant mode: {mode} ===")
|
||||
mode_start = time.perf_counter()
|
||||
try:
|
||||
result = _run_one_mode(
|
||||
mode=mode,
|
||||
args=args,
|
||||
config=config,
|
||||
checkpoint_path=checkpoint_path,
|
||||
config_path=config_path,
|
||||
prompt_embeds=prompt_embeds,
|
||||
freqs_cis=freqs_cis,
|
||||
)
|
||||
except UnsupportedMLXQuantizationError as exc:
|
||||
print(f"skipping mode (unsupported by this MLX build): {exc}")
|
||||
rows.append({"mode": mode, "status": "unsupported_by_mlx", "error": str(exc)})
|
||||
continue
|
||||
cleanup_mlx(mx)
|
||||
latents = result["latents"]
|
||||
if baseline_latents is None:
|
||||
baseline_latents = latents
|
||||
latent_path = args.output_dir / f"latents_{mode}.npy"
|
||||
np.save(latent_path, latents)
|
||||
|
||||
decode_time = 0.0
|
||||
decode_metrics = {}
|
||||
output_path = None
|
||||
if args.decode_backend != "none":
|
||||
output_path = args.output_dir / f"video_{mode}_{args.decode_backend}_{args.height}x{args.width}x{args.num_frames}.mp4"
|
||||
decode_metrics = _decode_with_metrics(args=args, latents=latents, output_path=output_path)
|
||||
decode_time = cast(float, decode_metrics["decode_export_s"])
|
||||
|
||||
mode_total = time.perf_counter() - mode_start
|
||||
mlx_denoise_peak_bytes = int(result["metrics"]["mlx_denoise_peak_bytes"])
|
||||
mlx_active_bytes = int(result["metrics"]["mlx_active_after_denoise_bytes"])
|
||||
metrics = {
|
||||
"mode": mode,
|
||||
"status": "ok",
|
||||
"prompt_encode_shared_s": prompt_time,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_frames": args.num_frames,
|
||||
"decode_backend": args.decode_backend,
|
||||
"decode_export_s": decode_time,
|
||||
"mode_total_excluding_shared_prompt_s": mode_total,
|
||||
"mode_total_including_shared_prompt_s": mode_total + prompt_time,
|
||||
"latents_path": str(latent_path),
|
||||
"output_path": str(output_path) if output_path else None,
|
||||
"mlx_denoise_peak_gib": mlx_denoise_peak_bytes / (1024**3),
|
||||
"mlx_active_after_denoise_gib": mlx_active_bytes / (1024**3),
|
||||
"mlx_dit_peak_under_16gb": mlx_denoise_peak_bytes < 16 * 1024**3,
|
||||
"mlx_dit_active_under_16gb": mlx_active_bytes < 16 * 1024**3,
|
||||
"mac_16gb_status": (
|
||||
"dit_memory_fits_16gb_measured_decode_separately"
|
||||
if mlx_denoise_peak_bytes < 16 * 1024**3 else "dit_memory_exceeds_16gb"
|
||||
),
|
||||
**result["metrics"],
|
||||
**decode_metrics,
|
||||
**_latent_delta_metrics(latents, baseline_latents),
|
||||
}
|
||||
rows.append(metrics)
|
||||
print(json.dumps(metrics, indent=2))
|
||||
|
||||
metrics_path = args.output_dir / "metrics.json"
|
||||
metrics_path.write_text(json.dumps(rows, indent=2))
|
||||
print(f"Wrote benchmark metrics to: {metrics_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Compare generated MP4s against a reference MP4 with simple pixel metrics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def _read_video(path: Path) -> np.ndarray:
|
||||
"""
|
||||
Read all frames from a video file as an RGB NumPy array.
|
||||
|
||||
Parameters:
|
||||
path (Path): Path to the video file.
|
||||
|
||||
Returns:
|
||||
np.ndarray: Video frames stacked along the first axis.
|
||||
|
||||
Raises:
|
||||
ValueError: If the video contains no readable frames.
|
||||
"""
|
||||
import cv2
|
||||
|
||||
cap = cv2.VideoCapture(str(path))
|
||||
frames = []
|
||||
try:
|
||||
while True:
|
||||
ok, frame_bgr = cap.read()
|
||||
if not ok:
|
||||
break
|
||||
frame_rgb = cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame_rgb)
|
||||
finally:
|
||||
cap.release()
|
||||
if not frames:
|
||||
raise ValueError(f"No frames read from {path}")
|
||||
return np.stack(frames, axis=0)
|
||||
|
||||
|
||||
def _metrics(candidate: np.ndarray, reference: np.ndarray) -> dict[str, float | int | list[int]]:
|
||||
"""
|
||||
Compute pixel-level comparison metrics between candidate and reference video frames.
|
||||
|
||||
Parameters:
|
||||
candidate (np.ndarray): Candidate video frames in frame, height, width, and channel order.
|
||||
reference (np.ndarray): Reference video frames with the same shape as the candidate.
|
||||
|
||||
Returns:
|
||||
dict[str, float | int | list[int]]: Frame dimensions and pixel comparison metrics, including MSE, MAE, maximum absolute difference, and PSNR in decibels.
|
||||
|
||||
Raises:
|
||||
ValueError: If the candidate and reference arrays have different shapes.
|
||||
"""
|
||||
if candidate.shape != reference.shape:
|
||||
raise ValueError(f"Shape mismatch: candidate={candidate.shape}, reference={reference.shape}")
|
||||
candidate_f = candidate.astype(np.float32)
|
||||
reference_f = reference.astype(np.float32)
|
||||
diff = candidate_f - reference_f
|
||||
mse = float(np.mean(np.square(diff)))
|
||||
mae = float(np.mean(np.abs(diff)))
|
||||
max_abs = float(np.max(np.abs(diff)))
|
||||
psnr = float(20.0 * np.log10(255.0 / np.sqrt(mse))) if mse > 0 else float("inf")
|
||||
return {
|
||||
"frames": int(candidate.shape[0]),
|
||||
"height": int(candidate.shape[1]),
|
||||
"width": int(candidate.shape[2]),
|
||||
"channels": int(candidate.shape[3]),
|
||||
"mse_vs_reference": mse,
|
||||
"mae_vs_reference": mae,
|
||||
"max_abs_vs_reference": max_abs,
|
||||
"psnr_db_vs_reference": psnr,
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Compare candidate MP4 videos with a reference and write pixel-level metrics to a JSON file."""
|
||||
parser = argparse.ArgumentParser(description="Compare MP4s against a reference MP4.")
|
||||
parser.add_argument("--reference", type=Path, required=True)
|
||||
parser.add_argument("--candidates", type=Path, nargs="+", required=True)
|
||||
parser.add_argument("--metrics-json", type=Path, required=True)
|
||||
args = parser.parse_args()
|
||||
|
||||
reference = _read_video(args.reference)
|
||||
rows = []
|
||||
for candidate_path in args.candidates:
|
||||
candidate = _read_video(candidate_path)
|
||||
row = {
|
||||
"reference_path": str(args.reference),
|
||||
"candidate_path": str(candidate_path),
|
||||
**_metrics(candidate, reference),
|
||||
}
|
||||
rows.append(row)
|
||||
print(json.dumps(row, indent=2))
|
||||
|
||||
args.metrics_json.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.metrics_json.write_text(json.dumps(rows, indent=2))
|
||||
print(f"Wrote video quality metrics to: {args.metrics_json}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -35,13 +35,11 @@ import torch
|
||||
# Generation — make one LTX2 video to evaluate.
|
||||
# ---------------------------------------------------------------------
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\""
|
||||
)
|
||||
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\"")
|
||||
|
||||
OUTPUT_PATH = "fastvideo/tests/eval/asset/ltx2.mp4"
|
||||
N_DUP = 4 # how many times to duplicate the video for the gen/ref corpora
|
||||
@@ -77,6 +75,7 @@ def generate_one_ltx2_video() -> str:
|
||||
# Eval — the point of the script. 4 lines from "two paths" to results.
|
||||
# ---------------------------------------------------------------------
|
||||
|
||||
|
||||
def _all_registered_metrics() -> list[str]:
|
||||
"""Every metric in the registry, sorted. Combined with
|
||||
``skip_missing_deps=True`` this is the "run everything that works in
|
||||
@@ -103,8 +102,7 @@ def score_all_metrics(video_path: str) -> None:
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.reset_peak_memory_stats()
|
||||
t_init0 = time.perf_counter()
|
||||
ev = create_evaluator(metrics=_all_registered_metrics(),
|
||||
device="cuda:0", num_gpus=1, skip_missing_deps=True)
|
||||
ev = create_evaluator(metrics=_all_registered_metrics(), device="cuda:0", num_gpus=1, skip_missing_deps=True)
|
||||
t_init1 = time.perf_counter()
|
||||
samples = samples_from(video=gen_dir, reference=ref_dir, text_prompt=PROMPT, fps=24.0,
|
||||
extract_audio=True) # auto-extract audio track from videos
|
||||
|
||||
@@ -23,14 +23,12 @@ import torch
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.eval import create_evaluator
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\""
|
||||
)
|
||||
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\"")
|
||||
|
||||
METRICS = [
|
||||
"audio.clap_score",
|
||||
|
||||
@@ -25,19 +25,17 @@ from fastvideo import VideoGenerator
|
||||
from fastvideo.eval import Evaluator
|
||||
from fastvideo.eval.io import build_eval_kwargs
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic.")
|
||||
|
||||
# VBench sub-metrics meaningful for an arbitrary text→video sample
|
||||
# (just the generated frames, optionally fps + the source prompt).
|
||||
@@ -45,14 +43,14 @@ PROMPT = (
|
||||
# vbench.scene, ...) are excluded — they need prompts built to a
|
||||
# specific schema.
|
||||
METRICS = [
|
||||
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
|
||||
"vbench.subject_consistency", # DINO frame-to-first cosine
|
||||
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
|
||||
"vbench.subject_consistency", # DINO frame-to-first cosine
|
||||
"vbench.background_consistency", # DINO on background patches
|
||||
"vbench.imaging_quality", # pyiqa MUSIQ
|
||||
"vbench.temporal_flickering", # pixel-wise frame deltas
|
||||
"vbench.motion_smoothness", # AMT frame interpolator residual
|
||||
"vbench.dynamic_degree", # RAFT optical-flow magnitude (needs fps)
|
||||
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
|
||||
"vbench.imaging_quality", # pyiqa MUSIQ
|
||||
"vbench.temporal_flickering", # pixel-wise frame deltas
|
||||
"vbench.motion_smoothness", # AMT frame interpolator residual
|
||||
"vbench.dynamic_degree", # RAFT optical-flow magnitude (needs fps)
|
||||
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -41,9 +41,8 @@ def _expected_filename(row: dict) -> str:
|
||||
return row["auxiliary_info"]["expected_gen_filename"]
|
||||
|
||||
|
||||
def _generate_videos(rows: list[dict], videos_dir: Path,
|
||||
model: str, num_gpus: int,
|
||||
num_frames: int, height: int, width: int) -> None:
|
||||
def _generate_videos(rows: list[dict], videos_dir: Path, model: str, num_gpus: int, num_frames: int, height: int,
|
||||
width: int) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
@@ -59,36 +58,43 @@ def _generate_videos(rows: list[dict], videos_dir: Path,
|
||||
try:
|
||||
for row, out_path in todo:
|
||||
gen.generate_video(
|
||||
prompt=row["prompt"], output_path=str(out_path), save_video=True,
|
||||
num_frames=num_frames, height=height, width=width,
|
||||
prompt=row["prompt"],
|
||||
output_path=str(out_path),
|
||||
save_video=True,
|
||||
num_frames=num_frames,
|
||||
height=height,
|
||||
width=width,
|
||||
)
|
||||
finally:
|
||||
gen.shutdown()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--dataset-root", type=Path, default=None,
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--dataset-root",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Path to a pre-downloaded Physics-IQ release. "
|
||||
"Defaults to ${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq, "
|
||||
"auto-fetching missing assets from the public bucket.")
|
||||
p.add_argument("--videos-dir", type=Path,
|
||||
"Defaults to ${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq, "
|
||||
"auto-fetching missing assets from the public bucket.")
|
||||
p.add_argument("--videos-dir",
|
||||
type=Path,
|
||||
default=Path("outputs_video/bench_physics_iq"),
|
||||
help="Where to read/write generated videos.")
|
||||
p.add_argument("--limit", type=int, default=None,
|
||||
help="Truncate to first N scenarios for smoke runs.")
|
||||
p.add_argument("--limit", type=int, default=None, help="Truncate to first N scenarios for smoke runs.")
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
p.add_argument("--model",
|
||||
default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the text→video generator to use.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
help="Re-score existing videos under --videos-dir.")
|
||||
p.add_argument("--scores-out", type=Path, default=None,
|
||||
p.add_argument("--skip-generation", action="store_true", help="Re-score existing videos under --videos-dir.")
|
||||
p.add_argument("--scores-out",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Where to write per-scenario scores (JSON). "
|
||||
"Defaults to <videos-dir>/scores.json.")
|
||||
"Defaults to <videos-dir>/scores.json.")
|
||||
args = p.parse_args()
|
||||
|
||||
# 1. Walk the Physics-IQ corpus. Pass --limit to the dataset
|
||||
@@ -100,8 +106,13 @@ def main() -> None:
|
||||
# 2. Generate (or reuse) one mp4 per scenario.
|
||||
if not args.skip_generation:
|
||||
_generate_videos(
|
||||
rows, args.videos_dir, args.model, args.num_gpus,
|
||||
args.num_frames, args.height, args.width,
|
||||
rows,
|
||||
args.videos_dir,
|
||||
args.model,
|
||||
args.num_gpus,
|
||||
args.num_frames,
|
||||
args.height,
|
||||
args.width,
|
||||
)
|
||||
|
||||
# 3. Score each scenario. The metric reads file paths directly out
|
||||
@@ -126,28 +137,26 @@ def main() -> None:
|
||||
|
||||
# 4. Aggregate per the upstream scoring recipe.
|
||||
metric = get_metric("physics_iq")
|
||||
components = metric.aggregate_components(
|
||||
[r["physics_iq"] for r in all_results]
|
||||
)
|
||||
components = metric.aggregate_components([r["physics_iq"] for r in all_results])
|
||||
|
||||
print()
|
||||
print("=== Physics-IQ aggregate ===")
|
||||
for name, value in components.items():
|
||||
print(f" {name:24s} {value:.4f}")
|
||||
|
||||
detailed = [
|
||||
{
|
||||
"scenario": row["auxiliary_info"]["scenario_id"],
|
||||
"view": row["view"],
|
||||
"scenario_name": row["auxiliary_info"]["scenario_name"],
|
||||
"score": results["physics_iq"].score,
|
||||
}
|
||||
for row, results in zip(matched, all_results)
|
||||
]
|
||||
detailed = [{
|
||||
"scenario": row["auxiliary_info"]["scenario_id"],
|
||||
"view": row["view"],
|
||||
"scenario_name": row["auxiliary_info"]["scenario_name"],
|
||||
"score": results["physics_iq"].score,
|
||||
} for row, results in zip(matched, all_results)]
|
||||
out = args.scores_out or (args.videos_dir / "scores.json")
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(
|
||||
{"aggregate": components, "per_scenario": detailed},
|
||||
{
|
||||
"aggregate": components,
|
||||
"per_scenario": detailed
|
||||
},
|
||||
indent=2,
|
||||
))
|
||||
print(f"\n[done] per-scenario scores → {out}")
|
||||
|
||||
@@ -39,9 +39,8 @@ def _slugify(prompt: str, max_len: int = 100) -> str:
|
||||
return re.sub(r"\s+", " ", s) or "output"
|
||||
|
||||
|
||||
def _generate_videos(prompts: list[str], videos_dir: Path,
|
||||
model: str, num_gpus: int,
|
||||
num_frames: int, height: int, width: int) -> None:
|
||||
def _generate_videos(prompts: list[str], videos_dir: Path, model: str, num_gpus: int, num_frames: int, height: int,
|
||||
width: int) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
@@ -57,53 +56,60 @@ def _generate_videos(prompts: list[str], videos_dir: Path,
|
||||
try:
|
||||
for prompt, out_path in todo:
|
||||
gen.generate_video(
|
||||
prompt=prompt, output_path=str(out_path), save_video=True,
|
||||
num_frames=num_frames, height=height, width=width,
|
||||
prompt=prompt,
|
||||
output_path=str(out_path),
|
||||
save_video=True,
|
||||
num_frames=num_frames,
|
||||
height=height,
|
||||
width=width,
|
||||
)
|
||||
finally:
|
||||
gen.shutdown()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--dimensions", default="aesthetic_quality,subject_consistency",
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--dimensions",
|
||||
default="aesthetic_quality,subject_consistency",
|
||||
help="Comma-separated VBench dimensions (or 'all').")
|
||||
p.add_argument("--limit", type=int, default=None,
|
||||
help="Truncate to first N prompts for smoke runs.")
|
||||
p.add_argument("--videos-dir", type=Path,
|
||||
default=Path("outputs_video/bench_vbench"))
|
||||
p.add_argument("--limit", type=int, default=None, help="Truncate to first N prompts for smoke runs.")
|
||||
p.add_argument("--videos-dir", type=Path, default=Path("outputs_video/bench_vbench"))
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
p.add_argument("--model",
|
||||
default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the text→video generator to use.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--fps", type=float, default=24.0,
|
||||
help="Frame-rate annotation passed to fps-aware metrics.")
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
p.add_argument("--fps", type=float, default=24.0, help="Frame-rate annotation passed to fps-aware metrics.")
|
||||
p.add_argument("--skip-generation",
|
||||
action="store_true",
|
||||
help="Re-score existing videos under --videos-dir without "
|
||||
"regenerating.")
|
||||
p.add_argument("--scores-out", type=Path, default=None,
|
||||
"regenerating.")
|
||||
p.add_argument("--scores-out",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Where to dump per-prompt scores as JSON. "
|
||||
"Defaults to <videos-dir>/scores.json.")
|
||||
"Defaults to <videos-dir>/scores.json.")
|
||||
args = p.parse_args()
|
||||
|
||||
# 1. Pull prompts from VBench.
|
||||
dims_arg: list[str] | str = (
|
||||
args.dimensions if args.dimensions == "all"
|
||||
else [d.strip() for d in args.dimensions.split(",") if d.strip()]
|
||||
)
|
||||
dims_arg: list[str] | str = (args.dimensions if args.dimensions == "all" else
|
||||
[d.strip() for d in args.dimensions.split(",") if d.strip()])
|
||||
ds = get_dataset("vbench", dimensions=dims_arg)
|
||||
rows = list(ds)[: args.limit]
|
||||
rows = list(ds)[:args.limit]
|
||||
print(f"[load] VBench: {len(rows)} prompts across {ds.dimensions}")
|
||||
|
||||
# 2. Generate (or reuse) one mp4 per prompt.
|
||||
if not args.skip_generation:
|
||||
_generate_videos(
|
||||
[row["prompt"] for row in rows],
|
||||
args.videos_dir, args.model, args.num_gpus,
|
||||
args.num_frames, args.height, args.width,
|
||||
args.videos_dir,
|
||||
args.model,
|
||||
args.num_gpus,
|
||||
args.num_frames,
|
||||
args.height,
|
||||
args.width,
|
||||
)
|
||||
|
||||
# 3. Score each video against the requested vbench sub-metrics.
|
||||
@@ -122,7 +128,7 @@ def main() -> None:
|
||||
samples.append({
|
||||
"video": str(video_path),
|
||||
"fps": args.fps,
|
||||
**row, # prompt / aux / dims
|
||||
**row, # prompt / aux / dims
|
||||
})
|
||||
matched_rows.append(row)
|
||||
|
||||
|
||||
@@ -35,20 +35,29 @@ from fastvideo.eval import create_evaluator, samples_from
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--gen-dir", type=Path, required=True,
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--gen-dir",
|
||||
type=Path,
|
||||
required=True,
|
||||
help="Directory of generated videos (.mp4, .avi, .mov, .mkv, .gif).")
|
||||
p.add_argument("--reference-dir", type=Path, default=None,
|
||||
p.add_argument("--reference-dir",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Directory of reference videos. Omit to score against the cached "
|
||||
"reference features (built on a previous run).")
|
||||
"reference features (built on a previous run).")
|
||||
p.add_argument("--device", default="cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
p.add_argument("--num-gpus", type=int, default=1,
|
||||
p.add_argument("--num-gpus",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of GPU replicas. >1 fans extraction out across devices.")
|
||||
p.add_argument("--cache-path", type=Path, default=None,
|
||||
p.add_argument("--cache-path",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Override the reference-feature cache path. "
|
||||
"Defaults to ${FASTVIDEO_EVAL_CACHE}/fvd/real_features_i3d.pt.")
|
||||
p.add_argument("--output", type=Path, default=None,
|
||||
"Defaults to ${FASTVIDEO_EVAL_CACHE}/fvd/real_features_i3d.pt.")
|
||||
p.add_argument("--output",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Write the result as JSON to this path (default: stdout only).")
|
||||
args = p.parse_args()
|
||||
|
||||
|
||||
@@ -36,57 +36,55 @@ from fastvideo import VideoGenerator
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.io import load_video
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic.")
|
||||
|
||||
DEFAULT_METRICS = [
|
||||
# No-input metrics: just need the generated frames.
|
||||
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
|
||||
"vbench.subject_consistency", # DINO frame-to-first cosine
|
||||
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
|
||||
"vbench.subject_consistency", # DINO frame-to-first cosine
|
||||
"vbench.background_consistency", # DINO on background patches
|
||||
"vbench.imaging_quality", # pyiqa MUSIQ
|
||||
"vbench.temporal_flickering", # pixel-wise frame deltas
|
||||
"vbench.motion_smoothness", # AMT frame interpolator residual
|
||||
"vbench.imaging_quality", # pyiqa MUSIQ
|
||||
"vbench.temporal_flickering", # pixel-wise frame deltas
|
||||
"vbench.motion_smoothness", # AMT frame interpolator residual
|
||||
# Need fps annotation:
|
||||
"vbench.dynamic_degree", # RAFT optical-flow magnitude
|
||||
"vbench.dynamic_degree", # RAFT optical-flow magnitude
|
||||
# Need the source prompt:
|
||||
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
|
||||
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
|
||||
]
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the LTX2 checkpoint.")
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers", help="HF repo id of the LTX2 checkpoint.")
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--output", default="outputs_video/ltx2_eval/clip.mp4",
|
||||
help="Where to save the generated mp4.")
|
||||
p.add_argument("--output", default="outputs_video/ltx2_eval/clip.mp4", help="Where to save the generated mp4.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--prompt", default=PROMPT)
|
||||
p.add_argument("--fps", type=float, default=24.0,
|
||||
p.add_argument("--fps",
|
||||
type=float,
|
||||
default=24.0,
|
||||
help="Frame-rate annotation passed to fps-aware metrics "
|
||||
"(e.g. vbench.dynamic_degree). LTX2 outputs at 24 fps "
|
||||
"by default.")
|
||||
p.add_argument("--metrics", default=",".join(DEFAULT_METRICS),
|
||||
"(e.g. vbench.dynamic_degree). LTX2 outputs at 24 fps "
|
||||
"by default.")
|
||||
p.add_argument("--metrics",
|
||||
default=",".join(DEFAULT_METRICS),
|
||||
help="Comma-separated metric names. Pass 'all' for every "
|
||||
"registered metric, or e.g. 'vbench' for the whole group.")
|
||||
"registered metric, or e.g. 'vbench' for the whole group.")
|
||||
p.add_argument("--scores-out", default="outputs_video/ltx2_eval/scores.json")
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
p.add_argument("--skip-generation",
|
||||
action="store_true",
|
||||
help="Reuse an existing --output video instead of regenerating.")
|
||||
return p.parse_args()
|
||||
|
||||
@@ -115,11 +113,10 @@ def generate(args: argparse.Namespace) -> Path:
|
||||
return out
|
||||
|
||||
|
||||
def evaluate_video(video_path: Path, prompt: str, fps: float,
|
||||
metric_names) -> dict:
|
||||
def evaluate_video(video_path: Path, prompt: str, fps: float, metric_names) -> dict:
|
||||
print(f"[eval] loading video from {video_path}...")
|
||||
video = load_video(str(video_path)) # (T, C, H, W) in [0, 1]
|
||||
video = video.unsqueeze(0) # → (1, T, C, H, W)
|
||||
video = load_video(str(video_path)) # (T, C, H, W) in [0, 1]
|
||||
video = video.unsqueeze(0) # → (1, T, C, H, W)
|
||||
|
||||
print(f"[eval] building evaluator: {metric_names}")
|
||||
evaluator = create_evaluator(metrics=metric_names, device="cuda")
|
||||
@@ -132,12 +129,9 @@ def evaluate_video(video_path: Path, prompt: str, fps: float,
|
||||
)
|
||||
|
||||
if isinstance(results, list):
|
||||
results = results[0] # batch of 1
|
||||
results = results[0] # batch of 1
|
||||
|
||||
return {
|
||||
name: {"score": r.score, "details": r.details}
|
||||
for name, r in results.items()
|
||||
}
|
||||
return {name: {"score": r.score, "details": r.details} for name, r in results.items()}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
@@ -157,7 +151,11 @@ def main() -> None:
|
||||
out = Path(args.scores_out)
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(
|
||||
{"video": str(video_path), "prompt": args.prompt, "scores": scores},
|
||||
{
|
||||
"video": str(video_path),
|
||||
"prompt": args.prompt,
|
||||
"scores": scores
|
||||
},
|
||||
indent=2,
|
||||
))
|
||||
print(f"[done] scores written to {out}")
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user