feat: add adaptive GPU scheduling

This commit is contained in:
CoolBoy
2026-08-02 16:59:44 +08:00
parent 80a1b8518d
commit eab5ab6dce
14 changed files with 1216 additions and 52 deletions

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
import os
import threading
import time
from pathlib import PurePosixPath
from typing import Any
from urllib.parse import quote
@@ -44,7 +45,9 @@ class HuggingFaceDiscovery:
http_client: JsonHttpClient | None = None,
legacy_http_client: JsonHttpClient | None = None,
timeout: int = 30,
retries: int = 2,
retries: int = 5,
page_interval_seconds: float | None = None,
page_cache_ttl_seconds: float | None = None,
) -> None:
token = os.getenv("MODELSCOPE_API_TOKEN") or os.getenv("MODELSCOPE_TOKEN") or EMBEDDED_MODELSCOPE_TOKEN
headers = {"User-Agent": "modelhub-submmit-cli/0.1"}
@@ -56,6 +59,7 @@ class HuggingFaceDiscovery:
default_headers=headers,
timeout=timeout,
retries=retries,
backoff_seconds=2.0,
)
self.legacy_http_client = legacy_http_client or JsonHttpClient(
base_url=base_url,
@@ -65,6 +69,24 @@ class HuggingFaceDiscovery:
)
self._repo_tree_cache: dict[str, list[dict[str, Any]]] = {}
self._repo_tree_lock = threading.Lock()
self._model_page_cache: dict[tuple[str, int, int], tuple[float, list[dict[str, Any]]]] = {}
self._model_page_cache_ttl = max(
0.0,
float(
page_cache_ttl_seconds
if page_cache_ttl_seconds is not None
else os.getenv("MODELSCOPE_PAGE_CACHE_TTL_SECONDS", "900")
),
)
self._page_interval_seconds = max(
0.0,
float(
page_interval_seconds
if page_interval_seconds is not None
else os.getenv("MODELSCOPE_PAGE_INTERVAL_SECONDS", "0.25")
),
)
self._last_model_page_request_at = 0.0
def list_recent_models(
self,
@@ -111,22 +133,37 @@ class HuggingFaceDiscovery:
task_tag = MODELSCOPE_TASK_TAGS.get(pipeline_tag, pipeline_tag)
models: list[HFModelSummary] = []
for page_number in range(1, (max_items + page_size - 1) // page_size + 1):
try:
payload = self.http_client.request_json(
"GET",
"/models",
query={
"page_number": page_number,
"page_size": page_size,
"sort": "last_modified",
"filter.task": task_tag,
},
)
except HttpJsonError as exc:
print(f"[modelscope] list_models_error task={task_tag} page={page_number} error={exc}", flush=True)
break
cache_key = (task_tag, page_number, page_size)
cached = self._model_page_cache.get(cache_key)
if cached is not None and time.monotonic() - cached[0] < self._model_page_cache_ttl:
items = list(cached[1])
else:
elapsed = time.monotonic() - self._last_model_page_request_at
if self._last_model_page_request_at > 0 and elapsed < self._page_interval_seconds:
time.sleep(self._page_interval_seconds - elapsed)
try:
payload = self.http_client.request_json(
"GET",
"/models",
query={
"page_number": page_number,
"page_size": page_size,
"sort": "last_modified",
"filter.task": task_tag,
},
)
self._last_model_page_request_at = time.monotonic()
except HttpJsonError as exc:
self._last_model_page_request_at = time.monotonic()
print(
f"[modelscope] list_models_error task={task_tag} page={page_number} "
f"partial_models={len(models)} retry_next_cycle=true error={exc}",
flush=True,
)
break
items = self._extract_models(payload)
items = self._extract_models(payload)
self._model_page_cache[cache_key] = (time.monotonic(), list(items))
if not items:
break
for item in items: