Compare commits

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

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

GET /api/engine/logs: incremental tail of the server's stdout/stderr (ray
relays worker output to the driver, so this captures every rank).
2026-08-06 22:40:39 -07:00
SolitaryThinker 02a10268e9 studio ui: warm models panel on the inference page
lists resident/loading/failed generators, preload from the model dropdown
using the persisted default job options (same cache key as jobs), unload with
409 surfaced as a toast. polls only while something is loading. mock server
+ playwright spec + vitest coverage.
2026-08-06 22:29:03 -07:00
SolitaryThinker 9c9dd70214 studio: preload/list/unload resident generators
generators were already cached across jobs but only implicitly — the first
job per config silently paid the whole model load. adds:
- POST /api/generators/preload: load a model into memory ahead of time
  (async, idempotent; state loading/ready/failed)
- GET /api/generators: what's resident or in flight
- POST /api/generators/unload: shutdown + free VRAM (409 while inference
  jobs run)
cache key builder shared with the job path and pinned by a test so preload
and job lookup can't drift apart.
2026-08-06 22:20:28 -07:00
SolitaryThinker dc13c59dca fix ray executor: implement the log-queue abstract methods
RayDistributedExecutor became uninstantiable when set_log_queue/
clear_log_queue were added to the Executor ABC (only the mp executor got
implementations). No-op on ray: a Manager queue can't cross nodes.
2026-08-06 19:16:33 -07:00
SolitaryThinker 6dd75ef5ab studio: multi-node inference knobs
- forward num_gpus to from_pretrained (was dropped: multi-gpu only ever
  happened via sp_size)
- FASTVIDEO_STUDIO_EXECUTOR_BACKEND=ray env knob so a deployment on a ray
  cluster can span nodes
- FASTVIDEO_STUDIO_MODEL_PATHS="id=/local/dir" serves a registered model id
  from local weights (avoids re-downloading checkpoints already on shared fs)
