feat: add central slot scheduler for model requests

parent 3580ff1d
......@@ -307,9 +307,25 @@ async def api_status(username: str = Depends(require_auth)):
# Request stats from queue manager
req_total = 0
req_active = 0
req_waiting = 0
req_metrics = {
"max_parallel_requests": 0,
"queue_max_size": 0,
"active_by_model": {},
"waiting_by_model": {},
}
try:
from codai.queue.manager import queue_manager
req_active = 1 if queue_manager._processing else 0
metrics = queue_manager.get_metrics()
req_active = int(metrics.get("active", 0))
req_waiting = int(metrics.get("waiting", 0))
req_total = req_active + req_waiting
req_metrics = {
"max_parallel_requests": metrics.get("max_parallel_requests", 0),
"queue_max_size": metrics.get("queue_max_size", 0),
"active_by_model": metrics.get("active_by_model", {}),
"waiting_by_model": metrics.get("waiting_by_model", {}),
}
except Exception:
pass
......@@ -364,7 +380,15 @@ async def api_status(username: str = Depends(require_auth)):
"enabled_models": enabled_models,
"vram": vram,
"cuda": is_cuda,
"requests": {"total": req_total, "active": req_active},
"requests": {
"total": req_total,
"active": req_active,
"waiting": req_waiting,
"max_parallel_requests": req_metrics["max_parallel_requests"],
"queue_max_size": req_metrics["queue_max_size"],
"active_by_model": req_metrics["active_by_model"],
"waiting_by_model": req_metrics["waiting_by_model"],
},
"recent_activity": recent_activity,
"whisper_server": whisper_status,
}
......@@ -1423,6 +1447,7 @@ async def api_get_settings(username: str = Depends(require_admin)):
"https_key_path": c.server.https_key_path,
"https_cert_path": c.server.https_cert_path,
"queue_max_size": c.server.queue_max_size,
"max_parallel_requests": c.server.max_parallel_requests,
},
"backend": {
"type": c.backend.type,
......@@ -1478,6 +1503,10 @@ async def api_save_settings(request: Request, username: str = Depends(require_ad
c.server.queue_max_size = max(1, int(srv["queue_max_size"]))
from codai.queue.manager import queue_manager
queue_manager.max_size = c.server.queue_max_size
if "max_parallel_requests" in srv:
c.server.max_parallel_requests = int(srv["max_parallel_requests"])
from codai.queue.manager import queue_manager
queue_manager.max_parallel_requests = c.server.max_parallel_requests
if "backend" in data:
bk = data["backend"]
......
......@@ -92,6 +92,8 @@ from codai.api.tts import router as tts_router
from codai.api.text import router as text_router
from codai.api.video import router as video_router
from codai.api.audio_gen import router as audio_gen_router
from codai.api.audio_stems import router as audio_stems_router
from codai.api.audio_clean import router as audio_clean_router
from codai.api.embeddings import router as embeddings_router
from codai.api.pipelines import router as pipelines_router
from codai.api.custom_pipelines import router as custom_pipelines_router
......
......@@ -274,6 +274,42 @@ async def _run_step(step: Dict, context: Dict, http_request) -> Dict:
return _extract_output(step_type, result)
def _infer_step_model_key(step: Dict) -> Optional[str]:
step_type = step.get('type')
params = step.get('params', {})
if step_type == 'stt':
model = params.get('model') or params.get('audio_model')
return f"audio:{model}" if model else None
if step_type == 'text_gen':
return params.get('model')
if step_type in {'image_gen', 'image_edit', 'image_upscale', 'image_depth', 'image_segment'}:
model = params.get('model')
return f"image:{model}" if model else None
if step_type in {'embed', 'embedding'}:
model = params.get('model')
return f"embedding:{model}" if model else None
if step_type in {'video_gen', 'video'}:
model = params.get('model')
return f"video:{model}" if model else None
return None
async def _run_scheduled_step(step: Dict, context: Dict, http_request) -> Dict:
from codai.queue.manager import queue_manager
model_key = _infer_step_model_key(step)
if not model_key:
return await _run_step(step, context, http_request)
request_id = f"pipeline-step-{uuid.uuid4().hex[:8]}"
lease = await queue_manager.acquire(request_id, model_key)
try:
return await _run_step(step, context, http_request)
finally:
await queue_manager.release(lease)
async def _execute_pipeline(pipeline_def: Dict, pipeline_input: str, http_request) -> Dict:
"""Execute all steps of a pipeline definition."""
context = {'input': pipeline_input}
......@@ -281,7 +317,7 @@ async def _execute_pipeline(pipeline_def: Dict, pipeline_input: str, http_reques
for i, step in enumerate(pipeline_def.get('steps', [])):
try:
out = await _run_step(step, context, http_request)
out = await _run_scheduled_step(step, context, http_request)
context[f'step{i}'] = out
steps_output.append({'step': i, 'type': step['type'],
'label': step.get('label', step['type']), **out})
......@@ -446,7 +482,7 @@ async def run_audio_understanding(request: AudioUnderstandRequest, http_request:
'response_format': 'json',
},
}
stt_out = await _run_step(stt_step, {'input': request.input or ''}, http_request)
stt_out = await _run_scheduled_step(stt_step, {'input': request.input or ''}, http_request)
transcript = stt_out.get('text') or stt_out.get('output') or ''
steps.append({'step': 0, 'type': 'stt', 'label': 'Transcribe audio', **stt_out})
......@@ -460,7 +496,7 @@ async def run_audio_understanding(request: AudioUnderstandRequest, http_request:
'prompt': f"{request.input or 'Summarize this audio transcript clearly.'}\n\nTranscript:\n{{{{step0.output}}}}",
},
}
text_out = await _run_step(text_step, {'input': request.input or '', 'step0': {'output': transcript, 'text': transcript}}, http_request)
text_out = await _run_scheduled_step(text_step, {'input': request.input or '', 'step0': {'output': transcript, 'text': transcript}}, http_request)
summary = text_out.get('output')
steps.append({'step': 1, 'type': 'text_gen', 'label': 'Reason over transcript', **text_out})
......@@ -486,7 +522,7 @@ async def run_full_music_dub(request: AudioMusicDubRequest, http_request: Reques
'response_format': 'json',
},
}
stt_out = await _run_step(stt_step, {'input': request.notes or ''}, http_request)
stt_out = await _run_scheduled_step(stt_step, {'input': request.notes or ''}, http_request)
transcript = stt_out.get('text') or stt_out.get('output') or ''
translated = transcript if not request.target_lang else f"[{request.target_lang}] {transcript}"
steps = [
......
......@@ -31,6 +31,7 @@ class ServerConfig:
https_key_path: Optional[str] = None
https_cert_path: Optional[str] = None
queue_max_size: int = 6
max_parallel_requests: int = 2
@dataclass
......@@ -302,6 +303,7 @@ class ConfigManager:
"https_key_path": self.config.server.https_key_path,
"https_cert_path": self.config.server.https_cert_path,
"queue_max_size": self.config.server.queue_max_size,
"max_parallel_requests": self.config.server.max_parallel_requests,
},
"backend": {
"type": self.config.backend.type,
......
......@@ -646,9 +646,10 @@ def main():
# Apply queue max size from config
# Apply queue scheduler settings from config
from codai.queue.manager import queue_manager
queue_manager.max_size = config.server.queue_max_size
queue_manager.max_parallel_requests = config.server.max_parallel_requests
# Start the server
import uvicorn
......
......@@ -427,6 +427,14 @@ class MultiModelManager:
def image_model(self) -> Optional[str]:
"""Return the first image model or None."""
return self.image_models[0] if self.image_models else None
def get_loaded_model_keys(self) -> set:
"""Return the set of currently loaded model keys."""
return set(self.models.keys())
def has_loaded_model(self, model_key: str) -> bool:
"""Return True when the given model key is currently loaded."""
return model_key in self.models
def cleanup(self):
"""Cleanup all models and resources."""
......
......@@ -14,75 +14,243 @@
# You should have received a copy of the GNU General Public License
# along with this program. If not, see <https://www.gnu.org/licenses/>.
"""Queue manager module - manages request queues for model loading notifications."""
"""Central scheduler for model-backed request admission and queue reporting."""
from typing import Dict, Optional
from collections import deque
from dataclasses import dataclass, field
from typing import Deque, Dict, Optional, Set
import asyncio
import time
@dataclass
class SchedulerLease:
request_id: str
model_key: str
started_at: float = field(default_factory=time.time)
waited: bool = False
wait_time_seconds: float = 0.0
@dataclass
class WaitingRequest:
request_id: str
model_key: str
enqueued_at: float
sequence: int
event: asyncio.Event = field(default_factory=asyncio.Event)
bypassed_by: int = 0
class QueueManager:
"""
Manages request queue for model loading notifications.
When clients are waiting for a model to load, sends them progress updates.
"""
"""Central in-process request scheduler."""
def __init__(self):
self.waiting_requests: Dict[str, float] = {} # request_id -> start_time
self.lock = asyncio.Lock()
self.max_size: int = 6
self.max_parallel_requests: int = 2
self.waiting: Deque[WaitingRequest] = deque()
self.waiting_by_id: Dict[str, WaitingRequest] = {}
self.active_leases: Dict[str, SchedulerLease] = {}
self.active_by_model: Dict[str, int] = {}
self.loaded_models: Set[str] = set()
self.sequence: int = 0
self.fairness_bypass_limit: int = 2
self.current_request_id: Optional[str] = None
self.model_loading: bool = False
self.model_name: Optional[str] = None
self.lock = asyncio.Lock()
self.max_size: int = 6
self._processing: bool = False
self._ready_request_ids: Set[str] = set()
def set_loaded_models(self, model_keys: Set[str]) -> None:
self.loaded_models = set(model_keys)
def mark_model_loaded(self, model_key: str) -> None:
self.loaded_models.add(model_key)
def mark_model_unloaded(self, model_key: str) -> None:
self.loaded_models.discard(model_key)
def reset_for_tests(self) -> None:
self.waiting.clear()
self.waiting_by_id.clear()
self.active_leases.clear()
self.active_by_model.clear()
self.loaded_models.clear()
self.sequence = 0
self.current_request_id = None
self.model_loading = False
self.model_name = None
self._processing = False
self._ready_request_ids.clear()
async def is_full(self) -> bool:
"""Return True if the queue has reached max_size."""
async with self.lock:
return len(self.waiting_requests) >= self.max_size
async def add_waiting(self, request_id: str) -> None:
"""Add a request to the waiting queue."""
return len(self.waiting) >= self.max_size
async def acquire(self, request_id: str, model_key: str) -> SchedulerLease:
waiter = None
async with self.lock:
if self._can_start_now(model_key):
return self._grant_lease(request_id, model_key)
waiter = self._enqueue_waiter(request_id, model_key)
await waiter.event.wait()
async with self.lock:
self._ready_request_ids.discard(request_id)
lease = self._grant_lease(request_id, model_key)
lease.waited = True
lease.wait_time_seconds = max(0.0, time.time() - waiter.enqueued_at)
return lease
async def release(self, lease: SchedulerLease) -> None:
async with self.lock:
self.active_leases.pop(lease.request_id, None)
current = self.active_by_model.get(lease.model_key, 0)
if current <= 1:
self.active_by_model.pop(lease.model_key, None)
else:
self.active_by_model[lease.model_key] = current - 1
if self.current_request_id == lease.request_id:
self.current_request_id = None
self._processing = bool(self.active_leases)
self._wake_waiters_locked()
async def add_waiting(self, request_id: str, model_key: str = "") -> None:
async with self.lock:
self.waiting_requests[request_id] = time.time()
if request_id in self.waiting_by_id:
return
self._enqueue_waiter(request_id, model_key or request_id)
async def remove_waiting(self, request_id: str) -> None:
"""Remove a request from the waiting queue."""
async with self.lock:
self.waiting_requests.pop(request_id, None)
waiter = self.waiting_by_id.pop(request_id, None)
if waiter and waiter in self.waiting:
self.waiting.remove(waiter)
self._ready_request_ids.discard(request_id)
async def start_processing(self, request_id: str, model_name: str = None) -> None:
"""Mark a request as now processing (model loaded)."""
async with self.lock:
self.waiting_requests.pop(request_id, None)
waiter = self.waiting_by_id.pop(request_id, None)
if waiter and waiter in self.waiting:
self.waiting.remove(waiter)
self.current_request_id = request_id
self.model_name = model_name
self._processing = True
async def finish_processing(self) -> None:
"""Mark current request as finished."""
async with self.lock:
self.current_request_id = None
self._processing = bool(self.active_leases)
async def is_waiting(self, request_id: str) -> bool:
"""Check if a request is in the waiting queue."""
async with self.lock:
return request_id in self.waiting_requests
return request_id in self.waiting_by_id
async def get_wait_time(self, request_id: str) -> float:
"""Get how long a request has been waiting in seconds."""
async with self.lock:
if request_id in self.waiting_requests:
return time.time() - self.waiting_requests[request_id]
waiter = self.waiting_by_id.get(request_id)
if waiter:
return time.time() - waiter.enqueued_at
return 0.0
async def get_queue_position(self, request_id: str) -> int:
"""Get the position of a request in the queue (1-based)."""
async with self.lock:
keys = list(self.waiting_requests.keys())
try:
return keys.index(request_id) + 1
except ValueError:
return 0
for index, waiter in enumerate(self.waiting, start=1):
if waiter.request_id == request_id:
return index
return 0
def get_metrics(self) -> Dict[str, object]:
return {
"active": len(self.active_leases),
"waiting": len(self.waiting),
"max_parallel_requests": self.max_parallel_requests,
"queue_max_size": self.max_size,
"active_by_model": dict(self.active_by_model),
"waiting_by_model": self._waiting_counts_locked(),
"loaded_models": sorted(self.loaded_models),
}
def _enqueue_waiter(self, request_id: str, model_key: str) -> WaitingRequest:
self.sequence += 1
waiter = WaitingRequest(
request_id=request_id,
model_key=model_key,
enqueued_at=time.time(),
sequence=self.sequence,
)
self.waiting.append(waiter)
self.waiting_by_id[request_id] = waiter
return waiter
def _grant_lease(self, request_id: str, model_key: str) -> SchedulerLease:
lease = SchedulerLease(request_id=request_id, model_key=model_key)
self.active_leases[request_id] = lease
self.active_by_model[model_key] = self.active_by_model.get(model_key, 0) + 1
self.current_request_id = request_id
self.model_name = model_key
self._processing = True
return lease
def _can_start_now(self, model_key: str) -> bool:
if self.max_parallel_requests > 0:
in_flight = len(self.active_leases) + len(self._ready_request_ids)
if in_flight >= self.max_parallel_requests:
return False
if self.active_by_model.get(model_key, 0) > 0:
return False
return self._is_loaded_model(model_key) or self._can_schedule_model_switch(model_key)
def _waiter_can_start_locked(self, waiter: WaitingRequest) -> bool:
return self._can_start_now(waiter.model_key)
def _is_loaded_model(self, model_key: str) -> bool:
return model_key in self.loaded_models
def _can_schedule_model_switch(self, model_key: str) -> bool:
if not self.loaded_models:
return True
if self._waiting_counts_locked().get(model_key, 0) > 0 and self._is_loaded_model(model_key):
return True
for loaded_key in self.loaded_models:
if self.active_by_model.get(loaded_key, 0) > 0:
return False
if self._waiting_counts_locked().get(loaded_key, 0) > 0:
return False
return True
def _wake_waiters_locked(self) -> None:
while True:
candidate = self._pick_next_waiter_locked()
if candidate is None:
return
self.waiting.remove(candidate)
self.waiting_by_id.pop(candidate.request_id, None)
self._ready_request_ids.add(candidate.request_id)
candidate.event.set()
if self.max_parallel_requests > 0 and len(self.active_leases) + len(self._ready_request_ids) >= self.max_parallel_requests:
return
def _pick_next_waiter_locked(self) -> Optional[WaitingRequest]:
for waiter in self.waiting:
if self._waiter_can_start_locked(waiter):
older_blocked = [
other for other in self.waiting
if other.sequence < waiter.sequence and not self._waiter_can_start_locked(other)
]
if any(other.bypassed_by >= self.fairness_bypass_limit for other in older_blocked):
continue
for other in older_blocked:
other.bypassed_by += 1
return waiter
return None
def _waiting_counts_locked(self) -> Dict[str, int]:
counts: Dict[str, int] = {}
for waiter in self.waiting:
counts[waiter.model_key] = counts.get(waiter.model_key, 0) + 1
return counts
# Global queue manager instance
queue_manager = QueueManager()
\ No newline at end of file
queue_manager = QueueManager()
......@@ -104,6 +104,35 @@ def test_audio_understanding_returns_transcript_only_without_text_model(monkeypa
assert len(body["steps"]) == 1
def test_audio_understanding_pipeline_steps_release_scheduler_slots(monkeypatch, studio_client):
from codai.api import custom_pipelines
from codai.queue.manager import queue_manager
observed = []
async def fake_run_step(step, context, http_request):
observed.append(queue_manager.get_metrics()["active"])
return {"output": step["type"], "text": step["type"]}
monkeypatch.setattr(custom_pipelines, "_run_step", fake_run_step)
queue_manager.reset_for_tests()
queue_manager.set_loaded_models({"audio:whisper-small", "qwen-text"})
response = studio_client.post(
"/v1/pipelines/audio-understand",
json={
"input": "Summarize",
"audio_model": "whisper-small",
"text_model": "qwen-text",
"audio": "ZmFrZQ==",
},
)
assert response.status_code == 200
assert observed == [1, 1]
assert queue_manager.get_metrics()["active"] == 0
def test_audio_understanding_requires_audio_source(studio_client):
response = studio_client.post(
"/v1/pipelines/audio-understand",
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment