diff --git a/test/unit/test_viz.py b/test/unit/test_viz.py index 0d0044f21b..c33f4dbcc0 100644 --- a/test/unit/test_viz.py +++ b/test/unit/test_viz.py @@ -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__": diff --git a/tinygrad/viz/js/index.js b/tinygrad/viz/js/index.js index 63416fca16..29a4265d91 100644 --- a/tinygrad/viz/js/index.js +++ b/tinygrad/viz/js/index.js @@ -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; diff --git a/tinygrad/viz/serve.py b/tinygrad/viz/serve.py index a2cd46c819..7b58df7b83 100755 --- a/tinygrad/viz/serve.py +++ b/tinygrad/viz/serve.py @@ -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