"""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