diff --git a/tinygrad/helpers.py b/tinygrad/helpers.py index f6c595f1a2..d0ffd3f803 100644 --- a/tinygrad/helpers.py +++ b/tinygrad/helpers.py @@ -316,6 +316,12 @@ def capstone_flatdump(lib: bytes): print(f"{instr.address:#08x}: {instr.mnemonic}\t{instr.op_str}") sys.stdout.flush() +def wait_cond(cb, value=True, timeout_ms=10000, msg="") -> bool: + start_time = int(time.perf_counter() * 1000) + while int(time.perf_counter() * 1000) - start_time < timeout_ms: + if (val:=cb()) == value: return val + raise TimeoutError(f"{msg}. Timed out after {timeout_ms} ms, condition not met: {val} != {value}") + # *** ctypes helpers # TODO: make this work with read only memoryviews (if possible) diff --git a/tinygrad/runtime/support/am/amdev.py b/tinygrad/runtime/support/am/amdev.py index 7c7b615997..cacfc77f38 100644 --- a/tinygrad/runtime/support/am/amdev.py +++ b/tinygrad/runtime/support/am/amdev.py @@ -1,5 +1,5 @@ from __future__ import annotations -import ctypes, collections, time, dataclasses, functools, os, hashlib +import ctypes, collections, dataclasses, functools, os, hashlib from tinygrad.helpers import mv_address, getenv, DEBUG, fetch from tinygrad.runtime.autogen.am import am from tinygrad.runtime.support.hcq import MMIOInterface @@ -214,12 +214,6 @@ class AMDev(PCIDevImplBase): self.reg("regBIF_BX_PF0_RSMU_INDEX").write(reg) self.reg("regBIF_BX_PF0_RSMU_DATA").write(val) - def wait_reg(self, reg:AMRegister, value:int, mask=0xffffffff, timeout=10000) -> int: - start_time = int(time.perf_counter() * 1000) - while int(time.perf_counter() * 1000) - start_time < timeout: - if ((rval:=reg.read()) & mask) == value: return rval - raise RuntimeError(f'wait_reg timeout reg=0x{reg.addr:X} mask=0x{mask:X} value=0x{value:X} last_val=0x{rval}') - def _run_discovery(self): # NOTE: Fixed register to query memory size without known ip bases to find the discovery table. # The table is located at the end of VRAM - 64KB and is 10KB in size. diff --git a/tinygrad/runtime/support/am/ip.py b/tinygrad/runtime/support/am/ip.py index d4e1c46374..f205c32747 100644 --- a/tinygrad/runtime/support/am/ip.py +++ b/tinygrad/runtime/support/am/ip.py @@ -1,7 +1,7 @@ import ctypes, time, contextlib, importlib, functools from typing import Literal from tinygrad.runtime.autogen.am import am -from tinygrad.helpers import to_mv, data64, lo32, hi32, DEBUG +from tinygrad.helpers import to_mv, data64, lo32, hi32, DEBUG, wait_cond class AM_IP: def __init__(self, adev): self.adev = adev @@ -53,12 +53,12 @@ class AM_GMC(AM_IP): # Can't issue TLB invalidation if the hub isn't initialized. if not self.hub_initted[ip]: return - if ip == "MM": self.adev.wait_reg(self.adev.regMMVM_INVALIDATE_ENG17_SEM, mask=0x1, value=0x1) + if ip == "MM": wait_cond(lambda: self.adev.regMMVM_INVALIDATE_ENG17_SEM.read() & 0x1, value=1, msg="mm flush_tlb timeout") self.adev.reg(f"reg{ip}VM_INVALIDATE_ENG17_REQ").write(flush_type=flush_type, per_vmid_invalidate_req=(1 << vmid), invalidate_l2_ptes=1, invalidate_l2_pde0=1, invalidate_l2_pde1=1, invalidate_l2_pde2=1, invalidate_l1_ptes=1, clear_protection_fault_status_addr=0) - self.adev.wait_reg(self.adev.reg(f"reg{ip}VM_INVALIDATE_ENG17_ACK"), mask=(1 << vmid), value=(1 << vmid)) + wait_cond(lambda: self.adev.reg(f"reg{ip}VM_INVALIDATE_ENG17_ACK").read() & (1 << vmid), value=(1 << vmid), msg="flush_tlb timeout") if ip == "MM": self.adev.regMMVM_INVALIDATE_ENG17_SEM.write(0x0) @@ -176,7 +176,8 @@ class AM_SMU(AM_IP): def _send_msg(self, msg, param, read_back_arg=False, timeout=10000, debug=False): # 10s self._smu_cmn_send_msg(msg, param, debug=debug) - self.adev.wait_reg(self.adev.mmMP1_SMN_C2PMSG_90 if not debug else self.adev.mmMP1_SMN_C2PMSG_54, mask=0xFFFFFFFF, value=1, timeout=timeout) + wait_cond(lambda: (self.adev.mmMP1_SMN_C2PMSG_90 if not debug else self.adev.mmMP1_SMN_C2PMSG_54).read(), value=1, timeout_ms=timeout, + msg=f"SMU msg {msg:#x} timeout") return (self.adev.mmMP1_SMN_C2PMSG_82 if not debug else self.adev.mmMP1_SMN_C2PMSG_53).read() if read_back_arg else None class AM_GFX(AM_IP): @@ -260,7 +261,7 @@ class AM_GFX(AM_IP): if hasattr(self.adev, 'regMM_ATC_L2_MISC_CG'): self.adev.regMM_ATC_L2_MISC_CG.write(enable=1, mem_ls_enable=1) self.adev.regRLC_SAFE_MODE.write(message=1, cmd=1) - self.adev.wait_reg(self.adev.regRLC_SAFE_MODE, mask=0x1, value=0x0) + wait_cond(lambda: self.adev.regRLC_SAFE_MODE.read() & 0x1, value=0, msg="RLC safe mode timeout") self.adev.regRLC_CGCG_CGLS_CTRL.update(cgcg_gfx_idle_threshold=0x36, cgcg_en=1, cgls_rep_compansat_delay=0xf, cgls_en=1) @@ -411,7 +412,7 @@ class AM_PSP(AM_IP): def is_sos_alive(self): return self.adev.reg(f"{self.reg_pref}_81").read() != 0x0 - def _wait_for_bootloader(self): self.adev.wait_reg(self.adev.reg(f"{self.reg_pref}_35"), mask=0x80000000, value=0x80000000) + def _wait_for_bootloader(self): wait_cond(lambda: self.adev.reg(f"{self.reg_pref}_35").read() & 0x80000000, value=0x80000000, msg="BL not ready") def _prep_msg1(self, data): assert len(data) <= self.msg1_view.nbytes, f"msg1 buffer is too small {len(data):#x} > {self.msg1_view.nbytes:#x}" @@ -446,7 +447,7 @@ class AM_PSP(AM_IP): time.sleep(0.02) # Wait until the sOS is ready - self.adev.wait_reg(self.adev.reg(f"{self.reg_pref}_64"), mask=0x80000000, value=0x80000000) + wait_cond(lambda: self.adev.reg(f"{self.reg_pref}_64").read() & 0x80000000, value=0x80000000, msg="sOS not ready") self.adev.wreg_pair(self.reg_pref, "_69", "_70", self.adev.paddr2mc(self.ring_paddr)) self.adev.reg(f"{self.reg_pref}_71").write(self.ring_size) @@ -455,7 +456,7 @@ class AM_PSP(AM_IP): # There might be handshake issue with hardware which needs delay time.sleep(0.02) - self.adev.wait_reg(self.adev.reg(f"{self.reg_pref}_64"), mask=0x8000FFFF, value=0x80000000) + wait_cond(lambda: self.adev.reg(f"{self.reg_pref}_64").read() & 0x8000FFFF, value=0x80000000, msg="sOS ring not created") def _ring_submit(self, cmd): msg = am.struct_psp_gfx_rb_frame(fence_value=(prev_wptr:=self.adev.reg(f"{self.reg_pref}_67").read()),