- accept an existing local dir as model_id in POST /api/jobs
2026-08-06 18:42:44 -07:00
27 changed files with 2895 additions and 103 deletions
@@ -0,0 +1,39 @@
import { expect, test } from '@playwright/test';
import { skipWithoutMock } from './helpers';
/**
* Warm-model slot: load a model through the mock (which flips
* loading -> ready after ~2s, surfaced by the panel's 5s poll), then unload it.
*/
test.describe('generators', () => {
skipWithoutMock();
test('loads and unloads the resident model', async ({ page }) => {
await page.goto('/inference');
const panel = page.getByRole('region', { name: 'Warm models' });
await expect(panel).toBeVisible();
await expect(panel.getByText('No model loaded')).toBeVisible();
await panel
.getByLabel('Model to load')
.selectOption('Wan-AI/Wan2.1-T2V-1.3B-Diffusers');
await panel.getByRole('button', { name: 'Load model' }).click();
await expect(panel.getByText('Wan2.1 T2V 1.3B Diffusers')).toBeVisible();
await expect(panel.getByText('ready')).toBeVisible();
await panel.getByRole('button', { name: 'Unload' }).click();
await expect(panel.getByText('No model loaded')).toBeVisible();
});
test('engine console streams output while open', async ({ page }) => {
await page.goto('/inference');
const engineConsole = page.getByRole('region', { name: 'Engine output' });
await engineConsole.getByRole('button', { name: 'Engine output' }).click();
await expect(engineConsole.getByText(/\[engine\]/).first()).toBeVisible();
});
});
+87
View File
@@ -72,3 +72,90 @@ def get_gpu_snapshot() -> dict[str, Any]:
except Exception as exc: # NVMLError, driver issues, …
logger.warning("GPU snapshot failed: %s", exc)
return {"available": False, "gpus": [], "error": str(exc)}
def _remote_gpu_probe() -> dict[str, Any]:
"""Self-contained per-node NVML probe (runs as a ray task on each node;
no fastvideo_studio import — worker environments don't have apps/ on
their path, so cloudpickle must carry this by value)."""
import contextlib as _ctx
import socket as _socket
out: dict[str, Any] = {"hostname": _socket.gethostname(), "available": False, "gpus": [], "error": None}
try:
import pynvml
pynvml.nvmlInit()
for i in range(pynvml.nvmlDeviceGetCount()):
h = pynvml.nvmlDeviceGetHandleByIndex(i)
name = pynvml.nvmlDeviceGetName(h)
if isinstance(name, bytes):
name = name.decode()
mem = pynvml.nvmlDeviceGetMemoryInfo(h)
util = pynvml.nvmlDeviceGetUtilizationRates(h)
temp = power = plimit = None
with _ctx.suppress(pynvml.NVMLError):
temp = int(pynvml.nvmlDeviceGetTemperature(h, pynvml.NVML_TEMPERATURE_GPU))
with _ctx.suppress(pynvml.NVMLError):
power = pynvml.nvmlDeviceGetPowerUsage(h) / 1000.0
plimit = pynvml.nvmlDeviceGetEnforcedPowerLimit(h) / 1000.0
out["gpus"].append({
"index": i,
"name": name,
"utilization": int(util.gpu),
"memory_used_mib": int(mem.used / (1024 * 1024)),
"memory_total_mib": int(mem.total / (1024 * 1024)),
"temperature_c": temp,
"power_watts": power,
"power_limit_watts": plimit,
})
out["available"] = True
except Exception as exc: # noqa: BLE001 -- reported per node
out["error"] = str(exc)
return out
def get_cluster_snapshot() -> dict[str, Any]:
"""Cluster-wide GPU/host telemetry.
When this process is connected to a ray cluster (a model has been
loaded), probe every alive node via per-node ray tasks. Otherwise fall
back to this host's NVML snapshot. Never raises.
"""
import socket
local = get_gpu_snapshot()
local_node = {"hostname": socket.gethostname(), "ip": None, "is_this_host": True,
"cpus": None, "ray_gpus": None, **local}
out: dict[str, Any] = {"mode": "local", "nodes": [local_node],
"resources": None, "error": None}
try:
import ray
if not ray.is_initialized():
out["error"] = "not connected to a ray cluster yet (load a model first); showing the API host only"
return out
from ray.util.scheduling_strategies import NodeAffinitySchedulingStrategy
alive = [n for n in ray.nodes() if n.get("Alive")]
probe = ray.remote(num_cpus=0)(_remote_gpu_probe)
refs = [probe.options(scheduling_strategy=NodeAffinitySchedulingStrategy(
node_id=n["NodeID"], soft=True)).remote() for n in alive]
snaps = ray.get(refs, timeout=15)
nodes = []
for n, snap in zip(alive, snaps, strict=True):
nodes.append({
"ip": n.get("NodeManagerAddress"),
"is_this_host": snap.get("hostname") == socket.gethostname(),
"cpus": n.get("Resources", {}).get("CPU"),
"ray_gpus": n.get("Resources", {}).get("GPU"),
**snap,
})
out["mode"] = "ray"
out["nodes"] = nodes
out["resources"] = {
"gpus_total": ray.cluster_resources().get("GPU", 0.0),
"gpus_available": ray.available_resources().get("GPU", 0.0),
}
except Exception as exc: # noqa: BLE001 -- degrade to the local view
logger.warning("cluster snapshot failed: %s", exc)
out["mode"] = "local"
out["nodes"] = [local_node]
out["error"] = f"cluster probe failed: {exc}"
return out
+282 -61
View File
@@ -41,6 +41,11 @@ _TQDM_FRAC_RE = re.compile(r"\b(\d+)/(\d+)\b")
_MAX_LOG_LINES = 2000 # ring-buffer cap per job
# ray's log relay prefixes worker lines like "(RayWorkerWrapper pid=123, ip=…)"
# — usually wrapped in ANSI color codes, which must be stripped before matching.
_RAY_RELAY_RE = re.compile(r"^\(\w+ pid=")
_ANSI_RE = re.compile(r"\x1b\[[0-9;]*m")
class JobStatus(str, enum.Enum):
PENDING = "pending"
@@ -259,13 +264,36 @@ class JobRunner:
self._jobs_lock = threading.Lock()
self._load_jobs()
# Cache of loaded generators keyed by model config so that we only pay
# the model-loading cost once per model configuration.
self._generators: dict[tuple, Any] = {}
self._generators_lock = threading.Lock()
# Exactly one generator lives in memory at a time. Loading a new
# config always releases the old instance first (shutdown + placement
# group teardown); unload deletes it outright.
self._generator: Any | None = None
self._generator_config: dict[str, Any] | None = None
self._generator_state: str = "empty" # empty | loading | ready | failed
self._generator_error: str | None = None
self._generator_lock = threading.Lock() # guards the fields above
# Serializes every slot transition (preload, job-triggered replace,
# unload). Held for the full duration of a load.
self._load_lock = threading.Lock()
# All generator creations run on this ONE persistent thread: the mp
# executor arms prctl(PR_SET_PDEATHSIG, SIGKILL) in its workers, and
# on Linux that fires when the CREATING THREAD exits — a generator
# spawned from a short-lived thread loses all its workers (silent
# SIGKILL, zombies) the moment that thread finishes.
import queue as _queue
self._loader_queue: _queue.Queue = _queue.Queue()
threading.Thread(target=self._loader_loop, daemon=True,
name="generator-loader").start()
# The inference job currently generating, fed by the engine log tee
# (ray relays worker output to the driver; tqdm lines land there).
self._active_inference_job: Job | None = None
# Shared Manager for log queues (avoids spawning a new process per job)
self._mp_manager = get_mp_context().Manager()
# One queue for the generator's whole life: mp workers get it at spawn
# (creation-time attach). Sending a Manager proxy over the executor's
# worker pipes post-hoc (set_log_queue RPC) breaks the pipe.
self._worker_log_queue = self._mp_manager.Queue()
atexit.register(self._shutdown)
# Ensure directories exist
@@ -617,6 +645,211 @@ class JobRunner:
"phase": job._log_buf.phase,
}
@staticmethod
def _generator_config_dict(
model_id: str,
workload_type: str,
num_gpus: int,
dit_cpu_offload: bool = False,
text_encoder_cpu_offload: bool = False,
vae_cpu_offload: bool = False,
image_encoder_cpu_offload: bool = False,
use_fsdp_inference: bool = False,
enable_torch_compile: bool = False,
vsa_sparsity: float = 0.0,
tp_size: int = -1,
sp_size: int = -1,
) -> dict[str, Any]:
"""Canonical engine-config dict; equality here == same generator."""
return {
"model_id": model_id,
"workload_type": workload_type,
"num_gpus": num_gpus,
"dit_cpu_offload": dit_cpu_offload,
"text_encoder_cpu_offload": text_encoder_cpu_offload,
"vae_cpu_offload": vae_cpu_offload,
"image_encoder_cpu_offload": image_encoder_cpu_offload,
"use_fsdp_inference": use_fsdp_inference,
"enable_torch_compile": enable_torch_compile,
"vsa_sparsity": vsa_sparsity,
"tp_size": tp_size,
"sp_size": sp_size,
}
def _slot_entry(self) -> dict[str, Any]:
return {
"state": self._generator_state,
"error": self._generator_error,
**(self._generator_config or {}),
}
def _running_inference_jobs(self) -> list[str]:
with self._jobs_lock:
return [j.id for j in self._jobs.values()
if j.status == JobStatus.RUNNING and j.job_type == "inference"]
def _loader_loop(self) -> None:
while True:
fn = self._loader_queue.get()
try:
fn()
except BaseException: # noqa: BLE001 -- surfaced via the caller's box
pass
finally:
self._loader_queue.task_done()
def _run_on_loader(self, fn: Any) -> Any:
"""Run ``fn`` on the persistent loader thread and return its result."""
box: dict[str, Any] = {}
done = threading.Event()
def wrapped() -> None:
try:
box["r"] = fn()
except BaseException as exc: # noqa: BLE001 -- re-raised below
box["e"] = exc
finally:
done.set()
self._loader_queue.put(wrapped)
done.wait()
if "e" in box:
raise box["e"]
return box["r"]
def preload_generator(self, **params: Any) -> dict[str, Any]:
"""Load a model into memory ahead of time. One load at a time; loading
a different config always releases the current instance first."""
config = self._generator_config_dict(**params)
with self._generator_lock:
if self._generator_state == "ready" and self._generator_config == config:
return self._slot_entry()
if not self._load_lock.acquire(blocking=False):
raise RuntimeError("a model load is already in progress")
try:
if self._running_inference_jobs():
raise RuntimeError("cannot swap models while inference jobs are running")
with self._generator_lock:
self._generator_state = "loading"
self._generator_config = config
self._generator_error = None
entry = self._slot_entry()
except BaseException:
self._load_lock.release()
raise
def _load() -> None:
try: # the preload owns _load_lock until the load resolves
self._run_on_loader(lambda: self._load_into_slot_locked(config))
except Exception: # noqa: BLE001 -- state already set to failed
pass
finally:
self._load_lock.release()
threading.Thread(target=_load, daemon=True, name="generator-preload").start()
return entry
def _load_into_slot_locked(self, config: dict[str, Any]) -> Any:
"""Release whatever is resident and load ``config``. Caller MUST hold
``_load_lock``. State is 'loading' on entry or set here."""
with self._generator_lock:
gen = self._generator
self._generator = None
self._generator_state = "loading"
self._generator_config = config
self._generator_error = None
if gen is not None:
logger.info("Releasing resident generator before loading a new one")
gen.shutdown()
del gen
# Import lazily so starting the server is fast even without a GPU.
from fastvideo import VideoGenerator
# Deployment-level knobs (set where the API server is launched):
# FASTVIDEO_STUDIO_MODEL_PATHS="id=/local/dir,..." serves a registered
# model id from local weights instead of the HF hub;
# FASTVIDEO_STUDIO_EXECUTOR_BACKEND=ray runs workers on an existing
# Ray cluster (the multi-node path — "mp" spawns local processes only).
model_path = config["model_id"]
for pair in os.environ.get("FASTVIDEO_STUDIO_MODEL_PATHS", "").split(","):
mid, sep, path = pair.partition("=")
if sep and mid.strip() == config["model_id"]:
model_path = path.strip()
executor_kwargs: dict[str, Any] = {}
backend = os.environ.get("FASTVIDEO_STUDIO_EXECUTOR_BACKEND")
if backend:
executor_kwargs["distributed_executor_backend"] = backend
logger.info("Loading model %s (%s)", config["model_id"],
", ".join(f"{k}={v}" for k, v in config.items() if k != "model_id"))
try:
new_gen = VideoGenerator.from_pretrained(
model_path,
workload_type=config["workload_type"],
num_gpus=config["num_gpus"],
dit_cpu_offload=config["dit_cpu_offload"],
# FastVideoArgs defaults this True, which disables FSDP and
# parks a full DiT copy in host RAM per worker — the mp
# executor's 4 workers OOM-killed the node silently. The UI's
# offload toggles are the studio's contract; layerwise off.
dit_layerwise_offload=False,
text_encoder_cpu_offload=config["text_encoder_cpu_offload"],
vae_cpu_offload=config["vae_cpu_offload"],
image_encoder_cpu_offload=config["image_encoder_cpu_offload"],
use_fsdp_inference=config["use_fsdp_inference"],
enable_torch_compile=config["enable_torch_compile"],
VSA_sparsity=config["vsa_sparsity"],
tp_size=config["tp_size"],
sp_size=config["sp_size"],
log_queue=self._worker_log_queue,
**executor_kwargs,
)
except Exception as exc:
logger.exception("Model load failed for %s", config["model_id"])
with self._generator_lock:
self._generator_state = "failed"
self._generator_error = str(exc)
raise
with self._generator_lock:
self._generator = new_gen
self._generator_config = config
self._generator_state = "ready"
self._generator_error = None
return new_gen
def list_generators(self) -> list[dict[str, Any]]:
"""The resident slot, or empty when nothing is loaded."""
with self._generator_lock:
if self._generator_state == "empty":
return []
return [self._slot_entry()]
def unload_generator(self, **_ignored: Any) -> bool:
"""Shut down and delete the resident generator, freeing GPU memory."""
if not self._load_lock.acquire(blocking=False):
raise RuntimeError("cannot unload while a model load is in progress")
try:
running = self._running_inference_jobs()
if running:
raise RuntimeError(f"cannot unload while inference jobs are running: {running}")
with self._generator_lock:
gen = self._generator
empty = self._generator_state == "empty"
self._generator = None
self._generator_state = "empty"
self._generator_config = None
self._generator_error = None
if empty:
return False
if gen is not None:
logger.info("Releasing resident generator")
gen.shutdown()
del gen
return True
finally:
self._load_lock.release()
def _get_or_create_generator(
self,
model_id: str,
@@ -633,69 +866,53 @@ class JobRunner:
sp_size: int = -1,
log_queue: mp.Queue | None = None,
) -> Any:
cache_key = (
model_id,
workload_type,
num_gpus,
dit_cpu_offload,
text_encoder_cpu_offload,
vae_cpu_offload,
image_encoder_cpu_offload,
use_fsdp_inference,
enable_torch_compile,
vsa_sparsity,
tp_size,
sp_size,
)
# Generators are cached by model_id and configuration parameters
with self._generators_lock:
if cache_key in self._generators:
return self._generators[cache_key]
# Import lazily so starting the server is fast even without a GPU.
from fastvideo import VideoGenerator
logger.info(
"Loading model %s (workload=%s, num_gpus=%d, offloads: "
"dit=%s text_encoder=%s vae=%s image_encoder=%s, fsdp=%s, "
"torch_compile=%s, vsa_sparsity=%.2f, tp=%d sp=%d) …",
model_id,
workload_type,
num_gpus,
dit_cpu_offload,
text_encoder_cpu_offload,
vae_cpu_offload,
image_encoder_cpu_offload,
use_fsdp_inference,
enable_torch_compile,
vsa_sparsity,
tp_size,
sp_size,
)
gen = VideoGenerator.from_pretrained(
model_id,
"""Return the resident generator if the config matches; otherwise
replace the slot. Blocks behind any in-flight load — every slot
transition happens under ``_load_lock``, so a job can never observe a
half-replaced slot."""
del log_queue # single-slot generators log via the engine tee
config = self._generator_config_dict(
model_id=model_id,
workload_type=workload_type,
num_gpus=num_gpus,
dit_cpu_offload=dit_cpu_offload,
text_encoder_cpu_offload=text_encoder_cpu_offload,
vae_cpu_offload=vae_cpu_offload,
image_encoder_cpu_offload=image_encoder_cpu_offload,
use_fsdp_inference=use_fsdp_inference,
enable_torch_compile=enable_torch_compile,
VSA_sparsity=vsa_sparsity,
vsa_sparsity=vsa_sparsity,
tp_size=tp_size,
sp_size=sp_size,
log_queue=log_queue,
)
with self._load_lock: # waits out preloads / other jobs' replaces
with self._generator_lock:
if (self._generator_state == "ready"
and self._generator_config == config
and self._generator is not None):
return self._generator
return self._run_on_loader(lambda: self._load_into_slot_locked(config))
with self._generators_lock:
if cache_key not in self._generators:
self._generators[cache_key] = gen
else: # Another thread may have created it while we were loading.
gen.shutdown()
gen = self._generators[cache_key]
return gen
def feed_engine_line(self, line: str) -> None:
"""Bridge ray-relayed worker output into the running job's log buffer.
On the ray backend worker logs cannot cross nodes via the mp queue,
but ray already relays them to the driver's stdout — which the engine
tee captures. Lines with ray's actor prefix are attributed to the one
running inference job, whose buffer parses tqdm into UI progress.
Driver-side logging is excluded (it reaches the buffer via the
logging handlers already).
"""
job = self._active_inference_job
if job is None:
return
line = _ANSI_RE.sub("", line)
if not _RAY_RELAY_RE.match(line):
return
try:
job._log_buf.write(line)
except Exception: # noqa: BLE001 -- never break the tee
pass
def _run_job(self, job: Job):
if job.job_type == "inference":
@@ -841,10 +1058,12 @@ class JobRunner:
fastvideo_logger.addHandler(buffer_handler)
fastvideo_logger.addHandler(file_handler)
# Queue for worker process logs (fsdp_load, cuda, etc.)
# Use Manager().Queue() so it can be shared with spawned workers (spawn
# does not inherit memory; mp.Queue only works through inheritance).
log_queue = self._mp_manager.Queue()
# Worker logs flow through the runner-wide queue the generator was
# created with; drain anything stale, then listen for this job.
log_queue = self._worker_log_queue
with contextlib.suppress(Exception):
while True:
log_queue.get_nowait()
queue_listener = logging.handlers.QueueListener(log_queue,
buffer_handler,
file_handler,
@@ -926,6 +1145,7 @@ class JobRunner:
generator = _gen_result[0]
buf.phase = "generating"
self._active_inference_job = job # engine tee feeds tqdm from here
logger.info("Starting generation for job %s (model=%s)", job.id, job.model_id)
gen_kwargs: dict[str, Any] = {
@@ -941,7 +1161,6 @@ class JobRunner:
"fps": job.fps,
"seed": job.seed,
"negative_prompt": job.negative_prompt or "",
"log_queue": log_queue,
}
if job.image_path:
gen_kwargs["image_path"] = job.image_path
@@ -985,6 +1204,8 @@ class JobRunner:
buf.phase = "failed"
finally:
if self._active_inference_job is job:
self._active_inference_job = None
queue_listener.stop()
# Remove handlers and close file
fastvideo_logger.removeHandler(buffer_handler)
+187 -2
View File
@@ -36,13 +36,15 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, PlainTextResponse
from fastvideo_studio.database import default_settings_dict
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, SettingsUpdate, UpdateCaptionRequest,
model_label)
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, GeneratorRequest, SettingsUpdate,
UpdateCaptionRequest, model_label)
# --- Config -----------------------------------------------------------------
# How long a started job stays "running" before it flips to "completed".
COMPLETE_AFTER_SECONDS = 3.0
# How long a preloading generator stays "loading" before it flips to "ready".
GENERATOR_READY_AFTER_SECONDS = 2.0
FFMPEG_BIN = shutil.which(os.getenv("FASTVIDEO_FFMPEG_BIN", "ffmpeg"))
# A small catalogue of fake models keyed by workload type. Mirrors the real
@@ -88,6 +90,10 @@ _DEFAULT_SETTINGS: dict[str, Any] = {
_state_lock = threading.Lock()
_settings: dict[str, Any] = dict(_DEFAULT_SETTINGS)
_jobs: dict[str, dict[str, Any]] = {}
# The engine's single model slot: None when empty, else the state dict.
_generator_slot: dict[str, Any] | None = None
# Fake engine stdout/stderr tail; grows a little on every poll.
_engine_log_lines: list[str] = ["[engine] FastVideo studio mock engine started"]
_datasets: dict[str, dict[str, Any]] = {}
# dataset_id -> {"file_names": [...], "captions": {file_name: caption}}
_dataset_files: dict[str, dict[str, Any]] = {}
@@ -393,6 +399,56 @@ def list_models(workload_type: str | None = None) -> list[dict[str, Any]]:
return _models_for(workload_type)
# Per-model sampling presets, mirroring the real /api/models/presets shape
# (keys the model has no recommendation for are simply absent).
_MODEL_PRESETS: dict[str, dict[str, Any]] = {
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": {
"height": 480,
"width": 832,
"num_frames": 81,
"fps": 16,
"num_inference_steps": 50,
"guidance_scale": 3.0,
"guidance_rescale": 0.0,
"negative_prompt": "Bright tones, overexposed, static, blurred details",
"seed": 1024,
},
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": {
"height": 480,
"width": 832,
"num_frames": 81,
"fps": 16,
"num_inference_steps": 40,
"guidance_scale": 5.0,
"seed": 1024,
},
"black-forest-labs/FLUX.1-schnell": {
"height": 1024,
"width": 1024,
"num_frames": 1,
"num_inference_steps": 4,
"guidance_scale": 0.0,
"seed": 42,
},
}
_GENERIC_PRESETS: dict[str, Any] = {
"height": 720,
"width": 1280,
"num_frames": 81,
"fps": 24,
"num_inference_steps": 50,
"guidance_scale": 5.0,
"guidance_rescale": 0.0,
"seed": 1024,
}
@app.get("/api/models/presets")
def model_presets(model_id: str) -> dict[str, Any]:
"""Recommended sampling settings; unknown models get generic defaults."""
return _MODEL_PRESETS.get(model_id, _GENERIC_PRESETS)
# --- GPUs -------------------------------------------------------------------
@@ -414,6 +470,49 @@ def list_gpus() -> dict[str, Any]:
return {"available": True, "gpus": gpus, "error": None}
@app.get("/api/cluster")
def cluster_status() -> dict[str, Any]:
"""Two fake ray nodes x 4 GPUs with per-request jitter, mirroring the
real server's /api/cluster shape."""
nodes = []
for host_idx, (hostname, ip, is_this_host) in enumerate([
("mock-node-0", "10.0.0.10", True),
("mock-node-1", "10.0.0.11", False),
]):
gpus = []
for index in range(4):
base_util = (13 + 29 * index + 41 * host_idx) % 100
gpus.append({
"index": index,
"name": "NVIDIA Mock GPU 80GB",
"utilization": max(0, min(100, base_util + random.randint(-5, 5))),
"memory_used_mib": 6_144 + 17_408 * index + random.randint(-256, 256),
"memory_total_mib": 81_920,
"temperature_c": 42 + 6 * index + random.randint(-3, 3),
"power_watts": 110.0 + 140.0 * index + random.randint(-20, 20),
"power_limit_watts": 700.0,
})
nodes.append({
"hostname": hostname,
"ip": ip,
"is_this_host": is_this_host,
"cpus": 64.0,
"ray_gpus": 4.0,
"available": True,
"error": None,
"gpus": gpus,
})
return {
"mode": "ray",
"nodes": nodes,
"resources": {
"gpus_total": 8.0,
"gpus_available": 5.0
},
"error": None,
}
# --- Uploads ----------------------------------------------------------------
@@ -570,6 +669,92 @@ def download_log(job_id: str) -> PlainTextResponse:
return PlainTextResponse("\n".join(lines) + "\n", media_type="text/plain")
# --- Generators (warm models) -------------------------------------------------
def _advance_generator(entry: dict[str, Any]) -> None:
"""Flip a loading generator to ready once enough wall-clock time has passed.
Like job status, generator state is *computed on read*, so polling the
generators list naturally shows loading -> ready.
"""
if entry["state"] == "loading" and time.time() - entry["started_at"] >= GENERATOR_READY_AFTER_SECONDS:
entry["state"] = "ready"
def _running_inference_ids() -> list[str]:
return [
j["id"] for j in _jobs.values() if j.get("job_type") == "inference" and _public_job(j)["status"] == "running"
]
@app.get("/api/generators")
def list_generators() -> list[dict[str, Any]]:
with _state_lock:
if _generator_slot is None:
return []
_advance_generator(_generator_slot)
return [dict(_generator_slot)]
@app.post("/api/generators/preload", status_code=202)
def preload_generator(req: GeneratorRequest) -> dict[str, Any]:
global _generator_slot
valid_ids = {m["id"] for m in _models_for(None)}
if req.model_id not in valid_ids:
raise HTTPException(
status_code=400,
detail=f"Unknown model_id '{req.model_id}'. Valid options: {sorted(valid_ids)}",
)
with _state_lock:
if _generator_slot is not None:
_advance_generator(_generator_slot)
if _generator_slot["state"] == "loading":
raise HTTPException(status_code=409, detail="a model load is already in progress")
if _generator_slot["state"] == "ready" and all(
_generator_slot.get(k) == v for k, v in req.model_dump().items()):
return dict(_generator_slot)
running = _running_inference_ids()
if running:
raise HTTPException(status_code=409,
detail=f"cannot swap models while inference jobs are running: {running}")
# Loading a new model always replaces (releases) the resident one.
_generator_slot = {"state": "loading", "started_at": time.time(), "error": None, **req.model_dump()}
return dict(_generator_slot)
@app.post("/api/generators/unload")
def unload_generator() -> dict[str, Any]:
global _generator_slot
with _state_lock:
if _generator_slot is not None:
_advance_generator(_generator_slot)
if _generator_slot["state"] == "loading":
raise HTTPException(status_code=409, detail="cannot unload while a model load is in progress")
if _generator_slot is None:
raise HTTPException(status_code=404, detail="No model is loaded")
running = _running_inference_ids()
if running:
raise HTTPException(status_code=409, detail=f"cannot unload while inference jobs are running: {running}")
_generator_slot = None
return {"unloaded": True}
# --- Engine logs --------------------------------------------------------------
@app.get("/api/engine/logs")
def engine_logs(after: int = 0) -> dict[str, Any]:
"""Growing fake tail of the engine's stdout/stderr: every poll appends a
couple of lines so the console visibly streams."""
with _state_lock:
n = len(_engine_log_lines)
_engine_log_lines.append(f"[engine] step {n}: worker heartbeat ok")
_engine_log_lines.append(f"[engine] step {n + 1}: gpu mem {random.randint(20, 80)}% used")
total = len(_engine_log_lines)
return {"lines": _engine_log_lines[max(0, after):], "total": total}
# --- Datasets ---------------------------------------------------------------
+2
View File
@@ -3,6 +3,7 @@
usable by both the real server and the dependency-light mock server."""
from fastvideo_studio.models.create_job_request import CreateJobRequest
from fastvideo_studio.models.generator_request import GeneratorRequest
from fastvideo_studio.models.settings_update import SettingsUpdate
from fastvideo_studio.models.create_dataset_request import CreateDatasetRequest
from fastvideo_studio.models.update_caption_request import UpdateCaptionRequest
@@ -15,6 +16,7 @@ def model_label(model_path: str) -> str:
__all__ = [
"CreateJobRequest",
"GeneratorRequest",
"SettingsUpdate",
"CreateDatasetRequest",
"UpdateCaptionRequest",
@@ -0,0 +1,24 @@
# SPDX-License-Identifier: Apache-2.0
"""Request model for preloading/unloading a resident generator.
Field names and defaults mirror the engine subset of ``CreateJobRequest`` so
the UI can send exactly the values it would put on a job — guaranteeing the
job's generator lookup hits this cache entry.
"""
from pydantic import BaseModel
class GeneratorRequest(BaseModel):
model_id: str
workload_type: str = "t2v"
num_gpus: int = 1
dit_cpu_offload: bool = False
text_encoder_cpu_offload: bool = False
vae_cpu_offload: bool = False
image_encoder_cpu_offload: bool = False
use_fsdp_inference: bool = False
enable_torch_compile: bool = False
vsa_sparsity: float = 0.0
tp_size: int = -1
sp_size: int = -1
+151 -4
View File
@@ -32,10 +32,10 @@ from fastapi.responses import FileResponse
from fastvideo.registry import (get_registered_model_paths, get_registered_models_with_workloads)
from fastvideo_studio.database import Database, _get_db_path
from fastvideo_studio.gpu import get_gpu_snapshot
from fastvideo_studio.gpu import get_cluster_snapshot, get_gpu_snapshot
from fastvideo_studio.job_runner import JobRunner, JobStatus
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, SettingsUpdate, UpdateCaptionRequest,
model_label)
from fastvideo_studio.models import (CreateDatasetRequest, CreateJobRequest, GeneratorRequest, SettingsUpdate,
UpdateCaptionRequest, model_label)
logging.basicConfig(
level=logging.INFO,
@@ -45,6 +45,67 @@ logger = logging.getLogger("fastvideo.studio.api")
DEFAULT_OUTPUT_DIR = os.path.join(os.path.dirname(__file__), "..", "outputs", "ui_jobs")
class _EngineLogBuffer:
"""Thread-safe ring buffer over the server process's stdout/stderr.
Because ray relays worker output to the driver (log_to_driver), teeing the
server's own streams captures engine output from every rank on every node;
under the mp executor, worker logs arrive via the logging handlers which
also write to stderr.
"""
def __init__(self, maxlen: int = 5000) -> None:
import collections
import threading
self._lines: collections.deque[str] = collections.deque(maxlen=maxlen)
self._dropped = 0
self._lock = threading.Lock()
self._partial = ""
on_line = None # optional callable(str), set once at startup
def write(self, text: str) -> None:
with self._lock:
buf = self._partial + text
*complete, self._partial = buf.split("\n")
for line in complete:
if len(self._lines) == self._lines.maxlen:
self._dropped += 1
self._lines.append(line)
if self.on_line is not None:
for line in complete:
# never break stdout on a bad feed
with contextlib.suppress(Exception):
self.on_line(line)
def get_lines(self, after: int = 0) -> tuple[list[str], int]:
with self._lock:
total = self._dropped + len(self._lines)
start = max(0, after - self._dropped)
return list(self._lines)[start:], total
class _Tee:
"""File-like that forwards to the original stream and the ring buffer."""
def __init__(self, orig: Any, buffer: _EngineLogBuffer) -> None:
self._orig = orig
self._buffer = buffer
def write(self, text: str) -> int:
self._buffer.write(text)
return self._orig.write(text)
def flush(self) -> None:
self._orig.flush()
def __getattr__(self, name: str) -> Any:
return getattr(self._orig, name)
engine_log = _EngineLogBuffer()
_available_models: list[dict[str, str]] = [{
"id": path,
"label": model_label(path)
@@ -102,6 +163,13 @@ def list_gpus() -> dict[str, Any]:
return get_gpu_snapshot()
@app.get("/api/cluster")
def cluster_status() -> dict[str, Any]:
"""Cluster-wide GPU/host telemetry (per-node NVML via ray when connected,
the local host otherwise)."""
return get_cluster_snapshot()
@app.get("/api/models")
def list_models(workload_type: str | None = None) -> list[dict[str, Any]]:
"""Return the catalogue of available video-generation models.
@@ -116,6 +184,25 @@ def list_models(workload_type: str | None = None) -> list[dict[str, Any]]:
return _available_models
_PRESET_FIELDS = ("height", "width", "num_frames", "fps", "num_inference_steps", "guidance_scale", "guidance_rescale",
"negative_prompt", "seed")
_preset_cache: dict[str, dict[str, Any]] = {}
@app.get("/api/models/presets")
def model_presets(model_id: str) -> dict[str, Any]:
"""The model's recommended sampling settings (config-only — never loads
weights). The UI populates the job form from these on model selection."""
if model_id not in _preset_cache:
from fastvideo.api.sampling_param import SamplingParam
try:
sp = SamplingParam.from_pretrained(model_id)
except Exception as exc: # noqa: BLE001 -- unknown/unresolvable model
raise HTTPException(status_code=404, detail=f"No presets for '{model_id}': {exc}") from exc
_preset_cache[model_id] = {f: getattr(sp, f) for f in _PRESET_FIELDS if getattr(sp, f, None) is not None}
return _preset_cache[model_id]
ALLOWED_IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
@@ -253,13 +340,58 @@ def get_job(job_id: str) -> dict[str, Any]:
return job.to_dict()
@app.get("/api/generators")
def list_generators() -> list[dict[str, Any]]:
"""Generators resident in memory plus preloads in flight or failed."""
return job_runner.list_generators()
@app.post("/api/generators/preload", status_code=202)
def preload_generator(req: GeneratorRequest) -> dict[str, Any]:
"""Load a model into memory ahead of time (replacing whatever is
resident). One load at a time — 409 while another load is in flight."""
valid_ids = {m["id"] for m in _available_models}
if req.model_id not in valid_ids and not os.path.isdir(req.model_id):
raise HTTPException(
status_code=400,
detail=(f"Unknown model_id '{req.model_id}'. "
f"Valid options: {sorted(valid_ids)}"),
)
try:
return job_runner.preload_generator(**req.model_dump())
except RuntimeError as exc:
raise HTTPException(status_code=409, detail=str(exc)) from exc
@app.post("/api/generators/unload")
def unload_generator() -> dict[str, Any]:
"""Shut down and delete the resident generator, freeing GPU memory."""
try:
unloaded = job_runner.unload_generator()
except RuntimeError as exc:
raise HTTPException(status_code=409, detail=str(exc)) from exc
if not unloaded:
raise HTTPException(status_code=404, detail="No model is loaded")
return {"unloaded": True}
@app.get("/api/engine/logs")
def engine_logs(after: int = 0) -> dict[str, Any]:
"""Incremental tail of the engine's stdout/stderr (driver + relayed
worker output). Poll with ?after=<total from the previous response>."""
lines, total = engine_log.get_lines(after=after)
return {"lines": lines, "total": total}
@app.post("/api/jobs", status_code=201)
def create_job(req: CreateJobRequest) -> dict[str, Any]:
"""Create a new job (does **not** start it automatically)."""
job_type = req.job_type or "inference"
if job_type == "inference":
valid_ids = {m["id"] for m in _available_models}
if req.model_id not in valid_ids:
# A local weights directory is as valid as a registered hub id —
# the registry resolves the pipeline from its model_index.
if req.model_id not in valid_ids and not os.path.isdir(req.model_id):
raise HTTPException(
status_code=400,
detail=(f"Unknown model_id '{req.model_id}'. "
@@ -622,6 +754,14 @@ def create_local_env(host: str, port: int) -> None:
def main() -> None:
import sys
sys.stdout = _Tee(sys.stdout, engine_log)
sys.stderr = _Tee(sys.stderr, engine_log)
# handlers created before the tee (module-level basicConfig) hold the
# original stream objects — re-point them or their output bypasses the buffer
for h in logging.getLogger().handlers:
if isinstance(h, logging.StreamHandler) and h.stream in (sys.__stderr__, sys.__stdout__):
h.setStream(sys.stderr) # type: ignore[arg-type] # duck-typed file-like
global job_runner, database, upload_dir, verbose, datasets_upload_dir # noqa: PLW0603
# Set up signal handlers to prevent worker crashes from killing the server
@@ -686,6 +826,9 @@ def main() -> None:
verbose=args.verbose,
database=database,
)
# ray relays worker output (incl. denoising tqdm) to the driver's stdout;
# feed those lines to the running job so the UI progress bar moves.
engine_log.on_line = job_runner.feed_engine_line
logger.info("Output directory: %s", output_dir)
logger.info("Log directory: %s", log_dir)
@@ -695,6 +838,10 @@ def main() -> None:
host=args.host,
port=args.port,
log_level="info",
# The engine-output console tails this process's stdout/stderr; the
# frontend polls several endpoints every few seconds, so access-log
# lines are pure self-noise there. App/job logging is unaffected.
access_log=False,
)
@@ -0,0 +1,182 @@
import { render, screen } from '@testing-library/react';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import GpusPage from './page';
import { getClusterStatus } from '@/lib/api';
import type { ClusterSnapshot } from '@/lib/api';
vi.mock('@/lib/api', () => ({
getClusterStatus: vi.fn(),
}));
const RAY_SNAPSHOT: ClusterSnapshot = {
mode: 'ray',
error: null,
resources: { gpus_total: 8, gpus_available: 5 },
nodes: [
{
hostname: 'node-a',
ip: '10.0.0.10',
is_this_host: true,
cpus: 64,
ray_gpus: 4,
available: true,
error: null,
gpus: [
{
index: 0,
name: 'NVIDIA B200',
utilization: 62,
memory_used_mib: 40_960,
memory_total_mib: 81_920,
temperature_c: 41,
power_watts: 312.4,
power_limit_watts: 1000,
},
{
index: 1,
name: 'NVIDIA B200',
utilization: 0,
memory_used_mib: 1_024,
memory_total_mib: 81_920,
temperature_c: null,
power_watts: null,
power_limit_watts: null,
},
],
},
{
hostname: 'node-b',
ip: '10.0.0.11',
is_this_host: false,
cpus: 32,
ray_gpus: 2,
available: true,
error: null,
gpus: [
{
index: 0,
name: 'NVIDIA B200',
utilization: 90,
memory_used_mib: 20_480,
memory_total_mib: 81_920,
temperature_c: 70,
power_watts: 900,
power_limit_watts: 1000,
},
],
},
],
};
const LOCAL_SNAPSHOT: ClusterSnapshot = {
mode: 'local',
error:
'not connected to a ray cluster yet (load a model first); showing the API host only',
resources: null,
nodes: [
{
hostname: 'localhost',
ip: null,
is_this_host: true,
cpus: null,
ray_gpus: null,
available: true,
error: null,
gpus: [
{
index: 0,
name: 'NVIDIA RTX 5090',
utilization: 12,
memory_used_mib: 2_048,
memory_total_mib: 32_768,
temperature_c: 38,
power_watts: 80,
power_limit_watts: 575,
},
],
},
],
};
beforeEach(() => {
vi.mocked(getClusterStatus).mockResolvedValue(RAY_SNAPSHOT);
});
describe('GpusPage', () => {
it('renders the header with mode and GPU totals', async () => {
render(<GpusPage />);
expect(await screen.findByText('ray cluster')).toBeInTheDocument();
expect(screen.getByText(/5 \/\s*8 GPUs available/)).toBeInTheDocument();
});
it('renders a section per node with host details and GPU rows', async () => {
render(<GpusPage />);
// Each hostname appears twice: once in the strip, once as a section.
expect(await screen.findAllByText('node-a')).toHaveLength(2);
expect(screen.getAllByText('node-b')).toHaveLength(2);
expect(screen.getByText('10.0.0.10')).toBeInTheDocument();
// Only node-a is the API host.
expect(screen.getAllByText('API host')).toHaveLength(1);
expect(screen.getByText(/64 CPUs · 4 ray GPUs/)).toBeInTheDocument();
expect(screen.getAllByText('NVIDIA B200')).toHaveLength(3);
expect(screen.getByText('GPU 1')).toBeInTheDocument();
expect(screen.getByText('62%')).toBeInTheDocument();
expect(
screen.getByText('40960 / 81920 MiB (40.0 GiB / 80.0 GiB)'),
).toBeInTheDocument();
// Optional sensors render only when present.
expect(screen.getByText('41°C')).toBeInTheDocument();
expect(screen.getByText('312 W / 1000 W')).toBeInTheDocument();
});
it('bars reflect utilization and VRAM values', async () => {
render(<GpusPage />);
await screen.findAllByText('node-a');
const utilMeters = screen
.getAllByRole('meter', { name: 'Utilization' })
.map((m) => m.getAttribute('aria-valuenow'));
expect(utilMeters).toEqual(['62', '0', '90']);
const vramMeters = screen
.getAllByRole('meter', { name: 'VRAM' })
.map((m) => m.getAttribute('aria-valuenow'));
// 40960/81920 = 50%, 1024/81920 ≈ 1%, 20480/81920 = 25%
expect(vramMeters).toEqual(['50', '1', '25']);
});
it('renders the compact strip with per-GPU segments', async () => {
render(<GpusPage />);
await screen.findAllByText('node-a');
const segments = screen.getAllByRole('img');
expect(segments).toHaveLength(3);
expect(segments[0]).toHaveAccessibleName(
'GPU 0: 62% utilization, 40.0 GiB / 80.0 GiB VRAM',
);
});
it('shows the informational banner and local mode', async () => {
vi.mocked(getClusterStatus).mockResolvedValue(LOCAL_SNAPSHOT);
render(<GpusPage />);
expect(await screen.findByText('local host only')).toBeInTheDocument();
expect(
screen.getByText(/not connected to a ray cluster yet/),
).toBeInTheDocument();
// No resources in local mode.
expect(screen.queryByText(/GPUs available/)).not.toBeInTheDocument();
});
it('explains when the API server is unreachable', async () => {
vi.mocked(getClusterStatus).mockRejectedValue(new Error('network down'));
render(<GpusPage />);
expect(
await screen.findByText(/Could not reach the API server/),
).toBeInTheDocument();
});
});
+250 -3
View File
@@ -1,11 +1,258 @@
'use client';
import GpuGrid from '@/components/system/GpuGrid';
import * as React from 'react';
import { AlertTriangle, Info } from 'lucide-react';
export default function GpusPage() {
import ClusterStrip, {
clampPercent,
formatGib,
utilizationColor,
} from '@/components/cluster/ClusterStrip';
import { Badge } from '@/components/ui/badge';
import { Button } from '@/components/ui/button';
import { Card, CardContent } from '@/components/ui/card';
import {
getClusterStatus,
type ClusterGpu,
type ClusterNode,
type ClusterSnapshot,
} from '@/lib/api';
import { cn } from '@/lib/utils';
const POLL_INTERVAL_MS = 5000;
function Meter({
label,
percent,
detail,
fillClass,
}: {
label: string;
percent: number;
detail: string;
fillClass: string;
}) {
const clamped = clampPercent(percent);
return (
<div className="flex flex-col gap-1">
<div className="flex items-baseline justify-between gap-2 text-xs">
<span className="text-muted-foreground">{label}</span>
<span className="font-medium tabular-nums text-foreground">
{detail}
</span>
</div>
<div
role="meter"
aria-label={label}
aria-valuenow={Math.round(clamped)}
aria-valuemin={0}
aria-valuemax={100}
className="h-1.5 overflow-hidden rounded-full bg-muted"
>
<div
className={cn(
'h-full rounded-full transition-[width] duration-500',
fillClass,
)}
style={{ width: `${clamped}%` }}
/>
</div>
</div>
);
}
function GpuRow({ gpu }: { gpu: ClusterGpu }) {
const memPercent =
gpu.memory_total_mib > 0
? (gpu.memory_used_mib / gpu.memory_total_mib) * 100
: 0;
return (
<div className="grid items-center gap-x-6 gap-y-2 border-t border-border pt-3 first:border-t-0 first:pt-0 md:grid-cols-[minmax(0,1fr)_minmax(0,1.2fr)_minmax(0,1.6fr)_auto]">
<div className="flex min-w-0 items-baseline gap-2">
<span className="min-w-0 truncate text-sm font-semibold">
{gpu.name}
</span>
<span className="shrink-0 text-xs font-medium uppercase tracking-wider text-muted-foreground">
GPU {gpu.index}
</span>
</div>
<Meter
label="Utilization"
percent={gpu.utilization}
detail={`${gpu.utilization}%`}
fillClass={utilizationColor(gpu.utilization)}
/>
<Meter
label="VRAM"
percent={memPercent}
detail={`${gpu.memory_used_mib} / ${gpu.memory_total_mib} MiB (${formatGib(gpu.memory_used_mib)} / ${formatGib(gpu.memory_total_mib)})`}
fillClass={memPercent >= 90 ? 'bg-rose-500' : 'bg-accent-blue'}
/>
<div className="flex flex-wrap gap-x-4 gap-y-1 text-xs tabular-nums text-muted-foreground md:w-28 md:justify-end">
{gpu.temperature_c != null && <span>{gpu.temperature_c}°C</span>}
{gpu.power_watts != null && (
<span>
{Math.round(gpu.power_watts)} W
{gpu.power_limit_watts != null &&
` / ${Math.round(gpu.power_limit_watts)} W`}
</span>
)}
</div>
</div>
);
}
function NodeSection({ node }: { node: ClusterNode }) {
return (
<Card>
<CardContent className="flex flex-col gap-3 p-5">
<div className="flex flex-wrap items-center gap-x-3 gap-y-1">
<span className="min-w-0 truncate text-sm font-semibold">
{node.hostname}
</span>
{node.ip && (
<span className="text-xs tabular-nums text-muted-foreground">
{node.ip}
</span>
)}
{node.is_this_host && <Badge variant="secondary">API host</Badge>}
<span className="ml-auto text-xs tabular-nums text-muted-foreground">
{node.cpus != null && `${Math.round(node.cpus)} CPUs`}
{node.cpus != null && node.ray_gpus != null && ' · '}
{node.ray_gpus != null && `${Math.round(node.ray_gpus)} ray GPUs`}
</span>
</div>
{!node.available && (
<p className="text-sm text-muted-foreground">
GPU telemetry unavailable
{node.error ? `: ${node.error}` : '.'}
</p>
)}
{node.gpus.map((gpu) => (
<GpuRow key={gpu.index} gpu={gpu} />
))}
</CardContent>
</Card>
);
}
export default function GpusPage() {
const [snapshot, setSnapshot] = React.useState<ClusterSnapshot | null>(null);
const [fetchError, setFetchError] = React.useState<string | null>(null);
const [retryToken, setRetryToken] = React.useState(0);
React.useEffect(() => {
let mounted = true;
let inFlight = false;
async function poll() {
if (inFlight || document.hidden) return;
inFlight = true;
try {
const next = await getClusterStatus();
if (mounted) {
setSnapshot(next);
setFetchError(null);
}
} catch {
if (mounted) {
setFetchError(
'Cluster status could not be refreshed. The values below may be stale.',
);
}
} finally {
inFlight = false;
}
}
poll();
const interval = setInterval(poll, POLL_INTERVAL_MS);
// Refresh immediately when the tab becomes visible again (polls are
// skipped while hidden).
document.addEventListener('visibilitychange', poll);
return () => {
mounted = false;
clearInterval(interval);
document.removeEventListener('visibilitychange', poll);
};
}, [retryToken]);
let body: React.ReactNode;
if (fetchError && !snapshot) {
body = (
<div
role="alert"
className="flex flex-col items-center gap-3 py-8 text-center"
>
<AlertTriangle className="size-6 text-destructive" aria-hidden />
<p className="text-muted-foreground">
Could not reach the API server. Cluster status needs the Studio API
server running.
</p>
<Button
type="button"
variant="outline"
onClick={() => setRetryToken((token) => token + 1)}
>
Try Again
</Button>
</div>
);
} else if (!snapshot) {
body = <p className="py-8 text-center text-muted-foreground">Loading…</p>;
} else {
body = (
<div className="flex flex-col gap-4">
<header className="flex flex-wrap items-center gap-3">
<h1 className="text-lg font-semibold">Cluster</h1>
<Badge variant="outline">
{snapshot.mode === 'ray' ? 'ray cluster' : 'local host only'}
</Badge>
{snapshot.resources && (
<span className="text-sm tabular-nums text-muted-foreground">
{Math.round(snapshot.resources.gpus_available)} /{' '}
{Math.round(snapshot.resources.gpus_total)} GPUs available
</span>
)}
</header>
{snapshot.error && (
<div
role="status"
className="flex flex-wrap items-center gap-3 rounded-lg border border-blue-400/40 bg-blue-500/10 px-3 py-2 text-sm"
>
<Info className="size-4 shrink-0 text-blue-600" aria-hidden />
<span className="min-w-0 flex-1">{snapshot.error}</span>
</div>
)}
{fetchError && (
<div
role="status"
aria-live="polite"
className="flex flex-wrap items-center gap-3 rounded-lg border border-amber-500/50 bg-amber-500/10 px-3 py-2 text-sm"
>
<AlertTriangle className="size-4 text-amber-600" aria-hidden />
<span className="min-w-0 flex-1">{fetchError}</span>
<Button
type="button"
variant="outline"
size="sm"
onClick={() => setRetryToken((token) => token + 1)}
>
Refresh Now
</Button>
</div>
)}
<ClusterStrip nodes={snapshot.nodes} />
{snapshot.nodes.map((node, i) => (
<NodeSection key={`${node.hostname}-${i}`} node={node} />
))}
</div>
);
}
return (
<div className="mx-auto flex w-full max-w-[1100px] flex-col gap-6 px-4 pb-12 pt-6">
<GpuGrid />
{body}
</div>
);
}
@@ -1,8 +1,10 @@
'use client';
import CreateJobButton from '@/components/jobs/CreateJobButton';
import EngineConsole from '@/components/jobs/EngineConsole';
import { HeaderActions } from '@/components/shell/HeaderActionsContext';
import JobQueue from '@/components/jobs/JobQueue';
import WarmModelsPanel from '@/components/jobs/WarmModelsPanel';
export default function InferencePage() {
return (
@@ -10,6 +12,8 @@ export default function InferencePage() {
<HeaderActions>
<CreateJobButton jobType="inference" />
</HeaderActions>
<WarmModelsPanel />
<EngineConsole />
<JobQueue jobType="inference" />
</>
);
@@ -0,0 +1,93 @@
'use client';
import type { ClusterNode } from '@/lib/api';
import { cn } from '@/lib/utils';
export function formatGib(mib: number): string {
return `${(mib / 1024).toFixed(1)} GiB`;
}
/** Bar fill class by load: blue when idle, amber under pressure, rose hot. */
export function utilizationColor(percent: number): string {
if (percent >= 85) return 'bg-rose-500';
if (percent >= 50) return 'bg-amber-500';
return 'bg-accent-blue';
}
export function clampPercent(percent: number): number {
return Math.max(0, Math.min(100, percent));
}
function GpuSegment({
index,
utilization,
memUsedMib,
memTotalMib,
}: {
index: number;
utilization: number;
memUsedMib: number;
memTotalMib: number;
}) {
const memPercent =
memTotalMib > 0 ? clampPercent((memUsedMib / memTotalMib) * 100) : 0;
const label =
`GPU ${index}: ${utilization}% utilization, ` +
`${formatGib(memUsedMib)} / ${formatGib(memTotalMib)} VRAM`;
return (
<div
role="img"
aria-label={label}
title={label}
className="flex w-10 shrink-0 flex-col gap-0.5"
>
<div className="h-1.5 overflow-hidden rounded-full bg-muted">
<div
className={cn('h-full rounded-full', utilizationColor(utilization))}
style={{ width: `${clampPercent(utilization)}%` }}
/>
</div>
<div className="h-1.5 overflow-hidden rounded-full bg-muted">
<div
className={cn(
'h-full rounded-full',
memPercent >= 90 ? 'bg-rose-500' : 'bg-accent-blue',
)}
style={{ width: `${memPercent}%` }}
/>
</div>
</div>
);
}
/** One compact line per node: hostname + tiny util/VRAM bars per GPU. */
export default function ClusterStrip({ nodes }: { nodes: ClusterNode[] }) {
return (
<div className="flex flex-col gap-2">
{nodes.map((node, i) => (
<div
key={`${node.hostname}-${i}`}
className="flex items-center gap-3"
>
<span className="w-40 shrink-0 truncate text-xs font-medium">
{node.hostname}
</span>
<div className="flex min-w-0 flex-wrap items-center gap-1.5">
{node.gpus.map((gpu) => (
<GpuSegment
key={gpu.index}
index={gpu.index}
utilization={gpu.utilization}
memUsedMib={gpu.memory_used_mib}
memTotalMib={gpu.memory_total_mib}
/>
))}
{node.gpus.length === 0 && (
<span className="text-xs text-muted-foreground">no GPUs</span>
)}
</div>
</div>
))}
</div>
);
}
@@ -4,14 +4,24 @@ import userEvent from '@testing-library/user-event';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import CreateJobModal from './CreateJobModal';
import { createJob, getDatasets, getModels, uploadImage } from '@/lib/api';
import {
createJob,
getDatasets,
getModelPresets,
getModels,
listGenerators,
uploadImage,
type GeneratorInfo,
} from '@/lib/api';
import { defaultOptionsStore } from '@/stores/defaultOptions';
import { DEFAULT_OPTIONS } from '@/lib/defaultOptions';
vi.mock('@/lib/api', () => ({
createJob: vi.fn(),
getModels: vi.fn(),
getModelPresets: vi.fn(),
getDatasets: vi.fn(),
listGenerators: vi.fn(),
uploadImage: vi.fn(),
getSettings: vi.fn(),
updateSettings: vi.fn(),
@@ -22,11 +32,32 @@ const MODELS = [
{ id: 'wan/t2v-14b', label: 'Wan T2V Large', type: 't2v' },
];
// A resident engine slot whose engine config differs from the persisted
// defaults on every field the modal adopts.
const WARM_SLOT: GeneratorInfo = {
state: 'ready',
model_id: 'wan/t2v-14b',
workload_type: 't2v',
num_gpus: 8,
dit_cpu_offload: true,
text_encoder_cpu_offload: true,
vae_cpu_offload: true,
image_encoder_cpu_offload: false,
use_fsdp_inference: true,
enable_torch_compile: true,
vsa_sparsity: 0.5,
tp_size: 1,
sp_size: 8,
error: null,
};
beforeEach(() => {
// Reset the shared options store to a known baseline for test isolation.
defaultOptionsStore.set({ options: DEFAULT_OPTIONS });
vi.mocked(getModels).mockResolvedValue(MODELS);
vi.mocked(getModelPresets).mockResolvedValue({});
vi.mocked(getDatasets).mockResolvedValue([]);
vi.mocked(listGenerators).mockResolvedValue([]);
vi.mocked(uploadImage).mockResolvedValue({ path: '/uploads/x.png' });
vi.mocked(createJob).mockResolvedValue({ id: 'job-1' } as never);
});
@@ -206,4 +237,173 @@ describe('CreateJobModal', () => {
await waitFor(() => expect(onSuccess).toHaveBeenCalledTimes(1));
});
it('defaults to the warm resident model and adopts its engine config', async () => {
vi.mocked(listGenerators).mockResolvedValue([WARM_SLOT]);
const user = userEvent.setup();
renderModal();
// Without a warm model the default logic picks the first model
// (wan/t2v-1.3b, covered by the seeding test above); the warm slot wins.
await waitFor(() =>
expect(screen.getByLabelText('Model')).toHaveValue('wan/t2v-14b'),
);
// The warm selection also triggers its presets fetch.
await waitFor(() =>
expect(getModelPresets).toHaveBeenCalledWith('wan/t2v-14b'),
);
await user.type(screen.getByLabelText('Prompt'), 'warm run');
await user.click(screen.getByRole('button', { name: 'Create Job' }));
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
// The engine fields mirror the resident slot (not the persisted defaults:
// num_gpus 1, all offloads/fsdp/compile false) so the job reuses the warm
// instance instead of replacing it.
expect(vi.mocked(createJob).mock.calls[0][0]).toMatchObject({
model_id: 'wan/t2v-14b',
num_gpus: 8,
tp_size: 1,
sp_size: 8,
dit_cpu_offload: true,
text_encoder_cpu_offload: true,
vae_cpu_offload: true,
image_encoder_cpu_offload: false,
use_fsdp_inference: true,
enable_torch_compile: true,
vsa_sparsity: 0.5,
});
});
it('restores persisted-default engine fields when switching away from the warm model', async () => {
defaultOptionsStore.set({
options: { ...DEFAULT_OPTIONS, numGpus: 2, tpSize: 2 },
});
vi.mocked(listGenerators).mockResolvedValue([WARM_SLOT]);
const user = userEvent.setup();
renderModal();
await waitFor(() =>
expect(screen.getByLabelText('Model')).toHaveValue('wan/t2v-14b'),
);
await user.selectOptions(screen.getByLabelText('Model'), 'wan/t2v-1.3b');
await user.type(screen.getByLabelText('Prompt'), 'cold run');
await user.click(screen.getByRole('button', { name: 'Create Job' }));
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
expect(vi.mocked(createJob).mock.calls[0][0]).toMatchObject({
model_id: 'wan/t2v-1.3b',
num_gpus: 2,
tp_size: 2,
use_fsdp_inference: false,
enable_torch_compile: false,
});
});
it('applies resolution preset chips and the orientation toggle to the payload', async () => {
const user = userEvent.setup();
renderModal();
await screen.findByRole('option', { name: 'Wan T2V (wan/t2v-1.3b)' });
await user.click(screen.getByText('Options'));
// 720p chip sets the /32-rounded dims; the orientation toggle swaps them.
await user.click(screen.getByRole('button', { name: '720p' }));
await user.click(screen.getByRole('button', { name: 'Landscape' }));
expect(
screen.getByRole('button', { name: 'Portrait' }),
).toBeInTheDocument();
await user.type(screen.getByLabelText('Prompt'), 'portrait 720p');
await user.click(screen.getByRole('button', { name: 'Create Job' }));
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
expect(vi.mocked(createJob).mock.calls[0][0]).toMatchObject({
height: 1280,
width: 704,
});
});
it('restores the model preset resolution via the Native chip', async () => {
vi.mocked(getModelPresets).mockResolvedValue({ height: 720, width: 1280 });
const user = userEvent.setup();
renderModal();
await screen.findByRole('option', { name: 'Wan T2V (wan/t2v-1.3b)' });
// Native is disabled until the selected model's presets have loaded.
await waitFor(() =>
expect(screen.getByRole('button', { name: 'Native' })).toBeEnabled(),
);
await user.click(screen.getByText('Options'));
await user.click(screen.getByRole('button', { name: '1080p' }));
await user.click(screen.getByRole('button', { name: 'Native' }));
await user.type(screen.getByLabelText('Prompt'), 'native dims');
await user.click(screen.getByRole('button', { name: 'Create Job' }));
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
expect(vi.mocked(createJob).mock.calls[0][0]).toMatchObject({
height: 720,
width: 1280,
});
});
it('populates sampling fields from the selected model presets; engine fields stay from defaults', async () => {
defaultOptionsStore.set({
options: { ...DEFAULT_OPTIONS, numGpus: 4, tpSize: 2, seed: 999 },
});
vi.mocked(getModelPresets).mockImplementation(async (id) =>
id === 'wan/t2v-14b'
? {
height: 720,
width: 1280,
num_frames: 121,
fps: 30,
num_inference_steps: 40,
guidance_scale: 6,
guidance_rescale: 0.5,
negative_prompt: 'blurry, low quality',
seed: 7,
}
: {},
);
const user = userEvent.setup();
renderModal();
await screen.findByRole('option', { name: 'Wan T2V Large (wan/t2v-14b)' });
await user.selectOptions(screen.getByLabelText('Model'), 'wan/t2v-14b');
// The negative prompt is the easiest preset-populated field to observe.
await waitFor(() =>
expect(screen.getByLabelText('Negative Prompt')).toHaveValue(
'blurry, low quality',
),
);
await user.type(screen.getByLabelText('Prompt'), 'preset test');
await user.click(screen.getByRole('button', { name: 'Create Job' }));
await waitFor(() => expect(createJob).toHaveBeenCalledTimes(1));
const payload = vi.mocked(createJob).mock.calls[0][0];
expect(payload).toMatchObject({
model_id: 'wan/t2v-14b',
// Sampling fields come from the model presets…
height: 720,
width: 1280,
num_frames: 121,
fps: 30,
num_inference_steps: 40,
guidance_scale: 6,
guidance_rescale: 0.5,
negative_prompt: 'blurry, low quality',
seed: 7,
// …while engine fields still come from the persisted defaults.
num_gpus: 4,
tp_size: 2,
});
});
});
@@ -23,15 +23,27 @@ import { defaultOptionsStore } from '@/stores/defaultOptions';
import {
createJob,
getDatasets,
getModelPresets,
getModels,
listGenerators,
uploadImage,
type CreateJobRequest,
type GeneratorInfo,
type Model,
} from '@/lib/api';
import { getDefaultModelForWorkload } from '@/lib/defaultOptions';
import { WORKLOAD_OPTIONS } from '@/lib/jobConfig';
import type { JobType } from '@/lib/types';
// ponytail: dims pre-rounded to multiples of 32 so every model family accepts
// them (H3 requires /32) — hence 720p→704 and 1080p→1088. Tooltips show the
// exact dims; the labels stay 480p/720p/1080p.
const RESOLUTION_PRESETS = [
{ label: '480p', height: 480, width: 832 },
{ label: '720p', height: 704, width: 1280 },
{ label: '1080p', height: 1088, width: 1920 },
] as const;
export interface CreateJobModalProps {
isOpen: boolean;
onClose: () => void;
@@ -55,6 +67,9 @@ export default function CreateJobModal({
const [models, setModels] = React.useState<Model[]>([]);
const [modelId, setModelId] = React.useState('');
// The engine's resident slot (null when empty/failed), fetched on open so
// jobs against the warm model can adopt its engine config.
const [warmSlot, setWarmSlot] = React.useState<GeneratorInfo | null>(null);
const [prompt, setPrompt] = React.useState('');
const [imagePath, setImagePath] = React.useState('');
const [imageFileName, setImageFileName] = React.useState('');
@@ -64,6 +79,12 @@ export default function CreateJobModal({
const [numFrames, setNumFrames] = React.useState(81);
const [height, setHeight] = React.useState(480);
const [width, setWidth] = React.useState(832);
// The selected model's preset dims, kept so the "Native" chip can restore
// them after a resolution chip or manual slider edit.
const [nativeDims, setNativeDims] = React.useState<{
height: number;
width: number;
} | null>(null);
const [guidanceScale, setGuidanceScale] = React.useState(5);
const [guidanceRescale, setGuidanceRescale] = React.useState(0);
const [fps, setFps] = React.useState(24);
@@ -178,17 +199,34 @@ export default function CreateJobModal({
let stale = false;
setIsLoadingModels(true);
setModelLoadError(null);
getModels(inferenceWorkload)
.then((list) => {
// For inference jobs, prefer the warm resident model (if any) over the
// persisted default so new jobs hit the already-loaded engine. The warm
// slot is a nice-to-have: if the lookup fails, fall back silently.
const warmSlotPromise = isInference
? listGenerators()
.then((gens) =>
gens[0] && gens[0].state !== 'failed' ? gens[0] : null,
)
.catch(() => null)
: Promise.resolve(null);
Promise.all([getModels(inferenceWorkload), warmSlotPromise])
.then(([list, slot]) => {
if (stale) return;
setModels(list);
setWarmSlot(slot);
const ids = list.map((m) => m.id);
const opts = defaultOptionsStore.get().options;
const defaultId = getDefaultModelForWorkload(
opts,
inferenceWorkload as 't2v' | 'i2v' | 't2i',
);
const chosen = ids.includes(defaultId) ? defaultId : (list[0]?.id ?? '');
const warmId = slot?.model_id ?? '';
const chosen =
warmId && ids.includes(warmId)
? warmId
: ids.includes(defaultId)
? defaultId
: (list[0]?.id ?? '');
setModelId(chosen);
if (workloadType === 'dmd_t2v') {
setRealScoreModelPath(chosen);
@@ -210,7 +248,81 @@ export default function CreateJobModal({
return () => {
stale = true;
};
}, [isOpen, inferenceWorkload, workloadType]);
}, [isOpen, isInference, inferenceWorkload, workloadType]);
// Whenever the selected model changes (including the initial selection on
// open), populate the sampling fields from that model's presets. Missing
// keys leave the field as-is; engine fields (GPUs, parallelism, offloads)
// come from the warm slot or persisted defaults (effect below). Latest
// selection wins: the cleanup marks superseded fetches stale, same as the
// model-list effect.
React.useEffect(() => {
if (!isOpen || !isInference || !modelId) return;
let stale = false;
setNativeDims(null);
getModelPresets(modelId)
.then((p) => {
if (stale) return;
if (p.height !== undefined) setHeight(p.height);
if (p.width !== undefined) setWidth(p.width);
if (p.height !== undefined && p.width !== undefined)
setNativeDims({ height: p.height, width: p.width });
if (p.num_frames !== undefined)
setNumFrames(workloadType === 't2i' ? 1 : p.num_frames);
if (p.fps !== undefined) setFps(p.fps);
if (p.num_inference_steps !== undefined)
setNumInferenceSteps(p.num_inference_steps);
if (p.guidance_scale !== undefined) setGuidanceScale(p.guidance_scale);
if (p.guidance_rescale !== undefined)
setGuidanceRescale(p.guidance_rescale);
if (p.negative_prompt !== undefined)
setNegativePrompt(p.negative_prompt);
if (p.seed !== undefined) setSeed(p.seed);
})
.catch((e) => {
// Presets are a convenience; on failure keep the current values.
console.error('Failed to load model presets:', e);
});
return () => {
stale = true;
};
}, [isOpen, isInference, modelId, workloadType]);
// When the selected model IS the warm resident one, adopt the slot's engine
// config so the job reuses the loaded instance instead of silently replacing
// it (e.g. persisted num_gpus=1 vs a warm 8-GPU slot). Switching away from
// the warm model restores the persisted-default engine fields; non-warm to
// non-warm switches leave the user's engine edits alone (today's behavior).
const wasWarmRef = React.useRef(false);
React.useEffect(() => {
if (!isOpen || !isInference || !modelId) return;
const isWarm = warmSlot?.model_id === modelId;
if (isWarm && warmSlot) {
setNumGpus(warmSlot.num_gpus);
setTpSize(warmSlot.tp_size);
setSpSize(warmSlot.sp_size);
setDitCpuOffload(warmSlot.dit_cpu_offload);
setTextEncoderCpuOffload(warmSlot.text_encoder_cpu_offload);
setVaeCpuOffload(warmSlot.vae_cpu_offload);
setImageEncoderCpuOffload(warmSlot.image_encoder_cpu_offload);
setUseFsdpInference(warmSlot.use_fsdp_inference);
setEnableTorchCompile(warmSlot.enable_torch_compile);
setVsaSparsity(warmSlot.vsa_sparsity);
} else if (wasWarmRef.current) {
const opts = defaultOptionsStore.get().options;
setNumGpus(opts.numGpus);
setTpSize(opts.tpSize);
setSpSize(opts.spSize);
setDitCpuOffload(opts.ditCpuOffload);
setTextEncoderCpuOffload(opts.textEncoderCpuOffload);
setVaeCpuOffload(opts.vaeCpuOffload);
setImageEncoderCpuOffload(opts.imageEncoderCpuOffload);
setUseFsdpInference(opts.useFsdpInference);
setEnableTorchCompile(opts.enableTorchCompile);
setVsaSparsity(opts.vsaSparsity);
}
wasWarmRef.current = isWarm;
}, [isOpen, isInference, modelId, warmSlot]);
// Training jobs need a dataset; load the ready datasets when relevant.
React.useEffect(() => {
@@ -741,6 +853,64 @@ export default function CreateJobModal({
disabled={isSubmitting}
/>
)}
<div className="col-span-full flex flex-wrap items-center gap-1.5">
<span className="pl-0.5 text-xs font-normal tracking-wide text-muted-foreground">
Resolution
</span>
{RESOLUTION_PRESETS.map((preset) => (
<Button
key={preset.label}
type="button"
size="sm"
variant="outline"
title={`${preset.height}×${preset.width}`}
onClick={() => {
// Apply in the current orientation so a portrait
// setup stays portrait when switching resolution.
const portrait = height > width;
setHeight(portrait ? preset.width : preset.height);
setWidth(portrait ? preset.height : preset.width);
}}
disabled={isSubmitting}
className="h-6 px-2 text-xs"
>
{preset.label}
</Button>
))}
<Button
type="button"
size="sm"
variant="outline"
title={
nativeDims
? `${nativeDims.height}×${nativeDims.width}`
: 'Model preset resolution (unavailable)'
}
onClick={() => {
if (!nativeDims) return;
setHeight(nativeDims.height);
setWidth(nativeDims.width);
}}
disabled={isSubmitting || !nativeDims}
className="h-6 px-2 text-xs"
>
Native
</Button>
<Button
type="button"
size="sm"
variant="outline"
title="Swap height and width"
onClick={() => {
setHeight(width);
setWidth(height);
}}
disabled={isSubmitting}
className="h-6 px-2 text-xs"
>
{height > width ? 'Portrait' : 'Landscape'}
</Button>
</div>
<SliderRow
id="modal-height"
label="Height"
@@ -0,0 +1,71 @@
import { act, fireEvent, render, screen } from '@testing-library/react';
import userEvent from '@testing-library/user-event';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import EngineConsole from './EngineConsole';
import { getEngineLogs } from '@/lib/api';
vi.mock('@/lib/api', () => ({
getEngineLogs: vi.fn(),
}));
beforeEach(() => {
vi.mocked(getEngineLogs).mockResolvedValue({
lines: ['[engine] booted', '[engine] worker heartbeat ok'],
total: 2,
});
});
describe('EngineConsole', () => {
it('is collapsed by default and does not fetch', () => {
render(<EngineConsole />);
expect(
screen.getByRole('button', { name: 'Engine output' }),
).toHaveAttribute('aria-expanded', 'false');
expect(getEngineLogs).not.toHaveBeenCalled();
});
it('shows log lines and polls with the cursor while open', async () => {
vi.useFakeTimers();
try {
render(<EngineConsole />);
fireEvent.click(screen.getByRole('button', { name: 'Engine output' }));
// Flush the immediate poll fired on expand.
await act(async () => {
await vi.advanceTimersByTimeAsync(0);
});
expect(getEngineLogs).toHaveBeenCalledWith(0);
expect(screen.getByText(/\[engine\] booted/)).toBeInTheDocument();
// The 2s interval polls again, from the previous total.
await act(async () => {
await vi.advanceTimersByTimeAsync(2000);
});
expect(getEngineLogs).toHaveBeenLastCalledWith(2);
// Collapsing stops the polling.
fireEvent.click(screen.getByRole('button', { name: 'Engine output' }));
const calls = vi.mocked(getEngineLogs).mock.calls.length;
await act(async () => {
await vi.advanceTimersByTimeAsync(10000);
});
expect(getEngineLogs).toHaveBeenCalledTimes(calls);
} finally {
vi.useRealTimers();
}
});
it('clear view empties the scrollback locally', async () => {
const user = userEvent.setup();
render(<EngineConsole />);
await user.click(screen.getByRole('button', { name: 'Engine output' }));
expect(await screen.findByText(/\[engine\] booted/)).toBeInTheDocument();
await user.click(screen.getByRole('button', { name: 'Clear view' }));
expect(screen.queryByText(/\[engine\] booted/)).not.toBeInTheDocument();
expect(screen.getByText('Waiting for engine output…')).toBeInTheDocument();
});
});
@@ -0,0 +1,123 @@
'use client';
import * as React from 'react';
import { ChevronDown, ChevronRight } from 'lucide-react';
import { Button } from '@/components/ui/button';
import { getEngineLogs } from '@/lib/api';
const POLL_INTERVAL_MS = 2000;
// Cap the DOM at the last ~500 lines; the server keeps its own ring buffer.
const MAX_LINES = 500;
// The engine tail includes uvicorn's access log (until a server restart picks
// up access_log=False); the frontend's own polling would otherwise flood the
// console with "GET /api/... 200 OK" lines. Keep non-GET and non-2xx lines.
const ACCESS_LOG_NOISE = /^INFO:\s+[\d.:]+\s+- "(?:GET|HEAD) \S+ HTTP\/[\d.]+" 2\d\d/;
/**
* Collapsible tail of the engine's stdout/stderr (driver + relayed worker
* output). Polls only while open; sticks to the bottom unless the user has
* scrolled up.
*/
export default function EngineConsole() {
const [open, setOpen] = React.useState(false);
const [lines, setLines] = React.useState<string[]>([]);
// Poll cursor + stick-to-bottom flag live in refs: they must update
// synchronously from async polls / scroll events, outside React's cycle.
const afterRef = React.useRef(0);
const stickRef = React.useRef(true);
const consoleRef = React.useRef<HTMLPreElement | null>(null);
React.useEffect(() => {
if (!open) return;
let mounted = true;
let locked = false;
async function poll() {
if (!mounted || locked) return;
locked = true;
try {
const data = await getEngineLogs(afterRef.current);
afterRef.current = data.total;
const fresh = data.lines.filter((l) => !ACCESS_LOG_NOISE.test(l));
if (mounted && fresh.length > 0) {
setLines((prev) => [...prev, ...fresh].slice(-MAX_LINES));
}
} catch (e) {
console.error('Failed to fetch engine logs:', e);
} finally {
locked = false;
}
}
poll();
const interval = setInterval(poll, POLL_INTERVAL_MS);
return () => {
mounted = false;
clearInterval(interval);
};
}, [open]);
// Follow the tail after new lines land, unless the user scrolled up.
React.useEffect(() => {
const el = consoleRef.current;
if (el && stickRef.current) el.scrollTop = el.scrollHeight;
}, [lines]);
function handleScroll() {
const el = consoleRef.current;
if (!el) return;
stickRef.current = el.scrollHeight - el.scrollTop - el.clientHeight < 40;
}
return (
<section
aria-label="Engine output"
className="mx-auto w-full max-w-[850px] px-10 pt-3"
>
<div className="rounded-lg border border-border bg-background">
<div className="flex items-center gap-2 px-2 py-1.5">
<button
type="button"
onClick={() => setOpen((o) => !o)}
aria-expanded={open}
className="flex flex-1 items-center gap-2 rounded-md px-2 py-1 text-sm font-semibold text-foreground hover:bg-accent"
>
{open ? (
<ChevronDown className="h-4 w-4" />
) : (
<ChevronRight className="h-4 w-4" />
)}
Engine output
</button>
{open && (
<Button
size="sm"
variant="ghost"
// Resets the local view only; the server buffer is untouched.
onClick={() => setLines([])}
>
Clear view
</Button>
)}
</div>
{open && (
<pre
ref={consoleRef}
onScroll={handleScroll}
className="m-0 h-64 overflow-auto whitespace-pre-wrap break-words rounded-b-lg border-t border-border bg-zinc-950 p-3 font-mono text-xs leading-normal text-zinc-200"
>
{lines.length === 0 ? (
<span className="italic text-zinc-500">
Waiting for engine output…
</span>
) : (
lines.join('\n')
)}
</pre>
)}
</div>
</section>
);
}
@@ -142,7 +142,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
return (
<article
className={cn(
'mb-3 flex cursor-pointer flex-col gap-2.5 rounded-lg border bg-background p-4 transition-colors last:mb-0',
'mb-1.5 flex cursor-pointer flex-col gap-1 rounded-lg border bg-background px-3 py-1.5 transition-colors last:mb-0',
isSelected
? 'border-accent-blue bg-accent-blue/5'
: 'border-border hover:border-muted-foreground/40',
@@ -152,20 +152,16 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
type="button"
aria-pressed={isSelected}
onClick={handleSelectJob}
className="flex w-full flex-col gap-2.5 rounded-md text-left"
className="flex w-full min-w-0 flex-col gap-1 rounded-md text-left"
>
<span className="flex flex-wrap items-center justify-between gap-2">
<span className="text-[0.95rem] font-semibold text-foreground">
{job.model_id}
</span>
<Badge variant={BADGE_VARIANTS[job.status] ?? 'secondary'}>
{job.status}
</Badge>
<span className="flex w-full min-w-0 items-center gap-2">
<span className="shrink-0 text-sm font-semibold text-foreground">
{job.model_id}
</span>
<span className="max-w-full overflow-hidden text-ellipsis whitespace-nowrap text-sm text-muted-foreground">
{job.prompt}
</span>
<span className="flex flex-wrap items-center gap-4 text-xs text-muted-foreground">
<Badge variant={BADGE_VARIANTS[job.status] ?? 'secondary'}>
{job.status}
</Badge>
<span className="ml-auto flex shrink-0 items-center gap-3 text-xs text-muted-foreground">
{job.job_type === 'inference' ? (
<>
<span>{job.num_frames} frames</span>
@@ -183,6 +179,10 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
</span>
)}
</span>
</span>
<span className="w-full whitespace-pre-wrap break-words text-xs text-muted-foreground">
{job.prompt}
</span>
</button>
<div className="flex flex-wrap items-center gap-1.5">
{job.status === 'running' ? (
@@ -190,7 +190,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
size="sm"
onClick={handleStop}
disabled={isLoading}
className="border-transparent bg-amber-500 text-black shadow-md hover:bg-amber-400"
className="h-6 px-2 text-xs border-transparent bg-amber-500 text-black shadow-md hover:bg-amber-400"
>
Stop
</Button>
@@ -199,7 +199,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
size="sm"
onClick={handleStart}
disabled={isLoading}
className="border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
className="h-6 px-2 text-xs border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
>
Restart
</Button>
@@ -208,7 +208,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
size="sm"
onClick={handleStart}
disabled={isLoading}
className="border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
className="h-6 px-2 text-xs border-transparent bg-emerald-600 text-white shadow-md hover:bg-emerald-500"
>
Start
</Button>
@@ -222,6 +222,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
onClick={handleDownloadVideo}
disabled={isLoading}
title="Download video"
className="h-6 px-2 text-xs"
>
Download Video
</Button>
@@ -231,6 +232,7 @@ export default function JobCard({ job, onJobUpdated }: JobCardProps) {
variant="destructive"
onClick={handleDelete}
disabled={isLoading}
className="h-6 px-2 text-xs"
>
Delete
</Button>
@@ -9,6 +9,7 @@ import { makeJob as makeBaseJob } from '@/test/factories';
vi.mock('@/lib/api', () => ({
getJobLogs: vi.fn(),
downloadJobLog: vi.fn(),
getJobVideoUrl: (id: string) => `http://test.local/api/jobs/${id}/video`,
}));
const makeJob = (overrides: Partial<Job> = {}): Job =>
@@ -45,6 +46,43 @@ describe('JobDetailsSidebar', () => {
expect(onWidthChange).toHaveBeenCalledWith(0);
});
it('plays completed inference output inline; running jobs get no player', async () => {
vi.mocked(getJobLogs).mockResolvedValue({
lines: [],
total: 0,
progress: 0,
progress_msg: '',
phase: '',
});
const { rerender } = render(
<JobDetailsSidebar
job={makeJob({
status: 'completed',
output_path: '/outputs/job-1.mp4',
prompt: 'a cat surfing a wave',
})}
onClose={vi.fn()}
/>,
);
const video = screen.getByLabelText('Generated video: a cat surfing a wave');
expect(video.tagName).toBe('VIDEO');
expect(video).toHaveAttribute('controls');
expect(video).toHaveAttribute(
'src',
'http://test.local/api/jobs/job-1/video',
);
rerender(
<JobDetailsSidebar
job={makeJob({ status: 'running', output_path: null })}
onClose={vi.fn()}
/>,
);
expect(screen.queryByLabelText(/Generated video/)).not.toBeInTheDocument();
});
it('renders log lines streamed from the job log poll', async () => {
vi.mocked(getJobLogs).mockResolvedValue({
lines: ['boot sequence started', 'loading model weights'],
@@ -6,7 +6,7 @@ import { X } from 'lucide-react';
import { Button } from '@/components/ui/button';
import { useDrawerFocus } from '@/hooks/useDrawerFocus';
import { useResizable } from '@/hooks/useResizable';
import { downloadJobLog, getJobLogs } from '@/lib/api';
import { downloadJobLog, getJobLogs, getJobVideoUrl } from '@/lib/api';
import type { Job } from '@/lib/types';
import { cn, downloadBlob } from '@/lib/utils';
@@ -180,6 +180,41 @@ export default function JobDetailsSidebar({
</div>
</div>
{job.status === 'completed' &&
job.output_path &&
(job.job_type === 'inference' || !job.job_type) && (
<div className="border-b border-border px-5 py-4">
<span className="mb-2 block text-xs font-semibold uppercase tracking-wider text-muted-foreground">
Output
</span>
{job.output_path.toLowerCase().endsWith('.png') ? (
// eslint-disable-next-line @next/next/no-img-element
<img
src={getJobVideoUrl(job.id)}
alt={
job.prompt
? `Generated image: ${job.prompt}`
: 'Generated image'
}
className="block w-full rounded-lg border border-border bg-background object-contain"
/>
) : (
<video
src={getJobVideoUrl(job.id)}
aria-label={
job.prompt
? `Generated video: ${job.prompt}`
: 'Generated video'
}
controls
playsInline
preload="metadata"
className="block max-h-80 w-full rounded-lg border border-border bg-background object-contain"
/>
)}
</div>
)}
<div className="flex min-h-0 flex-1 flex-col px-5 py-4">
<div className="mb-2 flex items-center justify-between">
<span className="text-xs font-semibold uppercase tracking-wider text-muted-foreground">
@@ -0,0 +1,188 @@
import { render, screen, waitFor } from '@testing-library/react';
import userEvent from '@testing-library/user-event';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import WarmModelsPanel from './WarmModelsPanel';
import {
getModels,
listGenerators,
preloadGenerator,
unloadGenerator,
type GeneratorInfo,
} from '@/lib/api';
import { DEFAULT_OPTIONS } from '@/lib/defaultOptions';
import { defaultOptionsStore } from '@/stores/defaultOptions';
import { toast } from 'sonner';
vi.mock('@/lib/api', () => ({
getModels: vi.fn(),
listGenerators: vi.fn(),
preloadGenerator: vi.fn(),
unloadGenerator: vi.fn(),
getSettings: vi.fn(),
updateSettings: vi.fn(),
}));
vi.mock('sonner', () => ({
toast: { error: vi.fn() },
}));
const makeGenerator = (
overrides: Partial<GeneratorInfo> = {},
): GeneratorInfo => ({
state: 'ready',
model_id: 'Wan-AI/Wan2.1-T2V-1.3B-Diffusers',
workload_type: 't2v',
num_gpus: 1,
dit_cpu_offload: false,
text_encoder_cpu_offload: false,
vae_cpu_offload: false,
image_encoder_cpu_offload: false,
use_fsdp_inference: false,
enable_torch_compile: false,
vsa_sparsity: 0,
tp_size: -1,
sp_size: -1,
error: null,
...overrides,
});
beforeEach(() => {
// Reset the shared options store to a known baseline for test isolation.
defaultOptionsStore.set({ options: DEFAULT_OPTIONS });
vi.mocked(getModels).mockResolvedValue([
{ id: 'Wan-AI/Wan2.1-T2V-1.3B-Diffusers', label: 'Wan2.1 T2V 1.3B' },
]);
vi.mocked(listGenerators).mockResolvedValue([]);
vi.mocked(preloadGenerator).mockResolvedValue(
makeGenerator({ state: 'loading' }),
);
vi.mocked(unloadGenerator).mockResolvedValue(undefined);
});
describe('WarmModelsPanel', () => {
it('shows the empty slot when no model is loaded', async () => {
render(<WarmModelsPanel />);
expect(await screen.findByText('No model loaded')).toBeInTheDocument();
expect(
screen.queryByRole('button', { name: 'Unload' }),
).not.toBeInTheDocument();
});
it('renders the resident slot with its state and config summary', async () => {
vi.mocked(listGenerators).mockResolvedValue([
makeGenerator({ num_gpus: 8, enable_torch_compile: true }),
]);
render(<WarmModelsPanel />);
// Model label: last path segment with dashes/underscores as spaces.
expect(
await screen.findByText('Wan2.1 T2V 1.3B Diffusers'),
).toBeInTheDocument();
expect(screen.getByText('ready')).toBeInTheDocument();
expect(screen.getByText('8 GPU · compile')).toBeInTheDocument();
expect(screen.getByRole('button', { name: 'Unload' })).toBeInTheDocument();
});
it('disables loading a new model while a load is in flight', async () => {
vi.mocked(listGenerators).mockResolvedValue([
makeGenerator({ state: 'loading' }),
]);
render(<WarmModelsPanel />);
expect(await screen.findByText('loading')).toBeInTheDocument();
expect(screen.getByRole('button', { name: 'Load model' })).toBeDisabled();
expect(
screen.queryByRole('button', { name: 'Unload' }),
).not.toBeInTheDocument();
});
it('shows the error on a failed slot and keeps retry enabled', async () => {
vi.mocked(listGenerators).mockResolvedValue([
makeGenerator({ state: 'failed', error: 'CUDA out of memory' }),
]);
render(<WarmModelsPanel />);
expect(await screen.findByText('failed')).toHaveAttribute(
'title',
'CUDA out of memory',
);
await waitFor(() =>
expect(screen.getByRole('button', { name: 'Load model' })).toBeEnabled(),
);
});
it('labels the load button as a swap when a different model is resident', async () => {
vi.mocked(listGenerators).mockResolvedValue([
makeGenerator({ model_id: 'FastVideo/FastHunyuan-diffusers' }),
]);
render(<WarmModelsPanel />);
expect(
await screen.findByRole('button', { name: 'Load (replaces current)' }),
).toBeInTheDocument();
});
it('loads the selected model with the persisted default options', async () => {
defaultOptionsStore.set({
options: { ...DEFAULT_OPTIONS, numGpus: 4, enableTorchCompile: true },
});
const user = userEvent.setup();
render(<WarmModelsPanel />);
const button = await screen.findByRole('button', { name: 'Load model' });
await waitFor(() => expect(button).toBeEnabled());
await user.click(button);
await waitFor(() =>
expect(preloadGenerator).toHaveBeenCalledWith({
model_id: 'Wan-AI/Wan2.1-T2V-1.3B-Diffusers',
workload_type: 't2v',
num_gpus: 4,
dit_cpu_offload: false,
text_encoder_cpu_offload: false,
vae_cpu_offload: false,
image_encoder_cpu_offload: false,
use_fsdp_inference: false,
enable_torch_compile: true,
vsa_sparsity: 0,
tp_size: -1,
sp_size: -1,
}),
);
// The panel refetches so the new "loading" slot appears promptly.
await waitFor(() => expect(listGenerators).toHaveBeenCalledTimes(2));
});
it('surfaces the backend detail when a load is rejected', async () => {
vi.mocked(preloadGenerator).mockRejectedValue(
new Error('a model load is already in progress'),
);
const user = userEvent.setup();
render(<WarmModelsPanel />);
const button = await screen.findByRole('button', { name: 'Load model' });
await waitFor(() => expect(button).toBeEnabled());
await user.click(button);
await waitFor(() =>
expect(toast.error).toHaveBeenCalledWith('Model was not loaded', {
description: 'a model load is already in progress',
}),
);
});
it('unloads the resident model with no payload', async () => {
vi.mocked(listGenerators).mockResolvedValue([makeGenerator()]);
const user = userEvent.setup();
render(<WarmModelsPanel />);
await user.click(await screen.findByRole('button', { name: 'Unload' }));
await waitFor(() => expect(unloadGenerator).toHaveBeenCalledTimes(1));
expect(vi.mocked(unloadGenerator).mock.calls[0]).toEqual([]);
// The slot refreshes after the unload succeeds.
await waitFor(() => expect(listGenerators).toHaveBeenCalledTimes(2));
});
});
@@ -0,0 +1,237 @@
'use client';
import * as React from 'react';
import { toast } from 'sonner';
import { Badge, type BadgeProps } from '@/components/ui/badge';
import { Button } from '@/components/ui/button';
import { NativeSelect } from '@/components/ui/native-select';
import { useStore } from '@/hooks/useStore';
import {
getModels,
listGenerators,
preloadGenerator,
unloadGenerator,
type GeneratorInfo,
type Model,
} from '@/lib/api';
import { getDefaultModelForWorkload } from '@/lib/defaultOptions';
import { defaultOptionsStore } from '@/stores/defaultOptions';
// Mirrors the backend's model_label(): readable label from an HF-style path.
function modelLabel(modelId: string): string {
return (modelId.split('/').pop() ?? modelId).replace(/[-_]/g, ' ');
}
// Compact summary of only the non-default engine bits, e.g. "8 GPU · compile".
function configSummary(gen: GeneratorInfo): string {
const parts: string[] = [];
if (gen.num_gpus !== 1) parts.push(`${gen.num_gpus} GPU`);
if (gen.sp_size !== -1) parts.push(`SP ${gen.sp_size}`);
if (gen.tp_size !== -1) parts.push(`TP ${gen.tp_size}`);
if (gen.dit_cpu_offload) parts.push('DiT offload');
if (gen.text_encoder_cpu_offload) parts.push('TE offload');
if (gen.vae_cpu_offload) parts.push('VAE offload');
if (gen.image_encoder_cpu_offload) parts.push('image enc offload');
if (gen.use_fsdp_inference) parts.push('FSDP');
if (gen.enable_torch_compile) parts.push('compile');
if (gen.vsa_sparsity > 0) parts.push(`VSA ${gen.vsa_sparsity.toFixed(2)}`);
return parts.join(' · ');
}
const STATE_VARIANTS: Record<GeneratorInfo['state'], BadgeProps['variant']> = {
ready: 'success',
loading: 'warning',
failed: 'destructive',
};
/**
* Utility strip for the engine's single model slot: shows the resident model
* (ready/loading/failed), loads the selected model using the persisted
* default job options (replacing whatever is resident), and unloads it.
*/
export default function WarmModelsPanel() {
const { options } = useStore(defaultOptionsStore);
const [slot, setSlot] = React.useState<GeneratorInfo | null>(null);
const [models, setModels] = React.useState<Model[]>([]);
const [modelId, setModelId] = React.useState('');
const [isBusy, setIsBusy] = React.useState(false);
const fetchSlot = React.useCallback(async () => {
try {
const list = await listGenerators();
setSlot(list[0] ?? null);
} catch (e) {
console.error('Failed to fetch generators:', e);
}
}, []);
React.useEffect(() => {
fetchSlot();
}, [fetchSlot]);
// Poll every 5s while a load is in flight; stop otherwise.
const isLoading = slot?.state === 'loading';
React.useEffect(() => {
if (!isLoading) return;
const interval = setInterval(fetchSlot, 5000);
return () => clearInterval(interval);
}, [isLoading, fetchSlot]);
// Same model catalogue (and default selection) as the create-job modal.
React.useEffect(() => {
getModels('t2v')
.then((list) => {
setModels(list);
const defaultId = getDefaultModelForWorkload(
defaultOptionsStore.get().options,
't2v',
);
setModelId(
list.some((m) => m.id === defaultId)
? defaultId
: (list[0]?.id ?? ''),
);
})
.catch((e) => console.error('Failed to load models:', e));
}, []);
async function handleLoad() {
if (!modelId || isBusy || isLoading) return;
setIsBusy(true);
try {
await preloadGenerator({
model_id: modelId,
workload_type: 't2v',
num_gpus: options.numGpus,
dit_cpu_offload: options.ditCpuOffload,
text_encoder_cpu_offload: options.textEncoderCpuOffload,
vae_cpu_offload: options.vaeCpuOffload,
image_encoder_cpu_offload: options.imageEncoderCpuOffload,
use_fsdp_inference: options.useFsdpInference,
enable_torch_compile: options.enableTorchCompile,
vsa_sparsity: options.vsaSparsity,
tp_size: options.tpSize,
sp_size: options.spSize,
});
await fetchSlot();
} catch (err) {
console.error('Failed to load model:', err);
toast.error('Model was not loaded', {
description:
err instanceof Error
? err.message
: 'Check the Studio API, then retry.',
});
} finally {
setIsBusy(false);
}
}
async function handleUnload() {
if (isBusy) return;
setIsBusy(true);
try {
await unloadGenerator();
await fetchSlot();
} catch (err) {
console.error('Failed to unload model:', err);
toast.error('Model was not unloaded', {
description:
err instanceof Error
? err.message
: 'Check the Studio API, then retry.',
});
} finally {
setIsBusy(false);
}
}
// Loading a model always replaces the resident one — say so on the button.
const replaces = slot !== null && !!modelId && slot.model_id !== modelId;
return (
<section
aria-label="Warm models"
className="mx-auto w-full max-w-[850px] px-10 pt-6"
>
<div className="flex flex-col gap-3 rounded-lg border border-border bg-background p-4">
<div className="flex flex-wrap items-center gap-2">
<h2 className="mr-auto text-sm font-semibold text-foreground">
Warm Model
</h2>
<label htmlFor="warm-model-select" className="sr-only">
Model to load
</label>
<NativeSelect
id="warm-model-select"
value={modelId}
onChange={(e) => setModelId(e.target.value)}
disabled={isBusy || models.length === 0}
className="h-9 w-auto max-w-64 rounded-lg"
>
<option value="" disabled>
{models.length === 0 ? 'Loading models…' : 'Select a model…'}
</option>
{models.map((model) => (
<option key={model.id} value={model.id}>
{model.label}
</option>
))}
</NativeSelect>
<Button
size="sm"
onClick={handleLoad}
disabled={isBusy || !modelId || isLoading}
>
{replaces ? 'Load (replaces current)' : 'Load model'}
</Button>
</div>
<p className="text-xs text-muted-foreground">
One model at a time stays resident in GPU memory so jobs skip the
load wait; loading a new one replaces it (uses your default job
options).
</p>
<div className="flex min-h-8 flex-wrap items-center gap-2">
{slot ? (
<>
<Badge
variant={STATE_VARIANTS[slot.state]}
className={
slot.state === 'loading' ? 'animate-pulse' : undefined
}
title={
slot.state === 'failed' ? (slot.error ?? undefined) : undefined
}
>
{slot.state}
</Badge>
<span className="text-sm font-medium text-foreground">
{modelLabel(slot.model_id)}
</span>
<span className="text-xs text-muted-foreground">
{configSummary(slot)}
</span>
{slot.state === 'ready' && (
<Button
size="sm"
variant="outline"
className="ml-auto"
onClick={handleUnload}
disabled={isBusy}
>
Unload
</Button>
)}
</>
) : (
<span className="text-sm text-muted-foreground">
No model loaded
</span>
)}
</div>
</div>
</section>
);
}
@@ -8,9 +8,9 @@ import { Button } from '@/components/ui/button';
import { ThemeToggle } from '@/components/ui/theme-toggle';
const TAB_TITLES: Record<string, string> = {
'/inference': 'Jobs',
'/finetuning': 'Jobs',
'/distillation': 'Jobs',
'/inference': 'Studio',
'/finetuning': 'Studio',
'/distillation': 'Studio',
'/datasets': 'Datasets',
'/gallery': 'Gallery',
'/gpus': 'GPUs',
@@ -108,7 +108,7 @@ export default function PrimarySidebar({
isJobsActive && TAB_ACTIVE,
)}
>
<span>Jobs</span>
<span>Studio</span>
<svg
viewBox="0 0 24 24"
fill="none"
+162
View File
@@ -202,6 +202,33 @@ export async function getModels(workloadType?: string): Promise<Model[]> {
return response.json();
}
/**
* A model's recommended sampling settings. Keys the backend has no value
* for are absent; the UI leaves those form fields untouched.
*/
export interface ModelPresets {
height?: number;
width?: number;
num_frames?: number;
fps?: number;
num_inference_steps?: number;
guidance_scale?: number;
guidance_rescale?: number;
negative_prompt?: string;
seed?: number;
}
export async function getModelPresets(modelId: string): Promise<ModelPresets> {
const baseApiUrl = getApiBaseUrl();
const response = await fetch(
`${baseApiUrl}/models/presets?model_id=${encodeURIComponent(modelId)}`,
);
if (!response.ok) {
throw new Error("Failed to fetch model presets");
}
return response.json();
}
export async function getGpus(): Promise<GpuSnapshot> {
const baseApiUrl = getApiBaseUrl();
const response = await fetch(`${baseApiUrl}/gpus`);
@@ -319,6 +346,109 @@ export async function downloadJobVideo(id: string): Promise<Blob> {
return response.blob();
}
// --- Generators (warm models) ---
/**
* Engine subset of CreateJobRequest that identifies one resident generator.
* Mirrors the backend's GeneratorRequest defaults.
*/
export interface GeneratorRequest {
model_id: string;
workload_type?: string;
num_gpus?: number;
dit_cpu_offload?: boolean;
text_encoder_cpu_offload?: boolean;
vae_cpu_offload?: boolean;
image_encoder_cpu_offload?: boolean;
use_fsdp_inference?: boolean;
enable_torch_compile?: boolean;
vsa_sparsity?: number;
tp_size?: number;
sp_size?: number;
}
export interface GeneratorInfo {
state: "ready" | "loading" | "failed";
model_id: string;
workload_type: string;
num_gpus: number;
dit_cpu_offload: boolean;
text_encoder_cpu_offload: boolean;
vae_cpu_offload: boolean;
image_encoder_cpu_offload: boolean;
use_fsdp_inference: boolean;
enable_torch_compile: boolean;
vsa_sparsity: number;
tp_size: number;
sp_size: number;
error: string | null;
started_at?: number;
}
/** The single resident slot: [] when nothing is loaded, else one entry. */
export async function listGenerators(): Promise<GeneratorInfo[]> {
const baseApiUrl = getApiBaseUrl();
const response = await fetch(`${baseApiUrl}/generators`);
if (!response.ok) {
throw new Error("Failed to fetch generators");
}
return response.json();
}
export async function preloadGenerator(
req: GeneratorRequest,
): Promise<GeneratorInfo> {
const baseApiUrl = getApiBaseUrl();
const response = await fetch(`${baseApiUrl}/generators/preload`, {
method: "POST",
headers: {
"Content-Type": "application/json",
},
body: JSON.stringify(req),
});
if (!response.ok) {
const error = await response
.json()
.catch(() => ({ detail: "Failed to preload model" }));
throw new Error(error.detail || "Failed to preload model");
}
return response.json();
}
/** Unload the single resident generator (no body — there is only one slot). */
export async function unloadGenerator(): Promise<void> {
const baseApiUrl = getApiBaseUrl();
const response = await fetch(`${baseApiUrl}/generators/unload`, {
method: "POST",
});
if (!response.ok) {
const error = await response
.json()
.catch(() => ({ detail: "Failed to unload model" }));
throw new Error(error.detail || "Failed to unload model");
}
}
// --- Engine logs ---
export interface EngineLogs {
lines: string[];
total: number;
}
/**
* Incremental tail of the engine's stdout/stderr. Poll with
* `after=<total from the previous response>`.
*/
export async function getEngineLogs(after: number = 0): Promise<EngineLogs> {
const baseApiUrl = getApiBaseUrl();
const response = await fetch(`${baseApiUrl}/engine/logs?after=${after}`);
if (!response.ok) {
throw new Error("Failed to fetch engine logs");
}
return response.json();
}
// --- Datasets ---
export interface Dataset {
@@ -418,3 +548,35 @@ export function getDatasetMediaUrl(
const baseApiUrl = getApiBaseUrl();
return `${baseApiUrl}/datasets/${datasetId}/media/${encodeURIComponent(fileName)}`;
}
// --- Cluster ---
/** Per-GPU telemetry within a cluster node (same shape as GpuInfo). */
export type ClusterGpu = GpuInfo;
export interface ClusterNode {
hostname: string;
ip: string | null;
is_this_host: boolean;
cpus: number | null;
ray_gpus: number | null;
available: boolean;
error: string | null;
gpus: ClusterGpu[];
}
export interface ClusterSnapshot {
mode: "ray" | "local";
nodes: ClusterNode[];
resources: { gpus_total: number; gpus_available: number } | null;
error: string | null;
}
export async function getClusterStatus(): Promise<ClusterSnapshot> {
const baseApiUrl = getApiBaseUrl();
const response = await fetch(`${baseApiUrl}/cluster`);
if (!response.ok) {
throw new Error("Failed to fetch cluster status");
}
return response.json();
}
@@ -0,0 +1,273 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the single-slot resident generator: preload, list, unload.
Exactly one VideoGenerator lives in memory. One load at a time; loading a new
config always releases the old instance; unload deletes it. The generator is
faked at the ``_create_generator_into_slot``/VideoGenerator boundary.
"""
from __future__ import annotations
import threading
import time
from types import SimpleNamespace
from typing import Any
import pytest
from fastvideo_studio.job_runner import JobRunner, JobStatus
class _FakeGenerator:
def __init__(self) -> None:
self.shutdown_calls = 0
def shutdown(self) -> None:
self.shutdown_calls += 1
@pytest.fixture()
def runner():
r = JobRunner.__new__(JobRunner) # skip __init__: no DB/Manager needed
r._jobs = {}
r._jobs_lock = threading.Lock()
r._generator = None
r._generator_config = None
r._generator_state = "empty"
r._generator_error = None
r._generator_lock = threading.Lock()
r._load_lock = threading.Lock()
r._worker_log_queue = None # no Manager in unit tests
import queue as _queue
r._loader_queue = _queue.Queue()
threading.Thread(target=r._loader_loop, daemon=True).start()
return r
def _install_fake_loader(runner, monkeypatch, made: list | None = None,
gate: threading.Event | None = None,
fail: str | None = None):
"""Replace the VideoGenerator load inside _create_generator_into_slot."""
import fastvideo_studio.job_runner as jr
class _FakeVG:
@staticmethod
def from_pretrained(model_path, **kwargs):
if gate is not None:
assert gate.wait(5), "test gate never opened"
if fail is not None:
raise RuntimeError(fail)
gen = _FakeGenerator()
if made is not None:
made.append((model_path, gen))
return gen
monkeypatch.setitem(__import__("sys").modules, "fastvideo",
SimpleNamespace(VideoGenerator=_FakeVG))
return jr
def _wait_state(runner, state, timeout=5.0):
deadline = time.time() + timeout
while time.time() < deadline:
with runner._generator_lock:
if runner._generator_state == state:
return runner._slot_entry()
time.sleep(0.02)
raise AssertionError(f"never reached state {state}: {runner._generator_state}")
def test_preload_then_ready_then_idempotent(runner, monkeypatch):
made = []
_install_fake_loader(runner, monkeypatch, made=made)
entry = runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
assert entry["state"] in ("loading", "ready")
_wait_state(runner, "ready")
again = runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
assert again["state"] == "ready"
assert len(made) == 1 # same config never reloads
def test_only_one_load_at_a_time(runner, monkeypatch):
gate = threading.Event()
_install_fake_loader(runner, monkeypatch, gate=gate)
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
with pytest.raises(RuntimeError, match="already in progress"):
runner.preload_generator(model_id="org/b", workload_type="t2v", num_gpus=8)
gate.set()
_wait_state(runner, "ready")
def test_new_config_always_releases_old_instance(runner, monkeypatch):
made = []
_install_fake_loader(runner, monkeypatch, made=made)
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
_wait_state(runner, "ready")
first = made[0][1]
runner.preload_generator(model_id="org/b", workload_type="t2v", num_gpus=8)
_wait_state(runner, "ready")
assert first.shutdown_calls == 1 # old instance released, not stacked
assert [m[0] for m in made] == ["org/a", "org/b"]
assert len(runner.list_generators()) == 1
assert runner.list_generators()[0]["model_id"] == "org/b"
def test_failed_load_reports_error_and_allows_retry(runner, monkeypatch):
_install_fake_loader(runner, monkeypatch, fail="no CUDA on this box")
runner.preload_generator(model_id="org/broken", workload_type="t2v", num_gpus=1)
failed = _wait_state(runner, "failed")
assert "no CUDA" in failed["error"]
# retry after failure must be accepted (this was the reported bug)
made = []
_install_fake_loader(runner, monkeypatch, made=made)
runner.preload_generator(model_id="org/broken", workload_type="t2v", num_gpus=1)
_wait_state(runner, "ready")
assert len(made) == 1
def test_unload_deletes_instance(runner, monkeypatch):
made = []
_install_fake_loader(runner, monkeypatch, made=made)
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
_wait_state(runner, "ready")
assert runner.unload_generator() is True
assert made[0][1].shutdown_calls == 1
assert runner.list_generators() == []
assert runner._generator is None
assert runner.unload_generator() is False # nothing resident
def test_unload_refuses_while_inference_job_runs(runner, monkeypatch):
made = []
_install_fake_loader(runner, monkeypatch, made=made)
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
_wait_state(runner, "ready")
runner._jobs["j1"] = SimpleNamespace(id="j1", status=JobStatus.RUNNING, job_type="inference")
with pytest.raises(RuntimeError, match="j1"):
runner.unload_generator()
assert made[0][1].shutdown_calls == 0
runner._jobs["j1"].job_type = "finetune" # training doesn't block
assert runner.unload_generator() is True
def test_job_waits_for_matching_preload(runner, monkeypatch):
gate = threading.Event()
made = []
_install_fake_loader(runner, monkeypatch, made=made, gate=gate)
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
got = []
t = threading.Thread(
target=lambda: got.append(
runner._get_or_create_generator("org/a", "t2v", 8)),
daemon=True)
t.start()
time.sleep(0.2)
assert not got # job blocked on the in-flight load
gate.set()
t.join(5)
assert got and got[0] is made[0][1]
assert len(made) == 1 # the job reused the preloaded instance
def test_job_with_different_config_replaces_slot(runner, monkeypatch):
made = []
_install_fake_loader(runner, monkeypatch, made=made)
runner.preload_generator(model_id="org/a", workload_type="t2v", num_gpus=8)
_wait_state(runner, "ready")
gen = runner._get_or_create_generator("org/b", "t2v", 4)
assert gen is made[1][1]
assert made[0][1].shutdown_calls == 1 # old released before new load
assert runner.list_generators()[0]["model_id"] == "org/b"
def test_job_during_replace_never_sees_half_replaced_slot(runner, monkeypatch):
"""Regression: user restarted a stale 1-gpu job while an 8-gpu generator
was resident; the replace nulled the generator with state still 'ready',
and the next (correct) job grabbed None -> AttributeError on
generate_video. All transitions now serialize under the load lock."""
made = []
gate = threading.Event()
_install_fake_loader(runner, monkeypatch, made=made)
runner.preload_generator(model_id="org/h3", workload_type="t2v", num_gpus=8)
_wait_state(runner, "ready")
_install_fake_loader(runner, monkeypatch, made=made, gate=gate)
results: dict[str, Any] = {}
def job_a(): # stale job: mismatching config triggers a slow replace
results["a"] = runner._get_or_create_generator("org/h3", "t2v", 1)
def job_b(): # correct job arriving mid-replace
time.sleep(0.3)
results["b"] = runner._get_or_create_generator("org/h3", "t2v", 8)
ta = threading.Thread(target=job_a, daemon=True)
tb = threading.Thread(target=job_b, daemon=True)
ta.start(); tb.start()
time.sleep(0.6)
gate.set()
ta.join(10); tb.join(10)
assert results["a"] is not None and hasattr(results["a"], "shutdown")
assert results["b"] is not None and hasattr(results["b"], "shutdown")
# b arrived second, so the slot ends at b's 8-gpu config
assert runner.list_generators()[0]["num_gpus"] == 8
def test_config_dict_matches_request_defaults(runner):
"""GeneratorRequest defaults and CreateJobRequest defaults must resolve to
the same slot config — else preloading never matches the job."""
from fastvideo_studio.models import CreateJobRequest, GeneratorRequest
job = CreateJobRequest(model_id="m", prompt="p").model_dump()
pre = GeneratorRequest(model_id="m").model_dump()
assert runner._generator_config_dict(**pre) == runner._generator_config_dict(
**{k: job[k] for k in pre})
def test_engine_log_buffer_incremental_tail():
from fastvideo_studio.server import _EngineLogBuffer
buf = _EngineLogBuffer(maxlen=3)
buf.write("a\nb\n")
buf.write("c") # partial line: not visible yet
lines, total = buf.get_lines(0)
assert lines == ["a", "b"] and total == 2
buf.write("!\nd\ne\n") # completes "c!", then overflows the ring
lines, total = buf.get_lines(total)
assert lines == ["c!", "d", "e"] and total == 5
# reader far behind: dropped lines are skipped, no crash
lines, _ = buf.get_lines(0)
assert lines == ["c!", "d", "e"]
def test_engine_feed_drives_job_progress(runner):
from fastvideo_studio.job_runner import Job
job = Job(id="j-prog", model_id="m", prompt="p")
runner._active_inference_job = job
# ray wraps relayed lines in ANSI color codes — the bridge must strip them
runner.feed_engine_line("\x1b[36m(RayWorkerWrapper pid=123, ip=10.0.0.2)\x1b[0m denoising: 40%|████ | 20/50 [00:30<00:45, 1.5s/it]")
assert job._log_buf.progress == 40.0
assert job._log_buf.progress_msg == "20/50 steps"
# driver-side lines (no ray prefix) are NOT double-fed
before = job._log_buf.get_lines()[1]
runner.feed_engine_line("INFO 08-07 [video_generator.py] driver line")
assert job._log_buf.get_lines()[1] == before
# no active job: no-op
runner._active_inference_job = None
runner.feed_engine_line("(RayWorkerWrapper pid=123) 90%|████| 45/50")
assert job._log_buf.progress == 40.0
@@ -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()
+15 -6
View File
@@ -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