diff --git a/docs/decisions.md b/docs/decisions.md index 75dc030..a6919bc 100644 --- a/docs/decisions.md +++ b/docs/decisions.md @@ -22,6 +22,7 @@ | `up` / `tunnel` | `up` по умолчанию после ready открывает туннель `:17801`, печатает URL и ждёт. `--no-tunnel` — только облако. Ctrl+C на туннеле GPU не гасит (`stop` отдельно). Команда `tunnel` остаётся для повторного входа | | Git update | На каждом `up` по умолчанию: `git pull` SwarmUI + репы из `extensions.yaml` + уже установленные на data (`Extensions`/`DLNodes`). `--no-update` или `UPDATE_GIT=false` — не тянуть | | LLM (opt-in) | `none` по умолчанию. `ollama` / `llamacpp` — флаг `--ollama`/`--llamacpp`/`--llm`, `LLM_RUNTIME` в vars, или вопрос в interactive `up`/`setup`. Ollama-модели из `ollama-models.yaml`. Порты: Ollama 17811, llama.cpp 17812. Назначение: помощь с промптами, не замена SwarmUI | +| Perf auto-tune | На `up`: probe GPU → tier. Swarm/Comfy: sageattention ExtraArgs на Ampere+ ≥16 GiB. Ollama: flash/KV/keep-alive + `GPU_OVERHEAD` чтобы оставить VRAM под Krea | | Data-диск | Старт **100 GB**, рост через resize вверх (вниз Selectel не умеет) | | SSH | CLI генерирует `/.gpu-rent/id_ed25519` без passphrase и сам регистрирует keypair | | Локальные файлы | Всё в корне репозитория: `.env`, `models.yaml`, `extensions.yaml`; runtime (`state.json`, lock, SSH) в `/.gpu-rent/`. Не `%USERPROFILE%\.gpu-rent` | diff --git a/docs/llm.md b/docs/llm.md index 561a564..00e2db2 100644 --- a/docs/llm.md +++ b/docs/llm.md @@ -84,6 +84,23 @@ Community abliterate-модели без гарантий безопасност --- +## Автотюнинг Ollama под GPU + +На установке `gpu-rent-ollama` читает `/mnt/swarm_data/.gpu-rent-gpu.json` и пишет env в systemd unit (одна карта вместе со SwarmUI): + +| Tier (VRAM) | Flash Attn | KEEP_ALIVE | KV cache | GPU_OVERHEAD (запас под Swarm/Krea) | +| --- | --- | --- | --- | --- | +| low (<16 GiB) | off | 2m | q4_0 | 6 GiB | +| mid (16–23) | on* | 5m | q8_0 | 10 GiB | +| high (24–47) | on* | 15m | q8_0 | 14 GiB | +| ultra (≥48) | on* | 30m | q8_0 | 20 GiB | + +\*Flash на Ampere+ (compute ≥ 8.0). Всегда `NUM_PARALLEL=1`, `MAX_LOADED_MODELS=1`. + +Файл: `/mnt/swarm_data/.gpu-rent-ollama.env` (пересоздаётся на каждом install Ollama). + +--- + ## llama.cpp 1. `LLM_RUNTIME=llamacpp` или `up --llamacpp` @@ -93,6 +110,8 @@ Community abliterate-модели без гарантий безопасност Без GGUF unit может стартовать, но API бесполезен — смотри `gpu-rent logs` / `journalctl -u gpu-rent-llamacpp`. +Параметры `-ngl` / `-c` ставятся по тому же GPU probe (full offload на mid+, меньше слоёв и ctx на low). + --- ## Idle-killer и LLM diff --git a/docs/swarmui.md b/docs/swarmui.md index 4731845..bdfb9af 100644 --- a/docs/swarmui.md +++ b/docs/swarmui.md @@ -20,9 +20,11 @@ 6. Если есть Civitai-токен и манифест моделей: seed весов **до** первого старта SwarmUI; иначе — дефолтная модель установщика. 7. Push непустых `Models/` / `Wildcards/` / `CustomWorkflows/`. 8. systemd unit `swarmui`: `./launch-linux.sh --launch_mode none --host 127.0.0.1 --port 7801`. -9. Авторизация SwarmUI включена, токен в `Data` на диске. -10. systemd unit idle-killer **после** seed и backend Idle. -11. Один snapshot boot volume `gpu-rent-boot-ok`, если ещё нет. +9. **GPU probe** → `/mnt/swarm_data/.gpu-rent-gpu.json` (VRAM / compute cap / tier) — до старта UI; Ollama/llama.cpp читают его при install. +10. Старт SwarmUI → seed LLM → idle-killer → **wait backend Idle**. +11. **Perf tune после Idle** — `triton`+`sageattention` в Comfy venv и `--use-sage-attention` в `Data/Backends.fds` (маркер `.gpu-rent-perf-tuned`; повтор при смене GPU). +12. Авторизация SwarmUI включена, токен в `Data` на диске. +13. Один snapshot boot volume `gpu-rent-boot-ok`, если ещё нет. Рестарт UI: `systemctl restart swarmui`, не `docker restart`. @@ -30,6 +32,8 @@ Первый запуск качает backend в `/opt/swarmui/dlbackend` (это bind на data volume). Иначе каждый recreate потеряет часы. +**Скорость без потери качества (gpu-rent):** после Idle backend, на GPU с compute capability ≥ 8.0 и ≥16 GiB VRAM — SageAttention. `Performance.AllowGpuSpecificOptimizations` у Swarm по умолчанию уже включает `--fast` для 30xx+. Если venv ещё не появился — маркер без `pip_ok`, догонит на следующем `up`. + Пока CUDA/ComfyUI поднимаются, UI уже может отвечать. Для MCP и `/API/` — рано: `ready` после Idle backend (подтвердить на spike). ## Доступ с ноутбука diff --git a/src/gpu_rent/capture.py b/src/gpu_rent/capture.py index a39539b..4122c4a 100644 --- a/src/gpu_rent/capture.py +++ b/src/gpu_rent/capture.py @@ -15,7 +15,9 @@ import yaml from gpu_rent.civitai import ( civitai_model_url, + fetch_model_version, fetch_model_version_by_hash, + fetch_model_versions_by_hashes, version_ids_from_payload, ) from gpu_rent.config import Config @@ -62,12 +64,23 @@ class ExtCaptureItem: directory: str +@dataclass +class ResolveOutcome: + """ok → item; unknown → not on Civitai; api_error → network/HTTP (retry later).""" + + item: ModelCaptureItem | None = None + status: str = "unknown" # ok | unknown | api_error + detail: str = "" + + @dataclass class CaptureReport: models_new: list[ModelCaptureItem] = field(default_factory=list) models_skip: list[str] = field(default_factory=list) models_unknown: list[str] = field(default_factory=list) + models_api_errors: list[str] = field(default_factory=list) ext_new: list[ExtCaptureItem] = field(default_factory=list) + ext_updated: list[ExtCaptureItem] = field(default_factory=list) ext_skip: list[str] = field(default_factory=list) ext_unknown: list[str] = field(default_factory=list) models_path: Path | None = None @@ -102,10 +115,10 @@ def resolve_model_item( token: str, api_host: str, link_host: str, -) -> ModelCaptureItem | None: +) -> ResolveOutcome: kind = str(raw.get("kind") or "") if kind not in MODEL_TYPES: - return None + return ResolveOutcome(status="unknown", detail="bad kind") rel = str(raw.get("rel") or raw.get("name") or "?") version_id = raw.get("version_id") model_id = raw.get("model_id") @@ -120,35 +133,74 @@ def resolve_model_item( except (TypeError, ValueError): mid = None + def _ok(v: int, m: int, name: str = title) -> ResolveOutcome: + return ResolveOutcome( + item=ModelCaptureItem( + kind=kind, + version_id=v, + model_id=m, + url=civitai_model_url(m, v, link_host), + title=name, + rel=rel, + ), + status="ok", + ) + if vid is not None and mid is not None: - return ModelCaptureItem( - kind=kind, - version_id=vid, - model_id=mid, - url=civitai_model_url(mid, vid, link_host), - title=title, - rel=rel, + return _ok(vid, mid) + + # Partial sidecar: have version id → GET /model-versions/{id} for modelId. + if vid is not None and mid is None: + try: + _h, version = fetch_model_version(token or "", api_host, vid) + except CloudError as exc: + msg = str(exc) + if "404" in msg or "пустой files" in msg: + return ResolveOutcome( + status="unknown", + detail=f"{rel} version_id={vid} ({msg})", + ) + return ResolveOutcome( + status="api_error", + detail=f"{rel} version_id={vid} {msg}", + ) + vid2, mid2 = version_ids_from_payload(version) + if vid2 is not None and mid2 is not None: + name = str(version.get("name") or title) + return _ok(vid2, mid2, name) + return ResolveOutcome( + status="unknown", + detail=f"{rel} version_id={vid} (нет modelId в ответе)", ) sha = raw.get("sha256") if not sha: - return None + return ResolveOutcome( + status="unknown", + detail=f"{rel} (нет sidecar и нет sha256)", + ) try: _host, version = fetch_model_version_by_hash(token or None, api_host, str(sha)) - except CloudError: - return None + except CloudError as exc: + msg = str(exc) + # 404 = genuinely not on Civitai; other → api_error + if "HTTP 404" in msg or msg.rstrip().endswith("404"): + return ResolveOutcome( + status="unknown", + detail=f"{rel} sha={str(sha)[:12]}…", + ) + return ResolveOutcome( + status="api_error", + detail=f"{rel} sha={str(sha)[:12]}… {msg}", + ) vid2, mid2 = version_ids_from_payload(version) if vid2 is None or mid2 is None: - return None - name = version.get("name") or title - return ModelCaptureItem( - kind=kind, - version_id=vid2, - model_id=mid2, - url=civitai_model_url(mid2, vid2, link_host), - title=str(name), - rel=rel, - ) + return ResolveOutcome( + status="unknown", + detail=f"{rel} sha={str(sha)[:12]}… (пустой payload)", + ) + name = str(version.get("name") or title) + return _ok(vid2, mid2, name) def _backup(path: Path) -> None: @@ -156,37 +208,46 @@ def _backup(path: Path) -> None: shutil.copy2(path, path.with_suffix(path.suffix + ".bak")) +def _keep_model_entry(it: dict) -> bool: + if not isinstance(it, dict): + return False + url = it.get("url") + if url: + return True + vid = it.get("version_id") + return vid not in (None, "", 0, "0") + + def merge_models_yaml( path: Path, new_items: list[ModelCaptureItem], *, dry_run: bool, ) -> tuple[list[ModelCaptureItem], list[str]]: - """Return (actually_new, skip_msgs). Writes path unless dry_run.""" + """Return (actually_new, skip_msgs). Dedupe by (kind, version_id).""" existing = parse_models(path) if path.is_file() else [] - have: set[int] = set() + have: set[tuple[str, int]] = set() for e in existing: - if e.version_id is not None: - have.add(e.version_id) - elif e.url: + vid = e.version_id + if vid is None and e.url: vid = extract_version_id(e.url) - if vid is not None: - have.add(vid) + if vid is not None: + have.add((e.kind, vid)) added: list[ModelCaptureItem] = [] skipped: list[str] = [] - seen_new: set[int] = set() + seen_new: set[tuple[str, int]] = set() for item in new_items: - if item.version_id in have or item.version_id in seen_new: + key = (item.kind, item.version_id) + if key in have or key in seen_new: skipped.append(f"{item.kind} {item.title} modelVersionId={item.version_id}") continue - seen_new.add(item.version_id) + seen_new.add(key) added.append(item) if dry_run or not added: return added, skipped - # Rebuild full mapping: keep existing entries, append new urls. data: dict[str, list[dict[str, str]]] = {k: [] for k in MODEL_TYPES} if path.is_file(): raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {} @@ -195,13 +256,12 @@ def merge_models_yaml( items = raw.get(kind) or [] if isinstance(items, list): for it in items: - if isinstance(it, dict) and (it.get("url") or it.get("version_id") not in (None, 0, "0")): + if _keep_model_entry(it): data[kind].append(dict(it)) for item in added: data[item.kind].append({"url": item.url}) - # Drop empty kinds for cleaner file out = {k: v for k, v in data.items() if v} _backup(path) path.parent.mkdir(parents=True, exist_ok=True) @@ -217,32 +277,9 @@ def merge_extensions_yaml( new_items: list[ExtCaptureItem], *, dry_run: bool, -) -> tuple[list[ExtCaptureItem], list[str]]: - existing = parse_extensions(path) if path.is_file() else [] - have_urls: set[tuple[str, str]] = set() - have_dirs: set[tuple[str, str]] = set() - for e in existing: - have_urls.add((e.kind, strip_git_auth(e.url).rstrip("/").lower())) - dirname = e.directory or "" - if dirname: - have_dirs.add((e.kind, dirname.lower())) - - added: list[ExtCaptureItem] = [] - skipped: list[str] = [] - for item in new_items: - key_url = (item.kind, strip_git_auth(item.url).rstrip("/").lower()) - key_dir = (item.kind, item.directory.lower()) - if key_url in have_urls or key_dir in have_dirs: - skipped.append(f"{item.kind} {item.directory} {item.url}") - continue - have_urls.add(key_url) - have_dirs.add(key_dir) - added.append(item) - - if dry_run or not added: - return added, skipped - - data: dict[str, list[dict[str, str]]] = {"swarmui": [], "comfy": []} +) -> tuple[list[ExtCaptureItem], list[ExtCaptureItem], list[str]]: + """Return (added, updated, skipped). Same dir + different URL → update url/ref.""" + data: dict[str, list[dict[str, Any]]] = {"swarmui": [], "comfy": []} if path.is_file(): raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {} if isinstance(raw, dict): @@ -253,13 +290,59 @@ def merge_extensions_yaml( if isinstance(it, dict) and it.get("url"): data[kind].append(dict(it)) - for item in added: - entry: dict[str, str] = { - "url": strip_git_auth(item.url), - "ref": item.ref, - "dir": item.directory, - } - data[item.kind].append(entry) + def _dir_of(it: dict) -> str: + d = it.get("dir") + if d: + return str(d).lower() + url = strip_git_auth(str(it.get("url") or "")) + name = url.rstrip("/").rsplit("/", 1)[-1] + if name.endswith(".git"): + name = name[:-4] + return name.lower() + + added: list[ExtCaptureItem] = [] + updated: list[ExtCaptureItem] = [] + skipped: list[str] = [] + + for item in new_items: + clean = strip_git_auth(item.url) + key_url = (item.kind, clean.rstrip("/").lower()) + key_dir = (item.kind, item.directory.lower()) + bucket = data.setdefault(item.kind, []) + + matched_url = False + matched_dir_idx: int | None = None + for idx, it in enumerate(bucket): + it_url = strip_git_auth(str(it.get("url") or "")).rstrip("/").lower() + if (item.kind, it_url) == key_url: + matched_url = True + break + if (item.kind, _dir_of(it)) == key_dir: + matched_dir_idx = idx + + if matched_url: + skipped.append(f"{item.kind} {item.directory} {clean}") + continue + + if matched_dir_idx is not None: + old = bucket[matched_dir_idx] + old_url = strip_git_auth(str(old.get("url") or "")) + old_ref = str(old.get("ref") or "main") + if old_url.rstrip("/").lower() == clean.rstrip("/").lower() and old_ref == item.ref: + skipped.append(f"{item.kind} {item.directory} {clean}") + continue + old["url"] = clean + old["ref"] = item.ref + if not old.get("dir"): + old["dir"] = item.directory + updated.append(item) + continue + + bucket.append({"url": clean, "ref": item.ref, "dir": item.directory}) + added.append(item) + + if dry_run or (not added and not updated): + return added, updated, skipped out = {k: v for k, v in data.items() if v} _backup(path) @@ -268,7 +351,7 @@ def merge_extensions_yaml( yaml.safe_dump(out, allow_unicode=True, default_flow_style=False, sort_keys=False), encoding="utf-8", ) - return added, skipped + return added, updated, skipped def capture_models( @@ -287,24 +370,91 @@ def capture_models( raise CloudError("inventory.models: ожидался list") resolved: list[ModelCaptureItem] = [] + need_hash: list[dict] = [] link_host = cfg.civitai_api_host or "civitai.red" + token = cfg.civitai_api_token + api_host = cfg.civitai_api_host + for raw in raw_models: if not isinstance(raw, dict): continue if kind_filter and str(raw.get("kind")) != kind_filter: continue - item = resolve_model_item( - raw, - token=cfg.civitai_api_token, - api_host=cfg.civitai_api_host, - link_host=link_host, - ) - if item is None: - rel = str(raw.get("rel") or raw.get("name") or "?") - sha = str(raw.get("sha256") or "")[:12] - report.models_unknown.append(f"{rel} sha={sha}…") + vid = raw.get("version_id") + mid = raw.get("model_id") + try: + vid_i = int(vid) if vid is not None else None + except (TypeError, ValueError): + vid_i = None + try: + mid_i = int(mid) if mid is not None else None + except (TypeError, ValueError): + mid_i = None + sha = raw.get("sha256") + + # Full sidecar / ids → no API. Partial vid → GET version. Else batch by-hash. + if vid_i is not None and mid_i is not None: + outcome = resolve_model_item( + raw, token=token, api_host=api_host, link_host=link_host + ) + elif vid_i is not None: + outcome = resolve_model_item( + raw, token=token, api_host=api_host, link_host=link_host + ) + elif sha: + need_hash.append(raw) continue - resolved.append(item) + else: + outcome = resolve_model_item( + raw, token=token, api_host=api_host, link_host=link_host + ) + + if outcome.status == "ok" and outcome.item is not None: + resolved.append(outcome.item) + elif outcome.status == "api_error": + report.models_api_errors.append(outcome.detail or "?") + else: + report.models_unknown.append(outcome.detail or str(raw.get("rel") or "?")) + + if need_hash: + digests = [str(r["sha256"]).strip().lower() for r in need_hash if r.get("sha256")] + log(f"capture: by-hash batch {len(digests)} файл(ов)…") + try: + by_hash = fetch_model_versions_by_hashes(token or None, api_host, digests) + except CloudError as exc: + report.models_api_errors.append(f"by-hash batch: {exc}") + by_hash = {} + for raw in need_hash: + rel = str(raw.get("rel") or "?") + sha = str(raw.get("sha256") or "")[:12] + report.models_api_errors.append(f"{rel} sha={sha}… (batch failed)") + need_hash = [] + + for raw in need_hash: + sha = str(raw.get("sha256") or "").strip().lower() + rel = str(raw.get("rel") or raw.get("name") or "?") + version = by_hash.get(sha) + if not version: + report.models_unknown.append(f"{rel} sha={sha[:12]}…") + continue + vid2, mid2 = version_ids_from_payload(version) + if vid2 is None or mid2 is None: + report.models_unknown.append(f"{rel} sha={sha[:12]}… (пустой payload)") + continue + kind = str(raw.get("kind") or "") + if kind not in MODEL_TYPES: + continue + title = str(version.get("name") or Path(str(raw.get("name") or rel)).stem) + resolved.append( + ModelCaptureItem( + kind=kind, + version_id=vid2, + model_id=mid2, + url=civitai_model_url(mid2, vid2, link_host), + title=title, + rel=rel, + ) + ) added, skipped = merge_models_yaml( cfg.models_manifest, resolved, dry_run=dry_run @@ -350,12 +500,13 @@ def capture_extensions( ) ) - added, skipped = merge_extensions_yaml( + added, updated, skipped = merge_extensions_yaml( cfg.extensions_manifest, resolved, dry_run=dry_run ) report.ext_new = added + report.ext_updated = updated report.ext_skip = skipped - report.wrote_extensions = bool(added) and not dry_run + report.wrote_extensions = (bool(added) or bool(updated)) and not dry_run return report @@ -364,7 +515,9 @@ def merge_reports(a: CaptureReport, b: CaptureReport) -> CaptureReport: models_new=a.models_new + b.models_new, models_skip=a.models_skip + b.models_skip, models_unknown=a.models_unknown + b.models_unknown, + models_api_errors=a.models_api_errors + b.models_api_errors, ext_new=a.ext_new + b.ext_new, + ext_updated=a.ext_updated + b.ext_updated, ext_skip=a.ext_skip + b.ext_skip, ext_unknown=a.ext_unknown + b.ext_unknown, models_path=a.models_path or b.models_path, @@ -390,27 +543,42 @@ def capture_all( return merge_reports(m, e) -def print_report(report: CaptureReport, log: Log, *, dry_run: bool) -> None: +def print_report( + report: CaptureReport, + log: Log, + *, + dry_run: bool, + show_models: bool = True, + show_extensions: bool = True, +) -> None: prefix = "[dry-run] " if dry_run else "" - log( - f"{prefix}models: +{len(report.models_new)} new, " - f"{len(report.models_skip)} already in yaml, " - f"{len(report.models_unknown)} unknown (no Civitai)" - ) - for item in report.models_new: - log(f" + {item.kind} {item.title} modelVersionId={item.version_id}") - for line in report.models_unknown: - log(f" ? {line}") + if show_models: + log( + f"{prefix}models: +{len(report.models_new)} new, " + f"{len(report.models_skip)} already in yaml, " + f"{len(report.models_unknown)} unknown (no Civitai), " + f"{len(report.models_api_errors)} api errors" + ) + for item in report.models_new: + log(f" + {item.kind} {item.title} modelVersionId={item.version_id}") + for line in report.models_unknown: + log(f" ? {line}") + for line in report.models_api_errors: + log(f" ! {line}") - log( - f"{prefix}extensions: +{len(report.ext_new)} new, " - f"{len(report.ext_skip)} already, " - f"{len(report.ext_unknown)} unknown" - ) - for item in report.ext_new: - log(f" + {item.kind} {item.directory} {item.url} @{item.ref}") - for line in report.ext_unknown: - log(f" ? {line}") + if show_extensions: + log( + f"{prefix}extensions: +{len(report.ext_new)} new, " + f"~{len(report.ext_updated)} updated, " + f"{len(report.ext_skip)} already, " + f"{len(report.ext_unknown)} unknown" + ) + for item in report.ext_new: + log(f" + {item.kind} {item.directory} {item.url} @{item.ref}") + for item in report.ext_updated: + log(f" ~ {item.kind} {item.directory} {item.url} @{item.ref}") + for line in report.ext_unknown: + log(f" ? {line}") if dry_run: log("dry-run — файлы не записаны") @@ -423,7 +591,4 @@ def print_report(report: CaptureReport, log: Log, *, dry_run: bool) -> None: f"(backup {report.extensions_path.name}.bak)" ) if not report.wrote_models and not report.wrote_extensions: - if report.models_new or report.ext_new: - pass - else: - log("нечего добавлять") + log("нечего добавлять") diff --git a/src/gpu_rent/civitai.py b/src/gpu_rent/civitai.py index f530a83..a817bc9 100644 --- a/src/gpu_rent/civitai.py +++ b/src/gpu_rent/civitai.py @@ -2,6 +2,7 @@ from __future__ import annotations +import time from dataclasses import dataclass import httpx @@ -147,6 +148,133 @@ def fetch_model_version_by_hash( raise CloudError(f"Civitai by-hash {digest[:12]}…: {last_error} (хосты {', '.join(seen)})") +def _sha_from_version(version: dict) -> list[str]: + """All SHA256 hashes listed on a version's files (lower).""" + out: list[str] = [] + for f in version.get("files") or []: + if not isinstance(f, dict): + continue + hashes = f.get("hashes") or {} + if not isinstance(hashes, dict): + continue + sha = hashes.get("SHA256") or hashes.get("sha256") + if sha: + out.append(str(sha).strip().lower()) + return out + + +def fetch_model_versions_by_hashes( + token: str | None, + host: str, + sha256_list: list[str], + *, + timeout: float = 60.0, + chunk_size: int = 100, +) -> dict[str, dict]: + """POST /api/v1/model-versions/by-hash (≤100). Map sha256(lower) → version. + + Unmatched hashes are omitted (caller treats as unknown). Retries once on 429. + If response versions lack files[].hashes, falls back to /by-hash/ids + GET version. + """ + cleaned: list[str] = [] + seen_h: set[str] = set() + for raw in sha256_list: + digest = str(raw).strip().lower() + if len(digest) == 64 and all(c in "0123456789abcdef" for c in digest): + if digest not in seen_h: + cleaned.append(digest) + seen_h.add(digest) + if not cleaned: + return {} + + first = _normalize_host(host) + order = [first, other_host(first)] + headers: dict[str, str] = {"Content-Type": "application/json"} + if token: + headers["Authorization"] = f"Bearer {token}" + result: dict[str, dict] = {} + + def _post(candidate: str, path: str, body: list[str]) -> httpx.Response | None: + url = f"https://{candidate}/api/v1/{path}" + last: httpx.Response | None = None + for attempt in range(2): + try: + with httpx.Client(timeout=timeout, follow_redirects=True) as client: + last = client.post(url, headers=headers, json=body) + except httpx.HTTPError: + return None + if last.status_code == 429 and attempt == 0: + time.sleep(2.0) + continue + return last + return last + + for i in range(0, len(cleaned), chunk_size): + chunk = cleaned[i : i + chunk_size] + last_error = "нет ответа" + ok = False + for candidate in order: + if candidate not in ALLOWED_HOSTS: + continue + response = _post(candidate, "model-versions/by-hash", chunk) + if response is None: + last_error = "network" + continue + if response.status_code == 200: + data = response.json() + if isinstance(data, list): + for version in data: + if not isinstance(version, dict) or version.get("id") is None: + continue + for sha in _sha_from_version(version): + if sha in seen_h: + result.setdefault(sha, version) + ok = True + break + last_error = "не list" + continue + last_error = f"HTTP {response.status_code}" + if response.status_code not in {404, 400}: + continue + if not ok: + raise CloudError( + f"Civitai by-hash batch ({len(chunk)}): {last_error} " + f"(хосты {', '.join(order)})" + ) + + missing = [h for h in chunk if h not in result] + if not missing: + continue + # Fallback: ids endpoint maps hash → versionId, then GET version for modelId. + for candidate in order: + if candidate not in ALLOWED_HOSTS: + continue + response = _post(candidate, "model-versions/by-hash/ids", missing) + if response is None or response.status_code != 200: + continue + pairs = response.json() + if not isinstance(pairs, list): + continue + for pair in pairs: + if not isinstance(pair, dict): + continue + h = str(pair.get("hash") or "").strip().lower() + try: + vid = int(pair.get("modelVersionId")) + except (TypeError, ValueError): + continue + if h not in seen_h: + continue + try: + _host, version = fetch_model_version(token or "", candidate, vid) + except CloudError: + continue + result.setdefault(h, version) + break + + return result + + def civitai_model_url(model_id: int, version_id: int, host: str = "civitai.red") -> str: """Canonical manifest URL (links only — no download).""" h = _normalize_host(host) diff --git a/src/gpu_rent/cli.py b/src/gpu_rent/cli.py index f24ab30..9caf6a8 100644 --- a/src/gpu_rent/cli.py +++ b/src/gpu_rent/cli.py @@ -746,7 +746,13 @@ def capture_models_cmd( kind_filter=kind, log=lambda m: console.print(m), ) - print_report(report, lambda m: console.print(m), dry_run=dry_run) + print_report( + report, + lambda m: console.print(m), + dry_run=dry_run, + show_models=True, + show_extensions=False, + ) except GpuRentError as exc: _die(exc) @@ -767,7 +773,13 @@ def capture_extensions_cmd( dry_run=dry_run, log=lambda m: console.print(m), ) - print_report(report, lambda m: console.print(m), dry_run=dry_run) + print_report( + report, + lambda m: console.print(m), + dry_run=dry_run, + show_models=False, + show_extensions=True, + ) except GpuRentError as exc: _die(exc) diff --git a/src/gpu_rent/perf_tiers.py b/src/gpu_rent/perf_tiers.py new file mode 100644 index 0000000..02afeaa --- /dev/null +++ b/src/gpu_rent/perf_tiers.py @@ -0,0 +1,147 @@ +"""GPU performance tiers for SwarmUI + Ollama auto-tune.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + + +# Shared GPU with SwarmUI: reserve VRAM so image gen still fits (Krea2 Turbo fp8 ~10–14GB). +TIER_LOW = "low" # <16 GiB +TIER_MID = "mid" # 16–23 +TIER_HIGH = "high" # 24–47 +TIER_ULTRA = "ultra" # ≥48 + + +@dataclass(frozen=True) +class GpuInfo: + name: str + vram_mib: int + compute_cap: str # e.g. "8.9" + uuid: str + tier: str + + +@dataclass(frozen=True) +class OllamaTune: + flash_attention: bool + keep_alive: str + num_parallel: int + max_loaded_models: int + kv_cache_type: str | None + gpu_overhead_bytes: int + notes: str + + +@dataclass(frozen=True) +class SwarmTune: + use_sage_attention: bool + install_triton_sage: bool + comfy_extra_args: str + notes: str + + +def tier_for_vram_mib(vram_mib: int) -> str: + gib = vram_mib / 1024.0 + if gib < 16: + return TIER_LOW + if gib < 24: + return TIER_MID + if gib < 48: + return TIER_HIGH + return TIER_ULTRA + + +def compute_cap_at_least(cap: str, major: int, minor: int = 0) -> bool: + """True if NVIDIA compute capability >= major.minor (Ampere=8.0).""" + try: + parts = str(cap).strip().split(".") + maj = int(parts[0]) + mnr = int(parts[1]) if len(parts) > 1 else 0 + return (maj, mnr) >= (major, minor) + except (TypeError, ValueError): + return False + + +def ollama_tune_for(info: GpuInfo) -> OllamaTune: + """Tune Ollama for prompt-help beside SwarmUI (share one GPU).""" + ampere_plus = compute_cap_at_least(info.compute_cap, 8, 0) + flash = ampere_plus or "A100" in info.name.upper() or "H100" in info.name.upper() + # Reserve VRAM for Comfy/Krea so Ollama does not fill the card. + if info.tier == TIER_ULTRA: + overhead = 20 * 1024**3 + keep = "30m" + kv = "q8_0" + note = "ultra: flash+q8 KV, 20GiB reserved for Swarm, keep 30m" + elif info.tier == TIER_HIGH: + overhead = 14 * 1024**3 + keep = "15m" + kv = "q8_0" + note = "high: flash+q8 KV, 14GiB reserved for Swarm, keep 15m" + elif info.tier == TIER_MID: + overhead = 10 * 1024**3 + keep = "5m" + kv = "q8_0" + note = "mid: flash+q8 KV, 10GiB reserved for Swarm, keep 5m" + else: + overhead = 6 * 1024**3 + keep = "2m" + kv = "q4_0" + flash = False # prefer stability on tiny cards + note = "low: conservative, 6GiB reserved, short keep-alive" + + return OllamaTune( + flash_attention=flash, + keep_alive=keep, + num_parallel=1, + max_loaded_models=1, + kv_cache_type=kv, + gpu_overhead_bytes=overhead, + notes=note, + ) + + +def swarm_tune_for(info: GpuInfo) -> SwarmTune: + """Swarm/Comfy launch: speed without quality loss (sage on Ampere+).""" + ampere_plus = compute_cap_at_least(info.compute_cap, 8, 0) + # SageAttention: Linux + Triton; safe quality, faster attention (Swarm docs). + use_sage = ampere_plus and info.tier in {TIER_MID, TIER_HIGH, TIER_ULTRA} + if info.tier == TIER_LOW: + use_sage = False + args = "--use-sage-attention" if use_sage else "" + return SwarmTune( + use_sage_attention=use_sage, + install_triton_sage=use_sage, + comfy_extra_args=args, + notes=( + "sage-attention + triton (Ampere+)" + if use_sage + else "stock Comfy attention (low VRAM or pre-Ampere)" + ), + ) + + +def parse_probe_dict(data: dict[str, Any]) -> GpuInfo: + vram = int(data.get("vram_mib") or 0) + tier = str(data.get("tier") or tier_for_vram_mib(vram)) + return GpuInfo( + name=str(data.get("name") or "unknown"), + vram_mib=vram, + compute_cap=str(data.get("compute_cap") or "0.0"), + uuid=str(data.get("uuid") or ""), + tier=tier, + ) + + +def ollama_env_lines(tune: OllamaTune) -> list[str]: + lines = [ + f"Environment=OLLAMA_NUM_PARALLEL={tune.num_parallel}", + f"Environment=OLLAMA_MAX_LOADED_MODELS={tune.max_loaded_models}", + f"Environment=OLLAMA_KEEP_ALIVE={tune.keep_alive}", + f"Environment=OLLAMA_GPU_OVERHEAD={tune.gpu_overhead_bytes}", + ] + if tune.flash_attention: + lines.append("Environment=OLLAMA_FLASH_ATTENTION=1") + if tune.kv_cache_type: + lines.append(f"Environment=OLLAMA_KV_CACHE_TYPE={tune.kv_cache_type}") + return lines diff --git a/src/gpu_rent/provision.py b/src/gpu_rent/provision.py index 44b1705..fa0f191 100644 --- a/src/gpu_rent/provision.py +++ b/src/gpu_rent/provision.py @@ -35,6 +35,50 @@ def _pkg_text(name: str) -> str: return files("gpu_rent.remote").joinpath(name).read_text(encoding="utf-8") +def probe_gpu(cfg: Config, host: str, log: Log) -> dict: + """Write /mnt/swarm_data/.gpu-rent-gpu.json; return parsed dict.""" + out = run_python( + cfg, + host, + _pkg_text("gpu_probe.py"), + remote_path="/tmp/gpu-rent-gpu_probe.py", + timeout=60, + log=log, + ) + # Last JSON line from script stdout + data: dict = {} + for line in reversed(out.splitlines()): + line = line.strip() + if line.startswith("{"): + try: + data = json.loads(line) + break + except json.JSONDecodeError: + continue + if data.get("ok"): + log( + f"GPU: {data.get('name')} " + f"{data.get('vram_mib')} MiB cap={data.get('compute_cap')} " + f"tier={data.get('tier')}" + ) + else: + log(f"GPU probe: {data.get('error') or 'нет данных'}") + return data + + +def tune_swarm_perf(cfg: Config, host: str, log: Log) -> bool: + """Install sage/triton into Comfy venv + ExtraArgs. Returns True if SwarmUI restart needed.""" + out = run_python( + cfg, + host, + _pkg_text("tune_swarm_perf.py"), + remote_path="/tmp/gpu-rent-tune_swarm_perf.py", + timeout=900, + log=log, + ) + return "RESTART_SWARMUI=1" in out + + def seed_extensions(cfg: Config, host: str, log: Log, *, update: bool = True) -> bool: from gpu_rent.llm_runtime import normalize_runtime @@ -372,6 +416,12 @@ def provision_vm( push_tree(cfg, host, cfg.local_workflows_dir, f"{DATA}/CustomWorkflows", log, models=False) if cfg.pull_output: pull_tree(cfg, host, f"{DATA}/Output", cfg.local_output_dir, log) + + try: + probe_gpu(cfg, host, log) + except Exception as exc: + log(f"GPU probe: {exc}") + ensure_swarmui_running(cfg, host, log, restart=restart) from gpu_rent.state import load_state, save_state diff --git a/src/gpu_rent/remote/gpu_probe.py b/src/gpu_rent/remote/gpu_probe.py new file mode 100644 index 0000000..1ae6219 --- /dev/null +++ b/src/gpu_rent/remote/gpu_probe.py @@ -0,0 +1,97 @@ +#!/usr/bin/env python3 +"""Probe NVIDIA GPU → /mnt/swarm_data/.gpu-rent-gpu.json (stdlib only).""" + +from __future__ import annotations + +import json +import re +import subprocess +import sys +from pathlib import Path + +OUT = Path("/mnt/swarm_data/.gpu-rent-gpu.json") + + +def tier_for_vram_mib(vram_mib: int) -> str: + gib = vram_mib / 1024.0 + if gib < 16: + return "low" + if gib < 24: + return "mid" + if gib < 48: + return "high" + return "ultra" + + +def run(argv: list[str]) -> str: + try: + return subprocess.check_output(argv, text=True, stderr=subprocess.DEVNULL).strip() + except (OSError, subprocess.CalledProcessError): + return "" + + +def main() -> int: + if not run(["bash", "-lc", "command -v nvidia-smi"]): + data = { + "ok": False, + "error": "nvidia-smi missing", + "name": "", + "vram_mib": 0, + "compute_cap": "0.0", + "uuid": "", + "tier": "low", + } + OUT.parent.mkdir(parents=True, exist_ok=True) + OUT.write_text(json.dumps(data, indent=2) + "\n", encoding="utf-8") + print(json.dumps(data)) + return 0 + + # name, uuid, memory.total [MiB], compute_cap + q = run( + [ + "nvidia-smi", + "--query-gpu=name,uuid,memory.total,compute_cap", + "--format=csv,noheader,nounits", + ] + ) + if not q: + data = { + "ok": False, + "error": "nvidia-smi query failed", + "name": "", + "vram_mib": 0, + "compute_cap": "0.0", + "uuid": "", + "tier": "low", + } + OUT.parent.mkdir(parents=True, exist_ok=True) + OUT.write_text(json.dumps(data, indent=2) + "\n", encoding="utf-8") + print(json.dumps(data)) + return 0 + + line = q.splitlines()[0] + # CSV: "NVIDIA A100-SXM4-40GB, GPU-..., 40960, 8.0" + parts = [p.strip() for p in line.split(",")] + name = parts[0] if parts else "unknown" + uuid = parts[1] if len(parts) > 1 else "" + vram_raw = parts[2] if len(parts) > 2 else "0" + cap = parts[3] if len(parts) > 3 else "0.0" + m = re.search(r"(\d+)", vram_raw.replace(" ", "")) + vram_mib = int(m.group(1)) if m else 0 + tier = tier_for_vram_mib(vram_mib) + data = { + "ok": True, + "name": name, + "vram_mib": vram_mib, + "compute_cap": cap, + "uuid": uuid, + "tier": tier, + } + OUT.parent.mkdir(parents=True, exist_ok=True) + OUT.write_text(json.dumps(data, indent=2) + "\n", encoding="utf-8") + print(json.dumps(data)) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/src/gpu_rent/remote/install_llamacpp.sh b/src/gpu_rent/remote/install_llamacpp.sh index 6c4fb01..155c2f3 100644 --- a/src/gpu_rent/remote/install_llamacpp.sh +++ b/src/gpu_rent/remote/install_llamacpp.sh @@ -110,9 +110,32 @@ else log "нет GGUF в ${MODELS_DIR} — положи файл вручную и systemctl restart ${UNIT}" fi +# GPU layers: share card with Swarm — full offload on mid+, leave headroom on low. +NGL=99 +CTX=8192 +if [[ -f "${DATA_ROOT}/.gpu-rent-gpu.json" ]]; then + eval "$(python3 - <<'PY' +import json +from pathlib import Path +gpu=json.loads(Path("/mnt/swarm_data/.gpu-rent-gpu.json").read_text()) +vram=int(gpu.get("vram_mib") or 0) +gib=vram/1024.0 +if gib < 16: + print("NGL=40"); print("CTX=4096") +elif gib < 24: + print("NGL=99"); print("CTX=8192") +elif gib < 48: + print("NGL=99"); print("CTX=16384") +else: + print("NGL=99"); print("CTX=32768") +PY +)" || true +fi +log "llama.cpp -ngl ${NGL} -c ${CTX}" + cat >/etc/systemd/system/${UNIT}.service < 1 else 0 + ampere = (maj, mnr) >= (8, 0) +except ValueError: + ampere = False +if gib >= 48: + tier, overhead, keep, kv, flash = "ultra", 20 * 1024**3, "30m", "q8_0", True +elif gib >= 24: + tier, overhead, keep, kv, flash = "high", 14 * 1024**3, "15m", "q8_0", True +elif gib >= 16: + tier, overhead, keep, kv, flash = "mid", 10 * 1024**3, "5m", "q8_0", True +else: + tier, overhead, keep, kv, flash = "low", 6 * 1024**3, "2m", "q4_0", False +flash = bool(flash and (ampere or "A100" in name.upper() or "H100" in name.upper() or gib >= 16)) +lines = [ + f"# auto gpu-rent ollama tune tier={tier} gpu={name!r} vram_mib={vram}", + "OLLAMA_NUM_PARALLEL=1", + "OLLAMA_MAX_LOADED_MODELS=1", + f"OLLAMA_KEEP_ALIVE={keep}", + f"OLLAMA_GPU_OVERHEAD={overhead}", +] +if flash: + lines.append("OLLAMA_FLASH_ATTENTION=1") +if kv: + lines.append(f"OLLAMA_KV_CACHE_TYPE={kv}") +out.write_text("\n".join(lines) + "\n", encoding="utf-8") +print(f"tier={tier} flash={int(flash)} keep={keep} overhead_gib={overhead/1024**3:.0f}") +PY +} + +# Prefer probe file written by gpu_probe.py; if missing, try nvidia-smi quickly. +if [[ ! -f "$GPU_JSON" ]] && command -v nvidia-smi >/dev/null 2>&1; then + log "нет ${GPU_JSON} — быстрый probe" + python3 - <<'PY' || true +import json, re, subprocess +from pathlib import Path +out = Path("/mnt/swarm_data/.gpu-rent-gpu.json") +try: + q = subprocess.check_output( + ["nvidia-smi", "--query-gpu=name,uuid,memory.total,compute_cap", "--format=csv,noheader,nounits"], + text=True, stderr=subprocess.DEVNULL, + ).strip().splitlines()[0] + parts = [p.strip() for p in q.split(",")] + vram = int(re.search(r"(\d+)", parts[2]).group(1)) + gib = vram / 1024.0 + tier = "low" if gib < 16 else "mid" if gib < 24 else "high" if gib < 48 else "ultra" + data = {"ok": True, "name": parts[0], "uuid": parts[1], "vram_mib": vram, "compute_cap": parts[3], "tier": tier} +except Exception as e: + data = {"ok": False, "error": str(e), "name": "", "uuid": "", "vram_mib": 0, "compute_cap": "0.0", "tier": "low"} +out.parent.mkdir(parents=True, exist_ok=True) +out.write_text(json.dumps(data, indent=2) + "\n") +print(data) +PY +fi + +write_ollama_env +log "ollama env: $(tr '\n' ' ' < "$OLLAMA_ENV_FILE")" +chown "${SWARM_USER}:${SWARM_USER}" "$OLLAMA_ENV_FILE" "$GPU_JSON" 2>/dev/null || true + if ! command -v ollama >/dev/null 2>&1; then # Supply-chain: prefer a pinned GitHub release. Official install.sh is curl|sh without checksum. # Override: OLLAMA_VERSION=0.6.5 OLLAMA_SHA256= @@ -55,9 +136,16 @@ fi systemctl stop ollama 2>/dev/null || true systemctl disable ollama 2>/dev/null || true +# Build Environment= lines from env file +ENV_LINES="" +while IFS= read -r line || [[ -n "$line" ]]; do + [[ -z "$line" || "$line" =~ ^# ]] && continue + ENV_LINES+="Environment=${line}"$'\n' +done < "$OLLAMA_ENV_FILE" + cat >/etc/systemd/system/${UNIT}.service </dev/null -log "ok — OLLAMA_HOST=127.0.0.1:11434 models=${OLLAMA_HOME}" +log "ok — OLLAMA_HOST=127.0.0.1:11434 models=${OLLAMA_HOME} (GPU-tuned)" diff --git a/src/gpu_rent/remote/scan_inventory.py b/src/gpu_rent/remote/scan_inventory.py index 4811a96..d83929b 100644 --- a/src/gpu_rent/remote/scan_inventory.py +++ b/src/gpu_rent/remote/scan_inventory.py @@ -13,7 +13,8 @@ DATA = Path("/mnt/swarm_data") MODELS = DATA / "Models" OUT = Path("/tmp/gpu-rent-inventory.json") -WEIGHT_SUFFIXES = {".safetensors", ".ckpt", ".pt", ".pth", ".bin", ".gguf", ".sft", ".onnx"} +WEIGHT_SUFFIXES = {".safetensors", ".ckpt", ".gguf", ".sft", ".onnx"} +# .pt/.pth/.bin намеренно вне capture — слишком много ложных «моделей» (torch scripts). # SwarmUI folder name → models.yaml kind FOLDER_TO_KIND = { @@ -51,19 +52,13 @@ def sha256_file(path: Path, chunk: int = 1024 * 1024) -> str: return h.hexdigest() -def read_sidecar_ids(weight: Path) -> tuple[int | None, int | None]: - """Return (version_id, model_id) from {stem}.civitai.json if present.""" - sidecar = weight.parent / f"{weight.stem}.civitai.json" - if not sidecar.is_file(): - return None, None - try: - data = json.loads(sidecar.read_text(encoding="utf-8")) - except (OSError, json.JSONDecodeError): - return None, None - if not isinstance(data, dict): - return None, None +def _ids_from_dict(data: dict) -> tuple[int | None, int | None]: vid = data.get("id") + if vid is None: + vid = data.get("modelVersionId") or data.get("versionId") mid = data.get("modelId") + if mid is None and isinstance(data.get("model"), dict): + mid = data["model"].get("id") try: version_id = int(vid) if vid is not None else None except (TypeError, ValueError): @@ -75,6 +70,28 @@ def read_sidecar_ids(weight: Path) -> tuple[int | None, int | None]: return version_id, model_id +def read_sidecar_ids(weight: Path) -> tuple[int | None, int | None]: + """Return (version_id, model_id) from civitai/swarm sidecars if present.""" + candidates = [ + weight.parent / f"{weight.stem}.civitai.json", + weight.parent / f"{weight.stem}.swarm.json", + weight.parent / f"{weight.stem}.json", + ] + for sidecar in candidates: + if not sidecar.is_file(): + continue + try: + data = json.loads(sidecar.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + continue + if not isinstance(data, dict): + continue + version_id, model_id = _ids_from_dict(data) + if version_id is not None or model_id is not None: + return version_id, model_id + return None, None + + def scan_models() -> list[dict]: items: list[dict] = [] if not MODELS.is_dir(): @@ -92,6 +109,20 @@ def scan_models() -> list[dict]: continue rel = path.relative_to(MODELS).as_posix() version_id, model_id = read_sidecar_ids(path) + # Full sidecar → skip expensive SHA256 (capture resolves URL locally). + if version_id is not None and model_id is not None: + items.append( + { + "kind": kind, + "rel": rel, + "name": path.name, + "sha256": None, + "version_id": version_id, + "model_id": model_id, + "sha_skipped": True, + } + ) + continue try: digest = sha256_file(path) except OSError as exc: diff --git a/src/gpu_rent/remote/tune_swarm_perf.py b/src/gpu_rent/remote/tune_swarm_perf.py new file mode 100644 index 0000000..939822a --- /dev/null +++ b/src/gpu_rent/remote/tune_swarm_perf.py @@ -0,0 +1,165 @@ +#!/usr/bin/env python3 +"""First-boot Swarm/Comfy performance tune: sageattention ExtraArgs + pip libs. + +Idempotent. Marker: /mnt/swarm_data/.gpu-rent-perf-tuned +Re-runs if GPU uuid changed. +""" + +from __future__ import annotations + +import json +import os +import re +import subprocess +import sys +from pathlib import Path + +DATA = Path("/mnt/swarm_data") +GPU_JSON = DATA / ".gpu-rent-gpu.json" +MARKER = DATA / ".gpu-rent-perf-tuned" +BACKENDS = DATA / "Data" / "Backends.fds" +COMFY_VENV_CANDIDATES = [ + DATA / "dlbackend" / "comfy" / "venv" / "bin" / "pip", + DATA / "dlbackend" / "comfy" / "ComfyUI" / "venv" / "bin" / "pip", + Path("/opt/swarmui/dlbackend/comfy/venv/bin/pip"), + Path("/opt/swarmui/dlbackend/comfy/ComfyUI/venv/bin/pip"), +] + + +def load_gpu() -> dict: + if not GPU_JSON.is_file(): + return {} + try: + return json.loads(GPU_JSON.read_text(encoding="utf-8")) + except json.JSONDecodeError: + return {} + + +def tier_notes(gpu: dict) -> dict: + """Mirror gpu_rent.perf_tiers.swarm_tune_for (stdlib-only on VM).""" + vram = int(gpu.get("vram_mib") or 0) + gib = vram / 1024.0 + if gib < 16: + tier = "low" + elif gib < 24: + tier = "mid" + elif gib < 48: + tier = "high" + else: + tier = "ultra" + cap = str(gpu.get("compute_cap") or "0.0") + try: + parts = cap.split(".") + maj, mnr = int(parts[0]), int(parts[1]) if len(parts) > 1 else 0 + ampere = (maj, mnr) >= (8, 0) + except ValueError: + ampere = False + use_sage = ampere and tier in {"mid", "high", "ultra"} + return { + "tier": tier, + "use_sage": use_sage, + "extra_args": "--use-sage-attention" if use_sage else "", + "uuid": str(gpu.get("uuid") or ""), + "name": str(gpu.get("name") or ""), + } + + +def find_pip() -> Path | None: + for p in COMFY_VENV_CANDIDATES: + if p.is_file(): + return p + # glob + for p in DATA.glob("dlbackend/**/venv/bin/pip"): + return p + return None + + +def patch_backends_extra_args(extra: str) -> bool: + """Ensure ExtraArgs contains sage flag for Comfy self-start backends.""" + if not extra: + return False + if not BACKENDS.is_file(): + print(f"no {BACKENDS} yet — skip ExtraArgs (Comfy not registered)") + return False + text = BACKENDS.read_text(encoding="utf-8") + if "--use-sage-attention" in text: + print("Backends.fds already has --use-sage-attention") + return False + lines = text.splitlines() + changed = False + out = [] + for line in lines: + if re.match(r"^(\s*)ExtraArgs:\s*$", line) or re.match(r"^(\s*)ExtraArgs:\s*\"\"\s*$", line): + indent = re.match(r"^(\s*)", line).group(1) + out.append(f"{indent}ExtraArgs: {extra}") + changed = True + elif re.match(r"^(\s*)ExtraArgs:\s+", line) and "--use-sage-attention" not in line: + out.append(line.rstrip() + f" {extra}") + changed = True + else: + out.append(line) + if not changed: + print("Backends.fds: no ExtraArgs field patched") + return False + BACKENDS.write_text("\n".join(out) + "\n", encoding="utf-8") + print(f"patched {BACKENDS} ExtraArgs += {extra}") + return True + + +def pip_install_sage(pip: Path) -> None: + print(f"pip install triton sageattention via {pip}") + env = dict(os.environ) + env["PIP_DISABLE_PIP_VERSION_CHECK"] = "1" + # Best-effort: do not fail whole tune if wheels missing for this torch. + cmd = [str(pip), "install", "-U", "triton", "sageattention"] + try: + subprocess.check_call(cmd, env=env) + print("triton + sageattention installed") + except subprocess.CalledProcessError as exc: + print(f"WARN: pip install failed ({exc}) — ExtraArgs may no-op until fixed") + + +def main() -> int: + gpu = load_gpu() + plan = tier_notes(gpu) + prev = {} + if MARKER.is_file(): + try: + prev = json.loads(MARKER.read_text(encoding="utf-8")) + except json.JSONDecodeError: + prev = {} + if prev.get("uuid") and prev.get("uuid") == plan["uuid"] and prev.get("extra_args") == plan["extra_args"]: + if prev.get("pip_ok") or not plan["use_sage"]: + print(f"perf tune already applied for {plan['name']} ({plan['tier']})") + return 0 + + print(f"perf tune: {plan['name']} tier={plan['tier']} sage={plan['use_sage']}") + pip_ok = False + restarted_needed = False + if plan["use_sage"]: + pip = find_pip() + if pip: + pip_install_sage(pip) + pip_ok = True + else: + print("Comfy venv pip not found yet — will retry next up") + if patch_backends_extra_args(plan["extra_args"]): + restarted_needed = True + + marker = { + "uuid": plan["uuid"], + "name": plan["name"], + "tier": plan["tier"], + "extra_args": plan["extra_args"], + "pip_ok": pip_ok or not plan["use_sage"], + "restart_needed": restarted_needed, + } + MARKER.write_text(json.dumps(marker, indent=2) + "\n", encoding="utf-8") + print("wrote", MARKER) + if restarted_needed: + print("RESTART_SWARMUI=1") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/src/gpu_rent/session.py b/src/gpu_rent/session.py index 6aefaad..4af50f4 100644 --- a/src/gpu_rent/session.py +++ b/src/gpu_rent/session.py @@ -21,7 +21,7 @@ from gpu_rent.cloud import ( wait_volume, ) from gpu_rent.bootstrap import run_bootstrap -from gpu_rent.provision import provision_vm +from gpu_rent.provision import provision_vm, tune_swarm_perf from gpu_rent.ready import wait_backend_idle from gpu_rent.snapshot import ensure_boot_snapshot from gpu_rent.notify import notify_ready @@ -119,6 +119,13 @@ def _bind_access( wait_backend_idle(cfg, ip, log) except CloudError as exc: log(f"ready: {exc}") + # Comfy venv + Backends.fds exist after Idle — sage/triton + ExtraArgs. + try: + if tune_swarm_perf(cfg, ip, log): + log("systemctl restart swarmui (perf ExtraArgs)") + run_ssh(cfg, ip, "sudo -n systemctl restart swarmui", timeout=120) + except Exception as exc: + log(f"perf tune: {exc}") try: ensure_boot_snapshot( conn, diff --git a/tests/test_capture_merge.py b/tests/test_capture_merge.py index 1d1be1a..7ed2841 100644 --- a/tests/test_capture_merge.py +++ b/tests/test_capture_merge.py @@ -1,6 +1,7 @@ """Unit tests for capture merge / URL builders (no SSH).""" from pathlib import Path +from unittest.mock import patch from gpu_rent.capture import ( ExtCaptureItem, @@ -11,6 +12,7 @@ from gpu_rent.capture import ( strip_git_auth, ) from gpu_rent.civitai import civitai_model_url, version_ids_from_payload +from gpu_rent.errors import CloudError from gpu_rent.manifests import parse_extensions, parse_models @@ -29,22 +31,90 @@ def test_version_ids_from_payload(): def test_resolve_from_sidecar(): - item = resolve_model_item( + out = resolve_model_item( { "kind": "lora", "rel": "Lora/foo.safetensors", "name": "foo.safetensors", "version_id": 3107521, "model_id": 2187487, - "sha256": "a" * 64, + "sha256": None, }, token="", api_host="civitai.red", link_host="civitai.red", ) - assert item is not None - assert item.version_id == 3107521 - assert "modelVersionId=3107521" in item.url + assert out.status == "ok" + assert out.item is not None + assert out.item.version_id == 3107521 + assert "modelVersionId=3107521" in out.item.url + + +def test_resolve_vid_only_fetches_model_id(): + with patch( + "gpu_rent.capture.fetch_model_version", + return_value=("civitai.red", {"id": 10, "modelId": 20, "name": "X"}), + ) as mocked: + out = resolve_model_item( + { + "kind": "lora", + "rel": "Lora/x.safetensors", + "name": "x.safetensors", + "version_id": 10, + "model_id": None, + "sha256": None, + }, + token="tok", + api_host="civitai.red", + link_host="civitai.red", + ) + mocked.assert_called_once() + assert out.status == "ok" + assert out.item is not None + assert out.item.model_id == 20 + assert "models/20?modelVersionId=10" in out.item.url + + +def test_resolve_by_hash_404_is_unknown(): + with patch( + "gpu_rent.capture.fetch_model_version_by_hash", + side_effect=CloudError("Civitai by-hash abc: HTTP 404 (хосты civitai.red)"), + ): + out = resolve_model_item( + { + "kind": "lora", + "rel": "Lora/m.safetensors", + "name": "m.safetensors", + "sha256": "a" * 64, + }, + token="", + api_host="civitai.red", + link_host="civitai.red", + ) + assert out.status == "unknown" + assert out.item is None + + +def test_resolve_by_hash_network_is_api_error(): + with patch( + "gpu_rent.capture.fetch_model_version_by_hash", + side_effect=CloudError( + "Civitai by-hash abc: Connection timeout (хосты civitai.red)" + ), + ): + out = resolve_model_item( + { + "kind": "lora", + "rel": "Lora/m.safetensors", + "name": "m.safetensors", + "sha256": "b" * 64, + }, + token="", + api_host="civitai.red", + link_host="civitai.red", + ) + assert out.status == "api_error" + assert "timeout" in out.detail.lower() or "Connection" in out.detail def test_merge_models_dedupe(tmp_path: Path): @@ -54,9 +124,19 @@ def test_merge_models_dedupe(tmp_path: Path): encoding="utf-8", ) items = [ - ModelCaptureItem("lora", 100, 1, "https://civitai.red/models/1?modelVersionId=100", "old"), - ModelCaptureItem("lora", 200, 2, "https://civitai.red/models/2?modelVersionId=200", "new"), - ModelCaptureItem("checkpoint", 300, 3, "https://civitai.red/models/3?modelVersionId=300", "ckpt"), + ModelCaptureItem( + "lora", 100, 1, "https://civitai.red/models/1?modelVersionId=100", "old" + ), + ModelCaptureItem( + "lora", 200, 2, "https://civitai.red/models/2?modelVersionId=200", "new" + ), + ModelCaptureItem( + "checkpoint", + 300, + 3, + "https://civitai.red/models/3?modelVersionId=300", + "ckpt", + ), ] added, skipped = merge_models_yaml(path, items, dry_run=False) assert len(added) == 2 @@ -72,13 +152,38 @@ def test_merge_models_dry_run(tmp_path: Path): path.write_text("lora: []\n", encoding="utf-8") before = path.read_text(encoding="utf-8") items = [ - ModelCaptureItem("lora", 1, 1, "https://civitai.red/models/1?modelVersionId=1", "x"), + ModelCaptureItem( + "lora", 1, 1, "https://civitai.red/models/1?modelVersionId=1", "x" + ), ] added, _ = merge_models_yaml(path, items, dry_run=True) assert len(added) == 1 assert path.read_text(encoding="utf-8") == before +def test_merge_models_kind_scoped_dedupe(tmp_path: Path): + """Same version_id under different kinds can both be kept.""" + path = tmp_path / "models.yaml" + path.write_text( + "lora:\n - url: https://civitai.red/models/1?modelVersionId=100\n", + encoding="utf-8", + ) + items = [ + ModelCaptureItem( + "checkpoint", + 100, + 1, + "https://civitai.red/models/1?modelVersionId=100", + "as-ckpt", + ), + ] + added, skipped = merge_models_yaml(path, items, dry_run=False) + assert len(added) == 1 + assert skipped == [] + entries = parse_models(path) + assert {(e.kind, e.version_id) for e in entries} == {("lora", 100), ("checkpoint", 100)} + + def test_merge_extensions_dedupe(tmp_path: Path): path = tmp_path / "extensions.yaml" path.write_text( @@ -95,8 +200,9 @@ def test_merge_extensions_dedupe(tmp_path: Path): "C", ), ] - added, skipped = merge_extensions_yaml(path, items, dry_run=False) + added, updated, skipped = merge_extensions_yaml(path, items, dry_run=False) assert len(added) == 2 + assert updated == [] assert len(skipped) == 1 repos = parse_extensions(path) urls = [r.url for r in repos] @@ -104,6 +210,25 @@ def test_merge_extensions_dedupe(tmp_path: Path): assert all("SECRET" not in u for u in urls) +def test_merge_extensions_updates_url_same_dir(tmp_path: Path): + path = tmp_path / "extensions.yaml" + path.write_text( + "swarmui:\n - url: https://github.com/old/A.git\n ref: main\n dir: A\n", + encoding="utf-8", + ) + items = [ + ExtCaptureItem("swarmui", "https://github.com/new/A.git", "develop", "A"), + ] + added, updated, skipped = merge_extensions_yaml(path, items, dry_run=False) + assert added == [] + assert len(updated) == 1 + assert skipped == [] + repos = parse_extensions(path) + assert len(repos) == 1 + assert repos[0].url == "https://github.com/new/A.git" + assert repos[0].ref == "develop" + + def test_strip_git_auth(): assert ( strip_git_auth("https://x-access-token:tok@github.com/org/r.git") diff --git a/tests/test_perf_tiers.py b/tests/test_perf_tiers.py new file mode 100644 index 0000000..6acc4c9 --- /dev/null +++ b/tests/test_perf_tiers.py @@ -0,0 +1,70 @@ +from gpu_rent.perf_tiers import ( + GpuInfo, + compute_cap_at_least, + ollama_env_lines, + ollama_tune_for, + swarm_tune_for, + tier_for_vram_mib, +) + + +def test_tier_boundaries(): + assert tier_for_vram_mib(8 * 1024) == "low" + assert tier_for_vram_mib(16 * 1024) == "mid" + assert tier_for_vram_mib(23 * 1024) == "mid" + assert tier_for_vram_mib(24 * 1024) == "high" + assert tier_for_vram_mib(40 * 1024) == "high" + assert tier_for_vram_mib(48 * 1024) == "ultra" + assert tier_for_vram_mib(80 * 1024) == "ultra" + + +def test_compute_cap(): + assert compute_cap_at_least("8.0", 8, 0) + assert compute_cap_at_least("8.9", 8, 0) + assert not compute_cap_at_least("7.5", 8, 0) + + +def test_ollama_high_reserves_vram_for_swarm(): + info = GpuInfo( + name="NVIDIA A100-SXM4-40GB", + vram_mib=40960, + compute_cap="8.0", + uuid="GPU-1", + tier="high", + ) + tune = ollama_tune_for(info) + assert tune.flash_attention + assert tune.num_parallel == 1 + assert tune.max_loaded_models == 1 + assert tune.kv_cache_type == "q8_0" + assert tune.gpu_overhead_bytes == 14 * 1024**3 + env = "\n".join(ollama_env_lines(tune)) + assert "OLLAMA_FLASH_ATTENTION=1" in env + assert "OLLAMA_GPU_OVERHEAD=" in env + + +def test_swarm_sage_on_ampere_mid(): + info = GpuInfo( + name="NVIDIA GeForce RTX 4090", + vram_mib=24576, + compute_cap="8.9", + uuid="GPU-2", + tier="high", + ) + st = swarm_tune_for(info) + assert st.use_sage_attention + assert "--use-sage-attention" in st.comfy_extra_args + assert st.install_triton_sage + + +def test_swarm_no_sage_on_low(): + info = GpuInfo( + name="NVIDIA T4", + vram_mib=15360, + compute_cap="7.5", + uuid="GPU-3", + tier="low", + ) + st = swarm_tune_for(info) + assert not st.use_sage_attention + assert st.comfy_extra_args == ""