forked from tinygrad/tinygrad
viz: switch to TracingKey (#11100)
* viz: switch to TracingKey * tuple * order is name, keys, fmt * add test_tracing_key
This commit is contained in:
+13
-2
@@ -1,7 +1,7 @@
|
||||
import unittest, decimal, json
|
||||
from dataclasses import dataclass
|
||||
|
||||
from tinygrad.uop.ops import UOp, UPat, Ops, PatternMatcher, TrackedPatternMatcher
|
||||
from tinygrad.uop.ops import UOp, UPat, Ops, PatternMatcher, TrackedPatternMatcher, TracingKey
|
||||
from tinygrad.uop.ops import graph_rewrite, track_rewrites, TRACK_MATCH_STATS
|
||||
from tinygrad.uop.symbolic import sym
|
||||
from tinygrad.dtype import dtypes
|
||||
@@ -94,14 +94,25 @@ class TestViz(unittest.TestCase):
|
||||
lst = get_viz_list()
|
||||
self.assertEqual(lst[0]["name"], "name_default n1")
|
||||
|
||||
# name can also come from a function
|
||||
# name can also come from a function that returns a string
|
||||
def test_dyn_name_fxn(self):
|
||||
@track_rewrites(name=lambda a,ret: a.render())
|
||||
def name_from_fxn(s:UOp): return graph_rewrite(s, PatternMatcher([]))
|
||||
name_from_fxn(UOp.variable("a", 1, 10)+1)
|
||||
lst = get_viz_list()
|
||||
# name gets deduped by the function call counter
|
||||
self.assertEqual(lst[0]["name"], "(a+1) n1")
|
||||
|
||||
# name can also come from a function that returns a TracingKey
|
||||
def test_tracing_key(self):
|
||||
@track_rewrites(name=lambda inp,ret: TracingKey("custom_name", fmt=f"input={inp.render()}"))
|
||||
def test(s:UOp): return graph_rewrite(s, PatternMatcher([]))
|
||||
test(UOp.variable("a", 1, 10)+1)
|
||||
lst = get_viz_list()
|
||||
# NOTE: names from TracingKey do not get deduped
|
||||
self.assertEqual(lst[0]["name"], "custom_name")
|
||||
self.assertEqual(lst[0]["kernel_code"], "input=(a+1)")
|
||||
|
||||
def test_colored_label(self):
|
||||
# NOTE: dataclass repr prints literal escape codes instead of unicode chars
|
||||
@dataclass(frozen=True)
|
||||
|
||||
@@ -3,7 +3,7 @@ import time, pprint
|
||||
from dataclasses import dataclass, replace, field
|
||||
from tinygrad.helpers import all_same, colored, DEBUG, GlobalCounters, ansilen, BEAM, NOOPT, all_int, CAPTURING, Metadata, TRACEMETA
|
||||
from tinygrad.helpers import DEVECTORIZE, time_to_str, VALIDATE_WITH_CPU, getenv
|
||||
from tinygrad.uop.ops import Ops, PatternMatcher, UOp, UPat, Variable, sym_infer, graph_rewrite, print_uops, track_rewrites
|
||||
from tinygrad.uop.ops import Ops, PatternMatcher, UOp, UPat, Variable, sym_infer, graph_rewrite, print_uops, track_rewrites, TracingKey
|
||||
from tinygrad.device import Device, Buffer
|
||||
from tinygrad.renderer import Renderer, ProgramSpec, Estimates
|
||||
from tinygrad.engine.schedule import ScheduleItem
|
||||
@@ -13,7 +13,7 @@ from tinygrad.uop.spec import type_verify
|
||||
|
||||
# **************** Program Creation ****************
|
||||
|
||||
@track_rewrites(name=lambda _ast,_renderer,ret:ret)
|
||||
@track_rewrites(name=lambda _ast,_renderer,ret: TracingKey(ret.name, (ret.function_name, ret.ast), ret.src))
|
||||
def get_program(ast:UOp, renderer:Renderer) -> ProgramSpec:
|
||||
"""
|
||||
Transform an AST into a ProgramSpec. May trigger BEAM search.
|
||||
|
||||
+6
-5
@@ -772,6 +772,7 @@ class TrackedGraphRewrite:
|
||||
class TracingKey:
|
||||
display_name:str # display name of this trace event
|
||||
keys:tuple[str, ...]=() # optional keys to search for related traces
|
||||
fmt:str|None=None # optional detailed formatting
|
||||
cat:str|None=None # optional category to color this by
|
||||
|
||||
tracked_keys:list[Any] = []
|
||||
@@ -785,11 +786,11 @@ if getenv("CAPTURE_PROCESS_REPLAY"):
|
||||
def save_to_diskcache():
|
||||
for k,v in replay_capture.items(): diskcache_put("process_replay", k, v, prepickled=True)
|
||||
|
||||
def track_rewrites(name:Callable|bool=True):
|
||||
def track_rewrites(name:Callable[..., str|TracingKey]|bool=True):
|
||||
def _decorator(func):
|
||||
def __wrapper(*args, **kwargs):
|
||||
if TRACK_MATCH_STATS >= 2:
|
||||
tracked_keys.append((fn:=func.__name__)+f" n{next(_name_cnt.setdefault(fn, itertools.count(1)))}")
|
||||
tracked_keys.append(TracingKey((fn:=func.__name__)+f" n{next(_name_cnt.setdefault(fn, itertools.count(1)))}", cat=fn))
|
||||
tracked_ctxs.append([])
|
||||
# late import!
|
||||
from tinygrad.device import cpu_profile
|
||||
@@ -797,9 +798,9 @@ def track_rewrites(name:Callable|bool=True):
|
||||
ret = func(*args, **kwargs)
|
||||
if TRACK_MATCH_STATS >= 2 and callable(name):
|
||||
name_ret = name(*args, **kwargs, ret=ret)
|
||||
tracked_keys[-1] = key = tracked_keys[-1].replace(fn, name_ret) if isinstance(name_ret, str) else name_ret
|
||||
if isinstance(key, str): e.name = TracingKey(key, (key,), func.__name__)
|
||||
else: e.name = TracingKey(f"{func.__name__} for {name_ret.name}", (name_ret.name,), func.__name__)
|
||||
assert isinstance(name_ret, TracingKey) or isinstance(name_ret, str)
|
||||
tracked_keys[-1] = k = TracingKey(n:=tracked_keys[-1].display_name.replace(fn, name_ret), (n,)) if isinstance(name_ret, str) else name_ret
|
||||
e.name = TracingKey(k.display_name if isinstance(name_ret, str) else f"{func.__name__} for {k.display_name}", k.keys, cat=func.__name__)
|
||||
if getenv("CAPTURE_PROCESS_REPLAY"):
|
||||
# find the unittest frame we're capturing in
|
||||
frm = sys._getframe(1)
|
||||
|
||||
@@ -5,7 +5,6 @@ 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, TracingKey, UOp, Ops, printable, GroupOp, srender, sint
|
||||
from tinygrad.renderer import ProgramSpec
|
||||
from tinygrad.device import ProfileEvent, ProfileDeviceEvent, ProfileRangeEvent, ProfileGraphEvent, ProfileGraphEntry, ProfilePointEvent
|
||||
from tinygrad.dtype import dtypes
|
||||
|
||||
@@ -22,13 +21,12 @@ 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]:
|
||||
def get_metadata(keys:list[TracingKey], contexts:list[list[TrackedGraphRewrite]]) -> list[dict]:
|
||||
ret = []
|
||||
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]
|
||||
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
|
||||
for key in k.keys: ref_map[key] = i
|
||||
ret.append({"name":k.display_name, "kernel_code":k.fmt, "steps":steps})
|
||||
return ret
|
||||
|
||||
# ** Complete rewrite details for a graph_rewrite call
|
||||
|
||||
Reference in New Issue
Block a user