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
This commit is contained in:
qazal
2025-07-04 20:35:25 +03:00
committed by GitHub
parent 3d8569f6d8
commit 41aa54eb5a
3 changed files with 15 additions and 15 deletions
+4 -7
View File
@@ -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) => {
+1 -3
View File
@@ -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;
+10 -5
View File
@@ -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<ISSUE GETTING SHAPE>"
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: