Skip to content

Commit 562ff9e

Browse files
committed
Feat: enable B200
1 parent 059ea0f commit 562ff9e

2 files changed

Lines changed: 12 additions & 31 deletions

File tree

src/discord-cluster-manager/modal_runner.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,8 @@
1919
tag = f"{cuda_version}-{flavor}-{operating_sys}"
2020

2121
# Move this to another file later:
22+
23+
2224
cuda_image = (
2325
Image.from_registry(f"nvidia/cuda:{tag}", add_python="3.11")
2426
.apt_install(
@@ -50,13 +52,20 @@
5052
.pip_install("requests")
5153
)
5254

55+
5356
cuda_image = cuda_image.add_local_python_source(
5457
"consts",
5558
"modal_runner",
5659
"modal_runner_archs",
5760
"run_eval",
5861
)
5962

63+
cuda_image_b200 = (
64+
Image.from_registry("nvidia/cuda:12.8.0-devel-ubuntu24.04", add_python="3.11")
65+
.pip_install("ninja", "packaging", "requests")
66+
.pip_install("torch==2.7.0", extra_index_url="https://download.pytorch.org/whl/cu128")
67+
)
68+
6069

6170
class TimeoutException(Exception):
6271
pass
Lines changed: 3 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -1,42 +1,14 @@
11
# This file contains wrapper functions for running
22
# Modal apps on specific devices. We will fix this later.
3-
from modal_runner import app, cuda_image, modal_run_config
4-
from modal_utils import deserialize_full_result
5-
from run_eval import FullResult, SystemInfo
3+
from modal_runner import app, cuda_image, cuda_image_b200, modal_run_config
64

75
gpus = ["T4", "L4", "A100-80GB", "H100!", "B200", "H200"]
86
for gpu in gpus:
97
gpu_slug = gpu.lower().split("-")[0].strip("!")
108
app.function(gpu=gpu, image=cuda_image, name=f"run_cuda_script_{gpu_slug}", serialized=True)(
119
modal_run_config
1210
)
13-
app.function(gpu=gpu, image=cuda_image, name=f"run_pytorch_script_{gpu_slug}", serialized=True)(
11+
img = cuda_image if gpu != "B200" else cuda_image_b200
12+
app.function(gpu=gpu, image=img, name=f"run_pytorch_script_{gpu_slug}", serialized=True)(
1413
modal_run_config
1514
)
16-
17-
18-
@app.function(image=cuda_image, max_containers=1, timeout=600)
19-
def run_pytorch_script_b200(config: dict, timeout: int = 300):
20-
"""Send a config and timeout to the server and return the response."""
21-
import requests
22-
23-
ip_addr = "34.59.196.5"
24-
port = "33001"
25-
26-
payload = {"config": config, "timeout": timeout}
27-
28-
try:
29-
response = requests.post(f"http://{ip_addr}:{port}", json=payload, timeout=timeout + 5)
30-
response.raise_for_status()
31-
print("ORIGINAL", response.json())
32-
33-
print("DESERIALIZED", deserialize_full_result(response.json()))
34-
return deserialize_full_result(response.json())
35-
except requests.RequestException as e:
36-
return FullResult(success=False, error=str(e), runs={}, system=SystemInfo())
37-
38-
39-
@app.local_entrypoint()
40-
def test_b200(timeout: int = 300):
41-
config = {}
42-
run_pytorch_script_b200.remote(config, timeout)

0 commit comments

Comments
 (0)