From 795559b3d4fa486a3c2a8261e9b82eefef07ee3a Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Tue, 30 Jun 2026 16:06:32 +0800 Subject: [PATCH] viz/cli: support passing only a start marker (#16802) * failing test * support end interval --- test/null/test_viz.py | 2 ++ tinygrad/viz/cli.py | 5 +++-- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/test/null/test_viz.py b/test/null/test_viz.py index e1a0d6afce..b6adf9c30b 100644 --- a/test/null/test_viz.py +++ b/test/null/test_viz.py @@ -1111,8 +1111,10 @@ class TestCLI(unittest.TestCase): with write_files(viz) as files, Context(NO_COLOR=1): flat = run_cli(*files, "-s", "NULL", "--interval", "interval_start", "interval_end") aggregate = run_cli(*files, "-s", "NULL", "--interval", "interval_start", "interval_end", "-t") + final = run_cli(*files, "-s", "NULL", "--interval", "interval_end", "-t") self.assertEqual([s["name"] for s in flat], ["interval_start", "target_1", "target_2", "interval_end"]) self.assertEqual(sorted(s["name"] for s in aggregate), ["target_1", "target_2"]) + assert all(s["name"].startswith("post_") for s in final), f"post_* kernels must be present in final, got {final}" if __name__ == "__main__": unittest.main() diff --git a/tinygrad/viz/cli.py b/tinygrad/viz/cli.py index c1aab8b385..1657e76f3e 100755 --- a/tinygrad/viz/cli.py +++ b/tinygrad/viz/cli.py @@ -155,7 +155,8 @@ 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] markers = profile.get("markers", []) - interval:tuple[int, int]|None = None if not args.interval else (marker_st(markers, args.interval[0]), marker_st(markers, args.interval[1])) + interval:tuple[int, int]|None = None + if (rng:=args.interval): interval = (marker_st(markers, rng[0]), marker_st(markers, rng[1]) if len(rng) > 1 else profile["dur"]) def produce_top_kernels() -> Iterator[dict]: 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) @@ -219,7 +220,7 @@ def get_arg_parser() -> argparse.ArgumentParser: 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("--interval", nargs=2, metavar=("START", "END"), help="Optional start and end marker") + parser.add_argument("--interval", nargs="+", metavar=("START", "END"), help="Optional start and end marker") 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))