Compare commits

..
Author SHA1 Message Date
SolitaryThinker 8f670ce294 [bugfix]: install the FA4 CUDA 13 runtime
Select flash-attn-4's cu13 extra on Blackwell install surfaces so revision 82d6441 receives its exact CuTe DSL 4.6.0.dev0 runtime. Enable the opt-in overlay for arm64 CUDA 13 images, including GB200, while CUDA 12.6 continues to use FA3/FA2.
2026-07-21 04:33:39 -07:00
SolitaryThinker 4943d43fb9 [bugfix]: align FA4 dependency stacks
Keep every dense FA4 install surface on the CuTe DSL 4.6-compatible upstream revision. Restore the private FP4 overlay's exact CuTe DSL 4.5.2 and Quack 0.5.0 pair, and document that it belongs in a separate inference environment because both distributions own flash_attn.cute.
2026-07-21 03:54:13 -07:00
968 changed files with 25144 additions and 73915 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)
+2 -13
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"
@@ -415,7 +404,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 25m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Training Tests"
env:
- TEST_TYPE=training
@@ -426,7 +415,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 25m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Distillation DMD Tests"
env:
- TEST_TYPE=distillation_dmd
-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
+7 -14
View File
@@ -13,23 +13,16 @@ on:
required: false
default: false
type: boolean
# Auto-rebuild the CUDA images when a repository-controlled image input
# changes on main. This includes the trusted SM89 kernel artifact's source,
# metadata/key helper, ABI dependency metadata, and build orchestration.
# Dreamverse (apps/dreamverse/docker/Dockerfile) and the ROCm Dockerfile stay
# manual-dispatch only.
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
# push trigger is a sufficient change detector on its own -- no separate
# detect-changes/paths-filter job is needed now that there is a single
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
# rocm Dockerfile stay manual-dispatch only.
push:
branches: [main]
paths:
- '.dockerignore'
- '.github/workflows/_template-build-image.yml'
- '.github/workflows/infra-build-image.yml'
- '.gitmodules'
- 'docker/Dockerfile'
- 'docker/uv-excludes'
- 'fastvideo-kernel/**'
- 'fastvideo/tests/modal/kernel_build_cache.py'
- 'pyproject.toml'
permissions:
@@ -57,7 +50,7 @@ jobs:
# 2.8.3 comes from the architecture-specific prebuilt releases.
build-cuda-images:
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
# on an in-scope main push (inputs are null on push). The
# on a push that changed docker/Dockerfile (inputs are null on push). The
# repository guard keeps fork syncs from auto-building; manual dispatch
# still works in forks.
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
-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 -3
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
@@ -76,7 +75,6 @@ docs/distillation/examples/
*.pkl
# Reference videos (negations must come after the catch-all on line below)
!fastvideo/tests/nightly/reference_video_*.mp4
# Static images
!docs/assets/images/**/*.png
-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/).
+5 -1
View File
@@ -47,7 +47,11 @@ COPY . /opt/FastVideo
RUN --mount=type=cache,target=/opt/uv/cache \
source /opt/venv/bin/activate \
&& uv pip install "/opt/FastVideo[dreamverse]"
&& if [[ "${UV_TORCH_BACKEND}" == "cu130" ]]; then \
uv pip install "/opt/FastVideo[dreamverse]" "nvidia-cutlass-dsl[cu13]"; \
else \
uv pip install "/opt/FastVideo[dreamverse]"; \
fi
# Standard docker build does not expose GPUs, while fastvideo-kernel/build.sh
# detects the CUDA architecture with torch at build time. The FastVideo package
@@ -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
+1 -1
View File
@@ -14,7 +14,7 @@ dependencies = [
[project.optional-dependencies]
server = [
"cerebras-cloud-sdk",
"flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@940cd9680f3315f2f06b43ab5bea2c2cf2d96806#subdirectory=flash_attn/cute",
"flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@82d6441eec5d4dfec120153db2c0145ae855a083#subdirectory=flash_attn/cute",
"flashinfer-python",
"openai>=1.40",
]
+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"])
+4 -3
View File
@@ -13,9 +13,10 @@ test.describe('create inference job', () => {
test('creates a T2V job and shows it in the queue', async ({ page }) => {
await page.goto('/inference');
// The trigger opens a real menu on click, so this path works for touch,
// mouse, and keyboard users.
await page.getByRole('button', { name: /create job/i }).click();
// The "Create Job" button reveals a workload menu on hover; wait for the
// T2V item to become visible before clicking so the CSS hover transition
// can't race the click.
await page.getByRole('button', { name: /create job/i }).hover();
const t2vItem = page.getByRole('menuitem', { name: /T2V/i });
await expect(t2vItem).toBeVisible();
await t2vItem.click();
+7 -10
View File
@@ -4,7 +4,7 @@ import { API_BASE, skipWithoutMock } from './helpers';
/**
* Gallery page: the seeded completed inference job surfaces as a media tile
* with playback controls or an explicit media-error fallback.
* (an <article> wrapping a <video>) captioned with its prompt.
*/
test.describe('gallery', () => {
skipWithoutMock();
@@ -30,15 +30,12 @@ test.describe('gallery', () => {
page.getByRole('heading', { level: 1, name: 'Gallery' }),
).toBeVisible();
const tile = page.locator('article').filter({ hasText: completed!.prompt });
await expect(tile).toBeVisible();
await expect(
tile.locator('video').or(tile.getByText('Preview unavailable')),
).toBeVisible();
// The completed job renders as an <article> containing a <video> tile.
const tile = page
.locator('article')
.filter({ has: page.locator('video') });
await expect(tile.first()).toBeVisible();
const video = tile.locator('video');
if (await video.isVisible()) {
await expect(video).toHaveAttribute('controls', '');
}
await expect(page.getByText(completed!.prompt)).toBeVisible();
});
});
-68
View File
@@ -42,74 +42,6 @@ test.describe('app shell', () => {
await expect(
page.getByRole('heading', { level: 1, name: section.title }),
).toBeVisible();
await expect(page.getByRole('main')).toHaveCount(1);
}
});
test('keeps navigation and content usable at responsive breakpoints', async ({
page,
}) => {
for (const width of [320, 375, 414, 768]) {
await page.setViewportSize({ width, height: 800 });
await page.goto('/inference');
const main = page.getByRole('main');
await expect(main).toBeVisible();
await expect(
page.getByRole('button', { name: /Create Job/i }),
).toBeVisible();
const initialBox = await main.boundingBox();
expect(initialBox?.x).toBe(width < 768 ? 0 : 220);
expect(initialBox?.width).toBe(width < 768 ? width : width - 220);
const navigation = page.getByRole('navigation', {
name: 'Primary navigation',
});
if (width < 768) {
await expect(
page.getByRole('button', { name: 'Open navigation' }),
).toBeVisible();
await page.getByRole('button', { name: 'Open navigation' }).click();
}
await expect(navigation).toBeVisible();
await navigation.getByRole('link', { name: 'Datasets' }).click();
await expect(page).toHaveURL(/\/datasets$/);
expect(
await page.evaluate(
() => document.documentElement.scrollWidth <= window.innerWidth,
),
).toBe(true);
}
});
test('uses full-width detail drawers on mobile', async ({ page }) => {
await page.setViewportSize({ width: 320, height: 800 });
await page.goto('/inference');
await page
.locator('article button[aria-pressed="false"]')
.first()
.click();
const jobDrawer = page.getByRole('dialog', { name: 'Job details' });
await expect(jobDrawer).toBeVisible();
expect(await jobDrawer.boundingBox()).toMatchObject({ x: 0, width: 320 });
await jobDrawer.getByRole('button', { name: 'Close' }).click();
await page.goto('/datasets');
await page
.locator('article button[aria-pressed="false"]')
.first()
.click();
const datasetDrawer = page.getByRole('dialog', {
name: /dataset details$/,
});
await expect(datasetDrawer).toBeVisible();
expect(await datasetDrawer.boundingBox()).toMatchObject({
x: 0,
width: 320,
});
});
});
-681
View File
@@ -9,7 +9,6 @@
"version": "0.1.0",
"dependencies": {
"@radix-ui/react-dialog": "^1.1.0",
"@radix-ui/react-dropdown-menu": "^2.1.24",
"@radix-ui/react-label": "^2.1.8",
"@radix-ui/react-scroll-area": "^1.2.10",
"@radix-ui/react-select": "^2.2.6",
@@ -1937,183 +1936,6 @@
}
}
},
"node_modules/@radix-ui/react-dropdown-menu": {
"version": "2.1.24",
"resolved": "https://registry.npmjs.org/@radix-ui/react-dropdown-menu/-/react-dropdown-menu-2.1.24.tgz",
"integrity": "sha512-geq8l2rJkxvkXsT9RMgtUE3P8pITFpTsvYpbySi1IH4fZEABD/Gp85myayFgxk0ktljGMJnCbeFkyTusvSvv7g==",
"license": "MIT",
"dependencies": {
"@radix-ui/primitive": "1.1.7",
"@radix-ui/react-compose-refs": "1.1.5",
"@radix-ui/react-context": "1.2.2",
"@radix-ui/react-id": "1.1.4",
"@radix-ui/react-menu": "2.1.24",
"@radix-ui/react-primitive": "2.1.10",
"@radix-ui/react-use-controllable-state": "1.2.6"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/primitive": {
"version": "1.1.7",
"resolved": "https://registry.npmjs.org/@radix-ui/primitive/-/primitive-1.1.7.tgz",
"integrity": "sha512-rqWnm76nYT8HoNNqEjpgJ7Pw/DrBj5iBTrmEPo6HTX5+VJyBNOqTdv4g89G63HuR5g0AaENoAcH7Is5fF2kZ8Q==",
"license": "MIT"
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/react-compose-refs": {
"version": "1.1.5",
"resolved": "https://registry.npmjs.org/@radix-ui/react-compose-refs/-/react-compose-refs-1.1.5.tgz",
"integrity": "sha512-+48PbAAbq3didjJxa+OaWY2ZwgAKsNiRGyeHKszblZMQ+kcpd9pAaT11cMkGEie0vsOi3QdeTE6d5Fe3Gn61kA==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/react-context": {
"version": "1.2.2",
"resolved": "https://registry.npmjs.org/@radix-ui/react-context/-/react-context-1.2.2.tgz",
"integrity": "sha512-RHCUGwKHDr0hDGg4X7ma4JG4/+12qxw8rkh5QKdDldlCvtja6nUx1Ef/8HVrJze81lEsgLQlqjzjGNHantgnQA==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/react-id": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-id/-/react-id-1.1.4.tgz",
"integrity": "sha512-TMQp2llA+RYn7JcjnrMnz7wN4pcVttPZnRZo52PLQsoLVKzNlVwUeHmfePgTgRluXFvlD3GD5g5MOVVTJCO0qA==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/react-primitive": {
"version": "2.1.10",
"resolved": "https://registry.npmjs.org/@radix-ui/react-primitive/-/react-primitive-2.1.10.tgz",
"integrity": "sha512-MucOnzh6hR5mid6VpkbglRAMYMjKLqRnGBbjXkzjK52fuQDd1qbkx78a5P40mkcnVXJdEVxm26E9OPAiUq7nBg==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-slot": "1.3.3"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/react-slot": {
"version": "1.3.3",
"resolved": "https://registry.npmjs.org/@radix-ui/react-slot/-/react-slot-1.3.3.tgz",
"integrity": "sha512-qx7oqnYbxnK9kYI9m317qmFmEgo6ywqWvbTogdj7cL9p3/yx4M48p7Rnw5z3H890cL/ow/EeWJsuTykeZVXP5Q==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-compose-refs": "1.1.5"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/react-use-controllable-state": {
"version": "1.2.6",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-controllable-state/-/react-use-controllable-state-1.2.6.tgz",
"integrity": "sha512-uEQJGT97ZA/TgP/Hydw47lHu+/vQj6z/0jA+WeTbK1o9Rx45GImjpD0tc3W5ad3D6XTSR6e1yEO0FvGq6WQfVQ==",
"license": "MIT",
"dependencies": {
"@radix-ui/primitive": "1.1.7",
"@radix-ui/react-use-effect-event": "0.0.5",
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/react-use-effect-event": {
"version": "0.0.5",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-effect-event/-/react-use-effect-event-0.0.5.tgz",
"integrity": "sha512-7cshFL8HGS/7HEiHH+9kL9HBwp2sa9yX18Knwek6KYWmXwM7pegMgta2AXMQKI+rq3JnfSj9x8wYqFMTdG1Jgg==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-dropdown-menu/node_modules/@radix-ui/react-use-layout-effect": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-layout-effect/-/react-use-layout-effect-1.1.4.tgz",
"integrity": "sha512-K20DkRkUwDnxEYMBPcg3Y6voLkEy5p5QQmszZgLngKKiC7dzBR/aEuK3w1qlx2JWDUNH6FluahYdgR3BP+QbYw==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-focus-guards": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-focus-guards/-/react-focus-guards-1.1.4.tgz",
@@ -2195,494 +2017,6 @@
}
}
},
"node_modules/@radix-ui/react-menu": {
"version": "2.1.24",
"resolved": "https://registry.npmjs.org/@radix-ui/react-menu/-/react-menu-2.1.24.tgz",
"integrity": "sha512-uW7RVuU6Lp/ZtfeY4b3kL32zccgEWvPv1+cf17ubYzHa9cL8AHokmk36cG/XEiH/smbQvumnieXX9j/e9RqJWA==",
"license": "MIT",
"dependencies": {
"@radix-ui/primitive": "1.1.7",
"@radix-ui/react-collection": "1.1.15",
"@radix-ui/react-compose-refs": "1.1.5",
"@radix-ui/react-context": "1.2.2",
"@radix-ui/react-direction": "1.1.4",
"@radix-ui/react-dismissable-layer": "1.1.19",
"@radix-ui/react-focus-guards": "1.1.6",
"@radix-ui/react-focus-scope": "1.1.16",
"@radix-ui/react-id": "1.1.4",
"@radix-ui/react-popper": "1.3.7",
"@radix-ui/react-portal": "1.1.17",
"@radix-ui/react-presence": "1.1.10",
"@radix-ui/react-primitive": "2.1.10",
"@radix-ui/react-roving-focus": "1.1.19",
"@radix-ui/react-slot": "1.3.3",
"@radix-ui/react-use-callback-ref": "1.1.4",
"aria-hidden": "^1.2.4",
"react-remove-scroll": "^2.7.2"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/primitive": {
"version": "1.1.7",
"resolved": "https://registry.npmjs.org/@radix-ui/primitive/-/primitive-1.1.7.tgz",
"integrity": "sha512-rqWnm76nYT8HoNNqEjpgJ7Pw/DrBj5iBTrmEPo6HTX5+VJyBNOqTdv4g89G63HuR5g0AaENoAcH7Is5fF2kZ8Q==",
"license": "MIT"
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-arrow": {
"version": "1.1.15",
"resolved": "https://registry.npmjs.org/@radix-ui/react-arrow/-/react-arrow-1.1.15.tgz",
"integrity": "sha512-v4zggRcjadnI+ClKDuijlQEW4tw3NoaeHc/PwpKnLoLLKNUG4InLegkstooLcRIUWCs+8L22dGURCVuFfOKfnA==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-primitive": "2.1.10"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-collection": {
"version": "1.1.15",
"resolved": "https://registry.npmjs.org/@radix-ui/react-collection/-/react-collection-1.1.15.tgz",
"integrity": "sha512-9W+B9NPF0NaaPh/1NJd3+KqsnlLqU9H7T2rvww+fp+T/evVXdNAyYcnfRQZFOjkR1ajQp3yORlqnI8soawLvNA==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-compose-refs": "1.1.5",
"@radix-ui/react-context": "1.2.2",
"@radix-ui/react-primitive": "2.1.10",
"@radix-ui/react-slot": "1.3.3"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-compose-refs": {
"version": "1.1.5",
"resolved": "https://registry.npmjs.org/@radix-ui/react-compose-refs/-/react-compose-refs-1.1.5.tgz",
"integrity": "sha512-+48PbAAbq3didjJxa+OaWY2ZwgAKsNiRGyeHKszblZMQ+kcpd9pAaT11cMkGEie0vsOi3QdeTE6d5Fe3Gn61kA==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-context": {
"version": "1.2.2",
"resolved": "https://registry.npmjs.org/@radix-ui/react-context/-/react-context-1.2.2.tgz",
"integrity": "sha512-RHCUGwKHDr0hDGg4X7ma4JG4/+12qxw8rkh5QKdDldlCvtja6nUx1Ef/8HVrJze81lEsgLQlqjzjGNHantgnQA==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-direction": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-direction/-/react-direction-1.1.4.tgz",
"integrity": "sha512-5pzg4FGQNpExhnhT2zlrP1wZFaYCd1K0nYWoFAdcYoYK868IEigqMX3B3f8yIoRlAhAeDWciLI6ZdCKHF9P4Vg==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-dismissable-layer": {
"version": "1.1.19",
"resolved": "https://registry.npmjs.org/@radix-ui/react-dismissable-layer/-/react-dismissable-layer-1.1.19.tgz",
"integrity": "sha512-8g4pfOL9HoKKLWGiypT+dphVqjFfmcXO5GBnhsG6zI+lxAx/8feQpr+1LSN8Re3hiZ+XkLNS4O9ztK11/LzQ6w==",
"license": "MIT",
"dependencies": {
"@radix-ui/primitive": "1.1.7",
"@radix-ui/react-compose-refs": "1.1.5",
"@radix-ui/react-primitive": "2.1.10",
"@radix-ui/react-use-callback-ref": "1.1.4",
"@radix-ui/react-use-effect-event": "0.0.5"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-focus-guards": {
"version": "1.1.6",
"resolved": "https://registry.npmjs.org/@radix-ui/react-focus-guards/-/react-focus-guards-1.1.6.tgz",
"integrity": "sha512-RNOJjfZMTyBM6xYmV3IVGXkPjIhcBAuv48POevAXwrGJhkWZ9p1rFoIS1JFooPuT193AZmRsCPhpoVJxx6OPoQ==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-focus-scope": {
"version": "1.1.16",
"resolved": "https://registry.npmjs.org/@radix-ui/react-focus-scope/-/react-focus-scope-1.1.16.tgz",
"integrity": "sha512-wmRZ2WWLvmt6KHy2rNPOdPUjwq5xOHY02+m+udwJTn0aNIox/rkskAvJTyTLGhPK6KgrUjlJUJpgmx/+wFiFIQ==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-compose-refs": "1.1.5",
"@radix-ui/react-primitive": "2.1.10",
"@radix-ui/react-use-callback-ref": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-id": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-id/-/react-id-1.1.4.tgz",
"integrity": "sha512-TMQp2llA+RYn7JcjnrMnz7wN4pcVttPZnRZo52PLQsoLVKzNlVwUeHmfePgTgRluXFvlD3GD5g5MOVVTJCO0qA==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-popper": {
"version": "1.3.7",
"resolved": "https://registry.npmjs.org/@radix-ui/react-popper/-/react-popper-1.3.7.tgz",
"integrity": "sha512-UsJrrd7w4wuKKTdvd/DNERVlwSlUcyXzjhyDwBk+3aPOsCjOY6ZSbxuw8E6lZTjjfP8Cpd0J8VVkrYUWyGYXyg==",
"license": "MIT",
"dependencies": {
"@floating-ui/react-dom": "^2.0.0",
"@radix-ui/react-arrow": "1.1.15",
"@radix-ui/react-compose-refs": "1.1.5",
"@radix-ui/react-context": "1.2.2",
"@radix-ui/react-primitive": "2.1.10",
"@radix-ui/react-use-callback-ref": "1.1.4",
"@radix-ui/react-use-layout-effect": "1.1.4",
"@radix-ui/react-use-rect": "1.1.4",
"@radix-ui/react-use-size": "1.1.4",
"@radix-ui/rect": "1.1.3"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-portal": {
"version": "1.1.17",
"resolved": "https://registry.npmjs.org/@radix-ui/react-portal/-/react-portal-1.1.17.tgz",
"integrity": "sha512-vKQLcWypUnwZVvfV7UkGahH2g6ySe8M8R+zYBwPrv5byZ9QAW6cQVvNKo7GgmD+p8aYb6D9JBuvy8/WhOno2wQ==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-primitive": "2.1.10",
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-presence": {
"version": "1.1.10",
"resolved": "https://registry.npmjs.org/@radix-ui/react-presence/-/react-presence-1.1.10.tgz",
"integrity": "sha512-3wyzCQ6+ubRA+D4uv9m95JYLXxmOHp05qjrkjeA7uKHHtjpPggQzc6DAb0URl7j67oR0K2foO4ip27TiX037Bw==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-primitive": {
"version": "2.1.10",
"resolved": "https://registry.npmjs.org/@radix-ui/react-primitive/-/react-primitive-2.1.10.tgz",
"integrity": "sha512-MucOnzh6hR5mid6VpkbglRAMYMjKLqRnGBbjXkzjK52fuQDd1qbkx78a5P40mkcnVXJdEVxm26E9OPAiUq7nBg==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-slot": "1.3.3"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-roving-focus": {
"version": "1.1.19",
"resolved": "https://registry.npmjs.org/@radix-ui/react-roving-focus/-/react-roving-focus-1.1.19.tgz",
"integrity": "sha512-V9jI6hDjT7l3jsCQD9bLNvDLM3tH/gdbOTp7Tefp3hbbgCGQoK7tUvrWiRlcoBHIZ809ElXwNQwVo0B98LuTXQ==",
"license": "MIT",
"dependencies": {
"@radix-ui/primitive": "1.1.7",
"@radix-ui/react-collection": "1.1.15",
"@radix-ui/react-compose-refs": "1.1.5",
"@radix-ui/react-context": "1.2.2",
"@radix-ui/react-direction": "1.1.4",
"@radix-ui/react-id": "1.1.4",
"@radix-ui/react-primitive": "2.1.10",
"@radix-ui/react-use-callback-ref": "1.1.4",
"@radix-ui/react-use-controllable-state": "1.2.6",
"@radix-ui/react-use-is-hydrated": "0.1.3",
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"@types/react-dom": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc",
"react-dom": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"@types/react-dom": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-slot": {
"version": "1.3.3",
"resolved": "https://registry.npmjs.org/@radix-ui/react-slot/-/react-slot-1.3.3.tgz",
"integrity": "sha512-qx7oqnYbxnK9kYI9m317qmFmEgo6ywqWvbTogdj7cL9p3/yx4M48p7Rnw5z3H890cL/ow/EeWJsuTykeZVXP5Q==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-compose-refs": "1.1.5"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-use-callback-ref": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-callback-ref/-/react-use-callback-ref-1.1.4.tgz",
"integrity": "sha512-R6OUY2e2fA6Yn6s+VSx5KBV6Nx8LQEhu+cz7LCej18rQ1HLyg9PSC9jP/ZNx0o6FAIK9c0F1kHylzSxKsdlkrQ==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-use-controllable-state": {
"version": "1.2.6",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-controllable-state/-/react-use-controllable-state-1.2.6.tgz",
"integrity": "sha512-uEQJGT97ZA/TgP/Hydw47lHu+/vQj6z/0jA+WeTbK1o9Rx45GImjpD0tc3W5ad3D6XTSR6e1yEO0FvGq6WQfVQ==",
"license": "MIT",
"dependencies": {
"@radix-ui/primitive": "1.1.7",
"@radix-ui/react-use-effect-event": "0.0.5",
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-use-effect-event": {
"version": "0.0.5",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-effect-event/-/react-use-effect-event-0.0.5.tgz",
"integrity": "sha512-7cshFL8HGS/7HEiHH+9kL9HBwp2sa9yX18Knwek6KYWmXwM7pegMgta2AXMQKI+rq3JnfSj9x8wYqFMTdG1Jgg==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-use-layout-effect": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-layout-effect/-/react-use-layout-effect-1.1.4.tgz",
"integrity": "sha512-K20DkRkUwDnxEYMBPcg3Y6voLkEy5p5QQmszZgLngKKiC7dzBR/aEuK3w1qlx2JWDUNH6FluahYdgR3BP+QbYw==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-use-rect": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-rect/-/react-use-rect-1.1.4.tgz",
"integrity": "sha512-cSOCh6JlkmfjLyNcLiu2nB4v+nm+dkZ+Q5KHWk/soo4U7ZLiEQFKHK9/YmtBHjfCEaU43IBKQOc4/uJmCaiCTQ==",
"license": "MIT",
"dependencies": {
"@radix-ui/rect": "1.1.3"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/react-use-size": {
"version": "1.1.4",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-size/-/react-use-size-1.1.4.tgz",
"integrity": "sha512-D3anSY15EJoxrihpsXI6SMrmmonnQtR2ni7arO+Lfdg3O95b9hNXxONk8jA5C8ANdF/h5HMAxejgs8PWJ6rlhw==",
"license": "MIT",
"dependencies": {
"@radix-ui/react-use-layout-effect": "1.1.4"
},
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-menu/node_modules/@radix-ui/rect": {
"version": "1.1.3",
"resolved": "https://registry.npmjs.org/@radix-ui/rect/-/rect-1.1.3.tgz",
"integrity": "sha512-JtyZR+mqgBibTo8xea3B6ZRmzZiM/YeVBtUkas6zMuXjAlfIFIW2FgqeM9eLyvEaYX66vr6DJMK+4U6LV0KhNw==",
"license": "MIT"
},
"node_modules/@radix-ui/react-popper": {
"version": "1.3.2",
"resolved": "https://registry.npmjs.org/@radix-ui/react-popper/-/react-popper-1.3.2.tgz",
@@ -3076,21 +2410,6 @@
}
}
},
"node_modules/@radix-ui/react-use-is-hydrated": {
"version": "0.1.3",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-is-hydrated/-/react-use-is-hydrated-0.1.3.tgz",
"integrity": "sha512-umO/aJ+82CpOnhDZUTbILCQf7kU/g0iv+oGs/Q8jw7IkhWBzaEP4sA268PhFAJTFetbwp3ICc6ktpI4TqtxcIw==",
"license": "MIT",
"peerDependencies": {
"@types/react": "*",
"react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
}
}
},
"node_modules/@radix-ui/react-use-layout-effect": {
"version": "1.1.2",
"resolved": "https://registry.npmjs.org/@radix-ui/react-use-layout-effect/-/react-use-layout-effect-1.1.2.tgz",
-1
View File
@@ -18,7 +18,6 @@
},
"dependencies": {
"@radix-ui/react-dialog": "^1.1.0",
"@radix-ui/react-dropdown-menu": "^2.1.24",
"@radix-ui/react-label": "^2.1.8",
"@radix-ui/react-scroll-area": "^1.2.10",
"@radix-ui/react-select": "^2.2.6",
@@ -1,64 +0,0 @@
import { act, fireEvent, render, screen } from '@testing-library/react';
import { describe, expect, it, vi } from 'vitest';
import { HeaderActionsProvider } from '@/components/shell/HeaderActionsContext';
import { getDatasets, type Dataset } from '@/lib/api';
import DatasetsPage from './page';
vi.mock('@/lib/api', () => ({
getDatasets: vi.fn(),
}));
vi.mock('@/components/datasets/AddDatasetButton', () => ({
default: () => null,
}));
vi.mock('@/components/datasets/CreateDatasetModal', () => ({
default: () => null,
}));
vi.mock('@/components/datasets/DatasetCard', () => ({
default: ({ dataset }: { dataset: Dataset }) => <div>{dataset.name}</div>,
}));
function renderPage() {
return render(
<HeaderActionsProvider>
<DatasetsPage />
</HeaderActionsProvider>,
);
}
describe('DatasetsPage', () => {
it('shows loading content before the initial request settles', async () => {
let resolveDatasets: (datasets: Dataset[]) => void = () => {};
vi.mocked(getDatasets).mockReturnValue(
new Promise<Dataset[]>((resolve) => {
resolveDatasets = resolve;
}),
);
renderPage();
expect(screen.getByLabelText('Loading datasets')).toBeInTheDocument();
expect(screen.queryByText('No datasets yet.')).not.toBeInTheDocument();
act(() => resolveDatasets([]));
expect(await screen.findByText('No datasets yet.')).toBeInTheDocument();
});
it('shows API failures separately from an empty list and retries', async () => {
vi.spyOn(console, 'error').mockImplementation(() => {});
vi.mocked(getDatasets).mockRejectedValueOnce(new Error('network down'));
renderPage();
expect(
await screen.findByText(/Could not load datasets from the Studio API/),
).toBeInTheDocument();
expect(screen.queryByText('No datasets yet.')).not.toBeInTheDocument();
vi.mocked(getDatasets).mockResolvedValueOnce([]);
fireEvent.click(screen.getByRole('button', { name: 'Try Again' }));
expect(await screen.findByText('No datasets yet.')).toBeInTheDocument();
});
});
+22 -73
View File
@@ -1,14 +1,12 @@
'use client';
import * as React from 'react';
import { AlertTriangle } from 'lucide-react';
import AddDatasetButton from '@/components/datasets/AddDatasetButton';
import CreateDatasetModal from '@/components/datasets/CreateDatasetModal';
import DatasetCard from '@/components/datasets/DatasetCard';
import { HeaderActions } from '@/components/shell/HeaderActionsContext';
import { Card } from '@/components/ui/card';
import { Button } from '@/components/ui/button';
import { useStore } from '@/hooks/useStore';
import { getDatasets } from '@/lib/api';
import type { Dataset } from '@/lib/api';
@@ -23,28 +21,18 @@ import {
export default function DatasetsPage() {
const [datasets, setDatasets] = React.useState<Dataset[]>([]);
const [isInitialLoading, setIsInitialLoading] = React.useState(true);
const [error, setError] = React.useState<string | null>(null);
const { open } = useStore(createDatasetModalStore);
const fetchSequence = React.useRef(0);
const fetchDatasets = React.useCallback(async () => {
const sequence = ++fetchSequence.current;
try {
const next = await getDatasets();
if (sequence === fetchSequence.current) {
setDatasets(next);
setError(null);
}
setDatasets(await getDatasets());
setError(null);
} catch (err) {
console.error('Failed to fetch datasets:', err);
if (sequence === fetchSequence.current) {
setError(
'Could not load datasets from the Studio API. Check the server and try again.',
);
}
} finally {
if (sequence === fetchSequence.current) setIsInitialLoading(false);
// Distinguish an API outage from a genuinely empty list, so the user
// isn't told they have no datasets when the server is unreachable.
setError(err instanceof Error ? err.message : 'Failed to load datasets');
}
}, []);
@@ -62,67 +50,28 @@ export default function DatasetsPage() {
<HeaderActions>
<AddDatasetButton />
</HeaderActions>
<div className="mx-auto flex w-full max-w-[850px] flex-col gap-6 px-4 pb-12">
<main className="mx-auto flex w-full max-w-[850px] flex-col gap-6 px-4 pb-12">
<Card className="p-6">
<div aria-busy={isInitialLoading}>
{isInitialLoading ? (
<div
aria-label="Loading datasets"
className="flex flex-col gap-3 py-2"
>
{[0, 1, 2].map((item) => (
<div
key={item}
className="h-24 animate-pulse rounded-lg border border-border bg-muted/50"
/>
))}
</div>
) : error && datasets.length === 0 ? (
<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={fetchDatasets}>
Try Again
</Button>
</div>
<div>
{error ? (
<p className="py-8 text-center text-destructive">{error}</p>
) : datasets.length === 0 ? (
<p className="py-8 text-center text-muted-foreground">
No datasets yet.
</p>
) : (
<>
{error && (
<p
role="status"
className="mb-3 rounded-lg border border-amber-500/50 bg-amber-500/10 px-3 py-2 text-sm text-foreground"
>
Dataset updates are temporarily unavailable. Showing the
most recent results.
</p>
)}
{datasets.length === 0 ? (
<p className="py-8 text-center text-muted-foreground">
No datasets yet.
</p>
) : (
datasets.map((ds) => (
<DatasetCard
key={ds.id}
dataset={ds}
onUpdated={fetchDatasets}
onSelect={() => handleSelectDataset(ds)}
/>
))
)}
</>
datasets.map((ds) => (
<DatasetCard
key={ds.id}
dataset={ds}
onUpdated={fetchDatasets}
onSelect={() => handleSelectDataset(ds)}
/>
))
)}
</div>
</Card>
</div>
</main>
<CreateDatasetModal
isOpen={open}
onClose={() => setCreateDatasetModalOpen(false)}
@@ -1,4 +1,4 @@
import { fireEvent, render, screen } from '@testing-library/react';
import { render, screen } from '@testing-library/react';
import { describe, expect, it, vi } from 'vitest';
import GalleryPage from './page';
@@ -43,22 +43,6 @@ describe('GalleryPage', () => {
expect(getJobsList).toHaveBeenCalledWith('inference');
});
it('provides video controls and a visible fallback when media fails', async () => {
vi.mocked(getJobsList).mockResolvedValue([makeJob()]);
renderGallery();
const video = await screen.findByLabelText(
'Generated video: a cat surfing a wave',
);
expect(video).toHaveAttribute('controls');
fireEvent.error(video);
expect(screen.getByText('Preview unavailable')).toBeInTheDocument();
expect(
screen.getByText('The generated file could not be loaded.'),
).toBeInTheDocument();
});
it('shows the empty state when no completed videos exist', async () => {
vi.mocked(getJobsList).mockResolvedValue([
makeJob({ status: 'running', output_path: null }),
+21 -70
View File
@@ -1,9 +1,8 @@
'use client';
import { AlertTriangle, ImageOff, Loader2 } from 'lucide-react';
import { 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';
@@ -12,61 +11,11 @@ function isImage(job: Job): boolean {
return job.output_path?.toLowerCase().endsWith('.png') ?? false;
}
function GalleryMedia({ job }: { job: Job }) {
const [failed, setFailed] = useState(false);
if (failed) {
return (
<div
role="status"
className="flex h-full flex-col items-center justify-center gap-2 px-4 text-center text-muted-foreground"
>
<ImageOff className="size-7" aria-hidden />
<span className="text-sm font-medium">Preview unavailable</span>
<span className="text-xs">
The generated file could not be loaded.
</span>
</div>
);
}
if (isImage(job)) {
return (
// eslint-disable-next-line @next/next/no-img-element
<img
src={getJobVideoUrl(job.id)}
alt={job.prompt}
className="block h-full w-full object-contain"
loading="lazy"
onError={() => setFailed(true)}
/>
);
}
return (
<video
src={getJobVideoUrl(job.id)}
aria-label={
job.prompt ? `Generated video: ${job.prompt}` : 'Generated video'
}
className="block h-full w-full object-contain"
controls
muted
loop
playsInline
preload="metadata"
onError={() => setFailed(true)}
/>
);
}
export default function GalleryPage() {
const [jobs, setJobs] = useState<Job[]>([]);
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 +40,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 +64,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
@@ -143,7 +77,24 @@ export default function GalleryPage() {
className="flex flex-col overflow-hidden rounded-lg border border-border bg-background"
>
<div className="relative aspect-video overflow-hidden bg-muted">
<GalleryMedia job={job} />
{isImage(job) ? (
// eslint-disable-next-line @next/next/no-img-element
<img
src={getJobVideoUrl(job.id)}
alt={job.prompt}
className="block h-full w-full object-contain"
loading="lazy"
/>
) : (
<video
src={getJobVideoUrl(job.id)}
className="block h-full w-full object-contain"
muted
loop
playsInline
preload="metadata"
/>
)}
</div>
<p
className="line-clamp-3 border-t border-border px-4 py-3 text-sm text-muted-foreground"
+2 -19
View File
@@ -41,7 +41,7 @@
--border: #e2e8f0;
--input: #cbd5e1;
--ring: #1d4ed8;
--ring: #94a3b8;
--radius: 0.5rem;
}
@@ -77,7 +77,7 @@
--border: #334155;
--input: #334155;
--ring: #7dd3fc;
--ring: #cbd5e1;
}
@theme inline {
@@ -125,7 +125,6 @@
html,
body {
min-height: 100%;
overflow-x: clip;
}
html {
@@ -164,22 +163,6 @@ a {
color: inherit;
}
:where(
a,
button,
input,
textarea,
select,
summary,
[role="button"],
[role="menuitem"],
[role="slider"],
[tabindex]
):focus-visible {
outline: 3px solid var(--ring) !important;
outline-offset: 2px !important;
}
summary {
list-style: none;
}
@@ -1,55 +0,0 @@
import { readFileSync } from 'node:fs';
import { join } from 'node:path';
import { describe, expect, it } from 'vitest';
const css = readFileSync(join(process.cwd(), 'src/app/globals.css'), 'utf8');
function token(block: string, name: string): string {
const match = block.match(new RegExp(`--${name}:\\s*(#[0-9a-fA-F]{6})`));
if (!match) throw new Error(`Missing --${name} token`);
return match[1];
}
function luminance(hex: string): number {
const channels = hex
.slice(1)
.match(/.{2}/g)!
.map((channel) => parseInt(channel, 16) / 255)
.map((channel) =>
channel <= 0.04045
? channel / 12.92
: ((channel + 0.055) / 1.055) ** 2.4,
);
return (
0.2126 * channels[0] + 0.7152 * channels[1] + 0.0722 * channels[2]
);
}
function contrast(first: string, second: string): number {
const firstLuminance = luminance(first);
const secondLuminance = luminance(second);
return (
(Math.max(firstLuminance, secondLuminance) + 0.05) /
(Math.min(firstLuminance, secondLuminance) + 0.05)
);
}
describe('global focus styles', () => {
it('keeps focus tokens above 3:1 against both page themes', () => {
const light = css.match(/:root\s*{([\s\S]*?)\n}/)?.[1] ?? '';
const dark = css.match(/\.dark\s*{([\s\S]*?)\n}/)?.[1] ?? '';
expect(contrast(token(light, 'ring'), token(light, 'background'))).toBeGreaterThanOrEqual(
3,
);
expect(contrast(token(dark, 'ring'), token(dark, 'background'))).toBeGreaterThanOrEqual(
3,
);
});
it('applies a non-animated three-pixel outline to focus-visible controls', () => {
expect(css).toContain('):focus-visible {');
expect(css).toContain('outline: 3px solid var(--ring) !important;');
expect(css).toContain('outline-offset: 2px !important;');
});
});
+2 -2
View File
@@ -4,8 +4,8 @@ import GpuGrid from '@/components/system/GpuGrid';
export default function GpusPage() {
return (
<div className="mx-auto flex w-full max-w-[1100px] flex-col gap-6 px-4 pb-12 pt-6">
<main className="mx-auto flex w-full max-w-[1100px] flex-col gap-6 px-4 pb-12 pt-6">
<GpuGrid />
</div>
</main>
);
}
@@ -52,20 +52,6 @@ describe('Settings page', () => {
expect(updateOption).toHaveBeenCalledWith('numFrames', expect.any(Number));
});
it('gives every slider an accessible name', () => {
renderPage();
const sliders = screen.getAllByRole('slider');
expect(sliders).toHaveLength(11);
for (const slider of sliders) {
expect(slider).toHaveAccessibleName();
}
expect(screen.getByRole('slider', { name: 'Frames' })).toBeInTheDocument();
expect(
screen.getByRole('slider', { name: 'Guidance Scale' }),
).toBeInTheDocument();
});
it('calls resetToDefaults when Reset to Defaults is clicked', () => {
renderPage();
fireEvent.click(screen.getByRole('button', { name: 'Reset to Defaults' }));
@@ -25,16 +25,6 @@ beforeEach(() => {
});
describe('DatasetCard', () => {
it('keeps selection and delete buttons as semantic siblings', () => {
render(<DatasetCard dataset={dataset} onUpdated={() => {}} />);
const selectButton = screen.getByRole('button', { pressed: false });
const deleteButton = screen.getByRole('button', { name: 'Delete' });
expect(selectButton).toHaveTextContent('My Dataset');
expect(selectButton).not.toContainElement(deleteButton);
});
it('renders the name, file count and human-readable size', () => {
render(<DatasetCard dataset={dataset} onUpdated={() => {}} />);
expect(screen.getByText('My Dataset')).toBeInTheDocument();
@@ -81,7 +71,7 @@ describe('DatasetCard', () => {
expect(onSelect).not.toHaveBeenCalled();
});
it('keeps the selection and delete actions separate', () => {
it('selects on keyboard activation of the card body but not of the Delete button', () => {
const onSelect = vi.fn();
render(
<DatasetCard dataset={dataset} onUpdated={() => {}} onSelect={onSelect} />,
@@ -93,12 +83,8 @@ describe('DatasetCard', () => {
});
expect(onSelect).not.toHaveBeenCalled();
// Activating the dedicated selection button selects the dataset.
fireEvent.click(
screen.getByRole('button', {
name: /My Dataset.*3 files.*2.0 KB/,
}),
);
// Activating the card body itself does select.
fireEvent.keyDown(screen.getByText('My Dataset'), { key: 'Enter' });
expect(onSelect).toHaveBeenCalledTimes(1);
});
@@ -55,33 +55,43 @@ export default function DatasetCard({
}
}
function handleKeyDown(e: React.KeyboardEvent) {
if ((e.target as HTMLElement).closest('button')) return;
if (e.key === 'Enter' || e.key === ' ') {
e.preventDefault();
onSelect();
}
}
return (
<article
<div
className={cn(
'mb-3 flex items-start gap-3 rounded-lg border border-border bg-background px-[1.15rem] py-4',
'mb-3 flex cursor-pointer flex-col gap-[0.6rem] rounded-lg border border-border bg-background px-[1.15rem] py-4',
isSelected && 'border-accent-blue bg-accent-blue/5',
)}
onClick={(e) => {
if ((e.target as HTMLElement).closest('button')) return;
onSelect();
}}
onKeyDown={handleKeyDown}
role="button"
tabIndex={0}
>
<button
type="button"
aria-pressed={isSelected}
onClick={onSelect}
className="flex min-w-0 flex-1 cursor-pointer flex-col gap-[0.6rem] rounded-md text-left"
>
<div className="flex flex-wrap items-center justify-between gap-2">
<span className="text-[0.95rem] font-semibold">{dataset.name}</span>
<span className="text-sm text-muted-foreground">
{fileCount} {fileCount === 1 ? 'file' : 'files'} · {sizeLabel}
</span>
</button>
<Button
type="button"
variant="destructive"
size="sm"
onClick={handleDelete}
disabled={isLoading}
>
Delete
</Button>
</article>
<Button
type="button"
variant="destructive"
size="sm"
onClick={handleDelete}
disabled={isLoading}
>
Delete
</Button>
</div>
<div className="text-sm text-muted-foreground">
{fileCount} {fileCount === 1 ? 'file' : 'files'} · {sizeLabel}
</div>
</div>
);
}
@@ -6,9 +6,6 @@ import * as api from '@/lib/api';
import type { Dataset } from '@/lib/api';
vi.mock('@/lib/api');
vi.mock('sonner', () => ({
toast: { error: vi.fn() },
}));
const mockedApi = vi.mocked(api);
@@ -30,27 +27,6 @@ beforeEach(() => {
});
describe('DatasetSidebar', () => {
it('fills the mobile viewport without reserving main-content width', async () => {
const onWidthChange = vi.fn();
render(
<DatasetSidebar
dataset={dataset}
isMobile
onClose={() => {}}
onWidthChange={onWidthChange}
/>,
);
const drawer = screen.getByRole('dialog', {
name: 'My Dataset dataset details',
});
expect(drawer).toHaveStyle({ width: '100%', maxWidth: 'none' });
expect(drawer).toHaveAttribute('aria-modal', 'true');
expect(drawer).toHaveFocus();
expect(onWidthChange).toHaveBeenCalledWith(0);
});
it('lists dataset files after loading', async () => {
render(<DatasetSidebar dataset={dataset} onClose={() => {}} />);
@@ -67,16 +43,6 @@ describe('DatasetSidebar', () => {
expect(mockedApi.getDatasetMediaUrl).toHaveBeenCalledWith('ds-1', 'b.mp4');
});
it('shows a fallback when a dataset preview cannot load', async () => {
render(<DatasetSidebar dataset={dataset} onClose={() => {}} />);
const preview = await screen.findByLabelText('Preview of a.mp4');
fireEvent.error(preview);
expect(screen.getByText('Preview unavailable')).toBeInTheDocument();
expect(screen.queryByLabelText('Preview of a.mp4')).not.toBeInTheDocument();
});
it('debounces caption save by 500ms', async () => {
render(<DatasetSidebar dataset={dataset} onClose={() => {}} />);
const textarea = await screen.findByDisplayValue('cap a');
@@ -133,39 +99,6 @@ describe('DatasetSidebar', () => {
}
});
it('shows a failed save and lets the user retry it', async () => {
vi.spyOn(console, 'error').mockImplementation(() => {});
mockedApi.updateDatasetCaption
.mockRejectedValueOnce(new Error('network down'))
.mockResolvedValueOnce(undefined);
render(<DatasetSidebar dataset={dataset} onClose={() => {}} />);
const textarea = await screen.findByDisplayValue('cap a');
vi.useFakeTimers();
try {
fireEvent.change(textarea, { target: { value: 'needs retry' } });
await act(async () => {
await vi.advanceTimersByTimeAsync(500);
});
expect(screen.getByText(/Not saved/)).toBeInTheDocument();
fireEvent.click(screen.getByRole('button', { name: 'Retry' }));
await act(async () => {
await Promise.resolve();
});
expect(mockedApi.updateDatasetCaption).toHaveBeenCalledTimes(2);
expect(mockedApi.updateDatasetCaption).toHaveBeenLastCalledWith(
'ds-1',
'a.mp4',
'needs retry',
);
expect(screen.getByText('Saved')).toBeInTheDocument();
} finally {
vi.useRealTimers();
}
});
it('debounces per file: editing another caption does not cancel a pending save', async () => {
render(<DatasetSidebar dataset={dataset} onClose={() => {}} />);
await screen.findByDisplayValue('cap a');
@@ -1,8 +1,7 @@
'use client';
import * as React from 'react';
import { ImageOff, X } from 'lucide-react';
import { toast } from 'sonner';
import { X } from 'lucide-react';
import DownloadCaptions from '@/components/datasets/DownloadCaptions';
import { Textarea } from '@/components/ui/textarea';
@@ -14,14 +13,12 @@ import {
type Dataset,
} from '@/lib/api';
import { cn } from '@/lib/utils';
import { useDrawerFocus } from '@/hooks/useDrawerFocus';
const SIDEBAR_MIN_WIDTH = 320;
const SIDEBAR_MAX_WIDTH = 900;
const INITIAL_PAGE_SIZE = 24;
const PAGE_SIZE = 24;
const SCROLL_THRESHOLD = 200;
type CaptionSaveState = 'idle' | 'saving' | 'saved' | 'error';
// Memoized so a caption keystroke re-renders only the edited card, not every
// visible <video> in the grid (visibleCount grows unbounded with scrolling).
@@ -30,101 +27,54 @@ const DatasetFileCard = React.memo(function DatasetFileCard({
mediaUrl,
caption,
thumbLoaded,
saveState,
onCaptionChange,
onCaptionRetry,
onThumbLoaded,
}: {
fileName: string;
mediaUrl: string;
caption: string;
thumbLoaded: boolean;
saveState: CaptionSaveState;
onCaptionChange: (fileName: string, value: string) => void;
onCaptionRetry: (fileName: string, value: string) => void;
onThumbLoaded: (fileName: string) => void;
}) {
const [mediaFailed, setMediaFailed] = React.useState(false);
React.useEffect(() => {
setMediaFailed(false);
}, [mediaUrl]);
return (
<div className="relative flex flex-col overflow-hidden rounded-lg border border-border bg-background">
{!thumbLoaded && !mediaFailed && (
{!thumbLoaded && (
<div className="pointer-events-none absolute inset-0 flex items-center justify-center bg-background/70">
<div className="h-6 w-6 animate-spin rounded-full border-2 border-muted-foreground/40 border-t-accent-blue" />
</div>
)}
{mediaFailed ? (
<div
role="status"
className="flex aspect-video w-full flex-col items-center justify-center gap-1 bg-muted px-2 text-center text-muted-foreground"
>
<ImageOff className="size-5" aria-hidden />
<span className="text-xs">Preview unavailable</span>
</div>
) : (
// eslint-disable-next-line jsx-a11y/media-has-caption
<video
src={mediaUrl}
aria-label={`Preview of ${fileName}`}
className="aspect-video w-full bg-border object-cover"
muted
autoPlay
loop
playsInline
onLoadedData={() => onThumbLoaded(fileName)}
onError={() => {
setMediaFailed(true);
onThumbLoaded(fileName);
}}
/>
)}
{/* eslint-disable-next-line jsx-a11y/media-has-caption */}
<video
src={mediaUrl}
className="aspect-video w-full bg-border object-cover"
muted
autoPlay
loop
playsInline
onLoadedData={() => onThumbLoaded(fileName)}
onError={() => onThumbLoaded(fileName)}
/>
<Textarea
aria-label={`Caption for ${fileName}`}
value={caption}
onChange={(e) => onCaptionChange(fileName, e.target.value)}
placeholder="Caption"
rows={2}
className="min-h-[2.5rem] resize-y rounded-none border-0 bg-transparent p-1.5 text-xs shadow-none focus-visible:border-transparent focus-visible:ring-0"
/>
<div
aria-live="polite"
className="flex min-h-6 items-center px-1.5 pb-1 text-[0.7rem] text-muted-foreground"
>
{saveState === 'saving' && <span>Saving…</span>}
{saveState === 'saved' && <span>Saved</span>}
{saveState === 'error' && (
<span role="alert" className="text-destructive">
Not saved.{' '}
<button
type="button"
onClick={() => onCaptionRetry(fileName, caption)}
className="inline-flex min-h-11 items-center font-medium underline underline-offset-2"
>
Retry
</button>
</span>
)}
</div>
</div>
);
});
export default function DatasetSidebar({
dataset,
isMobile = false,
onClose,
onWidthChange,
}: {
dataset: Dataset;
isMobile?: boolean;
onClose: () => void;
onWidthChange?: (w: number) => void;
}) {
const drawerRef = useDrawerFocus<HTMLElement>(isMobile);
const [width, setWidth] = React.useState(400);
const [isDragging, setIsDragging] = React.useState(false);
const [fileNames, setFileNames] = React.useState<string[]>([]);
@@ -134,21 +84,17 @@ export default function DatasetSidebar({
const [thumbLoaded, setThumbLoaded] = React.useState<
Record<string, boolean>
>({});
const [captionSaveStates, setCaptionSaveStates] = React.useState<
Record<string, CaptionSaveState>
>({});
// Pending debounced caption saves, keyed per file so editing one caption
// can't cancel another file's pending save.
const pendingSaves = React.useRef(
new Map<string, { timer: ReturnType<typeof setTimeout>; save: () => void }>(),
);
const captionVersions = React.useRef(new Map<string, number>());
const scrollRef = React.useRef<HTMLDivElement>(null);
React.useEffect(() => {
onWidthChange?.(isMobile ? 0 : width);
}, [isMobile, width, onWidthChange]);
onWidthChange?.(width);
}, [width, onWidthChange]);
React.useEffect(() => {
let cancelled = false;
@@ -160,8 +106,6 @@ export default function DatasetSidebar({
setCaptions(data.captions);
setVisibleCount(INITIAL_PAGE_SIZE);
setThumbLoaded({});
setCaptionSaveStates({});
captionVersions.current.clear();
})
.catch((err) => console.error('Failed to load dataset files:', err))
.finally(() => {
@@ -194,51 +138,23 @@ export default function DatasetSidebar({
});
const datasetId = dataset.id;
const persistCaption = React.useCallback(
(fileName: string, value: string, version: number) => {
setCaptionSaveStates((prev) => ({ ...prev, [fileName]: 'saving' }));
void updateDatasetCaption(datasetId, fileName, value)
.then(() => {
if (captionVersions.current.get(fileName) !== version) return;
setCaptionSaveStates((prev) => ({ ...prev, [fileName]: 'saved' }));
})
.catch((error) => {
if (captionVersions.current.get(fileName) !== version) return;
console.error('Failed to save caption:', error);
setCaptionSaveStates((prev) => ({ ...prev, [fileName]: 'error' }));
toast.error('Caption was not saved', {
description: `${fileName}: check the Studio API, then retry.`,
});
});
},
[datasetId],
);
const handleCaptionChange = React.useCallback(
(fileName: string, value: string) => {
setCaptions((prev) => ({ ...prev, [fileName]: value }));
setCaptionSaveStates((prev) => ({ ...prev, [fileName]: 'idle' }));
const pending = pendingSaves.current.get(fileName);
if (pending) clearTimeout(pending.timer);
const version = (captionVersions.current.get(fileName) ?? 0) + 1;
captionVersions.current.set(fileName, version);
const save = () => persistCaption(fileName, value, version);
const save = () => {
updateDatasetCaption(datasetId, fileName, value).catch((err) =>
console.error('Failed to save caption:', err),
);
};
const timer = setTimeout(() => {
pendingSaves.current.delete(fileName);
save();
}, 500);
pendingSaves.current.set(fileName, { timer, save });
},
[persistCaption],
);
const handleCaptionRetry = React.useCallback(
(fileName: string, value: string) => {
const version = (captionVersions.current.get(fileName) ?? 0) + 1;
captionVersions.current.set(fileName, version);
persistCaption(fileName, value, version);
},
[persistCaption],
[datasetId],
);
function handleScroll() {
@@ -273,16 +189,8 @@ export default function DatasetSidebar({
return (
<aside
ref={drawerRef}
tabIndex={-1}
role="dialog"
aria-label={`${dataset.name} dataset details`}
aria-modal={isMobile || undefined}
className="fixed bottom-0 right-0 top-[var(--header-height)] z-50 flex max-h-[calc(100dvh-var(--header-height))] min-w-0 shrink-0 flex-col border-l border-border bg-card md:min-w-[320px]"
style={{
width: isMobile ? '100%' : width,
maxWidth: isMobile ? 'none' : SIDEBAR_MAX_WIDTH,
}}
className="fixed bottom-0 right-0 top-[var(--header-height)] z-50 flex max-h-[calc(100vh-var(--header-height))] min-w-[320px] shrink-0 flex-col border-l border-border bg-card"
style={{ width, maxWidth: SIDEBAR_MAX_WIDTH }}
>
<div className="flex shrink-0 items-center justify-between border-b border-border px-5 py-4">
<h2 className="m-0 min-w-0 truncate text-base font-semibold text-foreground">
@@ -295,7 +203,7 @@ export default function DatasetSidebar({
onClick={onClose}
title="Close"
aria-label="Close"
className="flex size-11 items-center justify-center rounded-lg text-muted-foreground transition-colors hover:bg-accent hover:text-foreground"
className="flex items-center justify-center rounded-lg p-1.5 text-muted-foreground transition-colors hover:bg-accent hover:text-foreground"
>
<X className="h-[18px] w-[18px]" />
</button>
@@ -323,9 +231,7 @@ export default function DatasetSidebar({
mediaUrl={getDatasetMediaUrl(dataset.id, fileName)}
caption={captions[fileName] ?? ''}
thumbLoaded={!!thumbLoaded[fileName]}
saveState={captionSaveStates[fileName] ?? 'idle'}
onCaptionChange={handleCaptionChange}
onCaptionRetry={handleCaptionRetry}
onThumbLoaded={markThumbLoaded}
/>
))}
@@ -334,14 +240,14 @@ export default function DatasetSidebar({
</div>
</div>
{!isMobile && <div
<div
role="presentation"
onMouseDown={onMouseDown}
className={cn(
'absolute bottom-0 left-0 top-0 z-[1] w-1.5 cursor-col-resize hover:bg-accent-blue/25',
isDragging && 'bg-accent-blue/25',
)}
/>}
/>
</aside>
);
}
@@ -6,7 +6,7 @@ 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 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';
export default function DownloadCaptions({
fileNames,
@@ -1,53 +0,0 @@
import { render, screen } from '@testing-library/react';
import userEvent from '@testing-library/user-event';
import { describe, expect, it, vi } from 'vitest';
import CreateJobButton from './CreateJobButton';
vi.mock('./CreateJobModal', () => ({
default: ({
isOpen,
workloadType,
}: {
isOpen: boolean;
workloadType: string;
}) =>
isOpen ? (
<div role="dialog" data-workload-type={workloadType}>
Create job form
</div>
) : null,
}));
describe('CreateJobButton', () => {
it('opens the workload menu on click and selects an item', async () => {
const user = userEvent.setup();
render(<CreateJobButton jobType="inference" />);
await user.click(screen.getByRole('button', { name: 'Create Job' }));
await user.click(screen.getByRole('menuitem', { name: /I2V/i }));
expect(screen.getByRole('dialog')).toHaveAttribute(
'data-workload-type',
'i2v',
);
});
it('opens and operates the workload menu from the keyboard', async () => {
const user = userEvent.setup();
render(<CreateJobButton jobType="inference" />);
const trigger = screen.getByRole('button', { name: 'Create Job' });
trigger.focus();
await user.keyboard('{Enter}');
const firstItem = await screen.findByRole('menuitem', { name: /T2V/i });
expect(firstItem).toHaveFocus();
await user.keyboard('{Enter}');
expect(screen.getByRole('dialog')).toHaveAttribute(
'data-workload-type',
't2v',
);
});
});
@@ -2,7 +2,6 @@
import * as React from 'react';
import { ChevronDown } from 'lucide-react';
import * as DropdownMenu from '@radix-ui/react-dropdown-menu';
import CreateJobModal from '@/components/jobs/CreateJobModal';
import { Button } from '@/components/ui/button';
@@ -34,35 +33,31 @@ export default function CreateJobButton({ jobType }: CreateJobButtonProps) {
return (
<>
<DropdownMenu.Root>
<DropdownMenu.Trigger asChild>
<Button type="button" className="gap-1.5">
Create Job
<ChevronDown className="size-3.5 opacity-85" aria-hidden />
</Button>
</DropdownMenu.Trigger>
<DropdownMenu.Portal>
<DropdownMenu.Content
align="end"
sideOffset={4}
collisionPadding={8}
className="z-[200] min-w-48 overflow-hidden rounded-lg border border-border bg-popover py-1 text-popover-foreground shadow-lg"
>
{options.map((opt) => (
<DropdownMenu.Item
key={opt.type}
onSelect={() => openModal(opt.type)}
className="flex min-h-11 cursor-pointer select-none flex-col justify-center px-4 py-2 text-left text-sm font-medium outline-none data-[highlighted]:bg-secondary"
>
{opt.label}
<span className="mt-0.5 block text-xs font-normal text-muted-foreground">
{opt.desc}
</span>
</DropdownMenu.Item>
))}
</DropdownMenu.Content>
</DropdownMenu.Portal>
</DropdownMenu.Root>
<div className="group relative inline-block">
<Button type="button" className="gap-1.5">
Create Job
<ChevronDown className="size-3.5 opacity-85" aria-hidden />
</Button>
<div
role="menu"
className="invisible absolute right-0 top-full z-[200] mt-1 min-w-full -translate-y-1 rounded-lg border border-border bg-popover py-1 opacity-0 shadow-lg transition-all duration-150 group-hover:visible group-hover:translate-y-0 group-hover:opacity-100"
>
{options.map((opt) => (
<button
key={opt.type}
type="button"
role="menuitem"
onClick={() => openModal(opt.type)}
className="block w-full whitespace-nowrap px-4 py-2 text-left text-sm font-medium text-popover-foreground transition-colors hover:bg-secondary"
>
{opt.label}
<span className="mt-0.5 block text-xs font-normal text-muted-foreground">
{opt.desc}
</span>
</button>
))}
</div>
</div>
<CreateJobModal
isOpen={modalOpen}
onClose={() => setModalOpen(false)}
@@ -50,57 +50,6 @@ function renderModal(
}
describe('CreateJobModal', () => {
it('shows a model loading error instead of an empty model list', async () => {
vi.spyOn(console, 'error').mockImplementation(() => {});
vi.mocked(getModels).mockRejectedValueOnce(new Error('network down'));
renderModal();
expect(
await screen.findByText(/Models could not be loaded/),
).toBeInTheDocument();
expect(screen.getByLabelText('Model')).toHaveAttribute(
'aria-invalid',
'true',
);
});
it('keeps the form open and reports job creation failures', async () => {
vi.spyOn(console, 'error').mockImplementation(() => {});
vi.mocked(createJob).mockRejectedValueOnce(new Error('API rejected job'));
const user = userEvent.setup();
const { onClose, onSuccess } = renderModal();
await screen.findByRole('option', { name: 'Wan T2V (wan/t2v-1.3b)' });
await user.type(screen.getByLabelText('Prompt'), 'a careful test prompt');
await user.click(screen.getByRole('button', { name: 'Create Job' }));
expect(
await screen.findByText(/API rejected job.*then try again/),
).toBeInTheDocument();
expect(onSuccess).not.toHaveBeenCalled();
expect(onClose).not.toHaveBeenCalled();
});
it('reports image upload failures next to the file input', async () => {
vi.spyOn(console, 'error').mockImplementation(() => {});
vi.mocked(uploadImage).mockRejectedValueOnce(new Error('Upload failed'));
const user = userEvent.setup();
renderModal({ workloadType: 'i2v' });
await screen.findByRole('option', { name: 'Wan T2V (wan/t2v-1.3b)' });
const input = screen.getByLabelText('Image');
await user.upload(
input,
new File(['image'], 'input.png', { type: 'image/png' }),
);
expect(
await screen.findByText(/Upload failed.*Choose the image again/),
).toBeInTheDocument();
expect(input).toHaveAttribute('aria-invalid', 'true');
});
it('renders the form fields for an inference job', async () => {
renderModal();
@@ -103,17 +103,6 @@ export default function CreateJobModal({
const [fakeScoreModelPath, setFakeScoreModelPath] = React.useState('');
const [isSubmitting, setIsSubmitting] = React.useState(false);
const [isLoadingModels, setIsLoadingModels] = React.useState(false);
const [isLoadingDatasets, setIsLoadingDatasets] = React.useState(false);
const [modelLoadError, setModelLoadError] = React.useState<string | null>(
null,
);
const [datasetLoadError, setDatasetLoadError] = React.useState<string | null>(
null,
);
const [imageUploadError, setImageUploadError] = React.useState<string | null>(
null,
);
const [submitError, setSubmitError] = React.useState<string | null>(null);
const imageInputRef = React.useRef<HTMLInputElement>(null);
// Seed field values from the persisted default options each time the modal
@@ -155,10 +144,6 @@ export default function CreateJobModal({
setImageFileName('');
setSelectedDatasetId('');
setSelectedValidationDatasetId('');
setModelLoadError(null);
setDatasetLoadError(null);
setImageUploadError(null);
setSubmitError(null);
if (workloadType === 'dmd_t2v') {
setDmdUseVsa(false);
setDmdVsaSparsity(0.8);
@@ -177,7 +162,6 @@ export default function CreateJobModal({
// can't overwrite the current workload's model list/selection.
let stale = false;
setIsLoadingModels(true);
setModelLoadError(null);
getModels(inferenceWorkload)
.then((list) => {
if (stale) return;
@@ -196,13 +180,7 @@ export default function CreateJobModal({
}
})
.catch((e) => {
if (stale) return;
console.error('Failed to load models:', e);
setModels([]);
setModelId('');
setModelLoadError(
'Models could not be loaded. Check the Studio API and reopen this form to try again.',
);
if (!stale) console.error('Failed to load models:', e);
})
.finally(() => {
if (!stale) setIsLoadingModels(false);
@@ -215,22 +193,11 @@ export default function CreateJobModal({
// Training jobs need a dataset; load the ready datasets when relevant.
React.useEffect(() => {
if (isOpen && !isInference) {
setIsLoadingDatasets(true);
setDatasetLoadError(null);
getDatasets()
.then(setReadyDatasets)
.catch((error) => {
console.error('Failed to load datasets:', error);
setReadyDatasets([]);
setDatasetLoadError(
'Datasets could not be loaded. Check the Studio API and reopen this form to try again.',
);
})
.finally(() => setIsLoadingDatasets(false));
.catch(() => setReadyDatasets([]));
} else {
setReadyDatasets([]);
setIsLoadingDatasets(false);
setDatasetLoadError(null);
}
}, [isOpen, isInference]);
@@ -239,24 +206,16 @@ export default function CreateJobModal({
if (!file) {
setImagePath('');
setImageFileName('');
setImageUploadError(null);
return;
}
setIsUploadingImage(true);
setImageFileName(file.name);
setImageUploadError(null);
try {
const { path } = await uploadImage(file);
setImagePath(path);
} catch (error) {
console.error('Failed to upload image:', error);
} catch {
setImagePath('');
setImageFileName('');
setImageUploadError(
error instanceof Error
? `${error.message}. Choose the image again to retry.`
: 'The image could not be uploaded. Choose it again to retry.',
);
} finally {
setIsUploadingImage(false);
}
@@ -265,7 +224,6 @@ export default function CreateJobModal({
function clearImage() {
setImagePath('');
setImageFileName('');
setImageUploadError(null);
if (imageInputRef.current) imageInputRef.current.value = '';
}
@@ -281,7 +239,6 @@ export default function CreateJobModal({
workloadType === 'lora_t2v' ? 'lora' : jobType
) as JobType;
setIsSubmitting(true);
setSubmitError(null);
try {
const payload: CreateJobRequest = {
model_id: modelId,
@@ -339,11 +296,6 @@ export default function CreateJobModal({
onClose();
} catch (err) {
console.error('Failed to create job:', err);
setSubmitError(
err instanceof Error
? `${err.message}. Check the form and Studio API, then try again.`
: 'The job could not be created. Check the form and Studio API, then try again.',
);
} finally {
setIsSubmitting(false);
}
@@ -391,11 +343,7 @@ export default function CreateJobModal({
value={modelId}
onChange={(e) => setModelId(e.target.value)}
required
aria-describedby={
modelLoadError ? 'modal-model-error' : undefined
}
aria-invalid={modelLoadError ? true : undefined}
disabled={isSubmitting || isLoadingModels || !!modelLoadError}
disabled={isSubmitting || isLoadingModels}
>
<option value="" disabled>
{isLoadingModels
@@ -410,15 +358,6 @@ export default function CreateJobModal({
</option>
))}
</NativeSelect>
{modelLoadError && (
<p
id="modal-model-error"
role="alert"
className="text-sm text-destructive"
>
{modelLoadError}
</p>
)}
</FieldRow>
{isInference && workloadType === 'i2v' && (
@@ -430,10 +369,6 @@ export default function CreateJobModal({
accept=".png,.jpg,.jpeg,.webp,.bmp"
onChange={handleImageChange}
disabled={isSubmitting || isUploadingImage}
aria-describedby={
imageUploadError ? 'modal-image-error' : undefined
}
aria-invalid={imageUploadError ? true : undefined}
required
className="h-auto py-2 file:mr-3 file:cursor-pointer file:rounded-md file:border-0 file:bg-secondary file:px-2 file:py-1 file:text-sm file:text-secondary-foreground"
/>
@@ -450,15 +385,6 @@ export default function CreateJobModal({
</button>
</span>
)}
{imageUploadError && (
<p
id="modal-image-error"
role="alert"
className="text-sm text-destructive"
>
{imageUploadError}
</p>
)}
</FieldRow>
)}
@@ -506,22 +432,12 @@ export default function CreateJobModal({
id="modal-dataset"
value={selectedDatasetId}
onChange={(e) => setSelectedDatasetId(e.target.value)}
aria-describedby={
datasetLoadError ? 'modal-dataset-error' : undefined
}
aria-invalid={datasetLoadError ? true : undefined}
disabled={
isSubmitting || isLoadingDatasets || !!datasetLoadError
}
disabled={isSubmitting}
>
<option value="" disabled>
{isLoadingDatasets
? 'Loading datasets…'
: datasetLoadError
? 'Datasets unavailable'
: readyDatasets.length === 0
? 'No datasets (add in Datasets tab)'
: 'Select a dataset…'}
{readyDatasets.length === 0
? 'No datasets (add in Datasets tab)'
: 'Select a dataset…'}
</option>
{readyDatasets.map((d) => (
<option key={d.id} value={d.id}>
@@ -529,15 +445,6 @@ export default function CreateJobModal({
</option>
))}
</NativeSelect>
{datasetLoadError && (
<p
id="modal-dataset-error"
role="alert"
className="text-sm text-destructive"
>
{datasetLoadError}
</p>
)}
</FieldRow>
<FieldRow
htmlFor="modal-validation-dataset"
@@ -550,9 +457,7 @@ export default function CreateJobModal({
onChange={(e) =>
setSelectedValidationDatasetId(e.target.value)
}
disabled={
isSubmitting || isLoadingDatasets || !!datasetLoadError
}
disabled={isSubmitting}
>
<option value="">None</option>
{readyDatasets.map((d) => (
@@ -906,22 +811,9 @@ export default function CreateJobModal({
</details>
)}
<div className="flex flex-col items-start gap-2">
{submitError && (
<p role="alert" className="text-sm text-destructive">
{submitError}
</p>
)}
<Button
type="submit"
disabled={
isSubmitting ||
isUploadingImage ||
!!modelLoadError ||
!!datasetLoadError
}
>
{isSubmitting ? 'Creating…' : 'Create Job'}
<div>
<Button type="submit" disabled={isSubmitting}>
{isSubmitting ? 'Creating...' : 'Create Job'}
</Button>
</div>
</form>
@@ -20,14 +20,6 @@ vi.mock('@/lib/api', () => ({
downloadJobVideo: vi.fn(),
}));
vi.mock('@/lib/utils', async (importOriginal) => {
const actual = await importOriginal<typeof import('@/lib/utils')>();
return {
...actual,
downloadBlob: vi.fn(),
};
});
const makeJob = (overrides: Partial<Job> = {}): Job =>
makeBaseJob({
model_id: 'Wan2.1-T2V',
@@ -54,16 +46,6 @@ beforeEach(() => {
});
describe('JobCard', () => {
it('keeps selection and job action buttons as semantic siblings', () => {
render(<JobCard job={makeJob()} />);
const selectButton = screen.getByRole('button', { pressed: false });
const deleteButton = screen.getByRole('button', { name: 'Delete' });
expect(selectButton).toHaveTextContent('Wan2.1-T2V');
expect(selectButton).not.toContainElement(deleteButton);
});
it('renders the model, prompt, status and inference meta', () => {
render(<JobCard job={makeJob()} />);
expect(screen.getByText('Wan2.1-T2V')).toBeInTheDocument();
@@ -119,10 +119,18 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
}
}
function handleSelectJob() {
function handleSelectJob(e: React.MouseEvent | React.KeyboardEvent) {
if ((e.target as HTMLElement).closest('button')) return;
setActiveJobId(isSelected ? null : job.id);
}
function handleKeyDown(e: React.KeyboardEvent) {
if (e.key === 'Enter' || e.key === ' ') {
e.preventDefault();
handleSelectJob(e);
}
}
async function handleDownloadVideo(e: React.MouseEvent) {
e.preventDefault();
e.stopPropagation();
@@ -140,7 +148,11 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
}
return (
<article
<div
role="button"
tabIndex={0}
onClick={handleSelectJob}
onKeyDown={handleKeyDown}
className={cn(
'mb-3 flex cursor-pointer flex-col gap-2.5 rounded-lg border bg-background p-4 transition-colors last:mb-0',
isSelected
@@ -148,42 +160,35 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
: 'border-border hover:border-muted-foreground/40',
)}
>
<button
type="button"
aria-pressed={isSelected}
onClick={handleSelectJob}
className="flex w-full flex-col gap-2.5 rounded-md text-left"
>
<span className="flex flex-wrap items-center justify-between gap-2">
<span className="text-[0.95rem] font-semibold text-foreground">
{job.model_id}
</span>
<Badge variant={BADGE_VARIANTS[job.status] ?? 'secondary'}>
{job.status}
</Badge>
<div className="flex flex-wrap items-center justify-between gap-2">
<span className="text-[0.95rem] font-semibold text-foreground">
{job.model_id}
</span>
<span className="max-w-full overflow-hidden text-ellipsis whitespace-nowrap text-sm text-muted-foreground">
{job.prompt}
</span>
<span className="flex flex-wrap items-center gap-4 text-xs text-muted-foreground">
{job.job_type === 'inference' ? (
<>
<span>{job.num_frames} frames</span>
<span>
{job.height}×{job.width}
</span>
</>
) : (
<span>{job.workload_type?.replace(/_/g, ' ') ?? job.job_type}</span>
)}
{elapsedTime && (
<span className="inline-flex items-center gap-1">
<Timer className="size-3.5" aria-hidden />
{elapsedTime}
<Badge variant={BADGE_VARIANTS[job.status] ?? 'secondary'}>
{job.status}
</Badge>
</div>
<p className="max-w-full overflow-hidden text-ellipsis whitespace-nowrap text-sm text-muted-foreground">
{job.prompt}
</p>
<div className="flex flex-wrap items-center gap-4 text-xs text-muted-foreground">
{job.job_type === 'inference' ? (
<>
<span>{job.num_frames} frames</span>
<span>
{job.height}×{job.width}
</span>
)}
</span>
</button>
</>
) : (
<span>{job.workload_type?.replace(/_/g, ' ') ?? job.job_type}</span>
)}
{elapsedTime && (
<span className="inline-flex items-center gap-1">
<Timer className="size-3.5" aria-hidden />
{elapsedTime}
</span>
)}
</div>
<div className="flex flex-wrap items-center gap-1.5">
{job.status === 'running' ? (
<Button
@@ -235,6 +240,6 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
Delete
</Button>
</div>
</article>
</div>
);
}
@@ -19,32 +19,6 @@ const makeJob = (overrides: Partial<Job> = {}): Job =>
});
describe('JobDetailsSidebar', () => {
it('fills the mobile viewport without reserving main-content width', async () => {
vi.mocked(getJobLogs).mockResolvedValue({
lines: [],
total: 0,
progress: 0,
progress_msg: '',
phase: '',
});
const onWidthChange = vi.fn();
render(
<JobDetailsSidebar
job={makeJob({ status: 'completed' })}
isMobile
onClose={vi.fn()}
onWidthChange={onWidthChange}
/>,
);
const drawer = screen.getByRole('dialog', { name: 'Job details' });
expect(drawer).toHaveStyle({ width: '100%', maxWidth: 'none' });
expect(drawer).toHaveAttribute('aria-modal', 'true');
expect(drawer).toHaveFocus();
expect(onWidthChange).toHaveBeenCalledWith(0);
});
it('renders log lines streamed from the job log poll', async () => {
vi.mocked(getJobLogs).mockResolvedValue({
lines: ['boot sequence started', 'loading model weights'],
@@ -4,7 +4,6 @@ import * as React from 'react';
import { X } from 'lucide-react';
import { Button } from '@/components/ui/button';
import { useDrawerFocus } from '@/hooks/useDrawerFocus';
import { useResizable } from '@/hooks/useResizable';
import { downloadJobLog, getJobLogs } from '@/lib/api';
import type { Job } from '@/lib/types';
@@ -16,16 +15,13 @@ const POLL_INTERVAL_MS = 2000;
export default function JobDetailsSidebar({
job,
isMobile = false,
onClose,
onWidthChange,
}: {
job: Job;
isMobile?: boolean;
onClose: () => void;
onWidthChange?: (w: number) => void;
}) {
const drawerRef = useDrawerFocus<HTMLElement>(isMobile);
const [width, setWidth] = React.useState(360);
const [isDragging, setIsDragging] = React.useState(false);
const [isLoading, setIsLoading] = React.useState(false);
@@ -57,8 +53,8 @@ export default function JobDetailsSidebar({
});
React.useEffect(() => {
onWidthChange?.(isMobile ? 0 : width);
}, [isMobile, width, onWidthChange]);
onWidthChange?.(width);
}, [width, onWidthChange]);
// Auto-scroll the console to the bottom whenever new lines land. Runs after
// commit so scrollHeight reflects the freshly-rendered output.
@@ -141,16 +137,8 @@ export default function JobDetailsSidebar({
return (
<aside
ref={drawerRef}
tabIndex={-1}
role="dialog"
aria-label="Job details"
aria-modal={isMobile || undefined}
className="fixed bottom-0 right-0 top-[var(--header-height)] z-50 flex max-h-[calc(100dvh-var(--header-height))] min-w-0 shrink-0 flex-col border-l border-border bg-card md:min-w-[280px]"
style={{
width: isMobile ? '100%' : width,
maxWidth: isMobile ? 'none' : SIDEBAR_MAX_WIDTH,
}}
className="fixed bottom-0 right-0 top-[var(--header-height)] z-50 flex max-h-[calc(100vh-var(--header-height))] min-w-[280px] shrink-0 flex-col border-l border-border bg-card"
style={{ width, maxWidth: SIDEBAR_MAX_WIDTH }}
>
<div className="flex items-center justify-between border-b border-border px-5 py-4">
<h2 className="m-0 text-base font-semibold text-foreground">
@@ -207,14 +195,14 @@ export default function JobDetailsSidebar({
</pre>
</div>
{!isMobile && <div
<div
role="presentation"
onMouseDown={onMouseDown}
className={cn(
'absolute bottom-0 left-0 top-0 z-[1] w-1.5 cursor-col-resize hover:bg-accent-blue/25',
isDragging && 'bg-accent-blue/25',
)}
/>}
/>
</aside>
);
}
@@ -1,4 +1,4 @@
import { act, fireEvent, render, screen, waitFor } from '@testing-library/react';
import { act, render, screen, waitFor } from '@testing-library/react';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import JobQueue from '@/components/jobs/JobQueue';
@@ -37,46 +37,6 @@ beforeEach(() => {
});
describe('JobQueue', () => {
it('shows a loading placeholder before the initial request settles', async () => {
let resolveJobs: (jobs: Job[]) => void = () => {};
vi.mocked(getJobsList).mockReturnValue(
new Promise<Job[]>((resolve) => {
resolveJobs = resolve;
}),
);
render(<JobQueue jobType="inference" />);
expect(screen.getByLabelText('Loading jobs')).toBeInTheDocument();
expect(
screen.queryByText('No inference jobs yet. Create one above.'),
).not.toBeInTheDocument();
act(() => resolveJobs([]));
expect(
await screen.findByText('No inference jobs yet. Create one above.'),
).toBeInTheDocument();
});
it('shows request failures separately from an empty queue and retries', async () => {
vi.spyOn(console, 'error').mockImplementation(() => {});
vi.mocked(getJobsList).mockRejectedValueOnce(new Error('network down'));
render(<JobQueue jobType="inference" />);
expect(
await screen.findByText(/Could not load jobs from the Studio API/),
).toBeInTheDocument();
expect(
screen.queryByText('No inference jobs yet. Create one above.'),
).not.toBeInTheDocument();
vi.mocked(getJobsList).mockResolvedValueOnce([]);
fireEvent.click(screen.getByRole('button', { name: 'Try Again' }));
expect(
await screen.findByText('No inference jobs yet. Create one above.'),
).toBeInTheDocument();
});
it('shows an empty placeholder and fetches for the single job type', async () => {
render(<JobQueue jobType="inference" />);
expect(
@@ -1,10 +1,8 @@
'use client';
import * as React from 'react';
import { AlertTriangle } from 'lucide-react';
import JobCard from '@/components/jobs/JobCard';
import { Button } from '@/components/ui/button';
import { useStore } from '@/hooks/useStore';
import { getJobsList } from '@/lib/api';
import type { Job, JobType } from '@/lib/types';
@@ -34,8 +32,6 @@ function jobsShallowEqual(a: Job | null, b: Job | null): boolean {
export default function JobQueue({ jobType, jobTypesForList }: JobQueueProps) {
const [jobs, setJobs] = React.useState<Job[]>([]);
const [isInitialLoading, setIsInitialLoading] = React.useState(true);
const [error, setError] = React.useState<string | null>(null);
const { nonce } = useStore(jobsRefreshStore);
const { activeJobId } = useStore(activeJobStore);
@@ -75,22 +71,11 @@ export default function JobQueue({ jobType, jobTypesForList }: JobQueueProps) {
new Date(a.created_at ?? 0).getTime(),
);
}
if (seq === fetchSeq.current) {
setJobs(next);
setError(null);
}
if (seq === fetchSeq.current) setJobs(next);
} catch (e) {
console.error('Failed to fetch jobs:', e);
if (seq === fetchSeq.current) {
setError(
'Could not load jobs from the Studio API. Check the server and try again.',
);
}
} finally {
if (seq === fetchSeq.current) {
inFlight.current = false;
setIsInitialLoading(false);
}
if (seq === fetchSeq.current) inFlight.current = false;
}
}, [typesKey]);
@@ -137,59 +122,20 @@ export default function JobQueue({ jobType, jobTypesForList }: JobQueueProps) {
const multiType = typesToFetch.length > 1;
return (
<div className="mx-auto flex w-full max-w-[850px] flex-col gap-6 px-4 pb-12">
<main className="mx-auto flex w-full max-w-[850px] flex-col gap-6 px-4 pb-12">
<section className="p-6">
<div aria-busy={isInitialLoading}>
{isInitialLoading ? (
<div
aria-label="Loading jobs"
className="flex flex-col gap-3 py-2"
>
{[0, 1, 2].map((item) => (
<div
key={item}
className="h-32 animate-pulse rounded-lg border border-border bg-muted/50"
/>
))}
</div>
) : error && jobs.length === 0 ? (
<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={fetchJobs}>
Try Again
</Button>
</div>
) : (
<>
{error && (
<p
role="status"
className="mb-3 rounded-lg border border-amber-500/50 bg-amber-500/10 px-3 py-2 text-sm text-foreground"
>
Job updates are temporarily unavailable. Showing the most
recent results.
</p>
)}
{jobs.length === 0 ? (
<div>
{jobs.length === 0 ? (
<p className="py-8 text-center text-muted-foreground">
No {multiType ? 'jobs' : `${jobType} jobs`} yet. Create one above.
</p>
) : (
jobs.map((job) => (
<JobCard key={job.id} job={job} onJobUpdated={fetchJobs} />
))
)}
</>
) : (
jobs.map((job) => (
<JobCard key={job.id} job={job} onJobUpdated={fetchJobs} />
))
)}
</div>
</section>
</div>
</main>
);
}
@@ -9,7 +9,6 @@ import { HeaderActionsProvider } from '@/components/shell/HeaderActionsContext';
import PrimarySidebar from '@/components/shell/PrimarySidebar';
import JobDetailsSidebar from '@/components/jobs/JobDetailsSidebar';
import { Toaster } from '@/components/ui/sonner';
import { useMediaQuery } from '@/hooks/useMediaQuery';
import { useStore } from '@/hooks/useStore';
import {
activeDatasetStore,
@@ -24,80 +23,46 @@ export function AppShell({ children }: { children: React.ReactNode }) {
const pathname = usePathname();
const { activeJob } = useStore(activeJobStore);
const { activeDataset } = useStore(activeDatasetStore);
const isMobile = useMediaQuery('(max-width: 767px)');
const [primaryWidth, setPrimaryWidth] = React.useState(220);
const [secondaryWidth, setSecondaryWidth] = React.useState(0);
const [primaryOpen, setPrimaryOpen] = React.useState(false);
const jobSidebarOpen = JOB_ROUTES.includes(pathname) && activeJob != null;
const datasetSidebarOpen =
pathname === '/datasets' && activeDataset != null;
const secondaryOpen = jobSidebarOpen || datasetSidebarOpen;
// Mobile detail drawers claim aria-modal, so everything behind them must
// actually be inert — the platform enforces what the ARIA claims.
const drawerModal = isMobile && secondaryOpen;
React.useEffect(() => {
initDefaultOptions();
}, []);
React.useEffect(() => {
setPrimaryOpen(false);
}, [pathname]);
React.useEffect(() => {
function handleKeyDown(e: KeyboardEvent) {
if (e.key !== 'Escape' || document.querySelector('[data-modal]')) return;
if (primaryOpen) {
setPrimaryOpen(false);
return;
if (e.key === 'Escape' && !document.querySelector('[data-modal]')) {
if (activeJobStore.get().activeJob) setActiveJobId(null);
if (activeDatasetStore.get().activeDataset) setActiveDatasetId(null);
}
if (activeJobStore.get().activeJob) setActiveJobId(null);
if (activeDatasetStore.get().activeDataset) setActiveDatasetId(null);
}
document.addEventListener('keydown', handleKeyDown);
return () => document.removeEventListener('keydown', handleKeyDown);
}, [primaryOpen]);
}, []);
return (
<HeaderActionsProvider>
<div
style={{ display: 'contents' }}
inert={drawerModal ? true : undefined}
>
<Header
navigationOpen={primaryOpen}
onNavigationToggle={() => setPrimaryOpen((open) => !open)}
/>
</div>
<Header />
<div
className="flex overflow-hidden"
style={{
marginTop: 'var(--header-height)',
height: 'calc(100dvh - var(--header-height))',
height: 'calc(100vh - var(--header-height))',
}}
>
<PrimarySidebar
isMobile={isMobile}
mobileOpen={primaryOpen}
onMobileClose={() => setPrimaryOpen(false)}
onWidthChange={setPrimaryWidth}
/>
{primaryOpen && (
<button
type="button"
aria-label="Close navigation"
onClick={() => setPrimaryOpen(false)}
className="fixed inset-x-0 bottom-0 top-[var(--header-height)] z-40 bg-black/55 md:hidden"
/>
)}
<PrimarySidebar onWidthChange={setPrimaryWidth} />
<main
className="flex min-w-0 flex-1 flex-col overflow-auto"
inert={drawerModal ? true : undefined}
style={{
marginLeft: isMobile ? 0 : primaryWidth,
marginRight: isMobile || !secondaryOpen ? 0 : secondaryWidth,
marginLeft: primaryWidth,
marginRight: secondaryOpen ? secondaryWidth : 0,
}}
>
{children}
@@ -105,7 +70,6 @@ export function AppShell({ children }: { children: React.ReactNode }) {
{jobSidebarOpen && activeJob && (
<JobDetailsSidebar
job={activeJob}
isMobile={isMobile}
onClose={() => setActiveJobId(null)}
onWidthChange={setSecondaryWidth}
/>
@@ -113,7 +77,6 @@ export function AppShell({ children }: { children: React.ReactNode }) {
{datasetSidebarOpen && activeDataset && (
<DatasetSidebar
dataset={activeDataset}
isMobile={isMobile}
onClose={() => setActiveDatasetId(null)}
onWidthChange={setSecondaryWidth}
/>
@@ -1,10 +1,8 @@
'use client';
import { Menu, X } from 'lucide-react';
import { usePathname } from 'next/navigation';
import { useHeaderActions } from '@/components/shell/HeaderActionsContext';
import { Button } from '@/components/ui/button';
import { ThemeToggle } from '@/components/ui/theme-toggle';
const TAB_TITLES: Record<string, string> = {
@@ -17,47 +15,25 @@ const TAB_TITLES: Record<string, string> = {
'/settings': 'Settings',
};
export default function Header({
navigationOpen,
onNavigationToggle,
}: {
navigationOpen: boolean;
onNavigationToggle: () => void;
}) {
export default function Header() {
const pathname = usePathname();
const { actions } = useHeaderActions();
const title = TAB_TITLES[pathname] ?? 'FastVideo';
return (
<header className="fixed inset-x-0 top-0 z-[100] flex h-[var(--header-height)] items-center gap-2 border-b border-border bg-background/80 px-2 backdrop-blur sm:px-4 md:gap-6 md:px-6">
<Button
type="button"
variant="outline"
size="icon"
aria-label={navigationOpen ? 'Close navigation' : 'Open navigation'}
aria-controls="primary-navigation"
aria-expanded={navigationOpen}
onClick={onNavigationToggle}
className="shrink-0 md:hidden"
>
{navigationOpen ? (
<X className="size-5" aria-hidden />
) : (
<Menu className="size-5" aria-hidden />
)}
</Button>
<header className="fixed inset-x-0 top-0 z-[100] flex h-[var(--header-height)] items-center gap-6 border-b border-border bg-background/80 px-6 backdrop-blur">
{/* eslint-disable-next-line @next/next/no-img-element */}
<img
src="/logo.svg"
alt="FastVideo Logo"
width={100}
height={42}
className="hidden h-[42px] w-[78px] shrink-0 object-contain min-[361px]:block md:w-[100px]"
className="block h-[42px] w-[100px]"
/>
<h1 className="sr-only m-0 flex-1 text-xl font-semibold tracking-tight md:not-sr-only">
<h1 className="m-0 flex-1 text-xl font-semibold tracking-tight">
{title}
</h1>
<div className="ml-auto flex min-w-0 items-center gap-2 md:gap-3">
<div className="flex items-center gap-3">
{actions}
<ThemeToggle />
</div>
@@ -1,7 +1,6 @@
'use client';
import * as React from 'react';
import { X } from 'lucide-react';
import Link from 'next/link';
import { usePathname } from 'next/navigation';
@@ -20,18 +19,12 @@ const JOB_ROUTES = [
] as const;
const TAB_BASE =
'block min-h-11 px-5 py-[0.65rem] text-left text-sm text-muted-foreground transition-colors hover:bg-accent/60 hover:text-foreground';
'block px-5 py-[0.65rem] text-left text-sm text-muted-foreground transition-colors hover:bg-accent/60 hover:text-foreground';
const TAB_ACTIVE = 'bg-accent-blue/10 font-medium text-accent-blue';
export default function PrimarySidebar({
isMobile,
mobileOpen,
onMobileClose,
onWidthChange,
}: {
isMobile: boolean;
mobileOpen: boolean;
onMobileClose: () => void;
onWidthChange?: (w: number) => void;
}) {
const pathname = usePathname();
@@ -45,8 +38,8 @@ export default function PrimarySidebar({
const isJobsActive = JOB_ROUTES.some((r) => pathname === r.href);
React.useEffect(() => {
onWidthChange?.(isMobile ? 0 : layoutWidth);
}, [isMobile, layoutWidth, onWidthChange]);
onWidthChange?.(layoutWidth);
}, [layoutWidth, onWidthChange]);
React.useEffect(() => {
if (JOB_ROUTES.some((r) => pathname === r.href)) {
@@ -65,37 +58,11 @@ export default function PrimarySidebar({
return (
<aside
id="primary-navigation"
aria-hidden={isMobile && !mobileOpen}
inert={isMobile && !mobileOpen ? true : undefined}
className={cn(
'fixed bottom-0 left-0 top-[var(--header-height)] z-50 flex max-h-[calc(100dvh-var(--header-height))] shrink-0 flex-col border-r border-border bg-card transition-transform duration-200 md:translate-x-0',
mobileOpen ? 'translate-x-0' : '-translate-x-full',
)}
style={{
width: isMobile
? 'min(18rem, calc(100vw - 3rem))'
: effectiveWidth,
}}
className="fixed bottom-0 left-0 top-[var(--header-height)] z-50 flex max-h-[calc(100vh-var(--header-height))] shrink-0 flex-col border-r border-border bg-card"
style={{ width: effectiveWidth }}
>
{isMobile && (
<div className="flex h-14 items-center justify-between border-b border-border px-4">
<span className="text-sm font-semibold">Navigation</span>
<button
type="button"
onClick={onMobileClose}
aria-label="Close navigation"
className="flex size-11 items-center justify-center rounded-lg text-muted-foreground hover:bg-accent hover:text-foreground"
>
<X className="size-5" aria-hidden />
</button>
</div>
)}
{!isCollapsed && (
<nav
aria-label="Primary navigation"
className="flex flex-col overflow-y-auto py-2"
>
<nav className="flex flex-col py-2">
<div className="flex flex-col">
<button
type="button"
@@ -128,8 +95,6 @@ export default function PrimarySidebar({
<Link
key={route.href}
href={route.href}
aria-current={pathname === route.href ? 'page' : undefined}
onClick={onMobileClose}
className={cn(
TAB_BASE,
'px-4 py-2 text-[0.85rem]',
@@ -144,32 +109,24 @@ export default function PrimarySidebar({
</div>
<Link
href="/datasets"
aria-current={pathname === '/datasets' ? 'page' : undefined}
onClick={onMobileClose}
className={cn(TAB_BASE, pathname === '/datasets' && TAB_ACTIVE)}
>
Datasets
</Link>
<Link
href="/gallery"
aria-current={pathname === '/gallery' ? 'page' : undefined}
onClick={onMobileClose}
className={cn(TAB_BASE, pathname === '/gallery' && TAB_ACTIVE)}
>
Gallery
</Link>
<Link
href="/gpus"
aria-current={pathname === '/gpus' ? 'page' : undefined}
onClick={onMobileClose}
className={cn(TAB_BASE, pathname === '/gpus' && TAB_ACTIVE)}
>
GPUs
</Link>
<Link
href="/settings"
aria-current={pathname === '/settings' ? 'page' : undefined}
onClick={onMobileClose}
className={cn(TAB_BASE, pathname === '/settings' && TAB_ACTIVE)}
>
Settings
@@ -177,7 +134,7 @@ export default function PrimarySidebar({
</nav>
)}
{!isMobile && <div
<div
className={cn(
'absolute bottom-0 p-2',
isCollapsed ? '-right-[60px] top-0' : 'right-0',
@@ -188,7 +145,8 @@ export default function PrimarySidebar({
onClick={() => setIsCollapsed((v) => !v)}
title={isCollapsed ? 'Expand sidebar' : 'Collapse sidebar'}
className={cn(
'flex size-11 items-center justify-center rounded-lg text-muted-foreground transition-colors hover:bg-accent hover:text-foreground',
'flex items-center justify-center rounded-lg text-muted-foreground transition-colors hover:bg-accent hover:text-foreground',
isCollapsed ? 'p-3' : 'p-2',
)}
>
<svg
@@ -201,9 +159,9 @@ export default function PrimarySidebar({
<path d={isCollapsed ? 'M9 18l6-6-6-6' : 'M15 18l-6-6 6-6'} />
</svg>
</button>
</div>}
</div>
{!isMobile && !isCollapsed && (
{!isCollapsed && (
<div
role="presentation"
onMouseDown={onMouseDown}
@@ -1,4 +1,4 @@
import { act, render, screen } from '@testing-library/react';
import { render, screen } from '@testing-library/react';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import GpuGrid from './GpuGrid';
@@ -73,30 +73,4 @@ describe('GpuGrid', () => {
await screen.findByText(/Could not reach the API server/),
).toBeInTheDocument();
});
it('keeps the last snapshot visible and warns when a refresh fails', async () => {
vi.useFakeTimers();
try {
vi.mocked(getGpus)
.mockResolvedValueOnce(SNAPSHOT)
.mockRejectedValueOnce(new Error('network down'));
render(<GpuGrid />);
await act(async () => {
await vi.advanceTimersByTimeAsync(0);
});
expect(screen.getAllByText('NVIDIA B200')).toHaveLength(2);
await act(async () => {
await vi.advanceTimersByTimeAsync(3000);
});
expect(
screen.getByText(/values below may be stale/),
).toBeInTheDocument();
expect(screen.getAllByText('NVIDIA B200')).toHaveLength(2);
} finally {
vi.useRealTimers();
}
});
});
@@ -1,9 +1,7 @@
'use client';
import * as React from 'react';
import { AlertTriangle } from 'lucide-react';
import { Button } from '@/components/ui/button';
import { Card, CardContent } from '@/components/ui/card';
import { getGpus, type GpuInfo, type GpuSnapshot } from '@/lib/api';
import { cn } from '@/lib/utils';
@@ -99,8 +97,7 @@ function GpuCard({ gpu }: { gpu: GpuInfo }) {
export default function GpuGrid() {
const [snapshot, setSnapshot] = React.useState<GpuSnapshot | null>(null);
const [fetchError, setFetchError] = React.useState<string | null>(null);
const [retryToken, setRetryToken] = React.useState(0);
const [fetchError, setFetchError] = React.useState(false);
React.useEffect(() => {
let mounted = true;
@@ -113,14 +110,10 @@ export default function GpuGrid() {
const next = await getGpus();
if (mounted) {
setSnapshot(next);
setFetchError(null);
setFetchError(false);
}
} catch {
if (mounted) {
setFetchError(
'GPU status could not be refreshed. The values below may be stale.',
);
}
if (mounted) setFetchError(true);
} finally {
inFlight = false;
}
@@ -132,27 +125,14 @@ export default function GpuGrid() {
mounted = false;
clearInterval(interval);
};
}, [retryToken]);
}, []);
if (fetchError && !snapshot) {
return (
<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="text-muted-foreground">
Could not reach the API server. GPU status needs the Studio API server
running.
</p>
<Button
type="button"
variant="outline"
onClick={() => setRetryToken((token) => token + 1)}
>
Try Again
</Button>
</div>
<p className="py-8 text-center text-muted-foreground">
Could not reach the API server. GPU status needs the studio API server
running.
</p>
);
}
if (!snapshot) {
@@ -170,22 +150,9 @@ export default function GpuGrid() {
return (
<div className="flex flex-col gap-4">
{fetchError && (
<div
role="status"
aria-live="polite"
className="flex flex-wrap items-center gap-3 rounded-lg border border-amber-500/50 bg-amber-500/10 px-3 py-2 text-sm"
>
<AlertTriangle className="size-4 text-amber-600" aria-hidden />
<span className="min-w-0 flex-1">{fetchError}</span>
<Button
type="button"
variant="outline"
size="sm"
onClick={() => setRetryToken((token) => token + 1)}
>
Refresh Now
</Button>
</div>
<p className="rounded-md border border-amber-500/40 bg-amber-500/10 px-3 py-2 text-sm text-amber-600 dark:text-amber-400">
Lost contact with the API server — showing the last known values.
</p>
)}
<div className="grid gap-4 [grid-template-columns:repeat(auto-fill,minmax(280px,1fr))]">
{snapshot.gpus.map((gpu) => (
@@ -1,48 +0,0 @@
import { render, screen } from '@testing-library/react';
import { describe, expect, it } from 'vitest';
import { Button } from './button';
import { Input } from './input';
import { NativeSelect } from './native-select';
import { Slider } from './slider';
import { Switch } from './switch';
describe('shared control accessibility', () => {
it('keeps button, input, and select targets at least 44px tall', () => {
render(
<>
<Button size="sm">Small action</Button>
<Input aria-label="Text value" />
<NativeSelect aria-label="Choice" defaultValue="one">
<option value="one">One</option>
</NativeSelect>
</>,
);
expect(screen.getByRole('button', { name: 'Small action' })).toHaveClass(
'h-11',
);
expect(screen.getByRole('textbox', { name: 'Text value' })).toHaveClass(
'h-11',
);
expect(screen.getByRole('combobox', { name: 'Choice' })).toHaveClass(
'h-11',
);
});
it('uses 44px switch and slider interaction surfaces', () => {
render(
<>
<Switch aria-label="Enabled" />
<Slider aria-label="Amount" defaultValue={[50]} />
</>,
);
expect(screen.getByRole('switch', { name: 'Enabled' })).toHaveClass(
'size-11',
);
expect(screen.getByRole('slider', { name: 'Amount' })).toHaveClass(
'size-11',
);
});
});
@@ -29,12 +29,11 @@ const badgeVariants = cva(
);
export interface BadgeProps
extends React.HTMLAttributes<HTMLSpanElement>,
extends React.HTMLAttributes<HTMLDivElement>,
VariantProps<typeof badgeVariants> {}
// A span (phrasing content), so badges stay valid inside buttons and links.
function Badge({ className, variant, ...props }: BadgeProps) {
return <span className={cn(badgeVariants({ variant }), className)} {...props} />;
return <div className={cn(badgeVariants({ variant }), className)} {...props} />;
}
export { Badge, badgeVariants };
@@ -7,7 +7,7 @@ import { cva, type VariantProps } from "class-variance-authority";
import { cn } from "@/lib/utils";
const buttonVariants = cva(
"inline-flex items-center justify-center gap-2 whitespace-nowrap rounded-xl border !text-sm !font-semibold transition-colors duration-150 disabled:pointer-events-none disabled:cursor-not-allowed disabled:opacity-50",
"inline-flex items-center justify-center gap-2 whitespace-nowrap rounded-xl border !text-sm !font-semibold transition-colors duration-150 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-sky-400/40 disabled:pointer-events-none disabled:cursor-not-allowed disabled:opacity-50",
{
variants: {
variant: {
@@ -18,11 +18,11 @@ const buttonVariants = cva(
destructive: "border-rose-500/60 bg-rose-600/90 text-white hover:bg-rose-500",
},
size: {
default: "h-11 px-4 py-2",
sm: "h-11 rounded-lg px-3 !text-xs",
lg: "h-12 px-5 !text-sm",
icon: "size-11",
"icon-sm": "size-11",
default: "h-10 px-4 py-2",
sm: "h-9 rounded-lg px-3 !text-xs",
lg: "h-11 px-5 !text-sm",
icon: "size-10",
"icon-sm": "size-8",
},
},
defaultVariants: {
@@ -42,7 +42,7 @@ const DialogContent = React.forwardRef<
{...props}
>
{children}
<DialogPrimitive.Close className="absolute right-2 top-2 flex size-11 items-center justify-center rounded-lg text-muted-foreground opacity-70 transition-opacity hover:bg-secondary hover:opacity-100 disabled:pointer-events-none sm:right-4 sm:top-4">
<DialogPrimitive.Close className="absolute right-4 top-4 rounded-lg p-1 text-muted-foreground opacity-70 transition-opacity hover:bg-secondary hover:opacity-100 focus:outline-none focus:ring-2 focus:ring-sky-400/40 disabled:pointer-events-none">
<X className="h-4 w-4" />
<span className="sr-only">Close</span>
</DialogPrimitive.Close>
@@ -9,7 +9,7 @@ const Input = React.forwardRef<HTMLInputElement, React.ComponentProps<"input">>(
<input
type={type}
className={cn(
"flex h-11 w-full rounded-xl border border-input bg-card/60 px-3 py-2 text-sm text-foreground shadow-sm backdrop-blur-md transition-colors placeholder:text-muted-foreground focus-visible:border-ring disabled:cursor-not-allowed disabled:opacity-50",
"flex h-10 w-full rounded-xl border border-input bg-card/60 px-3 py-2 text-sm text-foreground shadow-sm backdrop-blur-md transition-colors placeholder:text-muted-foreground focus-visible:border-sky-400/70 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-sky-400/25 disabled:cursor-not-allowed disabled:opacity-50",
className,
)}
ref={ref}
@@ -11,7 +11,7 @@ const NativeSelect = React.forwardRef<
<select
ref={ref}
className={cn(
'flex h-11 w-full appearance-none rounded-xl border border-input bg-card px-3 py-2 text-sm text-foreground shadow-sm transition-colors focus-visible:border-ring disabled:cursor-not-allowed disabled:opacity-50',
'flex h-10 w-full appearance-none rounded-xl border border-input bg-card px-3 py-2 text-sm text-foreground shadow-sm transition-colors focus-visible:border-sky-400/70 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-sky-400/25 disabled:cursor-not-allowed disabled:opacity-50',
className,
)}
{...props}
@@ -17,7 +17,7 @@ const SelectTrigger = React.forwardRef<
<SelectPrimitive.Trigger
ref={ref}
className={cn(
'flex h-11 w-full items-center justify-between gap-2 rounded-xl border border-input bg-card px-3 py-2 text-sm text-foreground shadow-sm transition-colors placeholder:text-muted-foreground focus-visible:border-ring disabled:cursor-not-allowed disabled:opacity-50 [&>span]:line-clamp-1',
'flex h-10 w-full items-center justify-between gap-2 rounded-xl border border-input bg-card px-3 py-2 text-sm text-foreground shadow-sm outline-none transition-colors placeholder:text-muted-foreground focus:border-sky-400/70 focus:ring-2 focus:ring-sky-400/25 disabled:cursor-not-allowed disabled:opacity-50 [&>span]:line-clamp-1',
className,
)}
{...props}
@@ -117,7 +117,7 @@ const SelectItem = React.forwardRef<
<SelectPrimitive.Item
ref={ref}
className={cn(
'relative flex min-h-11 w-full cursor-default select-none items-center rounded-xl py-2 pl-8 pr-3 text-sm text-foreground outline-none data-[disabled]:pointer-events-none data-[disabled]:opacity-50 data-[highlighted]:bg-accent data-[highlighted]:text-accent-foreground',
'relative flex w-full cursor-default select-none items-center rounded-xl py-2 pl-8 pr-3 text-sm text-foreground outline-none data-[disabled]:pointer-events-none data-[disabled]:opacity-50 data-[highlighted]:bg-accent data-[highlighted]:text-accent-foreground',
className,
)}
{...props}
@@ -8,39 +8,21 @@ import { cn } from '@/lib/utils';
const Slider = React.forwardRef<
React.ElementRef<typeof SliderPrimitive.Root>,
React.ComponentPropsWithoutRef<typeof SliderPrimitive.Root>
>(
(
{
>(({ className, ...props }, ref) => (
<SliderPrimitive.Root
ref={ref}
className={cn(
'relative flex w-full touch-none select-none items-center',
className,
id,
'aria-label': ariaLabel,
'aria-labelledby': ariaLabelledBy,
'aria-describedby': ariaDescribedBy,
...props
},
ref,
) => (
<SliderPrimitive.Root
ref={ref}
className={cn(
'relative flex h-11 w-full touch-none select-none items-center',
className,
)}
{...props}
>
<SliderPrimitive.Track className="relative h-1.5 w-full grow overflow-hidden rounded-full bg-border">
<SliderPrimitive.Range className="absolute h-full bg-accent-blue" />
</SliderPrimitive.Track>
<SliderPrimitive.Thumb
id={id}
aria-label={ariaLabel}
aria-labelledby={ariaLabelledBy}
aria-describedby={ariaDescribedBy}
className="relative block size-11 rounded-full bg-transparent after:absolute after:left-1/2 after:top-1/2 after:size-4 after:-translate-x-1/2 after:-translate-y-1/2 after:rounded-full after:border-2 after:border-accent-blue after:bg-accent-blue after:shadow after:content-[''] hover:after:border-accent-blue/80 disabled:pointer-events-none disabled:opacity-50"
/>
</SliderPrimitive.Root>
),
);
)}
{...props}
>
<SliderPrimitive.Track className="relative h-1.5 w-full grow overflow-hidden rounded-full bg-border">
<SliderPrimitive.Range className="absolute h-full bg-accent-blue" />
</SliderPrimitive.Track>
<SliderPrimitive.Thumb className="block h-4 w-4 rounded-full border-2 border-accent-blue bg-accent-blue shadow transition-colors hover:border-accent-blue/80 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-sky-400/40 disabled:pointer-events-none disabled:opacity-50" />
</SliderPrimitive.Root>
));
Slider.displayName = SliderPrimitive.Root.displayName;
export { Slider };
@@ -12,14 +12,14 @@ const Switch = React.forwardRef<
<SwitchPrimitives.Root
ref={ref}
className={cn(
'peer relative inline-flex size-11 shrink-0 cursor-pointer items-center justify-center rounded-xl bg-transparent transition-colors before:absolute before:h-5 before:w-9 before:rounded-full before:border before:border-border before:bg-background data-[state=checked]:before:border-accent-blue data-[state=checked]:before:bg-accent-blue disabled:cursor-not-allowed disabled:opacity-50',
'peer inline-flex h-5 w-9 shrink-0 cursor-pointer items-center rounded-full border border-border transition-colors focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-sky-400/40 disabled:cursor-not-allowed disabled:opacity-50 data-[state=checked]:border-accent-blue data-[state=checked]:bg-accent-blue data-[state=unchecked]:bg-background',
className,
)}
{...props}
>
<SwitchPrimitives.Thumb
className={cn(
'pointer-events-none absolute left-1.5 top-[15px] z-[1] block h-3.5 w-3.5 rounded-full bg-muted-foreground shadow-lg ring-0 transition-transform data-[state=checked]:translate-x-4 data-[state=checked]:bg-white',
'pointer-events-none block h-3.5 w-3.5 rounded-full bg-muted-foreground shadow-lg ring-0 transition-transform data-[state=checked]:translate-x-4 data-[state=checked]:bg-white data-[state=unchecked]:translate-x-0.5',
)}
/>
</SwitchPrimitives.Root>
@@ -29,7 +29,7 @@ const TabsTrigger = React.forwardRef<
<TabsPrimitive.Trigger
ref={ref}
className={cn(
'inline-flex items-center justify-center whitespace-nowrap rounded-md px-3 py-1.5 text-sm font-medium text-muted-foreground transition-colors hover:text-foreground disabled:pointer-events-none disabled:opacity-50 data-[state=active]:bg-card data-[state=active]:text-foreground data-[state=active]:shadow-sm',
'inline-flex items-center justify-center whitespace-nowrap rounded-md px-3 py-1.5 text-sm font-medium text-muted-foreground transition-colors hover:text-foreground focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-sky-400/40 disabled:pointer-events-none disabled:opacity-50 data-[state=active]:bg-card data-[state=active]:text-foreground data-[state=active]:shadow-sm',
className,
)}
{...props}
@@ -8,7 +8,7 @@ const Textarea = React.forwardRef<HTMLTextAreaElement, React.ComponentProps<"tex
return (
<textarea
className={cn(
"flex min-h-24 w-full resize-y rounded-xl border border-input bg-card px-3 py-2 text-sm text-foreground shadow-sm transition-colors placeholder:text-muted-foreground focus-visible:border-ring disabled:cursor-not-allowed disabled:opacity-50",
"flex min-h-24 w-full resize-y rounded-xl border border-input bg-card px-3 py-2 text-sm text-foreground shadow-sm transition-colors placeholder:text-muted-foreground focus-visible:border-sky-400/70 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-sky-400/25 disabled:cursor-not-allowed disabled:opacity-50",
className,
)}
ref={ref}
@@ -1,23 +0,0 @@
'use client';
import * as React from 'react';
/**
* Move focus into a drawer when it opens as a modal (mobile), and hand it
* back to the previously focused element on close. Pairs with `inert` on
* the background content — together they make `aria-modal` truthful.
*/
export function useDrawerFocus<T extends HTMLElement>(active: boolean) {
const ref = React.useRef<T | null>(null);
React.useEffect(() => {
if (!active) return;
const previous = document.activeElement;
ref.current?.focus();
return () => {
if (previous instanceof HTMLElement) previous.focus();
};
}, [active]);
return ref;
}
@@ -1,18 +0,0 @@
'use client';
import * as React from 'react';
export function useMediaQuery(query: string): boolean {
const [matches, setMatches] = React.useState(false);
React.useEffect(() => {
const mediaQuery = window.matchMedia(query);
const updateMatch = () => setMatches(mediaQuery.matches);
updateMatch();
mediaQuery.addEventListener('change', updateMatch);
return () => mediaQuery.removeEventListener('change', updateMatch);
}, [query]);
return matches;
}
+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"
+21 -40
View File
@@ -65,10 +65,10 @@ ARG FLASH_ATTN_WHEEL_RELEASE=https://github.com/mjun0812/flash-attention-prebuil
ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.22
# The prebuilt flash-attn wheel ships an FA4 `flash_attn.cute` built against the
# 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
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the CuTe DSL 4.6 that
# flashinfer/quack pull in. After the wheel install we overlay this 4.6.0.dev0-compatible
# 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
@@ -169,57 +169,38 @@ RUN --mount=type=cache,target=/opt/uv/cache \
# flash_attn/__init__.py), so FA2/varlen/bert_padding stay from the install above;
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
# Then verify both survive so a broken overlay fails the build instead of shipping
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
# arm64 and the image runs FA3/FA2 as usual.
# an FA2-less image. CUDA 13 selects the matching runtime extra on both amd64 and
# arm64 (validated on GB200). CUDA 12.6 amd64 keeps the bare/default runtime;
# CUDA 12.6 arm64 remains unvalidated and skips the overlay. FA4 stays opt-in, so
# installing it does not auto-enable the backend on GB10 (sm_121).
RUN --mount=type=cache,target=/opt/uv/cache \
source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
if [ "${TARGETARCH}" = "arm64" ]; then \
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
if [ "${UV_TORCH_BACKEND}" = "cu130" ]; then \
FA4_PACKAGE="flash-attn-4[cu13]"; \
elif [ "${TARGETARCH}" = "arm64" ]; then \
echo "Skipping FA4 cute overlay on CUDA 12.6 arm64 (do not set FASTVIDEO_FA4)"; \
FA4_PACKAGE=""; \
else \
FA4_PACKAGE="flash-attn-4"; \
fi && \
if [ -n "${FA4_PACKAGE}" ]; then \
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
uv pip install "${FA4_PACKAGE} @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
python -c "import flash_attn; assert hasattr(flash_attn, 'flash_attn_func'), 'FA2 was clobbered by the cute overlay'; import flash_attn.cute; print('FA2 + FA4 cute OK')"; \
fi
COPY . .
# Build immutable FastVideo kernel wheels for the published image. The requested
# architecture remains installed for normal image users; amd64 images also carry
# an SM89 artifact so the predominant L40S Modal lanes can reuse it exactly.
ARG FASTVIDEO_KERNEL_PREBUILT_DIR=/opt/fastvideo-kernel-prebuilt
# Install FastVideo Unified Kernel exactly once. build.sh initializes only its
# CUTLASS/ThunderKittens submodules, then compiles for TORCH_CUDA_ARCH_LIST
# (default Hopper sm_90a) without requiring a live GPU.
RUN --mount=type=cache,target=/opt/uv/cache \
source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
default_arch="${TORCH_CUDA_ARCH_LIST}" && \
default_wheel_dir="${FASTVIDEO_KERNEL_PREBUILT_DIR}/${default_arch}" && \
export TORCH_CUDA_ARCH_LIST="${default_arch}" && \
cd fastvideo-kernel && \
CMAKE_ARGS= CMAKE_BUILD_PARALLEL_LEVEL=${CMAKE_BUILD_PARALLEL_LEVEL} \
./build.sh --wheel-dir "${default_wheel_dir}" && \
cd /FastVideo && \
CMAKE_ARGS= python fastvideo/tests/modal/kernel_build_cache.py write-build-info \
--wheel-dir "${default_wheel_dir}" \
--output "${default_wheel_dir}/metadata.json" && \
if [[ "${TARGETARCH:-amd64}" == "amd64" && "${default_arch}" != "8.9" ]]; then \
export TORCH_CUDA_ARCH_LIST=8.9 && \
l40s_wheel_dir="${FASTVIDEO_KERNEL_PREBUILT_DIR}/8.9" && \
cd /FastVideo/fastvideo-kernel && \
CMAKE_ARGS= CMAKE_BUILD_PARALLEL_LEVEL=${CMAKE_BUILD_PARALLEL_LEVEL} \
./build.sh --wheel-dir "${l40s_wheel_dir}" && \
cd /FastVideo && \
CMAKE_ARGS= python fastvideo/tests/modal/kernel_build_cache.py write-build-info \
--wheel-dir "${l40s_wheel_dir}" \
--output "${l40s_wheel_dir}/metadata.json"; \
fi && \
export TORCH_CUDA_ARCH_LIST="${default_arch}" && \
default_wheel="$(find "${default_wheel_dir}" -maxdepth 1 -type f \
\( -name 'fastvideo_kernel-*.whl' -o -name 'fastvideo-kernel-*.whl' \) \
| sort | tail -n 1)" && \
uv pip install "${default_wheel}" \
--reinstall-package fastvideo-kernel --no-deps
CMAKE_BUILD_PARALLEL_LEVEL=${CMAKE_BUILD_PARALLEL_LEVEL} \
TORCH_CUDA_ARCH_LIST=${TORCH_CUDA_ARCH_LIST} ./build.sh
# Install FastVideo itself (editable) now that the source is present, and set up
# shell configuration. Dependencies and the local kernel are already installed,
-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;
+10 -20
View File
@@ -269,26 +269,16 @@ The docs job:
### Docker Images
`.github/workflows/infra-build-image.yml` supports manual `workflow_dispatch`
runs and automatically rebuilds the CUDA matrix when a repository-controlled
image input changes on `main` in the canonical repository. Those inputs include
the CUDA Dockerfile and reusable workflow, dependency metadata, Docker context
policy, `fastvideo-kernel/**`, and the kernel artifact metadata/key helper.
Manual runs let maintainers choose which image families to build. The
`fastvideo-dev` matrix builds Python 3.12 images for CUDA 12.6 and CUDA 13 on
native `amd64` and `arm64` runners, then publishes one multi-platform manifest
per CUDA version. CUDA 12.6 owns the `py3.12-latest` and global `latest` tags, as
well as the explicit `py3.12-cuda12.6.3-latest` alias. CUDA 13 is published under
the explicit `py3.12-cuda13.0.0-latest` tag. This publication policy does not
change the unparameterized `docker/Dockerfile` build defaults, which remain CUDA
13 and `cu130`.
Published amd64 development images keep their configured Hopper kernel wheel
installed and also carry an immutable SM89 wheel under
`/opt/fastvideo-kernel-prebuilt`. Modal PR and SSIM jobs select the exact
source, ABI, and GPU-architecture match from that directory, so L40S jobs reuse
the trusted image artifact while kernel-changing PRs still build locally. Once
a kernel or artifact-key change reaches `main`, the image workflow republishes
the matching trusted artifact before later jobs consume the updated image tag.
runs and automatically rebuilds the CUDA matrix when `docker/Dockerfile`
changes on `main` in the canonical repository. Manual runs let maintainers
choose which image families to build. The `fastvideo-dev` matrix builds Python
3.12 images for CUDA 12.6 and CUDA 13 on native `amd64` and `arm64` runners,
then publishes one multi-platform manifest per CUDA version. CUDA 12.6 owns the
`py3.12-latest` and global `latest` tags, as well as the explicit
`py3.12-cuda12.6.3-latest` alias. CUDA 13 is published under the explicit
`py3.12-cuda13.0.0-latest` tag. This publication policy does not change the
unparameterized `docker/Dockerfile` build defaults, which remain CUDA 13 and
`cu130`.
The optional Dreamverse matrix builds backend and UI images for CUDA 12.6 and
CUDA 13 on `amd64`. Dreamverse remains `amd64`-only because its FA4 dependency
-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
+1 -1
View File
@@ -134,4 +134,4 @@ If you're planning to contribute to FastVideo please see the following page:
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ) for additional support.
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
+5 -7
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
@@ -102,4 +100,4 @@ If you're planning to contribute to FastVideo please see the following page:
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ) for additional support.
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
+1 -8
View File
@@ -134,16 +134,9 @@ uv pip install "https://github.com/mjun0812/flash-attention-prebuild-wheels/rele
If you hit other issues, please open an issue on our
[GitHub repository](https://github.com/hao-ai-lab/FastVideo). You can also join
our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)
our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ)
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
+1 -1
View File
@@ -11,7 +11,7 @@ This guide explains how to implement a custom diffusion pipeline in FastVideo, l
4. **Register Your Pipeline** - Make it discoverable by the framework
5. **Configure Your Pipeline** - (Coming soon)
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ).
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
## Step 1: Pipeline Modules
+5 -5
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
@@ -135,4 +134,5 @@ If the generated video doesn't match your prompt:
- Learn about [Advanced Inference Configurations](configuration.md)
- Learn about using [Optimizations](optimizations.md)
- See [Examples](examples/examples_inference_index.md) for more usage scenarios
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ).
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
+11 -29
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
@@ -84,8 +78,8 @@ python setup.py install
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
runtime on some GPU/shape combinations. To use FA4, install the pinned
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
runtime on some GPU/shape combinations. To use FA4 on CUDA 13, install the pinned
`flash-attn-4[cu13]` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
```bash
export FASTVIDEO_FA4=1
@@ -116,30 +110,18 @@ See the [Attn-QAT paper](https://arxiv.org/abs/2603.00040) and [flash-attention-
Install the FP4 flash attention kernel (without upgrading your existing torch):
```bash
# branch fix/cutlass-dsl-4.5 carries the cutlass-dsl 4.5 fix (cute.core.ThrMma
# -> cute.ThrMma); switch back to @fp4 once hao-ai-lab/flash-attention-fp4#2 merges.
pip install --no-deps "git+ssh://git@github.com/hao-ai-lab/flash-attention-fp4.git@fix/cutlass-dsl-4.5#subdirectory=flash_attn/cute"
pip install "nvidia-cutlass-dsl>=4.5.2" apache-tvm-ffi flashinfer-python
# Commit 940bf7e5 carries the CuTe DSL 4.5 fix (cute.core.ThrMma -> cute.ThrMma).
pip install --no-deps "git+https://github.com/hao-ai-lab/flash-attention-fp4.git@940bf7e511375ec160bc2d7188bef35915ded1e3#subdirectory=flash_attn/cute"
pip install "nvidia-cutlass-dsl[cu13]==4.5.2" "quack-kernels==0.5.0" apache-tvm-ffi flashinfer-python torch-c-dlpack-ext
```
The `--no-deps` flag prevents upgrading torch/torchvision. Use the supported
PyTorch 2.12.0 and CUDA 13 environment for this kernel.
Branch-to-`nvidia-cutlass-dsl` compatibility (the fork tracks the CuTe DSL API
surface closely):
| fork branch | cutlass-dsl | notes |
|---|---|---|
| `fp4` | `==4.4.2` (+ `nvidia-cutlass-dsl-libs-base==4.4.2`) | validated set on GB200: `quack-kernels==0.4.1`, `flashinfer-python==0.6.8`, `CUTE_DSL_ENABLE_TVM_FFI=1`, `FASTVIDEO_FA4=1` |
| `fix/cutlass-dsl-4.5` | `>=4.5.2` | carries the `cute.core.ThrMma` -> `cute.ThrMma` fix |
| any | 4.6-era | unsupported: `cute.make_fragment` was removed at module level; fails at CuTe JIT trace |
`FASTVIDEO_FA4=1` is required alongside the fork: it ships no compiled
FlashAttention-2, so dense attention paths raise ImportError without the FA4
opt-in. The same kernel also serves `ATTN_QAT_INFER` on sm_100a/sm_103a
(datacenter Blackwell) — the selection log's receipt line
(`ATTN_QAT_INFER resolved: ...`) records the arch, kernel, and quantization
mode that actually bound.
PyTorch 2.12.0 and CUDA 13 environment for this kernel. Keep this private FP4
overlay in a separate inference environment from FastVideo's upstream dense FA4
dependency: the two distributions provide the same `flash_attn.cute` package but
require different CuTe DSL versions. Keep the CUTLASS and Quack pins together;
Quack 0.5.1 and newer use CuTe DSL 4.6, whose API is incompatible with this FP4
kernel revision.
#### Usage
+5 -130
View File
@@ -13,102 +13,6 @@ For the canonical, code-level list of model IDs recognized by
We do this because we believe VSA is strictly better than STA for the
actively maintained `main` inference path.
## Registered Model IDs
Every Hugging Face model ID registered in `fastvideo/registry.py` on `main`
(commit `8d89f30d`), grouped by family. Any ID below can
be passed to `VideoGenerator.from_pretrained(...)`; FastVideo resolves the
matching pipeline and sampling defaults. The **Family** column is a
documentation grouping: it follows each registration's declared `model_family`,
except `black-forest-labs/FLUX.1-dev`, which declares none and is listed under
`flux` for readability. The **Workloads** column shows each
registration's declared `workload_types`; `—` means the entry is registered
without a UI workload option but is still loadable by ID. The **Example**
column links a runnable script in `examples/inference/basic/` where one exists.
| Family | HuggingFace Model ID | Workloads | Example |
|--------|----------------------|-----------|---------|
| cosmos | `nvidia/Cosmos-Predict2-2B-Video2World` | T2V | — |
| cosmos25 | `KyleShao/Cosmos-Predict2.5-2B-Diffusers` | T2V | [basic_cosmos2_5_t2w.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_cosmos2_5_t2w.py) |
| cosmos25 | `nvidia/Cosmos-Predict2.5-14B` | T2V | [basic_cosmos2_5_t2w.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_cosmos2_5_t2w.py) |
| dreamx_world | `FastVideo/DreamX-World-5B-Cam-Diffusers` | I2V | [basic_dreamx_world.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_dreamx_world.py) |
| dreamx_world | `FastVideo/DreamX-World-5B-Diffusers` | I2V | [basic_dreamx_world.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_dreamx_world.py) |
| flux | `black-forest-labs/FLUX.1-dev` | T2I | [basic_flux_dev.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_flux_dev.py) |
| flux2 | `black-forest-labs/FLUX.2-klein-4B`<br>`black-forest-labs/FLUX.2-klein-9B` | T2I | [basic_flux2_klein.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_flux2_klein.py) |
| flux2 | `black-forest-labs/FLUX.2-dev` | T2I | [basic_flux2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_flux2.py) |
| gamecraft | `FastVideo/HunyuanGameCraft-Diffusers` | I2V | [basic_gamecraft.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_gamecraft.py) |
| gen3c | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | T2V | [basic_gen3c.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_gen3c.py) |
| glm_image | `zai-org/GLM-Image` | T2I | [basic_glm_image.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_glm_image.py) |
| hunyuan | `hunyuanvideo-community/HunyuanVideo` | T2V | — |
| hunyuan | `FastVideo/FastHunyuan-diffusers` | T2V | — |
| hunyuan15 | `hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v` | T2V | [basic_hy15.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_hy15.py) |
| hunyuan15 | `hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled` | I2V | [basic_hy15.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_hy15.py) |
| hunyuan15 | `hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v` | T2V | [basic_hy15.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_hy15.py) |
| hunyuan15 | `hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_i2v_distilled` | I2V | [basic_hy15.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_hy15.py) |
| hunyuan15 | `weizhou03/HunyuanVideo-1.5-Diffusers-1080p`<br>`weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR` | — | [basic_hy15_1080p.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_hy15_1080p.py) |
| hyworld | `FastVideo/HY-WorldPlay-Bidirectional-Diffusers` | — | [basic_hyworld.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_hyworld.py) |
| kandinsky5 | `kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers` | T2V | [basic_kandinsky5_t2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_kandinsky5_t2v.py) |
| kandinsky5 | `kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers` | T2V | [basic_kandinsky5_t2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_kandinsky5_t2v.py) |
| kandinsky5 | `kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers` | T2V | [basic_kandinsky5_t2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_kandinsky5_t2v.py) |
| kandinsky5 | `kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers` | T2V | [basic_kandinsky5_t2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_kandinsky5_t2v.py) |
| kandinsky5 | `kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers` | I2V | [basic_kandinsky5_i2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_kandinsky5_i2v.py) |
| kandinsky5 | `kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers` | I2V | [basic_kandinsky5_i2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_kandinsky5_i2v.py) |
| kandinsky5 | `kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers` | I2V | [basic_kandinsky5_i2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_kandinsky5_i2v.py) |
| lingbot_video | `FastVideo/LingBot-Video-MoE-30B-A3B-Diffusers` | T2V | [basic_lingbot_video.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_lingbot_video.py) |
| lingbot_video | `FastVideo/LingBot-Video-Dense-1.3B-Diffusers` | T2V | [basic_lingbot_video.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_lingbot_video.py) |
| lingbotworld | `FastVideo/LingBot-World-Base-Cam-Diffusers` | I2V | [basic_lingbotworld_base_cam.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_lingbotworld_base_cam.py) |
| lingbotworld2 | `robbyant/lingbot-world-v2-14b-causal-fast` | I2V | [basic_lingbotworld2_causal_fast.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_lingbotworld2_causal_fast.py) |
| longcat | `FastVideo/LongCat-Video-T2V-Diffusers` | T2V | [basic_longcat_t2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_longcat_t2v.py) |
| longcat | `FastVideo/LongCat-Video-I2V-Diffusers` | I2V | [basic_longcat_i2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_longcat_i2v.py) |
| longcat | `FastVideo/LongCat-Video-VC-Diffusers` | — | [basic_longcat_vc.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_longcat_vc.py) |
| 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) |
| turbodiffusion | `loayrashid/TurboWan2.1-T2V-1.3B-Diffusers` | T2V | [basic_turbodiffusion.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_turbodiffusion.py) |
| turbodiffusion | `loayrashid/TurboWan2.1-T2V-14B-Diffusers` | T2V | [basic_turbodiffusion_14b.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_turbodiffusion_14b.py) |
| turbodiffusion | `loayrashid/TurboWan2.2-I2V-A14B-Diffusers` | I2V | [basic_turbodiffusion_i2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_turbodiffusion_i2v.py) |
| wan | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | T2V | [basic.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic.py) |
| wan | `Wan-AI/Wan2.1-T2V-14B-Diffusers`<br>`FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers` | T2V | — |
| wan | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | I2V | — |
| wan | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | I2V | — |
| wan | `weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers` | I2V | — |
| wan | `IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers` | — | [basic_wan2_2_Fun.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_wan2_2_Fun.py) |
| wan | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`<br>`FastVideo/FastWan2.1-T2V-14B-480P-Diffusers` | T2V | [basic_dmd.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_dmd.py) |
| wan | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | T2V, I2V | [basic_wan2_2_ti2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_wan2_2_ti2v.py) |
| wan | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers`<br>`FastVideo/FastWan2.2-TI2V-5B-Diffusers` | T2V, I2V | [basic_dmd.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_dmd.py) |
| wan | `decart-ai/Lucy-Edit-Dev`<br>`decart-ai/Lucy-Edit-1.1-Dev` | — | [basic_lucy_edit.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_lucy_edit.py) |
| wan | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | T2V | [basic_wan2_2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_wan2_2.py) |
| wan | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | I2V | [basic_wan2_2_i2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_wan2_2_i2v.py) |
| wan | `wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers` | T2V | [basic_self_forcing_causal.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal.py) |
| wan | `rand0nmr/SFWan2.2-T2V-A14B-Diffusers` | T2V | [basic_self_forcing_causal_wan2_2_t2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_t2v.py) |
| wan | `FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers` | I2V | [basic_self_forcing_causal_wan2_2_i2v.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py) |
| zimage | `Tongyi-MAI/Z-Image-Turbo` | T2I | [basic_zimage.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_zimage.py) |
**Note (stable_audio)**: the Stable Audio Open pipelines generate audio
(`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
supported path is the Wan2.1-Fun control pipeline
(`IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers`).
The symbols used have the following meanings:
- ✅ = Full compatibility
@@ -121,9 +25,6 @@ The `HuggingFace Model ID` can be passed directly to
`from_pretrained()`. FastVideo then uses model-specific default settings for
pipeline initialization and sampling.
Registered models absent from this table have not been validated against these
optimizations: absence means **untested**, not incompatible.
<style>
/* Target tables in this section */
#models-x-optimization + p + table {
@@ -155,7 +56,7 @@ optimizations: absence means **untested**, not incompatible.
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn (Legacy Branch) | Sage Attn | VSA | BSA |
|------------|---------------------|-------------|----------|-------------------|-----------|-----|-----|
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
@@ -164,31 +65,20 @@ optimizations: absence means **untested**, not incompatible.
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ | ⭕ |
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅ | ✅ | ⭕ | ⭕ |
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅ | ✅ | ⭕ | ⭕ |
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅ | ✅ | ⭕ | ⭕ |
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ | ⭕ |
| TurboWan2.1 T2V 1.3B | `loayrashid/TurboWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| TurboWan2.1 T2V 14B | `loayrashid/TurboWan2.1-T2V-14B-Diffusers` | 480P, 720P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| TurboWan2.2 I2V A14B | `loayrashid/TurboWan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| LongCat T2V 13.6B | `FastVideo/LongCat-Video-T2V-Diffusers` | 480P<br>720P | ❌ | ❌ | ❌ | ⭕ | ✅ |
| LongCat T2V 13.6B | See note** | 480P<br>720P | ❌ | ❌ | ❌ | ⭕ | ✅ |
| Matrix Game 2.0 Base Distilled | `FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Matrix Game 2.0 GTA Distilled | `FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Matrix Game 2.0 TempleRun Distilled | `FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| 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
@@ -208,21 +98,6 @@ resolve default pipeline and sampling configuration for it.
`FastVideo/GEN3C-Cosmos-7B-Diffusers`) or convert locally with
`scripts/checkpoint_conversion/convert_gen3c_to_fastvideo.py`.
## Hardware and OS
Per the installation guides:
- **NVIDIA GPU (x86_64)** — CUDA 12.6 or 13.0; see the
[GPU install guide](../getting_started/installation/gpu.md).
- **NVIDIA DGX Spark (GB10, aarch64)** — CUDA 13, from-source kernel build; see
the [DGX Spark install guide](../getting_started/installation/spark.md).
- **Apple silicon (MPS)** — macOS 14 or newer; see the
[MPS install guide](../getting_started/installation/mps.md) and
[`basic_mps.py`](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mps.py).
Optimization-specific hardware constraints (e.g. STA requiring Hopper) are
listed under [Special requirements](#special-requirements).
## Special requirements
### Sliding Tile Attention
+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

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