feat: add adaptive GPU scheduling
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user