feat: add per-model GPU and RAM placement

This commit is contained in:
Hermes Agent
2026-08-25 23:17:30 +10:00
parent 238de79607
commit 1029d23605
6 changed files with 60 additions and 17 deletions
+29 -9
View File
@@ -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}