Compare commits

..
Author SHA1 Message Date
SolitaryThinker be7cc7a1dd studio: keep frontend polling noise out of the engine console
uvicorn's access log lands in the engine tail, and the UI polls several
endpoints every few seconds — the console was mostly its own GET 200 lines.
access_log=False at the server (applies on restart) + a client-side filter
for 2xx GET/HEAD access-log lines (immediate, and covers servers started
before the flag).
2026-08-08 18:36:37 -07:00
SolitaryThinker 8c5778fa44 studio: create generators on one persistent loader thread
the mp executor's workers arm prctl(PR_SET_PDEATHSIG, SIGKILL), which on
linux fires when the CREATING THREAD exits — not the process. generators
spawned from short-lived threads (preload, per-job loader helpers) lose
every worker to a silent SIGKILL the moment that thread finishes: zombie
workers, broken pipes on the first rpc. all creations now run on a single
long-lived loader thread.
2026-08-07 18:28:36 -07:00
SolitaryThinker 0fbdbe9174 studio: pin dit_layerwise_offload=False at generator creation
FastVideoArgs defaults it True, which disables FSDP and holds a full DiT
copy in host RAM per worker — 4 mp workers x 66GB DiT (+ the Qwen3-VL text
encoder) OOM-killed the node silently right after 'workers ready' (defunct
workers, broken pipes on first RPC). the example scripts always overrode
this; the studio never did.
2026-08-07 18:20:09 -07:00
SolitaryThinker 659cf4fbfe studio: attach the worker log queue at generator creation, not via RPC
sending the Manager-queue proxy over the mp executor's worker pipes
(set_log_queue RPC on each generate) breaks the pipe — that path was never
exercised while the queue was also passed at creation. single-slot now keeps
one runner-wide queue for the generator's whole life (drained per job), and
generate_video no longer sends log_queue at all.
2026-08-07 18:11:48 -07:00
SolitaryThinker 609648b114 studio ui: resolution presets, inline video, Studio nav label, cluster panel
- create-job modal: 480p/720p/1080p preset chips (/32-safe dims) + native
  chip restoring the model preset + portrait/landscape toggle; sliders stay
- completed jobs stream their clip inline in the details sidebar (video tag
  over the existing /video endpoint; png for t2i)
- top bar/nav: Jobs -> Studio (labels only)
- gpus page is now a cluster view over GET /api/cluster: per-node cards with
  util/vram meters, temp/power, mode badge, available/total gpus; compact
  per-node strip; 5s polling, paused while hidden
