diff --git a/test/test_jit.py b/test/test_jit.py index f229a8b7c0..b6d0b0e2e4 100644 --- a/test/test_jit.py +++ b/test/test_jit.py @@ -718,8 +718,6 @@ class TestJitGraphSplit(unittest.TestCase): def ji_xfer(self): return {"type": "xfer"} def test_jit_split_simple(self): - if Device.DEFAULT == "REMOTE": raise unittest.SkipTest("REMOTE gpu is broken") - @TinyJit def f(inp): op0 = self.compute(Device.DEFAULT, inp) @@ -792,9 +790,9 @@ class TestJitGraphSplit(unittest.TestCase): multigraph=[self.ji_graph(5)], hcqgraph=[self.ji_graph(5)]) - @unittest.skip("flaky") def test_jit_multidev_xfer(self): if Device.DEFAULT in {"CPU", "LLVM"}: raise unittest.SkipTest("CPU/LLVM is not a valid default device for this test (zero-copies)") + if Device.DEFAULT == "METAL" or REAL_DEV == "METAL": raise unittest.SkipTest("Metal is flaky, with multidevice (same as metal llama 4gpu?)") try: Device[f"{Device.DEFAULT}:1"] except Exception: raise unittest.SkipTest("no multidevice") @@ -819,7 +817,6 @@ class TestJitGraphSplit(unittest.TestCase): @unittest.skipIf(getenv("MOCKGPU"), "MockGPU does not support parallel copies") def test_jit_multidev_copy(self): if Device.DEFAULT in {"CPU", "LLVM"}: raise unittest.SkipTest("CPU/LLVM is not a valid default device for this test (zero-copies)") - if Device.DEFAULT == "REMOTE": raise unittest.SkipTest("REMOTE gpu is broken") @TinyJit def f(inp): diff --git a/tinygrad/runtime/ops_remote.py b/tinygrad/runtime/ops_remote.py index 4c3b748a63..4227f274cc 100644 --- a/tinygrad/runtime/ops_remote.py +++ b/tinygrad/runtime/ops_remote.py @@ -9,6 +9,7 @@ from typing import Callable, Iterator, Any, cast from collections import defaultdict from dataclasses import dataclass, field, replace import multiprocessing, threading, functools, itertools, asyncio, http, http.client, hashlib, time, os, binascii, struct, ast, contextlib, weakref +import traceback, builtins from tinygrad.renderer import Renderer, ProgramSpec from tinygrad.dtype import DTYPES_DICT, dtypes from tinygrad.uop.ops import UOp, Ops, Variable, sint @@ -38,6 +39,11 @@ class RemoteProperties: graph_supports_multi: bool ib_gid: bytes|None +@dataclass(frozen=True) +class RemoteException: + exc: Exception + trace: str = "" + @dataclass(frozen=True) class GetProperties(RemoteRequest): pass @@ -119,9 +125,10 @@ class GraphExec(RemoteRequest): wait: bool # for safe deserialization +eval_excs = [v for k,v in builtins.__dict__.items() if isinstance(v, type) and issubclass(v, Exception) and not k.endswith("Warning")] eval_globals = {x.__name__:x for x in [SessionKey, SessionFree, RemoteProperties, GetProperties, Event, Wait, BufferAlloc, BufferOffset, BufferIOVAS, BufferFree, CopyIn, CopyOut, Transfer, BatchTransfer, IBConnect, ProgramAlloc, ProgramFree, ProgramExec, - GraphComputeItem, GraphAlloc, GraphFree, GraphExec, BufferSpec, UOp, Ops, dtypes]} + GraphComputeItem, GraphAlloc, GraphFree, GraphExec, BufferSpec, UOp, Ops, dtypes, RemoteException] + eval_excs} attribute_whitelist: dict[Any, set[str]] = {dtypes: {*DTYPES_DICT.keys(), 'imagef', 'imageh'}, Ops: {x.name for x in Ops}} eval_fxns = {ast.Constant: lambda x: x.value, ast.Tuple: lambda x: tuple(map(safe_eval, x.elts)), ast.List: lambda x: list(map(safe_eval, x.elts)), ast.Dict: lambda x: {safe_eval(k):safe_eval(v) for k,v in zip(x.keys, x.values)}, @@ -182,7 +189,10 @@ class RemoteHandler: key, value = hdr.split(':', 1) req_headers[key.lower()] = value.strip() req_body = await reader.readexactly(int(req_headers.get("content-length", "0"))) - res_status, res_body = await self.handle(req_method, req_path, req_body) + try: res_status, res_body = await self.handle(req_method, req_path, req_body) + except Exception as e: + res_status, res_body = http.HTTPStatus.INTERNAL_SERVER_ERROR, repr(RemoteException(e, traceback.format_exc())).encode() + print(f"{traceback.format_exc()}", flush=True) writer.write(f"HTTP/1.1 {res_status.value} {res_status.phrase}\r\nContent-Length: {len(res_body)}\r\n\r\n".encode() + res_body) async def ib_connect(self, ssession:SessionKey, dsession:SessionKey) -> IBConn|None: @@ -415,10 +425,14 @@ class RemoteConnection: for conn,data in datas.items(): conn.conn.request("POST", "/batch", data) for conn in datas.keys(): response = conn.conn.getresponse() - assert response.status == 200, f"POST /batch failed: {response}" resp = response.read() + conn.req = BatchRequest() # no matter what response, reset conn + if response.status == http.HTTPStatus.INTERNAL_SERVER_ERROR: + exc_wrapper = safe_eval(ast.parse(resp.decode(), mode="eval").body) + exc_wrapper.exc.add_note(exc_wrapper.trace) + raise exc_wrapper.exc + assert response.status == http.HTTPStatus.OK, f"POST /batch failed: {resp.decode()}" if conn == self: ret = resp - conn.req = BatchRequest() if take_q: RemoteConnection.q_lock.release() return ret