diff --git a/backend/app/ollama_client.py b/backend/app/ollama_client.py index b91a2d7..e93233e 100644 --- a/backend/app/ollama_client.py +++ b/backend/app/ollama_client.py @@ -118,19 +118,27 @@ class OllamaClient: raise OllamaError(str(chunk["error"])) yield chunk - async def warm(self, model: str, keep_alive: str = "30m") -> dict: + async def warm( + self, model: str, keep_alive: str = "30m", options: dict | None = None + ) -> dict: """Précharge un modèle en VRAM sans générer (/api/generate sans prompt). - Le paramètre keep_alive fixe la durée de rétention en mémoire. + Le paramètre keep_alive fixe la durée de rétention en mémoire. Les + `options` (num_ctx, num_batch, num_gpu…) DOIVENT correspondre à celles du + chat : sinon Ollama chargerait un runner distinct puis en rechargerait un + autre au premier message — un double chargement très coûteux pour un gros + modèle. """ + payload: dict = {"model": model, "keep_alive": keep_alive, "stream": False} + if options: + payload["options"] = options # Le chargement se fait désormais en tâche de fond côté API Loki. On lui # laisse jusqu'à dix minutes pour les gros modèles ou un stockage lent. timeout = httpx.Timeout(connect=10.0, read=600.0, write=30.0, pool=10.0) async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client: for attempt in range(2): resp = await client.post( - f"{self.host}/api/generate", - json={"model": model, "keep_alive": keep_alive, "stream": False}, + f"{self.host}/api/generate", json=payload ) try: # Inclut le corps JSON d'Ollama dans l'erreur (OOM, runner…), diff --git a/backend/app/routes/models.py b/backend/app/routes/models.py index 874356f..6151440 100644 --- a/backend/app/routes/models.py +++ b/backend/app/routes/models.py @@ -9,6 +9,7 @@ from fastapi import APIRouter, HTTPException, Query from fastapi.responses import StreamingResponse from pydantic import BaseModel +from .. import agent_config from ..config import settings from ..ollama_client import ollama @@ -70,10 +71,29 @@ class WarmRequest(BaseModel): keep_alive: str = "30m" -async def _warm_in_background(name: str, keep_alive: str) -> None: +async def _placement_of(name: str) -> dict: + """Où le modèle vient-il d'être chargé ? (GPU / CPU / mixte), via /api/ps.""" try: - await ollama.warm(name, keep_alive) - _warm_states[name] = {"state": "loaded"} + for m in await ollama.ps(): + if m.get("name") == name or m.get("model") == name: + size = m.get("size", 0) or 0 + vram = m.get("size_vram", 0) or 0 + if size <= 0: + return {} + pct = int(vram / size * 100) + where = "gpu" if pct >= 99 else "cpu" if pct <= 1 else "mixte" + return {"processor": where, "gpu_percent": str(pct)} + except (httpx.HTTPError, OSError): + pass + return {} + + +async def _warm_in_background(name: str, keep_alive: str, options: dict | None) -> None: + try: + await ollama.warm(name, keep_alive, options) + state = {"state": "loaded"} + state.update(await _placement_of(name)) + _warm_states[name] = state except Exception as exc: _warm_states[name] = { "state": "error", @@ -82,12 +102,19 @@ async def _warm_in_background(name: str, keep_alive: str) -> None: def start_model_warm(name: str, keep_alive: str) -> None: - """Démarre au plus une tâche de préchargement par modèle.""" + """Démarre au plus une tâche de préchargement par modèle. + + Les options de génération du modèle (num_ctx, num_batch, num_gpu…) sont + envoyées au préchargement pour qu'Ollama charge exactement le runner que le + chat utilisera — évite un rechargement complet au premier message. + """ current = _warm_tasks.get(name) if current and not current.done(): return + cfg = agent_config.get_config(name) + options = agent_config.ollama_options(cfg) _warm_states[name] = {"state": "loading"} - task = asyncio.create_task(_warm_in_background(name, keep_alive)) + task = asyncio.create_task(_warm_in_background(name, keep_alive, options)) _warm_tasks[name] = task def forget(done: asyncio.Task[None]) -> None: diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index 8bca677..3262d5f 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -99,6 +99,8 @@ async function apiError(res: Response, fallback: string): Promise { interface WarmStatus { state: "idle" | "loading" | "loaded" | "error"; error?: string; + processor?: "gpu" | "cpu" | "mixte"; + gpu_percent?: string; } const wait = (ms: number) => new Promise((resolve) => setTimeout(resolve, ms)); @@ -107,7 +109,10 @@ const wait = (ms: number) => new Promise((resolve) => setTimeout(resolve, ms)); * Lance le préchargement en arrière-plan puis suit son état. Chaque requête * reste courte afin qu'un reverse proxy ne puisse plus interrompre le warm-up. */ -export async function warmModel(name: string, keepAlive = "30m"): Promise { +export async function warmModel( + name: string, + keepAlive = "30m" +): Promise { const res = await fetch("/api/models/warm", { method: "POST", headers: { "Content-Type": "application/json" }, @@ -126,7 +131,7 @@ export async function warmModel(name: string, keepAlive = "30m"): Promise throw await apiError(statusRes, `suivi du préchargement refusé (${statusRes.status})`); } const status = (await statusRes.json()) as WarmStatus; - if (status.state === "loaded") return; + if (status.state === "loaded") return status; if (status.state === "error") { throw new Error(status.error ?? "préchargement impossible"); } diff --git a/frontend/src/store/useStore.ts b/frontend/src/store/useStore.ts index a7353f7..1038e56 100644 --- a/frontend/src/store/useStore.ts +++ b/frontend/src/store/useStore.ts @@ -159,8 +159,20 @@ export const useStore = create((set, get) => ({ // La configuration est propre au modèle : on attend son chargement avant // d'utiliser keep_alive, sinon la valeur du modèle précédent est envoyée. const ka = get().config?.keep_alive ?? "30m"; - await warmModel(name, ka); + const st = await warmModel(name, ka); await get().refreshLoadedModels(); + // Modèle chargé hors GPU : prévenir (souvent trop gros pour la VRAM). + if ( + (st.processor === "cpu" || st.processor === "mixte") && + get().selectedModel === name + ) { + set({ + warmError: + st.processor === "cpu" + ? "Modèle chargé sur le CPU (trop gros pour la VRAM) — lent. Choisis un modèle plus petit ou une quantization plus légère." + : `Modèle en partie sur GPU (${st.gpu_percent ?? "?"}%) : dépasse la VRAM, chargement plus lent.`, + }); + } } catch (err) { if (get().selectedModel === name) { set({