Skip to content

Commit 1848765

Browse files
authored
Merge pull request #9 from runpod-workers/luke/prewarm-first-request
Pre-load the model into GPU memory at boot so the first request doesn't
2 parents 126fde8 + 71a5e2e commit 1848765

4 files changed

Lines changed: 359 additions & 3 deletions

File tree

‎README.md‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -179,6 +179,8 @@ Non-streaming responses return Ollama's native response object:
179179
| `OLLAMA_KEEP_ALIVE` | `-1` (forever) | How long models stay loaded in VRAM |
180180
| `OLLAMA_LOAD_TIMEOUT` | `60m` | How long Ollama waits for a model to load into memory before failing the request. Ollama's own default is `5m`, which a large model on a fresh worker can exceed |
181181

182+
**Boot warm-up:** after registering the configured model, the worker loads it into GPU memory before accepting jobs (an empty-prompt `/api/generate`, bounded by `OLLAMA_LOAD_TIMEOUT`). Registering only writes the model to disk — Ollama loads VRAM lazily on the first inference — so without the warm-up the *first user request* pays the multi-minute load and can hit the execution timeout. With `OLLAMA_KEEP_ALIVE=-1` the weights then stay resident for the worker's lifetime. If the warm-up fails, the worker logs the cause and still starts; the first request retries the load.
183+
182184
## Storage and disk sizing
183185

184186
| Configuration | Ollama store | GGUF acquisition | Disk needed |

‎handler.py‎

Lines changed: 131 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import os
44
import re
55
import subprocess
6+
import time
67
import uuid
78

89
import requests
@@ -852,6 +853,136 @@ def ensure_default_model():
852853
return model
853854

854855

