Skip to content
6 changes: 5 additions & 1 deletion .github/workflows/test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ jobs:
strategy:
matrix:
os: [ubuntu-latest]
python-version: ["3.10", "3.11", "3.12", "3.13"]
python-version: ["3.10", "3.11", "3.12", "3.13", "3.14"]
steps:
- uses: actions/checkout@v4
- name: Set up Python ${{ matrix.python-version }}
Expand All @@ -43,6 +43,10 @@ jobs:
run: |
python -m pip install --upgrade pip
pip install '.[all,dev]'
- name: Run pyright
uses: jakebailey/pyright-action@v3
with:
pylance-version: latest-release
- name: Run doctest
run: pytest --doctest-modules src/gpuhunt
- name: Run pytest
Expand Down
7 changes: 7 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@ dev = [
"pre-commit",
"pytest~=7.0",
"pytest-mock",
"pyright==1.1.403", # Should match the pinned version in CI
"ruff==0.5.3", # Should match .pre-commit-config.yaml
"requests-mock",
]
Expand All @@ -74,3 +75,9 @@ ignore = [
[tool.ruff.lint.isort]
known-first-party = ["gpuhunt"]
combine-as-imports = true

[tool.pyright]
typeCheckingMode = "standard"
include = [
"src/"
]
6 changes: 5 additions & 1 deletion src/gpuhunt/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,11 @@
default_catalog as default_catalog,
query as query,
)
from gpuhunt._internal.errors import (
GPUHuntError as GPUHuntError,
MissingCredsError as MissingCredsError,
ProviderError as ProviderError,
)
from gpuhunt._internal.models import (
AcceleratorInfo as AcceleratorInfo,
AcceleratorVendor as AcceleratorVendor,
Expand All @@ -24,7 +29,6 @@
IntelAcceleratorInfo as IntelAcceleratorInfo,
NvidiaGPUInfo as NvidiaGPUInfo,
QueryFilter as QueryFilter,
RawCatalogItem as RawCatalogItem,
TenstorrentAcceleratorInfo as TenstorrentAcceleratorInfo,
TPUInfo as TPUInfo,
)
44 changes: 19 additions & 25 deletions src/gpuhunt/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import gpuhunt._internal.storage as storage
from gpuhunt._internal.utils import configure_logging
from gpuhunt.providers.base import OfflineProvider


def main():
Expand Down Expand Up @@ -41,49 +42,42 @@ def main():
elif args.provider == "azure":
from gpuhunt.providers.azure import AzureProvider

provider = AzureProvider(os.getenv("AZURE_SUBSCRIPTION_ID"))
provider = AzureProvider(os.environ["AZURE_SUBSCRIPTION_ID"])
elif args.provider == "crusoe":
from gpuhunt.providers.crusoe import CrusoeProvider

provider = CrusoeProvider(
access_key=os.getenv("CRUSOE_ACCESS_KEY"),
secret_key=os.getenv("CRUSOE_SECRET_KEY"),
project_id=os.getenv("CRUSOE_PROJECT_ID"),
)
provider = CrusoeProvider.from_env()
elif args.provider == "cloudrift":
from gpuhunt.providers.cloudrift import CloudRiftProvider

provider = CloudRiftProvider()
elif args.provider == "verda":
from gpuhunt.providers.verda import VerdaProvider

provider = VerdaProvider(os.getenv("VERDA_CLIENT_ID"), os.getenv("VERDA_CLIENT_SECRET"))
provider = VerdaProvider(
client_id=os.environ["VERDA_CLIENT_ID"],
client_secret=os.environ["VERDA_CLIENT_SECRET"],
)
elif args.provider == "digitalocean":
from gpuhunt.providers.digitalocean import DigitalOceanProvider

provider = DigitalOceanProvider(
api_key=os.getenv("DIGITAL_OCEAN_API_KEY"), api_url=os.getenv("DIGITAL_OCEAN_API_URL")
)
provider = DigitalOceanProvider.from_env()
elif args.provider == "gcp":
from gpuhunt.providers.gcp import GCPProvider

provider = GCPProvider(os.getenv("GCP_PROJECT_ID"))
provider = GCPProvider(project=os.environ["GCP_PROJECT_ID"])
elif args.provider == "hotaisle":
from gpuhunt.providers.hotaisle import HotAisleProvider

provider = HotAisleProvider(
api_key=os.getenv("HOTAISLE_API_KEY"), team_handle=os.getenv("HOTAISLE_TEAM_HANDLE")
)
provider = HotAisleProvider.from_env()
elif args.provider == "jarvislabs":
from gpuhunt.providers.jarvislabs import JarvisLabsProvider

