From 2d53e47b1448b7cf2219bfb08c64d2a6d30faa16 Mon Sep 17 00:00:00 2001
From: qazal <77887910+Qazalin@users.noreply.github.com>
Date: Sun, 15 Sep 2024 14:08:55 +0800
Subject: [PATCH] refactor viz saved context (prereq for tree view) (#6516)
* more styling
* warns
* refactor viz ctx to dataclass
* meh, fine for now
* name ctx
* allow smaller zooms
* more work
* fixup ctx.diffs
---
tinygrad/ops.py | 13 +++++++++----
viz/index.html | 25 +++++++++++--------------
viz/serve.py | 33 +++++++++++++++++----------------
3 files changed, 37 insertions(+), 34 deletions(-)
diff --git a/tinygrad/ops.py b/tinygrad/ops.py
index 8418552313..23386676e6 100644
--- a/tinygrad/ops.py
+++ b/tinygrad/ops.py
@@ -721,7 +721,12 @@ class PatternMatcher:
TRACK_MATCH_STATS = getenv("TRACK_MATCH_STATS", 2 if getenv("VIZ") else 0)
match_stats:Dict[UPat, List[Union[int, float]]] = dict()
-contexts: List[Tuple[Tuple[str, int], UOp, List[Tuple[UOp, UOp, str]]]] = []
+@dataclass(frozen=True)
+class TrackedRewriteContext:
+ loc: str # location that called graph_rewrite
+ sink: UOp # the sink passed into the rewrite
+ rewrites: List[Tuple[UOp, UOp, str]] # all rewrites of sparents. (before, after, UPat printable)
+contexts: List[TrackedRewriteContext] = []
class TrackedPattenMatcher(PatternMatcher):
def __init__(self, patterns:List[Tuple[UPat, Callable]]):
super().__init__(patterns)
@@ -742,7 +747,7 @@ class TrackedPattenMatcher(PatternMatcher):
match_stats[p][2] += (et:=time.perf_counter()-st)
match_stats[p][3] += et
if TRACK_MATCH_STATS >= 3: print(f"{et*1e6:7.2f} us -- ", p.printable())
- if TRACK_MATCH_STATS >= 2: contexts[-1][2].append((uop, ret, p.printable()))
+ if TRACK_MATCH_STATS >= 2: contexts[-1].rewrites.append((uop, ret, p.printable()))
return ret # NOTE: if it returns None, we keep trying to match
match_stats[p][2] += time.perf_counter()-st
return None
@@ -760,7 +765,7 @@ if TRACK_MATCH_STATS:
print(f"{ret[0]:6d} / {ret[1]:7d} -- {ret[3]*1000.:9.2f} / {ret[2]*1000.:9.2f} ms -- TOTAL")
if TRACK_MATCH_STATS >= 2:
with open("/tmp/rewrites.pkl", "wb") as f:
- print(f"rewrote {len(contexts)} graphs and applied {sum(len(x[2]) for x in contexts)} rules, saved to /tmp/rewrites.pkl")
+ print(f"rewrote {len(contexts)} graphs and applied {sum(len(x.rewrites) for x in contexts)} rules, saved to /tmp/rewrites.pkl")
pickle.dump(contexts, f)
if getenv("VIZ"):
import viz.serve
@@ -782,5 +787,5 @@ class RewriteContext:
self.nodes[replace_source] = self.replace[n] = found = self.rewrite(new_x) if (new_x := self.pm.rewrite(x)) else x
return found
def graph_rewrite(sink:UOp, pm:PatternMatcher) -> UOp:
- if TRACK_MATCH_STATS >= 2: contexts.append((get_location(), sink, []))
+ if TRACK_MATCH_STATS >= 2: contexts.append(TrackedRewriteContext(f"{(l:=get_location())[0].split('/')[-1]}:{l[1]}", sink, []))
return RewriteContext(pm).rewrite(sink)
diff --git a/viz/index.html b/viz/index.html
index 4341eab17b..a588a85a33 100644
--- a/viz/index.html
+++ b/viz/index.html
@@ -68,8 +68,6 @@
}
.metadata {
grid-column: span 3;
- display: flex;
- flex-direction: column;
overflow-y: auto;
}
.uop-list {
@@ -102,9 +100,8 @@
}
.code-block {
max-height: 30%;
- min-height: 30%;
- background-color: #191919;
overflow-y: auto;
+ background-color: #191919;
border-radius: 8px;
padding: 8px;
}
@@ -127,12 +124,12 @@
var totalRewrites = 0;
async function main() {
const ret = await (await fetch("/"+currentUOp)).json()
- const [location, graphs, blocks, rest] = ret;
+ const [ctx, rest] = ret;
totalUOps = rest.length-1;
- totalRewrites = graphs.length-1;
+ totalRewrites = ctx.graphs.length-1;
// graph
const g = new dagreD3.graphlib.Graph().setGraph({ rankdir: "LR" }).setDefaultEdgeLabel(function() { return {}; });
- const [graph, pattern] = graphs[currentRewrite];
+ const graph = ctx.graphs[currentRewrite];
for ([k,u] of Object.entries(graph)) {
g.setNode(k, {label: u[0], style: `fill: ${u[4]}; rx: 8; ry: 8;` });
for (src of u[2]) {
@@ -142,7 +139,7 @@
const svg = d3.select("svg");
const inner = svg.select("g");
var zoom = d3.zoom()
- .scaleExtent([0.25, 2])
+ .scaleExtent([0.05, 2])
.on("zoom", () => {
const transform = d3.event.transform;
inner.attr("transform", transform);
@@ -153,21 +150,21 @@
// metadata
const container = document.querySelector(".container.metadata");
container.innerHTML = "";
- container.appendChild(Object.assign(document.createElement("pre"), { textContent: location }));
- blocks.forEach((b) => {
+ container.appendChild(Object.assign(document.createElement("pre"), { textContent: ctx.loc }));
+ ctx.extra.forEach((b) => {
if (b.length == 0) return;
const pre = Object.assign(document.createElement("pre"), { innerHTML: `${b}`, className: "code-block" });
container.appendChild(pre);
})
- if (graphs.length > 1) {
+ if (ctx.graphs.length > 1) {
const rewriteCounter = Object.assign(document.createElement("div"), { className: "rewrite-counter" });
container.appendChild(rewriteCounter)
- graphs.forEach((g, i) => {
+ ctx.graphs.forEach((g, i) => {
const rewriteDiv = Object.assign(document.createElement("div"), { textContent: i, className: "uop-el" });
if (i === currentRewrite) {
rewriteDiv.classList.add("active-uop-el");
- if (g[1] != null) {
- const [pattern, diff] = g[1]
+ if (i !== 0) {
+ const [pattern, diff] = ctx.diffs[i-1];
container.appendChild(Object.assign(document.createElement("pre"), { innerHTML: `${pattern}`, className: "wrap" }));
const diffHtml = diff.map((line) => {
if (line.startsWith("+")) return `${line}`;
diff --git a/viz/serve.py b/viz/serve.py
index 33dd4a0bfb..371cc58ec7 100755
--- a/viz/serve.py
+++ b/viz/serve.py
@@ -1,12 +1,12 @@
#!/usr/bin/env python3
from dataclasses import asdict, dataclass
-from typing import Dict, List, Optional, Tuple
+from typing import Dict, List, Tuple
import pickle, re, os, sys, time, threading, webbrowser, json, difflib
from tinygrad.codegen.uopgraph import linearize_uop
from tinygrad.device import Device
from tinygrad.engine.realize import get_runner
from tinygrad.helpers import getenv
-from tinygrad.ops import UOp, UOps
+from tinygrad.ops import TrackedRewriteContext, UOp, UOps
from tinygrad.engine.graph import uops_colors, word_wrap
from http.server import HTTPServer, BaseHTTPRequestHandler
@@ -37,22 +37,24 @@ def uop_to_prg(ast:UOp) -> str:
@dataclass(frozen=True)
class UOpRet:
- loc: str
- graphs: List[Tuple[Dict[int, Tuple[str, str, List[int], str, str]], Optional[Tuple[str, List[str]]]]]
- extra: List[str]
+ loc: str # location that called graph_rewrite
+ graphs: List[Dict[int, Tuple[str, str, List[int], str, str]]] # a seralized version of UOp graphs
+ diffs: List[Tuple[str, List[str]]] # the diffs for each rewrite
+ extra: List[str] # these become code blocks in the UI
-def create_graph(ctx:Tuple[Tuple[str, int], UOp, List[Tuple[UOp, UOp, str]]]) -> UOpRet:
- loc, start, matches = ctx
- graphs: List[Tuple[Dict, Optional[Tuple[str, List[str]]]]] = [(uop_to_json(start), None)]
- for first, rewritten, pattern in matches:
+def create_graph(ctx:TrackedRewriteContext) -> UOpRet:
+ graphs = [uop_to_json(ctx.sink)]
+ diffs = []
+ for first, rewritten, pattern in ctx.rewrites:
diff = list(difflib.unified_diff(str(first).splitlines(), str(rewritten).splitlines()))
- graph = {**graphs[-1][0], **uop_to_json(rewritten)}
+ graph = {**graphs[-1], **uop_to_json(rewritten)}
for k,v in graph.copy().items():
if any(x == id(first) for x in v[2]):
graph[k] = v[:2]+([id(rewritten) if x == id(first) else x for x in v[2]],)+v[3:]
if k == id(first): del graph[k]
- graphs.append((graph, (pattern, diff)))
- return UOpRet(f"{loc[0].split('/')[-1]}:{loc[1]}", graphs, [str(start), uop_to_prg(start)] if start.op is UOps.SINK else [str(start)])
+ graphs.append(graph)
+ diffs.append((pattern, diff))
+ return UOpRet(ctx.loc, graphs, diffs, [str(ctx.sink), uop_to_prg(ctx.sink)] if ctx.sink.op is UOps.SINK else [str(ctx.sink)])
class Handler(BaseHTTPRequestHandler):
def do_GET(self):
@@ -72,11 +74,10 @@ class Handler(BaseHTTPRequestHandler):
self.send_response(200)
self.send_header("Content-type", "application/json")
self.end_headers()
- with open("/tmp/rewrites.pkl", "rb") as f: contexts = pickle.load(f)
- # TOOD: unify this loc_str logic
- rest = [f"{x[0][0].split('/')[-1]}:{x[0][1]}" for x in contexts]
+ with open("/tmp/rewrites.pkl", "rb") as f: contexts: List[TrackedRewriteContext] = pickle.load(f)
+ rest = [x.loc for x in contexts]
current_graph = create_graph(contexts[int(self.path.split("/")[-1])])
- ret = json.dumps(tuple(asdict(current_graph).values())+(rest,)).encode()
+ ret = json.dumps((asdict(current_graph), rest)).encode()
else:
self.send_response(404)
ret = b""