Compare commits
57
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0a861031c7 | ||
|
|
f08c5ee8af | ||
|
|
755f4a7967 | ||
|
|
d543a67b10 | ||
|
|
741aa8d289 | ||
|
|
ad9cd63122 | ||
|
|
68e6ffca9e | ||
|
|
64cdcf6be4 | ||
|
|
aa0d98a6b8 | ||
|
|
93b03bc14d | ||
|
|
160f0c9ccf | ||
|
|
e5d1110a0f | ||
|
|
622217ff2a | ||
|
|
cbab605eff | ||
|
|
ac98869aa1 | ||
|
|
56d4a6074f | ||
|
|
9713ea1275 | ||
|
|
37aa382cce | ||
|
|
3f00983287 | ||
|
|
2dc57f4070 | ||
|
|
9df19be719 | ||
|
|
089eea3970 | ||
|
|
dd8447ecc5 | ||
|
|
dca423fd31 | ||
|
|
1b43af8e8e | ||
|
|
628591b620 | ||
|
|
0980ca563f | ||
|
|
b158388733 | ||
|
|
c4ad4227c0 | ||
|
|
a63ccce73d | ||
|
|
ac56806aff | ||
|
|
0462e1b0e7 | ||
|
|
907f2100ec | ||
|
|
e0a3db5651 | ||
|
|
fca45bc8e1 | ||
|
|
aa95a4c18e | ||
|
|
942f7db3db | ||
|
|
aadb23f409 | ||
|
|
f56f567042 | ||
|
|
15a164a052 | ||
|
|
aaaa7a14a3 | ||
|
|
86d639c848 | ||
|
|
00338aa9ca | ||
|
|
74b409d7cf | ||
|
|
8537dcd6de | ||
|
|
528cef02c4 | ||
|
|
8208536cd1 | ||
|
|
0653f8f3af | ||
|
|
e0d702decb | ||
|
|
541ef014ee | ||
|
|
ffc1a7a58b | ||
|
|
9028953625 | ||
|
|
6eb95693a1 | ||
|
|
c3567eb468 | ||
|
|
3c3da4d057 | ||
|
|
0399713e7b | ||
|
|
0af2e9e8ef |
@@ -10,9 +10,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Clone a reference repo for FastVideo parity tests."
|
||||
)
|
||||
parser = argparse.ArgumentParser(description="Clone a reference repo for FastVideo parity tests.")
|
||||
parser.add_argument("repo_url", help="Official reference repository URL")
|
||||
parser.add_argument("target_dir", help="Directory to clone into")
|
||||
parser.add_argument("--branch", help="Branch or tag to clone")
|
||||
@@ -62,9 +60,7 @@ def gitignore_entry_for(target: Path) -> str:
|
||||
try:
|
||||
relative = resolved.relative_to(root)
|
||||
except ValueError as exc:
|
||||
raise ValueError(
|
||||
"--update-gitignore requires target_dir to be under the current directory"
|
||||
) from exc
|
||||
raise ValueError("--update-gitignore requires target_dir to be under the current directory") from exc
|
||||
|
||||
text = relative.as_posix().rstrip("/")
|
||||
return "/" + text + "/"
|
||||
|
||||
@@ -8,14 +8,12 @@ import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Download a HF model snapshot or selected files into a local directory."
|
||||
)
|
||||
description="Download a HF model snapshot or selected files into a local directory.")
|
||||
parser.add_argument("repo_id", help="HF repo id, for example Org/Model")
|
||||
parser.add_argument("local_dir", help="Destination directory")
|
||||
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
|
||||
|
||||
@@ -10,7 +10,6 @@ import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
|
||||
RAW_WEIGHT_SUFFIXES = (".safetensors", ".pt", ".pth", ".ckpt", ".bin")
|
||||
KNOWN_COMPONENTS = {
|
||||
@@ -34,8 +33,7 @@ KNOWN_COMPONENTS = {
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown."
|
||||
)
|
||||
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown.")
|
||||
parser.add_argument("source", help="HF repo id or local weights directory")
|
||||
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
|
||||
parser.add_argument("--revision", help="HF revision to inspect")
|
||||
@@ -94,14 +92,12 @@ def load_remote_files(
|
||||
) -> list[str]:
|
||||
from huggingface_hub import list_repo_files
|
||||
|
||||
return sorted(
|
||||
list_repo_files(
|
||||
repo_id,
|
||||
repo_type=repo_type,
|
||||
revision=revision,
|
||||
token=token,
|
||||
)
|
||||
)
|
||||
return sorted(list_repo_files(
|
||||
repo_id,
|
||||
repo_type=repo_type,
|
||||
revision=revision,
|
||||
token=token,
|
||||
))
|
||||
|
||||
|
||||
def load_remote_model_index(
|
||||
@@ -215,24 +211,24 @@ def build_result(args: argparse.Namespace) -> dict[str, Any]:
|
||||
"components_seen": components,
|
||||
"file_count": len(files),
|
||||
"file_scan_truncated": truncated,
|
||||
"files_sample": files[: args.sample_limit],
|
||||
"files_sample": files[:args.sample_limit],
|
||||
}
|
||||
|
||||
|
||||
def print_human(result: dict[str, Any]) -> None:
|
||||
for key in (
|
||||
"source",
|
||||
"source_kind",
|
||||
"repo_type",
|
||||
"revision",
|
||||
"token_env",
|
||||
"source_layout",
|
||||
"needs_conversion",
|
||||
"model_index_class",
|
||||
"model_index_diffusers_version",
|
||||
"model_index_error",
|
||||
"file_count",
|
||||
"file_scan_truncated",
|
||||
"source",
|
||||
"source_kind",
|
||||
"repo_type",
|
||||
"revision",
|
||||
"token_env",
|
||||
"source_layout",
|
||||
"needs_conversion",
|
||||
"model_index_class",
|
||||
"model_index_diffusers_version",
|
||||
"model_index_error",
|
||||
"file_count",
|
||||
"file_scan_truncated",
|
||||
):
|
||||
value = result.get(key)
|
||||
if value is not None:
|
||||
|
||||
@@ -18,7 +18,6 @@ import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
os.environ.setdefault("DISABLE_SP", "1")
|
||||
@@ -35,15 +34,10 @@ FASTVIDEO_CONFIG_CLASS = "<FastVideoConfig>" # TODO.
|
||||
FASTVIDEO_MODEL_MODULE = "fastvideo.models.<bucket>.<module>" # TODO.
|
||||
FASTVIDEO_MODEL_CLASS = "<FastVideoModel>" # TODO.
|
||||
|
||||
OFFICIAL_REF_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>")
|
||||
)
|
||||
LOCAL_WEIGHTS_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY)
|
||||
)
|
||||
CONVERTED_WEIGHTS_DIR = Path(
|
||||
os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / FAMILY)
|
||||
)
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>"))
|
||||
LOCAL_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY))
|
||||
CONVERTED_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR",
|
||||
REPO_ROOT / "converted_weights" / FAMILY))
|
||||
|
||||
|
||||
def _resolve_hf_token() -> str | None:
|
||||
@@ -99,18 +93,14 @@ def _load_official_model(device: torch.device, dtype: torch.dtype) -> torch.nn.M
|
||||
model = OfficialClass() # TODO: pass official config kwargs.
|
||||
state_dict = {} # TODO: load official state dict from LOCAL_WEIGHTS_DIR.
|
||||
missing, unexpected = model.load_state_dict(state_dict, strict=True)
|
||||
assert not missing and not unexpected, (
|
||||
f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
|
||||
)
|
||||
assert not missing and not unexpected, (f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
"""Load the FastVideo component with the same tensor content."""
|
||||
if not CONVERTED_WEIGHTS_DIR.exists() and not LOCAL_WEIGHTS_DIR.exists():
|
||||
pytest.skip(
|
||||
f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}"
|
||||
)
|
||||
pytest.skip(f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}")
|
||||
|
||||
# TODO: replace with the bucket-specific FastVideo config/class/loader.
|
||||
# DiT examples:
|
||||
@@ -127,8 +117,7 @@ def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.
|
||||
state_dict = {} # TODO: load converted or directly mapped state dict.
|
||||
missing, unexpected = model.load_state_dict(state_dict, strict=True)
|
||||
assert not missing and not unexpected, (
|
||||
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
|
||||
)
|
||||
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
@@ -187,11 +176,9 @@ def test_component_parity():
|
||||
|
||||
assert official_out.shape == fastvideo_out.shape
|
||||
diff = (official_out - fastvideo_out).abs()
|
||||
print(
|
||||
f"official abs_mean={official_out.abs().mean().item():.6f} "
|
||||
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
|
||||
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}"
|
||||
)
|
||||
print(f"official abs_mean={official_out.abs().mean().item():.6f} "
|
||||
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
|
||||
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
|
||||
|
||||
# TODO: pick tolerance by scope:
|
||||
# - single block / same kernel: 1e-4
|
||||
|
||||
@@ -27,7 +27,6 @@ try:
|
||||
except ImportError: # pragma: no cover - optional local conversion dependency
|
||||
snapshot_download = None
|
||||
|
||||
|
||||
# TODO: fill with authoritative component prefixes for monolithic checkpoints.
|
||||
# Example: {"model.model.": "transformer", "pretransform.model.": "vae"}
|
||||
COMPONENT_PREFIXES: dict[str, str] = {}
|
||||
@@ -47,10 +46,7 @@ SKIP_PATTERNS: tuple[str, ...] = ()
|
||||
|
||||
|
||||
def _hf_token() -> str | None:
|
||||
return (
|
||||
os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN")
|
||||
or os.environ.get("HF_API_KEY")
|
||||
)
|
||||
return (os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN") or os.environ.get("HF_API_KEY"))
|
||||
|
||||
|
||||
def resolve_src(src: str, revision: str | None) -> Path:
|
||||
@@ -95,11 +91,10 @@ def apply_mapping(key: str) -> str | None:
|
||||
return key
|
||||
|
||||
|
||||
def split_monolithic(
|
||||
state: dict[str, torch.Tensor],
|
||||
) -> dict[str, OrderedDict[str, torch.Tensor]]:
|
||||
def split_monolithic(state: dict[str, torch.Tensor], ) -> dict[str, OrderedDict[str, torch.Tensor]]:
|
||||
components: dict[str, OrderedDict[str, torch.Tensor]] = {
|
||||
name: OrderedDict() for name in set(COMPONENT_PREFIXES.values())
|
||||
name: OrderedDict()
|
||||
for name in set(COMPONENT_PREFIXES.values())
|
||||
}
|
||||
intentionally_skipped: list[str] = []
|
||||
unowned: list[str] = []
|
||||
@@ -117,10 +112,8 @@ def split_monolithic(
|
||||
unowned.append(key)
|
||||
if unowned:
|
||||
sample = ", ".join(unowned[:10])
|
||||
raise ValueError(
|
||||
f"Unowned monolithic keys: {len(unowned)}. "
|
||||
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}"
|
||||
)
|
||||
raise ValueError(f"Unowned monolithic keys: {len(unowned)}. "
|
||||
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}")
|
||||
if intentionally_skipped:
|
||||
print(f"Intentionally skipped {len(intentionally_skipped)} keys")
|
||||
return {name: weights for name, weights in components.items() if weights}
|
||||
@@ -143,8 +136,12 @@ def build_component_configs(_src_dir: Path) -> dict[str, dict[str, Any]]:
|
||||
# TODO: emit config content accepted by FastVideo loaders. Most components use
|
||||
# config.json; schedulers use scheduler_config.json.
|
||||
return {
|
||||
"transformer": {"_class_name": "<FastVideoTransformerClass>"},
|
||||
"vae": {"_class_name": "<FastVideoVAEClass>"},
|
||||
"transformer": {
|
||||
"_class_name": "<FastVideoTransformerClass>"
|
||||
},
|
||||
"vae": {
|
||||
"_class_name": "<FastVideoVAEClass>"
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -177,19 +174,13 @@ def build_model_index(
|
||||
}
|
||||
if revision:
|
||||
index["_fastvideo_converted_revision"] = revision
|
||||
return {
|
||||
key: value
|
||||
for key, value in index.items()
|
||||
if key.startswith("_") or key in available_components
|
||||
}
|
||||
return {key: value for key, value in index.items() if key.startswith("_") or key in available_components}
|
||||
|
||||
|
||||
def validate_component_configs(configs: dict[str, dict[str, Any]]) -> None:
|
||||
# TODO: instantiate each FastVideo config and call update_model_arch(...) or
|
||||
# update_model_config(...) with this JSON so unknown emitted keys fail here.
|
||||
placeholder_configs = [
|
||||
name for name, config in configs.items() if "<" in json.dumps(config)
|
||||
]
|
||||
placeholder_configs = [name for name, config in configs.items() if "<" in json.dumps(config)]
|
||||
if placeholder_configs:
|
||||
raise ValueError(f"Replace config placeholders for: {placeholder_configs}")
|
||||
|
||||
@@ -201,9 +192,7 @@ def verify_conversion(
|
||||
del dst_dir, components
|
||||
# TODO: load each emitted stateful component through its production loader and
|
||||
# assert strict load, or document exact allowed missing/unexpected keys.
|
||||
raise NotImplementedError(
|
||||
"Implement production config validation and strict-load checks"
|
||||
)
|
||||
raise NotImplementedError("Implement production config validation and strict-load checks")
|
||||
|
||||
|
||||
def write_component(
|
||||
@@ -216,9 +205,7 @@ def write_component(
|
||||
if component_dir.exists() and any(component_dir.iterdir()):
|
||||
shutil.rmtree(component_dir)
|
||||
component_dir.mkdir(parents=True, exist_ok=True)
|
||||
save_file(
|
||||
dict(state), str(component_dir / "diffusion_pytorch_model.safetensors")
|
||||
)
|
||||
save_file(dict(state), str(component_dir / "diffusion_pytorch_model.safetensors"))
|
||||
if config is not None:
|
||||
config_path = component_dir / config_filename(name)
|
||||
with config_path.open("w", encoding="utf-8") as f:
|
||||
@@ -261,9 +248,7 @@ def convert(
|
||||
|
||||
if layout in {"monolithic", "raw_official"}:
|
||||
# TODO: replace model.safetensors with the official monolithic file name.
|
||||
components = split_monolithic(
|
||||
load_checkpoint(default_monolithic_checkpoint(src_path))
|
||||
)
|
||||
components = split_monolithic(load_checkpoint(default_monolithic_checkpoint(src_path)))
|
||||
elif layout in {"separate_components", "mixed"}:
|
||||
if not src_path.is_dir():
|
||||
raise ValueError(f"{layout} layout requires a source directory: {src_path}")
|
||||
@@ -271,9 +256,7 @@ def convert(
|
||||
else:
|
||||
raise ValueError(f"Unsupported template layout: {layout}")
|
||||
|
||||
copied = (
|
||||
copy_passthrough(src_path, dst_dir) if src_path.is_dir() else []
|
||||
)
|
||||
copied = (copy_passthrough(src_path, dst_dir) if src_path.is_dir() else [])
|
||||
configs = build_component_configs(src_path if src_path.is_dir() else src_path.parent)
|
||||
validate_component_configs(configs)
|
||||
for name, state in components.items():
|
||||
@@ -289,9 +272,7 @@ def convert(
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--src", required=True, help="HF repo id, local dir, or checkpoint path"
|
||||
)
|
||||
parser.add_argument("--src", required=True, help="HF repo id, local dir, or checkpoint path")
|
||||
parser.add_argument("--revision", help="HF branch, tag, or commit for repo sources")
|
||||
parser.add_argument(
|
||||
"--dst",
|
||||
|
||||
@@ -24,8 +24,8 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
|
||||
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
|
||||
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
|
||||
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
|
||||
DRILL_LAYER_ENV: str = "<FAMILY>_DEBUG_DRILL_LAYER"
|
||||
HYPOTHESIS_ENV: str = "<FAMILY>_DEBUG_PATCH_<HYPOTHESIS>"
|
||||
REL_THRESHOLD: float = 0.005 # 0.5% abs_mean drift flags a block as divergent
|
||||
@@ -94,6 +94,7 @@ def _attach_block_hooks(
|
||||
handles: list[Any] = []
|
||||
|
||||
def _hook(name: str):
|
||||
|
||||
def fn(_module, _inputs, outputs):
|
||||
t = outputs[0] if isinstance(outputs, tuple) else outputs
|
||||
if not torch.is_tensor(t):
|
||||
@@ -101,6 +102,7 @@ def _attach_block_hooks(
|
||||
log.append({"side": label, **_stat(name, t)})
|
||||
if tensors is not None:
|
||||
tensors[name] = t.detach().float().cpu()
|
||||
|
||||
return fn
|
||||
|
||||
def _pre_hook(name: str):
|
||||
@@ -114,6 +116,7 @@ def _attach_block_hooks(
|
||||
log.append({"side": label, **_stat(key, t)})
|
||||
if tensors is not None:
|
||||
tensors[key] = t.detach().float().cpu()
|
||||
|
||||
return fn
|
||||
|
||||
# TODO: adapt attribute paths to your model. Remove adapter block if absent.
|
||||
@@ -131,43 +134,21 @@ def _attach_block_hooks(
|
||||
# magi-human uses: attention, mlp.pre_norm, mlp.up_gate_proj,
|
||||
# mlp.down_proj (pre+post), mlp, attn_post_norm, mlp_post_norm.
|
||||
if hasattr(layer, "attention"):
|
||||
handles.append(
|
||||
layer.attention.register_forward_hook(_hook(f"{tag}.attention"))
|
||||
)
|
||||
handles.append(layer.attention.register_forward_hook(_hook(f"{tag}.attention")))
|
||||
if hasattr(layer, "mlp"):
|
||||
mlp = layer.mlp
|
||||
if hasattr(mlp, "pre_norm"):
|
||||
handles.append(
|
||||
mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm"))
|
||||
)
|
||||
handles.append(mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm")))
|
||||
if hasattr(mlp, "up_gate_proj"):
|
||||
handles.append(
|
||||
mlp.up_gate_proj.register_forward_hook(
|
||||
_hook(f"{tag}.mlp.up_gate_proj")
|
||||
)
|
||||
)
|
||||
handles.append(mlp.up_gate_proj.register_forward_hook(_hook(f"{tag}.mlp.up_gate_proj")))
|
||||
if hasattr(mlp, "down_proj"):
|
||||
handles.append(
|
||||
mlp.down_proj.register_forward_pre_hook(
|
||||
_pre_hook(f"{tag}.mlp.down_proj")
|
||||
)
|
||||
)
|
||||
handles.append(
|
||||
mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj"))
|
||||
)
|
||||
handles.append(mlp.down_proj.register_forward_pre_hook(_pre_hook(f"{tag}.mlp.down_proj")))
|
||||
handles.append(mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj")))
|
||||
handles.append(mlp.register_forward_hook(_hook(f"{tag}.mlp")))
|
||||
if hasattr(layer, "attn_post_norm"):
|
||||
handles.append(
|
||||
layer.attn_post_norm.register_forward_hook(
|
||||
_hook(f"{tag}.attn_post_norm")
|
||||
)
|
||||
)
|
||||
handles.append(layer.attn_post_norm.register_forward_hook(_hook(f"{tag}.attn_post_norm")))
|
||||
if hasattr(layer, "mlp_post_norm"):
|
||||
handles.append(
|
||||
layer.mlp_post_norm.register_forward_hook(
|
||||
_hook(f"{tag}.mlp_post_norm")
|
||||
)
|
||||
)
|
||||
handles.append(layer.mlp_post_norm.register_forward_hook(_hook(f"{tag}.mlp_post_norm")))
|
||||
return handles
|
||||
|
||||
|
||||
@@ -193,11 +174,9 @@ def _write_log(entries: list[dict], path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(path, "w") as f:
|
||||
for e in entries:
|
||||
f.write(
|
||||
f"{e['name']} {e['shape']} "
|
||||
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
|
||||
f"{e['min']:.6f} {e['max']:.6f}\n"
|
||||
)
|
||||
f.write(f"{e['name']} {e['shape']} "
|
||||
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
|
||||
f"{e['min']:.6f} {e['max']:.6f}\n")
|
||||
|
||||
|
||||
def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
@@ -205,9 +184,14 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
return (0, "")
|
||||
if name.startswith(f"L{drill_layer:02d}."):
|
||||
sub_order = {
|
||||
"attention": 0, "attn_post_norm": 1, "mlp.pre_norm": 2,
|
||||
"mlp.up_gate_proj": 3, "mlp.down_proj<in>": 4,
|
||||
"mlp.down_proj": 5, "mlp": 6, "mlp_post_norm": 7,
|
||||
"attention": 0,
|
||||
"attn_post_norm": 1,
|
||||
"mlp.pre_norm": 2,
|
||||
"mlp.up_gate_proj": 3,
|
||||
"mlp.down_proj<in>": 4,
|
||||
"mlp.down_proj": 5,
|
||||
"mlp": 6,
|
||||
"mlp_post_norm": 7,
|
||||
}.get(name.split(".", 1)[1], 9)
|
||||
return (1, f"block[{drill_layer:02d}]", sub_order)
|
||||
if name.startswith("block["):
|
||||
@@ -216,10 +200,8 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
|
||||
|
||||
|
||||
def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
|
||||
hdr = (
|
||||
f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
|
||||
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}"
|
||||
)
|
||||
hdr = (f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
|
||||
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}")
|
||||
print(f"\n{hdr}\n{'-' * len(hdr)}")
|
||||
first_div: int | None = None
|
||||
for name in sorted(by_name.keys(), key=lambda n: _sort_key(n, drill_layer)):
|
||||
@@ -235,11 +217,9 @@ def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
|
||||
flag = " <<< DIVERGE"
|
||||
if first_div is None:
|
||||
first_div = int(name[len("block["):-1])
|
||||
print(
|
||||
f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
|
||||
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
|
||||
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}"
|
||||
)
|
||||
print(f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
|
||||
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
|
||||
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}")
|
||||
return first_div
|
||||
|
||||
|
||||
@@ -255,10 +235,8 @@ def _print_elementwise(up_t: dict[str, torch.Tensor], fv_t: dict[str, torch.Tens
|
||||
continue
|
||||
diff = (a - b).abs()
|
||||
rel = (diff.mean().item() / max(a.abs().mean().item(), 1e-9)) * 100
|
||||
print(
|
||||
f"{name:<30} {str(tuple(a.shape)):<22} "
|
||||
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%"
|
||||
)
|
||||
print(f"{name:<30} {str(tuple(a.shape)):<22} "
|
||||
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
|
||||
@@ -43,12 +43,10 @@ def _add_official_to_path() -> Path:
|
||||
|
||||
def _log_tensor_stats(label: str, tensor: torch.Tensor) -> None:
|
||||
value = tensor.detach().float()
|
||||
print(
|
||||
f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
|
||||
f"dtype={tensor.dtype} device={tensor.device} "
|
||||
f"min={value.min().item():.6f} max={value.max().item():.6f} "
|
||||
f"mean={value.mean().item():.6f} std={value.std().item():.6f}"
|
||||
)
|
||||
print(f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
|
||||
f"dtype={tensor.dtype} device={tensor.device} "
|
||||
f"min={value.min().item():.6f} max={value.max().item():.6f} "
|
||||
f"mean={value.mean().item():.6f} std={value.std().item():.6f}")
|
||||
|
||||
|
||||
def _extract_tensor(output: Any, key: str) -> torch.Tensor:
|
||||
@@ -73,10 +71,8 @@ def _run_official_pipeline(
|
||||
device: torch.device,
|
||||
) -> Any:
|
||||
del official_path, params, device
|
||||
pytest.skip(
|
||||
"TODO: import the official pipeline/factory, load official weights, "
|
||||
"run with params, and return the comparison target."
|
||||
)
|
||||
pytest.skip("TODO: import the official pipeline/factory, load official weights, "
|
||||
"run with params, and return the comparison target.")
|
||||
|
||||
|
||||
def _run_fastvideo_pipeline(model_path: Path, params: dict[str, Any]) -> Any:
|
||||
@@ -146,8 +142,6 @@ def test_todo_model_family_pipeline_official_parity() -> None:
|
||||
assert official_tensor.shape == fastvideo_tensor.shape
|
||||
|
||||
diff = (official_tensor - fastvideo_tensor).abs()
|
||||
print(
|
||||
f"diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
|
||||
)
|
||||
print(f"diff max={diff.max().item():.6f} "
|
||||
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}")
|
||||
assert_close(fastvideo_tensor, official_tensor, atol=1e-2, rtol=1e-2)
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
name: macOS MLX Smoke
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
paths:
|
||||
- ".github/workflows/ci-macos-mlx.yml"
|
||||
- "fastvideo/mlx_runtime/**"
|
||||
- "fastvideo/tests/mlx/**"
|
||||
- "fastvideo/tests/platforms/test_mps_vsa_error.py"
|
||||
- "fastvideo/platforms/mps.py"
|
||||
- "fastvideo/platforms/__init__.py"
|
||||
- "fastvideo/__init__.py"
|
||||
- "examples/inference/basic/mlx_*.py"
|
||||
- "fastvideo/benchmarks/mlx_*.py"
|
||||
- "pyproject.toml"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: macos-mlx-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
mlx-smoke:
|
||||
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
|
||||
runs-on: macos-15
|
||||
timeout-minutes: 25
|
||||
env:
|
||||
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
|
||||
TOKENIZERS_PARALLELISM: "false"
|
||||
MASTER_ADDR: localhost
|
||||
MASTER_PORT: "29513"
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
cache: pip
|
||||
|
||||
- uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install lightweight MLX smoke dependencies
|
||||
run: |
|
||||
uv pip install --system \
|
||||
--index-url https://download.pytorch.org/whl/cpu \
|
||||
torch==2.11.0 torchvision torchaudio
|
||||
uv pip install --system \
|
||||
pytest numpy scipy pillow imageio einops cloudpickle filelock \
|
||||
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru mlx \
|
||||
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
|
||||
|
||||
- name: Show Apple runtime
|
||||
run: |
|
||||
python - <<'PY'
|
||||
import platform
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
print("machine:", platform.machine())
|
||||
print("processor:", platform.processor())
|
||||
print("mlx default device:", mx.default_device())
|
||||
memory_size = mx.metal.device_info().get("memory_size") if mx.metal.is_available() else "metal unavailable"
|
||||
print("mlx memory_size:", memory_size)
|
||||
print("torch:", torch.__version__)
|
||||
print("torch mps available:", torch.backends.mps.is_available())
|
||||
PY
|
||||
|
||||
- name: Run MLX smoke tests
|
||||
run: |
|
||||
python -m pytest \
|
||||
fastvideo/tests/mlx/test_dmd_sampling.py \
|
||||
fastvideo/tests/mlx/test_memory_limits.py \
|
||||
fastvideo/tests/mlx/test_quant_capability.py \
|
||||
fastvideo/tests/mlx/test_mlx_dit_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_compile_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_checkpoint.py \
|
||||
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
|
||||
fastvideo/tests/mlx/test_taehv_decode.py \
|
||||
fastvideo/tests/mlx/test_frame_upsample.py \
|
||||
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
|
||||
fastvideo/tests/mlx/test_mlx_refine.py \
|
||||
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
|
||||
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
|
||||
fastvideo/tests/mlx/test_wan22_sample.py \
|
||||
fastvideo/tests/mlx/test_windowed_attention.py \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
|
||||
fastvideo/tests/platforms/test_mps_vsa_error.py \
|
||||
-q
|
||||
|
||||
# Same tests on MLX's CPU backend. Hosted macOS runners are scarce and
|
||||
# slower to schedule; this Linux job gives fast PR signal on the identical
|
||||
# graph (the parity tests were designed to be backend-agnostic), while the
|
||||
# macOS job above stays the source of truth for Metal behavior.
|
||||
mlx-smoke-linux-cpu:
|
||||
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
env:
|
||||
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
|
||||
TOKENIZERS_PARALLELISM: "false"
|
||||
MASTER_ADDR: localhost
|
||||
MASTER_PORT: "29513"
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
cache: pip
|
||||
|
||||
- uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install lightweight MLX smoke dependencies (CPU backend)
|
||||
run: |
|
||||
uv pip install --system \
|
||||
--index-url https://download.pytorch.org/whl/cpu \
|
||||
torch==2.11.0 torchvision torchaudio
|
||||
uv pip install --system \
|
||||
pytest numpy scipy pillow imageio einops cloudpickle filelock \
|
||||
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru "mlx[cpu]" \
|
||||
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
|
||||
|
||||
- name: Run MLX smoke tests (CPU backend)
|
||||
run: |
|
||||
python -m pytest \
|
||||
fastvideo/tests/mlx/test_dmd_sampling.py \
|
||||
fastvideo/tests/mlx/test_memory_limits.py \
|
||||
fastvideo/tests/mlx/test_quant_capability.py \
|
||||
fastvideo/tests/mlx/test_mlx_dit_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_compile_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_checkpoint.py \
|
||||
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
|
||||
fastvideo/tests/mlx/test_taehv_decode.py \
|
||||
fastvideo/tests/mlx/test_frame_upsample.py \
|
||||
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
|
||||
fastvideo/tests/mlx/test_mlx_refine.py \
|
||||
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
|
||||
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
|
||||
fastvideo/tests/mlx/test_wan22_sample.py \
|
||||
fastvideo/tests/mlx/test_windowed_attention.py \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
|
||||
fastvideo/tests/platforms/test_mps_vsa_error.py \
|
||||
-q
|
||||
@@ -6,6 +6,7 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
@@ -16,6 +17,7 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
|
||||
@@ -23,6 +23,7 @@ Miniconda3-latest-Linux-x86_64.sh
|
||||
*validation/
|
||||
data/
|
||||
outputs/
|
||||
outputs_audio/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
|
||||
@@ -22,6 +22,7 @@ repos:
|
||||
hooks:
|
||||
- id: yapf
|
||||
args: [--in-place, --verbose]
|
||||
language_version: python3.12
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.12
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
|
||||
|
||||
## NEWS
|
||||
- `2026/08/19`: FastVideo now supports MLX on Apple Silicon with [FastMetal-QAD](https://huggingface.co/collections/FastVideo/fastmetal), a family of 1.3B, 5B, and 14B models optimized for Mac—follow the [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/) and read the [Blog](https://haoailab.com/blogs/fastmetal/).
|
||||
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
|
||||
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
|
||||
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
|
||||
@@ -62,6 +63,11 @@ UV_TORCH_BACKEND=cu126 uv pip install fastvideo
|
||||
Use `UV_TORCH_BACKEND=cu130` on CUDA 13. Apple silicon users should follow the
|
||||
[MPS installation guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
|
||||
|
||||
> **On an Apple Silicon Mac?** FastVideo runs FastWan text-to-video natively
|
||||
> through an MLX runtime — a 5-second 480p clip generated locally, no cloud,
|
||||
> no discrete GPU. Install with `uv pip install -e '.[mlx]'` and follow the
|
||||
> [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
|
||||
|
||||
> **On an NVIDIA DGX Spark (GB10 / ARM64 + CUDA 13)?** There's no prebuilt ARM wheel for the FastVideo CUDA kernel, so it's an editable from-source install (`UV_TORCH_BACKEND=cu130 uv pip install -e .`, which compiles that kernel for you) rather than `UV_TORCH_BACKEND=cu130 uv pip install fastvideo`. A compatible prebuilt ARM64 FlashAttention wheel is available separately. Follow the [DGX Spark install guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/spark/).
|
||||
|
||||
@@ -3,7 +3,6 @@ from __future__ import annotations
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
TESTS_DIR = Path(__file__).resolve().parent
|
||||
DREAMVERSE_PACKAGE_DIR = TESTS_DIR.parent
|
||||
DREAMVERSE_APP_DIR = DREAMVERSE_PACKAGE_DIR.parent
|
||||
|
||||
@@ -5,7 +5,6 @@ from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
SERVER_DIR = Path(__file__).resolve().parents[1]
|
||||
|
||||
|
||||
@@ -53,9 +52,7 @@ def test_config_defaults_to_cerebras_with_parallel_groq_fallback_stage(monkeypat
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.PROMPT_PROVIDER == "cerebras"
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
|
||||
("cerebras", "groq"),
|
||||
)
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
|
||||
assert module.PROMPT_PROVIDER_PRIORITY == (
|
||||
"cerebras",
|
||||
"groq",
|
||||
@@ -86,9 +83,7 @@ def test_config_ignores_legacy_groq_primary_override(monkeypatch):
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.PROMPT_PROVIDER == "cerebras"
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
|
||||
("cerebras", "groq"),
|
||||
)
|
||||
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
|
||||
assert module.PROMPT_PROVIDER_PRIORITY == (
|
||||
"cerebras",
|
||||
"groq",
|
||||
@@ -106,24 +101,17 @@ def test_config_uses_local_overlay_paths_when_devtools_enabled(monkeypatch, tmp_
|
||||
|
||||
assert module.DEVTOOLS_ENABLED is True
|
||||
assert module.FRONTEND_ROOT.as_posix().endswith("apps/dreamverse/web")
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith(
|
||||
"dreamverse/prompts.local/next_segment_system_prompt.md"
|
||||
)
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith("dreamverse/prompts.local/next_segment_system_prompt.md")
|
||||
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
|
||||
"dreamverse/prompts/next_segment_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts/next_segment_system_prompt.md")
|
||||
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_PATH.endswith(
|
||||
"dreamverse/prompts.local/rewrite_user_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts.local/rewrite_user_system_prompt.md")
|
||||
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
|
||||
"dreamverse/prompts/rewrite_user_system_prompt.md"
|
||||
)
|
||||
"dreamverse/prompts/rewrite_user_system_prompt.md")
|
||||
assert module.CURATED_PRESETS_FILE_PATH.endswith(
|
||||
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json")
|
||||
assert module.CURATED_PRESETS_FALLBACK_FILE_PATH.endswith(
|
||||
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json")
|
||||
assert module.FRONTEND_STATIC_DIR_CANDIDATES[:2] == (
|
||||
str(module.FRONTEND_ROOT / "out"),
|
||||
str(module.FRONTEND_ROOT / "dist"),
|
||||
|
||||
@@ -9,6 +9,7 @@ from fastapi.testclient import TestClient
|
||||
import fastvideo.entrypoints.streaming as streaming_entrypoints
|
||||
import pytest
|
||||
|
||||
|
||||
def _install_stack03_import_stubs(monkeypatch):
|
||||
"""Keep entrypoint tests focused while later-stack runtime modules are absent."""
|
||||
if not hasattr(streaming_entrypoints, "build_health_router"):
|
||||
@@ -17,6 +18,7 @@ def _install_stack03_import_stubs(monkeypatch):
|
||||
gpu_pool_stub = types.ModuleType("dreamverse.gpu_pool")
|
||||
|
||||
class GPUPool:
|
||||
|
||||
def __init__(self, _gpu_ids):
|
||||
pass
|
||||
|
||||
@@ -49,6 +51,7 @@ def _install_stack03_import_stubs(monkeypatch):
|
||||
controller_stub = types.ModuleType("dreamverse.session.controller")
|
||||
|
||||
class SessionController:
|
||||
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
@@ -76,13 +79,11 @@ def _run_cli(module, monkeypatch, argv: list[str]) -> list[dict[str, object]]:
|
||||
uvicorn_stub = types.ModuleType("uvicorn")
|
||||
|
||||
def run(app, host: str, port: int) -> None:
|
||||
calls.append(
|
||||
{
|
||||
"app": app,
|
||||
"host": host,
|
||||
"port": port,
|
||||
}
|
||||
)
|
||||
calls.append({
|
||||
"app": app,
|
||||
"host": host,
|
||||
"port": port,
|
||||
})
|
||||
|
||||
uvicorn_stub.run = run
|
||||
monkeypatch.setitem(sys.modules, "uvicorn", uvicorn_stub)
|
||||
@@ -99,13 +100,11 @@ def test_server_cli_defaults_to_local_web_port(monkeypatch):
|
||||
server_main = _import_server_main(monkeypatch)
|
||||
calls = _run_cli(server_main, monkeypatch, ["dreamverse-server"])
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": server_main.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": server_main.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}]
|
||||
|
||||
|
||||
def test_server_cli_allows_explicit_host_and_port(monkeypatch):
|
||||
@@ -116,13 +115,11 @@ def test_server_cli_allows_explicit_host_and_port(monkeypatch):
|
||||
["dreamverse-server", "--host", "127.0.0.1", "--port", "8123"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": server_main.app,
|
||||
"host": "127.0.0.1",
|
||||
"port": 8123,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": server_main.app,
|
||||
"host": "127.0.0.1",
|
||||
"port": 8123,
|
||||
}]
|
||||
|
||||
|
||||
def test_server_does_not_expose_backend_source_as_static_assets(monkeypatch):
|
||||
@@ -142,13 +139,11 @@ def test_mock_server_cli_defaults_to_local_web_port(monkeypatch):
|
||||
["dreamverse-mock-server"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8009,
|
||||
}]
|
||||
|
||||
|
||||
def test_mock_server_cli_updates_latency(monkeypatch):
|
||||
@@ -161,13 +156,11 @@ def test_mock_server_cli_updates_latency(monkeypatch):
|
||||
["dreamverse-mock-server", "--latency", "321", "--port", "8111"],
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8111,
|
||||
}
|
||||
]
|
||||
assert calls == [{
|
||||
"app": mock_server.app,
|
||||
"host": "0.0.0.0",
|
||||
"port": 8111,
|
||||
}]
|
||||
assert mock_server.LATENCY_MS == 321
|
||||
finally:
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
|
||||
@@ -7,7 +7,6 @@ from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
import dreamverse.gpu_pool as gpu_pool
|
||||
|
||||
|
||||
@@ -85,9 +84,7 @@ def test_send_command_raises_on_worker_death():
|
||||
cmd_q = ctx.Queue()
|
||||
resp_q = ctx.Queue()
|
||||
|
||||
proc = ctx.Process(
|
||||
target=_child_consume_and_exit, args=(cmd_q, resp_q)
|
||||
)
|
||||
proc = ctx.Process(target=_child_consume_and_exit, args=(cmd_q, resp_q))
|
||||
proc.start()
|
||||
|
||||
# Wait for the spawn child to fully boot. Allow generous time —
|
||||
|
||||
@@ -7,7 +7,7 @@ ALLOWED_PREFIXES = (
|
||||
"fastvideo.entrypoints.video_generator",
|
||||
"fastvideo.configs",
|
||||
)
|
||||
ALLOWED_EXACT = ("fastvideo",)
|
||||
ALLOWED_EXACT = ("fastvideo", )
|
||||
FORBIDDEN_PREFIXES = (
|
||||
"fastvideo.pipelines",
|
||||
"fastvideo.models",
|
||||
@@ -38,19 +38,13 @@ def test_dreamverse_server_imports_only_public_fastvideo_surfaces() -> None:
|
||||
except SyntaxError as task_exc:
|
||||
raise AssertionError(f"Failed to parse {path}") from task_exc
|
||||
for node in ast.walk(tree):
|
||||
names = (
|
||||
[a.name for a in node.names] if isinstance(node, ast.Import)
|
||||
else [node.module] if isinstance(node, ast.ImportFrom) and node.module
|
||||
else []
|
||||
)
|
||||
names = ([a.name for a in node.names] if isinstance(node, ast.Import) else
|
||||
[node.module] if isinstance(node, ast.ImportFrom) and node.module else [])
|
||||
for name in names:
|
||||
if not name:
|
||||
continue
|
||||
rel_path = str(path.relative_to(root))
|
||||
if (
|
||||
name.startswith(FORBIDDEN_PREFIXES)
|
||||
and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS
|
||||
):
|
||||
if (name.startswith(FORBIDDEN_PREFIXES) and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS):
|
||||
bad.append((str(path.relative_to(root)), getattr(node, "lineno", 0), name))
|
||||
|
||||
assert bad == [], f"Forbidden internal imports: {bad}"
|
||||
|
||||
@@ -6,7 +6,6 @@ import os
|
||||
|
||||
from fastapi import WebSocketDisconnect
|
||||
|
||||
|
||||
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
|
||||
os.environ.setdefault("GROQ_API_KEY", "dummy")
|
||||
|
||||
@@ -14,6 +13,7 @@ import dreamverse.mock_server as mock_server
|
||||
|
||||
|
||||
class _FakeWebSocket:
|
||||
|
||||
def __init__(self, messages: list[tuple[float, dict[str, object]]]):
|
||||
self._messages = messages
|
||||
self._index = 0
|
||||
@@ -49,34 +49,34 @@ def test_mock_server_matches_current_single5s_protocol():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "simple_prompt_1",
|
||||
"curated_prompts": ["selected prompt"],
|
||||
"single_clip_mode": True,
|
||||
"enhancement_enabled": False,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.01,
|
||||
{
|
||||
"type": "simple_generate",
|
||||
"preset_id": "simple_custom_prompt",
|
||||
"prompt_id": "simple_custom_prompt",
|
||||
"prompt": "custom prompt",
|
||||
"enhancement_enabled": True,
|
||||
"initial_image": None,
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "simple_prompt_1",
|
||||
"curated_prompts": ["selected prompt"],
|
||||
"single_clip_mode": True,
|
||||
"enhancement_enabled": False,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.01,
|
||||
{
|
||||
"type": "simple_generate",
|
||||
"preset_id": "simple_custom_prompt",
|
||||
"prompt_id": "simple_custom_prompt",
|
||||
"prompt": "custom prompt",
|
||||
"enhancement_enabled": True,
|
||||
"initial_image": None,
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -92,24 +92,14 @@ def test_mock_server_matches_current_single5s_protocol():
|
||||
assert message_types.count("ltx2_stream_complete") == 2
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
gpu_assigned_event = next(
|
||||
payload for payload in ws.sent_json if payload["type"] == "gpu_assigned"
|
||||
)
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
gpu_assigned_event = next(payload for payload in ws.sent_json if payload["type"] == "gpu_assigned")
|
||||
assert gpu_assigned_event["session_timeout"] == mock_server.SESSION_TIMEOUT_SECONDS
|
||||
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
|
||||
assert segment_start_events[0]["prompt"] == "selected prompt"
|
||||
assert segment_start_events[1]["prompt"] == "custom prompt"
|
||||
|
||||
step_complete_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "step_complete"
|
||||
]
|
||||
step_complete_events = [payload for payload in ws.sent_json if payload["type"] == "step_complete"]
|
||||
assert len(step_complete_events) == 2
|
||||
assert step_complete_events[0]["latency_ms"] == {
|
||||
"total": 121.0,
|
||||
@@ -134,29 +124,29 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
|
||||
mock_server.LATENCY_MS = 1
|
||||
mock_server.GENERATION_SEGMENT_CAP = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "start a new rollout",
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "start a new rollout",
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -166,11 +156,7 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
|
||||
assert "generation_cap_reached" not in message_types
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
|
||||
assert segment_start_events[0]["prompt"] == "segment one"
|
||||
assert segment_start_events[1]["prompt"] == "segment one [start a new rollout]"
|
||||
@@ -187,54 +173,40 @@ def test_mock_server_rewrite_during_active_segment_restarts_from_first_rewritten
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 100
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one", "segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "restart from rewrite",
|
||||
},
|
||||
),
|
||||
(0.40, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one", "segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
0.02,
|
||||
{
|
||||
"type": "rewrite_seed_prompts",
|
||||
"rewrite_instruction": "restart from rewrite",
|
||||
},
|
||||
),
|
||||
(0.40, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
|
||||
"segment one",
|
||||
"segment one [restart from rewrite]",
|
||||
]
|
||||
assert all(
|
||||
payload["prompt"] != "segment two"
|
||||
for payload in segment_start_events[1:]
|
||||
)
|
||||
reset_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload.get("type") == "seed_prompts_reset_applied"
|
||||
]
|
||||
assert any(
|
||||
payload.get("reason") == "rewrite_during_generation"
|
||||
for payload in reset_events
|
||||
)
|
||||
assert all(payload["prompt"] != "segment two" for payload in segment_start_events[1:])
|
||||
reset_events = [payload for payload in ws.sent_json if payload.get("type") == "seed_prompts_reset_applied"]
|
||||
assert any(payload.get("reason") == "rewrite_during_generation" for payload in reset_events)
|
||||
finally:
|
||||
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
@@ -247,24 +219,24 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 1
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "custom_editable",
|
||||
"preset_label": "Custom rollout",
|
||||
"curated_prompts": [],
|
||||
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.20, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "custom_editable",
|
||||
"preset_label": "Custom rollout",
|
||||
"curated_prompts": [],
|
||||
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.20, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -276,15 +248,9 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
|
||||
assert "ltx2_stream_start" in message_types
|
||||
assert "prompt_sources_blocked" not in message_types
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert segment_start_events
|
||||
assert segment_start_events[0]["prompt"] == (
|
||||
"A moonbase corridor thriller with flooding [segment 1]"
|
||||
)
|
||||
assert segment_start_events[0]["prompt"] == ("A moonbase corridor thriller with flooding [segment 1]")
|
||||
finally:
|
||||
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
|
||||
mock_server.LATENCY_MS = old_latency_ms
|
||||
@@ -297,35 +263,37 @@ def test_mock_server_can_start_new_project_without_reconnecting():
|
||||
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
|
||||
mock_server.LATENCY_MS = 40
|
||||
|
||||
ws = _FakeWebSocket(
|
||||
[
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.02, {"type": "end_project_keep_session"}),
|
||||
(
|
||||
0.20,
|
||||
{
|
||||
"type": "project_init_v1",
|
||||
"preset_id": "test_preset_2",
|
||||
"preset_label": "Test Preset 2",
|
||||
"curated_prompts": ["segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.40, {"type": "leave"}),
|
||||
]
|
||||
)
|
||||
ws = _FakeWebSocket([
|
||||
(
|
||||
0.0,
|
||||
{
|
||||
"type": "session_init_v2",
|
||||
"preset_id": "test_preset",
|
||||
"curated_prompts": ["segment one"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.02, {
|
||||
"type": "end_project_keep_session"
|
||||
}),
|
||||
(
|
||||
0.20,
|
||||
{
|
||||
"type": "project_init_v1",
|
||||
"preset_id": "test_preset_2",
|
||||
"preset_label": "Test Preset 2",
|
||||
"curated_prompts": ["segment two"],
|
||||
"enhancement_enabled": True,
|
||||
"auto_extension_enabled": False,
|
||||
"loop_generation_enabled": False,
|
||||
},
|
||||
),
|
||||
(0.40, {
|
||||
"type": "leave"
|
||||
}),
|
||||
])
|
||||
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
|
||||
@@ -336,16 +304,11 @@ def test_mock_server_can_start_new_project_without_reconnecting():
|
||||
|
||||
project_idle_index = message_types.index("project_idle")
|
||||
stream_start_indexes = [
|
||||
index for index, message_type in enumerate(message_types)
|
||||
if message_type == "ltx2_stream_start"
|
||||
index for index, message_type in enumerate(message_types) if message_type == "ltx2_stream_start"
|
||||
]
|
||||
assert stream_start_indexes[0] < project_idle_index < stream_start_indexes[1]
|
||||
|
||||
segment_start_events = [
|
||||
payload
|
||||
for payload in ws.sent_json
|
||||
if payload["type"] == "ltx2_segment_start"
|
||||
]
|
||||
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
|
||||
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
|
||||
"segment one",
|
||||
"segment two",
|
||||
|
||||
@@ -6,7 +6,6 @@ import os
|
||||
import re
|
||||
import time
|
||||
|
||||
|
||||
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
|
||||
os.environ.setdefault("GROQ_API_KEY", "dummy")
|
||||
|
||||
@@ -22,6 +21,7 @@ from dreamverse.prompt_enhancer import (
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self._payload = payload
|
||||
|
||||
@@ -30,6 +30,7 @@ class _FakeResponse:
|
||||
|
||||
|
||||
class _FakeSyncCompletions:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self._payload = payload
|
||||
|
||||
@@ -38,6 +39,7 @@ class _FakeSyncCompletions:
|
||||
|
||||
|
||||
class _FakeSyncClient:
|
||||
|
||||
def __init__(self, payload: dict):
|
||||
self.chat = type(
|
||||
"_FakeChat",
|
||||
@@ -47,6 +49,7 @@ class _FakeSyncClient:
|
||||
|
||||
|
||||
class _DelayedSyncCompletions:
|
||||
|
||||
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
|
||||
self._payload = payload
|
||||
self._delay_s = delay_s
|
||||
@@ -61,29 +64,26 @@ class _DelayedSyncCompletions:
|
||||
|
||||
|
||||
class _DelayedSyncClient:
|
||||
|
||||
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
|
||||
self.chat = type(
|
||||
"_FakeChat",
|
||||
(),
|
||||
{
|
||||
"completions": _DelayedSyncCompletions(
|
||||
payload,
|
||||
delay_s=delay_s,
|
||||
exc=exc,
|
||||
)
|
||||
},
|
||||
{"completions": _DelayedSyncCompletions(
|
||||
payload,
|
||||
delay_s=delay_s,
|
||||
exc=exc,
|
||||
)},
|
||||
)()
|
||||
|
||||
|
||||
def _chat_payload_with_content(content: str) -> dict:
|
||||
return {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"content": content,
|
||||
}
|
||||
"choices": [{
|
||||
"message": {
|
||||
"content": content,
|
||||
}
|
||||
]
|
||||
}]
|
||||
}
|
||||
|
||||
|
||||
@@ -172,6 +172,7 @@ def _build_staged_enhancer(
|
||||
|
||||
|
||||
class _FakeOpenAIClient:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
self.chat = type(
|
||||
@@ -182,6 +183,7 @@ class _FakeOpenAIClient:
|
||||
|
||||
|
||||
class _FakeCerebrasClient:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
self.chat = type(
|
||||
@@ -192,16 +194,12 @@ class _FakeCerebrasClient:
|
||||
|
||||
|
||||
def test_parse_json_response_accepts_fenced_json_with_prose():
|
||||
parsed = _parse_json_response(
|
||||
"Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks."
|
||||
)
|
||||
parsed = _parse_json_response("Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks.")
|
||||
assert parsed == {"segment_prompts": ["A", "B"]}
|
||||
|
||||
|
||||
def test_parse_json_response_extracts_first_embedded_object():
|
||||
parsed = _parse_json_response(
|
||||
"Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)"
|
||||
)
|
||||
parsed = _parse_json_response("Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)")
|
||||
assert parsed == {"segment_prompts": ["A", "B"]}
|
||||
|
||||
|
||||
@@ -268,16 +266,12 @@ def test_build_client_supports_groq_provider(monkeypatch):
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "preset_a"
|
||||
@@ -286,15 +280,12 @@ def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"rewritten_prompts":["A","B"]}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"rewritten_prompts":["A","B"]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "current_rollout"
|
||||
@@ -303,19 +294,14 @@ def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_accepts_segment_dicts_without_top_level_rollout_metadata():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"segments":[{"prompt":"A"},{"text":"B"}]}'
|
||||
)
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segments":[{"prompt":"A"},{"text":"B"}]}'))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
preset_id="preset_a",
|
||||
preset_label="Preset A",
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "preset_a"
|
||||
@@ -329,14 +315,12 @@ def test_rewrite_prompt_sequence_accepts_numbered_prose_output():
|
||||
"The user is asking for a cinematic rewrite.\n\n"
|
||||
"1. A dog bounds across the moon's dusty surface, kicking up silver regolith as it chases a rabbit beneath the black sky.\n"
|
||||
"2. The rabbit darts around a crater rim while the dog lunges after it, Earth glowing blue in the distance.\n"
|
||||
)
|
||||
)
|
||||
))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.rollout_id == "current_rollout"
|
||||
@@ -355,12 +339,10 @@ def test_enhance_prompt_prefers_cerebras_before_groq_fallback():
|
||||
groq_delay_s=0.01,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -381,12 +363,10 @@ def test_enhance_prompt_uses_groq_when_cerebras_fails():
|
||||
groq_delay_s=0.01,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -408,13 +388,11 @@ def test_enhance_prompt_can_use_groq_when_cerebras_times_out():
|
||||
enhancer.http_timeout_ms = 50
|
||||
enhancer.default_timeout_ms = 50
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
timeout_ms=50,
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
timeout_ms=50,
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -434,12 +412,10 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
|
||||
groq_delay_s=0.08,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
@@ -453,15 +429,12 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("I cannot comply with JSON right now.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("I cannot comply with JSON right now."))
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
assert result.raw_response_text == "I cannot comply with JSON right now."
|
||||
@@ -473,9 +446,7 @@ def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
|
||||
def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'))
|
||||
captured = {
|
||||
"body": None,
|
||||
"timeout_seconds": None,
|
||||
@@ -486,8 +457,7 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
captured["timeout_seconds"] = timeout_seconds
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
|
||||
),
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'),
|
||||
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}',
|
||||
)
|
||||
|
||||
@@ -502,8 +472,7 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
rewrite_model="gpt-test",
|
||||
rewrite_temperature=0.2,
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert captured["body"]["messages"][0] == {
|
||||
@@ -512,12 +481,12 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
}
|
||||
assert captured["body"]["messages"][1]["role"] == "user"
|
||||
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
|
||||
"mode": "edit_existing_rollout",
|
||||
"request": (
|
||||
"Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."
|
||||
),
|
||||
"user_instruction": "make it cinematic",
|
||||
"mode":
|
||||
"edit_existing_rollout",
|
||||
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."),
|
||||
"user_instruction":
|
||||
"make it cinematic",
|
||||
"current_rollout": {
|
||||
"id": "preset_a",
|
||||
"label": "Preset A",
|
||||
@@ -528,11 +497,8 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
|
||||
|
||||
def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'))
|
||||
captured = {
|
||||
"body": None,
|
||||
}
|
||||
@@ -541,10 +507,8 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
del timeout_seconds
|
||||
captured["body"] = body
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'
|
||||
),
|
||||
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}'),
|
||||
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
|
||||
'"A","B","C","D","E","F"]}',
|
||||
)
|
||||
@@ -560,30 +524,29 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
|
||||
rewrite_model="gpt-test",
|
||||
rewrite_temperature=0.2,
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.prompts == ["A", "B", "C", "D", "E", "F"]
|
||||
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
|
||||
"mode": "new_rollout",
|
||||
"request": (
|
||||
"Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."
|
||||
),
|
||||
"user_instruction": "A moonbase corridor thriller with flooding and red alarms",
|
||||
"desired_segment_count": 6,
|
||||
"rollout_id_hint": "custom_editable",
|
||||
"rollout_label_hint": "Custom rollout",
|
||||
"mode":
|
||||
"new_rollout",
|
||||
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
|
||||
"Keep count and ordering identical."),
|
||||
"user_instruction":
|
||||
"A moonbase corridor thriller with flooding and red alarms",
|
||||
"desired_segment_count":
|
||||
6,
|
||||
"rollout_id_hint":
|
||||
"custom_editable",
|
||||
"rollout_label_hint":
|
||||
"Custom rollout",
|
||||
}
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared system prompt"
|
||||
captured = {
|
||||
"body": None,
|
||||
@@ -593,9 +556,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
del timeout_seconds
|
||||
captured["body"] = body
|
||||
return (
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
),
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'),
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}',
|
||||
)
|
||||
|
||||
@@ -609,8 +570,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
rewrite_instruction="make it cinematic",
|
||||
rewrite_model="gpt-test",
|
||||
system_prompt_override="session specific system prompt",
|
||||
)
|
||||
)
|
||||
))
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert captured["body"]["messages"][0] == {
|
||||
@@ -621,10 +581,7 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
|
||||
|
||||
def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
|
||||
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
|
||||
|
||||
@@ -635,24 +592,17 @@ def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
|
||||
|
||||
def test_resolve_rewrite_new_rollout_system_prompt_prefers_override():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
|
||||
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
|
||||
|
||||
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt(
|
||||
"session specific system prompt"
|
||||
)
|
||||
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt("session specific system prompt")
|
||||
|
||||
assert resolved == "session specific system prompt"
|
||||
|
||||
|
||||
def test_generate_auto_prompt_uses_selected_model():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"next_prompt":"Auto next"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Auto next"}'))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -680,8 +630,7 @@ def test_generate_auto_prompt_uses_selected_model():
|
||||
next_segment_idx=2,
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Auto next"
|
||||
@@ -690,9 +639,7 @@ def test_generate_auto_prompt_uses_selected_model():
|
||||
|
||||
|
||||
def test_enhance_prompt_uses_selected_model():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"next_prompt":"Enhanced next"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Enhanced next"}'))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
@@ -722,8 +669,7 @@ def test_enhance_prompt_uses_selected_model():
|
||||
next_segment_idx=2,
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Enhanced next"
|
||||
@@ -732,9 +678,7 @@ def test_enhance_prompt_uses_selected_model():
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"prompt":"Extended single clip"}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"prompt":"Extended single clip"}'))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
@@ -764,14 +708,12 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
|
||||
|
||||
enhancer._request_content = _fake_request_content # type: ignore[attr-defined]
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"short 5s idea",
|
||||
mode="single_clip",
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
result = asyncio.run(enhancer.enhance_prompt(
|
||||
"short 5s idea",
|
||||
mode="single_clip",
|
||||
model="gpt-alt",
|
||||
timeout_ms=800,
|
||||
))
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.prompt == "Extended single clip"
|
||||
@@ -784,17 +726,15 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
|
||||
"single 5-second LTX-2.3 video clip. Respond with "
|
||||
'valid JSON only as {"prompt": "..."}.' # noqa: E501
|
||||
),
|
||||
"user_prompt": "short 5s idea",
|
||||
"user_prompt":
|
||||
"short 5s idea",
|
||||
}
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_rejects_plain_text_response():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
"Medium shot of a woman by a rainy cafe window as she lifts her "
|
||||
"phone, exhales softly, and the camera makes a slow push in."
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content("Medium shot of a woman by a rainy cafe window as she lifts her "
|
||||
"phone, exhales softly, and the camera makes a slow push in."))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
|
||||
result = asyncio.run(
|
||||
@@ -803,17 +743,14 @@ def test_enhance_prompt_single_clip_rejects_plain_text_response():
|
||||
mode="single_clip",
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No JSON object found in assistant response." in result.error
|
||||
assert result.prompt == ""
|
||||
|
||||
|
||||
def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content('{"segment_prompts":["A","B"]}')
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segment_prompts":["A","B"]}'))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -824,17 +761,14 @@ def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
|
||||
mode="single_clip",
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "Missing prompt string." in (result.error or "")
|
||||
|
||||
|
||||
def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("A cinematic continuation with slow dolly movement.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("A cinematic continuation with slow dolly movement."))
|
||||
enhancer.enhance_system_prompt = "enhance system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -846,17 +780,14 @@ def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
next_segment_idx=2,
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
|
||||
|
||||
def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content("A calm, grounded continuation with subtle motion.")
|
||||
)
|
||||
enhancer = _build_test_enhancer(_chat_payload_with_content("A calm, grounded continuation with subtle motion."))
|
||||
enhancer.auto_system_prompt = "auto system prompt"
|
||||
enhancer.rewrite_model_options = ["gpt-test"]
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
@@ -867,34 +798,30 @@ def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
|
||||
next_segment_idx=2,
|
||||
model="gpt-test",
|
||||
timeout_ms=800,
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert result.prompt == ""
|
||||
assert "No JSON object found in assistant response." in (result.error or "")
|
||||
|
||||
|
||||
def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
|
||||
enhancer = _build_test_enhancer(
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "length",
|
||||
"message": {
|
||||
"content": [],
|
||||
"refusal": None,
|
||||
},
|
||||
}
|
||||
],
|
||||
"usage": {"completion_tokens": 0},
|
||||
}
|
||||
)
|
||||
enhancer = _build_test_enhancer({
|
||||
"choices": [{
|
||||
"finish_reason": "length",
|
||||
"message": {
|
||||
"content": [],
|
||||
"refusal": None,
|
||||
},
|
||||
}],
|
||||
"usage": {
|
||||
"completion_tokens": 0
|
||||
},
|
||||
})
|
||||
result = asyncio.run(
|
||||
enhancer.rewrite_prompt_sequence(
|
||||
["prompt one", "prompt two"],
|
||||
rewrite_instruction="make it cinematic",
|
||||
)
|
||||
)
|
||||
))
|
||||
assert result.fallback_used is True
|
||||
assert "No rewrite segment prompts found in assistant response." in (result.error or "")
|
||||
assert isinstance(result.raw_response_text, str)
|
||||
@@ -903,10 +830,7 @@ def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
|
||||
|
||||
def test_get_rewrite_model_config_returns_fixed_defaults():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.rewrite_default_model = "gpt-oss-120b"
|
||||
enhancer.rewrite_model_options = ["gpt-oss-120b"]
|
||||
|
||||
@@ -918,10 +842,7 @@ def test_get_rewrite_model_config_returns_fixed_defaults():
|
||||
|
||||
def test_get_prompt_config_includes_auto_extension_prompt():
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
enhancer.enhance_system_prompt_path = "/tmp/next.md"
|
||||
enhancer.auto_system_prompt_path = "/tmp/auto.md"
|
||||
enhancer.rewrite_all_system_prompt_path = "/tmp/rewrite.md"
|
||||
@@ -948,19 +869,14 @@ def test_get_prompt_config_includes_auto_extension_prompt():
|
||||
|
||||
def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
rewrite_fallback_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
rewrite_fallback_path.write_text("rewrite prompt\n", encoding="utf-8")
|
||||
next_path = tmp_path / "next.md"
|
||||
next_path.write_text("next prompt\n", encoding="utf-8")
|
||||
auto_path = tmp_path / "auto.md"
|
||||
auto_path.write_text("auto prompt\n", encoding="utf-8")
|
||||
enhancer.rewrite_all_system_prompt_path = str(
|
||||
tmp_path / "prompts.local" / "rewrite_window_system_prompt.md"
|
||||
)
|
||||
enhancer.rewrite_all_system_prompt_path = str(tmp_path / "prompts.local" / "rewrite_window_system_prompt.md")
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = str(rewrite_fallback_path)
|
||||
enhancer.enhance_system_prompt_path = str(next_path)
|
||||
enhancer.auto_system_prompt_path = str(auto_path)
|
||||
@@ -973,14 +889,9 @@ def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
|
||||
assert config["rewrite_window_system_prompt_path"] == str(rewrite_fallback_path)
|
||||
|
||||
|
||||
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(
|
||||
tmp_path,
|
||||
):
|
||||
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(tmp_path, ):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
@@ -1007,10 +918,7 @@ def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_emp
|
||||
|
||||
def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1021,9 +929,7 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
enhancer.auto_system_prompt_path = str(auto_path)
|
||||
enhancer.rewrite_all_system_prompt_path = str(rewrite_path)
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
auto_extension_system_prompt="auto updated",
|
||||
)
|
||||
config = enhancer.save_prompt_config(auto_extension_system_prompt="auto updated", )
|
||||
|
||||
assert auto_path.read_text(encoding="utf-8").strip() == "auto updated"
|
||||
assert config["auto_extension_system_prompt"] == "auto updated"
|
||||
@@ -1031,10 +937,7 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1052,9 +955,7 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_user_system_prompt_fallback_path = None
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_user_system_prompt="rewrite user updated",
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_user_system_prompt="rewrite user updated", )
|
||||
|
||||
assert rewrite_user_path.read_text(encoding="utf-8").strip() == "rewrite user updated"
|
||||
assert config["rewrite_user_system_prompt"] == "rewrite user updated"
|
||||
@@ -1062,10 +963,7 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1081,9 +979,7 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
enhancer.rewrite_default_model = "gpt-test"
|
||||
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_model="gpt-alt",
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_model="gpt-alt", )
|
||||
|
||||
assert enhancer.rewrite_default_model == "gpt-alt"
|
||||
assert config["rewrite_model"] == "gpt-alt"
|
||||
@@ -1092,10 +988,7 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
|
||||
|
||||
def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite.md"
|
||||
@@ -1109,9 +1002,7 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
enhancer.auto_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
|
||||
config = enhancer.save_prompt_config(
|
||||
rewrite_temperature=1.3,
|
||||
)
|
||||
config = enhancer.save_prompt_config(rewrite_temperature=1.3, )
|
||||
|
||||
assert enhancer.rewrite_default_temperature == 1.3
|
||||
assert config["rewrite_temperature"] == 1.3
|
||||
@@ -1119,10 +1010,7 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
|
||||
|
||||
def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_path):
|
||||
enhancer = _build_test_enhancer(
|
||||
_chat_payload_with_content(
|
||||
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
|
||||
)
|
||||
)
|
||||
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
|
||||
next_path = tmp_path / "next.md"
|
||||
auto_path = tmp_path / "auto.md"
|
||||
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
|
||||
@@ -1136,13 +1024,9 @@ def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_pat
|
||||
enhancer.auto_system_prompt_fallback_path = None
|
||||
enhancer.rewrite_all_system_prompt_fallback_path = None
|
||||
|
||||
enhancer.save_prompt_config(
|
||||
rewrite_window_system_prompt="rewrite updated",
|
||||
)
|
||||
enhancer.save_prompt_config(rewrite_window_system_prompt="rewrite updated", )
|
||||
|
||||
backup_paths = sorted(
|
||||
tmp_path.glob("rewrite_window_system_prompt.*.bak.md")
|
||||
)
|
||||
backup_paths = sorted(tmp_path.glob("rewrite_window_system_prompt.*.bak.md"))
|
||||
|
||||
assert rewrite_path.read_text(encoding="utf-8").strip() == "rewrite updated"
|
||||
assert len(backup_paths) == 1
|
||||
|
||||
@@ -27,13 +27,8 @@ try:
|
||||
except ModuleNotFoundError:
|
||||
websockets = None # type: ignore[assignment]
|
||||
|
||||
|
||||
DEFAULT_PRESET_FILE = (
|
||||
Path(__file__).resolve().parents[2]
|
||||
/ "web"
|
||||
/ "prompts"
|
||||
/ "selected_ltx2_continuation_story_presets.json"
|
||||
)
|
||||
DEFAULT_PRESET_FILE = (Path(__file__).resolve().parents[2] / "web" / "prompts" /
|
||||
"selected_ltx2_continuation_story_presets.json")
|
||||
|
||||
|
||||
def utc_now_iso() -> str:
|
||||
@@ -65,10 +60,7 @@ def safe_percentile(values: list[float], percentile: float) -> float | None:
|
||||
if lower == upper:
|
||||
return sorted_values[lower]
|
||||
fraction = rank - lower
|
||||
return (
|
||||
sorted_values[lower]
|
||||
+ (sorted_values[upper] - sorted_values[lower]) * fraction
|
||||
)
|
||||
return (sorted_values[lower] + (sorted_values[upper] - sorted_values[lower]) * fraction)
|
||||
|
||||
|
||||
def summarize_series(values: list[float]) -> dict[str, float | int | None]:
|
||||
@@ -145,24 +137,16 @@ def load_curated_prompts(
|
||||
selected_id = str(selected.get("id", "")).strip() or "unknown_preset"
|
||||
raw_prompts = selected.get("segment_prompts", [])
|
||||
if not isinstance(raw_prompts, list):
|
||||
raise ValueError(
|
||||
f"Preset {selected_id} has invalid segment_prompts (must be list)."
|
||||
)
|
||||
raise ValueError(f"Preset {selected_id} has invalid segment_prompts (must be list).")
|
||||
|
||||
prompts = [
|
||||
str(prompt).strip()
|
||||
for prompt in raw_prompts
|
||||
if isinstance(prompt, str) and str(prompt).strip()
|
||||
]
|
||||
prompts = [str(prompt).strip() for prompt in raw_prompts if isinstance(prompt, str) and str(prompt).strip()]
|
||||
if not prompts:
|
||||
raise ValueError(f"Preset {selected_id} has no non-empty prompts.")
|
||||
|
||||
limited = prompts[:curated_limit]
|
||||
if not limited:
|
||||
raise ValueError(
|
||||
f"curated_limit={curated_limit} produced no prompts for preset "
|
||||
f"{selected_id}."
|
||||
)
|
||||
raise ValueError(f"curated_limit={curated_limit} produced no prompts for preset "
|
||||
f"{selected_id}.")
|
||||
return selected_id, limited, len(prompts)
|
||||
|
||||
|
||||
@@ -224,11 +208,11 @@ async def run_single_session(
|
||||
|
||||
try:
|
||||
async with websockets.connect(
|
||||
url,
|
||||
max_size=None,
|
||||
ping_interval=None,
|
||||
open_timeout=connect_timeout_s,
|
||||
close_timeout=2.0,
|
||||
url,
|
||||
max_size=None,
|
||||
ping_interval=None,
|
||||
open_timeout=connect_timeout_s,
|
||||
close_timeout=2.0,
|
||||
) as ws:
|
||||
connect_finish_monotonic = time.monotonic()
|
||||
session_data["connect_finish_ts_utc"] = utc_now_iso()
|
||||
@@ -249,9 +233,7 @@ async def run_single_session(
|
||||
timeout_remaining = session_timeout_s - elapsed_s
|
||||
if timeout_remaining <= 0:
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = (
|
||||
f"Session timed out after {session_timeout_s:.1f}s."
|
||||
)
|
||||
session_data["error"] = (f"Session timed out after {session_timeout_s:.1f}s.")
|
||||
break
|
||||
|
||||
recv_start_epoch = time.time()
|
||||
@@ -265,9 +247,7 @@ async def run_single_session(
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = (
|
||||
"Timed out waiting for websocket message."
|
||||
)
|
||||
session_data["error"] = ("Timed out waiting for websocket message.")
|
||||
break
|
||||
except Exception as exc:
|
||||
session_data["status"] = "failed"
|
||||
@@ -288,20 +268,16 @@ async def run_single_session(
|
||||
|
||||
chunk_gap_ms: float | None = None
|
||||
if last_chunk_finish_monotonic is not None:
|
||||
chunk_gap_ms = (
|
||||
recv_finish_monotonic - last_chunk_finish_monotonic
|
||||
) * 1000.0
|
||||
chunk_gap_ms = (recv_finish_monotonic - last_chunk_finish_monotonic) * 1000.0
|
||||
|
||||
session_data["chunks"].append(
|
||||
{
|
||||
"segment_idx": current_segment_idx,
|
||||
"chunk_idx": session_data["total_chunks"],
|
||||
"size_bytes": len(message),
|
||||
"chunk_start_ts_utc": recv_start_iso,
|
||||
"chunk_finish_ts_utc": recv_finish_iso,
|
||||
"chunk_gap_ms": chunk_gap_ms,
|
||||
}
|
||||
)
|
||||
session_data["chunks"].append({
|
||||
"segment_idx": current_segment_idx,
|
||||
"chunk_idx": session_data["total_chunks"],
|
||||
"size_bytes": len(message),
|
||||
"chunk_start_ts_utc": recv_start_iso,
|
||||
"chunk_finish_ts_utc": recv_finish_iso,
|
||||
"chunk_gap_ms": chunk_gap_ms,
|
||||
})
|
||||
last_chunk_finish_monotonic = recv_finish_monotonic
|
||||
last_chunk_finish_epoch = recv_finish_epoch
|
||||
session_data["last_chunk_finish_ts_utc"] = recv_finish_iso
|
||||
@@ -321,9 +297,7 @@ async def run_single_session(
|
||||
if msg_type == "gpu_assigned":
|
||||
session_data["gpu_assigned_ts_utc"] = recv_finish_iso
|
||||
if connect_finish_monotonic is not None:
|
||||
session_data["queue_wait_ms"] = (
|
||||
recv_finish_monotonic - connect_finish_monotonic
|
||||
) * 1000.0
|
||||
session_data["queue_wait_ms"] = (recv_finish_monotonic - connect_finish_monotonic) * 1000.0
|
||||
elif msg_type == "ltx2_stream_start":
|
||||
if initial_total_segments is None:
|
||||
parsed_total = parse_int(data.get("total_segments"))
|
||||
@@ -338,20 +312,13 @@ async def run_single_session(
|
||||
session_data["media_segments_completed"] += 1
|
||||
if first_media_segment_complete_epoch is None:
|
||||
first_media_segment_complete_epoch = recv_finish_epoch
|
||||
session_data[
|
||||
"first_media_segment_complete_ts_utc"
|
||||
] = recv_finish_iso
|
||||
session_data["first_media_segment_complete_ts_utc"] = recv_finish_iso
|
||||
elif msg_type == "ltx2_segment_complete":
|
||||
session_data["segments_completed"] += 1
|
||||
seg_idx = parse_int(data.get("segment_idx"))
|
||||
if (
|
||||
initial_total_segments is not None
|
||||
and seg_idx is not None
|
||||
and seg_idx >= initial_total_segments
|
||||
):
|
||||
session_data[
|
||||
"target_segment_complete_ts_utc"
|
||||
] = recv_finish_iso
|
||||
if (initial_total_segments is not None and seg_idx is not None
|
||||
and seg_idx >= initial_total_segments):
|
||||
session_data["target_segment_complete_ts_utc"] = recv_finish_iso
|
||||
await asyncio.sleep(post_complete_wait_s)
|
||||
session_data["leave_sent_ts_utc"] = utc_now_iso()
|
||||
try:
|
||||
@@ -362,15 +329,11 @@ async def run_single_session(
|
||||
break
|
||||
elif msg_type == "session_timeout":
|
||||
session_data["status"] = "timeout"
|
||||
session_data["error"] = str(
|
||||
data.get("message") or "Backend session timeout"
|
||||
)
|
||||
session_data["error"] = str(data.get("message") or "Backend session timeout")
|
||||
break
|
||||
elif msg_type == "error":
|
||||
session_data["status"] = "failed"
|
||||
session_data["error"] = str(
|
||||
data.get("message") or "Backend error message"
|
||||
)
|
||||
session_data["error"] = str(data.get("message") or "Backend error message")
|
||||
break
|
||||
|
||||
if session_data["status"] == "failed" and session_data["error"] is None:
|
||||
@@ -379,29 +342,18 @@ async def run_single_session(
|
||||
session_data["status"] = "failed"
|
||||
session_data["error"] = f"WebSocket connect/run failed: {exc}"
|
||||
|
||||
if (
|
||||
first_chunk_finish_epoch is not None
|
||||
and last_chunk_finish_epoch is not None
|
||||
and session_data["total_chunk_bytes"] > 0
|
||||
):
|
||||
if (first_chunk_finish_epoch is not None and last_chunk_finish_epoch is not None
|
||||
and session_data["total_chunk_bytes"] > 0):
|
||||
duration_s = last_chunk_finish_epoch - first_chunk_finish_epoch
|
||||
if duration_s > 0:
|
||||
session_data["session_goodput_mbps"] = (
|
||||
session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0
|
||||
)
|
||||
session_data["session_goodput_mbps"] = (session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0)
|
||||
|
||||
if (
|
||||
first_chunk_finish_epoch is not None
|
||||
and first_media_segment_complete_epoch is not None
|
||||
):
|
||||
session_data["first_chunk_before_first_media_complete"] = (
|
||||
first_chunk_finish_epoch < first_media_segment_complete_epoch
|
||||
)
|
||||
if (first_chunk_finish_epoch is not None and first_media_segment_complete_epoch is not None):
|
||||
session_data["first_chunk_before_first_media_complete"] = (first_chunk_finish_epoch
|
||||
< first_media_segment_complete_epoch)
|
||||
|
||||
session_data["close_ts_utc"] = utc_now_iso()
|
||||
session_data["duration_ms"] = (
|
||||
time.monotonic() - session_start_monotonic
|
||||
) * 1000.0
|
||||
session_data["duration_ms"] = (time.monotonic() - session_start_monotonic) * 1000.0
|
||||
return session_data
|
||||
|
||||
|
||||
@@ -412,14 +364,11 @@ async def run_worker_sessions(
|
||||
config: dict[str, Any],
|
||||
) -> list[dict[str, Any]]:
|
||||
tasks = [
|
||||
asyncio.create_task(
|
||||
run_single_session(
|
||||
worker_id=worker_id,
|
||||
worker_session_idx=idx,
|
||||
config=config,
|
||||
)
|
||||
)
|
||||
for idx in range(session_count)
|
||||
asyncio.create_task(run_single_session(
|
||||
worker_id=worker_id,
|
||||
worker_session_idx=idx,
|
||||
config=config,
|
||||
)) for idx in range(session_count)
|
||||
]
|
||||
if not tasks:
|
||||
return []
|
||||
@@ -437,29 +386,23 @@ def worker_entry(
|
||||
try:
|
||||
ready_queue.put({"worker_id": worker_id, "status": "ready"})
|
||||
start_event.wait()
|
||||
sessions = asyncio.run(
|
||||
run_worker_sessions(
|
||||
worker_id=worker_id,
|
||||
session_count=session_count,
|
||||
config=config,
|
||||
)
|
||||
)
|
||||
result_queue.put(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"status": "ok",
|
||||
"sessions": sessions,
|
||||
}
|
||||
)
|
||||
sessions = asyncio.run(run_worker_sessions(
|
||||
worker_id=worker_id,
|
||||
session_count=session_count,
|
||||
config=config,
|
||||
))
|
||||
result_queue.put({
|
||||
"worker_id": worker_id,
|
||||
"status": "ok",
|
||||
"sessions": sessions,
|
||||
})
|
||||
except Exception as exc:
|
||||
result_queue.put(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"status": "error",
|
||||
"error": str(exc),
|
||||
"traceback": traceback.format_exc(),
|
||||
}
|
||||
)
|
||||
result_queue.put({
|
||||
"worker_id": worker_id,
|
||||
"status": "error",
|
||||
"error": str(exc),
|
||||
"traceback": traceback.format_exc(),
|
||||
})
|
||||
|
||||
|
||||
def build_summary(
|
||||
@@ -517,33 +460,22 @@ def build_summary(
|
||||
if len(all_chunk_finish_epochs) >= 2 and total_chunk_bytes > 0:
|
||||
duration_s = max(all_chunk_finish_epochs) - min(all_chunk_finish_epochs)
|
||||
if duration_s > 0:
|
||||
global_goodput_mbps = (
|
||||
total_chunk_bytes * 8.0 / duration_s / 1_000_000.0
|
||||
)
|
||||
global_goodput_mbps = (total_chunk_bytes * 8.0 / duration_s / 1_000_000.0)
|
||||
|
||||
bucket_throughputs_mbps = [
|
||||
(bytes_count * 8.0) / 1_000_000.0
|
||||
for _, bytes_count in sorted(bucket_bytes.items())
|
||||
]
|
||||
bucket_throughputs_mbps = [(bytes_count * 8.0) / 1_000_000.0 for _, bytes_count in sorted(bucket_bytes.items())]
|
||||
bucket_stats = summarize_series(bucket_throughputs_mbps)
|
||||
|
||||
chunk_gap_threshold_breaches = [
|
||||
value for value in chunk_gaps if value >= chunk_gap_threshold_ms
|
||||
]
|
||||
chunk_gap_threshold_breaches = [value for value in chunk_gaps if value >= chunk_gap_threshold_ms]
|
||||
non_success = len(sessions) - status_counts.get("success", 0)
|
||||
|
||||
fail_reasons: list[str] = []
|
||||
if non_success > 0:
|
||||
fail_reasons.append(
|
||||
f"{non_success} session(s) did not complete successfully."
|
||||
)
|
||||
fail_reasons.append(f"{non_success} session(s) did not complete successfully.")
|
||||
if not chunk_gaps:
|
||||
fail_reasons.append("No chunk gap data collected.")
|
||||
if chunk_gap_threshold_breaches:
|
||||
fail_reasons.append(
|
||||
f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
|
||||
f"{chunk_gap_threshold_ms:.0f}ms."
|
||||
)
|
||||
fail_reasons.append(f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
|
||||
f"{chunk_gap_threshold_ms:.0f}ms.")
|
||||
|
||||
passed = len(fail_reasons) == 0
|
||||
progressive_ratio = None
|
||||
@@ -554,20 +486,18 @@ def build_summary(
|
||||
"passed": passed,
|
||||
"fail_reasons": fail_reasons,
|
||||
"sessions": {
|
||||
"total": len(sessions),
|
||||
"success": status_counts.get("success", 0),
|
||||
"failed": status_counts.get("failed", 0),
|
||||
"timeout": status_counts.get("timeout", 0),
|
||||
"protocol_error": status_counts.get("protocol_error", 0),
|
||||
"other": (
|
||||
len(sessions)
|
||||
- (
|
||||
status_counts.get("success", 0)
|
||||
+ status_counts.get("failed", 0)
|
||||
+ status_counts.get("timeout", 0)
|
||||
+ status_counts.get("protocol_error", 0)
|
||||
)
|
||||
),
|
||||
"total":
|
||||
len(sessions),
|
||||
"success":
|
||||
status_counts.get("success", 0),
|
||||
"failed":
|
||||
status_counts.get("failed", 0),
|
||||
"timeout":
|
||||
status_counts.get("timeout", 0),
|
||||
"protocol_error":
|
||||
status_counts.get("protocol_error", 0),
|
||||
"other": (len(sessions) - (status_counts.get("success", 0) + status_counts.get("failed", 0) +
|
||||
status_counts.get("timeout", 0) + status_counts.get("protocol_error", 0))),
|
||||
},
|
||||
"chunk_gap_ms": {
|
||||
**chunk_gap_stats,
|
||||
@@ -606,51 +536,39 @@ def print_summary(
|
||||
bucket_bw = bandwidth["bucketed_1s"]
|
||||
|
||||
print("=== LTX2 Realtime Stress Test Summary ===")
|
||||
print(
|
||||
"Run: "
|
||||
f"url={run_info['url']} clients={run_info['clients']} "
|
||||
f"processes={run_info['processes']} "
|
||||
f"preset={run_info['preset_id']} "
|
||||
f"curated_limit={run_info['curated_limit']}"
|
||||
)
|
||||
print(
|
||||
"Sessions: "
|
||||
f"total={sessions['total']} success={sessions['success']} "
|
||||
f"failed={sessions['failed']} timeout={sessions['timeout']} "
|
||||
f"protocol_error={sessions['protocol_error']}"
|
||||
)
|
||||
print(
|
||||
"Chunk gap ms: "
|
||||
f"min={format_num(chunk_gap['min'])} "
|
||||
f"p50={format_num(chunk_gap['p50'])} "
|
||||
f"p95={format_num(chunk_gap['p95'])} "
|
||||
f"p99={format_num(chunk_gap['p99'])} "
|
||||
f"max={format_num(chunk_gap['max'])} "
|
||||
f"threshold={format_num(chunk_gap['threshold_ms'])} "
|
||||
f"breaches={chunk_gap['breach_count']}"
|
||||
)
|
||||
print(
|
||||
"Queue wait ms: "
|
||||
f"min={format_num(queue_wait['min'])} "
|
||||
f"p50={format_num(queue_wait['p50'])} "
|
||||
f"p95={format_num(queue_wait['p95'])} "
|
||||
f"max={format_num(queue_wait['max'])}"
|
||||
)
|
||||
print("Run: "
|
||||
f"url={run_info['url']} clients={run_info['clients']} "
|
||||
f"processes={run_info['processes']} "
|
||||
f"preset={run_info['preset_id']} "
|
||||
f"curated_limit={run_info['curated_limit']}")
|
||||
print("Sessions: "
|
||||
f"total={sessions['total']} success={sessions['success']} "
|
||||
f"failed={sessions['failed']} timeout={sessions['timeout']} "
|
||||
f"protocol_error={sessions['protocol_error']}")
|
||||
print("Chunk gap ms: "
|
||||
f"min={format_num(chunk_gap['min'])} "
|
||||
f"p50={format_num(chunk_gap['p50'])} "
|
||||
f"p95={format_num(chunk_gap['p95'])} "
|
||||
f"p99={format_num(chunk_gap['p99'])} "
|
||||
f"max={format_num(chunk_gap['max'])} "
|
||||
f"threshold={format_num(chunk_gap['threshold_ms'])} "
|
||||
f"breaches={chunk_gap['breach_count']}")
|
||||
print("Queue wait ms: "
|
||||
f"min={format_num(queue_wait['min'])} "
|
||||
f"p50={format_num(queue_wait['p50'])} "
|
||||
f"p95={format_num(queue_wait['p95'])} "
|
||||
f"max={format_num(queue_wait['max'])}")
|
||||
ratio = progressive["ratio"]
|
||||
ratio_text = "n/a" if ratio is None else f"{ratio * 100:.2f}%"
|
||||
print(
|
||||
"Progressive streaming: "
|
||||
f"{progressive['success_sessions']}/"
|
||||
f"{progressive['eligible_sessions']} ({ratio_text})"
|
||||
)
|
||||
print(
|
||||
"Bandwidth Mbps: "
|
||||
f"per_session_avg={format_num(per_session_bw['avg'])} "
|
||||
f"per_session_p95={format_num(per_session_bw['p95'])} "
|
||||
f"global={format_num(bandwidth['global_goodput_mbps'])} "
|
||||
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
|
||||
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}"
|
||||
)
|
||||
print("Progressive streaming: "
|
||||
f"{progressive['success_sessions']}/"
|
||||
f"{progressive['eligible_sessions']} ({ratio_text})")
|
||||
print("Bandwidth Mbps: "
|
||||
f"per_session_avg={format_num(per_session_bw['avg'])} "
|
||||
f"per_session_p95={format_num(per_session_bw['p95'])} "
|
||||
f"global={format_num(bandwidth['global_goodput_mbps'])} "
|
||||
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
|
||||
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}")
|
||||
print(f"VERDICT: {'PASS' if summary['passed'] else 'FAIL'}")
|
||||
if summary["fail_reasons"]:
|
||||
print("Fail reasons:")
|
||||
@@ -670,10 +588,8 @@ def distribute_sessions(total_clients: int, process_count: int) -> list[int]:
|
||||
|
||||
def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
if websockets is None:
|
||||
raise RuntimeError(
|
||||
"Missing dependency: websockets. Install it before running this "
|
||||
"stress test."
|
||||
)
|
||||
raise RuntimeError("Missing dependency: websockets. Install it before running this "
|
||||
"stress test.")
|
||||
|
||||
preset_file = Path(args.preset_file).expanduser().resolve()
|
||||
selected_preset_id, curated_prompts, total_prompt_count = load_curated_prompts(
|
||||
@@ -735,13 +651,8 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
start_event.set()
|
||||
|
||||
result_deadline = (
|
||||
time.monotonic()
|
||||
+ args.connect_timeout_s
|
||||
+ args.session_timeout_s
|
||||
+ args.post_complete_wait_s
|
||||
+ 180.0
|
||||
)
|
||||
result_deadline = (time.monotonic() + args.connect_timeout_s + args.session_timeout_s +
|
||||
args.post_complete_wait_s + 180.0)
|
||||
worker_results: list[dict[str, Any]] = []
|
||||
while len(worker_results) < len(processes):
|
||||
timeout_s = max(0.1, result_deadline - time.monotonic())
|
||||
@@ -765,24 +676,20 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
if result.get("status") == "ok":
|
||||
sessions.extend(result.get("sessions", []))
|
||||
else:
|
||||
worker_errors.append(
|
||||
{
|
||||
"worker_id": result.get("worker_id"),
|
||||
"error": result.get("error"),
|
||||
"traceback": result.get("traceback"),
|
||||
}
|
||||
)
|
||||
worker_errors.append({
|
||||
"worker_id": result.get("worker_id"),
|
||||
"error": result.get("error"),
|
||||
"traceback": result.get("traceback"),
|
||||
})
|
||||
|
||||
received_workers = {result.get("worker_id") for result in worker_results}
|
||||
expected_workers = set(range(len(processes)))
|
||||
missing_workers = sorted(expected_workers - received_workers)
|
||||
for worker_id in missing_workers:
|
||||
worker_errors.append(
|
||||
{
|
||||
"worker_id": worker_id,
|
||||
"error": "No worker result received.",
|
||||
}
|
||||
)
|
||||
worker_errors.append({
|
||||
"worker_id": worker_id,
|
||||
"error": "No worker result received.",
|
||||
})
|
||||
|
||||
run_end_epoch = time.time()
|
||||
run_end_iso = iso_from_epoch(run_end_epoch)
|
||||
@@ -795,9 +702,8 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
if worker_errors:
|
||||
summary["passed"] = False
|
||||
summary["fail_reasons"] = list(summary["fail_reasons"]) + [
|
||||
f"{len(worker_errors)} worker error(s) occurred."
|
||||
]
|
||||
summary["fail_reasons"] = list(
|
||||
summary["fail_reasons"]) + [f"{len(worker_errors)} worker error(s) occurred."]
|
||||
|
||||
output_payload = {
|
||||
"run_info": {
|
||||
@@ -833,9 +739,7 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Multiprocess realtime stress test for LTX2 streaming.",
|
||||
)
|
||||
parser = argparse.ArgumentParser(description="Multiprocess realtime stress test for LTX2 streaming.", )
|
||||
parser.add_argument(
|
||||
"-u",
|
||||
"--url",
|
||||
|
||||
@@ -47,13 +47,11 @@ def test_persist_session_init_image_returns_none_when_missing_data():
|
||||
|
||||
def test_persist_session_init_image_rejects_unsupported_mime():
|
||||
with pytest.raises(ValueError, match="PNG, JPEG, or WebP"):
|
||||
persist_session_init_image(
|
||||
{
|
||||
"name": "frame.gif",
|
||||
"mime_type": "image/gif",
|
||||
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
|
||||
}
|
||||
)
|
||||
persist_session_init_image({
|
||||
"name": "frame.gif",
|
||||
"mime_type": "image/gif",
|
||||
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
|
||||
})
|
||||
|
||||
|
||||
def test_persist_session_init_image_rejects_large_payload(monkeypatch):
|
||||
@@ -66,10 +64,8 @@ def test_persist_session_init_image_rejects_large_payload(monkeypatch):
|
||||
monkeypatch.setattr(base64, "b64decode", fake_b64decode)
|
||||
|
||||
with pytest.raises(ValueError, match="15 MB or smaller"):
|
||||
persist_session_init_image(
|
||||
{
|
||||
"name": "frame.png",
|
||||
"mime_type": "image/png",
|
||||
"data_url": data_url,
|
||||
}
|
||||
)
|
||||
persist_session_init_image({
|
||||
"name": "frame.png",
|
||||
"mime_type": "image/png",
|
||||
"data_url": data_url,
|
||||
})
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -8,12 +8,10 @@ import modal
|
||||
|
||||
IMAGE = os.environ.get("DREAMVERSE_IMAGE")
|
||||
if not IMAGE:
|
||||
raise RuntimeError(
|
||||
"DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
|
||||
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
|
||||
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
|
||||
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag."
|
||||
)
|
||||
raise RuntimeError("DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
|
||||
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
|
||||
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
|
||||
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag.")
|
||||
|
||||
# ``@modal.web_server`` invokes ``serve()`` directly and bypasses the image
|
||||
# ENTRYPOINT (``docker/docker_entrypoint.sh``). That entrypoint normally
|
||||
@@ -65,14 +63,10 @@ def serve():
|
||||
# ``or ""`` collapses ``None`` (unset) into an empty string, ``.strip()``
|
||||
# collapses whitespace-only values (e.g. ``" "``) — both should be
|
||||
# treated as missing.
|
||||
missing = [
|
||||
k for k in _REQUIRED_SECRET_KEYS
|
||||
if not (os.environ.get(k) or "").strip()
|
||||
]
|
||||
missing = [k for k in _REQUIRED_SECRET_KEYS if not (os.environ.get(k) or "").strip()]
|
||||
if missing:
|
||||
raise RuntimeError(
|
||||
"dreamverse-api-keys secret is missing required entries: "
|
||||
f"{', '.join(missing)}. Add them with `modal secret create "
|
||||
"dreamverse-api-keys ... --force` and redeploy "
|
||||
"(see apps/dreamverse/scripts/modal/README.md).")
|
||||
raise RuntimeError("dreamverse-api-keys secret is missing required entries: "
|
||||
f"{', '.join(missing)}. Add them with `modal secret create "
|
||||
"dreamverse-api-keys ... --force` and redeploy "
|
||||
"(see apps/dreamverse/scripts/modal/README.md).")
|
||||
subprocess.Popen(["dreamverse-server", "--host", "0.0.0.0", "--port", "8009"])
|
||||
|
||||
@@ -1,39 +0,0 @@
|
||||
import { expect, test } from '@playwright/test';
|
||||
|
||||
import { skipWithoutMock } from './helpers';
|
||||
|
||||
/**
|
||||
* Warm-model slot: load a model through the mock (which flips
|
||||
* loading -> ready after ~2s, surfaced by the panel's 5s poll), then unload it.
|
||||
*/
|
||||
test.describe('generators', () => {
|
||||
skipWithoutMock();
|
||||
|
||||
test('loads and unloads the resident model', async ({ page }) => {
|
||||
await page.goto('/inference');
|
||||
|
||||
const panel = page.getByRole('region', { name: 'Warm models' });
|
||||
await expect(panel).toBeVisible();
|
||||
await expect(panel.getByText('No model loaded')).toBeVisible();
|
||||
|
||||
await panel
|
||||
.getByLabel('Model to load')
|
||||
.selectOption('Wan-AI/Wan2.1-T2V-1.3B-Diffusers');
|
||||
await panel.getByRole('button', { name: 'Load model' }).click();
|
||||
|
||||
await expect(panel.getByText('Wan2.1 T2V 1.3B Diffusers')).toBeVisible();
|
||||
await expect(panel.getByText('ready')).toBeVisible();
|
||||
|
||||
await panel.getByRole('button', { name: 'Unload' }).click();
|
||||
await expect(panel.getByText('No model loaded')).toBeVisible();
|
||||
});
|
||||
|
||||
test('engine console streams output while open', async ({ page }) => {
|
||||
await page.goto('/inference');
|
||||
|
||||
const engineConsole = page.getByRole('region', { name: 'Engine output' });
|
||||
await engineConsole.getByRole('button', { name: 'Engine output' }).click();
|
||||
|
||||
await expect(engineConsole.getByText(/\[engine\]/).first()).toBeVisible();
|
||||
});
|
||||
});
|
||||
@@ -72,90 +72,3 @@ def get_gpu_snapshot() -> dict[str, Any]:
|
||||
except Exception as exc: # NVMLError, driver issues, …
|
||||
logger.warning("GPU snapshot failed: %s", exc)
|
||||
return {"available": False, "gpus": [], "error": str(exc)}
|
||||
|
||||
|
||||
def _remote_gpu_probe() -> dict[str, Any]:
|
||||
"""Self-contained per-node NVML probe (runs as a ray task on each node;
|
||||
no fastvideo_studio import — worker environments don't have apps/ on
|
||||
their path, so cloudpickle must carry this by value)."""
|
||||
import contextlib as _ctx
|
||||
import socket as _socket
|
||||
out: dict[str, Any] = {"hostname": _socket.gethostname(), "available": False, "gpus": [], "error": None}
|
||||
try:
|
||||
import pynvml
|
||||
pynvml.nvmlInit()
|
||||
for i in range(pynvml.nvmlDeviceGetCount()):
|
||||
h = pynvml.nvmlDeviceGetHandleByIndex(i)
|
||||
name = pynvml.nvmlDeviceGetName(h)
|
||||
if isinstance(name, bytes):
|
||||
name = name.decode()
|
||||
mem = pynvml.nvmlDeviceGetMemoryInfo(h)
|
||||
util = pynvml.nvmlDeviceGetUtilizationRates(h)
|
||||
temp = power = plimit = None
|
||||
with _ctx.suppress(pynvml.NVMLError):
|
||||
temp = int(pynvml.nvmlDeviceGetTemperature(h, pynvml.NVML_TEMPERATURE_GPU))
|
||||
with _ctx.suppress(pynvml.NVMLError):
|
||||
power = pynvml.nvmlDeviceGetPowerUsage(h) / 1000.0
|
||||
plimit = pynvml.nvmlDeviceGetEnforcedPowerLimit(h) / 1000.0
|
||||
out["gpus"].append({
|
||||
"index": i,
|
||||
"name": name,
|
||||
"utilization": int(util.gpu),
|
||||
"memory_used_mib": int(mem.used / (1024 * 1024)),
|
||||
"memory_total_mib": int(mem.total / (1024 * 1024)),
|
||||
"temperature_c": temp,
|
||||
"power_watts": power,
|
||||
"power_limit_watts": plimit,
|
||||
})
|
||||
out["available"] = True
|
||||
except Exception as exc: # noqa: BLE001 -- reported per node
|
||||
out["error"] = str(exc)
|
||||
return out
|
||||
|
||||
|
||||
def get_cluster_snapshot() -> dict[str, Any]:
|
||||
"""Cluster-wide GPU/host telemetry.
|
||||
|
||||
When this process is connected to a ray cluster (a model has been
|
||||
loaded), probe every alive node via per-node ray tasks. Otherwise fall
|
||||
back to this host's NVML snapshot. Never raises.
|
||||
"""
|
||||
import socket
|
||||
|
||||
local = get_gpu_snapshot()
|
||||
local_node = {"hostname": socket.gethostname(), "ip": None, "is_this_host": True,
|
||||
"cpus": None, "ray_gpus": None, **local}
|
||||
out: dict[str, Any] = {"mode": "local", "nodes": [local_node],
|
||||
"resources": None, "error": None}
|
||||
try:
|
||||
import ray
|
||||
if not ray.is_initialized():
|
||||
out["error"] = "not connected to a ray cluster yet (load a model first); showing the API host only"
|
||||
return out
|
||||
from ray.util.scheduling_strategies import NodeAffinitySchedulingStrategy
|
||||
alive = [n for n in ray.nodes() if n.get("Alive")]
|
||||
probe = ray.remote(num_cpus=0)(_remote_gpu_probe)
|
||||
refs = [probe.options(scheduling_strategy=NodeAffinitySchedulingStrategy(
|
||||
node_id=n["NodeID"], soft=True)).remote() for n in alive]
|
||||
snaps = ray.get(refs, timeout=15)
|
||||
nodes = []
|
||||
for n, snap in zip(alive, snaps, strict=True):
|
||||
nodes.append({
|
||||
"ip": n.get("NodeManagerAddress"),
|
||||
"is_this_host": snap.get("hostname") == socket.gethostname(),
|
||||
"cpus": n.get("Resources", {}).get("CPU"),
|
||||
"ray_gpus": n.get("Resources", {}).get("GPU"),
|
||||
**snap,
|
||||
})
|
||||
out["mode"] = "ray"
|
||||
out["nodes"] = nodes
|
||||
out["resources"] = {
|
||||
"gpus_total": ray.cluster_resources().get("GPU", 0.0),
|
||||
"gpus_available": ray.available_resources().get("GPU", 0.0),
|
||||
}
|
||||
except Exception as exc: # noqa: BLE001 -- degrade to the local view
|
||||
logger.warning("cluster snapshot failed: %s", exc)
|
||||
out["mode"] = "local"
|
||||
out["nodes"] = [local_node]
|
||||
out["error"] = f"cluster probe failed: {exc}"
|
||||
return out
|
||||
|
||||
@@ -41,11 +41,6 @@ _TQDM_FRAC_RE = re.compile(r"\b(\d+)/(\d+)\b")
|
||||
|
||||
_MAX_LOG_LINES = 2000 # ring-buffer cap per job
|
||||
|
||||
# ray's log relay prefixes worker lines like "(RayWorkerWrapper pid=123, ip=…)"
|
||||
# — usually wrapped in ANSI color codes, which must be stripped before matching.
|
||||
_RAY_RELAY_RE = re.compile(r"^\(\w+ pid=")
|
||||
_ANSI_RE = re.compile(r"\x1b\[[0-9;]*m")
|
||||
|
||||
|
||||
class JobStatus(str, enum.Enum):
|
||||
PENDING = "pending"
|
||||
@@ -264,36 +259,13 @@ class JobRunner:
|
||||
self._jobs_lock = threading.Lock()
|
||||
self._load_jobs()
|
||||
|
||||
# Exactly one generator lives in memory at a time. Loading a new
|
||||
# config always releases the old instance first (shutdown + placement
|
||||
# group teardown); unload deletes it outright.
|
||||
self._generator: Any | None = None
|
||||
self._generator_config: dict[str, Any] | None = None
|
||||
self._generator_state: str = "empty" # empty | loading | ready | failed
|
||||
self._generator_error: str | None = None
|
||||
self._generator_lock = threading.Lock() # guards the fields above
|
||||
# Serializes every slot transition (preload, job-triggered replace,
|
||||
# unload). Held for the full duration of a load.
|
||||
self._load_lock = threading.Lock()
|
||||
# All generator creations run on this ONE persistent thread: the mp
|
||||
# executor arms prctl(PR_SET_PDEATHSIG, SIGKILL) in its workers, and
|
||||
# on Linux that fires when the CREATING THREAD exits — a generator
|
||||
# spawned from a short-lived thread loses all its workers (silent
|
||||
# SIGKILL, zombies) the moment that thread finishes.
|
||||
import queue as _queue
|
||||
self._loader_queue: _queue.Queue = _queue.Queue()
|
||||
threading.Thread(target=self._loader_loop, daemon=True,
|
||||
name="generator-loader").start()
|
||||
# The inference job currently generating, fed by the engine log tee
|
||||
# (ray relays worker output to the driver; tqdm lines land there).
|
||||
self._active_inference_job: Job | None = None
|
||||
# Cache of loaded generators keyed by model config so that we only pay
|
||||
# the model-loading cost once per model configuration.
|
||||
self._generators: dict[tuple, Any] = {}
|
||||
self._generators_lock = threading.Lock()
|
||||
|
||||
# Shared Manager for log queues (avoids spawning a new process per job)
|
||||
self._mp_manager = get_mp_context().Manager()
|
||||
# One queue for the generator's whole life: mp workers get it at spawn
|
||||
# (creation-time attach). Sending a Manager proxy over the executor's
|
||||
# worker pipes post-hoc (set_log_queue RPC) breaks the pipe.
|
||||
self._worker_log_queue = self._mp_manager.Queue()
|
||||
atexit.register(self._shutdown)
|
||||
|
||||
# Ensure directories exist
|
||||
@@ -645,211 +617,6 @@ class JobRunner:
|
||||
"phase": job._log_buf.phase,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _generator_config_dict(
|
||||
model_id: str,
|
||||
workload_type: str,
|
||||
num_gpus: int,
|
||||
dit_cpu_offload: bool = False,
|
||||
text_encoder_cpu_offload: bool = False,
|
||||
vae_cpu_offload: bool = False,
|
||||
image_encoder_cpu_offload: bool = False,
|
||||
use_fsdp_inference: bool = False,
|
||||
enable_torch_compile: bool = False,
|
||||
vsa_sparsity: float = 0.0,
|
||||
tp_size: int = -1,
|
||||
sp_size: int = -1,
|
||||
) -> dict[str, Any]:
|
||||
"""Canonical engine-config dict; equality here == same generator."""
|
||||
return {
|
||||
"model_id": model_id,
|
||||
"workload_type": workload_type,
|
||||
"num_gpus": num_gpus,
|
||||
"dit_cpu_offload": dit_cpu_offload,
|
||||
"text_encoder_cpu_offload": text_encoder_cpu_offload,
|
||||
"vae_cpu_offload": vae_cpu_offload,
|
||||
"image_encoder_cpu_offload": image_encoder_cpu_offload,
|
||||
"use_fsdp_inference": use_fsdp_inference,
|
||||
"enable_torch_compile": enable_torch_compile,
|
||||
"vsa_sparsity": vsa_sparsity,
|
||||
"tp_size": tp_size,
|
||||
"sp_size": sp_size,
|
||||
}
|
||||
|
||||
def _slot_entry(self) -> dict[str, Any]:
|
||||
return {
|
||||
"state": self._generator_state,
|
||||
"error": self._generator_error,
|
||||
**(self._generator_config or {}),
|
||||
}
|
||||
|
||||
def _running_inference_jobs(self) -> list[str]:
|
||||
with self._jobs_lock:
|
||||
return [j.id for j in self._jobs.values()
|
||||
if j.status == JobStatus.RUNNING and j.job_type == "inference"]
|
||||
|
||||
def _loader_loop(self) -> None:
|
||||
while True:
|
||||
fn = self._loader_queue.get()
|
||||
try:
|
||||
fn()
|
||||
except BaseException: # noqa: BLE001 -- surfaced via the caller's box
|
||||
pass
|
||||
finally:
|
||||
self._loader_queue.task_done()
|
||||
|
||||
def _run_on_loader(self, fn: Any) -> Any:
|
||||
"""Run ``fn`` on the persistent loader thread and return its result."""
|
||||
box: dict[str, Any] = {}
|
||||
done = threading.Event()
|
||||
|
||||
def wrapped() -> None:
|
||||
try:
|
||||
box["r"] = fn()
|
||||
except BaseException as exc: # noqa: BLE001 -- re-raised below
|
||||
box["e"] = exc
|
||||
finally:
|
||||
done.set()
|
||||
|
||||
self._loader_queue.put(wrapped)
|
||||
done.wait()
|
||||
if "e" in box:
|
||||
raise box["e"]
|
||||
return box["r"]
|
||||
|
||||
def preload_generator(self, **params: Any) -> dict[str, Any]:
|
||||
"""Load a model into memory ahead of time. One load at a time; loading
|
||||
a different config always releases the current instance first."""
|
||||
config = self._generator_config_dict(**params)
|
||||
with self._generator_lock:
|
||||
if self._generator_state == "ready" and self._generator_config == config:
|
||||
return self._slot_entry()
|
||||
if not self._load_lock.acquire(blocking=False):
|
||||
raise RuntimeError("a model load is already in progress")
|
||||
try:
|
||||
if self._running_inference_jobs():
|
||||
raise RuntimeError("cannot swap models while inference jobs are running")
|
||||
with self._generator_lock:
|
||||
self._generator_state = "loading"
|
||||
self._generator_config = config
|
||||
self._generator_error = None
|
||||
entry = self._slot_entry()
|
||||
except BaseException:
|
||||
self._load_lock.release()
|
||||
raise
|
||||
|
||||
def _load() -> None:
|
||||
try: # the preload owns _load_lock until the load resolves
|
||||
self._run_on_loader(lambda: self._load_into_slot_locked(config))
|
||||
except Exception: # noqa: BLE001 -- state already set to failed
|
||||
pass
|
||||
finally:
|
||||
self._load_lock.release()
|
||||
|
||||
threading.Thread(target=_load, daemon=True, name="generator-preload").start()
|
||||
return entry
|
||||
|
||||
def _load_into_slot_locked(self, config: dict[str, Any]) -> Any:
|
||||
"""Release whatever is resident and load ``config``. Caller MUST hold
|
||||
``_load_lock``. State is 'loading' on entry or set here."""
|
||||
with self._generator_lock:
|
||||
gen = self._generator
|
||||
self._generator = None
|
||||
self._generator_state = "loading"
|
||||
self._generator_config = config
|
||||
self._generator_error = None
|
||||
if gen is not None:
|
||||
logger.info("Releasing resident generator before loading a new one")
|
||||
gen.shutdown()
|
||||
del gen
|
||||
|
||||
# Import lazily so starting the server is fast even without a GPU.
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Deployment-level knobs (set where the API server is launched):
|
||||
# FASTVIDEO_STUDIO_MODEL_PATHS="id=/local/dir,..." serves a registered
|
||||
# model id from local weights instead of the HF hub;
|
||||
# FASTVIDEO_STUDIO_EXECUTOR_BACKEND=ray runs workers on an existing
|
||||
# Ray cluster (the multi-node path — "mp" spawns local processes only).
|
||||
model_path = config["model_id"]
|
||||
for pair in os.environ.get("FASTVIDEO_STUDIO_MODEL_PATHS", "").split(","):
|
||||
mid, sep, path = pair.partition("=")
|
||||
if sep and mid.strip() == config["model_id"]:
|
||||
model_path = path.strip()
|
||||
executor_kwargs: dict[str, Any] = {}
|
||||
backend = os.environ.get("FASTVIDEO_STUDIO_EXECUTOR_BACKEND")
|
||||
if backend:
|
||||
executor_kwargs["distributed_executor_backend"] = backend
|
||||
|
||||
logger.info("Loading model %s (%s)", config["model_id"],
|
||||
", ".join(f"{k}={v}" for k, v in config.items() if k != "model_id"))
|
||||
try:
|
||||
new_gen = VideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
workload_type=config["workload_type"],
|
||||
num_gpus=config["num_gpus"],
|
||||
dit_cpu_offload=config["dit_cpu_offload"],
|
||||
# FastVideoArgs defaults this True, which disables FSDP and
|
||||
# parks a full DiT copy in host RAM per worker — the mp
|
||||
# executor's 4 workers OOM-killed the node silently. The UI's
|
||||
# offload toggles are the studio's contract; layerwise off.
|
||||
dit_layerwise_offload=False,
|
||||
text_encoder_cpu_offload=config["text_encoder_cpu_offload"],
|
||||
vae_cpu_offload=config["vae_cpu_offload"],
|
||||
image_encoder_cpu_offload=config["image_encoder_cpu_offload"],
|
||||
use_fsdp_inference=config["use_fsdp_inference"],
|
||||
enable_torch_compile=config["enable_torch_compile"],
|
||||
VSA_sparsity=config["vsa_sparsity"],
|
||||
tp_size=config["tp_size"],
|
||||
sp_size=config["sp_size"],
|
||||
log_queue=self._worker_log_queue,
|
||||
**executor_kwargs,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception("Model load failed for %s", config["model_id"])
|
||||
with self._generator_lock:
|
||||
self._generator_state = "failed"
|
||||
self._generator_error = str(exc)
|
||||
raise
|
||||
with self._generator_lock:
|
||||
self._generator = new_gen
|
||||
self._generator_config = config
|
||||
self._generator_state = "ready"
|
||||
self._generator_error = None
|
||||
return new_gen
|
||||
|
||||
def list_generators(self) -> list[dict[str, Any]]:
|
||||
"""The resident slot, or empty when nothing is loaded."""
|
||||
with self._generator_lock:
|
||||
if self._generator_state == "empty":
|
||||
return []
|
||||
return [self._slot_entry()]
|
||||
|
||||
def unload_generator(self, **_ignored: Any) -> bool:
|
||||
"""Shut down and delete the resident generator, freeing GPU memory."""
|
||||
if not self._load_lock.acquire(blocking=False):
|
||||
raise RuntimeError("cannot unload while a model load is in progress")
|
||||
try:
|
||||
running = self._running_inference_jobs()
|
||||
if running:
|
||||
raise RuntimeError(f"cannot unload while inference jobs are running: {running}")
|
||||
with self._generator_lock:
|
||||
gen = self._generator
|
||||
empty = self._generator_state == "empty"
|
||||
self._generator = None
|
||||
self._generator_state = "empty"
|
||||
self._generator_config = None
|
||||
self._generator_error = None
|
||||
if empty:
|
||||
return False
|
||||
if gen is not None:
|
||||
logger.info("Releasing resident generator")
|
||||
gen.shutdown()
|
||||
del gen
|
||||
return True
|
||||
finally:
|
||||
self._load_lock.release()
|
||||
|
||||
def _get_or_create_generator(
|
||||
self,
|
||||
model_id: str,
|
||||
@@ -866,53 +633,69 @@ class JobRunner:
|
||||
sp_size: int = -1,
|
||||
log_queue: mp.Queue | None = None,
|
||||
) -> Any:
|
||||
"""Return the resident generator if the config matches; otherwise
|
||||
replace the slot. Blocks behind any in-flight load — every slot
|
||||
transition happens under ``_load_lock``, so a job can never observe a
|
||||
half-replaced slot."""
|
||||
del log_queue # single-slot generators log via the engine tee
|
||||
config = self._generator_config_dict(
|
||||
model_id=model_id,
|
||||
cache_key = (
|
||||
model_id,
|
||||
workload_type,
|
||||
num_gpus,
|
||||
dit_cpu_offload,
|
||||
text_encoder_cpu_offload,
|
||||
vae_cpu_offload,
|
||||
image_encoder_cpu_offload,
|
||||
use_fsdp_inference,
|
||||
enable_torch_compile,
|
||||
vsa_sparsity,
|
||||
tp_size,
|
||||
sp_size,
|
||||
)
|
||||
|
||||
# Generators are cached by model_id and configuration parameters
|
||||
with self._generators_lock:
|
||||
if cache_key in self._generators:
|
||||
return self._generators[cache_key]
|
||||
|
||||
# Import lazily so starting the server is fast even without a GPU.
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
logger.info(
|
||||
"Loading model %s (workload=%s, num_gpus=%d, offloads: "
|
||||
"dit=%s text_encoder=%s vae=%s image_encoder=%s, fsdp=%s, "
|
||||
"torch_compile=%s, vsa_sparsity=%.2f, tp=%d sp=%d) …",
|
||||
model_id,
|
||||
workload_type,
|
||||
num_gpus,
|
||||
dit_cpu_offload,
|
||||
text_encoder_cpu_offload,
|
||||
vae_cpu_offload,
|
||||
image_encoder_cpu_offload,
|
||||
use_fsdp_inference,
|
||||
enable_torch_compile,
|
||||
vsa_sparsity,
|
||||
tp_size,
|
||||
sp_size,
|
||||
)
|
||||
|
||||
gen = VideoGenerator.from_pretrained(
|
||||
model_id,
|
||||
workload_type=workload_type,
|
||||
num_gpus=num_gpus,
|
||||
dit_cpu_offload=dit_cpu_offload,
|
||||
text_encoder_cpu_offload=text_encoder_cpu_offload,
|
||||
vae_cpu_offload=vae_cpu_offload,
|
||||
image_encoder_cpu_offload=image_encoder_cpu_offload,
|
||||
use_fsdp_inference=use_fsdp_inference,
|
||||
enable_torch_compile=enable_torch_compile,
|
||||
vsa_sparsity=vsa_sparsity,
|
||||
VSA_sparsity=vsa_sparsity,
|
||||
tp_size=tp_size,
|
||||
sp_size=sp_size,
|
||||
log_queue=log_queue,
|
||||
)
|
||||
with self._load_lock: # waits out preloads / other jobs' replaces
|
||||
with self._generator_lock:
|
||||
if (self._generator_state == "ready"
|
||||
and self._generator_config == config
|
||||
and self._generator is not None):
|
||||
return self._generator
|
||||
return self._run_on_loader(lambda: self._load_into_slot_locked(config))
|
||||
|
||||
def feed_engine_line(self, line: str) -> None:
|
||||
"""Bridge ray-relayed worker output into the running job's log buffer.
|
||||
|
||||
On the ray backend worker logs cannot cross nodes via the mp queue,
|
||||
but ray already relays them to the driver's stdout — which the engine
|
||||
tee captures. Lines with ray's actor prefix are attributed to the one
|
||||
running inference job, whose buffer parses tqdm into UI progress.
|
||||
Driver-side logging is excluded (it reaches the buffer via the
|
||||
logging handlers already).
|
||||
"""
|
||||
job = self._active_inference_job
|
||||
if job is None:
|
||||
return
|
||||
line = _ANSI_RE.sub("", line)
|
||||
if not _RAY_RELAY_RE.match(line):
|
||||
return
|
||||
try:
|
||||
job._log_buf.write(line)
|
||||
except Exception: # noqa: BLE001 -- never break the tee
|
||||
pass
|
||||
with self._generators_lock:
|
||||
if cache_key not in self._generators:
|
||||
self._generators[cache_key] = gen
|
||||
else: # Another thread may have created it while we were loading.
|
||||
gen.shutdown()
|
||||
gen = self._generators[cache_key]
|
||||
return gen
|
||||
|
||||
def _run_job(self, job: Job):
|
||||
if job.job_type == "inference":
|
||||
@@ -1058,12 +841,10 @@ class JobRunner:
|
||||
fastvideo_logger.addHandler(buffer_handler)
|
||||
fastvideo_logger.addHandler(file_handler)
|
||||
|
||||
# Worker logs flow through the runner-wide queue the generator was
|
||||
# created with; drain anything stale, then listen for this job.
|
||||
log_queue = self._worker_log_queue
|
||||
with contextlib.suppress(Exception):
|
||||
while True:
|
||||
log_queue.get_nowait()
|
||||
# Queue for worker process logs (fsdp_load, cuda, etc.)
|
||||
# Use Manager().Queue() so it can be shared with spawned workers (spawn
|
||||
# does not inherit memory; mp.Queue only works through inheritance).
|
||||
log_queue = self._mp_manager.Queue()
|
||||
queue_listener = logging.handlers.QueueListener(log_queue,
|
||||
buffer_handler,
|
||||
file_handler,
|
||||
@@ -1145,7 +926,6 @@ class JobRunner:
|
||||
|
||||
generator = _gen_result[0]
|
||||
buf.phase = "generating"
|
||||
self._active_inference_job = job # engine tee feeds tqdm from here
|
||||
logger.info("Starting generation for job %s (model=%s)", job.id, job.model_id)
|
||||
|
||||
gen_kwargs: dict[str, Any] = {
|
||||
@@ -1161,6 +941,7 @@ class JobRunner:
|
||||
"fps": job.fps,
|
||||
"seed": job.seed,
|
||||
"negative_prompt": job.negative_prompt or "",
|
||||
"log_queue": log_queue,
|
||||
}
|
||||
if job.image_path:
|
||||
gen_kwargs["image_path"] = job.image_path
|
||||
@@ -1204,8 +985,6 @@ class JobRunner:
|
||||
buf.phase = "failed"
|
||||
|
||||
finally:
|
||||
if self._active_inference_job is job:
|
||||
self._active_inference_job = None
|
||||
queue_listener.stop()
|
||||
# Remove handlers and close file
|
||||
fastvideo_logger.removeHandler(buffer_handler)
|
||||
|
||||
@@ -36,15 +36,13 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse, PlainTextResponse
|
||||
|
||||
from fastvideo_studio.database import default_settings_dict
|
||||
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, GeneratorRequest, SettingsUpdate,
|
||||
UpdateCaptionRequest, model_label)
|
||||
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, SettingsUpdate, UpdateCaptionRequest,
|
||||
model_label)
|
||||
|
||||
# --- Config -----------------------------------------------------------------
|
||||
|
||||
# How long a started job stays "running" before it flips to "completed".
|
||||
COMPLETE_AFTER_SECONDS = 3.0
|
||||
# How long a preloading generator stays "loading" before it flips to "ready".
|
||||
GENERATOR_READY_AFTER_SECONDS = 2.0
|
||||
FFMPEG_BIN = shutil.which(os.getenv("FASTVIDEO_FFMPEG_BIN", "ffmpeg"))
|
||||
|
||||
# A small catalogue of fake models keyed by workload type. Mirrors the real
|
||||
@@ -90,10 +88,6 @@ _DEFAULT_SETTINGS: dict[str, Any] = {
|
||||
_state_lock = threading.Lock()
|
||||
_settings: dict[str, Any] = dict(_DEFAULT_SETTINGS)
|
||||
_jobs: dict[str, dict[str, Any]] = {}
|
||||
# The engine's single model slot: None when empty, else the state dict.
|
||||
_generator_slot: dict[str, Any] | None = None
|
||||
# Fake engine stdout/stderr tail; grows a little on every poll.
|
||||
_engine_log_lines: list[str] = ["[engine] FastVideo studio mock engine started"]
|
||||
_datasets: dict[str, dict[str, Any]] = {}
|
||||
# dataset_id -> {"file_names": [...], "captions": {file_name: caption}}
|
||||
_dataset_files: dict[str, dict[str, Any]] = {}
|
||||
@@ -399,56 +393,6 @@ def list_models(workload_type: str | None = None) -> list[dict[str, Any]]:
|
||||
return _models_for(workload_type)
|
||||
|
||||
|
||||
# Per-model sampling presets, mirroring the real /api/models/presets shape
|
||||
# (keys the model has no recommendation for are simply absent).
|
||||
_MODEL_PRESETS: dict[str, dict[str, Any]] = {
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"fps": 16,
|
||||
"num_inference_steps": 50,
|
||||
"guidance_scale": 3.0,
|
||||
"guidance_rescale": 0.0,
|
||||
"negative_prompt": "Bright tones, overexposed, static, blurred details",
|
||||
"seed": 1024,
|
||||
},
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"fps": 16,
|
||||
"num_inference_steps": 40,
|
||||
"guidance_scale": 5.0,
|
||||
"seed": 1024,
|
||||
},
|
||||
"black-forest-labs/FLUX.1-schnell": {
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 0.0,
|
||||
"seed": 42,
|
||||
},
|
||||
}
|
||||
_GENERIC_PRESETS: dict[str, Any] = {
|
||||
"height": 720,
|
||||
"width": 1280,
|
||||
"num_frames": 81,
|
||||
"fps": 24,
|
||||
"num_inference_steps": 50,
|
||||
"guidance_scale": 5.0,
|
||||
"guidance_rescale": 0.0,
|
||||
"seed": 1024,
|
||||
}
|
||||
|
||||
|
||||
@app.get("/api/models/presets")
|
||||
def model_presets(model_id: str) -> dict[str, Any]:
|
||||
"""Recommended sampling settings; unknown models get generic defaults."""
|
||||
return _MODEL_PRESETS.get(model_id, _GENERIC_PRESETS)
|
||||
|
||||
|
||||
# --- GPUs -------------------------------------------------------------------
|
||||
|
||||
|
||||
@@ -470,49 +414,6 @@ def list_gpus() -> dict[str, Any]:
|
||||
return {"available": True, "gpus": gpus, "error": None}
|
||||
|
||||
|
||||
@app.get("/api/cluster")
|
||||
def cluster_status() -> dict[str, Any]:
|
||||
"""Two fake ray nodes x 4 GPUs with per-request jitter, mirroring the
|
||||
real server's /api/cluster shape."""
|
||||
nodes = []
|
||||
for host_idx, (hostname, ip, is_this_host) in enumerate([
|
||||
("mock-node-0", "10.0.0.10", True),
|
||||
("mock-node-1", "10.0.0.11", False),
|
||||
]):
|
||||
gpus = []
|
||||
for index in range(4):
|
||||
base_util = (13 + 29 * index + 41 * host_idx) % 100
|
||||
gpus.append({
|
||||
"index": index,
|
||||
"name": "NVIDIA Mock GPU 80GB",
|
||||
"utilization": max(0, min(100, base_util + random.randint(-5, 5))),
|
||||
"memory_used_mib": 6_144 + 17_408 * index + random.randint(-256, 256),
|
||||
"memory_total_mib": 81_920,
|
||||
"temperature_c": 42 + 6 * index + random.randint(-3, 3),
|
||||
"power_watts": 110.0 + 140.0 * index + random.randint(-20, 20),
|
||||
"power_limit_watts": 700.0,
|
||||
})
|
||||
nodes.append({
|
||||
"hostname": hostname,
|
||||
"ip": ip,
|
||||
"is_this_host": is_this_host,
|
||||
"cpus": 64.0,
|
||||
"ray_gpus": 4.0,
|
||||
"available": True,
|
||||
"error": None,
|
||||
"gpus": gpus,
|
||||
})
|
||||
return {
|
||||
"mode": "ray",
|
||||
"nodes": nodes,
|
||||
"resources": {
|
||||
"gpus_total": 8.0,
|
||||
"gpus_available": 5.0
|
||||
},
|
||||
"error": None,
|
||||
}
|
||||
|
||||
|
||||
# --- Uploads ----------------------------------------------------------------
|
||||
|
||||
|
||||
@@ -669,92 +570,6 @@ def download_log(job_id: str) -> PlainTextResponse:
|
||||
return PlainTextResponse("\n".join(lines) + "\n", media_type="text/plain")
|
||||
|
||||
|
||||
# --- Generators (warm models) -------------------------------------------------
|
||||
|
||||
|
||||
def _advance_generator(entry: dict[str, Any]) -> None:
|
||||
"""Flip a loading generator to ready once enough wall-clock time has passed.
|
||||
|
||||
Like job status, generator state is *computed on read*, so polling the
|
||||
generators list naturally shows loading -> ready.
|
||||
"""
|
||||
if entry["state"] == "loading" and time.time() - entry["started_at"] >= GENERATOR_READY_AFTER_SECONDS:
|
||||
entry["state"] = "ready"
|
||||
|
||||
|
||||
def _running_inference_ids() -> list[str]:
|
||||
return [
|
||||
j["id"] for j in _jobs.values() if j.get("job_type") == "inference" and _public_job(j)["status"] == "running"
|
||||
]
|
||||
|
||||
|
||||
@app.get("/api/generators")
|
||||
def list_generators() -> list[dict[str, Any]]:
|
||||
with _state_lock:
|
||||
if _generator_slot is None:
|
||||
return []
|
||||
_advance_generator(_generator_slot)
|
||||
return [dict(_generator_slot)]
|
||||
|
||||
|
||||
@app.post("/api/generators/preload", status_code=202)
|
||||
def preload_generator(req: GeneratorRequest) -> dict[str, Any]:
|
||||
global _generator_slot
|
||||
valid_ids = {m["id"] for m in _models_for(None)}
|
||||
if req.model_id not in valid_ids:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Unknown model_id '{req.model_id}'. Valid options: {sorted(valid_ids)}",
|
||||
)
|
||||
with _state_lock:
|
||||
if _generator_slot is not None:
|
||||
_advance_generator(_generator_slot)
|
||||
if _generator_slot["state"] == "loading":
|
||||
raise HTTPException(status_code=409, detail="a model load is already in progress")
|
||||
if _generator_slot["state"] == "ready" and all(
|
||||
_generator_slot.get(k) == v for k, v in req.model_dump().items()):
|
||||
return dict(_generator_slot)
|
||||
running = _running_inference_ids()
|
||||
if running:
|
||||
raise HTTPException(status_code=409,
|
||||
detail=f"cannot swap models while inference jobs are running: {running}")
|
||||
# Loading a new model always replaces (releases) the resident one.
|
||||
_generator_slot = {"state": "loading", "started_at": time.time(), "error": None, **req.model_dump()}
|
||||
return dict(_generator_slot)
|
||||
|
||||
|
||||
@app.post("/api/generators/unload")
|
||||
def unload_generator() -> dict[str, Any]:
|
||||
global _generator_slot
|
||||
with _state_lock:
|
||||
if _generator_slot is not None:
|
||||
_advance_generator(_generator_slot)
|
||||
if _generator_slot["state"] == "loading":
|
||||
raise HTTPException(status_code=409, detail="cannot unload while a model load is in progress")
|
||||
if _generator_slot is None:
|
||||
raise HTTPException(status_code=404, detail="No model is loaded")
|
||||
running = _running_inference_ids()
|
||||
if running:
|
||||
raise HTTPException(status_code=409, detail=f"cannot unload while inference jobs are running: {running}")
|
||||
_generator_slot = None
|
||||
return {"unloaded": True}
|
||||
|
||||
|
||||
# --- Engine logs --------------------------------------------------------------
|
||||
|
||||
|
||||
@app.get("/api/engine/logs")
|
||||
def engine_logs(after: int = 0) -> dict[str, Any]:
|
||||
"""Growing fake tail of the engine's stdout/stderr: every poll appends a
|
||||
couple of lines so the console visibly streams."""
|
||||
with _state_lock:
|
||||
n = len(_engine_log_lines)
|
||||
_engine_log_lines.append(f"[engine] step {n}: worker heartbeat ok")
|
||||
_engine_log_lines.append(f"[engine] step {n + 1}: gpu mem {random.randint(20, 80)}% used")
|
||||
total = len(_engine_log_lines)
|
||||
return {"lines": _engine_log_lines[max(0, after):], "total": total}
|
||||
|
||||
|
||||
# --- Datasets ---------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
usable by both the real server and the dependency-light mock server."""
|
||||
|
||||
from fastvideo_studio.models.create_job_request import CreateJobRequest
|
||||
from fastvideo_studio.models.generator_request import GeneratorRequest
|
||||
from fastvideo_studio.models.settings_update import SettingsUpdate
|
||||
from fastvideo_studio.models.create_dataset_request import CreateDatasetRequest
|
||||
from fastvideo_studio.models.update_caption_request import UpdateCaptionRequest
|
||||
@@ -16,7 +15,6 @@ def model_label(model_path: str) -> str:
|
||||
|
||||
__all__ = [
|
||||
"CreateJobRequest",
|
||||
"GeneratorRequest",
|
||||
"SettingsUpdate",
|
||||
"CreateDatasetRequest",
|
||||
"UpdateCaptionRequest",
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Request model for preloading/unloading a resident generator.
|
||||
|
||||
Field names and defaults mirror the engine subset of ``CreateJobRequest`` so
|
||||
the UI can send exactly the values it would put on a job — guaranteeing the
|
||||
job's generator lookup hits this cache entry.
|
||||
"""
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class GeneratorRequest(BaseModel):
|
||||
model_id: str
|
||||
workload_type: str = "t2v"
|
||||
num_gpus: int = 1
|
||||
dit_cpu_offload: bool = False
|
||||
text_encoder_cpu_offload: bool = False
|
||||
vae_cpu_offload: bool = False
|
||||
image_encoder_cpu_offload: bool = False
|
||||
use_fsdp_inference: bool = False
|
||||
enable_torch_compile: bool = False
|
||||
vsa_sparsity: float = 0.0
|
||||
tp_size: int = -1
|
||||
sp_size: int = -1
|
||||
@@ -32,10 +32,10 @@ from fastapi.responses import FileResponse
|
||||
|
||||
from fastvideo.registry import (get_registered_model_paths, get_registered_models_with_workloads)
|
||||
from fastvideo_studio.database import Database, _get_db_path
|
||||
from fastvideo_studio.gpu import get_cluster_snapshot, get_gpu_snapshot
|
||||
from fastvideo_studio.gpu import get_gpu_snapshot
|
||||
from fastvideo_studio.job_runner import JobRunner, JobStatus
|
||||
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, GeneratorRequest, SettingsUpdate,
|
||||
UpdateCaptionRequest, model_label)
|
||||
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, SettingsUpdate, UpdateCaptionRequest,
|
||||
model_label)
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
@@ -45,67 +45,6 @@ logger = logging.getLogger("fastvideo.studio.api")
|
||||
|
||||
DEFAULT_OUTPUT_DIR = os.path.join(os.path.dirname(__file__), "..", "outputs", "ui_jobs")
|
||||
|
||||
|
||||
class _EngineLogBuffer:
|
||||
"""Thread-safe ring buffer over the server process's stdout/stderr.
|
||||
|
||||
Because ray relays worker output to the driver (log_to_driver), teeing the
|
||||
server's own streams captures engine output from every rank on every node;
|
||||
under the mp executor, worker logs arrive via the logging handlers which
|
||||
also write to stderr.
|
||||
"""
|
||||
|
||||
def __init__(self, maxlen: int = 5000) -> None:
|
||||
import collections
|
||||
import threading
|
||||
self._lines: collections.deque[str] = collections.deque(maxlen=maxlen)
|
||||
self._dropped = 0
|
||||
self._lock = threading.Lock()
|
||||
self._partial = ""
|
||||
|
||||
on_line = None # optional callable(str), set once at startup
|
||||
|
||||
def write(self, text: str) -> None:
|
||||
with self._lock:
|
||||
buf = self._partial + text
|
||||
*complete, self._partial = buf.split("\n")
|
||||
for line in complete:
|
||||
if len(self._lines) == self._lines.maxlen:
|
||||
self._dropped += 1
|
||||
self._lines.append(line)
|
||||
if self.on_line is not None:
|
||||
for line in complete:
|
||||
# never break stdout on a bad feed
|
||||
with contextlib.suppress(Exception):
|
||||
self.on_line(line)
|
||||
|
||||
def get_lines(self, after: int = 0) -> tuple[list[str], int]:
|
||||
with self._lock:
|
||||
total = self._dropped + len(self._lines)
|
||||
start = max(0, after - self._dropped)
|
||||
return list(self._lines)[start:], total
|
||||
|
||||
|
||||
class _Tee:
|
||||
"""File-like that forwards to the original stream and the ring buffer."""
|
||||
|
||||
def __init__(self, orig: Any, buffer: _EngineLogBuffer) -> None:
|
||||
self._orig = orig
|
||||
self._buffer = buffer
|
||||
|
||||
def write(self, text: str) -> int:
|
||||
self._buffer.write(text)
|
||||
return self._orig.write(text)
|
||||
|
||||
def flush(self) -> None:
|
||||
self._orig.flush()
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._orig, name)
|
||||
|
||||
|
||||
engine_log = _EngineLogBuffer()
|
||||
|
||||
_available_models: list[dict[str, str]] = [{
|
||||
"id": path,
|
||||
"label": model_label(path)
|
||||
@@ -163,13 +102,6 @@ def list_gpus() -> dict[str, Any]:
|
||||
return get_gpu_snapshot()
|
||||
|
||||
|
||||
@app.get("/api/cluster")
|
||||
def cluster_status() -> dict[str, Any]:
|
||||
"""Cluster-wide GPU/host telemetry (per-node NVML via ray when connected,
|
||||
the local host otherwise)."""
|
||||
return get_cluster_snapshot()
|
||||
|
||||
|
||||
@app.get("/api/models")
|
||||
def list_models(workload_type: str | None = None) -> list[dict[str, Any]]:
|
||||
"""Return the catalogue of available video-generation models.
|
||||
@@ -184,25 +116,6 @@ def list_models(workload_type: str | None = None) -> list[dict[str, Any]]:
|
||||
return _available_models
|
||||
|
||||
|
||||
_PRESET_FIELDS = ("height", "width", "num_frames", "fps", "num_inference_steps", "guidance_scale", "guidance_rescale",
|
||||
"negative_prompt", "seed")
|
||||
_preset_cache: dict[str, dict[str, Any]] = {}
|
||||
|
||||
|
||||
@app.get("/api/models/presets")
|
||||
def model_presets(model_id: str) -> dict[str, Any]:
|
||||
"""The model's recommended sampling settings (config-only — never loads
|
||||
weights). The UI populates the job form from these on model selection."""
|
||||
if model_id not in _preset_cache:
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
try:
|
||||
sp = SamplingParam.from_pretrained(model_id)
|
||||
except Exception as exc: # noqa: BLE001 -- unknown/unresolvable model
|
||||
raise HTTPException(status_code=404, detail=f"No presets for '{model_id}': {exc}") from exc
|
||||
_preset_cache[model_id] = {f: getattr(sp, f) for f in _PRESET_FIELDS if getattr(sp, f, None) is not None}
|
||||
return _preset_cache[model_id]
|
||||
|
||||
|
||||
ALLOWED_IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
|
||||
|
||||
|
||||
@@ -340,58 +253,13 @@ def get_job(job_id: str) -> dict[str, Any]:
|
||||
return job.to_dict()
|
||||
|
||||
|
||||
@app.get("/api/generators")
|
||||
def list_generators() -> list[dict[str, Any]]:
|
||||
"""Generators resident in memory plus preloads in flight or failed."""
|
||||
return job_runner.list_generators()
|
||||
|
||||
|
||||
@app.post("/api/generators/preload", status_code=202)
|
||||
def preload_generator(req: GeneratorRequest) -> dict[str, Any]:
|
||||
"""Load a model into memory ahead of time (replacing whatever is
|
||||
resident). One load at a time — 409 while another load is in flight."""
|
||||
valid_ids = {m["id"] for m in _available_models}
|
||||
if req.model_id not in valid_ids and not os.path.isdir(req.model_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(f"Unknown model_id '{req.model_id}'. "
|
||||
f"Valid options: {sorted(valid_ids)}"),
|
||||
)
|
||||
try:
|
||||
return job_runner.preload_generator(**req.model_dump())
|
||||
except RuntimeError as exc:
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@app.post("/api/generators/unload")
|
||||
def unload_generator() -> dict[str, Any]:
|
||||
"""Shut down and delete the resident generator, freeing GPU memory."""
|
||||
try:
|
||||
unloaded = job_runner.unload_generator()
|
||||
except RuntimeError as exc:
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
if not unloaded:
|
||||
raise HTTPException(status_code=404, detail="No model is loaded")
|
||||
return {"unloaded": True}
|
||||
|
||||
|
||||
@app.get("/api/engine/logs")
|
||||
def engine_logs(after: int = 0) -> dict[str, Any]:
|
||||
"""Incremental tail of the engine's stdout/stderr (driver + relayed
|
||||
worker output). Poll with ?after=<total from the previous response>."""
|
||||
lines, total = engine_log.get_lines(after=after)
|
||||
return {"lines": lines, "total": total}
|
||||
|
||||
|
||||
@app.post("/api/jobs", status_code=201)
|
||||
def create_job(req: CreateJobRequest) -> dict[str, Any]:
|
||||
"""Create a new job (does **not** start it automatically)."""
|
||||
job_type = req.job_type or "inference"
|
||||
if job_type == "inference":
|
||||
valid_ids = {m["id"] for m in _available_models}
|
||||
# A local weights directory is as valid as a registered hub id —
|
||||
# the registry resolves the pipeline from its model_index.
|
||||
if req.model_id not in valid_ids and not os.path.isdir(req.model_id):
|
||||
if req.model_id not in valid_ids:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(f"Unknown model_id '{req.model_id}'. "
|
||||
@@ -754,14 +622,6 @@ def create_local_env(host: str, port: int) -> None:
|
||||
|
||||
|
||||
def main() -> None:
|
||||
import sys
|
||||
sys.stdout = _Tee(sys.stdout, engine_log)
|
||||
sys.stderr = _Tee(sys.stderr, engine_log)
|
||||
# handlers created before the tee (module-level basicConfig) hold the
|
||||
# original stream objects — re-point them or their output bypasses the buffer
|
||||
for h in logging.getLogger().handlers:
|
||||
if isinstance(h, logging.StreamHandler) and h.stream in (sys.__stderr__, sys.__stdout__):
|
||||
h.setStream(sys.stderr) # type: ignore[arg-type] # duck-typed file-like
|
||||
global job_runner, database, upload_dir, verbose, datasets_upload_dir # noqa: PLW0603
|
||||
|
||||
# Set up signal handlers to prevent worker crashes from killing the server
|
||||
@@ -826,9 +686,6 @@ def main() -> None:
|
||||
verbose=args.verbose,
|
||||
database=database,
|
||||
)
|
||||
# ray relays worker output (incl. denoising tqdm) to the driver's stdout;
|
||||
# feed those lines to the running job so the UI progress bar moves.
|
||||
engine_log.on_line = job_runner.feed_engine_line
|
||||
|
||||
logger.info("Output directory: %s", output_dir)
|
||||
logger.info("Log directory: %s", log_dir)
|
||||
@@ -838,10 +695,6 @@ def main() -> None:
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
log_level="info",
|
||||
# The engine-output console tails this process's stdout/stderr; the
|
||||
# frontend polls several endpoints every few seconds, so access-log
|
||||
# lines are pure self-noise there. App/job logging is unaffected.
|
||||
access_log=False,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,182 +0,0 @@
|
||||
import { render, screen } from '@testing-library/react';
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
|
||||
import GpusPage from './page';
|
||||
import { getClusterStatus } from '@/lib/api';
|
||||
import type { ClusterSnapshot } from '@/lib/api';
|
||||
|
||||
vi.mock('@/lib/api', () => ({
|
||||
getClusterStatus: vi.fn(),
|
||||
}));
|
||||
|
||||
const RAY_SNAPSHOT: ClusterSnapshot = {
|
||||
mode: 'ray',
|
||||
error: null,
|
||||
resources: { gpus_total: 8, gpus_available: 5 },
|
||||
nodes: [
|
||||
{
|
||||
hostname: 'node-a',
|
||||
ip: '10.0.0.10',
|
||||
is_this_host: true,
|
||||
cpus: 64,
|
||||
ray_gpus: 4,
|
||||
available: true,
|
||||
error: null,
|
||||
gpus: [
|
||||
{
|
||||
index: 0,
|
||||
name: 'NVIDIA B200',
|
||||
utilization: 62,
|
||||
memory_used_mib: 40_960,
|
||||
memory_total_mib: 81_920,
|
||||
temperature_c: 41,
|
||||
power_watts: 312.4,
|
||||
power_limit_watts: 1000,
|
||||
},
|
||||
{
|
||||
index: 1,
|
||||
name: 'NVIDIA B200',
|
||||
utilization: 0,
|
||||
memory_used_mib: 1_024,
|
||||
memory_total_mib: 81_920,
|
||||
temperature_c: null,
|
||||
power_watts: null,
|
||||
power_limit_watts: null,
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
hostname: 'node-b',
|
||||
ip: '10.0.0.11',
|
||||
is_this_host: false,
|
||||
cpus: 32,
|
||||
ray_gpus: 2,
|
||||
available: true,
|
||||
error: null,
|
||||
gpus: [
|
||||
{
|
||||
index: 0,
|
||||
name: 'NVIDIA B200',
|
||||
utilization: 90,
|
||||
memory_used_mib: 20_480,
|
||||
memory_total_mib: 81_920,
|
||||
temperature_c: 70,
|
||||
power_watts: 900,
|
||||
power_limit_watts: 1000,
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
const LOCAL_SNAPSHOT: ClusterSnapshot = {
|
||||
mode: 'local',
|
||||
error:
|
||||
'not connected to a ray cluster yet (load a model first); showing the API host only',
|
||||
resources: null,
|
||||
nodes: [
|
||||
{
|
||||
hostname: 'localhost',
|
||||
ip: null,
|
||||
is_this_host: true,
|
||||
cpus: null,
|
||||
ray_gpus: null,
|
||||
available: true,
|
||||
error: null,
|
||||
gpus: [
|
||||
{
|
||||
index: 0,
|
||||
name: 'NVIDIA RTX 5090',
|
||||
utilization: 12,
|
||||
memory_used_mib: 2_048,
|
||||
memory_total_mib: 32_768,
|
||||
temperature_c: 38,
|
||||
power_watts: 80,
|
||||
power_limit_watts: 575,
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
beforeEach(() => {
|
||||
vi.mocked(getClusterStatus).mockResolvedValue(RAY_SNAPSHOT);
|
||||
});
|
||||
|
||||
describe('GpusPage', () => {
|
||||
it('renders the header with mode and GPU totals', async () => {
|
||||
render(<GpusPage />);
|
||||
|
||||
expect(await screen.findByText('ray cluster')).toBeInTheDocument();
|
||||
expect(screen.getByText(/5 \/\s*8 GPUs available/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('renders a section per node with host details and GPU rows', async () => {
|
||||
render(<GpusPage />);
|
||||
|
||||
// Each hostname appears twice: once in the strip, once as a section.
|
||||
expect(await screen.findAllByText('node-a')).toHaveLength(2);
|
||||
expect(screen.getAllByText('node-b')).toHaveLength(2);
|
||||
expect(screen.getByText('10.0.0.10')).toBeInTheDocument();
|
||||
// Only node-a is the API host.
|
||||
expect(screen.getAllByText('API host')).toHaveLength(1);
|
||||
expect(screen.getByText(/64 CPUs · 4 ray GPUs/)).toBeInTheDocument();
|
||||
|
||||
expect(screen.getAllByText('NVIDIA B200')).toHaveLength(3);
|
||||
expect(screen.getByText('GPU 1')).toBeInTheDocument();
|
||||
expect(screen.getByText('62%')).toBeInTheDocument();
|
||||
expect(
|
||||
screen.getByText('40960 / 81920 MiB (40.0 GiB / 80.0 GiB)'),
|
||||
).toBeInTheDocument();
|
||||
// Optional sensors render only when present.
|
||||
expect(screen.getByText('41°C')).toBeInTheDocument();
|
||||
expect(screen.getByText('312 W / 1000 W')).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('bars reflect utilization and VRAM values', async () => {
|
||||
render(<GpusPage />);
|
||||
await screen.findAllByText('node-a');
|
||||
|
||||
const utilMeters = screen
|
||||
.getAllByRole('meter', { name: 'Utilization' })
|
||||
.map((m) => m.getAttribute('aria-valuenow'));
|
||||
expect(utilMeters).toEqual(['62', '0', '90']);
|
||||
|
||||
const vramMeters = screen
|
||||
.getAllByRole('meter', { name: 'VRAM' })
|
||||
.map((m) => m.getAttribute('aria-valuenow'));
|
||||
// 40960/81920 = 50%, 1024/81920 ≈ 1%, 20480/81920 = 25%
|
||||
expect(vramMeters).toEqual(['50', '1', '25']);
|
||||
});
|
||||
|
||||
it('renders the compact strip with per-GPU segments', async () => {
|
||||
render(<GpusPage />);
|
||||
await screen.findAllByText('node-a');
|
||||
|
||||
const segments = screen.getAllByRole('img');
|
||||
expect(segments).toHaveLength(3);
|
||||
expect(segments[0]).toHaveAccessibleName(
|
||||
'GPU 0: 62% utilization, 40.0 GiB / 80.0 GiB VRAM',
|
||||
);
|
||||
});
|
||||
|
||||
it('shows the informational banner and local mode', async () => {
|
||||
vi.mocked(getClusterStatus).mockResolvedValue(LOCAL_SNAPSHOT);
|
||||
render(<GpusPage />);
|
||||
|
||||
expect(await screen.findByText('local host only')).toBeInTheDocument();
|
||||
expect(
|
||||
screen.getByText(/not connected to a ray cluster yet/),
|
||||
).toBeInTheDocument();
|
||||
// No resources in local mode.
|
||||
expect(screen.queryByText(/GPUs available/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('explains when the API server is unreachable', async () => {
|
||||
vi.mocked(getClusterStatus).mockRejectedValue(new Error('network down'));
|
||||
render(<GpusPage />);
|
||||
expect(
|
||||
await screen.findByText(/Could not reach the API server/),
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
@@ -1,258 +1,11 @@
|
||||
'use client';
|
||||
|
||||
import * as React from 'react';
|
||||
import { AlertTriangle, Info } from 'lucide-react';
|
||||
|
||||
import ClusterStrip, {
|
||||
clampPercent,
|
||||
formatGib,
|
||||
utilizationColor,
|
||||
} from '@/components/cluster/ClusterStrip';
|
||||
import { Badge } from '@/components/ui/badge';
|
||||
import { Button } from '@/components/ui/button';
|
||||
import { Card, CardContent } from '@/components/ui/card';
|
||||
import {
|
||||
getClusterStatus,
|
||||
type ClusterGpu,
|
||||
type ClusterNode,
|
||||
type ClusterSnapshot,
|
||||
} from '@/lib/api';
|
||||
import { cn } from '@/lib/utils';
|
||||
|
||||
const POLL_INTERVAL_MS = 5000;
|
||||
|
||||
function Meter({
|
||||
label,
|
||||
percent,
|
||||
detail,
|
||||
fillClass,
|
||||
}: {
|
||||
label: string;
|
||||
percent: number;
|
||||
detail: string;
|
||||
fillClass: string;
|
||||
}) {
|
||||
const clamped = clampPercent(percent);
|
||||
return (
|
||||
<div className="flex flex-col gap-1">
|
||||
<div className="flex items-baseline justify-between gap-2 text-xs">
|
||||
<span className="text-muted-foreground">{label}</span>
|
||||
<span className="font-medium tabular-nums text-foreground">
|
||||
{detail}
|
||||
</span>
|
||||
</div>
|
||||
<div
|
||||
role="meter"
|
||||
aria-label={label}
|
||||
aria-valuenow={Math.round(clamped)}
|
||||
aria-valuemin={0}
|
||||
aria-valuemax={100}
|
||||
className="h-1.5 overflow-hidden rounded-full bg-muted"
|
||||
>
|
||||
<div
|
||||
className={cn(
|
||||
'h-full rounded-full transition-[width] duration-500',
|
||||
fillClass,
|
||||
)}
|
||||
style={{ width: `${clamped}%` }}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
function GpuRow({ gpu }: { gpu: ClusterGpu }) {
|
||||
const memPercent =
|
||||
gpu.memory_total_mib > 0
|
||||
? (gpu.memory_used_mib / gpu.memory_total_mib) * 100
|
||||
: 0;
|
||||
return (
|
||||
<div className="grid items-center gap-x-6 gap-y-2 border-t border-border pt-3 first:border-t-0 first:pt-0 md:grid-cols-[minmax(0,1fr)_minmax(0,1.2fr)_minmax(0,1.6fr)_auto]">
|
||||
<div className="flex min-w-0 items-baseline gap-2">
|
||||
<span className="min-w-0 truncate text-sm font-semibold">
|
||||
{gpu.name}
|
||||
</span>
|
||||
<span className="shrink-0 text-xs font-medium uppercase tracking-wider text-muted-foreground">
|
||||
GPU {gpu.index}
|
||||
</span>
|
||||
</div>
|
||||
<Meter
|
||||
label="Utilization"
|
||||
percent={gpu.utilization}
|
||||
detail={`${gpu.utilization}%`}
|
||||
fillClass={utilizationColor(gpu.utilization)}
|
||||
/>
|
||||
<Meter
|
||||
label="VRAM"
|
||||
percent={memPercent}
|
||||
detail={`${gpu.memory_used_mib} / ${gpu.memory_total_mib} MiB (${formatGib(gpu.memory_used_mib)} / ${formatGib(gpu.memory_total_mib)})`}
|
||||
fillClass={memPercent >= 90 ? 'bg-rose-500' : 'bg-accent-blue'}
|
||||
/>
|
||||
<div className="flex flex-wrap gap-x-4 gap-y-1 text-xs tabular-nums text-muted-foreground md:w-28 md:justify-end">
|
||||
{gpu.temperature_c != null && <span>{gpu.temperature_c}°C</span>}
|
||||
{gpu.power_watts != null && (
|
||||
<span>
|
||||
{Math.round(gpu.power_watts)} W
|
||||
{gpu.power_limit_watts != null &&
|
||||
` / ${Math.round(gpu.power_limit_watts)} W`}
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
function NodeSection({ node }: { node: ClusterNode }) {
|
||||
return (
|
||||
<Card>
|
||||
<CardContent className="flex flex-col gap-3 p-5">
|
||||
<div className="flex flex-wrap items-center gap-x-3 gap-y-1">
|
||||
<span className="min-w-0 truncate text-sm font-semibold">
|
||||
{node.hostname}
|
||||
</span>
|
||||
{node.ip && (
|
||||
<span className="text-xs tabular-nums text-muted-foreground">
|
||||
{node.ip}
|
||||
</span>
|
||||
)}
|
||||
{node.is_this_host && <Badge variant="secondary">API host</Badge>}
|
||||
<span className="ml-auto text-xs tabular-nums text-muted-foreground">
|
||||
{node.cpus != null && `${Math.round(node.cpus)} CPUs`}
|
||||
{node.cpus != null && node.ray_gpus != null && ' · '}
|
||||
{node.ray_gpus != null && `${Math.round(node.ray_gpus)} ray GPUs`}
|
||||
</span>
|
||||
</div>
|
||||
{!node.available && (
|
||||
<p className="text-sm text-muted-foreground">
|
||||
GPU telemetry unavailable
|
||||
{node.error ? `: ${node.error}` : '.'}
|
||||
</p>
|
||||
)}
|
||||
{node.gpus.map((gpu) => (
|
||||
<GpuRow key={gpu.index} gpu={gpu} />
|
||||
))}
|
||||
</CardContent>
|
||||
</Card>
|
||||
);
|
||||
}
|
||||
import GpuGrid from '@/components/system/GpuGrid';
|
||||
|
||||
export default function GpusPage() {
|
||||
const [snapshot, setSnapshot] = React.useState<ClusterSnapshot | null>(null);
|
||||
const [fetchError, setFetchError] = React.useState<string | null>(null);
|
||||
const [retryToken, setRetryToken] = React.useState(0);
|
||||
|
||||
React.useEffect(() => {
|
||||
let mounted = true;
|
||||
let inFlight = false;
|
||||
|
||||
async function poll() {
|
||||
if (inFlight || document.hidden) return;
|
||||
inFlight = true;
|
||||
try {
|
||||
const next = await getClusterStatus();
|
||||
if (mounted) {
|
||||
setSnapshot(next);
|
||||
setFetchError(null);
|
||||
}
|
||||
} catch {
|
||||
if (mounted) {
|
||||
setFetchError(
|
||||
'Cluster status could not be refreshed. The values below may be stale.',
|
||||
);
|
||||
}
|
||||
} finally {
|
||||
inFlight = false;
|
||||
}
|
||||
}
|
||||
|
||||
poll();
|
||||
const interval = setInterval(poll, POLL_INTERVAL_MS);
|
||||
// Refresh immediately when the tab becomes visible again (polls are
|
||||
// skipped while hidden).
|
||||
document.addEventListener('visibilitychange', poll);
|
||||
return () => {
|
||||
mounted = false;
|
||||
clearInterval(interval);
|
||||
document.removeEventListener('visibilitychange', poll);
|
||||
};
|
||||
}, [retryToken]);
|
||||
|
||||
let body: React.ReactNode;
|
||||
if (fetchError && !snapshot) {
|
||||
body = (
|
||||
<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. Cluster status needs the Studio API
|
||||
server running.
|
||||
</p>
|
||||
<Button
|
||||
type="button"
|
||||
variant="outline"
|
||||
onClick={() => setRetryToken((token) => token + 1)}
|
||||
>
|
||||
Try Again
|
||||
</Button>
|
||||
</div>
|
||||
);
|
||||
} else if (!snapshot) {
|
||||
body = <p className="py-8 text-center text-muted-foreground">Loading…</p>;
|
||||
} else {
|
||||
body = (
|
||||
<div className="flex flex-col gap-4">
|
||||
<header className="flex flex-wrap items-center gap-3">
|
||||
<h1 className="text-lg font-semibold">Cluster</h1>
|
||||
<Badge variant="outline">
|
||||
{snapshot.mode === 'ray' ? 'ray cluster' : 'local host only'}
|
||||
</Badge>
|
||||
{snapshot.resources && (
|
||||
<span className="text-sm tabular-nums text-muted-foreground">
|
||||
{Math.round(snapshot.resources.gpus_available)} /{' '}
|
||||
{Math.round(snapshot.resources.gpus_total)} GPUs available
|
||||
</span>
|
||||
)}
|
||||
</header>
|
||||
{snapshot.error && (
|
||||
<div
|
||||
role="status"
|
||||
className="flex flex-wrap items-center gap-3 rounded-lg border border-blue-400/40 bg-blue-500/10 px-3 py-2 text-sm"
|
||||
>
|
||||
<Info className="size-4 shrink-0 text-blue-600" aria-hidden />
|
||||
<span className="min-w-0 flex-1">{snapshot.error}</span>
|
||||
</div>
|
||||
)}
|
||||
{fetchError && (
|
||||
<div
|
||||
role="status"
|
||||
aria-live="polite"
|
||||
className="flex flex-wrap items-center gap-3 rounded-lg border border-amber-500/50 bg-amber-500/10 px-3 py-2 text-sm"
|
||||
>
|
||||
<AlertTriangle className="size-4 text-amber-600" aria-hidden />
|
||||
<span className="min-w-0 flex-1">{fetchError}</span>
|
||||
<Button
|
||||
type="button"
|
||||
variant="outline"
|
||||
size="sm"
|
||||
onClick={() => setRetryToken((token) => token + 1)}
|
||||
>
|
||||
Refresh Now
|
||||
</Button>
|
||||
</div>
|
||||
)}
|
||||
<ClusterStrip nodes={snapshot.nodes} />
|
||||
{snapshot.nodes.map((node, i) => (
|
||||
<NodeSection key={`${node.hostname}-${i}`} node={node} />
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="mx-auto flex w-full max-w-[1100px] flex-col gap-6 px-4 pb-12 pt-6">
|
||||
{body}
|
||||
<GpuGrid />
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
'use client';
|
||||
|
||||
import CreateJobButton from '@/components/jobs/CreateJobButton';
|
||||
import EngineConsole from '@/components/jobs/EngineConsole';
|
||||
import { HeaderActions } from '@/components/shell/HeaderActionsContext';
|
||||
import JobQueue from '@/components/jobs/JobQueue';
|
||||
import WarmModelsPanel from '@/components/jobs/WarmModelsPanel';
|
||||
|
||||
export default function InferencePage() {
|
||||
return (
|
||||
@@ -12,8 +10,6 @@ export default function InferencePage() {
|
||||
<HeaderActions>
|
||||
<CreateJobButton jobType="inference" />
|
||||
</HeaderActions>
|
||||
<WarmModelsPanel />
|
||||
<EngineConsole />
|
||||
<JobQueue jobType="inference" />
|
||||
</>
|
||||
);
|
||||
|
||||
@@ -1,93 +0,0 @@
|
||||
'use client';
|
||||
|
||||
import type { ClusterNode } from '@/lib/api';
|
||||
import { cn } from '@/lib/utils';
|
||||
|
||||
export function formatGib(mib: number): string {
|
||||
return `${(mib / 1024).toFixed(1)} GiB`;
|
||||
}
|
||||
|
||||
/** Bar fill class by load: blue when idle, amber under pressure, rose hot. */
|
||||
export function utilizationColor(percent: number): string {
|
||||
if (percent >= 85) return 'bg-rose-500';
|
||||
if (percent >= 50) return 'bg-amber-500';
|
||||
return 'bg-accent-blue';
|
||||
}
|
||||
|
||||
export function clampPercent(percent: number): number {
|
||||
return Math.max(0, Math.min(100, percent));
|
||||
}
|
||||
|
||||
function GpuSegment({
|
||||
index,
|
||||
utilization,
|
||||
memUsedMib,
|
||||
memTotalMib,
|
||||
}: {
|
||||
index: number;
|
||||
utilization: number;
|
||||
memUsedMib: number;
|
||||
memTotalMib: number;
|
||||
}) {
|
||||
const memPercent =
|
||||
memTotalMib > 0 ? clampPercent((memUsedMib / memTotalMib) * 100) : 0;
|
||||
const label =
|
||||
`GPU ${index}: ${utilization}% utilization, ` +
|
||||
`${formatGib(memUsedMib)} / ${formatGib(memTotalMib)} VRAM`;
|
||||
return (
|
||||
<div
|
||||
role="img"
|
||||
aria-label={label}
|
||||
title={label}
|
||||
className="flex w-10 shrink-0 flex-col gap-0.5"
|
||||
>
|
||||
<div className="h-1.5 overflow-hidden rounded-full bg-muted">
|
||||
<div
|
||||
className={cn('h-full rounded-full', utilizationColor(utilization))}
|
||||
style={{ width: `${clampPercent(utilization)}%` }}
|
||||
/>
|
||||
</div>
|
||||
<div className="h-1.5 overflow-hidden rounded-full bg-muted">
|
||||
<div
|
||||
className={cn(
|
||||
'h-full rounded-full',
|
||||
memPercent >= 90 ? 'bg-rose-500' : 'bg-accent-blue',
|
||||
)}
|
||||
style={{ width: `${memPercent}%` }}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
/** One compact line per node: hostname + tiny util/VRAM bars per GPU. */
|
||||
export default function ClusterStrip({ nodes }: { nodes: ClusterNode[] }) {
|
||||
return (
|
||||
<div className="flex flex-col gap-2">
|
||||
{nodes.map((node, i) => (
|
||||
<div
|
||||
key={`${node.hostname}-${i}`}
|
||||
className="flex items-center gap-3"
|
||||
>
|
||||
<span className="w-40 shrink-0 truncate text-xs font-medium">
|
||||
{node.hostname}
|
||||
</span>
|
||||
<div className="flex min-w-0 flex-wrap items-center gap-1.5">
|
||||
{node.gpus.map((gpu) => (
|
||||
<GpuSegment
|
||||
key={gpu.index}
|
||||
index={gpu.index}
|
||||
utilization={gpu.utilization}
|
||||
memUsedMib={gpu.memory_used_mib}
|
||||
memTotalMib={gpu.memory_total_mib}
|
||||
/>
|
||||
))}
|
||||
{node.gpus.length === 0 && (
|
||||
<span className="text-xs text-muted-foreground">no GPUs</span>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
@@ -4,24 +4,14 @@ import userEvent from '@testing-library/user-event';
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
|
||||
import CreateJobModal from './CreateJobModal';
|
||||
import {
|
||||
createJob,
|
||||
getDatasets,
|
||||
getModelPresets,
|
||||
getModels,
|
||||
listGenerators,
|
||||
uploadImage,
|
||||
type GeneratorInfo,
|
||||
} from '@/lib/api';
|
||||
import { createJob, getDatasets, getModels, uploadImage } from '@/lib/api';
|
||||
import { defaultOptionsStore } from '@/stores/defaultOptions';
|
||||
import { DEFAULT_OPTIONS } from '@/lib/defaultOptions';
|
||||
|
||||
vi.mock('@/lib/api', () => ({
|
||||
createJob: vi.fn(),
|
||||
getModels: vi.fn(),
|
||||
getModelPresets: vi.fn(),
|
||||
getDatasets: vi.fn(),
|
||||
listGenerators: vi.fn(),
|
||||
uploadImage: vi.fn(),
|
||||
getSettings: vi.fn(),
|
||||
updateSettings: vi.fn(),
|
||||
@@ -32,32 +22,11 @@ const MODELS = [
|
||||
{ id: 'wan/t2v-14b', label: 'Wan T2V Large', type: 't2v' },
|
||||
];
|
||||
|
||||
// A resident engine slot whose engine config differs from the persisted
|
||||
// defaults on every field the modal adopts.
|
||||
const WARM_SLOT: GeneratorInfo = {
|
||||
state: 'ready',
|
||||
model_id: 'wan/t2v-14b',
|
||||
workload_type: 't2v',
|
||||
num_gpus: 8,
|
||||
dit_cpu_offload: true,
|
||||
text_encoder_cpu_offload: true,
|
||||
vae_cpu_offload: true,
|
||||
image_encoder_cpu_offload: false,
|
||||
use_fsdp_inference: true,
|
||||
enable_torch_compile: true,
|
||||
vsa_sparsity: 0.5,
|
||||
tp_size: 1,
|
||||
sp_size: 8,
|
||||
error: null,
|
||||
};
|
||||
|
||||
beforeEach(() => {
|
||||
// Reset the shared options store to a known baseline for test isolation.
|
||||
defaultOptionsStore.set({ options: DEFAULT_OPTIONS });
|
||||
vi.mocked(getModels).mockResolvedValue(MODELS);
|
||||
vi.mocked(getModelPresets).mockResolvedValue({});
|
||||
vi.mocked(getDatasets).mockResolvedValue([]);
|
||||
vi.mocked(listGenerators).mockResolvedValue([]);
|
||||
vi.mocked(uploadImage).mockResolvedValue({ path: '/uploads/x.png' });
|
||||
vi.mocked(createJob).mockResolvedValue({ id: 'job-1' } as never);
|
||||
});
|
||||
@@ -237,173 +206,4 @@ describe('CreateJobModal', () => {
|
||||
|
||||
await waitFor(() => expect(onSuccess).toHaveBeenCalledTimes(1));
|
||||
});
|
||||
|
||||
it('defaults to the warm resident model and adopts its engine config', async () => {
|
||||
vi.mocked(listGenerators).mockResolvedValue([WARM_SLOT]);
|
||||
|
||||
const user = userEvent.setup();
|
||||
renderModal();
|
||||
|
||||
// Without a warm model the default logic picks the first model
|
||||
// (wan/t2v-1.3b, covered by the seeding test above); the warm slot wins.
|
||||
await waitFor(() =>
|
||||
expect(screen.getByLabelText('Model')).toHaveValue('wan/t2v-14b'),
|
||||
);
|
||||
// The warm selection also triggers its presets fetch.
|
||||
await waitFor(() =>
|
||||
expect(getModelPresets).toHaveBeenCalledWith('wan/t2v-14b'),
|
||||
);
|
||||
|
||||
await user.type(screen.getByLabelText('Prompt'), 'warm run');
|
||||
await user.click(screen.getByRole('button', { name: 'Create Job' }));
|
||||
|
||||
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
|
||||
// The engine fields mirror the resident slot (not the persisted defaults:
|
||||
// num_gpus 1, all offloads/fsdp/compile false) so the job reuses the warm
|
||||
// instance instead of replacing it.
|
||||
expect(vi.mocked(createJob).mock.calls[0][0]).toMatchObject({
|
||||
model_id: 'wan/t2v-14b',
|
||||
num_gpus: 8,
|
||||
tp_size: 1,
|
||||
sp_size: 8,
|
||||
dit_cpu_offload: true,
|
||||
text_encoder_cpu_offload: true,
|
||||
vae_cpu_offload: true,
|
||||
image_encoder_cpu_offload: false,
|
||||
use_fsdp_inference: true,
|
||||
enable_torch_compile: true,
|
||||
vsa_sparsity: 0.5,
|
||||
});
|
||||
});
|
||||
|
||||
it('restores persisted-default engine fields when switching away from the warm model', async () => {
|
||||
defaultOptionsStore.set({
|
||||
options: { ...DEFAULT_OPTIONS, numGpus: 2, tpSize: 2 },
|
||||
});
|
||||
vi.mocked(listGenerators).mockResolvedValue([WARM_SLOT]);
|
||||
|
||||
const user = userEvent.setup();
|
||||
renderModal();
|
||||
|
||||
await waitFor(() =>
|
||||
expect(screen.getByLabelText('Model')).toHaveValue('wan/t2v-14b'),
|
||||
);
|
||||
await user.selectOptions(screen.getByLabelText('Model'), 'wan/t2v-1.3b');
|
||||
await user.type(screen.getByLabelText('Prompt'), 'cold run');
|
||||
await user.click(screen.getByRole('button', { name: 'Create Job' }));
|
||||
|
||||
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
|
||||
expect(vi.mocked(createJob).mock.calls[0][0]).toMatchObject({
|
||||
model_id: 'wan/t2v-1.3b',
|
||||
num_gpus: 2,
|
||||
tp_size: 2,
|
||||
use_fsdp_inference: false,
|
||||
enable_torch_compile: false,
|
||||
});
|
||||
});
|
||||
|
||||
it('applies resolution preset chips and the orientation toggle to the payload', async () => {
|
||||
const user = userEvent.setup();
|
||||
renderModal();
|
||||
|
||||
await screen.findByRole('option', { name: 'Wan T2V (wan/t2v-1.3b)' });
|
||||
await user.click(screen.getByText('Options'));
|
||||
|
||||
// 720p chip sets the /32-rounded dims; the orientation toggle swaps them.
|
||||
await user.click(screen.getByRole('button', { name: '720p' }));
|
||||
await user.click(screen.getByRole('button', { name: 'Landscape' }));
|
||||
expect(
|
||||
screen.getByRole('button', { name: 'Portrait' }),
|
||||
).toBeInTheDocument();
|
||||
|
||||
await user.type(screen.getByLabelText('Prompt'), 'portrait 720p');
|
||||
await user.click(screen.getByRole('button', { name: 'Create Job' }));
|
||||
|
||||
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
|
||||
expect(vi.mocked(createJob).mock.calls[0][0]).toMatchObject({
|
||||
height: 1280,
|
||||
width: 704,
|
||||
});
|
||||
});
|
||||
|
||||
it('restores the model preset resolution via the Native chip', async () => {
|
||||
vi.mocked(getModelPresets).mockResolvedValue({ height: 720, width: 1280 });
|
||||
|
||||
const user = userEvent.setup();
|
||||
renderModal();
|
||||
|
||||
await screen.findByRole('option', { name: 'Wan T2V (wan/t2v-1.3b)' });
|
||||
// Native is disabled until the selected model's presets have loaded.
|
||||
await waitFor(() =>
|
||||
expect(screen.getByRole('button', { name: 'Native' })).toBeEnabled(),
|
||||
);
|
||||
await user.click(screen.getByText('Options'));
|
||||
|
||||
await user.click(screen.getByRole('button', { name: '1080p' }));
|
||||
await user.click(screen.getByRole('button', { name: 'Native' }));
|
||||
|
||||
await user.type(screen.getByLabelText('Prompt'), 'native dims');
|
||||
await user.click(screen.getByRole('button', { name: 'Create Job' }));
|
||||
|
||||
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
|
||||
expect(vi.mocked(createJob).mock.calls[0][0]).toMatchObject({
|
||||
height: 720,
|
||||
width: 1280,
|
||||
});
|
||||
});
|
||||
|
||||
it('populates sampling fields from the selected model presets; engine fields stay from defaults', async () => {
|
||||
defaultOptionsStore.set({
|
||||
options: { ...DEFAULT_OPTIONS, numGpus: 4, tpSize: 2, seed: 999 },
|
||||
});
|
||||
vi.mocked(getModelPresets).mockImplementation(async (id) =>
|
||||
id === 'wan/t2v-14b'
|
||||
? {
|
||||
height: 720,
|
||||
width: 1280,
|
||||
num_frames: 121,
|
||||
fps: 30,
|
||||
num_inference_steps: 40,
|
||||
guidance_scale: 6,
|
||||
guidance_rescale: 0.5,
|
||||
negative_prompt: 'blurry, low quality',
|
||||
seed: 7,
|
||||
}
|
||||
: {},
|
||||
);
|
||||
|
||||
const user = userEvent.setup();
|
||||
renderModal();
|
||||
|
||||
await screen.findByRole('option', { name: 'Wan T2V Large (wan/t2v-14b)' });
|
||||
await user.selectOptions(screen.getByLabelText('Model'), 'wan/t2v-14b');
|
||||
// The negative prompt is the easiest preset-populated field to observe.
|
||||
await waitFor(() =>
|
||||
expect(screen.getByLabelText('Negative Prompt')).toHaveValue(
|
||||
'blurry, low quality',
|
||||
),
|
||||
);
|
||||
|
||||
await user.type(screen.getByLabelText('Prompt'), 'preset test');
|
||||
await user.click(screen.getByRole('button', { name: 'Create Job' }));
|
||||
|
||||
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
|
||||
const payload = vi.mocked(createJob).mock.calls[0][0];
|
||||
expect(payload).toMatchObject({
|
||||
model_id: 'wan/t2v-14b',
|
||||
// Sampling fields come from the model presets…
|
||||
height: 720,
|
||||
width: 1280,
|
||||
num_frames: 121,
|
||||
fps: 30,
|
||||
num_inference_steps: 40,
|
||||
guidance_scale: 6,
|
||||
guidance_rescale: 0.5,
|
||||
negative_prompt: 'blurry, low quality',
|
||||
seed: 7,
|
||||
// …while engine fields still come from the persisted defaults.
|
||||
num_gpus: 4,
|
||||
tp_size: 2,
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
@@ -23,27 +23,15 @@ import { defaultOptionsStore } from '@/stores/defaultOptions';
|
||||
import {
|
||||
createJob,
|
||||
getDatasets,
|
||||
getModelPresets,
|
||||
getModels,
|
||||
listGenerators,
|
||||
uploadImage,
|
||||
type CreateJobRequest,
|
||||
type GeneratorInfo,
|
||||
type Model,
|
||||
} from '@/lib/api';
|
||||
import { getDefaultModelForWorkload } from '@/lib/defaultOptions';
|
||||
import { WORKLOAD_OPTIONS } from '@/lib/jobConfig';
|
||||
import type { JobType } from '@/lib/types';
|
||||
|
||||
// ponytail: dims pre-rounded to multiples of 32 so every model family accepts
|
||||
// them (H3 requires /32) — hence 720p→704 and 1080p→1088. Tooltips show the
|
||||
// exact dims; the labels stay 480p/720p/1080p.
|
||||
const RESOLUTION_PRESETS = [
|
||||
{ label: '480p', height: 480, width: 832 },
|
||||
{ label: '720p', height: 704, width: 1280 },
|
||||
{ label: '1080p', height: 1088, width: 1920 },
|
||||
] as const;
|
||||
|
||||
export interface CreateJobModalProps {
|
||||
isOpen: boolean;
|
||||
onClose: () => void;
|
||||
@@ -67,9 +55,6 @@ export default function CreateJobModal({
|
||||
|
||||
const [models, setModels] = React.useState<Model[]>([]);
|
||||
const [modelId, setModelId] = React.useState('');
|
||||
// The engine's resident slot (null when empty/failed), fetched on open so
|
||||
// jobs against the warm model can adopt its engine config.
|
||||
const [warmSlot, setWarmSlot] = React.useState<GeneratorInfo | null>(null);
|
||||
const [prompt, setPrompt] = React.useState('');
|
||||
const [imagePath, setImagePath] = React.useState('');
|
||||
const [imageFileName, setImageFileName] = React.useState('');
|
||||
@@ -79,12 +64,6 @@ export default function CreateJobModal({
|
||||
const [numFrames, setNumFrames] = React.useState(81);
|
||||
const [height, setHeight] = React.useState(480);
|
||||
const [width, setWidth] = React.useState(832);
|
||||
// The selected model's preset dims, kept so the "Native" chip can restore
|
||||
// them after a resolution chip or manual slider edit.
|
||||
const [nativeDims, setNativeDims] = React.useState<{
|
||||
height: number;
|
||||
width: number;
|
||||
} | null>(null);
|
||||
const [guidanceScale, setGuidanceScale] = React.useState(5);
|
||||
const [guidanceRescale, setGuidanceRescale] = React.useState(0);
|
||||
const [fps, setFps] = React.useState(24);
|
||||
@@ -199,34 +178,17 @@ export default function CreateJobModal({
|
||||
let stale = false;
|
||||
setIsLoadingModels(true);
|
||||
setModelLoadError(null);
|
||||
// For inference jobs, prefer the warm resident model (if any) over the
|
||||
// persisted default so new jobs hit the already-loaded engine. The warm
|
||||
// slot is a nice-to-have: if the lookup fails, fall back silently.
|
||||
const warmSlotPromise = isInference
|
||||
? listGenerators()
|
||||
.then((gens) =>
|
||||
gens[0] && gens[0].state !== 'failed' ? gens[0] : null,
|
||||
)
|
||||
.catch(() => null)
|
||||
: Promise.resolve(null);
|
||||
Promise.all([getModels(inferenceWorkload), warmSlotPromise])
|
||||
.then(([list, slot]) => {
|
||||
getModels(inferenceWorkload)
|
||||
.then((list) => {
|
||||
if (stale) return;
|
||||
setModels(list);
|
||||
setWarmSlot(slot);
|
||||
const ids = list.map((m) => m.id);
|
||||
const opts = defaultOptionsStore.get().options;
|
||||
const defaultId = getDefaultModelForWorkload(
|
||||
opts,
|
||||
inferenceWorkload as 't2v' | 'i2v' | 't2i',
|
||||
);
|
||||
const warmId = slot?.model_id ?? '';
|
||||
const chosen =
|
||||
warmId && ids.includes(warmId)
|
||||
? warmId
|
||||
: ids.includes(defaultId)
|
||||
? defaultId
|
||||
: (list[0]?.id ?? '');
|
||||
const chosen = ids.includes(defaultId) ? defaultId : (list[0]?.id ?? '');
|
||||
setModelId(chosen);
|
||||
if (workloadType === 'dmd_t2v') {
|
||||
setRealScoreModelPath(chosen);
|
||||
@@ -248,81 +210,7 @@ export default function CreateJobModal({
|
||||
return () => {
|
||||
stale = true;
|
||||
};
|
||||
}, [isOpen, isInference, inferenceWorkload, workloadType]);
|
||||
|
||||
// Whenever the selected model changes (including the initial selection on
|
||||
// open), populate the sampling fields from that model's presets. Missing
|
||||
// keys leave the field as-is; engine fields (GPUs, parallelism, offloads)
|
||||
// come from the warm slot or persisted defaults (effect below). Latest
|
||||
// selection wins: the cleanup marks superseded fetches stale, same as the
|
||||
// model-list effect.
|
||||
React.useEffect(() => {
|
||||
if (!isOpen || !isInference || !modelId) return;
|
||||
let stale = false;
|
||||
setNativeDims(null);
|
||||
getModelPresets(modelId)
|
||||
.then((p) => {
|
||||
if (stale) return;
|
||||
if (p.height !== undefined) setHeight(p.height);
|
||||
if (p.width !== undefined) setWidth(p.width);
|
||||
if (p.height !== undefined && p.width !== undefined)
|
||||
setNativeDims({ height: p.height, width: p.width });
|
||||
if (p.num_frames !== undefined)
|
||||
setNumFrames(workloadType === 't2i' ? 1 : p.num_frames);
|
||||
if (p.fps !== undefined) setFps(p.fps);
|
||||
if (p.num_inference_steps !== undefined)
|
||||
setNumInferenceSteps(p.num_inference_steps);
|
||||
if (p.guidance_scale !== undefined) setGuidanceScale(p.guidance_scale);
|
||||
if (p.guidance_rescale !== undefined)
|
||||
setGuidanceRescale(p.guidance_rescale);
|
||||
if (p.negative_prompt !== undefined)
|
||||
setNegativePrompt(p.negative_prompt);
|
||||
if (p.seed !== undefined) setSeed(p.seed);
|
||||
})
|
||||
.catch((e) => {
|
||||
// Presets are a convenience; on failure keep the current values.
|
||||
console.error('Failed to load model presets:', e);
|
||||
});
|
||||
return () => {
|
||||
stale = true;
|
||||
};
|
||||
}, [isOpen, isInference, modelId, workloadType]);
|
||||
|
||||
// When the selected model IS the warm resident one, adopt the slot's engine
|
||||
// config so the job reuses the loaded instance instead of silently replacing
|
||||
// it (e.g. persisted num_gpus=1 vs a warm 8-GPU slot). Switching away from
|
||||
// the warm model restores the persisted-default engine fields; non-warm to
|
||||
// non-warm switches leave the user's engine edits alone (today's behavior).
|
||||
const wasWarmRef = React.useRef(false);
|
||||
React.useEffect(() => {
|
||||
if (!isOpen || !isInference || !modelId) return;
|
||||
const isWarm = warmSlot?.model_id === modelId;
|
||||
if (isWarm && warmSlot) {
|
||||
setNumGpus(warmSlot.num_gpus);
|
||||
setTpSize(warmSlot.tp_size);
|
||||
setSpSize(warmSlot.sp_size);
|
||||
setDitCpuOffload(warmSlot.dit_cpu_offload);
|
||||
setTextEncoderCpuOffload(warmSlot.text_encoder_cpu_offload);
|
||||
setVaeCpuOffload(warmSlot.vae_cpu_offload);
|
||||
setImageEncoderCpuOffload(warmSlot.image_encoder_cpu_offload);
|
||||
setUseFsdpInference(warmSlot.use_fsdp_inference);
|
||||
setEnableTorchCompile(warmSlot.enable_torch_compile);
|
||||
setVsaSparsity(warmSlot.vsa_sparsity);
|
||||
} else if (wasWarmRef.current) {
|
||||
const opts = defaultOptionsStore.get().options;
|
||||
setNumGpus(opts.numGpus);
|
||||
setTpSize(opts.tpSize);
|
||||
setSpSize(opts.spSize);
|
||||
setDitCpuOffload(opts.ditCpuOffload);
|
||||
setTextEncoderCpuOffload(opts.textEncoderCpuOffload);
|
||||
setVaeCpuOffload(opts.vaeCpuOffload);
|
||||
setImageEncoderCpuOffload(opts.imageEncoderCpuOffload);
|
||||
setUseFsdpInference(opts.useFsdpInference);
|
||||
setEnableTorchCompile(opts.enableTorchCompile);
|
||||
setVsaSparsity(opts.vsaSparsity);
|
||||
}
|
||||
wasWarmRef.current = isWarm;
|
||||
}, [isOpen, isInference, modelId, warmSlot]);
|
||||
}, [isOpen, inferenceWorkload, workloadType]);
|
||||
|
||||
// Training jobs need a dataset; load the ready datasets when relevant.
|
||||
React.useEffect(() => {
|
||||
@@ -853,64 +741,6 @@ export default function CreateJobModal({
|
||||
disabled={isSubmitting}
|
||||
/>
|
||||
)}
|
||||
<div className="col-span-full flex flex-wrap items-center gap-1.5">
|
||||
<span className="pl-0.5 text-xs font-normal tracking-wide text-muted-foreground">
|
||||
Resolution
|
||||
</span>
|
||||
{RESOLUTION_PRESETS.map((preset) => (
|
||||
<Button
|
||||
key={preset.label}
|
||||
type="button"
|
||||
size="sm"
|
||||
variant="outline"
|
||||
title={`${preset.height}×${preset.width}`}
|
||||
onClick={() => {
|
||||
// Apply in the current orientation so a portrait
|
||||
// setup stays portrait when switching resolution.
|
||||
const portrait = height > width;
|
||||
setHeight(portrait ? preset.width : preset.height);
|
||||
setWidth(portrait ? preset.height : preset.width);
|
||||
}}
|
||||
disabled={isSubmitting}
|
||||
className="h-6 px-2 text-xs"
|
||||
>
|
||||
{preset.label}
|
||||
</Button>
|
||||
))}
|
||||
<Button
|
||||
type="button"
|
||||
size="sm"
|
||||
variant="outline"
|
||||
title={
|
||||
nativeDims
|
||||
? `${nativeDims.height}×${nativeDims.width}`
|
||||
: 'Model preset resolution (unavailable)'
|
||||
}
|
||||
onClick={() => {
|
||||
if (!nativeDims) return;
|
||||
setHeight(nativeDims.height);
|
||||
setWidth(nativeDims.width);
|
||||
}}
|
||||
disabled={isSubmitting || !nativeDims}
|
||||
className="h-6 px-2 text-xs"
|
||||
>
|
||||
Native
|
||||
</Button>
|
||||
<Button
|
||||
type="button"
|
||||
size="sm"
|
||||
variant="outline"
|
||||
title="Swap height and width"
|
||||
onClick={() => {
|
||||
setHeight(width);
|
||||
setWidth(height);
|
||||
}}
|
||||
disabled={isSubmitting}
|
||||
className="h-6 px-2 text-xs"
|
||||
>
|
||||
{height > width ? 'Portrait' : 'Landscape'}
|
||||
</Button>
|
||||
</div>
|
||||
<SliderRow
|
||||
id="modal-height"
|
||||
label="Height"
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
import { act, fireEvent, render, screen } from '@testing-library/react';
|
||||
import userEvent from '@testing-library/user-event';
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
|
||||
import EngineConsole from './EngineConsole';
|
||||
import { getEngineLogs } from '@/lib/api';
|
||||
|
||||
vi.mock('@/lib/api', () => ({
|
||||
getEngineLogs: vi.fn(),
|
||||
}));
|
||||
|
||||
beforeEach(() => {
|
||||
vi.mocked(getEngineLogs).mockResolvedValue({
|
||||
lines: ['[engine] booted', '[engine] worker heartbeat ok'],
|
||||
total: 2,
|
||||
});
|
||||
});
|
||||
|
||||
describe('EngineConsole', () => {
|
||||
it('is collapsed by default and does not fetch', () => {
|
||||
render(<EngineConsole />);
|
||||
|
||||
expect(
|
||||
screen.getByRole('button', { name: 'Engine output' }),
|
||||
).toHaveAttribute('aria-expanded', 'false');
|
||||
expect(getEngineLogs).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it('shows log lines and polls with the cursor while open', async () => {
|
||||
vi.useFakeTimers();
|
||||
try {
|
||||
render(<EngineConsole />);
|
||||
fireEvent.click(screen.getByRole('button', { name: 'Engine output' }));
|
||||
|
||||
// Flush the immediate poll fired on expand.
|
||||
await act(async () => {
|
||||
await vi.advanceTimersByTimeAsync(0);
|
||||
});
|
||||
expect(getEngineLogs).toHaveBeenCalledWith(0);
|
||||
expect(screen.getByText(/\[engine\] booted/)).toBeInTheDocument();
|
||||
|
||||
// The 2s interval polls again, from the previous total.
|
||||
await act(async () => {
|
||||
await vi.advanceTimersByTimeAsync(2000);
|
||||
});
|
||||
expect(getEngineLogs).toHaveBeenLastCalledWith(2);
|
||||
|
||||
// Collapsing stops the polling.
|
||||
fireEvent.click(screen.getByRole('button', { name: 'Engine output' }));
|
||||
const calls = vi.mocked(getEngineLogs).mock.calls.length;
|
||||
await act(async () => {
|
||||
await vi.advanceTimersByTimeAsync(10000);
|
||||
});
|
||||
expect(getEngineLogs).toHaveBeenCalledTimes(calls);
|
||||
} finally {
|
||||
vi.useRealTimers();
|
||||
}
|
||||
});
|
||||
|
||||
it('clear view empties the scrollback locally', async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<EngineConsole />);
|
||||
|
||||
await user.click(screen.getByRole('button', { name: 'Engine output' }));
|
||||
expect(await screen.findByText(/\[engine\] booted/)).toBeInTheDocument();
|
||||
|
||||
await user.click(screen.getByRole('button', { name: 'Clear view' }));
|
||||
expect(screen.queryByText(/\[engine\] booted/)).not.toBeInTheDocument();
|
||||
expect(screen.getByText('Waiting for engine output…')).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
@@ -1,123 +0,0 @@
|
||||
'use client';
|
||||
|
||||
import * as React from 'react';
|
||||
import { ChevronDown, ChevronRight } from 'lucide-react';
|
||||
|
||||
import { Button } from '@/components/ui/button';
|
||||
import { getEngineLogs } from '@/lib/api';
|
||||
|
||||
const POLL_INTERVAL_MS = 2000;
|
||||
// Cap the DOM at the last ~500 lines; the server keeps its own ring buffer.
|
||||
const MAX_LINES = 500;
|
||||
|
||||
// The engine tail includes uvicorn's access log (until a server restart picks
|
||||
// up access_log=False); the frontend's own polling would otherwise flood the
|
||||
// console with "GET /api/... 200 OK" lines. Keep non-GET and non-2xx lines.
|
||||
const ACCESS_LOG_NOISE = /^INFO:\s+[\d.:]+\s+- "(?:GET|HEAD) \S+ HTTP\/[\d.]+" 2\d\d/;
|
||||
|
||||
/**
|
||||
* Collapsible tail of the engine's stdout/stderr (driver + relayed worker
|
||||
* output). Polls only while open; sticks to the bottom unless the user has
|
||||
* scrolled up.
|
||||
*/
|
||||
export default function EngineConsole() {
|
||||
const [open, setOpen] = React.useState(false);
|
||||
const [lines, setLines] = React.useState<string[]>([]);
|
||||
// Poll cursor + stick-to-bottom flag live in refs: they must update
|
||||
// synchronously from async polls / scroll events, outside React's cycle.
|
||||
const afterRef = React.useRef(0);
|
||||
const stickRef = React.useRef(true);
|
||||
const consoleRef = React.useRef<HTMLPreElement | null>(null);
|
||||
|
||||
React.useEffect(() => {
|
||||
if (!open) return;
|
||||
let mounted = true;
|
||||
let locked = false;
|
||||
|
||||
async function poll() {
|
||||
if (!mounted || locked) return;
|
||||
locked = true;
|
||||
try {
|
||||
const data = await getEngineLogs(afterRef.current);
|
||||
afterRef.current = data.total;
|
||||
const fresh = data.lines.filter((l) => !ACCESS_LOG_NOISE.test(l));
|
||||
if (mounted && fresh.length > 0) {
|
||||
setLines((prev) => [...prev, ...fresh].slice(-MAX_LINES));
|
||||
}
|
||||
} catch (e) {
|
||||
console.error('Failed to fetch engine logs:', e);
|
||||
} finally {
|
||||
locked = false;
|
||||
}
|
||||
}
|
||||
|
||||
poll();
|
||||
const interval = setInterval(poll, POLL_INTERVAL_MS);
|
||||
return () => {
|
||||
mounted = false;
|
||||
clearInterval(interval);
|
||||
};
|
||||
}, [open]);
|
||||
|
||||
// Follow the tail after new lines land, unless the user scrolled up.
|
||||
React.useEffect(() => {
|
||||
const el = consoleRef.current;
|
||||
if (el && stickRef.current) el.scrollTop = el.scrollHeight;
|
||||
}, [lines]);
|
||||
|
||||
function handleScroll() {
|
||||
const el = consoleRef.current;
|
||||
if (!el) return;
|
||||
stickRef.current = el.scrollHeight - el.scrollTop - el.clientHeight < 40;
|
||||
}
|
||||
|
||||
return (
|
||||
<section
|
||||
aria-label="Engine output"
|
||||
className="mx-auto w-full max-w-[850px] px-10 pt-3"
|
||||
>
|
||||
<div className="rounded-lg border border-border bg-background">
|
||||
<div className="flex items-center gap-2 px-2 py-1.5">
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => setOpen((o) => !o)}
|
||||
aria-expanded={open}
|
||||
className="flex flex-1 items-center gap-2 rounded-md px-2 py-1 text-sm font-semibold text-foreground hover:bg-accent"
|
||||
>
|
||||
{open ? (
|
||||
<ChevronDown className="h-4 w-4" />
|
||||
) : (
|
||||
<ChevronRight className="h-4 w-4" />
|
||||
)}
|
||||
Engine output
|
||||
</button>
|
||||
{open && (
|
||||
<Button
|
||||
size="sm"
|
||||
variant="ghost"
|
||||
// Resets the local view only; the server buffer is untouched.
|
||||
onClick={() => setLines([])}
|
||||
>
|
||||
Clear view
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
{open && (
|
||||
<pre
|
||||
ref={consoleRef}
|
||||
onScroll={handleScroll}
|
||||
className="m-0 h-64 overflow-auto whitespace-pre-wrap break-words rounded-b-lg border-t border-border bg-zinc-950 p-3 font-mono text-xs leading-normal text-zinc-200"
|
||||
>
|
||||
{lines.length === 0 ? (
|
||||
<span className="italic text-zinc-500">
|
||||
Waiting for engine output…
|
||||
</span>
|
||||
) : (
|
||||
lines.join('\n')
|
||||
)}
|
||||
</pre>
|
||||
)}
|
||||
</div>
|
||||
</section>
|
||||
);
|
||||
}
|
||||
@@ -142,7 +142,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
return (
|
||||
<article
|
||||
className={cn(
|
||||
'mb-1.5 flex cursor-pointer flex-col gap-1 rounded-lg border bg-background px-3 py-1.5 transition-colors last:mb-0',
|
||||
'mb-3 flex cursor-pointer flex-col gap-2.5 rounded-lg border bg-background p-4 transition-colors last:mb-0',
|
||||
isSelected
|
||||
? 'border-accent-blue bg-accent-blue/5'
|
||||
: 'border-border hover:border-muted-foreground/40',
|
||||
@@ -152,16 +152,20 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
type="button"
|
||||
aria-pressed={isSelected}
|
||||
onClick={handleSelectJob}
|
||||
className="flex w-full min-w-0 flex-col gap-1 rounded-md text-left"
|
||||
className="flex w-full flex-col gap-2.5 rounded-md text-left"
|
||||
>
|
||||
<span className="flex w-full min-w-0 items-center gap-2">
|
||||
<span className="shrink-0 text-sm font-semibold text-foreground">
|
||||
{job.model_id}
|
||||
<span className="flex flex-wrap items-center justify-between gap-2">
|
||||
<span className="text-[0.95rem] font-semibold text-foreground">
|
||||
{job.model_id}
|
||||
</span>
|
||||
<Badge variant={BADGE_VARIANTS[job.status] ?? 'secondary'}>
|
||||
{job.status}
|
||||
</Badge>
|
||||
</span>
|
||||
<Badge variant={BADGE_VARIANTS[job.status] ?? 'secondary'}>
|
||||
{job.status}
|
||||
</Badge>
|
||||
<span className="ml-auto flex shrink-0 items-center gap-3 text-xs text-muted-foreground">
|
||||
<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>
|
||||
@@ -179,10 +183,6 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
</span>
|
||||
)}
|
||||
</span>
|
||||
</span>
|
||||
<span className="w-full whitespace-pre-wrap break-words text-xs text-muted-foreground">
|
||||
{job.prompt}
|
||||
</span>
|
||||
</button>
|
||||
<div className="flex flex-wrap items-center gap-1.5">
|
||||
{job.status === 'running' ? (
|
||||
@@ -190,7 +190,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
size="sm"
|
||||
onClick={handleStop}
|
||||
disabled={isLoading}
|
||||
className="h-6 px-2 text-xs border-transparent bg-amber-500 text-black shadow-md hover:bg-amber-400"
|
||||
className="border-transparent bg-amber-500 text-black shadow-md hover:bg-amber-400"
|
||||
>
|
||||
Stop
|
||||
</Button>
|
||||
@@ -199,7 +199,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
size="sm"
|
||||
onClick={handleStart}
|
||||
disabled={isLoading}
|
||||
className="h-6 px-2 text-xs border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
|
||||
className="border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
|
||||
>
|
||||
Restart
|
||||
</Button>
|
||||
@@ -208,7 +208,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
size="sm"
|
||||
onClick={handleStart}
|
||||
disabled={isLoading}
|
||||
className="h-6 px-2 text-xs border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
|
||||
className="border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
|
||||
>
|
||||
Start
|
||||
</Button>
|
||||
@@ -222,7 +222,6 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
onClick={handleDownloadVideo}
|
||||
disabled={isLoading}
|
||||
title="Download video"
|
||||
className="h-6 px-2 text-xs"
|
||||
>
|
||||
Download Video
|
||||
</Button>
|
||||
@@ -232,7 +231,6 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
|
||||
variant="destructive"
|
||||
onClick={handleDelete}
|
||||
disabled={isLoading}
|
||||
className="h-6 px-2 text-xs"
|
||||
>
|
||||
Delete
|
||||
</Button>
|
||||
|
||||
@@ -9,7 +9,6 @@ import { makeJob as makeBaseJob } from '@/test/factories';
|
||||
vi.mock('@/lib/api', () => ({
|
||||
getJobLogs: vi.fn(),
|
||||
downloadJobLog: vi.fn(),
|
||||
getJobVideoUrl: (id: string) => `http://test.local/api/jobs/${id}/video`,
|
||||
}));
|
||||
|
||||
const makeJob = (overrides: Partial<Job> = {}): Job =>
|
||||
@@ -46,43 +45,6 @@ describe('JobDetailsSidebar', () => {
|
||||
expect(onWidthChange).toHaveBeenCalledWith(0);
|
||||
});
|
||||
|
||||
it('plays completed inference output inline; running jobs get no player', async () => {
|
||||
vi.mocked(getJobLogs).mockResolvedValue({
|
||||
lines: [],
|
||||
total: 0,
|
||||
progress: 0,
|
||||
progress_msg: '',
|
||||
phase: '',
|
||||
});
|
||||
|
||||
const { rerender } = render(
|
||||
<JobDetailsSidebar
|
||||
job={makeJob({
|
||||
status: 'completed',
|
||||
output_path: '/outputs/job-1.mp4',
|
||||
prompt: 'a cat surfing a wave',
|
||||
})}
|
||||
onClose={vi.fn()}
|
||||
/>,
|
||||
);
|
||||
|
||||
const video = screen.getByLabelText('Generated video: a cat surfing a wave');
|
||||
expect(video.tagName).toBe('VIDEO');
|
||||
expect(video).toHaveAttribute('controls');
|
||||
expect(video).toHaveAttribute(
|
||||
'src',
|
||||
'http://test.local/api/jobs/job-1/video',
|
||||
);
|
||||
|
||||
rerender(
|
||||
<JobDetailsSidebar
|
||||
job={makeJob({ status: 'running', output_path: null })}
|
||||
onClose={vi.fn()}
|
||||
/>,
|
||||
);
|
||||
expect(screen.queryByLabelText(/Generated video/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('renders log lines streamed from the job log poll', async () => {
|
||||
vi.mocked(getJobLogs).mockResolvedValue({
|
||||
lines: ['boot sequence started', 'loading model weights'],
|
||||
|
||||
@@ -6,7 +6,7 @@ import { X } from 'lucide-react';
|
||||
import { Button } from '@/components/ui/button';
|
||||
import { useDrawerFocus } from '@/hooks/useDrawerFocus';
|
||||
import { useResizable } from '@/hooks/useResizable';
|
||||
import { downloadJobLog, getJobLogs, getJobVideoUrl } from '@/lib/api';
|
||||
import { downloadJobLog, getJobLogs } from '@/lib/api';
|
||||
import type { Job } from '@/lib/types';
|
||||
import { cn, downloadBlob } from '@/lib/utils';
|
||||
|
||||
@@ -180,41 +180,6 @@ export default function JobDetailsSidebar({
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{job.status === 'completed' &&
|
||||
job.output_path &&
|
||||
(job.job_type === 'inference' || !job.job_type) && (
|
||||
<div className="border-b border-border px-5 py-4">
|
||||
<span className="mb-2 block text-xs font-semibold uppercase tracking-wider text-muted-foreground">
|
||||
Output
|
||||
</span>
|
||||
{job.output_path.toLowerCase().endsWith('.png') ? (
|
||||
// eslint-disable-next-line @next/next/no-img-element
|
||||
<img
|
||||
src={getJobVideoUrl(job.id)}
|
||||
alt={
|
||||
job.prompt
|
||||
? `Generated image: ${job.prompt}`
|
||||
: 'Generated image'
|
||||
}
|
||||
className="block w-full rounded-lg border border-border bg-background object-contain"
|
||||
/>
|
||||
) : (
|
||||
<video
|
||||
src={getJobVideoUrl(job.id)}
|
||||
aria-label={
|
||||
job.prompt
|
||||
? `Generated video: ${job.prompt}`
|
||||
: 'Generated video'
|
||||
}
|
||||
controls
|
||||
playsInline
|
||||
preload="metadata"
|
||||
className="block max-h-80 w-full rounded-lg border border-border bg-background object-contain"
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div className="flex min-h-0 flex-1 flex-col px-5 py-4">
|
||||
<div className="mb-2 flex items-center justify-between">
|
||||
<span className="text-xs font-semibold uppercase tracking-wider text-muted-foreground">
|
||||
|
||||
@@ -1,188 +0,0 @@
|
||||
import { render, screen, waitFor } from '@testing-library/react';
|
||||
import userEvent from '@testing-library/user-event';
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
|
||||
import WarmModelsPanel from './WarmModelsPanel';
|
||||
import {
|
||||
getModels,
|
||||
listGenerators,
|
||||
preloadGenerator,
|
||||
unloadGenerator,
|
||||
type GeneratorInfo,
|
||||
} from '@/lib/api';
|
||||
import { DEFAULT_OPTIONS } from '@/lib/defaultOptions';
|
||||
import { defaultOptionsStore } from '@/stores/defaultOptions';
|
||||
import { toast } from 'sonner';
|
||||
|
||||
vi.mock('@/lib/api', () => ({
|
||||
getModels: vi.fn(),
|
||||
listGenerators: vi.fn(),
|
||||
preloadGenerator: vi.fn(),
|
||||
unloadGenerator: vi.fn(),
|
||||
getSettings: vi.fn(),
|
||||
updateSettings: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock('sonner', () => ({
|
||||
toast: { error: vi.fn() },
|
||||
}));
|
||||
|
||||
const makeGenerator = (
|
||||
overrides: Partial<GeneratorInfo> = {},
|
||||
): GeneratorInfo => ({
|
||||
state: 'ready',
|
||||
model_id: 'Wan-AI/Wan2.1-T2V-1.3B-Diffusers',
|
||||
workload_type: 't2v',
|
||||
num_gpus: 1,
|
||||
dit_cpu_offload: false,
|
||||
text_encoder_cpu_offload: false,
|
||||
vae_cpu_offload: false,
|
||||
image_encoder_cpu_offload: false,
|
||||
use_fsdp_inference: false,
|
||||
enable_torch_compile: false,
|
||||
vsa_sparsity: 0,
|
||||
tp_size: -1,
|
||||
sp_size: -1,
|
||||
error: null,
|
||||
...overrides,
|
||||
});
|
||||
|
||||
beforeEach(() => {
|
||||
// Reset the shared options store to a known baseline for test isolation.
|
||||
defaultOptionsStore.set({ options: DEFAULT_OPTIONS });
|
||||
vi.mocked(getModels).mockResolvedValue([
|
||||
{ id: 'Wan-AI/Wan2.1-T2V-1.3B-Diffusers', label: 'Wan2.1 T2V 1.3B' },
|
||||
]);
|
||||
vi.mocked(listGenerators).mockResolvedValue([]);
|
||||
vi.mocked(preloadGenerator).mockResolvedValue(
|
||||
makeGenerator({ state: 'loading' }),
|
||||
);
|
||||
vi.mocked(unloadGenerator).mockResolvedValue(undefined);
|
||||
});
|
||||
|
||||
describe('WarmModelsPanel', () => {
|
||||
it('shows the empty slot when no model is loaded', async () => {
|
||||
render(<WarmModelsPanel />);
|
||||
|
||||
expect(await screen.findByText('No model loaded')).toBeInTheDocument();
|
||||
expect(
|
||||
screen.queryByRole('button', { name: 'Unload' }),
|
||||
).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('renders the resident slot with its state and config summary', async () => {
|
||||
vi.mocked(listGenerators).mockResolvedValue([
|
||||
makeGenerator({ num_gpus: 8, enable_torch_compile: true }),
|
||||
]);
|
||||
render(<WarmModelsPanel />);
|
||||
|
||||
// Model label: last path segment with dashes/underscores as spaces.
|
||||
expect(
|
||||
await screen.findByText('Wan2.1 T2V 1.3B Diffusers'),
|
||||
).toBeInTheDocument();
|
||||
expect(screen.getByText('ready')).toBeInTheDocument();
|
||||
expect(screen.getByText('8 GPU · compile')).toBeInTheDocument();
|
||||
expect(screen.getByRole('button', { name: 'Unload' })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('disables loading a new model while a load is in flight', async () => {
|
||||
vi.mocked(listGenerators).mockResolvedValue([
|
||||
makeGenerator({ state: 'loading' }),
|
||||
]);
|
||||
render(<WarmModelsPanel />);
|
||||
|
||||
expect(await screen.findByText('loading')).toBeInTheDocument();
|
||||
expect(screen.getByRole('button', { name: 'Load model' })).toBeDisabled();
|
||||
expect(
|
||||
screen.queryByRole('button', { name: 'Unload' }),
|
||||
).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('shows the error on a failed slot and keeps retry enabled', async () => {
|
||||
vi.mocked(listGenerators).mockResolvedValue([
|
||||
makeGenerator({ state: 'failed', error: 'CUDA out of memory' }),
|
||||
]);
|
||||
render(<WarmModelsPanel />);
|
||||
|
||||
expect(await screen.findByText('failed')).toHaveAttribute(
|
||||
'title',
|
||||
'CUDA out of memory',
|
||||
);
|
||||
await waitFor(() =>
|
||||
expect(screen.getByRole('button', { name: 'Load model' })).toBeEnabled(),
|
||||
);
|
||||
});
|
||||
|
||||
it('labels the load button as a swap when a different model is resident', async () => {
|
||||
vi.mocked(listGenerators).mockResolvedValue([
|
||||
makeGenerator({ model_id: 'FastVideo/FastHunyuan-diffusers' }),
|
||||
]);
|
||||
render(<WarmModelsPanel />);
|
||||
|
||||
expect(
|
||||
await screen.findByRole('button', { name: 'Load (replaces current)' }),
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('loads the selected model with the persisted default options', async () => {
|
||||
defaultOptionsStore.set({
|
||||
options: { ...DEFAULT_OPTIONS, numGpus: 4, enableTorchCompile: true },
|
||||
});
|
||||
const user = userEvent.setup();
|
||||
render(<WarmModelsPanel />);
|
||||
|
||||
const button = await screen.findByRole('button', { name: 'Load model' });
|
||||
await waitFor(() => expect(button).toBeEnabled());
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() =>
|
||||
expect(preloadGenerator).toHaveBeenCalledWith({
|
||||
model_id: 'Wan-AI/Wan2.1-T2V-1.3B-Diffusers',
|
||||
workload_type: 't2v',
|
||||
num_gpus: 4,
|
||||
dit_cpu_offload: false,
|
||||
text_encoder_cpu_offload: false,
|
||||
vae_cpu_offload: false,
|
||||
image_encoder_cpu_offload: false,
|
||||
use_fsdp_inference: false,
|
||||
enable_torch_compile: true,
|
||||
vsa_sparsity: 0,
|
||||
tp_size: -1,
|
||||
sp_size: -1,
|
||||
}),
|
||||
);
|
||||
// The panel refetches so the new "loading" slot appears promptly.
|
||||
await waitFor(() => expect(listGenerators).toHaveBeenCalledTimes(2));
|
||||
});
|
||||
|
||||
it('surfaces the backend detail when a load is rejected', async () => {
|
||||
vi.mocked(preloadGenerator).mockRejectedValue(
|
||||
new Error('a model load is already in progress'),
|
||||
);
|
||||
const user = userEvent.setup();
|
||||
render(<WarmModelsPanel />);
|
||||
|
||||
const button = await screen.findByRole('button', { name: 'Load model' });
|
||||
await waitFor(() => expect(button).toBeEnabled());
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() =>
|
||||
expect(toast.error).toHaveBeenCalledWith('Model was not loaded', {
|
||||
description: 'a model load is already in progress',
|
||||
}),
|
||||
);
|
||||
});
|
||||
|
||||
it('unloads the resident model with no payload', async () => {
|
||||
vi.mocked(listGenerators).mockResolvedValue([makeGenerator()]);
|
||||
const user = userEvent.setup();
|
||||
render(<WarmModelsPanel />);
|
||||
|
||||
await user.click(await screen.findByRole('button', { name: 'Unload' }));
|
||||
|
||||
await waitFor(() => expect(unloadGenerator).toHaveBeenCalledTimes(1));
|
||||
expect(vi.mocked(unloadGenerator).mock.calls[0]).toEqual([]);
|
||||
// The slot refreshes after the unload succeeds.
|
||||
await waitFor(() => expect(listGenerators).toHaveBeenCalledTimes(2));
|
||||
});
|
||||
});
|
||||
@@ -1,237 +0,0 @@
|
||||
'use client';
|
||||
|
||||
import * as React from 'react';
|
||||
import { toast } from 'sonner';
|
||||
|
||||
import { Badge, type BadgeProps } from '@/components/ui/badge';
|
||||
import { Button } from '@/components/ui/button';
|
||||
import { NativeSelect } from '@/components/ui/native-select';
|
||||
import { useStore } from '@/hooks/useStore';
|
||||
import {
|
||||
getModels,
|
||||
listGenerators,
|
||||
preloadGenerator,
|
||||
unloadGenerator,
|
||||
type GeneratorInfo,
|
||||
type Model,
|
||||
} from '@/lib/api';
|
||||
import { getDefaultModelForWorkload } from '@/lib/defaultOptions';
|
||||
import { defaultOptionsStore } from '@/stores/defaultOptions';
|
||||
|
||||
// Mirrors the backend's model_label(): readable label from an HF-style path.
|
||||
function modelLabel(modelId: string): string {
|
||||
return (modelId.split('/').pop() ?? modelId).replace(/[-_]/g, ' ');
|
||||
}
|
||||
|
||||
// Compact summary of only the non-default engine bits, e.g. "8 GPU · compile".
|
||||
function configSummary(gen: GeneratorInfo): string {
|
||||
const parts: string[] = [];
|
||||
if (gen.num_gpus !== 1) parts.push(`${gen.num_gpus} GPU`);
|
||||
if (gen.sp_size !== -1) parts.push(`SP ${gen.sp_size}`);
|
||||
if (gen.tp_size !== -1) parts.push(`TP ${gen.tp_size}`);
|
||||
if (gen.dit_cpu_offload) parts.push('DiT offload');
|
||||
if (gen.text_encoder_cpu_offload) parts.push('TE offload');
|
||||
if (gen.vae_cpu_offload) parts.push('VAE offload');
|
||||
if (gen.image_encoder_cpu_offload) parts.push('image enc offload');
|
||||
if (gen.use_fsdp_inference) parts.push('FSDP');
|
||||
if (gen.enable_torch_compile) parts.push('compile');
|
||||
if (gen.vsa_sparsity > 0) parts.push(`VSA ${gen.vsa_sparsity.toFixed(2)}`);
|
||||
return parts.join(' · ');
|
||||
}
|
||||
|
||||
const STATE_VARIANTS: Record<GeneratorInfo['state'], BadgeProps['variant']> = {
|
||||
ready: 'success',
|
||||
loading: 'warning',
|
||||
failed: 'destructive',
|
||||
};
|
||||
|
||||
/**
|
||||
* Utility strip for the engine's single model slot: shows the resident model
|
||||
* (ready/loading/failed), loads the selected model using the persisted
|
||||
* default job options (replacing whatever is resident), and unloads it.
|
||||
*/
|
||||
export default function WarmModelsPanel() {
|
||||
const { options } = useStore(defaultOptionsStore);
|
||||
|
||||
const [slot, setSlot] = React.useState<GeneratorInfo | null>(null);
|
||||
const [models, setModels] = React.useState<Model[]>([]);
|
||||
const [modelId, setModelId] = React.useState('');
|
||||
const [isBusy, setIsBusy] = React.useState(false);
|
||||
|
||||
const fetchSlot = React.useCallback(async () => {
|
||||
try {
|
||||
const list = await listGenerators();
|
||||
setSlot(list[0] ?? null);
|
||||
} catch (e) {
|
||||
console.error('Failed to fetch generators:', e);
|
||||
}
|
||||
}, []);
|
||||
|
||||
React.useEffect(() => {
|
||||
fetchSlot();
|
||||
}, [fetchSlot]);
|
||||
|
||||
// Poll every 5s while a load is in flight; stop otherwise.
|
||||
const isLoading = slot?.state === 'loading';
|
||||
React.useEffect(() => {
|
||||
if (!isLoading) return;
|
||||
const interval = setInterval(fetchSlot, 5000);
|
||||
return () => clearInterval(interval);
|
||||
}, [isLoading, fetchSlot]);
|
||||
|
||||
// Same model catalogue (and default selection) as the create-job modal.
|
||||
React.useEffect(() => {
|
||||
getModels('t2v')
|
||||
.then((list) => {
|
||||
setModels(list);
|
||||
const defaultId = getDefaultModelForWorkload(
|
||||
defaultOptionsStore.get().options,
|
||||
't2v',
|
||||
);
|
||||
setModelId(
|
||||
list.some((m) => m.id === defaultId)
|
||||
? defaultId
|
||||
: (list[0]?.id ?? ''),
|
||||
);
|
||||
})
|
||||
.catch((e) => console.error('Failed to load models:', e));
|
||||
}, []);
|
||||
|
||||
async function handleLoad() {
|
||||
if (!modelId || isBusy || isLoading) return;
|
||||
setIsBusy(true);
|
||||
try {
|
||||
await preloadGenerator({
|
||||
model_id: modelId,
|
||||
workload_type: 't2v',
|
||||
num_gpus: options.numGpus,
|
||||
dit_cpu_offload: options.ditCpuOffload,
|
||||
text_encoder_cpu_offload: options.textEncoderCpuOffload,
|
||||
vae_cpu_offload: options.vaeCpuOffload,
|
||||
image_encoder_cpu_offload: options.imageEncoderCpuOffload,
|
||||
use_fsdp_inference: options.useFsdpInference,
|
||||
enable_torch_compile: options.enableTorchCompile,
|
||||
vsa_sparsity: options.vsaSparsity,
|
||||
tp_size: options.tpSize,
|
||||
sp_size: options.spSize,
|
||||
});
|
||||
await fetchSlot();
|
||||
} catch (err) {
|
||||
console.error('Failed to load model:', err);
|
||||
toast.error('Model was not loaded', {
|
||||
description:
|
||||
err instanceof Error
|
||||
? err.message
|
||||
: 'Check the Studio API, then retry.',
|
||||
});
|
||||
} finally {
|
||||
setIsBusy(false);
|
||||
}
|
||||
}
|
||||
|
||||
async function handleUnload() {
|
||||
if (isBusy) return;
|
||||
setIsBusy(true);
|
||||
try {
|
||||
await unloadGenerator();
|
||||
await fetchSlot();
|
||||
} catch (err) {
|
||||
console.error('Failed to unload model:', err);
|
||||
toast.error('Model was not unloaded', {
|
||||
description:
|
||||
err instanceof Error
|
||||
? err.message
|
||||
: 'Check the Studio API, then retry.',
|
||||
});
|
||||
} finally {
|
||||
setIsBusy(false);
|
||||
}
|
||||
}
|
||||
|
||||
// Loading a model always replaces the resident one — say so on the button.
|
||||
const replaces = slot !== null && !!modelId && slot.model_id !== modelId;
|
||||
|
||||
return (
|
||||
<section
|
||||
aria-label="Warm models"
|
||||
className="mx-auto w-full max-w-[850px] px-10 pt-6"
|
||||
>
|
||||
<div className="flex flex-col gap-3 rounded-lg border border-border bg-background p-4">
|
||||
<div className="flex flex-wrap items-center gap-2">
|
||||
<h2 className="mr-auto text-sm font-semibold text-foreground">
|
||||
Warm Model
|
||||
</h2>
|
||||
<label htmlFor="warm-model-select" className="sr-only">
|
||||
Model to load
|
||||
</label>
|
||||
<NativeSelect
|
||||
id="warm-model-select"
|
||||
value={modelId}
|
||||
onChange={(e) => setModelId(e.target.value)}
|
||||
disabled={isBusy || models.length === 0}
|
||||
className="h-9 w-auto max-w-64 rounded-lg"
|
||||
>
|
||||
<option value="" disabled>
|
||||
{models.length === 0 ? 'Loading models…' : 'Select a model…'}
|
||||
</option>
|
||||
{models.map((model) => (
|
||||
<option key={model.id} value={model.id}>
|
||||
{model.label}
|
||||
</option>
|
||||
))}
|
||||
</NativeSelect>
|
||||
<Button
|
||||
size="sm"
|
||||
onClick={handleLoad}
|
||||
disabled={isBusy || !modelId || isLoading}
|
||||
>
|
||||
{replaces ? 'Load (replaces current)' : 'Load model'}
|
||||
</Button>
|
||||
</div>
|
||||
<p className="text-xs text-muted-foreground">
|
||||
One model at a time stays resident in GPU memory so jobs skip the
|
||||
load wait; loading a new one replaces it (uses your default job
|
||||
options).
|
||||
</p>
|
||||
<div className="flex min-h-8 flex-wrap items-center gap-2">
|
||||
{slot ? (
|
||||
<>
|
||||
<Badge
|
||||
variant={STATE_VARIANTS[slot.state]}
|
||||
className={
|
||||
slot.state === 'loading' ? 'animate-pulse' : undefined
|
||||
}
|
||||
title={
|
||||
slot.state === 'failed' ? (slot.error ?? undefined) : undefined
|
||||
}
|
||||
>
|
||||
{slot.state}
|
||||
</Badge>
|
||||
<span className="text-sm font-medium text-foreground">
|
||||
{modelLabel(slot.model_id)}
|
||||
</span>
|
||||
<span className="text-xs text-muted-foreground">
|
||||
{configSummary(slot)}
|
||||
</span>
|
||||
{slot.state === 'ready' && (
|
||||
<Button
|
||||
size="sm"
|
||||
variant="outline"
|
||||
className="ml-auto"
|
||||
onClick={handleUnload}
|
||||
disabled={isBusy}
|
||||
>
|
||||
Unload
|
||||
</Button>
|
||||
)}
|
||||
</>
|
||||
) : (
|
||||
<span className="text-sm text-muted-foreground">
|
||||
No model loaded
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
);
|
||||
}
|
||||
@@ -8,9 +8,9 @@ import { Button } from '@/components/ui/button';
|
||||
import { ThemeToggle } from '@/components/ui/theme-toggle';
|
||||
|
||||
const TAB_TITLES: Record<string, string> = {
|
||||
'/inference': 'Studio',
|
||||
'/finetuning': 'Studio',
|
||||
'/distillation': 'Studio',
|
||||
'/inference': 'Jobs',
|
||||
'/finetuning': 'Jobs',
|
||||
'/distillation': 'Jobs',
|
||||
'/datasets': 'Datasets',
|
||||
'/gallery': 'Gallery',
|
||||
'/gpus': 'GPUs',
|
||||
|
||||
@@ -108,7 +108,7 @@ export default function PrimarySidebar({
|
||||
isJobsActive && TAB_ACTIVE,
|
||||
)}
|
||||
>
|
||||
<span>Studio</span>
|
||||
<span>Jobs</span>
|
||||
<svg
|
||||
viewBox="0 0 24 24"
|
||||
fill="none"
|
||||
|
||||
@@ -202,33 +202,6 @@ export async function getModels(workloadType?: string): Promise<Model[]> {
|
||||
return response.json();
|
||||
}
|
||||
|
||||
/**
|
||||
* A model's recommended sampling settings. Keys the backend has no value
|
||||
* for are absent; the UI leaves those form fields untouched.
|
||||
*/
|
||||
export interface ModelPresets {
|
||||
height?: number;
|
||||
width?: number;
|
||||
num_frames?: number;
|
||||
fps?: number;
|
||||
num_inference_steps?: number;
|
||||
guidance_scale?: number;
|
||||
guidance_rescale?: number;
|
||||
negative_prompt?: string;
|
||||
seed?: number;
|
||||
}
|
||||
|
||||
export async function getModelPresets(modelId: string): Promise<ModelPresets> {
|
||||
const baseApiUrl = getApiBaseUrl();
|
||||
const response = await fetch(
|
||||
`${baseApiUrl}/models/presets?model_id=${encodeURIComponent(modelId)}`,
|
||||
);
|
||||
if (!response.ok) {
|
||||
throw new Error("Failed to fetch model presets");
|
||||
}
|
||||
return response.json();
|
||||
}
|
||||
|
||||
export async function getGpus(): Promise<GpuSnapshot> {
|
||||
const baseApiUrl = getApiBaseUrl();
|
||||
const response = await fetch(`${baseApiUrl}/gpus`);
|
||||
@@ -346,109 +319,6 @@ export async function downloadJobVideo(id: string): Promise<Blob> {
|
||||
return response.blob();
|
||||
}
|
||||
|
||||
// --- Generators (warm models) ---
|
||||
|
||||
/**
|
||||
* Engine subset of CreateJobRequest that identifies one resident generator.
|
||||
* Mirrors the backend's GeneratorRequest defaults.
|
||||
*/
|
||||
export interface GeneratorRequest {
|
||||
model_id: string;
|
||||
workload_type?: string;
|
||||
num_gpus?: number;
|
||||
dit_cpu_offload?: boolean;
|
||||
text_encoder_cpu_offload?: boolean;
|
||||
vae_cpu_offload?: boolean;
|
||||
image_encoder_cpu_offload?: boolean;
|
||||
use_fsdp_inference?: boolean;
|
||||
enable_torch_compile?: boolean;
|
||||
vsa_sparsity?: number;
|
||||
tp_size?: number;
|
||||
sp_size?: number;
|
||||
}
|
||||
|
||||
export interface GeneratorInfo {
|
||||
state: "ready" | "loading" | "failed";
|
||||
model_id: string;
|
||||
workload_type: string;
|
||||
num_gpus: number;
|
||||
dit_cpu_offload: boolean;
|
||||
text_encoder_cpu_offload: boolean;
|
||||
vae_cpu_offload: boolean;
|
||||
image_encoder_cpu_offload: boolean;
|
||||
use_fsdp_inference: boolean;
|
||||
enable_torch_compile: boolean;
|
||||
vsa_sparsity: number;
|
||||
tp_size: number;
|
||||
sp_size: number;
|
||||
error: string | null;
|
||||
started_at?: number;
|
||||
}
|
||||
|
||||
/** The single resident slot: [] when nothing is loaded, else one entry. */
|
||||
export async function listGenerators(): Promise<GeneratorInfo[]> {
|
||||
const baseApiUrl = getApiBaseUrl();
|
||||
const response = await fetch(`${baseApiUrl}/generators`);
|
||||
if (!response.ok) {
|
||||
throw new Error("Failed to fetch generators");
|
||||
}
|
||||
return response.json();
|
||||
}
|
||||
|
||||
export async function preloadGenerator(
|
||||
req: GeneratorRequest,
|
||||
): Promise<GeneratorInfo> {
|
||||
const baseApiUrl = getApiBaseUrl();
|
||||
const response = await fetch(`${baseApiUrl}/generators/preload`, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(req),
|
||||
});
|
||||
if (!response.ok) {
|
||||
const error = await response
|
||||
.json()
|
||||
.catch(() => ({ detail: "Failed to preload model" }));
|
||||
throw new Error(error.detail || "Failed to preload model");
|
||||
}
|
||||
return response.json();
|
||||
}
|
||||
|
||||
/** Unload the single resident generator (no body — there is only one slot). */
|
||||
export async function unloadGenerator(): Promise<void> {
|
||||
const baseApiUrl = getApiBaseUrl();
|
||||
const response = await fetch(`${baseApiUrl}/generators/unload`, {
|
||||
method: "POST",
|
||||
});
|
||||
if (!response.ok) {
|
||||
const error = await response
|
||||
.json()
|
||||
.catch(() => ({ detail: "Failed to unload model" }));
|
||||
throw new Error(error.detail || "Failed to unload model");
|
||||
}
|
||||
}
|
||||
|
||||
// --- Engine logs ---
|
||||
|
||||
export interface EngineLogs {
|
||||
lines: string[];
|
||||
total: number;
|
||||
}
|
||||
|
||||
/**
|
||||
* Incremental tail of the engine's stdout/stderr. Poll with
|
||||
* `after=<total from the previous response>`.
|
||||
*/
|
||||
export async function getEngineLogs(after: number = 0): Promise<EngineLogs> {
|
||||
const baseApiUrl = getApiBaseUrl();
|
||||
const response = await fetch(`${baseApiUrl}/engine/logs?after=${after}`);
|
||||
if (!response.ok) {
|
||||
throw new Error("Failed to fetch engine logs");
|
||||
}
|
||||
return response.json();
|
||||
}
|
||||
|
||||
// --- Datasets ---
|
||||
|
||||
export interface Dataset {
|
||||
@@ -548,35 +418,3 @@ export function getDatasetMediaUrl(
|
||||
const baseApiUrl = getApiBaseUrl();
|
||||
return `${baseApiUrl}/datasets/${datasetId}/media/${encodeURIComponent(fileName)}`;
|
||||
}
|
||||
|
||||
// --- Cluster ---
|
||||
|
||||
/** Per-GPU telemetry within a cluster node (same shape as GpuInfo). */
|
||||
export type ClusterGpu = GpuInfo;
|
||||
|
||||
export interface ClusterNode {
|
||||
hostname: string;
|
||||
ip: string | null;
|
||||
is_this_host: boolean;
|
||||
cpus: number | null;
|
||||
ray_gpus: number | null;
|
||||
available: boolean;
|
||||
error: string | null;
|
||||
gpus: ClusterGpu[];
|
||||
}
|
||||
|
||||
export interface ClusterSnapshot {
|
||||
mode: "ray" | "local";
|
||||
nodes: ClusterNode[];
|
||||
resources: { gpus_total: number; gpus_available: number } | null;
|
||||
error: string | null;
|
||||
}
|
||||
|
||||
export async function getClusterStatus(): Promise<ClusterSnapshot> {
|
||||
const baseApiUrl = getApiBaseUrl();
|
||||
const response = await fetch(`${baseApiUrl}/cluster`);
|
||||
if (!response.ok) {
|
||||
throw new Error("Failed to fetch cluster status");
|
||||
}
|
||||
return response.json();
|
||||
}
|
||||
|
||||
@@ -1,273 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the single-slot resident generator: preload, list, unload.
|
||||
|
||||
Exactly one VideoGenerator lives in memory. One load at a time; loading a new
|
||||
config always releases the old instance; unload deletes it. The generator is
|
||||
faked at the ``_create_generator_into_slot``/VideoGenerator boundary.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo_studio.job_runner import JobRunner, JobStatus
|
||||
|
||||
|
||||
class _FakeGenerator:
|
||||
def __init__(self) -> None:
|
||||
self.shutdown_calls = 0
|
||||
|
||||
def shutdown(self) -> None:
|
||||
self.shutdown_calls += 1
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def runner():
|
||||
r = JobRunner.__new__(JobRunner) # skip __init__: no DB/Manager needed
|
||||
r._jobs = {}
|
||||
r._jobs_lock = threading.Lock()
|
||||
r._generator = None
|
||||
r._generator_config = None
|
||||
r._generator_state = "empty"
|
||||
r._generator_error = None
|
||||
r._generator_lock = threading.Lock()
|
||||
r._load_lock = threading.Lock()
|
||||
r._worker_log_queue = None # no Manager in unit tests
|
||||
import queue as _queue
|
||||
r._loader_queue = _queue.Queue()
|
||||
threading.Thread(target=r._loader_loop, daemon=True).start()
|
||||
return r
|
||||
|
||||
|
||||
def _install_fake_loader(runner, monkeypatch, made: list | None = None,
|
||||
gate: threading.Event | None = None,
|
||||
fail: str | None = None):
|
||||
"""Replace the VideoGenerator load inside _create_generator_into_slot."""
|
||||
import fastvideo_studio.job_runner as jr
|
||||
|
||||
class _FakeVG:
|
||||
@staticmethod
|
||||
def from_pretrained(model_path, **kwargs):
|
||||
if gate is not None:
|
||||
assert gate.wait(5), "test gate never opened"
|
||||
if fail is not None:
|
||||
raise RuntimeError(fail)
|
||||
gen = _FakeGenerator()
|
||||
if made is not None:
|
||||
made.append((model_path, gen))
|
||||
return gen
|
||||
|
||||
monkeypatch.setitem(__import__("sys").modules, "fastvideo",
|
||||
SimpleNamespace(VideoGenerator=_FakeVG))
|
||||
return jr
|
||||
|
||||
|
||||
def _wait_state(runner, state, timeout=5.0):
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
with runner._generator_lock:
|
||||
if runner._generator_state == state:
|
||||
return runner._slot_entry()
|
||||
time.sleep(0.02)
|
||||
raise AssertionError(f"never reached state {state}: {runner._generator_state}")
|
||||
|
||||
|
||||
def test_preload_then_ready_then_idempotent(runner, monkeypatch):
|
||||
made = []
|
||||
_install_fake_loader(runner, monkeypatch, made=made)
|
||||
|
||||
entry = runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
|
||||
assert entry["state"] in ("loading", "ready")
|
||||
_wait_state(runner, "ready")
|
||||
|
||||
again = runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
|
||||
assert again["state"] == "ready"
|
||||
assert len(made) == 1 # same config never reloads
|
||||
|
||||
|
||||
def test_only_one_load_at_a_time(runner, monkeypatch):
|
||||
gate = threading.Event()
|
||||
_install_fake_loader(runner, monkeypatch, gate=gate)
|
||||
|
||||
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
|
||||
with pytest.raises(RuntimeError, match="already in progress"):
|
||||
runner.preload_generator(model_id="org/b", workload_type="t2v", num_gpus=8)
|
||||
gate.set()
|
||||
_wait_state(runner, "ready")
|
||||
|
||||
|
||||
def test_new_config_always_releases_old_instance(runner, monkeypatch):
|
||||
made = []
|
||||
_install_fake_loader(runner, monkeypatch, made=made)
|
||||
|
||||
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
|
||||
_wait_state(runner, "ready")
|
||||
first = made[0][1]
|
||||
|
||||
runner.preload_generator(model_id="org/b", workload_type="t2v", num_gpus=8)
|
||||
_wait_state(runner, "ready")
|
||||
assert first.shutdown_calls == 1 # old instance released, not stacked
|
||||
assert [m[0] for m in made] == ["org/a", "org/b"]
|
||||
assert len(runner.list_generators()) == 1
|
||||
assert runner.list_generators()[0]["model_id"] == "org/b"
|
||||
|
||||
|
||||
def test_failed_load_reports_error_and_allows_retry(runner, monkeypatch):
|
||||
_install_fake_loader(runner, monkeypatch, fail="no CUDA on this box")
|
||||
runner.preload_generator(model_id="org/broken", workload_type="t2v", num_gpus=1)
|
||||
failed = _wait_state(runner, "failed")
|
||||
assert "no CUDA" in failed["error"]
|
||||
|
||||
# retry after failure must be accepted (this was the reported bug)
|
||||
made = []
|
||||
_install_fake_loader(runner, monkeypatch, made=made)
|
||||
runner.preload_generator(model_id="org/broken", workload_type="t2v", num_gpus=1)
|
||||
_wait_state(runner, "ready")
|
||||
assert len(made) == 1
|
||||
|
||||
|
||||
def test_unload_deletes_instance(runner, monkeypatch):
|
||||
made = []
|
||||
_install_fake_loader(runner, monkeypatch, made=made)
|
||||
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
|
||||
_wait_state(runner, "ready")
|
||||
|
||||
assert runner.unload_generator() is True
|
||||
assert made[0][1].shutdown_calls == 1
|
||||
assert runner.list_generators() == []
|
||||
assert runner._generator is None
|
||||
assert runner.unload_generator() is False # nothing resident
|
||||
|
||||
|
||||
def test_unload_refuses_while_inference_job_runs(runner, monkeypatch):
|
||||
made = []
|
||||
_install_fake_loader(runner, monkeypatch, made=made)
|
||||
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
|
||||
_wait_state(runner, "ready")
|
||||
runner._jobs["j1"] = SimpleNamespace(id="j1", status=JobStatus.RUNNING, job_type="inference")
|
||||
|
||||
with pytest.raises(RuntimeError, match="j1"):
|
||||
runner.unload_generator()
|
||||
assert made[0][1].shutdown_calls == 0
|
||||
|
||||
runner._jobs["j1"].job_type = "finetune" # training doesn't block
|
||||
assert runner.unload_generator() is True
|
||||
|
||||
|
||||
def test_job_waits_for_matching_preload(runner, monkeypatch):
|
||||
gate = threading.Event()
|
||||
made = []
|
||||
_install_fake_loader(runner, monkeypatch, made=made, gate=gate)
|
||||
|
||||
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
|
||||
got = []
|
||||
t = threading.Thread(
|
||||
target=lambda: got.append(
|
||||
runner._get_or_create_generator("org/a", "t2v", 8)),
|
||||
daemon=True)
|
||||
t.start()
|
||||
time.sleep(0.2)
|
||||
assert not got # job blocked on the in-flight load
|
||||
gate.set()
|
||||
t.join(5)
|
||||
assert got and got[0] is made[0][1]
|
||||
assert len(made) == 1 # the job reused the preloaded instance
|
||||
|
||||
|
||||
def test_job_with_different_config_replaces_slot(runner, monkeypatch):
|
||||
made = []
|
||||
_install_fake_loader(runner, monkeypatch, made=made)
|
||||
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
|
||||
_wait_state(runner, "ready")
|
||||
|
||||
gen = runner._get_or_create_generator("org/b", "t2v", 4)
|
||||
assert gen is made[1][1]
|
||||
assert made[0][1].shutdown_calls == 1 # old released before new load
|
||||
assert runner.list_generators()[0]["model_id"] == "org/b"
|
||||
|
||||
|
||||
def test_job_during_replace_never_sees_half_replaced_slot(runner, monkeypatch):
|
||||
"""Regression: user restarted a stale 1-gpu job while an 8-gpu generator
|
||||
was resident; the replace nulled the generator with state still 'ready',
|
||||
and the next (correct) job grabbed None -> AttributeError on
|
||||
generate_video. All transitions now serialize under the load lock."""
|
||||
made = []
|
||||
gate = threading.Event()
|
||||
_install_fake_loader(runner, monkeypatch, made=made)
|
||||
runner.preload_generator(model_id="org/h3", workload_type="t2v", num_gpus=8)
|
||||
_wait_state(runner, "ready")
|
||||
|
||||
_install_fake_loader(runner, monkeypatch, made=made, gate=gate)
|
||||
results: dict[str, Any] = {}
|
||||
|
||||
def job_a(): # stale job: mismatching config triggers a slow replace
|
||||
results["a"] = runner._get_or_create_generator("org/h3", "t2v", 1)
|
||||
|
||||
def job_b(): # correct job arriving mid-replace
|
||||
time.sleep(0.3)
|
||||
results["b"] = runner._get_or_create_generator("org/h3", "t2v", 8)
|
||||
|
||||
ta = threading.Thread(target=job_a, daemon=True)
|
||||
tb = threading.Thread(target=job_b, daemon=True)
|
||||
ta.start(); tb.start()
|
||||
time.sleep(0.6)
|
||||
gate.set()
|
||||
ta.join(10); tb.join(10)
|
||||
|
||||
assert results["a"] is not None and hasattr(results["a"], "shutdown")
|
||||
assert results["b"] is not None and hasattr(results["b"], "shutdown")
|
||||
# b arrived second, so the slot ends at b's 8-gpu config
|
||||
assert runner.list_generators()[0]["num_gpus"] == 8
|
||||
|
||||
|
||||
def test_config_dict_matches_request_defaults(runner):
|
||||
"""GeneratorRequest defaults and CreateJobRequest defaults must resolve to
|
||||
the same slot config — else preloading never matches the job."""
|
||||
from fastvideo_studio.models import CreateJobRequest, GeneratorRequest
|
||||
|
||||
job = CreateJobRequest(model_id="m", prompt="p").model_dump()
|
||||
pre = GeneratorRequest(model_id="m").model_dump()
|
||||
assert runner._generator_config_dict(**pre) == runner._generator_config_dict(
|
||||
**{k: job[k] for k in pre})
|
||||
|
||||
|
||||
def test_engine_log_buffer_incremental_tail():
|
||||
from fastvideo_studio.server import _EngineLogBuffer
|
||||
|
||||
buf = _EngineLogBuffer(maxlen=3)
|
||||
buf.write("a\nb\n")
|
||||
buf.write("c") # partial line: not visible yet
|
||||
lines, total = buf.get_lines(0)
|
||||
assert lines == ["a", "b"] and total == 2
|
||||
buf.write("!\nd\ne\n") # completes "c!", then overflows the ring
|
||||
lines, total = buf.get_lines(total)
|
||||
assert lines == ["c!", "d", "e"] and total == 5
|
||||
# reader far behind: dropped lines are skipped, no crash
|
||||
lines, _ = buf.get_lines(0)
|
||||
assert lines == ["c!", "d", "e"]
|
||||
|
||||
|
||||
def test_engine_feed_drives_job_progress(runner):
|
||||
from fastvideo_studio.job_runner import Job
|
||||
|
||||
job = Job(id="j-prog", model_id="m", prompt="p")
|
||||
runner._active_inference_job = job
|
||||
# ray wraps relayed lines in ANSI color codes — the bridge must strip them
|
||||
runner.feed_engine_line("\x1b[36m(RayWorkerWrapper pid=123, ip=10.0.0.2)\x1b[0m denoising: 40%|████ | 20/50 [00:30<00:45, 1.5s/it]")
|
||||
assert job._log_buf.progress == 40.0
|
||||
assert job._log_buf.progress_msg == "20/50 steps"
|
||||
|
||||
# driver-side lines (no ray prefix) are NOT double-fed
|
||||
before = job._log_buf.get_lines()[1]
|
||||
runner.feed_engine_line("INFO 08-07 [video_generator.py] driver line")
|
||||
assert job._log_buf.get_lines()[1] == before
|
||||
|
||||
# no active job: no-op
|
||||
runner._active_inference_job = None
|
||||
runner.feed_engine_line("(RayWorkerWrapper pid=123) 90%|████| 45/50")
|
||||
assert job._log_buf.progress == 40.0
|
||||
@@ -74,10 +74,8 @@ def test_snapshot_shapes_devices(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
}
|
||||
|
||||
|
||||
def test_snapshot_tolerates_missing_sensors(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setitem(sys.modules, "pynvml",
|
||||
_make_fake_pynvml(broken_sensors=True))
|
||||
def test_snapshot_tolerates_missing_sensors(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setitem(sys.modules, "pynvml", _make_fake_pynvml(broken_sensors=True))
|
||||
snap = gpu_mod.get_gpu_snapshot()
|
||||
assert snap["available"] is True
|
||||
g = snap["gpus"][0]
|
||||
@@ -86,8 +84,7 @@ def test_snapshot_tolerates_missing_sensors(
|
||||
assert g["power_limit_watts"] is None
|
||||
|
||||
|
||||
def test_snapshot_reports_nvml_failure(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_snapshot_reports_nvml_failure(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
fake = _make_fake_pynvml()
|
||||
fake.nvmlInit = lambda: (_ for _ in ()).throw(_NVMLError("driver gone"))
|
||||
monkeypatch.setitem(sys.modules, "pynvml", fake)
|
||||
|
||||
@@ -130,8 +130,7 @@ def test_dmd_builds_three_role_models_and_method_knobs() -> None:
|
||||
|
||||
|
||||
def test_dmd_vsa_maps_to_training_vsa_sparsity() -> None:
|
||||
config = build_training_config(
|
||||
_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
|
||||
config = build_training_config(_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
|
||||
assert config["training"]["vsa"]["sparsity"] == 0.9
|
||||
|
||||
|
||||
@@ -161,31 +160,27 @@ def test_validation_callback_only_when_file_given() -> None:
|
||||
without = build_training_config(_job("full_t2v"), "out")
|
||||
assert "validation" not in without["callbacks"]
|
||||
|
||||
with_file = build_training_config(
|
||||
_job("full_t2v", validation_dataset_file="val.json"), "out")
|
||||
with_file = build_training_config(_job("full_t2v", validation_dataset_file="val.json"), "out")
|
||||
validation = with_file["callbacks"]["validation"]
|
||||
assert validation["dataset_file"] == "val.json"
|
||||
assert validation["pipeline_target"].endswith(".WanPipeline")
|
||||
assert validation["sampling_steps"] == [50]
|
||||
|
||||
dmd = build_training_config(
|
||||
_job("dmd_t2v", validation_dataset_file="val.json"), "out")
|
||||
dmd = build_training_config(_job("dmd_t2v", validation_dataset_file="val.json"), "out")
|
||||
validation = dmd["callbacks"]["validation"]
|
||||
assert validation["pipeline_target"].endswith(".WanDMDPipeline")
|
||||
assert validation["sampling_steps"] == [3]
|
||||
assert validation["sampling_timesteps"] == [1000, 757, 522]
|
||||
|
||||
# KD/ODE-init has no sampling-based validation pipeline.
|
||||
ode = build_training_config(
|
||||
_job("ode_init", validation_dataset_file="val.json"), "out")
|
||||
ode = build_training_config(_job("ode_init", validation_dataset_file="val.json"), "out")
|
||||
assert "validation" not in ode["callbacks"]
|
||||
|
||||
|
||||
def test_ltx2_models_are_rejected() -> None:
|
||||
assert is_ltx2_model("Lightricks/LTX-2-19B")
|
||||
with pytest.raises(ValueError, match="LTX-2 training is not supported"):
|
||||
build_training_config(
|
||||
_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
|
||||
build_training_config(_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
|
||||
|
||||
|
||||
def test_unknown_workload_is_rejected() -> None:
|
||||
@@ -195,8 +190,7 @@ def test_unknown_workload_is_rejected() -> None:
|
||||
|
||||
def test_invalid_denoising_steps_are_rejected() -> None:
|
||||
with pytest.raises(ValueError, match="Invalid DMD denoising steps"):
|
||||
build_training_config(
|
||||
_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
|
||||
build_training_config(_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
|
||||
|
||||
|
||||
def test_training_env_has_no_backend_override() -> None:
|
||||
@@ -211,8 +205,7 @@ def test_workloads_match_frontend_job_config() -> None:
|
||||
(src/lib/jobConfig.ts) — drift means creatable-but-unrunnable jobs."""
|
||||
import re
|
||||
|
||||
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" /
|
||||
"jobConfig.ts").read_text(encoding="utf-8")
|
||||
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" / "jobConfig.ts").read_text(encoding="utf-8")
|
||||
all_types = set(re.findall(r'type:\s*"([^"]+)"', job_config))
|
||||
inference_types = {"t2v", "i2v", "t2i"}
|
||||
assert inference_types <= all_types, "jobConfig.ts parse failed"
|
||||
|
||||
+1
-1
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
|
||||
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
|
||||
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
|
||||
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
|
||||
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
|
||||
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
|
||||
|
||||
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
|
||||
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
{
|
||||
"recipes": [
|
||||
{
|
||||
"id": "fastwan21-t2v",
|
||||
"task": "Text to video",
|
||||
"label": "FastWan2.1 1.3B (distilled + VSA)",
|
||||
"model": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"source": "scripts/inference/inference_wan_VSA_DMD_1_3B.yaml",
|
||||
"command": "FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml"
|
||||
},
|
||||
{
|
||||
"id": "wan22-t2v",
|
||||
"task": "Text to video",
|
||||
"label": "Wan2.2 A14B",
|
||||
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_wan2_2.py",
|
||||
"command": "python examples/inference/basic/basic_wan2_2.py"
|
||||
},
|
||||
{
|
||||
"id": "wan21-i2v",
|
||||
"task": "Image to video",
|
||||
"label": "Wan2.1 14B 480P",
|
||||
"model": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
"source": "scripts/inference/inference_wan_i2v.yaml",
|
||||
"command": "fastvideo generate --config scripts/inference/inference_wan_i2v.yaml"
|
||||
},
|
||||
{
|
||||
"id": "turbowan22-i2v",
|
||||
"task": "Image to video",
|
||||
"label": "TurboWan2.2 A14B",
|
||||
"model": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_turbodiffusion_i2v.py",
|
||||
"command": "python examples/inference/basic/basic_turbodiffusion_i2v.py"
|
||||
},
|
||||
{
|
||||
"id": "wan22-ti2v",
|
||||
"task": "Text or image to video",
|
||||
"label": "Wan2.2 TI2V 5B",
|
||||
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_wan2_2_ti2v.py",
|
||||
"command": "python examples/inference/basic/basic_wan2_2_ti2v.py"
|
||||
},
|
||||
{
|
||||
"id": "matrix-game-2",
|
||||
"task": "Interactive world",
|
||||
"label": "Matrix Game 2.0",
|
||||
"model": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"source": "examples/inference/basic/basic_matrixgame2.py",
|
||||
"command": "python examples/inference/basic/basic_matrixgame2.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
(() => {
|
||||
let recipesPromise;
|
||||
|
||||
const loadRecipes = (url) => {
|
||||
recipesPromise ||= fetch(url).then((response) => {
|
||||
if (!response.ok) throw new Error(`HTTP ${response.status}`);
|
||||
return response.json();
|
||||
});
|
||||
return recipesPromise;
|
||||
};
|
||||
|
||||
const init = () => {
|
||||
document.querySelectorAll("[data-cookbook]").forEach(async (root) => {
|
||||
if (root.dataset.initialized) return;
|
||||
root.dataset.initialized = "true";
|
||||
|
||||
const select = root.querySelector("[data-cookbook-recipe]");
|
||||
const model = root.querySelector("[data-cookbook-model]");
|
||||
const source = root.querySelector("[data-cookbook-source]");
|
||||
const command = root.querySelector("[data-cookbook-command]");
|
||||
const status = root.querySelector("[data-cookbook-status]");
|
||||
|
||||
try {
|
||||
const { recipes } = await loadRecipes(root.dataset.recipes);
|
||||
const byId = new Map(recipes.map((recipe) => [recipe.id, recipe]));
|
||||
const groups = new Map();
|
||||
|
||||
select.replaceChildren();
|
||||
recipes.forEach((recipe) => {
|
||||
if (!groups.has(recipe.task)) {
|
||||
const group = document.createElement("optgroup");
|
||||
group.label = recipe.task;
|
||||
groups.set(recipe.task, group);
|
||||
select.append(group);
|
||||
}
|
||||
groups.get(recipe.task).append(new Option(recipe.label, recipe.id));
|
||||
});
|
||||
|
||||
const render = () => {
|
||||
const recipe = byId.get(select.value);
|
||||
model.textContent = recipe.model;
|
||||
source.textContent = recipe.source;
|
||||
source.href = `https://github.com/hao-ai-lab/FastVideo/blob/main/${recipe.source}`;
|
||||
command.textContent = recipe.command;
|
||||
status.textContent = `${recipe.label} selected.`;
|
||||
};
|
||||
|
||||
select.addEventListener("change", render);
|
||||
select.disabled = false;
|
||||
render();
|
||||
} catch (error) {
|
||||
status.textContent = "Recipes could not be loaded. Use the examples link below.";
|
||||
console.error("Failed to load FastVideo cookbook recipes", error);
|
||||
}
|
||||
});
|
||||
};
|
||||
|
||||
if (window.document$) window.document$.subscribe(init);
|
||||
else document.addEventListener("DOMContentLoaded", init);
|
||||
})();
|
||||
@@ -42,6 +42,46 @@ img {
|
||||
margin: 0 auto;
|
||||
}
|
||||
|
||||
.cookbook-picker {
|
||||
padding: 1rem;
|
||||
border: 0.05rem solid var(--md-default-fg-color--lightest);
|
||||
border-radius: 0.2rem;
|
||||
}
|
||||
|
||||
.cookbook-picker select {
|
||||
width: 100%;
|
||||
padding: 0.6rem;
|
||||
color: var(--md-default-fg-color);
|
||||
background: var(--md-default-bg-color);
|
||||
border: 0.05rem solid var(--md-default-fg-color--lighter);
|
||||
border-radius: 0.2rem;
|
||||
}
|
||||
|
||||
.cookbook-picker dl {
|
||||
display: grid;
|
||||
grid-template-columns: max-content 1fr;
|
||||
gap: 0.25rem 1rem;
|
||||
}
|
||||
|
||||
.cookbook-picker dt {
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.cookbook-picker dd {
|
||||
margin: 0;
|
||||
min-width: 0;
|
||||
overflow-wrap: anywhere;
|
||||
}
|
||||
|
||||
.cookbook-picker__status {
|
||||
position: absolute;
|
||||
width: 1px;
|
||||
height: 1px;
|
||||
overflow: hidden;
|
||||
clip: rect(0, 0, 0, 0);
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.md-typeset .copy-page-button.md-button {
|
||||
float: right;
|
||||
margin: 0 0 1rem 1rem;
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# Inference Cookbook
|
||||
|
||||
Choose a complete recipe maintained in the FastVideo repository. Each command
|
||||
runs its checked-in source directly, so coupled model, GPU, offload, and
|
||||
attention settings do not drift into unsupported combinations.
|
||||
|
||||
The commands expect a local clone:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git
|
||||
cd FastVideo
|
||||
```
|
||||
|
||||
<div class="cookbook-picker" data-cookbook data-recipes="../assets/cookbook-recipes.json">
|
||||
<label for="cookbook-recipe"><strong>Recipe</strong></label>
|
||||
<select id="cookbook-recipe" data-cookbook-recipe disabled>
|
||||
<option>Loading recipes…</option>
|
||||
</select>
|
||||
<dl>
|
||||
<dt>Model</dt>
|
||||
<dd data-cookbook-model>Loading…</dd>
|
||||
<dt>Source</dt>
|
||||
<dd><a data-cookbook-source href="../inference/examples/basic/">Browse maintained examples</a></dd>
|
||||
</dl>
|
||||
<pre><code class="language-bash" data-cookbook-command>Loading…</code></pre>
|
||||
<p class="cookbook-picker__status" role="status" aria-live="polite" data-cookbook-status></p>
|
||||
<noscript>
|
||||
JavaScript is needed for the recipe picker. Browse the
|
||||
<a href="../inference/examples/examples_inference_index/">inference examples</a>
|
||||
instead.
|
||||
</noscript>
|
||||
</div>
|
||||
|
||||
## Customize a recipe
|
||||
|
||||
Start from the checked-in source, then change only the settings your model
|
||||
supports:
|
||||
|
||||
- [Configuration](../inference/configuration.md) covers the Python and CLI
|
||||
config surfaces.
|
||||
- [Optimizations](../inference/optimizations.md) covers attention backends,
|
||||
compilation, and memory tradeoffs.
|
||||
- [Support matrix](../inference/support_matrix.md) lists supported models and
|
||||
optimizations.
|
||||
@@ -0,0 +1,128 @@
|
||||
# Fast mode (RIFE) — Apple Silicon
|
||||
|
||||
`--fast` makes local generation ~2.7× faster by **generating fewer frames and
|
||||
interpolating the rest** with an Apple-Silicon-native RIFE model, instead of
|
||||
denoising every frame. Video-diffusion denoise is dominated by self-attention,
|
||||
which is O(tokens²); halving the frames cuts the token count ~2× and the denoise
|
||||
compute ~3.7×, so the wall-clock drops far more than 2×. RIFE (which estimates
|
||||
its own optical flow — no motion vectors needed) fills the dropped frames back
|
||||
in for ~1.4 s, and a light unsharp pass counters its softening.
|
||||
|
||||
Measured on the 1.3B INT8 QAD model (fox, 480×832×81, M4): generate 41 + RIFE→81
|
||||
runs in ~35 s of denoise vs ~90 s full, at reconstruction MS-SSIM **0.97**.
|
||||
Reproduce with `python -m fastvideo.benchmarks.eval_metalfx_rife --mode int8`.
|
||||
|
||||
> **Note:** Apple's *MetalFX* frame interpolation is **not** usable here — it
|
||||
> requires game-engine motion vectors + depth, which diffusion output lacks. We
|
||||
> use the video-native **`rife-mlx`** model instead (Metal-backed, torch-free).
|
||||
|
||||
## Install
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[mlx]" # RIFE ships vendored; this only needs MLX
|
||||
```
|
||||
|
||||
## Use
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/mlx_wan_prompt_to_video.py \
|
||||
--mlx-checkpoint <FastWan2.1-T2V-1.3B-INT8-QAD> \
|
||||
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
|
||||
--num-frames 81 --fast \
|
||||
--output-path video_samples/fox_fast.mp4
|
||||
```
|
||||
|
||||
`--num-frames` stays the *target* length; fast mode generates the smallest
|
||||
VAE-aligned keyframe count that RIFE can interpolate to that target.
|
||||
|
||||
| Flag | Default | Meaning |
|
||||
|---|---|---|
|
||||
| `--fast` / `--no-fast` | off | enable fast mode |
|
||||
| `--fast-factor` | 2 | generate 1/factor of the frames (2 = half) |
|
||||
| `--fast-sharpen` | 0.6 | light unsharp strength to counter RIFE softness (0 disables) |
|
||||
|
||||
Fast mode composes with everything else (`--mlx-quantization int8`,
|
||||
`--mlx-compile`, TAEHV vs `--decode-backend wan-vae`). Keep `--fast-factor` at 2
|
||||
for quality — larger temporal gaps are where RIFE starts inventing motion.
|
||||
|
||||
## Spatial fast mode (`--fast-spatial`)
|
||||
|
||||
The spatial twin of `--fast`: instead of dropping frames, drop pixels. Denoise
|
||||
*and decode* at `height/width // fast-spatial-scale`, then resample the decoded
|
||||
frames up to the requested size. Self-attention is O(tokens²), so halving each
|
||||
spatial axis cuts the token count 4× and the denoise time far more than that —
|
||||
measured on the 1.3B INT8 QAD model at 480×832×81, M4 Max: **86.1 s → 10.3 s**
|
||||
of denoise. It composes with `--fast`; both together run the same clip in
|
||||
**4.5 s** of denoise.
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/mlx_wan_prompt_to_video.py \
|
||||
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
|
||||
--height 480 --width 832 --num-frames 81 --fast-spatial \
|
||||
--output-path video_samples/fox_fast_spatial.mp4
|
||||
```
|
||||
|
||||
| Flag | Default | Meaning |
|
||||
|---|---|---|
|
||||
| `--fast-spatial` / `--no-fast-spatial` | off | enable spatial fast mode |
|
||||
| `--fast-spatial-scale` | 2 | denoise at 1/scale of each spatial axis |
|
||||
| `--fast-spatial-upsample-mode` | `lanczos` | pixel interpolation kernel (`lanczos`, `cubic`, `bilinear`, `nearest`) |
|
||||
| `--fast-spatial-sharpen` | 0.4 | light unsharp strength to counter resampling softness (0 disables) |
|
||||
|
||||
### The upsample must happen in pixel space
|
||||
|
||||
This is the one thing to get right. The obvious implementation — bilinearly
|
||||
upsample the finished latents and decode at the target size — **does not work**,
|
||||
and produces a distinctive failure: correct composition and silhouette under a
|
||||
smeared, hazy veil, with ringing along strong edges.
|
||||
|
||||
A Wan latent cell is a *learned code* for an 8×8 (Wan2.1) or 16×16 (Wan2.2)
|
||||
pixel block, not a low-pass sample of the image. The average of two adjacent
|
||||
codes is not the code of the averaged blocks; it is a vector the decoder was
|
||||
never trained on. Measured on Wan2.1-1.3B at 480×832, a 2× bilinear latent
|
||||
upsample destroys **62%** of the latent's high-frequency energy while leaving
|
||||
its overall magnitude intact — exactly the signature of that veil. At Wan2.2-5B
|
||||
the same operation degrades to black or noise.
|
||||
|
||||
Decoded RGB frames have no such problem: an image *is* a sampled 2-D signal, so
|
||||
Lanczos interpolation is the operation it was defined for. The result is soft —
|
||||
it carries stage-1's real detail budget and no more — but clean and coherent.
|
||||
|
||||
`--refine` gets away with a latent-space upsample only because a second DMD pass
|
||||
re-denoises the hand-off; spatial fast mode passes the latent straight to the
|
||||
decoder, so it cannot.
|
||||
|
||||
## Refine (`--refine`) stage-2 timesteps
|
||||
|
||||
`--refine` hands stage 1 to stage 2 as `(1 - sigma) * upsampled + sigma * noise`,
|
||||
where `sigma` comes from the *first* stage-2 timestep. FastWan's DMD grid opens
|
||||
at `t=1000`, which is `sigma == 1` exactly — so a stage-2 grid that starts there
|
||||
weights the stage-1 result at zero and refine silently degrades into a plain
|
||||
full-resolution run at twice the cost.
|
||||
|
||||
Left unset, `--refine-dmd-denoising-steps` now derives the stage-2 grid from the
|
||||
stage-1 one with leading full-noise steps dropped (`1000,757,522` → `757,522`).
|
||||
That keeps the pass on timesteps the distilled student was trained on while
|
||||
letting stage-1 structure through: hand-off `sigma = 0.757`, stage-1 weight
|
||||
`0.243`. Passing a grid that starts at full noise is now an error rather than a
|
||||
silently wasted pass.
|
||||
|
||||
The run prints the resolved hand-off so it is visible:
|
||||
|
||||
```
|
||||
[refine] stage-2 hand-off sigma=0.7568 (stage-1 weight 0.2432)
|
||||
```
|
||||
|
||||
There is a trade-off in choosing that grid. Later start = more of the draft
|
||||
survives, but fewer stage-2 steps. On Wan2.1 the default `757,522` gives weight
|
||||
0.243 with two steps; `--refine-dmd-denoising-steps 522` gives weight 0.478 with
|
||||
one. `--refine-sigma` decouples the noise level from the timestep entirely — it
|
||||
logs a warning, because the DiT is then told a timestep that does not match the
|
||||
noise it receives.
|
||||
|
||||
**Wan2.2-5B has a lower ceiling.** Its warped schedule maps `1000,757,522` to
|
||||
sigmas `1.000, 0.940, 0.845`, so the best available stage-1 weight is **0.060**
|
||||
(vs 0.243 at 1.3B). Un-warped (`--no-warp`) the same grid gives `1.000, 0.757,
|
||||
0.522` and a weight of 0.243 — but warping is what matches the FastVideo
|
||||
sampling schedule, so turning it off changes the timesteps the distilled student
|
||||
sees. Which is better at 5B is unresolved and needs a run on real 5B weights.
|
||||
@@ -76,6 +76,10 @@ surfaces:
|
||||
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
|
||||
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
|
||||
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
|
||||
VSA_tile_size: "VSA-H3 tile geometry request; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_decode: "MiniMax-H3 sequence-parallel VAE decode opt-in; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_encode: "MiniMax-H3 sequence-parallel reference VAE encode opt-in; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_decode_strategy: "Chunk-transport collective for vae_parallel_decode; model-specific optimization not yet represented in the typed public schema."
|
||||
attention_backend: "Process-wide default attention-backend request applied per component at load time; kernel-selection knob not yet represented in the typed public schema."
|
||||
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
|
||||
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
|
||||
|
||||
import itertools
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
@@ -19,6 +20,40 @@ GENERATED_DOC_PREFIXES = (
|
||||
"training/examples/",
|
||||
"distillation/examples/",
|
||||
)
|
||||
COOKBOOK_DATA = ROOT_DIR / "docs/assets/cookbook-recipes.json"
|
||||
COOKBOOK_SOURCE_ROOTS = (
|
||||
ROOT_DIR / "examples/inference",
|
||||
ROOT_DIR / "scripts/inference",
|
||||
)
|
||||
|
||||
|
||||
def validate_cookbook() -> None:
|
||||
"""Keep cookbook entries tied to checked-in runnable sources."""
|
||||
recipes = json.loads(COOKBOOK_DATA.read_text(encoding="utf-8")).get("recipes")
|
||||
if not isinstance(recipes, list) or not recipes:
|
||||
raise ValueError(f"{COOKBOOK_DATA}: recipes must be a non-empty list")
|
||||
|
||||
seen: set[str] = set()
|
||||
for recipe in recipes:
|
||||
required = ("id", "task", "label", "model", "source", "command")
|
||||
missing = {key for key in required if not recipe.get(key)}
|
||||
if missing:
|
||||
raise ValueError(f"Cookbook recipe is missing: {', '.join(sorted(missing))}")
|
||||
if recipe["id"] in seen:
|
||||
raise ValueError(f"Duplicate cookbook recipe id: {recipe['id']}")
|
||||
seen.add(recipe["id"])
|
||||
|
||||
source = (ROOT_DIR / recipe["source"]).resolve()
|
||||
if not any(source.is_relative_to(root.resolve()) for root in COOKBOOK_SOURCE_ROOTS):
|
||||
raise ValueError(f"Cookbook source is outside an approved directory: {recipe['source']}")
|
||||
if not source.is_file():
|
||||
raise ValueError(f"Cookbook source does not exist: {recipe['source']}")
|
||||
|
||||
source_text = source.read_text(encoding="utf-8")
|
||||
if recipe["model"] not in source_text:
|
||||
raise ValueError(f"Cookbook model is not present in {recipe['source']}: {recipe['model']}")
|
||||
if recipe["source"] not in recipe["command"]:
|
||||
raise ValueError(f"Cookbook command does not invoke its source: {recipe['id']}")
|
||||
|
||||
|
||||
def fix_case(text: str) -> str:
|
||||
@@ -536,6 +571,7 @@ def on_pre_build(config, **kwargs):
|
||||
MkDocs hook to generate examples before building the documentation.
|
||||
This function is called automatically by MkDocs' native hook system.
|
||||
"""
|
||||
validate_cookbook()
|
||||
print("Generating example documentation...")
|
||||
generate_examples(generate_main_index=True)
|
||||
print("Example documentation generated successfully!")
|
||||
@@ -549,6 +585,7 @@ def on_page_context(context, page, **kwargs):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
validate_cookbook()
|
||||
print("Generating example documentation...")
|
||||
generate_examples(generate_main_index=True)
|
||||
print("Example documentation generated successfully!")
|
||||
|
||||
@@ -65,6 +65,7 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
|
||||
- [Quick Start](quick_start.md) - Generate your first video
|
||||
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
|
||||
|
||||
@@ -49,10 +49,12 @@ brew install ffmpeg
|
||||
|
||||
### Installation
|
||||
|
||||
FastWan's native Apple Silicon runtime requires the `mlx` extra.
|
||||
|
||||
#### With uv (recommended)
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
uv pip install "fastvideo[mlx]"
|
||||
```
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
@@ -60,7 +62,7 @@ uv pip install fastvideo
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
uv pip install "fastvideo[mlx]"
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -76,13 +78,13 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
uv pip install -e ".[mlx]"
|
||||
```
|
||||
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
uv pip install -e ".[mlx]"
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
@@ -23,61 +23,21 @@ Also optionally install flash-attn:
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Basic Usage
|
||||
## Choose a maintained recipe
|
||||
|
||||
### Text-to-Video Generation
|
||||
The cookbook selects complete, checked-in recipes instead of mixing model,
|
||||
parallelism, offload, and attention settings independently.
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
|
||||
|
||||
def main():
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
|
||||
# Define a prompt for your video
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
|
||||
|
||||
# Generate the video
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path="my_videos/", # Controls where videos are saved
|
||||
save_video=True
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
### Image-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
def main():
|
||||
# Create the generator
|
||||
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
|
||||
|
||||
# Set up parameters with an initial image
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
generator.generate_video(prompt, sampling_param=sampling_param,
|
||||
output_path="my_videos/",
|
||||
save_video=True)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
!!! tip "Need more control?"
|
||||
Start from a maintained recipe, then use the
|
||||
[configuration](../inference/configuration.md) and
|
||||
[optimization](../inference/optimizations.md) guides for supported changes.
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
|
||||
- [Installation Guide](installation.md) - Detailed installation instructions
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
|
||||
|
||||
@@ -64,6 +64,7 @@ column links a runnable script in `examples/inference/basic/` where one exists.
|
||||
| ltx2 | `FastVideo/LTX2-Distilled-Diffusers`<br>`FastVideo/LTX2.3-Distilled-Diffusers`<br>`FastVideo/LTX-2.3-Distilled-Diffusers` | T2V | [basic_ltx2_distilled.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2_distilled.py) |
|
||||
| ltx2 | `Lightricks/LTX-2.3`<br>`FastVideo/LTX2.3-base`<br>`FastVideo/LTX2.3-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
|
||||
| ltx2 | `Lightricks/LTX-2`<br>`FastVideo/LTX2-base`<br>`FastVideo/LTX2-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
|
||||
| mmaudio | `FastVideo/MMAudio-large-44k-v2-Diffusers` | V2A, T2A | [basic_mmaudio.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mmaudio.py) |
|
||||
| matrixgame | `FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-Base-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Diffusers`<br>`mignonjia/mg_longtuning_distilled_zelda`<br>`mignonjia/mg_sf_distilled_zelda_1k_steps`<br>`mignonjia/mg_sf_distilled_zelda`<br>`mignonjia/mg_causal_zelda`<br>`mignonjia/mg_bidirectional_zelda` | I2V | [basic_matrixgame2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame2.py) |
|
||||
| matrixgame | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | I2V | [basic_matrixgame3.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame3.py) |
|
||||
| minimax_h3 | `MiniMaxAI/MiniMax-H3` | T2V, I2V | [T2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_t2v.py)<br>[FL2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_fl2va.py)<br>[Ref2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_ref2va.py) |
|
||||
@@ -94,6 +95,10 @@ column links a runnable script in `examples/inference/basic/` where one exists.
|
||||
(`StableAudioT2AConfig` / `StableAudioOpenSmallConfig`); they are registered
|
||||
under the generic T2V workload option in the registry.
|
||||
|
||||
**Note (MMAudio)**: the registered Hugging Face model ID is reserved but not
|
||||
yet public. Follow the [MMAudio inference guide](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/pipelines/basic/mmaudio/README.md)
|
||||
to convert the official weights locally and set `MMAUDIO_MODEL_PATH`.
|
||||
|
||||
**Note (MiniMax H3)**: T2VA, FL2VA, and Ref2VA all generate video with stereo
|
||||
audio. Use the Ref2VA example when passing ordered image, video, or audio
|
||||
references.
|
||||
@@ -173,6 +178,17 @@ optimizations: absence means **untested**, not incompatible.
|
||||
| Matrix Game 3.0 Base Distilled | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | 720x1280 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| GEN3C Cosmos 7B | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | 704px1280p | ❌ | ❌ | ❌ | ⭕ | ⭕ |
|
||||
|
||||
## Apple Silicon native runtime
|
||||
|
||||
| Release path | Model | Mode | Validated hardware | Status |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| MLX FastWan T2V | FastWan-QAD-INT8-1.3B `[release model ID pending]` | 480x832, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 36 GB unified-memory class, MLX 0.31.2 | Release candidate; requires release-owner visual sign-off |
|
||||
|
||||
This is a text-to-video-only source-install release. It is validated on the
|
||||
hardware listed above; MLX allocator caps are not evidence of support for a
|
||||
physical 16 GB Mac. See [Apple Silicon FastWan](../getting_started/installation/mps.md)
|
||||
for the supported command and release gates.
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
***Lucy Edit Dev uses a non-commercial model license. FastVideo support is
|
||||
|
||||
@@ -33,6 +33,14 @@ For the typed config/request path added during the inference API refactor:
|
||||
python examples/inference/basic/basic_dmd_new_api.py
|
||||
```
|
||||
|
||||
For the few-step (4-step, DMD2-distilled) MiniMax-H3 preview, generating synchronized video and audio, optionally with block-sparse VSA attention:
|
||||
```
|
||||
python examples/inference/basic/basic_fasth3.py --prompt "your prompt" [--vsa-sparsity 0.9]
|
||||
```
|
||||
The default checkpoint `FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1` is private on the Hub while its license review completes; until it flips public, pass `--model-path` with a local snapshot of the release.
|
||||
|
||||
On Blackwell (sm_100) GPUs with a `fastvideo-kernel` build that carries the sm_100a block-sparse extension, `--vsa-kernel sm100a` routes the tile-64 attention forwards through the CUDA kernel instead of Triton (it sets `FASTVIDEO_VSA_SM100A=1` before the pipeline boots); if the extension or the arch is missing, the run warns once and falls back to Triton.
|
||||
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -12,11 +14,11 @@ def main():
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
@@ -24,22 +26,19 @@ def main():
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
|
||||
@@ -30,8 +30,7 @@ def main():
|
||||
"and casting reflections onto adjacent vehicles. "
|
||||
"The motion creates space in the lineup, signaling activity within the otherwise quiet station. "
|
||||
"It then comes to a smooth stop, resuming its position in line. "
|
||||
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene."
|
||||
)
|
||||
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene.")
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
@@ -47,4 +46,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
@@ -31,8 +31,7 @@ def main():
|
||||
"The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. "
|
||||
"The metal surface beneath the torch shows ongoing signs of heating and melting. "
|
||||
"The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, "
|
||||
"underscoring the ongoing nature of the welding operation."
|
||||
)
|
||||
"underscoring the ongoing nature of the welding operation.")
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
@@ -46,6 +45,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -50,4 +50,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
@@ -5,6 +5,8 @@ from fastvideo import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2"
|
||||
|
||||
|
||||
def main():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
@@ -14,10 +16,10 @@ def main():
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
@@ -25,7 +27,6 @@ def main():
|
||||
load_end_time = time.perf_counter()
|
||||
load_time = load_end_time - load_start_time
|
||||
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.num_frames = 81
|
||||
|
||||
@@ -39,18 +40,16 @@ def main():
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
start_time = time.perf_counter()
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
|
||||
end_time = time.perf_counter()
|
||||
gen_time2 = end_time - start_time
|
||||
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Time taken to generate video2: {gen_time2} seconds")
|
||||
|
||||
@@ -32,11 +32,9 @@ def main():
|
||||
),
|
||||
# PR 2 still routes a few advanced inference knobs through the
|
||||
# compatibility bridge until they get first-class typed fields.
|
||||
pipeline=PipelineSelection(
|
||||
experimental={
|
||||
"VSA_sparsity": 0.8,
|
||||
},
|
||||
),
|
||||
pipeline=PipelineSelection(experimental={
|
||||
"VSA_sparsity": 0.8,
|
||||
}, ),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
@@ -44,14 +42,12 @@ def main():
|
||||
load_end_time = time.perf_counter()
|
||||
load_time = load_end_time - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
|
||||
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
|
||||
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
|
||||
"LED umbrella. Steam rises from a street food cart, and a cat darts "
|
||||
"across the screen. Raindrops are visible on the camera lens, creating "
|
||||
"a cinematic bokeh effect."
|
||||
)
|
||||
prompt = ("A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
|
||||
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
|
||||
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
|
||||
"LED umbrella. Steam rises from a street food cart, and a cat darts "
|
||||
"across the screen. Raindrops are visible on the camera lens, creating "
|
||||
"a cinematic bokeh effect.")
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
output=OutputConfig(
|
||||
@@ -66,13 +62,11 @@ def main():
|
||||
end_time = time.perf_counter()
|
||||
gen_time = end_time - start_time
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently "
|
||||
"in the breeze, enhancing the lion's commanding presence. The tone is "
|
||||
"vibrant, embodying the raw energy of the wild. Low angle, steady "
|
||||
"tracking shot, cinematic."
|
||||
)
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently "
|
||||
"in the breeze, enhancing the lion's commanding presence. The tone is "
|
||||
"vibrant, embodying the raw energy of the wild. Low angle, steady "
|
||||
"tracking shot, cinematic.")
|
||||
request2 = GenerationRequest(
|
||||
prompt=prompt2,
|
||||
output=OutputConfig(
|
||||
|
||||
@@ -2,7 +2,6 @@ import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
|
||||
|
||||
|
||||
@@ -46,10 +45,8 @@ def main():
|
||||
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
|
||||
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
|
||||
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
|
||||
"action_speed_list": [
|
||||
float(value)
|
||||
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
|
||||
],
|
||||
"action_speed_list":
|
||||
[float(value) for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")],
|
||||
}
|
||||
if image_path:
|
||||
kwargs["image_path"] = image_path
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
|
||||
|
||||
FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1 is a 4-step distillation of
|
||||
MiniMaxAI/MiniMax-H3 (data-free DMD2): it walks a 4-step grid on the release
|
||||
sampler's shift-12 schedule instead of the base model's 50 steps, generating
|
||||
synchronized video and audio in one pipeline call.
|
||||
|
||||
The student was trained with block-sparse video attention (VSA, 64-token
|
||||
tiles) and its checkpoint carries the trained sparse-gate parameters
|
||||
(``attn.to_gate_compress``), so this script always runs the VSA-H3 attention
|
||||
backend. At the default ``--vsa-sparsity 0.0`` the attention math is exactly
|
||||
dense (every tile is selected); raise the sparsity for additional speedup.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
CompileConfig,
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1")
|
||||
# The HF repo is private while the MiniMax H3 Community License review
|
||||
# completes; until it flips public, pass --model-path with a local
|
||||
# snapshot of the release instead (e.g. the team export at
|
||||
# /mnt/lustre/vlm-wlsaidhi/fastvideo/exports/FastVideo-Minimax-FastH3-Preview-v0.1).
|
||||
parser.add_argument("--prompt", required=True)
|
||||
parser.add_argument("--output", default="outputs/fasth3")
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
parser.add_argument("--width", type=int, default=1344)
|
||||
parser.add_argument("--num-frames", type=int, default=124)
|
||||
# num_inference_steps counts sigma-GRID POINTS, matching the base model's
|
||||
# convention ("50 steps" = a 50-point grid = 49 transformer forwards). The
|
||||
# student's distilled 4-step grid is 4 FORWARDS, i.e. a 5-point grid
|
||||
# (t = 1000, 750, 500, 250 -> 0 on the shift-12 schedule) — so the correct
|
||||
# default here is 5. Other grids are off-distribution.
|
||||
parser.add_argument("--steps",
|
||||
type=int,
|
||||
default=5,
|
||||
help="num_inference_steps = sigma-grid points; N points run N-1 denoising "
|
||||
"forwards. 5 (default) is the distilled 4-forward grid")
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--num-gpus", type=int, default=4)
|
||||
parser.add_argument("--vsa-sparsity",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Run-level VSA sparsity in [0, 1). 0.0 (default) selects every tile, which is "
|
||||
"exactly dense attention; the student was trained at 0.9")
|
||||
# 64 is the trained contract: the student was TRAINED with 64-token
|
||||
# (4,4,4) tiles, and its to_gate_compress gates were learned against
|
||||
# pooling at that granularity — keep 64 unless you are ablating.
|
||||
parser.add_argument("--vsa-tile-size",
|
||||
type=int,
|
||||
choices=(64, 256),
|
||||
default=64,
|
||||
help="VSA-H3 tile size in tokens; 64 (default) is what the student was trained "
|
||||
"with and runs the native Triton block-sparse path, 256 is the FA4-CuTe-capable "
|
||||
"geometry for ablations")
|
||||
parser.add_argument("--vsa-kernel",
|
||||
choices=("triton", "sm100a"),
|
||||
default="triton",
|
||||
help="Block-sparse kernel for the tile-64 attention forward: triton (default, "
|
||||
"portable fwd+bwd) or sm100a — the opt-in Blackwell CUDA forward "
|
||||
"(fastvideo_kernel.block_sparse_attn_sm100a). sm100a needs an sm_100 GPU and a "
|
||||
"fastvideo-kernel build that carries the extension; if a precondition fails at "
|
||||
"run time the attention layer logs one warning and falls back to Triton. Only "
|
||||
"meaningful with --vsa-tile-size 64")
|
||||
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
|
||||
parser.add_argument("--compile-mode",
|
||||
default=None,
|
||||
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
|
||||
parser.add_argument("--repeats",
|
||||
type=int,
|
||||
default=1,
|
||||
help="generate N times; with --torch-compile the first run pays "
|
||||
"compilation, so steady-state is the last repeat")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
output_dir = Path(args.output)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if args.vsa_kernel == "sm100a":
|
||||
# The attention backend reads FASTVIDEO_VSA_SM100A per forward; set it
|
||||
# before the pipeline boots so spawned GPU workers inherit it. The
|
||||
# kernel is forward-only and inference runs under no-grad, so every
|
||||
# denoising forward qualifies for the CUDA route.
|
||||
os.environ["FASTVIDEO_VSA_SM100A"] = "1"
|
||||
|
||||
# Boot-time run configuration, folded into FastVideoArgs (the same route
|
||||
# examples/inference/basic/basic_minimax_h3_t2v.py uses for sparsity):
|
||||
# - attention_backend: the checkpoint carries trained to_gate_compress
|
||||
# gates, which only exist under the VSA-H3 backend — a dense-backend
|
||||
# load would reject them as unexpected weights. Layers that do not
|
||||
# support VSA-H3 (e.g. the token refiner) fall back to flash attention.
|
||||
# - VSA_tile_size: forwarded even at sparsity 0.0 because the gate-compress
|
||||
# branch pools per tile, and the gates were trained at 64 tokens/tile.
|
||||
experimental: dict[str, object] = {
|
||||
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
|
||||
"VSA_tile_size": args.vsa_tile_size,
|
||||
}
|
||||
if args.vsa_sparsity > 0.0:
|
||||
experimental["VSA_sparsity"] = args.vsa_sparsity
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
pipeline=PipelineSelection(experimental=experimental),
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=args.num_gpus > 1,
|
||||
parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus),
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
dit_layerwise=False,
|
||||
text_encoder=True,
|
||||
vae=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
compile=CompileConfig(
|
||||
enabled=args.torch_compile,
|
||||
mode=args.compile_mode,
|
||||
),
|
||||
),
|
||||
))
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=args.steps,
|
||||
# the base model is guidance-distilled; the student inherits it
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_dir / "fasth3.mp4"),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
print(f"Output written to: {result.video_path}")
|
||||
if result.generation_time is not None:
|
||||
# machine-readable: benchmark harnesses parse this line to separate
|
||||
# generation from model-load time (last occurrence = steady state)
|
||||
print(f"Generation time: {result.generation_time:.2f}s")
|
||||
for _ in range(args.repeats - 1):
|
||||
result = generator.generate(request)
|
||||
if result.generation_time is not None:
|
||||
print(f"Generation time: {result.generation_time:.2f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -64,12 +64,8 @@ def main() -> None:
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (
|
||||
args.num_gpus if args.num_gpus > 1 else 1
|
||||
)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (
|
||||
1 if args.num_gpus > 1 else args.num_gpus
|
||||
)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
|
||||
@@ -21,7 +21,6 @@ from fastvideo.api import (
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
DEFAULT_PROMPT = "a brushed steel espresso machine on a marble counter, morning window light"
|
||||
|
||||
|
||||
|
||||
@@ -9,11 +9,9 @@ import re
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
(
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
|
||||
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
|
||||
"35mm, bokeh"
|
||||
),
|
||||
("a cinematic photo of a red panda wearing a tiny backpack, standing on a "
|
||||
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
|
||||
"35mm, bokeh"),
|
||||
]
|
||||
|
||||
|
||||
@@ -42,9 +40,7 @@ def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(
|
||||
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
|
||||
)
|
||||
p = argparse.ArgumentParser(description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.", )
|
||||
p.add_argument(
|
||||
"--model-path",
|
||||
default="official_weights/FLUX.1-dev",
|
||||
@@ -108,9 +104,7 @@ def main() -> None:
|
||||
try:
|
||||
for i, prompt in enumerate(prompts):
|
||||
seed = args.seed + i
|
||||
filename_base = (
|
||||
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
|
||||
)
|
||||
filename_base = (f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}")
|
||||
_remove_existing_outputs(args.out_dir, filename_base)
|
||||
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
|
||||
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
|
||||
|
||||
@@ -33,21 +33,20 @@ MODEL_PATH = os.environ.get("GAMECRAFT_MODEL_PATH", "FastVideo/HunyuanGameCraft-
|
||||
|
||||
# Default prompts for demo
|
||||
DEFAULT_PROMPTS = {
|
||||
"village": "A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
|
||||
"temple": "A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
|
||||
"forest": "A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
|
||||
"village":
|
||||
"A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
|
||||
"temple":
|
||||
"A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
|
||||
"forest":
|
||||
"A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
|
||||
"beach": "A tropical beach with crystal clear turquoise water, white sand, and palm trees swaying in the breeze.",
|
||||
}
|
||||
|
||||
# I2V: default reference image (URL). Can override with a local path.
|
||||
DEFAULT_I2V_IMAGE_URL = (
|
||||
"https://huggingface.co/datasets/huggingface/documentation-images/"
|
||||
"resolve/main/diffusers/astronaut.jpg"
|
||||
)
|
||||
DEFAULT_I2V_PROMPT = (
|
||||
"An astronaut hatching from an egg, on the surface of the moon, "
|
||||
"the darkness and depth of space realised in the background."
|
||||
)
|
||||
DEFAULT_I2V_IMAGE_URL = ("https://huggingface.co/datasets/huggingface/documentation-images/"
|
||||
"resolve/main/diffusers/astronaut.jpg")
|
||||
DEFAULT_I2V_PROMPT = ("An astronaut hatching from an egg, on the surface of the moon, "
|
||||
"the darkness and depth of space realised in the background.")
|
||||
|
||||
OUTPUT_PATH = "video_samples_gamecraft"
|
||||
|
||||
|
||||
@@ -26,51 +26,35 @@ from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="GEN3C video generation")
|
||||
parser.add_argument("--model_path",
|
||||
type=str,
|
||||
default="converted_weights/GEN3C-Cosmos-7B")
|
||||
parser.add_argument("--image_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Input image for 3D cache conditioning")
|
||||
parser.add_argument("--prompt",
|
||||
type=str,
|
||||
default="A slow camera pan over a sunlit landscape.")
|
||||
parser.add_argument("--model_path", type=str, default="converted_weights/GEN3C-Cosmos-7B")
|
||||
parser.add_argument("--image_path", type=str, default=None, help="Input image for 3D cache conditioning")
|
||||
parser.add_argument("--prompt", type=str, default="A slow camera pan over a sunlit landscape.")
|
||||
parser.add_argument(
|
||||
"--negative_prompt",
|
||||
type=str,
|
||||
default=(
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
|
||||
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
|
||||
"flickering. Overall, the video is of poor quality."
|
||||
),
|
||||
default=("The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
|
||||
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
|
||||
"flickering. Overall, the video is of poor quality."),
|
||||
)
|
||||
parser.add_argument("--trajectory",
|
||||
type=str,
|
||||
default="left",
|
||||
choices=[
|
||||
"left", "right", "up", "down", "zoom_in",
|
||||
"zoom_out", "clockwise", "counterclockwise", "none"
|
||||
])
|
||||
parser.add_argument(
|
||||
"--trajectory",
|
||||
type=str,
|
||||
default="left",
|
||||
choices=["left", "right", "up", "down", "zoom_in", "zoom_out", "clockwise", "counterclockwise", "none"])
|
||||
parser.add_argument("--movement_distance", type=float, default=0.3)
|
||||
parser.add_argument("--camera_rotation",
|
||||
type=str,
|
||||
default="center_facing",
|
||||
choices=[
|
||||
"center_facing", "no_rotation",
|
||||
"trajectory_aligned"
|
||||
])
|
||||
choices=["center_facing", "no_rotation", "trajectory_aligned"])
|
||||
parser.add_argument("--height", type=int, default=704)
|
||||
parser.add_argument("--width", type=int, default=1280)
|
||||
parser.add_argument("--num_frames", type=int, default=121)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=35)
|
||||
parser.add_argument("--guidance_scale", type=float, default=1.0)
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs_video/gen3c.mp4")
|
||||
parser.add_argument("--output_path", type=str, default="outputs_video/gen3c.mp4")
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@ import json
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -12,31 +14,38 @@ def main():
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
video = generator.generate_video(prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt="",
|
||||
num_frames=81,
|
||||
fps=16)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
|
||||
video2 = generator.generate_video(prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt="",
|
||||
num_frames=81,
|
||||
fps=16)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -3,38 +3,37 @@ import json
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15_1080p"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
|
||||
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
|
||||
# or "weizhou03/HunyuanVideo-1.5-Diffusers-1080p" # 720p -> 1080p
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
|
||||
|
||||
|
||||
@@ -6,6 +6,8 @@ DEFAULT_PROMPT = 'A paved pathway leads towards a stone arch bridge spanning a c
|
||||
DEFAULT_IMAGE = 'https://raw.githubusercontent.com/Tencent-Hunyuan/HY-WorldPlay/main/assets/img/test.png'
|
||||
|
||||
OUTPUT_PATH = "video_samples_hyworld"
|
||||
|
||||
|
||||
def main():
|
||||
import argparse
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ IMAGE_PATH = "assets/girl.png"
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
|
||||
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
@@ -19,9 +19,7 @@ def main():
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A woman stands up and walks away"
|
||||
)
|
||||
prompt = ("A woman stands up and walks away")
|
||||
_ = generator.generate_video(
|
||||
prompt,
|
||||
image_path=IMAGE_PATH,
|
||||
|
||||
@@ -2,6 +2,7 @@ from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
@@ -17,21 +18,28 @@ def main():
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
_ = generator.generate_video(prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=512,
|
||||
width=768,
|
||||
num_frames=121)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=512,
|
||||
width=768,
|
||||
num_frames=121)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -6,7 +6,6 @@ from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
DATASET_DIR = REPO_ROOT / "examples" / "dataset" / "lingbotworld2"
|
||||
OUTPUT_PATH = REPO_ROOT / "outputs" / "lingbotworld2_causal_fast.mp4"
|
||||
|
||||
@@ -3,17 +3,19 @@ from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embeddin
|
||||
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
OUTPUT_PATH = "video_samples_lingbotworld"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
|
||||
@@ -21,20 +21,16 @@ import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
"A woman sits at a wooden table by the window in a cozy café. She reaches out "
|
||||
"with her right hand, picks up the white coffee cup from the saucer, and gently "
|
||||
"brings it to her lips to take a sip. After drinking, she places the cup back on "
|
||||
"the table and looks out the window, enjoying the peaceful atmosphere."
|
||||
)
|
||||
PROMPT = ("A woman sits at a wooden table by the window in a cozy café. She reaches out "
|
||||
"with her right hand, picks up the white coffee cup from the saucer, and gently "
|
||||
"brings it to her lips to take a sip. After drinking, she places the cup back on "
|
||||
"the table and looks out the window, enjoying the peaceful atmosphere.")
|
||||
|
||||
NEGATIVE_PROMPT = (
|
||||
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards"
|
||||
)
|
||||
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards")
|
||||
|
||||
# Input image path
|
||||
IMAGE_PATH = "assets/girl.png"
|
||||
@@ -51,20 +47,20 @@ def basic_generation():
|
||||
print("=" * 60)
|
||||
print("LongCat I2V: Basic Generation (50 steps, 480p)")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_i2v_basic"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -79,7 +75,7 @@ def basic_generation():
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -94,11 +90,11 @@ def distill_refine_generation():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat I2V: Distill + Refine Pipeline")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Stage 1: Distilled generation (16 steps at 480p)
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
num_gpus=1,
|
||||
@@ -111,9 +107,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_i2v_distill"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -128,14 +124,14 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
# Stage 2: Refinement (480p -> 768p)
|
||||
print("\n[Stage 2] Refinement (480p -> 768p with BSA)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
# Find the actual saved video file from stage 1
|
||||
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
|
||||
if not video_files:
|
||||
@@ -143,7 +139,7 @@ def distill_refine_generation():
|
||||
# Use the most recently created video file
|
||||
distill_video_path = max(video_files, key=os.path.getmtime)
|
||||
print(f"Using stage 1 video: {distill_video_path}")
|
||||
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
# Note: Refinement uses the T2V model (not I2V) since it's upscaling the generated video
|
||||
# For BSA [4, 4, 8]: latent must be divisible by 8
|
||||
@@ -163,9 +159,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_i2v_refine_720p"
|
||||
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -182,7 +178,7 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
@@ -192,13 +188,13 @@ def main():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat Image-to-Video Example")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
|
||||
# Run basic generation
|
||||
basic_generation()
|
||||
|
||||
|
||||
# Run distill+refine pipeline
|
||||
distill_refine_generation()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All generations complete!")
|
||||
print("=" * 60)
|
||||
@@ -206,5 +202,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -15,22 +15,18 @@ import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
"In a realistic photography style, a white boy around seven or eight years old "
|
||||
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
|
||||
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
|
||||
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
|
||||
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
|
||||
"features a green lawn and several tall trees, creating a warm and loving scene."
|
||||
)
|
||||
PROMPT = ("In a realistic photography style, a white boy around seven or eight years old "
|
||||
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
|
||||
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
|
||||
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
|
||||
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
|
||||
"features a green lawn and several tall trees, creating a warm and loving scene.")
|
||||
|
||||
NEGATIVE_PROMPT = (
|
||||
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards"
|
||||
)
|
||||
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards")
|
||||
|
||||
SEED = 42
|
||||
|
||||
@@ -44,20 +40,20 @@ def basic_generation():
|
||||
print("=" * 60)
|
||||
print("LongCat T2V: Basic Generation (50 steps, 480p)")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_t2v_basic"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -71,7 +67,7 @@ def basic_generation():
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -86,11 +82,11 @@ def distill_refine_generation():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat T2V: Distill + Refine Pipeline")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Stage 1: Distilled generation (16 steps at 480p)
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
num_gpus=1,
|
||||
@@ -103,9 +99,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_t2v_distill"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -119,14 +115,14 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
# Stage 2: Refinement (480p -> 720p)
|
||||
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
# Find the actual saved video file from stage 1
|
||||
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
|
||||
if not video_files:
|
||||
@@ -134,7 +130,7 @@ def distill_refine_generation():
|
||||
# Use the most recently created video file
|
||||
distill_video_path = max(video_files, key=os.path.getmtime)
|
||||
print(f"Using stage 1 video: {distill_video_path}")
|
||||
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
refine_generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
@@ -151,9 +147,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_t2v_refine_720p"
|
||||
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -170,7 +166,7 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
@@ -180,13 +176,13 @@ def main():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat Text-to-Video Example")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
|
||||
# Run basic generation
|
||||
basic_generation()
|
||||
|
||||
|
||||
# Run distill+refine pipeline
|
||||
distill_refine_generation()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All generations complete!")
|
||||
print("=" * 60)
|
||||
@@ -194,5 +190,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -21,21 +21,17 @@ import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# Common prompts and settings matching the shell script examples
|
||||
PROMPT = (
|
||||
"A person rides a motorcycle along a long, straight road that stretches between "
|
||||
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
|
||||
"the motorcycle centered between the guardrails, while the scenery passes by on "
|
||||
"both sides. The video captures the journey from the rider's perspective, emphasizing "
|
||||
"the sense of motion and adventure."
|
||||
)
|
||||
PROMPT = ("A person rides a motorcycle along a long, straight road that stretches between "
|
||||
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
|
||||
"the motorcycle centered between the guardrails, while the scenery passes by on "
|
||||
"both sides. The video captures the journey from the rider's perspective, emphasizing "
|
||||
"the sense of motion and adventure.")
|
||||
|
||||
NEGATIVE_PROMPT = (
|
||||
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards"
|
||||
)
|
||||
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
|
||||
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
|
||||
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
|
||||
"three legs, many people in the background, walking backwards")
|
||||
|
||||
# Input video path
|
||||
VIDEO_PATH = "assets/motorcycle.mp4"
|
||||
@@ -55,27 +51,25 @@ def basic_generation():
|
||||
print("=" * 60)
|
||||
print("LongCat VC: Basic Generation (50 steps, 480p)")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Check if video exists
|
||||
if not os.path.exists(VIDEO_PATH):
|
||||
raise FileNotFoundError(
|
||||
f"Video not found at {VIDEO_PATH}. "
|
||||
"Please provide a valid video path."
|
||||
)
|
||||
|
||||
raise FileNotFoundError(f"Video not found at {VIDEO_PATH}. "
|
||||
"Please provide a valid video path.")
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-VC-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
enable_bsa=False,
|
||||
)
|
||||
|
||||
|
||||
output_path = "outputs_video/longcat_vc_basic"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -91,7 +85,7 @@ def basic_generation():
|
||||
guidance_scale=4.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"\nBasic generation complete! Video saved to: {output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
@@ -106,18 +100,16 @@ def distill_refine_generation():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat VC: Distill + Refine Pipeline")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
# Check if video exists
|
||||
if not os.path.exists(VIDEO_PATH):
|
||||
raise FileNotFoundError(
|
||||
f"Video not found at {VIDEO_PATH}. "
|
||||
"Please provide a valid video path."
|
||||
)
|
||||
|
||||
raise FileNotFoundError(f"Video not found at {VIDEO_PATH}. "
|
||||
"Please provide a valid video path.")
|
||||
|
||||
# Stage 1: Distilled generation (16 steps at 480p)
|
||||
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LongCat-Video-VC-Diffusers",
|
||||
num_gpus=1,
|
||||
@@ -130,9 +122,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
|
||||
lora_nickname="distilled",
|
||||
)
|
||||
|
||||
|
||||
distill_output_path = "outputs_video/longcat_vc_distill"
|
||||
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -148,14 +140,14 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
# Stage 2: Refinement (480p -> 720p)
|
||||
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
|
||||
print("-" * 40)
|
||||
|
||||
|
||||
# Find the actual saved video file from stage 1
|
||||
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
|
||||
if not video_files:
|
||||
@@ -163,7 +155,7 @@ def distill_refine_generation():
|
||||
# Use the most recently created video file
|
||||
distill_video_path = max(video_files, key=os.path.getmtime)
|
||||
print(f"Using stage 1 video: {distill_video_path}")
|
||||
|
||||
|
||||
# Create a new generator with refinement LoRA and BSA enabled
|
||||
# Note: Refinement uses the T2V model (not VC) since it's upscaling the generated video
|
||||
refine_generator = VideoGenerator.from_pretrained(
|
||||
@@ -181,9 +173,9 @@ def distill_refine_generation():
|
||||
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
|
||||
lora_nickname="refinement",
|
||||
)
|
||||
|
||||
|
||||
refine_output_path = "outputs_video/longcat_vc_refine_720p"
|
||||
|
||||
|
||||
refine_generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
negative_prompt=NEGATIVE_PROMPT,
|
||||
@@ -200,7 +192,7 @@ def distill_refine_generation():
|
||||
guidance_scale=1.0,
|
||||
seed=SEED,
|
||||
)
|
||||
|
||||
|
||||
print(f"Refinement complete! Video saved to: {refine_output_path}")
|
||||
refine_generator.shutdown()
|
||||
|
||||
@@ -210,13 +202,13 @@ def main():
|
||||
print("\n" + "=" * 60)
|
||||
print("LongCat Video Continuation Example")
|
||||
print("=" * 60 + "\n")
|
||||
|
||||
|
||||
# Run basic generation
|
||||
basic_generation()
|
||||
|
||||
|
||||
# Run distill+refine pipeline
|
||||
distill_refine_generation()
|
||||
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All generations complete!")
|
||||
print("=" * 60)
|
||||
@@ -224,5 +216,3 @@ def main():
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -1,19 +1,16 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic.")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
@@ -36,4 +33,4 @@ def main() -> None:
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -67,25 +67,19 @@ _inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v")
|
||||
)
|
||||
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX23_MODEL_PATH",
|
||||
"FastVideo/LTX-2.3-Distilled-Diffusers")))
|
||||
OUTPUT_DIR = Path(os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v"))
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
DEFAULT_PROMPT = ("A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel.")
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
# Per-stage timing helpers --------------------------------------------------
|
||||
|
||||
|
||||
def _print_stage_breakdown(result: dict, label: str) -> float | None:
|
||||
"""Print stage execution times and return the sum, or None if missing."""
|
||||
logging_info = result.get("logging_info")
|
||||
@@ -114,9 +108,7 @@ def _collect_stage_times(
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
stage_times.setdefault(name, []).append(float(metrics.get("execution_time", 0.0)))
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
@@ -125,21 +117,18 @@ def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
raise FileNotFoundError(f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`.")
|
||||
|
||||
|
||||
# Main ---------------------------------------------------------------------
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py"
|
||||
)
|
||||
raise SystemExit("LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py")
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
@@ -201,11 +190,13 @@ def main() -> None:
|
||||
|
||||
common_kwargs = dict(
|
||||
prompt=PROMPT,
|
||||
negative_prompt="", # distilled is CFG-free; no negative needed
|
||||
guidance_scale=1.0, # CFG=1 for distilled
|
||||
height=1280, width=832, # portrait runway aspect
|
||||
num_frames=121, fps=24, # ~5s clip
|
||||
num_inference_steps=8, # distilled denoise steps
|
||||
negative_prompt="", # distilled is CFG-free; no negative needed
|
||||
guidance_scale=1.0, # CFG=1 for distilled
|
||||
height=1280,
|
||||
width=832, # portrait runway aspect
|
||||
num_frames=121,
|
||||
fps=24, # ~5s clip
|
||||
num_inference_steps=8, # distilled denoise steps
|
||||
# i2v: anchor the input image at frame 0 with full strength.
|
||||
# `ltx2_image_crf=0.0` skips an extra JPEG re-encode of an already
|
||||
# JPEG conditioning image.
|
||||
@@ -251,10 +242,7 @@ def main() -> None:
|
||||
**common_kwargs,
|
||||
)
|
||||
wall = time.perf_counter() - t0
|
||||
e2e = (
|
||||
result.get("e2e_latency")
|
||||
if isinstance(result, dict) else None
|
||||
) or wall
|
||||
e2e = (result.get("e2e_latency") if isinstance(result, dict) else None) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
if isinstance(result, dict):
|
||||
@@ -266,10 +254,8 @@ def main() -> None:
|
||||
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
print(f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s")
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
|
||||
@@ -100,23 +100,14 @@ _inductor.coordinate_descent_tuning = True
|
||||
_inductor.coordinate_descent_check_all_directions = True
|
||||
_inductor.epilogue_fusion = False
|
||||
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(
|
||||
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
|
||||
)
|
||||
)
|
||||
OUTPUT_DIR = Path(
|
||||
os.getenv(
|
||||
"LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"
|
||||
)
|
||||
)
|
||||
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX23_MODEL_PATH",
|
||||
"FastVideo/LTX-2.3-Distilled-Diffusers")))
|
||||
OUTPUT_DIR = Path(os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"))
|
||||
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
|
||||
DEFAULT_PROMPT = (
|
||||
"A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel."
|
||||
)
|
||||
DEFAULT_PROMPT = ("A fashion model takes a slow step forward and shifts her weight, "
|
||||
"the soft fabric of her clothing swaying and rippling with the "
|
||||
"motion, her hair shifting gently, soft even studio lighting on a "
|
||||
"clean light background, elegant slow-motion runway feel.")
|
||||
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
|
||||
|
||||
|
||||
@@ -147,9 +138,7 @@ def _collect_stage_times(
|
||||
return
|
||||
for name, metrics in stages.items():
|
||||
stage_order.setdefault(name, None)
|
||||
stage_times.setdefault(name, []).append(
|
||||
float(metrics.get("execution_time", 0.0))
|
||||
)
|
||||
stage_times.setdefault(name, []).append(float(metrics.get("execution_time", 0.0)))
|
||||
|
||||
|
||||
def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
@@ -157,20 +146,16 @@ def _resolve_refine_upsampler(model_root: str) -> Path:
|
||||
cand = Path(model_root) / name
|
||||
if (cand / "config.json").is_file():
|
||||
return cand
|
||||
raise FileNotFoundError(
|
||||
f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`."
|
||||
)
|
||||
raise FileNotFoundError(f"No refine upsampler directory under {model_root}. "
|
||||
f"Expected `{model_root}/spatial_upscaler/config.json`.")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not I2V_IMAGE:
|
||||
raise SystemExit(
|
||||
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_typed.py"
|
||||
)
|
||||
raise SystemExit("LTX23_I2V_IMAGE is required for i2v. Example:\n"
|
||||
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
|
||||
" python examples/inference/basic/"
|
||||
"basic_ltx2_3_distilled_i2v_typed.py")
|
||||
if not Path(I2V_IMAGE).is_file():
|
||||
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
|
||||
|
||||
@@ -220,10 +205,9 @@ def main() -> None:
|
||||
# model-specific VAE precision / decoder defaults are picked up
|
||||
# the same way the legacy example's
|
||||
# ``PipelineConfig.from_pretrained(model_root)`` did them.
|
||||
components=ComponentConfig(
|
||||
upsampler_weights=str(refine_upsampler_path),
|
||||
# Distilled has no refine LoRA — omit ``lora_path``.
|
||||
),
|
||||
components=ComponentConfig(upsampler_weights=str(refine_upsampler_path),
|
||||
# Distilled has no refine LoRA — omit ``lora_path``.
|
||||
),
|
||||
vae_tiling=False,
|
||||
preset_overrides={
|
||||
"refine": {
|
||||
@@ -278,11 +262,7 @@ def main() -> None:
|
||||
for w in range(warmup_runs):
|
||||
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
|
||||
t0 = time.perf_counter()
|
||||
generator.generate(
|
||||
build_request(
|
||||
OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7
|
||||
)
|
||||
)
|
||||
generator.generate(build_request(OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7))
|
||||
dt = time.perf_counter() - t0
|
||||
warmup_secs.append(dt)
|
||||
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
|
||||
@@ -291,46 +271,30 @@ def main() -> None:
|
||||
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
|
||||
|
||||
for m in range(measured_runs):
|
||||
out_path = (
|
||||
OUTPUT_DIR
|
||||
/ f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4"
|
||||
)
|
||||
print(
|
||||
f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}"
|
||||
)
|
||||
out_path = (OUTPUT_DIR / f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4")
|
||||
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
|
||||
t0 = time.perf_counter()
|
||||
result = generator.generate(
|
||||
build_request(out_path, seed=2002 + m)
|
||||
)
|
||||
result = generator.generate(build_request(out_path, seed=2002 + m))
|
||||
wall = time.perf_counter() - t0
|
||||
# ``e2e_latency`` is currently surfaced via ``result.extra``;
|
||||
# ``GenerationResult`` exposes ``generation_time`` as a
|
||||
# first-class field but the LTX-2 pipeline only fills the
|
||||
# legacy ``e2e_latency`` key. Prefer the explicit one, fall
|
||||
# back to wall-clock.
|
||||
e2e = (
|
||||
result.extra.get("e2e_latency")
|
||||
if hasattr(result, "extra") else None
|
||||
) or wall
|
||||
e2e = (result.extra.get("e2e_latency") if hasattr(result, "extra") else None) or wall
|
||||
measured_secs.append(e2e)
|
||||
print(
|
||||
f"[measured {m + 1}/{measured_runs}] "
|
||||
f"e2e={e2e:.2f}s wall={wall:.2f}s"
|
||||
)
|
||||
print(f"[measured {m + 1}/{measured_runs}] "
|
||||
f"e2e={e2e:.2f}s wall={wall:.2f}s")
|
||||
_print_stage_breakdown(result, f"measured {m + 1}")
|
||||
_collect_stage_times(result, stage_times, stage_order)
|
||||
|
||||
print("\n=== summary ===")
|
||||
print(
|
||||
f"warmup wall-times: "
|
||||
f"{[round(x, 1) for x in warmup_secs]}"
|
||||
)
|
||||
print(f"warmup wall-times: "
|
||||
f"{[round(x, 1) for x in warmup_secs]}")
|
||||
if measured_secs:
|
||||
avg = sum(measured_secs) / len(measured_secs)
|
||||
print(
|
||||
f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
|
||||
)
|
||||
print(f"measured e2e (n={len(measured_secs)}): "
|
||||
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s")
|
||||
if stage_times:
|
||||
print(f"average stage times over {measured_runs} measured runs:")
|
||||
avg_total = 0.0
|
||||
|
||||
@@ -1,21 +1,21 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic.")
|
||||
import os
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
|
||||
@@ -12,15 +12,11 @@ from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.layers.quantization.nvfp4_config import NVFP4Config
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
VALIDATION_JSON = (
|
||||
Path(__file__).resolve().parents[2] / "training" / "finetune" / "ltx2" / "validation.json"
|
||||
)
|
||||
VALIDATION_JSON = (Path(__file__).resolve().parents[2] / "training" / "finetune" / "ltx2" / "validation.json")
|
||||
|
||||
# Override with a local snapshot or converted directory when needed, e.g.
|
||||
# export LTX2_MODEL_PATH=/raid/$USER/hf/FastVideo/LTX2-Distilled-Diffusers
|
||||
MODEL_ID = os.path.expandvars(
|
||||
os.path.expanduser(os.getenv("LTX2_MODEL_PATH", "FastVideo/LTX2-Distilled-Diffusers"))
|
||||
)
|
||||
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX2_MODEL_PATH", "FastVideo/LTX2-Distilled-Diffusers")))
|
||||
OUTPUT_DIR = Path("outputs_video/ltx2_distilled_fast_profile")
|
||||
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
@@ -69,9 +65,7 @@ def print_stage_breakdown(
|
||||
return total
|
||||
|
||||
|
||||
def extract_sr_forward_latency(
|
||||
result: dict,
|
||||
) -> tuple[float | None, list[tuple[str, float]], list[str]]:
|
||||
def extract_sr_forward_latency(result: dict, ) -> tuple[float | None, list[tuple[str, float]], list[str]]:
|
||||
logging_info = result.get("logging_info")
|
||||
if logging_info is None:
|
||||
return None, [], []
|
||||
@@ -89,12 +83,8 @@ def extract_sr_forward_latency(
|
||||
if sr_match_substr:
|
||||
is_sr_stage = sr_match_substr in stage_name_l
|
||||
else:
|
||||
is_sr_stage = (
|
||||
"srdenoisingstage" in stage_name_l
|
||||
or "sr_denoising" in stage_name_l
|
||||
or "upsample" in stage_name_l
|
||||
or ("refine" in stage_name_l and "denois" in stage_name_l)
|
||||
)
|
||||
is_sr_stage = ("srdenoisingstage" in stage_name_l or "sr_denoising" in stage_name_l
|
||||
or "upsample" in stage_name_l or ("refine" in stage_name_l and "denois" in stage_name_l))
|
||||
if not is_sr_stage:
|
||||
continue
|
||||
exec_time = float(stage_metrics.get("execution_time", 0.0))
|
||||
@@ -161,11 +151,9 @@ def resolve_refine_upsampler_path(model_root: str) -> Path:
|
||||
return candidate
|
||||
|
||||
checked = "\n".join(f" - {candidate}" for candidate in candidates)
|
||||
raise FileNotFoundError(
|
||||
"Could not find an LTX2 refine upsampler directory.\n"
|
||||
"Checked:\n"
|
||||
f"{checked}"
|
||||
)
|
||||
raise FileNotFoundError("Could not find an LTX2 refine upsampler directory.\n"
|
||||
"Checked:\n"
|
||||
f"{checked}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
@@ -314,19 +302,15 @@ def main() -> None:
|
||||
|
||||
measured_times = run_times[measured_start_idx:]
|
||||
avg_time = sum(measured_times) / len(measured_times)
|
||||
print(
|
||||
f"Average video generation time over {len(measured_times)} runs "
|
||||
f"(runs {measured_start_idx + 1}-{len(run_times)}, skipping first {warmup_runs} warmup runs): "
|
||||
f"{avg_time:.2f}s"
|
||||
)
|
||||
print(f"Average video generation time over {len(measured_times)} runs "
|
||||
f"(runs {measured_start_idx + 1}-{len(run_times)}, skipping first {warmup_runs} warmup runs): "
|
||||
f"{avg_time:.2f}s")
|
||||
|
||||
measured_e2e_times = e2e_times[measured_start_idx:]
|
||||
avg_e2e_time = sum(measured_e2e_times) / len(measured_e2e_times)
|
||||
print(
|
||||
f"Average end-to-end latency over {len(measured_e2e_times)} runs "
|
||||
f"(runs {measured_start_idx + 1}-{len(e2e_times)}, skipping first {warmup_runs} warmup runs): "
|
||||
f"{avg_e2e_time:.2f}s"
|
||||
)
|
||||
print(f"Average end-to-end latency over {len(measured_e2e_times)} runs "
|
||||
f"(runs {measured_start_idx + 1}-{len(e2e_times)}, skipping first {warmup_runs} warmup runs): "
|
||||
f"{avg_e2e_time:.2f}s")
|
||||
|
||||
if sr_forward_times:
|
||||
avg_sr_forward = sum(sr_forward_times) / len(sr_forward_times)
|
||||
@@ -338,10 +322,8 @@ def main() -> None:
|
||||
|
||||
if non_stage_overhead_times:
|
||||
avg_non_stage_overhead = sum(non_stage_overhead_times) / len(non_stage_overhead_times)
|
||||
print(
|
||||
"Average non-stage overhead over "
|
||||
f"{len(non_stage_overhead_times)} measured runs: {avg_non_stage_overhead:.3f}s"
|
||||
)
|
||||
print("Average non-stage overhead over "
|
||||
f"{len(non_stage_overhead_times)} measured runs: {avg_non_stage_overhead:.3f}s")
|
||||
else:
|
||||
print("Average non-stage overhead unavailable (no stage timings).")
|
||||
finally:
|
||||
|
||||
@@ -13,24 +13,34 @@ MODEL_VARIANT = "base_distilled_model"
|
||||
# Variant-specific settings
|
||||
VARIANT_CONFIG = {
|
||||
"base_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
4,
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"gta_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
2,
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"templerun_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
7,
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples_matrixgame2"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -42,8 +52,8 @@ def main():
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
|
||||
@@ -14,27 +14,40 @@ MODEL_VARIANT = "base_distilled_model"
|
||||
# Variant-specific settings
|
||||
VARIANT_CONFIG = {
|
||||
"base_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"keyboard_dim": 4,
|
||||
"mode": "universal",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
4,
|
||||
"mode":
|
||||
"universal",
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
|
||||
},
|
||||
"gta_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"keyboard_dim": 2,
|
||||
"mode": "gta_drive",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
2,
|
||||
"mode":
|
||||
"gta_drive",
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
|
||||
},
|
||||
"templerun_distilled_model": {
|
||||
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
"keyboard_dim": 7,
|
||||
"mode": "templerun",
|
||||
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
"model_path":
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
"keyboard_dim":
|
||||
7,
|
||||
"mode":
|
||||
"templerun",
|
||||
"image_url":
|
||||
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples_matrixgame2"
|
||||
|
||||
|
||||
async def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -46,8 +59,8 @@ async def main():
|
||||
config["model_path"],
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
@@ -56,11 +69,8 @@ async def main():
|
||||
)
|
||||
|
||||
max_blocks = 50
|
||||
num_frames = 597
|
||||
actions = {
|
||||
"keyboard": torch.zeros((num_frames, config["keyboard_dim"])),
|
||||
"mouse": torch.zeros((num_frames, 2))
|
||||
}
|
||||
num_frames = 597
|
||||
actions = {"keyboard": torch.zeros((num_frames, config["keyboard_dim"])), "mouse": torch.zeros((num_frames, 2))}
|
||||
grid_sizes = torch.tensor([150, 44, 80])
|
||||
mode = config["mode"]
|
||||
|
||||
@@ -81,11 +91,11 @@ async def main():
|
||||
|
||||
for block_id in range(max_blocks):
|
||||
print(f"\n=== Block {block_id + 1}/{max_blocks} ===")
|
||||
|
||||
|
||||
action = await get_current_action_async(mode)
|
||||
keyboard_cond, mouse_cond = expand_action_to_frames(action, 12)
|
||||
await generator.step_async(keyboard_cond, mouse_cond)
|
||||
|
||||
|
||||
if (await asyncio.to_thread(input, "\nContinue? (y/n): ")).lower() == 'n':
|
||||
break
|
||||
|
||||
|
||||
@@ -25,6 +25,13 @@ from fastvideo.pipelines.basic.minimax_h3.packing import resolve_canvas_size
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
parser.add_argument("--image", required=True, help="First-frame image path.")
|
||||
parser.add_argument("--last-image", help="Optional last-frame image path.")
|
||||
parser.add_argument("--output", default="outputs/minimax_h3_fl2va")
|
||||
|
||||
@@ -25,6 +25,13 @@ from fastvideo.pipelines.basic.minimax_h3 import MiniMaxH3Reference
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
parser.add_argument("--reference-video", required=True)
|
||||
parser.add_argument("--reference-audio", help="Optional additional audio reference.")
|
||||
parser.add_argument("--output", default="outputs/minimax_h3_ref2va")
|
||||
|
||||
@@ -22,6 +22,13 @@ from fastvideo.api import (
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
parser.add_argument("--prompt", required=True)
|
||||
parser.add_argument("--output", default="outputs/minimax_h3_t2v")
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
@@ -30,13 +37,15 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--num-gpus", type=int, default=4)
|
||||
parser.add_argument("--torch-compile", action="store_true",
|
||||
help="torch.compile the DiT transformer path")
|
||||
parser.add_argument("--compile-mode", default=None,
|
||||
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
|
||||
parser.add_argument("--compile-mode",
|
||||
default=None,
|
||||
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
|
||||
parser.add_argument("--repeats", type=int, default=1,
|
||||
parser.add_argument("--repeats",
|
||||
type=int,
|
||||
default=1,
|
||||
help="generate N times; with --torch-compile the first run pays "
|
||||
"compilation, so steady-state is the last repeat")
|
||||
"compilation, so steady-state is the last repeat")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
@@ -67,24 +76,24 @@ def main() -> None:
|
||||
))
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_dir / "minimax_h3_t2v.mp4"),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_dir / "minimax_h3_t2v.mp4"),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
print(f"Output written to: {result.video_path}")
|
||||
if result.generation_time is not None:
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MMAudio large-44k-v2 video-to-audio example."""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--video-path", required=True)
|
||||
parser.add_argument("--output-path", default="outputs_audio/mmaudio.wav")
|
||||
parser.add_argument("--duration-seconds", type=float, default=8.0)
|
||||
parser.add_argument("--prompt", default="")
|
||||
parser.add_argument("--negative-prompt", default="music")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
os.environ.get(
|
||||
"MMAUDIO_MODEL_PATH",
|
||||
"converted_weights/mmaudio/large_44k_v2",
|
||||
),
|
||||
workload_type="v2a",
|
||||
num_gpus=1,
|
||||
)
|
||||
result = generator.generate_video(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
video_path=args.video_path,
|
||||
audio_end_in_s=args.duration_seconds,
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
)
|
||||
print(result["video_path"])
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,19 +1,20 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
config.text_encoder_precisions = ["fp16"]
|
||||
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
pipeline_config=config,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
dit_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
dit_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
# Create sampling parameters with reduced number of frames
|
||||
@@ -23,18 +24,19 @@ def main():
|
||||
sampling_param.width = 256
|
||||
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param)
|
||||
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, sampling_param=sampling_param)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -3,6 +3,8 @@ from fastvideo import VideoGenerator
|
||||
# from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -16,27 +18,24 @@ def main():
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
distributed_executor_backend="ray",
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
|
||||
@@ -7,7 +7,6 @@ import os
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a rainy neon-lit street at night, shallow depth of field, sharp focus, 35mm, bokeh",
|
||||
@@ -49,7 +48,9 @@ def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description="Run SD3.5 Medium text-to-image with FastVideo VideoGenerator.")
|
||||
p.add_argument("--model-path", default="stabilityai/stable-diffusion-3.5-medium", help="Path to local diffusers-format SD3.5 weights directory.")
|
||||
p.add_argument("--model-path",
|
||||
default="stabilityai/stable-diffusion-3.5-medium",
|
||||
help="Path to local diffusers-format SD3.5 weights directory.")
|
||||
p.add_argument(
|
||||
"--out-dir",
|
||||
"--outdir",
|
||||
|
||||
@@ -3,6 +3,8 @@ import time
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_causal"
|
||||
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -13,19 +15,18 @@ def main():
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user