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