Compare commits
52
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a02634fc72 | ||
|
|
1ec849a4e3 | ||
|
|
8208536cd1 | ||
|
|
0653f8f3af | ||
|
|
e0d702decb | ||
|
|
541ef014ee | ||
|
|
ffc1a7a58b | ||
|
|
9028953625 | ||
|
|
6eb95693a1 | ||
|
|
c3567eb468 | ||
|
|
15568f27db | ||
|
|
126a75ad63 | ||
|
|
a2bfc7cdb2 | ||
|
|
b963a24612 | ||
|
|
fb7be2fe2c | ||
|
|
ab00392664 | ||
|
|
9f1e7c19d2 | ||
|
|
e8b0e4c61e | ||
|
|
9145ffdc46 | ||
|
|
c3d07c870b | ||
|
|
e8812bef0b | ||
|
|
b9be2449dc | ||
|
|
bc7a804618 | ||
|
|
7b094c945b | ||
|
|
eeb3e8a597 | ||
|
|
05406c5d1b | ||
|
|
99d04a7f98 | ||
|
|
1b2b2a0161 | ||
|
|
98d65835b5 | ||
|
|
422585d08f | ||
|
|
e59a1ce16a | ||
|
|
d71acc0eb5 | ||
|
|
af2934dd6b | ||
|
|
1801512818 | ||
|
|
5ae05b032e | ||
|
|
7a592ff09a | ||
|
|
8b23984c79 | ||
|
|
bf18371afe | ||
|
|
69349dd2aa | ||
|
|
8d89f30d3f | ||
|
|
10546353da | ||
|
|
9fb74b9732 | ||
|
|
521dee0e82 | ||
|
|
229419208e | ||
|
|
65f3b946b9 | ||
|
|
191fcbf46c | ||
|
|
755a4e4470 | ||
|
|
9709b7513b | ||
|
|
32cd603515 | ||
|
|
d4bdd3621a | ||
|
|
e2f8322842 | ||
|
|
6966f9e0bc |
@@ -10,9 +10,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Clone a reference repo for FastVideo parity tests."
|
||||
)
|
||||
parser = argparse.ArgumentParser(description="Clone a reference repo for FastVideo parity tests.")
|
||||
parser.add_argument("repo_url", help="Official reference repository URL")
|
||||
parser.add_argument("target_dir", help="Directory to clone into")
|
||||
parser.add_argument("--branch", help="Branch or tag to clone")
|
||||
@@ -62,9 +60,7 @@ def gitignore_entry_for(target: Path) -> str:
|
||||
try:
|
||||
relative = resolved.relative_to(root)
|
||||
except ValueError as exc:
|
||||
raise ValueError(
|
||||
"--update-gitignore requires target_dir to be under the current directory"
|
||||
) from exc
|
||||
raise ValueError("--update-gitignore requires target_dir to be under the current directory") from exc
|
||||
|
||||
text = relative.as_posix().rstrip("/")
|
||||
return "/" + text + "/"
|
||||
|
||||
@@ -8,14 +8,12 @@ import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Download a HF model snapshot or selected files into a local directory."
|
||||
)
|
||||
description="Download a HF model snapshot or selected files into a local directory.")
|
||||
parser.add_argument("repo_id", help="HF repo id, for example Org/Model")
|
||||
parser.add_argument("local_dir", help="Destination directory")
|
||||
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
|
||||
|
||||
@@ -10,7 +10,6 @@ import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
|
||||
RAW_WEIGHT_SUFFIXES = (".safetensors", ".pt", ".pth", ".ckpt", ".bin")
|
||||
KNOWN_COMPONENTS = {
|
||||
@@ -34,8 +33,7 @@ KNOWN_COMPONENTS = {
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown."
|
||||
)
|
||||
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown.")
|
||||
parser.add_argument("source", help="HF repo id or local weights directory")
|
||||
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
|
||||
parser.add_argument("--revision", help="HF revision to inspect")
|
||||
@@ -94,14 +92,12 @@ def load_remote_files(
|
||||
) -> list[str]:
|
||||
from huggingface_hub import list_repo_files
|
||||
|
||||
return sorted(
|
||||
list_repo_files(
|
||||
repo_id,
|
||||
repo_type=repo_type,
|
||||
revision=revision,
|
||||
token=token,
|
||||
)
|
||||
)
|
||||
return sorted(list_repo_files(
|
||||
repo_id,
|
||||
repo_type=repo_type,
|
||||
revision=revision,
|
||||
token=token,
|
||||
))
|
||||
|
||||
|
||||
def load_remote_model_index(
|
||||
@@ -215,24 +211,24 @@ def build_result(args: argparse.Namespace) -> dict[str, Any]:
|
||||
"components_seen": components,
|
||||
"file_count": len(files),
|
||||
"file_scan_truncated": truncated,
|
||||
"files_sample": files[: args.sample_limit],
|
||||
"files_sample": files[:args.sample_limit],
|
||||
}
|
||||
|
||||
|
||||
def print_human(result: dict[str, Any]) -> None:
|
||||
for key in (
|
||||
"source",
|
||||
"source_kind",
|
||||
"repo_type",
|
||||
"revision",
|
||||
"token_env",
|
||||
"source_layout",
|
||||
"needs_conversion",
|
||||
"model_index_class",
|
||||
"model_index_diffusers_version",
|
||||
"model_index_error",
|
||||
"file_count",
|
||||
"file_scan_truncated",
|
||||
"source",
|
||||
"source_kind",
|
||||
"repo_type",
|
||||
"revision",
|
||||
"token_env",
|
||||
"source_layout",
|
||||
"needs_conversion",
|
||||
"model_index_class",
|
||||
"model_index_diffusers_version",
|
||||
"model_index_error",
|
||||
"file_count",
|
||||
"file_scan_truncated",
|
||||
):
|
||||
value = result.get(key)
|
||||
if value is not None:
|
||||
|
||||
@@ -18,7 +18,6 @@ import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
@@ -35,15 +34,10 @@ FASTVIDEO_CONFIG_CLASS = "<FastVideoConfig>" # TODO.
|
||||
FASTVIDEO_MODEL_MODULE = "fastvideo.models.<bucket>.<module>" # TODO.
|
||||
FASTVIDEO_MODEL_CLASS = "<FastVideoModel>" # TODO.
|
||||
|
||||
OFFICIAL_REF_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>")
|
||||
)
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY)
|
||||
)
|
||||
CONVERTED_WEIGHTS_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / FAMILY)
|
||||
)
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>"))
|
||||
LOCAL_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY))
|
||||
CONVERTED_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR",
|
||||
REPO_ROOT / "converted_weights" / FAMILY))
|
||||
|
||||
|
||||
def _resolve_hf_token() -> str | None:
|
||||
@@ -99,18 +93,14 @@ def _load_official_model(device: torch.device, dtype: torch.dtype) -> torch.nn.M
|
||||
model = OfficialClass() # TODO: pass official config kwargs.
|
||||
state_dict = {} # TODO: load official state dict from LOCAL_WEIGHTS_DIR.
|
||||
missing, unexpected = model.load_state_dict(state_dict, strict=True)
|
||||
assert not missing and not unexpected, (
|
||||
f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
|
||||
)
|
||||
assert not missing and not unexpected, (f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
"""Load the FastVideo component with the same tensor content."""
|
||||
if not CONVERTED_WEIGHTS_DIR.exists() and not LOCAL_WEIGHTS_DIR.exists():
|
||||
pytest.skip(
|
||||
f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}"
|
||||
)
|
||||
pytest.skip(f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}")
|
||||
|
||||
# TODO: replace with the bucket-specific FastVideo config/class/loader.
|
||||
# DiT examples:
|
||||
@@ -127,8 +117,7 @@ def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.
|
||||
state_dict = {} # TODO: load converted or directly mapped state dict.
|
||||
missing, unexpected = model.load_state_dict(state_dict, strict=True)
|
||||
assert not missing and not unexpected, (
|
||||
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
|
||||
)
|
||||
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
@@ -187,11 +176,9 @@ def test_component_parity():
|
||||
|
||||
assert official_out.shape == fastvideo_out.shape
|
||||
diff = (official_out - fastvideo_out).abs()
|
||||
print(
|
||||
f"official abs_mean={official_out.abs().mean().item():.6f} "
|
||||
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
|
||||
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}"
|
||||
)
|
||||
print(f"official abs_mean={official_out.abs().mean().item():.6f} "
|
||||
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
|
||||
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
|
||||
|
||||
# TODO: pick tolerance by scope:
|
||||
# - single block / same kernel: 1e-4
|
||||
|
||||
@@ -27,7 +27,6 @@ try:
|
||||
except ImportError: # pragma: no cover - optional local conversion dependency
|
||||
snapshot_download = None
|
||||
|
||||
|
||||
# TODO: fill with authoritative component prefixes for monolithic checkpoints.
|
||||
# Example: {"model.model.": "transformer", "pretransform.model.": "vae"}
|
||||
COMPONENT_PREFIXES: dict[str, str] = {}
|
||||
@@ -47,10 +46,7 @@ SKIP_PATTERNS: tuple[str, ...] = ()
|
||||
|
||||
|
||||
def _hf_token() -> str | None:
|
||||
return (
|
||||
os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN")
|
||||
or os.environ.get("HF_API_KEY")
|
||||
)
|
||||
return (os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN") or os.environ.get("HF_API_KEY"))
|
||||
|
||||
|
||||
def resolve_src(src: str, revision: str | None) -> Path:
|
||||
@@ -95,11 +91,10 @@ def apply_mapping(key: str) -> str | None:
|
||||
return key
|
||||
|
||||
|
||||
def split_monolithic(
|
||||
state: dict[str, torch.Tensor],
|
||||
) -> dict[str, OrderedDict[str, torch.Tensor]]:
|
||||
def split_monolithic(state: dict[str, torch.Tensor], ) -> dict[str, OrderedDict[str, torch.Tensor]]:
|
||||
components: dict[str, OrderedDict[str, torch.Tensor]] = {
|
||||
name: OrderedDict() for name in set(COMPONENT_PREFIXES.values())
|
||||
name: OrderedDict()
|
||||
for name in set(COMPONENT_PREFIXES.values())
|
||||
}
|
||||
intentionally_skipped: list[str] = []
|
||||
unowned: list[str] = []
|
||||
@@ -117,10 +112,8 @@ def split_monolithic(
|
||||
unowned.append(key)
|
||||
if unowned:
|
||||
sample = ", ".join(unowned[:10])
|
||||
raise ValueError(
|
||||
f"Unowned monolithic keys: {len(unowned)}. "
|
||||
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}"
|
||||
)
|
||||
raise ValueError(f"Unowned monolithic keys: {len(unowned)}. "
|
||||
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}")
|
||||
if intentionally_skipped:
|
||||
print(f"Intentionally skipped {len(intentionally_skipped)} keys")
|
||||
return {name: weights for name, weights in components.items() if weights}
|
||||
@@ -143,8 +136,12 @@ def build_component_configs(_src_dir: Path) -> dict[str, dict[str, Any]]:
|
||||
# TODO: emit config content accepted by FastVideo loaders. Most components use
|
||||
# config.json; schedulers use scheduler_config.json.
|
||||
return {
|
||||
"transformer": {"_class_name": "<FastVideoTransformerClass>"},
|
||||
"vae": {"_class_name": "<FastVideoVAEClass>"},
|
||||
"transformer": {
|
||||
"_class_name": "<FastVideoTransformerClass>"
|
||||
},
|
||||
"vae": {
|
||||
"_class_name": "<FastVideoVAEClass>"
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -177,19 +174,13 @@ def build_model_index(
|
||||
}
|
||||
if revision:
|
||||
index["_fastvideo_converted_revision"] = revision
|
||||
return {
|
||||
key: value
|
||||
for key, value in index.items()
|
||||
if key.startswith("_") or key in available_components
|
||||
}
|
||||
return {key: value for key, value in index.items() if key.startswith("_") or key in available_components}
|
||||
|
||||
|
||||
def validate_component_configs(configs: dict[str, dict[str, Any]]) -> None:
|
||||
# TODO: instantiate each FastVideo config and call update_model_arch(...) or
|
||||
# update_model_config(...) with this JSON so unknown emitted keys fail here.
|
||||
placeholder_configs = [
|
||||
name for name, config in configs.items() if "<" in json.dumps(config)
|
||||
]
|
||||
placeholder_configs = [name for name, config in configs.items() if "<" in json.dumps(config)]
|
||||
if placeholder_configs:
|
||||
raise ValueError(f"Replace config placeholders for: {placeholder_configs}")
|
||||
|
||||
@@ -201,9 +192,7 @@ def verify_conversion(
|
||||
del dst_dir, components
|
||||
# TODO: load each emitted stateful component through its production loader and
|
||||
# assert strict load, or document exact allowed missing/unexpected keys.
|
||||
raise NotImplementedError(
|
||||
"Implement production config validation and strict-load checks"
|
||||
)
|
||||
raise NotImplementedError("Implement production config validation and strict-load checks")
|
||||
|
||||
|
||||
def write_component(
|
||||
@@ -216,9 +205,7 @@ def write_component(
|
||||
if component_dir.exists() and any(component_dir.iterdir()):
|
||||
shutil.rmtree(component_dir)
|
||||
component_dir.mkdir(parents=True, exist_ok=True)
|
||||
save_file(
|
||||
dict(state), str(component_dir / "diffusion_pytorch_model.safetensors")
|
||||
)
|
||||
save_file(dict(state), str(component_dir / "diffusion_pytorch_model.safetensors"))
|
||||
if config is not None:
|
||||
config_path = component_dir / config_filename(name)
|
||||
with config_path.open("w", encoding="utf-8") as f:
|
||||
@@ -261,9 +248,7 @@ def convert(
|
||||
|
||||
if layout in {"monolithic", "raw_official"}:
|
||||
# TODO: replace model.safetensors with the official monolithic file name.
|
||||
components = split_monolithic(
|
||||
load_checkpoint(default_monolithic_checkpoint(src_path))
|
||||
)
|
||||
components = split_monolithic(load_checkpoint(default_monolithic_checkpoint(src_path)))
|
||||
elif layout in {"separate_components", "mixed"}:
|
||||
if not src_path.is_dir():
|
||||
raise ValueError(f"{layout} layout requires a source directory: {src_path}")
|
||||
@@ -271,9 +256,7 @@ def convert(
|
||||
else:
|
||||
raise ValueError(f"Unsupported template layout: {layout}")
|
||||
|
||||
copied = (
|
||||
copy_passthrough(src_path, dst_dir) if src_path.is_dir() else []
|
||||
)
|
||||
copied = (copy_passthrough(src_path, dst_dir) if src_path.is_dir() else [])
|
||||
configs = build_component_configs(src_path if src_path.is_dir() else src_path.parent)
|
||||
validate_component_configs(configs)
|
||||
for name, state in components.items():
|
||||
@@ -289,9 +272,7 @@ def convert(
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--src", required=True, help="HF repo id, local dir, or checkpoint path"
|
||||
)
|
||||
parser.add_argument("--src", required=True, help="HF repo id, local dir, or checkpoint path")
|
||||
parser.add_argument("--revision", help="HF branch, tag, or commit for repo sources")
|
||||
parser.add_argument(
|
||||
"--dst",
|
||||
|
||||
@@ -24,8 +24,8 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
|
||||
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
|
||||
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
|
||||
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
|
||||
DRILL_LAYER_ENV: str = "<FAMILY>_DEBUG_DRILL_LAYER"
|
||||
HYPOTHESIS_ENV: str = "<FAMILY>_DEBUG_PATCH_<HYPOTHESIS>"
|
||||
REL_THRESHOLD: float = 0.005 # 0.5% abs_mean drift flags a block as divergent
|
||||
@@ -94,6 +94,7 @@ def _attach_block_hooks(
|
||||
handles: list[Any] = []
|
||||
|
||||
def _hook(name: str):
|
||||
|
||||
def fn(_module, _inputs, outputs):
|
||||
t = outputs[0] if isinstance(outputs, tuple) else outputs
|
||||
if not torch.is_tensor(t):
|
||||
@@ -101,6 +102,7 @@ def _attach_block_hooks(
|
||||
log.append({"side": label, **_stat(name, t)})
|
||||
if tensors is not None:
|
||||
tensors[name] = t.detach().float().cpu()
|
||||
|
||||
return fn
|
||||
|
||||
def _pre_hook(name: str):
|
||||
@@ -114,6 +116,7 @@ def _attach_block_hooks(
|
||||
log.append({"side": label, **_stat(key, t)})
|
||||
if tensors is not None:
|
||||
tensors[key] = t.detach().float().cpu()
|
||||
|
||||
return fn
|
||||
|
||||
# TODO: adapt attribute paths to your model. Remove adapter block if absent.
|
||||
@@ -131,43 +134,21 @@ def _attach_block_hooks(
|
||||
# magi-human uses: attention, mlp.pre_norm, mlp.up_gate_proj,
|
||||
# mlp.down_proj (pre+post), mlp, attn_post_norm, mlp_post_norm.
|
||||
if hasattr(layer, "attention"):
|
||||
handles.append(
|
||||
layer.attention.register_forward_hook(_hook(f"{tag}.attention"))
|
||||
)
|
||||
handles.append(layer.attention.register_forward_hook(_hook(f"{tag}.attention")))
|
||||
if hasattr(layer, "mlp"):
|
||||
mlp = layer.mlp
|
||||
if hasattr(mlp, "pre_norm"):
|
||||
handles.append(
|
||||
mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm"))
|
||||
)
|
||||
handles.append(mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm")))
|
||||
if hasattr(mlp, "up_gate_proj"):
|
||||
handles.append(
|
||||
mlp.up_gate_proj.register_forward_hook(
|
||||
_hook(f"{tag}.mlp.up_gate_proj")
|
||||
)
|
||||
)
|
||||
handles.append(mlp.up_gate_proj.register_forward_hook(_hook(f"{tag}.mlp.up_gate_proj")))
|
||||
if hasattr(mlp, "down_proj"):
|
||||
handles.append(
|
||||
mlp.down_proj.register_forward_pre_hook(
|
||||
_pre_hook(f"{tag}.mlp.down_proj")
|
||||
)
|
||||
)
|
||||
handles.append(
|
||||
mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj"))
|
||||
)
|
||||
handles.append(mlp.down_proj.register_forward_pre_hook(_pre_hook(f"{tag}.mlp.down_proj")))
|
||||
handles.append(mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj")))
|
||||
handles.append(mlp.register_forward_hook(_hook(f"{tag}.mlp")))
|
||||
if hasattr(layer, "attn_post_norm"):
|
||||
handles.append(
|
||||
layer.attn_post_norm.register_forward_hook(
|
||||
_hook(f"{tag}.attn_post_norm")
|
||||
)
|
||||
)
|
||||
handles.append(layer.attn_post_norm.register_forward_hook(_hook(f"{tag}.attn_post_norm")))
|
||||
if hasattr(layer, "mlp_post_norm"):
|
||||
handles.append(
|
||||
layer.mlp_post_norm.register_forward_hook(
|
||||
_hook(f"{tag}.mlp_post_norm")
|
||||
)
|
||||
)
|
||||
handles.append(layer.mlp_post_norm.register_forward_hook(_hook(f"{tag}.mlp_post_norm")))
|
||||
return handles
|
||||
|
||||
|
||||
@@ -193,11 +174,9 @@ def _write_log(entries: list[dict], path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(path, "w") as f:
|
||||
for e in entries:
|
||||
f.write(
|
||||
f"{e['name']} {e['shape']} "
|
||||
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
|
||||
f"{e['min']:.6f} {e['max']:.6f}\n"
|
||||
)
|
||||
f.write(f"{e['name']} {e['shape']} "
|
||||
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
|
||||
f"{e['min']:.6f} {e['max']:.6f}\n")
|
||||
|
||||
|
||||
def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
@@ -205,9 +184,14 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
return (0, "")
|
||||
if name.startswith(f"L{drill_layer:02d}."):
|
||||
sub_order = {
|
||||
"attention": 0, "attn_post_norm": 1, "mlp.pre_norm": 2,
|
||||
"mlp.up_gate_proj": 3, "mlp.down_proj<in>": 4,
|
||||
"mlp.down_proj": 5, "mlp": 6, "mlp_post_norm": 7,
|
||||
"attention": 0,
|
||||
"attn_post_norm": 1,
|
||||
"mlp.pre_norm": 2,
|
||||
"mlp.up_gate_proj": 3,
|
||||
"mlp.down_proj<in>": 4,
|
||||
"mlp.down_proj": 5,
|
||||
"mlp": 6,
|
||||
"mlp_post_norm": 7,
|
||||
}.get(name.split(".", 1)[1], 9)
|
||||
return (1, f"block[{drill_layer:02d}]", sub_order)
|
||||
if name.startswith("block["):
|
||||
@@ -216,10 +200,8 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
|
||||
|
||||
def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
|
||||
hdr = (
|
||||
f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
|
||||
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}"
|
||||
)
|
||||
hdr = (f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
|
||||
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}")
|
||||
print(f"\n{hdr}\n{'-' * len(hdr)}")
|
||||
first_div: int | None = None
|
||||
for name in sorted(by_name.keys(), key=lambda n: _sort_key(n, drill_layer)):
|
||||
@@ -235,11 +217,9 @@ def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
|
||||
flag = " <<< DIVERGE"
|
||||
if first_div is None:
|
||||
first_div = int(name[len("block["):-1])
|
||||
print(
|
||||
f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
|
||||
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
|
||||
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}"
|
||||
)
|
||||
print(f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
|
||||
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
|
||||
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}")
|
||||
return first_div
|
||||
|
||||
|
||||
@@ -255,10 +235,8 @@ def _print_elementwise(up_t: dict[str, torch.Tensor], fv_t: dict[str, torch.Tens
|
||||
continue
|
||||
diff = (a - b).abs()
|
||||
rel = (diff.mean().item() / max(a.abs().mean().item(), 1e-9)) * 100
|
||||
print(
|
||||
f"{name:<30} {str(tuple(a.shape)):<22} "
|
||||
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%"
|
||||
)
|
||||
print(f"{name:<30} {str(tuple(a.shape)):<22} "
|
||||
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
|
||||
@@ -43,12 +43,10 @@ def _add_official_to_path() -> Path:
|
||||
|
||||
def _log_tensor_stats(label: str, tensor: torch.Tensor) -> None:
|
||||
value = tensor.detach().float()
|
||||
print(
|
||||
f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
|
||||
f"dtype={tensor.dtype} device={tensor.device} "
|
||||
f"min={value.min().item():.6f} max={value.max().item():.6f} "
|
||||
f"mean={value.mean().item():.6f} std={value.std().item():.6f}"
|
||||
)
|
||||
print(f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
|
||||
f"dtype={tensor.dtype} device={tensor.device} "
|
||||
f"min={value.min().item():.6f} max={value.max().item():.6f} "
|
||||
f"mean={value.mean().item():.6f} std={value.std().item():.6f}")
|
||||
|
||||
|
||||
def _extract_tensor(output: Any, key: str) -> torch.Tensor:
|
||||
@@ -73,10 +71,8 @@ def _run_official_pipeline(
|
||||
device: torch.device,
|
||||
) -> Any:
|
||||
del official_path, params, device
|
||||
pytest.skip(
|
||||
"TODO: import the official pipeline/factory, load official weights, "
|
||||
"run with params, and return the comparison target."
|
||||
)
|
||||
pytest.skip("TODO: import the official pipeline/factory, load official weights, "
|
||||
"run with params, and return the comparison target.")
|
||||
|
||||
|
||||
def _run_fastvideo_pipeline(model_path: Path, params: dict[str, Any]) -> Any:
|
||||
@@ -146,8 +142,6 @@ def test_todo_model_family_pipeline_official_parity() -> None:
|
||||
assert official_tensor.shape == fastvideo_tensor.shape
|
||||
|
||||
diff = (official_tensor - fastvideo_tensor).abs()
|
||||
print(
|
||||
f"diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
|
||||
)
|
||||
print(f"diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}")
|
||||
assert_close(fastvideo_tensor, official_tensor, atol=1e-2, rtol=1e-2)
|
||||
|
||||
+15
-2
@@ -1,6 +1,8 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
BUILDKITE_CLEAN_CHECKOUT: true
|
||||
# Buildkite only launches Modal; remote jobs initialize their own submodules.
|
||||
BUILDKITE_GIT_SUBMODULES: false
|
||||
|
||||
notify:
|
||||
- github_commit_status:
|
||||
@@ -66,6 +68,17 @@ 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"
|
||||
@@ -402,7 +415,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 25m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training
|
||||
@@ -413,7 +426,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 25m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: Distillation DMD Tests"
|
||||
env:
|
||||
- TEST_TYPE=distillation_dmd
|
||||
|
||||
@@ -187,6 +187,10 @@ 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)
|
||||
|
||||
@@ -27,14 +27,25 @@ jobs:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
# For PR events, lint the PR head — but keep the hook definitions from
|
||||
# the base branch so an untrusted PR cannot alter what gets executed.
|
||||
- name: Save trusted hook config
|
||||
# The gate scripts are saved too: the self-test step below executes them,
|
||||
# so it must run the base-branch copies, not the PR head's.
|
||||
- name: Save trusted hook config and gate scripts
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
- uses: actions/checkout@v4
|
||||
run: |
|
||||
cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
cp -a .github/scripts "$RUNNER_TEMP/trusted-scripts"
|
||||
echo "GATE_SCRIPTS_DIR=$RUNNER_TEMP/trusted-scripts" >> "$GITHUB_ENV"
|
||||
# allow-unsafe-pr-checkout acknowledges checkout's pull_request_target
|
||||
# guard: the head is data for the trusted hooks to lint; nothing from it
|
||||
# is executed (config and gate scripts are pinned to the base branch
|
||||
# above) and credentials are not persisted. SHA-pinned to v4.4.0 because
|
||||
# actionlint's action schema does not know the new input yet.
|
||||
- uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0
|
||||
if: github.event_name == 'pull_request_target'
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha }}
|
||||
persist-credentials: false
|
||||
allow-unsafe-pr-checkout: true
|
||||
- name: Restore trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
|
||||
@@ -48,5 +59,7 @@ jobs:
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
# After pre-commit so a self-test failure cannot mask lint failures.
|
||||
# GATE_SCRIPTS_DIR points at the base-branch copy on fork PRs (set above);
|
||||
# push / workflow_call runs use the checked-out tree directly.
|
||||
- name: Full-suite gate self-test
|
||||
run: bash .github/scripts/test_gate_full_suite.sh
|
||||
run: bash "${GATE_SCRIPTS_DIR:-.github/scripts}/test_gate_full_suite.sh"
|
||||
|
||||
@@ -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 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 golden-gate 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 [training]=training
|
||||
[ssim]=ssim [golden-gate]=golden_gate [training]=training
|
||||
[lora-inference]=inference_lora [lora-training]=training_lora
|
||||
[lora-extraction]=lora_extraction
|
||||
[distillation]=distillation_dmd [self-forcing]=self_forcing
|
||||
|
||||
@@ -13,16 +13,23 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
# 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.
|
||||
# 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.
|
||||
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:
|
||||
@@ -50,7 +57,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 a push that changed docker/Dockerfile (inputs are null on push). The
|
||||
# on an in-scope main push (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' }}
|
||||
|
||||
@@ -11,9 +11,7 @@ on:
|
||||
- 'requirements-mkdocs.txt'
|
||||
- 'scripts/check_docs_links.py'
|
||||
- '.github/workflows/infra-docs.yml'
|
||||
# Run the trusted base-branch workflow so fork PRs can be skipped without
|
||||
# waiting for maintainer approval.
|
||||
pull_request_target:
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
@@ -26,19 +24,21 @@ on:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pages: write
|
||||
id-token: write
|
||||
|
||||
concurrency:
|
||||
group: "pages"
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
build:
|
||||
# MkDocs executes repository code; only trusted same-repository PRs run it.
|
||||
if: github.event_name == 'push' || github.event.pull_request.head.repo.full_name == github.repository
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha || github.sha }}
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Setup Python
|
||||
uses: actions/setup-python@v5
|
||||
@@ -52,7 +52,6 @@ jobs:
|
||||
run: uv pip install --system -r requirements-mkdocs.txt
|
||||
|
||||
- name: Setup Pages
|
||||
if: github.event_name == 'push'
|
||||
uses: actions/configure-pages@v4
|
||||
|
||||
- name: Build documentation
|
||||
@@ -62,22 +61,17 @@ jobs:
|
||||
run: python scripts/check_docs_links.py
|
||||
|
||||
- name: Upload artifact
|
||||
if: github.event_name == 'push'
|
||||
uses: actions/upload-pages-artifact@v3
|
||||
with:
|
||||
path: ./site
|
||||
|
||||
deploy:
|
||||
permissions:
|
||||
pages: write
|
||||
id-token: write
|
||||
concurrency: pages
|
||||
environment:
|
||||
name: github-pages
|
||||
url: ${{ steps.deployment.outputs.page_url }}
|
||||
runs-on: ubuntu-latest
|
||||
needs: build
|
||||
if: github.event_name == 'push'
|
||||
if: github.ref == 'refs/heads/main'
|
||||
steps:
|
||||
- name: Deploy to GitHub Pages
|
||||
id: deployment
|
||||
|
||||
@@ -6,6 +6,7 @@ results/
|
||||
wandb/
|
||||
*.ipynb
|
||||
*.jpg
|
||||
!examples/datasets/lingbotworld2/image.jpg
|
||||
*.safetensors
|
||||
*.mp4
|
||||
*.png
|
||||
@@ -22,6 +23,7 @@ Miniconda3-latest-Linux-x86_64.sh
|
||||
*validation/
|
||||
data/
|
||||
outputs/
|
||||
outputs_audio/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
@@ -34,6 +36,7 @@ env
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
/Z-Image/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
@@ -73,6 +76,7 @@ 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
|
||||
|
||||
@@ -22,6 +22,7 @@ repos:
|
||||
hooks:
|
||||
- id: yapf
|
||||
args: [--in-place, --verbose]
|
||||
language_version: python3.12
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.12
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
|
||||
|
||||
## NEWS
|
||||
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), check out the [Blog](https://haoailab.com/blogs/fastwan-qad/).
|
||||
- `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/).
|
||||
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py).
|
||||
@@ -33,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 full list of supported models and recipes.
|
||||
- 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.
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- Sequence Parallelism for distributed inference
|
||||
- Multiple state-of-the-art attention backends
|
||||
|
||||
@@ -3,7 +3,6 @@ from __future__ import annotations
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
TESTS_DIR = Path(__file__).resolve().parent
|
||||
DREAMVERSE_PACKAGE_DIR = TESTS_DIR.parent
|
||||
DREAMVERSE_APP_DIR = DREAMVERSE_PACKAGE_DIR.parent
|
||||
|
||||
@@ -5,7 +5,6 @@ from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
SERVER_DIR = Path(__file__).resolve().parents[1]
|
||||
|
||||
|
||||
@@ -53,9 +52,7 @@ def test_config_defaults_to_cerebras_with_parallel_groq_fallback_stage(monkeypat
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.PROMPT_PROVIDER == "cerebras"
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
|
||||
("cerebras", "groq"),
|
||||
)
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
|
||||
assert module.PROMPT_PROVIDER_PRIORITY == (
|
||||
"cerebras",
|
||||
"groq",
|
||||
@@ -86,9 +83,7 @@ def test_config_ignores_legacy_groq_primary_override(monkeypatch):
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.PROMPT_PROVIDER == "cerebras"
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
|
||||
("cerebras", "groq"),
|
||||
)
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
|
||||
assert module.PROMPT_PROVIDER_PRIORITY == (
|
||||
"cerebras",
|
||||
"groq",
|
||||
@@ -106,24 +101,17 @@ def test_config_uses_local_overlay_paths_when_devtools_enabled(monkeypatch, tmp_
|
||||
|
||||
assert module.DEVTOOLS_ENABLED is True
|
||||
assert module.FRONTEND_ROOT.as_posix().endswith("apps/dreamverse/web")
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith(
|
||||
"dreamverse/prompts.local/next_segment_system_prompt.md"
|
||||
)
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith("dreamverse/prompts.local/next_segment_system_prompt.md")
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
|
||||
"dreamverse/prompts/next_segment_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts/next_segment_system_prompt.md")
|
||||
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_PATH.endswith(
|
||||
"dreamverse/prompts.local/rewrite_user_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts.local/rewrite_user_system_prompt.md")
|
||||
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
|
||||
"dreamverse/prompts/rewrite_user_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts/rewrite_user_system_prompt.md")
|
||||
assert module.CURATED_PRESETS_FILE_PATH.endswith(
|
||||
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json")
|
||||
assert module.CURATED_PRESETS_FALLBACK_FILE_PATH.endswith(
|
||||
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json")
|
||||
assert module.FRONTEND_STATIC_DIR_CANDIDATES[:2] == (
|
||||
str(module.FRONTEND_ROOT / "out"),
|
||||
str(module.FRONTEND_ROOT / "dist"),
|
||||
|
||||
@@ -9,6 +9,7 @@ from fastapi.testclient import TestClient
|
||||
import fastvideo.entrypoints.streaming as streaming_entrypoints
|
||||
import pytest
|
||||
|
||||
|
||||
def _install_stack03_import_stubs(monkeypatch):
|
||||
"""Keep entrypoint tests focused while later-stack runtime modules are absent."""
|
||||
if not hasattr(streaming_entrypoints, "build_health_router"):
|
||||
@@ -17,6 +18,7 @@ def _install_stack03_import_stubs(monkeypatch):
|
||||
gpu_pool_stub = types.ModuleType("dreamverse.gpu_pool")
|
||||
|
||||
class GPUPool:
|
||||
|
||||
def __init__(self, _gpu_ids):
|
||||
pass
|
||||
|
||||
@@ -49,6 +51,7 @@ def _install_stack03_import_stubs(monkeypatch):
|
||||
controller_stub = types.ModuleType("dreamverse.session.controller")
|
||||
|
||||
class SessionController:
|
||||
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
@@ -76,13 +79,11 @@ def _run_cli(module, monkeypatch, argv: list[str]) -> list[dict[str, object]]:
|
||||
uvicorn_stub = types.ModuleType("uvicorn")
|
||||
|
||||
def run(app, host: str, port: int) -> None:
|
||||
calls.append(
|
||||
{
|
||||
"app": app,
|
||||
"host": host,
|
||||
"port": port,
|
||||
}
|
||||
)
|
||||
calls.append({
|
||||
"app": app,
|
||||
"host": host,
|
||||
"port": port,
|
||||
})
|
||||
|
||||
uvicorn_stub.run = run
|
||||
monkeypatch.setitem(sys.modules, "uvicorn", uvicorn_stub)
|
||||
@@ -99,13 +100,11 @@ def test_server_cli_defaults_to_local_web_port(monkeypatch):
|
||||
server_main = _import_server_main(monkeypatch)
|
||||
calls = _run_cli(server_main, monkeypatch, ["dreamverse-server"])
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": server_main.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": server_main.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}]
|
||||
|
||||
|
||||
def test_server_cli_allows_explicit_host_and_port(monkeypatch):
|
||||
@@ -116,13 +115,11 @@ def test_server_cli_allows_explicit_host_and_port(monkeypatch):
|
||||
["dreamverse-server", "--host", "127.0.0.1", "--port", "8123"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": server_main.app,
|
||||
"host": "127.0.0.1",
|
||||
"port": 8123,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": server_main.app,
|
||||
"host": "127.0.0.1",
|
||||
"port": 8123,
|
||||
}]
|
||||
|
||||
|
||||
def test_server_does_not_expose_backend_source_as_static_assets(monkeypatch):
|
||||
@@ -142,13 +139,11 @@ def test_mock_server_cli_defaults_to_local_web_port(monkeypatch):
|
||||
["dreamverse-mock-server"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}]
|
||||
|
||||
|
||||
def test_mock_server_cli_updates_latency(monkeypatch):
|
||||
@@ -161,13 +156,11 @@ def test_mock_server_cli_updates_latency(monkeypatch):
|
||||
["dreamverse-mock-server", "--latency", "321", "--port", "8111"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8111,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8111,
|
||||
}]
|
||||
assert mock_server.LATENCY_MS == 321
|
||||
finally:
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
|
||||
@@ -7,7 +7,6 @@ from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
import dreamverse.gpu_pool as gpu_pool
|
||||
|
||||
|
||||
@@ -85,9 +84,7 @@ def test_send_command_raises_on_worker_death():
|
||||
cmd_q = ctx.Queue()
|
||||
resp_q = ctx.Queue()
|
||||
|
||||
proc = ctx.Process(
|
||||
target=_child_consume_and_exit, args=(cmd_q, resp_q)
|
||||
)
|
||||
proc = ctx.Process(target=_child_consume_and_exit, args=(cmd_q, resp_q))
|
||||
proc.start()
|
||||
|
||||
# Wait for the spawn child to fully boot. Allow generous time —
|
||||
|
||||
@@ -7,7 +7,7 @@ ALLOWED_PREFIXES = (
|
||||
"fastvideo.entrypoints.video_generator",
|
||||
"fastvideo.configs",
|
||||
)
|
||||
ALLOWED_EXACT = ("fastvideo",)
|
||||
ALLOWED_EXACT = ("fastvideo", )
|
||||
FORBIDDEN_PREFIXES = (
|
||||
"fastvideo.pipelines",
|
||||
"fastvideo.models",
|
||||
@@ -38,19 +38,13 @@ def test_dreamverse_server_imports_only_public_fastvideo_surfaces() -> None:
|
||||
except SyntaxError as task_exc:
|
||||
raise AssertionError(f"Failed to parse {path}") from task_exc
|
||||
for node in ast.walk(tree):
|
||||
names = (
|
||||
[a.name for a in node.names] if isinstance(node, ast.Import)
|
||||
else [node.module] if isinstance(node, ast.ImportFrom) and node.module
|
||||
else []
|
||||
)
|
||||
names = ([a.name for a in node.names] if isinstance(node, ast.Import) else
|
||||
[node.module] if isinstance(node, ast.ImportFrom) and node.module else [])
|
||||
for name in names:
|
||||
if not name:
|
||||
continue
|
||||
rel_path = str(path.relative_to(root))
|
||||
if (
|
||||
name.startswith(FORBIDDEN_PREFIXES)
|
||||
and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS
|
||||
):
|
||||
if (name.startswith(FORBIDDEN_PREFIXES) and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS):
|
||||
bad.append((str(path.relative_to(root)), getattr(node, "lineno", 0), name))
|
||||
|
||||
assert bad == [], f"Forbidden internal imports: {bad}"
|
||||
|
||||
@@ -6,7 +6,6 @@ import os
|
||||
|
||||
from fastapi import WebSocketDisconnect
|
||||
|
||||
|
||||
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
|
||||
os.environ.setdefault("GROQ_API_KEY", "dummy")
|
||||
|
||||
@@ -14,6 +13,7 @@ import dreamverse.mock_server as mock_server
|
||||
|
||||
|
||||
class _FakeWebSocket:
|
||||
|
||||
def __init__(self, messages: list[tuple[float, dict[str, object]]]):
|
||||
self._messages = messages
|
||||
self._index = 0
|
||||
@@ -49,34 +49,34 @@ def test_mock_server_matches_current_single5s_protocol():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "simple_prompt_1",
|
||||
"curated_prompts": ["selected prompt"],
|
||||
"single_clip_mode": True,
|
||||
"enhancement_enabled": False,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.01,
|
||||
{
|
||||
"type": "simple_generate",
|
||||
"preset_id": "simple_custom_prompt",
|
||||
"prompt_id": "simple_custom_prompt",
|
||||
"prompt": "custom prompt",
|
||||
"enhancement_enabled": True,
|
||||
"initial_image": None,
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "simple_prompt_1",
|
||||
"curated_prompts": ["selected prompt"],
|
||||
"single_clip_mode": True,
|
||||
"enhancement_enabled": False,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.01,
|
||||
{
|
||||
"type": "simple_generate",
|
||||
"preset_id": "simple_custom_prompt",
|
||||
"prompt_id": "simple_custom_prompt",
|
||||
"prompt": "custom prompt",
|
||||
"enhancement_enabled": True,
|
||||
"initial_image": None,
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -92,24 +92,14 @@ def test_mock_server_matches_current_single5s_protocol():
|
||||
assert message_types.count("ltx2_stream_complete") == 2
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
gpu_assigned_event = next(
|
||||
payload for payload in ws.sent_json if payload["type"] == "gpu_assigned"
|
||||
)
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
gpu_assigned_event = next(payload for payload in ws.sent_json if payload["type"] == "gpu_assigned")
|
||||
assert gpu_assigned_event["session_timeout"] == mock_server.SESSION_TIMEOUT_SECONDS
|
||||
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
|
||||
assert segment_start_events[0]["prompt"] == "selected prompt"
|
||||
assert segment_start_events[1]["prompt"] == "custom prompt"
|
||||
|
||||
step_complete_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "step_complete"
|
||||
]
|
||||
step_complete_events = [payload for payload in ws.sent_json if payload["type"] == "step_complete"]
|
||||
assert len(step_complete_events) == 2
|
||||
assert step_complete_events[0]["latency_ms"] == {
|
||||
"total": 121.0,
|
||||
@@ -134,29 +124,29 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
|
||||
mock_server.LATENCY_MS = 1
|
||||
mock_server.GENERATION_SEGMENT_CAP = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "start a new rollout",
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "start a new rollout",
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -166,11 +156,7 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
|
||||
assert "generation_cap_reached" not in message_types
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
|
||||
assert segment_start_events[0]["prompt"] == "segment one"
|
||||
assert segment_start_events[1]["prompt"] == "segment one [start a new rollout]"
|
||||
@@ -187,54 +173,40 @@ def test_mock_server_rewrite_during_active_segment_restarts_from_first_rewritten
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 100
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one", "segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "restart from rewrite",
|
||||
},
|
||||
),
|
||||
(0.40, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one", "segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "restart from rewrite",
|
||||
},
|
||||
),
|
||||
(0.40, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
|
||||
"segment one",
|
||||
"segment one [restart from rewrite]",
|
||||
]
|
||||
assert all(
|
||||
payload["prompt"] != "segment two"
|
||||
for payload in segment_start_events[1:]
|
||||
)
|
||||
reset_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload.get("type") == "seed_prompts_reset_applied"
|
||||
]
|
||||
assert any(
|
||||
payload.get("reason") == "rewrite_during_generation"
|
||||
for payload in reset_events
|
||||
)
|
||||
assert all(payload["prompt"] != "segment two" for payload in segment_start_events[1:])
|
||||
reset_events = [payload for payload in ws.sent_json if payload.get("type") == "seed_prompts_reset_applied"]
|
||||
assert any(payload.get("reason") == "rewrite_during_generation" for payload in reset_events)
|
||||
finally:
|
||||
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
@@ -247,24 +219,24 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "custom_editable",
|
||||
"preset_label": "Custom rollout",
|
||||
"curated_prompts": [],
|
||||
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "custom_editable",
|
||||
"preset_label": "Custom rollout",
|
||||
"curated_prompts": [],
|
||||
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -276,15 +248,9 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
|
||||
assert "ltx2_stream_start" in message_types
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert segment_start_events
|
||||
assert segment_start_events[0]["prompt"] == (
|
||||
"A moonbase corridor thriller with flooding [segment 1]"
|
||||
)
|
||||
assert segment_start_events[0]["prompt"] == ("A moonbase corridor thriller with flooding [segment 1]")
|
||||
finally:
|
||||
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
@@ -297,35 +263,37 @@ def test_mock_server_can_start_new_project_without_reconnecting():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 40
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.02, {"type": "end_project_keep_session"}),
|
||||
(
|
||||
0.20,
|
||||
{
|
||||
"type": "project_init_v1",
|
||||
"preset_id": "test_preset_2",
|
||||
"preset_label": "Test Preset 2",
|
||||
"curated_prompts": ["segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.40, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.02, {
|
||||
"type": "end_project_keep_session"
|
||||
}),
|
||||
(
|
||||
0.20,
|
||||
{
|
||||
"type": "project_init_v1",
|
||||
"preset_id": "test_preset_2",
|
||||
"preset_label": "Test Preset 2",
|
||||
"curated_prompts": ["segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.40, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -336,16 +304,11 @@ def test_mock_server_can_start_new_project_without_reconnecting():
|
||||
|
||||
project_idle_index = message_types.index("project_idle")
|
||||
stream_start_indexes = [
|
||||
index for index, message_type in enumerate(message_types)
|
||||
if message_type == "ltx2_stream_start"
|
||||
index for index, message_type in enumerate(message_types) if message_type == "ltx2_stream_start"
|
||||
]
|
||||
assert stream_start_indexes[0] < project_idle_index < stream_start_indexes[1]
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
|
||||
"segment one",
|
||||
"segment two",
|
||||
|
||||
@@ -6,7 +6,6 @@ import os
|
||||
import re
|
||||
import time
|
||||
|
||||
|
||||
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
|
||||
os.environ.setdefault("GROQ_API_KEY", "dummy")
|
||||
|
||||
@@ -22,6 +21,7 @@ from dreamverse.prompt_enhancer import (
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self._payload = payload
|
||||
|
||||
@@ -30,6 +30,7 @@ class _FakeResponse:
|
||||
|
||||
|
||||
class _FakeSyncCompletions:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self._payload = payload
|
||||
|
||||
@@ -38,6 +39,7 @@ class _FakeSyncCompletions:
|
||||
|
||||
|
||||
class _FakeSyncClient:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self.chat = type(
|
||||
"_FakeChat",
|
||||
@@ -47,6 +49,7 @@ class _FakeSyncClient:
|
||||
|
||||
|
||||
class _DelayedSyncCompletions:
|
||||
|
||||
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
|
||||
self._payload = payload
|
||||
self._delay_s = delay_s
|
||||
@@ -61,29 +64,26 @@ class _DelayedSyncCompletions:
|
||||
|
||||
|
||||
class _DelayedSyncClient:
|
||||
|
||||
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
|
||||
self.chat = type(
|
||||
"_FakeChat",
|
||||
(),
|
||||
{
|
||||
"completions": _DelayedSyncCompletions(
|
||||
payload,
|
||||
delay_s=delay_s,
|
||||
exc=exc,
|
||||
)
|
||||
},
|
||||
{"completions": _DelayedSyncCompletions(
|
||||
payload,
|
||||
delay_s=delay_s,
|
||||
exc=exc,
|
||||
)},
|
||||
)()
|
||||
|
||||
|
||||
def _chat_payload_with_content(content: str) -> dict:
|
||||
return {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"content": content,
|
||||
}
|
||||
"choices": [{
|
||||
"message": {
|
||||
"content": content,
|
||||
}
|
||||
]
|
||||
}]
|
||||
}
|
||||
|
||||
|
||||
@@ -172,6 +172,7 @@ def _build_staged_enhancer(
|
||||
|
||||
|
||||
class _FakeOpenAIClient:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
self.chat = type(
|
||||
@@ -182,6 +183,7 @@ class _FakeOpenAIClient:
|
||||
|
||||
|
||||
class _FakeCerebrasClient:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
self.chat = type(
|
||||
@@ -192,16 +194,12 @@ class _FakeCerebrasClient:
|
||||
|
||||
|
||||
def test_parse_json_response_accepts_fenced_json_with_prose():
|
||||
parsed = _parse_json_response(
|
||||
"Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks."
|
||||
)
|
||||
parsed = _parse_json_response("Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks.")
|
||||
assert parsed == {"segment_prompts": ["A", "B"]}
|
||||
|
||||
|
||||
def test_parse_json_response_extracts_first_embedded_object():
|
||||
parsed = _parse_json_response(
|
||||
"Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)"
|
||||
)
|
||||
parsed = _parse_json_response("Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)")
|
||||
assert parsed == {"segment_prompts": ["A", "B"]}
|
||||
|
||||
|
||||
@@ -268,16 +266,12 @@ def test_build_client_supports_groq_provider(monkeypatch):
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "preset_a"
|
||||
@@ -286,15 +280,12 @@ def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"rewritten_prompts":["A","B"]}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"rewritten_prompts":["A","B"]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "current_rollout"
|
||||
@@ -303,19 +294,14 @@ def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_segment_dicts_without_top_level_rollout_metadata():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"segments":[{"prompt":"A"},{"text":"B"}]}'
|
||||
)
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segments":[{"prompt":"A"},{"text":"B"}]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
preset_id="preset_a",
|
||||
preset_label="Preset A",
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "preset_a"
|
||||
@@ -329,14 +315,12 @@ def test_rewrite_prompt_sequence_accepts_numbered_prose_output():
|
||||
"The user is asking for a cinematic rewrite.\n\n"
|
||||
"1. A dog bounds across the moon's dusty surface, kicking up silver regolith as it chases a rabbit beneath the black sky.\n"
|
||||
"2. The rabbit darts around a crater rim while the dog lunges after it, Earth glowing blue in the distance.\n"
|
||||
)
|
||||
)
|
||||
))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "current_rollout"
|
||||
@@ -355,12 +339,10 @@ def test_enhance_prompt_prefers_cerebras_before_groq_fallback():
|
||||
groq_delay_s=0.01,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -381,12 +363,10 @@ def test_enhance_prompt_uses_groq_when_cerebras_fails():
|
||||
groq_delay_s=0.01,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -408,13 +388,11 @@ def test_enhance_prompt_can_use_groq_when_cerebras_times_out():
|
||||
enhancer.http_timeout_ms = 50
|
||||
enhancer.default_timeout_ms = 50
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
timeout_ms=50,
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
timeout_ms=50,
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -434,12 +412,10 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
|
||||
groq_delay_s=0.08,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -453,15 +429,12 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("I cannot comply with JSON right now.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("I cannot comply with JSON right now."))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
assert result.raw_response_text == "I cannot comply with JSON right now."
|
||||
@@ -473,9 +446,7 @@ def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
|
||||
def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'))
|
||||
captured = {
|
||||
"body": None,
|
||||
"timeout_seconds": None,
|
||||
@@ -486,8 +457,7 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
captured["timeout_seconds"] = timeout_seconds
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
|
||||
),
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'),
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}',
|
||||
)
|
||||
|
||||
@@ -502,8 +472,7 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
rewrite_model="gpt-test",
|
||||
rewrite_temperature=0.2,
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert captured["body"]["messages"][0] == {
|
||||
@@ -512,12 +481,12 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
}
|
||||
assert captured["body"]["messages"][1]["role"] == "user"
|
||||
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
|
||||
"mode": "edit_existing_rollout",
|
||||
"request": (
|
||||
"Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."
|
||||
),
|
||||
"user_instruction": "make it cinematic",
|
||||
"mode":
|
||||
"edit_existing_rollout",
|
||||
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."),
|
||||
"user_instruction":
|
||||
"make it cinematic",
|
||||
"current_rollout": {
|
||||
"id": "preset_a",
|
||||
"label": "Preset A",
|
||||
@@ -528,11 +497,8 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
|
||||
def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'))
|
||||
captured = {
|
||||
"body": None,
|
||||
}
|
||||
@@ -541,10 +507,8 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
del timeout_seconds
|
||||
captured["body"] = body
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'
|
||||
),
|
||||
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'),
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}',
|
||||
)
|
||||
@@ -560,30 +524,29 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
rewrite_model="gpt-test",
|
||||
rewrite_temperature=0.2,
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.prompts == ["A", "B", "C", "D", "E", "F"]
|
||||
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
|
||||
"mode": "new_rollout",
|
||||
"request": (
|
||||
"Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."
|
||||
),
|
||||
"user_instruction": "A moonbase corridor thriller with flooding and red alarms",
|
||||
"desired_segment_count": 6,
|
||||
"rollout_id_hint": "custom_editable",
|
||||
"rollout_label_hint": "Custom rollout",
|
||||
"mode":
|
||||
"new_rollout",
|
||||
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."),
|
||||
"user_instruction":
|
||||
"A moonbase corridor thriller with flooding and red alarms",
|
||||
"desired_segment_count":
|
||||
6,
|
||||
"rollout_id_hint":
|
||||
"custom_editable",
|
||||
"rollout_label_hint":
|
||||
"Custom rollout",
|
||||
}
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared system prompt"
|
||||
captured = {
|
||||
"body": None,
|
||||
@@ -593,9 +556,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
del timeout_seconds
|
||||
captured["body"] = body
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
),
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'),
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}',
|
||||
)
|
||||
|
||||
@@ -609,8 +570,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
rewrite_instruction="make it cinematic",
|
||||
rewrite_model="gpt-test",
|
||||
system_prompt_override="session specific system prompt",
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert captured["body"]["messages"][0] == {
|
||||
@@ -621,10 +581,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
|
||||
def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
|
||||
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
|
||||
|
||||
@@ -635,24 +592,17 @@ def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
|
||||
|
||||
def test_resolve_rewrite_new_rollout_system_prompt_prefers_override():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
|
||||
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
|
||||
|
||||
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt(
|
||||
"session specific system prompt"
|
||||
)
|
||||
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt("session specific system prompt")
|
||||
|
||||
assert resolved == "session specific system prompt"
|
||||
|
||||
|
||||
def test_generate_auto_prompt_uses_selected_model():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"next_prompt":"Auto next"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Auto next"}'))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -680,8 +630,7 @@ def test_generate_auto_prompt_uses_selected_model():
|
||||
next_segment_idx=2,
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Auto next"
|
||||
@@ -690,9 +639,7 @@ def test_generate_auto_prompt_uses_selected_model():
|
||||
|
||||
|
||||
def test_enhance_prompt_uses_selected_model():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"next_prompt":"Enhanced next"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Enhanced next"}'))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
@@ -722,8 +669,7 @@ def test_enhance_prompt_uses_selected_model():
|
||||
next_segment_idx=2,
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Enhanced next"
|
||||
@@ -732,9 +678,7 @@ def test_enhance_prompt_uses_selected_model():
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"prompt":"Extended single clip"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"prompt":"Extended single clip"}'))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
@@ -764,14 +708,12 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
|
||||
|
||||
enhancer._request_content = _fake_request_content # type: ignore[attr-defined]
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"short 5s idea",
|
||||
mode="single_clip",
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"short 5s idea",
|
||||
mode="single_clip",
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Extended single clip"
|
||||
@@ -784,17 +726,15 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
|
||||
"single 5-second LTX-2.3 video clip. Respond with "
|
||||
'valid JSON only as {"prompt": "..."}.' # noqa: E501
|
||||
),
|
||||
"user_prompt": "short 5s idea",
|
||||
"user_prompt":
|
||||
"short 5s idea",
|
||||
}
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_rejects_plain_text_response():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
"Medium shot of a woman by a rainy cafe window as she lifts her "
|
||||
"phone, exhales softly, and the camera makes a slow push in."
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content("Medium shot of a woman by a rainy cafe window as she lifts her "
|
||||
"phone, exhales softly, and the camera makes a slow push in."))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
|
||||
result = asyncio.run(
|
||||
@@ -803,17 +743,14 @@ def test_enhance_prompt_single_clip_rejects_plain_text_response():
|
||||
mode="single_clip",
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No JSON object found in assistant response." in result.error
|
||||
assert result.prompt == ""
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"segment_prompts":["A","B"]}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segment_prompts":["A","B"]}'))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -824,17 +761,14 @@ def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
|
||||
mode="single_clip",
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "Missing prompt string." in (result.error or "")
|
||||
|
||||
|
||||
def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("A cinematic continuation with slow dolly movement.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("A cinematic continuation with slow dolly movement."))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -846,17 +780,14 @@ def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
next_segment_idx=2,
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
|
||||
|
||||
def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("A calm, grounded continuation with subtle motion.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("A calm, grounded continuation with subtle motion."))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -867,34 +798,30 @@ def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
next_segment_idx=2,
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
|
||||
enhancer = _build_test_enhancer(
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "length",
|
||||
"message": {
|
||||
"content": [],
|
||||
"refusal": None,
|
||||
},
|
||||
}
|
||||
],
|
||||
"usage": {"completion_tokens": 0},
|
||||
}
|
||||
)
|
||||
enhancer = _build_test_enhancer({
|
||||
"choices": [{
|
||||
"finish_reason": "length",
|
||||
"message": {
|
||||
"content": [],
|
||||
"refusal": None,
|
||||
},
|
||||
}],
|
||||
"usage": {
|
||||
"completion_tokens": 0
|
||||
},
|
||||
})
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No rewrite segment prompts found in assistant response." in (result.error or "")
|
||||
assert isinstance(result.raw_response_text, str)
|
||||
@@ -903,10 +830,7 @@ def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
|
||||
|
||||
def test_get_rewrite_model_config_returns_fixed_defaults():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_default_model = "gpt-oss-120b"
|
||||
enhancer.rewrite_model_options = ["gpt-oss-120b"]
|
||||
|
||||
@@ -918,10 +842,7 @@ def test_get_rewrite_model_config_returns_fixed_defaults():
|
||||
|
||||
def test_get_prompt_config_includes_auto_extension_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.enhance_system_prompt_path = "/tmp/next.md"
|
||||
enhancer.auto_system_prompt_path = "/tmp/auto.md"
|
||||
enhancer.rewrite_all_system_prompt_path = "/tmp/rewrite.md"
|
||||
@@ -948,19 +869,14 @@ def test_get_prompt_config_includes_auto_extension_prompt():
|
||||
|
||||
def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
rewrite_fallback_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
rewrite_fallback_path.write_text("rewrite prompt\n", encoding="utf-8")
|
||||
next_path = tmp_path / "next.md"
|
||||
next_path.write_text("next prompt\n", encoding="utf-8")
|
||||
auto_path = tmp_path / "auto.md"
|
||||
auto_path.write_text("auto prompt\n", encoding="utf-8")
|
||||
enhancer.rewrite_all_system_prompt_path = str(
|
||||
tmp_path / "prompts.local" / "rewrite_window_system_prompt.md"
|
||||
)
|
||||
enhancer.rewrite_all_system_prompt_path = str(tmp_path / "prompts.local" / "rewrite_window_system_prompt.md")
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = str(rewrite_fallback_path)
|
||||
enhancer.enhance_system_prompt_path = str(next_path)
|
||||
enhancer.auto_system_prompt_path = str(auto_path)
|
||||
@@ -973,14 +889,9 @@ def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
|
||||
assert config["rewrite_window_system_prompt_path"] == str(rewrite_fallback_path)
|
||||
|
||||
|
||||
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(
|
||||
tmp_path,
|
||||
):
|
||||
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(tmp_path, ):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
@@ -1007,10 +918,7 @@ def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_emp
|
||||
|
||||
def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1021,9 +929,7 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
enhancer.auto_system_prompt_path = str(auto_path)
|
||||
enhancer.rewrite_all_system_prompt_path = str(rewrite_path)
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
auto_extension_system_prompt="auto updated",
|
||||
)
|
||||
config = enhancer.save_prompt_config(auto_extension_system_prompt="auto updated", )
|
||||
|
||||
assert auto_path.read_text(encoding="utf-8").strip() == "auto updated"
|
||||
assert config["auto_extension_system_prompt"] == "auto updated"
|
||||
@@ -1031,10 +937,7 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1052,9 +955,7 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_user_system_prompt_fallback_path = None
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_user_system_prompt="rewrite user updated",
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_user_system_prompt="rewrite user updated", )
|
||||
|
||||
assert rewrite_user_path.read_text(encoding="utf-8").strip() == "rewrite user updated"
|
||||
assert config["rewrite_user_system_prompt"] == "rewrite user updated"
|
||||
@@ -1062,10 +963,7 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1081,9 +979,7 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_model="gpt-alt",
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_model="gpt-alt", )
|
||||
|
||||
assert enhancer.rewrite_default_model == "gpt-alt"
|
||||
assert config["rewrite_model"] == "gpt-alt"
|
||||
@@ -1092,10 +988,7 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1109,9 +1002,7 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
enhancer.auto_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_temperature=1.3,
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_temperature=1.3, )
|
||||
|
||||
assert enhancer.rewrite_default_temperature == 1.3
|
||||
assert config["rewrite_temperature"] == 1.3
|
||||
@@ -1119,10 +1010,7 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
|
||||
def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
@@ -1136,13 +1024,9 @@ def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_pat
|
||||
enhancer.auto_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
|
||||
enhancer.save_prompt_config(
|
||||
rewrite_window_system_prompt="rewrite updated",
|
||||
)
|
||||
enhancer.save_prompt_config(rewrite_window_system_prompt="rewrite updated", )
|
||||
|
||||
backup_paths = sorted(
|
||||
tmp_path.glob("rewrite_window_system_prompt.*.bak.md")
|
||||
)
|
||||
backup_paths = sorted(tmp_path.glob("rewrite_window_system_prompt.*.bak.md"))
|
||||
|
||||
assert rewrite_path.read_text(encoding="utf-8").strip() == "rewrite updated"
|
||||
assert len(backup_paths) == 1
|
||||
|
||||
@@ -27,13 +27,8 @@ try:
|
||||
except ModuleNotFoundError:
|
||||
websockets = None # type: ignore[assignment]
|
||||
|
||||
|
||||
DEFAULT_PRESET_FILE = (
|
||||
Path(__file__).resolve().parents[2]
|
||||
/ "web"
|
||||
/ "prompts"
|
||||
/ "selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
DEFAULT_PRESET_FILE = (Path(__file__).resolve().parents[2] / "web" / "prompts" /
|
||||
"selected_ltx2_continuation_story_presets.json")
|
||||
|
||||
|
||||
def utc_now_iso() -> str:
|
||||
@@ -65,10 +60,7 @@ def safe_percentile(values: list[float], percentile: float) -> float | None:
|
||||
if lower == upper:
|
||||
return sorted_values[lower]
|
||||
fraction = rank - lower
|
||||
return (
|
||||
sorted_values[lower]
|
||||
+ (sorted_values[upper] - sorted_values[lower]) * fraction
|
||||
)
|
||||
return (sorted_values[lower] + (sorted_values[upper] - sorted_values[lower]) * fraction)
|
||||
|
||||
|
||||
def summarize_series(values: list[float]) -> dict[str, float | int | None]:
|
||||
@@ -145,24 +137,16 @@ def load_curated_prompts(
|
||||
selected_id = str(selected.get("id", "")).strip() or "unknown_preset"
|
||||
raw_prompts = selected.get("segment_prompts", [])
|
||||
if not isinstance(raw_prompts, list):
|
||||
raise ValueError(
|
||||
f"Preset {selected_id} has invalid segment_prompts (must be list)."
|
||||
)
|
||||
raise ValueError(f"Preset {selected_id} has invalid segment_prompts (must be list).")
|
||||
|
||||
prompts = [
|
||||
str(prompt).strip()
|
||||
for prompt in raw_prompts
|
||||
if isinstance(prompt, str) and str(prompt).strip()
|
||||
]
|
||||
prompts = [str(prompt).strip() for prompt in raw_prompts if isinstance(prompt, str) and str(prompt).strip()]
|
||||
if not prompts:
|
||||
raise ValueError(f"Preset {selected_id} has no non-empty prompts.")
|
||||
|
||||
limited = prompts[:curated_limit]
|
||||
if not limited:
|
||||
raise ValueError(
|
||||
f"curated_limit={curated_limit} produced no prompts for preset "
|
||||
f"{selected_id}."
|
||||
)
|
||||
raise ValueError(f"curated_limit={curated_limit} produced no prompts for preset "
|
||||
f"{selected_id}.")
|
||||
return selected_id, limited, len(prompts)
|
||||
|
||||
|
||||
@@ -224,11 +208,11 @@ async def run_single_session(
|
||||
|
||||
try:
|
||||
async with websockets.connect(
|
||||
url,
|
||||
max_size=None,
|
||||
ping_interval=None,
|
||||
open_timeout=connect_timeout_s,
|
||||
close_timeout=2.0,
|
||||
url,
|
||||
max_size=None,
|
||||
ping_interval=None,
|
||||
open_timeout=connect_timeout_s,
|
||||
close_timeout=2.0,
|
||||
) as ws:
|
||||
connect_finish_monotonic = time.monotonic()
|
||||
session_data["connect_finish_ts_utc"] = utc_now_iso()
|
||||
@@ -249,9 +233,7 @@ async def run_single_session(
|
||||
timeout_remaining = session_timeout_s - elapsed_s
|
||||
if timeout_remaining <= 0:
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = (
|
||||
f"Session timed out after {session_timeout_s:.1f}s."
|
||||
)
|
||||
session_data["error"] = (f"Session timed out after {session_timeout_s:.1f}s.")
|
||||
break
|
||||
|
||||
recv_start_epoch = time.time()
|
||||
@@ -265,9 +247,7 @@ async def run_single_session(
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = (
|
||||
"Timed out waiting for websocket message."
|
||||
)
|
||||
session_data["error"] = ("Timed out waiting for websocket message.")
|
||||
break
|
||||
except Exception as exc:
|
||||
session_data["status"] = "failed"
|
||||
@@ -288,20 +268,16 @@ async def run_single_session(
|
||||
|
||||
chunk_gap_ms: float | None = None
|
||||
if last_chunk_finish_monotonic is not None:
|
||||
chunk_gap_ms = (
|
||||
recv_finish_monotonic - last_chunk_finish_monotonic
|
||||
) * 1000.0
|
||||
chunk_gap_ms = (recv_finish_monotonic - last_chunk_finish_monotonic) * 1000.0
|
||||
|
||||
session_data["chunks"].append(
|
||||
{
|
||||
"segment_idx": current_segment_idx,
|
||||
"chunk_idx": session_data["total_chunks"],
|
||||
"size_bytes": len(message),
|
||||
"chunk_start_ts_utc": recv_start_iso,
|
||||
"chunk_finish_ts_utc": recv_finish_iso,
|
||||
"chunk_gap_ms": chunk_gap_ms,
|
||||
}
|
||||
)
|
||||
session_data["chunks"].append({
|
||||
"segment_idx": current_segment_idx,
|
||||
"chunk_idx": session_data["total_chunks"],
|
||||
"size_bytes": len(message),
|
||||
"chunk_start_ts_utc": recv_start_iso,
|
||||
"chunk_finish_ts_utc": recv_finish_iso,
|
||||
"chunk_gap_ms": chunk_gap_ms,
|
||||
})
|
||||
last_chunk_finish_monotonic = recv_finish_monotonic
|
||||
last_chunk_finish_epoch = recv_finish_epoch
|
||||
session_data["last_chunk_finish_ts_utc"] = recv_finish_iso
|
||||
@@ -321,9 +297,7 @@ async def run_single_session(
|
||||
if msg_type == "gpu_assigned":
|
||||
session_data["gpu_assigned_ts_utc"] = recv_finish_iso
|
||||
if connect_finish_monotonic is not None:
|
||||
session_data["queue_wait_ms"] = (
|
||||
recv_finish_monotonic - connect_finish_monotonic
|
||||
) * 1000.0
|
||||
session_data["queue_wait_ms"] = (recv_finish_monotonic - connect_finish_monotonic) * 1000.0
|
||||
elif msg_type == "ltx2_stream_start":
|
||||
if initial_total_segments is None:
|
||||
parsed_total = parse_int(data.get("total_segments"))
|
||||
@@ -338,20 +312,13 @@ async def run_single_session(
|
||||
session_data["media_segments_completed"] += 1
|
||||
if first_media_segment_complete_epoch is None:
|
||||
first_media_segment_complete_epoch = recv_finish_epoch
|
||||
session_data[
|
||||
"first_media_segment_complete_ts_utc"
|
||||
] = recv_finish_iso
|
||||
session_data["first_media_segment_complete_ts_utc"] = recv_finish_iso
|
||||
elif msg_type == "ltx2_segment_complete":
|
||||
session_data["segments_completed"] += 1
|
||||
seg_idx = parse_int(data.get("segment_idx"))
|
||||
if (
|
||||
initial_total_segments is not None
|
||||
and seg_idx is not None
|
||||
and seg_idx >= initial_total_segments
|
||||
):
|
||||
session_data[
|
||||
"target_segment_complete_ts_utc"
|
||||
] = recv_finish_iso
|
||||
if (initial_total_segments is not None and seg_idx is not None
|
||||
and seg_idx >= initial_total_segments):
|
||||
session_data["target_segment_complete_ts_utc"] = recv_finish_iso
|
||||
await asyncio.sleep(post_complete_wait_s)
|
||||
session_data["leave_sent_ts_utc"] = utc_now_iso()
|
||||
try:
|
||||
@@ -362,15 +329,11 @@ async def run_single_session(
|
||||
break
|
||||
elif msg_type == "session_timeout":
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = str(
|
||||
data.get("message") or "Backend session timeout"
|
||||
)
|
||||
session_data["error"] = str(data.get("message") or "Backend session timeout")
|
||||
break
|
||||
elif msg_type == "error":
|
||||
session_data["status"] = "failed"
|
||||
session_data["error"] = str(
|
||||
data.get("message") or "Backend error message"
|
||||
)
|
||||
session_data["error"] = str(data.get("message") or "Backend error message")
|
||||
break
|
||||
|
||||
if session_data["status"] == "failed" and session_data["error"] is None:
|
||||
@@ -379,29 +342,18 @@ async def run_single_session(
|
||||
session_data["status"] = "failed"
|
||||
session_data["error"] = f"WebSocket connect/run failed: {exc}"
|
||||
|
||||
if (
|
||||
first_chunk_finish_epoch is not None
|
||||
and last_chunk_finish_epoch is not None
|
||||
and session_data["total_chunk_bytes"] > 0
|
||||
):
|
||||
if (first_chunk_finish_epoch is not None and last_chunk_finish_epoch is not None
|
||||
and session_data["total_chunk_bytes"] > 0):
|
||||
duration_s = last_chunk_finish_epoch - first_chunk_finish_epoch
|
||||
if duration_s > 0:
|
||||
session_data["session_goodput_mbps"] = (
|
||||
session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0
|
||||
)
|
||||
session_data["session_goodput_mbps"] = (session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0)
|
||||
|
||||
if (
|
||||
first_chunk_finish_epoch is not None
|
||||
and first_media_segment_complete_epoch is not None
|
||||
):
|
||||
session_data["first_chunk_before_first_media_complete"] = (
|
||||
first_chunk_finish_epoch < first_media_segment_complete_epoch
|
||||
)
|
||||
if (first_chunk_finish_epoch is not None and first_media_segment_complete_epoch is not None):
|
||||
session_data["first_chunk_before_first_media_complete"] = (first_chunk_finish_epoch
|
||||
< first_media_segment_complete_epoch)
|
||||
|
||||
session_data["close_ts_utc"] = utc_now_iso()
|
||||
session_data["duration_ms"] = (
|
||||
time.monotonic() - session_start_monotonic
|
||||
) * 1000.0
|
||||
session_data["duration_ms"] = (time.monotonic() - session_start_monotonic) * 1000.0
|
||||
return session_data
|
||||
|
||||
|
||||
@@ -412,14 +364,11 @@ async def run_worker_sessions(
|
||||
config: dict[str, Any],
|
||||
) -> list[dict[str, Any]]:
|
||||
tasks = [
|
||||
asyncio.create_task(
|
||||
run_single_session(
|
||||
worker_id=worker_id,
|
||||
worker_session_idx=idx,
|
||||
config=config,
|
||||
)
|
||||
)
|
||||
for idx in range(session_count)
|
||||
asyncio.create_task(run_single_session(
|
||||
worker_id=worker_id,
|
||||
worker_session_idx=idx,
|
||||
config=config,
|
||||
)) for idx in range(session_count)
|
||||
]
|
||||
if not tasks:
|
||||
return []
|
||||
@@ -437,29 +386,23 @@ def worker_entry(
|
||||
try:
|
||||
ready_queue.put({"worker_id": worker_id, "status": "ready"})
|
||||
start_event.wait()
|
||||
sessions = asyncio.run(
|
||||
run_worker_sessions(
|
||||
worker_id=worker_id,
|
||||
session_count=session_count,
|
||||
config=config,
|
||||
)
|
||||
)
|
||||
result_queue.put(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"status": "ok",
|
||||
"sessions": sessions,
|
||||
}
|
||||
)
|
||||
sessions = asyncio.run(run_worker_sessions(
|
||||
worker_id=worker_id,
|
||||
session_count=session_count,
|
||||
config=config,
|
||||
))
|
||||
result_queue.put({
|
||||
"worker_id": worker_id,
|
||||
"status": "ok",
|
||||
"sessions": sessions,
|
||||
})
|
||||
except Exception as exc:
|
||||
result_queue.put(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"status": "error",
|
||||
"error": str(exc),
|
||||
"traceback": traceback.format_exc(),
|
||||
}
|
||||
)
|
||||
result_queue.put({
|
||||
"worker_id": worker_id,
|
||||
"status": "error",
|
||||
"error": str(exc),
|
||||
"traceback": traceback.format_exc(),
|
||||
})
|
||||
|
||||
|
||||
def build_summary(
|
||||
@@ -517,33 +460,22 @@ def build_summary(
|
||||
if len(all_chunk_finish_epochs) >= 2 and total_chunk_bytes > 0:
|
||||
duration_s = max(all_chunk_finish_epochs) - min(all_chunk_finish_epochs)
|
||||
if duration_s > 0:
|
||||
global_goodput_mbps = (
|
||||
total_chunk_bytes * 8.0 / duration_s / 1_000_000.0
|
||||
)
|
||||
global_goodput_mbps = (total_chunk_bytes * 8.0 / duration_s / 1_000_000.0)
|
||||
|
||||
bucket_throughputs_mbps = [
|
||||
(bytes_count * 8.0) / 1_000_000.0
|
||||
for _, bytes_count in sorted(bucket_bytes.items())
|
||||
]
|
||||
bucket_throughputs_mbps = [(bytes_count * 8.0) / 1_000_000.0 for _, bytes_count in sorted(bucket_bytes.items())]
|
||||
bucket_stats = summarize_series(bucket_throughputs_mbps)
|
||||
|
||||
chunk_gap_threshold_breaches = [
|
||||
value for value in chunk_gaps if value >= chunk_gap_threshold_ms
|
||||
]
|
||||
chunk_gap_threshold_breaches = [value for value in chunk_gaps if value >= chunk_gap_threshold_ms]
|
||||
non_success = len(sessions) - status_counts.get("success", 0)
|
||||
|
||||
fail_reasons: list[str] = []
|
||||
if non_success > 0:
|
||||
fail_reasons.append(
|
||||
f"{non_success} session(s) did not complete successfully."
|
||||
)
|
||||
fail_reasons.append(f"{non_success} session(s) did not complete successfully.")
|
||||
if not chunk_gaps:
|
||||
fail_reasons.append("No chunk gap data collected.")
|
||||
if chunk_gap_threshold_breaches:
|
||||
fail_reasons.append(
|
||||
f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
|
||||
f"{chunk_gap_threshold_ms:.0f}ms."
|
||||
)
|
||||
fail_reasons.append(f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
|
||||
f"{chunk_gap_threshold_ms:.0f}ms.")
|
||||
|
||||
passed = len(fail_reasons) == 0
|
||||
progressive_ratio = None
|
||||
@@ -554,20 +486,18 @@ def build_summary(
|
||||
"passed": passed,
|
||||
"fail_reasons": fail_reasons,
|
||||
"sessions": {
|
||||
"total": len(sessions),
|
||||
"success": status_counts.get("success", 0),
|
||||
"failed": status_counts.get("failed", 0),
|
||||
"timeout": status_counts.get("timeout", 0),
|
||||
"protocol_error": status_counts.get("protocol_error", 0),
|
||||
"other": (
|
||||
len(sessions)
|
||||
- (
|
||||
status_counts.get("success", 0)
|
||||
+ status_counts.get("failed", 0)
|
||||
+ status_counts.get("timeout", 0)
|
||||
+ status_counts.get("protocol_error", 0)
|
||||
)
|
||||
),
|
||||
"total":
|
||||
len(sessions),
|
||||
"success":
|
||||
status_counts.get("success", 0),
|
||||
"failed":
|
||||
status_counts.get("failed", 0),
|
||||
"timeout":
|
||||
status_counts.get("timeout", 0),
|
||||
"protocol_error":
|
||||
status_counts.get("protocol_error", 0),
|
||||
"other": (len(sessions) - (status_counts.get("success", 0) + status_counts.get("failed", 0) +
|
||||
status_counts.get("timeout", 0) + status_counts.get("protocol_error", 0))),
|
||||
},
|
||||
"chunk_gap_ms": {
|
||||
**chunk_gap_stats,
|
||||
@@ -606,51 +536,39 @@ def print_summary(
|
||||
bucket_bw = bandwidth["bucketed_1s"]
|
||||
|
||||
print("=== LTX2 Realtime Stress Test Summary ===")
|
||||
print(
|
||||
"Run: "
|
||||
f"url={run_info['url']} clients={run_info['clients']} "
|
||||
f"processes={run_info['processes']} "
|
||||
f"preset={run_info['preset_id']} "
|
||||
f"curated_limit={run_info['curated_limit']}"
|
||||
)
|
||||
print(
|
||||
"Sessions: "
|
||||
f"total={sessions['total']} success={sessions['success']} "
|
||||
f"failed={sessions['failed']} timeout={sessions['timeout']} "
|
||||
f"protocol_error={sessions['protocol_error']}"
|
||||
)
|
||||
print(
|
||||
"Chunk gap ms: "
|
||||
f"min={format_num(chunk_gap['min'])} "
|
||||
f"p50={format_num(chunk_gap['p50'])} "
|
||||
f"p95={format_num(chunk_gap['p95'])} "
|
||||
f"p99={format_num(chunk_gap['p99'])} "
|
||||
f"max={format_num(chunk_gap['max'])} "
|
||||
f"threshold={format_num(chunk_gap['threshold_ms'])} "
|
||||
f"breaches={chunk_gap['breach_count']}"
|
||||
)
|
||||
print(
|
||||
"Queue wait ms: "
|
||||
f"min={format_num(queue_wait['min'])} "
|
||||
f"p50={format_num(queue_wait['p50'])} "
|
||||
f"p95={format_num(queue_wait['p95'])} "
|
||||
f"max={format_num(queue_wait['max'])}"
|
||||
)
|
||||
print("Run: "
|
||||
f"url={run_info['url']} clients={run_info['clients']} "
|
||||
f"processes={run_info['processes']} "
|
||||
f"preset={run_info['preset_id']} "
|
||||
f"curated_limit={run_info['curated_limit']}")
|
||||
print("Sessions: "
|
||||
f"total={sessions['total']} success={sessions['success']} "
|
||||
f"failed={sessions['failed']} timeout={sessions['timeout']} "
|
||||
f"protocol_error={sessions['protocol_error']}")
|
||||
print("Chunk gap ms: "
|
||||
f"min={format_num(chunk_gap['min'])} "
|
||||
f"p50={format_num(chunk_gap['p50'])} "
|
||||
f"p95={format_num(chunk_gap['p95'])} "
|
||||
f"p99={format_num(chunk_gap['p99'])} "
|
||||
f"max={format_num(chunk_gap['max'])} "
|
||||
f"threshold={format_num(chunk_gap['threshold_ms'])} "
|
||||
f"breaches={chunk_gap['breach_count']}")
|
||||
print("Queue wait ms: "
|
||||
f"min={format_num(queue_wait['min'])} "
|
||||
f"p50={format_num(queue_wait['p50'])} "
|
||||
f"p95={format_num(queue_wait['p95'])} "
|
||||
f"max={format_num(queue_wait['max'])}")
|
||||
ratio = progressive["ratio"]
|
||||
ratio_text = "n/a" if ratio is None else f"{ratio * 100:.2f}%"
|
||||
print(
|
||||
"Progressive streaming: "
|
||||
f"{progressive['success_sessions']}/"
|
||||
f"{progressive['eligible_sessions']} ({ratio_text})"
|
||||
)
|
||||
print(
|
||||
"Bandwidth Mbps: "
|
||||
f"per_session_avg={format_num(per_session_bw['avg'])} "
|
||||
f"per_session_p95={format_num(per_session_bw['p95'])} "
|
||||
f"global={format_num(bandwidth['global_goodput_mbps'])} "
|
||||
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
|
||||
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}"
|
||||
)
|
||||
print("Progressive streaming: "
|
||||
f"{progressive['success_sessions']}/"
|
||||
f"{progressive['eligible_sessions']} ({ratio_text})")
|
||||
print("Bandwidth Mbps: "
|
||||
f"per_session_avg={format_num(per_session_bw['avg'])} "
|
||||
f"per_session_p95={format_num(per_session_bw['p95'])} "
|
||||
f"global={format_num(bandwidth['global_goodput_mbps'])} "
|
||||
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
|
||||
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}")
|
||||
print(f"VERDICT: {'PASS' if summary['passed'] else 'FAIL'}")
|
||||
if summary["fail_reasons"]:
|
||||
print("Fail reasons:")
|
||||
@@ -670,10 +588,8 @@ def distribute_sessions(total_clients: int, process_count: int) -> list[int]:
|
||||
|
||||
def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
if websockets is None:
|
||||
raise RuntimeError(
|
||||
"Missing dependency: websockets. Install it before running this "
|
||||
"stress test."
|
||||
)
|
||||
raise RuntimeError("Missing dependency: websockets. Install it before running this "
|
||||
"stress test.")
|
||||
|
||||
preset_file = Path(args.preset_file).expanduser().resolve()
|
||||
selected_preset_id, curated_prompts, total_prompt_count = load_curated_prompts(
|
||||
@@ -735,13 +651,8 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
start_event.set()
|
||||
|
||||
result_deadline = (
|
||||
time.monotonic()
|
||||
+ args.connect_timeout_s
|
||||
+ args.session_timeout_s
|
||||
+ args.post_complete_wait_s
|
||||
+ 180.0
|
||||
)
|
||||
result_deadline = (time.monotonic() + args.connect_timeout_s + args.session_timeout_s +
|
||||
args.post_complete_wait_s + 180.0)
|
||||
worker_results: list[dict[str, Any]] = []
|
||||
while len(worker_results) < len(processes):
|
||||
timeout_s = max(0.1, result_deadline - time.monotonic())
|
||||
@@ -765,24 +676,20 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
if result.get("status") == "ok":
|
||||
sessions.extend(result.get("sessions", []))
|
||||
else:
|
||||
worker_errors.append(
|
||||
{
|
||||
"worker_id": result.get("worker_id"),
|
||||
"error": result.get("error"),
|
||||
"traceback": result.get("traceback"),
|
||||
}
|
||||
)
|
||||
worker_errors.append({
|
||||
"worker_id": result.get("worker_id"),
|
||||
"error": result.get("error"),
|
||||
"traceback": result.get("traceback"),
|
||||
})
|
||||
|
||||
received_workers = {result.get("worker_id") for result in worker_results}
|
||||
expected_workers = set(range(len(processes)))
|
||||
missing_workers = sorted(expected_workers - received_workers)
|
||||
for worker_id in missing_workers:
|
||||
worker_errors.append(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"error": "No worker result received.",
|
||||
}
|
||||
)
|
||||
worker_errors.append({
|
||||
"worker_id": worker_id,
|
||||
"error": "No worker result received.",
|
||||
})
|
||||
|
||||
run_end_epoch = time.time()
|
||||
run_end_iso = iso_from_epoch(run_end_epoch)
|
||||
@@ -795,9 +702,8 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
if worker_errors:
|
||||
summary["passed"] = False
|
||||
summary["fail_reasons"] = list(summary["fail_reasons"]) + [
|
||||
f"{len(worker_errors)} worker error(s) occurred."
|
||||
]
|
||||
summary["fail_reasons"] = list(
|
||||
summary["fail_reasons"]) + [f"{len(worker_errors)} worker error(s) occurred."]
|
||||
|
||||
output_payload = {
|
||||
"run_info": {
|
||||
@@ -833,9 +739,7 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Multiprocess realtime stress test for LTX2 streaming.",
|
||||
)
|
||||
parser = argparse.ArgumentParser(description="Multiprocess realtime stress test for LTX2 streaming.", )
|
||||
parser.add_argument(
|
||||
"-u",
|
||||
"--url",
|
||||
|
||||
@@ -47,13 +47,11 @@ def test_persist_session_init_image_returns_none_when_missing_data():
|
||||
|
||||
def test_persist_session_init_image_rejects_unsupported_mime():
|
||||
with pytest.raises(ValueError, match="PNG, JPEG, or WebP"):
|
||||
persist_session_init_image(
|
||||
{
|
||||
"name": "frame.gif",
|
||||
"mime_type": "image/gif",
|
||||
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
|
||||
}
|
||||
)
|
||||
persist_session_init_image({
|
||||
"name": "frame.gif",
|
||||
"mime_type": "image/gif",
|
||||
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
|
||||
})
|
||||
|
||||
|
||||
def test_persist_session_init_image_rejects_large_payload(monkeypatch):
|
||||
@@ -66,10 +64,8 @@ def test_persist_session_init_image_rejects_large_payload(monkeypatch):
|
||||
monkeypatch.setattr(base64, "b64decode", fake_b64decode)
|
||||
|
||||
with pytest.raises(ValueError, match="15 MB or smaller"):
|
||||
persist_session_init_image(
|
||||
{
|
||||
"name": "frame.png",
|
||||
"mime_type": "image/png",
|
||||
"data_url": data_url,
|
||||
}
|
||||
)
|
||||
persist_session_init_image({
|
||||
"name": "frame.png",
|
||||
"mime_type": "image/png",
|
||||
"data_url": data_url,
|
||||
})
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -39,6 +39,17 @@ export FASTVIDEO_GENERATION_SEGMENT_CAP="${FASTVIDEO_GENERATION_SEGMENT_CAP:-6}"
|
||||
export FASTVIDEO_PROMPT_AUTO_SLEEP_MS="${FASTVIDEO_PROMPT_AUTO_SLEEP_MS:-120}"
|
||||
export FASTVIDEO_PROMPT_AUTO_TIMEOUT_MS="${FASTVIDEO_PROMPT_AUTO_TIMEOUT_MS:-1800}"
|
||||
|
||||
if [[ "${ENABLE_TORCH_COMPILE}" == "1" ]]; then
|
||||
# Persist Inductor, AOTAutograd, and Triton artifacts across launches.
|
||||
export DREAMVERSE_TORCH_COMPILE_CACHE_ROOT="${DREAMVERSE_TORCH_COMPILE_CACHE_ROOT:-${HOME}/.cache/dreamverse/torch_compile}"
|
||||
export TORCHINDUCTOR_CACHE_DIR="${TORCHINDUCTOR_CACHE_DIR:-${DREAMVERSE_TORCH_COMPILE_CACHE_ROOT}/inductor}"
|
||||
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-${DREAMVERSE_TORCH_COMPILE_CACHE_ROOT}/triton}"
|
||||
export TORCHINDUCTOR_FX_GRAPH_CACHE="${TORCHINDUCTOR_FX_GRAPH_CACHE:-1}"
|
||||
export TORCHINDUCTOR_AUTOGRAD_CACHE="${TORCHINDUCTOR_AUTOGRAD_CACHE:-1}"
|
||||
mkdir -p "${TORCHINDUCTOR_CACHE_DIR}" "${TRITON_CACHE_DIR}"
|
||||
echo "[launch-demo] torch.compile cache: ${DREAMVERSE_TORCH_COMPILE_CACHE_ROOT}"
|
||||
fi
|
||||
|
||||
cd "${DREAMVERSE_ROOT}"
|
||||
|
||||
if ! command -v dreamverse-server >/dev/null 2>&1; then
|
||||
|
||||
@@ -8,12 +8,10 @@ import modal
|
||||
|
||||
IMAGE = os.environ.get("DREAMVERSE_IMAGE")
|
||||
if not IMAGE:
|
||||
raise RuntimeError(
|
||||
"DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
|
||||
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
|
||||
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
|
||||
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag."
|
||||
)
|
||||
raise RuntimeError("DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
|
||||
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
|
||||
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
|
||||
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag.")
|
||||
|
||||
# ``@modal.web_server`` invokes ``serve()`` directly and bypasses the image
|
||||
# ENTRYPOINT (``docker/docker_entrypoint.sh``). That entrypoint normally
|
||||
@@ -65,14 +63,10 @@ def serve():
|
||||
# ``or ""`` collapses ``None`` (unset) into an empty string, ``.strip()``
|
||||
# collapses whitespace-only values (e.g. ``" "``) — both should be
|
||||
# treated as missing.
|
||||
missing = [
|
||||
k for k in _REQUIRED_SECRET_KEYS
|
||||
if not (os.environ.get(k) or "").strip()
|
||||
]
|
||||
missing = [k for k in _REQUIRED_SECRET_KEYS if not (os.environ.get(k) or "").strip()]
|
||||
if missing:
|
||||
raise RuntimeError(
|
||||
"dreamverse-api-keys secret is missing required entries: "
|
||||
f"{', '.join(missing)}. Add them with `modal secret create "
|
||||
"dreamverse-api-keys ... --force` and redeploy "
|
||||
"(see apps/dreamverse/scripts/modal/README.md).")
|
||||
raise RuntimeError("dreamverse-api-keys secret is missing required entries: "
|
||||
f"{', '.join(missing)}. Add them with `modal secret create "
|
||||
"dreamverse-api-keys ... --force` and redeploy "
|
||||
"(see apps/dreamverse/scripts/modal/README.md).")
|
||||
subprocess.Popen(["dreamverse-server", "--host", "0.0.0.0", "--port", "8009"])
|
||||
|
||||
@@ -13,10 +13,9 @@ test.describe('create inference job', () => {
|
||||
test('creates a T2V job and shows it in the queue', async ({ page }) => {
|
||||
await page.goto('/inference');
|
||||
|
||||
// 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();
|
||||
// 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();
|
||||
const t2vItem = page.getByRole('menuitem', { name: /T2V/i });
|
||||
await expect(t2vItem).toBeVisible();
|
||||
await t2vItem.click();
|
||||
|
||||
@@ -4,7 +4,7 @@ import { API_BASE, skipWithoutMock } from './helpers';
|
||||
|
||||
/**
|
||||
* Gallery page: the seeded completed inference job surfaces as a media tile
|
||||
* (an <article> wrapping a <video>) captioned with its prompt.
|
||||
* with playback controls or an explicit media-error fallback.
|
||||
*/
|
||||
test.describe('gallery', () => {
|
||||
skipWithoutMock();
|
||||
@@ -30,12 +30,15 @@ test.describe('gallery', () => {
|
||||
page.getByRole('heading', { level: 1, name: 'Gallery' }),
|
||||
).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 tile = page.locator('article').filter({ hasText: completed!.prompt });
|
||||
await expect(tile).toBeVisible();
|
||||
await expect(
|
||||
tile.locator('video').or(tile.getByText('Preview unavailable')),
|
||||
).toBeVisible();
|
||||
|
||||
await expect(page.getByText(completed!.prompt)).toBeVisible();
|
||||
const video = tile.locator('video');
|
||||
if (await video.isVisible()) {
|
||||
await expect(video).toHaveAttribute('controls', '');
|
||||
}
|
||||
});
|
||||
});
|
||||
|
||||
@@ -42,6 +42,74 @@ 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,
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
Generated
+681
@@ -9,6 +9,7 @@
|
||||
"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",
|
||||
@@ -1936,6 +1937,183 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"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",
|
||||
@@ -2017,6 +2195,494 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"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",
|
||||
@@ -2410,6 +3076,21 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"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",
|
||||
|
||||
@@ -18,6 +18,7 @@
|
||||
},
|
||||
"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",
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
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();
|
||||
});
|
||||
});
|
||||
@@ -1,12 +1,14 @@
|
||||
'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';
|
||||
@@ -21,18 +23,28 @@ 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 {
|
||||
setDatasets(await getDatasets());
|
||||
setError(null);
|
||||
const next = await getDatasets();
|
||||
if (sequence === fetchSequence.current) {
|
||||
setDatasets(next);
|
||||
setError(null);
|
||||
}
|
||||
} catch (err) {
|
||||
console.error('Failed to fetch datasets:', err);
|
||||
// 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');
|
||||
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);
|
||||
}
|
||||
}, []);
|
||||
|
||||
@@ -50,28 +62,67 @@ export default function DatasetsPage() {
|
||||
<HeaderActions>
|
||||
<AddDatasetButton />
|
||||
</HeaderActions>
|
||||
<main className="mx-auto flex w-full max-w-[850px] flex-col gap-6 px-4 pb-12">
|
||||
<div className="mx-auto flex w-full max-w-[850px] flex-col gap-6 px-4 pb-12">
|
||||
<Card className="p-6">
|
||||
<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>
|
||||
) : (
|
||||
datasets.map((ds) => (
|
||||
<DatasetCard
|
||||
key={ds.id}
|
||||
dataset={ds}
|
||||
onUpdated={fetchDatasets}
|
||||
onSelect={() => handleSelectDataset(ds)}
|
||||
<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>
|
||||
) : (
|
||||
<>
|
||||
{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)}
|
||||
/>
|
||||
))
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
</Card>
|
||||
</main>
|
||||
</div>
|
||||
<CreateDatasetModal
|
||||
isOpen={open}
|
||||
onClose={() => setCreateDatasetModalOpen(false)}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { render, screen } from '@testing-library/react';
|
||||
import { fireEvent, render, screen } from '@testing-library/react';
|
||||
import { describe, expect, it, vi } from 'vitest';
|
||||
|
||||
import GalleryPage from './page';
|
||||
@@ -43,6 +43,22 @@ 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 }),
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
'use client';
|
||||
|
||||
import { Loader2 } from 'lucide-react';
|
||||
import { AlertTriangle, ImageOff, Loader2 } from 'lucide-react';
|
||||
import { useEffect, useState } from 'react';
|
||||
|
||||
import { Button } from '@/components/ui/button';
|
||||
import { Card } from '@/components/ui/card';
|
||||
import { getJobVideoUrl, getJobsList } from '@/lib/api';
|
||||
import type { Job } from '@/lib/types';
|
||||
@@ -11,11 +12,61 @@ 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() {
|
||||
@@ -40,7 +91,13 @@ export default function GalleryPage() {
|
||||
return () => {
|
||||
cancelled = true;
|
||||
};
|
||||
}, []);
|
||||
}, [reloadKey]);
|
||||
|
||||
function retry() {
|
||||
setError(null);
|
||||
setIsLoading(true);
|
||||
setReloadKey((k) => k + 1);
|
||||
}
|
||||
|
||||
const galleryJobs = jobs.filter(
|
||||
(j) =>
|
||||
@@ -64,7 +121,16 @@ export default function GalleryPage() {
|
||||
<span>Loading gallery…</span>
|
||||
</div>
|
||||
) : error ? (
|
||||
<p className="py-8 text-destructive">{error}</p>
|
||||
<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>
|
||||
) : galleryJobs.length === 0 ? (
|
||||
<p className="py-8 text-center text-muted-foreground">
|
||||
No completed videos yet
|
||||
@@ -77,24 +143,7 @@ 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">
|
||||
{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"
|
||||
/>
|
||||
)}
|
||||
<GalleryMedia job={job} />
|
||||
</div>
|
||||
<p
|
||||
className="line-clamp-3 border-t border-border px-4 py-3 text-sm text-muted-foreground"
|
||||
|
||||
@@ -41,7 +41,7 @@
|
||||
|
||||
--border: #e2e8f0;
|
||||
--input: #cbd5e1;
|
||||
--ring: #94a3b8;
|
||||
--ring: #1d4ed8;
|
||||
|
||||
--radius: 0.5rem;
|
||||
}
|
||||
@@ -77,7 +77,7 @@
|
||||
|
||||
--border: #334155;
|
||||
--input: #334155;
|
||||
--ring: #cbd5e1;
|
||||
--ring: #7dd3fc;
|
||||
}
|
||||
|
||||
@theme inline {
|
||||
@@ -125,6 +125,7 @@
|
||||
html,
|
||||
body {
|
||||
min-height: 100%;
|
||||
overflow-x: clip;
|
||||
}
|
||||
|
||||
html {
|
||||
@@ -163,6 +164,22 @@ 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;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
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;');
|
||||
});
|
||||
});
|
||||
@@ -4,8 +4,8 @@ import GpuGrid from '@/components/system/GpuGrid';
|
||||
|
||||
export default function GpusPage() {
|
||||
return (
|
||||
<main className="mx-auto flex w-full max-w-[1100px] flex-col gap-6 px-4 pb-12 pt-6">
|
||||
<div className="mx-auto flex w-full max-w-[1100px] flex-col gap-6 px-4 pb-12 pt-6">
|
||||
<GpuGrid />
|
||||
</main>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -52,6 +52,20 @@ 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,6 +25,16 @@ 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();
|
||||
@@ -71,7 +81,7 @@ describe('DatasetCard', () => {
|
||||
expect(onSelect).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it('selects on keyboard activation of the card body but not of the Delete button', () => {
|
||||
it('keeps the selection and delete actions separate', () => {
|
||||
const onSelect = vi.fn();
|
||||
render(
|
||||
<DatasetCard dataset={dataset} onUpdated={() => {}} onSelect={onSelect} />,
|
||||
@@ -83,8 +93,12 @@ describe('DatasetCard', () => {
|
||||
});
|
||||
expect(onSelect).not.toHaveBeenCalled();
|
||||
|
||||
// Activating the card body itself does select.
|
||||
fireEvent.keyDown(screen.getByText('My Dataset'), { key: 'Enter' });
|
||||
// Activating the dedicated selection button selects the dataset.
|
||||
fireEvent.click(
|
||||
screen.getByRole('button', {
|
||||
name: /My Dataset.*3 files.*2.0 KB/,
|
||||
}),
|
||||
);
|
||||
expect(onSelect).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
|
||||
@@ -55,43 +55,33 @@ 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 (
|
||||
<div
|
||||
<article
|
||||
className={cn(
|
||||
'mb-3 flex cursor-pointer flex-col gap-[0.6rem] rounded-lg border border-border bg-background px-[1.15rem] py-4',
|
||||
'mb-3 flex items-start gap-3 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}
|
||||
>
|
||||
<div className="flex flex-wrap items-center justify-between gap-2">
|
||||
<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"
|
||||
>
|
||||
<span className="text-[0.95rem] font-semibold">{dataset.name}</span>
|
||||
<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>
|
||||
<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>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -6,6 +6,9 @@ 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);
|
||||
|
||||
@@ -27,6 +30,27 @@ 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={() => {}} />);
|
||||
|
||||
@@ -43,6 +67,16 @@ 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');
|
||||
@@ -99,6 +133,39 @@ 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,7 +1,8 @@
|
||||
'use client';
|
||||
|
||||
import * as React from 'react';
|
||||
import { X } from 'lucide-react';
|
||||
import { ImageOff, X } from 'lucide-react';
|
||||
import { toast } from 'sonner';
|
||||
|
||||
import DownloadCaptions from '@/components/datasets/DownloadCaptions';
|
||||
import { Textarea } from '@/components/ui/textarea';
|
||||
@@ -13,12 +14,14 @@ 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).
|
||||
@@ -27,54 +30,101 @@ 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 && (
|
||||
{!thumbLoaded && !mediaFailed && (
|
||||
<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>
|
||||
)}
|
||||
{/* 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)}
|
||||
/>
|
||||
{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);
|
||||
}}
|
||||
/>
|
||||
)}
|
||||
<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[]>([]);
|
||||
@@ -84,17 +134,21 @@ 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?.(width);
|
||||
}, [width, onWidthChange]);
|
||||
onWidthChange?.(isMobile ? 0 : width);
|
||||
}, [isMobile, width, onWidthChange]);
|
||||
|
||||
React.useEffect(() => {
|
||||
let cancelled = false;
|
||||
@@ -106,6 +160,8 @@ 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(() => {
|
||||
@@ -138,23 +194,51 @@ 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 save = () => {
|
||||
updateDatasetCaption(datasetId, fileName, value).catch((err) =>
|
||||
console.error('Failed to save caption:', err),
|
||||
);
|
||||
};
|
||||
const version = (captionVersions.current.get(fileName) ?? 0) + 1;
|
||||
captionVersions.current.set(fileName, version);
|
||||
const save = () => persistCaption(fileName, value, version);
|
||||
const timer = setTimeout(() => {
|
||||
pendingSaves.current.delete(fileName);
|
||||
save();
|
||||
}, 500);
|
||||
pendingSaves.current.set(fileName, { timer, save });
|
||||
},
|
||||
[datasetId],
|
||||
[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],
|
||||
);
|
||||
|
||||
function handleScroll() {
|
||||
@@ -189,8 +273,16 @@ export default function DatasetSidebar({
|
||||
|
||||
return (
|
||||
<aside
|
||||
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 }}
|
||||
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,
|
||||
}}
|
||||
>
|
||||
<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">
|
||||
@@ -203,7 +295,7 @@ export default function DatasetSidebar({
|
||||
onClick={onClose}
|
||||
title="Close"
|
||||
aria-label="Close"
|
||||
className="flex items-center justify-center rounded-lg p-1.5 text-muted-foreground transition-colors hover:bg-accent hover:text-foreground"
|
||||
className="flex size-11 items-center justify-center rounded-lg text-muted-foreground transition-colors hover:bg-accent hover:text-foreground"
|
||||
>
|
||||
<X className="h-[18px] w-[18px]" />
|
||||
</button>
|
||||
@@ -231,7 +323,9 @@ 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}
|
||||
/>
|
||||
))}
|
||||
@@ -240,14 +334,14 @@ export default function DatasetSidebar({
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div
|
||||
{!isMobile && <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 w-full cursor-pointer px-4 py-2 text-left text-sm font-medium text-foreground transition-colors hover:bg-muted disabled:cursor-not-allowed disabled:opacity-50';
|
||||
'block min-h-11 w-full cursor-pointer 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,
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
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,6 +2,7 @@
|
||||
|
||||
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';
|
||||
@@ -33,31 +34,35 @@ export default function CreateJobButton({ jobType }: CreateJobButtonProps) {
|
||||
|
||||
return (
|
||||
<>
|
||||
<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>
|
||||
<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>
|
||||
<CreateJobModal
|
||||
isOpen={modalOpen}
|
||||
onClose={() => setModalOpen(false)}
|
||||
|
||||
@@ -50,6 +50,57 @@ 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,6 +103,17 @@ 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
|
||||
@@ -144,6 +155,10 @@ export default function CreateJobModal({
|
||||
setImageFileName('');
|
||||
setSelectedDatasetId('');
|
||||
setSelectedValidationDatasetId('');
|
||||
setModelLoadError(null);
|
||||
setDatasetLoadError(null);
|
||||
setImageUploadError(null);
|
||||
setSubmitError(null);
|
||||
if (workloadType === 'dmd_t2v') {
|
||||
setDmdUseVsa(false);
|
||||
setDmdVsaSparsity(0.8);
|
||||
@@ -162,6 +177,7 @@ 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;
|
||||
@@ -180,7 +196,13 @@ export default function CreateJobModal({
|
||||
}
|
||||
})
|
||||
.catch((e) => {
|
||||
if (!stale) console.error('Failed to load models:', 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.',
|
||||
);
|
||||
})
|
||||
.finally(() => {
|
||||
if (!stale) setIsLoadingModels(false);
|
||||
@@ -193,11 +215,22 @@ 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(() => 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));
|
||||
} else {
|
||||
setReadyDatasets([]);
|
||||
setIsLoadingDatasets(false);
|
||||
setDatasetLoadError(null);
|
||||
}
|
||||
}, [isOpen, isInference]);
|
||||
|
||||
@@ -206,16 +239,24 @@ 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 {
|
||||
} catch (error) {
|
||||
console.error('Failed to upload image:', error);
|
||||
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);
|
||||
}
|
||||
@@ -224,6 +265,7 @@ export default function CreateJobModal({
|
||||
function clearImage() {
|
||||
setImagePath('');
|
||||
setImageFileName('');
|
||||
setImageUploadError(null);
|
||||
if (imageInputRef.current) imageInputRef.current.value = '';
|
||||
}
|
||||
|
||||
@@ -239,6 +281,7 @@ export default function CreateJobModal({
|
||||
workloadType === 'lora_t2v' ? 'lora' : jobType
|
||||
) as JobType;
|
||||
setIsSubmitting(true);
|
||||
setSubmitError(null);
|
||||
try {
|
||||
const payload: CreateJobRequest = {
|
||||
model_id: modelId,
|
||||
@@ -296,6 +339,11 @@ 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);
|
||||
}
|
||||
@@ -343,7 +391,11 @@ export default function CreateJobModal({
|
||||
value={modelId}
|
||||
onChange={(e) => setModelId(e.target.value)}
|
||||
required
|
||||
disabled={isSubmitting || isLoadingModels}
|
||||
aria-describedby={
|
||||
modelLoadError ? 'modal-model-error' : undefined
|
||||
}
|
||||
aria-invalid={modelLoadError ? true : undefined}
|
||||
disabled={isSubmitting || isLoadingModels || !!modelLoadError}
|
||||
>
|
||||
<option value="" disabled>
|
||||
{isLoadingModels
|
||||
@@ -358,6 +410,15 @@ 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' && (
|
||||
@@ -369,6 +430,10 @@ 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"
|
||||
/>
|
||||
@@ -385,6 +450,15 @@ export default function CreateJobModal({
|
||||
</button>
|
||||
</span>
|
||||
)}
|
||||
{imageUploadError && (
|
||||
<p
|
||||
id="modal-image-error"
|
||||
role="alert"
|
||||
className="text-sm text-destructive"
|
||||
>
|
||||
{imageUploadError}
|
||||
</p>
|
||||
)}
|
||||
</FieldRow>
|
||||
)}
|
||||
|
||||
@@ -432,12 +506,22 @@ export default function CreateJobModal({
|
||||
id="modal-dataset"
|
||||
value={selectedDatasetId}
|
||||
onChange={(e) => setSelectedDatasetId(e.target.value)}
|
||||
disabled={isSubmitting}
|
||||
aria-describedby={
|
||||
datasetLoadError ? 'modal-dataset-error' : undefined
|
||||
}
|
||||
aria-invalid={datasetLoadError ? true : undefined}
|
||||
disabled={
|
||||
isSubmitting || isLoadingDatasets || !!datasetLoadError
|
||||
}
|
||||
>
|
||||
<option value="" disabled>
|
||||
{readyDatasets.length === 0
|
||||
? 'No datasets (add in Datasets tab)'
|
||||
: 'Select a dataset…'}
|
||||
{isLoadingDatasets
|
||||
? 'Loading datasets…'
|
||||
: datasetLoadError
|
||||
? 'Datasets unavailable'
|
||||
: readyDatasets.length === 0
|
||||
? 'No datasets (add in Datasets tab)'
|
||||
: 'Select a dataset…'}
|
||||
</option>
|
||||
{readyDatasets.map((d) => (
|
||||
<option key={d.id} value={d.id}>
|
||||
@@ -445,6 +529,15 @@ 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"
|
||||
@@ -457,7 +550,9 @@ export default function CreateJobModal({
|
||||
onChange={(e) =>
|
||||
setSelectedValidationDatasetId(e.target.value)
|
||||
}
|
||||
disabled={isSubmitting}
|
||||
disabled={
|
||||
isSubmitting || isLoadingDatasets || !!datasetLoadError
|
||||
}
|
||||
>
|
||||
<option value="">None</option>
|
||||
{readyDatasets.map((d) => (
|
||||
@@ -811,9 +906,22 @@ export default function CreateJobModal({
|
||||
</details>
|
||||
)}
|
||||
|
||||
<div>
|
||||
<Button type="submit" disabled={isSubmitting}>
|
||||
{isSubmitting ? 'Creating...' : 'Create Job'}
|
||||
<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'}
|
||||
</Button>
|
||||
</div>
|
||||
</form>
|
||||
|
||||
@@ -20,6 +20,14 @@ 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',
|
||||
@@ -46,6 +54,16 @@ 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,18 +119,10 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
}
|
||||
}
|
||||
|
||||
function handleSelectJob(e: React.MouseEvent | React.KeyboardEvent) {
|
||||
if ((e.target as HTMLElement).closest('button')) return;
|
||||
function handleSelectJob() {
|
||||
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();
|
||||
@@ -148,11 +140,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
}
|
||||
|
||||
return (
|
||||
<div
|
||||
role="button"
|
||||
tabIndex={0}
|
||||
onClick={handleSelectJob}
|
||||
onKeyDown={handleKeyDown}
|
||||
<article
|
||||
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
|
||||
@@ -160,35 +148,42 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
: 'border-border hover:border-muted-foreground/40',
|
||||
)}
|
||||
>
|
||||
<div 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>
|
||||
<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>{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}
|
||||
<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>
|
||||
)}
|
||||
</div>
|
||||
<Badge variant={BADGE_VARIANTS[job.status] ?? 'secondary'}>
|
||||
{job.status}
|
||||
</Badge>
|
||||
</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}
|
||||
</span>
|
||||
)}
|
||||
</span>
|
||||
</button>
|
||||
<div className="flex flex-wrap items-center gap-1.5">
|
||||
{job.status === 'running' ? (
|
||||
<Button
|
||||
@@ -240,6 +235,6 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
Delete
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
</article>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -19,6 +19,32 @@ 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,6 +4,7 @@ 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';
|
||||
@@ -15,13 +16,16 @@ 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);
|
||||
@@ -53,8 +57,8 @@ export default function JobDetailsSidebar({
|
||||
});
|
||||
|
||||
React.useEffect(() => {
|
||||
onWidthChange?.(width);
|
||||
}, [width, onWidthChange]);
|
||||
onWidthChange?.(isMobile ? 0 : width);
|
||||
}, [isMobile, width, onWidthChange]);
|
||||
|
||||
// Auto-scroll the console to the bottom whenever new lines land. Runs after
|
||||
// commit so scrollHeight reflects the freshly-rendered output.
|
||||
@@ -137,8 +141,16 @@ export default function JobDetailsSidebar({
|
||||
|
||||
return (
|
||||
<aside
|
||||
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 }}
|
||||
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,
|
||||
}}
|
||||
>
|
||||
<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">
|
||||
@@ -195,14 +207,14 @@ export default function JobDetailsSidebar({
|
||||
</pre>
|
||||
</div>
|
||||
|
||||
<div
|
||||
{!isMobile && <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, render, screen, waitFor } from '@testing-library/react';
|
||||
import { act, fireEvent, render, screen, waitFor } from '@testing-library/react';
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
|
||||
import JobQueue from '@/components/jobs/JobQueue';
|
||||
@@ -37,6 +37,46 @@ 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,8 +1,10 @@
|
||||
'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';
|
||||
@@ -32,6 +34,8 @@ 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);
|
||||
|
||||
@@ -71,11 +75,22 @@ export default function JobQueue({ jobType, jobTypesForList }: JobQueueProps) {
|
||||
new Date(a.created_at ?? 0).getTime(),
|
||||
);
|
||||
}
|
||||
if (seq === fetchSeq.current) setJobs(next);
|
||||
if (seq === fetchSeq.current) {
|
||||
setJobs(next);
|
||||
setError(null);
|
||||
}
|
||||
} 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;
|
||||
if (seq === fetchSeq.current) {
|
||||
inFlight.current = false;
|
||||
setIsInitialLoading(false);
|
||||
}
|
||||
}
|
||||
}, [typesKey]);
|
||||
|
||||
@@ -122,20 +137,59 @@ export default function JobQueue({ jobType, jobTypesForList }: JobQueueProps) {
|
||||
const multiType = typesToFetch.length > 1;
|
||||
|
||||
return (
|
||||
<main className="mx-auto flex w-full max-w-[850px] flex-col gap-6 px-4 pb-12">
|
||||
<div className="mx-auto flex w-full max-w-[850px] flex-col gap-6 px-4 pb-12">
|
||||
<section className="p-6">
|
||||
<div>
|
||||
{jobs.length === 0 ? (
|
||||
<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 ? (
|
||||
<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>
|
||||
</main>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ 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,
|
||||
@@ -23,46 +24,80 @@ 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]')) {
|
||||
if (activeJobStore.get().activeJob) setActiveJobId(null);
|
||||
if (activeDatasetStore.get().activeDataset) setActiveDatasetId(null);
|
||||
if (e.key !== 'Escape' || document.querySelector('[data-modal]')) return;
|
||||
if (primaryOpen) {
|
||||
setPrimaryOpen(false);
|
||||
return;
|
||||
}
|
||||
if (activeJobStore.get().activeJob) setActiveJobId(null);
|
||||
if (activeDatasetStore.get().activeDataset) setActiveDatasetId(null);
|
||||
}
|
||||
document.addEventListener('keydown', handleKeyDown);
|
||||
return () => document.removeEventListener('keydown', handleKeyDown);
|
||||
}, []);
|
||||
}, [primaryOpen]);
|
||||
|
||||
return (
|
||||
<HeaderActionsProvider>
|
||||
<Header />
|
||||
<div
|
||||
style={{ display: 'contents' }}
|
||||
inert={drawerModal ? true : undefined}
|
||||
>
|
||||
<Header
|
||||
navigationOpen={primaryOpen}
|
||||
onNavigationToggle={() => setPrimaryOpen((open) => !open)}
|
||||
/>
|
||||
</div>
|
||||
<div
|
||||
className="flex overflow-hidden"
|
||||
style={{
|
||||
marginTop: 'var(--header-height)',
|
||||
height: 'calc(100vh - var(--header-height))',
|
||||
height: 'calc(100dvh - var(--header-height))',
|
||||
}}
|
||||
>
|
||||
<PrimarySidebar onWidthChange={setPrimaryWidth} />
|
||||
<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"
|
||||
/>
|
||||
)}
|
||||
<main
|
||||
className="flex min-w-0 flex-1 flex-col overflow-auto"
|
||||
inert={drawerModal ? true : undefined}
|
||||
style={{
|
||||
marginLeft: primaryWidth,
|
||||
marginRight: secondaryOpen ? secondaryWidth : 0,
|
||||
marginLeft: isMobile ? 0 : primaryWidth,
|
||||
marginRight: isMobile || !secondaryOpen ? 0 : secondaryWidth,
|
||||
}}
|
||||
>
|
||||
{children}
|
||||
@@ -70,6 +105,7 @@ export function AppShell({ children }: { children: React.ReactNode }) {
|
||||
{jobSidebarOpen && activeJob && (
|
||||
<JobDetailsSidebar
|
||||
job={activeJob}
|
||||
isMobile={isMobile}
|
||||
onClose={() => setActiveJobId(null)}
|
||||
onWidthChange={setSecondaryWidth}
|
||||
/>
|
||||
@@ -77,6 +113,7 @@ export function AppShell({ children }: { children: React.ReactNode }) {
|
||||
{datasetSidebarOpen && activeDataset && (
|
||||
<DatasetSidebar
|
||||
dataset={activeDataset}
|
||||
isMobile={isMobile}
|
||||
onClose={() => setActiveDatasetId(null)}
|
||||
onWidthChange={setSecondaryWidth}
|
||||
/>
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
'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> = {
|
||||
@@ -15,25 +17,47 @@ const TAB_TITLES: Record<string, string> = {
|
||||
'/settings': 'Settings',
|
||||
};
|
||||
|
||||
export default function Header() {
|
||||
export default function Header({
|
||||
navigationOpen,
|
||||
onNavigationToggle,
|
||||
}: {
|
||||
navigationOpen: boolean;
|
||||
onNavigationToggle: () => void;
|
||||
}) {
|
||||
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-6 border-b border-border bg-background/80 px-6 backdrop-blur">
|
||||
<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>
|
||||
{/* eslint-disable-next-line @next/next/no-img-element */}
|
||||
<img
|
||||
src="/logo.svg"
|
||||
alt="FastVideo Logo"
|
||||
width={100}
|
||||
height={42}
|
||||
className="block h-[42px] w-[100px]"
|
||||
className="hidden h-[42px] w-[78px] shrink-0 object-contain min-[361px]:block md:w-[100px]"
|
||||
/>
|
||||
<h1 className="m-0 flex-1 text-xl font-semibold tracking-tight">
|
||||
<h1 className="sr-only m-0 flex-1 text-xl font-semibold tracking-tight md:not-sr-only">
|
||||
{title}
|
||||
</h1>
|
||||
<div className="flex items-center gap-3">
|
||||
<div className="ml-auto flex min-w-0 items-center gap-2 md:gap-3">
|
||||
{actions}
|
||||
<ThemeToggle />
|
||||
</div>
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
'use client';
|
||||
|
||||
import * as React from 'react';
|
||||
import { X } from 'lucide-react';
|
||||
import Link from 'next/link';
|
||||
import { usePathname } from 'next/navigation';
|
||||
|
||||
@@ -19,12 +20,18 @@ const JOB_ROUTES = [
|
||||
] as const;
|
||||
|
||||
const TAB_BASE =
|
||||
'block px-5 py-[0.65rem] text-left text-sm text-muted-foreground transition-colors hover:bg-accent/60 hover:text-foreground';
|
||||
'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';
|
||||
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();
|
||||
@@ -38,8 +45,8 @@ export default function PrimarySidebar({
|
||||
const isJobsActive = JOB_ROUTES.some((r) => pathname === r.href);
|
||||
|
||||
React.useEffect(() => {
|
||||
onWidthChange?.(layoutWidth);
|
||||
}, [layoutWidth, onWidthChange]);
|
||||
onWidthChange?.(isMobile ? 0 : layoutWidth);
|
||||
}, [isMobile, layoutWidth, onWidthChange]);
|
||||
|
||||
React.useEffect(() => {
|
||||
if (JOB_ROUTES.some((r) => pathname === r.href)) {
|
||||
@@ -58,11 +65,37 @@ export default function PrimarySidebar({
|
||||
|
||||
return (
|
||||
<aside
|
||||
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 }}
|
||||
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,
|
||||
}}
|
||||
>
|
||||
{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 className="flex flex-col py-2">
|
||||
<nav
|
||||
aria-label="Primary navigation"
|
||||
className="flex flex-col overflow-y-auto py-2"
|
||||
>
|
||||
<div className="flex flex-col">
|
||||
<button
|
||||
type="button"
|
||||
@@ -95,6 +128,8 @@ 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]',
|
||||
@@ -109,24 +144,32 @@ 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
|
||||
@@ -134,7 +177,7 @@ export default function PrimarySidebar({
|
||||
</nav>
|
||||
)}
|
||||
|
||||
<div
|
||||
{!isMobile && <div
|
||||
className={cn(
|
||||
'absolute bottom-0 p-2',
|
||||
isCollapsed ? '-right-[60px] top-0' : 'right-0',
|
||||
@@ -145,8 +188,7 @@ export default function PrimarySidebar({
|
||||
onClick={() => setIsCollapsed((v) => !v)}
|
||||
title={isCollapsed ? 'Expand sidebar' : 'Collapse sidebar'}
|
||||
className={cn(
|
||||
'flex items-center justify-center rounded-lg text-muted-foreground transition-colors hover:bg-accent hover:text-foreground',
|
||||
isCollapsed ? 'p-3' : 'p-2',
|
||||
'flex size-11 items-center justify-center rounded-lg text-muted-foreground transition-colors hover:bg-accent hover:text-foreground',
|
||||
)}
|
||||
>
|
||||
<svg
|
||||
@@ -159,9 +201,9 @@ export default function PrimarySidebar({
|
||||
<path d={isCollapsed ? 'M9 18l6-6-6-6' : 'M15 18l-6-6 6-6'} />
|
||||
</svg>
|
||||
</button>
|
||||
</div>
|
||||
</div>}
|
||||
|
||||
{!isCollapsed && (
|
||||
{!isMobile && !isCollapsed && (
|
||||
<div
|
||||
role="presentation"
|
||||
onMouseDown={onMouseDown}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { render, screen } from '@testing-library/react';
|
||||
import { act, render, screen } from '@testing-library/react';
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
|
||||
import GpuGrid from './GpuGrid';
|
||||
@@ -73,4 +73,30 @@ 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,7 +1,9 @@
|
||||
'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';
|
||||
@@ -97,7 +99,8 @@ function GpuCard({ gpu }: { gpu: GpuInfo }) {
|
||||
|
||||
export default function GpuGrid() {
|
||||
const [snapshot, setSnapshot] = React.useState<GpuSnapshot | null>(null);
|
||||
const [fetchError, setFetchError] = React.useState(false);
|
||||
const [fetchError, setFetchError] = React.useState<string | null>(null);
|
||||
const [retryToken, setRetryToken] = React.useState(0);
|
||||
|
||||
React.useEffect(() => {
|
||||
let mounted = true;
|
||||
@@ -110,10 +113,14 @@ export default function GpuGrid() {
|
||||
const next = await getGpus();
|
||||
if (mounted) {
|
||||
setSnapshot(next);
|
||||
setFetchError(false);
|
||||
setFetchError(null);
|
||||
}
|
||||
} catch {
|
||||
if (mounted) setFetchError(true);
|
||||
if (mounted) {
|
||||
setFetchError(
|
||||
'GPU status could not be refreshed. The values below may be stale.',
|
||||
);
|
||||
}
|
||||
} finally {
|
||||
inFlight = false;
|
||||
}
|
||||
@@ -125,14 +132,27 @@ export default function GpuGrid() {
|
||||
mounted = false;
|
||||
clearInterval(interval);
|
||||
};
|
||||
}, []);
|
||||
}, [retryToken]);
|
||||
|
||||
if (fetchError && !snapshot) {
|
||||
return (
|
||||
<p className="py-8 text-center text-muted-foreground">
|
||||
Could not reach the API server. GPU status needs the studio API server
|
||||
running.
|
||||
</p>
|
||||
<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>
|
||||
);
|
||||
}
|
||||
if (!snapshot) {
|
||||
@@ -150,9 +170,22 @@ export default function GpuGrid() {
|
||||
return (
|
||||
<div className="flex flex-col gap-4">
|
||||
{fetchError && (
|
||||
<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
|
||||
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>
|
||||
)}
|
||||
<div className="grid gap-4 [grid-template-columns:repeat(auto-fill,minmax(280px,1fr))]">
|
||||
{snapshot.gpus.map((gpu) => (
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
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,11 +29,12 @@ const badgeVariants = cva(
|
||||
);
|
||||
|
||||
export interface BadgeProps
|
||||
extends React.HTMLAttributes<HTMLDivElement>,
|
||||
extends React.HTMLAttributes<HTMLSpanElement>,
|
||||
VariantProps<typeof badgeVariants> {}
|
||||
|
||||
// A span (phrasing content), so badges stay valid inside buttons and links.
|
||||
function Badge({ className, variant, ...props }: BadgeProps) {
|
||||
return <div className={cn(badgeVariants({ variant }), className)} {...props} />;
|
||||
return <span 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 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",
|
||||
"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",
|
||||
{
|
||||
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-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",
|
||||
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",
|
||||
},
|
||||
},
|
||||
defaultVariants: {
|
||||
|
||||
@@ -42,7 +42,7 @@ const DialogContent = React.forwardRef<
|
||||
{...props}
|
||||
>
|
||||
{children}
|
||||
<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">
|
||||
<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">
|
||||
<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-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",
|
||||
"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",
|
||||
className,
|
||||
)}
|
||||
ref={ref}
|
||||
|
||||
@@ -11,7 +11,7 @@ const NativeSelect = React.forwardRef<
|
||||
<select
|
||||
ref={ref}
|
||||
className={cn(
|
||||
'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',
|
||||
'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',
|
||||
className,
|
||||
)}
|
||||
{...props}
|
||||
|
||||
@@ -17,7 +17,7 @@ const SelectTrigger = React.forwardRef<
|
||||
<SelectPrimitive.Trigger
|
||||
ref={ref}
|
||||
className={cn(
|
||||
'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',
|
||||
'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',
|
||||
className,
|
||||
)}
|
||||
{...props}
|
||||
@@ -117,7 +117,7 @@ const SelectItem = React.forwardRef<
|
||||
<SelectPrimitive.Item
|
||||
ref={ref}
|
||||
className={cn(
|
||||
'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',
|
||||
'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',
|
||||
className,
|
||||
)}
|
||||
{...props}
|
||||
|
||||
@@ -8,21 +8,39 @@ 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,
|
||||
)}
|
||||
{...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>
|
||||
));
|
||||
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>
|
||||
),
|
||||
);
|
||||
Slider.displayName = SliderPrimitive.Root.displayName;
|
||||
|
||||
export { Slider };
|
||||
|
||||
@@ -12,14 +12,14 @@ const Switch = React.forwardRef<
|
||||
<SwitchPrimitives.Root
|
||||
ref={ref}
|
||||
className={cn(
|
||||
'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',
|
||||
'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',
|
||||
className,
|
||||
)}
|
||||
{...props}
|
||||
>
|
||||
<SwitchPrimitives.Thumb
|
||||
className={cn(
|
||||
'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',
|
||||
'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',
|
||||
)}
|
||||
/>
|
||||
</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 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',
|
||||
'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',
|
||||
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-sky-400/70 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-sky-400/25 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-ring disabled:cursor-not-allowed disabled:opacity-50",
|
||||
className,
|
||||
)}
|
||||
ref={ref}
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
'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;
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
'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;
|
||||
}
|
||||
@@ -74,10 +74,8 @@ def test_snapshot_shapes_devices(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
}
|
||||
|
||||
|
||||
def test_snapshot_tolerates_missing_sensors(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setitem(sys.modules, "pynvml",
|
||||
_make_fake_pynvml(broken_sensors=True))
|
||||
def test_snapshot_tolerates_missing_sensors(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setitem(sys.modules, "pynvml", _make_fake_pynvml(broken_sensors=True))
|
||||
snap = gpu_mod.get_gpu_snapshot()
|
||||
assert snap["available"] is True
|
||||
g = snap["gpus"][0]
|
||||
@@ -86,8 +84,7 @@ def test_snapshot_tolerates_missing_sensors(
|
||||
assert g["power_limit_watts"] is None
|
||||
|
||||
|
||||
def test_snapshot_reports_nvml_failure(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_snapshot_reports_nvml_failure(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
fake = _make_fake_pynvml()
|
||||
fake.nvmlInit = lambda: (_ for _ in ()).throw(_NVMLError("driver gone"))
|
||||
monkeypatch.setitem(sys.modules, "pynvml", fake)
|
||||
|
||||
@@ -130,8 +130,7 @@ def test_dmd_builds_three_role_models_and_method_knobs() -> None:
|
||||
|
||||
|
||||
def test_dmd_vsa_maps_to_training_vsa_sparsity() -> None:
|
||||
config = build_training_config(
|
||||
_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
|
||||
config = build_training_config(_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
|
||||
assert config["training"]["vsa"]["sparsity"] == 0.9
|
||||
|
||||
|
||||
@@ -161,31 +160,27 @@ def test_validation_callback_only_when_file_given() -> None:
|
||||
without = build_training_config(_job("full_t2v"), "out")
|
||||
assert "validation" not in without["callbacks"]
|
||||
|
||||
with_file = build_training_config(
|
||||
_job("full_t2v", validation_dataset_file="val.json"), "out")
|
||||
with_file = build_training_config(_job("full_t2v", validation_dataset_file="val.json"), "out")
|
||||
validation = with_file["callbacks"]["validation"]
|
||||
assert validation["dataset_file"] == "val.json"
|
||||
assert validation["pipeline_target"].endswith(".WanPipeline")
|
||||
assert validation["sampling_steps"] == [50]
|
||||
|
||||
dmd = build_training_config(
|
||||
_job("dmd_t2v", validation_dataset_file="val.json"), "out")
|
||||
dmd = build_training_config(_job("dmd_t2v", validation_dataset_file="val.json"), "out")
|
||||
validation = dmd["callbacks"]["validation"]
|
||||
assert validation["pipeline_target"].endswith(".WanDMDPipeline")
|
||||
assert validation["sampling_steps"] == [3]
|
||||
assert validation["sampling_timesteps"] == [1000, 757, 522]
|
||||
|
||||
# KD/ODE-init has no sampling-based validation pipeline.
|
||||
ode = build_training_config(
|
||||
_job("ode_init", validation_dataset_file="val.json"), "out")
|
||||
ode = build_training_config(_job("ode_init", validation_dataset_file="val.json"), "out")
|
||||
assert "validation" not in ode["callbacks"]
|
||||
|
||||
|
||||
def test_ltx2_models_are_rejected() -> None:
|
||||
assert is_ltx2_model("Lightricks/LTX-2-19B")
|
||||
with pytest.raises(ValueError, match="LTX-2 training is not supported"):
|
||||
build_training_config(
|
||||
_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
|
||||
build_training_config(_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
|
||||
|
||||
|
||||
def test_unknown_workload_is_rejected() -> None:
|
||||
@@ -195,8 +190,7 @@ def test_unknown_workload_is_rejected() -> None:
|
||||
|
||||
def test_invalid_denoising_steps_are_rejected() -> None:
|
||||
with pytest.raises(ValueError, match="Invalid DMD denoising steps"):
|
||||
build_training_config(
|
||||
_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
|
||||
build_training_config(_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
|
||||
|
||||
|
||||
def test_training_env_has_no_backend_override() -> None:
|
||||
@@ -211,8 +205,7 @@ def test_workloads_match_frontend_job_config() -> None:
|
||||
(src/lib/jobConfig.ts) — drift means creatable-but-unrunnable jobs."""
|
||||
import re
|
||||
|
||||
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" /
|
||||
"jobConfig.ts").read_text(encoding="utf-8")
|
||||
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" / "jobConfig.ts").read_text(encoding="utf-8")
|
||||
all_types = set(re.findall(r'type:\s*"([^"]+)"', job_config))
|
||||
inference_types = {"t2v", "i2v", "t2i"}
|
||||
assert inference_types <= all_types, "jobConfig.ts parse failed"
|
||||
|
||||
+30
-5
@@ -186,15 +186,40 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
|
||||
COPY . .
|
||||
|
||||
# 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.
|
||||
# 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
|
||||
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_BUILD_PARALLEL_LEVEL=${CMAKE_BUILD_PARALLEL_LEVEL} \
|
||||
TORCH_CUDA_ARCH_LIST=${TORCH_CUDA_ARCH_LIST} ./build.sh
|
||||
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
|
||||
|
||||
# Install FastVideo itself (editable) now that the source is present, and set up
|
||||
# shell configuration. Dependencies and the local kernel are already installed,
|
||||
|
||||
@@ -8,13 +8,15 @@ FastVideo provides highly optimized custom attention kernels to accelerate video
|
||||
* **[Sliding Tile Attention (STA)](sta/index.md)**: STA kernel support is kept in
|
||||
`fastvideo-kernel`; full FastVideo STA pipeline workflow is archived in
|
||||
`sta_do_not_delete`.
|
||||
* **[Attn-QAT Training](../training/attn_qat.md)**: Runtime-JIT Triton forward
|
||||
and backward kernels for role-local quantization-aware training.
|
||||
* **Backend development guide**: See the developer guide at
|
||||
[Attention Backend Development](../contributing/attention_backend.md).
|
||||
|
||||
## General Build Instructions
|
||||
|
||||
These instructions apply to building the `fastvideo-kernel` package from
|
||||
source, which includes both STA and VSA kernels.
|
||||
source, which includes STA, VSA, and Attn-QAT kernels.
|
||||
|
||||
### Prerequisites
|
||||
|
||||
|
||||
@@ -255,10 +255,9 @@ If you add a new CI test category:
|
||||
|
||||
### Documentation
|
||||
|
||||
`.github/workflows/infra-docs.yml` builds documentation for same-repository PRs
|
||||
that touch `docs/**`, `mkdocs.yml`, `requirements-mkdocs.txt`, or the workflow
|
||||
itself. Fork PRs skip this executable build instead of waiting for maintainer
|
||||
approval. On pushes to `main`, it also deploys the built site to GitHub Pages.
|
||||
`.github/workflows/infra-docs.yml` builds documentation for PRs that touch
|
||||
`docs/**`, `mkdocs.yml`, `requirements-mkdocs.txt`, or the workflow itself. On
|
||||
pushes to `main`, it also deploys the built site to GitHub Pages.
|
||||
|
||||
The docs job:
|
||||
|
||||
@@ -270,16 +269,26 @@ The docs job:
|
||||
### Docker Images
|
||||
|
||||
`.github/workflows/infra-build-image.yml` supports manual `workflow_dispatch`
|
||||
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`.
|
||||
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.
|
||||
|
||||
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
|
||||
|
||||
@@ -76,6 +76,7 @@ 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."
|
||||
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."
|
||||
@@ -333,15 +334,25 @@ surfaces:
|
||||
use_distill:
|
||||
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
|
||||
scheduler_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
text_encoder_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
tokenizer_archs:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
transformer_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
vae_arch:
|
||||
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.sd35.SD35Config
|
||||
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
|
||||
expand_timesteps:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
|
||||
@@ -421,6 +432,8 @@ surfaces:
|
||||
frame_receptive_field: "MagiHuman internal data-proxy receptive-field setting."
|
||||
image_conditioning: "MagiHuman preset variant marker for reference-image conditioning."
|
||||
ref_audio_offset: "MagiHuman internal data-proxy audio alignment offset."
|
||||
scheduler_sigma_min: "Z-Image scheduler parity invariant; not part of the public typed inference API."
|
||||
scheduler_use_reference_discrete_timesteps: "Z-Image scheduler parity invariant; not part of the public typed inference API."
|
||||
sr_local_attn_layers: "MagiHuman SR internal sparse-attention layer selection."
|
||||
text_offset: "MagiHuman internal data-proxy text alignment offset."
|
||||
vae_stride: "MagiHuman internal VAE/data-proxy stride setting."
|
||||
@@ -432,12 +445,17 @@ 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
|
||||
pose: request.inputs.pose
|
||||
c2ws_plucker_emb: request.inputs.c2ws_plucker_emb
|
||||
action_path: request.inputs.action_path
|
||||
refine_from: request.inputs.refine_from
|
||||
stage1_video: request.inputs.stage1_video
|
||||
prompt: request.prompt
|
||||
@@ -447,6 +465,7 @@ surfaces:
|
||||
output_video_name: request.output.output_video_name
|
||||
num_videos_per_prompt: request.sampling.num_videos_per_prompt
|
||||
seed: request.sampling.seed
|
||||
max_sequence_length: request.sampling.max_sequence_length
|
||||
num_frames: request.sampling.num_frames
|
||||
height: request.sampling.height
|
||||
width: request.sampling.width
|
||||
@@ -456,7 +475,10 @@ surfaces:
|
||||
num_inference_steps: request.sampling.num_inference_steps
|
||||
num_inference_steps_sr: request.sampling.num_inference_steps_sr
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
batch_cfg: request.sampling.batch_cfg
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
cfg_normalization: request.sampling.cfg_normalization
|
||||
cfg_truncation: request.sampling.cfg_truncation
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
use_embedded_guidance: request.sampling.use_embedded_guidance
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
@@ -507,8 +529,6 @@ 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."
|
||||
max_sequence_length: "Model-specific text-encoder sequence cap; not part of the public typed inference API."
|
||||
|
||||
sampling_param_extensions: {}
|
||||
|
||||
|
||||
@@ -2,6 +2,12 @@
|
||||
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computation, enabling much faster video generation.
|
||||
|
||||
!!! tip "Attn-QAT DMD2 workflow"
|
||||
The modular trainer also provides a Wan2.1 MixKit recipe that first
|
||||
fine-tunes with fake-quantized attention, then distills the student to
|
||||
timesteps `[1000, 757, 522]` while teacher and critic remain on Flash
|
||||
Attention. See [Attn-QAT Training](../training/attn_qat.md).
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
We provide two distilled models:
|
||||
|
||||
@@ -5,6 +5,7 @@ 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
|
||||
@@ -54,8 +55,10 @@ UV_TORCH_BACKEND=cu126 uv pip install -e .
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Hardware Requirements
|
||||
## 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)
|
||||
|
||||
@@ -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-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ) for additional support.
|
||||
|
||||
@@ -100,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-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ) for additional support.
|
||||
|
||||
@@ -134,9 +134,16 @@ 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-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ)
|
||||
our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)
|
||||
for additional support.
|
||||
|
||||
## Next: performance & tuning
|
||||
|
||||
Installed and verified? See [DGX Spark: Performance & Tuning](spark_performance.md)
|
||||
for which models are practical on the GB10, what makes them faster, and what
|
||||
won't help on this hardware (and why) — so you don't spend a night tuning knobs
|
||||
that can't move here.
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
# 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.
|
||||
+17
-8
@@ -5,7 +5,7 @@
|
||||
</div>
|
||||
|
||||
<div style="text-align: center;">
|
||||
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.</strong>
|
||||
<strong>FastVideo is a unified post-training and real-time inference framework for accelerated video generation.</strong>
|
||||
</div>
|
||||
|
||||
<div style="text-align: center;">
|
||||
@@ -25,14 +25,23 @@ 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
|
||||
- [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.
|
||||
- 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/))
|
||||
|
||||
## Documentation
|
||||
|
||||
|
||||
@@ -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-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ).
|
||||
|
||||
## Step 1: Pipeline Modules
|
||||
|
||||
|
||||
@@ -4,10 +4,11 @@ This page contains step-by-step instructions to get you quickly started with vid
|
||||
|
||||
## Requirements
|
||||
|
||||
- **OS**: Linux (Tested on Ubuntu 22.04+)
|
||||
- **OS**: Linux (tested on Ubuntu 22.04+), or macOS on Apple silicon via the
|
||||
[MPS installation guide](../getting_started/installation/mps.md)
|
||||
- **Python**: 3.10-3.12
|
||||
- **CUDA**: 12.6 or 13.0
|
||||
- **GPU**: At least one NVIDIA GPU
|
||||
- **CUDA**: 12.6 or 13.0 (NVIDIA GPUs)
|
||||
- **GPU**: At least one NVIDIA GPU, or an Apple silicon chip with MPS
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -134,5 +135,4 @@ 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 Discord](https://discord.gg/JA7cksDz86).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ).
|
||||
|
||||
@@ -3,6 +3,12 @@
|
||||
|
||||
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
|
||||
@@ -119,6 +125,22 @@ pip install "nvidia-cutlass-dsl>=4.5.2" apache-tvm-ffi flashinfer-python
|
||||
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.
|
||||
|
||||
#### Usage
|
||||
|
||||
Enable FP4 attention via the `--nvfp4_fa4` flag:
|
||||
|
||||
@@ -13,6 +13,102 @@ 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
|
||||
@@ -25,6 +121,9 @@ 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 {
|
||||
@@ -56,7 +155,7 @@ pipeline initialization and sampling.
|
||||
| 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 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
@@ -65,14 +164,14 @@ pipeline initialization and sampling.
|
||||
| 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 | See note** | 480P<br>720P | ❌ | ❌ | ❌ | ⭕ | ✅ |
|
||||
| LongCat T2V 13.6B | `FastVideo/LongCat-Video-T2V-Diffusers` | 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 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
@@ -98,6 +197,21 @@ 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
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
# Attn-QAT Training
|
||||
|
||||
Attn-QAT simulates low-bit attention during training while keeping the rest of
|
||||
the training method unchanged. In the modular `fastvideo/train` framework it is
|
||||
a per-role model option, not a separate training method: supervised fine-tuning
|
||||
and DMD2 still own their losses and optimizer cadence.
|
||||
|
||||
This guide covers the QAD Wan2.1-T2V-1.3B MixKit workflow:
|
||||
|
||||
1. run a 4,000-step supervised Attn-QAT fine-tune;
|
||||
2. export the stage-1 DCP checkpoint to Diffusers format; and
|
||||
3. distill the student to three denoising steps with DMD2.
|
||||
|
||||
The ready-to-run configs and wrappers are in
|
||||
`examples/train/scenario/qad_wan2_1_mixkit/`.
|
||||
|
||||
## Role-local attention backends
|
||||
|
||||
A DMD2 run owns three independent model roles. Configure the attention backend
|
||||
on each role so fake quantization is applied only to the student:
|
||||
|
||||
```yaml
|
||||
models:
|
||||
student:
|
||||
attention_backend: ATTN_QAT_TRAIN
|
||||
teacher:
|
||||
attention_backend: FLASH_ATTN
|
||||
critic:
|
||||
attention_backend: FLASH_ATTN
|
||||
```
|
||||
|
||||
The override is active only while that role's transformer is constructed, then
|
||||
the previous process-wide backend is restored. This lets student, teacher, and
|
||||
critic use different implementations in one process. Invalid role-level names
|
||||
fail during configuration instead of silently selecting another backend.
|
||||
|
||||
See [Training Infrastructure](train_infra.md) for the complete model-role
|
||||
configuration reference.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- Install FastVideo and make the `fastvideo-kernel` Python package importable.
|
||||
`ATTN_QAT_TRAIN` intentionally fails instead of falling back to dense
|
||||
attention when its kernel cannot be loaded.
|
||||
- Prepare the precomputed MixKit VAE latents and text embeddings.
|
||||
- Run the commands below from the repository root. The supplied recipe expects
|
||||
four GPUs by default; set `NUM_GPUS` to override it.
|
||||
|
||||
Download the published preprocessed dataset:
|
||||
|
||||
```bash
|
||||
bash examples/datasets/mixkit/download_dataset.sh
|
||||
```
|
||||
|
||||
## Stage 1: supervised Attn-QAT fine-tuning
|
||||
|
||||
The stage-1 config uses `ATTN_QAT_TRAIN` on the student, sequence parallelism
|
||||
across four GPUs, FP32 master weights, and 4,000 optimizer steps:
|
||||
|
||||
```bash
|
||||
NUM_GPUS=4 \
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh
|
||||
```
|
||||
|
||||
Pass a dataset directory as the first positional argument when it differs from
|
||||
the default:
|
||||
|
||||
```bash
|
||||
NUM_GPUS=4 \
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh \
|
||||
/path/to/combined_parquet_dataset
|
||||
```
|
||||
|
||||
The wrapper calls `examples/train/run.sh`; the YAML file remains the source of
|
||||
truth for optimizer, validation, checkpointing, and distributed settings.
|
||||
|
||||
## Export the stage-1 checkpoint
|
||||
|
||||
Modular training checkpoints use Distributed Checkpoint (DCP) format. Export
|
||||
the student before using it to initialize stage 2:
|
||||
|
||||
```bash
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/export_stage1.sh \
|
||||
checkpoints/wan_t2v_qat_finetune/checkpoint-4000 \
|
||||
checkpoints/wan_t2v_qat_finetune/diffusers
|
||||
```
|
||||
|
||||
Both arguments are optional; the command above shows their defaults.
|
||||
|
||||
## Stage 2: three-step DMD2 distillation
|
||||
|
||||
Stage 2 loads the exported student weights, keeps Attn-QAT on the student, and
|
||||
uses Flash Attention for the teacher and critic:
|
||||
|
||||
```bash
|
||||
NUM_GPUS=4 \
|
||||
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage2.sh \
|
||||
data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset \
|
||||
checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors
|
||||
```
|
||||
|
||||
The migrated recipe preserves these behaviors:
|
||||
|
||||
| Behavior | Modular configuration |
|
||||
|---|---|
|
||||
| Student fake-quantized attention | `models.student.attention_backend: ATTN_QAT_TRAIN` |
|
||||
| Teacher and critic full-precision attention | Role-local `FLASH_ATTN` |
|
||||
| Generator update every five critic steps | `method.generator_update_interval: 5` |
|
||||
| Three-step rollout | `method.dmd_denoising_steps: [1000, 757, 522]` |
|
||||
| Score timestep range | `method.min_timestep_ratio: 0.02`, `max_timestep_ratio: 0.98` |
|
||||
| Legacy guidance `cond + 2(cond - uncond)` | Standard CFG scale `3.0` |
|
||||
| Stage handoff | DCP checkpoint to Diffusers export to student override weights |
|
||||
|
||||
The timestep ratios apply to randomly sampled teacher and critic score
|
||||
timesteps; `dmd_denoising_steps` separately controls the student rollout. See
|
||||
[DMD Distillation](../distillation/dmd.md) for general DMD concepts.
|
||||
|
||||
## Architecture-specific Triton routing
|
||||
|
||||
The training kernel is runtime-JIT-compiled Triton code and selects its route on
|
||||
every call. It supports different query and key/value sequence lengths for
|
||||
cross-attention; key and value must have the same sequence length.
|
||||
|
||||
| Hardware/configuration | Route |
|
||||
|---|---|
|
||||
| SM100, validated non-causal BF16 QAT configuration with head dimension 128 | Large-tile forward and split 64x64 backward; optimized backward requires a 16-aligned KV length |
|
||||
| SM120, including RTX 5090 | Previous forward tiling with joined quantized/STE P@V operations and a shallower backward pipeline for long sequences |
|
||||
| Unsupported configurations | Previous Triton implementation |
|
||||
|
||||
Warp specialization is disabled automatically on SM100 and SM120 because the
|
||||
Triton 3.7 NVWS compiler pass aborts for this kernel on Blackwell. No user
|
||||
setting is required.
|
||||
|
||||
The available tuning and comparison controls are:
|
||||
|
||||
| Environment variable | Default | Effect |
|
||||
|---|---|---|
|
||||
| `FASTVIDEO_ATTN_QAT_FWD_MODE` | `fast` | Selects `fast`, `balanced`, or `reference` forward tiling on the SM100 optimized route |
|
||||
| `FASTVIDEO_ATTN_QAT_FWD_EXACT_M` | `0` | Set to `1` to recompute reference-order softmax statistics and keep `dV` bitwise-compatible on the SM100 optimized route |
|
||||
| `FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED` | `1` | Set to `0` to force the previous SM100 forward and backward for comparison |
|
||||
| `FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV` | `1` | Set to `0` to compare SM120 against the split P@V path |
|
||||
|
||||
The first invocation JIT-compiles the selected configuration; later calls reuse
|
||||
the Triton cache. To measure the production shape, run
|
||||
`python benchmarks/benchmark_attn_qat_train.py` from `fastvideo-kernel/`.
|
||||
|
||||
For import and backend-selection failures, see [Debugging](../utilities/debugging.md).
|
||||
@@ -62,6 +62,18 @@ bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh
|
||||
- Gradient checkpointing: `--enable_gradient_checkpointing_type "full"`
|
||||
- Memory scaling: Increase `--sp_size` or reduce `--num_latent_t` to fit in memory
|
||||
|
||||
## Attention Quantization-Aware Training
|
||||
|
||||
Attn-QAT fine-tunes a model while simulating low-bit attention in the forward
|
||||
and backward passes. The modular trainer can select the backend per model role,
|
||||
so a later DMD2 stage can keep fake quantization on the student while the
|
||||
teacher and critic use Flash Attention.
|
||||
|
||||
The ready-to-run Wan2.1 MixKit workflow includes supervised fine-tuning,
|
||||
checkpoint export, and three-step DMD2 distillation:
|
||||
|
||||
**→ [Follow the Attn-QAT training guide](attn_qat.md)**
|
||||
|
||||
## LoRA Finetuning
|
||||
|
||||
LoRA (Low-Rank Adaptation) trains lightweight adapters while keeping the base model frozen. This significantly reduces memory usage and training time.
|
||||
@@ -166,10 +178,11 @@ Ready-to-run training scripts are available for multiple models:
|
||||
| Wan2.1 I2V 14B | I2V | `examples/training/finetune/wan_i2v_14B_480p/crush_smol/` |
|
||||
| Wan2.1-Fun 1.3B InP | I2V | `examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/` |
|
||||
| Wan2.1 VSA | T2V/I2V | `examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/` |
|
||||
| Wan2.1 T2V 1.3B Attn-QAT | QAT SFT + DMD2 | `examples/train/scenario/qad_wan2_1_mixkit/` |
|
||||
|
||||
Each example includes:
|
||||
|
||||
- `download_dataset.sh` — download sample data
|
||||
- a README pointing at the matching download script under `examples/datasets/`
|
||||
- `preprocess_*.sh` — run preprocessing
|
||||
- `finetune_*.sh` — full finetune launcher
|
||||
- `finetune_*_lora.sh` — LoRA finetune launcher
|
||||
|
||||
@@ -43,9 +43,12 @@ Ready-to-run examples with preprocessing scripts, training launchers, and valida
|
||||
|
||||
**→ [Browse all training examples](examples/examples_training_index.md)**
|
||||
|
||||
For the complete two-stage Wan2.1 MixKit quantization-aware workflow, see
|
||||
**[Attn-QAT Training](attn_qat.md)**.
|
||||
|
||||
Each example includes:
|
||||
|
||||
- `download_dataset.sh` — download sample data
|
||||
- a README pointing at the matching download script under `examples/datasets/`
|
||||
- `preprocess_*.sh` — run preprocessing
|
||||
- `finetune_*.sh` — launch training (full finetune or LoRA)
|
||||
- `validation.json` — validation prompts for checkpoints
|
||||
@@ -59,9 +62,11 @@ FastVideo supports several training approaches:
|
||||
| **Full finetune** | Adapt entire model to a new domain or style |
|
||||
| **LoRA finetune** | Lightweight adaptation with frozen base weights |
|
||||
| **VSA finetune** | Finetune with Variable Sparse Attention for efficiency |
|
||||
| **Attn-QAT** | Train with fake-quantized attention, optionally followed by DMD2 distillation |
|
||||
|
||||
## Next Steps
|
||||
|
||||
1. **Get started**: Pick an example from the [training examples index](examples/examples_training_index.md)
|
||||
2. **Prepare data**: Follow [data preprocessing](data_preprocess.md) for your own dataset
|
||||
3. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
|
||||
3. **Train with quantized attention**: Follow the [Attn-QAT two-stage recipe](attn_qat.md)
|
||||
4. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
|
||||
|
||||
@@ -81,6 +81,7 @@ Common model parameters:
|
||||
| `disable_custom_init_weights` | `false` | Skip custom weight initialization (use for teacher/critic) |
|
||||
| `flow_shift` | `3.0` | Timestep shifting factor |
|
||||
| `enable_gradient_checkpointing_type` | `null` | Gradient checkpointing (`"full"` or `null`) |
|
||||
| `attention_backend` | `null` | Optional role-local backend for Wan models (for example `ATTN_QAT_TRAIN`); overrides the process default only while this role's transformer is built |
|
||||
|
||||
Which roles are needed depends on the training method:
|
||||
|
||||
@@ -211,6 +212,29 @@ pipeline:
|
||||
flow_shift: 8
|
||||
```
|
||||
|
||||
Registered transformer linear-quantization configs can also be selected by
|
||||
name. For example, the LTX-2 NVFP4-QAT recipe applies real FP4 forward GEMMs
|
||||
with a straight-through-estimator backward to its deployment-targeted
|
||||
attention/FFN projections:
|
||||
|
||||
```yaml
|
||||
pipeline:
|
||||
dit_config:
|
||||
quant_config: nvfp4_qat_train
|
||||
```
|
||||
|
||||
The LTX-2 recipe in
|
||||
`examples/train/configs/overfit_ltx2_t2v_nvfp4_qat.yaml` combines that linear
|
||||
configuration with `models.student.attention_backend: ATTN_QAT_TRAIN` for
|
||||
video-attention forward/backward. On sm120, its validation callback temporarily
|
||||
switches those layers to `ATTN_QAT_INFER`.
|
||||
On GB200, set `callbacks.validation.attn_qat_infer: false` to keep validation on
|
||||
the train-time QAT backend; the inference kernel is sm120-only.
|
||||
|
||||
User-adaptable LTX-2 fine-tuning recipes (full, LoRA, and NVFP4 QAT) live in
|
||||
`examples/train/configs/fine_tuning/ltx2/`, alongside the other model
|
||||
families under `examples/train/configs/fine_tuning/`.
|
||||
|
||||
---
|
||||
|
||||
## Training Methods
|
||||
@@ -298,6 +322,8 @@ method:
|
||||
| `dmd_denoising_steps` | *(required)* | Timestep schedule for student rollout |
|
||||
| `generator_update_interval` | `1` | Update student every N critic steps |
|
||||
| `real_score_guidance_scale` | `1.0` | CFG scale for teacher predictions |
|
||||
| `min_timestep_ratio` | `0.0` | Lower bound for randomly sampled teacher/critic score timesteps |
|
||||
| `max_timestep_ratio` | `1.0` | Upper bound for randomly sampled teacher/critic score timesteps |
|
||||
| `fake_score_learning_rate` | *(required)* | Critic optimizer learning rate |
|
||||
| `fake_score_betas` | *(required)* | Critic optimizer Adam betas |
|
||||
| `fake_score_lr_scheduler` | *(required)* | Critic LR scheduler type |
|
||||
|
||||
@@ -77,8 +77,13 @@ If forcing a backend fails, verify optional dependencies are installed:
|
||||
- `SAGE_ATTN`: SageAttention package
|
||||
- `SAGE_ATTN_THREE`: upstream `sageattn3` package
|
||||
- `ATTN_QAT_INFER`: `fastvideo-kernel` checkout/source install that exposes
|
||||
`attn_qat_infer`
|
||||
- `ATTN_QAT_TRAIN`: `fastvideo-kernel` install exposing `fastvideo_kernel`
|
||||
`attn_qat_infer`, AND a consumer-Blackwell (sm_120/sm_121) GPU -- on any
|
||||
other device the backend reports unavailable (even if a CUDA 13 wheel
|
||||
bundles the extension) and selection falls back to FlashAttention
|
||||
- `ATTN_QAT_TRAIN`: `fastvideo-kernel`; its runtime-JIT Triton implementation
|
||||
selects an optimized route on SM100, joins the quantized and STE P@V paths on
|
||||
SM120, and retains the previous route for unsupported configurations. See
|
||||
[Attn-QAT Training](../training/attn_qat.md) for architecture controls.
|
||||
|
||||
As a fallback, use:
|
||||
|
||||
|
||||
Regular → Executable
@@ -0,0 +1,17 @@
|
||||
# LingBot World 2 Example Dataset
|
||||
|
||||
These files were copied unchanged from the LingBot World 2 repository for the
|
||||
FastVideo causal-fast inference example.
|
||||
|
||||
- Repository: `https://github.com/Robbyant/lingbot-world-v2.git`
|
||||
- Source commit: `94f43115de8d4a4f9f282126528c300a0b232c5f`
|
||||
- Source directory: `examples/03`
|
||||
|
||||
## Files
|
||||
|
||||
- `image.jpg`: source image for image-to-video generation. SHA-256:
|
||||
`6ee3dacfef32cfef504dd698adb8a660cf15f686535c52fed4903fef27c0edd0`
|
||||
- `poses.npy`: camera-to-world trajectory matrices. SHA-256:
|
||||
`bd0a23a696e184b0b43e7767eb432bfe644690560fe327fa96961affc941c404`
|
||||
- `intrinsics.npy`: camera intrinsic parameters. SHA-256:
|
||||
`821fca6cf957ae8fbb1181307f02479efb1705e04c9e05734cd02fb43462e082`
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.2 MiB |
Binary file not shown.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user