856+
# Go-style durations, which is what Ollama's env vars use. "ms" before "m" so
857+
# "500ms" doesn't half-match as "500m" plus a dangling "s".
858+
_DURATION_RE = re.compile(r"^([+-]?)((?:\d+(?:\.\d+)?(?:ms|us|ns|h|m|s))+)$")
859+
_DURATION_PART_RE = re.compile(r"(\d+(?:\.\d+)?)(ms|us|ns|h|m|s)")
860+
_UNIT_SECONDS = {"h": 3600.0, "m": 60.0, "s": 1.0, "ms": 1e-3, "us": 1e-6, "ns": 1e-9}
861+
862+
# Fallback and cap for the warm-up bound: never let boot hang on a load forever,
863+
# even when OLLAMA_LOAD_TIMEOUT is unset, unparseable, or "retry indefinitely".
864+
_WARMUP_FALLBACK_S = 3600.0
865+
# Client-side margin over Ollama's own load timeout, so the server's error — which
866+
# says *why* the load failed — wins the race against a blank client timeout.
867+
_WARMUP_MARGIN_S = 60.0
868+
869+
870+
def _duration_seconds(value, default):
871+
"""Parse a Go time.ParseDuration string ('60m', '1h30m', '90s') to seconds.
872+
873+
Ollama parses OLLAMA_LOAD_TIMEOUT with Go's parser; mirroring it here keeps
874+
one env var meaning one thing. Unparseable values return `default` rather
875+
than raising — a bad env var must not break boot.
876+
"""
877+
match = _DURATION_RE.match((value or "").strip().lower())
878+
if not match:
879+
return default
880+
total = sum(
881+
float(number) * _UNIT_SECONDS[unit]
882+
for number, unit in _DURATION_PART_RE.findall(match.group(2))
883+
)
884+
return -total if match.group(1) == "-" else total
885+
886+
887+
def _warmup_timeout_seconds():
888+
"""How long the boot warm-up may spend loading the model into memory.
889+
890+
Ollama itself gives up after OLLAMA_LOAD_TIMEOUT (its default is 5m; this
891+
image's Hub default is 60m; non-positive means retry forever), so the client
892+
bound sits just above that: Ollama's load-timeout error names the model and
893+
the cause, which beats a mute client-side timeout. Non-positive or
894+
unparseable values are capped at an hour so warm-up is never unbounded.
895+
"""
896+
timeout = _duration_seconds(os.environ.get("OLLAMA_LOAD_TIMEOUT"), _WARMUP_FALLBACK_S)
897+
if timeout <= 0:
898+
timeout = _WARMUP_FALLBACK_S
899+
return timeout + _WARMUP_MARGIN_S
900+
901+
902+
def _report_residency(model):
903+
"""Log where the loaded weights actually ended up (VRAM vs CPU spill).
904+
905+
/api/ps is the only place Ollama reports size_vram; a partial-VRAM load is
906+
the "first request is mysteriously slow" case worth flagging at boot.
907+
"""
908+
response = session.get(f"{OLLAMA_BASE_URL}/api/ps", timeout=30)
909+
if not response.ok:
910+
return
911+
for entry in response.json().get("models", []):
912+
if entry.get("name") not in (model, f"{model}:latest"):
913+
continue
914+
size = entry.get("size") or 0
915+
vram = entry.get("size_vram") or 0
916+
if size and vram < size:
917+
print(
918+
f"WARN: '{model}' loaded, but only {_human_size(vram)} of "
919+
f"{_human_size(size)} fits in VRAM — the rest is offloaded to CPU and "
920+
f"responses will be much slower. Pick a GPU with more VRAM, a smaller "
921+
f"quantization via HF_QUANTIZATION, or a lower OLLAMA_CONTEXT_LENGTH.",
922+
flush=True,
923+
)
924+
else:
925+
print(f"'{model}' is resident in GPU memory ({_human_size(vram)})", flush=True)
926+
return
927+
print(f"WARN: '{model}' answered the warm-up but is not listed by /api/ps", flush=True)
928+
929+
930+
def warm_model(model):
931+
"""Load `model`'s weights into GPU memory before any user request arrives.
932+
933+
Registering a model (ensure_default_model) only writes it to disk — Ollama
934+
loads weights into VRAM lazily, on the first inference. Without this, the
935+
first user request pays the whole multi-minute load and can blow the
936+
endpoint's execution timeout. An empty prompt is Ollama's documented
937+
"just load it" request: no tokens are generated, and the image's
938+
OLLAMA_KEEP_ALIVE=-1 default keeps the weights resident afterwards.
939+
"""
940+
timeout = _warmup_timeout_seconds()
941+
print(
942+
f"Loading '{model}' into GPU memory (bounded at {int(timeout)}s "
943+
f"by OLLAMA_LOAD_TIMEOUT)...",
944+
flush=True,
945+
)
946+
started = time.monotonic()
947+
response = session.post(
948+
f"{OLLAMA_BASE_URL}/api/generate",
949+
json={"model": model, "prompt": "", "stream": False},
950+
timeout=timeout,
951+
)
952+
error = ollama_error(response)
953+
if error:
954+
raise ValueError(error)
955+
print(f"Loaded '{model}' in {time.monotonic() - started:.1f}s", flush=True)
956+
try:
957+
_report_residency(model)
958+
except (requests.RequestException, ValueError) as err:
959+
print(f"WARN: could not read /api/ps after warming '{model}': {err}", flush=True)
960+
961+
962+
def warm_default_model():
963+
"""Best-effort boot-time warm-up of the model the endpoint is configured for.
964+
965+
Deliberately non-fatal: the model is already registered, so the handler still
966+
works if this fails — the first request just pays the load, as it did before
967+
warm-up existed. But say exactly what failed, so "too big for this GPU" or
968+
"load timed out" is diagnosable from the boot log instead of from a timed-out
969+
first request. Returns True when the model is warm.
970+
"""
971+
model = resolve_default_model()
972+
if not model:
973+
return False
974+
try:
975+
warm_model(model)
976+
return True
977+
except (requests.RequestException, ValueError, OSError) as err:
978+
print(
979+
f"WARN: could not pre-load '{model}' into GPU memory — the first request "
980+
f"will trigger the load instead and may be slow or time out. Cause: {err}",
981+
flush=True,
982+
)
983+
return False
984+
985+
855986
def handler(job):
856987
job_input = job.get("input") or {}
857988

‎start.sh‎

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -92,11 +92,15 @@ fi
9292

9393
# Always runs: resolve_default_model has a fallback, so there is always a model to
9494
# prepare. ensure_default_model resolves HF_MODEL over OLLAMA_MODEL, reads Runpod's
95-
# model store first, and uses HF_TOKEN for gated repos.
96-
if /opt/venv/bin/python -c 'import handler; print("Model ready:", handler.ensure_default_model())'; then
95+
# model store first, and uses HF_TOKEN for gated repos. warm_default_model then
96+
# loads the weights into GPU memory: registering only writes the model to disk,
97+
# and without the warm-up the *first user request* pays the multi-minute VRAM
98+
# load and can blow the execution timeout. Warm-up failure is non-fatal (the
99+
# handler retries the load on the first request) and logs its own cause.
100+
if /opt/venv/bin/python -c 'import handler; print("Model ready:", handler.ensure_default_model()); handler.warm_default_model()'; then
97101
:
98102
else
99-
echo "WARN: startup model preparation failed — the handler will retry on the first request"
103+
echo "WARN: startup model preparation failed (cause in the traceback above) — the handler will retry on the first request, which will be slow and may time out"
100104
fi
101105

