feat: add per-model GPU and RAM placement
This commit is contained in:
+29
-9
@@ -1116,6 +1116,7 @@ def _new_job(name: str, action: str, target: str = "local") -> str:
|
||||
class ModelRequest(BaseModel):
|
||||
name: str
|
||||
target: str = "local"
|
||||
placement: str = "gpu_ram"
|
||||
|
||||
|
||||
class ConnectionRequest(BaseModel):
|
||||
@@ -1125,6 +1126,7 @@ class ConnectionRequest(BaseModel):
|
||||
|
||||
class ModelsRequest(BaseModel):
|
||||
names: list[str] = Field(default_factory=list)
|
||||
placements: dict[str, str] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class ChatAttachment(BaseModel):
|
||||
@@ -1140,6 +1142,7 @@ class ChatRequest(BaseModel):
|
||||
message: str = ""
|
||||
history: list[dict[str, Any]] = Field(default_factory=list)
|
||||
attachments: list[ChatAttachment] = Field(default_factory=list)
|
||||
placements: dict[str, str] = Field(default_factory=dict)
|
||||
request_id: str = ""
|
||||
conversation_id: str = ""
|
||||
|
||||
@@ -1172,18 +1175,29 @@ def _ollama_error(exc: HTTPError) -> HTTPException:
|
||||
return HTTPException(502, f"Ollama request failed: {detail}")
|
||||
|
||||
|
||||
def _load_model(name: str) -> dict[str, Any]:
|
||||
def _model_placement(value: str | None) -> str:
|
||||
placement = str(value or "gpu_ram").strip().lower()
|
||||
if placement not in {"gpu_ram", "ram_only"}:
|
||||
raise HTTPException(400, "Model placement must be gpu_ram or ram_only")
|
||||
return placement
|
||||
|
||||
|
||||
def _load_model(name: str, placement: str = "gpu_ram") -> dict[str, Any]:
|
||||
name = _require_installed_model(name)
|
||||
placement = _model_placement(placement)
|
||||
options: dict[str, Any] = {"num_predict": 1}
|
||||
if placement == "ram_only":
|
||||
options["num_gpu"] = 0
|
||||
try:
|
||||
result = _json_request(
|
||||
LOCAL_OLLAMA + "/api/generate",
|
||||
method="POST",
|
||||
payload={"model": name, "prompt": "", "stream": False, "keep_alive": CHAT_KEEP_ALIVE, "options": {"num_predict": 1}},
|
||||
payload={"model": name, "prompt": "", "stream": False, "keep_alive": CHAT_KEEP_ALIVE, "options": options},
|
||||
timeout=900,
|
||||
)
|
||||
except HTTPError as exc:
|
||||
raise _ollama_error(exc) from exc
|
||||
return {"ok": True, "model": name, "response": result.get("response", ""), "runtime": _runtime_snapshot()}
|
||||
return {"ok": True, "model": name, "placement": placement, "response": result.get("response", ""), "runtime": _runtime_snapshot()}
|
||||
|
||||
|
||||
def _chat_payload(body: ChatRequest, model_name: str | None = None) -> dict[str, Any]:
|
||||
@@ -1208,7 +1222,10 @@ def _chat_payload(body: ChatRequest, model_name: str | None = None) -> dict[str,
|
||||
if images:
|
||||
user_message["images"] = images
|
||||
messages.append(user_message)
|
||||
return {"model": model, "messages": messages, "stream": False, "keep_alive": CHAT_KEEP_ALIVE}
|
||||
payload: dict[str, Any] = {"model": model, "messages": messages, "stream": False, "keep_alive": CHAT_KEEP_ALIVE}
|
||||
if _model_placement(body.placements.get(model)) == "ram_only":
|
||||
payload["options"] = {"num_gpu": 0}
|
||||
return payload
|
||||
|
||||
|
||||
class _ChatStopped(Exception):
|
||||
@@ -1420,7 +1437,7 @@ def runtime() -> dict[str, Any]:
|
||||
|
||||
@router.post("/chat/load")
|
||||
def chat_load(body: ModelRequest) -> dict[str, Any]:
|
||||
return _load_model(body.name)
|
||||
return _load_model(body.name, body.placement)
|
||||
|
||||
|
||||
@router.post("/models/load")
|
||||
@@ -1428,16 +1445,19 @@ def models_load(body: ModelsRequest) -> dict[str, Any]:
|
||||
names = list(dict.fromkeys(_valid_name(name) for name in body.names if str(name).strip()))[:12]
|
||||
if not names:
|
||||
raise HTTPException(400, "Select at least one model to load")
|
||||
placements = {name: _model_placement(body.placements.get(name)) for name in names}
|
||||
results = []
|
||||
for name in names:
|
||||
_model_load_update(name, active=True, state="queued", stage="Waiting for Ollama", started_at=time.time(), error="")
|
||||
_model_load_update(name, active=True, state="queued", stage="Waiting for Ollama", started_at=time.time(), error="", placement=placements[name])
|
||||
for name in names:
|
||||
_model_load_update(name, state="loading", stage="Loading model into Ollama memory")
|
||||
placement = placements[name]
|
||||
stage = "Loading into GPU + RAM" if placement == "gpu_ram" else "Loading into system RAM only"
|
||||
_model_load_update(name, state="loading", stage=stage)
|
||||
try:
|
||||
results.append({"name": name, "ok": True, "result": _load_model(name)})
|
||||
results.append({"name": name, "placement": placement, "ok": True, "result": _load_model(name, placement)})
|
||||
_model_load_update(name, state="checking", stage="Checking Ollama resident state")
|
||||
except Exception as exc:
|
||||
results.append({"name": name, "ok": False, "error": str(exc)})
|
||||
results.append({"name": name, "placement": placement, "ok": False, "error": str(exc)})
|
||||
_model_load_update(name, active=False, state="failed", stage="Ollama load failed", finished_at=time.time(), error=str(exc))
|
||||
resident_rows = _local_ps()
|
||||
resident_names = {str(row.get("name") or row.get("model")) for row in resident_rows}
|
||||
|
||||
Reference in New Issue
Block a user