sqtt: decode each instruction exec (#13093)

* sqtt: decode each instruction exec

* start tests

* run_asm

* capture sqtt per kernel

* chaining vgprs

* test things

* inst_execs in viz

* can also configure l and g

* 1l + cleanup

* test_sleep

* test_wmma

* work

* test sleep with llvm builtin
This commit is contained in:
qazal
2025-11-05 17:30:27 +08:00
committed by GitHub
parent 54141e9cb9
commit 8119d9f082
3 changed files with 110 additions and 5 deletions
+15 -2
View File
@@ -40,12 +40,21 @@ class InstInfo:
def on_ev(self, ev):
self.hit, self.lat, self.stall = self.hit + 1, self.lat + ev.duration, self.stall + ev.stall
@dataclasses.dataclass(frozen=True)
class InstExec:
typ:str
inst:str
stall:int
dur:int
time:int
class _ROCParseCtx:
def __init__(self, dev_evs:dict[str, ProfileDeviceEvent], sqtt_evs:list[ProfileSQTTEvent], prog_evs:list[ProfileProgramEvent]):
self.dev_evs, self.sqtt_evs, self.prog_evs = dev_evs, iter(sqtt_evs), prog_evs
self.wave_events:dict[tuple[str, int, int, int], dict[int, InstInfo]] = {}
self.disasms:dict[int, tuple[str, int]] = {}
self.addr2prg:dict[int, ProfileProgramEvent] = {}
self.inst_execs:dict[tuple[str, int, int, int], list[InstExec]] = {}
for prog in prog_evs:
for addr, info in llvm_disasm(dev_evs[prog.device].arch, unwrap(prog.lib)).items():
@@ -66,14 +75,18 @@ class _ROCParseCtx:
if DEBUG >= 5: print("WAVE", ev.wave_id, self.active_se, ev.cu, ev.simd, ev.contexts, ev.begin_time, ev.end_time)
asm:dict[int, InstInfo] = {}
inst_execs:list[InstExec] = []
for j in range(ev.instructions_size):
inst_ev = ev.instructions_array[j]
inst_typ = rocprof.rocprofiler_thread_trace_decoder_inst_category_t__enumvalues[inst_ev.category]
asm.setdefault(inst_ev.pc.address, InstInfo(typ=inst_typ, inst=self.disasms[inst_ev.pc.address][0]))
inst_disasm = self.disasms[inst_ev.pc.address][0]
asm.setdefault(inst_ev.pc.address, InstInfo(typ=inst_typ, inst=inst_disasm))
asm[inst_ev.pc.address].on_ev(inst_ev)
inst_execs.append(InstExec(inst_typ, inst_disasm, inst_ev.stall, inst_ev.duration, inst_ev.time))
if ev.instructions_size > 0:
self.wave_events[(self.find_program(ev.instructions_array[0].pc.address).name, ev.wave_id, ev.cu, ev.simd)] = asm
self.wave_events[key:=(self.find_program(ev.instructions_array[0].pc.address).name, ev.wave_id, ev.cu, ev.simd)] = asm
self.inst_execs[key] = inst_execs
def decode(profile:list[ProfileEvent]) -> _ROCParseCtx:
dev_events:dict[str, ProfileDeviceEvent] = {}
+91
View File
@@ -0,0 +1,91 @@
import os
os.environ["PYTHONPATH"] = "."
os.environ["SQTT"] = "1"
os.environ["AMD"] = "1"
os.environ["VIZ"] = "1"
os.environ["AMD_LLVM"] = "0"
import unittest
import sys
from tinygrad import Tensor
from tinygrad.dtype import dtypes
from tinygrad.renderer import ProgramSpec
from tinygrad.uop.ops import UOp, Ops, KernelInfo
from tinygrad.engine.realize import CompiledRunner
from tinygrad.device import Device, ProfileDeviceEvent
from extra.sqtt.roc import decode, InstExec
dev = Device["AMD"]
def get_sqtt(asm:list[str], l:int=1, g:int=1) -> list[InstExec]:
# clear the old traces
dev.profile_events.clear()
# setup custom_kernel
name = sys._getframe(1).f_code.co_name
def fxn(_):
L = UOp.special(l, "lidx0")
G = UOp.special(g, "gidx0")
ops:list[str] = [UOp(Ops.CUSTOM, arg="asm volatile (")]
for inst in asm: ops.append(UOp(Ops.CUSTOM, src=(ops[-1],), arg=f' "{inst}\\n\\t"'))
ops.append(UOp(Ops.CUSTOM, src=(ops[-1],), arg=");"))
return UOp.sink(*ops, L, G, arg=KernelInfo(name=name))
k = Tensor.custom_kernel(Tensor.empty(1), fxn=fxn)[0]
# exec and decode sqtt
k.realize()
rctx = decode(dev.profile_events+[ProfileDeviceEvent("AMD", arch=dev.device_info())])
assert len(rctx.inst_execs) > 0, "empty sqtt output"
return list(rctx.inst_execs.values())[0][:-1]
class TestTiming(unittest.TestCase):
def test_v_add(self):
sqtt = get_sqtt([f"v_add_f32 v{10+i} v{10+i+1} {10+i}" for i in range(3)])
assert all(s.dur == 1 for s in sqtt)
assert all(s.stall == 0 for s in sqtt)
def test_chain_v_add_1l(self):
sqtt = get_sqtt([
"v_add_f32_e32 v1 v0 v0",
"v_add_f32_e32 v2 v1 v1",
])
assert all(s.dur == 1 for s in sqtt)
assert all(s.stall == 0 for s in sqtt)
def test_multi_cycle_inst(self):
sqtt = get_sqtt([
"v_mov_b32_e32 v4 0x3f800000",
"v_rcp_f32_e32 v5 v4",
"v_mul_f32_e32 v6 v5 v4",
])
rcp, mul = sqtt[1], sqtt[2]
self.assertGreater(rcp.dur, 1) # 4 cycles on gfx11
self.assertEqual(mul.dur, 1)
# mul depends on v5, how can it run before rcp is done?
self.assertGreaterEqual(mul.time, rcp.time+rcp.dur)
def test_wmma(self):
sqtt = get_sqtt([
"v_wmma_f32_16x16x16_f16 v[16:23], v[0:7], v[8:15], v[16:23]",
"v_add_f32_e32 v0 v16 v0",
], 32*4)
wmma = sqtt[0]
self.assertGreater(wmma.dur, 1) # rgp says 32 clocks
def test_sleep(self):
n = 1
def sleep_kernel(data0):
assert data0.dtype.base == dtypes.ulong
ops:list[UOp] = []
ops.append(UOp(Ops.CUSTOM, arg="unsigned long long t0 = __builtin_readcyclecounter();"))
ops.append(UOp(Ops.CUSTOM, arg=f"__builtin_amdgcn_s_sleep({n});", src=(ops[-1],)))
ops.append(UOp(Ops.CUSTOM, arg="unsigned long long t1 = __builtin_readcyclecounter();", src=(ops[-1],)))
ops.append(UOp(Ops.CUSTOM, arg=f"data0_{data0.size}[0] = t1 - t0;", src=(ops[-1],)))
return UOp.sink(data0, *ops, arg=KernelInfo(name=f"sleep_{n}"))
diff_hw_reg = Tensor.empty(1, dtype=dtypes.ulong)
diff_hw_reg = Tensor.custom_kernel(diff_hw_reg, fxn=sleep_kernel)[0]
diff_hw_reg.realize()
rctx = decode(dev.profile_events+[ProfileDeviceEvent("AMD", arch=dev.device_info())])
diff_sqtt = list(rctx.inst_execs.values())[0][2]
self.assertEqual(diff_sqtt.dur, diff_hw_reg.item()-1) # 1 cycle for reading the counter register
if __name__ == "__main__":
unittest.main()
+4 -3
View File
@@ -203,9 +203,10 @@ def load_sqtt(profile:list[ProfileEvent]) -> None:
except Exception: return err("DECODER IMPORT ISSUE")
try:
rctx = decode(profile)
steps = [{"name":x[0], "depth":0, "data":{"rows":[(e.inst, e.hit, e.lat, e.stall, str(e.typ).split("_")[-1]) for e in x[1].values()],
"cols":["Instruction", "Hit Count", "Latency", "Stall", "Type"], "summary":[]},
"query":f"/render?ctx={len(ctxs)}&step={i}&fmt=counters"} for i,x in enumerate(rctx.wave_events.items())]
steps = [{"name":x[0], "depth":0, "data":{"rows":[(e.inst, e.time, e.time-x[1][i-1].time if i else 0, e.dur, e.stall, str(e.typ).split("_")[-1])
for i,e in enumerate(x[1])],
"cols":["Instruction", "Clk", "Wait", "Duration", "Stall", "Type"], "summary":[]},
"query":f"/render?ctx={len(ctxs)}&step={i}&fmt=counters"} for i,x in enumerate(rctx.inst_execs.items())]
if not steps: return err("EMPTY SQTT OUTPUT", f"{len(sqtt_events)} SQTT events recorded, none got decoded")
except Exception: return err("DECODER ERROR")
ctxs.append({"name":"Counters", "steps":steps})