provider = JarvisLabsProvider(
api_key=os.getenv("JL_API_KEY"), api_url=os.getenv("JARVISLABS_API_URL")
)
provider = JarvisLabsProvider.from_env()
elif args.provider == "lambdalabs":
from gpuhunt.providers.lambdalabs import LambdaLabsProvider

provider = LambdaLabsProvider(os.getenv("LAMBDALABS_TOKEN"))
provider = LambdaLabsProvider(token=os.environ["LAMBDALABS_TOKEN"])
elif args.provider == "nebius":
from nebius.base.service_account.pk_file import Reader as PKReader

Expand All @@ -95,9 +89,9 @@ def main():
os.getenv("NEBIUS_ACCESS_TOKEN")
# or service account credentials
or PKReader(
filename=os.getenv("NEBIUS_PRIVATE_KEY_FILE"),
public_key_id=os.getenv("NEBIUS_PUBLIC_KEY_ID"),
service_account_id=os.getenv("NEBIUS_SERVICE_ACCOUNT_ID"),
filename=os.environ["NEBIUS_PRIVATE_KEY_FILE"],
public_key_id=os.environ["NEBIUS_PUBLIC_KEY_ID"],
service_account_id=os.environ["NEBIUS_SERVICE_ACCOUNT_ID"],
)
)
)
Expand All @@ -120,21 +114,21 @@ def main():
elif args.provider == "seeweb":
from gpuhunt.providers.seeweb import SeewebProvider

provider = SeewebProvider(os.getenv("SEEWEB_API_TOKEN"))
provider = SeewebProvider.from_env()
elif args.provider == "vastai":
from gpuhunt.providers.vastai import VastAIProvider

provider = VastAIProvider()
provider = VastAIProvider.from_env()
elif args.provider == "vultr":
from gpuhunt.providers.vultr import VultrProvider

provider = VultrProvider()
provider = VultrProvider.from_env()
else:
exit(f"Unknown provider {args.provider}")

logging.info("Fetching offers for %s", args.provider)
offers = provider.get()
if not args.no_filter:
if not args.no_filter and isinstance(provider, OfflineProvider):
offers = provider.filter(offers)
storage.dump(offers, args.output)

Expand Down
21 changes: 10 additions & 11 deletions src/gpuhunt/_internal/catalog.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
import csv
import dataclasses
import heapq
import io
import logging
Expand All @@ -13,9 +11,10 @@
from pathlib import Path

