diff --git a/tinygrad/viz/index.html b/tinygrad/viz/index.html index 7616788726..a80c89dc35 100644 --- a/tinygrad/viz/index.html +++ b/tinygrad/viz/index.html @@ -307,6 +307,7 @@ var currentKernel = -1; var currentRewrite = 0; var expandKernel = true; + const evtSources = []; async function main() { const mainContainer = document.querySelector('.main-container'); // ***** LHS kernels list @@ -357,13 +358,23 @@ }); // ***** UOp graph if (currentKernel == -1) return; - const cacheKey = `${currentKernel}-${currentUOp}`; + const kernel = kernels[currentKernel][1][currentUOp]; + const cacheKey = `kernel=${currentKernel}&idx=${currentUOp}`; + // close any pending event sources + let activeSrc = null; + for (const e of evtSources) { + if (e.url.split("?")[1] !== cacheKey) e.close(); + else if (e.readyState === EventSource.OPEN) activeSrc = e; + } if (cacheKey in cache) { ret = cache[cacheKey]; } - else { + // if we don't have a complete cache yet we start streaming this kernel + if (!(cacheKey in cache) || (cache[cacheKey].length !== kernel.match_count+1 && activeSrc == null)) { ret = []; + cache[cacheKey] = ret; const eventSource = new EventSource(`/kernels?kernel=${currentKernel}&idx=${currentUOp}`); + evtSources.push(eventSource); eventSource.onmessage = (e) => { if (e.data === "END") return eventSource.close(); const chunk = JSON.parse(e.data); @@ -374,14 +385,12 @@ const gUl = document.getElementById(`rewrite-${ret.length-1}`); if (gUl != null) gUl.classList.remove("disabled"); }; - cache[cacheKey] = ret; } if (ret.length === 0) return; renderGraph(ret[currentRewrite].graph, ret[currentRewrite].changed_nodes || []); // ***** RHS metadata const metadata = document.querySelector(".container.metadata"); metadata.innerHTML = ""; - const kernel = kernels[currentKernel][1][currentUOp]; metadata.appendChild(vsCodeOpener(kernel.loc.join(":").split("/"))); metadata.appendChild(highlightedCodeBlock(kernel.code_line, "python", true)); // ** code blocks diff --git a/tinygrad/viz/serve.py b/tinygrad/viz/serve.py index a8ff0c6126..6be4da4007 100755 --- a/tinygrad/viz/serve.py +++ b/tinygrad/viz/serve.py @@ -132,16 +132,19 @@ class Handler(BaseHTTPRequestHandler): if "kernel" in (query:=parse_qs(url.query)): def getarg(k:str,default=0): return int(query[k][0]) if k in query else default kidx, ridx = getarg("kernel"), getarg("idx") - # stream details - self.send_response(200) - self.send_header("Content-Type", "text/event-stream") - self.send_header("Cache-Control", "no-cache") - self.end_headers() - for r in get_details(contexts[0][kidx], contexts[1][kidx][ridx]): - self.wfile.write(f"data: {json.dumps(r)}\n\n".encode("utf-8")) - self.wfile.flush() - self.wfile.write("data: END\n\n".encode("utf-8")) - return self.wfile.flush() + try: + # stream details + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Cache-Control", "no-cache") + self.end_headers() + for r in get_details(contexts[0][kidx], contexts[1][kidx][ridx]): + self.wfile.write(f"data: {json.dumps(r)}\n\n".encode("utf-8")) + self.wfile.flush() + self.wfile.write("data: END\n\n".encode("utf-8")) + return self.wfile.flush() + # pass if client closed connection + except (BrokenPipeError, ConnectionResetError): return ret, content_type = json.dumps(kernels).encode(), "application/json" elif url.path == "/get_profile" and perfetto_profile is not None: ret, content_type = perfetto_profile, "application/json" else: status_code = 404