import ctypes, struct, time, functools, itertools from tinygrad.runtime.autogen import libusb from tinygrad.helpers import DEBUG, DEV, to_mv, round_up, ceildiv from tinygrad.runtime.support.hcq import MMIOInterface from tinygrad.runtime.support import c def alloc_cbuffer(sz:int) -> tuple[ctypes.Array, memoryview]: return (buf:=(ctypes.c_ubyte * sz)()), to_mv(ctypes.addressof(buf), sz) def checked(fn, msg=None): @functools.wraps(fn) def wrapper(*args): if (rc:=fn(*args)) < 0: raise RuntimeError(f"{msg or fn.__name__}: {ctypes.string_at(libusb.libusb_strerror(rc)).decode()}") return rc return wrapper class USB3: @staticmethod @functools.cache def ctx(): ctx = c.init_c_var(ctypes.POINTER(libusb.struct_libusb_context), checked(libusb.libusb_init)) if DEBUG >= 6: checked(libusb.libusb_set_option)(ctx, libusb.LIBUSB_OPTION_LOG_LEVEL, 4) return ctx @classmethod @functools.cache def list_devices(cls, vendor:int, dev:int) -> list[tuple[c.POINTER[libusb.struct_libusb_device], str]]: ret = [] for i in range(checked(libusb.libusb_get_device_list)(cls.ctx(), devs:=ctypes.POINTER(ctypes.POINTER(libusb.struct_libusb_device))())): desc = c.init_c_var(libusb.struct_libusb_device_descriptor, lambda x: checked(libusb.libusb_get_device_descriptor)(devs[i], x)) if (desc.idVendor, desc.idProduct) == (vendor, dev): ret.append((libusb.libusb_ref_device(devs[i]), f"usb:{libusb.libusb_get_bus_number(devs[i])}-{libusb.libusb_get_device_address(devs[i])}")) libusb.libusb_free_device_list(devs, 1) return ret def __init__(self, dev:c.POINTER[libusb.struct_libusb_device], *args, **kwargs): self._tags, self._transferred = itertools.count(1), ctypes.c_int(0) self._bulk_buf, self._bulk_mv = alloc_cbuffer(4 << 20) self._ctrl_buf, self._ctrl_mv = alloc_cbuffer(0x1000) self.handle = c.init_c_var(c.POINTER[libusb.struct_libusb_device_handle], lambda x: checked(libusb.libusb_open)(dev, x)) # Read product string descriptor _buf = (ctypes.c_ubyte * 256)() _desc = libusb.struct_libusb_device_descriptor() checked(libusb.libusb_get_device_descriptor)(libusb.libusb_get_device(self.handle), ctypes.byref(_desc)) _ret = checked(libusb.libusb_get_string_descriptor_ascii)(self.handle, _desc.iProduct, _buf, 256) self.product = bytes(_buf[:_ret]).decode("ascii", errors="replace") assert self.product.startswith("custom") or self.product.startswith("AS2462") # Detach kernel driver if needed if checked(libusb.libusb_kernel_driver_active)(self.handle, 0): checked(libusb.libusb_detach_kernel_driver)(self.handle, 0) checked(libusb.libusb_reset_device)(self.handle) # Set configuration and claim interface checked(libusb.libusb_set_configuration)(self.handle, 1) checked(libusb.libusb_claim_interface)(self.handle, 0) checked(libusb.libusb_set_interface_alt_setting)(self.handle, 0, 0) def control_write(self, request:int, value:int=0, index:int=0, data:bytes=b'', timeout:int=1000): assert len(data) <= len(self._ctrl_mv) self._ctrl_mv[:len(data)] = data assert checked(libusb.libusb_control_transfer)(self.handle, 0x40, request, value, index, self._ctrl_buf, len(data), timeout) == len(data) def control_read(self, request:int, length:int, value:int=0, index:int=0, timeout:int=1000) -> memoryview: assert length <= len(self._ctrl_mv) assert checked(libusb.libusb_control_transfer)(self.handle, 0xC0, request, value, index, self._ctrl_buf, length, timeout) == length return self._ctrl_mv[:length] def bulk_write(self, payload:bytes, timeout:int=1000): if len(payload) > len(self._bulk_mv): self._bulk_buf, self._bulk_mv = alloc_cbuffer(len(payload)) self._bulk_mv[:len(payload)] = payload checked(libusb.libusb_bulk_transfer, "bulk OUT 0x02 failed") \ (self.handle, 0x02, self._bulk_buf, len(payload), self._transferred, timeout) assert self._transferred.value == len(payload), f"bulk OUT short write: {self._transferred.value}/{len(payload)} bytes" def bulk_read(self, length:int, timeout:int=1000) -> memoryview: if length > len(self._bulk_mv): self._bulk_buf, self._bulk_mv = alloc_cbuffer(length) checked(libusb.libusb_bulk_transfer, "bulk IN 0x81 failed")(self.handle, 0x81, self._bulk_buf, length, self._transferred, timeout) return self._bulk_mv[:self._transferred.value] # NOTE: keep it for flash.py def send_batch(self, cdbs:list[bytes], odata:list[bytes|None]|None=None): for cdb, data in zip(cdbs, odata or [None] * len(cdbs)): self.bulk_write(struct.pack("> 32, value), 5000) def _f0_in(self) -> tuple[int, int, int]: data = self.usb.control_read(0xF0, 8, timeout=5000) return struct.unpack_from('> 5) & 0x7, data[7] def pcie_request(self, fmt_type:int, address:int, value:int|None=None, size:int=4, cnt:int=10): assert size > 0 and size <= 4, f"Invalid size {size}" if DEBUG >= 5: print("pcie_request", hex(fmt_type), hex(address), value, size) offset = address & 0x3 byte_en = ((1 << size) - 1) << offset self._f0_out(fmt_type, byte_en, address & ~0x3, (value << (8 * offset)) if value is not None else 0) # Fast path: memory writes and messages don't return completions. if ((fmt_type & 0b11011111) == 0b01000000) or ((fmt_type & 0b10111000) == 0b00110000): return # Read TLPs and config writes: read completion via 0xF0 IN. Retry on error/timeout. data, cpl_status, ret_status = self._f0_in() if ret_status != 0: time.sleep(0.001) # TODO: this sleep is very picky if cnt > 0: return self.pcie_request(fmt_type, address, value, size, cnt=cnt-1) raise RuntimeError(f"TLP error after retries: ret_status={ret_status}, address={address:#x}") if cpl_status: status_map = {0b001: f"Unsupported Request: {address:#x}", 0b100: "Completer Abort", 0b010: "Config Retry"} raise RuntimeError(f"TLP completion status: {status_map.get(cpl_status, f'Reserved (0b{cpl_status:03b})')}") if value is None: return (data >> (8 * offset)) & ((1 << (8 * size)) - 1) def pcie_cfg_req(self, byte_addr:int, bus:int=1, dev:int=0, fn:int=0, value:int|None=None, size:int=4): assert byte_addr >> 12 == 0 and bus >> 8 == 0 and dev >> 5 == 0 and fn >> 3 == 0 fmt_type = (0x44 if value is not None else 0x4) | int(bus > 0) address = (bus << 24) | (dev << 19) | (fn << 16) | (byte_addr & 0xfff) return self.pcie_request(fmt_type, address, value, size) def pcie_mem_write(self, address:int, data:bytes): """Streaming PCIe memory write via 0xF0 mode 1 + bulk OUT. Data is little-endian dwords on the wire.""" if not data: return assert len(data) % 4 == 0, f"pcie_mem_write requires 4-byte aligned size, got {len(data)}" self._f0_out(0x60, 0x0F, address, len(data) // 4, mode=1) self.usb.bulk_write(data) def pcie_mem_read(self, address:int, nbytes:int) -> memoryview: """Streaming PCIe memory read via 0xF0 mode 2 + bulk IN. Returns little-endian bytes.""" assert nbytes % 4 == 0, f"pcie_mem_read requires 4-byte aligned size, got {nbytes}" self._f0_out(0x20, 0x0F, address, nbytes // 4, mode=2) return self.usb.bulk_read(nbytes, timeout=30000) def read(self, base_addr:int, length:int) -> bytes: """Read from chip XDATA via vendor control IN (bRequest=0xE4). wValue=addr, wLength=size.""" result = b'' for off in range(0, length, 0xFF): chunk = min(0xFF, length - off) result += self.usb.control_read(0xE4, chunk, value=base_addr + off) return result def write(self, base_addr:int, data:bytes): """Write to chip XDATA via vendor control OUT (bRequest=0xE5). wValue=addr, wIndex=val.""" for off, val in enumerate(data): self.usb.control_write(0xE5, value=base_addr + off, index=val) def scsi_write(self, buf:bytes): """Write to SRAM via 0xF2 vendor command + bulk OUT.""" buf_padded = buf + b'\x00' * (round_up(len(buf), 512) - len(buf)) sectors = len(buf_padded) // 512 num_slots = ceildiv(len(buf_padded), 0x4000) # 16KB per slot windex = (num_slots & 0xFF) << 8 self.usb.control_write(0xF2, value=sectors, index=windex) self.usb.bulk_write(buf_padded) def scsi_read_arm(self, size:int): windex = (ceildiv(size, 0x4000) & 0xFF) << 8 self.usb.control_write(0xF2, value=(ceildiv(size, 512) & 0x7FFF) | 0x8000, index=windex) def scsi_read(self, size:int) -> memoryview: return self.usb.bulk_read(round_up(size, 512), timeout=10000)[:size] class USBMMIOInterface(MMIOInterface): def __init__(self, usb, addr, size, fmt, pcimem=True): # pylint: disable=super-init-not-called self.usb, self.addr, self.nbytes, self.fmt, self.el_sz, self.pcimem = usb, addr, size, fmt, struct.calcsize(fmt), pcimem def _off_from_index(self, index): if isinstance(index, slice): return ((index.start or 0) * self.el_sz, ((index.stop or len(self))-(index.start or 0)) * self.el_sz) return (index * self.el_sz, self.el_sz) def __getitem__(self, index): off, sz = self._off_from_index(index) if self.pcimem: assert sz % 4 == 0 and off % 4 == 0, f"pcie_mem_read requires 4-byte aligned access, got off={off}, sz={sz}" data = self.usb.pcie_mem_read(self.addr + off, sz) else: data = self.usb.scsi_read(sz) if self.addr == 0xf000 else self.usb.read(self.addr + off, sz) return int.from_bytes(data, "little") if sz == self.el_sz else data def __setitem__(self, index, data): off, _ = self._off_from_index(index) data = struct.pack(self.fmt, data) if isinstance(data, int) else bytes(data) if not self.pcimem: self.usb.scsi_write(data) if self.addr == 0xf000 else self.usb.write(self.addr + off, data) else: self.usb.pcie_mem_write(self.addr+off, data) def view(self, offset:int=0, size:int|None=None, fmt=None): return USBMMIOInterface(self.usb, self.addr+offset, self.nbytes-offset if size is None else size, fmt=fmt or self.fmt, pcimem=self.pcimem) if DEV.interface.startswith("MOCK"): from test.mockgpu.usb import MockUSB3 as USB3 # type: ignore # noqa: F811