zero-gpu-training-space / job_queue.py
flozi00's picture
Upload folder using huggingface_hub
6010c50 verified
Raw History Blame Contribute Delete
6.69 kB
"""Timeout-resilient, method-priority job queue for the training Space.
Why a queue
-----------
Training calls are gated by ``@spaces.GPU``: the GPU is only attached while a
call runs. If several training requests arrive at once — or a client times out
mid-training — the requests must not be lost. They are held in this queue and
processed later, in method-priority order, as the GPU becomes available.
Ordering (the requested behaviour)
---------------------------------
* **Method priority:** all **SFT** jobs run first, then all **DPO**, then all
**KTO**. If 5 SFT and 6 DPO requests arrive, every SFT job is drained
before any DPO job starts.
* **FIFO within a method:** earlier requests of the same method run first.
* **Timeout-resilient:** the queue lives in RAM (and, optionally, mirrored to
the bucket), so a *client-side* timeout does not drop a request — the worker
keeps it queued and trains it when the GPU is next available.
Persistence
----------
The raw training *data* is not persisted to the bucket by default (per the
Space's "no data persistence" design). The queue is held in RAM, which
survives a client timeout (the Space process keeps running while a client
waits/abandons). Set ``QUEUE_PERSIST=true`` to also mirror queued jobs to the
bucket so they survive a Space *restart* too.
"""
from __future__ import annotations
import itertools
import time
from dataclasses import dataclass, field
from typing import Optional
# Method priority: lower runs first. SFT first, then DPO, then KTO. Unknown
# methods sort last.
METHOD_PRIORITY = {"sft": 0, "dpo": 1, "kto": 2}
# A monotonic submission counter, used for stable FIFO ordering within a method.
_seq = itertools.count(1)
# A module-level lock: the RAM queue is shared between the (fast) submit
# callback and the (slow, GPU) worker callback, which Gradio runs in separate
# threads.
import threading
_lock = threading.Lock()
@dataclass
class Job:
"""A single queued training request: a method + the raw uploaded rows."""
method: str
rows: list
submitted_at: float = field(default_factory=time.time)
seq: int = 0
job_id: str = ""
def __post_init__(self) -> None:
self.seq = next(_seq)
self.job_id = f"j{self.seq:06d}"
self.method = self.method.lower()
def as_dict(self) -> dict:
return {
"job_id": self.job_id,
"method": self.method,
"submitted_at": self.submitted_at,
"n_examples": len(self.rows),
}
class JobQueue:
"""An in-RAM, method-priority, FIFO-within-method job queue.
Optional bucket persistence (``QUEUE_PERSIST=true`` + a ``BUCKET_ID``)
mirrors the queue to the bucket so it survives a Space restart.
"""
def __init__(self, bucket_id: str = "", persist: bool = False):
self._bucket_id = bucket_id
self._persist = bool(persist and bucket_id)
self._jobs: list[Job] = []
# -- submission ------------------------------------------------------- #
def submit(self, method: str, rows: list) -> Job:
"""Enqueue a request and return its handle. Fast — no GPU, no
training; this is what makes a client timeout harmless (the request is
already cached before the call returns)."""
job = Job(method=method.lower(), rows=list(rows))
with _lock:
self._jobs.append(job)
if self._persist:
self._mirror(job, add=True)
return job
# -- draining --------------------------------------------------------- #
def pop(self) -> Optional[Job]:
"""Pop and return the highest-priority pending job (method priority,
then FIFO), or ``None`` if the queue is empty."""
with _lock:
if not self._jobs:
return None
job = self._sorted()[0]
self._jobs.remove(job)
if self._persist:
self._mirror(job, add=False)
return job
def drain(self, max_jobs: int = 0):
"""Yield pending jobs in execution order (method priority, then FIFO),
up to ``max_jobs`` (0 = no cap). Each yielded job is removed from the
queue as it is popped, so the queue is drained as the worker runs."""
n = 0
while True:
job = self.pop()
if job is None:
break
yield job
n += 1
if max_jobs and n >= max_jobs:
return
# -- inspection ------------------------------------------------------- #
def pending_count(self) -> int:
with _lock:
return len(self._jobs)
def status(self) -> dict:
with _lock:
order = self._sorted()
counts: dict = {}
for j in self._jobs:
counts[j.method] = counts.get(j.method, 0) + 1
return {
"pending": len(self._jobs),
"by_method": counts,
"persisted": self._persist,
"next": (f"{order[0].method} · {order[0].job_id}" if order else None),
"order": [f"{j.method}·{j.job_id}" for j in order[:12]],
}
# -- helpers ---------------------------------------------------------- #
def _sorted(self) -> list[Job]:
"""Order the current queue by method priority, then FIFO (seq)."""
return sorted(
self._jobs,
key=lambda j: (METHOD_PRIORITY.get(j.method, 99), j.seq),
)
def _mirror(self, job: Job, add: bool) -> None:
"""Best-effort bucket mirror (opt-in). Never raises into the caller."""
try:
from store import delete_job, save_job
if add:
save_job(self._bucket_id, job)
else:
delete_job(self._bucket_id, job.job_id)
except Exception: # pragma: no cover
pass
# A module-level queue instance shared by the Space. It is (re)configured
# from the env on first use via ``init_queue``.
_QUEUE: Optional[JobQueue] = None
def init_queue(bucket_id: str = "", persist: bool = False) -> JobQueue:
"""(Re)create the module-level queue. Called once at Space startup."""
global _QUEUE
_QUEUE = JobQueue(bucket_id=bucket_id, persist=persist)
return _QUEUE
def get_queue() -> JobQueue:
"""Return the module-level queue, creating it lazily if needed."""
global _QUEUE
if _QUEUE is None:
import os
bucket_id = (os.environ.get("BUCKET_ID") or "").strip()
persist = (os.environ.get("QUEUE_PERSIST") or "").lower() in ("1", "true", "yes")
_QUEUE = JobQueue(bucket_id=bucket_id, persist=persist)
return _QUEUE