Compare commits

..
Author SHA1 Message Date
SolitaryThinker 5f8e84d22e [bugfix]: enforce mobile drawer modality with inert and managed focus 2026-07-29 15:24:08 -07:00
SolitaryThinker 698173402f [bugfix]: render Badge as a span so buttons stay valid HTML 2026-07-29 15:24:08 -07:00
SolitaryThinker 627ca97fda [bugfix]: make Download Captions menu click and keyboard accessible 2026-07-29 15:24:07 -07:00
SolitaryThinker 117302d153 [bugfix]: declare the scoped radix dropdown-menu dependency 2026-07-29 15:24:07 -07:00
RazorCrest00 ec36f7dfd6 [test]: guard Studio focus contrast 2026-07-28 12:53:36 -07:00
RazorCrest00 61f213403d [test]: harden Gallery and responsive Studio flows 2026-07-28 12:53:19 -07:00
RazorCrest00 e6481b691c [bugfix]: use high-contrast Studio focus indicators 2026-07-28 12:52:36 -07:00
RazorCrest00 5d7b8690e1 [bugfix]: meet Studio minimum touch-target sizes 2026-07-28 12:52:21 -07:00
RazorCrest00 e43336d31c [bugfix]: name the actual Studio slider controls 2026-07-28 12:52:01 -07:00
RazorCrest00 a587f1d3e4 [bugfix]: keep one main landmark per Studio route 2026-07-28 12:51:45 -07:00
RazorCrest00 14ad388ee3 [test]: enforce independent Studio card interactions 2026-07-28 12:51:32 -07:00
RazorCrest00 f962352c90 [bugfix]: separate dataset selection from card actions 2026-07-28 12:51:11 -07:00
RazorCrest00 8c1170ca06 [bugfix]: separate job selection from card actions 2026-07-28 12:50:54 -07:00
RazorCrest00 562b1d2361 [bugfix]: expose caption autosave failures and recovery 2026-07-28 12:50:27 -07:00
RazorCrest00 a7c8d8bb79 [bugfix]: preserve and label stale GPU telemetry 2026-07-28 12:50:09 -07:00
RazorCrest00 e85c566f2f [bugfix]: surface create-job operational failures 2026-07-28 12:49:34 -07:00
RazorCrest00 b446a6d9cf [bugfix]: distinguish dataset loading errors from empty lists 2026-07-28 12:48:43 -07:00
RazorCrest00 1d4505509a [bugfix]: distinguish job loading errors from empty queues 2026-07-28 12:47:36 -07:00
RazorCrest00 ec93f73983 [bugfix]: show unavailable dataset previews explicitly 2026-07-28 12:47:28 -07:00
RazorCrest00 6ac2174602 [bugfix]: expose Gallery playback and media failures 2026-07-28 12:47:20 -07:00
RazorCrest00 dff6f7757f [bugfix]: make Studio creation menu click and keyboard accessible 2026-07-28 12:47:09 -07:00
RazorCrest00 0cd4a64aa3 [bugfix]: prevent narrow Studio viewport overflow 2026-07-28 12:47:00 -07:00
RazorCrest00 d65406c48f [bugfix]: use responsive Studio detail drawers 2026-07-28 12:46:52 -07:00
RazorCrest00 842f24d974 [bugfix]: make Studio shell usable on mobile 2026-07-28 12:46:39 -07:00
873 changed files with 24928 additions and 64294 deletions
@@ -10,7 +10,9 @@ 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")
@@ -60,7 +62,9 @@ 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,12 +8,14 @@ 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,6 +10,7 @@ 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 = {
@@ -33,7 +34,8 @@ 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")
@@ -92,12 +94,14 @@ 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(
@@ -211,24 +215,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,6 +18,7 @@ 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")
@@ -34,10 +35,15 @@ 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:
@@ -93,14 +99,18 @@ 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:
@@ -117,7 +127,8 @@ 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()
@@ -176,9 +187,11 @@ 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,6 +27,7 @@ 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] = {}
@@ -46,7 +47,10 @@ 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:
@@ -91,10 +95,11 @@ 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] = []
@@ -112,8 +117,10 @@ def split_monolithic(state: dict[str, torch.Tensor], ) -> dict[str, OrderedDict[
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}
@@ -136,12 +143,8 @@ 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>"},
}
@@ -174,13 +177,19 @@ 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}")
@@ -192,7 +201,9 @@ 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(
@@ -205,7 +216,9 @@ 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:
@@ -248,7 +261,9 @@ 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}")
@@ -256,7 +271,9 @@ 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():
@@ -272,7 +289,9 @@ 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,7 +94,6 @@ 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):
@@ -102,7 +101,6 @@ 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):
@@ -116,7 +114,6 @@ 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.
@@ -134,21 +131,43 @@ 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
@@ -174,9 +193,11 @@ 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:
@@ -184,14 +205,9 @@ 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["):
@@ -200,8 +216,10 @@ 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)):
@@ -217,9 +235,11 @@ 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
@@ -235,8 +255,10 @@ 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,10 +43,12 @@ 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:
@@ -71,8 +73,10 @@ 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:
@@ -142,6 +146,8 @@ 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)
-11
View File
@@ -68,17 +68,6 @@ steps:
limit: 2
agents:
queue: "default"
- label: ":vertical_traffic_light: Golden-Gate Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "golden_gate"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: Unit Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "unit_test"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
-4
View File
@@ -187,10 +187,6 @@ case "$TEST_TYPE" in
log "Running transformer tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
;;
"golden_gate")
log "Running golden-gate tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_golden_gate_tests"
;;
"ssim")
log "Running SSIM tests..."
SSIM_BOOTSTRAP_ARGS=$(ssim_bootstrap_args)
-149
View File
@@ -1,149 +0,0 @@
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
+2 -2
View File
@@ -129,7 +129,7 @@ jobs:
set -euo pipefail
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
VALID="encoder vae transformer kernel unit dreamverse ssim golden-gate training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
exit 1
@@ -138,7 +138,7 @@ jobs:
declare -A MAP=(
[encoder]=encoder [vae]=vae [transformer]=transformer
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
[ssim]=ssim [golden-gate]=golden_gate [training]=training
[ssim]=ssim [training]=training
[lora-inference]=inference_lora [lora-training]=training_lora
[lora-extraction]=lora_extraction
[distillation]=distillation_dmd [self-forcing]=self_forcing
-2
View File
@@ -6,7 +6,6 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
@@ -17,7 +16,6 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
+1 -2
View File
@@ -6,7 +6,7 @@ results/
wandb/
*.ipynb
*.jpg
!examples/datasets/lingbotworld2/image.jpg
!examples/dataset/lingbotworld2/image.jpg
*.safetensors
*.mp4
*.png
@@ -23,7 +23,6 @@ Miniconda3-latest-Linux-x86_64.sh
*validation/
data/
outputs/
outputs_audio/
outputs_video
checkpoints/
sbatch.sh
-1
View File
@@ -22,7 +22,6 @@ 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
+1 -7
View File
@@ -9,7 +9,6 @@
**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/).
@@ -34,7 +33,7 @@ FastVideo has the following features:
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achieve >50x denoising speedup
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing.
- Causal distillation through Self-Forcing
- See this [page](https://hao-ai-lab.github.io/FastVideo/training/overview/) for the supported training workflows, and the [support matrix](https://hao-ai-lab.github.io/FastVideo/inference/support_matrix/) for supported models.
- See this [page](https://hao-ai-lab.github.io/FastVideo/training/overview/) for full list of supported models and recipes.
- State-of-the-art performance optimizations for inference
- Sequence Parallelism for distributed inference
- Multiple state-of-the-art attention backends
@@ -63,11 +62,6 @@ 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,6 +3,7 @@ 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,6 +5,7 @@ from pathlib import Path
import pytest
SERVER_DIR = Path(__file__).resolve().parents[1]
@@ -52,7 +53,9 @@ 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",
@@ -83,7 +86,9 @@ 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",
@@ -101,17 +106,24 @@ 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,7 +9,6 @@ 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"):
@@ -18,7 +17,6 @@ def _install_stack03_import_stubs(monkeypatch):
gpu_pool_stub = types.ModuleType("dreamverse.gpu_pool")
class GPUPool:
def __init__(self, _gpu_ids):
pass
@@ -51,7 +49,6 @@ def _install_stack03_import_stubs(monkeypatch):
controller_stub = types.ModuleType("dreamverse.session.controller")
class SessionController:
def __init__(self, **_kwargs):
pass
@@ -79,11 +76,13 @@ 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)
@@ -100,11 +99,13 @@ 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):
@@ -115,11 +116,13 @@ 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):
@@ -139,11 +142,13 @@ 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):
@@ -156,11 +161,13 @@ 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,6 +7,7 @@ from types import SimpleNamespace
import pytest
import dreamverse.gpu_pool as gpu_pool
@@ -84,7 +85,9 @@ 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,13 +38,19 @@ 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,6 +6,7 @@ import os
from fastapi import WebSocketDisconnect
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
os.environ.setdefault("GROQ_API_KEY", "dummy")
@@ -13,7 +14,6 @@ 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,14 +92,24 @@ 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,
@@ -124,29 +134,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))
@@ -156,7 +166,11 @@ 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]"
@@ -173,40 +187,54 @@ 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
@@ -219,24 +247,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))
@@ -248,9 +276,15 @@ 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
@@ -263,37 +297,35 @@ 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))
@@ -304,11 +336,16 @@ 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,6 +6,7 @@ import os
import re
import time
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
os.environ.setdefault("GROQ_API_KEY", "dummy")
@@ -21,7 +22,6 @@ from dreamverse.prompt_enhancer import (
class _FakeResponse:
def __init__(self, payload: dict):
self._payload = payload
@@ -30,7 +30,6 @@ class _FakeResponse:
class _FakeSyncCompletions:
def __init__(self, payload: dict):
self._payload = payload
@@ -39,7 +38,6 @@ class _FakeSyncCompletions:
class _FakeSyncClient:
def __init__(self, payload: dict):
self.chat = type(
"_FakeChat",
@@ -49,7 +47,6 @@ 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
@@ -64,26 +61,29 @@ 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,7 +172,6 @@ def _build_staged_enhancer(
class _FakeOpenAIClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.chat = type(
@@ -183,7 +182,6 @@ class _FakeOpenAIClient:
class _FakeCerebrasClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.chat = type(
@@ -194,12 +192,16 @@ 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"]}
@@ -266,12 +268,16 @@ 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"
@@ -280,12 +286,15 @@ 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"
@@ -294,14 +303,19 @@ 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"
@@ -315,12 +329,14 @@ 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"
@@ -339,10 +355,12 @@ 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
@@ -363,10 +381,12 @@ 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
@@ -388,11 +408,13 @@ 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
@@ -412,10 +434,12 @@ 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
@@ -429,12 +453,15 @@ 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."
@@ -446,7 +473,9 @@ 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,
@@ -457,7 +486,8 @@ 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"]}',
)
@@ -472,7 +502,8 @@ 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] == {
@@ -481,12 +512,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",
@@ -497,8 +528,11 @@ 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,
}
@@ -507,8 +541,10 @@ 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"]}',
)
@@ -524,29 +560,30 @@ 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,
@@ -556,7 +593,9 @@ 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"]}',
)
@@ -570,7 +609,8 @@ 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] == {
@@ -581,7 +621,10 @@ 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"
@@ -592,17 +635,24 @@ 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"
@@ -630,7 +680,8 @@ 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"
@@ -639,7 +690,9 @@ 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"]
@@ -669,7 +722,8 @@ 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"
@@ -678,7 +732,9 @@ 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"]
@@ -708,12 +764,14 @@ 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"
@@ -726,15 +784,17 @@ 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(
@@ -743,14 +803,17 @@ 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"
@@ -761,14 +824,17 @@ 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"
@@ -780,14 +846,17 @@ 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"
@@ -798,30 +867,34 @@ 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)
@@ -830,7 +903,10 @@ 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"]
@@ -842,7 +918,10 @@ 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"
@@ -869,14 +948,19 @@ 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)
@@ -889,9 +973,14 @@ 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"
@@ -918,7 +1007,10 @@ 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"
@@ -929,7 +1021,9 @@ 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"
@@ -937,7 +1031,10 @@ 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"
@@ -955,7 +1052,9 @@ 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"
@@ -963,7 +1062,10 @@ 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"
@@ -979,7 +1081,9 @@ 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"
@@ -988,7 +1092,10 @@ 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"
@@ -1002,7 +1109,9 @@ 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
@@ -1010,7 +1119,10 @@ 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"
@@ -1024,9 +1136,13 @@ 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,8 +27,13 @@ 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:
@@ -60,7 +65,10 @@ 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]:
@@ -137,16 +145,24 @@ 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)
@@ -208,11 +224,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()
@@ -233,7 +249,9 @@ 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()
@@ -247,7 +265,9 @@ 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"
@@ -268,16 +288,20 @@ 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
@@ -297,7 +321,9 @@ 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"))
@@ -312,13 +338,20 @@ 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:
@@ -329,11 +362,15 @@ 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:
@@ -342,18 +379,29 @@ 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
@@ -364,11 +412,14 @@ 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 []
@@ -386,23 +437,29 @@ 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(
@@ -460,22 +517,33 @@ 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
@@ -486,18 +554,20 @@ 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,
@@ -536,39 +606,51 @@ 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:")
@@ -588,8 +670,10 @@ 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(
@@ -651,8 +735,13 @@ 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())
@@ -676,20 +765,24 @@ 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)
@@ -702,8 +795,9 @@ 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": {
@@ -739,7 +833,9 @@ 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,11 +47,13 @@ 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):
@@ -64,8 +66,10 @@ 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
+15 -9
View File
@@ -8,10 +8,12 @@ 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
@@ -63,10 +65,14 @@ 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"])
+3 -21
View File
@@ -1,9 +1,8 @@
'use client';
import { AlertTriangle, ImageOff, Loader2 } from 'lucide-react';
import { ImageOff, Loader2 } from 'lucide-react';
import { useEffect, useState } from 'react';
import { Button } from '@/components/ui/button';
import { Card } from '@/components/ui/card';
import { getJobVideoUrl, getJobsList } from '@/lib/api';
import type { Job } from '@/lib/types';
@@ -65,8 +64,6 @@ export default function GalleryPage() {
const [isLoading, setIsLoading] = useState(true);
const [error, setError] = useState<string | null>(null);
const [reloadKey, setReloadKey] = useState(0);
useEffect(() => {
let cancelled = false;
async function load() {
@@ -91,13 +88,7 @@ export default function GalleryPage() {
return () => {
cancelled = true;
};
}, [reloadKey]);
function retry() {
setError(null);
setIsLoading(true);
setReloadKey((k) => k + 1);
}
}, []);
const galleryJobs = jobs.filter(
(j) =>
@@ -121,16 +112,7 @@ export default function GalleryPage() {
<span>Loading gallery…</span>
</div>
) : error ? (
<div
role="alert"
className="flex flex-col items-center gap-3 py-8 text-center"
>
<AlertTriangle className="size-6 text-destructive" aria-hidden />
<p className="max-w-md text-sm text-muted-foreground">{error}</p>
<Button type="button" variant="outline" onClick={retry}>
Try Again
</Button>
</div>
<p className="py-8 text-destructive">{error}</p>
) : galleryJobs.length === 0 ? (
<p className="py-8 text-center text-muted-foreground">
No completed videos yet
@@ -0,0 +1,50 @@
import { render, screen } from '@testing-library/react';
import userEvent from '@testing-library/user-event';
import { describe, expect, it, vi } from 'vitest';
import { downloadBlob } from '@/lib/utils';
import DownloadCaptions from './DownloadCaptions';
vi.mock('@/lib/utils', async (importOriginal) => {
const actual = await importOriginal<typeof import('@/lib/utils')>();
return { ...actual, downloadBlob: vi.fn() };
});
const props = {
fileNames: ['b.mp4', 'a.mp4'],
captions: { 'a.mp4': 'cap a', 'b.mp4': 'cap b' },
};
describe('DownloadCaptions', () => {
it('opens the format menu on click and downloads the selection', async () => {
const user = userEvent.setup();
render(<DownloadCaptions {...props} />);
await user.click(
screen.getByRole('button', { name: 'Download Captions' }),
);
await user.click(screen.getByRole('menuitem', { name: 'JSON' }));
expect(vi.mocked(downloadBlob)).toHaveBeenCalledWith(
expect.any(Blob),
'videos2caption.json',
);
});
it('operates entirely from the keyboard', async () => {
const user = userEvent.setup();
render(<DownloadCaptions {...props} />);
screen.getByRole('button', { name: 'Download Captions' }).focus();
await user.keyboard('{Enter}');
const first = await screen.findByRole('menuitem', { name: 'JSON' });
expect(first).toHaveFocus();
await user.keyboard('{ArrowDown}{Enter}');
expect(vi.mocked(downloadBlob)).toHaveBeenCalledWith(
expect.any(Blob),
'videos.txt',
);
});
});
@@ -1,12 +1,13 @@
'use client';
import { ChevronDown } from 'lucide-react';
import * as DropdownMenu from '@radix-ui/react-dropdown-menu';
import { Button } from '@/components/ui/button';
import { downloadBlob } from '@/lib/utils';
const MENU_ITEM =
'block min-h-11 w-full cursor-pointer px-4 py-2 text-left text-sm font-medium text-foreground transition-colors hover:bg-muted disabled:cursor-not-allowed disabled:opacity-50';
'block min-h-11 w-full cursor-pointer select-none px-4 py-2 text-left text-sm font-medium text-foreground outline-none transition-colors data-[highlighted]:bg-muted';
export default function DownloadCaptions({
fileNames,
@@ -62,52 +63,40 @@ export default function DownloadCaptions({
downloadBlob(new Blob([csv], { type: 'text/csv' }), 'captions.csv');
}
// Same click/keyboard-accessible menu idiom as CreateJobButton — hover-only
// menus exclude touch and keyboard users.
return (
<div className="group relative inline-block">
<Button
type="button"
variant="outline"
size="sm"
disabled={disabled}
aria-haspopup="menu"
className="gap-1.5"
>
Download Captions
<ChevronDown className="h-3.5 w-3.5 opacity-85" />
</Button>
{!disabled && (
<div
role="menu"
className="invisible absolute right-0 top-full z-[200] min-w-full -translate-y-1 pt-1 opacity-0 transition-all group-focus-within:visible group-focus-within:translate-y-0 group-focus-within:opacity-100 group-hover:visible group-hover:translate-y-0 group-hover:opacity-100"
<DropdownMenu.Root>
<DropdownMenu.Trigger asChild>
<Button
type="button"
variant="outline"
size="sm"
disabled={disabled}
className="gap-1.5"
>
<div className="overflow-hidden rounded-lg border border-border bg-popover py-1 shadow-lg">
<button
type="button"
role="menuitem"
className={MENU_ITEM}
onClick={handleDownloadJson}
>
JSON
</button>
<button
type="button"
role="menuitem"
className={MENU_ITEM}
onClick={handleDownloadTxt}
>
TXT
</button>
<button
type="button"
role="menuitem"
className={MENU_ITEM}
onClick={handleDownloadCsv}
>
CSV
</button>
</div>
</div>
)}
</div>
Download Captions
<ChevronDown className="h-3.5 w-3.5 opacity-85" aria-hidden />
</Button>
</DropdownMenu.Trigger>
<DropdownMenu.Portal>
<DropdownMenu.Content
align="end"
sideOffset={4}
collisionPadding={8}
className="z-[200] min-w-40 overflow-hidden rounded-lg border border-border bg-popover py-1 text-popover-foreground shadow-lg"
>
<DropdownMenu.Item className={MENU_ITEM} onSelect={handleDownloadJson}>
JSON
</DropdownMenu.Item>
<DropdownMenu.Item className={MENU_ITEM} onSelect={handleDownloadTxt}>
TXT
</DropdownMenu.Item>
<DropdownMenu.Item className={MENU_ITEM} onSelect={handleDownloadCsv}>
CSV
</DropdownMenu.Item>
</DropdownMenu.Content>
</DropdownMenu.Portal>
</DropdownMenu.Root>
);
}
+6 -3
View File
@@ -74,8 +74,10 @@ 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]
@@ -84,7 +86,8 @@ def test_snapshot_tolerates_missing_sensors(monkeypatch: pytest.MonkeyPatch) ->
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,7 +130,8 @@ 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
@@ -160,27 +161,31 @@ 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:
@@ -190,7 +195,8 @@ 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:
@@ -205,7 +211,8 @@ 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
View File
@@ -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=14c377950125c70b7a9dabf9c561fca53715ac7d
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
# 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
-52
View File
@@ -1,52 +0,0 @@
{
"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"
}
]
}
-60
View File
@@ -1,60 +0,0 @@
(() => {
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);
})();
-40
View File
@@ -42,46 +42,6 @@ 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;
-44
View File
@@ -1,44 +0,0 @@
# 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.
-128
View File
@@ -1,128 +0,0 @@
# 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,11 +76,6 @@ 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."
refine_transformer_path: "Generic stage-2 refine transformer override; no typed equivalent yet."
@@ -449,11 +444,7 @@ surfaces:
moved:
image_path: request.inputs.image_path
pil_image: request.inputs.pil_image
last_image: request.inputs.last_image
references: request.inputs.references
video_path: request.inputs.video_path
latents: request.inputs.latents
audio_latents: request.inputs.audio_latents
mouse_cond: request.inputs.mouse_cond
keyboard_cond: request.inputs.keyboard_cond
grid_sizes: request.inputs.grid_sizes
@@ -533,6 +524,7 @@ surfaces:
inpaint_mask: request.extensions.stable_audio.inpaint_mask
internal_only:
data_type: "Derived from the request shape and not a public input."
latents: "Pre-generated diffusion latents supplied by parity/debug harnesses; not a public input."
sampling_param_extensions: {}
-37
View File
@@ -2,7 +2,6 @@
# 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
@@ -20,40 +19,6 @@ 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:
@@ -571,7 +536,6 @@ 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!")
@@ -585,7 +549,6 @@ 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!")
+3 -7
View File
@@ -5,7 +5,6 @@ FastVideo supports the following hardware platforms:
- [NVIDIA CUDA](installation/gpu.md)
- [NVIDIA DGX Spark / GB10 (ARM64 + CUDA 13)](installation/spark.md)
([performance & tuning](installation/spark_performance.md))
- [Apple silicon](installation/mps.md)
## Quick Installation
@@ -55,17 +54,14 @@ UV_TORCH_BACKEND=cu126 uv pip install -e .
uv pip install flash-attn --no-build-isolation -v
```
## Requirements
## Hardware Requirements
- **Python**: 3.10-3.12 is the tested and recommended range (the commands
above pin 3.12)
- **NVIDIA GPUs**: CUDA 12.6+ with compute capability 7.0+
- **Apple Silicon**: macOS 14.0+ with M1/M2/M3/M4 chips
- **CPU**: x86_64 architecture (for CPU-only inference)
## Next Steps
- [Quick Start](quick_start.md) - Generate your first video
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
+4 -6
View File
@@ -49,12 +49,10 @@ brew install ffmpeg
### Installation
FastWan's native Apple Silicon runtime requires the `mlx` extra.
#### With uv (recommended)
```bash
uv pip install "fastvideo[mlx]"
uv pip install fastvideo
```
#### With Conda environment (alternative)
@@ -62,7 +60,7 @@ uv pip install "fastvideo[mlx]"
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
```bash
uv pip install "fastvideo[mlx]"
uv pip install fastvideo
```
### Installation from Source
@@ -78,13 +76,13 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Basic installation:
```bash
uv pip install -e ".[mlx]"
uv pip install -e .
```
Alternative with Conda environment:
```bash
uv pip install -e ".[mlx]"
uv pip install -e .
```
## Development Environment Setup
@@ -137,13 +137,6 @@ If you hit other issues, please open an issue on our
our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)
for additional support.
## Next: performance & tuning
Installed and verified? See [DGX Spark: Performance & Tuning](spark_performance.md)
for which models are practical on the GB10, what makes them faster, and what
won't help on this hardware (and why) — so you don't spend a night tuning knobs
that can't move here.
## Development Environment Setup
If you're planning to contribute to FastVideo please see the
@@ -1,195 +0,0 @@
# DGX Spark (GB10): Performance & Tuning
You have FastVideo [installed on a DGX Spark](spark.md) — this page is what to
run next. It covers **which models are practical on the GB10, what actually
makes them faster, and what won't help (and why)**, so you don't burn a night
tuning knobs that can't move on this hardware.
!!! tip "TL;DR"
- **Use distilled few-step models** (e.g. `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`).
They run in ~40 s/video. Full-step models are 12–47 min on the GB10.
- On few-step models, **VAE decode is the bottleneck**, not attention — it's
bandwidth-bound on the Spark's unified memory.
- **bf16 VAE decode** is the real, lossless lever (FastVideo already turns it
on for Wan). **FlashAttention, linear quantization, and `torch.compile` of
the VAE give little or nothing here** — see the table below.
- Heavy runs can make the box unreachable — run generations with VAE tiling on
and `nice -n 19`. See [Running safely](#running-safely-dont-lock-the-box).
## The hardware reality (this explains everything below)
The GB10 pairs a Blackwell GPU (`sm_121`) with **128 GB of unified LPDDR5X memory
(~270 GB/s) shared between CPU and GPU**. That bandwidth is roughly **10× below a
datacenter GPU's HBM**. Two consequences drive every tuning decision:
1. **Memory-bandwidth-bound stages hurt disproportionately.** VAE decode moves a
lot of data and becomes the dominant cost on short (few-step) generations.
2. **Compute-bound stages scale with step count.** Full-step diffusion (50+
steps) is denoise-bound and simply takes a long time here.
## Use distilled few-step models
The single biggest lever on the GB10 is **model choice**. A 3-step distilled
model is ~18× faster than the full-step version of the same architecture:
| Model | Steps | Time / video | Bottleneck |
|---|---|---|---|
| FastWan2.1-T2V-1.3B (distilled) | 3 | **~40 s** | VAE decode |
| Wan2.1-T2V-1.3B (full-step) | 50 | ~12 min | denoise |
| Cosmos-Predict2.5-2B (full-step) | 51 | ~47 min | denoise |
| LTX2.3-distilled (+audio) | 8 | ~6 min | mixed |
The bottleneck flips from decode to denoise at around **4 steps**. Below that,
you're paying mostly for VAE decode; above it, mostly for the denoising loop.
!!! note "Few-step timings are noisy — measure in-process"
On a 3-step run, one-time per-process startup (Triton autotune, allocator
warmup) dominates and never amortizes, so single-run totals wobble ~±30%.
Compare levers **back-to-back in one process or as medians**, never as two
separate single runs. The [reproduction script](#reproduce-these-numbers)
does this for you.
## bf16 VAE decode — the real lever (already on for Wan)
Because few-step generation is decode-bound, VAE decode precision is where the
time is. Decoding in **bf16 instead of fp32 is essentially lossless** (MS-SSIM
~0.9999 vs fp32 on the identical latent) and ~1.14× faster — worth roughly
5–7% end-to-end on a decode-bound few-step model.
**FastVideo already defaults Wan's decode to bf16** (`vae_decode_precision="bf16"`,
with encode kept at fp32), so for the recommended Wan/FastWan models there's
nothing to set. If you run a model that still defaults to an fp32 decode, set the
decode-only override yourself:
```python
from fastvideo.configs.pipelines.base import PipelineConfig
pipeline_config = PipelineConfig.from_pretrained(model_id)
pipeline_config.vae_decode_precision = "bf16" # decode-only; leaves encode precision alone
```
Decode is output-only, so lowering its precision is safe. (Encode seeds the
denoising trajectory for I2V/causal models, so that stays at the pipeline's
default — don't lower `vae_precision` blindly for those.)
## Memory: one unified 128 GB pool
The GB10 has **no separate VRAM** — CPU and GPU share one 128 GB LPDDR5X pool
(~118 GB usable). Two practical consequences:
- **`nvidia-smi` reports memory as `[N/A]`** on the GB10, and the system "used"
figure conflates CPU + GPU + cache, so it's only a soft upper bound — treat the
whole 128 GB as one shared budget. For a per-run figure, use FastVideo's own
`peak_memory_mb` (reported on the generation result and by the performance
benchmark), which is measured inside the worker that runs the model.
- **The 128 GB is a *working-set* ceiling, not storage** — the model cache lives
on the NVMe (3.7 TB, ample). What has to fit in 128 GB is the weights,
activations, and KV cache — and, critically, the **VAE decode buffers**, which
is why tiling matters
(an untiled high-res decode can spike the pool into swap and lock the box).
The recommended few-step models are comfortable here: their weights are small
(1.3–2 B) and few-step generation keeps activations modest — a Wan2.1-1.3B
few-step generation peaks at **~8.4 GB** (measured), a small fraction of the pool.
The pressure comes from **decode resolution/frames**, not the model — a
1080p×121-frame untiled decode is what pushes the pool toward its ceiling, which
is why VAE tiling stays
on by default.
## What helps vs. what doesn't on the GB10
The honest summary — most "obvious" GPU optimizations don't move the needle on
this hardware, for reasons specific to it:
| Lever | Effect on the GB10 | Use it? |
|---|---|---|
| Distilled few-step model | ~18× vs full-step | ✅ **the primary lever** |
| bf16 VAE decode | ~1.14×, lossless; ~5–7% e2e on few-step | ✅ default for Wan |
| VSA (video sparse attention) | works out of the box (Triton kernel auto-selects on `sm_121`) | ✅ automatic |
| Building FlashAttention | **no speedup** — Torch SDPA already hits an efficient flash kernel on `sm_121`, and FA2 ties it | ❌ not worth building |
| `torch.compile` of the VAE decode | recompile storm (per-frame varying shapes) → ~1.1× | ❌ dead end |
| Linear (fp8 / nvfp4) quantization on long-sequence models (e.g. Cosmos) | ~nothing — see below | ❌ wrong lever here |
| FP4 attention (`ATTN_QAT_INFER`) | works on `sm_121` (runtime allowlist landed in #1647; kernel build is #1598); helps, but needs a QAT-trained checkpoint | ⚠️ opt-in — see below |
| FP4 linear on short-sequence models (LTX2) | up to −24% denoise at 1080p (#1594) | ⚠️ model/resolution-dependent |
### Why linear quantization is the wrong lever on long-sequence models
Quantizing the linear (GEMM) layers is a natural first instinct, but on a
long-sequence video model it buys almost nothing on the GB10. A video-DiT denoise
step is dominated by **O(N²) attention** at these sequence lengths (tens of
thousands of tokens); the linear layers are a small single-digit fraction of the
work. Quantizing them faster leaves the attention-bound total essentially
unchanged — measured at ~1% on Cosmos-2.5, i.e. noise, and full-step CFG models
also lose quality to per-step quantization error.
The same mechanism **does** help on **short-sequence** models: LTX2's aggressive
VAE compression gives it short attention sequences, so FP4 linear reaches −24%
there (#1594). The rule: **on the GB10, the lever that matters is attention
(sparse or FP4), not the linear layers** — unless the model has short sequences.
### FP4 on the GB10 (opt-in)
Block-scaled FP4 works on `sm_121` under CUDA 13:
- **FP4 attention** (`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER`, #1598) is
numerically correct on the GB10 and ~6% faster end-to-end generation, but it only preserves
quality on a **quantization-aware-distilled checkpoint** (e.g.
`FastVideo/FastWan-QAD-1.3B`) — stock weights aren't trained to tolerate it.
- **FP4 linear** helps only where sequences are short (LTX2, above).
The [`qad_fp4_ab.py`](#reproduce-these-numbers) harness reproduces the FP4
attention A/B on the QAD checkpoint.
## Running safely (don't lock the box)
The GB10 is easy to make **unreachable** — a heavy build or an untiled high-res
decode starves the ~20 ARM cores and unified memory, `sshd` can't get cycles, and
you're locked out at *"Connection timed out during banner exchange"* until the box
is power-cycled. To avoid it:
- **Inference:** keep **VAE tiling on** (the default), use sane resolution/frames,
and run under `nice -n 19`:
```bash
nice -n 19 nohup python your_script.py > run.log 2>&1 &
```
- **Builds** (flash-attn, kernel): `nice -n 19`, `MAX_JOBS=2`, `nohup`. Never a
bare foreground high-parallelism build.
- Leave `*_cpu_offload` at the example defaults — "CPU" offload is the *same*
unified RAM on the GB10, so the win is tiling + sane resolution, not offloading.
## Gotchas specific to the GB10
A few things that surprise people on this box (beyond the memory notes above):
- **Don't force `TORCH_SDPA` on a VSA checkpoint** (FastWan, LTX2.3-distilled).
The SDPA path builds a model without the gate weights the checkpoint carries and
fails to load. Run the model natively — VSA auto-routes to its Triton kernel on
`sm_121`.
- **Few-step timings are noisy run-to-run** (~±30%) — one-time startup dominates a
3-step run. Compare in-process / as medians, never two separate single runs (the
benchmark script does this).
- **`nvidia-smi` shows `[N/A]` for memory** — see [Memory](#memory-one-unified-128-gb-pool).
- **Cosmos-2.5** uses a Qwen2.5-VL text encoder; make sure you're on a FastVideo
build recent enough to include its `transformers`-compatibility handling before
running it.
## Reproduce these numbers
Two scripts under `examples/inference/optimizations/` reproduce the claims on
your own GB10:
```bash
# Headline: few-step generation timing (median) + the bf16-vs-fp32 decode A/B.
# FASTVIDEO_STAGE_LOGGING=1 also prints the denoise / decode / text split.
FASTVIDEO_STAGE_LOGGING=1 nice -n 19 \
python examples/inference/optimizations/spark_benchmark.py
# FP4 attention quality/speed A/B on the QAD checkpoint (one arm per run).
QAD_LINEAR=0 FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER nice -n 19 \
python examples/inference/optimizations/qad_fp4_ab.py
```
See also the [Optimizations](../../inference/optimizations.md) reference for the
full list of attention backends and quantization options.
+49 -9
View File
@@ -23,21 +23,61 @@ Also optionally install flash-attn:
uv pip install flash-attn --no-build-isolation -v
```
## Choose a maintained recipe
## Basic Usage
The cookbook selects complete, checked-in recipes instead of mixing model,
parallelism, offload, and attention settings independently.
### Text-to-Video Generation
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
```python
from fastvideo import VideoGenerator
!!! 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.
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()
```
## 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
+8 -17
View File
@@ -5,7 +5,7 @@
</div>
<div style="text-align: center;">
<strong>FastVideo is a unified post-training and real-time inference framework for accelerated video generation.</strong>
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.</strong>
</div>
<div style="text-align: center;">
@@ -25,23 +25,14 @@ FastVideo is an inference and post-training framework for diffusion models. It f
FastVideo has the following features:
- End-to-end post-training support for bidirectional and autoregressive models
- Full finetuning and LoRA [finetuning](training/finetune.md) for state-of-the-art open video DiTs
- [Data preprocessing pipeline](training/data_preprocess.md) for video, image, and text data
- [Distribution Matching Distillation (DMD2)](distillation/dmd.md) stepwise distillation
- Sparse attention with [Video Sparse Attention](attention/vsa/index.md)
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achieve >50x denoising speedup
- [Attn-QAT training](training/attn_qat.md) for quantization-aware post-training
- Causal distillation through Self-Forcing
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing
- See the [training overview](training/overview.md) for the full training workflow
- State-of-the-art performance optimizations for inference
- Sequence parallelism for distributed inference
- Multiple state-of-the-art [attention backends](attention/index.md)
- User-friendly [CLI](inference/cli.md) and Python API
- See the [support matrix](inference/support_matrix.md) for supported models and [optimizations](inference/optimizations.md) for the full list
- Realtime video generation and editing
- [Dreamverse](https://github.com/hao-ai-lab/FastVideo/tree/main/apps/dreamverse): stream and "vibe direct" video in realtime ([live demo](https://dreamverse.fastvideo.org/))
- [Sliding Tile Attention](attention/sta/index.md)
- [Sage Attention](https://arxiv.org/abs/2410.02367)
- E2E post-training support
- Data preprocessing pipeline for video data
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs.
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
## Documentation
+3 -4
View File
@@ -4,11 +4,10 @@ This page contains step-by-step instructions to get you quickly started with vid
## Requirements
- **OS**: Linux (tested on Ubuntu 22.04+), or macOS on Apple silicon via the
[MPS installation guide](../getting_started/installation/mps.md)
- **OS**: Linux (Tested on Ubuntu 22.04+)
- **Python**: 3.10-3.12
- **CUDA**: 12.6 or 13.0 (NVIDIA GPUs)
- **GPU**: At least one NVIDIA GPU, or an Apple silicon chip with MPS
- **CUDA**: 12.6 or 13.0
- **GPU**: At least one NVIDIA GPU
## Installation
-6
View File
@@ -3,12 +3,6 @@
This page describes the various options for speeding up generation times in FastVideo.
!!! note "On a DGX Spark (GB10)?"
Several options on this page behave differently on the GB10's unified-memory
hardware — some give little or nothing there. See
[DGX Spark: Performance & Tuning](../getting_started/installation/spark_performance.md)
for what actually helps on that platform and why.
## Table of Contents
- Optimized Attention Backends
-21
View File
@@ -64,10 +64,8 @@ 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) |
| sd35 | `stabilityai/stable-diffusion-3.5-medium` | T2I | [basic_sd35_t2i.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_sd35_t2i.py) |
| stable_audio | `FastVideo/stable-audio-open-1.0-Diffusers` | T2V | [basic_stable_audio.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_stable_audio.py) |
| stable_audio | `FastVideo/stable-audio-open-small-Diffusers` | T2V | [basic_stable_audio_small.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_stable_audio_small.py) |
@@ -95,14 +93,6 @@ 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.
**Note (Wan-VACE)**: not currently supported — no VACE pipeline or registered
model ID exists on `main`
([#1435](https://github.com/hao-ai-lab/FastVideo/issues/1435)). The closest
@@ -178,17 +168,6 @@ 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
+1 -1
View File
@@ -49,7 +49,7 @@ configuration reference.
Download the published preprocessed dataset:
```bash
bash examples/datasets/mixkit/download_dataset.sh
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
```
## Stage 1: supervised Attn-QAT fine-tuning
+1 -1
View File
@@ -182,7 +182,7 @@ Ready-to-run training scripts are available for multiple models:
Each example includes:
- a README pointing at the matching download script under `examples/datasets/`
- `download_dataset.sh` — download sample data
- `preprocess_*.sh` — run preprocessing
- `finetune_*.sh` — full finetune launcher
- `finetune_*_lora.sh` — LoRA finetune launcher
+1 -1
View File
@@ -48,7 +48,7 @@ For the complete two-stage Wan2.1 MixKit quantization-aware workflow, see
Each example includes:
- a README pointing at the matching download script under `examples/datasets/`
- `download_dataset.sh` — download sample data
- `preprocess_*.sh` — run preprocessing
- `finetune_*.sh` — launch training (full finetune or LoRA)
- `validation.json` — validation prompts for checkpoints

Before

Width:  |  Height:  |  Size: 1.2 MiB

After

Width:  |  Height:  |  Size: 1.2 MiB

@@ -1,3 +0,0 @@
#!/bin/bash
# Wan-Syn 720P dataset (77x768x1280, 250k clips).
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x768x1280_250k" --local_dir "data/Wan-Syn_77x768x1280_250k" --repo_type "dataset"
+1 -1
View File
@@ -5,7 +5,7 @@ These scripts demonstrate self-forcing distillation (SFwan) for the causal Wan2.
## Run the recipe
1. Download the preprocessed text-video dataset:
```bash
bash examples/datasets/crush-smol/download_dataset.sh
bash examples/distill/SFWan2.1-T2V/download_dataset.sh
```
2. (Optional) Regenerate parquet shards locally:
```bash
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -9,7 +9,7 @@ uv pip install vsa
### 1. Download dataset:
```bash
bash examples/datasets/wan-syn/download_dataset_480p.sh
bash examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/download_dataset.sh
```
### 2. Configure and run distillation:
@@ -1,3 +1,3 @@
#!/bin/bash
# Wan-Syn 480P dataset (77x448x832, 600k clips).
mkdir -p data
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "data/Wan-Syn_77x448x832_600k" --repo_type "dataset"
@@ -9,7 +9,7 @@ uv pip install vsa
### 1. Download dataset:
```bash
bash examples/datasets/crush-smol/download_dataset.sh
bash examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/download_dataset.sh
```
### 2. Configure and run distillation:
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
-8
View File
@@ -33,14 +33,6 @@ 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!
+13 -12
View File
@@ -3,8 +3,6 @@ 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.
@@ -14,11 +12,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,
)
@@ -26,19 +24,22 @@ 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,7 +30,8 @@ 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,
@@ -46,3 +47,4 @@ def main():
if __name__ == "__main__":
main()
@@ -31,7 +31,8 @@ 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,
@@ -45,3 +46,6 @@ def main():
if __name__ == "__main__":
main()
@@ -50,3 +50,4 @@ def main():
if __name__ == "__main__":
main()
+10 -9
View File
@@ -5,8 +5,6 @@ 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"
@@ -16,10 +14,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,
@@ -27,6 +25,7 @@ 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
@@ -40,16 +39,18 @@ 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")
+20 -14
View File
@@ -32,9 +32,11 @@ 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()
@@ -42,12 +44,14 @@ 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(
@@ -62,11 +66,13 @@ 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,6 +2,7 @@ import os
from fastvideo import VideoGenerator
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
@@ -45,8 +46,10 @@ 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
-180
View File
@@ -1,180 +0,0 @@
# 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()
+6 -2
View File
@@ -64,8 +64,12 @@ 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,6 +21,7 @@ from fastvideo.api import (
SamplingConfig,
)
DEFAULT_PROMPT = "a brushed steel espresso machine on a marble counter, morning window light"
+11 -5
View File
@@ -9,9 +9,11 @@ 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"
),
]
@@ -40,7 +42,9 @@ 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",
@@ -104,7 +108,9 @@ 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}")
+11 -10
View File
@@ -33,20 +33,21 @@ 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"
+32 -16
View File
@@ -26,35 +26,51 @@ 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()
+16 -25
View File
@@ -3,8 +3,6 @@ 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.
@@ -14,38 +12,31 @@ 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()
+14 -13
View File
@@ -3,37 +3,38 @@ 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,8 +6,6 @@ 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,7 +19,9 @@ 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,7 +2,6 @@ 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",
@@ -18,28 +17,21 @@ 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,6 +6,7 @@ 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,19 +3,17 @@ 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
+34 -28
View File
@@ -21,16 +21,20 @@ 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"
@@ -47,20 +51,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,
@@ -75,7 +79,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -90,11 +94,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,
@@ -107,9 +111,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,
@@ -124,14 +128,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:
@@ -139,7 +143,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
@@ -159,9 +163,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,
@@ -178,7 +182,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -188,13 +192,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)
@@ -202,3 +206,5 @@ def main():
if __name__ == "__main__":
main()
+36 -30
View File
@@ -15,18 +15,22 @@ 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
@@ -40,20 +44,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,
@@ -67,7 +71,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -82,11 +86,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,
@@ -99,9 +103,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,
@@ -115,14 +119,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:
@@ -130,7 +134,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",
@@ -147,9 +151,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,
@@ -166,7 +170,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -176,13 +180,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)
@@ -190,3 +194,5 @@ def main():
if __name__ == "__main__":
main()
+45 -35
View File
@@ -21,17 +21,21 @@ 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"
@@ -51,25 +55,27 @@ 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,
@@ -85,7 +91,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -100,16 +106,18 @@ 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,
@@ -122,9 +130,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,
@@ -140,14 +148,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:
@@ -155,7 +163,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(
@@ -173,9 +181,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,
@@ -192,7 +200,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -202,13 +210,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)
@@ -216,3 +224,5 @@ def main():
if __name__ == "__main__":
main()
+15 -12
View File
@@ -1,16 +1,19 @@
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:
@@ -33,4 +36,4 @@ def main() -> None:
if __name__ == "__main__":
main()
main()
@@ -67,19 +67,25 @@ _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")
@@ -108,7 +114,9 @@ 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:
@@ -117,18 +125,21 @@ 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}")
@@ -190,13 +201,11 @@ 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.
@@ -242,7 +251,10 @@ 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):
@@ -254,8 +266,10 @@ 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,14 +100,23 @@ _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)
@@ -138,7 +147,9 @@ 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:
@@ -146,16 +157,20 @@ 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}")
@@ -205,9 +220,10 @@ 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": {
@@ -262,7 +278,11 @@ 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")
@@ -271,30 +291,46 @@ 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,11 +12,15 @@ 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"
@@ -65,7 +69,9 @@ 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, [], []
@@ -83,8 +89,12 @@ def extract_sr_forward_latency(result: dict, ) -> tuple[float | None, list[tuple
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))
@@ -151,9 +161,11 @@ 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:
@@ -302,15 +314,19 @@ 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)
@@ -322,8 +338,10 @@ 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:
+12 -22
View File
@@ -13,34 +13,24 @@ 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.
@@ -52,8 +42,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,40 +14,27 @@ 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.
@@ -59,8 +46,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
@@ -69,8 +56,11 @@ 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"]
@@ -91,11 +81,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
@@ -1,98 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Generate synchronized video/audio from a first frame with MiniMax H3."""
from __future__ import annotations
import argparse
from pathlib import Path
from PIL import Image
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
InputConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
SamplingConfig,
)
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")
parser.add_argument("--prompt", required=True)
parser.add_argument("--num-frames", type=int, default=192, help="192 frames is exactly 8 seconds at 24 fps.")
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)
return parser.parse_args()
def main() -> None:
args = parse_args()
first_image = Image.open(args.image).convert("RGB")
last_image = Image.open(args.last_image).convert("RGB") if args.last_image else None
height, width = resolve_canvas_size(*first_image.size)
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
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,
),
),
))
try:
result = generator.generate(
GenerationRequest(
prompt=args.prompt,
negative_prompt="",
inputs=InputConfig(pil_image=first_image, last_image=last_image),
sampling=SamplingConfig(
height=height,
width=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_fl2va.mp4"),
save_video=True,
return_frames=False,
),
))
print(f"Output written to: {result.video_path}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,104 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Generate synchronized video/audio from ordered references with MiniMax H3."""
from __future__ import annotations
import argparse
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
ComponentConfig,
EngineConfig,
GenerationRequest,
GeneratorConfig,
InputConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
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")
parser.add_argument("--prompt", required=True)
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)
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)
return parser.parse_args()
def main() -> None:
args = parse_args()
references = [MiniMaxH3Reference(source=args.reference_video, media_type="video")]
if args.reference_audio:
references.append(MiniMaxH3Reference(source=args.reference_audio, media_type="audio"))
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
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,
),
),
pipeline=PipelineSelection(
workload_type="i2v",
components=ComponentConfig(override_pipeline_cls_name="MiniMaxH3Ref2VAModularPipeline"),
),
))
try:
result = generator.generate(
GenerationRequest(
prompt=args.prompt,
negative_prompt="",
inputs=InputConfig(references=references),
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_ref2va.mp4"),
save_video=True,
return_frames=False,
),
))
print(f"Output written to: {result.video_path}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,112 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Generate video and audio from text with MiniMax H3."""
from __future__ import annotations
import argparse
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
CompileConfig,
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
SamplingConfig,
)
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)
parser.add_argument("--width", type=int, default=1344)
parser.add_argument("--num-frames", type=int, default=124)
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,
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)
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
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,
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:
# 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()
-44
View File
@@ -1,44 +0,0 @@
# 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()
+15 -17
View File
@@ -1,20 +1,19 @@
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
@@ -24,19 +23,18 @@ 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()

Some files were not shown because too many files have changed in this diff Show More