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 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
@@ -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