feat: clean deterministic OOM tasks on startup
This commit is contained in:
160
tests/test_queue_cleanup.py
Normal file
160
tests/test_queue_cleanup.py
Normal file
@@ -0,0 +1,160 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
PACKAGE_DIR = Path(__file__).resolve().parents[1] / "modelhub_submmit_api"
|
||||
if str(PACKAGE_DIR) in sys.path:
|
||||
sys.path.remove(str(PACKAGE_DIR))
|
||||
sys.path.insert(0, str(PACKAGE_DIR))
|
||||
|
||||
from modelhub_client import ModelHubClient, ModelHubClientPool # noqa: E402
|
||||
from queue_cleanup import OwnedTask, cleanup_certain_oom_tasks, find_certain_oom_tasks # noqa: E402
|
||||
|
||||
|
||||
GIB = 1024**3
|
||||
|
||||
|
||||
class FakeQueueClient:
|
||||
def __init__(self, records: list[dict[str, Any]], *, disappear_on_recheck: bool = False) -> None:
|
||||
self.records = list(records)
|
||||
self.disappear_on_recheck = disappear_on_recheck
|
||||
self.waiting_reads = 0
|
||||
self.stopped: list[list[int]] = []
|
||||
|
||||
def list_tasks_page(self, *, status: str, **_kwargs: Any) -> dict[str, Any]:
|
||||
if status == "waiting":
|
||||
self.waiting_reads += 1
|
||||
if self.disappear_on_recheck and self.waiting_reads >= 2:
|
||||
records: list[dict[str, Any]] = []
|
||||
else:
|
||||
records = [record for record in self.records if record["status"] == "waiting"]
|
||||
else:
|
||||
records = [record for record in self.records if record["status"] == status]
|
||||
return {"code": 0, "data": {"records": records, "pages": 1}}
|
||||
|
||||
def stop_tasks(self, task_ids: list[int]) -> dict[str, Any]:
|
||||
self.stopped.append(list(task_ids))
|
||||
ids = set(task_ids)
|
||||
self.records = [record for record in self.records if int(record["taskId"]) not in ids]
|
||||
return {"code": 0, "data": None}
|
||||
|
||||
def count_active_tasks(self, **_kwargs: Any) -> int:
|
||||
return len(self.records)
|
||||
|
||||
|
||||
class FakeDiscovery:
|
||||
def __init__(self, sizes: dict[str, int | None]) -> None:
|
||||
self.sizes = sizes
|
||||
|
||||
def list_repo_tree(self, repo_id: str) -> list[dict[str, Any]]:
|
||||
size = self.sizes[repo_id]
|
||||
if size is None:
|
||||
return [{"Path": "model.safetensors"}]
|
||||
return [{"Path": "model.safetensors", "Size": size}]
|
||||
|
||||
|
||||
class RecordingHttpClient:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[str, str, dict[str, Any] | None, dict[str, Any] | None]] = []
|
||||
|
||||
def request_json(
|
||||
self,
|
||||
method: str,
|
||||
path: str,
|
||||
*,
|
||||
query: dict[str, Any] | None = None,
|
||||
data: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
self.calls.append((method, path, query, data))
|
||||
return {"code": 0, "data": None}
|
||||
|
||||
|
||||
class QueueCleanupTests(unittest.TestCase):
|
||||
def test_stop_tasks_uses_documented_put_endpoint_and_integer_ids(self) -> None:
|
||||
http = RecordingHttpClient()
|
||||
client = ModelHubClient(http_client=http) # type: ignore[arg-type]
|
||||
|
||||
client.stop_tasks(["12", 12, 13])
|
||||
|
||||
self.assertEqual(
|
||||
[
|
||||
(
|
||||
"PUT",
|
||||
"/api/async/task/stop-create-contest-task",
|
||||
None,
|
||||
{"taskIds": [12, 13]},
|
||||
)
|
||||
],
|
||||
http.calls,
|
||||
)
|
||||
|
||||
def test_only_exact_size_capacity_failures_are_selected(self) -> None:
|
||||
tasks = [
|
||||
OwnedTask(0, 1, "owner/too-large", "Iluvatar_bi-100", "waiting"),
|
||||
OwnedTask(0, 2, "owner/fits", "Iluvatar_bi-100", "waiting"),
|
||||
OwnedTask(0, 3, "owner/unknown-size", "Iluvatar_bi-100", "waiting"),
|
||||
OwnedTask(0, 4, "owner/unknown-gpu", "new-gpu", "waiting"),
|
||||
]
|
||||
|
||||
selected, skipped = find_certain_oom_tasks(
|
||||
tasks,
|
||||
repository_sizes={
|
||||
"owner/too-large": 30 * GIB,
|
||||
"owner/fits": 20 * GIB,
|
||||
"owner/unknown-gpu": 30 * GIB,
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual([1], [item["taskId"] for item in selected])
|
||||
self.assertEqual(1, skipped["repositorySizeUnknown"])
|
||||
self.assertEqual(1, skipped["gpuCapacityUnknown"])
|
||||
self.assertEqual(1, skipped["fitsKnownCapacity"])
|
||||
|
||||
def test_cleanup_stops_only_certain_oom_tasks_on_the_owning_account(self) -> None:
|
||||
first = FakeQueueClient(
|
||||
[
|
||||
{"taskId": 1, "modelId": "owner/large", "gpuType": "Iluvatar_bi-100", "status": "waiting"},
|
||||
{"taskId": 2, "modelId": "owner/small", "gpuType": "Iluvatar_bi-100", "status": "waiting"},
|
||||
]
|
||||
)
|
||||
second = FakeQueueClient(
|
||||
[{"taskId": 3, "modelId": "owner/large", "gpuType": "MetaX_c-500", "status": "waiting"}]
|
||||
)
|
||||
pool = ModelHubClientPool([first, second], active_task_cap=100) # type: ignore[list-item]
|
||||
|
||||
summary = cleanup_certain_oom_tasks(
|
||||
pool,
|
||||
FakeDiscovery({"owner/large": 40 * GIB, "owner/small": 20 * GIB}), # type: ignore[arg-type]
|
||||
log=lambda _message: None,
|
||||
)
|
||||
|
||||
self.assertEqual(1, summary["certainOomCount"])
|
||||
self.assertEqual(1, summary["cancelledCount"])
|
||||
self.assertEqual([[1]], first.stopped)
|
||||
self.assertEqual([], second.stopped)
|
||||
|
||||
def test_task_that_disappears_during_scan_is_not_stopped(self) -> None:
|
||||
client = FakeQueueClient(
|
||||
[{"taskId": 1, "modelId": "owner/large", "gpuType": "Iluvatar_bi-100", "status": "waiting"}],
|
||||
disappear_on_recheck=True,
|
||||
)
|
||||
pool = ModelHubClientPool([client], active_task_cap=100) # type: ignore[list-item]
|
||||
|
||||
summary = cleanup_certain_oom_tasks(
|
||||
pool,
|
||||
FakeDiscovery({"owner/large": 40 * GIB}), # type: ignore[arg-type]
|
||||
log=lambda _message: None,
|
||||
)
|
||||
|
||||
self.assertEqual(1, summary["certainOomCount"])
|
||||
self.assertEqual(0, summary["cancelledCount"])
|
||||
self.assertEqual(1, summary["noLongerActiveCount"])
|
||||
self.assertEqual([], client.stopped)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user