102106
cd /

‎test_warmup.py‎

Lines changed: 219 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,219 @@
1+
"""Unit tests for the boot-time warm-up logic in handler.py.
2+
3+
These need no GPU, no network and no running Ollama, which is the point: the
4+
warm-up only ever runs on real hardware at cold start, so getting its request
5+
shape, bounds or failure handling wrong is otherwise only discovered by a timed
6+
out first request in production.
7+
8+
python3 -m pytest test_warmup.py -v
9+
"""
10+
11+
import json
12+
import sys
13+
import types
14+
15+
import pytest
16+
17+
# handler.py imports the runpod SDK at module scope; stub it when it isn't
18+
# installed so the pure functions stay testable from a bare checkout.
19+
if "runpod" not in sys.modules:
20+
try:
21+
import runpod # noqa: F401
22+
except ImportError:
23+
sys.modules["runpod"] = types.ModuleType("runpod")
24+
25+
import requests
26+
27+
import handler
28+
29+
30+
class FakeResponse:
31+
def __init__(self, status_code=200, payload=None, text=None, path="/api/generate"):
32+
self.status_code = status_code
33+
self.ok = status_code < 400
34+
self._payload = payload
35+
self.text = text if text is not None else json.dumps(payload or {})
36+
self.request = types.SimpleNamespace(path_url=path)
37+
38+
def json(self):
39+
if self._payload is None:
40+
raise ValueError("no json")
41+
return self._payload
42+
43+
44+
class FakeSession:
45+
"""Records every call so tests can assert on the exact request shape."""
46+
47+
def __init__(self, post_response, ps_response=None):
48+
self.post_response = post_response
49+
self.ps_response = ps_response or FakeResponse(payload={"models": []}, path="/api/ps")
50+
self.posts = []
51+
self.gets = []
52+
53+
def post(self, url, **kwargs):
54+
self.posts.append((url, kwargs))
55+
return self.post_response
56+
57+
def get(self, url, **kwargs):
58+
self.gets.append((url, kwargs))
59+
return self.ps_response
60+
61+
62+
def loaded_ps(model, size=1000, size_vram=None):
63+
return FakeResponse(
64+
payload={
65+
"models": [
66+
{"name": model, "size": size, "size_vram": size if size_vram is None else size_vram}
67+
]
68+
},
69+
path="/api/ps",
70+
)
71+
72+
73+
# --- _duration_seconds: Go-style durations, as Ollama's env vars use ------------
74+
75+
76+
@pytest.mark.parametrize(
77+
"value, expected",
78+
[
79+
("60m", 3600.0),
80+
("1h30m", 5400.0),
81+
("90s", 90.0),
82+
("500ms", 0.5),
83+
("1.5h", 5400.0),
84+
("-5m", -300.0),
85+
],
86+
)
87+
def test_duration_seconds_parses(value, expected):
88+
assert handler._duration_seconds(value, 123.0) == pytest.approx(expected)
89+
90+
91+
@pytest.mark.parametrize("value", ["", None, "garbage", "10", "5x", "m5", "1h 30m"])
92+
def test_duration_seconds_falls_back(value):
93+
assert handler._duration_seconds(value, 123.0) == 123.0
94+
95+
96+
# --- _warmup_timeout_seconds: always bounded, tracks OLLAMA_LOAD_TIMEOUT --------
97+
98+
99+
def test_warmup_timeout_tracks_load_timeout(monkeypatch):
100+
monkeypatch.setenv("OLLAMA_LOAD_TIMEOUT", "10m")
101+
assert handler._warmup_timeout_seconds() == 600 + handler._WARMUP_MARGIN_S
102+
103+
104+
@pytest.mark.parametrize("value", [None, "0s", "-1s", "not-a-duration"])
105+
def test_warmup_timeout_never_unbounded(monkeypatch, value):
106+
if value is None:
107+
monkeypatch.delenv("OLLAMA_LOAD_TIMEOUT", raising=False)
108+
else:
109+
monkeypatch.setenv("OLLAMA_LOAD_TIMEOUT", value)
110+
assert (
111+
handler._warmup_timeout_seconds()
112+
== handler._WARMUP_FALLBACK_S + handler._WARMUP_MARGIN_S
113+
)
114+
115+
116+
# --- warm_model: the load request itself ----------------------------------------
117+
118+
119+
def test_warm_model_sends_empty_prompt_load(monkeypatch):
120+
monkeypatch.setenv("OLLAMA_LOAD_TIMEOUT", "10m")
121+
session = FakeSession(
122+
FakeResponse(payload={"done": True}), ps_response=loaded_ps("llama3.2:3b")
123+
)
124+
monkeypatch.setattr(handler, "session", session)
125+
126+
handler.warm_model("llama3.2:3b")
127+
128+
url, kwargs = session.posts[0]
129+
assert url.endswith("/api/generate")
130+
# An empty prompt is Ollama's "load into memory, generate nothing" request.
131+
assert kwargs["json"] == {"model": "llama3.2:3b", "prompt": "", "stream": False}
132+
assert kwargs["timeout"] == 600 + handler._WARMUP_MARGIN_S
133+
134+
135+
def test_warm_model_raises_on_server_error(monkeypatch):
136+
session = FakeSession(
137+
FakeResponse(
138+
status_code=500,
139+
payload={"error": "timed out waiting for llama runner to start"},
140+
)
141+
)
142+
monkeypatch.setattr(handler, "session", session)
143+
144+
with pytest.raises(ValueError, match="timed out waiting for llama runner"):
145+
handler.warm_model("llama3.2:3b")
146+
147+
148+
def test_warm_model_warns_on_cpu_offload(monkeypatch, capsys):
149+
session = FakeSession(
150+
FakeResponse(payload={"done": True}),
151+
ps_response=loaded_ps("big-model:latest", size=1000, size_vram=400),
152+
)
153+
monkeypatch.setattr(handler, "session", session)
154+
155+
handler.warm_model("big-model")
156+
157+
out = capsys.readouterr().out
158+
assert "offloaded to CPU" in out
159+
160+
161+
def test_warm_model_reports_full_residency(monkeypatch, capsys):
162+
session = FakeSession(
163+
FakeResponse(payload={"done": True}), ps_response=loaded_ps("llama3.2:3b")
164+
)
165+
monkeypatch.setattr(handler, "session", session)
166+
167+
handler.warm_model("llama3.2:3b")
168+
169+
assert "resident in GPU memory" in capsys.readouterr().out
170+
171+
172+
def test_warm_model_survives_ps_failure(monkeypatch, capsys):
173+
"""A broken /api/ps must not turn a successful load into a failure."""
174+
session = FakeSession(FakeResponse(payload={"done": True}))
175+
176+
def broken_get(url, **kwargs):
177+
raise requests.ConnectionError("ps is down")
178+
179+
session.get = broken_get
180+
monkeypatch.setattr(handler, "session", session)
181+
182+
handler.warm_model("llama3.2:3b") # must not raise
183+
184+
assert "could not read /api/ps" in capsys.readouterr().out
185+
186+
187+
# --- warm_default_model: non-fatal, but says what failed -------------------------
188+
189+
190+
def test_warm_default_model_success(monkeypatch):
191+
monkeypatch.setattr(handler, "resolve_default_model", lambda: "llama3.2:3b")
192+
warmed = []
193+
monkeypatch.setattr(handler, "warm_model", warmed.append)
194+
195+
assert handler.warm_default_model() is True
196+
assert warmed == ["llama3.2:3b"]
197+
198+
199+
@pytest.mark.parametrize(
200+
"err",
201+
[
202+
requests.ConnectionError("connection refused"),
203+
requests.Timeout("read timed out"),
204+
ValueError("HTTP 500 from /api/generate: model requires more system memory"),
205+
],
206+
)
207+
def test_warm_default_model_failure_is_nonfatal_but_loud(monkeypatch, capsys, err):
208+
monkeypatch.setattr(handler, "resolve_default_model", lambda: "big-model")
209+
210+
def failing_warm(model):
211+
raise err
212+
213+
monkeypatch.setattr(handler, "warm_model", failing_warm)
214+
215+
assert handler.warm_default_model() is False
216+
out = capsys.readouterr().out
217+
assert "could not pre-load 'big-model'" in out
218+
# The cause must be surfaced, not swallowed behind a bare WARN.
219+
assert str(err) in out

0 commit comments

Comments
 (0)