gpuhunt 0.0.1.dev1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- gpuhunt/__init__.py +14 -0
- gpuhunt/__main__.py +52 -0
- gpuhunt/_catalog.py +71 -0
- gpuhunt/_models.py +24 -0
- gpuhunt/_storage.py +26 -0
- gpuhunt/providers/__init__.py +13 -0
- gpuhunt/providers/aws.py +195 -0
- gpuhunt/providers/azure.py +222 -0
- gpuhunt/providers/gcp.py +230 -0
- gpuhunt/providers/lambdalabs.py +68 -0
- gpuhunt/providers/tensordock.py +128 -0
- gpuhunt/providers/vastai.py +52 -0
- gpuhunt/version.py +1 -0
- gpuhunt-0.0.1.dev1.dist-info/LICENSE +353 -0
- gpuhunt-0.0.1.dev1.dist-info/METADATA +54 -0
- gpuhunt-0.0.1.dev1.dist-info/RECORD +19 -0
- gpuhunt-0.0.1.dev1.dist-info/WHEEL +5 -0
- gpuhunt-0.0.1.dev1.dist-info/top_level.txt +2 -0
- version.py +1 -0
gpuhunt/__init__.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
from functools import lru_cache
|
|
2
|
+
|
|
3
|
+
from gpuhunt._catalog import Catalog
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
@lru_cache()
|
|
7
|
+
def default_catalog() -> Catalog:
|
|
8
|
+
catalog = Catalog()
|
|
9
|
+
catalog.load()
|
|
10
|
+
return catalog
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def query() -> list: # todo
|
|
14
|
+
return default_catalog().query()
|
gpuhunt/__main__.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
import sys
|
|
5
|
+
|
|
6
|
+
import gpuhunt._storage as storage
|
|
7
|
+
from gpuhunt.providers.aws import AWSProvider
|
|
8
|
+
from gpuhunt.providers.azure import AzureProvider
|
|
9
|
+
from gpuhunt.providers.gcp import GCPProvider
|
|
10
|
+
from gpuhunt.providers.lambdalabs import LambdaLabsProvider
|
|
11
|
+
from gpuhunt.providers.tensordock import TensorDockProvider
|
|
12
|
+
from gpuhunt.providers.vastai import VastAIProvider
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def main():
|
|
16
|
+
parser = argparse.ArgumentParser(prog="python3 -m gpuhunt")
|
|
17
|
+
parser.add_argument(
|
|
18
|
+
"provider",
|
|
19
|
+
choices=["aws", "azure", "gcp", "lambdalabs", "tensordock", "vastai"],
|
|
20
|
+
)
|
|
21
|
+
parser.add_argument("--output", required=True)
|
|
22
|
+
parser.add_argument("--no-filter", action="store_true")
|
|
23
|
+
args = parser.parse_args()
|
|
24
|
+
logging.basicConfig(
|
|
25
|
+
level=logging.INFO,
|
|
26
|
+
stream=sys.stdout,
|
|
27
|
+
format="%(asctime)s %(levelname)s %(message)s",
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
if args.provider == "aws":
|
|
31
|
+
provider = AWSProvider()
|
|
32
|
+
elif args.provider == "azure":
|
|
33
|
+
provider = AzureProvider(os.getenv("AZURE_SUBSCRIPTION_ID"))
|
|
34
|
+
elif args.provider == "gcp":
|
|
35
|
+
provider = GCPProvider(os.getenv("GCP_PROJECT_ID"))
|
|
36
|
+
elif args.provider == "lambdalabs":
|
|
37
|
+
provider = LambdaLabsProvider(os.getenv("LAMBDALABS_TOKEN"))
|
|
38
|
+
elif args.provider == "tensordock":
|
|
39
|
+
provider = TensorDockProvider()
|
|
40
|
+
elif args.provider == "vastai":
|
|
41
|
+
provider = VastAIProvider()
|
|
42
|
+
else:
|
|
43
|
+
exit(f"Unknown provider {args.provider}")
|
|
44
|
+
|
|
45
|
+
offers = provider.get()
|
|
46
|
+
if not args.no_filter:
|
|
47
|
+
offers = provider.filter(offers)
|
|
48
|
+
storage.dump(offers, args.output)
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
if __name__ == "__main__":
|
|
52
|
+
main()
|
gpuhunt/_catalog.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
import csv
|
|
2
|
+
import io
|
|
3
|
+
import logging
|
|
4
|
+
import urllib.request
|
|
5
|
+
import zipfile
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from typing import Iterable, Optional
|
|
8
|
+
|
|
9
|
+
logger = logging.getLogger(__name__)
|
|
10
|
+
version_url = "https://dstack-gpu-pricing.s3.eu-west-1.amazonaws.com/v1/version"
|
|
11
|
+
catalog_url = (
|
|
12
|
+
"https://dstack-gpu-pricing.s3.eu-west-1.amazonaws.com/v1/{version}/catalog.zip"
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@dataclass(frozen=True)
|
|
17
|
+
class CatalogItem:
|
|
18
|
+
provider: str
|
|
19
|
+
instance_name: str
|
|
20
|
+
location: str
|
|
21
|
+
price: float
|
|
22
|
+
cpus: int
|
|
23
|
+
memory: float
|
|
24
|
+
gpu_count: int
|
|
25
|
+
gpu_name: Optional[str]
|
|
26
|
+
gpu_memory: Optional[float]
|
|
27
|
+
spot: bool
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class Catalog:
|
|
31
|
+
def __init__(self):
|
|
32
|
+
self.catalog = None
|
|
33
|
+
|
|
34
|
+
def query(self) -> list[CatalogItem]:
|
|
35
|
+
return list(self._read_catalog())
|
|
36
|
+
|
|
37
|
+
def load(self, version: str = None):
|
|
38
|
+
if version is None:
|
|
39
|
+
version = self.get_latest_version()
|
|
40
|
+
logger.debug("Downloading catalog %s...", version)
|
|
41
|
+
with urllib.request.urlopen(catalog_url.format(version=version)) as f:
|
|
42
|
+
self.catalog = io.BytesIO(f.read())
|
|
43
|
+
|
|
44
|
+
@staticmethod
|
|
45
|
+
def get_latest_version() -> str:
|
|
46
|
+
with urllib.request.urlopen(version_url) as f:
|
|
47
|
+
return f.read().decode("utf-8").strip()
|
|
48
|
+
|
|
49
|
+
def _read_catalog(self) -> Iterable[CatalogItem]:
|
|
50
|
+
with zipfile.ZipFile(self.catalog) as zip_file:
|
|
51
|
+
providers = [f[:-4] for f in zip_file.namelist() if f.endswith(".csv")]
|
|
52
|
+
for provider in providers:
|
|
53
|
+
with zip_file.open(f"{provider}.csv", "r") as csv_file:
|
|
54
|
+
reader: Iterable[dict[str, str]] = csv.DictReader(
|
|
55
|
+
io.TextIOWrapper(csv_file, "utf-8")
|
|
56
|
+
)
|
|
57
|
+
for row in reader:
|
|
58
|
+
yield CatalogItem(
|
|
59
|
+
provider=provider,
|
|
60
|
+
instance_name=row["instance_name"],
|
|
61
|
+
location=row["location"],
|
|
62
|
+
price=float(row["price"]),
|
|
63
|
+
cpus=int(row["cpu"]),
|
|
64
|
+
memory=float(row["memory"]),
|
|
65
|
+
gpu_count=int(row["gpu_count"]),
|
|
66
|
+
gpu_name=row["gpu_name"] or None,
|
|
67
|
+
gpu_memory=float(row["gpu_memory"])
|
|
68
|
+
if row["gpu_memory"]
|
|
69
|
+
else None,
|
|
70
|
+
spot=row["spot"] == "True",
|
|
71
|
+
)
|
gpuhunt/_models.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
from typing import Any, Optional
|
|
2
|
+
|
|
3
|
+
from pydantic import BaseModel, model_validator
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class InstanceOffer(BaseModel):
|
|
7
|
+
instance_name: str
|
|
8
|
+
location: Optional[str] = None # region or zone
|
|
9
|
+
price: Optional[float] = None # $ per hour
|
|
10
|
+
cpu: Optional[int] = None
|
|
11
|
+
memory: Optional[float] = None # in GB
|
|
12
|
+
gpu_count: Optional[int] = None
|
|
13
|
+
gpu_name: Optional[str] = None
|
|
14
|
+
gpu_memory: Optional[float] = None # in GB
|
|
15
|
+
spot: Optional[bool] = None
|
|
16
|
+
|
|
17
|
+
@model_validator(mode="before")
|
|
18
|
+
@classmethod
|
|
19
|
+
def parse_empty_as_none(cls, data: Any) -> Any:
|
|
20
|
+
if isinstance(data, dict):
|
|
21
|
+
for key, value in data.items():
|
|
22
|
+
if value == "":
|
|
23
|
+
data[key] = None
|
|
24
|
+
return data
|
gpuhunt/_storage.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
import csv
|
|
2
|
+
from typing import Iterable
|
|
3
|
+
|
|
4
|
+
from gpuhunt._models import InstanceOffer
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def dump(offers: list[InstanceOffer], path: str):
|
|
8
|
+
with open(path, "w", newline="") as f:
|
|
9
|
+
writer = csv.DictWriter(f, fieldnames=list(InstanceOffer.model_fields.keys()))
|
|
10
|
+
writer.writeheader()
|
|
11
|
+
for offer in offers:
|
|
12
|
+
writer.writerow(offer.model_dump())
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def load(path: str) -> list[InstanceOffer]:
|
|
16
|
+
offers = []
|
|
17
|
+
with open(path, "r", newline="") as f:
|
|
18
|
+
reader: Iterable[dict[str, str]] = csv.DictReader(f)
|
|
19
|
+
for row in reader:
|
|
20
|
+
offer = InstanceOffer.model_validate(row)
|
|
21
|
+
offers.append(offer)
|
|
22
|
+
return offers
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def sort_key(offer: InstanceOffer):
|
|
26
|
+
return offer.gpu_count, offer.instance_name, offer.price, offer.location
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
|
|
3
|
+
from gpuhunt._models import InstanceOffer
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class AbstractProvider(ABC):
|
|
7
|
+
@abstractmethod
|
|
8
|
+
def get(self) -> list[InstanceOffer]:
|
|
9
|
+
pass
|
|
10
|
+
|
|
11
|
+
@classmethod
|
|
12
|
+
def filter(cls, offers: list[InstanceOffer]) -> list[InstanceOffer]:
|
|
13
|
+
return offers
|
gpuhunt/providers/aws.py
ADDED
|
@@ -0,0 +1,195 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
import csv
|
|
3
|
+
import datetime
|
|
4
|
+
import logging
|
|
5
|
+
import os
|
|
6
|
+
import re
|
|
7
|
+
import tempfile
|
|
8
|
+
from collections import defaultdict
|
|
9
|
+
from typing import Iterable, Optional
|
|
10
|
+
|
|
11
|
+
import boto3
|
|
12
|
+
import requests
|
|
13
|
+
from botocore.exceptions import ClientError, EndpointConnectionError
|
|
14
|
+
|
|
15
|
+
from gpuhunt._models import InstanceOffer
|
|
16
|
+
from gpuhunt.providers import AbstractProvider
|
|
17
|
+
|
|
18
|
+
logger = logging.getLogger(__name__)
|
|
19
|
+
ec2_pricing_url = "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonEC2/current/index.csv"
|
|
20
|
+
disclaimer_rows_skip = 5
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class AWSProvider(AbstractProvider):
|
|
24
|
+
"""
|
|
25
|
+
AWSProvider parses Bulk API index file for AmazonEC2 in all regions and fills missing GPU details
|
|
26
|
+
|
|
27
|
+
Required IAM permissions:
|
|
28
|
+
* `ec2:DescribeInstanceTypes`
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
def __init__(self, cache_path: Optional[str] = None):
|
|
32
|
+
if cache_path:
|
|
33
|
+
self.cache_path = cache_path
|
|
34
|
+
else:
|
|
35
|
+
self.temp_dir = tempfile.TemporaryDirectory()
|
|
36
|
+
self.cache_path = self.temp_dir.name + "/index.csv"
|
|
37
|
+
# todo aws creds
|
|
38
|
+
self.filters = {
|
|
39
|
+
"TermType": ["OnDemand"],
|
|
40
|
+
"Tenancy": ["Shared"],
|
|
41
|
+
"Operating System": ["Linux"],
|
|
42
|
+
"CapacityStatus": ["Used"],
|
|
43
|
+
"Unit": ["Hrs"],
|
|
44
|
+
"Currency": ["USD"],
|
|
45
|
+
"Pre Installed S/W": ["", "NA"],
|
|
46
|
+
}
|
|
47
|
+
self.preview_gpus = {
|
|
48
|
+
"p4de.24xlarge": ("A100", 80.0),
|
|
49
|
+
}
|
|
50
|
+
|
|
51
|
+
def get(self) -> list[InstanceOffer]:
|
|
52
|
+
if not os.path.exists(self.cache_path):
|
|
53
|
+
logger.info("Downloading EC2 prices to %s", self.cache_path)
|
|
54
|
+
with requests.get(ec2_pricing_url, stream=True) as r:
|
|
55
|
+
r.raise_for_status()
|
|
56
|
+
with open(self.cache_path, "wb") as f:
|
|
57
|
+
for chunk in r.iter_content(chunk_size=8192):
|
|
58
|
+
f.write(chunk)
|
|
59
|
+
|
|
60
|
+
offers = []
|
|
61
|
+
with open(self.cache_path, "r", newline="") as f:
|
|
62
|
+
for _ in range(disclaimer_rows_skip):
|
|
63
|
+
f.readline()
|
|
64
|
+
reader: Iterable[dict[str, str]] = csv.DictReader(f)
|
|
65
|
+
for row in reader:
|
|
66
|
+
if self.skip(row):
|
|
67
|
+
continue
|
|
68
|
+
offer = InstanceOffer(
|
|
69
|
+
instance_name=row["Instance Type"],
|
|
70
|
+
location=row["Region Code"],
|
|
71
|
+
price=float(row["PricePerUnit"]),
|
|
72
|
+
cpu=int(row["vCPU"]),
|
|
73
|
+
memory=parse_memory(row["Memory"]),
|
|
74
|
+
gpu_count=parse_optional_count(row["GPU"]),
|
|
75
|
+
spot=False,
|
|
76
|
+
)
|
|
77
|
+
offers.append(offer)
|
|
78
|
+
self.fill_gpu_details(offers)
|
|
79
|
+
return self.add_spots(offers)
|
|
80
|
+
|
|
81
|
+
def skip(self, row: dict[str, str]) -> bool:
|
|
82
|
+
for key, values in self.filters.items():
|
|
83
|
+
if row[key] not in values:
|
|
84
|
+
return True
|
|
85
|
+
return False
|
|
86
|
+
|
|
87
|
+
def fill_gpu_details(self, offers: list[InstanceOffer]):
|
|
88
|
+
regions = defaultdict(list)
|
|
89
|
+
for offer in offers:
|
|
90
|
+
if offer.gpu_count > 0 and offer.instance_name not in self.preview_gpus:
|
|
91
|
+
regions[offer.location].append(offer.instance_name)
|
|
92
|
+
|
|
93
|
+
gpus = copy.deepcopy(self.preview_gpus)
|
|
94
|
+
while regions:
|
|
95
|
+
region = max(regions, key=lambda r: len(regions[r]))
|
|
96
|
+
instance_types = regions.pop(region)
|
|
97
|
+
|
|
98
|
+
logger.info("Fetching GPU details for %s", region)
|
|
99
|
+
client = boto3.client("ec2", region_name=region)
|
|
100
|
+
paginator = client.get_paginator("describe_instance_types")
|
|
101
|
+
for page in paginator.paginate(InstanceTypes=instance_types):
|
|
102
|
+
for i in page["InstanceTypes"]:
|
|
103
|
+
gpu = i["GpuInfo"]["Gpus"][0]
|
|
104
|
+
gpus[i["InstanceType"]] = (
|
|
105
|
+
gpu["Name"],
|
|
106
|
+
gpu["MemoryInfo"]["SizeInMiB"] / 1024,
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
regions = {
|
|
110
|
+
region: left
|
|
111
|
+
for region, names in regions.items()
|
|
112
|
+
if (left := [i for i in names if i not in instance_types])
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
for offer in offers:
|
|
116
|
+
if offer.gpu_count > 0:
|
|
117
|
+
offer.gpu_name, offer.gpu_memory = gpus[offer.instance_name]
|
|
118
|
+
|
|
119
|
+
def add_spots(self, offers: list[InstanceOffer]) -> list[InstanceOffer]:
|
|
120
|
+
region_instances = defaultdict(set)
|
|
121
|
+
for offer in offers:
|
|
122
|
+
region_instances[offer.location].add(offer.instance_name)
|
|
123
|
+
|
|
124
|
+
spot_prices = dict()
|
|
125
|
+
for region, instance_types in region_instances.items():
|
|
126
|
+
logger.info("Fetching spot prices for %s", region)
|
|
127
|
+
try:
|
|
128
|
+
client = boto3.client("ec2", region_name=region) # todo creds
|
|
129
|
+
pages = client.get_paginator("describe_spot_price_history").paginate(
|
|
130
|
+
Filters=[
|
|
131
|
+
{
|
|
132
|
+
"Name": "product-description",
|
|
133
|
+
"Values": ["Linux/UNIX"],
|
|
134
|
+
}
|
|
135
|
+
],
|
|
136
|
+
InstanceTypes=list(instance_types),
|
|
137
|
+
StartTime=datetime.datetime.utcnow(),
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
instance_prices = defaultdict(list)
|
|
141
|
+
for page in pages:
|
|
142
|
+
for item in page["SpotPriceHistory"]:
|
|
143
|
+
instance_prices[item["InstanceType"]].append(
|
|
144
|
+
float(item["SpotPrice"])
|
|
145
|
+
)
|
|
146
|
+
for (
|
|
147
|
+
instance_type,
|
|
148
|
+
zone_prices,
|
|
149
|
+
) in instance_prices.items(): # reduce zone prices to a single value
|
|
150
|
+
spot_prices[(instance_type, region)] = min(zone_prices)
|
|
151
|
+
except (ClientError, EndpointConnectionError) as e:
|
|
152
|
+
pass
|
|
153
|
+
|
|
154
|
+
spot_offers = []
|
|
155
|
+
for offer in offers:
|
|
156
|
+
if (
|
|
157
|
+
price := spot_prices.get((offer.instance_name, offer.location))
|
|
158
|
+
) is None:
|
|
159
|
+
continue
|
|
160
|
+
spot_offer = copy.deepcopy(offer)
|
|
161
|
+
spot_offer.spot = True
|
|
162
|
+
spot_offer.price = price
|
|
163
|
+
spot_offers.append(spot_offer)
|
|
164
|
+
return offers + spot_offers
|
|
165
|
+
|
|
166
|
+
@classmethod
|
|
167
|
+
def filter(cls, offers: list[InstanceOffer]) -> list[InstanceOffer]:
|
|
168
|
+
return [
|
|
169
|
+
i
|
|
170
|
+
for i in offers
|
|
171
|
+
if any(
|
|
172
|
+
i.instance_name.startswith(family)
|
|
173
|
+
for family in [
|
|
174
|
+
"t2.small",
|
|
175
|
+
"c5.",
|
|
176
|
+
"m5.",
|
|
177
|
+
"p3.",
|
|
178
|
+
"g5.",
|
|
179
|
+
"g4dn.",
|
|
180
|
+
"p4d.",
|
|
181
|
+
"p4de.",
|
|
182
|
+
]
|
|
183
|
+
)
|
|
184
|
+
]
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
def parse_memory(s: str) -> float:
|
|
188
|
+
r = re.match(r"^([0-9.]+) GiB$", s)
|
|
189
|
+
return float(r.group(1))
|
|
190
|
+
|
|
191
|
+
|
|
192
|
+
def parse_optional_count(s: str) -> int:
|
|
193
|
+
if not s:
|
|
194
|
+
return 0
|
|
195
|
+
return int(s)
|
|
@@ -0,0 +1,222 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
import re
|
|
5
|
+
import urllib.parse
|
|
6
|
+
from collections import namedtuple
|
|
7
|
+
from queue import Queue
|
|
8
|
+
from threading import Thread
|
|
9
|
+
from typing import Iterable, Optional, Tuple
|
|
10
|
+
|
|
11
|
+
import requests
|
|
12
|
+
from azure.core.credentials import TokenCredential
|
|
13
|
+
from azure.identity import DefaultAzureCredential
|
|
14
|
+
from azure.mgmt.compute import ComputeManagementClient
|
|
15
|
+
|
|
16
|
+
from gpuhunt._models import InstanceOffer
|
|
17
|
+
from gpuhunt.providers import AbstractProvider
|
|
18
|
+
|
|
19
|
+
logger = logging.getLogger(__name__)
|
|
20
|
+
prices_url = "https://prices.azure.com/api/retail/prices"
|
|
21
|
+
prices_version = "2023-01-01-preview"
|
|
22
|
+
prices_filters = [
|
|
23
|
+
"serviceName eq 'Virtual Machines'",
|
|
24
|
+
"priceType eq 'Consumption'",
|
|
25
|
+
"contains(productName, 'Windows') eq false",
|
|
26
|
+
"contains(productName, 'Dedicated') eq false",
|
|
27
|
+
"contains(meterName, 'Low Priority') eq false", # retires in 2025
|
|
28
|
+
]
|
|
29
|
+
VMSeries = namedtuple("VMSeries", ["pattern", "gpu_name", "gpu_memory"])
|
|
30
|
+
gpu_vm_series = [
|
|
31
|
+
VMSeries(r"NC(\d+)ads_A100_v4", "A100", 80.0), # NC A100 v4-series [A100 80GB]
|
|
32
|
+
VMSeries(r"NC(\d+)ads_A10_v4", "A10", 24.0), # NC A10 v4-series [A10]
|
|
33
|
+
VMSeries(r"NC(\d+)as_T4_v3", "T4", 16.0), # NCasT4_v3-series [T4]
|
|
34
|
+
VMSeries(r"NC(\d+)r?s_v3", "V100", 16.0), # NCv3-series [V100 16GB]
|
|
35
|
+
VMSeries(r"ND(\d+)amsr_A100_v4", "A100", 80.0), # NDm A100 v4-series [8xA100 80GB]
|
|
36
|
+
VMSeries(r"ND(\d+)asr_v4", "A100", 40.0), # ND A100 v4-series [8xA100 40GB]
|
|
37
|
+
VMSeries(r"ND(\d+)rs_v2", "V100", 32.0), # NDv2-series [8xV100 32GB]
|
|
38
|
+
VMSeries(r"NG(\d+)adm?s_V620_v1", "V620", None), # NGads V620-series [V620] # todo
|
|
39
|
+
VMSeries(r"NV(\d+)adm?s_A10_v5", "A10", 24.0), # NVadsA10 v5-series [A10]
|
|
40
|
+
VMSeries(r"NV(\d+)as_v4", "MI25", None), # NVv4-series [MI25] # todo
|
|
41
|
+
VMSeries(r"NV(\d+)s_v3", "M60", None), # NVv3-series [M60] # todo
|
|
42
|
+
]
|
|
43
|
+
# https://learn.microsoft.com/en-us/azure/virtual-machines/sizes-previous-gen
|
|
44
|
+
retired_vm_series = [
|
|
45
|
+
r"Basic_A(\d+)",
|
|
46
|
+
r"Standard_A(\d+)",
|
|
47
|
+
r"Standard_D(\d+)",
|
|
48
|
+
r"Standard_DC(\d+)s",
|
|
49
|
+
r"Standard_DS(\d+)",
|
|
50
|
+
r"Standard_F(\d+)",
|
|
51
|
+
r"Standard_F(\d+)s",
|
|
52
|
+
r"Standard_G(\d+)",
|
|
53
|
+
r"Standard_GS(\d+)",
|
|
54
|
+
r"Standard_L(\d+)s",
|
|
55
|
+
r"Standard_NC(\d+)r?",
|
|
56
|
+
r"Standard_NC(\d+)r?s_v2",
|
|
57
|
+
r"Standard_ND(\d+)r?s",
|
|
58
|
+
r"Standard_NV(\d+)",
|
|
59
|
+
r"Standard_NV(\d+)s_v2",
|
|
60
|
+
]
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
class AzureProvider(AbstractProvider):
|
|
64
|
+
def __init__(
|
|
65
|
+
self,
|
|
66
|
+
subscription_id: str,
|
|
67
|
+
credential: Optional[TokenCredential] = None,
|
|
68
|
+
cache_dir: Optional[str] = None,
|
|
69
|
+
):
|
|
70
|
+
self.cache_dir = cache_dir
|
|
71
|
+
self.client = ComputeManagementClient(
|
|
72
|
+
credential=credential or DefaultAzureCredential(),
|
|
73
|
+
subscription_id=subscription_id,
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
def get_pages(self, threads: int = 8) -> Iterable[list[dict]]:
|
|
77
|
+
q = Queue()
|
|
78
|
+
workers = [
|
|
79
|
+
Thread(target=self._get_pages_worker, args=(q, threads, i), daemon=True)
|
|
80
|
+
for i in range(threads)
|
|
81
|
+
]
|
|
82
|
+
for worker in workers:
|
|
83
|
+
worker.start()
|
|
84
|
+
|
|
85
|
+
exited = 0
|
|
86
|
+
while exited < threads:
|
|
87
|
+
page = q.get()
|
|
88
|
+
if page is None:
|
|
89
|
+
exited += 1
|
|
90
|
+
else:
|
|
91
|
+
yield page
|
|
92
|
+
q.task_done()
|
|
93
|
+
|
|
94
|
+
def _get_pages_worker(self, q: Queue, stride: int, worker_id: int):
|
|
95
|
+
page_id = worker_id
|
|
96
|
+
while True:
|
|
97
|
+
cached_page = None
|
|
98
|
+
if self.cache_dir is not None:
|
|
99
|
+
cached_page = os.path.join(self.cache_dir, f"{page_id:04}.json")
|
|
100
|
+
if cached_page is not None and os.path.exists(cached_page):
|
|
101
|
+
with open(cached_page, "r") as f:
|
|
102
|
+
data = json.load(f)
|
|
103
|
+
else:
|
|
104
|
+
logger.info("Worker %s fetches pricing page %s", worker_id, page_id)
|
|
105
|
+
data = requests.get(
|
|
106
|
+
prices_url,
|
|
107
|
+
params={
|
|
108
|
+
"api-version": "2023-01-01-preview",
|
|
109
|
+
"$filter": " and ".join(prices_filters),
|
|
110
|
+
"$skip": page_id * 100,
|
|
111
|
+
},
|
|
112
|
+
).json()
|
|
113
|
+
if cached_page is not None:
|
|
114
|
+
with open(cached_page, "w") as f:
|
|
115
|
+
json.dump(data, f)
|
|
116
|
+
|
|
117
|
+
if not data["Items"]:
|
|
118
|
+
q.put(None)
|
|
119
|
+
logger.info("Worker %s exited", worker_id)
|
|
120
|
+
return
|
|
121
|
+
q.put(data["Items"])
|
|
122
|
+
page_id += stride
|
|
123
|
+
|
|
124
|
+
def get(self) -> list[InstanceOffer]:
|
|
125
|
+
offers = []
|
|
126
|
+
for page in self.get_pages():
|
|
127
|
+
for item in page:
|
|
128
|
+
if is_retired(item["armSkuName"]):
|
|
129
|
+
continue
|
|
130
|
+
if not item["armSkuName"]:
|
|
131
|
+
continue
|
|
132
|
+
offer = InstanceOffer(
|
|
133
|
+
instance_name=item["armSkuName"],
|
|
134
|
+
location=item["armRegionName"],
|
|
135
|
+
price=item["retailPrice"],
|
|
136
|
+
spot="Spot" in item["meterName"],
|
|
137
|
+
)
|
|
138
|
+
offers.append(offer)
|
|
139
|
+
return self.fill_details(offers)
|
|
140
|
+
|
|
141
|
+
def fill_details(self, offers: list[InstanceOffer]) -> list[InstanceOffer]:
|
|
142
|
+
logger.info("Fetching instance details")
|
|
143
|
+
instances = {}
|
|
144
|
+
resources = self.client.resource_skus.list()
|
|
145
|
+
for resource in resources:
|
|
146
|
+
if resource.resource_type != "virtualMachines":
|
|
147
|
+
continue
|
|
148
|
+
if is_retired(resource.name):
|
|
149
|
+
continue
|
|
150
|
+
capabilities = {pair.name: pair.value for pair in resource.capabilities}
|
|
151
|
+
gpu_count, gpu_name, gpu_memory = 0, None, None
|
|
152
|
+
if "GPUs" in capabilities:
|
|
153
|
+
gpu_count = int(capabilities["GPUs"])
|
|
154
|
+
gpu_name, gpu_memory = get_gpu_name_memory(resource.name)
|
|
155
|
+
instances[resource.name] = InstanceOffer(
|
|
156
|
+
instance_name=resource.name,
|
|
157
|
+
cpu=capabilities["vCPUs"],
|
|
158
|
+
memory=float(capabilities["MemoryGB"]),
|
|
159
|
+
gpu_count=gpu_count,
|
|
160
|
+
gpu_name=gpu_name,
|
|
161
|
+
gpu_memory=gpu_memory,
|
|
162
|
+
)
|
|
163
|
+
with_details = []
|
|
164
|
+
without_details = []
|
|
165
|
+
for offer in offers:
|
|
166
|
+
if (resources := instances.get(offer.instance_name)) is None:
|
|
167
|
+
without_details.append(offer)
|
|
168
|
+
continue
|
|
169
|
+
offer.cpu = resources.cpu
|
|
170
|
+
offer.memory = resources.memory
|
|
171
|
+
offer.gpu_count = resources.gpu_count
|
|
172
|
+
offer.gpu_name = resources.gpu_name
|
|
173
|
+
offer.gpu_memory = resources.gpu_memory
|
|
174
|
+
with_details.append(offer)
|
|
175
|
+
return with_details + without_details
|
|
176
|
+
|
|
177
|
+
@classmethod
|
|
178
|
+
def filter(cls, offers: list[InstanceOffer]) -> list[InstanceOffer]:
|
|
179
|
+
vm_series = [
|
|
180
|
+
VMSeries(r"D(\d+)s_v3", None, None), # Dsv3-series
|
|
181
|
+
VMSeries(r"E(\d+)i?s_v4", None, None), # Esv4-series
|
|
182
|
+
VMSeries(r"E(\d+)-(\d+)s_v4", None, None), # Esv4-series (constrained vCPU)
|
|
183
|
+
VMSeries(r"NC(\d+)s_v3", "V100", 16 * 1024), # NCv3-series [V100 16GB]
|
|
184
|
+
VMSeries(r"NC(\d+)as_T4_v3", "T4", 16 * 1024), # NCasT4_v3-series [T4]
|
|
185
|
+
VMSeries(r"ND(\d+)rs_v2", "V100", 32 * 1024), # NDv2-series [8xV100 32GB]
|
|
186
|
+
VMSeries(
|
|
187
|
+
r"NV(\d+)adm?s_A10_v5", "A10", 24 * 1024
|
|
188
|
+
), # NVadsA10 v5-series [A10]
|
|
189
|
+
VMSeries(
|
|
190
|
+
r"NC(\d+)ads_A100_v4", "A100", 80 * 1024
|
|
191
|
+
), # NC A100 v4-series [A100 80GB]
|
|
192
|
+
VMSeries(
|
|
193
|
+
r"ND(\d+)asr_v4", "A100", 40 * 1024
|
|
194
|
+
), # ND A100 v4-series [8xA100 40GB]
|
|
195
|
+
VMSeries(
|
|
196
|
+
r"ND(\d+)amsr_A100_v4", "A100", 80 * 1024
|
|
197
|
+
), # NDm A100 v4-series [8xA100 80GB]
|
|
198
|
+
]
|
|
199
|
+
vm_series_pattern = re.compile(
|
|
200
|
+
f"^Standard_({'|'.join(series.pattern for series in vm_series)})$"
|
|
201
|
+
)
|
|
202
|
+
return [i for i in offers if vm_series_pattern.match(i.instance_name)]
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
def get_gpu_name_memory(vm_name: str) -> Tuple[Optional[str], Optional[float]]:
|
|
206
|
+
for pattern, gpu_name, gpu_memory in gpu_vm_series:
|
|
207
|
+
m = re.match(f"^Standard_{pattern}$", vm_name)
|
|
208
|
+
if m is None:
|
|
209
|
+
continue
|
|
210
|
+
if gpu_name == "A10" and vm_name.endswith("_v4"):
|
|
211
|
+
gpu_memory = gpu_memory * min(1.0, int(m.group(1)) / 16)
|
|
212
|
+
elif gpu_name == "A10" and vm_name.endswith("_v5"):
|
|
213
|
+
gpu_memory = gpu_memory * min(1.0, int(m.group(1)) / 36)
|
|
214
|
+
|
|
215
|
+
return gpu_name, gpu_memory
|
|
216
|
+
return None, None
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
def is_retired(name: str) -> bool:
|
|
220
|
+
if re.match(f"^({'|'.join(retired_vm_series)})$", name):
|
|
221
|
+
return True
|
|
222
|
+
return False
|