import gpuhunt._internal.constraints as constraints
import gpuhunt._internal.storage as storage
from gpuhunt._internal.models import AcceleratorVendor, CatalogItem, CPUArchitecture, QueryFilter
from gpuhunt._internal.utils import parse_compute_capability
from gpuhunt.providers import AbstractProvider
from gpuhunt.providers.base import AbstractProvider

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -212,9 +211,10 @@ def _load(self, version: str | None = None):
for provider in OFFLINE_PROVIDERS:
try:
with zip_file.open(f"{provider}.csv", "r") as csv_file:
reader = csv.DictReader(io.TextIOWrapper(csv_file, "utf-8"))
for row in reader:
item = CatalogItem.from_dict(row, provider=provider)
items = storage.load(
io.TextIOWrapper(csv_file, "utf-8"), provider=provider
)
for item in items:
catalog.setdefault(provider, []).append(item)
except KeyError:
logger.error(
Expand Down Expand Up @@ -253,9 +253,9 @@ def _get_offline_provider_items(
catalog_dir = os.getenv("GPUHUNT_CATALOG_DIR")
if catalog_dir is not None:
with open(Path(catalog_dir) / f"{provider_name}.csv", "rb") as csv_file:
reader = csv.DictReader(io.TextIOWrapper(csv_file, "utf-8"))
for row in reader:
item = CatalogItem.from_dict(row, provider=provider_name)
for item in storage.load(
io.TextIOWrapper(csv_file, "utf-8"), provider=provider_name
):
if constraints.matches(item, query_filter):
items.append(item)
return items
Expand All @@ -282,10 +282,9 @@ def _get_online_provider_items(
if provider.NAME != provider_name:
continue
found = True
for i in provider.get(
for item in provider.get(
query_filter=query_filter, balance_resources=self.balance_resources
):
item = CatalogItem(provider=provider_name, **dataclasses.asdict(i))
if constraints.matches(item, query_filter):
items.append(item)
if not found:
Expand Down
73 changes: 38 additions & 35 deletions src/gpuhunt/_internal/constraints.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,30 +15,6 @@
TPUInfo,
)

# v5litepod = v5e
_TPU_VERSIONS = ["v2", "v3", "v4", "v5p", "v5litepod", "v6e"]


Comparable = TypeVar("Comparable", bound=int | float | tuple[int, int])


def is_between(value: Comparable, left: Comparable | None, right: Comparable | None) -> bool:
if is_below(value, left) or is_above(value, right):
return False
return True


def is_below(value: Comparable, limit: Comparable | None) -> bool:
if limit is not None and value < limit:
return True
return False


def is_above(value: Comparable, limit: Comparable | None) -> bool:
if limit is not None and value > limit:
return True
return False


def matches(i: CatalogItem, q: QueryFilter) -> bool:
"""
Expand All @@ -53,21 +29,21 @@ def matches(i: CatalogItem, q: QueryFilter) -> bool:
"""
if q.provider is not None and i.provider.lower() not in map(str.lower, q.provider):
return False
if not is_between(i.price, q.min_price, q.max_price):
if not _is_between(i.price, q.min_price, q.max_price):
return False
if q.spot is not None and i.spot != q.spot:
return False
if q.cpu_arch and q.cpu_arch != i.cpu_arch:
return False
if not is_between(i.cpu, q.min_cpu, q.max_cpu):
if not _is_between(i.cpu, q.min_cpu, q.max_cpu):
return False
if not is_between(i.memory, q.min_memory, q.max_memory):
if not _is_between(i.memory, q.min_memory, q.max_memory):
return False
if not (q.min_gpu_count == 0 and i.gpu_count == 0):
# GPU filters should not be applied to non-gpu offers if `q.min_gpu_count == 0`.
if q.gpu_vendor and q.gpu_vendor != i.gpu_vendor:
return False
if not is_between(i.gpu_count, q.min_gpu_count, q.max_gpu_count):
if not _is_between(i.gpu_count, q.min_gpu_count, q.max_gpu_count):
return False
if q.gpu_name is not None:
if i.gpu_name is None:
Expand All @@ -80,20 +56,22 @@ def matches(i: CatalogItem, q: QueryFilter) -> bool:
if not i.gpu_name:
return False
cc = get_compute_capability(i.gpu_name)
if not cc or not is_between(cc, q.min_compute_capability, q.max_compute_capability):
if not cc or not _is_between(cc, q.min_compute_capability, q.max_compute_capability):
return False
if not is_between(
i.gpu_memory if i.gpu_count > 0 else 0, q.min_gpu_memory, q.max_gpu_memory
if not _is_between(
i.gpu_memory if i.gpu_count > 0 and i.gpu_memory is not None else 0,
q.min_gpu_memory,
q.max_gpu_memory,
):
return False
if not is_between(
(i.gpu_count * i.gpu_memory) if i.gpu_count > 0 else 0,
if not _is_between(
(i.gpu_count * i.gpu_memory) if i.gpu_count > 0 and i.gpu_memory is not None else 0,
q.min_total_gpu_memory,
q.max_total_gpu_memory,
):
return False
if i.disk_size is not None:
if not is_between(i.disk_size, q.min_disk_size, q.max_disk_size):
if not _is_between(i.disk_size, q.min_disk_size, q.max_disk_size):
return False
if q.allowed_flags is not None:
if any(flag not in q.allowed_flags for flag in i.flags):
Expand All @@ -116,7 +94,7 @@ def find_accelerators(


def get_compute_capability(gpu_name: str) -> tuple[int, int] | None:
if accelerators := find_accelerators(names=[gpu_name], vendors=AcceleratorVendor.NVIDIA):
if accelerators := find_accelerators(names=[gpu_name], vendors=[AcceleratorVendor.NVIDIA]):
assert isinstance(accelerators[0], NvidiaGPUInfo)
return accelerators[0].compute_capability
return None
Expand Down Expand Up @@ -298,6 +276,10 @@ def is_nvidia_superchip(gpu_name: str) -> bool:
),
]


# v5litepod = v5e
_TPU_VERSIONS = ["v2", "v3", "v4", "v5p", "v5litepod", "v6e"]

KNOWN_TPUS: list[TPUInfo] = [TPUInfo(name=version, memory=0) for version in _TPU_VERSIONS]

KNOWN_INTEL_ACCELERATORS: list[IntelAcceleratorInfo] = [
Expand Down Expand Up @@ -326,3 +308,24 @@ def is_nvidia_superchip(gpu_name: str) -> bool:
+ KNOWN_INTEL_ACCELERATORS
+ KNOWN_TENSTORRENT_ACCELERATORS
)


Comparable = TypeVar("Comparable", int, float, tuple[int, int])


def _is_between(value: Comparable, left: Comparable | None, right: Comparable | None) -> bool:
if _is_below(value, left) or _is_above(value, right):
return False
return True


def _is_below(value: Comparable, limit: Comparable | None) -> bool:
if limit is not None and value < limit:
return True
return False


def _is_above(value: Comparable, limit: Comparable | None) -> bool:
if limit is not None and value > limit:
return True
return False
Loading
Loading