Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion frontend/src/pages/Playground.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -291,7 +291,11 @@ export default function Playground() {
<Pane tone="base" label="Base" text={base[i]?.content} />
</div>
))}
{busy && <div className="grid grid-cols-2 gap-3"><Pane tone="tuned" label="Shadow" /><Pane tone="base" label="Base" /></div>}
{/* until the shadow answers, both panes wait here; after, its row
carries the base's wait, so don't draw a second one */}
{busy && msgs[msgs.length - 1]?.role === "user" && (
<div className="grid grid-cols-2 gap-3"><Pane tone="tuned" label="Shadow" /><Pane tone="base" label="Base" /></div>
)}
</div>
) : (
<div className="mx-auto flex max-w-3xl flex-col gap-4 py-6">
Expand Down

Large diffs are not rendered by default.

2 changes: 1 addition & 1 deletion shadowlm/_static/index.html
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
href="https://fonts.googleapis.com/css2?family=Manrope:wght@400;500;600;700&family=Sora:wght@500;600;700&family=JetBrains+Mono:wght@400;500;600&display=swap"
/>
<title>openfinetuner</title>
<script type="module" crossorigin src="./assets/index-ChIABspn.js"></script>
<script type="module" crossorigin src="./assets/index-Dn6kMFwk.js"></script>
<link rel="stylesheet" crossorigin href="./assets/index-sV2ScQ20.css">
</head>
<body>
Expand Down
17 changes: 14 additions & 3 deletions shadowlm/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -836,6 +836,11 @@ def _on_eval(m, _tee=tee, _tag=tag):
job.logs.append(job.live)
job.live = ""
self._persist(job) # terminal state + final metrics + logs → disk
# The trained model would otherwise live on in this loop's locals
# until the next job reassigns them, holding the GPU between runs
# where Clean VRAM can't reach it: let it go now.
be = result = callbacks = None # noqa: F841
self._release_gpu_cache()

# ---- inference -------------------------------------------------------------
def _infer_key(self, model: str, adapter: str | None,
Expand Down Expand Up @@ -929,10 +934,16 @@ def _gpu_used_mb() -> int | None:
def _drop_inference_models(self) -> tuple[int, str | None]:
"""Empty the inference cache and release the GPU allocator's cache.
Caller holds ``_model_lock``. Returns (models dropped, release error)."""
import gc # noqa: PLC0415

n = len(self._infer_cache)
self._infer_cache.clear()
return n, self._release_gpu_cache()

@staticmethod
def _release_gpu_cache() -> str | None:
"""Collect dropped models and hand the allocator's cache back to the
GPU. Returns the release error, if any."""
import gc # noqa: PLC0415

gc.collect()
freed_error: str | None = None
try:
Expand All @@ -944,7 +955,7 @@ def _drop_inference_models(self) -> tuple[int, str | None]:
freed_error = f"{type(e).__name__}: {e}"
print(f"[serve] VRAM release failed ({freed_error})", flush=True)
gc.collect()
return n, freed_error
return freed_error

def clear_vram(self) -> dict:
"""Drop every cached inference model and release the GPU allocator's
Expand Down
49 changes: 49 additions & 0 deletions tests/test_serve_vram.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@

from __future__ import annotations

import threading

from shadowlm.serve import Server


Expand All @@ -27,3 +29,50 @@ def test_clear_vram_reports_what_it_unloaded(tmp_path):
out = s.clear_vram()
assert out["unloaded"] == 1 and "error" not in out
assert s._infer_cache == {}


def test_a_finished_run_lets_go_of_its_model(tmp_path, monkeypatch):
"""The trained model isn't kept alive in the runner between jobs, where
Clean VRAM can't reach it."""
import gc
import weakref

from shadowlm import backends

alive: list[weakref.ref] = []

class FakeResult:
checkpoint = None
final_loss = 0.5

class FakeBackend:
def load(self, *a, **k):
self.weights = bytearray(1024)

def finetune(self, *a, **k):
return FakeResult()

def fake_select(*a, **k):
be = FakeBackend()
alive.append(weakref.ref(be))
return be

monkeypatch.setattr(backends, "select_backend", fake_select)
s = _server(tmp_path)
job_id = s.submit({
"base_model": "fake/model", "load_in_4bit": False, "max_seq_length": 64,
"config": {"method": "lora", "max_steps": 1},
"dataset": {"rows": [{"text": "hi"}], "format": "text"},
"eval_dataset": None,
})
for _ in range(200):
if s.jobs[job_id].status in ("succeeded", "failed"):
break
threading.Event().wait(0.02)
assert s.jobs[job_id].status == "succeeded", s.jobs[job_id].error
for _ in range(50): # the runner drops it right after persisting
gc.collect()
if alive and alive[0]() is None:
break
threading.Event().wait(0.02)
assert alive and alive[0]() is None
Loading