|
1 | 1 | # This file contains wrapper functions for running |
2 | 2 | # 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 |
6 | 4 |
|
7 | 5 | gpus = ["T4", "L4", "A100-80GB", "H100!", "B200", "H200"] |
8 | 6 | for gpu in gpus: |
9 | 7 | gpu_slug = gpu.lower().split("-")[0].strip("!") |
10 | 8 | app.function(gpu=gpu, image=cuda_image, name=f"run_cuda_script_{gpu_slug}", serialized=True)( |
11 | 9 | modal_run_config |
12 | 10 | ) |
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)( |
14 | 13 | modal_run_config |
15 | 14 | ) |
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