From 41aa54eb5ad1bd61da7394e94873186f40270f04 Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Fri, 4 Jul 2025 20:35:25 +0300 Subject: [PATCH] viz: resolve all graph references in python (#11087) * viz: resolve all graph references in python * it just maps things to the index * always map the name * key on the uop * diff * close --- tinygrad/viz/js/index.js | 11 ++++------- tinygrad/viz/js/worker.js | 4 +--- tinygrad/viz/serve.py | 15 ++++++++++----- 3 files changed, 15 insertions(+), 15 deletions(-) diff --git a/tinygrad/viz/js/index.js b/tinygrad/viz/js/index.js index 60f4f99a5b..3b714974f0 100644 --- a/tinygrad/viz/js/index.js +++ b/tinygrad/viz/js/index.js @@ -120,8 +120,6 @@ async function renderProfiler() { const canvas = profiler.append("canvas").attr("id", "timeline").node(); if (profileRet == null) profileRet = await (await fetch("/get_profile")).json() const { layout, st, et } = profileRet; - const kernelMap = new Map(); - for (const [i, c] of ctxs.entries()) kernelMap.set(c.function_name, { name:c.name, i }); // place devices on the y axis and set vertical positions const [tickSize, padding] = [10, 8]; deviceList.style.paddingTop = `${tickSize+padding}px`; @@ -143,13 +141,12 @@ async function renderProfiler() { const levelHeight = baseHeight-padding; const offsetY = baseY-canvasTop+padding/2; for (const [i,e] of timeline.shapes.entries()) { - const kernel = kernelMap.get(e.name); if (!nameMap.has(e.name)) { - const label = parseColors(kernel?.name ?? e.name).map(({ color, st }) => ({ color, st, width:ctx.measureText(st).width })); + const label = parseColors(e.name).map(({ color, st }) => ({ color, st, width:ctx.measureText(st).width })); nameMap.set(e.name, { fillColor:colors[i%colors.length], label }); } // offset y by depth - data.shapes.push({ x:e.st-st, dur:e.dur, name:e.name, height:levelHeight, y:offsetY+levelHeight*e.depth, kernel, ...nameMap.get(e.name) }); + data.shapes.push({ x:e.st-st, dur:e.dur, name:e.name, height:levelHeight, y:offsetY+levelHeight*e.depth, ref:e.ref, ...nameMap.get(e.name) }); } // position shapes on the canvas and scale to fit fixed area const startY = offsetY+(levelHeight*timeline.maxDepth)+padding/2; @@ -205,7 +202,7 @@ async function renderProfiler() { const width = xscale(e.x+e.dur)-x; ctx.fillStyle = e.fillColor; ctx.fillRect(x, e.y, width, e.height); - rectLst.push({ y0:e.y, y1:e.y+e.height, x0:x, x1:x+width, ref:e.kernel?.i, tooltipText:formatTime(e.dur) }); + rectLst.push({ y0:e.y, y1:e.y+e.height, x0:x, x1:x+width, ref:e.ref, tooltipText:formatTime(e.dur) }); // add label ctx.textAlign = "left"; ctx.textBaseline = "middle"; @@ -407,7 +404,7 @@ function setCtxWithHistory(newCtx) { // NOTE: browser does a structured clone, passing a mutable object is safe. history.replaceState(state, ""); history.pushState(state, ""); - setState({ expandSteps:true, currentCtx:newCtx, currentStep:0, currentRewrite:0 }); + setState({ expandSteps:true, currentCtx:newCtx+1, currentStep:0, currentRewrite:0 }); } window.addEventListener("popstate", (e) => { diff --git a/tinygrad/viz/js/worker.js b/tinygrad/viz/js/worker.js index 92e1b50239..3f78445f70 100644 --- a/tinygrad/viz/js/worker.js +++ b/tinygrad/viz/js/worker.js @@ -10,15 +10,13 @@ onmessage = (e) => { g.setGraph({ rankdir: "LR" }).setDefaultEdgeLabel(function() { return {}; }); if (additions.length !== 0) g.setNode("addition", {label:"", style:"fill: rgba(26, 27, 38, 0.5);", padding:0}); for (let [k, {label, src, ref, ...rest }] of Object.entries(graph)) { - const idx = ref ? ctxs.findIndex(k => k.ref === ref) : -1; - if (idx != -1) label += `\ncodegen@${ctxs[idx].function_name}`; // adjust node dims by label size (excluding escape codes) + add padding let [width, height] = [0, 0]; for (line of label.replace(/\u001B\[(?:K|.*?m)/g, "").split("\n")) { width = Math.max(width, ctx.measureText(line).width); height += LINE_HEIGHT; } - g.setNode(k, {width:width+NODE_PADDING*2, height:height+NODE_PADDING*2, padding:NODE_PADDING, label, ref:idx==-1 ? null : idx, ...rest}); + g.setNode(k, {width:width+NODE_PADDING*2, height:height+NODE_PADDING*2, padding:NODE_PADDING, label, ref, ...rest}); // add edges const edgeCounts = {} for (const s of src) edgeCounts[s] = (edgeCounts[s] || 0)+1; diff --git a/tinygrad/viz/serve.py b/tinygrad/viz/serve.py index 285c0264d0..16c471d2f9 100755 --- a/tinygrad/viz/serve.py +++ b/tinygrad/viz/serve.py @@ -21,12 +21,14 @@ uops_colors = {Ops.LOAD: "#ffc0c0", Ops.STORE: "#87CEEB", Ops.CONST: "#e0e0e0", # ** Metadata for a track_rewrites scope +ref_map:dict[Any, int] = {} def get_metadata(keys:list[Any], contexts:list[list[TrackedGraphRewrite]]) -> list[dict]: ret = [] - for k,v in zip(keys, contexts): + for i,(k,v) in enumerate(zip(keys, contexts)): steps = [{"name":s.name, "loc":s.loc, "depth":s.depth, "match_count":len(s.matches), "code_line":printable(s.loc)} for s in v] - if isinstance(k, ProgramSpec): ret.append({"name":k.name, "kernel_code":k.src, "ref":id(k.ast), "function_name":k.function_name, "steps":steps}) - else: ret.append({"name":str(k), "steps":steps}) + for key in (refs:=[k.name, k.function_name, k.ast] if isinstance(k, ProgramSpec) else [str(k)]): ref_map[key] = i + ret.append({"name":refs[0], "steps":steps}) + if isinstance(k, ProgramSpec): ret[-1]["kernel_code"] = k.src return ret # ** Complete rewrite details for a graph_rewrite call @@ -68,10 +70,11 @@ def uop_to_json(x:UOp) -> dict[int, dict]: label += f"\n{shape_to_str(u.shape)}" except Exception: label += "\n" + if (ref:=ref_map.get(u.arg.ast) if u.op is Ops.KERNEL else None) is not None: label += f"\ncodegen@{ctxs[ref]['name']}" # NOTE: kernel already has metadata in arg if TRACEMETA >= 2 and u.metadata is not None and u.op is not Ops.KERNEL: label += "\n"+repr(u.metadata) graph[id(u)] = {"label":label, "src":[id(x) for x in u.src if x not in excluded], "color":uops_colors.get(u.op, "#ffffff"), - "ref":id(u.arg.ast) if u.op is Ops.KERNEL else None, "tag":u.tag} + "ref":ref, "tag":u.tag} return graph @functools.cache @@ -111,7 +114,9 @@ def timeline_layout(events:list[tuple[int, int, float, DevEvent]]) -> dict: depth = next((i for i,level_et in enumerate(levels) if st>=level_et), len(levels)) if depth < len(levels): levels[depth] = et else: levels.append(et) - shapes.append({"name":e.name, "st":st, "dur":dur, "depth":depth}) + name = e.name + if (ref:=ref_map.get(name)) is not None: name = ctxs[ref]["name"] + shapes.append({"name":name, "ref":ref, "st":st, "dur":dur, "depth":depth}) return {"shapes":shapes, "maxDepth":len(levels)} def mem_layout(events:list[tuple[int, int, float, DevEvent]]) -> dict: