mirror of
https://github.com/tinygrad/tinygrad.git
synced 2026-08-29 12:56:07 +00:00
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:
@@ -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) => {
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user