diff --git a/extra/assembly/rdna3/asm.py b/extra/assembly/rdna3/asm.py index ab1fe16d2a..65d57a63c7 100644 --- a/extra/assembly/rdna3/asm.py +++ b/extra/assembly/rdna3/asm.py @@ -1,11 +1,22 @@ # RDNA3 assembler and disassembler from __future__ import annotations import re -from extra.assembly.rdna3.lib import Inst, RawImm, Reg, SGPR, VGPR, TTMP, FLOAT_ENC, SRC_FIELDS, unwrap +from extra.assembly.rdna3.lib import Inst, RawImm, Reg, SGPR, VGPR, TTMP, s, v, ttmp, _RegFactory, FLOAT_ENC, SRC_FIELDS, unwrap # Decoding helpers SPECIAL_GPRS = {106: "vcc_lo", 107: "vcc_hi", 124: "null", 125: "m0", 126: "exec_lo", 127: "exec_hi", 253: "scc"} SPECIAL_DEC = {**SPECIAL_GPRS, **{v: str(k) for k, v in FLOAT_ENC.items()}} +SPECIAL_PAIRS = {106: "vcc", 126: "exec"} # Special register pairs (for 64-bit ops) +# GFX11 hwreg names (IDs 16-17 are TBA - not supported, IDs 18-19 are PERF_SNAPSHOT) +HWREG_NAMES = {1: 'HW_REG_MODE', 2: 'HW_REG_STATUS', 3: 'HW_REG_TRAPSTS', 4: 'HW_REG_HW_ID', 5: 'HW_REG_GPR_ALLOC', + 6: 'HW_REG_LDS_ALLOC', 7: 'HW_REG_IB_STS', 15: 'HW_REG_SH_MEM_BASES', 18: 'HW_REG_PERF_SNAPSHOT_PC_LO', + 19: 'HW_REG_PERF_SNAPSHOT_PC_HI', 20: 'HW_REG_FLAT_SCR_LO', 21: 'HW_REG_FLAT_SCR_HI', + 22: 'HW_REG_XNACK_MASK', 23: 'HW_REG_HW_ID1', 24: 'HW_REG_HW_ID2', 25: 'HW_REG_POPS_PACKER', 28: 'HW_REG_IB_STS2'} +HWREG_IDS = {v.lower(): k for k, v in HWREG_NAMES.items()} # Reverse map for assembler +MSG_NAMES = {128: 'MSG_RTN_GET_DOORBELL', 129: 'MSG_RTN_GET_DDID', 130: 'MSG_RTN_GET_TMA', + 131: 'MSG_RTN_GET_REALTIME', 132: 'MSG_RTN_SAVE_WAVE', 133: 'MSG_RTN_GET_TBA'} +_16BIT_TYPES = ('f16', 'i16', 'u16', 'b16') +def _is_16bit(s: str) -> bool: return any(s.endswith(x) for x in _16BIT_TYPES) def decode_src(val: int) -> str: if val <= 105: return f"s{val}" @@ -16,28 +27,39 @@ def decode_src(val: int) -> str: if 256 <= val <= 511: return f"v{val - 256}" return "lit" if val == 255 else f"?{val}" -def _sreg(base: int, cnt: int = 1) -> str: return f"s{base}" if cnt == 1 else f"s[{base}:{base+cnt-1}]" -def _vreg(base: int, cnt: int = 1) -> str: return f"v{base}" if cnt == 1 else f"v[{base}:{base+cnt-1}]" +def _reg(prefix: str, base: int, cnt: int = 1) -> str: return f"{prefix}{base}" if cnt == 1 else f"{prefix}[{base}:{base+cnt-1}]" +def _sreg(base: int, cnt: int = 1) -> str: return _reg("s", base, cnt) +def _vreg(base: int, cnt: int = 1) -> str: return _reg("v", base, cnt) def _fmt_sdst(v: int, cnt: int = 1) -> str: """Format SGPR destination with special register names.""" if v == 124: return "null" - if 108 <= v <= 123: return f"ttmp[{v-108}:{v-108+cnt-1}]" if cnt > 1 else f"ttmp{v-108}" - if cnt > 1: - if v == 126 and cnt == 2: return "exec" - if v == 106 and cnt == 2: return "vcc" - return _sreg(v, cnt) + if 108 <= v <= 123: return _reg("ttmp", v - 108, cnt) + if cnt > 1 and v in SPECIAL_PAIRS: return SPECIAL_PAIRS[v] + if cnt > 1: return _sreg(v, cnt) return {126: "exec_lo", 127: "exec_hi", 106: "vcc_lo", 107: "vcc_hi", 125: "m0"}.get(v, f"s{v}") def _fmt_ssrc(v: int, cnt: int = 1) -> str: """Format SGPR source with special register names and pairs.""" if cnt == 2: - if v == 126: return "exec" - if v == 106: return "vcc" + if v in SPECIAL_PAIRS: return SPECIAL_PAIRS[v] if v <= 105: return _sreg(v, 2) - if 108 <= v <= 123: return f"ttmp[{v-108}:{v-108+1}]" + if 108 <= v <= 123: return _reg("ttmp", v - 108, 2) return decode_src(v) +def _fmt_src_n(v: int, cnt: int) -> str: + """Format source with given register count (1, 2, or 4).""" + if cnt == 1: return decode_src(v) + if v >= 256: return _vreg(v - 256, cnt) + if v <= 105: return _sreg(v, cnt) + if cnt == 2 and v in SPECIAL_PAIRS: return SPECIAL_PAIRS[v] + if 108 <= v <= 123: return _reg("ttmp", v - 108, cnt) + return decode_src(v) + +def _fmt_src64(v: int) -> str: + """Format 64-bit source (VGPR pair, SGPR pair, or special pair).""" + return _fmt_src_n(v, 2) + def _parse_sop_sizes(op_name: str) -> tuple[int, ...]: """Parse dst and src sizes from SOP instruction name. Returns (dst_cnt, src0_cnt) or (dst_cnt, src0_cnt, src1_cnt).""" if op_name in ('s_bitset0_b64', 's_bitset1_b64'): return (2, 1) @@ -83,21 +105,15 @@ def disasm(inst: Inst) -> str: if op_name == 'v_nop': return 'v_nop' if op_name == 'v_pipeflush': return 'v_pipeflush' parts = op_name.split('_') - is_16bit_dst = ('cvt' not in op_name) and (any(p in ('f16', 'i16', 'u16', 'b16') for p in parts[-2:-1]) or (len(parts) >= 2 and parts[-1] in ('f16', 'i16', 'u16', 'b16'))) - is_16bit_src = parts[-1] in ('f16', 'i16', 'u16', 'b16') and 'sat_pk' not in op_name - is_f64_dst = op_name in ('v_ceil_f64', 'v_floor_f64', 'v_fract_f64', 'v_frexp_mant_f64', 'v_rcp_f64', 'v_rndne_f64', 'v_rsq_f64', 'v_sqrt_f64', 'v_trunc_f64', 'v_cvt_f64_f32', 'v_cvt_f64_i32', 'v_cvt_f64_u32') - is_f64_src = op_name in ('v_ceil_f64', 'v_floor_f64', 'v_fract_f64', 'v_frexp_mant_f64', 'v_rcp_f64', 'v_rndne_f64', 'v_rsq_f64', 'v_sqrt_f64', 'v_trunc_f64', 'v_cvt_f32_f64', 'v_cvt_i32_f64', 'v_cvt_u32_f64', 'v_frexp_exp_i32_f64') + is_16bit_dst = any(p in _16BIT_TYPES for p in parts[-2:-1]) or (len(parts) >= 2 and parts[-1] in _16BIT_TYPES and 'cvt' not in op_name) + is_16bit_src = parts[-1] in _16BIT_TYPES and 'sat_pk' not in op_name and 'cvt' not in op_name + _F64_OPS = ('v_ceil_f64', 'v_floor_f64', 'v_fract_f64', 'v_frexp_mant_f64', 'v_rcp_f64', 'v_rndne_f64', 'v_rsq_f64', 'v_sqrt_f64', 'v_trunc_f64') + is_f64_dst = op_name in _F64_OPS or op_name in ('v_cvt_f64_f32', 'v_cvt_f64_i32', 'v_cvt_f64_u32') + is_f64_src = op_name in _F64_OPS or op_name in ('v_cvt_f32_f64', 'v_cvt_i32_f64', 'v_cvt_u32_f64', 'v_frexp_exp_i32_f64') if op_name == 'v_readfirstlane_b32': return f"v_readfirstlane_b32 {decode_src(vdst)}, v{src0 - 256 if src0 >= 256 else src0}" dst_str = _vreg(vdst, 2) if is_f64_dst else f"v{vdst & 0x7f}.{'h' if vdst >= 128 else 'l'}" if is_16bit_dst else f"v{vdst}" - if is_f64_src: - src_str = _vreg(src0 - 256, 2) if src0 >= 256 else _sreg(src0, 2) if src0 <= 105 else "vcc" if src0 == 106 else "exec" if src0 == 126 else f"ttmp[{src0-108}:{src0-108+1}]" if 108 <= src0 <= 123 else fmt_src(src0) - elif is_16bit_src and src0 >= 256 and 'cvt' not in op_name: - # Add .l/.h suffix for 16-bit ops, but NOT for conversion instructions - # v_cvt_f32_f16 takes a 32-bit register (reads low 16 bits implicitly) - src_str = f"v{(src0 - 256) & 0x7f}.{'h' if src0 >= 384 else 'l'}" - else: - src_str = fmt_src(src0) + src_str = _fmt_src64(src0) if is_f64_src else f"v{(src0 - 256) & 0x7f}.{'h' if src0 >= 384 else 'l'}" if is_16bit_src and src0 >= 256 else fmt_src(src0) return f"{op_name}_e32 {dst_str}, {src_str}" # VOP2 @@ -120,16 +136,9 @@ def disasm(inst: Inst) -> str: is_64bit_vsrc1 = is_64bit and 'class' not in op_name is_16bit = any(x in op_name for x in ('_f16', '_i16', '_u16')) and 'f32' not in op_name is_cmpx = op_name.startswith('v_cmpx') # VOPCX writes to exec, no vcc destination - if is_64bit: - src0_str = _vreg(src0 - 256, 2) if src0 >= 256 else _sreg(src0, 2) if src0 <= 105 else "vcc" if src0 == 106 else "exec" if src0 == 126 else f"ttmp[{src0-108}:{src0-108+1}]" if 108 <= src0 <= 123 else fmt_src(src0) - elif is_16bit and src0 >= 256: - src0_str = f"v{(src0 - 256) & 0x7f}.{'h' if src0 >= 384 else 'l'}" - else: - src0_str = fmt_src(src0) + src0_str = _fmt_src64(src0) if is_64bit else f"v{(src0 - 256) & 0x7f}.{'h' if src0 >= 384 else 'l'}" if is_16bit and src0 >= 256 else fmt_src(src0) vsrc1_str = _vreg(vsrc1, 2) if is_64bit_vsrc1 else f"v{vsrc1 & 0x7f}.{'h' if vsrc1 >= 128 else 'l'}" if is_16bit else f"v{vsrc1}" - if is_cmpx: - return f"{op_name}_e32 {src0_str}, {vsrc1_str}" - return f"{op_name}_e32 vcc_lo, {src0_str}, {vsrc1_str}" + return f"{op_name}_e32 {src0_str}, {vsrc1_str}" if is_cmpx else f"{op_name}_e32 vcc_lo, {src0_str}, {vsrc1_str}" # SOPP if cls_name == 'SOPP': @@ -161,47 +170,17 @@ def disasm(inst: Inst) -> str: # SMEM if cls_name == 'SMEM': - # No-operand instructions if op_name in ('s_gl1_inv', 's_dcache_inv'): return op_name sdata, sbase, soffset, offset = unwrap(inst._values['sdata']), unwrap(inst._values['sbase']), unwrap(inst._values['soffset']), unwrap(inst._values.get('offset', 0)) glc, dlc = unwrap(inst._values.get('glc', 0)), unwrap(inst._values.get('dlc', 0)) - # s_atc_probe/s_atc_probe_buffer: sdata is the probe mode (0-7), not a register - if op_name in ('s_atc_probe', 's_atc_probe_buffer'): - sbase_idx = sbase * 2 - sbase_cnt = 4 if op_name == 's_atc_probe_buffer' else 2 - sbase_str = _sreg(sbase_idx, sbase_cnt) - if offset and soffset != 124: - off_str = f"{decode_src(soffset)} offset:0x{offset:x}" - elif offset: - off_str = f"0x{offset:x}" - else: - off_str = decode_src(soffset) - return f"{op_name} {sdata}, {sbase_str}, {off_str}" + # Format offset: "soffset offset:X" if both, "0x{offset:x}" if only imm, or decode_src(soffset) + off_str = f"{decode_src(soffset)} offset:0x{offset:x}" if offset and soffset != 124 else f"0x{offset:x}" if offset else decode_src(soffset) + sbase_idx, sbase_cnt = sbase * 2, 4 if (8 <= op_val <= 12 or op_name == 's_atc_probe_buffer') else 2 + sbase_str = _fmt_ssrc(sbase_idx, sbase_cnt) if sbase_cnt == 2 else _sreg(sbase_idx, sbase_cnt) if sbase_idx <= 105 else _reg("ttmp", sbase_idx - 108, sbase_cnt) + if op_name in ('s_atc_probe', 's_atc_probe_buffer'): return f"{op_name} {sdata}, {sbase_str}, {off_str}" width = {0:1, 1:2, 2:4, 3:8, 4:16, 8:1, 9:2, 10:4, 11:8, 12:16}.get(op_val, 1) - # Offset handling: if offset is set, we need "soffset offset:X" format, otherwise just soffset or imm - if offset and soffset != 124: # both soffset register and offset immediate - off_str = f"{decode_src(soffset)} offset:0x{offset:x}" - elif offset: # only offset immediate (soffset=null) - off_str = f"0x{offset:x}" - elif soffset == 124: # null - off_str = "null" - else: # only soffset register - off_str = decode_src(soffset) - # sbase is stored as register pair index, multiply by 2 for actual register number - # s_buffer_load_* (op 8-12) use 4-reg sbase (buffer descriptor), s_load_* (op 0-4) use 2-reg sbase - sbase_idx = sbase * 2 - sbase_cnt = 4 if 8 <= op_val <= 12 else 2 - # Format sbase with special register names - if sbase_idx == 106 and sbase_cnt == 2: sbase_str = "vcc" - elif sbase_idx == 126 and sbase_cnt == 2: sbase_str = "exec" - elif 108 <= sbase_idx <= 123: sbase_str = f"ttmp[{sbase_idx-108}:{sbase_idx-108+sbase_cnt-1}]" - else: sbase_str = _sreg(sbase_idx, sbase_cnt) - # Build modifiers - mods = [] - if glc: mods.append("glc") - if dlc: mods.append("dlc") - mod_str = " " + " ".join(mods) if mods else "" - return f"{op_name} {_fmt_sdst(sdata, width)}, {sbase_str}, {off_str}{mod_str}" + mods = [m for m in ["glc" if glc else "", "dlc" if dlc else ""] if m] + return f"{op_name} {_fmt_sdst(sdata, width)}, {sbase_str}, {off_str}" + (" " + " ".join(mods) if mods else "") # DS (LDS/GDS) if cls_name == 'DS': @@ -225,19 +204,13 @@ def disasm(inst: Inst) -> str: # FLAT if cls_name == 'FLAT': vdst, addr, data, saddr, offset, seg = [unwrap(inst._values.get(f, 0)) for f in ['vdst', 'addr', 'data', 'saddr', 'offset', 'seg']] - prefix = {0: 'flat', 1: 'scratch', 2: 'global'}.get(seg, 'flat') - op_suffix = op_name.split('_', 1)[1] if '_' in op_name else op_name - instr = f"{prefix}_{op_suffix}" - is_store = 'store' in op_name + instr = f"{['flat', 'scratch', 'global'][seg] if seg < 3 else 'flat'}_{op_name.split('_', 1)[1] if '_' in op_name else op_name}" width = {'b32':1, 'b64':2, 'b96':3, 'b128':4, 'u8':1, 'i8':1, 'u16':1, 'i16':1}.get(op_name.split('_')[-1], 1) - if saddr == 0x7F: - addr_str, saddr_str = _vreg(addr, 2), "" - else: - addr_str = _vreg(addr) - saddr_str = f", {_sreg(saddr, 2)}" if saddr < 106 else f", off" if saddr == 124 else f", {decode_src(saddr)}" + addr_str = _vreg(addr, 2) if saddr == 0x7F else _vreg(addr) + saddr_str = "" if saddr == 0x7F else f", {_sreg(saddr, 2)}" if saddr < 106 else ", off" if saddr == 124 else f", {decode_src(saddr)}" off_str = f" offset:{offset}" if offset else "" - if is_store: return f"{instr} {addr_str}, {_vreg(data, width)}{saddr_str}{off_str}" - return f"{instr} {_vreg(vdst, width)}, {addr_str}{saddr_str}{off_str}" + vdata_str = _vreg(data if 'store' in op_name else vdst, width) + return f"{instr} {addr_str}, {vdata_str}{saddr_str}{off_str}" if 'store' in op_name else f"{instr} {vdata_str}, {addr_str}{saddr_str}{off_str}" # VOP3: vector ops with modifiers (can be 1, 2, or 3 sources depending on opcode range) if cls_name == 'VOP3': @@ -258,18 +231,10 @@ def disasm(inst: Inst) -> str: # v_mad_i64_i32/v_mad_u64_u32: 64-bit dst and src2, 32-bit src0/src1 is_mad64 = 'mad_i64_i32' in op_name or 'mad_u64_u32' in op_name def fmt_sd_src(v, neg_bit, is_64bit=False): - s = fmt_src(v) - if is_64bit or is_f64: - if v >= 256: s = _vreg(v - 256, 2) - elif v <= 105: s = _sreg(v, 2) - elif v == 106: s = "vcc" - elif v == 126: s = "exec" - elif 108 <= v <= 123: s = f"ttmp[{v-108}:{v-108+1}]" - if neg_bit: s = f"-{s}" - return s - src0_str = fmt_sd_src(src0, neg & 1, False) # 32-bit for mad64 - src1_str = fmt_sd_src(src1, neg & 2, False) # 32-bit for mad64 - src2_str = fmt_sd_src(src2, neg & 4, is_mad64) # 64-bit for mad64 + s = _fmt_src64(v) if (is_64bit or is_f64) else fmt_src(v) + return f"-{s}" if neg_bit else s + src0_str, src1_str = fmt_sd_src(src0, neg & 1), fmt_sd_src(src1, neg & 2) + src2_str = fmt_sd_src(src2, neg & 4, is_mad64) dst_str = _vreg(vdst, 2) if (is_f64 or is_mad64) else f"v{vdst}" sdst_str = _fmt_sdst(sdst, 1) # v_add_co_u32, v_sub_co_u32, v_subrev_co_u32, v_add_co_ci_u32, etc. only use 2 sources @@ -299,56 +264,24 @@ def disasm(inst: Inst) -> str: # v_mqsad_u32_u8: 128-bit (4 reg) dst/src2, 64-bit src0, 32-bit src1 is_sad64 = any(x in op_name for x in ('qsad_pk', 'mqsad_pk')) is_mqsad_u32 = 'mqsad_u32' in op_name - # Detect conversion ops: v_cvt_{dst_type}_{src_type} - each side may have different size - # Also handle v_cvt_pk_* which packs two values into one + # Detect 16-bit and 64-bit operand sizes for various instruction patterns if 'cvt_pk' in op_name: - # Pack ops: dst is packed 16-bit, src is determined by last type in name - # e.g., v_cvt_pk_i16_f32, v_cvt_pk_norm_i16_f32 - is_f16_dst = is_f16_src = is_f16_src2 = False # dst is 32-bit, srcs depend on op - is_f16_src = op_name.endswith('16') # only if final type is 16-bit - elif m := re.match(r'v_cvt_([a-z0-9_]+)_([a-z0-9]+)', op_name): + is_f16_dst, is_f16_src, is_f16_src2 = False, op_name.endswith('16'), False + elif m := re.match(r'v_(?:cvt|frexp_exp)_([a-z0-9_]+)_([a-z0-9]+)', op_name): dst_type, src_type = m.group(1), m.group(2) - # Check if dst/src ends with a 16-bit type suffix - is_f16_dst = any(dst_type.endswith(x) for x in ('f16', 'i16', 'u16', 'b16')) - is_f16_src = is_f16_src2 = any(src_type.endswith(x) for x in ('f16', 'i16', 'u16', 'b16')) - # Override is_f64 for conversion ops - check if dst or src is 64-bit - is_f64_dst = '64' in dst_type - is_f64_src = '64' in src_type - is_f64 = False # Don't use default is_f64 detection for cvt ops - elif m := re.match(r'v_frexp_exp_([a-z0-9]+)_([a-z0-9]+)', op_name): - # v_frexp_exp_i32_f64: 32-bit dst (exponent), 64-bit src - # v_frexp_exp_i16_f16: 16-bit dst, 16-bit src - dst_type, src_type = m.group(1), m.group(2) - is_f16_dst = any(dst_type.endswith(x) for x in ('f16', 'i16', 'u16', 'b16')) - is_f16_src = is_f16_src2 = any(src_type.endswith(x) for x in ('f16', 'i16', 'u16', 'b16')) - is_f64_dst = '64' in dst_type - is_f64_src = '64' in src_type - is_f64 = False - elif m := re.match(r'v_mad_([iu])32_([iu])16', op_name): - # v_mad_i32_i16, v_mad_u32_u16: 32-bit dst, 16-bit src0/src1, 32-bit src2 - is_f16_dst = False - is_f16_src = True # src0 and src1 are 16-bit - is_f16_src2 = False # src2 is 32-bit + is_f16_dst, is_f16_src, is_f16_src2 = _is_16bit(dst_type), _is_16bit(src_type), _is_16bit(src_type) + is_f64_dst, is_f64_src, is_f64 = '64' in dst_type, '64' in src_type, False + elif re.match(r'v_mad_[iu]32_[iu]16', op_name): + is_f16_dst, is_f16_src, is_f16_src2 = False, True, False # 32-bit dst, 16-bit src0/src1, 32-bit src2 elif 'pack_b32' in op_name: - # v_pack_b32_f16: 32-bit dst, 16-bit sources - is_f16_dst = False - is_f16_src = is_f16_src2 = True + is_f16_dst, is_f16_src, is_f16_src2 = False, True, True # 32-bit dst, 16-bit sources else: - # 16-bit ops need .h/.l suffix, but packed ops (dot2, pk_, sad, msad, qsad, mqsad) don't - is_16bit_op = ('f16' in op_name or 'i16' in op_name or 'u16' in op_name or 'b16' in op_name) and not any(x in op_name for x in ('dot2', 'pk_', 'sad', 'msad', 'qsad', 'mqsad')) + is_16bit_op = any(x in op_name for x in _16BIT_TYPES) and not any(x in op_name for x in ('dot2', 'pk_', 'sad', 'msad', 'qsad', 'mqsad')) is_f16_dst = is_f16_src = is_f16_src2 = is_16bit_op def fmt_vop3_src(v, neg_bit, abs_bit, hi_bit=False, reg_cnt=1, is_16=False): - s = fmt_src(v) - # Add register pair/quad for 64/128-bit, or .h suffix for f16 VGPRs with opsel - if reg_cnt > 1 and v >= 256: s = _vreg(v - 256, reg_cnt) - elif reg_cnt > 1 and v <= 105: s = _sreg(v, reg_cnt) - elif reg_cnt == 2 and v == 106: s = "vcc" - elif reg_cnt == 2 and v == 126: s = "exec" - elif reg_cnt > 1 and 108 <= v <= 123: s = f"ttmp[{v-108}:{v-108+reg_cnt-1}]" - elif is_16 and v >= 256: s = f"v{v - 256}.h" if hi_bit else f"v{v - 256}.l" + s = _fmt_src_n(v, reg_cnt) if reg_cnt > 1 else f"v{v - 256}.h" if is_16 and v >= 256 and hi_bit else f"v{v - 256}.l" if is_16 and v >= 256 else fmt_src(v) if abs_bit: s = f"|{s}|" - if neg_bit: s = f"-{s}" - return s + return f"-{s}" if neg_bit else s # Determine register count for each source (check for cvt-specific 64-bit flags first) is_src0_64 = locals().get('is_f64_src', is_f64 and not is_shift64) or is_sad64 or is_mqsad_u32 is_src1_64 = is_f64 and not is_class and not is_ldexp64 and not is_trig_preop @@ -419,228 +352,94 @@ def disasm(inst: Inst) -> str: if cls_name == 'VOP3SD': vdst, sdst = unwrap(inst._values.get('vdst', 0)), unwrap(inst._values.get('sdst', 0)) src0, src1, src2 = [unwrap(inst._values.get(f, 0)) for f in ('src0', 'src1', 'src2')] - neg = unwrap(inst._values.get('neg', 0)) - omod = unwrap(inst._values.get('omod', 0)) - clmp = unwrap(inst._values.get('clmp', 0)) - is_f64 = 'f64' in op_name - is_mad64 = 'mad_i64_i32' in op_name or 'mad_u64_u32' in op_name - def fmt_sd_src(v, neg_bit, is_64bit=False): - s = fmt_src(v) - if is_64bit or is_f64: - if v >= 256: s = _vreg(v - 256, 2) - elif v <= 105: s = _sreg(v, 2) - elif v == 106: s = "vcc" - elif v == 126: s = "exec" - elif 108 <= v <= 123: s = f"ttmp[{v-108}:{v-108+1}]" - if neg_bit: s = f"-{s}" - return s - src0_str = fmt_sd_src(src0, neg & 1, False) - src1_str = fmt_sd_src(src1, neg & 2, False) - src2_str = fmt_sd_src(src2, neg & 4, is_mad64) - dst_str = _vreg(vdst, 2) if (is_f64 or is_mad64) else f"v{vdst}" - sdst_str = _fmt_sdst(sdst, 1) - clamp_str = " clamp" if clmp else "" - omod_str = {1: " mul:2", 2: " mul:4", 3: " div:2"}.get(omod, "") - # v_add_co_u32, v_sub_co_u32, v_subrev_co_u32 only use 2 sources - if op_name in ('v_add_co_u32', 'v_sub_co_u32', 'v_subrev_co_u32'): - return f"{op_name}_e64 {dst_str}, {sdst_str}, {src0_str}, {src1_str}" + clamp_str - # v_add_co_ci_u32, v_sub_co_ci_u32, v_subrev_co_ci_u32 use 3 sources (src2 is carry-in) - if op_name in ('v_add_co_ci_u32', 'v_sub_co_ci_u32', 'v_subrev_co_ci_u32'): - return f"{op_name}_e64 {dst_str}, {sdst_str}, {src0_str}, {src1_str}, {src2_str}" + clamp_str - # v_div_scale, v_mad_*64_*32 use 3 sources - return f"{op_name} {dst_str}, {sdst_str}, {src0_str}, {src1_str}, {src2_str}" + clamp_str + omod_str + neg, omod, clmp = unwrap(inst._values.get('neg', 0)), unwrap(inst._values.get('omod', 0)), unwrap(inst._values.get('clmp', 0)) + is_f64, is_mad64 = 'f64' in op_name, 'mad_i64_i32' in op_name or 'mad_u64_u32' in op_name + def fmt_neg(v, neg_bit, is_64=False): return f"-{_fmt_src64(v) if (is_64 or is_f64) else fmt_src(v)}" if neg_bit else _fmt_src64(v) if (is_64 or is_f64) else fmt_src(v) + srcs = [fmt_neg(src0, neg & 1), fmt_neg(src1, neg & 2), fmt_neg(src2, neg & 4, is_mad64)] + dst_str, sdst_str = _vreg(vdst, 2) if (is_f64 or is_mad64) else f"v{vdst}", _fmt_sdst(sdst, 1) + clamp_str, omod_str = " clamp" if clmp else "", {1: " mul:2", 2: " mul:4", 3: " div:2"}.get(omod, "") + is_2src = op_name in ('v_add_co_u32', 'v_sub_co_u32', 'v_subrev_co_u32') + suffix = "_e64" if op_name.startswith('v_') and 'co_' in op_name else "" + return f"{op_name}{suffix} {dst_str}, {sdst_str}, {', '.join(srcs[:2] if is_2src else srcs)}" + clamp_str + omod_str # VOPD: dual-issue instructions if cls_name == 'VOPD': from extra.assembly.rdna3 import autogen - opx, opy = unwrap(inst._values.get('opx', 0)), unwrap(inst._values.get('opy', 0)) - vdstx, vdsty_enc = unwrap(inst._values.get('vdstx', 0)), unwrap(inst._values.get('vdsty', 0)) - srcx0, vsrcx1 = unwrap(inst._values.get('srcx0', 0)), unwrap(inst._values.get('vsrcx1', 0)) - srcy0, vsrcy1 = unwrap(inst._values.get('srcy0', 0)), unwrap(inst._values.get('vsrcy1', 0)) - # Decode vdsty: actual = (encoded << 1) | ((vdstx & 1) ^ 1) - vdsty = (vdsty_enc << 1) | ((vdstx & 1) ^ 1) - try: - opx_name = autogen.VOPDOp(opx).name.lower() - opy_name = autogen.VOPDOp(opy).name.lower() - except (ValueError, KeyError): - opx_name, opy_name = f"opx_{opx}", f"opy_{opy}" - # v_dual_mov_b32 only has 1 source - opx_str = f"{opx_name} v{vdstx}, {fmt_src(srcx0)}" if 'mov' in opx_name else f"{opx_name} v{vdstx}, {fmt_src(srcx0)}, v{vsrcx1}" - opy_str = f"{opy_name} v{vdsty}, {fmt_src(srcy0)}" if 'mov' in opy_name else f"{opy_name} v{vdsty}, {fmt_src(srcy0)}, v{vsrcy1}" - return f"{opx_str} :: {opy_str}" + opx, opy, vdstx, vdsty_enc = [unwrap(inst._values.get(f, 0)) for f in ('opx', 'opy', 'vdstx', 'vdsty')] + srcx0, vsrcx1, srcy0, vsrcy1 = [unwrap(inst._values.get(f, 0)) for f in ('srcx0', 'vsrcx1', 'srcy0', 'vsrcy1')] + vdsty = (vdsty_enc << 1) | ((vdstx & 1) ^ 1) # Decode vdsty + def fmt_vopd(op, vdst, src0, vsrc1): + try: name = autogen.VOPDOp(op).name.lower() + except (ValueError, KeyError): name = f"op_{op}" + return f"{name} v{vdst}, {fmt_src(src0)}" if 'mov' in name else f"{name} v{vdst}, {fmt_src(src0)}, v{vsrc1}" + return f"{fmt_vopd(opx, vdstx, srcx0, vsrcx1)} :: {fmt_vopd(opy, vdsty, srcy0, vsrcy1)}" # VOP3P: packed vector ops if cls_name == 'VOP3P': - vdst = unwrap(inst._values.get('vdst', 0)) + vdst, clmp = unwrap(inst._values.get('vdst', 0)), unwrap(inst._values.get('clmp', 0)) src0, src1, src2 = [unwrap(inst._values.get(f, 0)) for f in ('src0', 'src1', 'src2')] - neg = unwrap(inst._values.get('neg', 0)) # neg_lo - neg_hi = unwrap(inst._values.get('neg_hi', 0)) - opsel = unwrap(inst._values.get('opsel', 0)) - opsel_hi = unwrap(inst._values.get('opsel_hi', 0)) - opsel_hi2 = unwrap(inst._values.get('opsel_hi2', 0)) - clmp = unwrap(inst._values.get('clmp', 0)) - # WMMA ops have special register widths - is_wmma = 'wmma' in op_name - # Determine number of sources (dot ops are 3-src, most are 2-src) - is_3src = any(x in op_name for x in ('fma', 'mad', 'dot', 'wmma')) - # Format source operands - def fmt_vop3p_src(v, reg_cnt=1): - if v >= 256: return _vreg(v - 256, reg_cnt) - if v <= 105: return _sreg(v, reg_cnt) if reg_cnt > 1 else f"s{v}" - if v == 106 and reg_cnt == 2: return "vcc" - if v == 126 and reg_cnt == 2: return "exec" - return fmt_src(v) + neg, neg_hi = unwrap(inst._values.get('neg', 0)), unwrap(inst._values.get('neg_hi', 0)) + opsel, opsel_hi, opsel_hi2 = unwrap(inst._values.get('opsel', 0)), unwrap(inst._values.get('opsel_hi', 0)), unwrap(inst._values.get('opsel_hi2', 0)) + is_wmma, is_3src = 'wmma' in op_name, any(x in op_name for x in ('fma', 'mad', 'dot', 'wmma')) + def fmt_bits(name, val, n): return f"{name}:[{','.join(str((val >> i) & 1) for i in range(n))}]" # WMMA: f16/bf16 use 8-reg sources, iu8 uses 4-reg, iu4 uses 2-reg; all have 8-reg dst if is_wmma: src_cnt = 2 if 'iu4' in op_name else 4 if 'iu8' in op_name else 8 - src0_str = _vreg(src0 - 256, src_cnt) if src0 >= 256 else fmt_vop3p_src(src0, src_cnt) - src1_str = _vreg(src1 - 256, src_cnt) if src1 >= 256 else fmt_vop3p_src(src1, src_cnt) - src2_str = _vreg(src2 - 256, 8) if src2 >= 256 else fmt_vop3p_src(src2, 8) + src0_str, src1_str, src2_str = _fmt_src_n(src0, src_cnt), _fmt_src_n(src1, src_cnt), _fmt_src_n(src2, 8) dst_str = _vreg(vdst, 8) else: - src0_str = fmt_vop3p_src(src0) - src1_str = fmt_vop3p_src(src1) - src2_str = fmt_vop3p_src(src2) + src0_str, src1_str, src2_str = _fmt_src_n(src0, 1), _fmt_src_n(src1, 1), _fmt_src_n(src2, 1) dst_str = f"v{vdst}" - # Build modifiers - VOP3P uses op_sel, op_sel_hi, neg_lo, neg_hi - mods = [] - # op_sel: selects high/low half of each source - if opsel: - if is_3src: - mods.append(f"op_sel:[{opsel & 1},{(opsel >> 1) & 1},{(opsel >> 2) & 1}]") - else: - mods.append(f"op_sel:[{opsel & 1},{(opsel >> 1) & 1}]") - # op_sel_hi: selects high half for upper result lane (default [1,1] or [1,1,1]) - # opsel_hi is bits 0-1, opsel_hi2 is bit 2 (for src2) + n = 3 if is_3src else 2 full_opsel_hi = opsel_hi | (opsel_hi2 << 2) - default_opsel_hi = 0b111 if is_3src else 0b11 - if full_opsel_hi != default_opsel_hi: - if is_3src: - mods.append(f"op_sel_hi:[{full_opsel_hi & 1},{(full_opsel_hi >> 1) & 1},{(full_opsel_hi >> 2) & 1}]") - else: - mods.append(f"op_sel_hi:[{full_opsel_hi & 1},{(full_opsel_hi >> 1) & 1}]") - # neg_lo: negate lower half of source - if neg: - if is_3src: - mods.append(f"neg_lo:[{neg & 1},{(neg >> 1) & 1},{(neg >> 2) & 1}]") - else: - mods.append(f"neg_lo:[{neg & 1},{(neg >> 1) & 1}]") - # neg_hi: negate upper half of source - if neg_hi: - if is_3src: - mods.append(f"neg_hi:[{neg_hi & 1},{(neg_hi >> 1) & 1},{(neg_hi >> 2) & 1}]") - else: - mods.append(f"neg_hi:[{neg_hi & 1},{(neg_hi >> 1) & 1}]") + mods = [fmt_bits("op_sel", opsel, n)] if opsel else [] + if full_opsel_hi != (0b111 if is_3src else 0b11): mods.append(fmt_bits("op_sel_hi", full_opsel_hi, n)) + if neg: mods.append(fmt_bits("neg_lo", neg, n)) + if neg_hi: mods.append(fmt_bits("neg_hi", neg_hi, n)) if clmp: mods.append("clamp") mod_str = " " + " ".join(mods) if mods else "" - if is_3src: - return f"{op_name} {dst_str}, {src0_str}, {src1_str}, {src2_str}{mod_str}" - return f"{op_name} {dst_str}, {src0_str}, {src1_str}{mod_str}" + return f"{op_name} {dst_str}, {src0_str}, {src1_str}, {src2_str}{mod_str}" if is_3src else f"{op_name} {dst_str}, {src0_str}, {src1_str}{mod_str}" # VINTERP: interpolation instructions if cls_name == 'VINTERP': vdst = unwrap(inst._values.get('vdst', 0)) src0, src1, src2 = [unwrap(inst._values.get(f, 0)) for f in ('src0', 'src1', 'src2')] - waitexp = unwrap(inst._values.get('waitexp', 0)) - neg = unwrap(inst._values.get('neg', 0)) - clmp = unwrap(inst._values.get('clmp', 0)) - opsel = unwrap(inst._values.get('opsel', 0)) - def fmt_vi_src(v, neg_bit): - s = f"v{v - 256}" if v >= 256 else fmt_src(v) - if neg_bit: s = f"-{s}" - return s - src0_str = fmt_vi_src(src0, neg & 1) - src1_str = fmt_vi_src(src1, neg & 2) - src2_str = fmt_vi_src(src2, neg & 4) - # LLVM doesn't use .l/.h suffix for vinterp dst - dst_str = f"v{vdst}" - mods = [] - if waitexp: mods.append(f"wait_exp:{waitexp}") - if clmp: mods.append("clamp") - mod_str = " " + " ".join(mods) if mods else "" - return f"{op_name} {dst_str}, {src0_str}, {src1_str}, {src2_str}{mod_str}" + neg, waitexp, clmp = unwrap(inst._values.get('neg', 0)), unwrap(inst._values.get('waitexp', 0)), unwrap(inst._values.get('clmp', 0)) + def fmt_neg_vi(v, neg_bit): return f"-{v}" if neg_bit else v + srcs = [fmt_neg_vi(f"v{s - 256}" if s >= 256 else fmt_src(s), neg & (1 << i)) for i, s in enumerate([src0, src1, src2])] + mods = [m for m in [f"wait_exp:{waitexp}" if waitexp else "", "clamp" if clmp else ""] if m] + return f"{op_name} v{vdst}, {', '.join(srcs)}" + (" " + " ".join(mods) if mods else "") + + # MUBUF/MTBUF helpers + def _buf_vaddr(vaddr, offen, idxen): return _vreg(vaddr, 2) if offen and idxen else f"v{vaddr}" if offen or idxen else "off" + def _buf_srsrc(srsrc): srsrc_base = srsrc * 4; return _reg("ttmp", srsrc_base - 108, 4) if 108 <= srsrc_base <= 123 else _sreg(srsrc_base, 4) # MUBUF: buffer load/store if cls_name == 'MUBUF': - vdata, vaddr = unwrap(inst._values.get('vdata', 0)), unwrap(inst._values.get('vaddr', 0)) - srsrc, soffset = unwrap(inst._values.get('srsrc', 0)), unwrap(inst._values.get('soffset', 0)) - offset = unwrap(inst._values.get('offset', 0)) - offen, idxen = unwrap(inst._values.get('offen', 0)), unwrap(inst._values.get('idxen', 0)) - glc, dlc, slc = unwrap(inst._values.get('glc', 0)), unwrap(inst._values.get('dlc', 0)), unwrap(inst._values.get('slc', 0)) - tfe = unwrap(inst._values.get('tfe', 0)) - # Special ops with no operands + vdata, vaddr, srsrc, soffset = [unwrap(inst._values.get(f, 0)) for f in ('vdata', 'vaddr', 'srsrc', 'soffset')] + offset, offen, idxen = unwrap(inst._values.get('offset', 0)), unwrap(inst._values.get('offen', 0)), unwrap(inst._values.get('idxen', 0)) + glc, dlc, slc, tfe = [unwrap(inst._values.get(f, 0)) for f in ('glc', 'dlc', 'slc', 'tfe')] if op_name in ('buffer_gl0_inv', 'buffer_gl1_inv'): return op_name # Determine data width from op name - # d16 formats: _x and _xy use 1 reg, _xyz and _xyzw use 2 regs - # regular formats: _x=1, _xy=2, _xyz=3, _xyzw=4 - # atomic u64 uses 2 regs, cmpswap doubles width (compare + swap) - if 'd16' in op_name: - width = 2 if any(x in op_name for x in ('xyz', 'xyzw')) else 1 + if 'd16' in op_name: width = 2 if any(x in op_name for x in ('xyz', 'xyzw')) else 1 elif 'atomic' in op_name: - # cmpswap uses 2 regs for b32, 4 for b64; other atomics use 1 for b32, 2 for b64/u64/i64 base_width = 2 if any(x in op_name for x in ('b64', 'u64', 'i64')) else 1 width = base_width * 2 if 'cmpswap' in op_name else base_width - else: - width = {'b32':1, 'b64':2, 'b96':3, 'b128':4, 'b16':1, 'x':1, 'xy':2, 'xyz':3, 'xyzw':4}.get(op_name.split('_')[-1], 1) - # tfe adds 1 extra VGPR for texture fault status + else: width = {'b32':1, 'b64':2, 'b96':3, 'b128':4, 'b16':1, 'x':1, 'xy':2, 'xyz':3, 'xyzw':4}.get(op_name.split('_')[-1], 1) if tfe: width += 1 - is_store = 'store' in op_name - # Format vaddr - if offen and idxen: vaddr_str = f"v[{vaddr}:{vaddr+1}]" - elif offen or idxen: vaddr_str = f"v{vaddr}" - else: vaddr_str = "off" - # Format srsrc (4-aligned SGPR quad) - srsrc_base = srsrc * 4 - srsrc_str = f"s[{srsrc_base}:{srsrc_base+3}]" - # Format soffset - use decode_src for proper constant handling - soff_str = decode_src(soffset) - # Build modifiers - mods = [] - if offen: mods.append("offen") - if idxen: mods.append("idxen") - if offset: mods.append(f"offset:{offset}") - if glc: mods.append("glc") - if dlc: mods.append("dlc") - if slc: mods.append("slc") - if tfe: mods.append("tfe") - mod_str = " " + " ".join(mods) if mods else "" - if is_store: - return f"{op_name} {_vreg(vdata, width)}, {vaddr_str}, {srsrc_str}, {soff_str}{mod_str}" - return f"{op_name} {_vreg(vdata, width)}, {vaddr_str}, {srsrc_str}, {soff_str}{mod_str}" + mods = [m for m in ["offen" if offen else "", "idxen" if idxen else "", f"offset:{offset}" if offset else "", + "glc" if glc else "", "dlc" if dlc else "", "slc" if slc else "", "tfe" if tfe else ""] if m] + return f"{op_name} {_vreg(vdata, width)}, {_buf_vaddr(vaddr, offen, idxen)}, {_buf_srsrc(srsrc)}, {decode_src(soffset)}" + (" " + " ".join(mods) if mods else "") # MTBUF: typed buffer load/store if cls_name == 'MTBUF': - vdata, vaddr = unwrap(inst._values.get('vdata', 0)), unwrap(inst._values.get('vaddr', 0)) - srsrc, soffset = unwrap(inst._values.get('srsrc', 0)), unwrap(inst._values.get('soffset', 0)) - offset, fmt = unwrap(inst._values.get('offset', 0)), unwrap(inst._values.get('format', 0)) - offen, idxen = unwrap(inst._values.get('offen', 0)), unwrap(inst._values.get('idxen', 0)) - glc, dlc, slc = unwrap(inst._values.get('glc', 0)), unwrap(inst._values.get('dlc', 0)), unwrap(inst._values.get('slc', 0)) - # Format vaddr - if offen and idxen: vaddr_str = f"v[{vaddr}:{vaddr+1}]" - elif offen or idxen: vaddr_str = f"v{vaddr}" - else: vaddr_str = "off" - # Format srsrc (4-aligned SGPR quad, or ttmp) - srsrc_base = srsrc * 4 - if 108 <= srsrc_base <= 123: - srsrc_str = f"ttmp[{srsrc_base-108}:{srsrc_base-108+3}]" - else: - srsrc_str = f"s[{srsrc_base}:{srsrc_base+3}]" - # Format soffset - use decode_src for proper special register handling - soff_str = decode_src(soffset) - # Build modifiers - idxen must come before offen for LLVM - mods = [f"format:{fmt}"] - if idxen: mods.append("idxen") - if offen: mods.append("offen") - if offset: mods.append(f"offset:{offset}") - if glc: mods.append("glc") - if dlc: mods.append("dlc") - if slc: mods.append("slc") - # Determine vdata width: d16 xyz/xyzw use 2 regs, d16 x/xy use 1 reg - if 'd16' in op_name: - width = 2 if any(x in op_name for x in ('xyz', 'xyzw')) else 1 - else: - width = {'x':1, 'xy':2, 'xyz':3, 'xyzw':4}.get(op_name.split('_')[-1], 1) - return f"{op_name} {_vreg(vdata, width)}, {vaddr_str}, {srsrc_str}, {soff_str} {' '.join(mods)}" + vdata, vaddr, srsrc, soffset = [unwrap(inst._values.get(f, 0)) for f in ('vdata', 'vaddr', 'srsrc', 'soffset')] + offset, tbuf_fmt, offen, idxen = [unwrap(inst._values.get(f, 0)) for f in ('offset', 'format', 'offen', 'idxen')] + glc, dlc, slc = [unwrap(inst._values.get(f, 0)) for f in ('glc', 'dlc', 'slc')] + mods = [f"format:{tbuf_fmt}"] + [m for m in ["idxen" if idxen else "", "offen" if offen else "", f"offset:{offset}" if offset else "", + "glc" if glc else "", "dlc" if dlc else "", "slc" if slc else ""] if m] + width = 2 if 'd16' in op_name and any(x in op_name for x in ('xyz', 'xyzw')) else 1 if 'd16' in op_name else {'x':1, 'xy':2, 'xyz':3, 'xyzw':4}.get(op_name.split('_')[-1], 1) + return f"{op_name} {_vreg(vdata, width)}, {_buf_vaddr(vaddr, offen, idxen)}, {_buf_srsrc(srsrc)}, {decode_src(soffset)} {' '.join(mods)}" # SOP1/SOP2/SOPC/SOPK if cls_name in ('SOP1', 'SOP2', 'SOPC', 'SOPK'): @@ -648,15 +447,13 @@ def disasm(inst: Inst) -> str: dst_cnt, src0_cnt = sizes[0], sizes[1] src1_cnt = sizes[2] if len(sizes) > 2 else src0_cnt if cls_name == 'SOP1': - if op_name == 's_getpc_b64': return f"{op_name} {_fmt_sdst(unwrap(inst._values.get('sdst', 0)), 2)}" - if op_name in ('s_setpc_b64', 's_rfe_b64'): return f"{op_name} {_fmt_ssrc(unwrap(inst._values.get('ssrc0', 0)), 2)}" - if op_name == 's_swappc_b64': return f"{op_name} {_fmt_sdst(unwrap(inst._values.get('sdst', 0)), 2)}, {_fmt_ssrc(unwrap(inst._values.get('ssrc0', 0)), 2)}" + sdst, ssrc0 = unwrap(inst._values.get('sdst', 0)), unwrap(inst._values.get('ssrc0', 0)) + if op_name == 's_getpc_b64': return f"{op_name} {_fmt_sdst(sdst, 2)}" + if op_name in ('s_setpc_b64', 's_rfe_b64'): return f"{op_name} {_fmt_ssrc(ssrc0, 2)}" + if op_name == 's_swappc_b64': return f"{op_name} {_fmt_sdst(sdst, 2)}, {_fmt_ssrc(ssrc0, 2)}" if op_name in ('s_sendmsg_rtn_b32', 's_sendmsg_rtn_b64'): - msg_id = unwrap(inst._values.get('ssrc0', 0)) - msg_names = {128: 'MSG_RTN_GET_DOORBELL', 129: 'MSG_RTN_GET_DDID', 130: 'MSG_RTN_GET_TMA', 131: 'MSG_RTN_GET_REALTIME', 132: 'MSG_RTN_SAVE_WAVE', 133: 'MSG_RTN_GET_TBA'} - msg = msg_names.get(msg_id, str(msg_id)) - return f"{op_name} {_fmt_sdst(unwrap(inst._values.get('sdst', 0)), 2 if 'b64' in op_name else 1)}, sendmsg({msg})" - return f"{op_name} {_fmt_sdst(unwrap(inst._values.get('sdst', 0)), dst_cnt)}, {_fmt_ssrc(unwrap(inst._values.get('ssrc0', 0)), src0_cnt)}" + return f"{op_name} {_fmt_sdst(sdst, 2 if 'b64' in op_name else 1)}, sendmsg({MSG_NAMES.get(ssrc0, str(ssrc0))})" + return f"{op_name} {_fmt_sdst(sdst, dst_cnt)}, {_fmt_ssrc(ssrc0, src0_cnt)}" if cls_name == 'SOP2': sdst, ssrc0, ssrc1 = [unwrap(inst._values.get(f, 0)) for f in ('sdst', 'ssrc0', 'ssrc1')] return f"{op_name} {_fmt_sdst(sdst, dst_cnt)}, {_fmt_ssrc(ssrc0, src0_cnt)}, {_fmt_ssrc(ssrc1, src1_cnt)}" @@ -666,41 +463,24 @@ def disasm(inst: Inst) -> str: sdst, simm16 = unwrap(inst._values.get('sdst', 0)), unwrap(inst._values.get('simm16', 0)) if op_name == 's_version': return f"{op_name} 0x{simm16:x}" if op_name in ('s_setreg_b32', 's_getreg_b32'): - # Decode hwreg: (size-1) << 11 | offset << 6 | id hwreg_id, hwreg_offset, hwreg_size = simm16 & 0x3f, (simm16 >> 6) & 0x1f, ((simm16 >> 11) & 0x1f) + 1 - # GFX11+ hwreg names (IDs 16-17 are TBA which are not supported on GFX11, IDs 18-19 are PERF_SNAPSHOT) - hwreg_names = {1: 'HW_REG_MODE', 2: 'HW_REG_STATUS', 3: 'HW_REG_TRAPSTS', 4: 'HW_REG_HW_ID', - 5: 'HW_REG_GPR_ALLOC', 6: 'HW_REG_LDS_ALLOC', 7: 'HW_REG_IB_STS', - 15: 'HW_REG_SH_MEM_BASES', - 18: 'HW_REG_PERF_SNAPSHOT_PC_LO', 19: 'HW_REG_PERF_SNAPSHOT_PC_HI', - 20: 'HW_REG_FLAT_SCR_LO', 21: 'HW_REG_FLAT_SCR_HI', 22: 'HW_REG_XNACK_MASK', - 23: 'HW_REG_HW_ID1', 24: 'HW_REG_HW_ID2', 25: 'HW_REG_POPS_PACKER', - 28: 'HW_REG_IB_STS2'} - # For unsupported registers (TBA_LO/HI, TMA_LO/HI on GFX11), output raw simm16 value - if hwreg_id in (16, 17, 18, 19) and hwreg_id not in hwreg_names: - # Unsupported on GFX11 - use raw encoding - hwreg_str = f"0x{simm16:x}" - else: - hwreg_name = hwreg_names.get(hwreg_id, str(hwreg_id)) - hwreg_str = f"hwreg({hwreg_name}, {hwreg_offset}, {hwreg_size})" - if op_name == 's_setreg_b32': - return f"{op_name} {hwreg_str}, {_fmt_sdst(sdst, 1)}" - return f"{op_name} {_fmt_sdst(sdst, 1)}, {hwreg_str}" + hwreg_str = f"0x{simm16:x}" if hwreg_id in (16, 17) else f"hwreg({HWREG_NAMES.get(hwreg_id, str(hwreg_id))}, {hwreg_offset}, {hwreg_size})" + return f"{op_name} {hwreg_str}, {_fmt_sdst(sdst, 1)}" if op_name == 's_setreg_b32' else f"{op_name} {_fmt_sdst(sdst, 1)}, {hwreg_str}" return f"{op_name} {_fmt_sdst(sdst, dst_cnt)}, 0x{simm16:x}" # Generic fallback - def fmt(n, v): + def fmt_field(n, v): v = unwrap(v) if n in SRC_FIELDS: return fmt_src(v) if v != 255 else "0xff" if n in ('sdst', 'vdst'): return f"{'s' if n == 'sdst' else 'v'}{v}" return f"v{v}" if n == 'vsrc1' else f"0x{v:x}" if n == 'simm16' else str(v) - ops = [fmt(n, inst._values.get(n, 0)) for n in inst._fields if n not in ('encoding', 'op')] + ops = [fmt_field(n, inst._values.get(n, 0)) for n in inst._fields if n not in ('encoding', 'op')] return f"{op_name} {', '.join(ops)}" if ops else op_name # Assembler SPECIAL_REGS = {'vcc_lo': RawImm(106), 'vcc_hi': RawImm(107), 'null': RawImm(124), 'off': RawImm(124), 'm0': RawImm(125), 'exec_lo': RawImm(126), 'exec_hi': RawImm(127), 'scc': RawImm(253)} FLOAT_CONSTS = {'0.5': 0.5, '-0.5': -0.5, '1.0': 1.0, '-1.0': -1.0, '2.0': 2.0, '-2.0': -2.0, '4.0': 4.0, '-4.0': -4.0} -REG_MAP = {'s': SGPR, 'v': VGPR, 't': TTMP, 'ttmp': TTMP} +REG_MAP: dict[str, _RegFactory] = {'s': s, 'v': v, 't': ttmp, 'ttmp': ttmp} def parse_operand(op: str) -> tuple: op = op.strip().lower() @@ -717,23 +497,16 @@ def parse_operand(op: str) -> tuple: if op in SPECIAL_REGS: return (SPECIAL_REGS[op], neg, abs_, hi_half) if m := re.match(r'^([svt](?:tmp)?)\[(\d+):(\d+)\]$', op): return (REG_MAP[m.group(1)][int(m.group(2)):int(m.group(3))+1], neg, abs_, hi_half) if m := re.match(r'^([svt](?:tmp)?)(\d+)$', op): - return (REG_MAP[m.group(1)](int(m.group(2)), 1, hi_half), neg, abs_, hi_half) + reg = REG_MAP[m.group(1)][int(m.group(2))] + reg.hi = hi_half + return (reg, neg, abs_, hi_half) # hwreg(name, offset, size) or hwreg(name) -> simm16 encoding if m := re.match(r'^hwreg\((\w+)(?:,\s*(\d+),\s*(\d+))?\)$', op): - # GFX11 hwreg names - note IDs 18-19 are PERF_SNAPSHOT on GFX11, not TMA - hwreg_names = {'hw_reg_mode': 1, 'hw_reg_status': 2, 'hw_reg_trapsts': 3, 'hw_reg_hw_id': 4, - 'hw_reg_gpr_alloc': 5, 'hw_reg_lds_alloc': 6, 'hw_reg_ib_sts': 7, - 'hw_reg_sh_mem_bases': 15, - 'hw_reg_perf_snapshot_pc_lo': 18, 'hw_reg_perf_snapshot_pc_hi': 19, - 'hw_reg_flat_scr_lo': 20, 'hw_reg_flat_scr_hi': 21, 'hw_reg_xnack_mask': 22, - 'hw_reg_hw_id1': 23, 'hw_reg_hw_id2': 24, 'hw_reg_pops_packer': 25, 'hw_reg_ib_sts2': 28} name_str = m.group(1).lower() - hwreg_id = hwreg_names.get(name_str, int(name_str) if name_str.isdigit() else None) + hwreg_id = HWREG_IDS.get(name_str, int(name_str) if name_str.isdigit() else None) if hwreg_id is None: raise ValueError(f"unknown hwreg name: {name_str}") - offset = int(m.group(2)) if m.group(2) else 0 - size = int(m.group(3)) if m.group(3) else 32 - simm16 = ((size - 1) << 11) | (offset << 6) | hwreg_id - return (simm16, neg, abs_, hi_half) + offset, size = int(m.group(2)) if m.group(2) else 0, int(m.group(3)) if m.group(3) else 32 + return (((size - 1) << 11) | (offset << 6) | hwreg_id, neg, abs_, hi_half) raise ValueError(f"cannot parse operand: {op}") SMEM_OPS = {'s_load_b32', 's_load_b64', 's_load_b128', 's_load_b256', 's_load_b512', @@ -801,6 +574,13 @@ def asm(text: str) -> Inst: if mnemonic.replace('_e32', '') in vcc_ops and len(values) >= 5: values = [values[0], values[2], values[3]] if mnemonic.startswith('v_cmp') and len(values) >= 3 and operands[0].strip().lower() in ('vcc_lo', 'vcc_hi', 'vcc'): values = values[1:] + # CMPX instructions with _e64 suffix: prepend implicit EXEC_LO destination (vdst=126) + if 'cmpx' in mnemonic and mnemonic.endswith('_e64') and len(values) == 2: + values = [VGPR(126, 1)] + values + # Recalculate modifiers: parsed[0]=src0, parsed[1]=src1 (no vdst in user input) + neg_bits = sum((1 << i) for i, p in enumerate(parsed[:3]) if p[1]) + abs_bits = sum((1 << i) for i, p in enumerate(parsed[:3]) if p[2]) + opsel_bits = sum((1 << i) for i, p in enumerate(parsed[:2]) if p[3]) vop3sd_ops = {'v_div_scale_f32', 'v_div_scale_f64'} if mnemonic in vop3sd_ops and len(parsed) >= 5: neg_bits = sum((1 << i) for i, p in enumerate(parsed[2:5]) if p[1]) diff --git a/test/test_tiny.py b/test/test_tiny.py index d83221a77a..21c6d300a0 100644 --- a/test/test_tiny.py +++ b/test/test_tiny.py @@ -129,6 +129,23 @@ class TestTiny(unittest.TestCase): probs = Tensor.rand(1, 1, 28, 28).sequential(layers).tolist() self.assertEqual(len(probs[0]), 10) + def test_conv2d_backward_weight(self): + # Simple test for conv2d backward weight gradient - this exercises a kernel that was causing GPU hangs + conv = nn.Conv2d(1, 8, 5) + Tensor.realize(*[p.replace(Tensor.ones_like(p).contiguous()) for p in nn.state.get_parameters([conv])]) + for x in nn.state.get_parameters([conv]): x.requires_grad_() + out = Tensor.empty(4, 1, 14, 14).sequential([conv, Tensor.relu]) + out.sum().backward() + Tensor.realize(*[x.grad for x in nn.state.get_parameters([conv]) if x.grad is not None]) + + def test_conv2d_backward_weight_two_layers(self): + # Same as above but with 2 conv layers - this was causing GPU hangs + layers = [nn.Conv2d(1, 8, 5), Tensor.relu, nn.Conv2d(8, 8, 5), Tensor.relu] + Tensor.realize(*[p.replace(Tensor.ones_like(p).contiguous()) for p in nn.state.get_parameters(layers)]) + for x in nn.state.get_parameters(layers): x.requires_grad_() + Tensor.empty(4, 1, 14, 14).sequential(layers).sum().backward() + Tensor.realize(*[x.grad for x in nn.state.get_parameters(layers) if x.grad is not None]) + # TODO: this is failing because of how swizzling rewrites the ShapeTracker of the final STORE @unittest.skipIf(CI and Device.DEFAULT == "DSP", "failing because of make things that can't be images not images") def test_mnist_backward(self): diff --git a/tinygrad/renderer/rdna_new.py b/tinygrad/renderer/rdna_new.py index 62bb2bad6f..edc3e787fd 100644 --- a/tinygrad/renderer/rdna_new.py +++ b/tinygrad/renderer/rdna_new.py @@ -86,6 +86,16 @@ class RDNARenderer(Renderer): lds_size = 0 labels: dict[str, int] = {} # Label -> instruction index pending_waits: set[UOp] = set() # Track loads that need waits before use + # Allocate dedicated SGPRs for exec mask saving to avoid conflicts with kernel arguments + # These will be allocated on first use, which happens after all DEFINE_GLOBAL/VAR ops + exec_save_load: list[SGPR | None] = [None] # Wrapped in list for nonlocal mutation + exec_save_if: list[SGPR | None] = [None] + def get_exec_save_load() -> SGPR: + if exec_save_load[0] is None: exec_save_load[0] = ra.alloc_sgpr(None) + return s[exec_save_load[0]] if isinstance(exec_save_load[0], int) else exec_save_load[0] + def get_exec_save_if() -> SGPR: + if exec_save_if[0] is None: exec_save_if[0] = ra.alloc_sgpr(None) + return s[exec_save_if[0]] if isinstance(exec_save_if[0], int) else exec_save_if[0] def maybe_wait(srcs): """Emit waitcnt if any source (or transitive source) is pending from an async load.""" @@ -805,7 +815,7 @@ class RDNARenderer(Renderer): for j in range(itemsize // 4): code.append(v_mov_b32_e32(v[dst.idx + j] if hasattr(dst, 'idx') else dst, 0)) - # Set up exec mask based on condition (use s[13] to avoid conflict with IF/ENDIF which uses s[12]) + # Set up exec mask based on condition (use dynamically allocated SGPR) cond_reg = get_reg(cond_uop) code.append(v_cmp_ne_i32_e32(0, cond_reg)) # VCC = (cond != 0) # Clamp address to 0 for masked lanes to prevent invalid memory accesses @@ -815,7 +825,7 @@ class RDNARenderer(Renderer): clamped_addr = ra.alloc_vgpr(u) code.append(v_cndmask_b32_e64(clamped_addr, 0, addr, VCC_LO)) # clamped = cond ? addr : 0 addr = clamped_addr # Use clamped address for this load only - code.append(s_and_saveexec_b32(s[13], VCC_LO)) # Save exec, mask with condition + code.append(s_and_saveexec_b32(get_exec_save_load(), VCC_LO)) # Save exec, mask with condition # Now do the load (only executed by lanes where condition is true) if itemsize == 1: @@ -831,7 +841,7 @@ class RDNARenderer(Renderer): # Restore exec mask if we masked it if cond_uop is not None: - code.append(s_mov_b32(EXEC_LO, s[13])) + code.append(s_mov_b32(EXEC_LO, get_exec_save_load())) r[u] = dst pending_waits.add(u) # Track that this load result needs wait before use @@ -886,7 +896,7 @@ class RDNARenderer(Renderer): buf_result = r.get(buf_uop) if buf_uop in r else get_reg(buf_uop) buf_reg = buf_result[0] if isinstance(buf_result, tuple) else buf_result - # Handle conditional store (mask exec for lanes where condition is false - use s[13]) + # Handle conditional store (mask exec for lanes where condition is false) if cond_uop is not None: cond_reg = get_reg(cond_uop) code.append(v_cmp_ne_i32_e32(0, cond_reg)) # VCC = (cond != 0) @@ -896,7 +906,7 @@ class RDNARenderer(Renderer): clamped_addr = ra.alloc_vgpr(val_uop) code.append(v_cndmask_b32_e64(clamped_addr, 0, addr, VCC_LO)) # clamped = cond ? addr : 0 addr = clamped_addr # Use clamped address for this store only - code.append(s_and_saveexec_b32(s[13], VCC_LO)) # Save exec, mask with condition + code.append(s_and_saveexec_b32(get_exec_save_load(), VCC_LO)) # Save exec, mask with condition if itemsize == 1: code.append(global_store_b8(addr=addr, data=val, saddr=buf_reg)) @@ -911,7 +921,7 @@ class RDNARenderer(Renderer): # Restore exec mask if we masked it if cond_uop is not None: - code.append(s_mov_b32(EXEC_LO, s[13])) + code.append(s_mov_b32(EXEC_LO, get_exec_save_load())) elif u.op is Ops.RANGE: loop_var = ra.alloc_vgpr(u) @@ -1020,7 +1030,7 @@ class RDNARenderer(Renderer): # Save exec and mask with condition cond = get_reg(u.src[0]) code.append(v_cmp_ne_i32_e32(0, cond)) # condition != 0 - code.append(s_and_saveexec_b32(s[12], VCC_LO)) # Save exec, AND with condition + code.append(s_and_saveexec_b32(get_exec_save_if(), VCC_LO)) # Save exec, AND with condition code.append(f"s_cbranch_execz .L_ENDIF_{i}") # Skip if all lanes masked elif u.op is Ops.ENDIF: @@ -1028,7 +1038,7 @@ class RDNARenderer(Renderer): if_uop = u.src[0] if_idx = uops.index(if_uop) code.append(f".L_ENDIF_{if_idx}:") - code.append(s_mov_b32(EXEC_LO, s[12])) # exec_lo = saved + code.append(s_mov_b32(EXEC_LO, get_exec_save_if())) # exec_lo = saved # Emit kernel prologue (load kernargs) prologue: list[Inst] = []