forked from tinygrad/tinygrad
prereq viz cleanups for unique profile keys (#17649)
* cleaner * just use VIZ=-2 * better
This commit is contained in:
@@ -504,7 +504,7 @@ jobs:
|
||||
- name: Run AMD renderer tests (AMD:LLVM)
|
||||
run: DEV=MOCKKFD+AMD:LLVM python -m pytest -n=auto test/amd/ --durations 20
|
||||
- name: Run SQTT profiling tests
|
||||
run: PROFILE=1 SQTT=1 python3 -m pytest -n=auto test/amd/test_sqtt_profiler.py
|
||||
run: VIZ=-2 python3 -m pytest -n=auto test/amd/test_sqtt_profiler.py
|
||||
- name: Run AMD emulated tests on NULL backend
|
||||
env:
|
||||
AMD: 0
|
||||
|
||||
@@ -522,12 +522,16 @@ class TestVizIntegration(unittest.TestCase):
|
||||
def one(A:UOp): return A[0].store(UOp.const(1.0, dtypes.float)).sink(arg=KernelInfo(kernel_name))
|
||||
def zero(A:UOp): return A[0].store(UOp.const(0.0, dtypes.float)).sink(arg=KernelInfo(kernel_name))
|
||||
with save_viz() as viz:
|
||||
Tensor.custom_kernel(Tensor.empty(4, device="NULL"), fxn=one)[0].realize()
|
||||
Tensor.custom_kernel(Tensor.empty(4, device="NULL"), fxn=zero)[0].realize()
|
||||
ctx_refs = [i for i,c in enumerate(viz.list_items()) if c["name"] == kernel_name]
|
||||
@TinyJit
|
||||
def f(a:Tensor, b:Tensor): return Tensor.custom_kernel(a, fxn=one)[0], Tensor.custom_kernel(b, fxn=zero)[0]
|
||||
a, b = Tensor.empty(4, device="NULL"), Tensor.empty(4, device="NULL")
|
||||
# warmup
|
||||
for _ in range(2): Tensor.realize(*f(a, b))
|
||||
Tensor.realize(*f(a, b))
|
||||
kernels = {i for i,c in enumerate(viz.list_items()) if c["name"] == kernel_name}
|
||||
profile = decode_profile(unwrap(get_profile(viz.data, cpu_events)))
|
||||
events = [e for e in profile["layout"]["NULL"]["events"] if e["name"] == kernel_name]
|
||||
self.assertEqual([e["ref"] for e in events], ctx_refs)
|
||||
self.assertEqual({e["ref"] for e in events}, kernels)
|
||||
|
||||
from tinygrad.device import ProfileDeviceEvent, ProfileGraphEvent, ProfileGraphEntry
|
||||
from tinygrad.viz.serve import get_profile
|
||||
|
||||
@@ -459,7 +459,7 @@ pm_to_program = PatternMatcher([
|
||||
(UPat(Ops.PROGRAM, src=(UPat(), UPat(Ops.LINEAR), UPat(Ops.SOURCE, name="source")), name="prg"), do_compile),
|
||||
])
|
||||
|
||||
@rewrite_group(name=lambda ast,renderer,ret,**kwargs: TracingKey(ret.src[0].arg.name,(ret.src[0].arg.function_name, ast), ret=renderer), replay=True)
|
||||
@rewrite_group(name=lambda ast,renderer,ret,**_: TracingKey((k:=ret.src[0].arg).name,(k.function_name, ast),ret=renderer), replay=True)
|
||||
@Context(ALLOW_DEVICE_USAGE=0)
|
||||
def do_to_program(ast:UOp, renderer:Renderer) -> UOp:
|
||||
"""
|
||||
|
||||
Reference in New Issue
Block a user