From 7eb763a3c290c96972dbabfa24acfb87d05d80e2 Mon Sep 17 00:00:00 2001 From: nimlgen <138685161+nimlgen@users.noreply.github.com> Date: Thu, 27 Aug 2026 17:29:17 +0300 Subject: [PATCH] bnxt to extra (#17772) * bnxt to extra * x * x * les --- .github/workflows/autogen.yml | 2 +- extra/bnxt_driver/bnxtdev.py | 238 ++ extra/bnxt_driver/connect.py | 118 + extra/bnxt_driver/loopback.py | 46 + test/unit/test_bnxt.py | 115 + test/unit/test_bnxt_transport.py | 40 + tinygrad/runtime/autogen/__init__.py | 17 + tinygrad/runtime/autogen/bnxt.py | 5407 ++++++++++++++++++++++++ tinygrad/runtime/support/mlx/mlxdev.py | 6 +- tinygrad/runtime/support/system.py | 2 + 10 files changed, 5986 insertions(+), 5 deletions(-) create mode 100644 extra/bnxt_driver/bnxtdev.py create mode 100644 extra/bnxt_driver/connect.py create mode 100644 extra/bnxt_driver/loopback.py create mode 100644 test/unit/test_bnxt.py create mode 100644 test/unit/test_bnxt_transport.py create mode 100644 tinygrad/runtime/autogen/bnxt.py diff --git a/.github/workflows/autogen.yml b/.github/workflows/autogen.yml index 6b3fb964b7..a6d0e7eb23 100644 --- a/.github/workflows/autogen.yml +++ b/.github/workflows/autogen.yml @@ -54,7 +54,7 @@ jobs: python3 -c "from tinygrad.runtime.autogen import mesa" python3 -c "from tinygrad.runtime.autogen import avcodec" python3 -c "from tinygrad.runtime.autogen import llvm_qcom" - python3 -c "from tinygrad.runtime.autogen import mlx5" + python3 -c "from tinygrad.runtime.autogen import mlx5, bnxt" python3 -c "from tinygrad.runtime.autogen import ggml_common" REGEN=1 python3 -c "from tinygrad.runtime.autogen import libclang" - name: Check for differences diff --git a/extra/bnxt_driver/bnxtdev.py b/extra/bnxt_driver/bnxtdev.py new file mode 100644 index 0000000000..bf898625d0 --- /dev/null +++ b/extra/bnxt_driver/bnxtdev.py @@ -0,0 +1,238 @@ +import ctypes, struct +from tinygrad.helpers import ceildiv, getenv, wait_cond, DEBUG +from tinygrad.runtime.autogen import bnxt, pci +from tinygrad.runtime.support.system import PCIDevice, System, ipv4_to_gid + +BNXT_DEBUG = getenv("BNXT_DEBUG", 0) +BNXT_ACCESS, BNXT_INIT_MASK, BNXT_RTR_MASK, BNXT_RTS_MASK = 3, 0xd, 0x41515ad, 0xae005 +BNXT_CHIMP_COMM, BNXT_CHIMP_COMM_TRIGGER = 0x0, 0x100 +BNXT_BACKING_STORE = ((0, 2), (1, 0), (2, 2), (3, 0), (4, 2), (5, 0), (6, 0), (14, 2), (15, 0)) + +def db_value(xid, typ, index, epoch): + return (xid & bnxt.DBC_DBC_XID_MASK | bnxt.DBC_DBC_PATH_ROCE | typ | bnxt.BNXT_QPLIB_DBR_VALID) << 32 | \ + index & bnxt.DBC_DBC_INDEX_MASK | epoch << bnxt.BNXT_QPLIB_DBR_EPOCH_SHIFT + +def _pbl(dev, paddrs, queue=False): + if len(paddrs) == 1: return 0, paddrs[0] + values = [p | bnxt.PTU_PTE_VALID for p in paddrs] + if queue: + values[-1] |= bnxt.PTU_PTE_LAST + if len(values) > 1: values[-2] |= bnxt.PTU_PTE_NEXT_TO_LAST + table, table_paddrs = dev.pci_dev.alloc_sysmem(ceildiv(len(values), 512) * 0x1000) + table[:len(values) * 8] = struct.pack(f"<{len(values)}Q", *values) + if len(table_paddrs) == 1: return 1, table_paddrs[0] + top, top_paddrs = dev.pci_dev.alloc_sysmem(0x1000) + top[:len(table_paddrs) * 8] = struct.pack(f"<{len(table_paddrs)}Q", *(p | bnxt.PTU_PTE_VALID for p in table_paddrs)) + return 2, top_paddrs[0] + +def _queue(dev, stride:int=16, aux=False): + mem, paddrs = dev.pci_dev.alloc_sysmem(0x1000 + aux * 0x400) + level, base = _pbl(dev, paddrs, queue=True) + return {"mem":mem, "paddrs":paddrs, "stride":stride, "prod":0, "cons":0, "level":level, "base":base} + +def _qread(q, i): + off = (i & 15) * q["stride"] + return q["mem"][off:off + q["stride"]] + +def _qwrite(q, i, data, aux=False): + off = 0x1000 + i % 128 * 8 if aux else (i & 15) * q["stride"] + q["mem"][off:off + len(data)] = data + +class BNXTDev: + def __init__(self, pci_dev:PCIDevice, ip:str=getenv("BNXT_IP", "10.0.0.1")): + self.pci_dev, self.devfmt = pci_dev, pci_dev.pcibus + self.bar0, self.db = pci_dev.map_bar(0, fmt='I'), pci_dev.map_bar(2, fmt='Q') + pci_dev.write_config(pci.PCI_COMMAND, pci_dev.read_config(pci.PCI_COMMAND, 2) | pci.PCI_COMMAND_MASTER, 2) + self.resp, self.resp_pa = pci_dev.alloc_sysmem(0x1000) + self.seq = 0 + + ver = self.hwrm("ver_get") + if DEBUG >= 2: print(f"bnxt {self.devfmt}: firmware {ver.hwrm_fw_maj_8b}.{ver.hwrm_fw_min_8b}.{ver.hwrm_fw_bld_8b}") + self.hwrm("func_reset", timeout_ms=40000) + caps = self.hwrm("func_qcaps", fid=0xffff) + self.mac, self.port_id = int.from_bytes(bytes(caps.mac_address), 'big'), caps.port_id + self.hwrm("func_drv_rgtr") + self.db_off = self.hwrm("func_qcfg", fid=0xffff).legacy_l2_db_size_kb * 1024 + + self.setup_backing_store() + self._open_rcfw() + self._open_l2() + self.local_gid = ipv4_to_gid(ip) + gids, mac = (ctypes.c_uint32 * 4)(*(int.from_bytes(self.local_gid[i:i + 4], 'big') for i in (12, 8, 4, 0))), self.mac.to_bytes(6, 'big') + smac = (ctypes.c_uint16 * 3)(*(int.from_bytes(mac[i:i + 2], 'big') for i in (0, 2, 4))) + self.gid_id = self.rcfw("add_gid", gid=gids, src_mac=smac).xid + + if DEBUG >= 2: print(f"bnxt {self.devfmt}: booted mac={self.mac.to_bytes(6, 'big').hex(':')} gid={self.local_gid.hex()}") + + def hwrm(self, name, timeout_ms=10000, **fields): + inp, out = getattr(bnxt, f"struct_hwrm_{name}_input"), getattr(bnxt, f"struct_hwrm_{name}_output") + opcode = getattr(bnxt, f"HWRM_{name.upper()}") + self.seq = (self.seq + 1) & 0xffff + data = bytes(inp(req_type=opcode, cmpl_ring=bnxt.BNXT_HWRM_NO_CMPL_RING, seq_id=self.seq, target_id=bnxt.BNXT_HWRM_TARGET, + resp_addr=self.resp_pa[0], **fields)) + self.resp[:] = bytes(len(self.resp)) + System.memory_barrier() + for i, w in enumerate(memoryview(bytearray(data.ljust(bnxt.HWRM_MAX_REQ_LEN, b'\0'))).cast('I')): + self.bar0[BNXT_CHIMP_COMM // 4 + i] = w + self.bar0[BNXT_CHIMP_COMM_TRIGGER // 4] = 1 + def hdr(): return bnxt.struct_hwrm_resp_hdr.from_buffer_copy(bytes(self.resp[:8])) + wait_cond(lambda: (n := hdr().resp_len) and hdr().seq_id == self.seq and self.resp[n - 1], timeout_ms=timeout_ms, msg=f"HWRM {name}") + ret = out.from_buffer_copy(bytes(self.resp[:ctypes.sizeof(out)])) + assert ret.error_code == 0, f"HWRM {name}: {ret.error_code}" + return ret + + def setup_backing_store(self): + counts: dict[int, int] = {} + for typ, extra in BNXT_BACKING_STORE: + caps = self.hwrm("func_backing_store_qcaps_v2", type=typ) + size, splits = caps.entry_size, tuple(getattr(caps, f"split_entry_{j}") for j in range(caps.subtype_valid_cnt)) + counts[typ] = n = counts[0] if typ == 15 else max(caps.min_num_entries, sum(splits) + extra) + # a zero bitmap means the type has a single instance 0 + for instance in [i for i in range(8) if caps.instance_bit_map >> i & 1] or [0]: + mem, paddrs = self.pci_dev.alloc_sysmem(ceildiv(n * size, 0x1000) * 0x1000) + if caps.ctx_init_value: + for off in range(caps.ctx_init_offset, len(mem), size): mem[off] = caps.ctx_init_value + lvl, base = _pbl(self, paddrs) + self.hwrm("func_backing_store_cfg_v2", type=typ, instance=instance, entry_size=size, num_entries=n, page_dir=base, + page_size_pbl_level=lvl, subtype_valid_cnt=len(splits), + flags=bnxt.FUNC_BACKING_STORE_CFG_V2_REQ_FLAGS_BS_CFG_ALL_DONE if typ == 15 else 0, + **{f"split_entry_{j}": v for j, v in enumerate(splits)}) + + def _open_rcfw(self): + self.rcfw_first = True + + self.creq = _queue(self) + self.creq_id = self.hwrm("ring_alloc", ring_type=bnxt.RING_ALLOC_REQ_RING_TYPE_NQ, page_tbl_addr=self.creq["base"], + page_size=12, page_tbl_depth=self.creq["level"], length=16, int_mode=bnxt.RING_ALLOC_REQ_INT_MODE_MSIX).ring_id + + self.cmdq = _queue(self) + self.doorbell(self.creq_id, bnxt.DBC_DBC_TYPE_NQ_ARM, 0, 0) + init = bnxt.struct_cmdq_init(cmdq_pbl=self.cmdq["base"], creq_ring_id=self.creq_id, + cmdq_size_cmdq_lvl=16 << bnxt.CMDQ_INIT_CMDQ_SIZE_SFT) + + System.memory_barrier() + for i, w in enumerate(memoryview(bytearray(bytes(init))).cast('I')): self.bar0[bnxt.RCFW_COMM_BASE_OFFSET // 4 + i] = w + + _, p = self.pci_dev.alloc_sysmem(0x1000) + self.rcfw("initialize_fw", stat_ctx_id=self.hwrm("stat_ctx_alloc", stats_dma_addr=p[0], stats_dma_length=176).stat_ctx_id, + flags=bnxt.CMDQ_INITIALIZE_FW_FLAGS_HW_REQUESTER_RETX_SUPPORTED) + + # RoCE notification ring: never armed or serviced, but CQ and L2 ring allocation require one + nq = _queue(self) + self.nq_id = self.hwrm("ring_alloc", ring_type=bnxt.RING_ALLOC_REQ_RING_TYPE_NQ, page_tbl_addr=nq["base"], + page_size=12, page_tbl_depth=nq["level"], length=16, logical_id=1, int_mode=bnxt.RING_ALLOC_REQ_INT_MODE_MSIX).ring_id + + def rcfw(self, name, timeout_ms=20000, **fields): + req_t, resp_t = getattr(bnxt, f"struct_cmdq_{name}"), getattr(bnxt, f"struct_creq_{name}_resp") + op = getattr(bnxt, f"CMDQ_BASE_OPCODE_{name.upper()}") + data = bytes(req_t(opcode=op, cmd_size=(slots := ceildiv(ctypes.sizeof(req_t), 16)), **fields)).ljust(slots * 16, b'\0') + for i in range(slots): _qwrite(self.cmdq, self.cmdq["prod"] + i, data[i * 16:(i + 1) * 16]) + + self.cmdq["prod"] += slots + prod = self.cmdq["prod"] & 0xffff + if self.rcfw_first: prod, self.rcfw_first = prod | 1 << bnxt.FIRMWARE_FIRST_FLAG, False + + System.memory_barrier() + + self.bar0[(bnxt.RCFW_COMM_BASE_OFFSET + bnxt.RCFW_PF_VF_COMM_PROD_OFFSET) // 4] = prod + self.bar0[(bnxt.RCFW_COMM_BASE_OFFSET + bnxt.RCFW_COMM_TRIG_OFFSET) // 4] = bnxt.RCFW_CMDQ_TRIG_VAL + + def poll(): + h = bnxt.struct_creq_base.from_buffer_copy(bytes(_qread(self.creq, self.creq["cons"]))) + return bool(h.v & bnxt.CREQ_BASE_V) != bool((self.creq["cons"] // 16) & 1) + wait_cond(poll, timeout_ms=timeout_ms, msg=f"RCFW {name}") + + ret = resp_t.from_buffer_copy(bytes(_qread(self.creq, self.creq["cons"]))) + self.creq["cons"] += 1 + + # NQ_ARM also publishes the CREQ consumer index, which is what frees ring space for the next command + self.doorbell(self.creq_id, bnxt.DBC_DBC_TYPE_NQ_ARM, self.creq["cons"] & 15, (self.creq["cons"] // 16) & 1) + assert ret.status == 0, f"RCFW {name}: {ret.status}" + + if BNXT_DEBUG >= 1: print(f"bnxt {self.devfmt}: rcfw {name} xid={getattr(ret, 'xid', 0):#x}") + return ret + + def doorbell(self, xid, typ, index, epoch): + System.memory_barrier() + self.db[self.db_off // 8] = db_value(xid, typ, index, epoch) + + # L2 receive path, required for RoCE ingress even though no ethernet receive buffers are posted + def _open_l2(self): + cq = _queue(self) + ci = self.hwrm("ring_alloc", enables=bnxt.RING_ALLOC_REQ_ENABLES_NQ_RING_ID_VALID, ring_type=bnxt.RING_ALLOC_REQ_RING_TYPE_L2_CMPL, + page_tbl_addr=cq["base"], page_size=12, page_tbl_depth=cq["level"], length=16, nq_ring_id=self.nq_id).ring_id + rx = _queue(self) + ri = self.hwrm("ring_alloc", enables=bnxt.RING_ALLOC_REQ_ENABLES_NQ_RING_ID_VALID | + bnxt.RING_ALLOC_REQ_ENABLES_RX_BUF_SIZE_VALID, ring_type=bnxt.RING_ALLOC_REQ_RING_TYPE_RX, page_tbl_addr=rx["base"], + page_size=12, page_tbl_depth=rx["level"], length=16, rx_buf_size=640, nq_ring_id=self.nq_id).ring_id + vi = self.hwrm("vnic_alloc").vnic_id + self.hwrm("vnic_cfg", enables=bnxt.VNIC_CFG_REQ_ENABLES_MRU | bnxt.VNIC_CFG_REQ_ENABLES_DEFAULT_RX_RING_ID | + bnxt.VNIC_CFG_REQ_ENABLES_DEFAULT_CMPL_RING_ID, vnic_id=vi, mru=9018, + default_rx_ring_id=ri, default_cmpl_ring_id=ci) + self.hwrm("cfa_l2_filter_alloc", flags=bnxt.CFA_L2_FILTER_ALLOC_REQ_FLAGS_PATH_RX, + enables=bnxt.CFA_L2_FILTER_ALLOC_REQ_ENABLES_L2_ADDR | bnxt.CFA_L2_FILTER_ALLOC_REQ_ENABLES_L2_ADDR_MASK | + bnxt.CFA_L2_FILTER_ALLOC_REQ_ENABLES_DST_ID, l2_addr=tuple(self.mac.to_bytes(6, 'big')), l2_addr_mask=(0xff,) * 6, dst_id=vi) + + def register_mem(self, paddrs:list[int], size:int, log_page_size:int=12) -> int: + level, base = _pbl(self, paddrs[:ceildiv(size, 1 << log_page_size)]) + return self.rcfw("register_mr", flags=bnxt.CMDQ_REGISTER_MR_FLAGS_ALLOC_MR, + log2_pg_size_lvl=level << bnxt.CMDQ_REGISTER_MR_LVL_SFT | log_page_size << bnxt.CMDQ_REGISTER_MR_LOG2_PG_SIZE_SFT, + access=bnxt.CMDQ_REGISTER_MR_ACCESS_LOCAL_WRITE | bnxt.CMDQ_REGISTER_MR_ACCESS_REMOTE_WRITE, + log2_pbl_pg_size=12, pbl=base, va=paddrs[0], mr_size=size).xid + +class BNXTQP: + def __init__(self, dev:BNXTDev): + self.dev, self.sq_psn, self.msn = dev, 0, 0 + + self.cqq = _queue(dev, ctypes.sizeof(bnxt.struct_cq_base)) + self.cq_id = dev.rcfw("create_cq", cq_size=16, pbl=self.cqq["base"], + pg_size_lvl=self.cqq["level"], cq_fco_cnq_id=dev.nq_id).xid + + self.sq = _queue(dev, aux=True) + self.qpn = dev.rcfw("create_qp", type=bnxt.CMDQ_CREATE_QP_TYPE_RC, + sq_size=16, sq_fwo_sq_sge=1, scq_cid=self.cq_id, rcq_cid=self.cq_id, + sq_pbl=self.sq["base"], sq_pg_size_sq_lvl=self.sq["level"]).xid + self.qp_op(1, BNXT_INIT_MASK, access=BNXT_ACCESS, pkey=0xffff) + + def qp_op(self, state, mask, network_type=0, **fields): + self.dev.rcfw("modify_qp", qp_cid=self.qpn, modify_mask=mask, + network_type_en_sqd_async_notify_new_state=state | network_type, **fields) + + def connect(self, qpn:int, gid:bytes, mac:int): + network_type = bnxt.CMDQ_MODIFY_QP_NETWORK_TYPE_ROCEV2_IPV4 + dgid = (ctypes.c_uint32 * 4)(*(int.from_bytes(gid[i:i + 4], 'little') for i in (0, 4, 8, 12))) + dmac = (ctypes.c_uint16 * 3)(*(int.from_bytes(mac.to_bytes(6, 'big')[i:i + 2], 'little') for i in (0, 2, 4))) + + self.qp_op(2, BNXT_RTR_MASK, network_type=network_type, qp_type=bnxt.CMDQ_MODIFY_QP_QP_TYPE_RC, access=BNXT_ACCESS, + pkey=0xffff, dgid=dgid, sgid_index=self.dev.gid_id, hop_limit=64, dest_mac=dmac, + path_mtu_pingpong_push_enable=bnxt.CMDQ_MODIFY_QP_PATH_MTU_MTU_1024, max_dest_rd_atomic=4, + dest_qp_id=qpn) + self.qp_op(3, BNXT_RTS_MASK, network_type=network_type, qp_type=bnxt.CMDQ_MODIFY_QP_QP_TYPE_RC, access=BNXT_ACCESS, + max_rd_atomic=1) + + if BNXT_DEBUG >= 1: print(f"bnxt: QP {self.qpn:#x} connected (remote={qpn:#x})") + + def _poll(self, timeout): + def poll(): + base = bnxt.struct_cq_base.from_buffer_copy(bytes(_qread(self.cqq, self.cqq["cons"]))) + return bool(base.cqe_type_toggle & bnxt.CQ_BASE_TOGGLE) == (not bool((self.cqq["cons"] // 16) & 1)) + wait_cond(poll, timeout_ms=timeout, msg="BNXT CQ") + raw = bytes(_qread(self.cqq, self.cqq["cons"])) + self.cqq["cons"] += 1 + self.dev.doorbell(self.cq_id, bnxt.DBC_DBC_TYPE_CQ, self.cqq["cons"] & 15, (self.cqq["cons"] // 16) & 1) + return raw + + def rdma_write(self, rva, rkey, lva, lkey, size, timeout_ms=20000): + start = self.sq["prod"] & 15 + hdr = bytes(bnxt.struct_sq_rdma_hdr(wqe_type=bnxt.SQ_RDMA_HDR_WQE_TYPE_WRITE_WQE, + flags=bnxt.SQ_SEND_FLAGS_SIGNAL_COMP, wqe_size=3, length=size, remote_va=rva, remote_key=rkey)) + for i, data in enumerate((hdr[:16], hdr[16:32], bytes(bnxt.struct_sq_sge(va_or_pa=lva, l_key=lkey, size=size)))): + _qwrite(self.sq, start + i, data) + nxt = (self.sq_psn + max(1, ceildiv(size, 1024))) & 0xffffff + value = start << bnxt.SQ_MSN_SEARCH_START_IDX_SFT | nxt << bnxt.SQ_MSN_SEARCH_NEXT_PSN_SFT | self.sq_psn + _qwrite(self.sq, self.msn, struct.pack(" dict[str, Any]: + for line in iter(stream.readline, ""): + print(f" [remote] {line}", end="") + try: value = json.loads(line) + except json.JSONDecodeError: continue + if isinstance(value, dict): return value + raise RuntimeError(f"remote exited before publishing {what}") + +def wait_line(stream:IO[str], text:str) -> str: + for line in iter(stream.readline, ""): + print(f" [remote] {line}", end="") + if text in line: return line + raise RuntimeError(f"remote exited before reporting {text!r}") + +def send_line(stream:IO[str], value:str|dict[str, Any]): + stream.write((json.dumps(value) if isinstance(value, dict) else value) + "\n") + stream.flush() + +def qp_info(dev:BNXTDev, qp:BNXTQP) -> dict[str, Any]: + return {"qpn":qp.qpn, "mac":dev.mac.to_bytes(6, "big").hex(), "gid":dev.local_gid.hex()} + +def server(): + dev = BNXTDev(PCIDevice("bnxt", os.getenv("BNXT_PCI", "0000:41:00.0")), ip=os.getenv("BNXT_IP", REMOTE_IP)) + qp = BNXTQP(dev) + print(json.dumps(qp_info(dev, qp)), flush=True) + + peer = json.loads(sys.stdin.readline()) + qp.connect(peer["qpn"], bytes.fromhex(peer["gid"]), int(peer["mac"], 16)) + print("connected", flush=True) + + target, target_paddrs = dev.pci_dev.alloc_sysmem(0x1000) + target[:0x1000] = bytes(0x1000) + rkey = dev.register_mem(target_paddrs, 0x1000) + print(json.dumps({"target_addr":target_paddrs[0], "rkey":rkey}), flush=True) + + assert sys.stdin.readline().strip() == "done" + received = bytes(target).rstrip(b"\0") + print(f"AS TEXT: {received.decode(errors='replace')!r}", flush=True) + print(json.dumps({"data":received.hex()}), flush=True) + +def sync_remote(): + if os.getenv("SYNC", "1") == "0": return + print("syncing BNXT driver to remote") + subprocess.run(["rsync", "-azR", *SYNC_FILES, f"{REMOTE}:~/tinygrad/"], cwd=TINYGRAD, check=True) + +def start_remote() -> subprocess.Popen[str]: + print("booting remote") + command = (f"cd ~/tinygrad && sudo env PYTHONPATH=. PYTHONUNBUFFERED=1 BNXT_DEBUG={os.getenv('BNXT_DEBUG', '0')} " + f"BNXT_PCI={REMOTE_PCI} BNXT_IP={REMOTE_IP} python3 extra/bnxt_driver/connect.py --server") + return subprocess.Popen(SSH + [command], stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=sys.stderr, text=True) + +def client(): + assert 0 < len(MESSAGE) <= 0x1000 + sync_remote() + remote = start_remote() + assert remote.stdin is not None and remote.stdout is not None + remote_info = read_json(remote.stdout, "QP information") + print("booting local") + dev = BNXTDev(PCIDevice("bnxt", LOCAL_PCI), ip=LOCAL_IP) + qp = BNXTQP(dev) + + send_line(remote.stdin, qp_info(dev, qp)) + wait_line(remote.stdout, "connected") + qp.connect(remote_info["qpn"], bytes.fromhex(remote_info["gid"]), int(remote_info["mac"], 16)) + print("both QPs in RTS") + + remote_target = read_json(remote.stdout, "MR information") + source, source_paddrs = dev.pci_dev.alloc_sysmem(0x1000) + source[:len(MESSAGE)] = MESSAGE + lkey = dev.register_mem(source_paddrs, 0x1000) + print(f"RDMA WRITE {len(MESSAGE)}B to remote phys 0x{remote_target['target_addr']:x}") + qp.rdma_write(remote_target["target_addr"], remote_target["rkey"], source_paddrs[0], lkey, len(MESSAGE)) + + send_line(remote.stdin, "done") + wait_line(remote.stdout, "AS TEXT") + result = read_json(remote.stdout, "RDMA result") + assert bytes.fromhex(result["data"]) == MESSAGE + print("RDMA WRITE data verified") + + remote.stdin.close() + assert remote.wait() == 0 + print("RDMA WRITE test complete") + +if __name__ == "__main__": + server() if "--server" in sys.argv else client() diff --git a/extra/bnxt_driver/loopback.py b/extra/bnxt_driver/loopback.py new file mode 100644 index 0000000000..42b25a44a3 --- /dev/null +++ b/extra/bnxt_driver/loopback.py @@ -0,0 +1,46 @@ +#!/usr/bin/env python3 +"""Local BNXT RoCEv2 RDMA WRITE loopback using the firmware's PHY loopback mode. + +The kernel bnxt_en/bnxt_re modules must be unloaded first. + + sudo PYTHONPATH=. BNXT_PCI=0000:41:00.0 BNXT_IP=10.0.200.5 python3 extra/bnxt_driver/loopback.py +""" +import os +import sys +import time + +sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "../..")) + +from extra.bnxt_driver.bnxtdev import BNXTDev, BNXTQP +from tinygrad.runtime.autogen import bnxt +from tinygrad.runtime.support.system import PCIDevice + +BUF_SIZE = 0x1000 +BNXT_PCI = os.getenv("BNXT_PCI", "0000:41:00.0") +BNXT_IP = os.getenv("BNXT_IP", "10.0.200.5") + +if __name__ == "__main__": + print(f"[init] BNXT at {BNXT_PCI}") + dev = BNXTDev(PCIDevice("bnxt", BNXT_PCI), ip=BNXT_IP) + tx_qp, rx_qp = BNXTQP(dev), BNXTQP(dev) + print(f"[init] loopback-connect TX QP 0x{tx_qp.qpn:x} <-> RX QP 0x{rx_qp.qpn:x}") + tx_qp.connect(rx_qp.qpn, dev.local_gid, dev.mac) + rx_qp.connect(tx_qp.qpn, dev.local_gid, dev.mac) + + src, src_paddrs = dev.pci_dev.alloc_sysmem(BUF_SIZE) + dst, dst_paddrs = dev.pci_dev.alloc_sysmem(BUF_SIZE) + message = b"Hello from BNXT RoCE PHY loopback!" + src[:BUF_SIZE], dst[:BUF_SIZE] = bytes(BUF_SIZE), bytes(BUF_SIZE) + src[:len(message)] = message + lkey = dev.register_mem(src_paddrs, BUF_SIZE) + rkey = dev.register_mem(dst_paddrs, BUF_SIZE) + + print("[loopback] enabling local PHY loopback") + dev.hwrm("port_phy_cfg", port_id=dev.port_id, enables=bnxt.PORT_PHY_CFG_REQ_ENABLES_LPBK, lpbk=bnxt.PORT_PHY_CFG_REQ_LPBK_LOCAL) + time.sleep(1) + tx_qp.rdma_write(dst_paddrs[0], rkey, src_paddrs[0], lkey, len(message)) + got = bytes(dst[:len(message)]) + print(f"[result] {got!r}") + assert got == message + print("BNXT RoCE PHY loopback RDMA WRITE passed") + dev.hwrm("port_phy_cfg", port_id=dev.port_id, enables=bnxt.PORT_PHY_CFG_REQ_ENABLES_LPBK, lpbk=bnxt.PORT_PHY_CFG_REQ_LPBK_NONE) diff --git a/test/unit/test_bnxt.py b/test/unit/test_bnxt.py new file mode 100644 index 0000000000..0716da3dae --- /dev/null +++ b/test/unit/test_bnxt.py @@ -0,0 +1,115 @@ +import struct, unittest +from types import SimpleNamespace +from unittest.mock import patch + +from tinygrad.runtime.autogen import bnxt +from extra.bnxt_driver.bnxtdev import BNXT_BACKING_STORE, BNXTDev, BNXTQP, _queue, _qwrite, ipv4_to_gid + +class FakePCI: + def __init__(self): self.next_addr, self.allocations = 0x100000, [] + def alloc_sysmem(self, size, contiguous=False): + pages = [self.next_addr+i*0x1000 for i in range((size+0xfff)//0x1000)] + self.next_addr += len(pages)*0x1000 + self.allocations.append(mem := bytearray(size)) + return mem, pages + +class FakeDev: + def __init__(self): self.pci_dev, self.calls = FakePCI(), [] + def hwrm(self, name, **fields): + self.calls.append((name, fields)) + typ = fields.get("type", 0) + return SimpleNamespace(ctx_init_value=0x5a, ctx_init_offset=4, entry_size=16 if typ == 0 else 4, + subtype_valid_cnt=typ == 0, split_entry_0=2, instance_bit_map=5 if typ == 0 else 1, min_num_entries=0) + +class FakeRCFW: + def __init__(self): self.calls, self.doorbells = [], [] + def exec(self, name, **fields): + self.calls.append((name, fields)) + return SimpleNamespace(xid={"create_cq":77, "create_qp":88, "register_mr":0x5678}.get(name, 0)) + def doorbell(self, *args, **kwargs): self.doorbells.append((args, kwargs)) + +class FakeQPDev: + def __init__(self): self.pci_dev, self.fw, self.gid_id, self.nq_id = FakePCI(), FakeRCFW(), 9, 41 + def rcfw(self, *args, **kwargs): return self.fw.exec(*args, **kwargs) + def doorbell(self, *args, **kwargs): self.fw.doorbell(*args, **kwargs) + +class TestMemory(unittest.TestCase): + def test_cmdq_and_sq_aux(self): + dev = FakeDev() + cmdq, sq = _queue(dev), _queue(dev, aux=True) + self.assertEqual((cmdq["level"], cmdq["base"]), (0, 0x100000)) + _qwrite(sq, 3, b"ABCDEFGH", aux=True) + self.assertEqual(bytes(sq["mem"][0x1018:0x1020]), b"ABCDEFGH") + + def test_f320_backing_layout_and_final_marker(self): + self.assertEqual(len(BNXT_BACKING_STORE), 9) + dev = FakeDev() + small = ((0, 6), (15, 0)) + with patch("extra.bnxt_driver.bnxtdev.BNXT_BACKING_STORE", small): BNXTDev.setup_backing_store(dev) + cfg = [fields for name, fields in dev.calls if name == "func_backing_store_cfg_v2"] + self.assertEqual([(x["type"], x["instance"]) for x in cfg], [(0, 0), (0, 2), (15, 0)]) + self.assertTrue(all(not x["flags"] for x in cfg[:-1])) + self.assertEqual(cfg[-1]["flags"], bnxt.FUNC_BACKING_STORE_CFG_V2_REQ_FLAGS_BS_CFG_ALL_DONE) + self.assertEqual((dev.pci_dev.allocations[0][4], dev.pci_dev.allocations[0][20]), (0x5a, 0x5a)) + +class TestRCFW(unittest.TestCase): + def setUp(self): + patch("extra.bnxt_driver.bnxtdev.System.memory_barrier").start() + self.addCleanup(patch.stopall) + + def test_doorbell_encodes_xid_type_and_index(self): + dev = BNXTDev.__new__(BNXTDev) + dev.db, dev.db_off = [0]*1024, 0x1000 + dev.doorbell(0x123456, bnxt.DBC_DBC_TYPE_CQ_ARMALL, 0x456, epoch=1) + key = dev.db[0x1000//8] + self.assertEqual(key >> 32, + 0x123456 & bnxt.DBC_DBC_XID_MASK | bnxt.DBC_DBC_PATH_ROCE | bnxt.DBC_DBC_TYPE_CQ_ARMALL | bnxt.BNXT_QPLIB_DBR_VALID) + self.assertEqual(key & 0xffffffff, 0x456 | 1<> 20) ^ ((lqpn * rqpn) >> 40)) & 0xFFFFF return ((v & 0x3FFF) ^ ((v & 0xFC000) >> 14)) | 0xC000 diff --git a/tinygrad/runtime/support/system.py b/tinygrad/runtime/support/system.py index 1f11403d30..b6bd184aad 100644 --- a/tinygrad/runtime/support/system.py +++ b/tinygrad/runtime/support/system.py @@ -10,6 +10,8 @@ from tinygrad.runtime.support.usb import USB3, CustomASM24Controller, USBMMIOInt MAP_FIXED, MAP_FIXED_NOREPLACE = 0x10, 0x100000 MAP_LOCKED, MAP_POPULATE, MAP_NORESERVE = 0 if OSX else 0x2000, getattr(mmap, "MAP_POPULATE", 0 if OSX else 0x008000), 0x400 +def ipv4_to_gid(ip:str) -> bytes: return bytes(10) + b'\xff\xff' + socket.inet_aton(ip) + class _System: def write_sysfs(self, path:str, value:str, msg:str, expected:str|None=None): if FileIOInterface(path, os.O_RDONLY).read().splitlines()[0] != (expected or value):