2026-08-07 00:59:27 -07:00
SolitaryThinker 9a2eda427f studio: GET /api/cluster — per-node GPU telemetry via ray
self-contained NVML probe scheduled on every alive ray node (worker envs
can't import apps/, so cloudpickle carries it by value); falls back to the
local host when not connected. includes ray gpu totals/available.
2026-08-07 00:51:09 -07:00
SolitaryThinker cb7e8e960a studio ui: show the full prompt on job cards (wrap, don't truncate) 2026-08-07 00:41:16 -07:00
SolitaryThinker 4a4e9b3f31 studio ui: compact job cards to a two-line layout
model / status / prompt / meta on one line, slim action row under it;
tighter padding and margins — roughly half the previous card height.
2026-08-07 00:38:58 -07:00
SolitaryThinker 0f2bbcda10 h3: rank-0 denoising progress bar; studio: strip ANSI in the progress bridge
h3's denoising loop emitted no per-step progress at all, and ray wraps
relayed worker lines in ANSI color codes the bridge's prefix match never
survived — together the UI bar sat at 0 until completion. tqdm on rank 0
(mirroring the shared denoising stage; non-tty refreshes are plain lines ray
relays per step) + ANSI strip before matching/feeding.
2026-08-07 00:33:03 -07:00
SolitaryThinker 58f7b913b9 studio: drive job progress from ray-relayed worker output
on the ray backend worker tqdm can't reach the job buffer via the mp queue,
but ray already relays worker stdout to the driver — which the engine tee
captures. lines with ray's actor prefix now feed the running inference job's
log buffer, whose existing tqdm parser moves the UI progress bar. no UI
changes needed.
2026-08-07 00:26:52 -07:00
SolitaryThinker 07f1cfc356 ray executor: wait out async GPU release on shutdown; retry node GPU probe on init
killed actors free their devices asynchronously — a load fired right after a
release could still see 0 GPUs on the node and hard-fail. shutdown now waits
(<=60s) for the GPUs to return to ray's ledger, and initialize_ray_cluster's
node-local GPU probe retries for 60s instead of failing instantly.
2026-08-07 00:05:38 -07:00
SolitaryThinker 3101331020 studio: serialize all generator-slot transitions; wait for ray PG removal
two races hit by real UI use:
- a job-triggered replace nulled the resident generator while state stayed
  'ready', so a concurrent job with the old config grabbed None and died on
  generate_video. every slot transition (preload / job replace / unload) now
  runs under one load lock; jobs block behind in-flight loads instead of
  peeking at half-replaced state.
- ray placement-group removal is async: a load right after a release raced
  the teardown and failed 'no GPU available'. executor shutdown now waits
  (<=30s) for the PG to actually reach REMOVED.
2026-08-07 00:00:50 -07:00
SolitaryThinker 4a63de71d7 studio ui: default to the warm model, populate from model presets
create-job modal now defaults its model to the resident one and, for any
selected model, populates every sampling field from GET /api/models/presets
(fixes generating H3 with wan-style defaults: guidance 5.0 and 81 frames both
fail H3's input validation). selecting the warm model also adopts the slot's
engine config so the job actually reuses the resident instance instead of
silently replacing it.
2026-08-06 23:09:14 -07:00
SolitaryThinker 99df34a78d studio: GET /api/models/presets — per-model sampling defaults
config-only via SamplingParam.from_pretrained (never loads weights, cached);
the create-job modal populates its sampling fields from these.
2026-08-06 22:59:03 -07:00
SolitaryThinker 7b94433092 studio ui: single-slot model panel + engine output console
panel shows the one resident model (load/replace/unload, one load at a
time); engine console tails the server's stdout/stderr via /api/engine/logs
(collapsible, polls only while open, stick-to-bottom).
2026-08-06 22:50:26 -07:00
SolitaryThinker 3a7e7fc69f studio: single resident generator + engine log tail; fix ray PG leak
reload-after-unload failed because RayDistributedExecutor.shutdown() killed
the actors but never removed the placement group — the GPUs stayed reserved
forever. remove it on shutdown.

studio generator management simplified to a single slot: one VideoGenerator
instance at a time, one load in flight at a time (409 otherwise), loading a
new config always releases the old instance, unload shuts it down and
deletes it. jobs wait out an in-flight preload and reuse it on config match.

GET /api/engine/logs: incremental tail of the server's stdout/stderr (ray
relays worker output to the driver, so this captures every rank).
2026-08-06 22:40:39 -07:00
SolitaryThinker 02a10268e9 studio ui: warm models panel on the inference page
lists resident/loading/failed generators, preload from the model dropdown
using the persisted default job options (same cache key as jobs), unload with
409 surfaced as a toast. polls only while something is loading. mock server
+ playwright spec + vitest coverage.
2026-08-06 22:29:03 -07:00
SolitaryThinker 9c9dd70214 studio: preload/list/unload resident generators
generators were already cached across jobs but only implicitly — the first
job per config silently paid the whole model load. adds:
- POST /api/generators/preload: load a model into memory ahead of time
  (async, idempotent; state loading/ready/failed)
- GET /api/generators: what's resident or in flight
- POST /api/generators/unload: shutdown + free VRAM (409 while inference
  jobs run)
cache key builder shared with the job path and pinned by a test so preload
and job lookup can't drift apart.
2026-08-06 22:20:28 -07:00
SolitaryThinker dc13c59dca fix ray executor: implement the log-queue abstract methods
RayDistributedExecutor became uninstantiable when set_log_queue/
clear_log_queue were added to the Executor ABC (only the mp executor got
implementations). No-op on ray: a Manager queue can't cross nodes.
2026-08-06 19:16:33 -07:00
SolitaryThinker 6dd75ef5ab studio: multi-node inference knobs
- forward num_gpus to from_pretrained (was dropped: multi-gpu only ever
  happened via sp_size)
- FASTVIDEO_STUDIO_EXECUTOR_BACKEND=ray env knob so a deployment on a ray
  cluster can span nodes
- FASTVIDEO_STUDIO_MODEL_PATHS="id=/local/dir" serves a registered model id
  from local weights (avoids re-downloading checkpoints already on shared fs)
- accept an existing local dir as model_id in POST /api/jobs
2026-08-06 18:42:44 -07:00
832 changed files with 28186 additions and 53062 deletions
@@ -10,7 +10,9 @@ from pathlib import Path
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Clone a reference repo for FastVideo parity tests.")
parser = argparse.ArgumentParser(
description="Clone a reference repo for FastVideo parity tests."
)
parser.add_argument("repo_url", help="Official reference repository URL")
parser.add_argument("target_dir", help="Directory to clone into")
parser.add_argument("--branch", help="Branch or tag to clone")
@@ -60,7 +62,9 @@ def gitignore_entry_for(target: Path) -> str:
try:
relative = resolved.relative_to(root)
except ValueError as exc:
raise ValueError("--update-gitignore requires target_dir to be under the current directory") from exc
raise ValueError(
"--update-gitignore requires target_dir to be under the current directory"
) from exc
text = relative.as_posix().rstrip("/")
return "/" + text + "/"
@@ -8,12 +8,14 @@ import os
import sys
from pathlib import Path
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Download a HF model snapshot or selected files into a local directory.")
description="Download a HF model snapshot or selected files into a local directory."
)
parser.add_argument("repo_id", help="HF repo id, for example Org/Model")
parser.add_argument("local_dir", help="Destination directory")
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
@@ -10,6 +10,7 @@ import sys
from pathlib import Path
from typing import Any
HF_TOKEN_ENV_KEYS = ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY")
RAW_WEIGHT_SUFFIXES = (".safetensors", ".pt", ".pth", ".ckpt", ".bin")
KNOWN_COMPONENTS = {
@@ -33,7 +34,8 @@ KNOWN_COMPONENTS = {
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown.")
description="Classify a HF repo or local directory as Diffusers, raw, custom, or unknown."
)
parser.add_argument("source", help="HF repo id or local weights directory")
parser.add_argument("--repo-type", default="model", help="HF repo type (default: model)")
parser.add_argument("--revision", help="HF revision to inspect")
@@ -92,12 +94,14 @@ def load_remote_files(
) -> list[str]:
from huggingface_hub import list_repo_files
return sorted(list_repo_files(
repo_id,
repo_type=repo_type,
revision=revision,
token=token,
))
return sorted(
list_repo_files(
repo_id,
repo_type=repo_type,
revision=revision,
token=token,
)
)
def load_remote_model_index(
@@ -211,24 +215,24 @@ def build_result(args: argparse.Namespace) -> dict[str, Any]:
"components_seen": components,
"file_count": len(files),
"file_scan_truncated": truncated,
"files_sample": files[:args.sample_limit],
"files_sample": files[: args.sample_limit],
}
def print_human(result: dict[str, Any]) -> None:
for key in (
"source",
"source_kind",
"repo_type",
"revision",
"token_env",
"source_layout",
"needs_conversion",
"model_index_class",
"model_index_diffusers_version",
"model_index_error",
"file_count",
"file_scan_truncated",
"source",
"source_kind",
"repo_type",
"revision",
"token_env",
"source_layout",
"needs_conversion",
"model_index_class",
"model_index_diffusers_version",
"model_index_error",
"file_count",
"file_scan_truncated",
):
value = result.get(key)
if value is not None:
@@ -18,6 +18,7 @@ import pytest
import torch
from torch.testing import assert_close
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29519")
os.environ.setdefault("DISABLE_SP", "1")
@@ -34,10 +35,15 @@ FASTVIDEO_CONFIG_CLASS = "<FastVideoConfig>" # TODO.
FASTVIDEO_MODEL_MODULE = "fastvideo.models.<bucket>.<module>" # TODO.
FASTVIDEO_MODEL_CLASS = "<FastVideoModel>" # TODO.
OFFICIAL_REF_DIR = Path(os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>"))
LOCAL_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY))
CONVERTED_WEIGHTS_DIR = Path(os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR",
REPO_ROOT / "converted_weights" / FAMILY))
OFFICIAL_REF_DIR = Path(
os.getenv("<FAMILY_UPPER>_OFFICIAL_REF_DIR", REPO_ROOT / "<ReferenceDir>")
)
LOCAL_WEIGHTS_DIR = Path(
os.getenv("<FAMILY_UPPER>_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / FAMILY)
)
CONVERTED_WEIGHTS_DIR = Path(
os.getenv("<FAMILY_UPPER>_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / FAMILY)
)
def _resolve_hf_token() -> str | None:
@@ -93,14 +99,18 @@ def _load_official_model(device: torch.device, dtype: torch.dtype) -> torch.nn.M
model = OfficialClass() # TODO: pass official config kwargs.
state_dict = {} # TODO: load official state dict from LOCAL_WEIGHTS_DIR.
missing, unexpected = model.load_state_dict(state_dict, strict=True)
assert not missing and not unexpected, (f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
assert not missing and not unexpected, (
f"official load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
)
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
"""Load the FastVideo component with the same tensor content."""
if not CONVERTED_WEIGHTS_DIR.exists() and not LOCAL_WEIGHTS_DIR.exists():
pytest.skip(f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}")
pytest.skip(
f"No FastVideo loadable weights: {CONVERTED_WEIGHTS_DIR} or {LOCAL_WEIGHTS_DIR}"
)
# TODO: replace with the bucket-specific FastVideo config/class/loader.
# DiT examples:
@@ -117,7 +127,8 @@ def _load_fastvideo_model(device: torch.device, dtype: torch.dtype) -> torch.nn.
state_dict = {} # TODO: load converted or directly mapped state dict.
missing, unexpected = model.load_state_dict(state_dict, strict=True)
assert not missing and not unexpected, (
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}")
f"FastVideo load mismatch missing={missing[:5]} unexpected={unexpected[:5]}"
)
return model.to(device=device, dtype=dtype).eval()
@@ -176,9 +187,11 @@ def test_component_parity():
assert official_out.shape == fastvideo_out.shape
diff = (official_out - fastvideo_out).abs()
print(f"official abs_mean={official_out.abs().mean().item():.6f} "
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
print(
f"official abs_mean={official_out.abs().mean().item():.6f} "
f"fastvideo abs_mean={fastvideo_out.abs().mean().item():.6f} "
f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}"
)
# TODO: pick tolerance by scope:
# - single block / same kernel: 1e-4
@@ -27,6 +27,7 @@ try:
except ImportError: # pragma: no cover - optional local conversion dependency
snapshot_download = None
# TODO: fill with authoritative component prefixes for monolithic checkpoints.
# Example: {"model.model.": "transformer", "pretransform.model.": "vae"}
COMPONENT_PREFIXES: dict[str, str] = {}
@@ -46,7 +47,10 @@ SKIP_PATTERNS: tuple[str, ...] = ()
def _hf_token() -> str | None:
return (os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN") or os.environ.get("HF_API_KEY"))
return (
os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN")
or os.environ.get("HF_API_KEY")
)
def resolve_src(src: str, revision: str | None) -> Path:
@@ -91,10 +95,11 @@ def apply_mapping(key: str) -> str | None:
return key
def split_monolithic(state: dict[str, torch.Tensor], ) -> dict[str, OrderedDict[str, torch.Tensor]]:
def split_monolithic(
state: dict[str, torch.Tensor],
) -> dict[str, OrderedDict[str, torch.Tensor]]:
components: dict[str, OrderedDict[str, torch.Tensor]] = {
name: OrderedDict()
for name in set(COMPONENT_PREFIXES.values())
name: OrderedDict() for name in set(COMPONENT_PREFIXES.values())
}
intentionally_skipped: list[str] = []
unowned: list[str] = []
@@ -112,8 +117,10 @@ def split_monolithic(state: dict[str, torch.Tensor], ) -> dict[str, OrderedDict[
unowned.append(key)
if unowned:
sample = ", ".join(unowned[:10])
raise ValueError(f"Unowned monolithic keys: {len(unowned)}. "
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}")
raise ValueError(
f"Unowned monolithic keys: {len(unowned)}. "
f"Add COMPONENT_PREFIXES or SKIP_PATTERNS entries. Sample: {sample}"
)
if intentionally_skipped:
print(f"Intentionally skipped {len(intentionally_skipped)} keys")
return {name: weights for name, weights in components.items() if weights}
@@ -136,12 +143,8 @@ def build_component_configs(_src_dir: Path) -> dict[str, dict[str, Any]]:
# TODO: emit config content accepted by FastVideo loaders. Most components use
# config.json; schedulers use scheduler_config.json.
return {
"transformer": {
"_class_name": "<FastVideoTransformerClass>"
},
"vae": {
"_class_name": "<FastVideoVAEClass>"
},
"transformer": {"_class_name": "<FastVideoTransformerClass>"},
"vae": {"_class_name": "<FastVideoVAEClass>"},
}
@@ -174,13 +177,19 @@ def build_model_index(
}
if revision:
index["_fastvideo_converted_revision"] = revision
return {key: value for key, value in index.items() if key.startswith("_") or key in available_components}
return {
key: value
for key, value in index.items()
if key.startswith("_") or key in available_components
}
def validate_component_configs(configs: dict[str, dict[str, Any]]) -> None:
# TODO: instantiate each FastVideo config and call update_model_arch(...) or
# update_model_config(...) with this JSON so unknown emitted keys fail here.
placeholder_configs = [name for name, config in configs.items() if "<" in json.dumps(config)]
placeholder_configs = [
name for name, config in configs.items() if "<" in json.dumps(config)
]
if placeholder_configs:
raise ValueError(f"Replace config placeholders for: {placeholder_configs}")
@@ -192,7 +201,9 @@ def verify_conversion(
del dst_dir, components
# TODO: load each emitted stateful component through its production loader and
# assert strict load, or document exact allowed missing/unexpected keys.
raise NotImplementedError("Implement production config validation and strict-load checks")
raise NotImplementedError(
"Implement production config validation and strict-load checks"
)
def write_component(
@@ -205,7 +216,9 @@ def write_component(
if component_dir.exists() and any(component_dir.iterdir()):
shutil.rmtree(component_dir)
component_dir.mkdir(parents=True, exist_ok=True)
save_file(dict(state), str(component_dir / "diffusion_pytorch_model.safetensors"))
save_file(
dict(state), str(component_dir / "diffusion_pytorch_model.safetensors")
)
if config is not None:
config_path = component_dir / config_filename(name)
with config_path.open("w", encoding="utf-8") as f:
@@ -248,7 +261,9 @@ def convert(
if layout in {"monolithic", "raw_official"}:
# TODO: replace model.safetensors with the official monolithic file name.
components = split_monolithic(load_checkpoint(default_monolithic_checkpoint(src_path)))
components = split_monolithic(
load_checkpoint(default_monolithic_checkpoint(src_path))
)
elif layout in {"separate_components", "mixed"}:
if not src_path.is_dir():
raise ValueError(f"{layout} layout requires a source directory: {src_path}")
@@ -256,7 +271,9 @@ def convert(
else:
raise ValueError(f"Unsupported template layout: {layout}")
copied = (copy_passthrough(src_path, dst_dir) if src_path.is_dir() else [])
copied = (
copy_passthrough(src_path, dst_dir) if src_path.is_dir() else []
)
configs = build_component_configs(src_path if src_path.is_dir() else src_path.parent)
validate_component_configs(configs)
for name, state in components.items():
@@ -272,7 +289,9 @@ def convert(
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--src", required=True, help="HF repo id, local dir, or checkpoint path")
parser.add_argument(
"--src", required=True, help="HF repo id, local dir, or checkpoint path"
)
parser.add_argument("--revision", help="HF branch, tag, or commit for repo sources")
parser.add_argument(
"--dst",
@@ -24,8 +24,8 @@ from typing import Any
import torch
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
FAMILY: str = "<family>" # e.g. "magi_human", "ltx2", "wan"
COMPONENT: str = "<component>" # e.g. "dit", "vae", "encoder"
DRILL_LAYER_ENV: str = "<FAMILY>_DEBUG_DRILL_LAYER"
HYPOTHESIS_ENV: str = "<FAMILY>_DEBUG_PATCH_<HYPOTHESIS>"
REL_THRESHOLD: float = 0.005 # 0.5% abs_mean drift flags a block as divergent
@@ -94,7 +94,6 @@ def _attach_block_hooks(
handles: list[Any] = []
def _hook(name: str):
def fn(_module, _inputs, outputs):
t = outputs[0] if isinstance(outputs, tuple) else outputs
if not torch.is_tensor(t):
@@ -102,7 +101,6 @@ def _attach_block_hooks(
log.append({"side": label, **_stat(name, t)})
if tensors is not None:
tensors[name] = t.detach().float().cpu()
return fn
def _pre_hook(name: str):
@@ -116,7 +114,6 @@ def _attach_block_hooks(
log.append({"side": label, **_stat(key, t)})
if tensors is not None:
tensors[key] = t.detach().float().cpu()
return fn
# TODO: adapt attribute paths to your model. Remove adapter block if absent.
@@ -134,21 +131,43 @@ def _attach_block_hooks(
# magi-human uses: attention, mlp.pre_norm, mlp.up_gate_proj,
# mlp.down_proj (pre+post), mlp, attn_post_norm, mlp_post_norm.
if hasattr(layer, "attention"):
handles.append(layer.attention.register_forward_hook(_hook(f"{tag}.attention")))
handles.append(
layer.attention.register_forward_hook(_hook(f"{tag}.attention"))
)
if hasattr(layer, "mlp"):
mlp = layer.mlp
if hasattr(mlp, "pre_norm"):
handles.append(mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm")))
handles.append(
mlp.pre_norm.register_forward_hook(_hook(f"{tag}.mlp.pre_norm"))
)
if hasattr(mlp, "up_gate_proj"):
handles.append(mlp.up_gate_proj.register_forward_hook(_hook(f"{tag}.mlp.up_gate_proj")))
handles.append(
mlp.up_gate_proj.register_forward_hook(
_hook(f"{tag}.mlp.up_gate_proj")
)
)
if hasattr(mlp, "down_proj"):
handles.append(mlp.down_proj.register_forward_pre_hook(_pre_hook(f"{tag}.mlp.down_proj")))
handles.append(mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj")))
handles.append(
mlp.down_proj.register_forward_pre_hook(
_pre_hook(f"{tag}.mlp.down_proj")
)
)
handles.append(
mlp.down_proj.register_forward_hook(_hook(f"{tag}.mlp.down_proj"))
)
handles.append(mlp.register_forward_hook(_hook(f"{tag}.mlp")))
if hasattr(layer, "attn_post_norm"):
handles.append(layer.attn_post_norm.register_forward_hook(_hook(f"{tag}.attn_post_norm")))
handles.append(
layer.attn_post_norm.register_forward_hook(
_hook(f"{tag}.attn_post_norm")
)
)
if hasattr(layer, "mlp_post_norm"):
handles.append(layer.mlp_post_norm.register_forward_hook(_hook(f"{tag}.mlp_post_norm")))
handles.append(
layer.mlp_post_norm.register_forward_hook(
_hook(f"{tag}.mlp_post_norm")
)
)
return handles
@@ -174,9 +193,11 @@ def _write_log(entries: list[dict], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w") as f:
for e in entries:
f.write(f"{e['name']} {e['shape']} "
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
f"{e['min']:.6f} {e['max']:.6f}\n")
f.write(
f"{e['name']} {e['shape']} "
f"{e['abs_mean']:.8f} {e['sum']:.4f} "
f"{e['min']:.6f} {e['max']:.6f}\n"
)
def _sort_key(name: str, drill_layer: int) -> tuple:
@@ -184,14 +205,9 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
return (0, "")
if name.startswith(f"L{drill_layer:02d}."):
sub_order = {
"attention": 0,
"attn_post_norm": 1,
"mlp.pre_norm": 2,
"mlp.up_gate_proj": 3,
"mlp.down_proj<in>": 4,
"mlp.down_proj": 5,
"mlp": 6,
"mlp_post_norm": 7,
"attention": 0, "attn_post_norm": 1, "mlp.pre_norm": 2,
"mlp.up_gate_proj": 3, "mlp.down_proj<in>": 4,
"mlp.down_proj": 5, "mlp": 6, "mlp_post_norm": 7,
}.get(name.split(".", 1)[1], 9)
return (1, f"block[{drill_layer:02d}]", sub_order)
if name.startswith("block["):
@@ -200,8 +216,10 @@ def _sort_key(name: str, drill_layer: int) -> tuple:
def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
hdr = (f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}")
hdr = (
f"{'name':<18} {'up_shape':<22} {'up_absmean':>12} {'fv_absmean':>12} "
f"{'absmean_diff':>14} {'rel%':>8} {'up_sum':>14} {'fv_sum':>14} {'sum_diff':>12}"
)
print(f"\n{hdr}\n{'-' * len(hdr)}")
first_div: int | None = None
for name in sorted(by_name.keys(), key=lambda n: _sort_key(n, drill_layer)):
@@ -217,9 +235,11 @@ def _print_table(by_name: dict[str, dict], drill_layer: int) -> int | None:
flag = " <<< DIVERGE"
if first_div is None:
first_div = int(name[len("block["):-1])
print(f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}")
print(
f"{name:<18} {str(up['shape']):<22} {up['abs_mean']:>12.6f} "
f"{fv['abs_mean']:>12.6f} {am_diff:>14.6f} {am_rel * 100:>7.3f}% "
f"{up['sum']:>14.4f} {fv['sum']:>14.4f} {sum_diff:>12.4f}{flag}"
)
return first_div
@@ -235,8 +255,10 @@ def _print_elementwise(up_t: dict[str, torch.Tensor], fv_t: dict[str, torch.Tens
continue
diff = (a - b).abs()
rel = (diff.mean().item() / max(a.abs().mean().item(), 1e-9)) * 100
print(f"{name:<30} {str(tuple(a.shape)):<22} "
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%")
print(
f"{name:<30} {str(tuple(a.shape)):<22} "
f"{diff.max().item():>12.6f} {diff.mean().item():>12.6f} {rel:>9.4f}%"
)
def main() -> None:
@@ -43,10 +43,12 @@ def _add_official_to_path() -> Path:
def _log_tensor_stats(label: str, tensor: torch.Tensor) -> None:
value = tensor.detach().float()
print(f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
f"dtype={tensor.dtype} device={tensor.device} "
f"min={value.min().item():.6f} max={value.max().item():.6f} "
f"mean={value.mean().item():.6f} std={value.std().item():.6f}")
print(
f"[{_MODEL_FAMILY} PIPELINE] {label}: shape={tuple(tensor.shape)} "
f"dtype={tensor.dtype} device={tensor.device} "
f"min={value.min().item():.6f} max={value.max().item():.6f} "
f"mean={value.mean().item():.6f} std={value.std().item():.6f}"
)
def _extract_tensor(output: Any, key: str) -> torch.Tensor:
@@ -71,8 +73,10 @@ def _run_official_pipeline(
device: torch.device,
) -> Any:
del official_path, params, device
pytest.skip("TODO: import the official pipeline/factory, load official weights, "
"run with params, and return the comparison target.")
pytest.skip(
"TODO: import the official pipeline/factory, load official weights, "
"run with params, and return the comparison target."
)
def _run_fastvideo_pipeline(model_path: Path, params: dict[str, Any]) -> Any:
@@ -142,6 +146,8 @@ def test_todo_model_family_pipeline_official_parity() -> None:
assert official_tensor.shape == fastvideo_tensor.shape
diff = (official_tensor - fastvideo_tensor).abs()
print(f"diff max={diff.max().item():.6f} "
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}")
print(
f"diff max={diff.max().item():.6f} "
f"mean={diff.mean().item():.6f} median={diff.median().item():.6f}"
)
assert_close(fastvideo_tensor, official_tensor, atol=1e-2, rtol=1e-2)
-149
View File
@@ -1,149 +0,0 @@
name: macOS MLX Smoke
on:
pull_request:
branches: [main]
paths:
- ".github/workflows/ci-macos-mlx.yml"
- "fastvideo/mlx_runtime/**"
- "fastvideo/tests/mlx/**"
- "fastvideo/tests/platforms/test_mps_vsa_error.py"
- "fastvideo/platforms/mps.py"
- "fastvideo/platforms/__init__.py"
- "fastvideo/__init__.py"
- "examples/inference/basic/mlx_*.py"
- "fastvideo/benchmarks/mlx_*.py"
- "pyproject.toml"
workflow_dispatch:
permissions:
contents: read
concurrency:
group: macos-mlx-${{ github.ref }}
cancel-in-progress: true
jobs:
mlx-smoke:
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
runs-on: macos-15
timeout-minutes: 25
env:
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
TOKENIZERS_PARALLELISM: "false"
MASTER_ADDR: localhost
MASTER_PORT: "29513"
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
cache: pip
- uses: astral-sh/setup-uv@v3
- name: Install lightweight MLX smoke dependencies
run: |
uv pip install --system \
--index-url https://download.pytorch.org/whl/cpu \
torch==2.11.0 torchvision torchaudio
uv pip install --system \
pytest numpy scipy pillow imageio einops cloudpickle filelock \
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru mlx \
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
- name: Show Apple runtime
run: |
python - <<'PY'
import platform
import mlx.core as mx
import torch
print("machine:", platform.machine())
print("processor:", platform.processor())
print("mlx default device:", mx.default_device())
memory_size = mx.metal.device_info().get("memory_size") if mx.metal.is_available() else "metal unavailable"
print("mlx memory_size:", memory_size)
print("torch:", torch.__version__)
print("torch mps available:", torch.backends.mps.is_available())
PY
- name: Run MLX smoke tests
run: |
python -m pytest \
fastvideo/tests/mlx/test_dmd_sampling.py \
fastvideo/tests/mlx/test_memory_limits.py \
fastvideo/tests/mlx/test_quant_capability.py \
fastvideo/tests/mlx/test_mlx_dit_parity.py \
fastvideo/tests/mlx/test_mlx_compile_parity.py \
fastvideo/tests/mlx/test_mlx_checkpoint.py \
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
fastvideo/tests/mlx/test_taehv_decode.py \
fastvideo/tests/mlx/test_frame_upsample.py \
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
fastvideo/tests/mlx/test_mlx_refine.py \
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
fastvideo/tests/mlx/test_wan22_sample.py \
fastvideo/tests/mlx/test_windowed_attention.py \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
fastvideo/tests/platforms/test_mps_vsa_error.py \
-q
# Same tests on MLX's CPU backend. Hosted macOS runners are scarce and
# slower to schedule; this Linux job gives fast PR signal on the identical
# graph (the parity tests were designed to be backend-agnostic), while the
# macOS job above stays the source of truth for Metal behavior.
mlx-smoke-linux-cpu:
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
runs-on: ubuntu-latest
timeout-minutes: 20
env:
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
TOKENIZERS_PARALLELISM: "false"
MASTER_ADDR: localhost
MASTER_PORT: "29513"
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
cache: pip
- uses: astral-sh/setup-uv@v3
- name: Install lightweight MLX smoke dependencies (CPU backend)
run: |
uv pip install --system \
--index-url https://download.pytorch.org/whl/cpu \
torch==2.11.0 torchvision torchaudio
uv pip install --system \
pytest numpy scipy pillow imageio einops cloudpickle filelock \
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru "mlx[cpu]" \
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
- name: Run MLX smoke tests (CPU backend)
run: |
python -m pytest \
fastvideo/tests/mlx/test_dmd_sampling.py \
fastvideo/tests/mlx/test_memory_limits.py \
fastvideo/tests/mlx/test_quant_capability.py \
fastvideo/tests/mlx/test_mlx_dit_parity.py \
fastvideo/tests/mlx/test_mlx_compile_parity.py \
fastvideo/tests/mlx/test_mlx_checkpoint.py \
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
fastvideo/tests/mlx/test_taehv_decode.py \
fastvideo/tests/mlx/test_frame_upsample.py \
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
fastvideo/tests/mlx/test_mlx_refine.py \
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
fastvideo/tests/mlx/test_wan22_sample.py \
fastvideo/tests/mlx/test_windowed_attention.py \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
fastvideo/tests/platforms/test_mps_vsa_error.py \
-q
-2
View File
@@ -6,7 +6,6 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
@@ -17,7 +16,6 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
-1
View File
@@ -23,7 +23,6 @@ Miniconda3-latest-Linux-x86_64.sh
*validation/
data/
outputs/
outputs_audio/
outputs_video
checkpoints/
sbatch.sh
-1
View File
@@ -22,7 +22,6 @@ repos:
hooks:
- id: yapf
args: [--in-place, --verbose]
language_version: python3.12
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.11.12
-6
View File
@@ -9,7 +9,6 @@
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
## NEWS
- `2026/08/19`: FastVideo now supports MLX on Apple Silicon with [FastMetal-QAD](https://huggingface.co/collections/FastVideo/fastmetal), a family of 1.3B, 5B, and 14B models optimized for Mac—follow the [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/) and read the [Blog](https://haoailab.com/blogs/fastmetal/).
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
@@ -63,11 +62,6 @@ UV_TORCH_BACKEND=cu126 uv pip install fastvideo
Use `UV_TORCH_BACKEND=cu130` on CUDA 13. Apple silicon users should follow the
[MPS installation guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
> **On an Apple Silicon Mac?** FastVideo runs FastWan text-to-video natively
> through an MLX runtime — a 5-second 480p clip generated locally, no cloud,
> no discrete GPU. Install with `uv pip install -e '.[mlx]'` and follow the
> [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
> **On an NVIDIA DGX Spark (GB10 / ARM64 + CUDA 13)?** There's no prebuilt ARM wheel for the FastVideo CUDA kernel, so it's an editable from-source install (`UV_TORCH_BACKEND=cu130 uv pip install -e .`, which compiles that kernel for you) rather than `UV_TORCH_BACKEND=cu130 uv pip install fastvideo`. A compatible prebuilt ARM64 FlashAttention wheel is available separately. Follow the [DGX Spark install guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/spark/).
@@ -3,6 +3,7 @@ from __future__ import annotations
import sys
from pathlib import Path
TESTS_DIR = Path(__file__).resolve().parent
DREAMVERSE_PACKAGE_DIR = TESTS_DIR.parent
DREAMVERSE_APP_DIR = DREAMVERSE_PACKAGE_DIR.parent
@@ -5,6 +5,7 @@ from pathlib import Path
import pytest
SERVER_DIR = Path(__file__).resolve().parents[1]
@@ -52,7 +53,9 @@ def test_config_defaults_to_cerebras_with_parallel_groq_fallback_stage(monkeypat
module = _load_config_module()
assert module.PROMPT_PROVIDER == "cerebras"
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
("cerebras", "groq"),
)
assert module.PROMPT_PROVIDER_PRIORITY == (
"cerebras",
"groq",
@@ -83,7 +86,9 @@ def test_config_ignores_legacy_groq_primary_override(monkeypatch):
module = _load_config_module()
assert module.PROMPT_PROVIDER == "cerebras"
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (("cerebras", "groq"), )
assert module.PROMPT_PROVIDER_RUNTIME_STAGES == (
("cerebras", "groq"),
)
assert module.PROMPT_PROVIDER_PRIORITY == (
"cerebras",
"groq",
@@ -101,17 +106,24 @@ def test_config_uses_local_overlay_paths_when_devtools_enabled(monkeypatch, tmp_
assert module.DEVTOOLS_ENABLED is True
assert module.FRONTEND_ROOT.as_posix().endswith("apps/dreamverse/web")
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith("dreamverse/prompts.local/next_segment_system_prompt.md")
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_PATH.endswith(
"dreamverse/prompts.local/next_segment_system_prompt.md"
)
assert module.PROMPT_ENHANCE_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
"dreamverse/prompts/next_segment_system_prompt.md")
"dreamverse/prompts/next_segment_system_prompt.md"
)
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_PATH.endswith(
"dreamverse/prompts.local/rewrite_user_system_prompt.md")
"dreamverse/prompts.local/rewrite_user_system_prompt.md"
)
assert module.PROMPT_REWRITE_USER_SYSTEM_PROMPT_FALLBACK_PATH.endswith(
"dreamverse/prompts/rewrite_user_system_prompt.md")
"dreamverse/prompts/rewrite_user_system_prompt.md"
)
assert module.CURATED_PRESETS_FILE_PATH.endswith(
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json")
"apps/dreamverse/web/prompts.local/selected_ltx2_continuation_story_presets.json"
)
assert module.CURATED_PRESETS_FALLBACK_FILE_PATH.endswith(
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json")
"apps/dreamverse/web/prompts/selected_ltx2_continuation_story_presets.json"
)
assert module.FRONTEND_STATIC_DIR_CANDIDATES[:2] == (
str(module.FRONTEND_ROOT / "out"),
str(module.FRONTEND_ROOT / "dist"),
@@ -9,7 +9,6 @@ from fastapi.testclient import TestClient
import fastvideo.entrypoints.streaming as streaming_entrypoints
import pytest
def _install_stack03_import_stubs(monkeypatch):
"""Keep entrypoint tests focused while later-stack runtime modules are absent."""
if not hasattr(streaming_entrypoints, "build_health_router"):
@@ -18,7 +17,6 @@ def _install_stack03_import_stubs(monkeypatch):
gpu_pool_stub = types.ModuleType("dreamverse.gpu_pool")
class GPUPool:
def __init__(self, _gpu_ids):
pass
@@ -51,7 +49,6 @@ def _install_stack03_import_stubs(monkeypatch):
controller_stub = types.ModuleType("dreamverse.session.controller")
class SessionController:
def __init__(self, **_kwargs):
pass
@@ -79,11 +76,13 @@ def _run_cli(module, monkeypatch, argv: list[str]) -> list[dict[str, object]]:
uvicorn_stub = types.ModuleType("uvicorn")
def run(app, host: str, port: int) -> None:
calls.append({
"app": app,
"host": host,
"port": port,
})
calls.append(
{
"app": app,
"host": host,
"port": port,
}
)
uvicorn_stub.run = run
monkeypatch.setitem(sys.modules, "uvicorn", uvicorn_stub)
@@ -100,11 +99,13 @@ def test_server_cli_defaults_to_local_web_port(monkeypatch):
server_main = _import_server_main(monkeypatch)
calls = _run_cli(server_main, monkeypatch, ["dreamverse-server"])
assert calls == [{
"app": server_main.app,
"host": "0.0.0.0",
"port": 8009,
}]
assert calls == [
{
"app": server_main.app,
"host": "0.0.0.0",
"port": 8009,
}
]
def test_server_cli_allows_explicit_host_and_port(monkeypatch):
@@ -115,11 +116,13 @@ def test_server_cli_allows_explicit_host_and_port(monkeypatch):
["dreamverse-server", "--host", "127.0.0.1", "--port", "8123"],
)
assert calls == [{
"app": server_main.app,
"host": "127.0.0.1",
"port": 8123,
}]
assert calls == [
{
"app": server_main.app,
"host": "127.0.0.1",
"port": 8123,
}
]
def test_server_does_not_expose_backend_source_as_static_assets(monkeypatch):
@@ -139,11 +142,13 @@ def test_mock_server_cli_defaults_to_local_web_port(monkeypatch):
["dreamverse-mock-server"],
)
assert calls == [{
"app": mock_server.app,
"host": "0.0.0.0",
"port": 8009,
}]
assert calls == [
{
"app": mock_server.app,
"host": "0.0.0.0",
"port": 8009,
}
]
def test_mock_server_cli_updates_latency(monkeypatch):
@@ -156,11 +161,13 @@ def test_mock_server_cli_updates_latency(monkeypatch):
["dreamverse-mock-server", "--latency", "321", "--port", "8111"],
)
assert calls == [{
"app": mock_server.app,
"host": "0.0.0.0",
"port": 8111,
}]
assert calls == [
{
"app": mock_server.app,
"host": "0.0.0.0",
"port": 8111,
}
]
assert mock_server.LATENCY_MS == 321
finally:
mock_server.LATENCY_MS = old_latency_ms
@@ -7,6 +7,7 @@ from types import SimpleNamespace
import pytest
import dreamverse.gpu_pool as gpu_pool
@@ -84,7 +85,9 @@ def test_send_command_raises_on_worker_death():
cmd_q = ctx.Queue()
resp_q = ctx.Queue()
proc = ctx.Process(target=_child_consume_and_exit, args=(cmd_q, resp_q))
proc = ctx.Process(
target=_child_consume_and_exit, args=(cmd_q, resp_q)
)
proc.start()
# Wait for the spawn child to fully boot. Allow generous time —
@@ -7,7 +7,7 @@ ALLOWED_PREFIXES = (
"fastvideo.entrypoints.video_generator",
"fastvideo.configs",
)
ALLOWED_EXACT = ("fastvideo", )
ALLOWED_EXACT = ("fastvideo",)
FORBIDDEN_PREFIXES = (
"fastvideo.pipelines",
"fastvideo.models",
@@ -38,13 +38,19 @@ def test_dreamverse_server_imports_only_public_fastvideo_surfaces() -> None:
except SyntaxError as task_exc:
raise AssertionError(f"Failed to parse {path}") from task_exc
for node in ast.walk(tree):
names = ([a.name for a in node.names] if isinstance(node, ast.Import) else
[node.module] if isinstance(node, ast.ImportFrom) and node.module else [])
names = (
[a.name for a in node.names] if isinstance(node, ast.Import)
else [node.module] if isinstance(node, ast.ImportFrom) and node.module
else []
)
for name in names:
if not name:
continue
rel_path = str(path.relative_to(root))
if (name.startswith(FORBIDDEN_PREFIXES) and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS):
if (
name.startswith(FORBIDDEN_PREFIXES)
and (rel_path, name) not in ALLOWED_INTERNAL_IMPORTS
):
bad.append((str(path.relative_to(root)), getattr(node, "lineno", 0), name))
assert bad == [], f"Forbidden internal imports: {bad}"
@@ -6,6 +6,7 @@ import os
from fastapi import WebSocketDisconnect
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
os.environ.setdefault("GROQ_API_KEY", "dummy")
@@ -13,7 +14,6 @@ import dreamverse.mock_server as mock_server
class _FakeWebSocket:
def __init__(self, messages: list[tuple[float, dict[str, object]]]):
self._messages = messages
self._index = 0
@@ -49,34 +49,34 @@ def test_mock_server_matches_current_single5s_protocol():
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
mock_server.LATENCY_MS = 1
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "simple_prompt_1",
"curated_prompts": ["selected prompt"],
"single_clip_mode": True,
"enhancement_enabled": False,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.01,
{
"type": "simple_generate",
"preset_id": "simple_custom_prompt",
"prompt_id": "simple_custom_prompt",
"prompt": "custom prompt",
"enhancement_enabled": True,
"initial_image": None,
},
),
(0.20, {
"type": "leave"
}),
])
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "simple_prompt_1",
"curated_prompts": ["selected prompt"],
"single_clip_mode": True,
"enhancement_enabled": False,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.01,
{
"type": "simple_generate",
"preset_id": "simple_custom_prompt",
"prompt_id": "simple_custom_prompt",
"prompt": "custom prompt",
"enhancement_enabled": True,
"initial_image": None,
},
),
(0.20, {"type": "leave"}),
]
)
asyncio.run(mock_server.websocket_endpoint(ws))
@@ -92,14 +92,24 @@ def test_mock_server_matches_current_single5s_protocol():
assert message_types.count("ltx2_stream_complete") == 2
assert "prompt_sources_blocked" not in message_types
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
gpu_assigned_event = next(payload for payload in ws.sent_json if payload["type"] == "gpu_assigned")
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
gpu_assigned_event = next(
payload for payload in ws.sent_json if payload["type"] == "gpu_assigned"
)
assert gpu_assigned_event["session_timeout"] == mock_server.SESSION_TIMEOUT_SECONDS
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
assert segment_start_events[0]["prompt"] == "selected prompt"
assert segment_start_events[1]["prompt"] == "custom prompt"
step_complete_events = [payload for payload in ws.sent_json if payload["type"] == "step_complete"]
step_complete_events = [
payload
for payload in ws.sent_json
if payload["type"] == "step_complete"
]
assert len(step_complete_events) == 2
assert step_complete_events[0]["latency_ms"] == {
"total": 121.0,
@@ -124,29 +134,29 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
mock_server.LATENCY_MS = 1
mock_server.GENERATION_SEGMENT_CAP = 1
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.02,
{
"type": "rewrite_seed_prompts",
"rewrite_instruction": "start a new rollout",
},
),
(0.20, {
"type": "leave"
}),
])
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.02,
{
"type": "rewrite_seed_prompts",
"rewrite_instruction": "start a new rollout",
},
),
(0.20, {"type": "leave"}),
]
)
asyncio.run(mock_server.websocket_endpoint(ws))
@@ -156,7 +166,11 @@ def test_mock_server_regular_cap_waits_for_rewrite_rollout():
assert "generation_cap_reached" not in message_types
assert "prompt_sources_blocked" not in message_types
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
assert [payload["segment_idx"] for payload in segment_start_events] == [1, 1]
assert segment_start_events[0]["prompt"] == "segment one"
assert segment_start_events[1]["prompt"] == "segment one [start a new rollout]"
@@ -173,40 +187,54 @@ def test_mock_server_rewrite_during_active_segment_restarts_from_first_rewritten
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
mock_server.LATENCY_MS = 100
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one", "segment two"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.02,
{
"type": "rewrite_seed_prompts",
"rewrite_instruction": "restart from rewrite",
},
),
(0.40, {
"type": "leave"
}),
])
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one", "segment two"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(
0.02,
{
"type": "rewrite_seed_prompts",
"rewrite_instruction": "restart from rewrite",
},
),
(0.40, {"type": "leave"}),
]
)
asyncio.run(mock_server.websocket_endpoint(ws))
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
"segment one",
"segment one [restart from rewrite]",
]
assert all(payload["prompt"] != "segment two" for payload in segment_start_events[1:])
reset_events = [payload for payload in ws.sent_json if payload.get("type") == "seed_prompts_reset_applied"]
assert any(payload.get("reason") == "rewrite_during_generation" for payload in reset_events)
assert all(
payload["prompt"] != "segment two"
for payload in segment_start_events[1:]
)
reset_events = [
payload
for payload in ws.sent_json
if payload.get("type") == "seed_prompts_reset_applied"
]
assert any(
payload.get("reason") == "rewrite_during_generation"
for payload in reset_events
)
finally:
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
mock_server.LATENCY_MS = old_latency_ms
@@ -219,24 +247,24 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
mock_server.LATENCY_MS = 1
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "custom_editable",
"preset_label": "Custom rollout",
"curated_prompts": [],
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.20, {
"type": "leave"
}),
])
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "custom_editable",
"preset_label": "Custom rollout",
"curated_prompts": [],
"initial_rollout_prompt": "A moonbase corridor thriller with flooding",
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.20, {"type": "leave"}),
]
)
asyncio.run(mock_server.websocket_endpoint(ws))
@@ -248,9 +276,15 @@ def test_mock_server_supports_initial_custom_rollout_prompt():
assert "ltx2_stream_start" in message_types
assert "prompt_sources_blocked" not in message_types
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
assert segment_start_events
assert segment_start_events[0]["prompt"] == ("A moonbase corridor thriller with flooding [segment 1]")
assert segment_start_events[0]["prompt"] == (
"A moonbase corridor thriller with flooding [segment 1]"
)
finally:
mock_server.MOCK_SEGMENT_BYTES = old_segment_bytes
mock_server.LATENCY_MS = old_latency_ms
@@ -263,37 +297,35 @@ def test_mock_server_can_start_new_project_without_reconnecting():
mock_server.MOCK_SEGMENT_BYTES = b"mock-fmp4-bytes"
mock_server.LATENCY_MS = 40
ws = _FakeWebSocket([
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.02, {
"type": "end_project_keep_session"
}),
(
0.20,
{
"type": "project_init_v1",
"preset_id": "test_preset_2",
"preset_label": "Test Preset 2",
"curated_prompts": ["segment two"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.40, {
"type": "leave"
}),
])
ws = _FakeWebSocket(
[
(
0.0,
{
"type": "session_init_v2",
"preset_id": "test_preset",
"curated_prompts": ["segment one"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.02, {"type": "end_project_keep_session"}),
(
0.20,
{
"type": "project_init_v1",
"preset_id": "test_preset_2",
"preset_label": "Test Preset 2",
"curated_prompts": ["segment two"],
"enhancement_enabled": True,
"auto_extension_enabled": False,
"loop_generation_enabled": False,
},
),
(0.40, {"type": "leave"}),
]
)
asyncio.run(mock_server.websocket_endpoint(ws))
@@ -304,11 +336,16 @@ def test_mock_server_can_start_new_project_without_reconnecting():
project_idle_index = message_types.index("project_idle")
stream_start_indexes = [
index for index, message_type in enumerate(message_types) if message_type == "ltx2_stream_start"
index for index, message_type in enumerate(message_types)
if message_type == "ltx2_stream_start"
]
assert stream_start_indexes[0] < project_idle_index < stream_start_indexes[1]
segment_start_events = [payload for payload in ws.sent_json if payload["type"] == "ltx2_segment_start"]
segment_start_events = [
payload
for payload in ws.sent_json
if payload["type"] == "ltx2_segment_start"
]
assert [payload["prompt"] for payload in segment_start_events[:2]] == [
"segment one",
"segment two",
@@ -6,6 +6,7 @@ import os
import re
import time
os.environ.setdefault("CEREBRAS_API_KEY", "dummy")
os.environ.setdefault("GROQ_API_KEY", "dummy")
@@ -21,7 +22,6 @@ from dreamverse.prompt_enhancer import (
class _FakeResponse:
def __init__(self, payload: dict):
self._payload = payload
@@ -30,7 +30,6 @@ class _FakeResponse:
class _FakeSyncCompletions:
def __init__(self, payload: dict):
self._payload = payload
@@ -39,7 +38,6 @@ class _FakeSyncCompletions:
class _FakeSyncClient:
def __init__(self, payload: dict):
self.chat = type(
"_FakeChat",
@@ -49,7 +47,6 @@ class _FakeSyncClient:
class _DelayedSyncCompletions:
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
self._payload = payload
self._delay_s = delay_s
@@ -64,26 +61,29 @@ class _DelayedSyncCompletions:
class _DelayedSyncClient:
def __init__(self, payload: dict, delay_s: float = 0.0, exc: Exception | None = None):
self.chat = type(
"_FakeChat",
(),
{"completions": _DelayedSyncCompletions(
payload,
delay_s=delay_s,
exc=exc,
)},
{
"completions": _DelayedSyncCompletions(
payload,
delay_s=delay_s,
exc=exc,
)
},
)()
def _chat_payload_with_content(content: str) -> dict:
return {
"choices": [{
"message": {
"content": content,
"choices": [
{
"message": {
"content": content,
}
}
}]
]
}
@@ -172,7 +172,6 @@ def _build_staged_enhancer(
class _FakeOpenAIClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.chat = type(
@@ -183,7 +182,6 @@ class _FakeOpenAIClient:
class _FakeCerebrasClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.chat = type(
@@ -194,12 +192,16 @@ class _FakeCerebrasClient:
def test_parse_json_response_accepts_fenced_json_with_prose():
parsed = _parse_json_response("Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks.")
parsed = _parse_json_response(
"Here is the rewrite:\n```json\n{\"segment_prompts\":[\"A\",\"B\"]}\n```\nThanks."
)
assert parsed == {"segment_prompts": ["A", "B"]}
def test_parse_json_response_extracts_first_embedded_object():
parsed = _parse_json_response("Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)")
parsed = _parse_json_response(
"Model output:\n{\"segment_prompts\":[\"A\",\"B\"]}\n(complete)"
)
assert parsed == {"segment_prompts": ["A", "B"]}
@@ -266,12 +268,16 @@ def test_build_client_supports_groq_provider(monkeypatch):
def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
))
)
)
assert result.fallback_used is False
assert result.error is None
assert result.rollout_id == "preset_a"
@@ -280,12 +286,15 @@ def test_rewrite_prompt_sequence_accepts_segment_prompts_output():
def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
enhancer = _build_test_enhancer(_chat_payload_with_content('{"rewritten_prompts":["A","B"]}'))
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"rewritten_prompts":["A","B"]}')
)
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
))
)
)
assert result.fallback_used is False
assert result.error is None
assert result.rollout_id == "current_rollout"
@@ -294,14 +303,19 @@ def test_rewrite_prompt_sequence_accepts_legacy_rewritten_prompts_output():
def test_rewrite_prompt_sequence_accepts_segment_dicts_without_top_level_rollout_metadata():
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segments":[{"prompt":"A"},{"text":"B"}]}'))
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"segments":[{"prompt":"A"},{"text":"B"}]}'
)
)
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
preset_id="preset_a",
preset_label="Preset A",
rewrite_instruction="make it cinematic",
))
)
)
assert result.fallback_used is False
assert result.error is None
assert result.rollout_id == "preset_a"
@@ -315,12 +329,14 @@ def test_rewrite_prompt_sequence_accepts_numbered_prose_output():
"The user is asking for a cinematic rewrite.\n\n"
"1. A dog bounds across the moon's dusty surface, kicking up silver regolith as it chases a rabbit beneath the black sky.\n"
"2. The rabbit darts around a crater rim while the dog lunges after it, Earth glowing blue in the distance.\n"
))
)
)
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
))
)
)
assert result.fallback_used is False
assert result.error is None
assert result.rollout_id == "current_rollout"
@@ -339,10 +355,12 @@ def test_enhance_prompt_prefers_cerebras_before_groq_fallback():
groq_delay_s=0.01,
)
result = asyncio.run(enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
))
result = asyncio.run(
enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
)
)
assert result.fallback_used is False
assert result.error is None
@@ -363,10 +381,12 @@ def test_enhance_prompt_uses_groq_when_cerebras_fails():
groq_delay_s=0.01,
)
result = asyncio.run(enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
))
result = asyncio.run(
enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
)
)
assert result.fallback_used is False
assert result.error is None
@@ -388,11 +408,13 @@ def test_enhance_prompt_can_use_groq_when_cerebras_times_out():
enhancer.http_timeout_ms = 50
enhancer.default_timeout_ms = 50
result = asyncio.run(enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
timeout_ms=50,
))
result = asyncio.run(
enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
timeout_ms=50,
)
)
assert result.fallback_used is False
assert result.error is None
@@ -412,10 +434,12 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
groq_delay_s=0.08,
)
result = asyncio.run(enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
))
result = asyncio.run(
enhancer.enhance_prompt(
"A rainy alley at night",
mode="single_clip",
)
)
assert result.fallback_used is False
assert result.error is None
@@ -429,12 +453,15 @@ def test_enhance_prompt_can_use_cerebras_when_it_returns_first():
def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
enhancer = _build_test_enhancer(_chat_payload_with_content("I cannot comply with JSON right now."))
enhancer = _build_test_enhancer(
_chat_payload_with_content("I cannot comply with JSON right now.")
)
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
))
)
)
assert result.fallback_used is True
assert "No JSON object found in assistant response." in (result.error or "")
assert result.raw_response_text == "I cannot comply with JSON right now."
@@ -446,7 +473,9 @@ def test_rewrite_prompt_sequence_keeps_raw_output_on_parse_error():
def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
enhancer = _build_test_enhancer(
_chat_payload_with_content(
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'))
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
)
)
captured = {
"body": None,
"timeout_seconds": None,
@@ -457,7 +486,8 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
captured["timeout_seconds"] = timeout_seconds
return (
_chat_payload_with_content(
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'),
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}'
),
'{"id":"rewritten_rollout","label":"Rewritten Rollout","segment_prompts":["A","B"]}',
)
@@ -472,7 +502,8 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
rewrite_model="gpt-test",
rewrite_temperature=0.2,
timeout_ms=800,
))
)
)
assert result.fallback_used is False
assert captured["body"]["messages"][0] == {
@@ -481,12 +512,12 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
}
assert captured["body"]["messages"][1]["role"] == "user"
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
"mode":
"edit_existing_rollout",
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
"Keep count and ordering identical."),
"user_instruction":
"make it cinematic",
"mode": "edit_existing_rollout",
"request": (
"Rewrite all segment prompts with improved continuity and cinematic detail. "
"Keep count and ordering identical."
),
"user_instruction": "make it cinematic",
"current_rollout": {
"id": "preset_a",
"label": "Preset A",
@@ -497,8 +528,11 @@ def test_rewrite_prompt_sequence_uses_current_rollout_payload_shape():
def test_rewrite_prompt_sequence_supports_new_rollout_mode():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}'))
_chat_payload_with_content(
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}'
)
)
captured = {
"body": None,
}
@@ -507,8 +541,10 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
del timeout_seconds
captured["body"] = body
return (
_chat_payload_with_content('{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}'),
_chat_payload_with_content(
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}'
),
'{"id":"custom_editable","label":"Custom rollout","segment_prompts":['
'"A","B","C","D","E","F"]}',
)
@@ -524,29 +560,30 @@ def test_rewrite_prompt_sequence_supports_new_rollout_mode():
rewrite_model="gpt-test",
rewrite_temperature=0.2,
timeout_ms=800,
))
)
)
assert result.fallback_used is False
assert result.prompts == ["A", "B", "C", "D", "E", "F"]
assert prompt_enhancer_module.json.loads(captured["body"]["messages"][1]["content"]) == {
"mode":
"new_rollout",
"request": ("Rewrite all segment prompts with improved continuity and cinematic detail. "
"Keep count and ordering identical."),
"user_instruction":
"A moonbase corridor thriller with flooding and red alarms",
"desired_segment_count":
6,
"rollout_id_hint":
"custom_editable",
"rollout_label_hint":
"Custom rollout",
"mode": "new_rollout",
"request": (
"Rewrite all segment prompts with improved continuity and cinematic detail. "
"Keep count and ordering identical."
),
"user_instruction": "A moonbase corridor thriller with flooding and red alarms",
"desired_segment_count": 6,
"rollout_id_hint": "custom_editable",
"rollout_label_hint": "Custom rollout",
}
def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
enhancer.rewrite_all_system_prompt = "shared system prompt"
captured = {
"body": None,
@@ -556,7 +593,9 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
del timeout_seconds
captured["body"] = body
return (
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'),
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
),
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}',
)
@@ -570,7 +609,8 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
rewrite_instruction="make it cinematic",
rewrite_model="gpt-test",
system_prompt_override="session specific system prompt",
))
)
)
assert result.fallback_used is False
assert captured["body"]["messages"][0] == {
@@ -581,7 +621,10 @@ def test_rewrite_prompt_sequence_uses_session_override_system_prompt():
def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
@@ -592,17 +635,24 @@ def test_resolve_rewrite_new_rollout_system_prompt_uses_dedicated_prompt():
def test_resolve_rewrite_new_rollout_system_prompt_prefers_override():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
enhancer.rewrite_all_system_prompt = "shared rewrite system prompt"
enhancer.rewrite_user_system_prompt = "new rollout rewrite system prompt"
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt("session specific system prompt")
resolved = enhancer.resolve_rewrite_new_rollout_system_prompt(
"session specific system prompt"
)
assert resolved == "session specific system prompt"
def test_generate_auto_prompt_uses_selected_model():
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Auto next"}'))
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"next_prompt":"Auto next"}')
)
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
enhancer.rewrite_default_model = "gpt-test"
@@ -630,7 +680,8 @@ def test_generate_auto_prompt_uses_selected_model():
next_segment_idx=2,
model="gpt-alt",
timeout_ms=800,
))
)
)
assert result.fallback_used is False
assert result.error is None
assert result.prompt == "Auto next"
@@ -639,7 +690,9 @@ def test_generate_auto_prompt_uses_selected_model():
def test_enhance_prompt_uses_selected_model():
enhancer = _build_test_enhancer(_chat_payload_with_content('{"next_prompt":"Enhanced next"}'))
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"next_prompt":"Enhanced next"}')
)
enhancer.enhance_system_prompt = "enhance system prompt"
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
@@ -669,7 +722,8 @@ def test_enhance_prompt_uses_selected_model():
next_segment_idx=2,
model="gpt-alt",
timeout_ms=800,
))
)
)
assert result.fallback_used is False
assert result.error is None
assert result.prompt == "Enhanced next"
@@ -678,7 +732,9 @@ def test_enhance_prompt_uses_selected_model():
def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field():
enhancer = _build_test_enhancer(_chat_payload_with_content('{"prompt":"Extended single clip"}'))
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"prompt":"Extended single clip"}')
)
enhancer.enhance_system_prompt = "enhance system prompt"
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
@@ -708,12 +764,14 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
enhancer._request_content = _fake_request_content # type: ignore[attr-defined]
result = asyncio.run(enhancer.enhance_prompt(
"short 5s idea",
mode="single_clip",
model="gpt-alt",
timeout_ms=800,
))
result = asyncio.run(
enhancer.enhance_prompt(
"short 5s idea",
mode="single_clip",
model="gpt-alt",
timeout_ms=800,
)
)
assert result.fallback_used is False
assert result.error is None
assert result.prompt == "Extended single clip"
@@ -726,15 +784,17 @@ def test_enhance_prompt_single_clip_uses_auto_extension_prompt_and_prompt_field(
"single 5-second LTX-2.3 video clip. Respond with "
'valid JSON only as {"prompt": "..."}.' # noqa: E501
),
"user_prompt":
"short 5s idea",
"user_prompt": "short 5s idea",
}
def test_enhance_prompt_single_clip_rejects_plain_text_response():
enhancer = _build_test_enhancer(
_chat_payload_with_content("Medium shot of a woman by a rainy cafe window as she lifts her "
"phone, exhales softly, and the camera makes a slow push in."))
_chat_payload_with_content(
"Medium shot of a woman by a rainy cafe window as she lifts her "
"phone, exhales softly, and the camera makes a slow push in."
)
)
enhancer.auto_system_prompt = "auto system prompt"
result = asyncio.run(
@@ -743,14 +803,17 @@ def test_enhance_prompt_single_clip_rejects_plain_text_response():
mode="single_clip",
model="gpt-test",
timeout_ms=800,
))
)
)
assert result.fallback_used is True
assert "No JSON object found in assistant response." in result.error
assert result.prompt == ""
def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
enhancer = _build_test_enhancer(_chat_payload_with_content('{"segment_prompts":["A","B"]}'))
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"segment_prompts":["A","B"]}')
)
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test"]
enhancer.rewrite_default_model = "gpt-test"
@@ -761,14 +824,17 @@ def test_enhance_prompt_single_clip_rejects_segment_prompts_json():
mode="single_clip",
model="gpt-test",
timeout_ms=800,
))
)
)
assert result.fallback_used is True
assert result.prompt == ""
assert "Missing prompt string." in (result.error or "")
def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
enhancer = _build_test_enhancer(_chat_payload_with_content("A cinematic continuation with slow dolly movement."))
enhancer = _build_test_enhancer(
_chat_payload_with_content("A cinematic continuation with slow dolly movement.")
)
enhancer.enhance_system_prompt = "enhance system prompt"
enhancer.rewrite_model_options = ["gpt-test"]
enhancer.rewrite_default_model = "gpt-test"
@@ -780,14 +846,17 @@ def test_enhance_prompt_requires_json_and_does_not_fallback_to_raw_text():
next_segment_idx=2,
model="gpt-test",
timeout_ms=800,
))
)
)
assert result.fallback_used is True
assert result.prompt == ""
assert "No JSON object found in assistant response." in (result.error or "")
def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
enhancer = _build_test_enhancer(_chat_payload_with_content("A calm, grounded continuation with subtle motion."))
enhancer = _build_test_enhancer(
_chat_payload_with_content("A calm, grounded continuation with subtle motion.")
)
enhancer.auto_system_prompt = "auto system prompt"
enhancer.rewrite_model_options = ["gpt-test"]
enhancer.rewrite_default_model = "gpt-test"
@@ -798,30 +867,34 @@ def test_generate_auto_prompt_requires_json_and_does_not_fallback_to_raw_text():
next_segment_idx=2,
model="gpt-test",
timeout_ms=800,
))
)
)
assert result.fallback_used is True
assert result.prompt == ""
assert "No JSON object found in assistant response." in (result.error or "")
def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
enhancer = _build_test_enhancer({
"choices": [{
"finish_reason": "length",
"message": {
"content": [],
"refusal": None,
},
}],
"usage": {
"completion_tokens": 0
},
})
enhancer = _build_test_enhancer(
{
"choices": [
{
"finish_reason": "length",
"message": {
"content": [],
"refusal": None,
},
}
],
"usage": {"completion_tokens": 0},
}
)
result = asyncio.run(
enhancer.rewrite_prompt_sequence(
["prompt one", "prompt two"],
rewrite_instruction="make it cinematic",
))
)
)
assert result.fallback_used is True
assert "No rewrite segment prompts found in assistant response." in (result.error or "")
assert isinstance(result.raw_response_text, str)
@@ -830,7 +903,10 @@ def test_rewrite_prompt_sequence_includes_raw_json_when_content_empty():
def test_get_rewrite_model_config_returns_fixed_defaults():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
enhancer.rewrite_default_model = "gpt-oss-120b"
enhancer.rewrite_model_options = ["gpt-oss-120b"]
@@ -842,7 +918,10 @@ def test_get_rewrite_model_config_returns_fixed_defaults():
def test_get_prompt_config_includes_auto_extension_prompt():
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
enhancer.enhance_system_prompt_path = "/tmp/next.md"
enhancer.auto_system_prompt_path = "/tmp/auto.md"
enhancer.rewrite_all_system_prompt_path = "/tmp/rewrite.md"
@@ -869,14 +948,19 @@ def test_get_prompt_config_includes_auto_extension_prompt():
def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
rewrite_fallback_path = tmp_path / "rewrite_window_system_prompt.md"
rewrite_fallback_path.write_text("rewrite prompt\n", encoding="utf-8")
next_path = tmp_path / "next.md"
next_path.write_text("next prompt\n", encoding="utf-8")
auto_path = tmp_path / "auto.md"
auto_path.write_text("auto prompt\n", encoding="utf-8")
enhancer.rewrite_all_system_prompt_path = str(tmp_path / "prompts.local" / "rewrite_window_system_prompt.md")
enhancer.rewrite_all_system_prompt_path = str(
tmp_path / "prompts.local" / "rewrite_window_system_prompt.md"
)
enhancer.rewrite_all_system_prompt_fallback_path = str(rewrite_fallback_path)
enhancer.enhance_system_prompt_path = str(next_path)
enhancer.auto_system_prompt_path = str(auto_path)
@@ -889,9 +973,14 @@ def test_get_prompt_config_reports_loaded_fallback_prompt_path(tmp_path):
assert config["rewrite_window_system_prompt_path"] == str(rewrite_fallback_path)
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(tmp_path, ):
def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_empty(
tmp_path,
):
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
@@ -918,7 +1007,10 @@ def test_reload_system_prompts_falls_back_to_rewrite_window_when_user_prompt_emp
def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite.md"
@@ -929,7 +1021,9 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
enhancer.auto_system_prompt_path = str(auto_path)
enhancer.rewrite_all_system_prompt_path = str(rewrite_path)
config = enhancer.save_prompt_config(auto_extension_system_prompt="auto updated", )
config = enhancer.save_prompt_config(
auto_extension_system_prompt="auto updated",
)
assert auto_path.read_text(encoding="utf-8").strip() == "auto updated"
assert config["auto_extension_system_prompt"] == "auto updated"
@@ -937,7 +1031,10 @@ def test_save_prompt_config_updates_auto_extension_prompt(tmp_path):
def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite.md"
@@ -955,7 +1052,9 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
enhancer.rewrite_all_system_prompt_fallback_path = None
enhancer.rewrite_user_system_prompt_fallback_path = None
config = enhancer.save_prompt_config(rewrite_user_system_prompt="rewrite user updated", )
config = enhancer.save_prompt_config(
rewrite_user_system_prompt="rewrite user updated",
)
assert rewrite_user_path.read_text(encoding="utf-8").strip() == "rewrite user updated"
assert config["rewrite_user_system_prompt"] == "rewrite user updated"
@@ -963,7 +1062,10 @@ def test_save_prompt_config_updates_rewrite_user_prompt(tmp_path):
def test_save_prompt_config_updates_rewrite_model(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite.md"
@@ -979,7 +1081,9 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
enhancer.rewrite_default_model = "gpt-test"
enhancer.rewrite_model_options = ["gpt-test", "gpt-alt"]
config = enhancer.save_prompt_config(rewrite_model="gpt-alt", )
config = enhancer.save_prompt_config(
rewrite_model="gpt-alt",
)
assert enhancer.rewrite_default_model == "gpt-alt"
assert config["rewrite_model"] == "gpt-alt"
@@ -988,7 +1092,10 @@ def test_save_prompt_config_updates_rewrite_model(tmp_path):
def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite.md"
@@ -1002,7 +1109,9 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
enhancer.auto_system_prompt_fallback_path = None
enhancer.rewrite_all_system_prompt_fallback_path = None
config = enhancer.save_prompt_config(rewrite_temperature=1.3, )
config = enhancer.save_prompt_config(
rewrite_temperature=1.3,
)
assert enhancer.rewrite_default_temperature == 1.3
assert config["rewrite_temperature"] == 1.3
@@ -1010,7 +1119,10 @@ def test_save_prompt_config_updates_rewrite_temperature(tmp_path):
def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_path):
enhancer = _build_test_enhancer(
_chat_payload_with_content('{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'))
_chat_payload_with_content(
'{"id":"preset_a","label":"Preset A","segment_prompts":["A","B"]}'
)
)
next_path = tmp_path / "next.md"
auto_path = tmp_path / "auto.md"
rewrite_path = tmp_path / "rewrite_window_system_prompt.md"
@@ -1024,9 +1136,13 @@ def test_save_prompt_config_creates_versioned_backup_for_existing_prompt(tmp_pat
enhancer.auto_system_prompt_fallback_path = None
enhancer.rewrite_all_system_prompt_fallback_path = None
enhancer.save_prompt_config(rewrite_window_system_prompt="rewrite updated", )
enhancer.save_prompt_config(
rewrite_window_system_prompt="rewrite updated",
)
backup_paths = sorted(tmp_path.glob("rewrite_window_system_prompt.*.bak.md"))
backup_paths = sorted(
tmp_path.glob("rewrite_window_system_prompt.*.bak.md")
)
assert rewrite_path.read_text(encoding="utf-8").strip() == "rewrite updated"
assert len(backup_paths) == 1
@@ -27,8 +27,13 @@ try:
except ModuleNotFoundError:
websockets = None # type: ignore[assignment]
DEFAULT_PRESET_FILE = (Path(__file__).resolve().parents[2] / "web" / "prompts" /
"selected_ltx2_continuation_story_presets.json")
DEFAULT_PRESET_FILE = (
Path(__file__).resolve().parents[2]
/ "web"
/ "prompts"
/ "selected_ltx2_continuation_story_presets.json"
)
def utc_now_iso() -> str:
@@ -60,7 +65,10 @@ def safe_percentile(values: list[float], percentile: float) -> float | None:
if lower == upper:
return sorted_values[lower]
fraction = rank - lower
return (sorted_values[lower] + (sorted_values[upper] - sorted_values[lower]) * fraction)
return (
sorted_values[lower]
+ (sorted_values[upper] - sorted_values[lower]) * fraction
)
def summarize_series(values: list[float]) -> dict[str, float | int | None]:
@@ -137,16 +145,24 @@ def load_curated_prompts(
selected_id = str(selected.get("id", "")).strip() or "unknown_preset"
raw_prompts = selected.get("segment_prompts", [])
if not isinstance(raw_prompts, list):
raise ValueError(f"Preset {selected_id} has invalid segment_prompts (must be list).")
raise ValueError(
f"Preset {selected_id} has invalid segment_prompts (must be list)."
)
prompts = [str(prompt).strip() for prompt in raw_prompts if isinstance(prompt, str) and str(prompt).strip()]
prompts = [
str(prompt).strip()
for prompt in raw_prompts
if isinstance(prompt, str) and str(prompt).strip()
]
if not prompts:
raise ValueError(f"Preset {selected_id} has no non-empty prompts.")
limited = prompts[:curated_limit]
if not limited:
raise ValueError(f"curated_limit={curated_limit} produced no prompts for preset "
f"{selected_id}.")
raise ValueError(
f"curated_limit={curated_limit} produced no prompts for preset "
f"{selected_id}."
)
return selected_id, limited, len(prompts)
@@ -208,11 +224,11 @@ async def run_single_session(
try:
async with websockets.connect(
url,
max_size=None,
ping_interval=None,
open_timeout=connect_timeout_s,
close_timeout=2.0,
url,
max_size=None,
ping_interval=None,
open_timeout=connect_timeout_s,
close_timeout=2.0,
) as ws:
connect_finish_monotonic = time.monotonic()
session_data["connect_finish_ts_utc"] = utc_now_iso()
@@ -233,7 +249,9 @@ async def run_single_session(
timeout_remaining = session_timeout_s - elapsed_s
if timeout_remaining <= 0:
session_data["status"] = "timeout"
session_data["error"] = (f"Session timed out after {session_timeout_s:.1f}s.")
session_data["error"] = (
f"Session timed out after {session_timeout_s:.1f}s."
)
break
recv_start_epoch = time.time()
@@ -247,7 +265,9 @@ async def run_single_session(
)
except asyncio.TimeoutError:
session_data["status"] = "timeout"
session_data["error"] = ("Timed out waiting for websocket message.")
session_data["error"] = (
"Timed out waiting for websocket message."
)
break
except Exception as exc:
session_data["status"] = "failed"
@@ -268,16 +288,20 @@ async def run_single_session(
chunk_gap_ms: float | None = None
if last_chunk_finish_monotonic is not None:
chunk_gap_ms = (recv_finish_monotonic - last_chunk_finish_monotonic) * 1000.0
chunk_gap_ms = (
recv_finish_monotonic - last_chunk_finish_monotonic
) * 1000.0
session_data["chunks"].append({
"segment_idx": current_segment_idx,
"chunk_idx": session_data["total_chunks"],
"size_bytes": len(message),
"chunk_start_ts_utc": recv_start_iso,
"chunk_finish_ts_utc": recv_finish_iso,
"chunk_gap_ms": chunk_gap_ms,
})
session_data["chunks"].append(
{
"segment_idx": current_segment_idx,
"chunk_idx": session_data["total_chunks"],
"size_bytes": len(message),
"chunk_start_ts_utc": recv_start_iso,
"chunk_finish_ts_utc": recv_finish_iso,
"chunk_gap_ms": chunk_gap_ms,
}
)
last_chunk_finish_monotonic = recv_finish_monotonic
last_chunk_finish_epoch = recv_finish_epoch
session_data["last_chunk_finish_ts_utc"] = recv_finish_iso
@@ -297,7 +321,9 @@ async def run_single_session(
if msg_type == "gpu_assigned":
session_data["gpu_assigned_ts_utc"] = recv_finish_iso
if connect_finish_monotonic is not None:
session_data["queue_wait_ms"] = (recv_finish_monotonic - connect_finish_monotonic) * 1000.0
session_data["queue_wait_ms"] = (
recv_finish_monotonic - connect_finish_monotonic
) * 1000.0
elif msg_type == "ltx2_stream_start":
if initial_total_segments is None:
parsed_total = parse_int(data.get("total_segments"))
@@ -312,13 +338,20 @@ async def run_single_session(
session_data["media_segments_completed"] += 1
if first_media_segment_complete_epoch is None:
first_media_segment_complete_epoch = recv_finish_epoch
session_data["first_media_segment_complete_ts_utc"] = recv_finish_iso
session_data[
"first_media_segment_complete_ts_utc"
] = recv_finish_iso
elif msg_type == "ltx2_segment_complete":
session_data["segments_completed"] += 1
seg_idx = parse_int(data.get("segment_idx"))
if (initial_total_segments is not None and seg_idx is not None
and seg_idx >= initial_total_segments):
session_data["target_segment_complete_ts_utc"] = recv_finish_iso
if (
initial_total_segments is not None
and seg_idx is not None
and seg_idx >= initial_total_segments
):
session_data[
"target_segment_complete_ts_utc"
] = recv_finish_iso
await asyncio.sleep(post_complete_wait_s)
session_data["leave_sent_ts_utc"] = utc_now_iso()
try:
@@ -329,11 +362,15 @@ async def run_single_session(
break
elif msg_type == "session_timeout":
session_data["status"] = "timeout"
session_data["error"] = str(data.get("message") or "Backend session timeout")
session_data["error"] = str(
data.get("message") or "Backend session timeout"
)
break
elif msg_type == "error":
session_data["status"] = "failed"
session_data["error"] = str(data.get("message") or "Backend error message")
session_data["error"] = str(
data.get("message") or "Backend error message"
)
break
if session_data["status"] == "failed" and session_data["error"] is None:
@@ -342,18 +379,29 @@ async def run_single_session(
session_data["status"] = "failed"
session_data["error"] = f"WebSocket connect/run failed: {exc}"
if (first_chunk_finish_epoch is not None and last_chunk_finish_epoch is not None
and session_data["total_chunk_bytes"] > 0):
if (
first_chunk_finish_epoch is not None
and last_chunk_finish_epoch is not None
and session_data["total_chunk_bytes"] > 0
):
duration_s = last_chunk_finish_epoch - first_chunk_finish_epoch
if duration_s > 0:
session_data["session_goodput_mbps"] = (session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0)
session_data["session_goodput_mbps"] = (
session_data["total_chunk_bytes"] * 8.0 / duration_s / 1_000_000.0
)
if (first_chunk_finish_epoch is not None and first_media_segment_complete_epoch is not None):
session_data["first_chunk_before_first_media_complete"] = (first_chunk_finish_epoch
< first_media_segment_complete_epoch)
if (
first_chunk_finish_epoch is not None
and first_media_segment_complete_epoch is not None
):
session_data["first_chunk_before_first_media_complete"] = (
first_chunk_finish_epoch < first_media_segment_complete_epoch
)
session_data["close_ts_utc"] = utc_now_iso()
session_data["duration_ms"] = (time.monotonic() - session_start_monotonic) * 1000.0
session_data["duration_ms"] = (
time.monotonic() - session_start_monotonic
) * 1000.0
return session_data
@@ -364,11 +412,14 @@ async def run_worker_sessions(
config: dict[str, Any],
) -> list[dict[str, Any]]:
tasks = [
asyncio.create_task(run_single_session(
worker_id=worker_id,
worker_session_idx=idx,
config=config,
)) for idx in range(session_count)
asyncio.create_task(
run_single_session(
worker_id=worker_id,
worker_session_idx=idx,
config=config,
)
)
for idx in range(session_count)
]
if not tasks:
return []
@@ -386,23 +437,29 @@ def worker_entry(
try:
ready_queue.put({"worker_id": worker_id, "status": "ready"})
start_event.wait()
sessions = asyncio.run(run_worker_sessions(
worker_id=worker_id,
session_count=session_count,
config=config,
))
result_queue.put({
"worker_id": worker_id,
"status": "ok",
"sessions": sessions,
})
sessions = asyncio.run(
run_worker_sessions(
worker_id=worker_id,
session_count=session_count,
config=config,
)
)
result_queue.put(
{
"worker_id": worker_id,
"status": "ok",
"sessions": sessions,
}
)
except Exception as exc:
result_queue.put({
"worker_id": worker_id,
"status": "error",
"error": str(exc),
"traceback": traceback.format_exc(),
})
result_queue.put(
{
"worker_id": worker_id,
"status": "error",
"error": str(exc),
"traceback": traceback.format_exc(),
}
)
def build_summary(
@@ -460,22 +517,33 @@ def build_summary(
if len(all_chunk_finish_epochs) >= 2 and total_chunk_bytes > 0:
duration_s = max(all_chunk_finish_epochs) - min(all_chunk_finish_epochs)
if duration_s > 0:
global_goodput_mbps = (total_chunk_bytes * 8.0 / duration_s / 1_000_000.0)
global_goodput_mbps = (
total_chunk_bytes * 8.0 / duration_s / 1_000_000.0
)
bucket_throughputs_mbps = [(bytes_count * 8.0) / 1_000_000.0 for _, bytes_count in sorted(bucket_bytes.items())]
bucket_throughputs_mbps = [
(bytes_count * 8.0) / 1_000_000.0
for _, bytes_count in sorted(bucket_bytes.items())
]
bucket_stats = summarize_series(bucket_throughputs_mbps)
chunk_gap_threshold_breaches = [value for value in chunk_gaps if value >= chunk_gap_threshold_ms]
chunk_gap_threshold_breaches = [
value for value in chunk_gaps if value >= chunk_gap_threshold_ms
]
non_success = len(sessions) - status_counts.get("success", 0)
fail_reasons: list[str] = []
if non_success > 0:
fail_reasons.append(f"{non_success} session(s) did not complete successfully.")
fail_reasons.append(
f"{non_success} session(s) did not complete successfully."
)
if not chunk_gaps:
fail_reasons.append("No chunk gap data collected.")
if chunk_gap_threshold_breaches:
fail_reasons.append(f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
f"{chunk_gap_threshold_ms:.0f}ms.")
fail_reasons.append(
f"{len(chunk_gap_threshold_breaches)} chunk gap(s) were >= "
f"{chunk_gap_threshold_ms:.0f}ms."
)
passed = len(fail_reasons) == 0
progressive_ratio = None
@@ -486,18 +554,20 @@ def build_summary(
"passed": passed,
"fail_reasons": fail_reasons,
"sessions": {
"total":
len(sessions),
"success":
status_counts.get("success", 0),
"failed":
status_counts.get("failed", 0),
"timeout":
status_counts.get("timeout", 0),
"protocol_error":
status_counts.get("protocol_error", 0),
"other": (len(sessions) - (status_counts.get("success", 0) + status_counts.get("failed", 0) +
status_counts.get("timeout", 0) + status_counts.get("protocol_error", 0))),
"total": len(sessions),
"success": status_counts.get("success", 0),
"failed": status_counts.get("failed", 0),
"timeout": status_counts.get("timeout", 0),
"protocol_error": status_counts.get("protocol_error", 0),
"other": (
len(sessions)
- (
status_counts.get("success", 0)
+ status_counts.get("failed", 0)
+ status_counts.get("timeout", 0)
+ status_counts.get("protocol_error", 0)
)
),
},
"chunk_gap_ms": {
**chunk_gap_stats,
@@ -536,39 +606,51 @@ def print_summary(
bucket_bw = bandwidth["bucketed_1s"]
print("=== LTX2 Realtime Stress Test Summary ===")
print("Run: "
f"url={run_info['url']} clients={run_info['clients']} "
f"processes={run_info['processes']} "
f"preset={run_info['preset_id']} "
f"curated_limit={run_info['curated_limit']}")
print("Sessions: "
f"total={sessions['total']} success={sessions['success']} "
f"failed={sessions['failed']} timeout={sessions['timeout']} "
f"protocol_error={sessions['protocol_error']}")
print("Chunk gap ms: "
f"min={format_num(chunk_gap['min'])} "
f"p50={format_num(chunk_gap['p50'])} "
f"p95={format_num(chunk_gap['p95'])} "
f"p99={format_num(chunk_gap['p99'])} "
f"max={format_num(chunk_gap['max'])} "
f"threshold={format_num(chunk_gap['threshold_ms'])} "
f"breaches={chunk_gap['breach_count']}")
print("Queue wait ms: "
f"min={format_num(queue_wait['min'])} "
f"p50={format_num(queue_wait['p50'])} "
f"p95={format_num(queue_wait['p95'])} "
f"max={format_num(queue_wait['max'])}")
print(
"Run: "
f"url={run_info['url']} clients={run_info['clients']} "
f"processes={run_info['processes']} "
f"preset={run_info['preset_id']} "
f"curated_limit={run_info['curated_limit']}"
)
print(
"Sessions: "
f"total={sessions['total']} success={sessions['success']} "
f"failed={sessions['failed']} timeout={sessions['timeout']} "
f"protocol_error={sessions['protocol_error']}"
)
print(
"Chunk gap ms: "
f"min={format_num(chunk_gap['min'])} "
f"p50={format_num(chunk_gap['p50'])} "
f"p95={format_num(chunk_gap['p95'])} "
f"p99={format_num(chunk_gap['p99'])} "
f"max={format_num(chunk_gap['max'])} "
f"threshold={format_num(chunk_gap['threshold_ms'])} "
f"breaches={chunk_gap['breach_count']}"
)
print(
"Queue wait ms: "
f"min={format_num(queue_wait['min'])} "
f"p50={format_num(queue_wait['p50'])} "
f"p95={format_num(queue_wait['p95'])} "
f"max={format_num(queue_wait['max'])}"
)
ratio = progressive["ratio"]
ratio_text = "n/a" if ratio is None else f"{ratio * 100:.2f}%"
print("Progressive streaming: "
f"{progressive['success_sessions']}/"
f"{progressive['eligible_sessions']} ({ratio_text})")
print("Bandwidth Mbps: "
f"per_session_avg={format_num(per_session_bw['avg'])} "
f"per_session_p95={format_num(per_session_bw['p95'])} "
f"global={format_num(bandwidth['global_goodput_mbps'])} "
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}")
print(
"Progressive streaming: "
f"{progressive['success_sessions']}/"
f"{progressive['eligible_sessions']} ({ratio_text})"
)
print(
"Bandwidth Mbps: "
f"per_session_avg={format_num(per_session_bw['avg'])} "
f"per_session_p95={format_num(per_session_bw['p95'])} "
f"global={format_num(bandwidth['global_goodput_mbps'])} "
f"bucket_avg={format_num(bucket_bw['avg_mbps'])} "
f"bucket_peak={format_num(bucket_bw['peak_mbps'])}"
)
print(f"VERDICT: {'PASS' if summary['passed'] else 'FAIL'}")
if summary["fail_reasons"]:
print("Fail reasons:")
@@ -588,8 +670,10 @@ def distribute_sessions(total_clients: int, process_count: int) -> list[int]:
def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
if websockets is None:
raise RuntimeError("Missing dependency: websockets. Install it before running this "
"stress test.")
raise RuntimeError(
"Missing dependency: websockets. Install it before running this "
"stress test."
)
preset_file = Path(args.preset_file).expanduser().resolve()
selected_preset_id, curated_prompts, total_prompt_count = load_curated_prompts(
@@ -651,8 +735,13 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
start_event.set()
result_deadline = (time.monotonic() + args.connect_timeout_s + args.session_timeout_s +
args.post_complete_wait_s + 180.0)
result_deadline = (
time.monotonic()
+ args.connect_timeout_s
+ args.session_timeout_s
+ args.post_complete_wait_s
+ 180.0
)
worker_results: list[dict[str, Any]] = []
while len(worker_results) < len(processes):
timeout_s = max(0.1, result_deadline - time.monotonic())
@@ -676,20 +765,24 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
if result.get("status") == "ok":
sessions.extend(result.get("sessions", []))
else:
worker_errors.append({
"worker_id": result.get("worker_id"),
"error": result.get("error"),
"traceback": result.get("traceback"),
})
worker_errors.append(
{
"worker_id": result.get("worker_id"),
"error": result.get("error"),
"traceback": result.get("traceback"),
}
)
received_workers = {result.get("worker_id") for result in worker_results}
expected_workers = set(range(len(processes)))
missing_workers = sorted(expected_workers - received_workers)
for worker_id in missing_workers:
worker_errors.append({
"worker_id": worker_id,
"error": "No worker result received.",
})
worker_errors.append(
{
"worker_id": worker_id,
"error": "No worker result received.",
}
)
run_end_epoch = time.time()
run_end_iso = iso_from_epoch(run_end_epoch)
@@ -702,8 +795,9 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
if worker_errors:
summary["passed"] = False
summary["fail_reasons"] = list(
summary["fail_reasons"]) + [f"{len(worker_errors)} worker error(s) occurred."]
summary["fail_reasons"] = list(summary["fail_reasons"]) + [
f"{len(worker_errors)} worker error(s) occurred."
]
output_payload = {
"run_info": {
@@ -739,7 +833,9 @@ def run_stress(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Multiprocess realtime stress test for LTX2 streaming.", )
parser = argparse.ArgumentParser(
description="Multiprocess realtime stress test for LTX2 streaming.",
)
parser.add_argument(
"-u",
"--url",
@@ -47,11 +47,13 @@ def test_persist_session_init_image_returns_none_when_missing_data():
def test_persist_session_init_image_rejects_unsupported_mime():
with pytest.raises(ValueError, match="PNG, JPEG, or WebP"):
persist_session_init_image({
"name": "frame.gif",
"mime_type": "image/gif",
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
})
persist_session_init_image(
{
"name": "frame.gif",
"mime_type": "image/gif",
"data_url": "data:image/gif;base64,R0lGODlhAQABAAAAACw=",
}
)
def test_persist_session_init_image_rejects_large_payload(monkeypatch):
@@ -64,8 +66,10 @@ def test_persist_session_init_image_rejects_large_payload(monkeypatch):
monkeypatch.setattr(base64, "b64decode", fake_b64decode)
with pytest.raises(ValueError, match="15 MB or smaller"):
persist_session_init_image({
"name": "frame.png",
"mime_type": "image/png",
"data_url": data_url,
})
persist_session_init_image(
{
"name": "frame.png",
"mime_type": "image/png",
"data_url": data_url,
}
)
File diff suppressed because it is too large Load Diff
+15 -9
View File
@@ -8,10 +8,12 @@ import modal
IMAGE = os.environ.get("DREAMVERSE_IMAGE")
if not IMAGE:
raise RuntimeError("DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag.")
raise RuntimeError(
"DREAMVERSE_IMAGE is required. Set it to a published SHA-specific Dreamverse image, "
"for example a dreamverse-backend-cuda13.0.0-sha-* tag or a "
"dreamverse-ui-cuda13.0.0-sha-* tag if serving the static UI. "
"CUDA 12 / cu126 images use the corresponding cuda12.6.3 tag."
)
# ``@modal.web_server`` invokes ``serve()`` directly and bypasses the image
# ENTRYPOINT (``docker/docker_entrypoint.sh``). That entrypoint normally
@@ -63,10 +65,14 @@ def serve():
# ``or ""`` collapses ``None`` (unset) into an empty string, ``.strip()``
# collapses whitespace-only values (e.g. ``" "``) — both should be
# treated as missing.
missing = [k for k in _REQUIRED_SECRET_KEYS if not (os.environ.get(k) or "").strip()]
missing = [
k for k in _REQUIRED_SECRET_KEYS
if not (os.environ.get(k) or "").strip()
]
if missing:
raise RuntimeError("dreamverse-api-keys secret is missing required entries: "
f"{', '.join(missing)}. Add them with `modal secret create "
"dreamverse-api-keys ... --force` and redeploy "
"(see apps/dreamverse/scripts/modal/README.md).")
raise RuntimeError(
"dreamverse-api-keys secret is missing required entries: "
f"{', '.join(missing)}. Add them with `modal secret create "
"dreamverse-api-keys ... --force` and redeploy "
"(see apps/dreamverse/scripts/modal/README.md).")
subprocess.Popen(["dreamverse-server", "--host", "0.0.0.0", "--port", "8009"])
@@ -0,0 +1,39 @@
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();
});
});
+87
View File
@@ -72,3 +72,90 @@ 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
+282 -61
View File
@@ -41,6 +41,11 @@ _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"
@@ -259,13 +264,36 @@ class JobRunner:
self._jobs_lock = threading.Lock()
self._load_jobs()
# 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()
# 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
# 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
@@ -617,6 +645,211 @@ 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,
@@ -633,69 +866,53 @@ class JobRunner:
sp_size: int = -1,
log_queue: mp.Queue | None = None,
) -> Any:
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,
"""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,
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))
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 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
def _run_job(self, job: Job):
if job.job_type == "inference":
@@ -841,10 +1058,12 @@ class JobRunner:
fastvideo_logger.addHandler(buffer_handler)
fastvideo_logger.addHandler(file_handler)
# 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()
# 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_listener = logging.handlers.QueueListener(log_queue,
buffer_handler,
file_handler,
@@ -926,6 +1145,7 @@ 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] = {
@@ -941,7 +1161,6 @@ 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
@@ -985,6 +1204,8 @@ 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)
+187 -2
View File
@@ -36,13 +36,15 @@ 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, SettingsUpdate, UpdateCaptionRequest,
model_label)
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, GeneratorRequest, 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
@@ -88,6 +90,10 @@ _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]] = {}
@@ -393,6 +399,56 @@ 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 -------------------------------------------------------------------
@@ -414,6 +470,49 @@ 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 ----------------------------------------------------------------
@@ -570,6 +669,92 @@ 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 ---------------------------------------------------------------
+2
View File
@@ -3,6 +3,7 @@
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
@@ -15,6 +16,7 @@ def model_label(model_path: str) -> str:
__all__ = [
"CreateJobRequest",
"GeneratorRequest",
"SettingsUpdate",
"CreateDatasetRequest",
"UpdateCaptionRequest",
@@ -0,0 +1,24 @@
# 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
+151 -4
View File
@@ -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_gpu_snapshot
from fastvideo_studio.gpu import get_cluster_snapshot, get_gpu_snapshot
from fastvideo_studio.job_runner import JobRunner, JobStatus
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, SettingsUpdate, UpdateCaptionRequest,
model_label)
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, GeneratorRequest, SettingsUpdate,
UpdateCaptionRequest, model_label)
logging.basicConfig(
level=logging.INFO,
@@ -45,6 +45,67 @@ 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)
@@ -102,6 +163,13 @@ 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.
@@ -116,6 +184,25 @@ 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"}
@@ -253,13 +340,58 @@ 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}
if req.model_id not in valid_ids:
# 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):
raise HTTPException(
status_code=400,
detail=(f"Unknown model_id '{req.model_id}'. "
@@ -622,6 +754,14 @@ 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
@@ -686,6 +826,9 @@ 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)
@@ -695,6 +838,10 @@ 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,
)
@@ -0,0 +1,182 @@
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();
});
});
+250 -3
View File
@@ -1,11 +1,258 @@
'use client';
import GpuGrid from '@/components/system/GpuGrid';
import * as React from 'react';
import { AlertTriangle, Info } from 'lucide-react';
export default function GpusPage() {
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>
);
}
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">
<GpuGrid />
{body}
</div>
);
}
@@ -1,8 +1,10 @@
'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 (
@@ -10,6 +12,8 @@ export default function InferencePage() {
<HeaderActions>
<CreateJobButton jobType="inference" />
</HeaderActions>
<WarmModelsPanel />
<EngineConsole />
<JobQueue jobType="inference" />
</>
);
@@ -0,0 +1,93 @@
'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,14 +4,24 @@ import userEvent from '@testing-library/user-event';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import CreateJobModal from './CreateJobModal';
import { createJob, getDatasets, getModels, uploadImage } from '@/lib/api';
import {
createJob,
getDatasets,
getModelPresets,
getModels,
listGenerators,
uploadImage,
type GeneratorInfo,
} 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(),
@@ -22,11 +32,32 @@ 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);
});
@@ -206,4 +237,173 @@ 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,15 +23,27 @@ 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;
@@ -55,6 +67,9 @@ 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('');
@@ -64,6 +79,12 @@ 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);
@@ -178,17 +199,34 @@ export default function CreateJobModal({
let stale = false;
setIsLoadingModels(true);
setModelLoadError(null);
getModels(inferenceWorkload)
.then((list) => {
// 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]) => {
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 chosen = ids.includes(defaultId) ? defaultId : (list[0]?.id ?? '');
const warmId = slot?.model_id ?? '';
const chosen =
warmId && ids.includes(warmId)
? warmId
: ids.includes(defaultId)
? defaultId
: (list[0]?.id ?? '');
setModelId(chosen);
if (workloadType === 'dmd_t2v') {
setRealScoreModelPath(chosen);
@@ -210,7 +248,81 @@ export default function CreateJobModal({
return () => {
stale = true;
};
}, [isOpen, inferenceWorkload, workloadType]);
}, [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]);
// Training jobs need a dataset; load the ready datasets when relevant.
React.useEffect(() => {
@@ -741,6 +853,64 @@ 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"
@@ -0,0 +1,71 @@
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();
});
});
@@ -0,0 +1,123 @@
'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-3 flex cursor-pointer flex-col gap-2.5 rounded-lg border bg-background p-4 transition-colors last:mb-0',
'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',
isSelected
? 'border-accent-blue bg-accent-blue/5'
: 'border-border hover:border-muted-foreground/40',
@@ -152,20 +152,16 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
type="button"
aria-pressed={isSelected}
onClick={handleSelectJob}
className="flex w-full flex-col gap-2.5 rounded-md text-left"
className="flex w-full min-w-0 flex-col gap-1 rounded-md text-left"
>
<span className="flex flex-wrap items-center justify-between gap-2">
<span className="text-[0.95rem] font-semibold text-foreground">
{job.model_id}
</span>
<Badge variant={BADGE_VARIANTS[job.status] ?? 'secondary'}>
{job.status}
</Badge>
<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>
<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">
<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">
{job.job_type === 'inference' ? (
<>
<span>{job.num_frames} frames</span>
@@ -183,6 +179,10 @@ 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="border-transparent bg-amber-500 text-black shadow-md hover:bg-amber-400"
className="h-6 px-2 text-xs 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="border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
className="h-6 px-2 text-xs 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="border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
className="h-6 px-2 text-xs border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
>
Start
</Button>
@@ -222,6 +222,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
onClick={handleDownloadVideo}
disabled={isLoading}
title="Download video"
className="h-6 px-2 text-xs"
>
Download Video
</Button>
@@ -231,6 +232,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
variant="destructive"
onClick={handleDelete}
disabled={isLoading}
className="h-6 px-2 text-xs"
>
Delete
</Button>
@@ -9,6 +9,7 @@ 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 =>
@@ -45,6 +46,43 @@ 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 } from '@/lib/api';
import { downloadJobLog, getJobLogs, getJobVideoUrl } from '@/lib/api';
import type { Job } from '@/lib/types';
import { cn, downloadBlob } from '@/lib/utils';
@@ -180,6 +180,41 @@ 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">
@@ -0,0 +1,188 @@
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));
});
});
@@ -0,0 +1,237 @@
'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': 'Jobs',
'/finetuning': 'Jobs',
'/distillation': 'Jobs',
'/inference': 'Studio',
'/finetuning': 'Studio',
'/distillation': 'Studio',
'/datasets': 'Datasets',
'/gallery': 'Gallery',
'/gpus': 'GPUs',
@@ -108,7 +108,7 @@ export default function PrimarySidebar({
isJobsActive && TAB_ACTIVE,
)}
>
<span>Jobs</span>
<span>Studio</span>
<svg
viewBox="0 0 24 24"
fill="none"
+162
View File
@@ -202,6 +202,33 @@ 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`);
@@ -319,6 +346,109 @@ 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 {
@@ -418,3 +548,35 @@ 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();
}
@@ -0,0 +1,273 @@
# 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
+6 -3
View File
@@ -74,8 +74,10 @@ def test_snapshot_shapes_devices(monkeypatch: pytest.MonkeyPatch) -> None:
}
def test_snapshot_tolerates_missing_sensors(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setitem(sys.modules, "pynvml", _make_fake_pynvml(broken_sensors=True))
def test_snapshot_tolerates_missing_sensors(
monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setitem(sys.modules, "pynvml",
_make_fake_pynvml(broken_sensors=True))
snap = gpu_mod.get_gpu_snapshot()
assert snap["available"] is True
g = snap["gpus"][0]
@@ -84,7 +86,8 @@ def test_snapshot_tolerates_missing_sensors(monkeypatch: pytest.MonkeyPatch) ->
assert g["power_limit_watts"] is None
def test_snapshot_reports_nvml_failure(monkeypatch: pytest.MonkeyPatch) -> None:
def test_snapshot_reports_nvml_failure(
monkeypatch: pytest.MonkeyPatch) -> None:
fake = _make_fake_pynvml()
fake.nvmlInit = lambda: (_ for _ in ()).throw(_NVMLError("driver gone"))
monkeypatch.setitem(sys.modules, "pynvml", fake)
@@ -130,7 +130,8 @@ def test_dmd_builds_three_role_models_and_method_knobs() -> None:
def test_dmd_vsa_maps_to_training_vsa_sparsity() -> None:
config = build_training_config(_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
config = build_training_config(
_job("dmd_t2v", dmd_use_vsa=True, dmd_vsa_sparsity=0.9), "out")
assert config["training"]["vsa"]["sparsity"] == 0.9
@@ -160,27 +161,31 @@ def test_validation_callback_only_when_file_given() -> None:
without = build_training_config(_job("full_t2v"), "out")
assert "validation" not in without["callbacks"]
with_file = build_training_config(_job("full_t2v", validation_dataset_file="val.json"), "out")
with_file = build_training_config(
_job("full_t2v", validation_dataset_file="val.json"), "out")
validation = with_file["callbacks"]["validation"]
assert validation["dataset_file"] == "val.json"
assert validation["pipeline_target"].endswith(".WanPipeline")
assert validation["sampling_steps"] == [50]
dmd = build_training_config(_job("dmd_t2v", validation_dataset_file="val.json"), "out")
dmd = build_training_config(
_job("dmd_t2v", validation_dataset_file="val.json"), "out")
validation = dmd["callbacks"]["validation"]
assert validation["pipeline_target"].endswith(".WanDMDPipeline")
assert validation["sampling_steps"] == [3]
assert validation["sampling_timesteps"] == [1000, 757, 522]
# KD/ODE-init has no sampling-based validation pipeline.
ode = build_training_config(_job("ode_init", validation_dataset_file="val.json"), "out")
ode = build_training_config(
_job("ode_init", validation_dataset_file="val.json"), "out")
assert "validation" not in ode["callbacks"]
def test_ltx2_models_are_rejected() -> None:
assert is_ltx2_model("Lightricks/LTX-2-19B")
with pytest.raises(ValueError, match="LTX-2 training is not supported"):
build_training_config(_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
build_training_config(
_job("full_t2v", model_id="Lightricks/LTX-2-19B"), "out")
def test_unknown_workload_is_rejected() -> None:
@@ -190,7 +195,8 @@ def test_unknown_workload_is_rejected() -> None:
def test_invalid_denoising_steps_are_rejected() -> None:
with pytest.raises(ValueError, match="Invalid DMD denoising steps"):
build_training_config(_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
build_training_config(
_job("dmd_t2v", dmd_denoising_steps="1000,abc"), "out")
def test_training_env_has_no_backend_override() -> None:
@@ -205,7 +211,8 @@ def test_workloads_match_frontend_job_config() -> None:
(src/lib/jobConfig.ts) — drift means creatable-but-unrunnable jobs."""
import re
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" / "jobConfig.ts").read_text(encoding="utf-8")
job_config = (Path(__file__).resolve().parents[1] / "src" / "lib" /
"jobConfig.ts").read_text(encoding="utf-8")
all_types = set(re.findall(r'type:\s*"([^"]+)"', job_config))
inference_types = {"t2v", "i2v", "t2i"}
assert inference_types <= all_types, "jobConfig.ts parse failed"
+1 -1
View File
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
-52
View File
@@ -1,52 +0,0 @@
{
"recipes": [
{
"id": "fastwan21-t2v",
"task": "Text to video",
"label": "FastWan2.1 1.3B (distilled + VSA)",
"model": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"source": "scripts/inference/inference_wan_VSA_DMD_1_3B.yaml",
"command": "FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml"
},
{
"id": "wan22-t2v",
"task": "Text to video",
"label": "Wan2.2 A14B",
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
"source": "examples/inference/basic/basic_wan2_2.py",
"command": "python examples/inference/basic/basic_wan2_2.py"
},
{
"id": "wan21-i2v",
"task": "Image to video",
"label": "Wan2.1 14B 480P",
"model": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
"source": "scripts/inference/inference_wan_i2v.yaml",
"command": "fastvideo generate --config scripts/inference/inference_wan_i2v.yaml"
},
{
"id": "turbowan22-i2v",
"task": "Image to video",
"label": "TurboWan2.2 A14B",
"model": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
"source": "examples/inference/basic/basic_turbodiffusion_i2v.py",
"command": "python examples/inference/basic/basic_turbodiffusion_i2v.py"
},
{
"id": "wan22-ti2v",
"task": "Text or image to video",
"label": "Wan2.2 TI2V 5B",
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
"source": "examples/inference/basic/basic_wan2_2_ti2v.py",
"command": "python examples/inference/basic/basic_wan2_2_ti2v.py"
},
{
"id": "matrix-game-2",
"task": "Interactive world",
"label": "Matrix Game 2.0",
"model": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"source": "examples/inference/basic/basic_matrixgame2.py",
"command": "python examples/inference/basic/basic_matrixgame2.py"
}
]
}
-60
View File
@@ -1,60 +0,0 @@
(() => {
let recipesPromise;
const loadRecipes = (url) => {
recipesPromise ||= fetch(url).then((response) => {
if (!response.ok) throw new Error(`HTTP ${response.status}`);
return response.json();
});
return recipesPromise;
};
const init = () => {
document.querySelectorAll("[data-cookbook]").forEach(async (root) => {
if (root.dataset.initialized) return;
root.dataset.initialized = "true";
const select = root.querySelector("[data-cookbook-recipe]");
const model = root.querySelector("[data-cookbook-model]");
const source = root.querySelector("[data-cookbook-source]");
const command = root.querySelector("[data-cookbook-command]");
const status = root.querySelector("[data-cookbook-status]");
try {
const { recipes } = await loadRecipes(root.dataset.recipes);
const byId = new Map(recipes.map((recipe) => [recipe.id, recipe]));
const groups = new Map();
select.replaceChildren();
recipes.forEach((recipe) => {
if (!groups.has(recipe.task)) {
const group = document.createElement("optgroup");
group.label = recipe.task;
groups.set(recipe.task, group);
select.append(group);
}
groups.get(recipe.task).append(new Option(recipe.label, recipe.id));
});
const render = () => {
const recipe = byId.get(select.value);
model.textContent = recipe.model;
source.textContent = recipe.source;
source.href = `https://github.com/hao-ai-lab/FastVideo/blob/main/${recipe.source}`;
command.textContent = recipe.command;
status.textContent = `${recipe.label} selected.`;
};
select.addEventListener("change", render);
select.disabled = false;
render();
} catch (error) {
status.textContent = "Recipes could not be loaded. Use the examples link below.";
console.error("Failed to load FastVideo cookbook recipes", error);
}
});
};
if (window.document$) window.document$.subscribe(init);
else document.addEventListener("DOMContentLoaded", init);
})();
-40
View File
@@ -42,46 +42,6 @@ img {
margin: 0 auto;
}
.cookbook-picker {
padding: 1rem;
border: 0.05rem solid var(--md-default-fg-color--lightest);
border-radius: 0.2rem;
}
.cookbook-picker select {
width: 100%;
padding: 0.6rem;
color: var(--md-default-fg-color);
background: var(--md-default-bg-color);
border: 0.05rem solid var(--md-default-fg-color--lighter);
border-radius: 0.2rem;
}
.cookbook-picker dl {
display: grid;
grid-template-columns: max-content 1fr;
gap: 0.25rem 1rem;
}
.cookbook-picker dt {
font-weight: 700;
}
.cookbook-picker dd {
margin: 0;
min-width: 0;
overflow-wrap: anywhere;
}
.cookbook-picker__status {
position: absolute;
width: 1px;
height: 1px;
overflow: hidden;
clip: rect(0, 0, 0, 0);
white-space: nowrap;
}
.md-typeset .copy-page-button.md-button {
float: right;
margin: 0 0 1rem 1rem;
-44
View File
@@ -1,44 +0,0 @@
# Inference Cookbook
Choose a complete recipe maintained in the FastVideo repository. Each command
runs its checked-in source directly, so coupled model, GPU, offload, and
attention settings do not drift into unsupported combinations.
The commands expect a local clone:
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo
```
<div class="cookbook-picker" data-cookbook data-recipes="../assets/cookbook-recipes.json">
<label for="cookbook-recipe"><strong>Recipe</strong></label>
<select id="cookbook-recipe" data-cookbook-recipe disabled>
<option>Loading recipes…</option>
</select>
<dl>
<dt>Model</dt>
<dd data-cookbook-model>Loading…</dd>
<dt>Source</dt>
<dd><a data-cookbook-source href="../inference/examples/basic/">Browse maintained examples</a></dd>
</dl>
<pre><code class="language-bash" data-cookbook-command>Loading…</code></pre>
<p class="cookbook-picker__status" role="status" aria-live="polite" data-cookbook-status></p>
<noscript>
JavaScript is needed for the recipe picker. Browse the
<a href="../inference/examples/examples_inference_index/">inference examples</a>
instead.
</noscript>
</div>
## Customize a recipe
Start from the checked-in source, then change only the settings your model
supports:
- [Configuration](../inference/configuration.md) covers the Python and CLI
config surfaces.
- [Optimizations](../inference/optimizations.md) covers attention backends,
compilation, and memory tradeoffs.
- [Support matrix](../inference/support_matrix.md) lists supported models and
optimizations.
-128
View File
@@ -1,128 +0,0 @@
# Fast mode (RIFE) — Apple Silicon
`--fast` makes local generation ~2.7× faster by **generating fewer frames and
interpolating the rest** with an Apple-Silicon-native RIFE model, instead of
denoising every frame. Video-diffusion denoise is dominated by self-attention,
which is O(tokens²); halving the frames cuts the token count ~2× and the denoise
compute ~3.7×, so the wall-clock drops far more than 2×. RIFE (which estimates
its own optical flow — no motion vectors needed) fills the dropped frames back
in for ~1.4 s, and a light unsharp pass counters its softening.
Measured on the 1.3B INT8 QAD model (fox, 480×832×81, M4): generate 41 + RIFE→81
runs in ~35 s of denoise vs ~90 s full, at reconstruction MS-SSIM **0.97**.
Reproduce with `python -m fastvideo.benchmarks.eval_metalfx_rife --mode int8`.
> **Note:** Apple's *MetalFX* frame interpolation is **not** usable here — it
> requires game-engine motion vectors + depth, which diffusion output lacks. We
> use the video-native **`rife-mlx`** model instead (Metal-backed, torch-free).
## Install
```bash
uv pip install -e ".[mlx]" # RIFE ships vendored; this only needs MLX
```
## Use
```bash
python examples/inference/basic/mlx_wan_prompt_to_video.py \
--mlx-checkpoint <FastWan2.1-T2V-1.3B-INT8-QAD> \
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
--num-frames 81 --fast \
--output-path video_samples/fox_fast.mp4
```
`--num-frames` stays the *target* length; fast mode generates the smallest
VAE-aligned keyframe count that RIFE can interpolate to that target.
| Flag | Default | Meaning |
|---|---|---|
| `--fast` / `--no-fast` | off | enable fast mode |
| `--fast-factor` | 2 | generate 1/factor of the frames (2 = half) |
| `--fast-sharpen` | 0.6 | light unsharp strength to counter RIFE softness (0 disables) |
Fast mode composes with everything else (`--mlx-quantization int8`,
`--mlx-compile`, TAEHV vs `--decode-backend wan-vae`). Keep `--fast-factor` at 2
for quality — larger temporal gaps are where RIFE starts inventing motion.
## Spatial fast mode (`--fast-spatial`)
The spatial twin of `--fast`: instead of dropping frames, drop pixels. Denoise
*and decode* at `height/width // fast-spatial-scale`, then resample the decoded
frames up to the requested size. Self-attention is O(tokens²), so halving each
spatial axis cuts the token count 4× and the denoise time far more than that —
measured on the 1.3B INT8 QAD model at 480×832×81, M4 Max: **86.1 s → 10.3 s**
of denoise. It composes with `--fast`; both together run the same clip in
**4.5 s** of denoise.
```bash
python examples/inference/basic/mlx_wan_prompt_to_video.py \
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
--height 480 --width 832 --num-frames 81 --fast-spatial \
--output-path video_samples/fox_fast_spatial.mp4
```
| Flag | Default | Meaning |
|---|---|---|
| `--fast-spatial` / `--no-fast-spatial` | off | enable spatial fast mode |
| `--fast-spatial-scale` | 2 | denoise at 1/scale of each spatial axis |
| `--fast-spatial-upsample-mode` | `lanczos` | pixel interpolation kernel (`lanczos`, `cubic`, `bilinear`, `nearest`) |
| `--fast-spatial-sharpen` | 0.4 | light unsharp strength to counter resampling softness (0 disables) |
### The upsample must happen in pixel space
This is the one thing to get right. The obvious implementation — bilinearly
upsample the finished latents and decode at the target size — **does not work**,
and produces a distinctive failure: correct composition and silhouette under a
smeared, hazy veil, with ringing along strong edges.
A Wan latent cell is a *learned code* for an 8×8 (Wan2.1) or 16×16 (Wan2.2)
pixel block, not a low-pass sample of the image. The average of two adjacent
codes is not the code of the averaged blocks; it is a vector the decoder was
never trained on. Measured on Wan2.1-1.3B at 480×832, a 2× bilinear latent
upsample destroys **62%** of the latent's high-frequency energy while leaving
its overall magnitude intact — exactly the signature of that veil. At Wan2.2-5B
the same operation degrades to black or noise.
Decoded RGB frames have no such problem: an image *is* a sampled 2-D signal, so
Lanczos interpolation is the operation it was defined for. The result is soft —
it carries stage-1's real detail budget and no more — but clean and coherent.
`--refine` gets away with a latent-space upsample only because a second DMD pass
re-denoises the hand-off; spatial fast mode passes the latent straight to the
decoder, so it cannot.
## Refine (`--refine`) stage-2 timesteps
`--refine` hands stage 1 to stage 2 as `(1 - sigma) * upsampled + sigma * noise`,
where `sigma` comes from the *first* stage-2 timestep. FastWan's DMD grid opens
at `t=1000`, which is `sigma == 1` exactly — so a stage-2 grid that starts there
weights the stage-1 result at zero and refine silently degrades into a plain
full-resolution run at twice the cost.
Left unset, `--refine-dmd-denoising-steps` now derives the stage-2 grid from the
stage-1 one with leading full-noise steps dropped (`1000,757,522` → `757,522`).
That keeps the pass on timesteps the distilled student was trained on while
letting stage-1 structure through: hand-off `sigma = 0.757`, stage-1 weight
`0.243`. Passing a grid that starts at full noise is now an error rather than a
silently wasted pass.
The run prints the resolved hand-off so it is visible:
```
[refine] stage-2 hand-off sigma=0.7568 (stage-1 weight 0.2432)
```
There is a trade-off in choosing that grid. Later start = more of the draft
survives, but fewer stage-2 steps. On Wan2.1 the default `757,522` gives weight
0.243 with two steps; `--refine-dmd-denoising-steps 522` gives weight 0.478 with
one. `--refine-sigma` decouples the noise level from the timestep entirely — it
logs a warning, because the DiT is then told a timestep that does not match the
noise it receives.
**Wan2.2-5B has a lower ceiling.** Its warped schedule maps `1000,757,522` to
sigmas `1.000, 0.940, 0.845`, so the best available stage-1 weight is **0.060**
(vs 0.243 at 1.3B). Un-warped (`--no-warp`) the same grid gives `1.000, 0.757,
0.522` and a weight of 0.243 — but warping is what matches the FastVideo
sampling schedule, so turning it off changes the timesteps the distilled student
sees. Which is better at 5B is unresolved and needs a run on real 5B weights.
@@ -76,10 +76,6 @@ surfaces:
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
VSA_tile_size: "VSA-H3 tile geometry request; model-specific optimization not yet represented in the typed public schema."
vae_parallel_decode: "MiniMax-H3 sequence-parallel VAE decode opt-in; model-specific optimization not yet represented in the typed public schema."
vae_parallel_encode: "MiniMax-H3 sequence-parallel reference VAE encode opt-in; model-specific optimization not yet represented in the typed public schema."
vae_parallel_decode_strategy: "Chunk-transport collective for vae_parallel_decode; model-specific optimization not yet represented in the typed public schema."
attention_backend: "Process-wide default attention-backend request applied per component at load time; kernel-selection knob not yet represented in the typed public schema."
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
-37
View File
@@ -2,7 +2,6 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import json
import os
import re
from dataclasses import dataclass, field
@@ -20,40 +19,6 @@ GENERATED_DOC_PREFIXES = (
"training/examples/",
"distillation/examples/",
)
COOKBOOK_DATA = ROOT_DIR / "docs/assets/cookbook-recipes.json"
COOKBOOK_SOURCE_ROOTS = (
ROOT_DIR / "examples/inference",
ROOT_DIR / "scripts/inference",
)
def validate_cookbook() -> None:
"""Keep cookbook entries tied to checked-in runnable sources."""
recipes = json.loads(COOKBOOK_DATA.read_text(encoding="utf-8")).get("recipes")
if not isinstance(recipes, list) or not recipes:
raise ValueError(f"{COOKBOOK_DATA}: recipes must be a non-empty list")
seen: set[str] = set()
for recipe in recipes:
required = ("id", "task", "label", "model", "source", "command")
missing = {key for key in required if not recipe.get(key)}
if missing:
raise ValueError(f"Cookbook recipe is missing: {', '.join(sorted(missing))}")
if recipe["id"] in seen:
raise ValueError(f"Duplicate cookbook recipe id: {recipe['id']}")
seen.add(recipe["id"])
source = (ROOT_DIR / recipe["source"]).resolve()
if not any(source.is_relative_to(root.resolve()) for root in COOKBOOK_SOURCE_ROOTS):
raise ValueError(f"Cookbook source is outside an approved directory: {recipe['source']}")
if not source.is_file():
raise ValueError(f"Cookbook source does not exist: {recipe['source']}")
source_text = source.read_text(encoding="utf-8")
if recipe["model"] not in source_text:
raise ValueError(f"Cookbook model is not present in {recipe['source']}: {recipe['model']}")
if recipe["source"] not in recipe["command"]:
raise ValueError(f"Cookbook command does not invoke its source: {recipe['id']}")
def fix_case(text: str) -> str:
@@ -571,7 +536,6 @@ def on_pre_build(config, **kwargs):
MkDocs hook to generate examples before building the documentation.
This function is called automatically by MkDocs' native hook system.
"""
validate_cookbook()
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
@@ -585,7 +549,6 @@ def on_page_context(context, page, **kwargs):
if __name__ == "__main__":
validate_cookbook()
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
+2 -3
View File
@@ -65,7 +65,6 @@ uv pip install flash-attn --no-build-isolation -v
## Next Steps
- [Quick Start](quick_start.md) - Generate your first video
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
+4 -6
View File
@@ -49,12 +49,10 @@ brew install ffmpeg
### Installation
FastWan's native Apple Silicon runtime requires the `mlx` extra.
#### With uv (recommended)
```bash
uv pip install "fastvideo[mlx]"
uv pip install fastvideo
```
#### With Conda environment (alternative)
@@ -62,7 +60,7 @@ uv pip install "fastvideo[mlx]"
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
```bash
uv pip install "fastvideo[mlx]"
uv pip install fastvideo
```
### Installation from Source
@@ -78,13 +76,13 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Basic installation:
```bash
uv pip install -e ".[mlx]"
uv pip install -e .
```
Alternative with Conda environment:
```bash
uv pip install -e ".[mlx]"
uv pip install -e .
```
## Development Environment Setup
+49 -9
View File
@@ -23,21 +23,61 @@ Also optionally install flash-attn:
uv pip install flash-attn --no-build-isolation -v
```
## Choose a maintained recipe
## Basic Usage
The cookbook selects complete, checked-in recipes instead of mixing model,
parallelism, offload, and attention settings independently.
### Text-to-Video Generation
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
```python
from fastvideo import VideoGenerator
!!! tip "Need more control?"
Start from a maintained recipe, then use the
[configuration](../inference/configuration.md) and
[optimization](../inference/optimizations.md) guides for supported changes.
def main():
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate the video
video = generator.generate_video(
prompt,
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
```
### Image-to-Video Generation
```python
from fastvideo import VideoGenerator, SamplingParam
def main():
# Create the generator
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
generator.generate_video(prompt, sampling_param=sampling_param,
output_path="my_videos/",
save_video=True)
if __name__ == '__main__':
main()
```
## Next Steps
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Installation Guide](installation.md) - Detailed installation instructions
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
-16
View File
@@ -64,7 +64,6 @@ 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) |
@@ -95,10 +94,6 @@ column links a runnable script in `examples/inference/basic/` where one exists.
(`StableAudioT2AConfig` / `StableAudioOpenSmallConfig`); they are registered
under the generic T2V workload option in the registry.
**Note (MMAudio)**: the registered Hugging Face model ID is reserved but not
yet public. Follow the [MMAudio inference guide](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/pipelines/basic/mmaudio/README.md)
to convert the official weights locally and set `MMAUDIO_MODEL_PATH`.
**Note (MiniMax H3)**: T2VA, FL2VA, and Ref2VA all generate video with stereo
audio. Use the Ref2VA example when passing ordered image, video, or audio
references.
@@ -178,17 +173,6 @@ optimizations: absence means **untested**, not incompatible.
| Matrix Game 3.0 Base Distilled | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | 720x1280 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| GEN3C Cosmos 7B | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | 704px1280p | ❌ | ❌ | ❌ | ⭕ | ⭕ |
## Apple Silicon native runtime
| Release path | Model | Mode | Validated hardware | Status |
| --- | --- | --- | --- | --- |
| MLX FastWan T2V | FastWan-QAD-INT8-1.3B `[release model ID pending]` | 480x832, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 36 GB unified-memory class, MLX 0.31.2 | Release candidate; requires release-owner visual sign-off |
This is a text-to-video-only source-install release. It is validated on the
hardware listed above; MLX allocator caps are not evidence of support for a
physical 16 GB Mac. See [Apple Silicon FastWan](../getting_started/installation/mps.md)
for the supported command and release gates.
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
***Lucy Edit Dev uses a non-commercial model license. FastVideo support is
-8
View File
@@ -33,14 +33,6 @@ For the typed config/request path added during the inference API refactor:
python examples/inference/basic/basic_dmd_new_api.py
```
For the few-step (4-step, DMD2-distilled) MiniMax-H3 preview, generating synchronized video and audio, optionally with block-sparse VSA attention:
```
python examples/inference/basic/basic_fasth3.py --prompt "your prompt" [--vsa-sparsity 0.9]
```
The default checkpoint `FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1` is private on the Hub while its license review completes; until it flips public, pass `--model-path` with a local snapshot of the release.
On Blackwell (sm_100) GPUs with a `fastvideo-kernel` build that carries the sm_100a block-sparse extension, `--vsa-kernel sm100a` routes the tile-64 attention forwards through the CUDA kernel instead of Triton (it sets `FASTVIDEO_VSA_SM100A=1` before the pipeline boots); if the extension or the arch is missing, the run warns once and falls back to Triton.
## Basic Walkthrough
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
+13 -12
View File
@@ -3,8 +3,6 @@ from fastvideo import VideoGenerator
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -14,11 +12,11 @@ def main():
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
@@ -26,19 +24,22 @@ def main():
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
@@ -30,7 +30,8 @@ def main():
"and casting reflections onto adjacent vehicles. "
"The motion creates space in the lineup, signaling activity within the otherwise quiet station. "
"It then comes to a smooth stop, resuming its position in line. "
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene.")
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene."
)
generator.generate_video(
prompt,
@@ -46,3 +47,4 @@ def main():
if __name__ == "__main__":
main()
@@ -31,7 +31,8 @@ def main():
"The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. "
"The metal surface beneath the torch shows ongoing signs of heating and melting. "
"The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, "
"underscoring the ongoing nature of the welding operation.")
"underscoring the ongoing nature of the welding operation."
)
generator.generate_video(
prompt,
@@ -45,3 +46,6 @@ def main():
if __name__ == "__main__":
main()
@@ -50,3 +50,4 @@ def main():
if __name__ == "__main__":
main()
+10 -9
View File
@@ -5,8 +5,6 @@ from fastvideo import VideoGenerator
from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_dmd2"
def main():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
@@ -16,10 +14,10 @@ def main():
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
# Adjust these offload parameters if you have < 32GB of VRAM
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
dit_cpu_offload=False,
vae_cpu_offload=False,
VSA_sparsity=0.8,
@@ -27,6 +25,7 @@ def main():
load_end_time = time.perf_counter()
load_time = load_end_time - load_start_time
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.num_frames = 81
@@ -40,16 +39,18 @@ def main():
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
start_time = time.perf_counter()
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
end_time = time.perf_counter()
gen_time2 = end_time - start_time
print(f"Time taken to load model: {load_time} seconds")
print(f"Time taken to generate video: {gen_time} seconds")
print(f"Time taken to generate video2: {gen_time2} seconds")
+20 -14
View File
@@ -32,9 +32,11 @@ def main():
),
# PR 2 still routes a few advanced inference knobs through the
# compatibility bridge until they get first-class typed fields.
pipeline=PipelineSelection(experimental={
"VSA_sparsity": 0.8,
}, ),
pipeline=PipelineSelection(
experimental={
"VSA_sparsity": 0.8,
},
),
)
load_start_time = time.perf_counter()
@@ -42,12 +44,14 @@ def main():
load_end_time = time.perf_counter()
load_time = load_end_time - load_start_time
prompt = ("A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
"LED umbrella. Steam rises from a street food cart, and a cat darts "
"across the screen. Raindrops are visible on the camera lens, creating "
"a cinematic bokeh effect.")
prompt = (
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. "
"The puddles reflect glowing signs in kanji, advertising ramen, karaoke, "
"and VR arcades. A woman in a translucent raincoat walks briskly with an "
"LED umbrella. Steam rises from a street food cart, and a cat darts "
"across the screen. Raindrops are visible on the camera lens, creating "
"a cinematic bokeh effect."
)
request = GenerationRequest(
prompt=prompt,
output=OutputConfig(
@@ -62,11 +66,13 @@ def main():
end_time = time.perf_counter()
gen_time = end_time - start_time
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently "
"in the breeze, enhancing the lion's commanding presence. The tone is "
"vibrant, embodying the raw energy of the wild. Low angle, steady "
"tracking shot, cinematic.")
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently "
"in the breeze, enhancing the lion's commanding presence. The tone is "
"vibrant, embodying the raw energy of the wild. Low angle, steady "
"tracking shot, cinematic."
)
request2 = GenerationRequest(
prompt=prompt2,
output=OutputConfig(
@@ -2,6 +2,7 @@ import os
from fastvideo import VideoGenerator
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
@@ -45,8 +46,10 @@ def main():
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
"action_speed_list":
[float(value) for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")],
"action_speed_list": [
float(value)
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
],
}
if image_path:
kwargs["image_path"] = image_path
-180
View File
@@ -1,180 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1 is a 4-step distillation of
MiniMaxAI/MiniMax-H3 (data-free DMD2): it walks a 4-step grid on the release
sampler's shift-12 schedule instead of the base model's 50 steps, generating
synchronized video and audio in one pipeline call.
The student was trained with block-sparse video attention (VSA, 64-token
tiles) and its checkpoint carries the trained sparse-gate parameters
(``attn.to_gate_compress``), so this script always runs the VSA-H3 attention
backend. At the default ``--vsa-sparsity 0.0`` the attention math is exactly
dense (every tile is selected); raise the sparsity for additional speedup.
"""
from __future__ import annotations
import argparse
import os
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
CompileConfig,
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default="FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1")
# The HF repo is private while the MiniMax H3 Community License review
# completes; until it flips public, pass --model-path with a local
# snapshot of the release instead (e.g. the team export at
# /mnt/lustre/vlm-wlsaidhi/fastvideo/exports/FastVideo-Minimax-FastH3-Preview-v0.1).
parser.add_argument("--prompt", required=True)
parser.add_argument("--output", default="outputs/fasth3")
parser.add_argument("--height", type=int, default=768)
parser.add_argument("--width", type=int, default=1344)
parser.add_argument("--num-frames", type=int, default=124)
# num_inference_steps counts sigma-GRID POINTS, matching the base model's
# convention ("50 steps" = a 50-point grid = 49 transformer forwards). The
# student's distilled 4-step grid is 4 FORWARDS, i.e. a 5-point grid
# (t = 1000, 750, 500, 250 -> 0 on the shift-12 schedule) — so the correct
# default here is 5. Other grids are off-distribution.
parser.add_argument("--steps",
type=int,
default=5,
help="num_inference_steps = sigma-grid points; N points run N-1 denoising "
"forwards. 5 (default) is the distilled 4-forward grid")
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--num-gpus", type=int, default=4)
parser.add_argument("--vsa-sparsity",
type=float,
default=0.0,
help="Run-level VSA sparsity in [0, 1). 0.0 (default) selects every tile, which is "
"exactly dense attention; the student was trained at 0.9")
# 64 is the trained contract: the student was TRAINED with 64-token
# (4,4,4) tiles, and its to_gate_compress gates were learned against
# pooling at that granularity — keep 64 unless you are ablating.
parser.add_argument("--vsa-tile-size",
type=int,
choices=(64, 256),
default=64,
help="VSA-H3 tile size in tokens; 64 (default) is what the student was trained "
"with and runs the native Triton block-sparse path, 256 is the FA4-CuTe-capable "
"geometry for ablations")
parser.add_argument("--vsa-kernel",
choices=("triton", "sm100a"),
default="triton",
help="Block-sparse kernel for the tile-64 attention forward: triton (default, "
"portable fwd+bwd) or sm100a — the opt-in Blackwell CUDA forward "
"(fastvideo_kernel.block_sparse_attn_sm100a). sm100a needs an sm_100 GPU and a "
"fastvideo-kernel build that carries the extension; if a precondition fails at "
"run time the attention layer logs one warning and falls back to Triton. Only "
"meaningful with --vsa-tile-size 64")
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--repeats",
type=int,
default=1,
help="generate N times; with --torch-compile the first run pays "
"compilation, so steady-state is the last repeat")
return parser.parse_args()
def main() -> None:
args = parse_args()
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
if args.vsa_kernel == "sm100a":
# The attention backend reads FASTVIDEO_VSA_SM100A per forward; set it
# before the pipeline boots so spawned GPU workers inherit it. The
# kernel is forward-only and inference runs under no-grad, so every
# denoising forward qualifies for the CUDA route.
os.environ["FASTVIDEO_VSA_SM100A"] = "1"
# Boot-time run configuration, folded into FastVideoArgs (the same route
# examples/inference/basic/basic_minimax_h3_t2v.py uses for sparsity):
# - attention_backend: the checkpoint carries trained to_gate_compress
# gates, which only exist under the VSA-H3 backend — a dense-backend
# load would reject them as unexpected weights. Layers that do not
# support VSA-H3 (e.g. the token refiner) fall back to flash attention.
# - VSA_tile_size: forwarded even at sparsity 0.0 because the gate-compress
# branch pools per tile, and the gates were trained at 64 tokens/tile.
experimental: dict[str, object] = {
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
"VSA_tile_size": args.vsa_tile_size,
}
if args.vsa_sparsity > 0.0:
experimental["VSA_sparsity"] = args.vsa_sparsity
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
pipeline=PipelineSelection(experimental=experimental),
engine=EngineConfig(
num_gpus=args.num_gpus,
use_fsdp_inference=args.num_gpus > 1,
parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus),
offload=OffloadConfig(
dit=False,
dit_layerwise=False,
text_encoder=True,
vae=True,
pin_cpu_memory=False,
),
compile=CompileConfig(
enabled=args.torch_compile,
mode=args.compile_mode,
),
),
))
try:
request = GenerationRequest(
prompt=args.prompt,
negative_prompt="",
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=args.num_frames,
fps=24,
num_inference_steps=args.steps,
# the base model is guidance-distilled; the student inherits it
guidance_scale=1.0,
batch_cfg=False,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output_dir / "fasth3.mp4"),
save_video=True,
return_frames=False,
),
)
result = generator.generate(request)
print(f"Output written to: {result.video_path}")
if result.generation_time is not None:
# machine-readable: benchmark harnesses parse this line to separate
# generation from model-load time (last occurrence = steady state)
print(f"Generation time: {result.generation_time:.2f}s")
for _ in range(args.repeats - 1):
result = generator.generate(request)
if result.generation_time is not None:
print(f"Generation time: {result.generation_time:.2f}s")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+6 -2
View File
@@ -64,8 +64,12 @@ def main() -> None:
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
tp_size = args.tp_size if args.tp_size is not None else (
args.num_gpus if args.num_gpus > 1 else 1
)
sp_size = args.sp_size if args.sp_size is not None else (
1 if args.num_gpus > 1 else args.num_gpus
)
generator_config = GeneratorConfig(
model_path=args.model_path,
@@ -21,6 +21,7 @@ from fastvideo.api import (
SamplingConfig,
)
DEFAULT_PROMPT = "a brushed steel espresso machine on a marble counter, morning window light"
+11 -5
View File
@@ -9,9 +9,11 @@ import re
DEFAULT_PROMPTS = [
"a photo of a cat",
("a cinematic photo of a red panda wearing a tiny backpack, standing on a "
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
"35mm, bokeh"),
(
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
"35mm, bokeh"
),
]
@@ -40,7 +42,9 @@ def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.", )
p = argparse.ArgumentParser(
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
)
p.add_argument(
"--model-path",
default="official_weights/FLUX.1-dev",
@@ -104,7 +108,9 @@ def main() -> None:
try:
for i, prompt in enumerate(prompts):
seed = args.seed + i
filename_base = (f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}")
filename_base = (
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
)
_remove_existing_outputs(args.out_dir, filename_base)
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
+11 -10
View File
@@ -33,20 +33,21 @@ MODEL_PATH = os.environ.get("GAMECRAFT_MODEL_PATH", "FastVideo/HunyuanGameCraft-
# Default prompts for demo
DEFAULT_PROMPTS = {
"village":
"A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
"temple":
"A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
"forest":
"A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
"village": "A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
"temple": "A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
"forest": "A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
"beach": "A tropical beach with crystal clear turquoise water, white sand, and palm trees swaying in the breeze.",
}
# I2V: default reference image (URL). Can override with a local path.
DEFAULT_I2V_IMAGE_URL = ("https://huggingface.co/datasets/huggingface/documentation-images/"
"resolve/main/diffusers/astronaut.jpg")
DEFAULT_I2V_PROMPT = ("An astronaut hatching from an egg, on the surface of the moon, "
"the darkness and depth of space realised in the background.")
DEFAULT_I2V_IMAGE_URL = (
"https://huggingface.co/datasets/huggingface/documentation-images/"
"resolve/main/diffusers/astronaut.jpg"
)
DEFAULT_I2V_PROMPT = (
"An astronaut hatching from an egg, on the surface of the moon, "
"the darkness and depth of space realised in the background."
)
OUTPUT_PATH = "video_samples_gamecraft"
+32 -16
View File
@@ -26,35 +26,51 @@ from fastvideo import VideoGenerator
def main():
parser = argparse.ArgumentParser(description="GEN3C video generation")
parser.add_argument("--model_path", type=str, default="converted_weights/GEN3C-Cosmos-7B")
parser.add_argument("--image_path", type=str, default=None, help="Input image for 3D cache conditioning")
parser.add_argument("--prompt", type=str, default="A slow camera pan over a sunlit landscape.")
parser.add_argument("--model_path",
type=str,
default="converted_weights/GEN3C-Cosmos-7B")
parser.add_argument("--image_path",
type=str,
default=None,
help="Input image for 3D cache conditioning")
parser.add_argument("--prompt",
type=str,
default="A slow camera pan over a sunlit landscape.")
parser.add_argument(
"--negative_prompt",
type=str,
default=("The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
"flickering. Overall, the video is of poor quality."),
default=(
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special "
"effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and "
"flickering. Overall, the video is of poor quality."
),
)
parser.add_argument(
"--trajectory",
type=str,
default="left",
choices=["left", "right", "up", "down", "zoom_in", "zoom_out", "clockwise", "counterclockwise", "none"])
parser.add_argument("--trajectory",
type=str,
default="left",
choices=[
"left", "right", "up", "down", "zoom_in",
"zoom_out", "clockwise", "counterclockwise", "none"
])
parser.add_argument("--movement_distance", type=float, default=0.3)
parser.add_argument("--camera_rotation",
type=str,
default="center_facing",
choices=["center_facing", "no_rotation", "trajectory_aligned"])
choices=[
"center_facing", "no_rotation",
"trajectory_aligned"
])
parser.add_argument("--height", type=int, default=704)
parser.add_argument("--width", type=int, default=1280)
parser.add_argument("--num_frames", type=int, default=121)
parser.add_argument("--num_inference_steps", type=int, default=35)
parser.add_argument("--guidance_scale", type=float, default=1.0)
parser.add_argument("--output_path", type=str, default="outputs_video/gen3c.mp4")
parser.add_argument("--output_path",
type=str,
default="outputs_video/gen3c.mp4")
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
+16 -25
View File
@@ -3,8 +3,6 @@ import json
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_hy15"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -14,38 +12,31 @@ def main():
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt,
output_path=OUTPUT_PATH,
save_video=True,
negative_prompt="",
num_frames=81,
fps=16)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2,
output_path=OUTPUT_PATH,
save_video=True,
negative_prompt="",
num_frames=81,
fps=16)
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="", num_frames=81, fps=16)
if __name__ == "__main__":
main()
main()
+14 -13
View File
@@ -3,37 +3,38 @@ import json
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_hy15_1080p"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
# or "weizhou03/HunyuanVideo-1.5-Diffusers-1080p" # 720p -> 1080p
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
@@ -6,8 +6,6 @@ DEFAULT_PROMPT = 'A paved pathway leads towards a stone arch bridge spanning a c
DEFAULT_IMAGE = 'https://raw.githubusercontent.com/Tencent-Hunyuan/HY-WorldPlay/main/assets/img/test.png'
OUTPUT_PATH = "video_samples_hyworld"
def main():
import argparse
@@ -7,7 +7,7 @@ IMAGE_PATH = "assets/girl.png"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
num_gpus=1,
@@ -19,7 +19,9 @@ def main():
# image_encoder_cpu_offload=False,
)
prompt = ("A woman stands up and walks away")
prompt = (
"A woman stands up and walks away"
)
_ = generator.generate_video(
prompt,
image_path=IMAGE_PATH,
@@ -2,7 +2,6 @@ from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
@@ -18,28 +17,21 @@ def main():
# image_encoder_cpu_offload=False,
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
_ = generator.generate_video(prompt,
output_path=OUTPUT_PATH,
save_video=True,
height=512,
width=768,
num_frames=121)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2,
output_path=OUTPUT_PATH,
save_video=True,
height=512,
width=768,
num_frames=121)
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
if __name__ == "__main__":
main()
main()
@@ -6,6 +6,7 @@ from pathlib import Path
from fastvideo import VideoGenerator
REPO_ROOT = Path(__file__).resolve().parents[3]
DATASET_DIR = REPO_ROOT / "examples" / "dataset" / "lingbotworld2"
OUTPUT_PATH = REPO_ROOT / "outputs" / "lingbotworld2_causal_fast.mp4"
@@ -3,19 +3,17 @@ from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embeddin
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_lingbotworld"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"FastVideo/LingBot-World-Base-Cam-Diffusers",
"FastVideo/LingBot-World-Base-Cam-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
+34 -28
View File
@@ -21,16 +21,20 @@ import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = ("A woman sits at a wooden table by the window in a cozy café. She reaches out "
"with her right hand, picks up the white coffee cup from the saucer, and gently "
"brings it to her lips to take a sip. After drinking, she places the cup back on "
"the table and looks out the window, enjoying the peaceful atmosphere.")
PROMPT = (
"A woman sits at a wooden table by the window in a cozy café. She reaches out "
"with her right hand, picks up the white coffee cup from the saucer, and gently "
"brings it to her lips to take a sip. After drinking, she places the cup back on "
"the table and looks out the window, enjoying the peaceful atmosphere."
)
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards")
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
# Input image path
IMAGE_PATH = "assets/girl.png"
@@ -47,20 +51,20 @@ def basic_generation():
print("=" * 60)
print("LongCat I2V: Basic Generation (50 steps, 480p)")
print("=" * 60)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-I2V-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_i2v_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -75,7 +79,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -90,11 +94,11 @@ def distill_refine_generation():
print("\n" + "=" * 60)
print("LongCat I2V: Distill + Refine Pipeline")
print("=" * 60)
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-I2V-Diffusers",
num_gpus=1,
@@ -107,9 +111,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_i2v_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -124,14 +128,14 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 768p)
print("\n[Stage 2] Refinement (480p -> 768p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
@@ -139,7 +143,7 @@ def distill_refine_generation():
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
# Note: Refinement uses the T2V model (not I2V) since it's upscaling the generated video
# For BSA [4, 4, 8]: latent must be divisible by 8
@@ -159,9 +163,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_i2v_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -178,7 +182,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -188,13 +192,13 @@ def main():
print("\n" + "=" * 60)
print("LongCat Image-to-Video Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
@@ -202,3 +206,5 @@ def main():
if __name__ == "__main__":
main()
+36 -30
View File
@@ -15,18 +15,22 @@ import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = ("In a realistic photography style, a white boy around seven or eight years old "
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
"features a green lawn and several tall trees, creating a warm and loving scene.")
PROMPT = (
"In a realistic photography style, a white boy around seven or eight years old "
"sits on a park bench, wearing a light blue T-shirt, denim shorts, and white sneakers. "
"He holds an ice cream cone with vanilla and chocolate flavors, and beside him is a "
"medium-sized golden Labrador. Smiling, the boy offers the ice cream to the dog, "
"who eagerly licks it with its tongue. The sun is shining brightly, and the background "
"features a green lawn and several tall trees, creating a warm and loving scene."
)
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards")
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
SEED = 42
@@ -40,20 +44,20 @@ def basic_generation():
print("=" * 60)
print("LongCat T2V: Basic Generation (50 steps, 480p)")
print("=" * 60)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_t2v_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -67,7 +71,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -82,11 +86,11 @@ def distill_refine_generation():
print("\n" + "=" * 60)
print("LongCat T2V: Distill + Refine Pipeline")
print("=" * 60)
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
num_gpus=1,
@@ -99,9 +103,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_t2v_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -115,14 +119,14 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 720p)
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
@@ -130,7 +134,7 @@ def distill_refine_generation():
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
refine_generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-T2V-Diffusers",
@@ -147,9 +151,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_t2v_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -166,7 +170,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -176,13 +180,13 @@ def main():
print("\n" + "=" * 60)
print("LongCat Text-to-Video Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
@@ -190,3 +194,5 @@ def main():
if __name__ == "__main__":
main()
+45 -35
View File
@@ -21,17 +21,21 @@ import os
from fastvideo import VideoGenerator
# Common prompts and settings matching the shell script examples
PROMPT = ("A person rides a motorcycle along a long, straight road that stretches between "
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
"the motorcycle centered between the guardrails, while the scenery passes by on "
"both sides. The video captures the journey from the rider's perspective, emphasizing "
"the sense of motion and adventure.")
PROMPT = (
"A person rides a motorcycle along a long, straight road that stretches between "
"a body of water and a forested hillside. The rider steadily accelerates, keeping "
"the motorcycle centered between the guardrails, while the scenery passes by on "
"both sides. The video captures the journey from the rider's perspective, emphasizing "
"the sense of motion and adventure."
)
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards")
NEGATIVE_PROMPT = (
"Bright tones, overexposed, static, blurred details, subtitles, style, works, "
"paintings, images, static, overall gray, worst quality, low quality, JPEG compression "
"residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, "
"deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, "
"three legs, many people in the background, walking backwards"
)
# Input video path
VIDEO_PATH = "assets/motorcycle.mp4"
@@ -51,25 +55,27 @@ def basic_generation():
print("=" * 60)
print("LongCat VC: Basic Generation (50 steps, 480p)")
print("=" * 60)
# Check if video exists
if not os.path.exists(VIDEO_PATH):
raise FileNotFoundError(f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path.")
raise FileNotFoundError(
f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path."
)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-VC-Diffusers",
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
enable_bsa=False,
)
output_path = "outputs_video/longcat_vc_basic"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -85,7 +91,7 @@ def basic_generation():
guidance_scale=4.0,
seed=SEED,
)
print(f"\nBasic generation complete! Video saved to: {output_path}")
generator.shutdown()
@@ -100,16 +106,18 @@ def distill_refine_generation():
print("\n" + "=" * 60)
print("LongCat VC: Distill + Refine Pipeline")
print("=" * 60)
# Check if video exists
if not os.path.exists(VIDEO_PATH):
raise FileNotFoundError(f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path.")
raise FileNotFoundError(
f"Video not found at {VIDEO_PATH}. "
"Please provide a valid video path."
)
# Stage 1: Distilled generation (16 steps at 480p)
print("\n[Stage 1] Distilled generation (16 steps, 480p)")
print("-" * 40)
generator = VideoGenerator.from_pretrained(
"FastVideo/LongCat-Video-VC-Diffusers",
num_gpus=1,
@@ -122,9 +130,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Distilled-LoRA",
lora_nickname="distilled",
)
distill_output_path = "outputs_video/longcat_vc_distill"
generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -140,14 +148,14 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Distilled generation complete! Video saved to: {distill_output_path}")
generator.shutdown()
# Stage 2: Refinement (480p -> 720p)
print("\n[Stage 2] Refinement (480p -> 720p with BSA)")
print("-" * 40)
# Find the actual saved video file from stage 1
video_files = glob.glob(os.path.join(distill_output_path, "*.mp4"))
if not video_files:
@@ -155,7 +163,7 @@ def distill_refine_generation():
# Use the most recently created video file
distill_video_path = max(video_files, key=os.path.getmtime)
print(f"Using stage 1 video: {distill_video_path}")
# Create a new generator with refinement LoRA and BSA enabled
# Note: Refinement uses the T2V model (not VC) since it's upscaling the generated video
refine_generator = VideoGenerator.from_pretrained(
@@ -173,9 +181,9 @@ def distill_refine_generation():
lora_path="FastVideo/LongCat-Video-T2V-Refinement-LoRA",
lora_nickname="refinement",
)
refine_output_path = "outputs_video/longcat_vc_refine_720p"
refine_generator.generate_video(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
@@ -192,7 +200,7 @@ def distill_refine_generation():
guidance_scale=1.0,
seed=SEED,
)
print(f"Refinement complete! Video saved to: {refine_output_path}")
refine_generator.shutdown()
@@ -202,13 +210,13 @@ def main():
print("\n" + "=" * 60)
print("LongCat Video Continuation Example")
print("=" * 60 + "\n")
# Run basic generation
basic_generation()
# Run distill+refine pipeline
distill_refine_generation()
print("\n" + "=" * 60)
print("All generations complete!")
print("=" * 60)
@@ -216,3 +224,5 @@ def main():
if __name__ == "__main__":
main()
+15 -12
View File
@@ -1,16 +1,19 @@
from fastvideo import VideoGenerator
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
def main() -> None:
@@ -33,4 +36,4 @@ def main() -> None:
if __name__ == "__main__":
main()
main()
@@ -67,19 +67,25 @@ _inductor.coordinate_descent_tuning = True
_inductor.coordinate_descent_check_all_directions = True
_inductor.epilogue_fusion = False
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX23_MODEL_PATH",
"FastVideo/LTX-2.3-Distilled-Diffusers")))
OUTPUT_DIR = Path(os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v"))
MODEL_ID = os.path.expandvars(
os.path.expanduser(
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
)
)
OUTPUT_DIR = Path(
os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v")
)
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
DEFAULT_PROMPT = ("A fashion model takes a slow step forward and shifts her weight, "
"the soft fabric of her clothing swaying and rippling with the "
"motion, her hair shifting gently, soft even studio lighting on a "
"clean light background, elegant slow-motion runway feel.")
DEFAULT_PROMPT = (
"A fashion model takes a slow step forward and shifts her weight, "
"the soft fabric of her clothing swaying and rippling with the "
"motion, her hair shifting gently, soft even studio lighting on a "
"clean light background, elegant slow-motion runway feel."
)
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
# Per-stage timing helpers --------------------------------------------------
def _print_stage_breakdown(result: dict, label: str) -> float | None:
"""Print stage execution times and return the sum, or None if missing."""
logging_info = result.get("logging_info")
@@ -108,7 +114,9 @@ def _collect_stage_times(
return
for name, metrics in stages.items():
stage_order.setdefault(name, None)
stage_times.setdefault(name, []).append(float(metrics.get("execution_time", 0.0)))
stage_times.setdefault(name, []).append(
float(metrics.get("execution_time", 0.0))
)
def _resolve_refine_upsampler(model_root: str) -> Path:
@@ -117,18 +125,21 @@ def _resolve_refine_upsampler(model_root: str) -> Path:
cand = Path(model_root) / name
if (cand / "config.json").is_file():
return cand
raise FileNotFoundError(f"No refine upsampler directory under {model_root}. "
f"Expected `{model_root}/spatial_upscaler/config.json`.")
raise FileNotFoundError(
f"No refine upsampler directory under {model_root}. "
f"Expected `{model_root}/spatial_upscaler/config.json`."
)
# Main ---------------------------------------------------------------------
def main() -> None:
if not I2V_IMAGE:
raise SystemExit("LTX23_I2V_IMAGE is required for i2v. Example:\n"
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py")
raise SystemExit(
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
" python examples/inference/basic/basic_ltx2_3_distilled_i2v.py"
)
if not Path(I2V_IMAGE).is_file():
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
@@ -190,13 +201,11 @@ def main() -> None:
common_kwargs = dict(
prompt=PROMPT,
negative_prompt="", # distilled is CFG-free; no negative needed
guidance_scale=1.0, # CFG=1 for distilled
height=1280,
width=832, # portrait runway aspect
num_frames=121,
fps=24, # ~5s clip
num_inference_steps=8, # distilled denoise steps
negative_prompt="", # distilled is CFG-free; no negative needed
guidance_scale=1.0, # CFG=1 for distilled
height=1280, width=832, # portrait runway aspect
num_frames=121, fps=24, # ~5s clip
num_inference_steps=8, # distilled denoise steps
# i2v: anchor the input image at frame 0 with full strength.
# `ltx2_image_crf=0.0` skips an extra JPEG re-encode of an already
# JPEG conditioning image.
@@ -242,7 +251,10 @@ def main() -> None:
**common_kwargs,
)
wall = time.perf_counter() - t0
e2e = (result.get("e2e_latency") if isinstance(result, dict) else None) or wall
e2e = (
result.get("e2e_latency")
if isinstance(result, dict) else None
) or wall
measured_secs.append(e2e)
print(f"[measured {m + 1}/{measured_runs}] e2e={e2e:.2f}s wall={wall:.2f}s")
if isinstance(result, dict):
@@ -254,8 +266,10 @@ def main() -> None:
print(f"warmup wall-times: {[round(x, 1) for x in warmup_secs]}")
if measured_secs:
avg = sum(measured_secs) / len(measured_secs)
print(f"measured e2e (n={len(measured_secs)}): "
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s")
print(
f"measured e2e (n={len(measured_secs)}): "
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
)
if stage_times:
print(f"average stage times over {measured_runs} measured runs:")
avg_total = 0.0
@@ -100,14 +100,23 @@ _inductor.coordinate_descent_tuning = True
_inductor.coordinate_descent_check_all_directions = True
_inductor.epilogue_fusion = False
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX23_MODEL_PATH",
"FastVideo/LTX-2.3-Distilled-Diffusers")))
OUTPUT_DIR = Path(os.getenv("LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"))
MODEL_ID = os.path.expandvars(
os.path.expanduser(
os.getenv("LTX23_MODEL_PATH", "FastVideo/LTX-2.3-Distilled-Diffusers")
)
)
OUTPUT_DIR = Path(
os.getenv(
"LTX23_OUTPUT_DIR", "outputs_video/ltx2_3_distilled_i2v_typed"
)
)
I2V_IMAGE = os.getenv("LTX23_I2V_IMAGE", "")
DEFAULT_PROMPT = ("A fashion model takes a slow step forward and shifts her weight, "
"the soft fabric of her clothing swaying and rippling with the "
"motion, her hair shifting gently, soft even studio lighting on a "
"clean light background, elegant slow-motion runway feel.")
DEFAULT_PROMPT = (
"A fashion model takes a slow step forward and shifts her weight, "
"the soft fabric of her clothing swaying and rippling with the "
"motion, her hair shifting gently, soft even studio lighting on a "
"clean light background, elegant slow-motion runway feel."
)
PROMPT = os.getenv("LTX23_I2V_PROMPT", DEFAULT_PROMPT)
@@ -138,7 +147,9 @@ def _collect_stage_times(
return
for name, metrics in stages.items():
stage_order.setdefault(name, None)
stage_times.setdefault(name, []).append(float(metrics.get("execution_time", 0.0)))
stage_times.setdefault(name, []).append(
float(metrics.get("execution_time", 0.0))
)
def _resolve_refine_upsampler(model_root: str) -> Path:
@@ -146,16 +157,20 @@ def _resolve_refine_upsampler(model_root: str) -> Path:
cand = Path(model_root) / name
if (cand / "config.json").is_file():
return cand
raise FileNotFoundError(f"No refine upsampler directory under {model_root}. "
f"Expected `{model_root}/spatial_upscaler/config.json`.")
raise FileNotFoundError(
f"No refine upsampler directory under {model_root}. "
f"Expected `{model_root}/spatial_upscaler/config.json`."
)
def main() -> None:
if not I2V_IMAGE:
raise SystemExit("LTX23_I2V_IMAGE is required for i2v. Example:\n"
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
" python examples/inference/basic/"
"basic_ltx2_3_distilled_i2v_typed.py")
raise SystemExit(
"LTX23_I2V_IMAGE is required for i2v. Example:\n"
" export LTX23_I2V_IMAGE=/path/to/portrait_or_product.jpg\n"
" python examples/inference/basic/"
"basic_ltx2_3_distilled_i2v_typed.py"
)
if not Path(I2V_IMAGE).is_file():
raise SystemExit(f"LTX23_I2V_IMAGE not found: {I2V_IMAGE}")
@@ -205,9 +220,10 @@ def main() -> None:
# model-specific VAE precision / decoder defaults are picked up
# the same way the legacy example's
# ``PipelineConfig.from_pretrained(model_root)`` did them.
components=ComponentConfig(upsampler_weights=str(refine_upsampler_path),
# Distilled has no refine LoRA — omit ``lora_path``.
),
components=ComponentConfig(
upsampler_weights=str(refine_upsampler_path),
# Distilled has no refine LoRA — omit ``lora_path``.
),
vae_tiling=False,
preset_overrides={
"refine": {
@@ -262,7 +278,11 @@ def main() -> None:
for w in range(warmup_runs):
print(f"\n[warmup {w + 1}/{warmup_runs}] compiling + generating…")
t0 = time.perf_counter()
generator.generate(build_request(OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7))
generator.generate(
build_request(
OUTPUT_DIR / f"_warmup_{w + 1}.mp4", seed=7
)
)
dt = time.perf_counter() - t0
warmup_secs.append(dt)
print(f"[warmup {w + 1}/{warmup_runs}] wall={dt:.1f}s")
@@ -271,30 +291,46 @@ def main() -> None:
(OUTPUT_DIR / f"_warmup_{w + 1}.mp4").unlink(missing_ok=True)
for m in range(measured_runs):
out_path = (OUTPUT_DIR / f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4")
print(f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}")
out_path = (
OUTPUT_DIR
/ f"output_ltx2_3_distilled_i2v_typed_run_{m + 1}.mp4"
)
print(
f"\n[measured {m + 1}/{measured_runs}] generating: {out_path}"
)
t0 = time.perf_counter()
result = generator.generate(build_request(out_path, seed=2002 + m))
result = generator.generate(
build_request(out_path, seed=2002 + m)
)
wall = time.perf_counter() - t0
# ``e2e_latency`` is currently surfaced via ``result.extra``;
# ``GenerationResult`` exposes ``generation_time`` as a
# first-class field but the LTX-2 pipeline only fills the
# legacy ``e2e_latency`` key. Prefer the explicit one, fall
# back to wall-clock.
e2e = (result.extra.get("e2e_latency") if hasattr(result, "extra") else None) or wall
e2e = (
result.extra.get("e2e_latency")
if hasattr(result, "extra") else None
) or wall
measured_secs.append(e2e)
print(f"[measured {m + 1}/{measured_runs}] "
f"e2e={e2e:.2f}s wall={wall:.2f}s")
print(
f"[measured {m + 1}/{measured_runs}] "
f"e2e={e2e:.2f}s wall={wall:.2f}s"
)
_print_stage_breakdown(result, f"measured {m + 1}")
_collect_stage_times(result, stage_times, stage_order)
print("\n=== summary ===")
print(f"warmup wall-times: "
f"{[round(x, 1) for x in warmup_secs]}")
print(
f"warmup wall-times: "
f"{[round(x, 1) for x in warmup_secs]}"
)
if measured_secs:
avg = sum(measured_secs) / len(measured_secs)
print(f"measured e2e (n={len(measured_secs)}): "
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s")
print(
f"measured e2e (n={len(measured_secs)}): "
f"{[round(x, 2) for x in measured_secs]} -> avg {avg:.2f}s"
)
if stage_times:
print(f"average stage times over {measured_runs} measured runs:")
avg_total = 0.0
@@ -1,21 +1,21 @@
from fastvideo import VideoGenerator
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
import os
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
@@ -12,11 +12,15 @@ from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.layers.quantization.nvfp4_config import NVFP4Config
from fastvideo.utils import maybe_download_model
VALIDATION_JSON = (Path(__file__).resolve().parents[2] / "training" / "finetune" / "ltx2" / "validation.json")
VALIDATION_JSON = (
Path(__file__).resolve().parents[2] / "training" / "finetune" / "ltx2" / "validation.json"
)
# Override with a local snapshot or converted directory when needed, e.g.
# export LTX2_MODEL_PATH=/raid/$USER/hf/FastVideo/LTX2-Distilled-Diffusers
MODEL_ID = os.path.expandvars(os.path.expanduser(os.getenv("LTX2_MODEL_PATH", "FastVideo/LTX2-Distilled-Diffusers")))
MODEL_ID = os.path.expandvars(
os.path.expanduser(os.getenv("LTX2_MODEL_PATH", "FastVideo/LTX2-Distilled-Diffusers"))
)
OUTPUT_DIR = Path("outputs_video/ltx2_distilled_fast_profile")
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
@@ -65,7 +69,9 @@ def print_stage_breakdown(
return total
def extract_sr_forward_latency(result: dict, ) -> tuple[float | None, list[tuple[str, float]], list[str]]:
def extract_sr_forward_latency(
result: dict,
) -> tuple[float | None, list[tuple[str, float]], list[str]]:
logging_info = result.get("logging_info")
if logging_info is None:
return None, [], []
@@ -83,8 +89,12 @@ def extract_sr_forward_latency(result: dict, ) -> tuple[float | None, list[tuple
if sr_match_substr:
is_sr_stage = sr_match_substr in stage_name_l
else:
is_sr_stage = ("srdenoisingstage" in stage_name_l or "sr_denoising" in stage_name_l
or "upsample" in stage_name_l or ("refine" in stage_name_l and "denois" in stage_name_l))
is_sr_stage = (
"srdenoisingstage" in stage_name_l
or "sr_denoising" in stage_name_l
or "upsample" in stage_name_l
or ("refine" in stage_name_l and "denois" in stage_name_l)
)
if not is_sr_stage:
continue
exec_time = float(stage_metrics.get("execution_time", 0.0))
@@ -151,9 +161,11 @@ def resolve_refine_upsampler_path(model_root: str) -> Path:
return candidate
checked = "\n".join(f" - {candidate}" for candidate in candidates)
raise FileNotFoundError("Could not find an LTX2 refine upsampler directory.\n"
"Checked:\n"
f"{checked}")
raise FileNotFoundError(
"Could not find an LTX2 refine upsampler directory.\n"
"Checked:\n"
f"{checked}"
)
def main() -> None:
@@ -302,15 +314,19 @@ def main() -> None:
measured_times = run_times[measured_start_idx:]
avg_time = sum(measured_times) / len(measured_times)
print(f"Average video generation time over {len(measured_times)} runs "
f"(runs {measured_start_idx + 1}-{len(run_times)}, skipping first {warmup_runs} warmup runs): "
f"{avg_time:.2f}s")
print(
f"Average video generation time over {len(measured_times)} runs "
f"(runs {measured_start_idx + 1}-{len(run_times)}, skipping first {warmup_runs} warmup runs): "
f"{avg_time:.2f}s"
)
measured_e2e_times = e2e_times[measured_start_idx:]
avg_e2e_time = sum(measured_e2e_times) / len(measured_e2e_times)
print(f"Average end-to-end latency over {len(measured_e2e_times)} runs "
f"(runs {measured_start_idx + 1}-{len(e2e_times)}, skipping first {warmup_runs} warmup runs): "
f"{avg_e2e_time:.2f}s")
print(
f"Average end-to-end latency over {len(measured_e2e_times)} runs "
f"(runs {measured_start_idx + 1}-{len(e2e_times)}, skipping first {warmup_runs} warmup runs): "
f"{avg_e2e_time:.2f}s"
)
if sr_forward_times:
avg_sr_forward = sum(sr_forward_times) / len(sr_forward_times)
@@ -322,8 +338,10 @@ def main() -> None:
if non_stage_overhead_times:
avg_non_stage_overhead = sum(non_stage_overhead_times) / len(non_stage_overhead_times)
print("Average non-stage overhead over "
f"{len(non_stage_overhead_times)} measured runs: {avg_non_stage_overhead:.3f}s")
print(
"Average non-stage overhead over "
f"{len(non_stage_overhead_times)} measured runs: {avg_non_stage_overhead:.3f}s"
)
else:
print("Average non-stage overhead unavailable (no stage timings).")
finally:
+12 -22
View File
@@ -13,34 +13,24 @@ MODEL_VARIANT = "base_distilled_model"
# Variant-specific settings
VARIANT_CONFIG = {
"base_distilled_model": {
"model_path":
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"keyboard_dim":
4,
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
"model_path": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"keyboard_dim": 4,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
},
"gta_distilled_model": {
"model_path":
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"keyboard_dim":
2,
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"keyboard_dim": 2,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
},
"templerun_distilled_model": {
"model_path":
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
"keyboard_dim":
7,
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
"keyboard_dim": 7,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
},
}
OUTPUT_PATH = "video_samples_matrixgame2"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -52,8 +42,8 @@ def main():
config["model_path"],
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
@@ -14,40 +14,27 @@ MODEL_VARIANT = "base_distilled_model"
# Variant-specific settings
VARIANT_CONFIG = {
"base_distilled_model": {
"model_path":
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"keyboard_dim":
4,
"mode":
"universal",
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
"model_path": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"keyboard_dim": 4,
"mode": "universal",
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
},
"gta_distilled_model": {
"model_path":
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"keyboard_dim":
2,
"mode":
"gta_drive",
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
"keyboard_dim": 2,
"mode": "gta_drive",
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
},
"templerun_distilled_model": {
"model_path":
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
"keyboard_dim":
7,
"mode":
"templerun",
"image_url":
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
"keyboard_dim": 7,
"mode": "templerun",
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
},
}
OUTPUT_PATH = "video_samples_matrixgame2"
async def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -59,8 +46,8 @@ async def main():
config["model_path"],
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
@@ -69,8 +56,11 @@ async def main():
)
max_blocks = 50
num_frames = 597
actions = {"keyboard": torch.zeros((num_frames, config["keyboard_dim"])), "mouse": torch.zeros((num_frames, 2))}
num_frames = 597
actions = {
"keyboard": torch.zeros((num_frames, config["keyboard_dim"])),
"mouse": torch.zeros((num_frames, 2))
}
grid_sizes = torch.tensor([150, 44, 80])
mode = config["mode"]
@@ -91,11 +81,11 @@ async def main():
for block_id in range(max_blocks):
print(f"\n=== Block {block_id + 1}/{max_blocks} ===")
action = await get_current_action_async(mode)
keyboard_cond, mouse_cond = expand_action_to_frames(action, 12)
await generator.step_async(keyboard_cond, mouse_cond)
if (await asyncio.to_thread(input, "\nContinue? (y/n): ")).lower() == 'n':
break
@@ -25,13 +25,6 @@ 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,13 +25,6 @@ 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,13 +22,6 @@ 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)
@@ -37,15 +30,13 @@ 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()
@@ -76,24 +67,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:
-44
View File
@@ -1,44 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""MMAudio large-44k-v2 video-to-audio example."""
import argparse
import os
from fastvideo import VideoGenerator
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--video-path", required=True)
parser.add_argument("--output-path", default="outputs_audio/mmaudio.wav")
parser.add_argument("--duration-seconds", type=float, default=8.0)
parser.add_argument("--prompt", default="")
parser.add_argument("--negative-prompt", default="music")
return parser.parse_args()
def main() -> None:
args = parse_args()
generator = VideoGenerator.from_pretrained(
os.environ.get(
"MMAUDIO_MODEL_PATH",
"converted_weights/mmaudio/large_44k_v2",
),
workload_type="v2a",
num_gpus=1,
)
result = generator.generate_video(
prompt=args.prompt,
negative_prompt=args.negative_prompt,
video_path=args.video_path,
audio_end_in_s=args.duration_seconds,
output_path=args.output_path,
save_video=True,
return_frames=False,
)
print(result["video_path"])
generator.shutdown()
if __name__ == "__main__":
main()
+15 -17
View File
@@ -1,20 +1,19 @@
from fastvideo import VideoGenerator, PipelineConfig
from fastvideo.api.sampling_param import SamplingParam
def main():
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
config.text_encoder_precisions = ["fp16"]
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
pipeline_config=config,
use_fsdp_inference=False, # Disable FSDP for MPS
dit_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
disable_autocast=False,
num_gpus=1,
use_fsdp_inference=False, # Disable FSDP for MPS
dit_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
disable_autocast=False,
num_gpus=1,
)
# Create sampling parameters with reduced number of frames
@@ -24,19 +23,18 @@ def main():
sampling_param.width = 256
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(prompt, sampling_param=sampling_param)
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, sampling_param=sampling_param)
if __name__ == "__main__":
main()
+12 -11
View File
@@ -3,8 +3,6 @@ from fastvideo import VideoGenerator
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -18,24 +16,27 @@ 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)
+2 -3
View File
@@ -7,6 +7,7 @@ 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",
@@ -48,9 +49,7 @@ 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,8 +3,6 @@ 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.
@@ -15,18 +13,19 @@ 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