mirror of
https://github.com/tinygrad/tinygrad.git
synced 2026-08-26 12:46:07 +00:00
Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7a1190b729 | ||
|
|
49d1bf93d6 | ||
|
|
433248c998 | ||
|
|
04c79505ec | ||
|
|
39f99b207a | ||
|
|
7e14cdcb06 | ||
|
|
69cdc8066d | ||
|
|
9c89be5235 | ||
|
|
2b838dc1d8 | ||
|
|
7f139a934f | ||
|
|
a19d21ea9c | ||
|
|
b557c46233 |
@@ -654,7 +654,7 @@ jobs:
|
||||
- name: Run process replay tests
|
||||
uses: ./.github/actions/process-replay
|
||||
|
||||
testrdna3:
|
||||
testamdasm:
|
||||
name: AMD ASM IDE
|
||||
runs-on: ubuntu-24.04
|
||||
timeout-minutes: 10
|
||||
@@ -677,8 +677,25 @@ jobs:
|
||||
run: cloc --by-file extra/assembly/amd/*.py
|
||||
- name: Run RDNA3 emulator tests
|
||||
run: python -m pytest -n=auto extra/assembly/amd/ --durations 20
|
||||
- name: Install pdfplumber
|
||||
run: pip install pdfplumber
|
||||
- name: Run RDNA3 emulator tests (AMD_LLVM=1)
|
||||
run: AMD_LLVM=1 python -m pytest -n=auto extra/assembly/amd/ --durations 20
|
||||
- name: Run RDNA3 dtype tests
|
||||
run: PYTHONPATH="." AMD=1 PYTHON_REMU=1 MOCKGPU=1 AMD_LLVM=0 pytest -n=auto test/test_dtype_alu.py test/test_dtype.py
|
||||
- name: Run RDNA3 dtype tests (AMD_LLVM=1)
|
||||
run: PYTHONPATH="." AMD=1 PYTHON_REMU=1 MOCKGPU=1 AMD_LLVM=1 pytest -n=auto test/test_dtype_alu.py test/test_dtype.py
|
||||
|
||||
testamdautogen:
|
||||
name: AMD autogen
|
||||
runs-on: ubuntu-24.04
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- name: Checkout Code
|
||||
uses: actions/checkout@v4
|
||||
- name: Setup Environment
|
||||
uses: ./.github/actions/setup-tinygrad
|
||||
with:
|
||||
key: rdna3-autogen
|
||||
pydeps: "pdfplumber"
|
||||
- name: Verify AMD autogen is up to date
|
||||
run: |
|
||||
python -m extra.assembly.amd.dsl --arch all
|
||||
|
||||
@@ -208,3 +208,9 @@ Key patterns to watch (from ResNet50 benchmark):
|
||||
- `vmin==vmax folding`: ~55ms, 0.33% match rate - checks 52K ops but rarely matches
|
||||
|
||||
Patterns with 0% match rate are workload-specific overhead. They may be useful in other workloads, so don't remove them without understanding their purpose.
|
||||
|
||||
## AMD Performance Counter Profiling
|
||||
|
||||
Set VIZ to `-2` to save performance counters traces for the AMD backend.
|
||||
|
||||
Use the CLI in `./extra/sqtt/roc.py` to explore the trace.
|
||||
|
||||
+603
-696
File diff suppressed because it is too large
Load Diff
@@ -2209,33 +2209,33 @@ buffer_atomic_xor_x2 = functools.partial(MUBUF, MUBUFOp.BUFFER_ATOMIC_XOR_X2)
|
||||
buffer_atomic_inc_x2 = functools.partial(MUBUF, MUBUFOp.BUFFER_ATOMIC_INC_X2)
|
||||
buffer_atomic_dec_x2 = functools.partial(MUBUF, MUBUFOp.BUFFER_ATOMIC_DEC_X2)
|
||||
cdna4 = functools.partial(MUBUF, MUBUFOp.CDNA4)
|
||||
scratch_load_ubyte = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_UBYTE, seg=2)
|
||||
scratch_load_sbyte = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SBYTE, seg=2)
|
||||
scratch_load_ushort = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_USHORT, seg=2)
|
||||
scratch_load_sshort = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SSHORT, seg=2)
|
||||
scratch_load_dword = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_DWORD, seg=2)
|
||||
scratch_load_dwordx2 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_DWORDX2, seg=2)
|
||||
scratch_load_dwordx3 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_DWORDX3, seg=2)
|
||||
scratch_load_dwordx4 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_DWORDX4, seg=2)
|
||||
scratch_store_byte = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_BYTE, seg=2)
|
||||
scratch_store_byte_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_BYTE_D16_HI, seg=2)
|
||||
scratch_store_short = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_SHORT, seg=2)
|
||||
scratch_store_short_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_SHORT_D16_HI, seg=2)
|
||||
scratch_store_dword = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_DWORD, seg=2)
|
||||
scratch_store_dwordx2 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_DWORDX2, seg=2)
|
||||
scratch_store_dwordx3 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_DWORDX3, seg=2)
|
||||
scratch_store_dwordx4 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_DWORDX4, seg=2)
|
||||
scratch_load_ubyte_d16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_UBYTE_D16, seg=2)
|
||||
scratch_load_ubyte_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_UBYTE_D16_HI, seg=2)
|
||||
scratch_load_sbyte_d16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SBYTE_D16, seg=2)
|
||||
scratch_load_sbyte_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SBYTE_D16_HI, seg=2)
|
||||
scratch_load_short_d16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SHORT_D16, seg=2)
|
||||
scratch_load_short_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SHORT_D16_HI, seg=2)
|
||||
scratch_load_lds_ubyte = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_UBYTE, seg=2)
|
||||
scratch_load_lds_sbyte = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_SBYTE, seg=2)
|
||||
scratch_load_lds_ushort = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_USHORT, seg=2)
|
||||
scratch_load_lds_sshort = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_SSHORT, seg=2)
|
||||
scratch_load_lds_dword = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_DWORD, seg=2)
|
||||
scratch_load_ubyte = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_UBYTE, seg=1)
|
||||
scratch_load_sbyte = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SBYTE, seg=1)
|
||||
scratch_load_ushort = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_USHORT, seg=1)
|
||||
scratch_load_sshort = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SSHORT, seg=1)
|
||||
scratch_load_dword = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_DWORD, seg=1)
|
||||
scratch_load_dwordx2 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_DWORDX2, seg=1)
|
||||
scratch_load_dwordx3 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_DWORDX3, seg=1)
|
||||
scratch_load_dwordx4 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_DWORDX4, seg=1)
|
||||
scratch_store_byte = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_BYTE, seg=1)
|
||||
scratch_store_byte_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_BYTE_D16_HI, seg=1)
|
||||
scratch_store_short = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_SHORT, seg=1)
|
||||
scratch_store_short_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_SHORT_D16_HI, seg=1)
|
||||
scratch_store_dword = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_DWORD, seg=1)
|
||||
scratch_store_dwordx2 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_DWORDX2, seg=1)
|
||||
scratch_store_dwordx3 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_DWORDX3, seg=1)
|
||||
scratch_store_dwordx4 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_DWORDX4, seg=1)
|
||||
scratch_load_ubyte_d16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_UBYTE_D16, seg=1)
|
||||
scratch_load_ubyte_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_UBYTE_D16_HI, seg=1)
|
||||
scratch_load_sbyte_d16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SBYTE_D16, seg=1)
|
||||
scratch_load_sbyte_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SBYTE_D16_HI, seg=1)
|
||||
scratch_load_short_d16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SHORT_D16, seg=1)
|
||||
scratch_load_short_d16_hi = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_SHORT_D16_HI, seg=1)
|
||||
scratch_load_lds_ubyte = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_UBYTE, seg=1)
|
||||
scratch_load_lds_sbyte = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_SBYTE, seg=1)
|
||||
scratch_load_lds_ushort = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_USHORT, seg=1)
|
||||
scratch_load_lds_sshort = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_SSHORT, seg=1)
|
||||
scratch_load_lds_dword = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_LDS_DWORD, seg=1)
|
||||
s_load_dword = functools.partial(SMEM, SMEMOp.S_LOAD_DWORD)
|
||||
s_load_dwordx2 = functools.partial(SMEM, SMEMOp.S_LOAD_DWORDX2)
|
||||
s_load_dwordx4 = functools.partial(SMEM, SMEMOp.S_LOAD_DWORDX4)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -56,6 +56,12 @@ class DSOp(IntEnum):
|
||||
DS_MAX_F32 = 19
|
||||
DS_NOP = 20
|
||||
DS_ADD_F32 = 21
|
||||
DS_GWS_SEMA_RELEASE_ALL = 24
|
||||
DS_GWS_INIT = 25
|
||||
DS_GWS_SEMA_V = 26
|
||||
DS_GWS_SEMA_BR = 27
|
||||
DS_GWS_SEMA_P = 28
|
||||
DS_GWS_BARRIER = 29
|
||||
DS_STORE_B8 = 30
|
||||
DS_STORE_B16 = 31
|
||||
DS_ADD_RTN_U32 = 32
|
||||
@@ -178,10 +184,13 @@ class FLATOp(IntEnum):
|
||||
FLAT_LOAD_D16_HI_B16 = 35
|
||||
FLAT_STORE_D16_HI_B8 = 36
|
||||
FLAT_STORE_D16_HI_B16 = 37
|
||||
GLOBAL_LOAD_ADDTID_B32 = 40
|
||||
GLOBAL_STORE_ADDTID_B32 = 41
|
||||
FLAT_ATOMIC_SWAP_B32 = 51
|
||||
FLAT_ATOMIC_CMPSWAP_B32 = 52
|
||||
FLAT_ATOMIC_ADD_U32 = 53
|
||||
FLAT_ATOMIC_SUB_U32 = 54
|
||||
FLAT_ATOMIC_CSUB_U32 = 55
|
||||
FLAT_ATOMIC_MIN_I32 = 56
|
||||
FLAT_ATOMIC_MIN_U32 = 57
|
||||
FLAT_ATOMIC_MAX_I32 = 58
|
||||
@@ -717,6 +726,7 @@ class SOPPOp(IntEnum):
|
||||
S_SET_INST_PREFETCH_DISTANCE = 4
|
||||
S_CLAUSE = 5
|
||||
S_DELAY_ALU = 7
|
||||
S_WAITCNT_DEPCTR = 8
|
||||
S_WAITCNT = 9
|
||||
S_WAIT_IDLE = 10
|
||||
S_WAIT_EVENT = 11
|
||||
@@ -1848,6 +1858,12 @@ ds_min_f32 = functools.partial(DS, DSOp.DS_MIN_F32)
|
||||
ds_max_f32 = functools.partial(DS, DSOp.DS_MAX_F32)
|
||||
ds_nop = functools.partial(DS, DSOp.DS_NOP)
|
||||
ds_add_f32 = functools.partial(DS, DSOp.DS_ADD_F32)
|
||||
ds_gws_sema_release_all = functools.partial(DS, DSOp.DS_GWS_SEMA_RELEASE_ALL)
|
||||
ds_gws_init = functools.partial(DS, DSOp.DS_GWS_INIT)
|
||||
ds_gws_sema_v = functools.partial(DS, DSOp.DS_GWS_SEMA_V)
|
||||
ds_gws_sema_br = functools.partial(DS, DSOp.DS_GWS_SEMA_BR)
|
||||
ds_gws_sema_p = functools.partial(DS, DSOp.DS_GWS_SEMA_P)
|
||||
ds_gws_barrier = functools.partial(DS, DSOp.DS_GWS_BARRIER)
|
||||
ds_store_b8 = functools.partial(DS, DSOp.DS_STORE_B8)
|
||||
ds_store_b16 = functools.partial(DS, DSOp.DS_STORE_B16)
|
||||
ds_add_rtn_u32 = functools.partial(DS, DSOp.DS_ADD_RTN_U32)
|
||||
@@ -1968,10 +1984,13 @@ flat_load_d16_hi_i8 = functools.partial(FLAT, FLATOp.FLAT_LOAD_D16_HI_I8)
|
||||
flat_load_d16_hi_b16 = functools.partial(FLAT, FLATOp.FLAT_LOAD_D16_HI_B16)
|
||||
flat_store_d16_hi_b8 = functools.partial(FLAT, FLATOp.FLAT_STORE_D16_HI_B8)
|
||||
flat_store_d16_hi_b16 = functools.partial(FLAT, FLATOp.FLAT_STORE_D16_HI_B16)
|
||||
global_load_addtid_b32 = functools.partial(FLAT, FLATOp.GLOBAL_LOAD_ADDTID_B32)
|
||||
global_store_addtid_b32 = functools.partial(FLAT, FLATOp.GLOBAL_STORE_ADDTID_B32)
|
||||
flat_atomic_swap_b32 = functools.partial(FLAT, FLATOp.FLAT_ATOMIC_SWAP_B32)
|
||||
flat_atomic_cmpswap_b32 = functools.partial(FLAT, FLATOp.FLAT_ATOMIC_CMPSWAP_B32)
|
||||
flat_atomic_add_u32 = functools.partial(FLAT, FLATOp.FLAT_ATOMIC_ADD_U32)
|
||||
flat_atomic_sub_u32 = functools.partial(FLAT, FLATOp.FLAT_ATOMIC_SUB_U32)
|
||||
flat_atomic_csub_u32 = functools.partial(FLAT, FLATOp.FLAT_ATOMIC_CSUB_U32)
|
||||
flat_atomic_min_i32 = functools.partial(FLAT, FLATOp.FLAT_ATOMIC_MIN_I32)
|
||||
flat_atomic_min_u32 = functools.partial(FLAT, FLATOp.FLAT_ATOMIC_MIN_U32)
|
||||
flat_atomic_max_i32 = functools.partial(FLAT, FLATOp.FLAT_ATOMIC_MAX_I32)
|
||||
@@ -2226,28 +2245,28 @@ buffer_atomic_cmpswap_f32 = functools.partial(MUBUF, MUBUFOp.BUFFER_ATOMIC_CMPSW
|
||||
buffer_atomic_min_f32 = functools.partial(MUBUF, MUBUFOp.BUFFER_ATOMIC_MIN_F32)
|
||||
buffer_atomic_max_f32 = functools.partial(MUBUF, MUBUFOp.BUFFER_ATOMIC_MAX_F32)
|
||||
buffer_atomic_add_f32 = functools.partial(MUBUF, MUBUFOp.BUFFER_ATOMIC_ADD_F32)
|
||||
scratch_load_u8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_U8, seg=2)
|
||||
scratch_load_i8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_I8, seg=2)
|
||||
scratch_load_u16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_U16, seg=2)
|
||||
scratch_load_i16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_I16, seg=2)
|
||||
scratch_load_b32 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_B32, seg=2)
|
||||
scratch_load_b64 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_B64, seg=2)
|
||||
scratch_load_b96 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_B96, seg=2)
|
||||
scratch_load_b128 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_B128, seg=2)
|
||||
scratch_store_b8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B8, seg=2)
|
||||
scratch_store_b16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B16, seg=2)
|
||||
scratch_store_b32 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B32, seg=2)
|
||||
scratch_store_b64 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B64, seg=2)
|
||||
scratch_store_b96 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B96, seg=2)
|
||||
scratch_store_b128 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B128, seg=2)
|
||||
scratch_load_d16_u8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_U8, seg=2)
|
||||
scratch_load_d16_i8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_I8, seg=2)
|
||||
scratch_load_d16_b16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_B16, seg=2)
|
||||
scratch_load_d16_hi_u8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_HI_U8, seg=2)
|
||||
scratch_load_d16_hi_i8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_HI_I8, seg=2)
|
||||
scratch_load_d16_hi_b16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_HI_B16, seg=2)
|
||||
scratch_store_d16_hi_b8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_D16_HI_B8, seg=2)
|
||||
scratch_store_d16_hi_b16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_D16_HI_B16, seg=2)
|
||||
scratch_load_u8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_U8, seg=1)
|
||||
scratch_load_i8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_I8, seg=1)
|
||||
scratch_load_u16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_U16, seg=1)
|
||||
scratch_load_i16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_I16, seg=1)
|
||||
scratch_load_b32 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_B32, seg=1)
|
||||
scratch_load_b64 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_B64, seg=1)
|
||||
scratch_load_b96 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_B96, seg=1)
|
||||
scratch_load_b128 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_B128, seg=1)
|
||||
scratch_store_b8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B8, seg=1)
|
||||
scratch_store_b16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B16, seg=1)
|
||||
scratch_store_b32 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B32, seg=1)
|
||||
scratch_store_b64 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B64, seg=1)
|
||||
scratch_store_b96 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B96, seg=1)
|
||||
scratch_store_b128 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_B128, seg=1)
|
||||
scratch_load_d16_u8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_U8, seg=1)
|
||||
scratch_load_d16_i8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_I8, seg=1)
|
||||
scratch_load_d16_b16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_B16, seg=1)
|
||||
scratch_load_d16_hi_u8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_HI_U8, seg=1)
|
||||
scratch_load_d16_hi_i8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_HI_I8, seg=1)
|
||||
scratch_load_d16_hi_b16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_LOAD_D16_HI_B16, seg=1)
|
||||
scratch_store_d16_hi_b8 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_D16_HI_B8, seg=1)
|
||||
scratch_store_d16_hi_b16 = functools.partial(FLAT, SCRATCHOp.SCRATCH_STORE_D16_HI_B16, seg=1)
|
||||
s_load_b32 = functools.partial(SMEM, SMEMOp.S_LOAD_B32)
|
||||
s_load_b64 = functools.partial(SMEM, SMEMOp.S_LOAD_B64)
|
||||
s_load_b128 = functools.partial(SMEM, SMEMOp.S_LOAD_B128)
|
||||
@@ -2485,6 +2504,7 @@ s_sleep = functools.partial(SOPP, SOPPOp.S_SLEEP)
|
||||
s_set_inst_prefetch_distance = functools.partial(SOPP, SOPPOp.S_SET_INST_PREFETCH_DISTANCE)
|
||||
s_clause = functools.partial(SOPP, SOPPOp.S_CLAUSE)
|
||||
s_delay_alu = functools.partial(SOPP, SOPPOp.S_DELAY_ALU)
|
||||
s_waitcnt_depctr = functools.partial(SOPP, SOPPOp.S_WAITCNT_DEPCTR)
|
||||
s_waitcnt = functools.partial(SOPP, SOPPOp.S_WAITCNT)
|
||||
s_wait_idle = functools.partial(SOPP, SOPPOp.S_WAIT_IDLE)
|
||||
s_wait_event = functools.partial(SOPP, SOPPOp.S_WAIT_EVENT)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+58
-15
@@ -6,22 +6,21 @@ from typing import overload, Annotated, TypeVar, Generic
|
||||
|
||||
# Bit field DSL
|
||||
class BitField:
|
||||
def __init__(self, hi: int, lo: int, name: str | None = None): self.hi, self.lo, self.name = hi, lo, name
|
||||
def __set_name__(self, owner, name): self.name, self._owner = name, owner
|
||||
def __init__(self, hi: int, lo: int, name: str | None = None): self.hi, self.lo, self.name, self._marker = hi, lo, name, None
|
||||
def __set_name__(self, owner, name):
|
||||
import typing
|
||||
self.name, self._owner = name, owner
|
||||
# Cache marker at class definition time
|
||||
hints = typing.get_type_hints(owner, include_extras=True)
|
||||
if name in hints:
|
||||
hint = hints[name]
|
||||
if typing.get_origin(hint) is Annotated:
|
||||
args = typing.get_args(hint)
|
||||
self._marker = args[1] if len(args) > 1 else None
|
||||
def __eq__(self, val: int) -> tuple[BitField, int]: return (self, val) # type: ignore
|
||||
def mask(self) -> int: return (1 << (self.hi - self.lo + 1)) - 1
|
||||
@property
|
||||
def marker(self) -> type | None:
|
||||
# Get marker from Annotated type hint if present
|
||||
import typing
|
||||
if hasattr(self, '_owner') and self.name:
|
||||
hints = typing.get_type_hints(self._owner, include_extras=True)
|
||||
if self.name in hints:
|
||||
hint = hints[self.name]
|
||||
if typing.get_origin(hint) is Annotated:
|
||||
args = typing.get_args(hint)
|
||||
return args[1] if len(args) > 1 else None
|
||||
return None
|
||||
def marker(self) -> type | None: return self._marker
|
||||
@overload
|
||||
def __get__(self, obj: None, objtype: type) -> BitField: ...
|
||||
@overload
|
||||
@@ -179,6 +178,21 @@ class Inst:
|
||||
raise ValueError(f"SOP1 {op_val.name} expects {expected} destination register(s), got {sdst_val.count}")
|
||||
if isinstance(ssrc0_val, Reg) and ssrc0_val.count != expected:
|
||||
raise ValueError(f"SOP1 {op_val.name} expects {expected} source register(s), got {ssrc0_val.count}")
|
||||
# FLAT: set sve=1 when addr is a VGPR for scratch only
|
||||
# For scratch (seg=1), sve=1 means addr VGPR is used; sve=0 means addr is "off"
|
||||
# For global (seg=2) and flat (seg=0), sve is always 0
|
||||
if self.__class__.__name__ == 'FLAT' and 'sve' in self._fields:
|
||||
seg_val = self._values.get('seg', 0)
|
||||
if isinstance(seg_val, RawImm): seg_val = seg_val.val
|
||||
addr_val = orig_args.get('addr')
|
||||
if seg_val == 1 and isinstance(addr_val, VGPR): self._values['sve'] = 1
|
||||
# VOP3P: v_fma_mix* instructions (opcodes 32-34) have opsel_hi default of 0, not 7
|
||||
if self.__class__.__name__ == 'VOP3P':
|
||||
op_val = orig_args.get(field_names[0]) if args else orig_args.get('op')
|
||||
if hasattr(op_val, 'value'): op_val = op_val.value
|
||||
if op_val in (32, 33, 34) and 'opsel_hi' not in orig_args and 'opsel_hi2' not in orig_args:
|
||||
self._values['opsel_hi'] = 0
|
||||
self._values['opsel_hi2'] = 0
|
||||
# Type check and encode values
|
||||
for name, val in list(self._values.items()):
|
||||
if name == 'encoding': continue
|
||||
@@ -283,6 +297,10 @@ class Inst:
|
||||
from extra.assembly.amd.autogen.rdna3 import VOP3Op
|
||||
try: op_name = VOP3Op(op).name
|
||||
except ValueError: pass
|
||||
if op_name is None and self.__class__.__name__ == 'VOPC':
|
||||
from extra.assembly.amd.autogen.rdna3 import VOPCOp
|
||||
try: op_name = VOPCOp(op).name
|
||||
except ValueError: pass
|
||||
if op_name is None: return False
|
||||
# V_LDEXP_F64 has 32-bit integer exponent in src1, so literal is 32-bit
|
||||
if op_name == 'V_LDEXP_F64': return False
|
||||
@@ -315,6 +333,9 @@ class Inst:
|
||||
op_val = inst._values.get('op', 0)
|
||||
has_literal = cls.__name__ == 'VOP2' and op_val in (44, 45, 55, 56)
|
||||
has_literal = has_literal or (cls.__name__ == 'SOP2' and op_val in (69, 70))
|
||||
# VOPD fmaak/fmamk always have a literal (opx/opy value 1 or 2)
|
||||
opx, opy = inst._values.get('opx', 0), inst._values.get('opy', 0)
|
||||
has_literal = has_literal or (cls.__name__ == 'VOPD' and (opx in (1, 2) or opy in (1, 2)))
|
||||
for n in SRC_FIELDS:
|
||||
if n in inst._values and isinstance(inst._values[n], RawImm) and inst._values[n].val == 255: has_literal = True
|
||||
if has_literal:
|
||||
@@ -333,6 +354,14 @@ class Inst:
|
||||
lit = f", literal={hex(self._literal)}" if self._literal is not None else ""
|
||||
return f"{self.__class__.__name__}({', '.join(f'{k}={v}' for k, v in items)}{lit})"
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
if name.startswith('_'): raise AttributeError(name)
|
||||
return unwrap(self._values.get(name, 0))
|
||||
|
||||
def lit(self, v: int) -> str:
|
||||
from extra.assembly.amd.asm import decode_src
|
||||
return f"0x{self._literal:x}" if v == 255 and self._literal else decode_src(v)
|
||||
|
||||
def __eq__(self, other):
|
||||
if not isinstance(other, Inst): return NotImplemented
|
||||
return self.__class__ == other.__class__ and self._values == other._values and self._literal == other._literal
|
||||
@@ -512,10 +541,24 @@ def _parse_single_pdf(url: str) -> dict:
|
||||
break
|
||||
formats[fmt_name] = fields
|
||||
|
||||
# fix known PDF errors
|
||||
# fix known PDF errors - assert if already present (so we know when the bug is fixed)
|
||||
if 'SMEM' in formats:
|
||||
formats['SMEM'] = [(n, 13 if n == 'DLC' else 14 if n == 'GLC' else h, 13 if n == 'DLC' else 14 if n == 'GLC' else l, e, t)
|
||||
for n, h, l, e, t in formats['SMEM']]
|
||||
# add missing opcodes not in PDF tables (RDNA3/RDNA3.5 specific)
|
||||
if doc_name in ('RDNA3', 'RDNA3.5'):
|
||||
if 'SOPPOp' in enums:
|
||||
assert 8 not in enums['SOPPOp'], "S_WAITCNT_DEPCTR now in PDF, remove workaround"
|
||||
enums['SOPPOp'][8] = 'S_WAITCNT_DEPCTR'
|
||||
if 'DSOp' in enums:
|
||||
gws_ops = {24: 'DS_GWS_SEMA_RELEASE_ALL', 25: 'DS_GWS_INIT', 26: 'DS_GWS_SEMA_V',
|
||||
27: 'DS_GWS_SEMA_BR', 28: 'DS_GWS_SEMA_P', 29: 'DS_GWS_BARRIER'}
|
||||
for k in gws_ops: assert k not in enums['DSOp'], f"{gws_ops[k]} now in PDF, remove workaround"
|
||||
enums['DSOp'].update(gws_ops)
|
||||
if 'FLATOp' in enums:
|
||||
flat_ops = {40: 'GLOBAL_LOAD_ADDTID_B32', 41: 'GLOBAL_STORE_ADDTID_B32', 55: 'FLAT_ATOMIC_CSUB_U32'}
|
||||
for k in flat_ops: assert k not in enums['FLATOp'], f"{flat_ops[k]} now in PDF, remove workaround"
|
||||
enums['FLATOp'].update(flat_ops)
|
||||
|
||||
return {"formats": formats, "enums": enums, "src_enum": src_enum, "doc_name": doc_name, "is_cdna": is_cdna}
|
||||
|
||||
@@ -601,7 +644,7 @@ def generate(output_path: str | None = None, arch: str = "rdna3") -> dict:
|
||||
for cls_name, ops in sorted(enums.items()):
|
||||
fmt = cls_name[:-2]
|
||||
for op_val, name in sorted(ops.items()):
|
||||
seg = {"GLOBAL": ", seg=2", "SCRATCH": ", seg=2"}.get(fmt, "")
|
||||
seg = {"GLOBAL": ", seg=2", "SCRATCH": ", seg=1"}.get(fmt, "")
|
||||
tgt = {"GLOBAL": "FLAT, GLOBALOp", "SCRATCH": "FLAT, SCRATCHOp"}.get(fmt, f"{fmt}, {cls_name}")
|
||||
if fmt in formats or fmt in ("GLOBAL", "SCRATCH"):
|
||||
if fmt in ("VOP1", "VOP2", "VOPC"):
|
||||
|
||||
+327
-337
@@ -1,9 +1,9 @@
|
||||
# RDNA3 emulator - executes compiled pseudocode from AMD ISA PDF
|
||||
# mypy: ignore-errors
|
||||
from __future__ import annotations
|
||||
import ctypes, os
|
||||
import ctypes
|
||||
from extra.assembly.amd.dsl import Inst, RawImm
|
||||
from extra.assembly.amd.pcode import _f32, _i32, _sext, _f16, _i16, _f64, _i64, Reg
|
||||
from extra.assembly.amd.pcode import _f32, _i32, _sext, _f16, _i16, _f64, _i64, Reg, SliceProxy
|
||||
from extra.assembly.amd.autogen.rdna3.gen_pcode import get_compiled_functions
|
||||
from extra.assembly.amd.autogen.rdna3 import (
|
||||
SOP1, SOP2, SOPC, SOPK, SOPP, SMEM, VOP1, VOP2, VOP3, VOP3SD, VOP3P, VOPC, DS, FLAT, VOPD, SrcEnum,
|
||||
@@ -89,85 +89,84 @@ def _get_compiled() -> dict:
|
||||
if _COMPILED is None: _COMPILED = get_compiled_functions()
|
||||
return _COMPILED
|
||||
|
||||
def _run_pcode(fn, op_cls, op, s0, s1, s2, d0, scc, vcc, lane, exec_mask, vdst_idx):
|
||||
"""Create Regs, run pseudocode, extract results."""
|
||||
# Determine flags from op_cls and op.name
|
||||
is_div_scale = 'DIV_SCALE' in op.name
|
||||
is_64 = op.name.endswith(('_B64', '_I64', '_U64', '_F64')) or op.name in ('V_MAD_U64_U32', 'V_MAD_I64_I32')
|
||||
is_cmp = op_cls.__name__ == 'VOPCOp' and not op.name.startswith('V_CMPX')
|
||||
is_cmpx = op_cls.__name__ == 'VOPCOp' and op.name.startswith('V_CMPX')
|
||||
has_sdst = op_cls.__name__ == 'VOP3SDOp'
|
||||
|
||||
# Create Regs - D0 gets s0 for DIV_SCALE (passthrough behavior)
|
||||
S0, S1, S2 = Reg(s0), Reg(s1), Reg(s2)
|
||||
D0, D1 = Reg(s0 if is_div_scale else d0), Reg(0)
|
||||
SCC, VCC, EXEC = Reg(scc), Reg(vcc), Reg(exec_mask)
|
||||
tmp = Reg(0)
|
||||
|
||||
# Call pseudocode
|
||||
ret = fn(S0, S1, S2, D0, D1, SCC, VCC, EXEC, tmp, lane)
|
||||
|
||||
# Build result
|
||||
result = {'d0': D0._val, 'scc': SCC._val & 1}
|
||||
if has_sdst or VCC._val != vcc: result['vcc_lane'] = (VCC._val >> lane) & 1
|
||||
if is_cmpx: result['exec_lane'] = (EXEC._val >> lane) & 1
|
||||
elif EXEC._val != exec_mask: result['exec'] = EXEC._val
|
||||
if is_cmp: result['vcc_lane'] = (D0._val >> lane) & 1
|
||||
if is_64: result['d0_64'] = True
|
||||
if D1._val: result['d1'] = D1._val & 1
|
||||
# V_WRITELANE_B32 returns (wr_lane, value) directly
|
||||
if ret is not None: result['vgpr_write'] = (ret[0], vdst_idx, ret[1])
|
||||
return result
|
||||
|
||||
class WaveState:
|
||||
__slots__ = ('sgpr', 'vgpr', 'scc', 'pc', 'literal', '_pend_sgpr', '_scc_reg', '_vcc_reg', '_exec_reg')
|
||||
__slots__ = ('sgpr', 'vgpr', 'scc', 'pc', 'literal', '_pend_sgpr')
|
||||
def __init__(self):
|
||||
self.sgpr = [Reg(0) for _ in range(SGPR_COUNT)]
|
||||
self.vgpr = [[Reg(0) for _ in range(VGPR_COUNT)] for _ in range(WAVE_SIZE)]
|
||||
self.sgpr[EXEC_LO]._val = 0xffffffff
|
||||
self.scc, self.pc, self.literal, self._pend_sgpr = 0, 0, 0, {}
|
||||
# Reg wrappers for pseudocode access
|
||||
self._scc_reg = Reg(0)
|
||||
self._vcc_reg = self.sgpr[VCC_LO]
|
||||
self._exec_reg = self.sgpr[EXEC_LO]
|
||||
self.sgpr, self.vgpr = [0] * SGPR_COUNT, [[0] * VGPR_COUNT for _ in range(WAVE_SIZE)]
|
||||
self.sgpr[EXEC_LO], self.scc, self.pc, self.literal, self._pend_sgpr = 0xffffffff, 0, 0, 0, {}
|
||||
|
||||
@property
|
||||
def vcc(self) -> int: return self.sgpr[VCC_LO]._val | (self.sgpr[VCC_HI]._val << 32)
|
||||
def vcc(self) -> int: return self.sgpr[VCC_LO] | (self.sgpr[VCC_HI] << 32)
|
||||
@vcc.setter
|
||||
def vcc(self, v: int): self.sgpr[VCC_LO]._val, self.sgpr[VCC_HI]._val = v & 0xffffffff, (v >> 32) & 0xffffffff
|
||||
def vcc(self, v: int): self.sgpr[VCC_LO], self.sgpr[VCC_HI] = v & 0xffffffff, (v >> 32) & 0xffffffff
|
||||
@property
|
||||
def exec_mask(self) -> int: return self.sgpr[EXEC_LO]._val | (self.sgpr[EXEC_HI]._val << 32)
|
||||
def exec_mask(self) -> int: return self.sgpr[EXEC_LO] | (self.sgpr[EXEC_HI] << 32)
|
||||
@exec_mask.setter
|
||||
def exec_mask(self, v: int): self.sgpr[EXEC_LO]._val, self.sgpr[EXEC_HI]._val = v & 0xffffffff, (v >> 32) & 0xffffffff
|
||||
def exec_mask(self, v: int): self.sgpr[EXEC_LO], self.sgpr[EXEC_HI] = v & 0xffffffff, (v >> 32) & 0xffffffff
|
||||
|
||||
def rsgpr(self, i: int) -> int: return 0 if i == NULL else self.scc if i == SCC else self.sgpr[i]._val if i < SGPR_COUNT else 0
|
||||
def rsgpr(self, i: int) -> int: return 0 if i == NULL else self.scc if i == SCC else self.sgpr[i] if i < SGPR_COUNT else 0
|
||||
def wsgpr(self, i: int, v: int):
|
||||
if i < SGPR_COUNT and i != NULL: self.sgpr[i]._val = v & 0xffffffff
|
||||
if i < SGPR_COUNT and i != NULL: self.sgpr[i] = v & 0xffffffff
|
||||
def rsgpr64(self, i: int) -> int: return self.rsgpr(i) | (self.rsgpr(i+1) << 32)
|
||||
def wsgpr64(self, i: int, v: int): self.wsgpr(i, v & 0xffffffff); self.wsgpr(i+1, (v >> 32) & 0xffffffff)
|
||||
|
||||
def rsrc(self, v: int, lane: int) -> int:
|
||||
if v < SGPR_COUNT: return self.sgpr[v]._val
|
||||
if v < SGPR_COUNT: return self.sgpr[v]
|
||||
if v == SCC: return self.scc
|
||||
if v < 255: return _INLINE_CONSTS[v - 128]
|
||||
if v == 255: return self.literal
|
||||
return self.vgpr[lane][v - 256]._val if v <= 511 else 0
|
||||
|
||||
def rsrc_reg(self, v: int, lane: int) -> Reg:
|
||||
"""Return the Reg object for a source operand."""
|
||||
if v < SGPR_COUNT: return self.sgpr[v]
|
||||
if v == SCC: self._scc_reg._val = self.scc; return self._scc_reg
|
||||
if v < 255: return Reg(_INLINE_CONSTS[v - 128])
|
||||
if v == 255: return Reg(self.literal)
|
||||
return self.vgpr[lane][v - 256] if v <= 511 else Reg(0)
|
||||
return self.vgpr[lane][v - 256] if v <= 511 else 0
|
||||
|
||||
def rsrc_f16(self, v: int, lane: int) -> int:
|
||||
"""Read source operand for VOP3P packed f16 operations. Uses f16 inline constants."""
|
||||
if v < SGPR_COUNT: return self.sgpr[v]._val
|
||||
if v < SGPR_COUNT: return self.sgpr[v]
|
||||
if v == SCC: return self.scc
|
||||
if v < 255: return _INLINE_CONSTS_F16[v - 128]
|
||||
if v == 255: return self.literal
|
||||
return self.vgpr[lane][v - 256]._val if v <= 511 else 0
|
||||
|
||||
def rsrc_reg_f16(self, v: int, lane: int) -> Reg:
|
||||
"""Return Reg for VOP3P source. Inline constants are f16 in low 16 bits only."""
|
||||
if v < SGPR_COUNT: return self.sgpr[v]
|
||||
if v == SCC: self._scc_reg._val = self.scc; return self._scc_reg
|
||||
if v < 255: return Reg(_INLINE_CONSTS_F16[v - 128]) # f16 inline constant
|
||||
if v == 255: return Reg(self.literal)
|
||||
return self.vgpr[lane][v - 256] if v <= 511 else Reg(0)
|
||||
return self.vgpr[lane][v - 256] if v <= 511 else 0
|
||||
|
||||
def rsrc64(self, v: int, lane: int) -> int:
|
||||
"""Read 64-bit source operand. For inline constants, returns 64-bit representation."""
|
||||
# Inline constants 128-254 need special handling for 64-bit ops
|
||||
if 128 <= v < 255: return _INLINE_CONSTS_F64[v - 128]
|
||||
if v == 255: return self.literal
|
||||
if v == 255: return self.literal # 32-bit literal, caller handles extension
|
||||
return self.rsrc(v, lane) | ((self.rsrc(v+1, lane) if v < VCC_LO or 256 <= v <= 511 else 0) << 32)
|
||||
|
||||
def rsrc_reg64(self, v: int, lane: int) -> Reg:
|
||||
"""Return Reg for 64-bit source operand. For inline constants, returns 64-bit f64 value."""
|
||||
if 128 <= v < 255: return Reg(_INLINE_CONSTS_F64[v - 128])
|
||||
if v == 255: return Reg(self.literal)
|
||||
if v < SGPR_COUNT: return Reg(self.sgpr[v]._val | (self.sgpr[v+1]._val << 32))
|
||||
if 256 <= v <= 511:
|
||||
vgpr_idx = v - 256
|
||||
return Reg(self.vgpr[lane][vgpr_idx]._val | (self.vgpr[lane][vgpr_idx + 1]._val << 32))
|
||||
return Reg(0)
|
||||
|
||||
def pend_sgpr_lane(self, reg: int, lane: int, val: int):
|
||||
if reg not in self._pend_sgpr: self._pend_sgpr[reg] = 0
|
||||
if val: self._pend_sgpr[reg] |= (1 << lane)
|
||||
def commit_pends(self):
|
||||
for reg, val in self._pend_sgpr.items(): self.sgpr[reg]._val = val
|
||||
for reg, val in self._pend_sgpr.items(): self.sgpr[reg] = val
|
||||
self._pend_sgpr.clear()
|
||||
|
||||
# Instruction decode
|
||||
@@ -281,245 +280,341 @@ def exec_scalar(st: WaveState, inst: Inst) -> int:
|
||||
fn = compiled.get(op_cls, {}).get(op)
|
||||
if fn is None: raise NotImplementedError(f"{op.name} not in pseudocode")
|
||||
|
||||
# Build context - handle 64-bit ops that need 64-bit source reads
|
||||
# 64-bit source ops: name ends with _B64, _I64, _U64 or contains _U64, _I64 before last underscore
|
||||
# Read sources - 64-bit ops need 64-bit source reads
|
||||
is_64bit_s0 = op.name.endswith(('_B64', '_I64', '_U64')) or '_U64_' in op.name or '_I64_' in op.name
|
||||
is_64bit_s0s1 = op_cls is SOPCOp and op in (SOPCOp.S_CMP_EQ_U64, SOPCOp.S_CMP_LG_U64)
|
||||
s0 = st.rsrc64(ssrc0, 0) if is_64bit_s0 or is_64bit_s0s1 else (st.rsrc(ssrc0, 0) if inst_type != SOPK else st.rsgpr(inst.sdst))
|
||||
is_64bit_sop2 = is_64bit_s0 and inst_type is SOP2
|
||||
s1 = st.rsrc64(inst.ssrc1, 0) if (is_64bit_sop2 or is_64bit_s0s1) else (st.rsrc(inst.ssrc1, 0) if inst_type in (SOP2, SOPC) else inst.simm16 if inst_type is SOPK else 0)
|
||||
s1 = st.rsrc64(inst.ssrc1, 0) if (is_64bit_sop2 or is_64bit_s0s1) else (st.rsrc(inst.ssrc1, 0) if inst_type in (SOP2, SOPC) else 0)
|
||||
s2 = inst.simm16 if inst_type is SOPK else 0 # SOPK: 16-bit immediate passed as S2
|
||||
d0 = st.rsgpr64(sdst) if (is_64bit_s0 or is_64bit_s0s1) and sdst is not None else (st.rsgpr(sdst) if sdst is not None else 0)
|
||||
literal = inst.simm16 if inst_type is SOPK else st.literal
|
||||
|
||||
# Create Reg objects for new calling convention
|
||||
S0, S1, S2, D0 = Reg(s0), Reg(s1), Reg(0), Reg(d0)
|
||||
SCC, VCC, EXEC = Reg(st.scc), Reg(st.vcc), Reg(st.exec_mask)
|
||||
|
||||
# Execute compiled function - fn(S0, S1, S2, D0, SCC, VCC, laneId, EXEC, SIMM16, VGPR, SRC0, VDST)
|
||||
fn(S0, S1, S2, D0, SCC, VCC, 0, EXEC, Reg(literal), None, 0, 0)
|
||||
|
||||
# Apply results from Reg objects
|
||||
is_64bit_d0 = is_64bit_s0 or is_64bit_s0s1
|
||||
# Execute and apply results
|
||||
result = _run_pcode(fn, op_cls, op, s0, s1, s2, d0, st.scc, st.vcc, 0, st.exec_mask, 0)
|
||||
if sdst is not None:
|
||||
if is_64bit_d0:
|
||||
st.wsgpr64(sdst, D0._val)
|
||||
else:
|
||||
st.wsgpr(sdst, D0._val)
|
||||
st.scc = SCC._val
|
||||
st.exec_mask = EXEC._val
|
||||
if result.get('d0_64'): st.wsgpr64(sdst, result['d0'])
|
||||
else: st.wsgpr(sdst, result['d0'])
|
||||
st.scc = result['scc']
|
||||
if 'exec' in result: st.exec_mask = result['exec']
|
||||
return 0
|
||||
|
||||
def exec_vector(st: WaveState, inst: Inst, lane: int, lds: bytearray | None = None,
|
||||
d0_override: 'Reg | None' = None, vcc_override: 'Reg | None' = None) -> None:
|
||||
"""Execute vector instruction for one lane.
|
||||
d0_override: For VOPC/VOP3-VOPC, use this Reg instead of st.sgpr[vdst] for D0 output.
|
||||
vcc_override: For VOP3SD, use this Reg instead of st.sgpr[sdst] for VCC output.
|
||||
"""
|
||||
def exec_vector(st: WaveState, inst: Inst, lane: int, lds: bytearray | None = None) -> None:
|
||||
"""Execute vector instruction for one lane."""
|
||||
compiled = _get_compiled()
|
||||
inst_type, V = type(inst), st.vgpr[lane]
|
||||
|
||||
# Memory ops (not ALU pseudocode)
|
||||
if inst_type is FLAT:
|
||||
op, addr_reg, data_reg, vdst, offset, saddr = inst.op, inst.addr, inst.data, inst.vdst, _sext(inst.offset, 13), inst.saddr
|
||||
addr = V[addr_reg]._val | (V[addr_reg+1]._val << 32)
|
||||
addr = (st.rsgpr64(saddr) + V[addr_reg]._val + offset) & 0xffffffffffffffff if saddr not in (NULL, 0x7f) else (addr + offset) & 0xffffffffffffffff
|
||||
addr = V[addr_reg] | (V[addr_reg+1] << 32)
|
||||
addr = (st.rsgpr64(saddr) + V[addr_reg] + offset) & 0xffffffffffffffff if saddr not in (NULL, 0x7f) else (addr + offset) & 0xffffffffffffffff
|
||||
if op in FLAT_LOAD:
|
||||
cnt, sz, sign = FLAT_LOAD[op]
|
||||
for i in range(cnt): val = mem_read(addr + i * sz, sz); V[vdst + i]._val = _sext(val, sz * 8) & 0xffffffff if sign else val
|
||||
for i in range(cnt): val = mem_read(addr + i * sz, sz); V[vdst + i] = _sext(val, sz * 8) & 0xffffffff if sign else val
|
||||
elif op in FLAT_STORE:
|
||||
cnt, sz = FLAT_STORE[op]
|
||||
for i in range(cnt): mem_write(addr + i * sz, sz, V[data_reg + i]._val & ((1 << (sz * 8)) - 1))
|
||||
for i in range(cnt): mem_write(addr + i * sz, sz, V[data_reg + i] & ((1 << (sz * 8)) - 1))
|
||||
elif op in FLAT_D16_LOAD:
|
||||
sz, sign, hi = FLAT_D16_LOAD[op]
|
||||
val = mem_read(addr, sz)
|
||||
if sign: val = _sext(val, sz * 8) & 0xffff
|
||||
if hi: V[vdst]._val = (V[vdst]._val & 0xffff) | (val << 16)
|
||||
else: V[vdst]._val = (V[vdst]._val & 0xffff0000) | (val & 0xffff)
|
||||
if hi: V[vdst] = (V[vdst] & 0xffff) | (val << 16) # upper 16 bits
|
||||
else: V[vdst] = (V[vdst] & 0xffff0000) | (val & 0xffff) # lower 16 bits
|
||||
elif op in FLAT_D16_STORE:
|
||||
sz, hi = FLAT_D16_STORE[op]
|
||||
val = (V[data_reg]._val >> 16) & 0xffff if hi else V[data_reg]._val & 0xffff
|
||||
val = (V[data_reg] >> 16) & 0xffff if hi else V[data_reg] & 0xffff
|
||||
mem_write(addr, sz, val & ((1 << (sz * 8)) - 1))
|
||||
else: raise NotImplementedError(f"FLAT op {op}")
|
||||
return
|
||||
|
||||
if inst_type is DS:
|
||||
op, addr, vdst = inst.op, (V[inst.addr]._val + inst.offset0) & 0xffff, inst.vdst
|
||||
op, addr, vdst = inst.op, (V[inst.addr] + inst.offset0) & 0xffff, inst.vdst
|
||||
if op in DS_LOAD:
|
||||
cnt, sz, sign = DS_LOAD[op]
|
||||
for i in range(cnt): val = int.from_bytes(lds[addr+i*sz:addr+i*sz+sz], 'little'); V[vdst + i]._val = _sext(val, sz * 8) & 0xffffffff if sign else val
|
||||
for i in range(cnt): val = int.from_bytes(lds[addr+i*sz:addr+i*sz+sz], 'little'); V[vdst + i] = _sext(val, sz * 8) & 0xffffffff if sign else val
|
||||
elif op in DS_STORE:
|
||||
cnt, sz = DS_STORE[op]
|
||||
for i in range(cnt): lds[addr+i*sz:addr+i*sz+sz] = (V[inst.data0 + i]._val & ((1 << (sz * 8)) - 1)).to_bytes(sz, 'little')
|
||||
for i in range(cnt): lds[addr+i*sz:addr+i*sz+sz] = (V[inst.data0 + i] & ((1 << (sz * 8)) - 1)).to_bytes(sz, 'little')
|
||||
else: raise NotImplementedError(f"DS op {op}")
|
||||
return
|
||||
|
||||
# VOPD: dual-issue, execute two ops using VOP2/VOP3 compiled functions
|
||||
# Both ops execute simultaneously using pre-instruction values, so read all inputs first
|
||||
if inst_type is VOPD:
|
||||
vdsty = (inst.vdsty << 1) | ((inst.vdstx & 1) ^ 1)
|
||||
# Read all source operands BEFORE any writes (dual-issue semantics)
|
||||
sx0, sx1 = Reg(st.rsrc(inst.srcx0, lane)), Reg(V[inst.vsrcx1]._val)
|
||||
sy0, sy1 = Reg(st.rsrc(inst.srcy0, lane)), Reg(V[inst.vsrcy1]._val)
|
||||
dx0, dy0 = Reg(V[inst.vdstx]._val), Reg(V[vdsty]._val)
|
||||
st._scc_reg._val = st.scc
|
||||
sx0, sx1 = st.rsrc(inst.srcx0, lane), V[inst.vsrcx1]
|
||||
sy0, sy1 = st.rsrc(inst.srcy0, lane), V[inst.vsrcy1]
|
||||
dx0, dy0 = V[inst.vdstx], V[vdsty]
|
||||
# FMAAK/FMAMK in VOPD use literal as S2
|
||||
literal = getattr(inst, '_literal', None) or 0
|
||||
res_x = res_y = None
|
||||
if (op_x := _VOPD_TO_VOP.get(inst.opx)):
|
||||
if (fn_x := compiled.get(type(op_x), {}).get(op_x)):
|
||||
fn_x(sx0, sx1, Reg(0), dx0, st._scc_reg, st.sgpr[VCC_LO], lane, st.sgpr[EXEC_LO], Reg(st.literal), None, Reg(0), Reg(inst.vdstx))
|
||||
if (fn := compiled.get(type(op_x), {}).get(op_x)):
|
||||
# opx 1=FMAMK, 2=FMAAK use literal
|
||||
sx2 = literal if inst.opx in (VOPDOp.V_DUAL_FMAMK_F32, VOPDOp.V_DUAL_FMAAK_F32) else 0
|
||||
res_x = _run_pcode(fn, type(op_x), op_x, sx0, sx1, sx2, dx0, st.scc, st.vcc, lane, st.exec_mask, 0)
|
||||
if (op_y := _VOPD_TO_VOP.get(inst.opy)):
|
||||
if (fn_y := compiled.get(type(op_y), {}).get(op_y)):
|
||||
fn_y(sy0, sy1, Reg(0), dy0, st._scc_reg, st.sgpr[VCC_LO], lane, st.sgpr[EXEC_LO], Reg(st.literal), None, Reg(0), Reg(vdsty))
|
||||
V[inst.vdstx]._val, V[vdsty]._val = dx0._val, dy0._val
|
||||
st.scc = st._scc_reg._val
|
||||
if (fn := compiled.get(type(op_y), {}).get(op_y)):
|
||||
# opy 1=FMAMK, 2=FMAAK use literal
|
||||
sy2 = literal if inst.opy in (VOPDOp.V_DUAL_FMAMK_F32, VOPDOp.V_DUAL_FMAAK_F32) else 0
|
||||
res_y = _run_pcode(fn, type(op_y), op_y, sy0, sy1, sy2, dy0, st.scc, st.vcc, lane, st.exec_mask, 0)
|
||||
# Write results after both ops complete
|
||||
if res_x: V[inst.vdstx] = res_x['d0']
|
||||
if res_y: V[vdsty] = res_y['d0']
|
||||
return
|
||||
|
||||
# Determine instruction format and get function
|
||||
is_vop3_vopc = False
|
||||
is_readlane = False
|
||||
# VOP3SD: has extra scalar dest for carry output
|
||||
if inst_type is VOP3SD:
|
||||
op = VOP3SDOp(inst.op)
|
||||
fn = compiled.get(VOP3SDOp, {}).get(op)
|
||||
if fn is None: raise NotImplementedError(f"{op.name} not in pseudocode")
|
||||
s0, s1, s2 = st.rsrc(inst.src0, lane), st.rsrc(inst.src1, lane), st.rsrc(inst.src2, lane)
|
||||
# For 64-bit src2 ops (V_MAD_U64_U32, V_MAD_I64_I32), read from consecutive registers
|
||||
mad64_ops = (VOP3SDOp.V_MAD_U64_U32, VOP3SDOp.V_MAD_I64_I32)
|
||||
if op in mad64_ops:
|
||||
s2 = (V[inst.src2 - 256] | (V[inst.src2 - 256 + 1] << 32)) if inst.src2 >= 256 else st.rsgpr64(inst.src2)
|
||||
d0 = V[inst.vdst]
|
||||
# For carry-in ops (V_*_CO_CI_*), src2 register contains the carry bitmask (not VCC).
|
||||
# The pseudocode uses VCC but in VOP3SD encoding, the actual carry source is inst.src2.
|
||||
carry_ops = (VOP3SDOp.V_ADD_CO_CI_U32, VOP3SDOp.V_SUB_CO_CI_U32, VOP3SDOp.V_SUBREV_CO_CI_U32)
|
||||
vcc_for_exec = st.rsgpr64(inst.src2) if op in carry_ops else st.vcc
|
||||
result = _run_pcode(fn, VOP3SDOp, op, s0, s1, s2, d0, st.scc, vcc_for_exec, lane, st.exec_mask, inst.vdst)
|
||||
if result.get('d0_64'):
|
||||
V[inst.vdst] = result['d0'] & 0xffffffff
|
||||
V[inst.vdst + 1] = (result['d0'] >> 32) & 0xffffffff
|
||||
else:
|
||||
V[inst.vdst] = result['d0'] & 0xffffffff
|
||||
if result.get('vcc_lane') is not None:
|
||||
st.pend_sgpr_lane(inst.sdst, lane, result['vcc_lane'])
|
||||
return
|
||||
|
||||
|
||||
|
||||
# Get op enum and sources (None means "no source" for that operand)
|
||||
if inst_type is VOP1:
|
||||
if inst.op == VOP1Op.V_NOP: return
|
||||
# V_READFIRSTLANE_B32: read from first active lane's VGPR -> SGPR (not in pseudocode - needs cross-lane access)
|
||||
if inst.op == VOP1Op.V_READFIRSTLANE_B32:
|
||||
first_lane = (st.exec_mask & -st.exec_mask).bit_length() - 1 if st.exec_mask else 0
|
||||
vgpr_idx = inst.src0 - 256 if inst.src0 >= 256 else inst.src0 # VGPR index
|
||||
st.wsgpr(inst.vdst, st.vgpr[first_lane][vgpr_idx])
|
||||
return
|
||||
op_cls, op, src0, src1, src2, vdst = VOP1Op, VOP1Op(inst.op), inst.src0, None, None, inst.vdst
|
||||
# V_READFIRSTLANE_B32 writes to SGPR, not VGPR
|
||||
is_readlane = inst.op == VOP1Op.V_READFIRSTLANE_B32
|
||||
elif inst_type is VOP2:
|
||||
op_cls, op, src0, src1, src2, vdst = VOP2Op, VOP2Op(inst.op), inst.src0, inst.vsrc1 + 256, None, inst.vdst
|
||||
op_cls, op = VOP2Op, VOP2Op(inst.op)
|
||||
# FMAAK/FMAMK use inline literal constant as S2
|
||||
literal = getattr(inst, '_literal', None)
|
||||
src0, src1, src2, vdst = inst.src0, inst.vsrc1 + 256, literal, inst.vdst
|
||||
elif inst_type is VOP3:
|
||||
# VOP3 ops 0-255 are VOPC comparisons encoded as VOP3 (use VOPCOp pseudocode)
|
||||
if inst.op < 256:
|
||||
# VOP3-encoded VOPC - destination is an SGPR (vdst field)
|
||||
op_cls, op, src0, src1, src2, vdst = VOPCOp, VOPCOp(inst.op), inst.src0, inst.src1, None, inst.vdst
|
||||
is_vop3_vopc = True
|
||||
else:
|
||||
op_cls, op, src0, src1, src2, vdst = VOP3Op, VOP3Op(inst.op), inst.src0, inst.src1, inst.src2, inst.vdst
|
||||
# V_READFIRSTLANE_B32 and V_READLANE_B32 write to SGPR
|
||||
is_readlane = inst.op in (VOP3Op.V_READFIRSTLANE_B32, VOP3Op.V_READLANE_B32)
|
||||
elif inst_type is VOP3SD:
|
||||
op_cls, op, src0, src1, src2, vdst = VOP3SDOp, VOP3SDOp(inst.op), inst.src0, inst.src1, inst.src2, inst.vdst
|
||||
# V_READFIRSTLANE_B32 in VOP3 encoding - same as VOP1 but with VOP3 format
|
||||
if op == VOP3Op.V_READFIRSTLANE_B32:
|
||||
first_lane = (st.exec_mask & -st.exec_mask).bit_length() - 1 if st.exec_mask else 0
|
||||
vgpr_idx = inst.src0 - 256 if inst.src0 >= 256 else inst.src0
|
||||
st.wsgpr(inst.vdst, st.vgpr[first_lane][vgpr_idx])
|
||||
return
|
||||
# V_READLANE_B32: read from specific lane's VGPR -> SGPR (lane specified in src1)
|
||||
if op == VOP3Op.V_READLANE_B32:
|
||||
read_lane = st.rsrc(inst.src1, lane) & 0x1f # Lane to read from (5 bits)
|
||||
vgpr_idx = inst.src0 - 256 if inst.src0 >= 256 else inst.src0
|
||||
st.wsgpr(inst.vdst, st.vgpr[read_lane][vgpr_idx])
|
||||
return
|
||||
# V_PERM_B32: byte permutation - not in pseudocode PDF, implement directly
|
||||
# D0[byte_i] = selector[byte_i] < 8 ? {src0, src1}[selector[byte_i]] : (selector[byte_i] >= 0xD ? 0xFF : 0x00)
|
||||
if op == VOP3Op.V_PERM_B32:
|
||||
s0, s1, s2 = st.rsrc(inst.src0, lane), st.rsrc(inst.src1, lane), st.rsrc(inst.src2, lane)
|
||||
# Combine src1 and src0 into 8-byte value: src1 is bytes 0-3, src0 is bytes 4-7
|
||||
combined = (s1 & 0xffffffff) | ((s0 & 0xffffffff) << 32)
|
||||
result = 0
|
||||
for i in range(4): # 4 result bytes
|
||||
sel = (s2 >> (i * 8)) & 0xff # byte selector for this position
|
||||
if sel <= 7: result |= (((combined >> (sel * 8)) & 0xff) << (i * 8)) # select byte from combined
|
||||
elif sel >= 0xd: result |= (0xff << (i * 8)) # 0xD-0xF: constant 0xFF
|
||||
# else 0x8-0xC: constant 0x00 (already 0)
|
||||
V[vdst] = result & 0xffffffff
|
||||
return
|
||||
elif inst_type is VOPC:
|
||||
op_cls, op, src0, src1, src2, vdst = VOPCOp, VOPCOp(inst.op), inst.src0, inst.vsrc1 + 256, None, VCC_LO
|
||||
elif inst_type is VOP3P:
|
||||
op_cls, op, src0, src1, src2, vdst = VOP3POp, VOP3POp(inst.op), inst.src0, inst.src1, inst.src2, inst.vdst
|
||||
# WMMA instructions are handled specially (only execute for lane 0)
|
||||
if op in (VOP3POp.V_WMMA_F32_16X16X16_F16, VOP3POp.V_WMMA_F16_16X16X16_F16):
|
||||
if lane == 0: exec_wmma(st, inst, op)
|
||||
# VOP3P: Packed 16-bit operations using compiled functions
|
||||
op = VOP3POp(inst.op)
|
||||
# WMMA: wave-level matrix multiply-accumulate (special handling - needs cross-lane access)
|
||||
if op in (VOP3POp.V_WMMA_F32_16X16X16_F16, VOP3POp.V_WMMA_F32_16X16X16_BF16, VOP3POp.V_WMMA_F16_16X16X16_F16):
|
||||
if lane == 0: # Only execute once per wave, write results for all lanes
|
||||
exec_wmma(st, inst, op)
|
||||
return
|
||||
# V_FMA_MIX: Mixed precision FMA - inputs can be f16 or f32 controlled by opsel
|
||||
if op in (VOP3POp.V_FMA_MIX_F32, VOP3POp.V_FMA_MIXLO_F16, VOP3POp.V_FMA_MIXHI_F16):
|
||||
opsel = getattr(inst, 'opsel', 0)
|
||||
opsel_hi = getattr(inst, 'opsel_hi', 0)
|
||||
neg = getattr(inst, 'neg', 0)
|
||||
neg_hi = getattr(inst, 'neg_hi', 0)
|
||||
vdst = inst.vdst
|
||||
# Read raw 32-bit values - for V_FMA_MIX, sources can be either f32 or f16
|
||||
s0_raw = st.rsrc(inst.src0, lane)
|
||||
s1_raw = st.rsrc(inst.src1, lane)
|
||||
s2_raw = st.rsrc(inst.src2, lane) if inst.src2 is not None else 0
|
||||
# opsel[i]=0: use as f32, opsel[i]=1: use hi f16 as f32
|
||||
# For src0: opsel[0], for src1: opsel[1], for src2: opsel[2]
|
||||
if opsel & 1: s0 = _f16((s0_raw >> 16) & 0xffff) # hi f16 -> f32
|
||||
else: s0 = _f32(s0_raw) # use as f32
|
||||
if opsel & 2: s1 = _f16((s1_raw >> 16) & 0xffff)
|
||||
else: s1 = _f32(s1_raw)
|
||||
if opsel & 4: s2 = _f16((s2_raw >> 16) & 0xffff)
|
||||
else: s2 = _f32(s2_raw)
|
||||
# Apply neg modifiers (for f32 values)
|
||||
if neg & 1: s0 = -s0
|
||||
if neg & 2: s1 = -s1
|
||||
if neg & 4: s2 = -s2
|
||||
# Compute FMA: d = s0 * s1 + s2
|
||||
result = s0 * s1 + s2
|
||||
V = st.vgpr[lane]
|
||||
if op == VOP3POp.V_FMA_MIX_F32:
|
||||
V[vdst] = _i32(result)
|
||||
elif op == VOP3POp.V_FMA_MIXLO_F16:
|
||||
lo = _i16(result) & 0xffff
|
||||
V[vdst] = (V[vdst] & 0xffff0000) | lo
|
||||
else: # V_FMA_MIXHI_F16
|
||||
hi = _i16(result) & 0xffff
|
||||
V[vdst] = (V[vdst] & 0x0000ffff) | (hi << 16)
|
||||
return
|
||||
# Use rsrc_f16 for VOP3P to get correct f16 inline constants
|
||||
s0_raw = st.rsrc_f16(inst.src0, lane)
|
||||
s1_raw = st.rsrc_f16(inst.src1, lane)
|
||||
s2_raw = st.rsrc_f16(inst.src2, lane) if inst.src2 is not None else 0
|
||||
# Handle opsel (which 16-bit halves to use for each source)
|
||||
opsel = getattr(inst, 'opsel', 0)
|
||||
opsel_hi = getattr(inst, 'opsel_hi', 3) # Default: use hi for hi result
|
||||
opsel_hi2 = getattr(inst, 'opsel_hi2', 1) # Default for src2
|
||||
# Handle neg modifiers for VOP3P
|
||||
# neg applies to lo result inputs, neg_hi applies to hi result inputs
|
||||
neg = getattr(inst, 'neg', 0)
|
||||
neg_hi = getattr(inst, 'neg_hi', 0)
|
||||
# Build "virtual" sources with halves arranged for pseudocode: lo half goes to [15:0], hi half goes to [31:16]
|
||||
# opsel bit 0/1/2 selects which half of src0/1/2 goes to the LO result
|
||||
# opsel_hi bit 0/1 selects which half of src0/1 goes to the HI result
|
||||
s0_lo = (s0_raw >> 16) & 0xffff if (opsel & 1) else s0_raw & 0xffff
|
||||
s1_lo = (s1_raw >> 16) & 0xffff if (opsel & 2) else s1_raw & 0xffff
|
||||
s2_lo = (s2_raw >> 16) & 0xffff if (opsel & 4) else s2_raw & 0xffff
|
||||
s0_hi = (s0_raw >> 16) & 0xffff if (opsel_hi & 1) else s0_raw & 0xffff
|
||||
s1_hi = (s1_raw >> 16) & 0xffff if (opsel_hi & 2) else s1_raw & 0xffff
|
||||
s2_hi = (s2_raw >> 16) & 0xffff if opsel_hi2 else s2_raw & 0xffff
|
||||
# Apply neg to lo result inputs (toggle f16 sign bit)
|
||||
if neg & 1: s0_lo ^= 0x8000
|
||||
if neg & 2: s1_lo ^= 0x8000
|
||||
if neg & 4: s2_lo ^= 0x8000
|
||||
# Apply neg_hi to hi result inputs
|
||||
if neg_hi & 1: s0_hi ^= 0x8000
|
||||
if neg_hi & 2: s1_hi ^= 0x8000
|
||||
if neg_hi & 4: s2_hi ^= 0x8000
|
||||
# Pack into format expected by pseudocode: [31:16] = hi input, [15:0] = lo input
|
||||
s0 = (s0_hi << 16) | s0_lo
|
||||
s1 = (s1_hi << 16) | s1_lo
|
||||
s2 = (s2_hi << 16) | s2_lo
|
||||
vdst = inst.vdst
|
||||
fn = compiled.get(VOP3POp, {}).get(op)
|
||||
if fn is None: raise NotImplementedError(f"{op.name} not in pseudocode")
|
||||
result = _run_pcode(fn, VOP3POp, op, s0, s1, s2, 0, st.scc, st.vcc, lane, st.exec_mask, vdst)
|
||||
st.vgpr[lane][vdst] = result['d0'] & 0xffffffff
|
||||
return
|
||||
else: raise NotImplementedError(f"Unknown vector type {inst_type}")
|
||||
|
||||
fn = compiled.get(op_cls, {}).get(op)
|
||||
if fn is None: raise NotImplementedError(f"{op.name} not in pseudocode")
|
||||
|
||||
# Build source Regs - get the actual register or create temp for inline constants
|
||||
# VOP3P uses f16 inline constants (16-bit value in low half only)
|
||||
if inst_type is VOP3P:
|
||||
S0 = st.rsrc_reg_f16(src0, lane)
|
||||
S1 = st.rsrc_reg_f16(src1, lane) if src1 is not None else Reg(0)
|
||||
S2 = st.rsrc_reg_f16(src2, lane) if src2 is not None else Reg(0)
|
||||
# Apply op_sel_hi modifiers: control which half is used for hi-half computation
|
||||
# opsel_hi[0]=0 means src0 hi comes from lo half, =1 means from hi half (default)
|
||||
# opsel_hi[1]=0 means src1 hi comes from lo half, =1 means from hi half (default)
|
||||
# opsel_hi2=0 means src2 hi comes from lo half, =1 means from hi half (default)
|
||||
opsel_hi = getattr(inst, 'opsel_hi', 3) # default 0b11
|
||||
opsel_hi2 = getattr(inst, 'opsel_hi2', 1) # default 1
|
||||
# If opsel_hi bit is 0, replicate lo half to hi half
|
||||
if not (opsel_hi & 1): # src0 hi from lo
|
||||
lo = S0._val & 0xffff
|
||||
S0 = Reg((lo << 16) | lo)
|
||||
if not (opsel_hi & 2): # src1 hi from lo
|
||||
lo = S1._val & 0xffff
|
||||
S1 = Reg((lo << 16) | lo)
|
||||
if not opsel_hi2: # src2 hi from lo
|
||||
lo = S2._val & 0xffff
|
||||
S2 = Reg((lo << 16) | lo)
|
||||
# Read sources (with VOP3 modifiers if applicable)
|
||||
neg, abs_ = (getattr(inst, 'neg', 0), getattr(inst, 'abs', 0)) if inst_type is VOP3 else (0, 0)
|
||||
opsel = getattr(inst, 'opsel', 0) if inst_type is VOP3 else 0
|
||||
def mod_src(val: int, idx: int) -> int:
|
||||
if (abs_ >> idx) & 1: val = _i32(abs(_f32(val)))
|
||||
if (neg >> idx) & 1: val = _i32(-_f32(val))
|
||||
return val
|
||||
def mod_src64(val: int, idx: int) -> int:
|
||||
if (abs_ >> idx) & 1: val = _i64(abs(_f64(val)))
|
||||
if (neg >> idx) & 1: val = _i64(-_f64(val))
|
||||
return val
|
||||
|
||||
# Determine if sources are 64-bit based on instruction type
|
||||
# For 64-bit shift ops: src0 is 32-bit (shift amount), src1 is 64-bit (value to shift)
|
||||
# For V_LDEXP_F64: src0 is 64-bit float, src1 is 32-bit integer exponent
|
||||
# For most other _B64/_I64/_U64/_F64 ops: all sources are 64-bit
|
||||
is_64bit_op = op.name.endswith(('_B64', '_I64', '_U64', '_F64'))
|
||||
is_ldexp_64 = op in (VOP3Op.V_LDEXP_F64,)
|
||||
is_shift_64 = op in (VOP3Op.V_LSHLREV_B64, VOP3Op.V_LSHRREV_B64, VOP3Op.V_ASHRREV_I64)
|
||||
is_16bit_src = op_cls is VOP3Op and op in _VOP3_16BIT_OPS and op not in _CVT_32_64_SRC_OPS
|
||||
is_vop2_16bit = op_cls is VOP2Op and op in _VOP2_16BIT_OPS # VOP2 16-bit ops use f16 inline constants
|
||||
|
||||
if is_shift_64:
|
||||
s0 = mod_src(st.rsrc(src0, lane), 0) # shift amount is 32-bit
|
||||
s1 = st.rsrc64(src1, lane) if src1 is not None else 0 # value to shift is 64-bit
|
||||
s2 = mod_src(st.rsrc(src2, lane), 2) if src2 is not None else 0
|
||||
elif is_ldexp_64:
|
||||
s0 = mod_src64(st.rsrc64(src0, lane), 0) # mantissa is 64-bit float
|
||||
s1 = mod_src(st.rsrc(src1, lane), 1) if src1 is not None else 0 # exponent is 32-bit int
|
||||
s2 = mod_src(st.rsrc(src2, lane), 2) if src2 is not None else 0
|
||||
elif is_64bit_op:
|
||||
s0 = mod_src64(st.rsrc64(src0, lane), 0)
|
||||
s1 = mod_src64(st.rsrc64(src1, lane), 1) if src1 is not None else 0
|
||||
s2 = mod_src64(st.rsrc64(src2, lane), 2) if src2 is not None else 0
|
||||
elif is_16bit_src:
|
||||
# For 16-bit source ops, opsel bits select which half to use
|
||||
s0_raw = mod_src(st.rsrc(src0, lane), 0)
|
||||
s1_raw = mod_src(st.rsrc(src1, lane), 1) if src1 is not None else 0
|
||||
s2_raw = mod_src(st.rsrc(src2, lane), 2) if src2 is not None else 0
|
||||
s0 = ((s0_raw >> 16) & 0xffff) if (opsel & 1) else (s0_raw & 0xffff)
|
||||
s1 = ((s1_raw >> 16) & 0xffff) if (opsel & 2) else (s1_raw & 0xffff)
|
||||
s2 = ((s2_raw >> 16) & 0xffff) if (opsel & 4) else (s2_raw & 0xffff)
|
||||
elif is_vop2_16bit:
|
||||
s0 = mod_src(st.rsrc_f16(src0, lane), 0)
|
||||
s1 = mod_src(st.rsrc(src1, lane), 1) if src1 is not None else 0
|
||||
s2 = mod_src(st.rsrc(src2, lane), 2) if src2 is not None else 0
|
||||
else:
|
||||
# Check if this is a 64-bit F64 op - needs 64-bit source reads for f64 operands
|
||||
# V_LDEXP_F64: S0 is f64, S1 is i32 (exponent)
|
||||
# V_ADD_F64, V_MUL_F64, etc: S0 and S1 are f64
|
||||
# VOP1 F64 ops (V_TRUNC_F64, V_FLOOR_F64, etc): S0 is f64
|
||||
is_f64_op = hasattr(op, 'name') and '_F64' in op.name
|
||||
is_ldexp_f64 = hasattr(op, 'name') and op.name == 'V_LDEXP_F64'
|
||||
if is_f64_op:
|
||||
S0 = st.rsrc_reg64(src0, lane)
|
||||
# V_LDEXP_F64: S1 is i32 exponent, not f64
|
||||
if is_ldexp_f64:
|
||||
S1 = st.rsrc_reg(src1, lane) if src1 is not None else Reg(0)
|
||||
else:
|
||||
S1 = st.rsrc_reg64(src1, lane) if src1 is not None else Reg(0)
|
||||
S2 = st.rsrc_reg64(src2, lane) if src2 is not None else Reg(0)
|
||||
s0 = mod_src(st.rsrc(src0, lane), 0)
|
||||
s1 = mod_src(st.rsrc(src1, lane), 1) if src1 is not None else 0
|
||||
# src2 can be a register index OR a raw literal value (for FMAAK/FMAMK)
|
||||
# If src2 > 511, it's a raw literal value, not a register index
|
||||
s2 = src2 if src2 is not None and src2 > 511 else (mod_src(st.rsrc(src2, lane), 2) if src2 is not None else 0)
|
||||
d0 = V[vdst] if not is_64bit_op else (V[vdst] | (V[vdst + 1] << 32))
|
||||
|
||||
# V_CNDMASK_B32: VOP3 encoding uses src2 as mask (not VCC); VOP2 uses VCC implicitly
|
||||
vcc_for_fn = st.rsgpr64(src2) if op in (VOP3Op.V_CNDMASK_B32,) and inst_type is VOP3 and src2 is not None and src2 < 256 else st.vcc
|
||||
|
||||
# Execute pseudocode
|
||||
result = _run_pcode(fn, op_cls, op, s0, s1, s2, d0, st.scc, vcc_for_fn, lane, st.exec_mask, vdst)
|
||||
|
||||
# Apply results
|
||||
if 'vgpr_write' in result:
|
||||
# Lane instruction wrote to VGPR: (lane, vgpr_idx, value)
|
||||
wr_lane, wr_idx, wr_val = result['vgpr_write']
|
||||
st.vgpr[wr_lane][wr_idx] = wr_val
|
||||
if 'vcc_lane' in result:
|
||||
# VOP2 carry instructions write carry to VCC implicitly; VOPC writes to vdst
|
||||
vcc_dst = VCC_LO if op_cls is VOP2Op and op in (VOP2Op.V_ADD_CO_CI_U32, VOP2Op.V_SUB_CO_CI_U32, VOP2Op.V_SUBREV_CO_CI_U32) else vdst
|
||||
st.pend_sgpr_lane(vcc_dst, lane, result['vcc_lane'])
|
||||
if 'exec_lane' in result:
|
||||
# V_CMPX instructions write to EXEC per-lane
|
||||
st.pend_sgpr_lane(EXEC_LO, lane, result['exec_lane'])
|
||||
if 'd0' in result and op_cls not in (VOPCOp,) and 'vgpr_write' not in result:
|
||||
# V_READFIRSTLANE_B32 and V_READLANE_B32 write to SGPR, not VGPR
|
||||
writes_to_sgpr = op in (VOP1Op.V_READFIRSTLANE_B32,) or \
|
||||
(op_cls is VOP3Op and op in (VOP3Op.V_READFIRSTLANE_B32, VOP3Op.V_READLANE_B32))
|
||||
is_16bit_dst = op in _VOP3_16BIT_DST_OPS or op in _VOP1_16BIT_DST_OPS
|
||||
if writes_to_sgpr:
|
||||
st.wsgpr(vdst, result['d0'] & 0xffffffff)
|
||||
elif result.get('d0_64') or is_64bit_op:
|
||||
V[vdst] = result['d0'] & 0xffffffff
|
||||
V[vdst + 1] = (result['d0'] >> 32) & 0xffffffff
|
||||
elif is_16bit_dst and inst_type is VOP3:
|
||||
# VOP3 16-bit ops: opsel[3] controls hi/lo destination
|
||||
if opsel & 8: V[vdst] = (V[vdst] & 0x0000ffff) | ((result['d0'] & 0xffff) << 16)
|
||||
else: V[vdst] = (V[vdst] & 0xffff0000) | (result['d0'] & 0xffff)
|
||||
else:
|
||||
S0 = st.rsrc_reg(src0, lane)
|
||||
S1 = st.rsrc_reg(src1, lane) if src1 is not None else Reg(0)
|
||||
S2 = st.rsrc_reg(src2, lane) if src2 is not None else Reg(0)
|
||||
# VOP3SD V_MAD_U64_U32 and V_MAD_I64_I32 need S2 as 64-bit from VGPR pair
|
||||
if inst_type is VOP3SD and op in (VOP3SDOp.V_MAD_U64_U32, VOP3SDOp.V_MAD_I64_I32) and src2 is not None:
|
||||
if 256 <= src2 <= 511: # VGPR
|
||||
vgpr_idx = src2 - 256
|
||||
S2 = Reg(V[vgpr_idx]._val | (V[vgpr_idx + 1]._val << 32))
|
||||
|
||||
# Apply source modifiers (neg, abs) for VOP3/VOP3SD
|
||||
if inst_type in (VOP3, VOP3SD):
|
||||
neg, abs_mod = getattr(inst, 'neg', 0), getattr(inst, 'abs', 0)
|
||||
if neg or abs_mod:
|
||||
# Apply to f32 values - need to handle as float
|
||||
import struct
|
||||
def apply_mods(reg, neg_bit, abs_bit):
|
||||
val = reg._val
|
||||
f = struct.unpack('<f', struct.pack('<I', val & 0xffffffff))[0]
|
||||
if abs_bit: f = abs(f)
|
||||
if neg_bit: f = -f
|
||||
return Reg(struct.unpack('<I', struct.pack('<f', f))[0])
|
||||
if neg & 1 or abs_mod & 1: S0 = apply_mods(S0, neg & 1, abs_mod & 1)
|
||||
if neg & 2 or abs_mod & 2: S1 = apply_mods(S1, neg & 2, abs_mod & 2)
|
||||
if neg & 4 or abs_mod & 4: S2 = apply_mods(S2, neg & 4, abs_mod & 4)
|
||||
|
||||
# Apply opsel for VOP3 f16 operations - select which half to use
|
||||
# opsel[0]: src0, opsel[1]: src1, opsel[2]: src2 (0=lo, 1=hi)
|
||||
if inst_type is VOP3:
|
||||
opsel = getattr(inst, 'opsel', 0)
|
||||
if opsel:
|
||||
# If opsel bit is set, swap lo and hi so that .f16 reads the hi half
|
||||
if opsel & 1: # src0 from hi
|
||||
S0 = Reg(((S0._val >> 16) & 0xffff) | (S0._val << 16))
|
||||
if opsel & 2: # src1 from hi
|
||||
S1 = Reg(((S1._val >> 16) & 0xffff) | (S1._val << 16))
|
||||
if opsel & 4: # src2 from hi
|
||||
S2 = Reg(((S2._val >> 16) & 0xffff) | (S2._val << 16))
|
||||
|
||||
# For VOPC and VOP3-encoded VOPC, D0 is an SGPR (VCC_LO for VOPC, vdst for VOP3 VOPC)
|
||||
# V_READFIRSTLANE_B32 and V_READLANE_B32 also write to SGPR
|
||||
# Use d0_override if provided (for batch execution with shared output register)
|
||||
is_vopc = inst_type is VOPC or (inst_type is VOP3 and is_vop3_vopc)
|
||||
if is_vopc:
|
||||
D0 = d0_override if d0_override is not None else st.sgpr[VCC_LO if inst_type is VOPC else vdst]
|
||||
elif is_readlane:
|
||||
D0 = st.sgpr[vdst]
|
||||
else:
|
||||
D0 = V[vdst]
|
||||
|
||||
# Execute compiled function - D0 is modified in place
|
||||
st._scc_reg._val = st.scc
|
||||
# For VOP3SD, pass sdst register as VCC parameter (carry-out destination)
|
||||
# Use vcc_override if provided (for batch execution with shared output register)
|
||||
# For VOP3 V_CNDMASK_B32, src2 specifies the condition selector (not VCC)
|
||||
if inst_type is VOP3SD:
|
||||
vcc_reg = vcc_override if vcc_override is not None else st.sgpr[inst.sdst]
|
||||
elif inst_type is VOP3 and op == VOP3Op.V_CNDMASK_B32 and src2 is not None:
|
||||
vcc_reg = st.rsrc_reg(src2, lane) # Use src2 as condition
|
||||
else:
|
||||
vcc_reg = st.sgpr[VCC_LO]
|
||||
# SRC0/VDST are VGPR indices (0-255), not hardware encoding (256-511)
|
||||
src0_idx = (src0 - 256) if src0 and src0 >= 256 else (src0 if src0 else 0)
|
||||
result = fn(S0, S1, S2, D0, st._scc_reg, vcc_reg, lane, st.sgpr[EXEC_LO], Reg(st.literal), st.vgpr, Reg(src0_idx), Reg(vdst))
|
||||
st.scc = st._scc_reg._val
|
||||
|
||||
# Handle special results
|
||||
if result:
|
||||
if 'vgpr_write' in result:
|
||||
wr_lane, wr_idx, wr_val = result['vgpr_write']
|
||||
st.vgpr[wr_lane][wr_idx]._val = wr_val
|
||||
|
||||
# 64-bit destination: write high 32 bits to next VGPR (determined from op name)
|
||||
is_64bit_dst = not is_vopc and not is_readlane and hasattr(op, 'name') and \
|
||||
any(s in op.name for s in ('_B64', '_I64', '_U64', '_F64'))
|
||||
if is_64bit_dst:
|
||||
V[vdst + 1]._val = (D0._val >> 32) & 0xffffffff
|
||||
D0._val = D0._val & 0xffffffff # Keep only low 32 bits in D0
|
||||
V[vdst] = result['d0'] & 0xffffffff
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# WMMA (Wave Matrix Multiply-Accumulate)
|
||||
@@ -570,12 +665,12 @@ def exec_wmma(st: WaveState, inst, op: VOP3POp) -> None:
|
||||
lane, reg = (i // 2) % 32, (i // 2) // 32
|
||||
lo = _i16(mat_d[i]) & 0xffff
|
||||
hi = _i16(mat_d[i + 1]) & 0xffff
|
||||
st.vgpr[lane][vdst + reg]._val = (hi << 16) | lo
|
||||
st.vgpr[lane][vdst + reg] = (hi << 16) | lo
|
||||
else:
|
||||
# Output is f32
|
||||
for i in range(256):
|
||||
lane, reg = i % 32, i // 32
|
||||
st.vgpr[lane][vdst + reg]._val = _i32(mat_d[i])
|
||||
st.vgpr[lane][vdst + reg] = _i32(mat_d[i])
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# MAIN EXECUTION LOOP
|
||||
@@ -584,113 +679,6 @@ def exec_wmma(st: WaveState, inst, op: VOP3POp) -> None:
|
||||
SCALAR_TYPES = {SOP1, SOP2, SOPC, SOPK, SOPP, SMEM}
|
||||
VECTOR_TYPES = {VOP1, VOP2, VOP3, VOP3SD, VOPC, FLAT, DS, VOPD, VOP3P}
|
||||
|
||||
# Pre-cache compiled functions for fast lookup
|
||||
_COMPILED_CACHE: dict | None = None
|
||||
def _get_fn(op_cls, op):
|
||||
global _COMPILED_CACHE
|
||||
if _COMPILED_CACHE is None: _COMPILED_CACHE = _get_compiled()
|
||||
return _COMPILED_CACHE.get(op_cls, {}).get(op)
|
||||
|
||||
def exec_vector_batch(st: WaveState, inst: Inst, exec_mask: int, n_lanes: int, lds: bytearray | None = None) -> None:
|
||||
"""Execute vector instruction for all active lanes at once."""
|
||||
compiled = _get_compiled()
|
||||
inst_type = type(inst)
|
||||
vgpr = st.vgpr
|
||||
|
||||
# Memory ops - still per-lane but inlined
|
||||
if inst_type is FLAT:
|
||||
op, addr_reg, data_reg, vdst, offset, saddr = inst.op, inst.addr, inst.data, inst.vdst, _sext(inst.offset, 13), inst.saddr
|
||||
if op in FLAT_LOAD:
|
||||
cnt, sz, sign = FLAT_LOAD[op]
|
||||
for lane in range(n_lanes):
|
||||
if not (exec_mask & (1 << lane)): continue
|
||||
V = vgpr[lane]
|
||||
addr = V[addr_reg]._val | (V[addr_reg+1]._val << 32)
|
||||
addr = (st.rsgpr64(saddr) + V[addr_reg]._val + offset) & 0xffffffffffffffff if saddr not in (NULL, 0x7f) else (addr + offset) & 0xffffffffffffffff
|
||||
for i in range(cnt): val = mem_read(addr + i * sz, sz); V[vdst + i]._val = _sext(val, sz * 8) & 0xffffffff if sign else val
|
||||
elif op in FLAT_STORE:
|
||||
cnt, sz = FLAT_STORE[op]
|
||||
for lane in range(n_lanes):
|
||||
if not (exec_mask & (1 << lane)): continue
|
||||
V = vgpr[lane]
|
||||
addr = V[addr_reg]._val | (V[addr_reg+1]._val << 32)
|
||||
addr = (st.rsgpr64(saddr) + V[addr_reg]._val + offset) & 0xffffffffffffffff if saddr not in (NULL, 0x7f) else (addr + offset) & 0xffffffffffffffff
|
||||
for i in range(cnt): mem_write(addr + i * sz, sz, V[data_reg + i]._val & ((1 << (sz * 8)) - 1))
|
||||
elif op in FLAT_D16_LOAD:
|
||||
sz, sign, hi = FLAT_D16_LOAD[op]
|
||||
for lane in range(n_lanes):
|
||||
if not (exec_mask & (1 << lane)): continue
|
||||
V = vgpr[lane]
|
||||
addr = V[addr_reg]._val | (V[addr_reg+1]._val << 32)
|
||||
addr = (st.rsgpr64(saddr) + V[addr_reg]._val + offset) & 0xffffffffffffffff if saddr not in (NULL, 0x7f) else (addr + offset) & 0xffffffffffffffff
|
||||
val = mem_read(addr, sz)
|
||||
if sign: val = _sext(val, sz * 8) & 0xffff
|
||||
if hi: V[vdst]._val = (V[vdst]._val & 0xffff) | (val << 16)
|
||||
else: V[vdst]._val = (V[vdst]._val & 0xffff0000) | (val & 0xffff)
|
||||
elif op in FLAT_D16_STORE:
|
||||
sz, hi = FLAT_D16_STORE[op]
|
||||
for lane in range(n_lanes):
|
||||
if not (exec_mask & (1 << lane)): continue
|
||||
V = vgpr[lane]
|
||||
addr = V[addr_reg]._val | (V[addr_reg+1]._val << 32)
|
||||
addr = (st.rsgpr64(saddr) + V[addr_reg]._val + offset) & 0xffffffffffffffff if saddr not in (NULL, 0x7f) else (addr + offset) & 0xffffffffffffffff
|
||||
val = (V[data_reg]._val >> 16) & 0xffff if hi else V[data_reg]._val & 0xffff
|
||||
mem_write(addr, sz, val & ((1 << (sz * 8)) - 1))
|
||||
else: raise NotImplementedError(f"FLAT op {op}")
|
||||
return
|
||||
|
||||
if inst_type is DS:
|
||||
op, vdst = inst.op, inst.vdst
|
||||
if op in DS_LOAD:
|
||||
cnt, sz, sign = DS_LOAD[op]
|
||||
for lane in range(n_lanes):
|
||||
if not (exec_mask & (1 << lane)): continue
|
||||
V = vgpr[lane]
|
||||
addr = (V[inst.addr]._val + inst.offset0) & 0xffff
|
||||
for i in range(cnt): val = int.from_bytes(lds[addr+i*sz:addr+i*sz+sz], 'little'); V[vdst + i]._val = _sext(val, sz * 8) & 0xffffffff if sign else val
|
||||
elif op in DS_STORE:
|
||||
cnt, sz = DS_STORE[op]
|
||||
for lane in range(n_lanes):
|
||||
if not (exec_mask & (1 << lane)): continue
|
||||
V = vgpr[lane]
|
||||
addr = (V[inst.addr]._val + inst.offset0) & 0xffff
|
||||
for i in range(cnt): lds[addr+i*sz:addr+i*sz+sz] = (V[inst.data0 + i]._val & ((1 << (sz * 8)) - 1)).to_bytes(sz, 'little')
|
||||
else: raise NotImplementedError(f"DS op {op}")
|
||||
return
|
||||
|
||||
# For VOPC, VOP3-encoded VOPC, and VOP3SD, we write per-lane bits to an SGPR.
|
||||
# The pseudocode does D0.u64[laneId] = bit or VCC.u64[laneId] = bit.
|
||||
# To avoid corrupting reads from the same SGPR, use a shared output Reg(0).
|
||||
# Exception: CMPX instructions write to EXEC (not D0/VCC).
|
||||
d0_override, vcc_override = None, None
|
||||
vopc_dst, vop3sd_dst = None, None
|
||||
is_cmpx = False
|
||||
if inst_type is VOPC:
|
||||
op = VOPCOp(inst.op)
|
||||
is_cmpx = 'CMPX' in op.name
|
||||
if not is_cmpx: # Regular CMP writes to VCC
|
||||
d0_override, vopc_dst = Reg(0), VCC_LO
|
||||
else: # CMPX writes to EXEC - clear it first, accumulate per-lane
|
||||
st.sgpr[EXEC_LO]._val = 0
|
||||
elif inst_type is VOP3 and inst.op < 256: # VOP3-encoded VOPC
|
||||
op = VOPCOp(inst.op)
|
||||
is_cmpx = 'CMPX' in op.name
|
||||
if not is_cmpx: # Regular CMP writes to destination SGPR
|
||||
d0_override, vopc_dst = Reg(0), inst.vdst
|
||||
else: # CMPX writes to EXEC - clear it first, accumulate per-lane
|
||||
st.sgpr[EXEC_LO]._val = 0
|
||||
if inst_type is VOP3SD:
|
||||
vcc_override, vop3sd_dst = Reg(0), inst.sdst
|
||||
|
||||
# For other vector ops, dispatch to exec_vector per lane (can optimize later)
|
||||
for lane in range(n_lanes):
|
||||
if exec_mask & (1 << lane): exec_vector(st, inst, lane, lds, d0_override, vcc_override)
|
||||
|
||||
# Write accumulated per-lane bit results to destination SGPRs
|
||||
# (CMPX writes directly to EXEC in the pseudocode, so no separate write needed)
|
||||
if vopc_dst is not None: st.sgpr[vopc_dst]._val = d0_override._val
|
||||
if vop3sd_dst is not None: st.sgpr[vop3sd_dst]._val = vcc_override._val
|
||||
|
||||
def step_wave(program: Program, st: WaveState, lds: bytearray, n_lanes: int) -> int:
|
||||
inst = program.get(st.pc)
|
||||
if inst is None: return 1
|
||||
@@ -708,7 +696,9 @@ def step_wave(program: Program, st: WaveState, lds: bytearray, n_lanes: int) ->
|
||||
if is_readlane:
|
||||
exec_vector(st, inst, 0, lds) # Execute once with lane 0
|
||||
else:
|
||||
exec_vector_batch(st, inst, st.exec_mask, n_lanes, lds)
|
||||
exec_mask = st.exec_mask
|
||||
for lane in range(n_lanes):
|
||||
if exec_mask & (1 << lane): exec_vector(st, inst, lane, lds)
|
||||
st.commit_pends()
|
||||
st.pc += inst_words
|
||||
return 0
|
||||
@@ -732,12 +722,12 @@ def exec_workgroup(program: Program, workgroup_id: tuple[int, int, int], local_s
|
||||
gx, gy, gz = workgroup_id
|
||||
# Set workgroup IDs in SGPRs based on USER_SGPR_COUNT and enable flags from COMPUTE_PGM_RSRC2
|
||||
sgpr_idx = wg_id_sgpr_base
|
||||
if wg_id_enables[0]: st.sgpr[sgpr_idx]._val = gx; sgpr_idx += 1
|
||||
if wg_id_enables[1]: st.sgpr[sgpr_idx]._val = gy; sgpr_idx += 1
|
||||
if wg_id_enables[2]: st.sgpr[sgpr_idx]._val = gz
|
||||
if wg_id_enables[0]: st.sgpr[sgpr_idx] = gx; sgpr_idx += 1
|
||||
if wg_id_enables[1]: st.sgpr[sgpr_idx] = gy; sgpr_idx += 1
|
||||
if wg_id_enables[2]: st.sgpr[sgpr_idx] = gz
|
||||
for i in range(n_lanes):
|
||||
tid = wave_start + i
|
||||
st.vgpr[i][0]._val = tid if local_size == (lx, 1, 1) else ((tid // (lx * ly)) << 20) | (((tid // lx) % ly) << 10) | (tid % lx)
|
||||
st.vgpr[i][0] = tid if local_size == (lx, 1, 1) else ((tid // (lx * ly)) << 20) | (((tid // lx) % ly) << 10) | (tid % lx)
|
||||
waves.append((st, n_lanes, wave_start))
|
||||
has_barrier = any(isinstance(inst, SOPP) and inst.op == SOPPOp.S_BARRIER for inst in program.values())
|
||||
for _ in range(2 if has_barrier else 1):
|
||||
|
||||
+95
-88
@@ -564,8 +564,19 @@ class Reg:
|
||||
def __ge__(s, o): return s._val >= int(o)
|
||||
def __eq__(s, o): return s._val == int(o)
|
||||
def __ne__(s, o): return s._val != int(o)
|
||||
def __format__(s, spec): return format(s._val, spec)
|
||||
def __repr__(s): return f"Reg(0x{s._val:x})"
|
||||
|
||||
class ExecContext:
|
||||
"""Execution context for running compiled pseudocode strings (for testing)."""
|
||||
def __init__(self, s0=0, s1=0, s2=0, d0=0, scc=0, vcc=0, lane=0, exec_mask=0xffffffff):
|
||||
self.S0, self.S1, self.S2, self.D0, self.D1 = Reg(s0), Reg(s1), Reg(s2), Reg(d0), Reg(0)
|
||||
self.SCC, self.VCC, self.EXEC = Reg(scc), Reg(vcc), Reg(exec_mask)
|
||||
self.tmp, self.saveexec, self.laneId = Reg(0), Reg(exec_mask), lane
|
||||
def run(self, code: str):
|
||||
if not code.strip(): return
|
||||
code = code if '\n' in code else compile_pseudocode(code)
|
||||
exec(code, {**globals(), 'Reg': Reg, 'SliceProxy': SliceProxy}, self.__dict__)
|
||||
def result(self) -> dict:
|
||||
return {'d0': self.D0._val, 'd1': self.D1._val, 'scc': self.SCC._val & 1, 'vcc': self.VCC._val, 'exec': self.EXEC._val}
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# COMPILER: pseudocode -> Python (minimal transforms)
|
||||
@@ -646,13 +657,14 @@ def compile_pseudocode(pseudocode: str) -> str:
|
||||
return '\n'.join(lines)
|
||||
|
||||
def _assign(lhs: str, rhs: str) -> str:
|
||||
"""Generate assignment. Bare tmp/SCC/etc get wrapped in Reg(). For params (SCC, VCC, EXEC, D0), modify in place."""
|
||||
# Parameters passed to function - modify in place using .b32 setter (or .b64 for 64-bit types)
|
||||
if lhs in ('SCC', 'VCC', 'EXEC', 'D0'):
|
||||
return f"{lhs}.b32 = {rhs}"
|
||||
# Local variables - create new Reg
|
||||
if lhs in ('tmp', 'D1', 'saveexec'):
|
||||
return f"{lhs} = Reg({rhs})"
|
||||
"""Generate assignment. Outputs modify Reg in-place via ._val."""
|
||||
# Output registers and tmp: modify in-place so caller sees changes
|
||||
if lhs in ('SCC', 'VCC', 'EXEC', 'D0', 'D1', 'tmp'):
|
||||
return f"{lhs}._val = int({rhs})"
|
||||
# saveexec needs to be a new Reg for typed accessor access
|
||||
if lhs == 'saveexec':
|
||||
return f"{lhs} = Reg(int({rhs}))"
|
||||
# Other locals: natural style
|
||||
return f"{lhs} = {rhs}"
|
||||
|
||||
def _expr(e: str) -> str:
|
||||
@@ -734,50 +746,7 @@ def _expr(e: str) -> str:
|
||||
e = f'(({t}) if ({cond}) else ({f}))'
|
||||
return e
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# EXECUTION CONTEXT
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
class ExecContext:
|
||||
"""Context for running compiled pseudocode."""
|
||||
def __init__(self, s0=0, s1=0, s2=0, d0=0, scc=0, vcc=0, lane=0, exec_mask=MASK32, literal=0, vgprs=None, src0_idx=0, vdst_idx=0):
|
||||
self.S0, self.S1, self.S2 = Reg(s0), Reg(s1), Reg(s2)
|
||||
self.D0, self.D1 = Reg(d0), Reg(0)
|
||||
self.SCC, self.VCC, self.EXEC = Reg(scc), Reg(vcc), Reg(exec_mask)
|
||||
self.tmp, self.saveexec = Reg(0), Reg(exec_mask)
|
||||
self.lane, self.laneId, self.literal = lane, lane, literal
|
||||
self.SIMM16, self.SIMM32 = Reg(literal), Reg(literal)
|
||||
self.VGPR = vgprs if vgprs is not None else {}
|
||||
self.SRC0, self.VDST = Reg(src0_idx), Reg(vdst_idx)
|
||||
|
||||
def run(self, code: str):
|
||||
"""Execute compiled code."""
|
||||
# Start with module globals (helpers, aliases), then add instance-specific bindings
|
||||
ns = dict(globals())
|
||||
ns.update({
|
||||
'S0': self.S0, 'S1': self.S1, 'S2': self.S2, 'D0': self.D0, 'D1': self.D1,
|
||||
'SCC': self.SCC, 'VCC': self.VCC, 'EXEC': self.EXEC,
|
||||
'EXEC_LO': SliceProxy(self.EXEC, 31, 0), 'EXEC_HI': SliceProxy(self.EXEC, 63, 32),
|
||||
'tmp': self.tmp, 'saveexec': self.saveexec,
|
||||
'lane': self.lane, 'laneId': self.laneId, 'literal': self.literal,
|
||||
'SIMM16': self.SIMM16, 'SIMM32': self.SIMM32,
|
||||
'VGPR': self.VGPR, 'SRC0': self.SRC0, 'VDST': self.VDST,
|
||||
})
|
||||
exec(code, ns)
|
||||
# Sync rebinds: if register was reassigned to new Reg or value, copy it back
|
||||
def _sync(ctx_reg, ns_val):
|
||||
if isinstance(ns_val, Reg): ctx_reg._val = ns_val._val
|
||||
else: ctx_reg._val = int(ns_val) & MASK64
|
||||
if ns.get('SCC') is not self.SCC: _sync(self.SCC, ns['SCC'])
|
||||
if ns.get('VCC') is not self.VCC: _sync(self.VCC, ns['VCC'])
|
||||
if ns.get('EXEC') is not self.EXEC: _sync(self.EXEC, ns['EXEC'])
|
||||
if ns.get('D0') is not self.D0: _sync(self.D0, ns['D0'])
|
||||
if ns.get('D1') is not self.D1: _sync(self.D1, ns['D1'])
|
||||
if ns.get('tmp') is not self.tmp: _sync(self.tmp, ns['tmp'])
|
||||
if ns.get('saveexec') is not self.saveexec: _sync(self.saveexec, ns['saveexec'])
|
||||
|
||||
def result(self) -> dict:
|
||||
return {"d0": self.D0._val, "scc": self.SCC._val & 1}
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════════
|
||||
# PDF EXTRACTION AND CODE GENERATION
|
||||
@@ -954,47 +923,86 @@ from extra.assembly.amd.pcode import *
|
||||
# CLZ/CTZ: The PDF pseudocode searches for the first 1 bit but doesn't break.
|
||||
# Hardware stops at first match. SOP1 uses tmp=i, VOP1/VOP3 use D0.i32=i
|
||||
if 'CLZ' in op.name or 'CTZ' in op.name:
|
||||
code = code.replace('tmp = Reg(i)', 'tmp._val = i; break')
|
||||
code = code.replace('tmp = Reg(i)', 'tmp = Reg(i); break')
|
||||
code = code.replace('D0.i32 = i', 'D0.i32 = i; break')
|
||||
# Detect flags for result handling
|
||||
is_64 = any(p in pc for p in ['D0.u64', 'D0.b64', 'D0.f64', 'D0.i64', 'D1.u64', 'D1.b64', 'D1.f64', 'D1.i64'])
|
||||
has_d1 = '{ D1' in pc
|
||||
if has_d1: is_64 = True
|
||||
is_cmp = cls_name == 'VOPCOp' and 'D0.u64[laneId]' in pc
|
||||
is_cmpx = cls_name == 'VOPCOp' and 'EXEC.u64[laneId]' in pc # V_CMPX writes to EXEC per-lane
|
||||
# V_DIV_SCALE passes through S0 if no branch taken
|
||||
is_div_scale = 'DIV_SCALE' in op.name
|
||||
# VOP3SD instructions that write VCC per-lane (either via VCC.u64[laneId] or by setting VCC = 0/1)
|
||||
has_sdst = cls_name == 'VOP3SDOp' and ('VCC.u64[laneId]' in pc or is_div_scale)
|
||||
|
||||
# Generate function that takes Reg objects directly - modifies D0 in place
|
||||
# V_DIV_FMAS_F32/F64: PDF page 449 says 2^32/2^64 but hardware behavior is more complex.
|
||||
# The scale direction depends on S2 (the addend): if exponent(S2) > 127 (i.e., S2 >= 2.0),
|
||||
# scale by 2^+64 (to unscale a numerator that was scaled). Otherwise scale by 2^-64
|
||||
# (to unscale a denominator that was scaled).
|
||||
if op.name == 'V_DIV_FMAS_F32':
|
||||
code = code.replace(
|
||||
'D0.f32 = 2.0 ** 32 * fma(S0.f32, S1.f32, S2.f32)',
|
||||
'D0.f32 = (2.0 ** 64 if exponent(S2.f32) > 127 else 2.0 ** -64) * fma(S0.f32, S1.f32, S2.f32)')
|
||||
if op.name == 'V_DIV_FMAS_F64':
|
||||
code = code.replace(
|
||||
'D0.f64 = 2.0 ** 64 * fma(S0.f64, S1.f64, S2.f64)',
|
||||
'D0.f64 = (2.0 ** 128 if exponent(S2.f64) > 1023 else 2.0 ** -128) * fma(S0.f64, S1.f64, S2.f64)')
|
||||
# V_DIV_SCALE_F32/F64: PDF page 463-464 has several bugs vs hardware behavior:
|
||||
# 1. Zero case: hardware sets VCC=1 (PDF doesn't)
|
||||
# 2. Denorm denom: hardware returns NaN (PDF says scale). VCC is set independently by exp diff check.
|
||||
# 3. Tiny numer (exp<=23): hardware sets VCC=1 (PDF doesn't)
|
||||
# 4. Result would be denorm: hardware doesn't scale, just sets VCC=1
|
||||
if op.name == 'V_DIV_SCALE_F32':
|
||||
# Fix 1: Set VCC=1 when zero operands produce NaN
|
||||
code = code.replace(
|
||||
'D0.f32 = float("nan")',
|
||||
'VCC._val = 0x1; D0.f32 = float("nan")')
|
||||
# Fix 2: Denorm denom returns NaN. Must check this AFTER all VCC-setting logic runs.
|
||||
# Insert at end of all branches, before the final result is used
|
||||
code = code.replace(
|
||||
'elif S1.f32 == DENORM.f32:\n D0.f32 = ldexp(S0.f32, 64)',
|
||||
'elif False:\n pass # denorm check moved to end')
|
||||
# Add denorm check at the very end - this overrides D0 but preserves VCC
|
||||
code += '\nif S1.f32 == DENORM.f32:\n D0.f32 = float("nan")'
|
||||
# Fix 3: Tiny numer should set VCC=1
|
||||
code = code.replace(
|
||||
'elif exponent(S2.f32) <= 23:\n D0.f32 = ldexp(S0.f32, 64)',
|
||||
'elif exponent(S2.f32) <= 23:\n VCC._val = 0x1; D0.f32 = ldexp(S0.f32, 64)')
|
||||
# Fix 4: S2/S1 would be denorm - don't scale, just set VCC
|
||||
code = code.replace(
|
||||
'elif S2.f32 / S1.f32 == DENORM.f32:\n VCC._val = int(0x1)\n if S0.f32 == S2.f32:\n D0.f32 = ldexp(S0.f32, 64)',
|
||||
'elif S2.f32 / S1.f32 == DENORM.f32:\n VCC._val = 0x1')
|
||||
if op.name == 'V_DIV_SCALE_F64':
|
||||
# Same fixes for f64 version
|
||||
code = code.replace(
|
||||
'D0.f64 = float("nan")',
|
||||
'VCC._val = 0x1; D0.f64 = float("nan")')
|
||||
code = code.replace(
|
||||
'elif S1.f64 == DENORM.f64:\n D0.f64 = ldexp(S0.f64, 128)',
|
||||
'elif False:\n pass # denorm check moved to end')
|
||||
code += '\nif S1.f64 == DENORM.f64:\n D0.f64 = float("nan")'
|
||||
code = code.replace(
|
||||
'elif exponent(S2.f64) <= 52:\n D0.f64 = ldexp(S0.f64, 128)',
|
||||
'elif exponent(S2.f64) <= 52:\n VCC._val = 0x1; D0.f64 = ldexp(S0.f64, 128)')
|
||||
code = code.replace(
|
||||
'elif S2.f64 / S1.f64 == DENORM.f64:\n VCC._val = int(0x1)\n if S0.f64 == S2.f64:\n D0.f64 = ldexp(S0.f64, 128)',
|
||||
'elif S2.f64 / S1.f64 == DENORM.f64:\n VCC._val = 0x1')
|
||||
# V_DIV_FIXUP_F32/F64: PDF doesn't check isNAN(S0), but hardware returns OVERFLOW if S0 is NaN.
|
||||
# When division fails (e.g., due to denorm denom), S0 becomes NaN, and fixup should return ±inf.
|
||||
if op.name == 'V_DIV_FIXUP_F32':
|
||||
code = code.replace(
|
||||
'D0.f32 = ((-abs(S0.f32)) if (sign_out) else (abs(S0.f32)))',
|
||||
'D0.f32 = ((-OVERFLOW_F32) if (sign_out) else (OVERFLOW_F32)) if isNAN(S0.f32) else ((-abs(S0.f32)) if (sign_out) else (abs(S0.f32)))')
|
||||
if op.name == 'V_DIV_FIXUP_F64':
|
||||
code = code.replace(
|
||||
'D0.f64 = ((-abs(S0.f64)) if (sign_out) else (abs(S0.f64)))',
|
||||
'D0.f64 = ((-OVERFLOW_F64) if (sign_out) else (OVERFLOW_F64)) if isNAN(S0.f64) else ((-abs(S0.f64)) if (sign_out) else (abs(S0.f64)))')
|
||||
# SIMM16/SIMM32 (inline literal constants) are passed as S2
|
||||
code = code.replace('SIMM16', 'S2').replace('SIMM32', 'S2')
|
||||
# Generate function with standard signature
|
||||
fn_name = f"_{cls_name}_{op.name}"
|
||||
lines.append(f"def {fn_name}(S0, S1, S2, D0, SCC, VCC, laneId, EXEC, SIMM16, VGPR, SRC0, VDST):")
|
||||
# Add original pseudocode as comment
|
||||
lines.append(f"def {fn_name}(S0, S1, S2, D0, D1, SCC, VCC, EXEC, tmp, laneId):")
|
||||
for pc_line in pc.split('\n'):
|
||||
lines.append(f" # {pc_line}")
|
||||
# Only create extra Reg objects for registers that need fresh state
|
||||
# Add EXEC_LO/EXEC_HI if needed
|
||||
combined = code + pc
|
||||
# D1 and tmp/saveexec need to be created fresh
|
||||
if 'D1' in combined: lines.append(" D1 = Reg(0)")
|
||||
if 'tmp' in combined: lines.append(" tmp = Reg(0)")
|
||||
if 'saveexec' in combined: lines.append(" saveexec = Reg(EXEC._val)")
|
||||
if 'SIMM32' in combined: lines.append(" SIMM32 = SIMM16")
|
||||
if 'EXEC_LO' in combined: lines.append(" EXEC_LO = SliceProxy(EXEC, 31, 0)")
|
||||
if 'EXEC_HI' in combined: lines.append(" EXEC_HI = SliceProxy(EXEC, 63, 32)")
|
||||
# For DIV_SCALE, D0 starts with S0's value
|
||||
if is_div_scale: lines.append(" D0._val = S0._val")
|
||||
# Add compiled pseudocode
|
||||
lines.append(" # --- compiled pseudocode ---")
|
||||
has_code = False
|
||||
for line in code.split('\n'):
|
||||
if line.strip():
|
||||
code_lines = [line for line in code.split('\n') if line.strip()]
|
||||
if code_lines:
|
||||
for line in code_lines:
|
||||
lines.append(f" {line}")
|
||||
has_code = True
|
||||
lines.append(" # --- end pseudocode ---")
|
||||
# All Reg objects (D0, SCC, VCC, EXEC) are modified in place
|
||||
# The emulator determines 64-bit ops from the opcode name
|
||||
if not has_code:
|
||||
else:
|
||||
lines.append(" pass")
|
||||
lines.append("")
|
||||
|
||||
@@ -1016,9 +1024,8 @@ from extra.assembly.amd.pcode import *
|
||||
if 'VOP3Op' in enum_names:
|
||||
lines.append('''
|
||||
# V_WRITELANE_B32: Write scalar to specific lane's VGPR (not in PDF pseudocode)
|
||||
def _VOP3Op_V_WRITELANE_B32(S0, S1, S2, D0, SCC, VCC, laneId, EXEC, SIMM16, VGPR, SRC0, VDST):
|
||||
wr_lane = S1._val & 0x1f # lane select (5 bits for wave32)
|
||||
return {'vgpr_write': (wr_lane, VDST._val, S0._val & 0xffffffff)}
|
||||
def _VOP3Op_V_WRITELANE_B32(S0, S1, S2, D0, D1, SCC, VCC, EXEC, tmp, laneId):
|
||||
return (int(S1) & 0x1f, int(S0) & 0xffffffff) # (wr_lane, value)
|
||||
VOP3Op_FUNCTIONS[VOP3Op.V_WRITELANE_B32] = _VOP3Op_V_WRITELANE_B32
|
||||
''')
|
||||
|
||||
|
||||
@@ -107,16 +107,16 @@ class PythonEmulator:
|
||||
return step_wave(self.program, self.state, self.lds, self.n_lanes)
|
||||
def set_sgpr(self, idx: int, val: int):
|
||||
assert self.state is not None
|
||||
self.state.sgpr[idx]._val = val & 0xffffffff
|
||||
self.state.sgpr[idx] = val & 0xffffffff
|
||||
def set_vgpr(self, lane: int, idx: int, val: int):
|
||||
assert self.state is not None
|
||||
self.state.vgpr[lane][idx]._val = val & 0xffffffff
|
||||
self.state.vgpr[lane][idx] = val & 0xffffffff
|
||||
|
||||
def get_snapshot(self) -> StateSnapshot:
|
||||
assert self.state is not None
|
||||
return StateSnapshot(pc=self.state.pc, scc=self.state.scc, vcc=self.state.vcc & 0xffffffff,
|
||||
exec_mask=self.state.exec_mask & 0xffffffff, sgpr=[r._val for r in self.state.sgpr],
|
||||
vgpr=[[r._val for r in self.state.vgpr[i]] for i in range(WAVE_SIZE)])
|
||||
exec_mask=self.state.exec_mask & 0xffffffff, sgpr=list(self.state.sgpr),
|
||||
vgpr=[list(self.state.vgpr[i]) for i in range(WAVE_SIZE)])
|
||||
|
||||
def run_single_kernel(kernel: bytes, n_lanes: int, args_ptr: int, global_size: tuple[int, int, int],
|
||||
program, max_steps: int, debug: bool, trace_len: int, kernel_idx: int = 0,
|
||||
@@ -191,6 +191,9 @@ def run_single_kernel(kernel: bytes, n_lanes: int, args_ptr: int, global_size: t
|
||||
python_result = python.step()
|
||||
|
||||
if rust_result != python_result:
|
||||
# Rust returns 1 for unsupported instructions - skip test
|
||||
if rust_result == 1 and python_result == 0:
|
||||
raise unittest.SkipTest(f"Rust emulator doesn't support instruction: {inst_str}")
|
||||
trace_str = "\n".join(f" step {s}: PC={pc:3d} {d}" for s, pc, d, _, _ in trace)
|
||||
return False, f"K{kernel_idx} WG({gidx},{gidy},{gidz}) Step {step}: different return codes: rust={rust_result}, python={python_result}, inst={inst_str}\n Recent instructions:\n{trace_str}", total_steps
|
||||
|
||||
@@ -361,6 +364,7 @@ class TestTinygradKernels(unittest.TestCase):
|
||||
|
||||
# Matmul
|
||||
def test_gemm(self): self._test_kernel(lambda T: T.empty(8, 8) @ T.empty(8, 8), max_steps=100000)
|
||||
@unittest.skip("Rust emulator crashes on this kernel (assertion failure in thread.rs)")
|
||||
def test_gemm_fp16(self): self._test_kernel(lambda T: T.empty(16, 16).half() @ T.empty(16, 16).half(), max_steps=100000)
|
||||
|
||||
# Complex ops
|
||||
|
||||
+1252
-46
File diff suppressed because it is too large
Load Diff
@@ -65,12 +65,18 @@ def parse_llvm_tests(text: str) -> list[tuple[str, bytes]]:
|
||||
if not asm_text: continue
|
||||
for j in range(i, min(i + 3, len(lines))):
|
||||
# Match GFX11, W32, or W64 encodings (all valid for gfx11)
|
||||
# Format 1: "// GFX11: v_foo ... ; encoding: [0x01,0x02,...]"
|
||||
# Format 2: "// GFX11: [0x01,0x02,...]" (used by DS, older files)
|
||||
if m := re.search(r'(?:GFX11|W32|W64)[^:]*:.*?encoding:\s*\[(.*?)\]', lines[j]):
|
||||
hex_bytes = m.group(1).replace('0x', '').replace(',', '').replace(' ', '')
|
||||
if hex_bytes:
|
||||
try: tests.append((asm_text, bytes.fromhex(hex_bytes)))
|
||||
except ValueError: pass
|
||||
break
|
||||
elif m := re.search(r'(?:GFX11|W32|W64)[^:]*:\s*\[(0x[0-9a-fA-F,x\s]+)\]', lines[j]):
|
||||
hex_bytes = m.group(1).replace('0x', '').replace(',', '').replace(' ', '')
|
||||
else:
|
||||
continue
|
||||
if hex_bytes:
|
||||
try: tests.append((asm_text, bytes.fromhex(hex_bytes)))
|
||||
except ValueError: pass
|
||||
break
|
||||
return tests
|
||||
|
||||
def try_assemble(text: str):
|
||||
|
||||
@@ -4,7 +4,9 @@ import unittest
|
||||
from extra.assembly.amd.pcode import (Reg, TypedView, SliceProxy, ExecContext, compile_pseudocode, _expr, MASK32, MASK64,
|
||||
_f32, _i32, _f16, _i16, f32_to_f16, _isnan, _bf16, _ibf16, bf16_to_f32, f32_to_bf16,
|
||||
BYTE_PERMUTE, v_sad_u8, v_msad_u8)
|
||||
from extra.assembly.amd.autogen.rdna3.gen_pcode import _VOP3SDOp_V_DIV_SCALE_F32, _VOPCOp_V_CMP_CLASS_F32
|
||||
from extra.assembly.amd.autogen.rdna3.gen_pcode import VOP3SDOp_FUNCTIONS, VOPCOp_FUNCTIONS
|
||||
from extra.assembly.amd.autogen.rdna3 import VOP3SDOp, VOPCOp
|
||||
from extra.assembly.amd.emu import _run_pcode
|
||||
|
||||
class TestReg(unittest.TestCase):
|
||||
def test_u32_read(self):
|
||||
@@ -208,6 +210,8 @@ D0.u32 = tmp.u32""")
|
||||
for i in 0 : 31 do
|
||||
if S0.u32[i] == 1 then
|
||||
tmp = i
|
||||
endif
|
||||
endfor
|
||||
D0.i32 = tmp""")
|
||||
ctx = ExecContext(s0=0b1000) # Bit 3 is set
|
||||
ctx.run(code)
|
||||
@@ -225,41 +229,39 @@ class TestPseudocodeRegressions(unittest.TestCase):
|
||||
"""Regression tests for pseudocode instruction emulation bugs."""
|
||||
|
||||
def test_v_div_scale_f32_vcc_always_returned(self):
|
||||
"""V_DIV_SCALE_F32 must set VCC bit for the lane when scaling is needed.
|
||||
The new calling convention uses Reg objects and modifies VCC in place."""
|
||||
"""V_DIV_SCALE_F32 must always return vcc_lane, even when VCC=0 (no scaling needed).
|
||||
Bug: when VCC._val == vcc (both 0), vcc_lane wasn't returned, so VCC bits weren't written.
|
||||
This caused division to produce wrong results for multiple lanes."""
|
||||
# Normal case: 1.0 / 3.0, no scaling needed, VCC should be 0
|
||||
S0 = Reg(0x3f800000) # 1.0
|
||||
S1 = Reg(0x40400000) # 3.0
|
||||
S2 = Reg(0x3f800000) # 1.0 (numerator)
|
||||
D0 = Reg(0)
|
||||
VCC = Reg(0)
|
||||
_VOP3SDOp_V_DIV_SCALE_F32(S0, S1, S2, D0, Reg(0), VCC, 0, Reg(0xffffffff), Reg(0), None, Reg(0), Reg(0))
|
||||
# VCC bit 0 should be 0 when no scaling needed
|
||||
self.assertEqual(VCC._val & 1, 0, "VCC bit should be 0 when no scaling needed")
|
||||
s0 = 0x3f800000 # 1.0
|
||||
s1 = 0x40400000 # 3.0
|
||||
s2 = 0x3f800000 # 1.0 (numerator)
|
||||
fn = VOP3SDOp_FUNCTIONS[VOP3SDOp.V_DIV_SCALE_F32]
|
||||
result = _run_pcode(fn, VOP3SDOp, VOP3SDOp.V_DIV_SCALE_F32, s0, s1, s2, 0, 0, 0, 0, 0xffffffff, 0)
|
||||
# Must always have vcc_lane in result
|
||||
self.assertIn('vcc_lane', result, "V_DIV_SCALE_F32 must always return vcc_lane")
|
||||
self.assertEqual(result['vcc_lane'], 0, "vcc_lane should be 0 when no scaling needed")
|
||||
|
||||
def test_v_cmp_class_f32_detects_quiet_nan(self):
|
||||
"""V_CMP_CLASS_F32 must correctly identify quiet NaN vs signaling NaN.
|
||||
Bug: isQuietNAN and isSignalNAN both used math.isnan which can't distinguish them."""
|
||||
quiet_nan = 0x7fc00000 # quiet NaN: exponent=255, bit22=1
|
||||
signal_nan = 0x7f800001 # signaling NaN: exponent=255, bit22=0
|
||||
fn = VOPCOp_FUNCTIONS[VOPCOp.V_CMP_CLASS_F32]
|
||||
# Test quiet NaN detection (bit 1 in mask)
|
||||
s1_quiet = 0b0000000010 # bit 1 = quiet NaN
|
||||
D0 = Reg(0)
|
||||
_VOPCOp_V_CMP_CLASS_F32(Reg(quiet_nan), Reg(s1_quiet), Reg(0), D0, Reg(0), Reg(0), 0, Reg(0xffffffff), Reg(0), None, Reg(0), Reg(0))
|
||||
self.assertEqual(D0._val & 1, 1, "Should detect quiet NaN with quiet NaN mask")
|
||||
result = _run_pcode(fn, VOPCOp, VOPCOp.V_CMP_CLASS_F32, quiet_nan, s1_quiet, 0, 0, 0, 0, 0, 0xffffffff, 0)
|
||||
self.assertEqual(result['vcc_lane'], 1, "Should detect quiet NaN with quiet NaN mask")
|
||||
# Test signaling NaN detection (bit 0 in mask)
|
||||
s1_signal = 0b0000000001 # bit 0 = signaling NaN
|
||||
D0 = Reg(0)
|
||||
_VOPCOp_V_CMP_CLASS_F32(Reg(signal_nan), Reg(s1_signal), Reg(0), D0, Reg(0), Reg(0), 0, Reg(0xffffffff), Reg(0), None, Reg(0), Reg(0))
|
||||
self.assertEqual(D0._val & 1, 1, "Should detect signaling NaN with signaling NaN mask")
|
||||
result = _run_pcode(fn, VOPCOp, VOPCOp.V_CMP_CLASS_F32, signal_nan, s1_signal, 0, 0, 0, 0, 0, 0xffffffff, 0)
|
||||
self.assertEqual(result['vcc_lane'], 1, "Should detect signaling NaN with signaling NaN mask")
|
||||
# Test that quiet NaN doesn't match signaling NaN mask
|
||||
D0 = Reg(0)
|
||||
_VOPCOp_V_CMP_CLASS_F32(Reg(quiet_nan), Reg(s1_signal), Reg(0), D0, Reg(0), Reg(0), 0, Reg(0xffffffff), Reg(0), None, Reg(0), Reg(0))
|
||||
self.assertEqual(D0._val & 1, 0, "Quiet NaN should not match signaling NaN mask")
|
||||
result = _run_pcode(fn, VOPCOp, VOPCOp.V_CMP_CLASS_F32, quiet_nan, s1_signal, 0, 0, 0, 0, 0, 0xffffffff, 0)
|
||||
self.assertEqual(result['vcc_lane'], 0, "Quiet NaN should not match signaling NaN mask")
|
||||
# Test that signaling NaN doesn't match quiet NaN mask
|
||||
D0 = Reg(0)
|
||||
_VOPCOp_V_CMP_CLASS_F32(Reg(signal_nan), Reg(s1_quiet), Reg(0), D0, Reg(0), Reg(0), 0, Reg(0xffffffff), Reg(0), None, Reg(0), Reg(0))
|
||||
self.assertEqual(D0._val & 1, 0, "Signaling NaN should not match quiet NaN mask")
|
||||
result = _run_pcode(fn, VOPCOp, VOPCOp.V_CMP_CLASS_F32, signal_nan, s1_quiet, 0, 0, 0, 0, 0, 0xffffffff, 0)
|
||||
self.assertEqual(result['vcc_lane'], 0, "Signaling NaN should not match quiet NaN mask")
|
||||
|
||||
def test_isnan_with_typed_view(self):
|
||||
"""_isnan must work with TypedView objects, not just Python floats.
|
||||
|
||||
@@ -4,46 +4,9 @@ import unittest, io, sys, re, subprocess, os
|
||||
from extra.assembly.amd.autogen.rdna3 import *
|
||||
from extra.assembly.amd.dsl import Inst
|
||||
from extra.assembly.amd.asm import asm
|
||||
from extra.assembly.amd.asm import detect_format
|
||||
from extra.assembly.amd.test.helpers import get_llvm_mc, get_llvm_objdump
|
||||
|
||||
# Instruction format detection based on encoding bits
|
||||
def detect_format(data: bytes) -> type[Inst] | None:
|
||||
"""Detect instruction format from machine code bytes."""
|
||||
if len(data) < 4: return None
|
||||
word = int.from_bytes(data[:4], 'little')
|
||||
enc_9bit = (word >> 23) & 0x1FF # 9-bit encoding for SOP1/SOPC/SOPP
|
||||
enc_8bit = (word >> 24) & 0xFF
|
||||
|
||||
# Check 9-bit encodings first (most specific)
|
||||
if enc_9bit == 0x17D: return SOP1 # bits 31:23 = 101111101
|
||||
if enc_9bit == 0x17E: return SOPC # bits 31:23 = 101111110
|
||||
if enc_9bit == 0x17F: return SOPP # bits 31:23 = 101111111
|
||||
# SOPK: bits 31:28 = 1011, bits 27:23 = opcode (check after SOP1/SOPC/SOPP)
|
||||
if enc_8bit in range(0xB0, 0xC0): return SOPK
|
||||
# SOP2: bits 31:23 in range 0x100-0x17C (0x80-0xBE in bits 31:24, but not SOPK)
|
||||
if 0x80 <= enc_8bit <= 0x9F: return SOP2
|
||||
# VOP1: bits 31:25 = 0111111 (0x3F)
|
||||
if (word >> 25) == 0x3F: return VOP1
|
||||
# VOPC: bits 31:25 = 0111110 (0x3E)
|
||||
if (word >> 25) == 0x3E: return VOPC
|
||||
# VOP2: bits 31:30 = 00
|
||||
if (word >> 30) == 0: return VOP2
|
||||
|
||||
# Check 64-bit formats
|
||||
if len(data) >= 8:
|
||||
if enc_8bit in (0xD4, 0xD5, 0xD7): return VOP3
|
||||
if enc_8bit == 0xD6: return VOP3SD
|
||||
if enc_8bit == 0xCC: return VOP3P
|
||||
if enc_8bit == 0xCD: return VINTERP
|
||||
if enc_8bit in (0xC8, 0xC9): return VOPD
|
||||
if enc_8bit == 0xF4: return SMEM
|
||||
if enc_8bit == 0xD8: return DS
|
||||
if enc_8bit in (0xDC, 0xDD, 0xDE, 0xDF): return FLAT
|
||||
if enc_8bit in (0xE0, 0xE1, 0xE2, 0xE3): return MUBUF
|
||||
if enc_8bit in (0xE8, 0xE9, 0xEA, 0xEB): return MTBUF
|
||||
|
||||
return None
|
||||
|
||||
def disassemble_lib(lib: bytes, compiler) -> list[tuple[str, bytes]]:
|
||||
"""Disassemble ELF binary and return list of (instruction_text, machine_code_bytes)."""
|
||||
old_stdout = sys.stdout
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,7 @@
|
||||
# Run assembly on the AMD runtime and check correctness
|
||||
# VIZ=2 to profile
|
||||
import pathlib
|
||||
from tinygrad import Tensor, Device, dtypes
|
||||
from tinygrad import Tensor, Device, dtypes, Context
|
||||
from tinygrad.engine.realize import ExecItem, CompiledRunner
|
||||
from tinygrad.renderer import ProgramSpec
|
||||
from tinygrad.uop.ops import track_rewrites, UOp
|
||||
@@ -55,9 +55,10 @@ def get_asm_prg() -> ProgramSpec:
|
||||
eis.append(ExecItem(ast, [C_asm.uop.buffer, from_torch(B).uop.buffer, from_torch(A).uop.buffer], fixedvars={"SZ":N, "NUM_WG":NUM_WG},
|
||||
prg=CompiledRunner(get_asm_prg())))
|
||||
|
||||
for ei in eis:
|
||||
et = ei.run(wait=True)
|
||||
print(f"{(N*N*N*2 / et)*1e-12:.2f} REAL TFLOPS")
|
||||
with Context(DEBUG=2):
|
||||
for ei in eis:
|
||||
et = ei.run(wait=True)
|
||||
print(f"{(N*N*N*2 / et)*1e-12:.2f} REAL TFLOPS")
|
||||
|
||||
# ** correctness
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import unittest, operator, math
|
||||
from tinygrad import Tensor, dtypes, Device
|
||||
from tinygrad.dtype import DType, truncate
|
||||
from tinygrad.helpers import CI, getenv, CPU_LLVM
|
||||
from tinygrad.helpers import CI, getenv
|
||||
from tinygrad.tensor import _to_np_dtype
|
||||
from tinygrad.device import is_dtype_supported
|
||||
from tinygrad.runtime.ops_python import from_storage_scalar
|
||||
@@ -48,7 +48,7 @@ class ht:
|
||||
int32 = strat.integers(-2147483648, 2147483647)
|
||||
int64 = strat.integers(-9223372036854775808, 9223372036854775807)
|
||||
bool = strat.booleans()
|
||||
ht.bfloat16 = ht.uint16
|
||||
ht.bfloat16 = ht.uint16.filter(lambda x: ((x >> 7) & 0xFF) != 0) # filter subnormal bfloat16
|
||||
ht.fp8e4m3 = ht.uint8
|
||||
ht.fp8e5m2 = ht.uint8
|
||||
|
||||
@@ -138,7 +138,6 @@ class TestDTypeALU(unittest.TestCase):
|
||||
def test_float16_unary(self, a, op): universal_test_unary(a, dtypes.float16, op)
|
||||
|
||||
@unittest.skipUnless(is_dtype_supported(dtypes.bfloat16), f"no bfloat16 on {Device.DEFAULT}")
|
||||
@unittest.skipIf(CPU_LLVM, "bfloat16 precision issues with CPU_LLVM")
|
||||
@given(ht.bfloat16, strat.sampled_from(unary_operations))
|
||||
def test_bfloat16_unary(self, a, op): universal_test_unary(from_storage_scalar(a, dtypes.bfloat16), dtypes.bfloat16, op)
|
||||
|
||||
|
||||
@@ -189,16 +189,18 @@ class AM_SMU(AM_IP):
|
||||
return table_t.from_buffer(bytearray(self.adev.vram.view(self.driver_table_paddr, ctypes.sizeof(table_t))[:]))
|
||||
|
||||
def set_clocks(self, level):
|
||||
if self.adev.ip_ver[am.MP0_HWIP] in {(13,0,6), (13,0,12)}: return # TODO
|
||||
|
||||
if not hasattr(self, 'clcks'):
|
||||
clks = [self.smu_mod.PPCLK_UCLK, self.smu_mod.PPCLK_FCLK, self.smu_mod.PPCLK_SOCCLK]
|
||||
if self.adev.ip_ver[am.MP0_HWIP] not in {(13,0,6), (13,0,12)}: clks.append(self.smu_mod.PPCLK_GFXCLK)
|
||||
|
||||
self.clcks = {}
|
||||
for clck in [self.smu_mod.PPCLK_GFXCLK, self.smu_mod.PPCLK_UCLK, self.smu_mod.PPCLK_FCLK, self.smu_mod.PPCLK_SOCCLK]:
|
||||
for clck in clks:
|
||||
cnt = self._send_msg(self.smu_mod.PPSMC_MSG_GetDpmFreqByIndex, (clck<<16)|0xff, read_back_arg=True)&0x7fffffff
|
||||
self.clcks[clck] = [self._send_msg(self.smu_mod.PPSMC_MSG_GetDpmFreqByIndex, (clck<<16)|i, read_back_arg=True)&0x7fffffff for i in range(cnt)]
|
||||
|
||||
for clck, vals in self.clcks.items():
|
||||
self._send_msg(self.smu_mod.PPSMC_MSG_SetSoftMinByFreq, clck << 16 | (vals[level]))
|
||||
if not vals: continue
|
||||
with contextlib.suppress(TimeoutError): self._send_msg(self.smu_mod.PPSMC_MSG_SetSoftMinByFreq, clck << 16 | (vals[level]), timeout=20)
|
||||
self._send_msg(self.smu_mod.PPSMC_MSG_SetSoftMaxByFreq, clck << 16 | (vals[level]))
|
||||
|
||||
def _smu_cmn_send_msg(self, msg:int, param=0, debug=False):
|
||||
|
||||
@@ -70,7 +70,7 @@ def validate_index(buf:UOp, idx:UOp, gate:UOp|None=None):
|
||||
# WEBGPU has a BITCAST in the index. TODO: fix
|
||||
if any(x.op is Ops.BITCAST for x in idx.toposort()): return True
|
||||
|
||||
if not z3_imported: raise ImportError("z3 >= 4.12.4 is required for bounds checking, try IGNORE_OOB=0 or \"pip install 'z3-solver>=4.12.4\"")
|
||||
if not z3_imported: raise ImportError("bounds checking requires z3 >= 4.12.4, use IGNORE_OOB=1 to disable, or \"pip install 'z3-solver>=4.12.4\"")
|
||||
solver = z3.Solver(ctx=z3.Context())
|
||||
z3_idx, z3_mask = uops_to_z3(solver, idx, gate)
|
||||
solver.add(z3_mask)
|
||||
|
||||
Reference in New Issue
Block a user