viz: move timeline layout to python (#10998)

* viz: move timeline layout to python

* DevEvent has a device and a name
This commit is contained in:
qazal
2025-06-27 13:06:00 +03:00
committed by GitHub
parent b4eb876d5a
commit a39343e39f
3 changed files with 48 additions and 43 deletions
+9 -9
View File
@@ -204,11 +204,11 @@ class TestVizProfiler(unittest.TestCase):
j = json.loads(get_profile(prof))
dev_events = j['devEvents']['NV']
dev_events = j['layout']['NV']['timeline']['shapes']
self.assertEqual(len(dev_events), 1)
event = dev_events[0]
self.assertEqual(event['name'], 'E_2')
self.assertEqual(event['ts'], 0)
self.assertEqual(event['st'], 0)
self.assertEqual(event['dur'], 10)
def test_perfetto_copy_node(self):
@@ -217,9 +217,9 @@ class TestVizProfiler(unittest.TestCase):
j = json.loads(get_profile(prof))
event = j['devEvents']['NV'][0]
event = j['layout']['NV']['timeline']['shapes'][0]
self.assertEqual(event['name'], 'COPYxx')
self.assertEqual(event['ts'], 900) # diff clock
self.assertEqual(event['st'], 900) # diff clock
self.assertEqual(event['dur'], 10)
def test_perfetto_graph(self):
@@ -232,19 +232,19 @@ class TestVizProfiler(unittest.TestCase):
j = json.loads(get_profile(prof))
devices = list(j['devEvents'])
devices = list(j['layout'])
self.assertEqual(devices[0], 'NV')
self.assertEqual(devices[1], 'NV:1')
nv_events = j['devEvents']['NV']
nv_events = j['layout']['NV']['timeline']['shapes']
self.assertEqual(nv_events[0]['name'], 'E_25_4n2')
self.assertEqual(nv_events[0]['ts'], 0)
self.assertEqual(nv_events[0]['st'], 0)
self.assertEqual(nv_events[0]['dur'], 2)
#self.assertEqual(j['devEvents'][6]['pid'], j['devEvents'][0]['pid'])
nv1_events = j['devEvents']['NV:1']
nv1_events = j['layout']['NV:1']['timeline']['shapes']
self.assertEqual(nv1_events[0]['name'], 'NV -> NV:1')
self.assertEqual(nv1_events[0]['ts'], 954)
self.assertEqual(nv1_events[0]['st'], 954)
#self.assertEqual(j['devEvents'][7]['pid'], j['devEvents'][3]['pid'])
if __name__ == "__main__":
+12 -24
View File
@@ -236,8 +236,7 @@ async function renderProfiler() {
displayGraph("profiler");
d3.select(".metadata").html("");
if (data != null) return;
const { devEvents, st, et } = await (await fetch("/get_profile")).json();
const events = new Map(Object.entries(devEvents));
const { layout, st, et } = await (await fetch("/get_profile")).json();
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
@@ -250,37 +249,26 @@ async function renderProfiler() {
// color by name
const nameMap = new Map();
data = [];
for (const [k, v] of events) {
if (v.length === 0) continue;
for (const [k, { timeline }] of Object.entries(layout)) {
if (timeline.shapes.length === 0) continue;
const div = deviceList.appendChild(document.createElement("div"));
div.id = k;
div.innerText = k;
div.style.padding = `${padding}px`;
const { y:baseY, height:baseHeight } = rect(`#${k}`);
// position events on the y axis, stack ones that overlap
const levels = [];
v.sort((a,b) => (a.ts-st) - (b.ts-st));
const levelHeight = baseHeight-padding;
const offsetY = baseY-canvasTop+padding/2;
for (const [i,e] of v.entries()) {
// assign to the first free depth
const start = e.ts-st;
const end = start+e.dur;
let depth = levels.findIndex(l => start >= l);
if (depth === -1) {
depth = levels.length;
levels.push(end);
} else levels[depth] = end;
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 }));
nameMap.set(e.name, { fillColor:colors[i%colors.length], label });
}
// offset y by depth
data.push({ x:start, dur:e.dur, name:e.name, height:levelHeight, y:offsetY+levelHeight*depth, kernel, ...nameMap.get(e.name) });
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 }));
nameMap.set(e.name, { fillColor:colors[i%colors.length], label });
}
// offset y by depth
data.push({ x:e.st-st, dur:e.dur, name:e.name, height:levelHeight, y:offsetY+levelHeight*e.depth, kernel, ...nameMap.get(e.name) });
}
// lastly, adjust device rect by number of levels
div.style.height = `${levelHeight*levels.length+padding}px`;
div.style.height = `${levelHeight*timeline.maxDepth+padding}px`;
}
// draw events on a timeline
const dpr = window.devicePixelRatio || 1;
+27 -10
View File
@@ -1,12 +1,12 @@
#!/usr/bin/env python3
import multiprocessing, pickle, difflib, os, threading, json, time, sys, webbrowser, socket, argparse, socketserver, functools
import multiprocessing, pickle, difflib, os, threading, json, time, sys, webbrowser, socket, argparse, socketserver, functools, decimal
from http.server import BaseHTTPRequestHandler
from urllib.parse import parse_qs, urlparse
from typing import Any, TypedDict, Generator
from tinygrad.helpers import colored, getenv, tqdm, unwrap, word_wrap, TRACEMETA
from tinygrad.uop.ops import TrackedGraphRewrite, UOp, Ops, lines, GroupOp, srender, sint
from tinygrad.renderer import ProgramSpec
from tinygrad.device import ProfileEvent, ProfileDeviceEvent, ProfileRangeEvent, ProfileGraphEvent
from tinygrad.device import ProfileEvent, ProfileDeviceEvent, ProfileRangeEvent, ProfileGraphEvent, ProfileGraphEntry
from tinygrad.dtype import dtypes
uops_colors = {Ops.LOAD: "#ffc0c0", Ops.STORE: "#87CEEB", Ops.CONST: "#e0e0e0", Ops.VCONST: "#e0e0e0", Ops.REDUCE: "#FF5B5B",
@@ -93,27 +93,44 @@ def get_details(ctx:TrackedGraphRewrite) -> Generator[GraphRewriteDetails, None,
# Profiler API
def events_to_json(profile:list[ProfileEvent]):
DevEvent = ProfileRangeEvent|ProfileGraphEntry
def flatten_events(profile:list[ProfileEvent]) -> Generator[tuple[decimal.Decimal, decimal.Decimal, DevEvent], None, None]:
for e in profile:
if isinstance(e, ProfileRangeEvent): yield (e.device, e.name, e.st, e.en, e.is_copy)
if isinstance(e, ProfileRangeEvent): yield (e.st, e.en, e)
if isinstance(e, ProfileGraphEvent):
for ent in e.ents: yield (ent.device, ent.name, e.sigs[ent.st_id], e.sigs[ent.en_id], ent.is_copy)
for ent in e.ents: yield (e.sigs[ent.st_id], e.sigs[ent.en_id], ent)
# timeline layout stacks events in a contiguous block. When a late starter finishes late, there is whitespace in the higher levels.
def timeline_layout(events:list[tuple[int, int, float, DevEvent]]) -> dict:
shapes:list[dict] = []
levels:list[int] = []
for st,et,dur,e in events:
if dur == 0: continue
# find a free level to put the event
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})
return {"shapes":shapes, "maxDepth":len(levels)}
def get_profile(profile:list[ProfileEvent]):
# start by getting the time diffs
devs = {e.device:(e.comp_tdiff, e.copy_tdiff if e.copy_tdiff is not None else e.comp_tdiff) for e in profile if isinstance(e,ProfileDeviceEvent)}
# map events per device
dev_events:dict[str, list] = {}
dev_events:dict[str, list[tuple[int, int, float, DevEvent]]] = {}
min_ts:int|None = None
max_ts:int|None = None
for device, name, ts, en, is_copy in events_to_json(profile):
time_diff = devs[device][is_copy]
for ts,en,e in flatten_events(profile):
time_diff = devs[e.device][e.__dict__.get("is_copy",False)] if e.device in devs else decimal.Decimal(0)
st = int(ts+time_diff)
et = st if en is None else int(en+time_diff)
dev_events.setdefault(device,[]).append({"name":name, "ts":st, "dur":et-st})
dev_events.setdefault(e.device,[]).append((st, et, float(en-ts), e))
if min_ts is None or st < min_ts: min_ts = st
if max_ts is None or et > max_ts: max_ts = et
return json.dumps({"devEvents":dev_events, "st":min_ts, "et":max_ts}).encode("utf-8")
# return layout of per device events
for events in dev_events.values(): events.sort(key=lambda v:v[0])
dev_layout = {k:{"timeline":timeline_layout(v)} for k,v in dev_events.items()}
return json.dumps({"layout":dev_layout, "st":min_ts, "et":max_ts}).encode("utf-8")
# ** HTTP server