diff --git a/test/mockgpu/amd/emu.py b/test/mockgpu/amd/emu.py index f6557177bc..6ba3d6de3a 100644 --- a/test/mockgpu/amd/emu.py +++ b/test/mockgpu/amd/emu.py @@ -52,7 +52,7 @@ class _MXCSRContext: lib.set_fpcr(self._saved) from tinygrad.uop.ops import UOp, Ops, KernelInfo, AxisType -from tinygrad.dtype import dtypes +from tinygrad.dtype import dtypes, AddrSpace from tinygrad.device import Buffer, BufferSpec from tinygrad.runtime.autogen import hsa from tinygrad.helpers import Context, DEBUG, PROFILE, colored @@ -230,13 +230,11 @@ VOPD_TO_VOP2 = { ir4.VOPDOp.V_DUAL_MOV_B32: ir3.VOP1Op.V_MOV_B32_E32, ir4.VOPDOp.V_DUAL_CNDMASK_B32: ir3.VOP2Op.V_CNDMASK_B32_E32, ir4.VOPDOp.V_DUAL_FMAAK_F32: ir3.VOP2Op.V_FMAAK_F32_E32, ir4.VOPDOp.V_DUAL_FMAMK_F32: ir3.VOP2Op.V_FMAMK_F32_E32, } -WAVE_SIZE = 32 +def _wave_size(arch: str) -> int: return 64 if arch.startswith("cdna") else 32 # Special registers stored after inline constants (256-259) PC_LO_IDX, PC_HI_IDX, SCRATCH_STRIDE_IDX = 256, 257, 259 # SGPR buffer: 0-127 = SGPRs, 128-255 = inline constants, 256-259 = special registers -SGPR_COUNT, VGPR_SIZE = 260, 256 * 32 -# Sentinel PC value for s_endpgm -ENDPGM_PC = 0xFFFFFFFFFFFFFFFF +SGPR_COUNT = 260 def _op_name(inst) -> str: if hasattr(inst, 'opx'): return f"{inst.opx.name}_{inst.opy.name}" # VOPD has opx/opy not op @@ -246,7 +244,9 @@ def _to_u32(val: UOp) -> UOp: if val.dtype == dtypes.uint32: return val if val.dtype.itemsize == 4: return val.bitcast(dtypes.uint32) # same size: bitcast (float32->uint32) return val.cast(dtypes.uint32) # different size: cast (bool, int16, etc) -def _lane_active(exec_mask: UOp, lane: UOp) -> UOp: return ((exec_mask >> lane.cast(dtypes.uint32)) & _c(1)).ne(_c(0)) +def _lane_active(exec_mask: UOp, lane: UOp) -> UOp: + if exec_mask.dtype == dtypes.uint64: return ((exec_mask >> lane.cast(dtypes.uint64)) & UOp.const(dtypes.uint64, 1)).ne(UOp.const(dtypes.uint64, 0)) + return ((exec_mask >> lane.cast(dtypes.uint32)) & _c(1)).ne(_c(0)) def _hi16(v: UOp) -> UOp: return (v >> _c(16)) & _c(0xFFFF) def _cond(cond, if_true, if_false): """Select between values based on condition (works with UOp or bool).""" @@ -255,7 +255,13 @@ def _cond_hi16(cond, val: UOp) -> UOp: return _cond(cond, _hi16(val), val) def _apply_opsel(val: UOp, sel_bit: int, opsel: int) -> UOp: return _hi16(val) if opsel & (1 << sel_bit) else val def _set_lane_bit(old: UOp, lane: UOp, val: UOp, exec_mask: UOp) -> UOp: - """Set/clear a single bit in a 32-bit mask based on lane index, respecting exec mask.""" + """Set/clear a single bit in a mask based on lane index, respecting exec mask.""" + if old.dtype in (dtypes.uint64, dtypes.int64): + dt = dtypes.uint64 + mask = UOp.const(dt, 1) << lane.cast(dt) + new_bit = _to_u32(val).cast(dt) << lane.cast(dt) + cleared = old.cast(dt) & (mask ^ UOp.const(dt, 0xFFFFFFFFFFFFFFFF)) + return _lane_active(exec_mask, lane).where(cleared | new_bit, old.cast(dt)) mask = _c(1) << lane.cast(dtypes.uint32) new_bit = _to_u32(val) << lane.cast(dtypes.uint32) cleared = old & (mask ^ _c(MASK32)) @@ -410,27 +416,33 @@ def _collect_data_slices(assigns: list[tuple[str, UOp]], data_prefix: str, pcode class _Ctx: """Context for instruction compilation - holds buffers and helpers.""" - __slots__ = ('inst_size', 'dyn_fields', '_axis_id') + __slots__ = ('inst_size', 'dyn_fields', '_axis_id', 'wave_size', 'vgpr', 'accvgpr') sgpr = UOp(Ops.PARAM, dtypes.uint32.ptr(SGPR_COUNT), arg=0) - vgpr = UOp(Ops.PARAM, dtypes.uint32.ptr(VGPR_SIZE), arg=1) vmem = UOp(Ops.PARAM, dtypes.uint32.ptr(1 << 46), arg=2) lds = UOp(Ops.PARAM, dtypes.uint32.ptr(16384), arg=3) scratch = UOp(Ops.PARAM, dtypes.uint8.ptr(1 << 30), arg=4) - def __init__(self, inst_size: int): - self.inst_size, self._axis_id = inst_size, 0 + def __init__(self, inst_size: int, wave_size: int = 32): + self.inst_size, self._axis_id, self.wave_size = inst_size, 0, wave_size self.dyn_fields: list[tuple[int, int]] = [] # (lo, hi) of fields read dynamically + self.vgpr = UOp(Ops.PARAM, dtypes.uint32.ptr(256 * wave_size), arg=1) + self.accvgpr = UOp(Ops.PARAM, dtypes.uint32.ptr(256 * wave_size), arg=5) if wave_size == 64 else self.vgpr - def range(self, n: int = 32) -> UOp: + def range(self, n: int | None = None) -> UOp: """Create a lane range UOp with unique axis ID.""" + if n is None: n = self.wave_size self._axis_id += 1 return UOp.range(n, self._axis_id, AxisType.LOOP, dtype=dtypes.int) def unroll_lanes(self, get_lane_bit, exec_mask: UOp, apply_exec: bool = True) -> UOp: - """Combine 32 lane bits into a 32-bit mask using RANGE+REDUCE.""" + """Combine lane bits into a mask using RANGE+REDUCE (32-bit for RDNA, 64-bit for CDNA).""" lane = self.range() - bit = get_lane_bit(lane).cast(dtypes.uint32) << lane.cast(dtypes.uint32) - result = bit.reduce(lane, arg=Ops.ADD) + if self.wave_size <= 32: + bit = get_lane_bit(lane).cast(dtypes.uint32) << lane.cast(dtypes.uint32) + result = bit.reduce(lane, arg=Ops.ADD) + else: + bit = get_lane_bit(lane).cast(dtypes.uint64) << lane.cast(dtypes.uint64) + result = bit.reduce(lane, arg=Ops.ADD) return result & exec_mask if apply_exec else result def inst_word(self, dword_idx: int) -> UOp: @@ -480,6 +492,13 @@ class _Ctx: mask &= ~field_mask # zero dynamic bits in mask return base, mask, size + def rexec(self) -> UOp: + """Read full EXEC mask (32-bit for RDNA, 64-bit for CDNA).""" + lo = self.rsgpr_dyn(_c(EXEC_LO.offset)) + if self.wave_size <= 32: return lo + hi = self.rsgpr_dyn(_c(EXEC_LO.offset + 1)) + return _u64(lo, hi) + # Dynamic register access (takes UOp index instead of int) def rsgpr_dyn(self, reg: UOp, valid: UOp | None = None) -> UOp: """Read SGPR with dynamic register index.""" @@ -490,15 +509,38 @@ class _Ctx: """Write SGPR with dynamic register index. Writes to NULL (124) are discarded.""" return self.sgpr.index(reg.cast(dtypes.int), reg.ne(_c(124))).store(val.cast(dtypes.uint32)) + def wmask(self, reg: UOp, val: UOp) -> list[UOp]: + """Write a lane mask (VCC/EXEC). Splits into lo/hi for wave64.""" + if self.wave_size > 32: + lo, hi = _split64(val) + return [self.wsgpr_dyn(reg, lo), self.wsgpr_dyn(reg + _c(1), hi)] + return [self.wsgpr_dyn(reg, val)] + + def rmask(self, reg: UOp) -> UOp: + """Read a lane mask (VCC/EXEC). Combines lo/hi for wave64.""" + if self.wave_size > 32: return _u64(self.rsgpr_dyn(reg), self.rsgpr_dyn(reg + _c(1))) + return self.rsgpr_dyn(reg) + def rvgpr_dyn(self, reg: UOp, lane: UOp, valid: UOp | None = None) -> UOp: """Read VGPR with dynamic register index.""" - idx = reg.cast(dtypes.int) * _c(32, dtypes.int) + lane.cast(dtypes.int) + idx = reg.cast(dtypes.int) * _c(self.wave_size, dtypes.int) + lane.cast(dtypes.int) return self.vgpr.index(idx, valid, ptr=True).load() if valid is not None else self.vgpr.index(idx, ptr=True).load() def wvgpr_dyn(self, reg: UOp, lane: UOp, val: UOp, exec_mask: UOp, after: UOp | None = None) -> UOp: """Write VGPR with dynamic register index.""" buf = self.vgpr.after(after) if after is not None else self.vgpr - offset = reg.cast(dtypes.int) * _c(32, dtypes.int) + lane.cast(dtypes.int) + offset = reg.cast(dtypes.int) * _c(self.wave_size, dtypes.int) + lane.cast(dtypes.int) + return buf.index(offset, _lane_active(exec_mask, lane)).store(val.cast(dtypes.uint32)) + + def raccvgpr_dyn(self, reg: UOp, lane: UOp, valid: UOp | None = None) -> UOp: + """Read ACCVGPR with dynamic register index (CDNA only).""" + idx = reg.cast(dtypes.int) * _c(self.wave_size, dtypes.int) + lane.cast(dtypes.int) + return self.accvgpr.index(idx, valid, ptr=True).load() if valid is not None else self.accvgpr.index(idx, ptr=True).load() + + def waccvgpr_dyn(self, reg: UOp, lane: UOp, val: UOp, exec_mask: UOp, after: UOp | None = None) -> UOp: + """Write ACCVGPR with dynamic register index (CDNA only).""" + buf = self.accvgpr.after(after) if after is not None else self.accvgpr + offset = reg.cast(dtypes.int) * _c(self.wave_size, dtypes.int) + lane.cast(dtypes.int) return buf.index(offset, _lane_active(exec_mask, lane)).store(val.cast(dtypes.uint32)) def rsrc_dyn(self, off: UOp, lane: UOp | None, bits: int = 32, literal: UOp | None = None, is_f64: bool = False, do_cast: bool = True) -> UOp: @@ -557,14 +599,19 @@ class _Ctx: stores.extend([self.wsgpr_dyn(sdst_reg, lo), self.wsgpr_dyn(sdst_reg + _c(1), hi)]) else: stores.append(self.wsgpr_dyn(sdst_reg, _val_to_u32(val))) elif dest.startswith('SCC'): stores.append(self.wsgpr_dyn(_c(SCC.offset), _to_u32(val))) - elif dest.startswith('EXEC'): stores.append(self.wsgpr_dyn(_c(EXEC_LO.offset), _to_u32(val))) - elif dest.startswith('VCC'): stores.append(self.wsgpr_dyn(_c(VCC_LO.offset), _to_u32(val))) + elif dest.startswith('EXEC'): + if self.wave_size > 32 and val.dtype in (dtypes.uint64, dtypes.int64): + lo, hi = _split64(val) + stores.extend([self.wsgpr_dyn(_c(EXEC_LO.offset), lo), self.wsgpr_dyn(_c(EXEC_LO.offset + 1), hi)]) + else: stores.append(self.wsgpr_dyn(_c(EXEC_LO.offset), _to_u32(val))) + elif dest.startswith('VCC'): stores.extend(self.wmask(_c(VCC_LO.offset), val)) return stores def compile_sop_pcode(self, op, srcs: dict[str, UOp], sdst_reg: UOp, sdst_size: int) -> UOp: """Compile a scalar instruction with dynamic destination register.""" pcode = get_pcode(op) - srcs.update({'VCC': self.rsgpr_dyn(_c(VCC_LO.offset)), 'EXEC': self.rsgpr_dyn(_c(EXEC_LO.offset)), 'SCC': self.rsgpr_dyn(_c(SCC.offset))}) + srcs.update({'VCC': self.rmask(_c(VCC_LO.offset)), 'EXEC': self.rexec(), 'SCC': self.rsgpr_dyn(_c(SCC.offset)), + '_wave_size': self.wave_size}) if 'D0' not in srcs: srcs['D0'] = self.rsgpr_dyn(sdst_reg) # D0 is current dest value for read-modify-write ops _, assigns = parse_pcode(pcode, srcs) return UOp.sink(*self.scalar_stores(assigns, sdst_reg, sdst_size), *self.inc_pc()) @@ -579,7 +626,7 @@ class _Ctx: src2_off = self.inst_field(type(inst).src2) if hasattr(type(inst), 'src2') else None exec_lo = self.rsgpr_dyn(_c(EXEC_LO.offset)) srcs = { - 'SRC0': src0_reg, 'VDST': vdst_off, 'EXEC_LO': exec_lo, 'EXEC': exec_lo.cast(dtypes.uint64), '_vgpr': self.vgpr, + 'SRC0': src0_reg, 'VDST': vdst_off, 'EXEC_LO': exec_lo, 'EXEC': exec_lo.cast(dtypes.uint64), '_vgpr': self.vgpr, '_wave_size': self.wave_size, 'S0': self.rsrc_dyn(src0_off, _c(0, dtypes.int)) if 'WRITELANE' in op_name else src0_reg, 'S1': self.rsrc_dyn(src1_off, _c(0, dtypes.int)) if src1_off is not None else _c(0), 'S2': self.rsrc_dyn(src2_off, _c(0, dtypes.int)) if src2_off is not None else _c(0), @@ -597,9 +644,9 @@ class _Ctx: """Compile VOP instruction. Returns sink with stores and inc_pc.""" pcode = get_pcode(op) vcc_reg = sdst_reg if sdst_reg is not None else VCC_LO.offset - if 'VCC' not in srcs: srcs['VCC'] = self.rsgpr_dyn(_c(vcc_reg)) + if 'VCC' not in srcs: srcs['VCC'] = self.rmask(_c(vcc_reg)) srcs.update({'EXEC': exec_mask, 'SCC': self.rsgpr_dyn(_c(SCC.offset)), 'laneId': lane, 'VDST': vdst_reg, - 'ROUND_MODE': _c(0), 'ROUND_TOWARD_ZERO': _c(0), 'ROUND_NEAREST_EVEN': _c(0), '_vgpr': self.vgpr, + 'ROUND_MODE': _c(0), 'ROUND_TOWARD_ZERO': _c(0), 'ROUND_NEAREST_EVEN': _c(0), '_vgpr': self.vgpr, '_wave_size': self.wave_size, # CDNA SDWA byte/word select constants (E32 always uses BYTE0/WORD0 defaults) 'SDWA_SRC0_SEL': _c(0), 'BYTE0': _c(0), 'BYTE1': _c(1), 'BYTE2': _c(2), 'BYTE3': _c(3), 'WORD0': _c(0), 'WORD1': _c(1)}) # rounding mode and SDWA constants @@ -648,7 +695,9 @@ class _Ctx: raw_stores.append(('vgpr_direct', self.vgpr.index(vgpr_idx.cast(dtypes.int), active).store(new_val))) continue if 'D0' in dest and '[laneId]' in dest: - raw_stores.append(('vcc', self.wsgpr_dyn(_c(VCC_LO.offset), _set_lane_bit(self.rsgpr_dyn(_c(VCC_LO.offset)), lane, val, exec_mask)))) + old_vcc = self.rmask(_c(VCC_LO.offset)) + new_vcc = _set_lane_bit(old_vcc, lane, val, exec_mask) + raw_stores.extend([('vcc', s) for s in self.wmask(_c(VCC_LO.offset), new_vcc)]) elif dest.startswith('D0'): if (slice_match := re.match(r'D0\[(\d+)\s*:\s*(\d+)\]', dest)): hi_bit, lo_bit = int(slice_match.group(1)), int(slice_match.group(2)) @@ -694,7 +743,7 @@ class _Ctx: for mask_val, reg in [(vcc_val, vcc_reg), (exec_val, EXEC_LO.offset)]: if mask_val is None: continue def get_bit(l, v=mask_val): return (_to_u32(v.substitute({lane: l})) & _c(1)).cast(dtypes.uint32) - stores.append(self.wsgpr_dyn(_c(reg), self.unroll_lanes(get_bit, exec_mask, apply_exec=False))) + stores.extend(self.wmask(_c(reg), self.unroll_lanes(get_bit, exec_mask, apply_exec=False))) if lane_stores: stores.append(UOp.sink(*lane_stores).end(lane)) stores.extend(scalar_stores) return UOp.sink(*stores, *self.inc_pc()) @@ -718,9 +767,10 @@ def _compile_sopp(inst: ir3.SOPP | ir4.SOPP, ctx: _Ctx) -> UOp: if inst.op in _get_pcode_dict(inst.op): pcode = get_pcode(inst.op) pc_bytes = ctx.rpc() # PC is already 64-bit byte address - vcc, exec_lo = ctx.rsgpr_dyn(_c(VCC_LO.offset)), ctx.rsgpr_dyn(_c(EXEC_LO.offset)) + vcc, exec_val = ctx.rmask(_c(VCC_LO.offset)), ctx.rexec() srcs = {'PC': pc_bytes.cast(dtypes.int64), 'SIMM16': simm16, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'VCC': vcc, - 'VCCZ': vcc.eq(UOp.const(dtypes.uint32, 0)).cast(dtypes.uint32), 'EXECZ': exec_lo.eq(UOp.const(dtypes.uint32, 0)).cast(dtypes.uint32)} + 'VCCZ': vcc.eq(UOp.const(vcc.dtype, 0)).cast(dtypes.uint32), + 'EXECZ': exec_val.eq(UOp.const(exec_val.dtype, 0)).cast(dtypes.uint32)} for dest, val in parse_pcode(pcode, srcs)[1]: if dest == 'PC' or dest.startswith('PC.'): lo, hi = _split64(val.cast(dtypes.uint64)) @@ -833,7 +883,7 @@ def _sdwa_write(old: UOp, val: UOp, dst_sel: UOp, dst_unused: UOp) -> UOp: def _compile_sdwa(inst: irc.VOP1_SDWA | irc.VOP2_SDWA | irc.VOP2_SDWA_SDST | irc.VOPC_SDWA_SDST, ctx: _Ctx) -> UOp: """Compile CDNA SDWA (Sub-Dword Access) VOP1/VOP2/VOPC instructions.""" is_vopc = isinstance(inst, irc.VOPC_SDWA_SDST) - exec_mask, bits = ctx.rsgpr_dyn(_c(EXEC_LO.offset)), inst.canonical_op_bits + exec_mask, bits = ctx.rexec(), inst.canonical_op_bits # sd=1 means use sdst register, sd=0 means use VCC (for VOPC_SDWA_SDST and VOP2_SDWA_SDST) has_sdst = isinstance(inst, (irc.VOP2_SDWA_SDST, irc.VOPC_SDWA_SDST)) sdst_off = _c(inst.sdst.offset) if has_sdst and getattr(inst, 'sd', 0) else _c(VCC_LO.offset) @@ -879,9 +929,9 @@ def _compile_sdwa(inst: irc.VOP1_SDWA | irc.VOP2_SDWA | irc.VOP2_SDWA_SDST | irc if has_dst_sel: dst_sel = ctx.inst_field(type(inst).dst_sel) dst_unused = ctx.inst_field(type(inst).dst_unused) - srcs.update({'VCC': ctx.rsgpr_dyn(_c(VCC_LO.offset)), 'EXEC': exec_mask, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), + srcs.update({'VCC': ctx.rmask(_c(VCC_LO.offset)), 'EXEC': exec_mask, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'laneId': lane, 'VDST': vdst_reg, 'ROUND_MODE': _c(0), 'ROUND_TOWARD_ZERO': _c(0), - 'ROUND_NEAREST_EVEN': _c(0), '_vgpr': ctx.vgpr, + 'ROUND_NEAREST_EVEN': _c(0), '_vgpr': ctx.vgpr, '_wave_size': ctx.wave_size, 'SDWA_SRC0_SEL': _c(0), 'BYTE0': _c(0), 'BYTE1': _c(1), 'BYTE2': _c(2), 'BYTE3': _c(3), 'WORD0': _c(0), 'WORD1': _c(1)}) _, assigns = parse_pcode(pcode, srcs) @@ -897,11 +947,13 @@ def _compile_sdwa(inst: irc.VOP1_SDWA | irc.VOP2_SDWA | irc.VOP2_SDWA_SDST | irc result = _sdwa_write(old, result, dst_sel, dst_unused) stores.append(ctx.wvgpr_dyn(vdst_reg, lane, result, exec_mask)) elif dest.startswith('VCC'): - stores.append(ctx.wsgpr_dyn(_c(VCC_LO.offset), _set_lane_bit(ctx.rsgpr_dyn(_c(VCC_LO.offset)), lane, val, exec_mask))) + old_vcc = ctx.rmask(_c(VCC_LO.offset)) + stores.extend(ctx.wmask(_c(VCC_LO.offset), _set_lane_bit(old_vcc, lane, val, exec_mask))) if vcc_val is not None: # Initialize sdst to 0 before lane loop (old value may be unrelated data), then set lane bits in loop init_stores = [ctx.wsgpr_dyn(sdst_off, _c(0)), ctx.wsgpr_dyn(sdst_off + _c(1), _c(0))] - stores.append(ctx.wsgpr_dyn(sdst_off, _set_lane_bit(ctx.rsgpr_dyn(sdst_off), lane, vcc_val, exec_mask))) + old_sdst = ctx.rmask(sdst_off) + stores.extend(ctx.wmask(sdst_off, _set_lane_bit(old_sdst, lane, vcc_val, exec_mask))) if stores: return UOp.sink(*init_stores, UOp.sink(*stores).end(lane), *ctx.inc_pc()) return UOp.sink(*init_stores, *ctx.inc_pc()) @@ -912,7 +964,7 @@ def _compile_sdwa(inst: irc.VOP1_SDWA | irc.VOP2_SDWA | irc.VOP2_SDWA_SDST | irc def _compile_vop12(inst: ir3.VOP1 | ir3.VOP1_SDST | ir3.VOP2 | ir4.VOP1 | ir4.VOP1_SDST | ir4.VOP2 | irc.VOP1 | irc.VOP2, ctx: _Ctx) -> UOp: op_name = _op_name(inst) if op_name in ('V_READFIRSTLANE_B32_E32', 'V_PERMLANE64_B32_E32'): return ctx.compile_lane_pcode(inst.op, inst) - lane, exec_mask, bits = ctx.range(), ctx.rsgpr_dyn(_c(EXEC_LO.offset)), inst.canonical_op_bits + lane, exec_mask, bits = ctx.range(), ctx.rexec(), inst.canonical_op_bits literal = ctx.inst_field(type(inst).literal) if hasattr(type(inst), 'literal') else None # type: ignore[union-attr] is_f64 = 'F64' in op_name and 'B64' not in op_name vdst_reg = ctx.inst_field(type(inst).vdst) @@ -957,7 +1009,7 @@ def _compile_vop12(inst: ir3.VOP1 | ir3.VOP1_SDST | ir3.VOP2 | ir4.VOP1 | ir4.VO def _compile_vopc(inst: ir3.VOPC|ir3.VOP3|ir4.VOPC|ir4.VOP3|irc.VOPC|irc.VOP3, ctx: _Ctx, opsel: int = 0, abs_bits: int = 0, neg_bits: int = 0) -> UOp: - exec_mask, op_name, bits = ctx.rsgpr_dyn(_c(EXEC_LO.offset)), _op_name(inst), inst.canonical_op_bits + exec_mask, op_name, bits = ctx.rexec(), _op_name(inst), inst.canonical_op_bits is_cmpx, is_vopc = 'CMPX' in op_name, hasattr(inst, 'vsrc1') # is_vopc: e32 vs e64 # Handle both VOPC (vsrc1) and VOP3 (src1) instruction formats - read operands dynamically @@ -998,10 +1050,10 @@ def _compile_vopc(inst: ir3.VOPC|ir3.VOP3|ir4.VOPC|ir4.VOP3|irc.VOPC|irc.VOP3, c # CMPX e32: writes EXEC only; CMPX e64: writes both EXEC and SDST; non-CMPX: writes dst only if is_cmpx: - stores = [ctx.wsgpr_dyn(_c(EXEC_LO.offset), new_result)] - if not is_vopc: stores.append(ctx.wsgpr_dyn(dst_off, new_result)) + stores = ctx.wmask(_c(EXEC_LO.offset), new_result) + if not is_vopc: stores.extend(ctx.wmask(dst_off, new_result)) else: - stores = [ctx.wsgpr_dyn(dst_off, new_result)] if not is_vopc else [ctx.wsgpr_dyn(_c(VCC_LO.offset), new_result)] + stores = ctx.wmask(dst_off, new_result) if not is_vopc else ctx.wmask(_c(VCC_LO.offset), new_result) return UOp.sink(*stores, *ctx.inc_pc()) @@ -1026,7 +1078,7 @@ def _compile_bitop3(inst, ctx: _Ctx, exec_mask: UOp, bits: dict, op_name: str) - return UOp.sink(ctx.wvgpr_dyn(vdst_reg, lane, result.cast(dtypes.uint32), exec_mask).end(lane), *ctx.inc_pc()) def _compile_vop3(inst: ir3.VOP3 | ir4.VOP3 | irc.VOP3, ctx: _Ctx) -> UOp: - exec_mask = ctx.rsgpr_dyn(_c(EXEC_LO.offset)) + exec_mask = ctx.rexec() bits = inst.canonical_op_bits opsel, op_name = getattr(inst, 'opsel', 0) or 0, _op_name(inst) @@ -1082,7 +1134,7 @@ def _compile_vop3(inst: ir3.VOP3 | ir4.VOP3 | irc.VOP3, ctx: _Ctx) -> UOp: return ctx.compile_vop_pcode(inst.op, srcs, lane, vdst_reg, exec_mask, opsel_dst_hi=opsel_dst_hi, clmp=getattr(inst, 'clmp', 0)) def _compile_vop3sd(inst: ir3.VOP3SD | ir4.VOP3SD | irc.VOP3SD, ctx: _Ctx) -> UOp: - exec_mask = ctx.rsgpr_dyn(_c(EXEC_LO.offset)) + exec_mask = ctx.rexec() bits, pcode, ops = inst.canonical_op_bits, get_pcode(inst.op), inst.canonical_operands # Read operands dynamically from instruction encoding @@ -1094,7 +1146,7 @@ def _compile_vop3sd(inst: ir3.VOP3SD | ir4.VOP3SD | irc.VOP3SD, ctx: _Ctx) -> UO vcc_in_off = src2_off if has_carry_in else sdst_off def load_srcs(lane_uop): - ret = {'VCC': ctx.rsgpr_dyn(vcc_in_off), 'EXEC': exec_mask, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'laneId': lane_uop} + ret = {'VCC': ctx.rmask(vcc_in_off), 'EXEC': exec_mask, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'laneId': lane_uop} ret['S0'] = ctx.rsrc_dyn(src0_off, lane_uop, bits['s0'], literal, ops['s0'][0] == Fmt.FMT_NUM_F64) ret['S1'] = ctx.rsrc_dyn(src1_off, lane_uop, bits['s1'], literal, ops['s1'][0] == Fmt.FMT_NUM_F64) if 's2' in ops: ret['S2'] = ctx.rsrc_dyn(src2_off, lane_uop, bits['s2'], literal, ops['s2'][0] == Fmt.FMT_NUM_F64) @@ -1134,15 +1186,109 @@ def _compile_vop3sd(inst: ir3.VOP3SD | ir4.VOP3SD | irc.VOP3SD, ctx: _Ctx) -> UO else: d0_u32 = d0_val.bitcast(dtypes.uint32) if d0_val.dtype in (dtypes.float32, dtypes.half) else d0_val.cast(dtypes.uint32) vgpr_stores.append(ctx.wvgpr_dyn(vdst_reg, lane3, d0_u32, exec_mask)) - # Write carry output (wsgpr_dyn handles NULL register 124) - vcc_write = ctx.wsgpr_dyn(sdst_off, final_vcc) - return UOp.sink(vcc_write, UOp.group(*vgpr_stores).end(lane3), *ctx.inc_pc()) + # Write carry output (wmask handles lo/hi split for wave64) + vcc_writes = ctx.wmask(sdst_off, final_vcc) + return UOp.sink(*vcc_writes, UOp.group(*vgpr_stores).end(lane3), *ctx.inc_pc()) else: return ctx.compile_vop_pcode(inst.op, srcs, lane, vdst_reg, exec_mask, sdst_reg=inst.sdst.offset) +def _compile_mfma(inst: irc.VOP3P, ctx: _Ctx) -> UOp: + """CDNA MFMA 16x16xK matrix multiply-accumulate. + + Uses local temp arrays to cache inputs, avoiding aliasing issues when vdst overlaps src0/src1. + Phase 1: Read all input f32 values from VGPRs into temp arrays (range loop over 64 lanes). + Phase 2: Compute 256 output values using temp arrays and write to VGPRs (range loop over 64 lanes). + """ + op_name = _op_name(inst) + exec_mask = ctx.rexec() + vdst_reg = ctx.inst_field(type(inst).vdst) + src0_r = ctx.inst_field(type(inst).src0) - _c(256) + src1_r = ctx.inst_field(type(inst).src1) - _c(256) + src2_off = ctx.inst_field(type(inst).src2) + is_bf16 = 'BF16' in op_name + is_fp8 = 'FP8' in op_name or 'F8' in op_name + import re as _re + m = _re.search(r'(\d+)X(\d+)X(\d+)', op_name) + M, N, K = int(m.group(1)), int(m.group(2)), int(m.group(3)) + assert M == 16 and N == 16, f"only 16x16 MFMA supported, got {M}x{N}" + cvt = _FUNCS['bf16_to_f32'] if is_bf16 else _FUNCS['f16_to_f32'] + vpg = 4 if is_fp8 else 2 + k_per_grp = K // 4 + n_regs = k_per_grp // vpg + n_elems = 16 * K # total input elements per matrix + + # src2 can be VGPR (>=256) or inline constant/SGPR (<256) + src2_is_vgpr = src2_off >= _c(256) + src2_r = src2_off - _c(256) + acc_scalar = ctx.rsgpr_dyn(src2_off, src2_is_vgpr.ne(True)).bitcast(dtypes.float32) + + # Phase 1: Read inputs into a single temp array using a range loop. + # Layout: tmp[0..n_elems-1] = A[row][k], tmp[n_elems..2*n_elems-1] = B^T[col][k] + # Each lane holds k_per_grp elements. lane_in_grp = row/col, grp gives k offset. + b_off = UOp.const(dtypes.int, n_elems) + tmp = UOp(Ops.DEFINE_LOCAL, dtypes.float32.ptr(n_elems * 2, addrspace=AddrSpace.LOCAL), arg=(n_elems * 2,)) + + read_lane = ctx.range() + read_row = read_lane % UOp.const(dtypes.int, 16) + read_grp = read_lane // UOp.const(dtypes.int, 16) + + read_stores = [] + for kl in range(k_per_grp): + reg_idx, sub_idx = kl // vpg, kl % vpg + # Read A: raw from src0 at this lane's VGPR + a_raw = ctx.rvgpr_dyn(src0_r + _c(reg_idx), read_lane) + if is_fp8: + a_f = ((a_raw >> UOp.const(dtypes.uint32, sub_idx * 8)) & UOp.const(dtypes.uint32, 0xFF)).cast(dtypes.uint32) + else: + a_f = cvt((a_raw >> UOp.const(dtypes.uint32, sub_idx * 16)) & UOp.const(dtypes.uint32, 0xFFFF)) + # Store to tmp[row * K + grp * k_per_grp + kl] + a_idx = read_row * UOp.const(dtypes.int, K) + read_grp * UOp.const(dtypes.int, k_per_grp) + UOp.const(dtypes.int, kl) + read_stores.append(tmp.index(a_idx).store(a_f)) + + # Read B: raw from src1 at this lane's VGPR + b_raw = ctx.rvgpr_dyn(src1_r + _c(reg_idx), read_lane) + if is_fp8: + b_f = ((b_raw >> UOp.const(dtypes.uint32, sub_idx * 8)) & UOp.const(dtypes.uint32, 0xFF)).cast(dtypes.uint32) + else: + b_f = cvt((b_raw >> UOp.const(dtypes.uint32, sub_idx * 16)) & UOp.const(dtypes.uint32, 0xFFFF)) + b_idx = b_off + read_row * UOp.const(dtypes.int, K) + read_grp * UOp.const(dtypes.int, k_per_grp) + UOp.const(dtypes.int, kl) + read_stores.append(tmp.index(b_idx).store(b_f)) + + read_phase = UOp.group(*read_stores).end(read_lane) + + # Phase 2: Compute dot products and write outputs using a range loop. + # Each lane computes 4 output values: D[m][n] where m = grp*4 + out_reg, n = n_idx. + tmp2 = tmp.after(read_phase) + + compute_lane = ctx.range() + n_idx = compute_lane % UOp.const(dtypes.int, 16) + c_grp = compute_lane // UOp.const(dtypes.int, 16) + + compute_stores = [] + for out_reg in range(4): + # Read accumulator from ACCVGPR (or scalar constant) + acc_v = ctx.raccvgpr_dyn(src2_r + _c(out_reg), compute_lane, src2_is_vgpr).bitcast(dtypes.float32) + acc = src2_is_vgpr.where(acc_v, acc_scalar) + + # m = c_grp*4 + out_reg + m_base = c_grp * UOp.const(dtypes.int, 4) + UOp.const(dtypes.int, out_reg) + for k in range(K): + # A[m][k] from tmp2[m * K + k] -- no .load(), pm_add_loads adds it + a_val = tmp2.index(m_base * UOp.const(dtypes.int, K) + UOp.const(dtypes.int, k)) + # B[n][k] from tmp2[n_elems + n * K + k] + b_val = tmp2.index(b_off + n_idx * UOp.const(dtypes.int, K) + UOp.const(dtypes.int, k)) + acc = acc + a_val * b_val + + # Write output to ACCVGPR (MFMA destination is ACCVGPR) + compute_stores.append(ctx.waccvgpr_dyn(vdst_reg + _c(out_reg), compute_lane, acc.bitcast(dtypes.uint32), exec_mask)) + + compute_phase = UOp.group(*compute_stores).end(compute_lane) + + return UOp.sink(read_phase, compute_phase, *ctx.inc_pc()) + def _compile_wmma(inst: ir3.VOP3P | ir4.VOP3P | irc.VOP3P, ctx: _Ctx) -> UOp: op_name = _op_name(inst) - exec_mask = ctx.rsgpr_dyn(_c(EXEC_LO.offset)) + exec_mask = ctx.rexec() vdst_reg = ctx.inst_field(type(inst).vdst) src0_r = ctx.inst_field(type(inst).src0) - _c(256) src1_r = ctx.inst_field(type(inst).src1) - _c(256) @@ -1196,9 +1342,30 @@ def _compile_wmma(inst: ir3.VOP3P | ir4.VOP3P | irc.VOP3P, ctx: _Ctx) -> UOp: def _compile_vop3p(inst: ir3.VOP3P | ir4.VOP3P | irc.VOP3P, ctx: _Ctx) -> UOp: op_name = _op_name(inst) if 'WMMA' in op_name and ('16X16X16_F16' in op_name or '16X16X16_BF16' in op_name): return _compile_wmma(inst, ctx) + if 'MFMA' in op_name and '16X16X' in op_name and isinstance(inst, irc.VOP3P): return _compile_mfma(inst, ctx) + + # ACCVGPR_WRITE/READ/MOV: copies between VGPR and ACCVGPR register files + if 'ACCVGPR' in op_name: + lane = ctx.range() + exec_mask = ctx.rexec() + vdst_reg = ctx.inst_field(type(inst).vdst) + if 'READ' in op_name: + # v_accvgpr_read: VGPR[vdst] = ACCVGPR[src0] (src0 encoded as VGPR offset, but reads from ACCVGPR file) + src0_off = ctx.inst_field(type(inst).src0) - _c(256) + val = ctx.raccvgpr_dyn(src0_off, lane) + return UOp.sink(ctx.wvgpr_dyn(vdst_reg, lane, val, exec_mask).end(lane), *ctx.inc_pc()) + elif 'WRITE' in op_name: + # v_accvgpr_write: ACCVGPR[vdst] = src0 (src0 can be VGPR or SGPR/const) + src0 = ctx.rsrc_dyn(ctx.inst_field(type(inst).src0), lane, 32) + return UOp.sink(ctx.waccvgpr_dyn(vdst_reg, lane, src0, exec_mask).end(lane), *ctx.inc_pc()) + else: + # v_accvgpr_mov: ACCVGPR[vdst] = ACCVGPR[src0] + src0_off = ctx.inst_field(type(inst).src0) - _c(256) + val = ctx.raccvgpr_dyn(src0_off, lane) + return UOp.sink(ctx.waccvgpr_dyn(vdst_reg, lane, val, exec_mask).end(lane), *ctx.inc_pc()) lane = ctx.range() - exec_mask = ctx.rsgpr_dyn(_c(EXEC_LO.offset)) + exec_mask = ctx.rexec() vdst_reg = ctx.inst_field(type(inst).vdst) is_pk_f32 = 'PK' in op_name and 'F32' in op_name and 'MOV' not in op_name # CDNA packed F32 ops do_cast = any(x in op_name for x in ('F16', 'F32', 'BF16')) and 'IU' not in op_name and not is_pk_f32 @@ -1270,7 +1437,7 @@ def _compile_vop3p(inst: ir3.VOP3P | ir4.VOP3P | irc.VOP3P, ctx: _Ctx) -> UOp: return ctx.compile_vop_pcode(inst.op, srcs, lane, vdst_reg, exec_mask) def _compile_vopd(inst: ir3.VOPD | ir4.VOPD, ctx: _Ctx) -> UOp: - exec_mask = ctx.rsgpr_dyn(_c(EXEC_LO.offset)) + exec_mask = ctx.rexec() # Read operands dynamically - use type(inst) to get correct field descriptors inst_type = type(inst) vdstx_reg = ctx.inst_field(inst_type.vdstx) @@ -1296,9 +1463,9 @@ def _compile_vopd(inst: ir3.VOPD | ir4.VOPD, ctx: _Ctx) -> UOp: if vop in (ir3.VOP2Op.V_FMAAK_F32_E32, ir3.VOP2Op.V_FMAMK_F32_E32, ir3.VOP2Op.V_FMAAK_F32_E32, ir3.VOP2Op.V_FMAMK_F32_E32): assert literal is not None srcs['SIMM32'] = literal - if op in (ir3.VOPDOp.V_DUAL_CNDMASK_B32, ir4.VOPDOp.V_DUAL_CNDMASK_B32): srcs['VCC'] = ctx.rsgpr_dyn(_c(VCC_LO.offset)) + if op in (ir3.VOPDOp.V_DUAL_CNDMASK_B32, ir4.VOPDOp.V_DUAL_CNDMASK_B32): srcs['VCC'] = ctx.rmask(_c(VCC_LO.offset)) pcode = get_pcode(vop) - srcs.update({'VCC': ctx.rsgpr_dyn(_c(VCC_LO.offset)), 'EXEC': exec_mask, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'laneId': lane}) + srcs.update({'VCC': ctx.rmask(_c(VCC_LO.offset)), 'EXEC': exec_mask, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'laneId': lane}) for dest, val in parse_pcode(pcode, srcs)[1]: if dest.startswith('D0'): all_stores.append(ctx.wvgpr_dyn(vdst_reg, lane, _val_to_u32(val), exec_mask, after=srcy1)) return UOp.sink(UOp.group(*all_stores).end(lane), *ctx.inc_pc()) @@ -1306,7 +1473,7 @@ def _compile_vopd(inst: ir3.VOPD | ir4.VOPD, ctx: _Ctx) -> UOp: def _compile_mem_op(inst: ir3.DS|ir3.FLAT|ir3.GLOBAL|ir3.SCRATCH|ir4.DS|ir4.VFLAT|ir4.VGLOBAL|ir4.VSCRATCH |irc.DS|irc.FLAT|irc.GLOBAL|irc.SCRATCH, ctx: _Ctx) -> UOp: """Unified memory operation compiler for DS, FLAT, GLOBAL, SCRATCH.""" - exec_mask, op_name = ctx.rsgpr_dyn(_c(EXEC_LO.offset)), _op_name(inst) + exec_mask, op_name = ctx.rexec(), _op_name(inst) pcode = get_pcode(inst.op) # CDNA pcode uses CalcGlobalAddr/CalcDsAddr to compute address from raw components, but make_addr already handles this. # Strip the addr computation line and use pre-computed ADDR directly (rename 'addr' -> 'ADDR' in remaining pcode). @@ -1352,7 +1519,7 @@ def _compile_mem_op(inst: ir3.DS|ir3.FLAT|ir3.GLOBAL|ir3.SCRATCH|ir4.DS|ir4.VFLA if is_lds and 'PERMUTE' in op_name: pcode = get_pcode(inst.op) srcs = {'ADDR': addr_reg, 'DATA0': vdata_reg, 'VDST': vdst_reg, 'OFFSET': offset, - 'EXEC': exec_mask.cast(dtypes.uint64), '_vgpr': ctx.vgpr} + 'EXEC': exec_mask.cast(dtypes.uint64), '_vgpr': ctx.vgpr, '_wave_size': ctx.wave_size} _, assigns = parse_pcode(pcode, srcs) stores = [ctx.vgpr.index(val[0].cast(dtypes.int)).store(val[1].cast(dtypes.uint32)) for dest, val in assigns if dest.startswith('VGPR[')] return UOp.sink(*stores, *ctx.inc_pc()) @@ -1517,7 +1684,7 @@ def _get_runner(inst_bytes: bytes, arch: str = "rdna3"): break if handler is None: raise RuntimeError(f"[emu] unimplemented instruction type: {type(inst).__name__} {_op_name(inst)}") - ctx = _Ctx(inst_size) + ctx = _Ctx(inst_size, _wave_size(arch)) sink = handler(inst, ctx) base, mask, size = ctx.canonical_mask(inst_bytes) canonical_name = f"{_op_name(inst).lower()}_{base.to_bytes(size, 'little').hex()}" @@ -1553,29 +1720,41 @@ F32_INLINE = {240: 0x3f000000, 241: 0xbf000000, 242: 0x3f800000, 243: 0xbf800000 244: 0x40000000, 245: 0xc0000000, 246: 0x40800000, 247: 0xc0800000, 248: 0x3e22f983} # 2.0, -2.0, 4.0, -4.0, 1/(2*pi) class WaveState: - __slots__ = ('vgpr_buf', 'sgpr_buf', '_vgpr_mv', '_sgpr_mv', 'n_lanes') + __slots__ = ('vgpr_buf', 'sgpr_buf', 'accvgpr_buf', '_vgpr_mv', '_sgpr_mv', 'n_lanes', 'wave_size') - def __init__(self, n_lanes: int = WAVE_SIZE): - self.n_lanes = n_lanes - self.vgpr_buf = Buffer('CPU', VGPR_SIZE, dtypes.uint32).ensure_allocated() + def __init__(self, n_lanes: int, wave_size: int = 32): + self.n_lanes, self.wave_size = n_lanes, wave_size + vgpr_size = 256 * wave_size + self.vgpr_buf = Buffer('CPU', vgpr_size, dtypes.uint32).ensure_allocated() self.sgpr_buf = Buffer('CPU', SGPR_COUNT, dtypes.uint32).ensure_allocated() + # CDNA (wave64) has separate ACCVGPR file; RDNA shares with VGPR + if wave_size == 64: + self.accvgpr_buf = Buffer('CPU', vgpr_size, dtypes.uint32).ensure_allocated() + ctypes.memset(self.accvgpr_buf._buf.va_addr, 0, vgpr_size * 4) + else: + self.accvgpr_buf = self.vgpr_buf self._vgpr_mv = self.vgpr_buf.as_memoryview(force_zero_copy=True).cast('I') self._sgpr_mv = self.sgpr_buf.as_memoryview(force_zero_copy=True).cast('I') # Zero memory using ctypes memset (much faster than Python loops) - ctypes.memset(self.vgpr_buf._buf.va_addr, 0, VGPR_SIZE * 4) + ctypes.memset(self.vgpr_buf._buf.va_addr, 0, vgpr_size * 4) ctypes.memset(self.sgpr_buf._buf.va_addr, 0, SGPR_COUNT * 4) # Pre-populate inline constants at indices 128-255 for i in range(65): self._write_sgpr(128 + i, i) # 128-192: integers 0-64 for i in range(16): self._write_sgpr(193 + i, (-(i + 1)) & MASK32) # 193-208: -1 to -16 for off, val in F32_INLINE.items(): self._write_sgpr(off, val) # 240-248: float constants - self._write_sgpr(EXEC_LO.offset, (1 << n_lanes) - 1) + # EXEC mask: for 64-lane waves, set both EXEC_LO and EXEC_HI + if wave_size == 64: + self._write_sgpr(EXEC_LO.offset, (1 << min(n_lanes, 32)) - 1) + self._write_sgpr(EXEC_LO.offset + 1, (1 << max(n_lanes - 32, 0)) - 1 if n_lanes > 32 else 0) + else: + self._write_sgpr(EXEC_LO.offset, (1 << n_lanes) - 1) self._write_sgpr(PC_LO_IDX, 0) self._write_sgpr(PC_HI_IDX, 0) def _write_sgpr(self, idx: int, val: int): self._sgpr_mv[idx] = val & MASK32 def _read_sgpr(self, idx: int) -> int: return self._sgpr_mv[idx] - def _write_vgpr(self, reg: int, lane: int, val: int): self._vgpr_mv[reg * 32 + lane] = val & MASK32 - def _read_vgpr(self, reg: int, lane: int) -> int: return self._vgpr_mv[reg * 32 + lane] + def _write_vgpr(self, reg: int, lane: int, val: int): self._vgpr_mv[reg * self.wave_size + lane] = val & MASK32 + def _read_vgpr(self, reg: int, lane: int) -> int: return self._vgpr_mv[reg * self.wave_size + lane] @property def pc(self) -> int: return self._read_sgpr(PC_LO_IDX) | (self._read_sgpr(PC_HI_IDX) << 32) @@ -1622,11 +1801,12 @@ def run_asm(lib: int, lib_sz: int, gx: int, gy: int, gz: int, lx: int, ly: int, program: dict[int, tuple[Callable, list[int], bool, Inst]] = {} # pc -> (fxn, globals, is_barrier, inst) lds_size = ((rsrc2 & hsa.AMD_COMPUTE_PGM_RSRC_TWO_GRANULATED_LDS_SIZE) >> hsa.AMD_COMPUTE_PGM_RSRC_TWO_GRANULATED_LDS_SIZE_SHIFT) * 512 total_threads = lx * ly * lz + wave_size = _wave_size(arch) # Use Buffer objects with external_ptr=0 for vmem vmem_buf = Buffer('CPU', 1 << 40, dtypes.uint32, options=BufferSpec(external_ptr=0)).ensure_allocated() lds_buf = Buffer('CPU', max(lds_size // 4, 1), dtypes.uint32).ensure_allocated() - scratch_buf = Buffer('CPU', scratch_size * WAVE_SIZE, dtypes.uint8).ensure_allocated() if scratch_size else None + scratch_buf = Buffer('CPU', scratch_size * wave_size, dtypes.uint8).ensure_allocated() if scratch_size else None # Initialize SQTT encoder — emits packets inline as instructions execute (only when profiling) if PROFILE: @@ -1650,42 +1830,53 @@ def run_asm(lib: int, lib_sz: int, gx: int, gy: int, gz: int, lx: int, ly: int, for gidz in range(gz): for gidy in range(gy): for gidx in range(gx): - # Initialize all wavefronts for this workgroup - waves: list[tuple[WaveState, list]] = [] - for wave_start in range(0, total_threads, WAVE_SIZE): - st = _init_wave(lib, wave_start, total_threads, lx, ly, lz, args_ptr, rsrc2, scratch_size, arch, gidx, gidy, gidz, user_data) + for wave_start in range(0, total_threads, wave_size): + n_lanes, st = min(wave_size, total_threads - wave_start), WaveState(min(wave_size, total_threads - wave_start), wave_size) + st.pc = lib # Set PC to code base address + # Initialize user SGPRs: hardware loads COMPUTE_USER_DATA registers directly into s[0:N] + if user_data: + for i, val in enumerate(user_data): st._write_sgpr(i, val) + else: + st._write_sgpr(0, args_ptr & MASK32) + st._write_sgpr(1, (args_ptr >> 32) & MASK32) + + # Workgroup IDs in SGPRs after user SGPRs + sgpr_idx = (rsrc2 & hsa.AMD_COMPUTE_PGM_RSRC_TWO_USER_SGPR_COUNT) >> hsa.AMD_COMPUTE_PGM_RSRC_TWO_USER_SGPR_COUNT_SHIFT + for enabled, gid in [(hsa.AMD_COMPUTE_PGM_RSRC_TWO_ENABLE_SGPR_WORKGROUP_ID_X, gidx), + (hsa.AMD_COMPUTE_PGM_RSRC_TWO_ENABLE_SGPR_WORKGROUP_ID_Y, gidy), + (hsa.AMD_COMPUTE_PGM_RSRC_TWO_ENABLE_SGPR_WORKGROUP_ID_Z, gidz)]: + if rsrc2 & enabled: + st._write_sgpr(sgpr_idx, gid) + sgpr_idx += 1 + + # RDNA4 uses TTMP registers for workgroup IDs: ttmp[9]=gidx, ttmp[10]=gidy, ttmp[11]=gidz + if arch == "rdna4": + st._write_sgpr(ttmp[9].offset, gidx) + st._write_sgpr(ttmp[10].offset, gidy) + st._write_sgpr(ttmp[11].offset, gidz) + + # v0 = packed workitem IDs, scratch stride in secret SGPR + for lane in range(n_lanes): + tid = wave_start + lane + st._write_vgpr(0, lane, ((tid // (lx * ly)) << 20) | (((tid // lx) % ly) << 10) | (tid % lx)) + st._write_sgpr(SCRATCH_STRIDE_IDX, scratch_size) + + # Pass buffer addresses via ctypes (pre-create to avoid allocation in loop) c_bufs = [ctypes.c_uint64(st.sgpr_buf._buf.va_addr), ctypes.c_uint64(st.vgpr_buf._buf.va_addr), ctypes.c_uint64(vmem_buf._buf.va_addr), ctypes.c_uint64(lds_buf._buf.va_addr), - ctypes.c_uint64(scratch_buf._buf.va_addr if scratch_buf else 0)] - waves.append((st, c_bufs)) - - # Execute wavefronts with barrier synchronization - # Each wave runs until it hits s_barrier or s_endpgm. When all waves have stopped, release barrier waves. - done = [False] * len(waves) - for total_inst in range(10_000_000): - if all(done): break - for wi, (st, c_bufs) in enumerate(waves): - if done[wi]: continue - # Run this wave until barrier or endpgm - for _ in range(1_000_000): - pc = st.pc - if pc == ENDPGM_PC: - done[wi] = True - if tracing: sqtt_finish(wi) - break - fxn, globals_list, is_barrier, inst = _ensure_compiled(pc) - fxn(*[c_bufs[g] for g in globals_list]) - if tracing: - inst_op = inst.op.value if hasattr(inst, 'op') else 0 - sqtt_emit(wi, inst, (st.pc != ENDPGM_PC and st.pc != pc + inst.size()) if inst_op in _BRANCH_OPS else None) - if is_barrier: break # s_barrier hit: PC already advanced past it, pause this wave - else: raise RuntimeError("exceeded 1M instructions in single wave, likely infinite loop") - # All waves have either hit barrier or endpgm — release barrier waves for next round - else: raise RuntimeError("exceeded 10M total scheduling rounds") - tracing = False # only trace the first workgroup - - # Reset LDS for next workgroup - if lds_size > 0: ctypes.memset(lds_buf._buf.va_addr, 0, max(lds_size, 4)) - - if PROFILE: sqtt_traces.append(sqtt_finalize()) + ctypes.c_uint64(scratch_buf._buf.va_addr if scratch_buf else 0), + ctypes.c_uint64(st.accvgpr_buf._buf.va_addr)] + for inst_count in range(1_000_000): + if (pc := st.pc) == 0xFFFFFFFFFFFFFFFF: break + if pc not in program: + prev_len = len(_canonical_runner_cache) + runner = _decode_at(pc, arch) + program[pc] = (runner._prg.fxn, runner.p.globals) + if DEBUG >= 3: + inst = decode_inst(bytes((ctypes.c_char * 16).from_address(pc).raw), arch) + msg = f"[emu] PC={pc - lib}: {inst!r}" + print(colored(msg, 'green') if len(_canonical_runner_cache) > prev_len else msg) + fxn, globals_list = program[pc] + fxn(*[c_bufs[g] for g in globals_list]) + else: raise RuntimeError("exceeded 1M instructions, likely infinite loop") return 0 diff --git a/test/mockgpu/amd/pcode.py b/test/mockgpu/amd/pcode.py index d2fbc279ab..2a6a0a08a8 100644 --- a/test/mockgpu/amd/pcode.py +++ b/test/mockgpu/amd/pcode.py @@ -41,9 +41,11 @@ def _bitreverse(v: UOp, bits: int) -> UOp: def _extract_bits(val: UOp, hi: int, lo: int) -> UOp: dt = dtypes.uint64 if val.dtype in (dtypes.uint64, dtypes.int64) else dtypes.uint32 - result = ((val >> _const(dt, lo)) if lo > 0 else val) & _const(val.dtype, (1 << (hi - lo + 1)) - 1) - # Downcast to uint32 when extracting <=32 bits from a 64-bit value, so .f32 bitcast works correctly - if dt == dtypes.uint64 and (hi - lo + 1) <= 32: result = result.cast(dtypes.uint32) + width = hi - lo + 1 + result = ((val >> _const(dt, lo)) if lo > 0 else val) & _const(val.dtype, (1 << width) - 1) + # Downcast to narrowest fitting type so { hi, lo } concatenation computes correct shift + narrow = {8: dtypes.uint8, 16: dtypes.uint16, 32: dtypes.uint32, 64: dtypes.uint64}.get(width) + if narrow and narrow != dt: result = result.cast(narrow) return result def _set_bit(old, pos, val): @@ -536,7 +538,8 @@ class Parser: self.eat('RBRACKET') vgpr = self.vars.get('_vgpr') if vgpr is None: return _u32(0) - return vgpr.index(_to_u32(reg) * _u32(32) + _to_u32(lane), ptr=True).load() + ws = self.vars.get('_wave_size', 32) + return vgpr.index(_to_u32(reg) * _u32(ws) + _to_u32(lane), ptr=True).load() if self.try_eat('LPAREN'): args = self._parse_args() self.eat('RPAREN') @@ -548,8 +551,8 @@ class Parser: if name == 'OVERFLOW_F32': return _const(dtypes.uint32, 0x7F7FFFFF).bitcast(dtypes.float32) if name == 'UNDERFLOW_F64': return _const(dtypes.uint64, 1).bitcast(dtypes.float64) if name == 'OVERFLOW_F64': return _const(dtypes.uint64, 0x7FEFFFFFFFFFFFFF).bitcast(dtypes.float64) - if name == 'WAVE32': return _const(dtypes.bool, True) - if name == 'WAVE64': return _const(dtypes.bool, False) + if name == 'WAVE32': return _const(dtypes.bool, self.vars.get('_wave_size', 32) <= 32) + if name == 'WAVE64': return _const(dtypes.bool, self.vars.get('_wave_size', 32) > 32) if name == 'WAVE_MODE' and self.try_eat('DOT') and self.try_eat_val('IEEE', 'IDENT'): return _u32(1) if self.try_eat('LBRACE'): idx = self.eat('NUM').val @@ -561,7 +564,8 @@ class Parser: self.eat('RBRACKET') vgpr = self.vars.get('_vgpr') if vgpr is None: return _u32(0) - return vgpr.index(_to_u32(reg) * _u32(32) + _u32(int(idx)), ptr=True).load() + ws = self.vars.get('_wave_size', 32) + return vgpr.index(_to_u32(reg) * _u32(ws) + _u32(int(idx)), ptr=True).load() elem = self.vars.get(f'{name}@{idx}', self.vars.get(f'{name}{idx}')) if elem is None: # Extract bit idx from base variable (like var[idx]) @@ -1037,7 +1041,8 @@ def parse_block(lines: list[str], start: int, env: dict[str, VarVal], funcs: dic if j < len(toks) and toks[j].type == 'EQUALS': j += 1 ln = parse_tokens(lane_toks, env, funcs) rg, val = parse_tokens(reg_toks, env, funcs), parse_tokens(toks[j:], env, funcs) - vgpr_idx = _to_u32(rg) * _u32(32) + _to_u32(ln) + ws = env.get('_wave_size', 32) + vgpr_idx = _to_u32(rg) * _u32(ws) + _to_u32(ln) if assigns is not None: assigns.append((f'VGPR[{_tok_str(lane_toks)}][{_tok_str(reg_toks)}][{hi_val}:{lo_val}]', (vgpr_idx, val, hi_val, lo_val))) i += 1 @@ -1047,7 +1052,8 @@ def parse_block(lines: list[str], start: int, env: dict[str, VarVal], funcs: dic ln = parse_tokens(lane_toks, env, funcs) rg, val = parse_tokens(reg_toks, env, funcs), parse_tokens(toks[j:], env, funcs) if assigns is not None: - assigns.append((f'VGPR[{_tok_str(lane_toks)}][{_tok_str(reg_toks)}]', (_to_u32(rg) * _u32(32) + _to_u32(ln), val))) + ws = env.get('_wave_size', 32) + assigns.append((f'VGPR[{_tok_str(lane_toks)}][{_tok_str(reg_toks)}]', (_to_u32(rg) * _u32(ws) + _to_u32(ln), val))) i += 1 continue