Compare commits
20
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
be7cc7a1dd | ||
|
|
8c5778fa44 | ||
|
|
0fbdbe9174 | ||
|
|
659cf4fbfe | ||
|
|
609648b114 | ||
|
|
9a2eda427f | ||
|
|
cb7e8e960a | ||
|
|
4a4e9b3f31 | ||
|
|
0f2bbcda10 | ||
|
|
58f7b913b9 | ||
|
|
07f1cfc356 | ||
|
|
3101331020 | ||
|
|
4a63de71d7 | ||
|
|
99df34a78d | ||
|
|
7b94433092 | ||
|
|
3a7e7fc69f | ||
|
|
02a10268e9 | ||
|
|
9c9dd70214 | ||
|
|
dc13c59dca | ||
|
|
6dd75ef5ab |
@@ -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();
|
||||
});
|
||||
});
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 ---------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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();
|
||||
});
|
||||
});
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
@@ -9,10 +9,11 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.distributed import get_local_torch_device, get_world_group
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.profiler import get_global_controller
|
||||
from tqdm.auto import tqdm
|
||||
from fastvideo.hooks.activation_trace import trace_step
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MINIMAX_H3_KEYFRAME_NOISE_AUG,
|
||||
@@ -104,6 +105,11 @@ class MiniMaxH3DenoisingStage(PipelineStage):
|
||||
controller = get_global_controller()
|
||||
denoise_region = (controller.region("profiler_region_inference_denoising")
|
||||
if controller is not None else contextlib.nullcontext())
|
||||
# rank-0 step progress, mirroring the shared denoising stage; on a
|
||||
# non-tty each refresh is a plain line, so ray's log relay (and the
|
||||
# studio progress parser behind it) sees per-step updates.
|
||||
steps_bar = (tqdm(total=len(video_timesteps), desc="denoising")
|
||||
if get_world_group().local_rank == 0 else None)
|
||||
try:
|
||||
with denoise_region:
|
||||
for index, (video_timestep, audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps,
|
||||
@@ -149,7 +155,11 @@ class MiniMaxH3DenoisingStage(PipelineStage):
|
||||
)[0]
|
||||
batch.step_index = index
|
||||
batch.timestep = video_timestep
|
||||
if steps_bar is not None:
|
||||
steps_bar.update()
|
||||
finally:
|
||||
if steps_bar is not None:
|
||||
steps_bar.close()
|
||||
if bool(getattr(fastvideo_args, "dit_layerwise_offload", False)):
|
||||
manager = getattr(self.transformer, "_layerwise_offload_manager", None)
|
||||
if manager is not None and getattr(manager, "enabled", False):
|
||||
|
||||
@@ -341,6 +341,16 @@ class RayDistributedExecutor(Executor):
|
||||
if response["status"] != "lora_adapter_merged":
|
||||
raise RuntimeError(f"Worker {i} failed to merge LoRA weights")
|
||||
|
||||
def set_log_queue(self, log_queue) -> None:
|
||||
# A multiprocessing.Manager queue cannot reach ray actors on other
|
||||
# nodes; worker logs stay in ray's per-worker log files instead.
|
||||
if log_queue is not None:
|
||||
logger.warning("set_log_queue is a no-op on the ray backend; "
|
||||
"worker logs remain in the ray session log dir")
|
||||
|
||||
def clear_log_queue(self) -> None:
|
||||
pass
|
||||
|
||||
def collective_rpc(self,
|
||||
method: str | Callable,
|
||||
timeout: float | None = None,
|
||||
@@ -375,6 +385,39 @@ class RayDistributedExecutor(Executor):
|
||||
ray.kill(worker)
|
||||
|
||||
self.workers = []
|
||||
# Killing the actors does not release the placement group — without
|
||||
# this, the GPUs stay reserved and no new engine can ever schedule
|
||||
# in the same ray cluster.
|
||||
pg = getattr(self.fastvideo_args, "ray_placement_group", None)
|
||||
if pg is not None:
|
||||
try:
|
||||
ray.util.remove_placement_group(pg)
|
||||
# Removal is async; a generator created right after shutdown
|
||||
# races the teardown and fails "no GPU available". Wait for
|
||||
# the resources to actually free.
|
||||
import time as _time
|
||||
from ray.util import placement_group_table
|
||||
deadline = _time.monotonic() + 30
|
||||
while _time.monotonic() < deadline:
|
||||
if placement_group_table(pg).get("state") == "REMOVED":
|
||||
break
|
||||
_time.sleep(0.5)
|
||||
else:
|
||||
logger.warning("Placement group still not removed after 30s")
|
||||
# The killed actors release their GPUs asynchronously as the
|
||||
# processes die; wait for the resources to reappear so the
|
||||
# next load doesn't fail its GPU-availability probe.
|
||||
want = float(self.fastvideo_args.num_gpus)
|
||||
deadline = _time.monotonic() + 60
|
||||
while _time.monotonic() < deadline:
|
||||
if ray.available_resources().get("GPU", 0.0) >= want:
|
||||
break
|
||||
_time.sleep(0.5)
|
||||
else:
|
||||
logger.warning("GPUs not back in ray ledger 60s after shutdown")
|
||||
except Exception: # noqa: BLE001 -- already-removed / cluster gone
|
||||
logger.warning("Failed to remove ray placement group", exc_info=True)
|
||||
self.fastvideo_args.ray_placement_group = None
|
||||
|
||||
def __del__(self):
|
||||
self.shutdown()
|
||||
|
||||
@@ -217,12 +217,21 @@ def initialize_ray_cluster(
|
||||
# the current node has at least one device.
|
||||
current_ip = get_ip()
|
||||
current_node_id = ray.get_runtime_context().get_node_id()
|
||||
current_node_resource = available_resources_per_node()[current_node_id]
|
||||
if current_node_resource.get(device_str, 0) < 1:
|
||||
raise ValueError(f"Current node has no {device_str} available. "
|
||||
f"{current_node_resource=}. FastVideo engine cannot start without "
|
||||
f"{device_str}. Make sure you have at least 1 {device_str} "
|
||||
f"available in a node {current_node_id=} {current_ip=}.")
|
||||
# A previous engine's actors may still be releasing their devices
|
||||
# (actor death and placement-group teardown are asynchronous) — retry
|
||||
# briefly instead of failing an otherwise-valid load.
|
||||
deadline = time.monotonic() + 60
|
||||
while True:
|
||||
current_node_resource = available_resources_per_node()[current_node_id]
|
||||
if current_node_resource.get(device_str, 0) >= 1:
|
||||
break
|
||||
if time.monotonic() >= deadline:
|
||||
raise ValueError(f"Current node has no {device_str} available. "
|
||||
f"{current_node_resource=}. FastVideo engine cannot start without "
|
||||
f"{device_str}. Make sure you have at least 1 {device_str} "
|
||||
f"available in a node {current_node_id=} {current_ip=}.")
|
||||
logger.info("Waiting for a %s to free on the current node...", device_str)
|
||||
time.sleep(2)
|
||||
# This way, at least bundle is required to be created in a current
|
||||
# node.
|
||||
placement_group_specs[0][f"node:{current_ip}"] = 0.001
|
||||
|
||||
Reference in New Issue
Block a user