put match times in viz (#12544)

* put match times in viz

* float
This commit is contained in:
qazal
2025-10-09 06:56:10 +03:00
committed by GitHub
parent 51420d1f99
commit baab7e334d
2 changed files with 9 additions and 9 deletions
+7 -7
View File
@@ -869,11 +869,11 @@ match_stats:dict[UPat, list[int|float]] = dict()
@dataclass(frozen=True)
class TrackedGraphRewrite:
loc:tuple[str, int] # location that called graph_rewrite
sink:int # the sink input to graph_rewrite
matches:list[tuple[int, int, tuple]] # before/after UOp, UPat location
name:str|None # optional name of the rewrite
depth:int # depth if it's a subrewrite
loc:tuple[str, int] # location that called graph_rewrite
sink:int # the sink input to graph_rewrite
matches:list[tuple[int, int, tuple, float]] # before/after UOp, UPat location and time
name:str|None # optional name of the rewrite
depth:int # depth if it's a subrewrite
bottom_up:bool
tracked_keys:list[TracingKey] = []
@@ -945,14 +945,14 @@ class TrackedPatternMatcher(PatternMatcher):
try: ret = match(uop, ctx)
except Exception:
if TRACK_MATCH_STATS >= 2 and active_rewrites:
active_rewrites[-1].matches.append((track_uop(uop), track_uop(UOp(Ops.REWRITE_ERROR, src=uop.src, arg=str(sys.exc_info()[1]))), p.location))
active_rewrites[-1].matches.append((track_uop(uop), track_uop(UOp(Ops.REWRITE_ERROR,src=uop.src,arg=str(sys.exc_info()[1]))),p.location,0))
raise
if ret is not None and ret is not uop:
match_stats[p][0] += 1
match_stats[p][3] += (et:=time.perf_counter()-st)
if TRACK_MATCH_STATS >= 3: print(f"{et*1e6:7.2f} us -- ", printable(p.location))
if TRACK_MATCH_STATS >= 2 and isinstance(ret, UOp) and active_rewrites:
active_rewrites[-1].matches.append((track_uop(uop), track_uop(ret), p.location))
active_rewrites[-1].matches.append((track_uop(uop), track_uop(ret), p.location, et))
return ret
match_stats[p][2] += time.perf_counter()-st
return None
+2 -2
View File
@@ -97,12 +97,12 @@ def _reconstruct(a:int, i:int):
def get_details(ctx:TrackedGraphRewrite, i:int=0) -> Generator[GraphRewriteDetails, None, None]:
yield {"graph":uop_to_json(next_sink:=_reconstruct(ctx.sink, i)), "uop":str(next_sink), "changed_nodes":None, "diff":None, "upat":None}
replaces: dict[UOp, UOp] = {}
for u0_num,u1_num,upat_loc in tqdm(ctx.matches):
for u0_num,u1_num,upat_loc,dur in tqdm(ctx.matches):
replaces[u0:=_reconstruct(u0_num, i)] = u1 = _reconstruct(u1_num, i)
try: new_sink = next_sink.substitute(replaces)
except RuntimeError as e: new_sink = UOp(Ops.NOOP, arg=str(e))
yield {"graph":(sink_json:=uop_to_json(new_sink)), "uop":str(new_sink), "changed_nodes":[id(x) for x in u1.toposort() if id(x) in sink_json],
"diff":list(difflib.unified_diff(str(u0).splitlines(), str(u1).splitlines())), "upat":(upat_loc, printable(upat_loc))}
"diff":list(difflib.unified_diff(str(u0).splitlines(),str(u1).splitlines())), "upat":(upat_loc,printable(upat_loc)+f"\n{dur*1e6:.2f} us")}
if not ctx.bottom_up: next_sink = new_sink
# encoder helpers