119 lines
5.6 KiB
Python
119 lines
5.6 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
from datetime import datetime, timezone
|
|
|
|
from app.clients.modelhub import ModelHubClient
|
|
from app.crawlers.download_success import DownloadSuccessCrawler
|
|
from app.crawlers.not_adapted import NotAdaptedCrawler
|
|
from app.scheduler.budget import Budget
|
|
from app.scheduler.downloader import Downloader
|
|
from app.scheduler.planner import Planner
|
|
from app.scheduler.submitter import Submitter
|
|
from app.settings import Settings
|
|
from app.storage.repositories import Repository, future_seconds
|
|
|
|
LOG = logging.getLogger(__name__)
|
|
|
|
|
|
class StrategyLoop:
|
|
def __init__(self, settings: Settings, repo: Repository, client: ModelHubClient):
|
|
self.settings = settings
|
|
self.repo = repo
|
|
self.client = client
|
|
self.budget = Budget(repo, settings.max_user_tasks)
|
|
self.planner = Planner(repo, self.budget, settings.default_gpu_alias_list)
|
|
self.submitter = Submitter(repo, client)
|
|
self.downloader = Downloader(repo, client, settings.max_self_downloads, settings.download_poll_interval_seconds)
|
|
self.not_adapted_crawler = NotAdaptedCrawler(settings, repo, client)
|
|
self.download_success_crawler = DownloadSuccessCrawler(settings, repo, client)
|
|
self._last_crawler_refresh = None
|
|
|
|
def refresh_crawlers_if_due(self) -> dict[str, int]:
|
|
now = datetime.now(timezone.utc)
|
|
if self._last_crawler_refresh is not None:
|
|
elapsed = (now - self._last_crawler_refresh).total_seconds()
|
|
if elapsed < self.settings.crawler_refresh_seconds:
|
|
return {"not_adapted": 0, "download_success": 0}
|
|
self._last_crawler_refresh = now
|
|
stats = {"not_adapted": 0, "download_success": 0}
|
|
try:
|
|
stats["not_adapted"] = self.not_adapted_crawler.refresh()
|
|
except Exception:
|
|
LOG.exception("not-adapted crawler failed")
|
|
try:
|
|
stats["download_success"] = self.download_success_crawler.refresh()
|
|
except Exception:
|
|
LOG.exception("download-success crawler failed")
|
|
return stats
|
|
|
|
def mark_ready_tasks(self, limit: int = 50) -> int:
|
|
changed = 0
|
|
for task in self.repo.due_tasks(("pending",), limit=limit):
|
|
if task["is_bounty"]:
|
|
if self.settings.bounty_allow_direct_submit:
|
|
self.repo.mark_task_status(task["id"], "ready_to_submit")
|
|
changed += 1
|
|
else:
|
|
status = self.client.get_download_status(task["model_id"])
|
|
if status == "SUCCESS":
|
|
self.repo.upsert_download(task["model_id"], task["source"], "SUCCESS", "others")
|
|
self.repo.mark_task_status(task["id"], "ready_to_submit")
|
|
changed += 1
|
|
else:
|
|
self.repo.mark_task_status(task["id"], "pending", "DOWNLOAD_STATUS", status, future_seconds(300))
|
|
elif task["is_promote"]:
|
|
absent = self.client.is_model_absent_from_modelhub(task["model_id"])
|
|
if not absent:
|
|
self.repo.mark_task_status(task["id"], "ready_to_submit")
|
|
changed += 1
|
|
else:
|
|
self.repo.mark_task_status(task["id"], "pending", "NOT_IN_MODELHUB", "model not adapted yet", future_seconds(1800))
|
|
elif task["known_downloaded_by_others"]:
|
|
status = self.client.get_download_status(task["model_id"])
|
|
if status == "SUCCESS":
|
|
self.repo.upsert_download(task["model_id"], task["source"], "SUCCESS", "others")
|
|
self.repo.mark_task_status(task["id"], "ready_to_submit")
|
|
changed += 1
|
|
elif task["self_download_allowed"]:
|
|
self.repo.mark_task_status(task["id"], "waiting_download", "DOWNLOAD_STATUS", status, future_seconds(300))
|
|
else:
|
|
self.repo.mark_task_status(task["id"], "pending", "DOWNLOAD_STATUS", status, future_seconds(300))
|
|
elif task["self_download_allowed"]:
|
|
self.repo.mark_task_status(task["id"], "waiting_download")
|
|
else:
|
|
self.repo.mark_task_status(task["id"], "skipped", "NO_READY_PATH", "candidate is not bounty/promote/downloaded/self-download")
|
|
return changed
|
|
|
|
def run_once(self) -> dict[str, int]:
|
|
crawler_stats = self.refresh_crawlers_if_due()
|
|
created = self.planner.materialize_tasks()
|
|
polled = self.downloader.poll_active_downloads()
|
|
released = self.submitter.release_queue_full_backoff()
|
|
ready = self.mark_ready_tasks()
|
|
submitted = self.submitter.submit_due()
|
|
started_downloads = 0
|
|
if submitted == 0 and ready == 0:
|
|
started_downloads = self.downloader.maybe_start_self_downloads()
|
|
return {
|
|
"crawled_not_adapted": crawler_stats["not_adapted"],
|
|
"crawled_download_success": crawler_stats["download_success"],
|
|
"created_tasks": created,
|
|
"polled_downloads": polled,
|
|
"released_backoff": released,
|
|
"ready_tasks": ready,
|
|
"submitted_tasks": submitted,
|
|
"started_downloads": started_downloads,
|
|
}
|
|
|
|
def run_forever(self, stop_event):
|
|
LOG.info("strategy loop started")
|
|
while not stop_event.is_set():
|
|
try:
|
|
stats = self.run_once()
|
|
LOG.info("loop stats: %s", stats)
|
|
except Exception:
|
|
LOG.exception("strategy loop iteration failed")
|
|
stop_event.wait(self.settings.poll_interval_seconds)
|
|
LOG.info("strategy loop stopped")
|