diff --git a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh index de9f641120..17884d2627 100755 --- a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh +++ b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh @@ -3,4 +3,4 @@ export BENCHMARK=5 export EVAL_BS=0 VIZ=${VIZ:--1} FULL_LAYERS=1 DEBUG=0 examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh SRC="AMD"; [[ $DEV == NULL* ]] && SRC="NULL" -python -m tinygrad.viz.cli -s "$SRC" --top 20 +python -m tinygrad.viz.cli -s "$SRC" -t diff --git a/tinygrad/viz/README b/tinygrad/viz/README index 3227159b52..858b22f4cd 100644 --- a/tinygrad/viz/README +++ b/tinygrad/viz/README @@ -30,11 +30,11 @@ user story: viewing code * schedule 2 (97) = main.py:97 * schedule 3 (10) = main.py:145 * web: click "schedule 1", get list of kernels (like DEBUG=2) -* cli: `python -m tinygrad.viz.cli -s TINY -i "Schedule 3 Kernels n1"` +* cli: `python -m tinygrad.viz.cli -s TINY "Schedule 3 Kernels n1"` * kernel 1 "E_34_34" -- 'sin' * kernel 2 "R_4545" * web: click "E_34_34" -* cli: `python -m tinygrad.viz.cli -s TINY -i "do_to_program for E_34_34" "initial symbolic"` +* cli: `python -m tinygrad.viz.cli -s TINY "do_to_program for E_34_34" "initial symbolic"` * pre-rewritten UOp graph (step through rewrite here) * post-rewritten UOp graph * UOp list diff --git a/tinygrad/viz/cli.py b/tinygrad/viz/cli.py index ff86664475..ac00abed03 100755 --- a/tinygrad/viz/cli.py +++ b/tinygrad/viz/cli.py @@ -83,11 +83,11 @@ def main(args) -> None: profile = decode_profile(profile_bytes) profile["layout"].update([(f'{c["name"][5:]}{" SQTT" if s["name"].endswith("PKTS") else ""} {s["name"]}', s["data"]) for c in viz_data.ctxs if c["name"].startswith("SQTT") for s in c["steps"] if s["name"].endswith(("PMC", "PKTS"))]) - if args.list and args.src == "ALL": return print("ALL\n"+"\n".join(fmt_colored(k) for k in profile["layout"])) + if args.list and not args.src: return print("ALL\n"+"\n".join(fmt_colored(k) for k in profile["layout"])) # ** SQTT printer - data = None if args.src == "ALL" else get(profile["layout"], args.src) - if "SQTT" in args.src: + data = None if not args.src else get(profile["layout"], args.src[0]) + if args.src and "SQTT" in args.src[0]: # modern terminals support 24-bit color def hex_colored(st:str, color:str) -> str: return f"\x1b[38;2;{int(color[1:3],16)};{int(color[3:5],16)};{int(color[5:7],16)}m{st}\x1b[0m" print(f"{'Clk':<12} {'Unit':<20} {'Op':<22} {'Dur':<4} {'Delay':<4} {'Info'}") @@ -119,7 +119,7 @@ def main(args) -> None: print(emit(row, lambda _: f"{row['clk']:<12} {unit:<20} {op_str}{' '*(22-ansilen(op_str))} {row['dur']:<4} {str(row['delay']):<4} {info}")) # ** PMC printer - elif "PMC" in args.src: + elif args.src and "PMC" in args.src[0]: pmc = viz.unpack_pmc(unwrap(data)) pmc_fmt:list[str] = [] for name,val,*detail in pmc["rows"]: @@ -138,7 +138,7 @@ def main(args) -> None: else: timelines = [(n,l) for n,l in profile["layout"].items() if isinstance(l, dict) and l.get("event_type") == 0] def produce_top_kernels() -> Iterator[dict]: - tagged = ((n,e) for n,l in timelines for e in l["events"]) if args.src == "ALL" else ((args.src,e) for e in unwrap(data)["events"]) + tagged = ((n,e) for n,l in timelines for e in l["events"]) if not args.src else ((args.src[0],e) for e in unwrap(data)["events"]) agg:dict[tuple[str,str], tuple[float, int, int|None, dict[str, float]]] = {} # map (device, kernel name) to (total time, count, ref, est) est_keys = ("FLOPS", "B/s mem", "B/s lds") total = 0 @@ -149,18 +149,18 @@ def main(args) -> None: agg[(dev,e["name"])] = (t+et, c+1, e["ref"], est) total += et items = sorted(agg.items(), key=lambda kv:kv[1][0], reverse=True) - num_rows = len(items) if args.top < 0 else args.top + num_rows = len(items) if args.t < 0 else args.t for (dev,name),(t,c,ref,est) in items[:num_rows]: - display = f"{dev[:7]:7s} {fmt_colored(name)}" if args.src == "ALL" else fmt_colored(name) + display = f"{dev[:7]:7s} {fmt_colored(name)}" if not args.src else fmt_colored(name) yield {"name":display, "dur_ms":t, "count":c, "pct":t/total*100.0, "ref":ref, "fmt":{k:int(est[k]/(t*1e-3)) for k in est_keys if k in est}} if num_rows > 0 and items[num_rows:]: other_t = sum(t for _,(t,_,_,_) in items[num_rows:]) other_c = sum(c for _,(_,c,_,_) in items[num_rows:]) yield {"name":"Other", "dur_ms":other_t, "count":other_c, "pct":other_t/total*100.0, "ref":None, "fmt":None} def produce_all_kernels() -> Iterator[dict]: - event_streams = [[(e["st"], n, e) for e in l["events"]] for n,l in timelines] if args.src == "ALL" \ - else [[(e["st"], args.src, e) for e in unwrap(data)["events"]]] - if args.src == "ALL": + event_streams = [[(e["st"], n, e) for e in l["events"]] for n,l in timelines] if not args.src \ + else [[(e["st"], args.src[0], e) for e in unwrap(data)["events"]]] + if not args.src: for n,l in profile["layout"].items(): if not isinstance(l, dict) or l.get("event_type") != 0: yield {"device":"SOURCE", "name":n, "st_ms":0, "ref":None, "ext":None} marker_stream = sorted([(m["ts"], "MARKER", m) for m in profile.get("markers", [])], key=lambda t:t[0]) @@ -185,7 +185,7 @@ def main(args) -> None: ptm = colored(time_to_str(k["dur_ms"]*1e-3, w=9), "yellow" if k["dur_ms"] > 10 else None) name = f"*** {k['device'][:7]:7s} "+k["name"]+" "*(46-ansilen(k["name"])) return f"{name} tm {ptm}/{k['st_ms']:9.2f}ms"+(f" ({fmt_data(k['fmt'])})" if k["fmt"] else "") - fmt_row = fmt_top if args.top else fmt_all + fmt_row = fmt_top if args.t else fmt_all seen_refs:set[int] = set() def render_event(k:dict, ls=args.list) -> None: print(emit(k, to_str=fmt_row)) @@ -196,29 +196,25 @@ def main(args) -> None: if DEBUG >= 4 and s["name"] == "View Source": print_step(s) if DEBUG >= 5 or ls: print(emit(" "*s["depth"]+s["name"]+(f" - {s['match_count']}" if s.get('match_count', 0) else ''))) if DEBUG >= 6: print_step(s) - if DEBUG >= 7 or (args.item and len(args.item) > 1 and s["name"] == args.item[1]): print_step(s, reconstruct_matches=True) + if DEBUG >= 7 or (len(args.src) > 2 and s["name"] == args.src[2]): print_step(s, reconstruct_matches=True) elif DEBUG >= 3 and k.get("ext"): print(emit(k["ext"])) - produce = produce_top_kernels if args.top else produce_all_kernels - if args.item: - if len(args.item) > 2: raise RuntimeError(f"-i takes at most 2 names (got {args.item})") - k = get({r["name"]:r for r in produce()}, args.item[0]) + produce = produce_top_kernels if args.t else produce_all_kernels + if len(args.src) > 1: + k = get({r["name"]:r for r in produce()}, args.src[1]) with Context(DEBUG=max(DEBUG.value, 3)): render_event(k, ls=True) else: for k in produce(): render_event(k) def get_arg_parser() -> argparse.ArgumentParser: - parser = argparse.ArgumentParser(add_help=False, prog="python -m tinygrad.viz.cli") - g_opts = parser.add_argument_group("optional args") - g_opts.add_argument("-s", "--src", type=str, default="ALL", metavar="NAME", help="Select a data source (default: ALL)") - g_opts.add_argument("-i", "--item", nargs="+", default=None, metavar="NAME", help="Select an item within the source (default: list all items)") - g_opts.add_argument("--list", "--ls", dest="list", action="store_true", help="List sources") - g_opts.add_argument("-t", "--top", nargs="?", type=int, const=20, metavar="COUNT", help="Aggregate top kernels (optional count, default 20)") - g_opts.add_argument("--profile-path", type=str, metavar="PATH", help="Optional path to profile.pkl (default: latest profile)", + parser = argparse.ArgumentParser(prog="python -m tinygrad.viz.cli") + parser.add_argument("-s", "--src", nargs="+", default=[], metavar="NAME", help="Select a data source (default: all)") + parser.add_argument("--list", "--ls", dest="list", action="store_true", help="List sources") + parser.add_argument("-t", nargs="?", type=int, const=20, metavar="COUNT", help="Aggregate top kernels (optional count, default 20)") + parser.add_argument("--profile-path", type=str, metavar="PATH", help="Optional path to profile.pkl (default: latest profile)", default=temp("profile.pkl", append_user=True)) - g_opts.add_argument("--rewrites-path", type=str, metavar="PATH", help="Optional path to rewrites.pkl (default: latest rewrites)", + parser.add_argument("--rewrites-path", type=str, metavar="PATH", help="Optional path to rewrites.pkl (default: latest rewrites)", default=temp("rewrites.pkl", append_user=True)) - g_opts.add_argument("--json", action="store_true", help="Emit profiler output as JSON") - g_opts.add_argument("-h", "--help", action="help", help="show this help message and exit") + parser.add_argument("--json", action="store_true", help="Emit profiler output as JSON") return parser if __name__ == "__main__":