From 4571979fac2521572010771fedd0f8de70a806d7 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Wed, 31 Dec 2025 21:33:37 +0000 Subject: [PATCH] refactor --- extra/assembly/amd/asm.py | 9 +- extra/assembly/amd/autogen/rdna4/ins.py | 17 +- extra/assembly/amd/test/test_llvm.py | 397 ++++++++---------------- 3 files changed, 141 insertions(+), 282 deletions(-) diff --git a/extra/assembly/amd/asm.py b/extra/assembly/amd/asm.py index 838b471249..2abb3cbad3 100644 --- a/extra/assembly/amd/asm.py +++ b/extra/assembly/amd/asm.py @@ -509,6 +509,13 @@ def _disasm_ldsdir(inst) -> str: if inst.op == 0: return f"lds_param_load v{inst.vdst}, attr{inst.attr}.{['x','y','z','w'][inst.attr_chan]}{wait}" raise ValueError(f"unknown LDSDIR op: {inst.op}") +def _disasm_vdsdir(inst) -> str: + wait_va = f" wait_va_vdst:{inst.wait_va}" if inst.wait_va != 0 else "" + wait_vm = f" wait_vm_vsrc:{inst.wait_vm}" if inst.wait_vm != 0 else "" + if inst.op == 1: return f"ds_direct_load v{inst.vdst}{wait_va}{wait_vm}" + if inst.op == 0: return f"ds_param_load v{inst.vdst}, attr{inst.attr}.{['x','y','z','w'][inst.attr_chan]}{wait_va}{wait_vm}" + raise ValueError(f"unknown VDSDIR op: {inst.op}") + def _disasm_vexport(inst) -> str: target = _EXP_TARGETS.get(inst.target, f"invalid_target_{inst.target}") en = inst.en @@ -584,7 +591,7 @@ _DISASM_BY_NAME = { 'VOPD': _disasm_vopd, 'VOP3P': _disasm_vop3p, 'VINTERP': _disasm_vinterp, 'SOPP': _disasm_sopp, 'SMEM': _disasm_smem, 'DS': _disasm_ds, 'VDS': _disasm_ds, 'FLAT': _disasm_flat, 'MUBUF': _disasm_buf, 'MTBUF': _disasm_buf, 'MIMG': _disasm_mimg, 'SOP1': _disasm_sop1, 'SOP2': _disasm_sop2, 'SOPC': _disasm_sopc, 'SOPK': _disasm_sopk, - 'VEXPORT': _disasm_vexport, 'EXP': _disasm_vexport, 'LDSDIR': _disasm_ldsdir, + 'VEXPORT': _disasm_vexport, 'EXP': _disasm_vexport, 'LDSDIR': _disasm_ldsdir, 'VDSDIR': _disasm_vdsdir, 'VBUFFER': _disasm_vbuffer, 'VFLAT': _disasm_vflat, 'VGLOBAL': _disasm_vflat, 'VSCRATCH': _disasm_vflat, 'VSAMPLE': _disasm_vsample, 'VIMAGE': _disasm_vimage, } diff --git a/extra/assembly/amd/autogen/rdna4/ins.py b/extra/assembly/amd/autogen/rdna4/ins.py index aff5662146..b1609ded0a 100644 --- a/extra/assembly/amd/autogen/rdna4/ins.py +++ b/extra/assembly/amd/autogen/rdna4/ins.py @@ -94,17 +94,14 @@ class VDS(Inst64): data1:VGPRField = bits[55:48] vdst:VGPRField = bits[63:56] -class VDSDIR(Inst64): - encoding = bits[31:24] == 0b11001101 +class VDSDIR(Inst32): + encoding = bits[31:24] == 0b11001110 vdst:VGPRField = bits[7:0] - waitexp = bits[10:8] - opsel = bits[14:11] - cm = bits[15] - op:Annotated[BitField, VDSDIROp] = bits[20:16] - src0:Src = bits[40:32] - src1:Src = bits[49:41] - src2:Src = bits[58:50] - neg = bits[63:61] + attr_chan = bits[9:8] + attr = bits[15:10] + wait_va = bits[19:16] + op:Annotated[BitField, VDSDIROp] = bits[21:20] + wait_vm = bits[23] # single bit, bit 22 is reserved class VEXPORT(Inst64): encoding = bits[31:26] == 0b111110 diff --git a/extra/assembly/amd/test/test_llvm.py b/extra/assembly/amd/test/test_llvm.py index 7ed4e331a8..1685e9ec36 100644 --- a/extra/assembly/amd/test/test_llvm.py +++ b/extra/assembly/amd/test/test_llvm.py @@ -1,123 +1,48 @@ #!/usr/bin/env python3 """Test RDNA3/RDNA4 assembler/disassembler against LLVM test vectors.""" -import unittest, re, subprocess, os +import unittest, re, subprocess from tinygrad.helpers import fetch from extra.assembly.amd.test.helpers import get_llvm_mc LLVM_BASE = "https://raw.githubusercontent.com/llvm/llvm-project/main/llvm/test/MC/AMDGPU" -# ═══════════════════════════════════════════════════════════════════════════════ -# RDNA3 (GFX11) TEST FILES -# ═══════════════════════════════════════════════════════════════════════════════ - RDNA3_TEST_FILES = { - # Scalar ALU - 'sop1': 'gfx11_asm_sop1.s', - 'sop2': 'gfx11_asm_sop2.s', - 'sopp': 'gfx11_asm_sopp.s', - 'sopk': 'gfx11_asm_sopk.s', - 'sopc': 'gfx11_asm_sopc.s', - # Vector ALU - 'vop1': 'gfx11_asm_vop1.s', - 'vop2': 'gfx11_asm_vop2.s', - 'vopc': 'gfx11_asm_vopc.s', - 'vop3': 'gfx11_asm_vop3.s', - 'vop3p': 'gfx11_asm_vop3p.s', - 'vinterp': 'gfx11_asm_vinterp.s', - 'vopd': 'gfx11_asm_vopd.s', - 'vopcx': 'gfx11_asm_vopcx.s', - # VOP3 promotions - 'vop3_from_vop1': 'gfx11_asm_vop3_from_vop1.s', - 'vop3_from_vop2': 'gfx11_asm_vop3_from_vop2.s', - 'vop3_from_vopc': 'gfx11_asm_vop3_from_vopc.s', - 'vop3_from_vopcx': 'gfx11_asm_vop3_from_vopcx.s', - # Memory - 'ds': 'gfx11_asm_ds.s', - 'smem': 'gfx11_asm_smem.s', - 'flat': 'gfx11_asm_flat.s', - 'mubuf': 'gfx11_asm_mubuf.s', - 'mtbuf': 'gfx11_asm_mtbuf.s', - 'mimg': 'gfx11_asm_mimg.s', - 'ldsdir': 'gfx11_asm_ldsdir.s', - # Export - 'exp': 'gfx11_asm_exp.s', - # WMMA - 'wmma': 'gfx11_asm_wmma.s', - # Features - 'vop3_features': 'gfx11_asm_vop3_features.s', - 'vop3p_features': 'gfx11_asm_vop3p_features.s', - 'vopd_features': 'gfx11_asm_vopd_features.s', - # Alias files - 'vop3_alias': 'gfx11_asm_vop3_alias.s', - 'vop3p_alias': 'gfx11_asm_vop3p_alias.s', - 'vopc_alias': 'gfx11_asm_vopc_alias.s', - 'vopcx_alias': 'gfx11_asm_vopcx_alias.s', - 'vinterp_alias': 'gfx11_asm_vinterp_alias.s', - 'smem_alias': 'gfx11_asm_smem_alias.s', - 'mubuf_alias': 'gfx11_asm_mubuf_alias.s', - 'mtbuf_alias': 'gfx11_asm_mtbuf_alias.s', + 'sop1': 'gfx11_asm_sop1.s', 'sop2': 'gfx11_asm_sop2.s', 'sopp': 'gfx11_asm_sopp.s', 'sopk': 'gfx11_asm_sopk.s', 'sopc': 'gfx11_asm_sopc.s', + 'vop1': 'gfx11_asm_vop1.s', 'vop2': 'gfx11_asm_vop2.s', 'vopc': 'gfx11_asm_vopc.s', 'vop3': 'gfx11_asm_vop3.s', 'vop3p': 'gfx11_asm_vop3p.s', + 'vinterp': 'gfx11_asm_vinterp.s', 'vopd': 'gfx11_asm_vopd.s', 'vopcx': 'gfx11_asm_vopcx.s', + 'vop3_from_vop1': 'gfx11_asm_vop3_from_vop1.s', 'vop3_from_vop2': 'gfx11_asm_vop3_from_vop2.s', + 'vop3_from_vopc': 'gfx11_asm_vop3_from_vopc.s', 'vop3_from_vopcx': 'gfx11_asm_vop3_from_vopcx.s', + 'ds': 'gfx11_asm_ds.s', 'smem': 'gfx11_asm_smem.s', 'flat': 'gfx11_asm_flat.s', + 'mubuf': 'gfx11_asm_mubuf.s', 'mtbuf': 'gfx11_asm_mtbuf.s', 'mimg': 'gfx11_asm_mimg.s', 'ldsdir': 'gfx11_asm_ldsdir.s', + 'exp': 'gfx11_asm_exp.s', 'wmma': 'gfx11_asm_wmma.s', + 'vop3_features': 'gfx11_asm_vop3_features.s', 'vop3p_features': 'gfx11_asm_vop3p_features.s', 'vopd_features': 'gfx11_asm_vopd_features.s', + 'vop3_alias': 'gfx11_asm_vop3_alias.s', 'vop3p_alias': 'gfx11_asm_vop3p_alias.s', 'vopc_alias': 'gfx11_asm_vopc_alias.s', + 'vopcx_alias': 'gfx11_asm_vopcx_alias.s', 'vinterp_alias': 'gfx11_asm_vinterp_alias.s', + 'smem_alias': 'gfx11_asm_smem_alias.s', 'mubuf_alias': 'gfx11_asm_mubuf_alias.s', 'mtbuf_alias': 'gfx11_asm_mtbuf_alias.s', } -# ═══════════════════════════════════════════════════════════════════════════════ -# RDNA4 (GFX12) TEST FILES -# ═══════════════════════════════════════════════════════════════════════════════ - RDNA4_TEST_FILES = { - # Scalar ALU - 'sop1': 'gfx12_asm_sop1.s', - 'sop2': 'gfx12_asm_sop2.s', - 'sop2_alias': 'gfx12_asm_sop2_alias.s', - 'sopp': 'gfx12_asm_sopp.s', - 'sopk': 'gfx12_asm_sopk.s', - 'sopk_alias': 'gfx12_asm_sopk_alias.s', - 'sopc': 'gfx12_asm_sopc.s', - # Vector ALU - 'vop1': 'gfx12_asm_vop1.s', - 'vop2': 'gfx12_asm_vop2.s', - 'vop2_aliases': 'gfx12_asm_vop2_aliases.s', - 'vopc': 'gfx12_asm_vopc.s', - 'vopcx': 'gfx12_asm_vopcx.s', - 'vop3': 'gfx12_asm_vop3.s', - 'vop3_aliases': 'gfx12_asm_vop3_aliases.s', - 'vop3c': 'gfx12_asm_vop3c.s', - 'vop3cx': 'gfx12_asm_vop3cx.s', - 'vop3p': 'gfx12_asm_vop3p.s', - 'vop3p_aliases': 'gfx12_asm_vop3p_aliases.s', - 'vop3p_features': 'gfx12_asm_vop3p_features.s', - 'vopd': 'gfx12_asm_vopd.s', - 'vopd_features': 'gfx12_asm_vopd_features.s', - # VOP3 promotions - 'vop3_from_vop1': 'gfx12_asm_vop3_from_vop1.s', - 'vop3_from_vop2': 'gfx12_asm_vop3_from_vop2.s', - # Memory - 'ds': 'gfx12_asm_ds.s', - 'ds_alias': 'gfx12_asm_ds_alias.s', - 'smem': 'gfx12_asm_smem.s', - 'vflat': 'gfx12_asm_vflat.s', - 'vflat_alias': 'gfx12_asm_vflat_alias.s', - 'vbuffer_mubuf': 'gfx12_asm_vbuffer_mubuf.s', - 'vbuffer_mubuf_alias': 'gfx12_asm_vbuffer_mubuf_alias.s', - 'vbuffer_mtbuf': 'gfx12_asm_vbuffer_mtbuf.s', - 'vbuffer_mtbuf_alias': 'gfx12_asm_vbuffer_mtbuf_alias.s', - 'vimage': 'gfx12_asm_vimage.s', - 'vimage_alias': 'gfx12_asm_vimage_alias.s', - 'vsample': 'gfx12_asm_vsample.s', - 'vdsdir': 'gfx12_asm_vdsdir.s', - 'vdsdir_alias': 'gfx12_asm_vdsdir_alias.s', - # Export - 'exp': 'gfx12_asm_exp.s', - # WMMA - 'wmma_w32': 'gfx12_asm_wmma_w32.s', - 'wmma_w64': 'gfx12_asm_wmma_w64.s', - # Features - 'features': 'gfx12_asm_features.s', - 'global_load_tr': 'gfx12_asm_global_load_tr.s', + 'sop1': 'gfx12_asm_sop1.s', 'sop2': 'gfx12_asm_sop2.s', 'sop2_alias': 'gfx12_asm_sop2_alias.s', + 'sopp': 'gfx12_asm_sopp.s', 'sopk': 'gfx12_asm_sopk.s', 'sopk_alias': 'gfx12_asm_sopk_alias.s', 'sopc': 'gfx12_asm_sopc.s', + 'vop1': 'gfx12_asm_vop1.s', 'vop2': 'gfx12_asm_vop2.s', 'vop2_aliases': 'gfx12_asm_vop2_aliases.s', + 'vopc': 'gfx12_asm_vopc.s', 'vopcx': 'gfx12_asm_vopcx.s', + 'vop3': 'gfx12_asm_vop3.s', 'vop3_aliases': 'gfx12_asm_vop3_aliases.s', 'vop3c': 'gfx12_asm_vop3c.s', 'vop3cx': 'gfx12_asm_vop3cx.s', + 'vop3p': 'gfx12_asm_vop3p.s', 'vop3p_aliases': 'gfx12_asm_vop3p_aliases.s', 'vop3p_features': 'gfx12_asm_vop3p_features.s', + 'vopd': 'gfx12_asm_vopd.s', 'vopd_features': 'gfx12_asm_vopd_features.s', + 'vop3_from_vop1': 'gfx12_asm_vop3_from_vop1.s', 'vop3_from_vop2': 'gfx12_asm_vop3_from_vop2.s', + 'ds': 'gfx12_asm_ds.s', 'ds_alias': 'gfx12_asm_ds_alias.s', 'smem': 'gfx12_asm_smem.s', + 'vflat': 'gfx12_asm_vflat.s', 'vflat_alias': 'gfx12_asm_vflat_alias.s', + 'vbuffer_mubuf': 'gfx12_asm_vbuffer_mubuf.s', 'vbuffer_mubuf_alias': 'gfx12_asm_vbuffer_mubuf_alias.s', + 'vbuffer_mtbuf': 'gfx12_asm_vbuffer_mtbuf.s', 'vbuffer_mtbuf_alias': 'gfx12_asm_vbuffer_mtbuf_alias.s', + 'vimage': 'gfx12_asm_vimage.s', 'vimage_alias': 'gfx12_asm_vimage_alias.s', 'vsample': 'gfx12_asm_vsample.s', + 'vdsdir': 'gfx12_asm_vdsdir.s', 'vdsdir_alias': 'gfx12_asm_vdsdir_alias.s', + 'exp': 'gfx12_asm_exp.s', 'wmma_w32': 'gfx12_asm_wmma_w32.s', 'wmma_w64': 'gfx12_asm_wmma_w64.s', + 'features': 'gfx12_asm_features.s', 'global_load_tr': 'gfx12_asm_global_load_tr.s', } -def parse_llvm_tests(text: str, gfx_prefix: str = "GFX11") -> list[tuple[str, bytes]]: +def parse_llvm_tests(text: str, gfx_prefix: str) -> list[tuple[str, bytes]]: """Parse LLVM test format into (asm, expected_bytes) pairs.""" tests, lines = [], text.split('\n') - # Support GFX11, GFX12, W32, W64 encodings pattern = rf'(?:{gfx_prefix}|W32|W64)[^:]*:.*?encoding:\s*\[(.*?)\]' pattern2 = rf'(?:{gfx_prefix}|W32|W64)[^:]*:\s*\[(0x[0-9a-fA-F,x\s]+)\]' for i, line in enumerate(lines): @@ -130,22 +55,19 @@ def parse_llvm_tests(text: str, gfx_prefix: str = "GFX11") -> list[tuple[str, by hex_bytes = m.group(1).replace('0x', '').replace(',', '').replace(' ', '') elif m := re.search(pattern2, lines[j]): hex_bytes = m.group(1).replace('0x', '').replace(',', '').replace(' ', '') - else: - continue + else: continue if hex_bytes: try: tests.append((asm_text, bytes.fromhex(hex_bytes))) except ValueError: pass break return tests -def compile_asm_batch(instrs: list[str], mcpu: str = 'gfx1100', mattr: str = '+real-true16,+wavefrontsize32') -> list[bytes]: +def compile_asm_batch(instrs: list[str], mcpu: str, mattr: str = '+real-true16,+wavefrontsize32') -> list[bytes]: """Compile multiple instructions with a single llvm-mc call.""" if not instrs: return [] - asm_text = ".text\n" + "\n".join(instrs) + "\n" - result = subprocess.run( - [get_llvm_mc(), '-triple=amdgcn', f'-mcpu={mcpu}', f'-mattr={mattr}', '-show-encoding'], - input=asm_text, capture_output=True, text=True, timeout=30) - if result.returncode != 0: raise RuntimeError(f"llvm-mc batch failed: {result.stderr.strip()}") + result = subprocess.run([get_llvm_mc(), '-triple=amdgcn', f'-mcpu={mcpu}', f'-mattr={mattr}', '-show-encoding'], + input=".text\n" + "\n".join(instrs) + "\n", capture_output=True, text=True, timeout=30) + if result.returncode != 0: raise RuntimeError(f"llvm-mc failed: {result.stderr.strip()}") results = [] for line in result.stdout.split('\n'): if 'encoding:' not in line: continue @@ -155,17 +77,76 @@ def compile_asm_batch(instrs: list[str], mcpu: str = 'gfx1100', mattr: str = '+r if len(results) != len(instrs): raise RuntimeError(f"expected {len(instrs)} encodings, got {len(results)}") return results -# ═══════════════════════════════════════════════════════════════════════════════ -# RDNA3 TESTS -# ═══════════════════════════════════════════════════════════════════════════════ +def matches_encoding(data: bytes, fmt) -> bool: + """Check if instruction bytes match format's expected encoding bits.""" + if not hasattr(fmt, '_encoding') or fmt._encoding is None: return True + bf, expected = fmt._encoding + val = int.from_bytes(data[:fmt._size()], 'little') + return ((val >> bf.lo) & bf.mask()) == expected -class TestLLVMRDNA3(unittest.TestCase): - """Test RDNA3 assembler against LLVM test vectors (disassembly only for now).""" +class TestLLVMBase(unittest.TestCase): + """Base class for LLVM assembler tests.""" tests: dict[str, list[tuple[str, bytes]]] = {} + formats: dict[str, type] = {} + gfx_prefix: str = "" + mcpu: str = "" + arch_name: str = "" + + @classmethod + def _load_tests(cls, test_files: dict[str, str]): + for name, filename in test_files.items(): + try: + data = fetch(f"{LLVM_BASE}/{filename}").read_bytes() + cls.tests[name] = parse_llvm_tests(data.decode('utf-8', errors='ignore'), cls.gfx_prefix) + except Exception as e: + print(f"Warning: couldn't fetch {filename}: {e}") + cls.tests[name] = [] + + def _test_disasm(self, name: str): + """Test decoding instructions and verify disassembly produces correct bytes.""" + if name not in self.tests or not self.tests[name]: self.skipTest(f"No test data for {name}") + fmt_cls = self.formats.get(name) + if fmt_cls is None: self.skipTest(f"No format class for {name}") + + to_test: list[tuple[str, bytes, str | None, str | None]] = [] + for asm_text, data in self.tests.get(name, []): + if len(data) > fmt_cls._size(): continue + if not matches_encoding(data, fmt_cls): continue + try: + decoded = fmt_cls.from_bytes(data) + if decoded.to_bytes()[:len(data)] != data: + to_test.append((asm_text, data, None, "decode roundtrip failed")) + continue + to_test.append((asm_text, data, decoded.disasm(), None)) + except Exception as e: + to_test.append((asm_text, data, None, f"exception: {e}")) + + disasm_strs = [(i, t[2]) for i, t in enumerate(to_test) if t[2] is not None] + llvm_map = {} + if disasm_strs: + llvm_results = compile_asm_batch([s for _, s in disasm_strs], self.mcpu) + llvm_map = {i: llvm_results[j] for j, (i, _) in enumerate(disasm_strs)} + + passed, failed, failures = 0, 0, [] + for idx, (asm_text, data, disasm_str, error) in enumerate(to_test): + if error: + failed += 1; failures.append(f"{error} for {data.hex()}") + elif disasm_str is not None and idx in llvm_map: + llvm_bytes = llvm_map[idx] + if llvm_bytes == data: passed += 1 + else: failed += 1; failures.append(f"'{disasm_str}': expected={data.hex()} got={llvm_bytes.hex()}") + + print(f"{self.arch_name} {name.upper()} disasm: {passed} passed, {failed} failed") + if failures[:5]: print(" " + "\n ".join(failures[:5])) + self.assertGreater(passed, 0, f"No tests passed for {name}") + +class TestLLVMRDNA3(TestLLVMBase): + """Test RDNA3 assembler against LLVM test vectors.""" + gfx_prefix, mcpu, arch_name = "GFX11", "gfx1100", "RDNA3" @classmethod def setUpClass(cls): - from extra.assembly.amd.autogen.rdna3.ins import SOP1, SOP2, SOPC, SOPK, SOPP, VOP1, VOP2, VOP3, VOP3SD, VOP3P, VOPC, VOPD, VINTERP, DS, SMEM, FLAT, MUBUF, MTBUF, MIMG, LDSDIR, EXP + from extra.assembly.amd.autogen.rdna3.ins import SOP1, SOP2, SOPC, SOPK, SOPP, VOP1, VOP2, VOP3, VOP3P, VOPC, VOPD, VINTERP, DS, SMEM, FLAT, MUBUF, MTBUF, MIMG, LDSDIR, EXP cls.formats = { 'sop1': SOP1, 'sop2': SOP2, 'sopc': SOPC, 'sopk': SOPK, 'sopp': SOPP, 'vop1': VOP1, 'vop2': VOP2, 'vopc': VOPC, 'vopcx': VOPC, 'vop3': VOP3, 'vop3p': VOP3P, @@ -176,167 +157,41 @@ class TestLLVMRDNA3(unittest.TestCase): 'vop3_alias': VOP3, 'vop3p_alias': VOP3P, 'vopc_alias': VOPC, 'vopcx_alias': VOPC, 'vinterp_alias': VINTERP, 'smem_alias': SMEM, 'mubuf_alias': MUBUF, 'mtbuf_alias': MTBUF, } - for name, filename in RDNA3_TEST_FILES.items(): - try: - data = fetch(f"{LLVM_BASE}/{filename}").read_bytes() - cls.tests[name] = parse_llvm_tests(data.decode('utf-8', errors='ignore'), "GFX11") - except Exception as e: - print(f"Warning: couldn't fetch {filename}: {e}") - cls.tests[name] = [] + cls._load_tests(RDNA3_TEST_FILES) - def _test_disasm(self, name: str): - """Test decoding instructions and verify disassembly produces correct bytes.""" - if name not in self.tests or not self.tests[name]: - self.skipTest(f"No test data for {name}") - fmt_cls = self.formats.get(name) - if fmt_cls is None: - self.skipTest(f"No format class for {name}") - - to_test: list[tuple[str, bytes, str | None, str | None]] = [] - for asm_text, data in self.tests.get(name, []): - if len(data) > fmt_cls._size(): continue - try: - decoded = fmt_cls.from_bytes(data) - if decoded.to_bytes()[:len(data)] != data: - to_test.append((asm_text, data, None, "decode roundtrip failed")) - continue - to_test.append((asm_text, data, decoded.disasm(), None)) - except Exception as e: - to_test.append((asm_text, data, None, f"exception: {e}")) - - # Batch compile disasm strings - disasm_strs = [(i, t[2]) for i, t in enumerate(to_test) if t[2] is not None] - if disasm_strs: - llvm_results = compile_asm_batch([s for _, s in disasm_strs], 'gfx1100', '+real-true16,+wavefrontsize32') - llvm_map = {i: llvm_results[j] for j, (i, _) in enumerate(disasm_strs)} - else: - llvm_map = {} - - passed, failed = 0, 0 - failures: list[str] = [] - for idx, (asm_text, data, disasm_str, error) in enumerate(to_test): - if error: - failed += 1; failures.append(f"{error} for {data.hex()}") - elif disasm_str is not None and idx in llvm_map: - llvm_bytes = llvm_map[idx] - if llvm_bytes is not None and llvm_bytes == data: passed += 1 - elif llvm_bytes is not None: failed += 1; failures.append(f"'{disasm_str}': expected={data.hex()} got={llvm_bytes.hex()}") - - print(f"RDNA3 {name.upper()} disasm: {passed} passed, {failed} failed") - if failures[:5]: print(" " + "\n ".join(failures[:5])) - self.assertGreater(passed, 0, f"No tests passed for {name}") - -# Generate test methods dynamically for RDNA3 -def _make_rdna3_disasm_test(name): - def test(self): self._test_disasm(name) - return test - -for name in RDNA3_TEST_FILES: - setattr(TestLLVMRDNA3, f'test_{name}_disasm', _make_rdna3_disasm_test(name)) - -# ═══════════════════════════════════════════════════════════════════════════════ -# RDNA4 TESTS -# ═══════════════════════════════════════════════════════════════════════════════ - -class TestLLVMRDNA4(unittest.TestCase): - """Test RDNA4 assembler against LLVM test vectors (disassembly only for now).""" - tests: dict[str, list[tuple[str, bytes]]] = {} +class TestLLVMRDNA4(TestLLVMBase): + """Test RDNA4 assembler against LLVM test vectors.""" + gfx_prefix, mcpu, arch_name = "GFX12", "gfx1200", "RDNA4" @classmethod def setUpClass(cls): import extra.assembly.amd.autogen.rdna4.ins as rdna4 - # Get available formats, some may not be generated yet - def get_fmt(name): return getattr(rdna4, name, None) - SOP1, SOP2, SOPC, SOPK, SOPP = get_fmt('SOP1'), get_fmt('SOP2'), get_fmt('SOPC'), get_fmt('SOPK'), get_fmt('SOPP') - VOP1, VOP2, VOP3, VOP3SD, VOP3P, VOPC, VOPD = get_fmt('VOP1'), get_fmt('VOP2'), get_fmt('VOP3'), get_fmt('VOP3SD'), get_fmt('VOP3P'), get_fmt('VOPC'), get_fmt('VOPD') - VINTERP, VDS, SMEM = get_fmt('VINTERP'), get_fmt('VDS'), get_fmt('SMEM') - VEXPORT, VBUFFER, VDSDIR = get_fmt('VEXPORT'), get_fmt('VBUFFER'), get_fmt('VDSDIR') - VFLAT, VGLOBAL, VSCRATCH = get_fmt('VFLAT'), get_fmt('VGLOBAL'), get_fmt('VSCRATCH') - VIMAGE, VSAMPLE = get_fmt('VIMAGE'), get_fmt('VSAMPLE') - # Note: RDNA4 uses different format names (VDS instead of DS, VBUFFER instead of MUBUF, etc.) + get = lambda n: getattr(rdna4, n, None) cls.formats = { - 'sop1': SOP1, 'sop2': SOP2, 'sop2_alias': SOP2, 'sopc': SOPC, 'sopk': SOPK, 'sopk_alias': SOPK, 'sopp': SOPP, - 'vop1': VOP1, 'vop2': VOP2, 'vop2_aliases': VOP2, 'vopc': VOPC, 'vopcx': VOPC, - 'vop3': VOP3, 'vop3_aliases': VOP3, 'vop3c': VOP3, 'vop3cx': VOP3, 'vop3p': VOP3P, 'vop3p_aliases': VOP3P, 'vop3p_features': VOP3P, - 'vopd': VOPD, 'vopd_features': VOPD, - 'vop3_from_vop1': VOP3, 'vop3_from_vop2': VOP3, - 'ds': VDS, 'ds_alias': VDS, 'smem': SMEM, - 'vinterp': VINTERP, 'exp': VEXPORT, - 'vbuffer_mubuf': VBUFFER, 'vbuffer_mubuf_alias': VBUFFER, 'vbuffer_mtbuf': VBUFFER, 'vbuffer_mtbuf_alias': VBUFFER, - 'vdsdir': None, 'vdsdir_alias': None, # VDSDIR is 64-bit but ds_direct_load is 32-bit - 'vflat': VFLAT, 'vflat_alias': VFLAT, - 'vimage': VIMAGE, 'vimage_alias': VIMAGE, - 'vsample': VSAMPLE, - 'wmma_w32': VOP3P, 'wmma_w64': VOP3P, - 'features': None, # Generic features file - 'global_load_tr': VGLOBAL, # Uses VGLOBAL format + 'sop1': get('SOP1'), 'sop2': get('SOP2'), 'sop2_alias': get('SOP2'), 'sopc': get('SOPC'), + 'sopk': get('SOPK'), 'sopk_alias': get('SOPK'), 'sopp': get('SOPP'), + 'vop1': get('VOP1'), 'vop2': get('VOP2'), 'vop2_aliases': get('VOP2'), 'vopc': get('VOPC'), 'vopcx': get('VOPC'), + 'vop3': get('VOP3'), 'vop3_aliases': get('VOP3'), 'vop3c': get('VOP3'), 'vop3cx': get('VOP3'), + 'vop3p': get('VOP3P'), 'vop3p_aliases': get('VOP3P'), 'vop3p_features': get('VOP3P'), + 'vopd': get('VOPD'), 'vopd_features': get('VOPD'), + 'vop3_from_vop1': get('VOP3'), 'vop3_from_vop2': get('VOP3'), + 'ds': get('VDS'), 'ds_alias': get('VDS'), 'smem': get('SMEM'), 'vinterp': get('VINTERP'), 'exp': get('VEXPORT'), + 'vbuffer_mubuf': get('VBUFFER'), 'vbuffer_mubuf_alias': get('VBUFFER'), + 'vbuffer_mtbuf': get('VBUFFER'), 'vbuffer_mtbuf_alias': get('VBUFFER'), + 'vdsdir': get('VDSDIR'), 'vdsdir_alias': get('VDSDIR'), + 'vflat': get('VFLAT'), 'vflat_alias': get('VFLAT'), + 'vimage': get('VIMAGE'), 'vimage_alias': get('VIMAGE'), 'vsample': get('VSAMPLE'), + 'wmma_w32': get('VOP3P'), 'wmma_w64': get('VOP3P'), + 'features': None, 'global_load_tr': get('VGLOBAL'), } - for name, filename in RDNA4_TEST_FILES.items(): - try: - data = fetch(f"{LLVM_BASE}/{filename}").read_bytes() - cls.tests[name] = parse_llvm_tests(data.decode('utf-8', errors='ignore'), "GFX12") - except Exception as e: - print(f"Warning: couldn't fetch {filename}: {e}") - cls.tests[name] = [] + cls._load_tests(RDNA4_TEST_FILES) - def _test_disasm(self, name: str): - """Test decoding instructions and verify disassembly produces correct bytes.""" - if name not in self.tests or not self.tests[name]: - self.skipTest(f"No test data for {name}") - fmt_cls = self.formats.get(name) - if fmt_cls is None: - self.skipTest(f"No format class for {name}") - - # Check if instruction matches format encoding - def matches_encoding(data: bytes, fmt) -> bool: - if not hasattr(fmt, '_encoding') or fmt._encoding is None: return True - bf, expected = fmt._encoding - val = int.from_bytes(data[:fmt._size()], 'little') - actual = (val >> bf.lo) & bf.mask() - return actual == expected - - to_test: list[tuple[str, bytes, str | None, str | None]] = [] - for asm_text, data in self.tests.get(name, []): - if len(data) > fmt_cls._size(): continue - if not matches_encoding(data, fmt_cls): continue # Skip instructions with wrong encoding - try: - decoded = fmt_cls.from_bytes(data) - if decoded.to_bytes()[:len(data)] != data: - to_test.append((asm_text, data, None, "decode roundtrip failed")) - continue - to_test.append((asm_text, data, decoded.disasm(), None)) - except Exception as e: - to_test.append((asm_text, data, None, f"exception: {e}")) - - # Batch compile disasm strings - disasm_strs = [(i, t[2]) for i, t in enumerate(to_test) if t[2] is not None] - if disasm_strs: - llvm_results = compile_asm_batch([s for _, s in disasm_strs], 'gfx1200', '+real-true16,+wavefrontsize32') - llvm_map = {i: llvm_results[j] for j, (i, _) in enumerate(disasm_strs)} - else: - llvm_map = {} - - passed, failed = 0, 0 - failures: list[str] = [] - for idx, (asm_text, data, disasm_str, error) in enumerate(to_test): - if error: - failed += 1; failures.append(f"{error} for {data.hex()}") - elif disasm_str is not None and idx in llvm_map: - llvm_bytes = llvm_map[idx] - if llvm_bytes is not None and llvm_bytes == data: passed += 1 - elif llvm_bytes is not None: failed += 1; failures.append(f"'{disasm_str}': expected={data.hex()} got={llvm_bytes.hex()}") - - print(f"RDNA4 {name.upper()} disasm: {passed} passed, {failed} failed") - if failures[:5]: print(" " + "\n ".join(failures[:5])) - self.assertGreater(passed, 0, f"No tests passed for {name}") - -# Generate test methods dynamically for RDNA4 -def _make_rdna4_disasm_test(name): +# Generate test methods dynamically +def _make_test(name): def test(self): self._test_disasm(name) return test -for name in RDNA4_TEST_FILES: - setattr(TestLLVMRDNA4, f'test_{name}_disasm', _make_rdna4_disasm_test(name)) +for name in RDNA3_TEST_FILES: setattr(TestLLVMRDNA3, f'test_{name}_disasm', _make_test(name)) +for name in RDNA4_TEST_FILES: setattr(TestLLVMRDNA4, f'test_{name}_disasm', _make_test(name)) -if __name__ == "__main__": - unittest.main() +if __name__ == "__main__": unittest.main()