From 134474ca975f545004005fe3f1869240baf55818 Mon Sep 17 00:00:00 2001 From: ZaneHam Date: Thu, 3 Sep 2026 21:09:01 +1200 Subject: [PATCH] variadic packs, multiple translation units, matrix cores and the i1 fixes --- .gitignore | 6 + CHANGELOG.md | 4 + Makefile | 7 +- lang/en.txt | 4 + mfrg.s | 40 ++ src/amdgpu/amdgpu.h | 30 +- src/amdgpu/emit.c | 29 +- src/amdgpu/enc_tab.c | 44 ++ src/amdgpu/encode.c | 6 + src/amdgpu/isel.c | 133 +++++-- src/amdgpu/ra_ssa.c | 1 + src/barracuda.h | 1 + src/cpu/cpu_emit.c | 49 +-- src/cpu/rv64_emit.c | 51 +-- src/fe/ast.h | 14 + src/fe/bc_err.c | 12 +- src/fe/bc_err.h | 14 +- src/fe/bc_render.c | 1 + src/fe/lexer.c | 16 +- src/fe/parser.c | 276 ++++++++++++- src/fe/parser.h | 3 + src/fe/preproc.c | 52 ++- src/fe/preproc.h | 2 + src/fe/sema.c | 51 ++- src/ir/bir.c | 130 ++++++ src/ir/bir.h | 26 +- src/ir/bir_lower.c | 926 ++++++++++++++++++++++++++++++++----------- src/ir/bir_lower.h | 4 + src/main.c | 412 ++++++++++--------- src/metal/emit.c | 5 + src/nvidia/emit.c | 55 ++- src/nvidia/isel.c | 338 ++++++++++++---- src/nvidia/nvidia.h | 23 +- src/tensix/isel.c | 6 + src/tensix/rv_isel.c | 69 +--- src/triton/lower.c | 35 +- src/triton/sema.c | 27 +- src/triton/triton.h | 3 + tests/bpad.cu | 7 + tests/bstr.cu | 5 + tests/bstride.cu | 27 ++ tests/gpu_mma.c | 129 ++++++ tests/i1esc.cu | 20 + tests/mfbf.cu | 2 + tests/mfbf.opt | 6 + tests/mfi8.cu | 2 + tests/mfi8.opt | 6 + tests/mfrg.cu | 2 + tests/mfrg.opt | 6 + tests/mma16.cu | 16 + tests/mma16.opt | 5 + tests/packs.cu | 81 ++++ tests/test_mfma.opt | 8 +- tests/tmain.c | 3 + tests/tmma.c | 185 +++++++++ tests/tmtu.c | 195 +++++++++ tests/tnv_bstr.c | 77 ++++ tests/tnv_i1.c | 109 +++++ tests/tpack.c | 255 ++++++++++++ tests/tphase.c | 167 +++++++- tests/trpi.c | 363 ++++++++++++++++- tests/tu_dup.cu | 1 + tests/tu_g1.cu | 3 + tests/tu_g2.cu | 3 + tests/tu_h1.cu | 3 + tests/tu_h2.cu | 3 + tests/tu_hdr.cuh | 5 + tests/tu_lib.cu | 3 + tests/tu_sta.cu | 3 + tests/tu_stb.cu | 3 + tests/tu_t1.cu | 4 + tests/tu_t2.cu | 4 + tests/tu_use.cu | 1 + tests/tu_use.opt | 6 + 74 files changed, 3847 insertions(+), 776 deletions(-) create mode 100644 mfrg.s create mode 100644 tests/bpad.cu create mode 100644 tests/bstr.cu create mode 100644 tests/bstride.cu create mode 100644 tests/gpu_mma.c create mode 100644 tests/i1esc.cu create mode 100644 tests/mfbf.cu create mode 100644 tests/mfbf.opt create mode 100644 tests/mfi8.cu create mode 100644 tests/mfi8.opt create mode 100644 tests/mfrg.cu create mode 100644 tests/mfrg.opt create mode 100644 tests/mma16.cu create mode 100644 tests/mma16.opt create mode 100644 tests/packs.cu create mode 100644 tests/tmma.c create mode 100644 tests/tmtu.c create mode 100644 tests/tnv_bstr.c create mode 100644 tests/tnv_i1.c create mode 100644 tests/tpack.c create mode 100644 tests/tu_dup.cu create mode 100644 tests/tu_g1.cu create mode 100644 tests/tu_g2.cu create mode 100644 tests/tu_h1.cu create mode 100644 tests/tu_h2.cu create mode 100644 tests/tu_hdr.cuh create mode 100644 tests/tu_lib.cu create mode 100644 tests/tu_sta.cu create mode 100644 tests/tu_stb.cu create mode 100644 tests/tu_t1.cu create mode 100644 tests/tu_t2.cu create mode 100644 tests/tu_use.cu create mode 100644 tests/tu_use.opt diff --git a/.gitignore b/.gitignore index 908e089..5ab984a 100644 --- a/.gitignore +++ b/.gitignore @@ -53,6 +53,12 @@ __pycache__/ /tri_*.ttinsn /tri_*.elf +# Line-splicing fixtures. The suite writes these itself because +# .gitattributes checks the tree out as eol=lf, so a committed CRLF +# fixture would arrive as an LF one and test nothing. +/tests/spl_lf.cu +/tests/spl_crlf.cu + # Ad-hoc test kernels and runners that piled up during the Moa GPU # port and the parameter-count and scratch-allocation investigations. # Listed by name rather than pattern so future intentional tests are diff --git a/CHANGELOG.md b/CHANGELOG.md index 40ff8e3..f79fe23 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,10 @@ Booth — Changelog ### Frontend +- variadic template parameter packs, several `.cu` files as separate + translation units, `mma.sync` and `mfma` lowering, and an i1 that no + longer strides by zero (Zane Hambly, 2026-09-03) + - `(a) + (b)` adds again; the parser treated any parenthesised identifier as a type name without asking whether it named one, so the left operand vanished into a cast with no diagnostic (Zane Hambly, 2026-09-03) diff --git a/Makefile b/Makefile index f5e7367..ffbc6c5 100644 --- a/Makefile +++ b/Makefile @@ -127,7 +127,7 @@ TARGET = kath # strict flags without libhsa or libcuda having to be present. Neither is a # trunner test; both carry their own main and want real hardware. EXAMPLE_SRC = $(wildcard examples/*.c) -HOSTCHK = $(OBJDIR)/tests/tnv_rt.o $(patsubst examples/%.c,$(OBJDIR)/examples/%.o,$(EXAMPLE_SRC)) +HOSTCHK = $(OBJDIR)/tests/tnv_rt.o $(OBJDIR)/tests/tnv_i1.o $(OBJDIR)/tests/tnv_bstr.o $(OBJDIR)/tests/gpu_mma.o $(patsubst examples/%.c,$(OBJDIR)/examples/%.o,$(EXAMPLE_SRC)) all: $(TARGET) $(ALT_RT) $(HOSTCHK) @@ -157,6 +157,7 @@ TSRC = tests/tmain.c tests/tsmoke.c tests/tcomp.c tests/tenc.c \ tests/tra_ssa.c \ tests/tguard.c \ tests/ttriton.c \ + tests/tmma.c \ tests/ttdf.c \ tests/ttmc.c \ tests/trv_enc.c tests/trv_buf.c tests/trv_elf.c tests/trv_isel.c \ @@ -168,7 +169,9 @@ TSRC = tests/tmain.c tests/tsmoke.c tests/tcomp.c tests/tenc.c \ tests/trpi.c \ tests/tmlir.c \ tests/tbir.c \ - tests/tocm.c + tests/tocm.c \ + tests/tpack.c \ + tests/tmtu.c TOBJS = $(TSRC:%.c=$(OBJDIR)/%.o) COBJS = $(OBJDIR)/src/kauri_impl.o $(OBJDIR)/src/ir/bir.o $(OBJDIR)/src/ir/bir_print.o $(OBJDIR)/src/ir/bir_lower.o $(OBJDIR)/src/ir/bir_mem2reg.o $(OBJDIR)/src/ir/bir_cfold.o $(OBJDIR)/src/ir/bir_dce.o $(OBJDIR)/src/ir/bir_struct.o $(OBJDIR)/src/ir/bir_insert.o $(OBJDIR)/src/ir/bir_sroa.o $(OBJDIR)/src/ir/bir_inline.o \ diff --git a/lang/en.txt b/lang/en.txt index a31f278..ec8a093 100644 --- a/lang/en.txt +++ b/lang/en.txt @@ -83,6 +83,10 @@ E123=expected identifier after 'import' E124=expected attribute name after '.' E125=unrecognised expression +# ---- Linking translation units (E126-E127) ---- +E126='%s' is defined in more than one translation unit +E127=too many globals to reference (max %d) + # ---- ABEND Messages (hex IDs) ---- A0C1=illegal GPU instruction A0C4=memory access violation (page not mapped) diff --git a/mfrg.s b/mfrg.s new file mode 100644 index 0000000..236f42e --- /dev/null +++ b/mfrg.s @@ -0,0 +1,40 @@ + .amdgcn_target "amdgcn-amd-amdhsa--gfx942" + .text + + .globl mf16 + .p2align 8 + .type mf16,@function +mf16: + ; 1 SGPRs, 248 VGPRs, 0 LDS bytes, 0 scratch bytes +.LBB0: + s_load_dwordx2 s[4:5], s[0:1], 0 + s_load_dwordx2 s[6:7], s[0:1], 8 + s_load_dwordx2 s[8:9], s[0:1], 16 + v_mov_b32 v6, 0 + v_mov_b32 v5, 0 + v_mov_b32 v0, 0 + v_add_u32 v4, 4, v5 + v_add_u32 v3, 4, v0 + s_waitcnt lgkmcnt(0) + global_load_dword v224, v6, s[8:9] + v_add_u32 v2, 4, v6 + global_load_dword v240, v0, s[4:5] + v_add_u32 v1, 8, v6 + global_load_dword v244, v5, s[6:7] + v_add_u32 v0, 12, v6 + global_load_dword v245, v4, s[6:7] + global_load_dword v227, v0, s[8:9] + global_load_dword v241, v3, s[4:5] + global_load_dword v225, v2, s[8:9] + global_load_dword v226, v1, s[8:9] + s_waitcnt vmcnt(0) + v_mfma_f32_16x16x16f16 v[208:211], v[240:241], v[244:245], v[224:227] + v_add_u32 v1, 4, v6 + global_store_dword v6, v208, s[8:9] + v_add_u32 v0, 8, v6 + global_store_dword v1, v209, s[8:9] + v_add_u32 v1, 12, v6 + global_store_dword v0, v210, s[8:9] + global_store_dword v1, v211, s[8:9] + s_endpgm + diff --git a/src/amdgpu/amdgpu.h b/src/amdgpu/amdgpu.h index 63b60ff..de0e067 100644 --- a/src/amdgpu/amdgpu.h +++ b/src/amdgpu/amdgpu.h @@ -20,10 +20,6 @@ #define AMD_WAVE_SIZE 32 #define AMD_WAVE64 64 -/* Pre-loaded SGPRs — layout depends on kernel needs. - * Default (no dispatch_ptr): s[0:1]=kernarg, s2+=TGID → reserved=3 - * With dispatch_ptr: s[0:1]=dispatch, s[2:3]=kernarg, s4+=TGID → reserved=5 - * Actual positions computed per-kernel in isel. */ #define AMD_KERN_MIN_RESERVED 2 /* absolute floor: kernarg pair */ /* Pre-loaded VGPRs */ @@ -287,6 +283,9 @@ typedef enum { AMD_V_MFMA_F32_32X32X16_BF8_FP8, AMD_V_MFMA_F32_32X32X16_BF8_BF8, /* F64 matrix (gfx942 CDNA3) */ + AMD_V_MFMA_I32_16X16X32_I8, + AMD_V_MFMA_I32_32X32X16_I8, + AMD_V_MFMA_F64_4X4X4_F64, AMD_V_MFMA_F64_16X16X4_F64, @@ -313,11 +312,25 @@ typedef enum { typedef struct { uint8_t kind; /* mop_kind_t */ - uint8_t pad; + uint8_t nreg; /* registers covered; 0 and 1 both mean one */ uint16_t reg_num; /* physical or virtual reg number */ int32_t imm; /* immediate value or special reg ID */ } moperand_t; /* 8 bytes */ +#define MFMA_NONE 0xFFu + +#define RA_MFMA_LO 208u /* v208..v247: the aligned fragment window */ +#define RA_MFMA_HI 248u + +typedef struct { + uint16_t op; /* amd_op_t */ + uint8_t wa, wb, wd; /* A, B and C/D widths in registers */ + uint8_t op90a, op942; +} amd_mfma_t; + +const amd_mfma_t *amd_mfsh(uint16_t op); +uint8_t amd_mfop(const amd_mfma_t *sh, int target); + /* ---- Machine Instruction ---- */ #define MINST_MAX_OPS 6 @@ -363,12 +376,10 @@ typedef struct { uint16_t bir_func; /* BIR func index (for rplan BIR scan) */ uint8_t needs_dispatch; /* 1 if kernel uses blockDim/gridDim (dispatch_ptr) */ uint8_t max_dim; /* highest dim used: 0=x, 1=xy, 2=xyz */ + uint8_t uses_mfma; /* reserve the aligned MFMA fragment window */ uint32_t launch_bounds_max; /* 0 = unconstrained. >0 = programmer's optimistic thread count */ uint32_t launch_bounds_min; /* 0 = not set */ - /* Resource plan — stamped by amd_rplan(), read by isel + emit. - * Target decisions made once. No is_cdna() downstream. - * Like pre-flight checks: argue with the checklist, not the runway. */ uint8_t exec_w; /* 0=B32 (Wave32), 1=B64 (Wave64) */ uint8_t smem_hz; /* 1=SMEM→SALU hazard, promote to VALU */ uint8_t scr_afs; /* 1=architected flat scratch (no prologue) */ @@ -399,6 +410,7 @@ typedef struct { #define AMD_MAX_MINSTS (1 << 18) /* 256K machine instructions */ #define AMD_MAX_MBLOCKS (1 << 16) #define AMD_MAX_MFUNCS (1 << 12) +#define AMD_MAX_ELFK 64 #define AMD_MAX_VREGS (1 << 16) #define AMD_CODE_SIZE (4*1024*1024) #define AMD_ASM_SIZE (4*1024*1024) @@ -445,6 +457,8 @@ typedef struct { char asm_buf[AMD_ASM_SIZE]; uint32_t asm_len; + + int enc_err; /* encoder refused something; fail the compile */ } amd_module_t; /* ---- Encoding Table Entry ---- */ diff --git a/src/amdgpu/emit.c b/src/amdgpu/emit.c index acd333a..d1e89a6 100644 --- a/src/amdgpu/emit.c +++ b/src/amdgpu/emit.c @@ -688,7 +688,8 @@ static void ra_lin(amd_module_t *A, uint32_t mf_idx) * for spill relays. GFX942's AccVGPRs are MFMA-only — the * 8-bit encoding fields in VOP/FLAT literally can't see them. * We tried. The hardware was unimpressed. */ - for (uint16_t r = RA_VGPR_CEIL; r-- > 0; ) + uint16_t vceil = F->uses_mfma ? RA_MFMA_LO : RA_VGPR_CEIL; + for (uint16_t r = vceil; r-- > 0; ) RA.vgpr_free[RA.num_vgpr_free++] = (uint8_t)r; ra_nspill = 0; @@ -791,6 +792,7 @@ static void ra_lin(amd_module_t *A, uint32_t mf_idx) if (F->is_kernel && F->num_sgprs < F->first_alloc_sgpr) F->num_sgprs = F->first_alloc_sgpr; F->num_vgprs = RA.max_vgpr; + if (F->uses_mfma && F->num_vgprs < RA_MFMA_HI) F->num_vgprs = RA_MFMA_HI; /* Match old regalloc_function order exactly: * 1. min SGPR/VGPR + launch_bounds @@ -1952,6 +1954,7 @@ static void ra_gc(amd_module_t *A, uint32_t mf_idx) F->num_sgprs = max_sgpr; F->num_vgprs = max_vgpr; + if (F->uses_mfma && F->num_vgprs < RA_MFMA_HI) F->num_vgprs = RA_MFMA_HI; gc_success = 1; break; /* success, no spills */ } @@ -1999,6 +2002,11 @@ static void asm_append(amd_module_t *A, const char *fmt, ...) static void print_operand(amd_module_t *A, const moperand_t *op) { + if (op->nreg > 1 && (op->kind == MOP_SGPR || op->kind == MOP_VGPR)) { + asm_append(A, "%c[%u:%u]", op->kind == MOP_SGPR ? 's' : 'v', + op->reg_num, op->reg_num + op->nreg - 1); + return; + } switch (op->kind) { case MOP_SGPR: asm_append(A, "s%u", op->reg_num); @@ -2464,13 +2472,17 @@ int amdgpu_emit_elf(amd_module_t *A, const char *path) * like a well-organised criminal enterprise. */ static uint8_t rodata[16384]; /* up to ~256 KDs with alignment */ uint32_t rodata_len = 0; - static uint32_t rodata_kd_off[64]; /* KD offset within .rodata */ - static uint32_t code_offsets[64]; /* code offset within .text */ + static uint32_t rodata_kd_off[AMD_MAX_ELFK]; /* KD offset within .rodata */ + static uint32_t code_offsets[AMD_MAX_ELFK]; /* code offset within .text */ uint32_t num_kernels = 0; for (uint32_t fi = 0; fi < A->num_mfuncs; fi++) { if (!A->mfuncs[fi].is_kernel) continue; - if (num_kernels >= 64) break; + if (num_kernels >= AMD_MAX_ELFK) { + fprintf(stderr, "amdgpu: more kernels than one object holds " + "(max %u)\n", AMD_MAX_ELFK); + return BC_ERR_AMDGPU; + } mfunc_t *F = &A->mfuncs[fi]; @@ -2718,8 +2730,7 @@ int amdgpu_emit_elf(amd_module_t *A, const char *path) uint32_t pt_idx = A->bir->type_fields[ft->count + pi]; const bir_type_t *pt = &A->bir->types[pt_idx]; is_ptr = (pt->kind == BIR_TYPE_PTR); - if (!is_ptr) - arg_sz = (pt->width > 0) ? (uint32_t)(pt->width / 8) : 4; + if (!is_ptr) arg_sz = bir_bsz(A->bir, pt_idx, 8); } if (is_ptr) { mp_fixmap(mp_buf, &mp_pos, 4); @@ -2831,7 +2842,8 @@ int amdgpu_emit_elf(amd_module_t *A, const char *path) static uint32_t sk_name[64], sf_name[64]; /* .strtab offsets */ ki = 0; - for (uint32_t fi = 0; fi < A->num_mfuncs && ki < num_kernels && ki < 64; fi++) { + for (uint32_t fi = 0; fi < A->num_mfuncs && ki < num_kernels + && ki < AMD_MAX_ELFK; fi++) { if (!A->mfuncs[fi].is_kernel) continue; const char *name = A->bir->strings + A->mfuncs[fi].name; char kd[256]; @@ -2908,7 +2920,8 @@ int amdgpu_emit_elf(amd_module_t *A, const char *path) uint32_t si = 1; ki = 0; - for (uint32_t fi = 0; fi < A->num_mfuncs && ki < num_kernels && ki < 64; fi++) { + for (uint32_t fi = 0; fi < A->num_mfuncs && ki < num_kernels + && ki < AMD_MAX_ELFK; fi++) { if (!A->mfuncs[fi].is_kernel) continue; /* .kd descriptor (STT_OBJECT) in .rodata (section 5) */ diff --git a/src/amdgpu/enc_tab.c b/src/amdgpu/enc_tab.c index f161a63..fe42d1d 100644 --- a/src/amdgpu/enc_tab.c +++ b/src/amdgpu/enc_tab.c @@ -597,6 +597,8 @@ const amd_enc_entry_t amd_enc_table_gfx9[AMD_OP_COUNT] = { [AMD_V_MFMA_F32_32X32X16_FP8_BF8] = { AMD_FMT_VOP3P_MAI, 0x76, "v_mfma_f32_32x32x16_fp8_bf8" }, [AMD_V_MFMA_F32_32X32X16_BF8_FP8] = { AMD_FMT_VOP3P_MAI, 0x75, "v_mfma_f32_32x32x16_bf8_fp8" }, [AMD_V_MFMA_F32_32X32X16_BF8_BF8] = { AMD_FMT_VOP3P_MAI, 0x74, "v_mfma_f32_32x32x16_bf8_bf8" }, + [AMD_V_MFMA_I32_16X16X32_I8] = { AMD_FMT_VOP3P_MAI, 0x57, "v_mfma_i32_16x16x32_i8" }, + [AMD_V_MFMA_I32_32X32X16_I8] = { AMD_FMT_VOP3P_MAI, 0x56, "v_mfma_i32_32x32x16_i8" }, /* F64 — for when your matrix really needs 52 bits of mantissa */ [AMD_V_MFMA_F64_4X4X4_F64] = { AMD_FMT_VOP3P_MAI, 0x6F, "v_mfma_f64_4x4x4f64" }, [AMD_V_MFMA_F64_16X16X4_F64] = { AMD_FMT_VOP3P_MAI, 0x6E, "v_mfma_f64_16x16x4f64" }, @@ -606,3 +608,45 @@ const amd_enc_entry_t amd_enc_table_gfx9[AMD_OP_COUNT] = { [AMD_PSEUDO_COPY] = { AMD_FMT_PSEUDO, 0, "PSEUDO_COPY" }, [AMD_PSEUDO_DEF] = { AMD_FMT_PSEUDO, 0, "PSEUDO_DEF" }, }; + +static const amd_mfma_t mfma_shapes[] = { + { AMD_V_MFMA_F32_4X4X4_F16, 2, 2, 4, 0x4A, 0x4A }, + { AMD_V_MFMA_F32_16X16X16_F16, 2, 2, 4, 0x4D, 0x4D }, + { AMD_V_MFMA_F32_32X32X8_F16, 2, 2, 16, 0x4C, 0x4C }, + { AMD_V_MFMA_F32_4X4X4_BF16_1K, 2, 2, 4, 0x65, 0x5F }, + { AMD_V_MFMA_F32_16X16X16_BF16_1K, 2, 2, 4, 0x67, 0x61 }, + { AMD_V_MFMA_F32_32X32X8_BF16_1K, 2, 2, 16, 0x66, 0x60 }, + { AMD_V_MFMA_F32_4X4X1_F32, 1, 1, 4, 0x42, 0x42 }, + { AMD_V_MFMA_F32_16X16X4_F32, 1, 1, 4, 0x45, 0x45 }, + { AMD_V_MFMA_F32_32X32X2_F32, 1, 1, 16, 0x44, 0x44 }, + { AMD_V_MFMA_I32_4X4X4_I8, 1, 1, 4, 0x52, 0x52 }, + { AMD_V_MFMA_I32_16X16X16_I8, 1, 1, 4, 0x55, MFMA_NONE }, + { AMD_V_MFMA_I32_32X32X8_I8, 1, 1, 16, 0x54, MFMA_NONE }, + { AMD_V_MFMA_I32_16X16X32_I8, 2, 2, 4, MFMA_NONE, 0x57 }, + { AMD_V_MFMA_I32_32X32X16_I8, 2, 2, 16, MFMA_NONE, 0x56 }, + { AMD_V_MFMA_F32_16X16X32_FP8_FP8, 2, 2, 4, MFMA_NONE, 0x73 }, + { AMD_V_MFMA_F32_16X16X32_FP8_BF8, 2, 2, 4, MFMA_NONE, 0x72 }, + { AMD_V_MFMA_F32_16X16X32_BF8_FP8, 2, 2, 4, MFMA_NONE, 0x71 }, + { AMD_V_MFMA_F32_16X16X32_BF8_BF8, 2, 2, 4, MFMA_NONE, 0x70 }, + { AMD_V_MFMA_F32_32X32X16_FP8_FP8, 2, 2, 16, MFMA_NONE, 0x77 }, + { AMD_V_MFMA_F32_32X32X16_FP8_BF8, 2, 2, 16, MFMA_NONE, 0x76 }, + { AMD_V_MFMA_F32_32X32X16_BF8_FP8, 2, 2, 16, MFMA_NONE, 0x75 }, + { AMD_V_MFMA_F32_32X32X16_BF8_BF8, 2, 2, 16, MFMA_NONE, 0x74 }, + { AMD_V_MFMA_F64_4X4X4_F64, 2, 2, 2, 0x6F, 0x6F }, + { AMD_V_MFMA_F64_16X16X4_F64, 2, 2, 8, 0x6E, 0x6E }, +}; + +const amd_mfma_t *amd_mfsh(uint16_t op) +{ + for (unsigned i = 0; i < sizeof mfma_shapes / sizeof mfma_shapes[0]; i++) + if (mfma_shapes[i].op == op) return &mfma_shapes[i]; + return NULL; +} + +uint8_t amd_mfop(const amd_mfma_t *sh, int target) +{ + if (!sh) return MFMA_NONE; + if (target == AMD_TARGET_GFX90A) return sh->op90a; + if (target == AMD_TARGET_GFX942) return sh->op942; + return MFMA_NONE; +} diff --git a/src/amdgpu/encode.c b/src/amdgpu/encode.c index 838360c..7dc41aa 100644 --- a/src/amdgpu/encode.c +++ b/src/amdgpu/encode.c @@ -456,6 +456,12 @@ static void encode_flat_global(amd_module_t *A, const minst_t *mi, uint16_t hw_o static void encode_vop3p_mai(amd_module_t *A, const minst_t *mi, uint16_t hw_op) { + const amd_mfma_t *sh = amd_mfsh(mi->op); + if (sh) { + uint8_t t = amd_mfop(sh, A->target); + if (t == MFMA_NONE) { A->enc_err = 1; return; } + hw_op = t; + } /* VOP3P-MAI: 64-bit encoding for MFMA matrix instructions. DW0: [31:23]=0x1A7(prefix) [22:16]=OP(7b) [15]=ACC_CD [14:11]=ABID [10:8]=CBSZ [7:0]=VDST diff --git a/src/amdgpu/isel.c b/src/amdgpu/isel.c index a053ca4..cf99c5e 100644 --- a/src/amdgpu/isel.c +++ b/src/amdgpu/isel.c @@ -240,6 +240,8 @@ static void divergence_analysis(const bir_func_t *F) case BIR_SHFL: case BIR_SHFL_UP: case BIR_SHFL_DOWN: case BIR_SHFL_XOR: case BIR_ALLOCA: /* per-thread scratch — inherently divergent */ + case BIR_MMA: + case BIR_MFRG: case BIR_MFMA: /* matrix result is a collective warp operation */ /* atomic RMW: lanes serialise, each sees a different old value. oh yeah fixin it now: ZH */ case BIR_ATOMIC_ADD: case BIR_ATOMIC_SUB: @@ -704,20 +706,6 @@ static int bir_type_width(uint32_t tidx) return 32; } -static uint32_t arrsz(uint32_t tidx) -{ - if (tidx >= S.bir->num_types) return 4; - const bir_type_t *t = &S.bir->types[tidx]; - if (t->kind == BIR_TYPE_ARRAY) - return t->count * arrsz(t->inner); - if (t->kind == BIR_TYPE_INT || t->kind == BIR_TYPE_FLOAT - || t->kind == BIR_TYPE_BFLOAT) - return t->width / 8; - if (t->kind == BIR_TYPE_PTR) return 8; - if (t->kind == BIR_TYPE_STRUCT) return t->num_fields * 4; - return 4; -} - /* Get type kind */ static int bir_type_kind(uint32_t tidx) { @@ -734,13 +722,17 @@ static int get_addrspace(uint32_t tidx) return BIR_AS_GLOBAL; } -/* Get pointee type's size in bytes */ static uint32_t pointee_size(uint32_t ptr_type) { - if (ptr_type >= S.bir->num_types) return 4; - const bir_type_t *pt = &S.bir->types[ptr_type]; - if (pt->kind != BIR_TYPE_PTR || pt->inner >= S.bir->num_types) return 4; - return arrsz(pt->inner); + return bir_gsz(S.bir, ptr_type, 8); +} + +static void amd_unsz(const char *what, uint32_t ty) +{ + char buf[128]; + if (bir_type_str(S.bir, ty, buf, (int)sizeof buf) <= 0) buf[0] = 0; + fprintf(stderr, "kath: %s of %s has no storage size\n", what, buf); + S.had_error = 1; } /* ---- Instruction Selection: Individual BIR Opcodes ---- */ @@ -1418,6 +1410,8 @@ static void isel_gep(uint32_t idx, const bir_inst_t *I, int div) uint32_t elem_sz = pointee_size(ptr_type); uint32_t base_val = get_op(I, 0); + if (!elem_sz) { amd_unsz("gep", ptr_type); return; } + /* Check if base pointer carries an SGPR pair (saddr mode) */ uint16_t sbase = 0xFFFF; if (!BIR_VAL_IS_CONST(base_val) && base_val != BIR_VAL_NONE) { @@ -1556,6 +1550,9 @@ static void isel_gep(uint32_t idx, const bir_inst_t *I, int div) static void isel_alloca(uint32_t idx, const bir_inst_t *I) { + uint32_t sz = pointee_size(I->type); + if (!sz) { amd_unsz("alloca", I->type); return; } + /* Compute scratch frame offset */ uint32_t align = 1u << I->subop; S.scratch_offset = (S.scratch_offset + align - 1) & ~(align - 1); @@ -1569,15 +1566,13 @@ static void isel_alloca(uint32_t idx, const bir_inst_t *I) /* v_mov_b32 vr, scratch_offset */ emit1(AMD_V_MOV_B32, mop_vreg_v((uint16_t)vr), mop_imm((int32_t)S.scratch_offset)); - uint32_t sz = pointee_size(I->type); - if (sz < 4) sz = 4; - S.scratch_offset += sz; + S.scratch_offset += (sz + 3u) & ~3u; } static void isel_shared_alloc(uint32_t idx, const bir_inst_t *I) { uint32_t sz = pointee_size(I->type); - if (sz < 1) sz = 4; + if (!sz) { amd_unsz("shared_alloc", I->type); return; } /* Align to 4 bytes */ S.lds_offset = (S.lds_offset + 3u) & ~3u; uint32_t vr = map_bir_val(idx, 0); @@ -2182,12 +2177,10 @@ static void isel_warp(uint32_t idx, const bir_inst_t *I) /* F64 matrix (gfx942) */ #define MFMA_F64_4x4x4 20 #define MFMA_F64_16x16x4 21 +#define MFMA_I8_16x16x32 22 +#define MFMA_I8_32x32x16 23 -static void isel_mfma(uint32_t idx, const bir_inst_t *I) -{ - /* Map subop → AMD machine opcode. The hardware does the hard - part; we just shuffle operands like a very expensive postman. */ - static const uint16_t mfma_ops[] = { +static const uint16_t mfma_op_tab[] = { [MFMA_F16_4x4x4] = AMD_V_MFMA_F32_4X4X4_F16, [MFMA_F16_16x16x16] = AMD_V_MFMA_F32_16X16X16_F16, [MFMA_F16_32x32x8] = AMD_V_MFMA_F32_32X32X8_F16, @@ -2210,23 +2203,77 @@ static void isel_mfma(uint32_t idx, const bir_inst_t *I) [MFMA_BF8_BF8_32x32] = AMD_V_MFMA_F32_32X32X16_BF8_BF8, [MFMA_F64_4x4x4] = AMD_V_MFMA_F64_4X4X4_F64, [MFMA_F64_16x16x4] = AMD_V_MFMA_F64_16X16X4_F64, - }; + [MFMA_I8_16x16x32] = AMD_V_MFMA_I32_16X16X32_I8, + [MFMA_I8_32x32x16] = AMD_V_MFMA_I32_32X32X16_I8, +}; - uint8_t var = I->subop; - if (var > MFMA_F64_16x16x4) return; - uint16_t mop = mfma_ops[var]; +#define MF_D 208u /* up to 16, 16-aligned */ +#define MF_C 224u /* up to 16, 16-aligned */ +#define MF_A 240u /* up to 4, 4-aligned */ +#define MF_B 244u /* up to 4, 4-aligned */ - /* MFMA: ops[0]=A, ops[1]=B, ops[2]=C(accum). All VGPR on gfx942. */ - moperand_t a = ensure_vgpr(resolve_val(I->operands[0], 1)); - moperand_t b = ensure_vgpr(resolve_val(I->operands[1], 1)); - moperand_t c = ensure_vgpr(resolve_val(I->operands[2], 1)); - uint32_t vr = map_bir_val(idx, 1); +static void isel_refuse(const char *what); - /* Wait for any pending VMEM before the matrix op — the hardware - won't schedule around these for you */ +static void mf_ldst(int store, uint32_t ptr_val, uint16_t phys, int32_t off) +{ + uint16_t sbase = 0xFFFF; + if (ptr_val != BIR_VAL_NONE && !BIR_VAL_IS_CONST(ptr_val)) { + uint32_t si = BIR_VAL_INDEX(ptr_val); + if (si < BIR_MAX_INSTS) sbase = S.amd->val_sbase[si]; + } + moperand_t addr = ensure_vgpr(resolve_val(ptr_val, 1)); + moperand_t ops[MINST_MAX_OPS]; + + if (sbase != 0xFFFF) { + moperand_t a = addr; + if (off != 0) { + a = mop_vreg_v((uint16_t)new_vreg(1)); + emit2(AMD_V_ADD_U32, a, mop_imm(off), addr); + } + if (store) { ops[0] = a; ops[1] = mop_vgpr(phys); ops[2] = mop_sgpr(sbase); + emit_minst(AMD_GLOBAL_STORE_DWORD, 0, 3, ops, 0); } + else { ops[0] = mop_vgpr(phys); ops[1] = a; ops[2] = mop_sgpr(sbase); + emit_minst(AMD_GLOBAL_LOAD_DWORD, 1, 2, ops, 0); } + return; + } + if (store) { ops[0] = addr; ops[1] = mop_vgpr(phys); ops[2] = mop_imm(off); + emit_minst(AMD_GLOBAL_STORE_DWORD, 0, 3, ops, 0); } + else emit2(AMD_GLOBAL_LOAD_DWORD, mop_vgpr(phys), addr, mop_imm(off)); +} + +static void isel_mfrg(const bir_inst_t *I) +{ + if (I->subop >= (uint8_t)(sizeof mfma_op_tab / sizeof mfma_op_tab[0])) { + isel_refuse("MFMA variant"); return; + } + uint16_t mop = mfma_op_tab[I->subop]; + const amd_mfma_t *sh = amd_mfsh(mop); + if (!sh || amd_mfop(sh, S.amd->target) == MFMA_NONE) { + isel_refuse("this MFMA shape on this CDNA target"); return; + } + S.mf->uses_mfma = 1; + + for (uint8_t i = 0; i < sh->wa; i++) + mf_ldst(0, I->operands[0], (uint16_t)(MF_A + i), (int32_t)i * 4); + for (uint8_t i = 0; i < sh->wb; i++) + mf_ldst(0, I->operands[1], (uint16_t)(MF_B + i), (int32_t)i * 4); + for (uint8_t i = 0; i < sh->wd; i++) + mf_ldst(0, I->operands[2], (uint16_t)(MF_C + i), (int32_t)i * 4); emit_wait_vm(); - emit3(mop, mop_vreg_v((uint16_t)vr), a, b, c); + moperand_t d = mop_vgpr(MF_D), a = mop_vgpr(MF_A); + moperand_t b = mop_vgpr(MF_B), c = mop_vgpr(MF_C); + d.nreg = sh->wd; a.nreg = sh->wa; b.nreg = sh->wb; c.nreg = sh->wd; + emit3(mop, d, a, b, c); + + for (uint8_t i = 0; i < sh->wd; i++) + mf_ldst(1, I->operands[2], (uint16_t)(MF_D + i), (int32_t)i * 4); +} + +static void isel_mfma(uint32_t idx, const bir_inst_t *I) +{ + (void)idx; (void)I; + isel_refuse("MFMA in register form (use __builtin_mfma_*)"); } static void isel_select(uint32_t idx, const bir_inst_t *I, int div) @@ -2667,6 +2714,12 @@ static void isel_function(uint32_t fi) case BIR_MFMA: isel_mfma(idx, I); break; + case BIR_MMA: + isel_refuse("warp-collective mma (NVIDIA PTX only for now)"); + break; + case BIR_MFRG: + isel_mfrg(I); + break; /* Misc */ case BIR_SELECT: diff --git a/src/amdgpu/ra_ssa.c b/src/amdgpu/ra_ssa.c index 3a199bb..9db7105 100644 --- a/src/amdgpu/ra_ssa.c +++ b/src/amdgpu/ra_ssa.c @@ -948,6 +948,7 @@ static uint32_t rs_aloc(amd_module_t *A, mfunc_t *F, if (F->is_kernel && F->num_sgprs < F->first_alloc_sgpr) F->num_sgprs = F->first_alloc_sgpr; F->num_vgprs = max_vgpr; + if (F->uses_mfma && F->num_vgprs < RA_MFMA_HI) F->num_vgprs = RA_MFMA_HI; return n_spill; } diff --git a/src/barracuda.h b/src/barracuda.h index ea58235..97905bb 100644 --- a/src/barracuda.h +++ b/src/barracuda.h @@ -26,6 +26,7 @@ /* Arguments in one call. Real ocean kernels pass 23, so 16 was not enough. */ #define BC_MAX_ARGS 64 #define BC_MAX_PATH 512 +#define BC_MAX_TUS 256 #define BC_MAX_DEPTH 256 #define CUDA_GLOBAL 0x0001 #define CUDA_DEVICE 0x0002 diff --git a/src/cpu/cpu_emit.c b/src/cpu/cpu_emit.c index befd389..15be45d 100644 --- a/src/cpu/cpu_emit.c +++ b/src/cpu/cpu_emit.c @@ -105,42 +105,12 @@ static void load_val(cpu_mod_t *X,int reg,uint32_t v){ else ld_slot(X,reg,slot(X,BIR_VAL_INDEX(v))); } -/* element size in bytes of a pointer's pointee, default 4 (i32/f32). - * Drives GEP stride, so it must be width-accurate: i8->1, i16->2, - * i32/f32->4, i64/f64->8, ptr-to-ptr->8. */ -static int type_size(const cpu_mod_t *X,uint32_t ty); - -static int pointee_sz(cpu_mod_t *X,uint32_t ty){ - if (tyM->num_types && X->M->types[ty].kind==BIR_TYPE_PTR){ - uint32_t in=X->M->types[ty].inner; - if (inM->num_types){ - uint8_t k=X->M->types[in].kind; - if (k==BIR_TYPE_PTR) return 8; - /* an array of structs strides by the whole struct, not by 4: a - * struct pointee has no width field, so size it properly or every - * index past the first lands in the wrong element. */ - if (k==BIR_TYPE_STRUCT || k==BIR_TYPE_ARRAY || k==BIR_TYPE_VECTOR) return type_size(X,in); - uint32_t w=X->M->types[in].width; - if (w>=8) return (int)(w/8); - } - } - return 4; +static int type_size(const cpu_mod_t *X,uint32_t ty){ + return (int)bir_bsz(X->M,ty,8); } -/* size in bytes of a type. Aggregate layout is naive (no padding), - * which is fine here: the only aggregates we size are tiles, and a tile - * is a run of one uniform scalar, so the plain sum lands exactly right. */ -static int type_size(const cpu_mod_t *X,uint32_t ty){ - if (ty>=X->M->num_types) return 8; - const bir_type_t *t=&X->M->types[ty]; - switch (t->kind){ - case BIR_TYPE_INT: case BIR_TYPE_FLOAT: case BIR_TYPE_BFLOAT: return t->width?(int)(t->width/8):4; - case BIR_TYPE_PTR: return 8; - case BIR_TYPE_ARRAY: return (int)t->count*type_size(X,t->inner); - case BIR_TYPE_VECTOR: return (int)t->width*type_size(X,t->inner); - case BIR_TYPE_STRUCT: { int s=0; for(uint16_t i=0;inum_fields;i++) s+=type_size(X,X->M->type_fields[t->count+i]); return s; } - default: return 8; - } +static int pointee_sz(cpu_mod_t *X,uint32_t ty){ + return (int)bir_gsz(X->M,ty,8); } /* type index of a value (const or inst result); 0 if unknown. */ @@ -369,7 +339,9 @@ static void cpu_func(cpu_mod_t *X,const bir_func_t *F){ const bir_inst_t*I=&X->M->insts[ix]; if ((I->op==BIR_ALLOCA||I->op==BIR_SHARED_ALLOC) && natypeM->num_types)?X->M->types[I->type].inner:0; - int sz=(type_size(X,pte)+7)&~7; if(sz<8)sz=8; + int sz=type_size(X,pte); + if(!sz){ fprintf(stderr,"kath: alloca of a type with no storage size\n"); X->n_errs++; } + sz=(sz+7)&~7; if(sz<8)sz=8; off-=sz; X->alloca_off[na++]=off; } } @@ -509,7 +481,9 @@ static void cpu_func(cpu_mod_t *X,const bir_func_t *F){ else { eb(X,0x31);modrm(X,3,X_RDX,X_RDX); eb(X,0xF7);modrm(X,3,6,X_RCX); } if (I->op==BIR_UREM){ rexw(X,X_RDX,X_RAX);eb(X,0x89);modrm(X,3,X_RDX,X_RAX); } st_slot(X,X_RAX,s); break; } - case BIR_GEP: { int sz=pointee_sz(X,I->type); load_val(X,X_RCX,I->operands[1]); mov_imm(X,X_RAX,sz); eb(X,0x48);eb(X,0x0F);eb(X,0xAF);modrm(X,3,X_RCX,X_RAX); load_val(X,X_RAX,I->operands[0]); rexw(X,X_RCX,X_RAX);eb(X,0x01);modrm(X,3,X_RCX,X_RAX); st_slot(X,X_RAX,s); break; } + case BIR_GEP: { int sz=pointee_sz(X,I->type); + if(!sz){ fprintf(stderr,"kath: gep through a pointer with no storage size\n"); X->n_errs++; break; } + load_val(X,X_RCX,I->operands[1]); mov_imm(X,X_RAX,sz); eb(X,0x48);eb(X,0x0F);eb(X,0xAF);modrm(X,3,X_RCX,X_RAX); load_val(X,X_RAX,I->operands[0]); rexw(X,X_RCX,X_RAX);eb(X,0x01);modrm(X,3,X_RCX,X_RAX); st_slot(X,X_RAX,s); break; } case BIR_LOAD: { load_val(X,X_RAX,I->operands[0]); /* addr in rax */ const bir_type_t *t=(I->typeM->num_types)?&X->M->types[I->type]:0; int isflt=t&&(t->kind==BIR_TYPE_FLOAT||t->kind==BIR_TYPE_BFLOAT); @@ -777,6 +751,9 @@ static void cpu_func(cpu_mod_t *X,const bir_func_t *F){ if (is_float_ty(X,I->type)){ int w64=(I->typeM->num_types&&X->M->types[I->type].width==64); st_xmm_slot(X,X_XMM0,s,w64); } else st_slot(X,X_RAX,s); break; } + case BIR_MMA: case BIR_MFRG: + fprintf(stderr,"kath: warp-collective mma not supported on the x86-64 backend\n"); + X->n_errs++; break; default: mov_imm(X,X_RAX,0); st_slot(X,X_RAX,s); break; }} } diff --git a/src/cpu/rv64_emit.c b/src/cpu/rv64_emit.c index b8a7fc9..54bb5f4 100644 --- a/src/cpu/rv64_emit.c +++ b/src/cpu/rv64_emit.c @@ -284,33 +284,12 @@ static void call_libm1(rv64_mod_t*V,const char *name,uint32_t op0,int32_t s){ fst_slot(V,V_FA0,s); } -/* element size in bytes of a pointer's pointee (drives GEP stride) */ -/* size in bytes of a type. Aggregates summed naively (no padding), which is - * right for these kernels: every field is 4 or 8 bytes, so nothing needs it. */ static int type_size(const rv64_mod_t*V,uint32_t ty){ - if (ty>=V->M->num_types) return 8; - const bir_type_t *t=&V->M->types[ty]; - switch (t->kind){ - case BIR_TYPE_INT: case BIR_TYPE_FLOAT: case BIR_TYPE_BFLOAT: return t->width?(int)(t->width/8):4; - case BIR_TYPE_PTR: return 8; - case BIR_TYPE_ARRAY: return (int)t->count*type_size(V,t->inner); - case BIR_TYPE_VECTOR: return (int)t->width*type_size(V,t->inner); - case BIR_TYPE_STRUCT: { int s=0; for(uint16_t i=0;inum_fields;i++) s+=type_size(V,V->M->type_fields[t->count+i]); return s; } - default: return 8; - } + return (int)bir_bsz(V->M,ty,8); } -static int pointee_sz(rv64_mod_t*V,uint32_t ty){ - if (tyM->num_types && V->M->types[ty].kind==BIR_TYPE_PTR){ - uint32_t in=V->M->types[ty].inner; - if (inM->num_types){ - uint8_t k=V->M->types[in].kind; - if (k==BIR_TYPE_PTR) return 8; - if (k==BIR_TYPE_STRUCT || k==BIR_TYPE_ARRAY || k==BIR_TYPE_VECTOR) return type_size(V,in); - uint32_t w=V->M->types[in].width; if (w>=8) return (int)(w/8); - } - } - return 4; +static int pointee_sz(const rv64_mod_t*V,uint32_t ty){ + return (int)bir_gsz(V->M,ty,8); } /* 64 for an i64, 32 for anything narrower; the shifts and the divider use * it to pick the *W word ops so a 32-bit value shifts like a 32-bit value. */ @@ -387,9 +366,8 @@ static void rv64_func(rv64_mod_t *V,const bir_func_t *F){ const bir_inst_t*I=&V->M->insts[ix]; if ((I->op==BIR_ALLOCA||I->op==BIR_SHARED_ALLOC) && natypeM->num_types)?V->M->types[I->type].inner:0; - int sz=8; if (pteM->num_types){ const bir_type_t*t=&V->M->types[pte]; - if (t->kind==BIR_TYPE_ARRAY) sz=(int)t->count*( (t->innerM->num_types&&V->M->types[t->inner].width)?(int)(V->M->types[t->inner].width/8):4 ); - else if (t->width) sz=(int)(t->width/8); } + int sz=type_size(V,pte); + if(!sz){ fprintf(stderr,"kath: alloca of a type with no storage size\n"); V->n_errs++; } sz=(sz+7)&~7; if(sz<8)sz=8; off-=sz; V->alloca_off[na++]=off; } } @@ -631,14 +609,20 @@ static void rv64_func(rv64_mod_t *V,const bir_func_t *F){ case BIR_SHFL: case BIR_SHFL_UP: case BIR_SHFL_DOWN: case BIR_SHFL_XOR: case BIR_BALLOT: case BIR_VOTE_ANY: case BIR_VOTE_ALL: load_val(V,V_T0,I->operands[1]); st_slot(V,V_T0,s); break; - case BIR_GEP: { int sz=pointee_sz(V,I->type); load_val(V,V_T0,I->operands[1]); e_li(V,V_T1,sz); e_mul(V,V_T0,V_T0,V_T1); load_val(V,V_T1,I->operands[0]); e_add(V,V_T0,V_T0,V_T1); st_slot(V,V_T0,s); break; } + case BIR_GEP: { int sz=pointee_sz(V,I->type); + if(!sz){ fprintf(stderr,"kath: gep through a pointer with no storage size\n"); V->n_errs++; break; } + load_val(V,V_T0,I->operands[1]); e_li(V,V_T1,sz); e_mul(V,V_T0,V_T0,V_T1); load_val(V,V_T1,I->operands[0]); e_add(V,V_T0,V_T0,V_T1); st_slot(V,V_T0,s); break; } case BIR_LOAD: { load_val(V,V_T0,I->operands[0]); - const bir_type_t*t=(I->typeM->num_types)?&V->M->types[I->type]:0; int w=t?(int)t->width:32; - if (w==64) e_ld(V,V_T0,V_T0,0); else if (w==16) e_lh(V,V_T0,V_T0,0); else if (w==8||w==1) e_lb(V,V_T0,V_T0,0); else e_lw(V,V_T0,V_T0,0); + int aw=pointee_sz(V,val_type(V,I->operands[0])); + if (aw==8) e_ld(V,V_T0,V_T0,0); else if (aw==4) e_lw(V,V_T0,V_T0,0); + else if (aw==2) e_lh(V,V_T0,V_T0,0); else if (aw==1) e_lb(V,V_T0,V_T0,0); + else { fprintf(stderr,"kath: %d-byte load has no single RV64 access\n",aw); V->n_errs++; break; } st_slot(V,V_T0,s); break; } case BIR_STORE: { load_val(V,V_T0,I->operands[0]); load_val(V,V_T1,I->operands[1]); - int w=pointee_sz(V,val_type(V,I->operands[1]))*8; - if (w==64) e_sd(V,V_T1,V_T0,0); else if (w==16) e_sh(V,V_T1,V_T0,0); else if (w==8) e_sb(V,V_T1,V_T0,0); else e_sw(V,V_T1,V_T0,0); + int aw=pointee_sz(V,val_type(V,I->operands[1])); + if (aw==8) e_sd(V,V_T1,V_T0,0); else if (aw==4) e_sw(V,V_T1,V_T0,0); + else if (aw==2) e_sh(V,V_T1,V_T0,0); else if (aw==1) e_sb(V,V_T1,V_T0,0); + else { fprintf(stderr,"kath: %d-byte store has no single RV64 access\n",aw); V->n_errs++; } break; } case BIR_ICMP: { load_val(V,V_T0,I->operands[0]); load_val(V,V_T1,I->operands[1]); switch(I->subop){ @@ -666,6 +650,9 @@ static void rv64_func(rv64_mod_t *V,const bir_func_t *F){ if (is_kernel){ if(n_retcodelen; e_jal(V,V_ZERO,0); } /* -> loop_cont */ else { e_ld(V,V_RA,V_SP,0); e_ld(V,V_S0,V_SP,8); e_addi(V,V_SP,V_SP,frame); e_jalr(V,V_ZERO,V_RA,0); } break; + case BIR_MMA: case BIR_MFRG: + fprintf(stderr,"kath: warp-collective mma not supported on the RV64 backend\n"); + V->n_errs++; break; default: e_li(V,V_T0,0); st_slot(V,V_T0,s); break; }} } diff --git a/src/fe/ast.h b/src/fe/ast.h index a6e6da2..65b00a2 100644 --- a/src/fe/ast.h +++ b/src/fe/ast.h @@ -29,6 +29,9 @@ typedef enum { AST_INIT_LIST, AST_SCOPE_RES, AST_TEMPLATE_ARGS, + AST_PACK_EXP, + AST_PACK_SIZE, + AST_FOLD, AST_EXPR_STMT, AST_BLOCK, @@ -66,6 +69,17 @@ typedef enum { AST_TYPE_COUNT } ast_type_t; +#define TP_NTYP 0x01 +#define TP_PACK 0x02 + +#define PRM_PACK 0x01 +#define PRM_VARG 0x02 + +#define FLD_UL 0x01 +#define FLD_UR 0x02 +#define FLD_BL 0x03 +#define FLD_BR 0x04 + #define QUAL_CONST 0x01 #define QUAL_VOLATILE 0x02 #define QUAL_STATIC 0x04 diff --git a/src/fe/bc_err.c b/src/fe/bc_err.c index ce8e120..e9feeff 100644 --- a/src/fe/bc_err.c +++ b/src/fe/bc_err.c @@ -32,7 +32,11 @@ static const char *bc_dflt[BC_EID_MAX] = { /* E025 */ "unexpected token in function body", /* E026 */ "unexpected token in block", /* E027 */ "unexpected token at top level", - /* E028-E039 */ NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL, + /* E028 */ "parameter pack must be the last template parameter here", + /* E029 */ "default argument on a template parameter pack", + /* E030 */ "unsupported: %s", + /* E031 */ "pack expansion has no unexpanded parameter pack", + /* E032-E039 */ NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL, /* ---- Preprocessor ---- */ /* E040 */ "macro string pool exhausted", @@ -65,7 +69,8 @@ static const char *bc_dflt[BC_EID_MAX] = { /* E080 */ "switch expression must be integer type", /* E081 */ "__global__ function must return void", /* E082 */ "'%s' passes more than %d arguments", - /* E083-E099 */ NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL, + /* E083 */ "parameter pack '%s' must be expanded", + /* E084-E099 */ NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL, /* ---- Lowering ---- */ /* E100 */ "too many labels (max 256)", @@ -95,7 +100,8 @@ static const char *bc_dflt[BC_EID_MAX] = { /* E123 */ "expected identifier after 'import'", /* E124 */ "expected attribute name after '.'", /* E125 */ "unrecognised expression", - /* E126-E149 */ + /* E126 */ "'%s' is defined in more than one translation unit", + /* E127 */ "too many globals to reference (max %d)", }; /* ---- ABEND compiled-in defaults ---- diff --git a/src/fe/bc_err.h b/src/fe/bc_err.h index 89bcac5..9814cfa 100644 --- a/src/fe/bc_err.h +++ b/src/fe/bc_err.h @@ -39,6 +39,10 @@ typedef enum { BC_E025 = 25, /* unexpected token in function body */ BC_E026 = 26, /* unexpected token in block */ BC_E027 = 27, /* unexpected token at top level */ + BC_E028 = 28, /* parameter pack must be the last template parameter here */ + BC_E029 = 29, /* default argument on a template parameter pack */ + BC_E030 = 30, /* unsupported: %s -- also raised by sema and lowering */ + BC_E031 = 31, /* pack expansion has no unexpanded parameter pack */ /* ---- Preprocessor (E040-E069) ---- */ BC_E040 = 40, /* macro string pool exhausted */ @@ -70,6 +74,7 @@ typedef enum { BC_E080 = 80, /* switch expression must be integer type */ BC_E081 = 81, /* __global__ function must return void */ BC_E082 = 82, /* '%s' passes more than %d arguments */ + BC_E083 = 83, /* parameter pack '%s' must be expanded */ /* ---- Lowering (E100-E129) ---- */ BC_E100 = 100, /* too many labels (max 256) */ @@ -85,9 +90,6 @@ typedef enum { BC_E110 = 110, /* unknown field in lvalue */ BC_E111 = 111, /* not an lvalue */ - /* ---- Triton frontend (E112-E140), block reserved to avoid the FE - collisions the Python path used to have. E112-E125 are catalogued - below; E126-E140 (sema/lower) carry runtime-built messages. ---- */ BC_E112 = 112, /* indentation too deeply nested */ BC_E113 = 113, /* inconsistent dedent */ BC_E114 = 114, /* unterminated string literal (triton) */ @@ -101,7 +103,11 @@ typedef enum { BC_E122 = 122, /* expected parameter name */ BC_E123 = 123, /* expected identifier after 'import' */ BC_E124 = 124, /* expected attribute name after '.' */ - BC_E125 = 125 /* unrecognised expression */ + BC_E125 = 125, /* unrecognised expression */ + + /* ---- Linking translation units (E126-E127) ---- */ + BC_E126 = 126, /* symbol defined in more than one translation unit */ + BC_E127 = 127 /* too many globals to reference (global_ref subop is 8 bits) */ } bc_eid_t; /* Returns format string for eid -- loaded translation or compiled-in English */ diff --git a/src/fe/bc_render.c b/src/fe/bc_render.c index 9b6fc3e..3c1b2d7 100644 --- a/src/fe/bc_render.c +++ b/src/fe/bc_render.c @@ -38,6 +38,7 @@ static const char *line_at(const char *src, uint32_t line, int *len) } const char *e = p; while (*e != '\0' && *e != '\n') e++; + if (e > p && e[-1] == '\r') e--; /* CRLF source, don't echo the CR */ *len = (int)(e - p); return p; } diff --git a/src/fe/lexer.c b/src/fe/lexer.c index 0114402..621686d 100644 --- a/src/fe/lexer.c +++ b/src/fe/lexer.c @@ -366,6 +366,14 @@ int lexer_token_text(const lexer_t *L, const token_t *tok, return len; } +static uint32_t lx_splc(const lexer_t *L) +{ + if (cur(L) != '\\') return 0; + if (peek(L, 1) == '\n') return 2; + if (peek(L, 1) == '\r' && peek(L, 2) == '\n') return 3; + return 0; +} + static void skip_whitespace(lexer_t *L) { while (!at_end(L)) { @@ -514,12 +522,10 @@ static void scan_pp_line(lexer_t *L) uint16_t start_col = (uint16_t)(L->pos - L->line_start + 1); while (!at_end(L) && cur(L) != '\n') { - if (cur(L) == '\\' && peek(L, 1) == '\n') { - advance(L); /* skip backslash */ - advance(L); /* skip newline (line continuation) */ - } else { + uint32_t n = lx_splc(L); + if (n == 0) n = 1; + for (uint32_t i = 0; i < n; i++) advance(L); - } } emit(L, TOK_PP_LINE, start, L->pos - start, start_line, start_col); diff --git a/src/fe/parser.c b/src/fe/parser.c index c2db378..be0e250 100644 --- a/src/fe/parser.c +++ b/src/fe/parser.c @@ -28,6 +28,9 @@ static const char *ast_names[] = { [AST_INIT_LIST] = "init_list", [AST_SCOPE_RES] = "scope", [AST_TEMPLATE_ARGS] = "template_args", + [AST_PACK_EXP] = "pack_exp", + [AST_PACK_SIZE] = "pack_size", + [AST_FOLD] = "fold", [AST_EXPR_STMT] = "expr_stmt", [AST_BLOCK] = "block", [AST_IF] = "if", @@ -133,6 +136,22 @@ static void parse_error(parser_t *P, bc_eid_t eid, ...) } } +static void nderr(parser_t *P, uint32_t node, bc_eid_t eid, ...) +{ + if (P->num_errors < BC_MAX_ERRORS) { + bc_error_t *e = &P->errors[P->num_errors++]; + e->loc.line = P->nodes[node].line; + e->loc.col = P->nodes[node].col; + e->loc.offset = 0; + e->code = BC_ERR_PARSE; + e->eid = (uint16_t)eid; + va_list ap; + va_start(ap, eid); + vsnprintf(e->msg, sizeof(e->msg), bc_efmt(eid), ap); + va_end(ap); + } +} + static int match(parser_t *P, int type) { if (cur_type(P) == type) { advance(P); return 1; } @@ -237,6 +256,88 @@ static int is_reg_type(const parser_t *P, uint32_t off, uint16_t len) return 0; } +/* ---- Parameter pack registry ---- */ + +static void pk_add(parser_t *P, uint32_t off, uint32_t len) +{ + if (P->npacks >= 32 || len == 0) return; + P->packs[P->npacks].off = off; + P->packs[P->npacks].len = len; + P->npacks++; +} + +static int pk_is(const parser_t *P, uint32_t off, uint32_t len) +{ + const char *q = tntxt(P, off); + for (int i = 0; i < P->npacks; i++) { + if (P->packs[i].len == len && + memcmp(tntxt(P, P->packs[i].off), q, (size_t)len) == 0) + return 1; + } + return 0; +} + +static int pk_tsp(const parser_t *P, uint32_t tsp) +{ + if (!tsp || P->nodes[tsp].type != AST_TYPE_SPEC) return 0; + if (P->nodes[tsp].d.btype.kind != TYPE_NAME) return 0; + uint32_t nm = P->nodes[tsp].first_child; + if (!nm || P->nodes[nm].type != AST_IDENT) return 0; + return pk_is(P, P->nodes[nm].d.text.offset, P->nodes[nm].d.text.len); +} + +static int pk_sub(const parser_t *P, uint32_t node, int depth) +{ + if (!node || depth > 64) return 0; + const ast_node_t *n = &P->nodes[node]; + if (n->type == AST_IDENT && + pk_is(P, n->d.text.offset, n->d.text.len)) return 1; + for (uint32_t c = n->first_child; c; c = P->nodes[c].next_sibling) + if (pk_sub(P, c, depth + 1)) return 1; + return 0; +} + +static void pk_chk(parser_t *P, uint32_t node, int in_exp, int depth) +{ + if (!node || depth > 64) return; + ast_node_t *n = &P->nodes[node]; + int t = n->type; + + if (t == AST_TEMPLATE_PARAM) return; + if (t == AST_PARAM && n->d.oper.op == PRM_PACK) return; + if (t == AST_PACK_EXP || t == AST_PACK_SIZE || t == AST_FOLD) in_exp = 1; + + if (t == AST_IDENT && !in_exp && + pk_is(P, n->d.text.offset, n->d.text.len)) { + char nm[64]; + uint32_t len = n->d.text.len < 63 ? n->d.text.len : 63; + memcpy(nm, tntxt(P, n->d.text.offset), (size_t)len); + nm[len] = 0; + nderr(P, node, BC_E083, nm); + return; + } + + for (uint32_t c = n->first_child; c; c = P->nodes[c].next_sibling) + pk_chk(P, c, in_exp, depth + 1); +} + +static int is_fop(int t) +{ + switch (t) { + case TOK_PLUS: case TOK_MINUS: case TOK_STAR: case TOK_SLASH: + case TOK_PERCENT: case TOK_CARET: case TOK_AMP: case TOK_PIPE: + case TOK_SHL: case TOK_SHR: + case TOK_PLUS_EQ: case TOK_MINUS_EQ: case TOK_STAR_EQ: + case TOK_SLASH_EQ: case TOK_PERCENT_EQ: case TOK_CARET_EQ: + case TOK_AMP_EQ: case TOK_PIPE_EQ: case TOK_SHL_EQ: case TOK_SHR_EQ: + case TOK_ASSIGN: case TOK_EQ: case TOK_NE: case TOK_LT: case TOK_GT: + case TOK_LE: case TOK_GE: case TOK_LAND: case TOK_LOR: case TOK_COMMA: + return 1; + default: + return 0; + } +} + static int is_type_keyword(int type) { switch (type) { @@ -467,6 +568,95 @@ static int looks_like_cast(parser_t *P) return 0; } +static uint32_t pexp(parser_t *P, uint32_t arg) +{ + if (cur_type(P) != TOK_ELLIPSIS) return arg; + advance(P); + if (cur_type(P) == TOK_LBRACKET) { + parse_error(P, BC_E030, "pack indexing"); + return arg; + } + if (!pk_sub(P, arg, 0)) parse_error(P, BC_E031); + uint32_t n = alloc_node(P, AST_PACK_EXP); + if (!n) return arg; + add_child(P, n, arg); + return n; +} + +static int foldp(const parser_t *P) +{ + int depth = 0; + for (uint32_t i = P->pos; i < P->num_tokens; i++) { + int t = P->tokens[i].type; + if (t == TOK_LPAREN || t == TOK_LBRACKET || t == TOK_LBRACE) depth++; + else if (t == TOK_RPAREN || t == TOK_RBRACKET || t == TOK_RBRACE) { + if (--depth <= 0) return 0; + } + else if (t == TOK_ELLIPSIS && depth == 1) return 1; + else if (t == TOK_SEMI || t == TOK_EOF) return 0; + } + return 0; +} + +static uint32_t pfold(parser_t *P) +{ + uint32_t n = alloc_node(P, AST_FOLD); + char got[64]; + + advance(P); + if (cur_type(P) == TOK_ELLIPSIS) { + advance(P); + int op = cur_type(P); + if (!is_fop(op)) { + parse_error(P, BC_E020, "fold operator", + cur_text(P, got, sizeof got)); + return n; + } + advance(P); + uint32_t e = parse_expr(P, 150); + add_child(P, n, e); + P->nodes[n].d.oper.op = op; + if (cur_type(P) == TOK_RPAREN) { + P->nodes[n].d.oper.flags = FLD_UL; + if (!pk_sub(P, e, 0)) parse_error(P, BC_E031); + advance(P); + return n; + } + parse_error(P, BC_E020, ")", cur_text(P, got, sizeof got)); + return n; + } + + uint32_t e1 = parse_expr(P, 150); + add_child(P, n, e1); + int op = cur_type(P); + if (!is_fop(op)) { + parse_error(P, BC_E020, "fold operator", cur_text(P, got, sizeof got)); + return n; + } + P->nodes[n].d.oper.op = op; + advance(P); + expect(P, TOK_ELLIPSIS); + if (cur_type(P) == TOK_RPAREN) { + P->nodes[n].d.oper.flags = FLD_UR; + if (!pk_sub(P, e1, 0)) parse_error(P, BC_E031); + advance(P); + return n; + } + if (cur_type(P) != op) { + parse_error(P, BC_E020, token_type_name(op), + cur_text(P, got, sizeof got)); + return n; + } + advance(P); + uint32_t e2 = parse_expr(P, 150); + add_child(P, n, e2); + int p1 = pk_sub(P, e1, 0), p2 = pk_sub(P, e2, 0); + if (p1 == p2) parse_error(P, BC_E031); + P->nodes[n].d.oper.flags = p1 ? FLD_BR : FLD_BL; + expect(P, TOK_RPAREN); + return n; +} + static uint32_t parse_primary(parser_t *P) { int t = cur_type(P); @@ -536,7 +726,7 @@ static uint32_t parse_primary(parser_t *P) uint32_t ilist = alloc_node(P, AST_INIT_LIST); advance(P); while (cur_type(P) != TOK_RBRACE && cur_type(P) != TOK_EOF) { - uint32_t elem = parse_expr(P, 21); + uint32_t elem = pexp(P, parse_expr(P, 21)); add_child(P, ilist, elem); if (!match(P, TOK_COMMA)) break; } @@ -549,6 +739,27 @@ static uint32_t parse_primary(parser_t *P) } return n; } + if (t == TOK_SIZEOF && peek_type(P, 1) == TOK_ELLIPSIS) { + uint32_t n = alloc_node(P, AST_PACK_SIZE); + advance(P); + advance(P); + expect(P, TOK_LPAREN); + if (cur_type(P) == TOK_IDENT) { + uint32_t id = alloc_node(P, AST_IDENT); + P->nodes[id].d.text.offset = cur(P)->offset; + P->nodes[id].d.text.len = cur(P)->len; + if (!pk_is(P, cur(P)->offset, cur(P)->len)) { + char got[64]; + parse_error(P, BC_E083, cur_text(P, got, sizeof got)); + } + advance(P); + add_child(P, n, id); + } else { + parse_error(P, BC_E022); + } + expect(P, TOK_RPAREN); + return n; + } if (t == TOK_SIZEOF) { uint32_t n = alloc_node(P, AST_SIZEOF); advance(P); @@ -590,6 +801,8 @@ static uint32_t parse_primary(parser_t *P) expect(P, TOK_RPAREN); return cast; } + if (t == TOK_LPAREN && foldp(P)) + return pfold(P); if (looks_like_cast(P)) { uint32_t saved = P->pos; advance(P); @@ -620,7 +833,7 @@ static uint32_t parse_primary(parser_t *P) uint32_t n = alloc_node(P, AST_INIT_LIST); advance(P); while (cur_type(P) != TOK_RBRACE && cur_type(P) != TOK_EOF) { - uint32_t elem = parse_expr(P, 21); + uint32_t elem = pexp(P, parse_expr(P, 21)); add_child(P, n, elem); if (!match(P, TOK_COMMA)) break; } @@ -651,6 +864,10 @@ static uint32_t parse_expr(parser_t *P, int min_prec) for (;;) { int t = cur_type(P); if (t == TOK_EOF || t == TOK_SEMI || t == TOK_RBRACE) break; + if (t == TOK_ELLIPSIS && peek_type(P, 1) == TOK_LBRACKET) { + parse_error(P, BC_E030, "pack indexing"); + break; + } int pbp = postfix_bp(t); if (pbp >= 0 && pbp >= min_prec) { @@ -667,7 +884,7 @@ static uint32_t parse_expr(parser_t *P, int min_prec) add_child(P, call, lhs); advance(P); while (cur_type(P) != TOK_RPAREN && cur_type(P) != TOK_EOF) { - uint32_t arg = parse_expr(P, 21); + uint32_t arg = pexp(P, parse_expr(P, 21)); add_child(P, call, arg); if (!match(P, TOK_COMMA)) break; } @@ -719,7 +936,7 @@ static uint32_t parse_expr(parser_t *P, int min_prec) expect(P, TOK_LAUNCH_CLOSE); expect(P, TOK_LPAREN); while (cur_type(P) != TOK_RPAREN && cur_type(P) != TOK_EOF) { - uint32_t arg = parse_expr(P, 21); + uint32_t arg = pexp(P, parse_expr(P, 21)); add_child(P, launch, arg); if (!match(P, TOK_COMMA)) break; } @@ -928,6 +1145,7 @@ static uint32_t parse_param_list(parser_t *P) if (cur_type(P) == TOK_ELLIPSIS) { uint32_t va = alloc_node(P, AST_PARAM); P->nodes[va].d.oper.flags = 1; + P->nodes[va].d.oper.op = PRM_VARG; advance(P); if (!first) first = va; else P->nodes[last].next_sibling = va; last = va; @@ -944,13 +1162,32 @@ static uint32_t parse_param_list(parser_t *P) { int ptr_depth = 0; while (cur_type(P) == TOK_STAR || cur_type(P) == TOK_AMP || - cur_type(P) == TOK_CONST || cur_type(P) == TOK_CU_RESTRICT) { + cur_type(P) == TOK_LAND || cur_type(P) == TOK_CONST || + cur_type(P) == TOK_CU_RESTRICT) { if (cur_type(P) == TOK_STAR) ptr_depth++; advance(P); } P->nodes[param].d.oper.flags = ptr_depth; } + if (cur_type(P) == TOK_ELLIPSIS) { + if (pk_tsp(P, type)) { + advance(P); + P->nodes[param].d.oper.op = PRM_PACK; + if (cur_type(P) == TOK_LBRACKET) + parse_error(P, BC_E030, "pack indexing"); + } else { + advance(P); + if (!first) first = param; else P->nodes[last].next_sibling = param; + last = param; + uint32_t va = alloc_node(P, AST_PARAM); + P->nodes[va].d.oper.flags = 1; + P->nodes[va].d.oper.op = PRM_VARG; + P->nodes[last].next_sibling = va; + break; + } + } + if (is_fnptr(P)) { int d = P->nodes[param].d.oper.flags; uint32_t fname = fnptr(P, &d); @@ -960,6 +1197,8 @@ static uint32_t parse_param_list(parser_t *P) uint32_t name = alloc_node(P, AST_IDENT); P->nodes[name].d.text.offset = cur(P)->offset; P->nodes[name].d.text.len = cur(P)->len; + if (P->nodes[param].d.oper.op == PRM_PACK) + pk_add(P, cur(P)->offset, cur(P)->len); advance(P); add_child(P, param, name); } @@ -1062,23 +1301,30 @@ static uint32_t parse_declaration(parser_t *P) if (cur_type(P) == TOK_TEMPLATE) { uint32_t tmpl = alloc_node(P, AST_TEMPLATE_DECL); + int sv_npk = P->npacks; + int nparm = 0, lastpk = 0; advance(P); expect(P, TOK_LT); while (cur_type(P) != TOK_GT && cur_type(P) != TOK_EOF) { uint32_t tp = alloc_node(P, AST_TEMPLATE_PARAM); + int fl = 0; if (cur_type(P) == TOK_TYPENAME || cur_type(P) == TOK_CLASS) { - P->nodes[tp].d.oper.flags = 0; advance(P); } else { - P->nodes[tp].d.oper.flags = 1; + fl = TP_NTYP; uint16_t q2, c2; uint32_t ptype = parse_type_spec(P, &q2, &c2); add_child(P, tp, ptype); } + if (match(P, TOK_ELLIPSIS)) fl |= TP_PACK; + P->nodes[tp].d.oper.flags = fl; + nparm++; + if (fl & TP_PACK) lastpk = nparm; if (cur_type(P) == TOK_IDENT) { uint32_t name = alloc_node(P, AST_IDENT); P->nodes[name].d.text.offset = cur(P)->offset; P->nodes[name].d.text.len = cur(P)->len; + if (fl & TP_PACK) pk_add(P, cur(P)->offset, cur(P)->len); advance(P); add_child(P, tp, name); /* A type parameter is a type name for the body below it, @@ -1088,6 +1334,7 @@ static uint32_t parse_declaration(parser_t *P) (uint16_t)P->nodes[name].d.text.len); } if (cur_type(P) == TOK_ASSIGN) { + if (fl & TP_PACK) parse_error(P, BC_E029); advance(P); uint32_t def = parse_expr(P, 21); add_child(P, tp, def); @@ -1097,7 +1344,15 @@ static uint32_t parse_declaration(parser_t *P) } expect(P, TOK_GT); uint32_t inner = parse_declaration(P); + if (lastpk && lastpk != nparm && inner && + (P->nodes[inner].type == AST_STRUCT_DEF || + P->nodes[inner].type == AST_VAR_DECL)) + parse_error(P, BC_E028); + if (lastpk && inner && P->nodes[inner].type == AST_USING) + nderr(P, inner, BC_E030, "variadic alias template"); add_child(P, tmpl, inner); + if (P->npacks > sv_npk) pk_chk(P, inner, 0, 0); + P->npacks = sv_npk; return tmpl; } @@ -1389,11 +1644,8 @@ static uint32_t parse_declaration(parser_t *P) while (cur_type(P) != TOK_GT && cur_type(P) != TOK_EOF) { uint16_t q2, c2; uint32_t ta = parse_type_spec(P, &q2, &c2); - if (ta) add_child(P, targs, ta); - else { - uint32_t ea = parse_expr(P, 21); - add_child(P, targs, ea); - } + if (!ta) ta = parse_expr(P, 21); + add_child(P, targs, pexp(P, ta)); if (!match(P, TOK_COMMA)) break; } expect(P, TOK_GT); diff --git a/src/fe/parser.h b/src/fe/parser.h index b0f9d27..b97c589 100644 --- a/src/fe/parser.h +++ b/src/fe/parser.h @@ -32,6 +32,9 @@ typedef struct parser_s { uint32_t anon_len; uint32_t anon_cnt; + struct { uint32_t off; uint32_t len; } packs[32]; + int npacks; + /* Enclosing struct name, so a constructor can be told apart from a * declaration that happens to start with a type name. len 0 = not in one. */ uint32_t cs_off; diff --git a/src/fe/preproc.c b/src/fe/preproc.c index 223a4b2..55855a8 100644 --- a/src/fe/preproc.c +++ b/src/fe/preproc.c @@ -72,6 +72,21 @@ static void pp_skip_to_eol(preproc_t *pp) pp_advance(pp); } +static uint32_t pp_nlsz(const preproc_t *pp) +{ + if (pp_cur(pp) == '\n') return 1; + if (pp_cur(pp) == '\r' && pp_peek(pp, 1) == '\n') return 2; + return 0; +} + +static uint32_t pp_splc(const preproc_t *pp) +{ + if (pp_cur(pp) != '\\') return 0; + if (pp_peek(pp, 1) == '\n') return 2; + if (pp_peek(pp, 1) == '\r' && pp_peek(pp, 2) == '\n') return 3; + return 0; +} + static int pp_is_ident_start(char c) { return (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || c == '_'; @@ -142,6 +157,16 @@ static void pp_emit_str(preproc_t *pp, const char *s, uint32_t len) pp->out_len += len; } +static void pp_eatnl(preproc_t *pp) +{ + uint32_t n = pp_nlsz(pp); + if (n == 0) return; + for (uint32_t i = 0; i < n; i++) + pp_advance(pp); + for (uint32_t i = 0; i <= pp->nspl; i++) + pp_emit_char(pp, '\n'); +} + /* ---- String pool ---- */ static uint32_t pool_add(preproc_t *pp, const char *s, uint32_t len) @@ -301,10 +326,12 @@ static uint32_t pp_read_ident(const preproc_t *pp, char *buf, uint32_t max) static uint32_t pp_collect_line(preproc_t *pp, char *buf, uint32_t max) { uint32_t len = 0; - while (!pp_at_end(pp) && pp_cur(pp) != '\n') { - if (pp_cur(pp) == '\\' && pp_peek(pp, 1) == '\n') { - pp_advance(pp); /* skip backslash */ - pp_advance(pp); /* skip newline (counts line) */ + while (!pp_at_end(pp) && pp_nlsz(pp) == 0) { + uint32_t n = pp_splc(pp); + if (n > 0) { + for (uint32_t i = 0; i < n; i++) + pp_advance(pp); + pp->nspl++; continue; } if (len + 1 < max) @@ -1471,17 +1498,14 @@ int pp_process(preproc_t *pp) /* Check for start of line — skip horizontal whitespace, look for '#' */ uint32_t line_start = pp->pos; + pp->nspl = 0; pp_skip_hspace(pp); if (pp_at_end(pp)) continue; if (pp_cur(pp) == '#' && pp_peek(pp, 1) != '#') { pp_process_directive(pp); - /* Eat the trailing newline after the directive */ - if (!pp_at_end(pp) && pp_cur(pp) == '\n') { - pp_emit_char(pp, '\n'); /* preserve line count */ - pp_advance(pp); - } + pp_eatnl(pp); continue; } @@ -1489,10 +1513,7 @@ int pp_process(preproc_t *pp) if (!pp_is_active(pp)) { /* In inactive conditional block — skip line, emit newline */ pp_skip_to_eol(pp); - if (!pp_at_end(pp) && pp_cur(pp) == '\n') { - pp_emit_char(pp, '\n'); - pp_advance(pp); - } + pp_eatnl(pp); continue; } @@ -1520,10 +1541,7 @@ int pp_process(preproc_t *pp) pp_expand_and_emit(pp, line, llen); while (joins--) pp_emit_char(pp, '\n'); - if (!pp_at_end(pp) && pp_cur(pp) == '\n') { - pp_emit_char(pp, '\n'); - pp_advance(pp); - } + pp_eatnl(pp); } /* Abandoning the run leaves the include stack loaded, and every entry diff --git a/src/fe/preproc.h b/src/fe/preproc.h index d66ed85..a79343a 100644 --- a/src/fe/preproc.h +++ b/src/fe/preproc.h @@ -57,6 +57,8 @@ typedef struct { uint32_t pos; uint32_t line; + uint32_t nspl; + /* Output */ char *out; uint32_t out_len; diff --git a/src/fe/sema.c b/src/fe/sema.c index 64786d8..fa6965f 100644 --- a/src/fe/sema.c +++ b/src/fe/sema.c @@ -975,9 +975,10 @@ static uint32_t check_expr(sema_ctx_t *S, uint32_t node) get_text(S, callee_n, cname, sizeof(cname)); uint32_t arg_types[BC_MAX_ARGS]; - int nargs = 0; + int nargs = 0, argvar = 0; uint32_t arg = ND(S, callee_n)->next_sibling; while (arg && nargs < BC_MAX_ARGS) { + if (ND(S, arg)->type == AST_PACK_EXP) argvar = 1; arg_types[nargs++] = check_expr(S, arg); arg = ND(S, arg)->next_sibling; } @@ -1033,7 +1034,7 @@ static uint32_t check_expr(sema_ctx_t *S, uint32_t node) if (strcmp(cname, vec_ctors[vi].name) == 0) { uint32_t elem_t = intern_type(S, vec_ctors[vi].elem, 0, 0, 0, 0); uint32_t vt = intern_type(S, STYPE_VECTOR, 0, vec_ctors[vi].lanes, elem_t, 0); - if (nargs != (int)vec_ctors[vi].lanes) + if (nargs != (int)vec_ctors[vi].lanes && !argvar) sema_error(S, node, BC_E073, cname, (int)vec_ctors[vi].lanes, nargs); return annotate(S, node, vt); @@ -1077,6 +1078,20 @@ static uint32_t check_expr(sema_ctx_t *S, uint32_t node) return annotate(S, node, st_int(S)); } + /* ---- Warp-collective 16x16x16 f16 matrix multiply ---- */ + if (strncmp(cname, "__builtin_mma_m16n16k", 21) == 0) { + if (nargs != 6) + sema_error(S, node, BC_E075, cname, nargs); + return annotate(S, node, st_void(S)); + } + + /* ---- MFMA over per-lane fragments in memory ---- */ + if (strncmp(cname, "__builtin_mfma_", 15) == 0) { + if (nargs != 3) + sema_error(S, node, BC_E075, cname, nargs); + return annotate(S, node, st_void(S)); + } + /* ---- MFMA intrinsics (CDNA matrix ops) ---- */ if (strncmp(cname, "__builtin_amdgcn_mfma_", 22) == 0) { if (nargs != 3) @@ -1090,7 +1105,7 @@ static uint32_t check_expr(sema_ctx_t *S, uint32_t node) if (strcmp(cname, cuda_builtins[i].name) != 0) continue; const cuda_builtin_t *b = &cuda_builtins[i]; - if (b->nargs >= 0 && nargs != b->nargs) { + if (b->nargs >= 0 && nargs != b->nargs && !argvar) { sema_error(S, node, BC_E073, cname, b->nargs, nargs); } @@ -1124,7 +1139,7 @@ static uint32_t check_expr(sema_ctx_t *S, uint32_t node) uint32_t ft = sym->type; if (ft < S->num_types && S->types[ft].kind == STYPE_FUNC) { int expected = (int)S->types[ft].width; - if (expected > 0 && nargs != expected) { + if (expected > 0 && nargs != expected && !argvar) { sema_error(S, node, BC_E073, cname, expected, nargs); } @@ -1180,6 +1195,28 @@ static uint32_t check_expr(sema_ctx_t *S, uint32_t node) return annotate(S, node, st_void(S)); } + case AST_PACK_SIZE: + return annotate(S, node, st_ulong(S)); + + case AST_PACK_EXP: { + uint32_t pat = n->first_child; + return annotate(S, node, pat ? check_expr(S, pat) : st_void(S)); + } + + case AST_FOLD: { + uint32_t e = n->first_child; + uint32_t t = st_void(S); + while (e) { t = check_expr(S, e); e = ND(S, e)->next_sibling; } + switch (n->d.oper.op) { + case TOK_LAND: case TOK_LOR: + case TOK_EQ: case TOK_NE: case TOK_LT: + case TOK_GT: case TOK_LE: case TOK_GE: + return annotate(S, node, st_bool(S)); + default: + return annotate(S, node, t); + } + } + default: return annotate(S, node, st_int(S)); } @@ -1504,7 +1541,8 @@ static void collect_func_decl(sema_ctx_t *S, uint32_t node) int nparams = 0; uint32_t c = ND(S, node)->first_child; while (c) { - if (ND(S, c)->type == AST_PARAM && nparams < 32) { + if (ND(S, c)->type == AST_PARAM && nparams < 32 + && ND(S, c)->d.oper.op != PRM_VARG) { uint32_t pt_n = ND(S, c)->first_child; int pdepth = ND(S, c)->d.oper.flags; param_types[nparams++] = resolve_typespec(S, pt_n, pdepth); @@ -1572,7 +1610,8 @@ static void check_func_def(sema_ctx_t *S, uint32_t node) uint32_t c = ND(S, node)->first_child; while (c) { - if (ND(S, c)->type == AST_PARAM) { + if (ND(S, c)->type == AST_PARAM + && ND(S, c)->d.oper.op != PRM_VARG) { uint32_t pt_n = ND(S, c)->first_child; int pdepth = ND(S, c)->d.oper.flags; uint32_t pt = resolve_typespec(S, pt_n, pdepth); diff --git a/src/ir/bir.c b/src/ir/bir.c index 4eb4be3..3fffae5 100644 --- a/src/ir/bir.c +++ b/src/ir/bir.c @@ -107,6 +107,8 @@ static const char *op_names[BIR_OP_COUNT] = { [BIR_FMIN] = "fmin", [BIR_MFMA] = "mfma", + [BIR_MMA] = "mma", + [BIR_MFRG] = "mfrg", [BIR_CALL] = "call", [BIR_SELECT] = "select", @@ -389,6 +391,83 @@ uint32_t bir_type_func(bir_module_t *M, uint32_t ret, return intern_compound(M, BIR_TYPE_FUNC, ret, params, nparams); } +/* ---- Type Sizes ---- */ + +static uint32_t bsz_al(uint32_t sz, uint32_t psz) +{ + uint32_t a = 1; + if (sz > psz) sz = psz; + while (a * 2u <= sz) a *= 2u; + return a; +} + +static uint32_t bsz_up(uint32_t x, uint32_t a) +{ + return (x + a - 1u) & ~(a - 1u); +} + +uint32_t bir_bsz(const bir_module_t *M, uint32_t ty, uint32_t psz) +{ + struct { uint32_t ty, mul, fld, sum, alg; } fr[BIR_BSZ_DEEP]; + uint32_t sp = 1, guard = 4u * BIR_MAX_TYPE_FIELDS; + + fr[0].ty = ty; fr[0].mul = 1; fr[0].fld = 0; fr[0].sum = 0; fr[0].alg = 1; + + while (guard--) { + uint32_t i = sp - 1, w, va, val; + const bir_type_t *T; + + if (fr[i].ty >= M->num_types) return 0; + T = &M->types[fr[i].ty]; + + if (T->kind == BIR_TYPE_ARRAY || T->kind == BIR_TYPE_VECTOR) { + uint32_t n = (T->kind == BIR_TYPE_ARRAY) ? T->count + : (uint32_t)T->width; + if (n && fr[i].mul > 0xFFFFFFFFu / n) return 0; + fr[i].mul *= n; fr[i].ty = T->inner; continue; + } + + if (T->kind == BIR_TYPE_STRUCT) { + if (fr[i].fld < (uint32_t)T->num_fields) { + uint32_t f = T->count + fr[i].fld++; + if (f >= M->num_type_fields || sp >= BIR_BSZ_DEEP) return 0; + fr[sp].ty = M->type_fields[f]; + fr[sp].mul = 1; fr[sp].fld = 0; fr[sp].sum = 0; fr[sp].alg = 1; + sp++; + continue; + } + w = bsz_up(fr[i].sum, fr[i].alg); + va = fr[i].alg; + } else { + switch (T->kind) { + case BIR_TYPE_INT: + case BIR_TYPE_FLOAT: + case BIR_TYPE_BFLOAT: w = ((uint32_t)T->width + 7u) / 8u; break; + case BIR_TYPE_PTR: w = psz; break; + default: return 0; + } + va = bsz_al(w, psz); + } + if (!w) return 0; + + if (fr[i].mul > 0xFFFFFFFFu / w) return 0; + val = fr[i].mul * w; + if (--sp == 0) return val; + + if (va > fr[sp - 1].alg) fr[sp - 1].alg = va; + if (bsz_up(fr[sp - 1].sum, va) > 0xFFFFFFFFu - val) return 0; + fr[sp - 1].sum = bsz_up(fr[sp - 1].sum, va) + val; + } + return 0; +} + +uint32_t bir_gsz(const bir_module_t *M, uint32_t ty, uint32_t psz) +{ + if (ty < M->num_types && M->types[ty].kind == BIR_TYPE_PTR) + return bir_bsz(M, M->types[ty].inner, psz); + return bir_bsz(M, ty, psz); +} + /* ---- String Table ---- */ uint32_t bir_add_string(bir_module_t *M, const char *s, uint32_t len) @@ -466,6 +545,57 @@ int bir_global_is_bytes(const bir_module_t *M, uint32_t gi) return M->consts[ci].kind == BIR_CONST_BYTES; } +int bir_mang(const char *name, uint16_t tu, char *out, int size) +{ + int n = snprintf(out, (size_t)size, "%s__%u", name, (unsigned)tu); + if (n < 0 || n >= size) { out[0] = '\0'; return 1; } + return 0; +} + +uint32_t bir_fsym(const bir_module_t *M, const char *name, uint16_t tu, + int nargs) +{ + char mng[BIR_SYM_MAX]; + + for (int pass = 0; pass < 2; pass++) { + const char *want = name; + uint16_t wtu = BIR_TU_EXT; + + if (pass == 0) { + if (tu == BIR_TU_EXT) continue; + if (bir_mang(name, tu, mng, (int)sizeof mng) != 0) continue; + want = mng; + wtu = tu; + } + for (uint32_t i = 0; i < M->num_funcs; i++) { + const bir_func_t *F = &M->funcs[i]; + if (F->tu != wtu || F->name >= M->string_len) continue; + if (strcmp(&M->strings[F->name], want) != 0) continue; + if (nargs < 0 || F->num_params == (uint16_t)nargs) return i; + } + } + return BIR_SYM_NONE; +} + +uint32_t bir_gsym(const bir_module_t *M, const char *name, uint16_t tu) +{ + char mng[BIR_SYM_MAX]; + + if (tu != BIR_TU_EXT && bir_mang(name, tu, mng, (int)sizeof mng) == 0) { + for (uint32_t i = 0; i < M->num_globals; i++) { + const bir_global_t *G = &M->globals[i]; + if (G->tu != tu || G->name >= M->string_len) continue; + if (strcmp(&M->strings[G->name], mng) == 0) return i; + } + } + for (uint32_t i = 0; i < M->num_globals; i++) { + const bir_global_t *G = &M->globals[i]; + if (G->tu != BIR_TU_EXT || G->name >= M->string_len) continue; + if (strcmp(&M->strings[G->name], name) == 0) return i; + } + return BIR_SYM_NONE; +} + uint32_t bir_const_float(bir_module_t *M, uint32_t type, double val) { uint32_t guard = M->num_consts; diff --git a/src/ir/bir.h b/src/ir/bir.h index 302bf9c..c243baf 100644 --- a/src/ir/bir.h +++ b/src/ir/bir.h @@ -20,6 +20,7 @@ #define BIR_MAX_FUNCS (1 << 12) #define BIR_MAX_GLOBALS (1 << 12) #define BIR_MAX_STRINGS (1 << 20) +#define BIR_BSZ_DEEP 16 /* aggregate nesting bir_bsz will walk */ /* ---- Value References ---- */ @@ -168,6 +169,10 @@ typedef enum { /* Matrix */ BIR_MFMA, /* subop = variant ID, ops: [0]=A, [1]=B, [2]=C(accum) */ + BIR_MMA, /* warp-collective D += A*B over a 16x16x16 f16 tile. + * ops: [0]=A [1]=lda [2]=B [3]=ldb [4]=D [5]=ldd */ + BIR_MFRG, /* MFMA over per-lane fragments in memory. + * subop = variant, ops: [0]=A [1]=B [2]=C/D */ /* Misc */ BIR_CALL, /* ops[0] = callee func index, rest = args */ @@ -258,7 +263,7 @@ typedef struct { uint16_t num_blocks; uint16_t num_params; uint16_t cuda_flags; /* CUDA_GLOBAL, CUDA_DEVICE, CUDA_HOST from ast.h */ - uint16_t pad; + uint16_t tu; /* owning TU for internal linkage, BIR_TU_EXT for external */ uint32_t launch_bounds_max; uint32_t launch_bounds_min; } bir_func_t; /* 32 bytes */ @@ -272,7 +277,14 @@ typedef struct { uint16_t cuda_flags; uint8_t addrspace; /* bir_addrspace_t */ uint8_t is_const; -} bir_global_t; /* 16 bytes */ + uint16_t tu; /* owning TU for internal linkage, BIR_TU_EXT for external */ + uint16_t pad; +} bir_global_t; /* 20 bytes */ + + +#define BIR_TU_EXT 0xFFFFu +#define BIR_SYM_NONE 0xFFFFFFFFu +#define BIR_SYM_MAX 160 /* room for a 128-char name plus the suffix */ /* ---- Pool overflow ---- */ @@ -343,6 +355,10 @@ uint32_t bir_type_struct(bir_module_t *M, const uint32_t *fields, int nfields uint32_t bir_type_func(bir_module_t *M, uint32_t ret, const uint32_t *params, int nparams); +uint32_t bir_bsz(const bir_module_t *M, uint32_t ty, uint32_t psz); + +uint32_t bir_gsz(const bir_module_t *M, uint32_t ty, uint32_t psz); + /* String table */ uint32_t bir_add_string(bir_module_t *M, const char *s, uint32_t len); @@ -359,6 +375,12 @@ uint32_t bir_const_bytes(bir_module_t *M, uint32_t type, * silent wrong output. */ int bir_global_is_bytes(const bir_module_t *M, uint32_t gi); +int bir_mang(const char *name, uint16_t tu, char *out, int size); + +uint32_t bir_fsym(const bir_module_t *M, const char *name, uint16_t tu, + int nargs); +uint32_t bir_gsym(const bir_module_t *M, const char *name, uint16_t tu); + /* Name tables */ const char *bir_op_name(int op); const char *bir_type_kind_name(int kind); diff --git a/src/ir/bir_lower.c b/src/ir/bir_lower.c index 11b838f..5ba8659 100644 --- a/src/ir/bir_lower.c +++ b/src/ir/bir_lower.c @@ -24,6 +24,9 @@ #define MAX_ENUMS 256 #define MAX_TYPEDEFS 64 #define MAX_TEMPLATES 32 +#define MAX_PKELM 16 +#define MAX_FPACKS 4 +#define MAX_INSTS 16 #define MAX_SCOPES 64 #define MAX_LOOPS 32 @@ -64,8 +67,16 @@ typedef struct { uint32_t type; /* BIR type index */ int64_t ival; /* for non-type params */ int is_type; /* 1 = typename, 0 = int param */ + int is_pack; + int npk; + uint32_t pk_t[MAX_PKELM]; } binding_t; +typedef struct { + char name[64]; + int n; +} fpack_t; + typedef struct { const parser_t *P; bir_module_t *M; @@ -87,6 +98,12 @@ typedef struct { int ntemplates; binding_t bindings[8]; int nbindings; + struct { char key[256]; char sym[128]; } inst[MAX_INSTS]; + int ninst; + fpack_t fpk[MAX_FPACKS]; + int nfpk; + struct { char name[64]; int idx; } pkact[MAX_FPACKS]; + int npkact; /* Current function state */ uint32_t cur_func; @@ -107,12 +124,25 @@ typedef struct { const sema_ctx_t *sema; /* NULL if sema didn't run */ + uint16_t tu; + uint32_t cur_node; /* AST node being lowered — for source loc tracking */ bc_error_t errors[BC_MAX_ERRORS]; int nerrors; } lower_t; +/* ---- Names ---- */ + +static void ncpy(char *d, size_t n, const char *s) +{ + size_t k = strlen(s); + + if (k >= n) k = n - 1; + memcpy(d, s, k); + d[k] = '\0'; +} + /* ---- AST Navigation ---- */ static const ast_node_t *ND(const lower_t *L, uint32_t i) @@ -171,6 +201,25 @@ static void lower_error(lower_t *L, uint32_t node, bc_eid_t eid, ...) } } + +static uint32_t sym_fx(const lower_t *L, const char *name) +{ + for (uint32_t i = 0; i < L->M->num_funcs; i++) + if (L->M->funcs[i].name < L->M->string_len + && strcmp(&L->M->strings[L->M->funcs[i].name], name) == 0) + return i; + return BIR_SYM_NONE; +} + +static uint32_t sym_gx(const lower_t *L, const char *name) +{ + for (uint32_t i = 0; i < L->M->num_globals; i++) + if (L->M->globals[i].name < L->M->string_len + && strcmp(&L->M->strings[L->M->globals[i].name], name) == 0) + return i; + return BIR_SYM_NONE; +} + /* ---- Scope ---- */ static void push_scope(lower_t *L) @@ -190,7 +239,7 @@ static void add_sym(lower_t *L, const char *name, uint32_t ref, { if (L->nsyms >= MAX_SYMS) return; sym_t *s = &L->syms[L->nsyms++]; - snprintf(s->name, sizeof(s->name), "%s", name); + ncpy(s->name, sizeof(s->name), name); s->ref = ref; s->type = type; s->is_alloca = is_alloca; @@ -229,7 +278,7 @@ static int find_typedef(lower_t *L, const char *name, uint32_t *type) static int find_binding(lower_t *L, const char *name, uint32_t *type) { for (int i = 0; i < L->nbindings; i++) - if (L->bindings[i].is_type + if (L->bindings[i].is_type && !L->bindings[i].is_pack && strcmp(L->bindings[i].name, name) == 0) { *type = L->bindings[i].type; return 1; @@ -240,7 +289,7 @@ static int find_binding(lower_t *L, const char *name, uint32_t *type) static int find_binding_int(lower_t *L, const char *name, int64_t *val) { for (int i = 0; i < L->nbindings; i++) - if (!L->bindings[i].is_type + if (!L->bindings[i].is_type && !L->bindings[i].is_pack && strcmp(L->bindings[i].name, name) == 0) { *val = L->bindings[i].ival; return 1; @@ -248,6 +297,47 @@ static int find_binding_int(lower_t *L, const char *name, int64_t *val) return 0; } +static void pk_nm(char *out, size_t n, const char *base, int i) +{ + snprintf(out, n, "%s#%d", base, i); +} + +static const fpack_t *find_fpk(const lower_t *L, const char *name) +{ + for (int i = 0; i < L->nfpk; i++) + if (strcmp(L->fpk[i].name, name) == 0) return &L->fpk[i]; + return NULL; +} + +static const binding_t *find_tpk(const lower_t *L, const char *name) +{ + for (int i = 0; i < L->nbindings; i++) + if (L->bindings[i].is_pack + && strcmp(L->bindings[i].name, name) == 0) + return &L->bindings[i]; + return NULL; +} + +static int pk_cnt(const lower_t *L, const char *name) +{ + const fpack_t *f = find_fpk(L, name); + if (f) return f->n; + const binding_t *b = find_tpk(L, name); + if (b) return b->npk; + return -1; +} + +static void pk_rw(const lower_t *L, char *name, size_t n) +{ + for (int i = L->npkact - 1; i >= 0; i--) + if (strcmp(L->pkact[i].name, name) == 0) { + char tmp[128]; + pk_nm(tmp, sizeof tmp, name, L->pkact[i].idx); + snprintf(name, n, "%s", tmp); + return; + } +} + static template_def_t *find_template(lower_t *L, const char *name) { for (int i = 0; i < L->ntemplates; i++) @@ -565,6 +655,13 @@ static uint32_t coerce_to(lower_t *L, uint32_t val, uint32_t dst_t, return BIR_MAKE_VAL(inst); } +static int is_i1(const lower_t *L, uint32_t t) +{ + return t < L->M->num_types + && L->M->types[t].kind == BIR_TYPE_INT + && L->M->types[t].width == 1; +} + static int is_ptr_type(const lower_t *L, uint32_t t) { return t < L->M->num_types && L->M->types[t].kind == BIR_TYPE_PTR; @@ -707,7 +804,222 @@ static int is_compound_assign(int tok) /* ---- Forward Declarations ---- */ +static uint32_t bin_val(lower_t *L, uint32_t node, int op, + uint32_t lhs, uint32_t lhs_n, + uint32_t rhs, uint32_t rhs_n) +{ + uint32_t lt = ref_type(L, lhs), rt = ref_type(L, rhs); + + uint32_t t32 = bir_type_int(L->M, 32); + if (is_i1(L, lt)) { lhs = coerce_to(L, lhs, t32, 1); lt = t32; } + if (is_i1(L, rt)) { rhs = coerce_to(L, rhs, t32, 1); rt = t32; } + + uint32_t res_t = lt; + int lf = is_float_type(L, lt), rf = is_float_type(L, rt); + if (lf && !rf) { + rhs = coerce_to(L, rhs, lt, node_is_unsigned(L, rhs_n)); + res_t = lt; + } else if (!lf && rf) { + lhs = coerce_to(L, lhs, rt, node_is_unsigned(L, lhs_n)); + res_t = rt; + } else if (lf && rf) { + if (lt < L->M->num_types && rt < L->M->num_types + && L->M->types[rt].width > L->M->types[lt].width) { + lhs = coerce_to(L, lhs, rt, 0); + res_t = rt; + } else if (lt < L->M->num_types && rt < L->M->num_types + && L->M->types[lt].width > L->M->types[rt].width) { + rhs = coerce_to(L, rhs, lt, 0); + res_t = lt; + } + } + int fp = is_float_type(L, res_t); + int opc = bin_op_code(op, fp, node_is_unsigned(L, node)); + if (opc < 0) { + lower_error(L, node, BC_E102); + return lhs; + } + uint32_t inst = emit(L, (uint16_t)opc, res_t, 2, 0); + set_op(L, inst, 0, lhs); + set_op(L, inst, 1, rhs); + return BIR_MAKE_VAL(inst); +} + static uint32_t lower_expr(lower_t *L, uint32_t node); + +/* ---- Pack expansion ---- */ + +static int pk_scan(const lower_t *L, uint32_t node, + char out[][64], int max, int nfound, int depth) +{ + if (!node || depth > 64 || nfound >= max) return nfound; + const ast_node_t *n = ND(L, node); + if (n->type == AST_IDENT) { + char nm[64]; + get_text(L, node, nm, sizeof(nm)); + if (pk_cnt(L, nm) >= 0) { + for (int i = 0; i < nfound; i++) + if (strcmp(out[i], nm) == 0) return nfound; + snprintf(out[nfound++], 64, "%s", nm); + return nfound; + } + } + for (uint32_t c = n->first_child; c; c = ND(L, c)->next_sibling) + nfound = pk_scan(L, c, out, max, nfound, depth + 1); + return nfound; +} + +static int pk_len(lower_t *L, uint32_t pat, char nms[][64], int *nn) +{ + *nn = pk_scan(L, pat, nms, MAX_FPACKS, 0, 0); + if (*nn == 0) return -1; + int len = pk_cnt(L, nms[0]); + for (int i = 1; i < *nn; i++) + if (pk_cnt(L, nms[i]) != len) { + lower_error(L, pat, BC_E030, "pack lengths differ in one pattern"); + return -1; + } + return len; +} + +static uint32_t pk_at(lower_t *L, uint32_t pat, char nms[][64], int nn, int i) +{ + int sv = L->npkact; + for (int k = 0; k < nn && L->npkact < MAX_FPACKS; k++) { + ncpy(L->pkact[L->npkact].name, + sizeof(L->pkact[0].name), nms[k]); + L->pkact[L->npkact].idx = i; + L->npkact++; + } + uint32_t v = lower_expr(L, pat); + L->npkact = sv; + return v; +} + +static int pk_exp(lower_t *L, uint32_t node, uint32_t *out, int max) +{ + uint32_t pat = ND(L, node)->first_child; + char nms[MAX_FPACKS][64]; + int nn = 0; + int len = pk_len(L, pat, nms, &nn); + if (len < 0) { + lower_error(L, node, BC_E030, "pack expansion over an unbound pack"); + return -1; + } + if (len > max) { + lower_error(L, node, BC_E030, "pack longer than the argument list"); + return -1; + } + for (int i = 0; i < len; i++) + out[i] = pk_at(L, pat, nms, nn, i); + return len; +} + +static uint32_t lfold(lower_t *L, uint32_t node) +{ + const ast_node_t *n = ND(L, node); + int op = n->d.oper.op; + int form = n->d.oper.flags; + uint32_t a = n->first_child; + uint32_t b = a ? ND(L, a)->next_sibling : 0; + uint32_t pat, init = 0; + int init_1st = 0; + + switch (form) { + case FLD_UL: case FLD_UR: pat = a; break; + case FLD_BL: pat = b; init = a; init_1st = 1; break; + case FLD_BR: pat = a; init = b; break; + default: lower_error(L, node, BC_E106); return BIR_VAL_NONE; + } + + char nms[MAX_FPACKS][64]; + int nn = 0; + int len = pk_len(L, pat, nms, &nn); + if (len < 0) { + lower_error(L, node, BC_E030, "fold over an unbound pack"); + return BIR_VAL_NONE; + } + + uint32_t t1 = bir_type_int(L->M, 1); + if (op == TOK_LAND || op == TOK_LOR) { + int is_and = (op == TOK_LAND); + uint32_t pt = bir_type_ptr(L->M, t1, BIR_AS_PRIVATE); + uint32_t al = emit(L, BIR_ALLOCA, pt, 0, 0); + uint32_t seed = BIR_MAKE_CONST(bir_const_int(L->M, t1, is_and ? 0 : 1)); + uint32_t s0 = emit(L, BIR_STORE, bir_type_void(L->M), 2, 0); + set_op(L, s0, 0, seed); + set_op(L, s0, 1, BIR_MAKE_VAL(al)); + uint32_t end_b = new_block(L, "fold.end"); + int m = len + (init ? 1 : 0); + for (int k = 0; k < m; k++) { + uint32_t v; + if (init && ((init_1st && k == 0) || (!init_1st && k == m - 1))) + v = lower_expr(L, init); + else + v = pk_at(L, pat, nms, nn, init_1st && init ? k - 1 : k); + uint32_t nxt = new_block(L, "fold.on"); + uint32_t br = emit(L, BIR_BR_COND, bir_type_void(L->M), 4, 0); + set_op(L, br, 0, v); + set_op(L, br, 1, is_and ? nxt : end_b); + set_op(L, br, 2, is_and ? end_b : nxt); + set_op(L, br, 3, end_b); + set_block(L, nxt); + } + uint32_t done = BIR_MAKE_CONST(bir_const_int(L->M, t1, is_and ? 1 : 0)); + uint32_t s1 = emit(L, BIR_STORE, bir_type_void(L->M), 2, 0); + set_op(L, s1, 0, done); + set_op(L, s1, 1, BIR_MAKE_VAL(al)); + uint32_t j = emit(L, BIR_BR, bir_type_void(L->M), 1, 0); + set_op(L, j, 0, end_b); + set_block(L, end_b); + uint32_t ld = emit(L, BIR_LOAD, t1, 1, 0); + set_op(L, ld, 0, BIR_MAKE_VAL(al)); + return BIR_MAKE_VAL(ld); + } + + if (op == TOK_COMMA) { + uint32_t last = BIR_VAL_NONE; + int m = len + (init ? 1 : 0); + for (int k = 0; k < m; k++) { + if (init && ((init_1st && k == 0) || (!init_1st && k == m - 1))) + last = lower_expr(L, init); + else + last = pk_at(L, pat, nms, nn, init_1st && init ? k - 1 : k); + } + return last; + } + + if (len == 0) { + if (init) return lower_expr(L, init); + lower_error(L, node, BC_E030, "empty pack folded over this operator"); + return BIR_VAL_NONE; + } + if (len > MAX_PKELM) { + lower_error(L, node, BC_E030, "fold longer than the pack limit"); + return BIR_VAL_NONE; + } + + uint32_t v[MAX_PKELM]; + for (int i = 0; i < len; i++) + v[i] = pk_at(L, pat, nms, nn, i); + + if (form == FLD_UL || form == FLD_BL) { + uint32_t acc = init ? lower_expr(L, init) : v[0]; + uint32_t acc_n = init ? init : pat; + for (int i = init ? 0 : 1; i < len; i++) { + acc = bin_val(L, node, op, acc, acc_n, v[i], pat); + acc_n = pat; + } + return acc; + } + uint32_t acc = init ? lower_expr(L, init) : v[len - 1]; + uint32_t acc_n = init ? init : pat; + for (int i = init ? len - 1 : len - 2; i >= 0; i--) { + acc = bin_val(L, node, op, v[i], pat, acc, acc_n); + acc_n = pat; + } + return acc; +} static uint32_t lower_lvalue(lower_t *L, uint32_t node); static void lower_stmt(lower_t *L, uint32_t node); static void lower_block_stmts(lower_t *L, uint32_t node); @@ -796,6 +1108,7 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) case AST_IDENT: { char name[128]; get_text(L, node, name, sizeof(name)); + pk_rw(L, name, sizeof(name)); /* Builtin constant: warpSize (HIP) */ if (strcmp(name, "warpSize") == 0 && L->sema) { @@ -819,23 +1132,23 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) sym_t *s = find_sym(L, name); if (!s) { /* Check file-scope globals (__shared__, __device__, __constant__) */ - for (uint32_t gi = 0; gi < L->M->num_globals; gi++) { + uint32_t gi = bir_gsym(L->M, name, L->tu); + if (gi != BIR_SYM_NONE) { bir_global_t *G = &L->M->globals[gi]; - if (G->name < L->M->string_len - && strcmp(&L->M->strings[G->name], name) == 0) { - int adrspc = G->addrspace; - uint32_t ptr_t = bir_type_ptr(L->M, G->type, adrspc); - if (G->cuda_flags & CUDA_SHARED) { - uint32_t sa = emit(L, BIR_SHARED_ALLOC, ptr_t, 0, 0); - add_sym(L, name, sa, G->type, 1); - } else { - uint32_t gr = emit(L, BIR_GLOBAL_REF, ptr_t, 0, - (uint8_t)gi); - add_sym(L, name, gr, G->type, 1); - } - s = find_sym(L, name); - break; + int adrspc = G->addrspace; + uint32_t ptr_t = bir_type_ptr(L->M, G->type, adrspc); + if (G->cuda_flags & CUDA_SHARED) { + uint32_t sa = emit(L, BIR_SHARED_ALLOC, ptr_t, 0, 0); + add_sym(L, name, sa, G->type, 1); + } else if (gi > 0xFFu) { + lower_error(L, node, BC_E127, 256); + return BIR_VAL_NONE; + } else { + uint32_t gr = emit(L, BIR_GLOBAL_REF, ptr_t, 0, + (uint8_t)gi); + add_sym(L, name, gr, G->type, 1); } + s = find_sym(L, name); } } if (!s) { @@ -1132,19 +1445,15 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) && L->M->types[lt].kind == BIR_TYPE_STRUCT) { char oname[32]; op_name_from_tok(op, oname, sizeof(oname)); - /* Look up operator function in module */ - for (uint32_t fi = 0; fi < L->M->num_funcs; fi++) { - if (L->M->funcs[fi].name < L->M->string_len - && strcmp(&L->M->strings[L->M->funcs[fi].name], - oname) == 0) { - uint32_t ftype = L->M->funcs[fi].type; - uint32_t ret_t = L->M->types[ftype].inner; - uint32_t inst = emit(L, BIR_CALL, ret_t, 3, 0); - set_op(L, inst, 0, fi); - set_op(L, inst, 1, lhs); - set_op(L, inst, 2, rhs); - return BIR_MAKE_VAL(inst); - } + uint32_t fi = bir_fsym(L->M, oname, L->tu, 2); + if (fi != BIR_SYM_NONE) { + uint32_t ftype = L->M->funcs[fi].type; + uint32_t ret_t = L->M->types[ftype].inner; + uint32_t inst = emit(L, BIR_CALL, ret_t, 3, 0); + set_op(L, inst, 0, fi); + set_op(L, inst, 1, lhs); + set_op(L, inst, 2, rhs); + return BIR_MAKE_VAL(inst); } } @@ -1176,41 +1485,7 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) } } - /* Usual arithmetic conversion: promote both operands - * to the wider/float type. C says int*double → double, - * not int*double → garbage. Without this, backends get - * mixed-type ops (mul.u32 with an f64 register) and - * the PTX JIT has strong opinions about that. */ - uint32_t res_t = lt; - int lf = is_float_type(L, lt), rf = is_float_type(L, rt); - if (lf && !rf) { - rhs = coerce_to(L, rhs, lt, node_is_unsigned(L, rhs_n)); - res_t = lt; - } else if (!lf && rf) { - lhs = coerce_to(L, lhs, rt, node_is_unsigned(L, lhs_n)); - res_t = rt; - } else if (lf && rf) { - /* Both float: promote narrower to wider */ - if (lt < L->M->num_types && rt < L->M->num_types - && L->M->types[rt].width > L->M->types[lt].width) { - lhs = coerce_to(L, lhs, rt, 0); - res_t = rt; - } else if (lt < L->M->num_types && rt < L->M->num_types - && L->M->types[lt].width > L->M->types[rt].width) { - rhs = coerce_to(L, rhs, lt, 0); - res_t = lt; - } - } - int fp = is_float_type(L, res_t); - int opc = bin_op_code(op, fp, node_is_unsigned(L, node)); - if (opc < 0) { - lower_error(L, node, BC_E102); - return lhs; - } - uint32_t inst = emit(L, (uint16_t)opc, res_t, 2, 0); - set_op(L, inst, 0, lhs); - set_op(L, inst, 1, rhs); - return BIR_MAKE_VAL(inst); + return bin_val(L, node, op, lhs, lhs_n, rhs, rhs_n); } } @@ -1883,6 +2158,64 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) } } + /* ---- Warp-collective 16x16x16 f16 matrix multiply ---- */ + if (strncmp(cname, "__builtin_mma_", 14) == 0) { + static const char *const mma_tab[] = { + "m16n16k16_f16", "m16n16k16_bf16", + "m16n16k8_f16", "m16n16k8_bf16", + }; + const char *sfx = cname + 14; + for (int mi = 0; mi < (int)(sizeof mma_tab / sizeof mma_tab[0]); mi++) { + if (strcmp(sfx, mma_tab[mi]) != 0) continue; + uint32_t an = ND(L, callee_n)->next_sibling; + uint32_t ops[6]; + int na = 0; + while (an != 0 && na < 6) { + ops[na++] = lower_expr(L, an); + an = ND(L, an)->next_sibling; + } + if (na != 6) return BIR_VAL_NONE; + uint32_t r = emit(L, BIR_MMA, bir_type_void(L->M), 6, + (uint8_t)mi); + for (int k = 0; k < 6; k++) set_op(L, r, k, ops[k]); + return BIR_MAKE_VAL(r); + } + return BIR_VAL_NONE; + } + + /* ---- MFMA over per-lane fragments in memory ---- */ + if (strncmp(cname, "__builtin_mfma_", 15) == 0) { + static const char *const mfrg_tab[] = { + "f32_4x4x4_f16", "f32_16x16x16_f16", "f32_32x32x8_f16", + "f32_4x4x4_bf16", "f32_16x16x16_bf16", "f32_32x32x8_bf16", + "f32_4x4x1_f32", "f32_16x16x4_f32", "f32_32x32x2_f32", + "i32_4x4x4_i8", "i32_16x16x16_i8", "i32_32x32x8_i8", + "f32_16x16x32_fp8_fp8", "f32_16x16x32_fp8_bf8", + "f32_16x16x32_bf8_fp8", "f32_16x16x32_bf8_bf8", + "f32_32x32x16_fp8_fp8", "f32_32x32x16_fp8_bf8", + "f32_32x32x16_bf8_fp8", "f32_32x32x16_bf8_bf8", + "f64_4x4x4_f64", "f64_16x16x4_f64", + "i32_16x16x32_i8", "i32_32x32x16_i8", + }; + const char *sfx = cname + 15; + for (int mi = 0; mi < (int)(sizeof mfrg_tab / sizeof mfrg_tab[0]); mi++) { + if (strcmp(sfx, mfrg_tab[mi]) != 0) continue; + uint32_t an = ND(L, callee_n)->next_sibling; + uint32_t ops[3]; + int na = 0; + while (an != 0 && na < 3) { + ops[na++] = lower_expr(L, an); + an = ND(L, an)->next_sibling; + } + if (na != 3) return BIR_VAL_NONE; + uint32_t r = emit(L, BIR_MFRG, bir_type_void(L->M), 3, + (uint8_t)mi); + for (int k = 0; k < 3; k++) set_op(L, r, k, ops[k]); + return BIR_MAKE_VAL(r); + } + return BIR_VAL_NONE; + } + /* ---- MFMA intrinsics (CDNA matrix multiply) ---- */ if (strncmp(cname, "__builtin_amdgcn_mfma_", 22) == 0) { static const struct { const char *sfx; uint8_t var; } mfma_tab[] = { @@ -2009,31 +2342,18 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) /* ---- Regular function call ---- */ - /* Find function in module */ - uint32_t fi = 0; - int found = 0; - for (uint32_t i = 0; i < L->M->num_funcs; i++) { - if (L->M->funcs[i].name < L->M->string_len - && strcmp(&L->M->strings[L->M->funcs[i].name], cname) == 0) { - fi = i; - found = 1; - break; - } - } - if (!found) { - lower_error(L, node, BC_E105); - return BIR_VAL_NONE; - } - - uint32_t ftype = L->M->funcs[fi].type; - uint32_t ret_t = L->M->types[ftype].inner; - /* Lower arguments */ uint32_t args[BC_MAX_ARGS]; int nargs = 0; uint32_t arg = ND(L, callee_n)->next_sibling; while (arg && nargs < BC_MAX_ARGS) { - args[nargs++] = lower_expr(L, arg); + if (ND(L, arg)->type == AST_PACK_EXP) { + int got = pk_exp(L, arg, args + nargs, BC_MAX_ARGS - nargs); + if (got < 0) return BIR_VAL_NONE; + nargs += got; + } else { + args[nargs++] = lower_expr(L, arg); + } arg = ND(L, arg)->next_sibling; } /* Sema rejects this first, so reaching it means the two caps have @@ -2044,6 +2364,20 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) return BIR_VAL_NONE; } + uint32_t fi = bir_fsym(L->M, cname, L->tu, nargs); + if (fi == BIR_SYM_NONE) { + uint32_t any = bir_fsym(L->M, cname, L->tu, -1); + if (any == BIR_SYM_NONE) + lower_error(L, node, BC_E105); + else + lower_error(L, node, BC_E073, cname, + (int)L->M->funcs[any].num_params, nargs); + return BIR_VAL_NONE; + } + + uint32_t ftype = L->M->funcs[fi].type; + uint32_t ret_t = L->M->types[ftype].inner; + if (1 + nargs <= BIR_OPERANDS_INLINE) { uint32_t inst = emit(L, BIR_CALL, ret_t, (uint8_t)(1+nargs), 0); set_op(L, inst, 0, fi); @@ -2144,6 +2478,7 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) int64_t sz = 4; if (inner && ND(L, inner)->type == AST_TYPE_SPEC) { switch (ND(L, inner)->d.btype.kind) { + case TYPE_BOOL: sz = 1; break; case TYPE_CHAR: sz = 1; break; case TYPE_SHORT: sz = 2; break; case TYPE_LONG: case TYPE_LLONG: sz = 8; break; @@ -2232,8 +2567,14 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) lower_error(L, node, BC_E106); return BIR_VAL_NONE; } - uint32_t gi = L->M->num_globals++; + uint32_t gi = L->M->num_globals; + if (gi > 0xFFu) { + lower_error(L, node, BC_E127, 256); + return BIR_VAL_NONE; + } + L->M->num_globals++; bir_global_t *G = &L->M->globals[gi]; + memset(G, 0, sizeof(*G)); char gname[24]; snprintf(gname, sizeof(gname), ".str.%u", gi); G->name = bir_add_string(L->M, gname, (uint32_t)strlen(gname)); @@ -2243,11 +2584,33 @@ static uint32_t lower_expr(lower_t *L, uint32_t node) G->cuda_flags = CUDA_CONSTANT; G->addrspace = BIR_AS_CONSTANT; G->is_const = 1; + G->tu = BIR_TU_EXT; - uint32_t inst = emit(L, BIR_GLOBAL_REF, ptr, 0, (uint8_t)(gi & 0xff)); + uint32_t inst = emit(L, BIR_GLOBAL_REF, ptr, 0, (uint8_t)gi); return BIR_MAKE_VAL(inst); } + case AST_PACK_SIZE: { + uint32_t id = n->first_child; + char nm[64]; + if (id) get_text(L, id, nm, sizeof(nm)); else nm[0] = 0; + int cnt = pk_cnt(L, nm); + if (cnt < 0) { + lower_error(L, node, BC_E030, "sizeof... of an unbound pack"); + return BIR_VAL_NONE; + } + return BIR_MAKE_CONST(bir_const_int(L->M, + bir_type_int(L->M, 32), cnt)); + } + + case AST_FOLD: + return lfold(L, node); + + case AST_PACK_EXP: + lower_error(L, node, BC_E030, + "pack expansion outside an argument list"); + return BIR_VAL_NONE; + default: lower_error(L, node, BC_E106); return BIR_VAL_NONE; @@ -2265,6 +2628,7 @@ static uint32_t lower_lvalue(lower_t *L, uint32_t node) case AST_IDENT: { char name[128]; get_text(L, node, name, sizeof(name)); + pk_rw(L, name, sizeof(name)); sym_t *s = find_sym(L, name); if (!s) { lower_error(L, node, BC_E107); @@ -2993,24 +3357,94 @@ static void lower_func_body(lower_t *L, uint32_t func_def, memcpy(fname, fname_raw, sizeof(fname)); } + uint16_t flnk = BIR_TU_EXT; + uint16_t quals = ND(L, func_def)->qualifiers; + if (L->tu != BIR_TU_EXT && (quals & QUAL_STATIC)) { + char mng[BIR_SYM_MAX]; + if (bir_mang(fname, L->tu, mng, (int)sizeof fname) != 0) { + lower_error(L, func_def, BC_E126, fname); + return; + } + memcpy(fname, mng, strlen(mng) + 1u); + flnk = L->tu; + } + + if (L->tu != BIR_TU_EXT) { + uint32_t dup = sym_fx(L, fname); + if (dup != BIR_SYM_NONE) { + if (!(quals & QUAL_INLINE) && !name_override) + lower_error(L, func_def, BC_E126, fname); + return; + } + } + uint32_t ret_t = resolve_type(L, type_n, ret_ptr, 0); /* Collect parameters */ uint32_t param_nodes[32]; int nparams = collect_params(L, func_def, param_nodes, 32); - /* Resolve param types and build function type */ uint32_t param_types[32]; - for (int i = 0; i < nparams; i++) { + char pnames[32][64]; + int np = 0; + L->nfpk = 0; + for (int i = 0; i < nparams && np < 32; i++) { const ast_node_t *pn = ND(L, param_nodes[i]); + if (pn->d.oper.op == PRM_VARG) continue; uint32_t pt_type_n = pn->first_child; int pdepth = pn->d.oper.flags; /* stored by parser */ uint16_t p_cuda = pn->cuda_flags; - param_types[i] = resolve_type(L, pt_type_n, pdepth, p_cuda); + + char base[64]; + base[0] = 0; + for (uint32_t pc = pn->first_child; pc; pc = ND(L, pc)->next_sibling) + if (ND(L, pc)->type == AST_IDENT) { + get_text(L, pc, base, sizeof(base)); + break; + } + + if (pn->d.oper.op == PRM_PACK) { + char tn[64]; + tn[0] = 0; + if (pt_type_n && ND(L, pt_type_n)->first_child) + get_text(L, ND(L, pt_type_n)->first_child, tn, sizeof(tn)); + const binding_t *b = find_tpk(L, tn); + if (!b) { + lower_error(L, param_nodes[i], BC_E030, + "function parameter pack with no bound types"); + return; + } + if (base[0] && L->nfpk < MAX_FPACKS) { + snprintf(L->fpk[L->nfpk].name, + sizeof(L->fpk[0].name), "%s", base); + L->fpk[L->nfpk].n = b->npk; + L->nfpk++; + } + for (int k = 0; k < b->npk && np < 32; k++) { + int svb = L->nbindings; + if (L->nbindings < 8) { + binding_t *t = &L->bindings[L->nbindings++]; + memset(t, 0, sizeof(*t)); + snprintf(t->name, sizeof(t->name), "%s", tn); + t->is_type = 1; + t->type = b->pk_t[k]; + } + param_types[np] = resolve_type(L, pt_type_n, pdepth, p_cuda); + L->nbindings = svb; + pk_nm(pnames[np], sizeof(pnames[0]), base, k); + np++; + } + continue; + } + + param_types[np] = resolve_type(L, pt_type_n, pdepth, p_cuda); + snprintf(pnames[np], sizeof(pnames[0]), "%s", base); + np++; } + int nparams_x = np; /* Create function type */ - uint32_t fn_type = bir_type_func(L->M, ret_t, param_types, nparams); + uint32_t fn_type = bir_type_func(L->M, ret_t, param_types, nparams_x); /* Create function. Bailing leaves cur_func on the previous one. */ if (L->M->num_funcs >= BIR_MAX_FUNCS) { @@ -3024,8 +3458,9 @@ static void lower_func_body(lower_t *L, uint32_t func_def, memset(F, 0, sizeof(*F)); F->name = bir_add_string(L->M, fname, (uint32_t)strlen(fname)); F->type = fn_type; + F->tu = flnk; F->cuda_flags = cuda_flags; - F->num_params = (uint16_t)nparams; + F->num_params = (uint16_t)nparams_x; F->first_block = L->M->num_blocks; F->num_blocks = 0; F->launch_bounds_max = ND(L, func_def)->launch_bounds_max; @@ -3041,18 +3476,9 @@ static void lower_func_body(lower_t *L, uint32_t func_def, /* Emit PARAM instructions */ push_scope(L); - for (int i = 0; i < nparams; i++) { + for (int i = 0; i < nparams_x; i++) { uint32_t inst = emit(L, BIR_PARAM, param_types[i], 0, (uint8_t)i); - /* Get param name */ - uint32_t pname_n = 0; - uint32_t pc = ND(L, param_nodes[i])->first_child; - while (pc) { - if (ND(L, pc)->type == AST_IDENT) { pname_n = pc; break; } - pc = ND(L, pc)->next_sibling; - } - if (pname_n) { - char pname[128]; - get_text(L, pname_n, pname, sizeof(pname)); + if (pnames[i][0]) { /* Promote all params to allocas so they're reassignable. mem2reg cleans up the ones that are never written. */ uint32_t pt = bir_type_ptr(L->M, param_types[i], BIR_AS_PRIVATE); @@ -3060,7 +3486,7 @@ static void lower_func_body(lower_t *L, uint32_t func_def, uint32_t st = emit(L, BIR_STORE, bir_type_void(L->M), 2, 0); set_op(L, st, 0, BIR_MAKE_VAL(inst)); set_op(L, st, 1, BIR_MAKE_VAL(al)); - add_sym(L, pname, al, param_types[i], 1); + add_sym(L, pnames[i], al, param_types[i], 1); } } @@ -3250,10 +3676,6 @@ static void collect_global_var(lower_t *L, uint32_t node) { uint16_t cuda = ND(L, node)->cuda_flags; if (!(cuda & (CUDA_SHARED | CUDA_CONSTANT | CUDA_DEVICE))) return; - if (L->M->num_globals >= BIR_MAX_GLOBALS) { - bir_pfull(L->M, BIR_P_GLOBALS); - return; - } uint32_t type_n = child_at(L, node, 0); uint32_t name_n = child_at(L, node, 1); @@ -3283,10 +3705,34 @@ static void collect_global_var(lower_t *L, uint32_t node) } } + uint16_t glnk = BIR_TU_EXT; + if (L->tu != BIR_TU_EXT && (ND(L, node)->qualifiers & QUAL_STATIC)) { + char mng[BIR_SYM_MAX]; + if (bir_mang(gname, L->tu, mng, (int)sizeof gname) != 0) { + lower_error(L, node, BC_E126, gname); + return; + } + memcpy(gname, mng, strlen(mng) + 1u); + glnk = L->tu; + } + + uint32_t old = sym_gx(L, gname); + if (old != BIR_SYM_NONE) { + if (L->M->globals[old].type != elem_t) + lower_error(L, node, BC_E126, gname); + return; + } + + if (L->M->num_globals >= BIR_MAX_GLOBALS) { + bir_pfull(L->M, BIR_P_GLOBALS); + return; + } uint32_t gi = L->M->num_globals++; bir_global_t *G = &L->M->globals[gi]; + memset(G, 0, sizeof(*G)); G->name = bir_add_string(L->M, gname, (uint32_t)strlen(gname)); G->type = elem_t; + G->tu = glnk; G->initializer = BIR_VAL_NONE; /* Check for literal initializer (skip past array size if present) */ { @@ -3346,6 +3792,23 @@ static void collect_template(lower_t *L, uint32_t node) * - callee = "scale" * - Deduce T from literal arguments */ +static uint32_t arg_type(lower_t *L, uint32_t arg) +{ + const ast_node_t *a = ND(L, arg); + switch (a->type) { + case AST_FLOAT_LIT: { + int f32; + parse_float_text(L->src + a->d.text.offset, (int)a->d.text.len, &f32); + return f32 ? bir_type_float(L->M, 32) : bir_type_float(L->M, 64); + } + case AST_INT_LIT: return bir_type_int(L->M, 32); + case AST_BOOL_LIT: return bir_type_int(L->M, 1); + case AST_CAST: return resolve_type(L, a->first_child, + a->d.oper.flags, 0); + default: return 0; + } +} + static void scan_launches(lower_t *L, uint32_t node) { if (!node) return; @@ -3360,136 +3823,144 @@ static void scan_launches(lower_t *L, uint32_t node) template_def_t *tmpl = find_template(L, cname); if (!tmpl) goto recurse; + uint32_t arg = ND(L, callee_n)->next_sibling; + if (arg) arg = ND(L, arg)->next_sibling; + if (arg) arg = ND(L, arg)->next_sibling; - /* Deduce template type from launch arguments. - Skip grid and block args (children 1 and 2 of launch). - Actual args start after that. */ - uint32_t arg = ND(L, callee_n)->next_sibling; /* grid */ - if (arg) arg = ND(L, arg)->next_sibling; /* block */ - /* Optional shared mem and stream */ - /* Skip to actual function args: after >>> comes () */ - /* In AST_LAUNCH, children after grid,block,[smem],[stream] are args */ - /* We need to find the function args. Let's skip 2 more possible. */ - if (arg) arg = ND(L, arg)->next_sibling; /* first real arg, or smem */ - - /* Collect template params from AST_TEMPLATE_DECL */ uint32_t tmpl_node = tmpl->ast; - uint32_t tp = ND(L, tmpl_node)->first_child; - int ntparams = 0; - binding_t new_bindings[8]; - - while (tp && ND(L, tp)->type == AST_TEMPLATE_PARAM && ntparams < 8) { - int is_type = (ND(L, tp)->d.oper.flags == 0); - uint32_t tpname_n = 0; - uint32_t tpc = ND(L, tp)->first_child; - while (tpc) { - if (ND(L, tpc)->type == AST_IDENT) { tpname_n = tpc; break; } - tpc = ND(L, tpc)->next_sibling; - } - - new_bindings[ntparams].is_type = is_type; - new_bindings[ntparams].type = 0; - new_bindings[ntparams].ival = 0; - if (tpname_n) - get_text(L, tpname_n, new_bindings[ntparams].name, - sizeof(new_bindings[0].name)); - else - new_bindings[ntparams].name[0] = '\0'; - - ntparams++; - tp = ND(L, tp)->next_sibling; + binding_t nb[8]; + int ntp = 0; + + for (uint32_t tp = ND(L, tmpl_node)->first_child; + tp && ND(L, tp)->type == AST_TEMPLATE_PARAM && ntp < 8; + tp = ND(L, tp)->next_sibling) { + int fl = ND(L, tp)->d.oper.flags; + memset(&nb[ntp], 0, sizeof(nb[0])); + nb[ntp].is_type = !(fl & TP_NTYP); + if (fl & TP_PACK) { + nb[ntp].is_pack = 1; + nb[ntp].is_type = 0; + } + for (uint32_t tpc = ND(L, tp)->first_child; tpc; + tpc = ND(L, tpc)->next_sibling) + if (ND(L, tpc)->type == AST_IDENT) { + get_text(L, tpc, nb[ntp].name, sizeof(nb[0].name)); + break; + } + ntp++; } - /* Deduce types from arguments. - For typename T: look at float literal args (2.0f → float) */ - /* Find the function def to match params with template params */ uint32_t func_n = 0; - { - uint32_t fc = ND(L, tmpl_node)->first_child; - while (fc) { - if (ND(L, fc)->type == AST_FUNC_DEF) { func_n = fc; break; } - fc = ND(L, fc)->next_sibling; - } - } - - if (func_n) { - /* Match function params against launch args to deduce T */ - uint32_t fparam_nodes[16]; - int nfparams = collect_params(L, func_n, fparam_nodes, 16); - - /* For each function param, check if its type is a template param */ - /* Then look at the corresponding launch arg to deduce */ - for (int i = 0; i < nfparams && arg; i++) { - uint32_t fpt = ND(L, fparam_nodes[i])->first_child; - if (fpt && ND(L, fpt)->type == AST_TYPE_SPEC - && ND(L, fpt)->d.btype.kind == TYPE_NAME) { - /* Named type — might be a template param */ - uint32_t fpt_name = ND(L, fpt)->first_child; - if (fpt_name) { - char tname[64]; - get_text(L, fpt_name, tname, sizeof(tname)); - for (int b = 0; b < ntparams; b++) { - if (new_bindings[b].is_type - && strcmp(new_bindings[b].name, tname) == 0 - && new_bindings[b].type == 0) { - /* Deduce from arg */ - if (ND(L, arg)->type == AST_FLOAT_LIT) { - int is_f32; - parse_float_text( - L->src + ND(L, arg)->d.text.offset, - (int)ND(L, arg)->d.text.len, &is_f32); - new_bindings[b].type = is_f32 - ? bir_type_float(L->M, 32) - : bir_type_float(L->M, 64); - } else if (ND(L, arg)->type == AST_INT_LIT) { - new_bindings[b].type = bir_type_int(L->M, 32); - } else if (ND(L, arg)->type == AST_IDENT) { - /* Look up variable type */ - /* For pointers like d_data (float*), - the template param T maps to pointee */ - /* Can't easily deduce without host sym table. - Default to float for now. */ - new_bindings[b].type = bir_type_float(L->M, 32); - } - } - } - } + for (uint32_t fc = ND(L, tmpl_node)->first_child; fc; + fc = ND(L, fc)->next_sibling) + if (ND(L, fc)->type == AST_FUNC_DEF) { func_n = fc; break; } + if (!func_n) return; + + uint32_t fpn[16]; + int nfp = collect_params(L, func_n, fpn, 16); + int ok = 1; + + for (int i = 0; i < nfp; i++) { + const ast_node_t *pn = ND(L, fpn[i]); + uint32_t ts = pn->first_child; + char tname[64]; + tname[0] = 0; + if (ts && ND(L, ts)->type == AST_TYPE_SPEC + && ND(L, ts)->d.btype.kind == TYPE_NAME + && ND(L, ts)->first_child) + get_text(L, ND(L, ts)->first_child, tname, sizeof(tname)); + + if (pn->d.oper.op == PRM_PACK) { + int b = -1; + for (int k = 0; k < ntp; k++) + if (nb[k].is_pack && strcmp(nb[k].name, tname) == 0) b = k; + if (b < 0) { ok = 0; break; } + while (arg && nb[b].npk < MAX_PKELM) { + uint32_t at = arg_type(L, arg); + if (!at) { ok = 0; break; } + nb[b].pk_t[nb[b].npk++] = at; + arg = ND(L, arg)->next_sibling; } - arg = ND(L, arg)->next_sibling; + if (arg) ok = 0; + break; + } + + if (tname[0]) { + for (int b = 0; b < ntp; b++) { + if (!nb[b].is_type || strcmp(nb[b].name, tname) != 0 + || nb[b].type != 0) + continue; + if (!arg) break; + uint32_t at = arg_type(L, arg); + nb[b].type = at ? at : bir_type_float(L->M, 32); + } + } + if (arg) arg = ND(L, arg)->next_sibling; + } + + if (!ok) { + lower_error(L, node, BC_E030, + "variadic template whose pack cannot be deduced here"); + return; + } + + for (int i = 0; i < ntp; i++) + if (nb[i].is_type && nb[i].type == 0) + nb[i].type = bir_type_float(L->M, 32); + + char key[256]; + int kl = snprintf(key, sizeof(key), "%s", cname); + for (int i = 0; i < ntp && kl > 0 && kl < (int)sizeof(key) - 32; i++) { + if (nb[i].is_pack) { + kl += snprintf(key + kl, sizeof(key) - (size_t)kl, + "|p%d", nb[i].npk); + for (int k = 0; k < nb[i].npk + && kl < (int)sizeof(key) - 16; k++) + kl += snprintf(key + kl, sizeof(key) - (size_t)kl, + ":%u", nb[i].pk_t[k]); + } else if (nb[i].is_type) { + kl += snprintf(key + kl, sizeof(key) - (size_t)kl, + "|t%u", nb[i].type); + } else { + kl += snprintf(key + kl, sizeof(key) - (size_t)kl, + "|n%lld", (long long)nb[i].ival); } } - /* Default unresolved type bindings to float */ - for (int i = 0; i < ntparams; i++) { - if (new_bindings[i].is_type && new_bindings[i].type == 0) - new_bindings[i].type = bir_type_float(L->M, 32); + int nsame = 0; + for (int i = 0; i < L->ninst; i++) { + if (strcmp(L->inst[i].key, key) == 0) return; + if (strncmp(L->inst[i].key, cname, strlen(cname)) == 0 + && L->inst[i].key[strlen(cname)] == '|') + nsame++; + } + if (L->ninst >= MAX_INSTS) { + lower_error(L, node, BC_E030, "too many template instantiations"); + return; } - /* Build mangled name */ - char mangled[256]; - snprintf(mangled, sizeof(mangled), "%s", cname); - /* For template instantiation, we could add suffix - but keep it simple — just use the base name if unique */ + char mangled[128]; + if (nsame == 0) { + ncpy(mangled, sizeof(mangled), cname); + } else { + char base[112]; - /* Check if already instantiated */ - int already = 0; - for (uint32_t i = 0; i < L->M->num_funcs; i++) { - if (L->M->funcs[i].name < L->M->string_len - && strcmp(&L->M->strings[L->M->funcs[i].name], mangled) == 0) { - already = 1; - break; - } + ncpy(base, sizeof(base), cname); + snprintf(mangled, sizeof(mangled), "%s$%d", base, nsame); } + snprintf(L->inst[L->ninst].key, sizeof(L->inst[0].key), "%s", key); + snprintf(L->inst[L->ninst].sym, sizeof(L->inst[0].sym), "%s", mangled); + L->ninst++; - if (!already && func_n) { - /* Set bindings and instantiate */ - int old_nb = L->nbindings; - for (int i = 0; i < ntparams && L->nbindings < 8; i++) - L->bindings[L->nbindings++] = new_bindings[i]; + if (bir_fsym(L->M, mangled, L->tu, -1) != BIR_SYM_NONE) + return; + { + int old_nb = L->nbindings; + for (int i = 0; i < ntp && L->nbindings < 8; i++) + L->bindings[L->nbindings++] = nb[i]; uint16_t cuda = ND(L, func_n)->cuda_flags; lower_func_body(L, func_n, cuda, mangled); - L->nbindings = old_nb; } @@ -3511,6 +3982,14 @@ static void scan_launches(lower_t *L, uint32_t node) int bir_lower(const parser_t *P, uint32_t ast_root, bir_module_t *M, const sema_ctx_t *sema, bc_error_t *out_errs, int *out_nerrs) +{ + bir_module_init(M); + return bir_ltu(P, ast_root, M, sema, BIR_TU_EXT, out_errs, out_nerrs); +} + +int bir_ltu(const parser_t *P, uint32_t ast_root, bir_module_t *M, + const sema_ctx_t *sema, uint16_t tu, + bc_error_t *out_errs, int *out_nerrs) { static lower_t L_storage; /* large struct — static to avoid stack overflow */ lower_t *L = &L_storage; @@ -3519,8 +3998,7 @@ int bir_lower(const parser_t *P, uint32_t ast_root, bir_module_t *M, L->M = M; L->src = P->src; L->sema = sema; - - bir_module_init(M); + L->tu = tu; /* Pass 1: collect declarations */ uint32_t c = P->nodes[ast_root].first_child; diff --git a/src/ir/bir_lower.h b/src/ir/bir_lower.h index 5cec334..0f5b9e9 100644 --- a/src/ir/bir_lower.h +++ b/src/ir/bir_lower.h @@ -18,4 +18,8 @@ int bir_lower(const struct parser_s *P, uint32_t ast_root, bir_module_t *M, const struct sema_ctx_s *sema, bc_error_t *out_errs, int *out_nerrs); +int bir_ltu(const struct parser_s *P, uint32_t ast_root, bir_module_t *M, + const struct sema_ctx_s *sema, uint16_t tu, + bc_error_t *out_errs, int *out_nerrs); + #endif /* BARRACUDA_BIR_LOWER_H */ diff --git a/src/main.c b/src/main.c index aba15cb..0ab6924 100644 --- a/src/main.c +++ b/src/main.c @@ -215,6 +215,158 @@ static void dump_tokens(const lexer_t *L) } } + +typedef struct { + int mode_pp, mode_lex, mode_parse, mode_sema, mode_hip, no_pp; + int want_ast, want_sema, want_bir; + const char *const *incs; int nincs; + const char *const *defs; int ndefs; + bir_module_t *M; + uint16_t tu; +} tuc_t; + +static int comp_tu(const char *file, const tuc_t *c) +{ + static bc_error_t lower_errs[BC_MAX_ERRORS]; + uint32_t src_len = 0; + + if (read_file(file, source_buf, BC_MAX_SOURCE, &src_len) != BC_OK) + return BC_ERR_IO; + + const char *lex_src = source_buf; + uint32_t lex_len = src_len; + + if (!c->no_pp) { + preproc_t *pp = (preproc_t *)malloc(sizeof(preproc_t)); + if (!pp) { + fprintf(stderr, "error: failed to allocate preprocessor\n"); + return BC_ERR_IO; + } + pp_init(pp, source_buf, src_len, pp_out_buf, BC_MAX_SOURCE, file); + + if (c->mode_hip) { + pp_define(pp, "__HIPCC__", "1"); + pp_define(pp, "__HIP_DEVICE_COMPILE__", "1"); + /* Which platform HIP thinks it is compiling for follows the + * selected target, so ask the registry rather than keep a + * copy of the mode flag here. */ + const be_desc_t *hb = be_active(); + if (hb != NULL && strcmp(hb->name, "nvptx") == 0) + pp_define(pp, "__HIP_PLATFORM_NVIDIA__", "1"); + else + pp_define(pp, "__HIP_PLATFORM_AMD__", "1"); + } + + for (int i = 0; i < c->nincs; i++) + pp_add_include_path(pp, c->incs[i]); + for (int i = 0; i < c->ndefs; i++) { + char dname[BC_MAX_IDENT]; + const char *eq = strchr(c->defs[i], '='); + if (eq) { + uint32_t nlen = (uint32_t)(eq - c->defs[i]); + if (nlen >= BC_MAX_IDENT) nlen = BC_MAX_IDENT - 1; + memcpy(dname, c->defs[i], nlen); + dname[nlen] = '\0'; + pp_define(pp, dname, eq + 1); + } else { + pp_define(pp, c->defs[i], "1"); + } + } + + int prc = pp_process(pp); + + bc_diag(file, source_buf, pp->errors, pp->num_errors); + + if (c->mode_pp) { + fwrite(pp_out_buf, 1, pp->out_len, stdout); + free(pp); + return prc; + } + + /* A truncated expansion is not shorter source, it is different + * source, and every phase after this would be reading a lie. */ + if (pp->ovflw) { + free(pp); + return 1; + } + + lex_src = pp_out_buf; + lex_len = pp->out_len; + free(pp); + } + + lexer_t L; + lexer_init(&L, lex_src, lex_len, token_buf, BC_MAX_TOKENS); + int rc = lexer_tokenize(&L); + + bc_diag(file, lex_src, L.errors, L.num_errors); + + if (c->mode_lex) { + dump_tokens(&L); + printf("\n%u tokens, %d error(s)\n", L.num_tokens, L.num_errors); + } + + if (!c->want_ast) return rc; + + parser_t P; + parser_init(&P, token_buf, L.num_tokens, lex_src, + node_buf, BC_MAX_NODES); + uint32_t root = parser_parse(&P); + + bc_diag(file, lex_src, P.errors, P.num_errors); + + if (c->mode_parse) { + ast_dump(&P, root, 0); + printf("\n%u nodes, %d parse error(s)\n", + P.num_nodes, P.num_errors); + } + + /* Semantic analysis */ + sema_ctx_t *sema_ctx = NULL; + if (c->want_sema && P.num_errors == 0) + { + sema_ctx = (sema_ctx_t *)malloc(sizeof(sema_ctx_t)); + if (!sema_ctx) { + fprintf(stderr, "error: failed to allocate sema context\n"); + return BC_ERR_IO; + } + sema_init(sema_ctx, &P, root, (int)be_warp_size()); + sema_check(sema_ctx, root); + + bc_diag(file, lex_src, sema_ctx->errors, sema_ctx->num_errors); + + if (c->mode_sema) { + sema_dump(sema_ctx, root); + int sema_rc = sema_ctx->num_errors > 0 ? BC_ERR_SEMA : BC_OK; + free(sema_ctx); + return sema_rc; + } + + if (sema_ctx->num_errors > 0) { + free(sema_ctx); + return BC_ERR_SEMA; + } + } + + if (c->want_bir && P.num_errors == 0) { + int num_lower_errs = 0; + int lrc = bir_ltu(&P, root, c->M, sema_ctx, c->tu, + lower_errs, &num_lower_errs); + for (int i = 0; i < num_lower_errs; i++) { + fprintf(stderr, "%s:%u:%u: E%03u: %s\n", + file, lower_errs[i].loc.line, + lower_errs[i].loc.col, + lower_errs[i].eid, lower_errs[i].msg); + } + if (lrc != BC_OK) rc = lrc; + } + + if (sema_ctx) free(sema_ctx); + if (P.num_errors > 0) rc = BC_ERR_PARSE; + + return rc; +} + static void usage(const char *prog) { fprintf(stderr, @@ -286,7 +438,8 @@ static void usage(const char *prog) int main(int argc, char *argv[]) { - const char *file = NULL; + const char *files[BC_MAX_TUS]; + int nfile = 0; const char *output_file = NULL; const char *lang_file = NULL; int mode_pp = 0; @@ -296,7 +449,7 @@ int main(int argc, char *argv[]) int mode_ir = 0; int mode_tdf = 0; int mode_tdf_fission = 0; - int mode_hip = 0; /* HIP frontend: see HIP NOTES below */ + int mode_hip = 0; int mode_triton = 0; /* Triton frontend: see TRITON NOTES below */ int mode_mlir = 0; /* MLIR frontend: see MLIR NOTES below */ int mode_bir_in = 0; /* BIR text in: see BIR NOTES below */ @@ -387,8 +540,13 @@ int main(int argc, char *argv[]) else if (strcmp(argv[i], "--help") == 0 || strcmp(argv[i], "-h") == 0) { usage(argv[0]); return 0; - } else if (argv[i][0] != '-') - file = argv[i]; + } else if (argv[i][0] != '-') { + if (nfile >= BC_MAX_TUS) { + fprintf(stderr, "too many input files (max %d)\n", BC_MAX_TUS); + return 1; + } + files[nfile++] = argv[i]; + } else { /* Not one of the driver's, so offer it round the backend * registry before calling it unknown. Target selection and @@ -411,7 +569,7 @@ int main(int argc, char *argv[]) } } - if (!file) { + if (nfile == 0) { usage(argv[0]); return 1; } @@ -438,32 +596,25 @@ int main(int argc, char *argv[]) int want_ast = mode_parse || want_sema; - /* ---- HIP NOTES (1 of 2) ------------------------------------------- - * HIP is a frontend-only mode, not a separate parser. The HIP source - * language is, in practice, a syntactic superset of CUDA: the same - * __global__ / __device__ / __shared__ qualifiers, the same - * threadIdx / blockIdx / blockDim builtins, and the same arithmetic - * happens with the same syntax. What changes when you flip --hip on - * is the set of preprocessor macros we predefine, so that any - * #if defined(__HIPCC__) or __HIP_PLATFORM_AMD__ guards in the source - * pick the HIP branch instead of falling through to whatever the - * source thought the default platform was. - * - * Auto-detection: if the filename ends in ".hip", we assume HIP mode - * without making the user say so on the command line, since the - * extension is a strong enough signal for anyone using a HIP build - * pipeline to drop their files into Booth unchanged. */ - if (file) { - size_t flen = strlen(file); - if (flen >= 4 && strcmp(file + flen - 4, ".hip") == 0) + for (int i = 0; i < nfile; i++) { + size_t flen = strlen(files[i]); + if (flen >= 4 && strcmp(files[i] + flen - 4, ".hip") == 0) mode_hip = 1; } /* Load translation file before any diagnostics fire */ if (lang_file) bc_eload(lang_file); + if (nfile > 1 && (mode_bir_in || mode_mlir || mode_triton)) { + fprintf(stderr, "only one input file for --bir-in, --mlir " + "and --triton (%d given)\n", nfile); + return 1; + } + + const char *file = files[0]; uint32_t src_len = 0; - if (read_file(file, source_buf, BC_MAX_SOURCE, &src_len) != BC_OK) + if (nfile == 1 && + read_file(file, source_buf, BC_MAX_SOURCE, &src_len) != BC_OK) return 1; /* ---- BIR NOTES ---------------------------------------------------- @@ -656,184 +807,55 @@ int main(int argc, char *argv[]) return 1; } - /* Preprocessing */ - const char *lex_src = source_buf; - uint32_t lex_len = src_len; - - if (!no_pp) { - preproc_t *pp = (preproc_t *)malloc(sizeof(preproc_t)); - if (!pp) { - fprintf(stderr, "error: failed to allocate preprocessor\n"); - return 1; - } - pp_init(pp, source_buf, src_len, pp_out_buf, BC_MAX_SOURCE, file); - - /* ---- HIP NOTES (2 of 2) --------------------------------------- - * This is the only spot in the pipeline that knows or cares - * whether we are compiling CUDA or HIP. pp_init has just defined - * the CUDA defaults (__BARRACUDA__, __CUDA_ARCH__, __CUDACC__) - * unconditionally, which is correct for CUDA and harmless for - * HIP because real HIP source files distinguish platforms with - * __HIP_PLATFORM_AMD__ versus __HIP_PLATFORM_NVIDIA__ rather - * than by the presence or absence of __CUDACC__. - * - * When --hip is on, we additively define the HIP-specific - * macros so that the preprocessor takes the HIP branch wherever - * the source asks for it: - * __HIPCC__ compiler identity, "we are a HIP compiler" - * __HIP_DEVICE_COMPILE__ we are compiling device code (always true here) - * __HIP_PLATFORM_AMD__ target is AMD silicon (the common case) - * __HIP_PLATFORM_NVIDIA__ target is NVIDIA via the HIP-on-CUDA path - * - * Beyond these macros, nothing in the parser, sema, IR, or - * backends needs to know about HIP. The pipeline downstream of - * here is identical to a CUDA compile. */ - if (mode_hip) { - pp_define(pp, "__HIPCC__", "1"); - pp_define(pp, "__HIP_DEVICE_COMPILE__", "1"); - /* Which platform HIP thinks it is compiling for follows the - * selected target, so ask the registry rather than keep a - * copy of the mode flag here. */ - const be_desc_t *hb = be_active(); - if (hb != NULL && strcmp(hb->name, "nvptx") == 0) - pp_define(pp, "__HIP_PLATFORM_NVIDIA__", "1"); - else - pp_define(pp, "__HIP_PLATFORM_AMD__", "1"); - } - - for (int i = 0; i < num_include_paths; i++) - pp_add_include_path(pp, include_paths[i]); - for (int i = 0; i < num_defines; i++) { - char dname[BC_MAX_IDENT]; - const char *eq = strchr(defines[i], '='); - if (eq) { - uint32_t nlen = (uint32_t)(eq - defines[i]); - if (nlen >= BC_MAX_IDENT) nlen = BC_MAX_IDENT - 1; - memcpy(dname, defines[i], nlen); - dname[nlen] = '\0'; - pp_define(pp, dname, eq + 1); - } else { - pp_define(pp, defines[i], "1"); - } - } - - int prc = pp_process(pp); - - bc_diag(file, source_buf, pp->errors, pp->num_errors); - - if (mode_pp) { - fwrite(pp_out_buf, 1, pp->out_len, stdout); - free(pp); - return prc != BC_OK ? 1 : 0; - } - - /* A truncated expansion is not shorter source, it is different - * source, and every phase after this would be reading a lie. */ - if (pp->ovflw) { - free(pp); + tuc_t tc; + tc.mode_pp = mode_pp; + tc.mode_lex = mode_lex; + tc.mode_parse = mode_parse; + tc.mode_sema = mode_sema; + tc.mode_hip = mode_hip; + tc.no_pp = no_pp; + tc.want_ast = want_ast; + tc.want_sema = want_sema; + tc.want_bir = want_bir; + tc.incs = include_paths; + tc.nincs = num_include_paths; + tc.defs = defines; + tc.ndefs = num_defines; + tc.M = NULL; + tc.tu = BIR_TU_EXT; + + if (want_bir) { + bir_module = (bir_module_t *)malloc(sizeof(bir_module_t)); + if (!bir_module) { + fprintf(stderr, "error: failed to allocate BIR module\n"); return 1; } - - lex_src = pp_out_buf; - lex_len = pp->out_len; - free(pp); + bir_module_init(bir_module); + tc.M = bir_module; } - lexer_t L; - lexer_init(&L, lex_src, lex_len, token_buf, BC_MAX_TOKENS); - int rc = lexer_tokenize(&L); - - bc_diag(file, lex_src, L.errors, L.num_errors); - - if (mode_lex) { - dump_tokens(&L); - printf("\n%u tokens, %d error(s)\n", L.num_tokens, L.num_errors); + int rc = BC_OK; + for (int t = 0; t < nfile; t++) { + tc.tu = (nfile > 1) ? (uint16_t)t : BIR_TU_EXT; + rc = comp_tu(files[t], &tc); + if (rc != BC_OK) break; } - if (want_ast) { - parser_t P; - parser_init(&P, token_buf, L.num_tokens, lex_src, - node_buf, BC_MAX_NODES); - uint32_t root = parser_parse(&P); - - bc_diag(file, lex_src, P.errors, P.num_errors); - - if (mode_parse) { - ast_dump(&P, root, 0); - printf("\n%u nodes, %d parse error(s)\n", - P.num_nodes, P.num_errors); - } - - /* Semantic analysis */ - sema_ctx_t *sema_ctx = NULL; - if (want_sema && P.num_errors == 0) - { - sema_ctx = (sema_ctx_t *)malloc(sizeof(sema_ctx_t)); - if (!sema_ctx) { - fprintf(stderr, "error: failed to allocate sema context\n"); - return 1; - } - sema_init(sema_ctx, &P, root, (int)be_warp_size()); - sema_check(sema_ctx, root); - - bc_diag(file, lex_src, sema_ctx->errors, sema_ctx->num_errors); - - if (mode_sema) { - sema_dump(sema_ctx, root); - int sema_rc = sema_ctx->num_errors > 0 ? 1 : 0; - free(sema_ctx); - return sema_rc; - } - - /* Only --sema used to act on these. Every other mode printed the - * errors, carried on into codegen, wrote an output file and exited - * zero, so a build system saw a clean compile and a kernel that - * had been lowered from source we had already rejected. */ - if (sema_ctx->num_errors > 0) { - free(sema_ctx); - return 1; - } - } - - if (want_bir && P.num_errors == 0) { - bc_error_t lower_errs[BC_MAX_ERRORS]; - int num_lower_errs = 0; - bir_module = (bir_module_t *)malloc(sizeof(bir_module_t)); - if (!bir_module) { - fprintf(stderr, "error: failed to allocate BIR module\n"); - return 1; - } - int lrc = bir_lower(&P, root, bir_module, sema_ctx, - lower_errs, &num_lower_errs); - if (num_lower_errs > 0) { - for (int i = 0; i < num_lower_errs; i++) { - fprintf(stderr, "%s:%u:%u: E%03u: %s\n", - file, lower_errs[i].loc.line, - lower_errs[i].loc.col, - lower_errs[i].eid, lower_errs[i].msg); - } - } - if (lrc == BC_OK) { - backend_cfg_t cfg = {0}; - cfg.no_mem2reg = no_mem2reg; - cfg.no_cfold = no_cfold; - cfg.no_dce = no_dce; - cfg.no_sched = no_sched; - cfg.no_sroa = no_sroa; - cfg.mode_ir = mode_ir; - cfg.mode_tdf = mode_tdf; - cfg.mode_tdf_fission = mode_tdf_fission; - cfg.output_file = output_file; - int brc = run_bir_backends(bir_module, &cfg); - if (brc != BC_OK) rc = brc; - } - free(bir_module); - if (lrc != BC_OK) rc = lrc; - } - - if (sema_ctx) free(sema_ctx); - if (P.num_errors > 0) rc = BC_ERR_PARSE; + if (rc == BC_OK && want_bir) { + backend_cfg_t cfg = {0}; + cfg.no_mem2reg = no_mem2reg; + cfg.no_cfold = no_cfold; + cfg.no_dce = no_dce; + cfg.no_sched = no_sched; + cfg.no_sroa = no_sroa; + cfg.mode_ir = mode_ir; + cfg.mode_tdf = mode_tdf; + cfg.mode_tdf_fission = mode_tdf_fission; + cfg.output_file = output_file; + rc = run_bir_backends(bir_module, &cfg); } + if (bir_module) free(bir_module); + return rc != BC_OK ? 1 : 0; } diff --git a/src/metal/emit.c b/src/metal/emit.c index 6edd3d0..f2d88a4 100644 --- a/src/metal/emit.c +++ b/src/metal/emit.c @@ -507,6 +507,11 @@ static int mt_stmt(metal_module_t *mm, uint32_t gi) if (!mt_lhs(mm, gi)) return 0; return mt_wfmt(mm, "gdim.%c;\n", mt_dim(I->subop)); + case BIR_MMA: case BIR_MFRG: + fprintf(stderr, "kath: warp-collective mma not supported on the " + "Metal backend\n"); + return 0; + case BIR_SELECT: if (!mt_lhs(mm, gi)) return 0; if (!mt_val(mm, I->operands[0])) return 0; diff --git a/src/nvidia/emit.c b/src/nvidia/emit.c index d38d9a0..87c4608 100644 --- a/src/nvidia/emit.c +++ b/src/nvidia/emit.c @@ -23,6 +23,13 @@ static void nv_apnd(nv_module_t *nv, const char *fmt, ...) nv->out_len = NV_MAX_OUT - 1; } +const nv_mmash_t nv_mmash[NV_MMA_NSHAPE] = { + { "m16n8k16.row.col.f32.f16.f16.f32", 8, 4 }, + { "m16n8k16.row.col.f32.bf16.bf16.f32", 8, 4 }, + { "m16n8k8.row.col.f32.f16.f16.f32", 4, 2 }, + { "m16n8k8.row.col.f32.bf16.bf16.f32", 4, 2 }, +}; + /* ---- Special Register Names ---- */ static const char *spec_name(int32_t id) @@ -40,6 +47,7 @@ static const char *spec_name(int32_t id) case NV_SPEC_NCTAID_X: return "%nctaid.x"; case NV_SPEC_NCTAID_Y: return "%nctaid.y"; case NV_SPEC_NCTAID_Z: return "%nctaid.z"; + case NV_SPEC_LANEID: return "%laneid"; default: return "%tid.x"; } } @@ -58,6 +66,7 @@ static void em_opnd(nv_module_t *nv, const nv_opnd_t *op) case NV_RF_PRED: nv_apnd(nv, "%%p%u", op->reg_num); break; case NV_RF_U16: nv_apnd(nv, "%%rh%u", op->reg_num); break; case NV_RF_F16: nv_apnd(nv, "%%h%u", op->reg_num); break; + case NV_RF_B32: nv_apnd(nv, "%%rb%u", op->reg_num); break; default: nv_apnd(nv, "%%r%u", op->reg_num); break; } break; @@ -640,15 +649,17 @@ static void em_inst(nv_module_t *nv, const nv_minst_t *I) } /* ---- Loads/Stores: Shared ---- */ - case NV_LD_SHR_U32: case NV_LD_SHR_F32: { - const char *tsuf = (I->op == NV_LD_SHR_F32) ? ".f32" : ".u32"; + case NV_LD_SHR_U32: case NV_LD_SHR_F32: case NV_LD_SHR_U8: { + const char *tsuf = (I->op == NV_LD_SHR_F32) ? ".f32" + : (I->op == NV_LD_SHR_U8) ? ".u8" : ".u32"; nv_apnd(nv, "ld.shared%s ", tsuf); em_opnd(nv, &I->ops[0]); nv_apnd(nv, ", ["); em_opnd(nv, &I->ops[1]); nv_apnd(nv, "]"); break; } - case NV_ST_SHR_U32: case NV_ST_SHR_F32: { - const char *tsuf = (I->op == NV_ST_SHR_F32) ? ".f32" : ".u32"; + case NV_ST_SHR_U32: case NV_ST_SHR_F32: case NV_ST_SHR_U8: { + const char *tsuf = (I->op == NV_ST_SHR_F32) ? ".f32" + : (I->op == NV_ST_SHR_U8) ? ".u8" : ".u32"; nv_apnd(nv, "st.shared%s [", tsuf); em_opnd(nv, &I->ops[0]); nv_apnd(nv, "], "); em_opnd(nv, &I->ops[1]); @@ -657,12 +668,13 @@ static void em_inst(nv_module_t *nv, const nv_minst_t *I) /* ---- Loads/Stores: Local ---- */ case NV_LD_LOC_U32: case NV_LD_LOC_U64: - case NV_LD_LOC_F32: case NV_LD_LOC_F64: { + case NV_LD_LOC_F32: case NV_LD_LOC_F64: case NV_LD_LOC_U8: { const char *tsuf; switch (I->op) { case NV_LD_LOC_U64: tsuf = ".u64"; break; case NV_LD_LOC_F32: tsuf = ".f32"; break; case NV_LD_LOC_F64: tsuf = ".f64"; break; + case NV_LD_LOC_U8: tsuf = ".u8"; break; default: tsuf = ".u32"; break; } nv_apnd(nv, "ld.local%s ", tsuf); @@ -671,12 +683,13 @@ static void em_inst(nv_module_t *nv, const nv_minst_t *I) break; } case NV_ST_LOC_U32: case NV_ST_LOC_U64: - case NV_ST_LOC_F32: case NV_ST_LOC_F64: { + case NV_ST_LOC_F32: case NV_ST_LOC_F64: case NV_ST_LOC_U8: { const char *tsuf; switch (I->op) { case NV_ST_LOC_U64: tsuf = ".u64"; break; case NV_ST_LOC_F32: tsuf = ".f32"; break; case NV_ST_LOC_F64: tsuf = ".f64"; break; + case NV_ST_LOC_U8: tsuf = ".u8"; break; default: tsuf = ".u32"; break; } nv_apnd(nv, "st.local%s [", tsuf); @@ -933,6 +946,34 @@ static void em_inst(nv_module_t *nv, const nv_minst_t *I) break; } + case NV_MOV_PK16: + nv_apnd(nv, "mov.b32 "); + em_opnd(nv, &I->ops[0]); nv_apnd(nv, ", {"); + em_opnd(nv, &I->ops[1]); nv_apnd(nv, ", "); + em_opnd(nv, &I->ops[2]); nv_apnd(nv, "}"); + break; + case NV_BARWARP: + nv_apnd(nv, "bar.warp.sync 0xffffffff"); + break; + case NV_MMA: { + const nv_mmash_t *sh = &nv_mmash[I->flags % NV_MMA_NSHAPE]; + uint8_t nreg[4]; + nreg[0] = 4; nreg[1] = (uint8_t)(sh->na / 2); + nreg[2] = (uint8_t)(sh->nb / 2); nreg[3] = 4; + nv_apnd(nv, "mma.sync.aligned.%s ", sh->sfx); + for (uint8_t g = 0; g < 4; g++) { + if (g > 0) nv_apnd(nv, ", "); + nv_apnd(nv, "{"); + for (uint8_t k = 0; k < nreg[g]; k++) { + nv_opnd_t o = I->ops[g]; + o.reg_num = (uint16_t)(o.reg_num + k); + if (k > 0) nv_apnd(nv, ", "); + em_opnd(nv, &o); + } + nv_apnd(nv, "}"); + } + break; + } default: nv_apnd(nv, "/* unknown op %u */", I->op); break; @@ -986,6 +1027,8 @@ static void em_func(nv_module_t *nv, uint32_t fi) nv_apnd(nv, "\t.reg .u16 %%rh<%u>;\n", MF->rc[NV_RF_U16]); if (MF->rc[NV_RF_F16] > 1) nv_apnd(nv, "\t.reg .f16 %%h<%u>;\n", MF->rc[NV_RF_F16]); + if (MF->rc[NV_RF_B32] > 1) + nv_apnd(nv, "\t.reg .b32 %%rb<%u>;\n", MF->rc[NV_RF_B32]); /* Local (stack) memory — without this declaration, ld.local/st.local * access unmapped memory and the driver gets very cross with us */ diff --git a/src/nvidia/isel.c b/src/nvidia/isel.c index af6ff0f..e7210fd 100644 --- a/src/nvidia/isel.c +++ b/src/nvidia/isel.c @@ -13,10 +13,19 @@ static struct { uint32_t lcl_off; /* local (alloca) byte offset */ uint32_t shr_off; /* shared memory byte offset */ uint32_t cur_func; /* current nv_mfunc_t index */ + int unsz; /* a type with no storage size was met */ + int had_error; /* an op we refuse to fake; fail the compile */ } S; +static void nv_refuse(const char *what) +{ + fprintf(stderr, "kath: %s not supported on this backend\n", what); + S.had_error = 1; +} + /* Forward declaration — rslv needs em1 for f64 constant materialisation */ static void em1(uint16_t op, nv_opnd_t d, nv_opnd_t a, nv_opnd_t b); +static nv_opnd_t mat_pred(nv_opnd_t op); /* ---- Deferred PHI Copies ---- * PHI elimination requires inserting MOV copies into predecessor @@ -105,7 +114,7 @@ static uint8_t bir_rfile(uint32_t type_idx) switch (T->kind) { case BIR_TYPE_INT: - if (T->width <= 1) return NV_RF_PRED; + if (T->width <= 1) return NV_RF_U32; if (T->width <= 16) return NV_RF_U16; if (T->width <= 32) return NV_RF_U32; return NV_RF_U64; @@ -122,6 +131,17 @@ static uint8_t bir_rfile(uint32_t type_idx) } } +static uint8_t def_rf(uint32_t idx, uint32_t type_idx) +{ + if (idx < S.bir->num_insts) { + uint16_t o = S.bir->insts[idx].op; + if (o == BIR_ICMP || o == BIR_FCMP + || o == BIR_VOTE_ANY || o == BIR_VOTE_ALL) + return NV_RF_PRED; + } + return bir_rfile(type_idx); +} + /* Map a BIR instruction to a vreg, creating one if needed */ static nv_opnd_t map_val(uint32_t idx, uint32_t type_idx) { @@ -129,15 +149,14 @@ static nv_opnd_t map_val(uint32_t idx, uint32_t type_idx) if (S.nv->val_vreg[idx] != 0) return mop_reg(S.nv->val_rfile[idx], S.nv->val_vreg[idx]); - uint8_t rf = bir_rfile(type_idx); + uint8_t rf = def_rf(idx, type_idx); uint16_t rn = new_vreg(rf); S.nv->val_vreg[idx] = rn; S.nv->val_rfile[idx] = rf; return mop_reg(rf, rn); } -/* Resolve a BIR operand (instruction ref or constant) */ -static nv_opnd_t rslv(uint32_t val) +static nv_opnd_t rslv_p(uint32_t val) { if (val == BIR_VAL_NONE) return mop_imm(0); @@ -185,6 +204,11 @@ static nv_opnd_t rslv(uint32_t val) return map_val(si, S.bir->insts[si].type); } +static nv_opnd_t rslv(uint32_t val) +{ + return mat_pred(rslv_p(val)); +} + /* Resolve, but return the register file for the BIR value */ static uint8_t rslv_rf(uint32_t val) { @@ -197,9 +221,10 @@ static uint8_t rslv_rf(uint32_t val) } uint32_t si = BIR_VAL_INDEX(val); if (si < S.bir->num_insts) { - if (S.nv->val_rfile[si] != 0 || S.nv->val_vreg[si] != 0) - return S.nv->val_rfile[si]; - return bir_rfile(S.bir->insts[si].type); + uint8_t rf = (S.nv->val_rfile[si] != 0 || S.nv->val_vreg[si] != 0) + ? S.nv->val_rfile[si] + : def_rf(si, S.bir->insts[si].type); + return rf == NV_RF_PRED ? NV_RF_U32 : rf; } return NV_RF_U32; } @@ -225,40 +250,27 @@ static uint32_t get_op(const bir_inst_t *I, uint32_t k) return BIR_VAL_NONE; } -/* Type size in bytes for load/store width selection. - * Structs need the full sum-of-fields treatment — GEP strides - * depend on it. Without this, parts[tid] computes base + tid*4 - * instead of base + tid*40 and you read someone else's neutron. - * Which is bad nuclear physics even by Monte Carlo standards. */ static uint32_t type_bytes(uint32_t type_idx) { - if (type_idx >= S.bir->num_types) return 4; - const bir_type_t *T = &S.bir->types[type_idx]; - switch (T->kind) { - case BIR_TYPE_INT: - case BIR_TYPE_FLOAT: - case BIR_TYPE_BFLOAT: - return ((uint32_t)T->width + 7u) / 8u; - case BIR_TYPE_PTR: - return 8; - case BIR_TYPE_STRUCT: { - uint32_t sz = 0; - int guard = 64; - for (uint16_t i = 0; i < T->num_fields && guard > 0; - i++, guard--) { - uint32_t fi = T->count + (uint32_t)i; - if (fi < S.bir->num_type_fields) - sz += type_bytes(S.bir->type_fields[fi]); - } - return sz ? sz : 4; - } - case BIR_TYPE_ARRAY: { - uint32_t esz = type_bytes(T->inner); - return T->count * esz; - } - default: - return 4; - } + return bir_bsz(S.bir, type_idx, 8); +} + +static uint32_t pntsz(uint32_t ptr_val) +{ + uint32_t si, pt; + if (ptr_val == BIR_VAL_NONE || BIR_VAL_IS_CONST(ptr_val)) return 0; + si = BIR_VAL_INDEX(ptr_val); + if (si >= S.bir->num_insts) return 0; + pt = S.bir->insts[si].type; + return bir_gsz(S.bir, pt, 8); +} + +static void unsz(const char *what, uint32_t ty) +{ + char buf[128]; + if (bir_type_str(S.bir, ty, buf, (int)sizeof buf) <= 0) buf[0] = 0; + fprintf(stderr, "kath: %s of %s has no storage size\n", what, buf); + S.unsz = 1; } /* ---- Emission ---- */ @@ -477,10 +489,6 @@ static void is_frem(uint32_t idx, const bir_inst_t *I) /* ---- Comparison ---- */ -/* Predicate registers are 1-bit — they can't appear as operands to - * setp.*.u32. If we're comparing a predicate (e.g. branch on icmp - * result), materialise it into a u32 register first. The PTX assembler - * is surprisingly strict about this. Who knew. */ static nv_opnd_t mat_pred(nv_opnd_t op) { if (op.kind == NV_MOP_REG && op.rfile == NV_RF_PRED) { @@ -496,17 +504,9 @@ static nv_opnd_t mat_pred(nv_opnd_t op) static void is_icmp(uint32_t idx, const bir_inst_t *I) { nv_opnd_t d = map_val(idx, I->type); - /* Override: comparison result is a predicate */ - S.nv->val_rfile[idx] = NV_RF_PRED; - d.rfile = NV_RF_PRED; - nv_opnd_t a = rslv(I->operands[0]); nv_opnd_t b = rslv(I->operands[1]); - /* Predicates can't be setp source operands — widen to u32 */ - a = mat_pred(a); - b = mat_pred(b); - uint16_t op; switch (I->subop) { case BIR_ICMP_EQ: op = NV_SETP_EQ_U32; break; @@ -527,9 +527,6 @@ static void is_icmp(uint32_t idx, const bir_inst_t *I) static void is_fcmp(uint32_t idx, const bir_inst_t *I) { nv_opnd_t d = map_val(idx, I->type); - S.nv->val_rfile[idx] = NV_RF_PRED; - d.rfile = NV_RF_PRED; - uint8_t rf = rslv_rf(I->operands[0]); nv_opnd_t a = rslv(I->operands[0]); nv_opnd_t b = rslv(I->operands[1]); @@ -561,7 +558,7 @@ static void is_selp(uint32_t idx, const bir_inst_t *I) { uint8_t rf = bir_rfile(I->type); nv_opnd_t d = map_val(idx, I->type); - nv_opnd_t cond = rslv(I->operands[0]); + nv_opnd_t cond = rslv_p(I->operands[0]); nv_opnd_t tv = rslv(I->operands[1]); nv_opnd_t fv = rslv(I->operands[2]); @@ -610,9 +607,24 @@ static void is_cvt(uint32_t idx, const bir_inst_t *I) case BIR_UITOFP: op = (bir_rfile(I->type) == NV_RF_F64) ? NV_CVT_F64_U32 : NV_CVT_F32_U32; break; case BIR_FPTRUNC: op = NV_CVT_F32_F64; break; case BIR_FPEXT: op = NV_CVT_F64_F32; break; - case BIR_ZEXT: op = NV_CVT_U64_U32; break; - case BIR_SEXT: op = NV_CVT_S64_S32; break; - case BIR_TRUNC: op = NV_CVT_U32_U64; break; + case BIR_ZEXT: + case BIR_SEXT: + case BIR_TRUNC: { + uint8_t drf = S.nv->val_rfile[idx]; + uint8_t srf = rslv_rf(I->operands[0]); + if (drf == srf) { + em1u(drf == NV_RF_U64 ? NV_MOV_U64 : NV_MOV_U32, d, s); + return; + } + if (drf == NV_RF_U64 && srf == NV_RF_U32) + op = (I->op == BIR_SEXT) ? NV_CVT_S64_S32 : NV_CVT_U64_U32; + else if (drf == NV_RF_U32 && srf == NV_RF_U64) + op = NV_CVT_U32_U64; + else + op = (I->op == BIR_TRUNC) ? NV_CVT_U32_U64 : + (I->op == BIR_SEXT) ? NV_CVT_S64_S32 : NV_CVT_U64_U32; + break; + } case BIR_PTRTOINT: case BIR_INTTOPTR: case BIR_BITCAST: @@ -648,17 +660,21 @@ static void is_load(uint32_t idx, const bir_inst_t *I) nv_opnd_t addr = mat_const(ptr_val, NV_RF_U64); + uint32_t psz = pntsz(ptr_val); uint16_t op; switch (as) { case BIR_AS_SHARED: - op = (drf == NV_RF_F32) ? NV_LD_SHR_F32 : NV_LD_SHR_U32; + op = (psz == 1) ? NV_LD_SHR_U8 : + (drf == NV_RF_F32) ? NV_LD_SHR_F32 : NV_LD_SHR_U32; break; case BIR_AS_PRIVATE: - op = (drf == NV_RF_F64) ? NV_LD_LOC_F64 : + op = (psz == 1) ? NV_LD_LOC_U8 : + (drf == NV_RF_F64) ? NV_LD_LOC_F64 : (drf == NV_RF_F32) ? NV_LD_LOC_F32 : (drf == NV_RF_U64) ? NV_LD_LOC_U64 : NV_LD_LOC_U32; break; default: /* global */ + if (psz == 1) { op = NV_LD_GLB_U8; break; } switch (drf) { case NV_RF_F32: op = NV_LD_GLB_F32; break; case NV_RF_F64: op = NV_LD_GLB_F64; break; @@ -698,17 +714,21 @@ static void is_store(const bir_inst_t *I) if (val.kind != NV_MOP_REG) val = mat_const(I->operands[0], vrf); + uint32_t psz = pntsz(ptr_val); uint16_t op; switch (as) { case BIR_AS_SHARED: - op = (vrf == NV_RF_F32) ? NV_ST_SHR_F32 : NV_ST_SHR_U32; + op = (psz == 1) ? NV_ST_SHR_U8 : + (vrf == NV_RF_F32) ? NV_ST_SHR_F32 : NV_ST_SHR_U32; break; case BIR_AS_PRIVATE: - op = (vrf == NV_RF_F64) ? NV_ST_LOC_F64 : + op = (psz == 1) ? NV_ST_LOC_U8 : + (vrf == NV_RF_F64) ? NV_ST_LOC_F64 : (vrf == NV_RF_F32) ? NV_ST_LOC_F32 : (vrf == NV_RF_U64) ? NV_ST_LOC_U64 : NV_ST_LOC_U32; break; default: + if (psz == 1) { op = NV_ST_GLB_U8; break; } switch (vrf) { case NV_RF_F32: op = NV_ST_GLB_F32; break; case NV_RF_F64: op = NV_ST_GLB_F64; break; @@ -733,12 +753,13 @@ static void is_alloca(uint32_t idx, const bir_inst_t *I) * reference these as local memory. */ nv_opnd_t d = map_val(idx, I->type); - uint32_t sz = 4; /* default 4 bytes */ + uint32_t sz = 0; if (I->type < S.bir->num_types) { const bir_type_t *T = &S.bir->types[I->type]; - if (T->kind == BIR_TYPE_PTR && T->inner < S.bir->num_types) + if (T->kind == BIR_TYPE_PTR) sz = type_bytes(T->inner); } + if (!sz) { unsz("alloca", I->type); return; } uint32_t off = S.lcl_off; S.lcl_off += sz; @@ -753,12 +774,13 @@ static void is_shralloc(uint32_t idx, const bir_inst_t *I) { nv_opnd_t d = map_val(idx, I->type); - uint32_t sz = 4; + uint32_t sz = 0; if (I->type < S.bir->num_types) { const bir_type_t *T = &S.bir->types[I->type]; - if (T->kind == BIR_TYPE_PTR && T->inner < S.bir->num_types) + if (T->kind == BIR_TYPE_PTR) sz = type_bytes(T->inner); } + if (!sz) { unsz("shared_alloc", I->type); return; } uint32_t off = S.shr_off; S.shr_off += sz; @@ -784,13 +806,8 @@ static void is_gep(uint32_t idx, const bir_inst_t *I) nv_opnd_t offset = rslv(I->operands[1]); - /* Compute stride from pointee type */ - uint32_t stride = 4; - if (I->type < S.bir->num_types) { - const bir_type_t *T = &S.bir->types[I->type]; - if (T->kind == BIR_TYPE_PTR && T->inner < S.bir->num_types) - stride = type_bytes(T->inner); - } + uint32_t stride = bir_gsz(S.bir, I->type, 8); + if (!stride) { unsz("gep", I->type); return; } /* mad.lo.u64 %rd, index, stride, base */ if (offset.kind == NV_MOP_REG) { @@ -875,7 +892,7 @@ static void is_br(const bir_inst_t *I) static void is_brcond(const bir_inst_t *I) { - nv_opnd_t cond = rslv(I->operands[0]); + nv_opnd_t cond = rslv_p(I->operands[0]); uint32_t true_bir = I->operands[1]; uint32_t false_bir = I->operands[2]; @@ -937,7 +954,7 @@ static void is_phi(uint32_t idx, const bir_inst_t *I, * Only reject genuinely out-of-range values. */ if (m_pred >= NV_MAX_MBLK) continue; - nv_opnd_t src = rslv(val); + nv_opnd_t src = rslv_p(val); /* Self-copy: skip */ if (src.kind == NV_MOP_REG && src.rfile == d.rfile && @@ -1036,7 +1053,7 @@ static void is_vote(uint32_t idx, const bir_inst_t *I) { nv_opnd_t d = map_val(idx, I->type); uint32_t pred_op = I->operands[1]; - nv_opnd_t pred = rslv(pred_op); + nv_opnd_t pred = rslv_p(pred_op); if (pred.kind != NV_MOP_REG || pred.rfile != NV_RF_PRED) { uint16_t prn = new_vreg(NV_RF_PRED); @@ -1130,6 +1147,121 @@ static void is_ret(const bir_inst_t *I) /* ---- Per-Block Instruction Selection ---- */ + +#define MMA_N 16 + +static nv_opnd_t mma_u32(void) +{ + return mop_reg(NV_RF_U32, new_vreg(NV_RF_U32)); +} + +static nv_opnd_t mma_off(nv_opnd_t a, int32_t k) +{ + if (k == 0) return a; + nv_opnd_t d = mma_u32(); + em1(NV_ADD_U32, d, a, mop_imm(k)); + return d; +} + +static nv_opnd_t mma_adr(nv_opnd_t base, nv_opnd_t ld, + nv_opnd_t row, nv_opnd_t col, int32_t esz) +{ + nv_opnd_t t = mma_u32(); + em1(NV_MUL_LO_U32, t, row, ld); + nv_opnd_t e = mma_u32(); + em1(NV_ADD_U32, e, t, col); + nv_opnd_t w = mop_reg(NV_RF_U64, new_vreg(NV_RF_U64)); + em1u(NV_CVT_U64_U32, w, e); + nv_opnd_t a = mop_reg(NV_RF_U64, new_vreg(NV_RF_U64)); + nv_opnd_t ops[4] = { a, w, mop_imm(esz), base }; + emit(NV_MAD_LO_U64, 1, 3, ops, 0); + return a; +} + +static uint16_t mma_frag(uint8_t rf, int n) +{ + uint16_t b = new_vreg(rf); + for (int i = 1; i < n; i++) (void)new_vreg(rf); + return b; +} + +static void is_mma(const bir_inst_t *I) +{ + const nv_mmash_t *sh = &nv_mmash[I->subop % NV_MMA_NSHAPE]; + nv_opnd_t ap = mat_const(I->operands[0], NV_RF_U64); + nv_opnd_t la = mat_const(I->operands[1], NV_RF_U32); + nv_opnd_t bp = mat_const(I->operands[2], NV_RF_U64); + nv_opnd_t lb = mat_const(I->operands[3], NV_RF_U32); + nv_opnd_t dp = mat_const(I->operands[4], NV_RF_U64); + nv_opnd_t ld = mat_const(I->operands[5], NV_RF_U32); + + em0(NV_BARWARP); + + nv_opnd_t lane = mma_u32(); + em1u(NV_MOV_U32, lane, mop_spec(NV_SPEC_LANEID)); + nv_opnd_t gid = mma_u32(); + em1(NV_SHR_U32, gid, lane, mop_imm(2)); + nv_opnd_t tig = mma_u32(); + em1(NV_AND_B32, tig, lane, mop_imm(3)); + nv_opnd_t t2 = mma_u32(); + em1(NV_SHL_B32, t2, tig, mop_imm(1)); + + uint16_t abase = mma_frag(NV_RF_B32, sh->na / 2); + for (uint8_t i = 0; i < sh->na; i += 2) { + nv_opnd_t h[2]; + for (uint8_t j = 0; j < 2; j++) { + uint8_t e = (uint8_t)(i + j); + nv_opnd_t row = mma_off(gid, ((e & 2) != 0) ? 8 : 0); + nv_opnd_t col = mma_off(t2, (e & 1) + ((e >= 4) ? 8 : 0)); + nv_opnd_t a = mma_adr(ap, la, row, col, 2); + h[j] = mop_reg(NV_RF_U16, new_vreg(NV_RF_U16)); + em1u(NV_LD_GLB_U16, h[j], a); + } + nv_opnd_t pk[3] = { mop_reg(NV_RF_B32, (uint16_t)(abase + i / 2)), + h[0], h[1] }; + emit(NV_MOV_PK16, 1, 2, pk, 0); + } + + for (int nb = 0; nb < MMA_N; nb += 8) { + uint16_t bbase = mma_frag(NV_RF_B32, sh->nb / 2); + for (uint8_t i = 0; i < sh->nb; i += 2) { + nv_opnd_t h[2]; + for (uint8_t j = 0; j < 2; j++) { + uint8_t e = (uint8_t)(i + j); + nv_opnd_t row = mma_off(t2, (e & 1) + ((e >= 2) ? 8 : 0)); + nv_opnd_t col = mma_off(gid, nb); + nv_opnd_t a = mma_adr(bp, lb, row, col, 2); + h[j] = mop_reg(NV_RF_U16, new_vreg(NV_RF_U16)); + em1u(NV_LD_GLB_U16, h[j], a); + } + nv_opnd_t pk[3] = { mop_reg(NV_RF_B32, (uint16_t)(bbase + i / 2)), + h[0], h[1] }; + emit(NV_MOV_PK16, 1, 2, pk, 0); + } + + nv_opnd_t da[4]; + uint16_t cbase = mma_frag(NV_RF_F32, 4); + for (int i = 0; i < 4; i++) { + nv_opnd_t row = mma_off(gid, (i >= 2) ? 8 : 0); + nv_opnd_t col = mma_off(t2, (i & 1) + nb); + da[i] = mma_adr(dp, ld, row, col, 4); + em1u(NV_LD_GLB_F32, mop_reg(NV_RF_F32, (uint16_t)(cbase + i)), + da[i]); + } + uint16_t dbase = mma_frag(NV_RF_F32, 4); + nv_opnd_t mo[4] = { mop_reg(NV_RF_F32, dbase), + mop_reg(NV_RF_B32, abase), + mop_reg(NV_RF_B32, bbase), + mop_reg(NV_RF_F32, cbase) }; + emit(NV_MMA, 1, 3, mo, (uint16_t)(I->subop % NV_MMA_NSHAPE)); + for (int i = 0; i < 4; i++) { + nv_opnd_t st[2] = { da[i], + mop_reg(NV_RF_F32, (uint16_t)(dbase + i)) }; + emit(NV_ST_GLB_F32, 0, 2, st, 0); + } + } +} + static void isel_blk(uint32_t bir_bi) { const bir_block_t *B = &S.bir->blocks[bir_bi]; @@ -1254,7 +1386,16 @@ static void isel_blk(uint32_t bir_bi) /* ---- Stubs ---- */ case BIR_CALL: case BIR_INLINE_ASM: - case BIR_GLOBAL_REF: case BIR_MFMA: + case BIR_GLOBAL_REF: + break; + + case BIR_MFMA: + case BIR_MFRG: + nv_refuse("MFMA (AMD matrix intrinsic on a PTX target)"); + break; + + case BIR_MMA: + is_mma(I); break; default: @@ -1319,11 +1460,38 @@ static nv_minst_t mk_setp(nv_opnd_t dst, nv_opnd_t a, nv_opnd_t b) return I; } +static nv_minst_t mk_selp(nv_opnd_t dst, nv_opnd_t src) +{ + nv_minst_t I; + memset(&I, 0, sizeof(I)); + I.op = NV_SELP_U32; + I.num_defs = 1; + I.num_uses = 3; + I.ops[0] = dst; + I.ops[1] = mop_imm(1); + I.ops[2] = mop_imm(0); + I.ops[3] = src; + return I; +} + /* ---- Emit one PHI copy into a bridge block ---- * Appends the copy instruction(s) to the bridge block. * Called BEFORE the bridge's terminating bra is emitted. */ static void brg_copy(nv_pcopy_t *pc) { + if (pc->src.kind == NV_MOP_REG && pc->src.rfile == NV_RF_PRED + && pc->mop != NV_MOV_PRED) { + nv_opnd_t tmp = pc->dst; + if (tmp.rfile != NV_RF_U32) + tmp = mop_reg(NV_RF_U32, new_vreg(NV_RF_U32)); + nv_opnd_t s_ops[4] = { tmp, mop_imm(1), mop_imm(0), pc->src }; + emit(NV_SELP_U32, 1, 3, s_ops, 0); + if (tmp.reg_num != pc->dst.reg_num || tmp.rfile != pc->dst.rfile) { + nv_opnd_t c_ops[2] = { pc->dst, tmp }; + emit(pc->mop, 1, 1, c_ops, 0); + } + return; + } if (pc->mop == NV_MOV_PRED && pc->src.kind == NV_MOP_IMM) { nv_opnd_t ops[3] = { pc->dst, pc->src, mop_imm(0) }; emit(NV_SETP_NE_U32, 1, 2, ops, 0); @@ -1497,7 +1665,19 @@ static void phi_fix(void) nv_minst_t insts[2]; uint32_t count = 0; - if (pc->mop == NV_MOV_PRED && pc->src.kind == NV_MOP_IMM) { + if (pc->src.kind == NV_MOP_REG && pc->src.rfile == NV_RF_PRED + && pc->mop != NV_MOV_PRED) { + nv_opnd_t tmp = pc->dst; + if (tmp.rfile != NV_RF_U32) + tmp = mop_reg(NV_RF_U32, new_vreg(NV_RF_U32)); + insts[0] = mk_selp(tmp, pc->src); + count = 1; + if (tmp.reg_num != pc->dst.reg_num + || tmp.rfile != pc->dst.rfile) { + insts[1] = mk_mov(pc->mop, pc->dst, tmp); + count = 2; + } + } else if (pc->mop == NV_MOV_PRED && pc->src.kind == NV_MOP_IMM) { insts[0] = mk_setp(pc->dst, pc->src, mop_imm(0)); count = 1; } else if (pc->mop == NV_MOV_PRED && @@ -1639,5 +1819,5 @@ int nv_compile(const bir_module_t *bir, nv_module_t *nv) if (rc != BC_OK) return rc; } - return BC_OK; + return (S.unsz || S.had_error) ? BC_ERR_NVIDIA : BC_OK; } diff --git a/src/nvidia/nvidia.h b/src/nvidia/nvidia.h index 3629998..40e8b03 100644 --- a/src/nvidia/nvidia.h +++ b/src/nvidia/nvidia.h @@ -90,14 +90,14 @@ typedef enum { NV_ST_GLB_U8, NV_ST_GLB_U16, /* Loads / stores — shared */ - NV_LD_SHR_U32, NV_LD_SHR_F32, - NV_ST_SHR_U32, NV_ST_SHR_F32, + NV_LD_SHR_U32, NV_LD_SHR_F32, NV_LD_SHR_U8, + NV_ST_SHR_U32, NV_ST_SHR_F32, NV_ST_SHR_U8, /* Loads / stores — local (scratch / alloca) */ NV_LD_LOC_U32, NV_LD_LOC_U64, - NV_LD_LOC_F32, NV_LD_LOC_F64, + NV_LD_LOC_F32, NV_LD_LOC_F64, NV_LD_LOC_U8, NV_ST_LOC_U32, NV_ST_LOC_U64, - NV_ST_LOC_F32, NV_ST_LOC_F64, + NV_ST_LOC_F32, NV_ST_LOC_F64, NV_ST_LOC_U8, /* Parameter loads */ NV_LD_PARAM_U32, NV_LD_PARAM_U64, @@ -150,6 +150,11 @@ typedef enum { NV_EXIT, /* exit; */ NV_MOV_F64_LIT, /* mov.f64 %fd, 0dXXXX — ops[0]=dst, ops[1].imm=hi32, ops[2].imm=lo32 */ NV_LEA_LOCAL, /* mov.u64 %rd, __local+off — ops[0]=dst, ops[1].imm=byte offset */ + NV_MOV_PK16, /* mov.b32 %rb, {%rh_lo, %rh_hi} — packs two halves */ + NV_MMA, /* mma.sync.aligned..row.col.f32...f32. + * ops are the BASE register of each fragment tuple: + * [0]=D [1]=A [2]=B [3]=C; flags = nv_mmash_t index */ + NV_BARWARP, /* bar.warp.sync 0xffffffff */ NV_OP_COUNT } nv_ptx_op_t; @@ -164,6 +169,7 @@ typedef enum { NV_RF_PRED = 4, /* %p — predicate */ NV_RF_U16 = 5, /* %rh — 16-bit integer */ NV_RF_F16 = 6, /* %h — 16-bit float */ + NV_RF_B32 = 7, /* %rb — untyped 32-bit, mma fragment halves */ NV_RF_COUNT } nv_rfile_t; @@ -192,6 +198,15 @@ typedef enum { #define NV_SPEC_NCTAID_X 9 #define NV_SPEC_NCTAID_Y 10 #define NV_SPEC_NCTAID_Z 11 +#define NV_SPEC_LANEID 12 + +typedef struct { + const char *sfx; /* everything between "aligned." and the operands */ + uint8_t na, nb; /* A and B elements per lane */ +} nv_mmash_t; + +#define NV_MMA_NSHAPE 4 +extern const nv_mmash_t nv_mmash[NV_MMA_NSHAPE]; typedef struct { uint8_t kind; /* nv_mop_t */ diff --git a/src/tensix/isel.c b/src/tensix/isel.c index 9b63747..d71cfe8 100644 --- a/src/tensix/isel.c +++ b/src/tensix/isel.c @@ -740,6 +740,12 @@ static void isel_block(uint32_t bir_bi) S.had_error = 1; break; + case BIR_MMA: case BIR_MFRG: + fprintf(stderr, "kath: warp-collective mma not supported on the " + "Tensix backend\n"); + S.had_error = 1; + break; + /* ---- Misc ---- */ case BIR_SELECT: isel_select(idx, I); diff --git a/src/tensix/rv_isel.c b/src/tensix/rv_isel.c index e5af9f2..e9dd142 100644 --- a/src/tensix/rv_isel.c +++ b/src/tensix/rv_isel.c @@ -938,20 +938,14 @@ static int sel_unrch(rv_buf_t *out) return emit(out, rv_ebreak()); } -static uint32_t tybytes(const bir_module_t *M, uint32_t ti); +static uint32_t gepbytes(const bir_module_t *M, uint32_t ti); -/* Access width in bytes from the pointer's pointee type. Returns 0 when the - * type is not a pointer we can size, which is what an untyped pointer into L1 - * looks like, and the caller then assumes a word. Widths we cannot express in - * one RV32 access come back as-is so the caller can refuse. */ static uint32_t accwid(const bir_module_t *M, uint32_t ptr_val) { if (BIR_VAL_IS_CONST(ptr_val)) return 0u; uint32_t idx = BIR_VAL_INDEX(ptr_val); if (idx >= M->num_insts) return 0u; - uint32_t pt = M->insts[idx].type; - if (pt >= M->num_types || M->types[pt].kind != BIR_TYPE_PTR) return 0u; - return tybytes(M, M->types[pt].inner); + return gepbytes(M, M->insts[idx].type); } /* Pick the load or store for an access width, or refuse. An i64 or an @@ -960,7 +954,6 @@ static uint32_t accwid(const bir_module_t *M, uint32_t ptr_val) static int winsn(uint32_t w, int is_load, uint32_t *insn) { switch (w) { - case 0u: /* unsizable pointee; assume a word, as the isel always has */ case 4u: *insn = is_load ? rv_lw (RV_T0, RV_T1, 0) : rv_sw(RV_T0, RV_T1, 0); return BC_OK; case 2u: *insn = is_load ? rv_lhu(RV_T0, RV_T1, 0) : rv_sh(RV_T0, RV_T1, 0); return BC_OK; case 1u: *insn = is_load ? rv_lbu(RV_T0, RV_T1, 0) : rv_sb(RV_T0, RV_T1, 0); return BC_OK; @@ -1031,49 +1024,14 @@ static uint32_t algnup(uint32_t x, uint32_t a) return (x + (a - 1u)) & ~(a - 1u); } -/* - * Byte size of a BIR type. Primitives are immediate. Aggregates - * recurse and apply natural alignment between fields, which is - * the layout C/CUDA front ends produce in the absence of - * explicit packing or alignment attributes. - * - * Returns 0 if any subcomponent is unknown (e.g. a struct that - * contains a FUNC or an as-yet-unhandled type kind). Callers must - * check and refuse rather than silently computing a wrong offset. - */ static uint32_t tybytes(const bir_module_t *M, uint32_t ti) { - if (ti >= M->num_types) return 0u; - const bir_type_t *t = &M->types[ti]; - switch (t->kind) { - case BIR_TYPE_INT: return (uint32_t)(t->width / 8u); - case BIR_TYPE_FLOAT: return (uint32_t)(t->width / 8u); - case BIR_TYPE_BFLOAT: return 2u; - case BIR_TYPE_PTR: return 4u; - case BIR_TYPE_ARRAY: { - uint32_t es = tybytes(M, t->inner); - if (es == 0u) return 0u; - return es * t->count; - } - case BIR_TYPE_VECTOR: { - uint32_t es = tybytes(M, t->inner); - if (es == 0u) return 0u; - return es * (uint32_t)t->width; /* width = lane count for VECTOR */ - } - case BIR_TYPE_STRUCT: { - uint32_t off = 0u; - for (uint16_t i = 0; i < t->num_fields; i++) { - if (t->count + i >= M->num_type_fields) return 0u; - uint32_t ft = M->type_fields[t->count + i]; - uint32_t fs = tybytes(M, ft); - if (fs == 0u) return 0u; - off = algnup(off, natalg(fs)) + fs; - } - return off; /* no tail padding for now */ - } - default: - return 0u; - } + return bir_bsz(M, ti, 4); +} + +static uint32_t gepbytes(const bir_module_t *M, uint32_t ti) +{ + return bir_gsz(M, ti, 4); } /* @@ -1227,7 +1185,7 @@ static int sel_gep(const bir_module_t *M, uint32_t inst_idx, off = (int64_t)sfldof(M, base_pointee, (uint32_t)idx_v); } else { - uint32_t elem_sz = tybytes(M, M->types[I->type].inner); + uint32_t elem_sz = gepbytes(M, I->type); if (elem_sz == 0u) { fprintf(stderr, "rv_isel: unknown GEP pointee size\n"); return BC_ERR_TDF; @@ -1248,7 +1206,7 @@ static int sel_gep(const bir_module_t *M, uint32_t inst_idx, "rv_isel: non-constant struct GEP index is illegal\n"); return BC_ERR_TDF; } - uint32_t elem_sz = tybytes(M, M->types[I->type].inner); + uint32_t elem_sz = gepbytes(M, I->type); if (elem_sz == 0u) { fprintf(stderr, "rv_isel: unknown GEP pointee size\n"); return BC_ERR_TDF; @@ -1411,7 +1369,12 @@ int rv_isel_func(const bir_module_t *M, uint32_t func_idx, M->types[I->type].kind == BIR_TYPE_PTR) { pointee_sz = tybytes(M, M->types[I->type].inner); } - if (pointee_sz == 0u) pointee_sz = 4u; /* default i32-shaped */ + if (pointee_sz == 0u) { + fprintf(stderr, + "rv_isel: alloca at idx %u has unsizable " + "pointee type\n", idx); + return BC_ERR_TDF; + } alctot = rndup(alctot, ISEL_ALLOCA_ALIGN); alcoff[lidx] = alctot; alctot += pointee_sz; diff --git a/src/triton/lower.c b/src/triton/lower.c index b6f8090..646932c 100644 --- a/src/triton/lower.c +++ b/src/triton/lower.c @@ -1117,6 +1117,18 @@ static uint32_t l_expr(tn_lower_t *L, uint32_t node_idx) static void l_stmt(tn_lower_t *L, uint32_t node_idx); +static void l_rebnd(tn_lower_t *L, uint32_t tgt, uint32_t decl, uint32_t v) +{ + if (decl >= TN_MAX_NODES) return; + if (L->node_val[decl] != BIR_VAL_NONE && L->vald[decl] < L->loopd) { + l_err(L, 141, l_tok(L, tgt), + "loop-carried reassignment not yet lowered"); + return; + } + L->node_val[decl] = v; + L->vald[decl] = (uint8_t)L->loopd; +} + static void l_assign(tn_lower_t *L, uint32_t node_idx) { uint32_t value_node = l_kid(L, node_idx, 1); @@ -1128,6 +1140,15 @@ static void l_assign(tn_lower_t *L, uint32_t node_idx) return; } uint32_t v = l_expr(L, value_node); + uint32_t tgt = l_kid(L, node_idx, 0); + if (v != BIR_VAL_NONE && tgt && + L->parser->nodes[tgt].kind == TN_NK_NAME) { + int k = L->sema->node_sym_kind[tgt]; + if (k == TN_SYM_LOCAL || k == TN_SYM_LOOPVAR) { + l_rebnd(L, tgt, L->sema->node_sym_aux[tgt], v); + return; + } + } L->node_val[node_idx] = v; } @@ -1189,7 +1210,10 @@ static void l_for(tn_lower_t *L, uint32_t node_idx) uint32_t bodyb = l_new_block(L, "for.body"); l_op(L, brc, bodyb); /* [1]=true */ L->cur_block = bodyb; L->node_val[node_idx] = kphi; /* bind the loop variable k */ + L->vald[node_idx] = (uint8_t)(L->loopd + 1); + L->loopd++; l_block(L, body); + L->loopd--; uint32_t kn = l_emit(L, BIR_ADD, L->t_i32, 0); l_op(L,kn,kphi); l_op(L,kn,step); l_op(L, kphi, bodyb); l_op(L, kphi, kn); /* phi back-edge pair [body: k+step] */ uint32_t brh = l_emit(L, BIR_BR, L->t_void, 0); l_op(L,brh,head); /* back-edge */ @@ -1319,14 +1343,11 @@ static void l_stmt(tn_lower_t *L, uint32_t node_idx) l_op(L, inst, old_val); l_op(L, inst, rhs_val); - /* Re-bind the local to the new value so subsequent - * references read the updated state. */ int kind = L->sema->node_sym_kind[target_idx]; - uint32_t aux = L->sema->node_sym_aux[target_idx]; - if ((kind == TN_SYM_LOCAL || kind == TN_SYM_LOOPVAR) && - aux < TN_MAX_NODES) { - L->node_val[aux] = inst; - } + if (kind == TN_SYM_LOCAL || kind == TN_SYM_LOOPVAR) + l_rebnd(L, target_idx, L->sema->node_sym_aux[target_idx], inst); + else if (kind == TN_SYM_PARAM) + l_rebnd(L, target_idx, node_idx, inst); break; } case TN_NK_IF: diff --git a/src/triton/sema.c b/src/triton/sema.c index c2908f4..07dfb99 100644 --- a/src/triton/sema.c +++ b/src/triton/sema.c @@ -433,15 +433,17 @@ static void s_bind_assign_target(tn_sema_t *S, uint32_t target_idx, if (t->kind == TN_NK_NAME) { const tn_tok_t *tk = s_name_tok(S, target_idx); if (!tk) return; - /* Bind the local with aux = the assign node that introduced - * it. Lowering uses aux to find the BIR value the RHS - * produced, so it must point at the declaring node even on a - * re-assign. */ - if (!s_lookup(S, P->lex->src, tk->off, (uint16_t)tk->len)) { + const tn_sym_t *prev = s_lookup(S, P->lex->src, tk->off, + (uint16_t)tk->len); + uint32_t decl = assign_node; + if (!prev || prev->kind == TN_SYM_PARAM) { s_bind(S, tk->off, (uint16_t)tk->len, TN_SYM_LOCAL, assign_node, assign_node); + } else if (prev->kind == TN_SYM_LOCAL || + prev->kind == TN_SYM_LOOPVAR) { + decl = prev->aux; } - s_annotate(S, target_idx, TN_SYM_LOCAL, assign_node); + s_annotate(S, target_idx, TN_SYM_LOCAL, decl); return; } @@ -566,13 +568,20 @@ static void s_walk(tn_sema_t *S, uint32_t node_idx) return; } - case TN_NK_AUG_ASSIGN: - /* a += b: a must already be bound (Python's rule), so walk - * both sides without introducing a new binding. */ + case TN_NK_AUG_ASSIGN: { for (uint32_t i = 0; i < nk; i++) { s_walk(S, s_kid(S, node_idx, i)); } + uint32_t atgt = (nk > 0) ? s_kid(S, node_idx, 0) : 0; + if (atgt && atgt < P->num_nodes && + P->nodes[atgt].kind == TN_NK_NAME && + S->node_sym_kind[atgt] == TN_SYM_PARAM) { + const tn_tok_t *tk = s_name_tok(S, atgt); + if (tk) s_bind(S, tk->off, (uint16_t)tk->len, + TN_SYM_LOCAL, node_idx, node_idx); + } return; + } case TN_NK_FOR: { /* kids: target, iter, body, optional else. */ diff --git a/src/triton/triton.h b/src/triton/triton.h index c73363e..0d20878 100644 --- a/src/triton/triton.h +++ b/src/triton/triton.h @@ -491,6 +491,9 @@ typedef struct { * declaring node generated. */ uint32_t node_val[TN_MAX_NODES]; + int loopd; + uint8_t vald[TN_MAX_NODES]; + /* Rank-2 (and rank-1) tiles in a kernel that uses tl.dot are * materialised and fully unrolled: each tile is an array of * per-element BIR scalar values. tile_mode turns the whole kernel diff --git a/tests/bpad.cu b/tests/bpad.cu new file mode 100644 index 0000000..2ec86fa --- /dev/null +++ b/tests/bpad.cu @@ -0,0 +1,7 @@ +struct Pad { char t; int v; }; + +__global__ void bpad(int *out, const Pad *in) +{ + int i = threadIdx.x; + out[i] = in[i].v; +} diff --git a/tests/bstr.cu b/tests/bstr.cu new file mode 100644 index 0000000..0cdcb84 --- /dev/null +++ b/tests/bstr.cu @@ -0,0 +1,5 @@ +__global__ void ccopy(char *out, const char *in) +{ + int i = threadIdx.x; + out[i] = in[i]; +} diff --git a/tests/bstride.cu b/tests/bstride.cu new file mode 100644 index 0000000..615a00c --- /dev/null +++ b/tests/bstride.cu @@ -0,0 +1,27 @@ +__global__ void bglb(bool *out, const bool *in) +{ + int i = threadIdx.x; + out[i] = in[i]; +} + +__global__ void bshr(bool *out, const bool *in) +{ + int i = threadIdx.x; + __shared__ bool sh[64]; + sh[i] = in[i]; + __syncthreads(); + out[i] = sh[i]; +} + +__global__ void bloc(bool *out, const bool *in) +{ + int i = threadIdx.x; + bool loc[8]; + loc[i & 7] = in[i]; + out[i] = loc[i & 7]; +} + +__global__ void bsize(unsigned long long *out) +{ + out[0] = sizeof(bool); +} diff --git a/tests/gpu_mma.c b/tests/gpu_mma.c new file mode 100644 index 0000000..6d14847 --- /dev/null +++ b/tests/gpu_mma.c @@ -0,0 +1,129 @@ +/* gpu_mma.c -- run the 16x16x16 f16 matrix multiply on a real NVIDIA GPU. + * PTX is JITed by the driver, so no CUDA SDK is needed. + * + * kath --nvidia-ptx tests/mma16.cu -o mma16.ptx + * gcc tests/gpu_mma.c runtime/host/cuda/nv_rt.c -Iruntime/include -o gpu_mma + * ./gpu_mma mma16.ptx + */ +#include "booth/nv_rt.h" + +#include +#include +#include + +#define MM 16 +#define NN 16 +#define KK 16 +#define WARP 32 + +/* Padded row strides, so a hard-coded 16 in the address maths shows up. */ +#define LDA (KK + 3) +#define LDB (NN + 2) +#define LDD (NN + 5) + +/* Asymmetric in both indices: a layout that transposed a fragment, or + * swapped the two N halves, would not survive this. */ +#define AV(i, k) ((float)((((i) * 3 + (k) * 5) % 7) - 3)) +#define BV(k, j) ((float)((((k) * 2 + (j) * 3) % 5) - 2)) + +/* bf16 is just the top half of the f32; exact for the small integers used + * here, so there is nothing to round. */ +static unsigned short f2b(float f) +{ + union { float f; unsigned int u; } p; + p.f = f; + return (unsigned short)(p.u >> 16); +} + +/* Only small integers go through here, so the exponent path is the whole + * story and there is nothing to round. */ +static unsigned short f2h(float f) +{ + union { float f; unsigned int u; } p; + p.f = f; + unsigned int s = (p.u >> 16) & 0x8000u; + int e = (int)((p.u >> 23) & 0xFFu) - 127 + 15; + unsigned int m = p.u & 0x7FFFFFu; + if (p.f == 0.0f) return (unsigned short)s; + if (e <= 0 || e >= 31) return (unsigned short)(s | 0x7C00u); + return (unsigned short)(s | ((unsigned int)e << 10) | (m >> 13)); +} + +struct shape { const char *kern; int kk; int bf; }; + +static int oneshape(nv_dev_t *dev, const char *ptx, const struct shape *sp) +{ + nv_kern_t k; + if (nv_rt_load(dev, ptx, sp->kern, &k) != NV_RT_OK) { + fprintf(stderr, "gpu_mma: load %s failed\n", sp->kern); + return 1; + } + static unsigned short ha[MM * LDA], hb[KK * LDB]; + static float hd[MM * LDD], ref[MM * LDD]; + memset(ha, 0, sizeof ha); + memset(hb, 0, sizeof hb); + for (int i = 0; i < MM; i++) + for (int q = 0; q < sp->kk; q++) + ha[i * LDA + q] = sp->bf ? f2b(AV(i, q)) : f2h(AV(i, q)); + for (int q = 0; q < sp->kk; q++) + for (int j = 0; j < NN; j++) + hb[q * LDB + j] = sp->bf ? f2b(BV(q, j)) : f2h(BV(q, j)); + for (int i = 0; i < MM * LDD; i++) + hd[i] = (float)(i * 2 + 1); + memcpy(ref, hd, sizeof ref); + for (int i = 0; i < MM; i++) + for (int j = 0; j < NN; j++) { + float acc = hd[i * LDD + j]; + for (int q = 0; q < sp->kk; q++) + acc += AV(i, q) * BV(q, j); + ref[i * LDD + j] = acc; + } + + CUdevptr da = nv_rt_alloc(dev, sizeof ha); + CUdevptr db = nv_rt_alloc(dev, sizeof hb); + CUdevptr dd = nv_rt_alloc(dev, sizeof hd); + nv_rt_h2d(dev, da, ha, sizeof ha); + nv_rt_h2d(dev, db, hb, sizeof hb); + nv_rt_h2d(dev, dd, hd, sizeof hd); + int lda = LDA, ldb = LDB, ldd = LDD; + void *args[6] = { &da, &db, &dd, &lda, &ldb, &ldd }; + int rc = nv_rt_launch(dev, &k, 1, 1, 1, WARP, 1, 1, 0, args); + nv_rt_sync(dev); + nv_rt_d2h(dev, hd, dd, sizeof hd); + nv_rt_free(dev, da); nv_rt_free(dev, db); nv_rt_free(dev, dd); + nv_rt_unload(dev, &k); + if (rc != NV_RT_OK) { fprintf(stderr, "gpu_mma: launch %s failed\n", sp->kern); return 1; } + + int bad = 0; + for (int i = 0; i < MM * LDD; i++) { + if (hd[i] != ref[i]) { + if (bad < 4) + fprintf(stderr, " %s [%d,%d] got %.1f want %.1f\n", sp->kern, + i / LDD, i % LDD, (double)hd[i], (double)ref[i]); + bad++; + } + } + printf(" %-7s k=%2d %-4s %s (%d elements)\n", sp->kern, sp->kk, + sp->bf ? "bf16" : "f16", bad ? "FAIL" : "PASS", MM * LDD); + return bad ? 1 : 0; +} + +int main(int argc, char **argv) +{ + const char *ptx = (argc > 1) ? argv[1] : "mma16.ptx"; + static const struct shape shapes[] = { + { "mma16", 16, 0 }, { "mmab16", 16, 1 }, + { "mma8", 8, 0 }, { "mmab8", 8, 1 }, + }; + nv_dev_t dev; + if (nv_rt_init(&dev) != NV_RT_OK) { + fprintf(stderr, "gpu_mma: no CUDA driver, skipping\n"); + return 77; + } + int bad = 0; + for (unsigned i = 0; i < sizeof shapes / sizeof shapes[0]; i++) + bad += oneshape(&dev, ptx, &shapes[i]); + nv_rt_shut(&dev); + printf("gpu_mma: %s\n", bad ? "FAIL" : "PASS - every shape matches the host reference"); + return bad ? 1 : 0; +} diff --git a/tests/i1esc.cu b/tests/i1esc.cu new file mode 100644 index 0000000..16096e9 --- /dev/null +++ b/tests/i1esc.cu @@ -0,0 +1,20 @@ +__device__ int ident(int x) { return x; } + +__global__ void i1esc(int *out, int a, int b) +{ + __shared__ int shr; + + out[0] = (a == 0) || (b == 0); + out[1] = (a == 0) && (b == 0); + out[2] = (a == 0) + 5; + out[3] = (b == 0) * 3; + out[4] = ident(a == 0); + out[5] = !a; + out[6] = (int)((float)(a == 0) * 4.0f); + + shr = (a == b); + __syncthreads(); + out[7] = shr; + + atomicAdd(&out[8], (b == 0)); +} diff --git a/tests/mfbf.cu b/tests/mfbf.cu new file mode 100644 index 0000000..0ed4a07 --- /dev/null +++ b/tests/mfbf.cu @@ -0,0 +1,2 @@ +extern "C" __global__ void mfbf(const float *a, const float *b, float *acc) +{ __builtin_mfma_f32_16x16x16_bf16(a, b, acc); } diff --git a/tests/mfbf.opt b/tests/mfbf.opt new file mode 100644 index 0000000..ceb2e4b --- /dev/null +++ b/tests/mfbf.opt @@ -0,0 +1,6 @@ +# Options for mfbf.cu. See test_diag.opt for the format. +# +# MFMA is a CDNA matrix core instruction. The default --amdgpu target is +# gfx1100 (RDNA 3), which has none, and PTX wants mma.sync instead. +xfail --amdgpu MFMA needs a CDNA target (--gfx90a or --gfx942) +xfail --nvidia-ptx MFMA is AMD; PTX has mma.sync diff --git a/tests/mfi8.cu b/tests/mfi8.cu new file mode 100644 index 0000000..2737b9b --- /dev/null +++ b/tests/mfi8.cu @@ -0,0 +1,2 @@ +extern "C" __global__ void mfi8(const int *a, const int *b, int *acc) +{ __builtin_mfma_i32_16x16x16_i8(a, b, acc); } diff --git a/tests/mfi8.opt b/tests/mfi8.opt new file mode 100644 index 0000000..6f0c0e7 --- /dev/null +++ b/tests/mfi8.opt @@ -0,0 +1,6 @@ +# Options for mfi8.cu. See test_diag.opt for the format. +# +# MFMA is a CDNA matrix core instruction. The default --amdgpu target is +# gfx1100 (RDNA 3), which has none, and PTX wants mma.sync instead. +xfail --amdgpu MFMA needs a CDNA target (--gfx90a or --gfx942) +xfail --nvidia-ptx MFMA is AMD; PTX has mma.sync diff --git a/tests/mfrg.cu b/tests/mfrg.cu new file mode 100644 index 0000000..eecb009 --- /dev/null +++ b/tests/mfrg.cu @@ -0,0 +1,2 @@ +extern "C" __global__ void mf16(const float *a, const float *b, float *acc) +{ __builtin_mfma_f32_16x16x16_f16(a, b, acc); } diff --git a/tests/mfrg.opt b/tests/mfrg.opt new file mode 100644 index 0000000..1b102bd --- /dev/null +++ b/tests/mfrg.opt @@ -0,0 +1,6 @@ +# Options for mfrg.cu. See test_diag.opt for the format. +# +# MFMA is a CDNA matrix core instruction. The default --amdgpu target is +# gfx1100 (RDNA 3), which has none, and PTX wants mma.sync instead. +xfail --amdgpu MFMA needs a CDNA target (--gfx90a or --gfx942) +xfail --nvidia-ptx MFMA is AMD; PTX has mma.sync diff --git a/tests/mma16.cu b/tests/mma16.cu new file mode 100644 index 0000000..b767ecc --- /dev/null +++ b/tests/mma16.cu @@ -0,0 +1,16 @@ +/* Every mma.sync shape Booth knows, one kernel each. */ +extern "C" __global__ void mma16(const short *a, const short *b, float *d, + int lda, int ldb, int ldd) +{ __builtin_mma_m16n16k16_f16(a, lda, b, ldb, d, ldd); } + +extern "C" __global__ void mmab16(const short *a, const short *b, float *d, + int lda, int ldb, int ldd) +{ __builtin_mma_m16n16k16_bf16(a, lda, b, ldb, d, ldd); } + +extern "C" __global__ void mma8(const short *a, const short *b, float *d, + int lda, int ldb, int ldd) +{ __builtin_mma_m16n16k8_f16(a, lda, b, ldb, d, ldd); } + +extern "C" __global__ void mmab8(const short *a, const short *b, float *d, + int lda, int ldb, int ldd) +{ __builtin_mma_m16n16k8_bf16(a, lda, b, ldb, d, ldd); } diff --git a/tests/mma16.opt b/tests/mma16.opt new file mode 100644 index 0000000..6adf4c8 --- /dev/null +++ b/tests/mma16.opt @@ -0,0 +1,5 @@ +# Options for mma16.cu. See test_diag.opt for the format. +# +# The warp-collective mma lowers on PTX only. AMD needs register tuples in +# the machine IR before MFMA can carry a real fragment, so it refuses. +xfail --amdgpu warp-collective mma is NVIDIA PTX only diff --git a/tests/packs.cu b/tests/packs.cu new file mode 100644 index 0000000..d23f248 --- /dev/null +++ b/tests/packs.cu @@ -0,0 +1,81 @@ +#include + +__device__ float add3(float a, float b, float c) +{ + return a + b + c; +} + +template +__global__ void pk_sum(float *out, A... a) +{ + out[0] = (0.0f + ... + (float)a); +} + +template +__global__ void pk_cnt(int *out, T t, A... a) +{ + out[0] = (int)sizeof...(a) + (int)sizeof...(A) + (int)t; +} + +template +__global__ void pk_call(float *out, A... a) +{ + out[0] = add3(a...); +} + +template +__global__ void pk_any(int *out, A... a) +{ + int r = 0; + if ((((int)a == 0) || ...)) r = 1; + out[0] = r; +} + +template +__global__ void pk_seq(int *out, A... a) +{ + int r = 0; + (..., (r = r * 10 + (int)a)); + out[0] = r; +} + +template +__device__ void pk_drop(A&&...) +{ +} + +template +__device__ int pk_none(void) +{ + return 0; +} + +template +__device__ T pk_first(T x) +{ + return x; +} + +__device__ int pk_varg(int a, ...) +{ + return a; +} + +int main(void) +{ + float *d; + int *n; + cudaMalloc(&d, 64); + cudaMalloc(&n, 64); + + pk_sum<<<1, 1>>>(d, 1.0f, 2.0f, 3.0f, 4.0f); + pk_sum<<<1, 1>>>(d, 1.0f, 2.0f); + pk_cnt<<<1, 1>>>(n, 1, 2, 3); + pk_call<<<1, 1>>>(d, 1.5f, 2.5f, 3.5f); + pk_any<<<1, 1>>>(n, 1, 0, 2); + pk_seq<<<1, 1>>>(n, 1, 2, 3); + + cudaFree(d); + cudaFree(n); + return 0; +} diff --git a/tests/test_mfma.opt b/tests/test_mfma.opt index 46163bf..ccd0e02 100644 --- a/tests/test_mfma.opt +++ b/tests/test_mfma.opt @@ -1,5 +1,7 @@ # Options for test_mfma.cu. See test_diag.opt for the format. # -# This one is a real bug rather than a missing feature. reprocheck used to -# count it as a skip and say nothing, which is how it stayed quiet. -xfail --amdgpu verify rejects VGPR in scalar source of v_mfma +# Both refusals are deliberate. MFMA operands are register tuples (v[0:3]) +# and a vreg here is a single 32-bit register, so the old single-register +# form llvm-mc rejects is no longer emitted; PTX has no MFMA at all. +xfail --amdgpu MFMA needs register tuples the machine IR cannot express +xfail --nvidia-ptx MFMA is an AMD intrinsic, PTX wants mma.sync instead diff --git a/tests/tmain.c b/tests/tmain.c index 5d825b7..a2112f9 100644 --- a/tests/tmain.c +++ b/tests/tmain.c @@ -82,6 +82,7 @@ static const tfam_t fam_order[] = { { "err", "terrs.c", "diagnostics", 2 }, { "typ", "ttypes.c", "type table", 2 }, { "tab", "ttabs.c", "static tables", 2 }, + { "pck", "tpack.c", "parameter packs", 2 }, { "dce", "tdce.c", "dead code elimination", 2 }, { "cfd", "tcfold.c", "constant folding", 2 }, @@ -105,6 +106,7 @@ static const tfam_t fam_order[] = { { "tdf", "ttdf.c", "Tensix dataflow", 2 }, { "tri", "ttriton.c", "Triton frontend", 2 }, + { "mma", "tmma.c", "warp matrix multiply", 2 }, { "mlr", "tmlir.c", "MLIR reader", 2 }, { "bir", "tbir.c", "BIR text frontend", 2 }, @@ -118,6 +120,7 @@ static const tfam_t fam_order[] = { { "ord", "tordr.c", "harness ordering", 2 }, { "ocm", "tocm.c", "OCaml frontend", 2 }, + { "mtu", "tmtu.c", "multiple translation units", 2 }, { "rpi", "trpi.c", "shipped-bug regressions", 2 }, }; diff --git a/tests/tmma.c b/tests/tmma.c new file mode 100644 index 0000000..edd7ccc --- /dev/null +++ b/tests/tmma.c @@ -0,0 +1,185 @@ +/* tmma.c -- warp-collective matrix multiply + * PTX gets mma.sync for every shape it knows, AMD gets MFMA in the register + * tuple form llvm-mc accepts, and everything else says no out loud. */ + +#include "tharns.h" + +static char obuf[TH_BUFSZ]; +static char ptx[131072]; +static unsigned char bin[65536]; + +static int mm_run(const char *args) +{ + char cmd[TH_BUFSZ]; + snprintf(cmd, TH_BUFSZ, BC_BIN " %s", args); + return th_run(cmd, obuf, TH_BUFSZ); +} + +static int mm_txt(const char *args, const char *path, char *buf, size_t cap) +{ + if (mm_run(args) != 0) return -1; + FILE *fp = fopen(path, "rb"); + if (!fp) return -1; + size_t n = fread(buf, 1, cap - 1, fp); + fclose(fp); + buf[n] = '\0'; + return 0; +} + +static int mm_ptx(void) +{ + return mm_txt("--nvidia-ptx tests/mma16.cu -o mma16.ptx", "mma16.ptx", + ptx, sizeof ptx); +} + +static int mm_cnt(const char *hay, const char *needle) +{ + int c = 0; + const char *p = hay; + while ((p = strstr(p, needle)) != NULL) { c++; p++; } + return c; +} + +/* first VOP3P-MAI word in an AMD code object, or -1 */ +static int mm_mfop(const char *args, const char *path) +{ + if (mm_run(args) != 0) return -1; + FILE *fp = fopen(path, "rb"); + if (!fp) return -1; + size_t n = fread(bin, 1, sizeof bin, fp); + fclose(fp); + for (size_t o = 0; o + 8 <= n; o++) { + unsigned w = (unsigned)bin[o] | (unsigned)bin[o+1] << 8 | + (unsigned)bin[o+2] << 16 | (unsigned)bin[o+3] << 24; + if ((w >> 23) == 0x1A7u) return (int)((w >> 16) & 0x7Fu); + } + return -1; +} + +/* ---- PTX ---- */ + +static void mma01(void) +{ + CHEQ(mm_ptx(), 0); + CHEQ(mm_cnt(ptx, "mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32"), 2); + CHEQ(mm_cnt(ptx, "mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32"), 2); + CHEQ(mm_cnt(ptx, "mma.sync.aligned.m16n8k8.row.col.f32.f16.f16.f32"), 2); + CHEQ(mm_cnt(ptx, "mma.sync.aligned.m16n8k8.row.col.f32.bf16.bf16.f32"), 2); + PASS(); +} + +/* k16 carries 4 A and 2 B registers, k8 carries 2 and 1. A tuple that lost + * a register would still assemble, so count the commas. */ +static void mma02(void) +{ + CHEQ(mm_ptx(), 0); + static const struct { const char *sfx; int commas; } want[] = { + { "m16n8k16.row.col.f32.f16.f16.f32", 13 }, /* 3 + 3+1+3 inside */ + { "m16n8k8.row.col.f32.f16.f16.f32", 10 }, /* 3 + 1+0+3 inside */ + }; + for (unsigned w = 0; w < sizeof want / sizeof want[0]; w++) { + const char *p = strstr(ptx, want[w].sfx); + CHECK(p != NULL); + if (!p) return; + const char *e = strchr(p, ';'); + CHECK(e != NULL); + if (!e) return; + int commas = 0, braces = 0; + for (const char *q = p; q < e; q++) { + if (*q == ',') commas++; + if (*q == '{') braces++; + } + CHEQ(braces, 4); + CHEQ(commas, want[w].commas); + } + PASS(); +} + +static void mma03(void) +{ + CHEQ(mm_ptx(), 0); + CHECK(strstr(ptx, "%laneid") != NULL); + CHECK(strstr(ptx, ".reg .b32") != NULL); + PASS(); +} + +/* mma.sync is undefined on a divergent warp, so each kernel converges first. */ +static void mma04(void) +{ + CHEQ(mm_ptx(), 0); + CHEQ(mm_cnt(ptx, "bar.warp.sync 0xffffffff"), 4); + PASS(); +} + +/* ---- AMD ---- */ + +/* CDNA3 ISA s7.1 wants contiguous, size-aligned operands. */ +static void mma05(void) +{ + CHEQ(mm_txt("--amdgpu --gfx942 tests/mfrg.cu -o mfrg.s", "mfrg.s", + ptx, sizeof ptx), 0); + CHECK(strstr(ptx, "v_mfma_f32_16x16x16f16 v[208:211], v[240:241], " + "v[244:245], v[224:227]") != NULL); + PASS(); +} + +/* The bf16 opcode moved between CDNA2 and CDNA3: 0x67 on gfx90a, 0x61 on + * gfx942. One table for both targets emitted the wrong instruction. */ +static void mma06(void) +{ + CHEQ(mm_mfop("--amdgpu-bin --gfx90a tests/mfbf.cu -o mfbf.hsaco", + "mfbf.hsaco"), 0x67); + CHEQ(mm_mfop("--amdgpu-bin --gfx942 tests/mfbf.cu -o mfbf.hsaco", + "mfbf.hsaco"), 0x61); + PASS(); +} + +static void mma07(void) +{ + CHEQ(mm_mfop("--amdgpu-bin --gfx942 tests/mfrg.cu -o mfrg.hsaco", + "mfrg.hsaco"), 0x4D); + CHEQ(mm_mfop("--amdgpu-bin --gfx90a tests/mfrg.cu -o mfrg.hsaco", + "mfrg.hsaco"), 0x4D); + PASS(); +} + +/* i8 16x16x16 is CDNA2 only; CDNA3 spells it 16x16x32. */ +static void mma08(void) +{ + CHNE(mm_run("--amdgpu --gfx942 tests/mfi8.cu -o mfi8.s"), 0); + CHECK(strstr(obuf, "not supported") != NULL); + CHEQ(mm_mfop("--amdgpu-bin --gfx90a tests/mfi8.cu -o mfi8.hsaco", + "mfi8.hsaco"), 0x55); + PASS(); +} + +/* ---- Refusals ---- */ + +static void mma09(void) +{ + CHNE(mm_run("--amdgpu --gfx942 tests/mma16.cu -o x.s"), 0); + CHECK(strstr(obuf, "not supported") != NULL); + CHNE(mm_run("--cpu tests/mma16.cu -o x.o"), 0); + CHECK(strstr(obuf, "not supported") != NULL); + PASS(); +} + +static void mma10(void) +{ + CHNE(mm_run("--nvidia-ptx tests/mfrg.cu -o x.ptx"), 0); + CHECK(strstr(obuf, "not supported") != NULL); + CHNE(mm_run("--amdgpu --gfx942 tests/test_mfma.cu -o x.s"), 0); + CHECK(strstr(obuf, "not supported") != NULL); + PASS(); +} + +TH_REG("mma", 1, "every mma.sync shape reaches PTX", mma01) +TH_REG("mma", 2, "fragment tuples are the right width", mma02) +TH_REG("mma", 3, "lane picks its own fragment", mma03) +TH_REG("mma", 4, "the warp converges before mma.sync", mma04) +TH_REG("mma", 5, "MFMA emits an aligned register tuple", mma05) +TH_REG("mma", 6, "bf16 opcode follows the CDNA target", mma06) +TH_REG("mma", 7, "f16 MFMA encodes the same on both", mma07) +TH_REG("mma", 8, "i8 16x16x16 is CDNA2 only", mma08) +TH_REG("mma", 9, "warp-collective mma refuses off PTX", mma09) +TH_REG("mma", 10, "AMD matrix ops refuse off CDNA", mma10) diff --git a/tests/tmtu.c b/tests/tmtu.c new file mode 100644 index 0000000..0988865 --- /dev/null +++ b/tests/tmtu.c @@ -0,0 +1,195 @@ +/* tmtu.c -- several .cu files in one compile + * + * The thing worth testing here is not that the driver loops over argv. It is + * that two files can hold their own `static` helper of the same name and each + * kernel still calls its own, because getting that wrong builds a program + * nobody wrote and says nothing about it. */ + +#include "tharns.h" + +static char obuf[TH_BUFSZ]; + +/* ---- Two files reach the backend ---- */ + +static void mtu01(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_sta.cu tests/tu_stb.cu", + obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "func @tu_ka") != NULL); + CHECK(strstr(obuf, "func @tu_kb") != NULL); + PASS(); +} +TH_REG("mtu", 1, "two files lower into one module", mtu01) + +/* ---- A static helper belongs to its own file ---- */ + +static void mtu02(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_sta.cu tests/tu_stb.cu", + obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "func @scale__0") != NULL); + CHECK(strstr(obuf, "func @scale__1") != NULL); + CHECK(strstr(obuf, "call i32 @scale__0") != NULL); + CHECK(strstr(obuf, "call i32 @scale__1") != NULL); + PASS(); +} +TH_REG("mtu", 2, "each file calls its own static helper", mtu02) + +/* The bodies differ by 1 versus 100, so if the two ever collapsed into one + * the surviving add would give it away. */ +static void mtu03(void) +{ + int rc = th_run(BC_BIN " --ir --no-cfold tests/tu_sta.cu tests/tu_stb.cu", + obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "add i32 %0, 1") != NULL); + CHECK(strstr(obuf, "add i32 %0, 100") != NULL); + PASS(); +} +TH_REG("mtu", 3, "the two static bodies both survive", mtu03) + +/* ---- Calling across files ---- */ + +static void mtu04(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_lib.cu tests/tu_use.cu", + obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "func @tu_mul7") != NULL); + CHECK(strstr(obuf, "func @tu_kuse") != NULL); + PASS(); +} +TH_REG("mtu", 4, "a __device__ function crosses files", mtu04) + +/* Symbols become visible in command-line order, so the caller listed first + * finds nothing. Refusing is the point: the alternative is a call to whatever + * happened to be lying around. */ +static void mtu05(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_use.cu tests/tu_lib.cu", + obuf, TH_BUFSZ); + CHNE(rc, 0); + CHECK(strstr(obuf, "E105") != NULL); + PASS(); +} +TH_REG("mtu", 5, "a callee listed later is not found", mtu05) + +/* ---- One symbol, two definitions ---- */ + +static void mtu06(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_lib.cu tests/tu_dup.cu", + obuf, TH_BUFSZ); + CHNE(rc, 0); + CHECK(strstr(obuf, "E126") != NULL); + PASS(); +} +TH_REG("mtu", 6, "a kernel defined twice is refused", mtu06) + +/* ---- A header both files include ---- */ + +static void mtu07(void) +{ + int rc = th_run(BC_BIN " --ir -Itests tests/tu_h1.cu tests/tu_h2.cu", + obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "func @tu_sq") != NULL); + CHECK(strstr(obuf, "func @tu_sq__0") == NULL); + CHECK(strstr(obuf, "3 functions, 1 globals") != NULL); + PASS(); +} +TH_REG("mtu", 7, "an inline header gives one copy, not two", mtu07) + +/* ---- Both kernels reach an emitter ---- */ + +static void mtu08(void) +{ + int rc = th_run(BC_BIN " --nvidia-ptx tests/tu_sta.cu tests/tu_stb.cu" + " -o mtu08.ptx", obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "2 kernels") != NULL); + remove("mtu08.ptx"); + PASS(); +} +TH_REG("mtu", 8, "two files give two PTX entry points", mtu08) + +static void mtu09(void) +{ + int rc = th_run(BC_BIN " --amdgpu-bin tests/tu_sta.cu tests/tu_stb.cu" + " -o mtu09.hsaco", obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "2 kernels") != NULL); + remove("mtu09.hsaco"); + PASS(); +} +TH_REG("mtu", 9, "two files give two AMD kernels", mtu09) + +/* ---- One file is still one file ---- */ + +/* Nothing is qualified when there is nothing to keep apart, so a single-file + * compile emits the names it always did. */ +static void mtu10(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_sta.cu", obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "func @scale") != NULL); + CHECK(strstr(obuf, "scale__0") == NULL); + PASS(); +} +TH_REG("mtu", 10, "one file keeps unqualified names", mtu10) + +/* ---- The frontends that read one document say so ---- */ + +static void mtu11(void) +{ + int rc = th_run(BC_BIN " --triton --ir tests/tri_vadd.py tests/tri_vadd.py", + obuf, TH_BUFSZ); + CHNE(rc, 0); + CHECK(strstr(obuf, "only one input file") != NULL); + PASS(); +} +TH_REG("mtu", 11, "--triton refuses a second file", mtu11) + +/* ---- The same template in both files ---- */ + +/* A template lives in a header, so both files carry their own copy of the + * declaration and both ask for the same instantiation. C++ says that is one + * function, and one is what comes out. */ +static void mtu12(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_t1.cu tests/tu_t2.cu", + obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "func @tu_tsc") != NULL); + CHECK(strstr(obuf, "1 functions") != NULL); + PASS(); +} +TH_REG("mtu", 12, "one instantiation, not one per file", mtu12) + +/* ---- A global declared in one file and defined in another ---- */ + +/* Globals are looked up where they are used rather than where the file is + * read, so unlike a call these do not mind which order the files came in. */ +static void mtu13(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_g1.cu tests/tu_g2.cu", + obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "1 globals") != NULL); + CHECK(strstr(obuf, "func @tu_kg1") != NULL); + CHECK(strstr(obuf, "func @tu_kg2") != NULL); + PASS(); +} +TH_REG("mtu", 13, "extern __constant__ resolves across files", mtu13) + +static void mtu14(void) +{ + int rc = th_run(BC_BIN " --ir tests/tu_g2.cu tests/tu_g1.cu", + obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "1 globals") != NULL); + PASS(); +} +TH_REG("mtu", 14, "and resolves in either order", mtu14) diff --git a/tests/tnv_bstr.c b/tests/tnv_bstr.c new file mode 100644 index 0000000..033e615 --- /dev/null +++ b/tests/tnv_bstr.c @@ -0,0 +1,77 @@ +/* tnv_bstr.c — does a byte array survive a round trip through the card + * + * The GEP stride for a one-byte element is one byte, and the access has to be + * one byte with it. On master PTX strode by one and loaded four, so element i + * read elements i..i+3 and a store scribbled over the three that followed. + * Sixteen distinct values in sixteen adjacent bytes is the smallest thing that + * catches it: any over-wide access loses the neighbours. + * + * Wants a real card. Not a trunner test. + * + * ./kath --nvidia-ptx tests/bstr.cu -o bstr.ptx + * ./tnv_bstr bstr.ptx + */ + +#include "booth/nv_rt.h" +#include + +#define N 16 + +int main(int argc, char **argv) +{ + const char *ptx = (argc > 1) ? argv[1] : "bstr.ptx"; + nv_dev_t dev; + nv_kern_t kern; + char src[N], dst[N]; + int errs = 0; + + for (int i = 0; i < N; i++) { src[i] = (char)(i * 7 + 1); dst[i] = 0; } + + if (nv_rt_init(&dev) != NV_RT_OK) { + fprintf(stderr, "tnv_bstr: no device\n"); + return 2; + } + printf("tnv_bstr: %s, sm_%d%d\n", dev.dev_name, dev.sm_major, dev.sm_minor); + + if (nv_rt_load(&dev, ptx, "ccopy", &kern) != NV_RT_OK) { + fprintf(stderr, "tnv_bstr: JIT refused %s\n", ptx); + nv_rt_shut(&dev); + return 1; + } + + CUdevptr d_in = nv_rt_alloc(&dev, sizeof src); + CUdevptr d_out = nv_rt_alloc(&dev, sizeof dst); + if (!d_in || !d_out) { + fprintf(stderr, "tnv_bstr: alloc failed\n"); + nv_rt_shut(&dev); + return 1; + } + + void *args[2]; + args[0] = &d_out; + args[1] = &d_in; + + if (nv_rt_h2d(&dev, d_in, src, sizeof src) != NV_RT_OK + || nv_rt_h2d(&dev, d_out, dst, sizeof dst) != NV_RT_OK + || nv_rt_launch(&dev, &kern, 1, 1, 1, N, 1, 1, 0, args) != NV_RT_OK + || nv_rt_sync(&dev) != NV_RT_OK + || nv_rt_d2h(&dev, dst, d_out, sizeof dst) != NV_RT_OK) { + fprintf(stderr, "tnv_bstr: launch or readback failed\n"); + nv_rt_free(&dev, d_in); nv_rt_free(&dev, d_out); + nv_rt_shut(&dev); + return 1; + } + + for (int i = 0; i < N; i++) { + printf(" [%2d] got %4d want %4d%s\n", i, dst[i], src[i], + dst[i] == src[i] ? "" : " <-- "); + if (dst[i] != src[i]) errs++; + } + printf("tnv_bstr: %s (%d mismatches)\n", errs ? "FAIL" : "PASS", errs); + + nv_rt_free(&dev, d_in); + nv_rt_free(&dev, d_out); + nv_rt_unload(&dev, &kern); + nv_rt_shut(&dev); + return errs ? 1 : 0; +} diff --git a/tests/tnv_i1.c b/tests/tnv_i1.c new file mode 100644 index 0000000..0347dd1 --- /dev/null +++ b/tests/tnv_i1.c @@ -0,0 +1,109 @@ +/* tnv_i1.c — does the i1 escape kernel load and give the right numbers + * + * Before the predicate materialisation fix the PTX for tests/i1esc.cu carried + * `st.global.u32 [%rd], %p` and its relatives, so the driver JIT refused the + * module and nv_rt_load failed outright. That makes this unambiguous: the load + * either happens or it does not. The values then check the other half, because + * C promises `(a==0)` is exactly 0 or 1 and `(a==0)+5` is 5 or 6, which a fix + * that clamps everything to a predicate would get wrong. + * + * Wants a real card. Not a trunner test. + * + * ./kath --nvidia-ptx tests/i1esc.cu -o i1esc.ptx + * ./tnv_i1 i1esc.ptx + */ + +#include "booth/nv_rt.h" +#include + +#define NSLOT 9 +#define ATBAS 100 + +static int host[NSLOT]; + +static int check(nv_dev_t *dev, nv_kern_t *kern, int a, int b) +{ + int want[NSLOT]; + int errs = 0; + + want[0] = (a == 0) || (b == 0); + want[1] = (a == 0) && (b == 0); + want[2] = (a == 0) + 5; + want[3] = (b == 0) * 3; + want[4] = (a == 0); + want[5] = !a; + want[6] = (int)((float)(a == 0) * 4.0f); + want[7] = (a == b); + want[8] = ATBAS + (b == 0); + + for (int i = 0; i < NSLOT; i++) host[i] = 0; + host[8] = ATBAS; + + CUdevptr d_out = nv_rt_alloc(dev, sizeof host); + if (!d_out) { fprintf(stderr, " alloc failed\n"); return 1; } + + if (nv_rt_h2d(dev, d_out, host, sizeof host) != NV_RT_OK) { + fprintf(stderr, " H2D failed\n"); + nv_rt_free(dev, d_out); + return 1; + } + + void *args[3]; + args[0] = &d_out; + args[1] = &a; + args[2] = &b; + + if (nv_rt_launch(dev, kern, 1, 1, 1, 1, 1, 1, 0, args) != NV_RT_OK + || nv_rt_sync(dev) != NV_RT_OK + || nv_rt_d2h(dev, host, d_out, sizeof host) != NV_RT_OK) { + fprintf(stderr, " launch or readback failed\n"); + nv_rt_free(dev, d_out); + return 1; + } + nv_rt_free(dev, d_out); + + printf(" a=%d b=%d ->", a, b); + for (int i = 0; i < NSLOT; i++) printf(" %d", host[i]); + printf("\n"); + + for (int i = 0; i < NSLOT; i++) { + if (host[i] != want[i]) { + fprintf(stderr, " slot %d: got %d, wanted %d\n", + i, host[i], want[i]); + errs++; + } + } + return errs; +} + +int main(int argc, char **argv) +{ + const char *ptx = (argc > 1) ? argv[1] : "i1esc.ptx"; + nv_dev_t dev; + nv_kern_t kern; + int errs = 0; + + if (nv_rt_init(&dev) != NV_RT_OK) { + fprintf(stderr, "tnv_i1: no device\n"); + return 2; + } + printf("tnv_i1: %s, sm_%d%d\n", dev.dev_name, dev.sm_major, dev.sm_minor); + + if (nv_rt_load(&dev, ptx, "i1esc", &kern) != NV_RT_OK) { + fprintf(stderr, "tnv_i1: JIT refused %s\n", ptx); + nv_rt_shut(&dev); + return 1; + } + printf("tnv_i1: %s loaded\n", ptx); + + errs += check(&dev, &kern, 0, 0); + errs += check(&dev, &kern, 0, 1); + errs += check(&dev, &kern, 1, 0); + errs += check(&dev, &kern, 1, 1); + + printf("tnv_i1: %s (%d mismatches)\n", errs ? "FAIL" : "PASS", errs); + + nv_rt_unload(&dev, &kern); + nv_rt_shut(&dev); + return errs ? 1 : 0; +} diff --git a/tests/tpack.c b/tests/tpack.c new file mode 100644 index 0000000..705304c --- /dev/null +++ b/tests/tpack.c @@ -0,0 +1,255 @@ +/* tpack.c -- template parameter packs + * The grammar, the standard's placement rules, and what a pack lowers to. */ + +#include "tharns.h" + +static char obuf[TH_BUFSZ]; + +static int wsrc(const char *path, const char *text) +{ + FILE *f = fopen(path, "wb"); + if (!f) return -1; + fputs(text, f); + fclose(f); + return 0; +} + +static int run_ir(const char *text) +{ + if (wsrc("pack_tmp.cu", text) != 0) return -1; + int rc = th_run(BC_BIN " --ir pack_tmp.cu", obuf, TH_BUFSZ); + remove("pack_tmp.cu"); + return rc; +} + +static int run_parse(const char *text) +{ + if (wsrc("pack_tmp.cu", text) != 0) return -1; + int rc = th_run(BC_BIN " --parse pack_tmp.cu", obuf, TH_BUFSZ); + remove("pack_tmp.cu"); + return rc; +} + +/* ---- packs: the fixture compiles ---- */ + +static void pck01(void) +{ + int rc = th_run(BC_BIN " --ir tests/packs.cu", obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "error") == NULL); + PASS(); +} +TH_REG("pck", 1, "the pack fixture reaches BIR", pck01) + +/* ---- packs: a pack parameter becomes one BIR parameter per member ---- */ + +static void pck02(void) +{ + int rc = th_run(BC_BIN " --ir tests/packs.cu", obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, + "func @pk_sum(ptr %0, f32 %1, f32 %2, f32 %3, f32 %4)") + != NULL); + PASS(); +} +TH_REG("pck", 2, "a pack widens the parameter list", pck02) + +/* ---- packs: two lengths are two functions ---- */ + +static void pck03(void) +{ + int rc = th_run(BC_BIN " --ir tests/packs.cu", obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, "func @pk_sum$1(") != NULL); + PASS(); +} +TH_REG("pck", 3, "two specialisations do not share a symbol", pck03) + +/* ---- packs: sizeof... is a constant by lowering ---- */ + +static void pck04(void) +{ + int rc = run_ir( + "template\n" + "__global__ void k(int *o, A... a)\n" + "{ o[0] = (int)sizeof...(a) + (int)sizeof...(A); }\n" + "int main(void){ int *o; cudaMalloc(&o,16);" + " k<<<1,1>>>(o, 1, 2, 3); return 0; }\n"); + CHEQ(rc, 0); + CHECK(strstr(obuf, "store i32 6,") != NULL); + PASS(); +} +TH_REG("pck", 4, "sizeof... folds to the pack length", pck04) + +/* ---- packs: a fold expands to the operator chain ---- */ + +static void pck05(void) +{ + int rc = run_ir( + "template\n" + "__global__ void k(float *o, A... a) { o[0] = (0.0f + ... + (float)a); }\n" + "int main(void){ float *o; cudaMalloc(&o,16);" + " k<<<1,1>>>(o, 1.0f, 2.0f, 3.0f); return 0; }\n"); + CHEQ(rc, 0); + /* seed plus one add per member */ + int n = 0; + for (const char *p = obuf; (p = strstr(p, "fadd")) != NULL; p++) n++; + CHEQ(n, 3); + PASS(); +} +TH_REG("pck", 5, "a binary left fold becomes one op per member", pck05) + +/* ---- packs: an expansion widens a call ---- */ + +static void pck06(void) +{ + int rc = th_run(BC_BIN " --ir --no-sroa tests/packs.cu", obuf, TH_BUFSZ); + CHEQ(rc, 0); + CHECK(strstr(obuf, + "func @pk_call(ptr %0, f32 %1, f32 %2, f32 %3)") != NULL); + PASS(); +} +TH_REG("pck", 6, "an expansion in a call supplies every member", pck06) + +/* ---- packs: || keeps its short circuit ---- */ + +static void pck07(void) +{ + int rc = th_run(BC_BIN " --ir tests/packs.cu", obuf, TH_BUFSZ); + CHEQ(rc, 0); + const char *f = strstr(obuf, "func @pk_any"); + CHECK(f != NULL); + CHECK(strstr(f, "br_cond") != NULL); + PASS(); +} +TH_REG("pck", 7, "a logical fold branches rather than chains", pck07) + +/* ---- packs: an unexpanded pack is ill-formed ---- */ + +static void pck08(void) +{ + int rc = run_parse( + "template\n" + "__device__ int k(A... a) { return f(a); }\n"); + (void)rc; + CHECK(strstr(obuf, "error[E083]") != NULL); + PASS(); +} +TH_REG("pck", 8, "a pack used without an expansion is rejected", pck08) + +/* ---- packs: [temp.param]/14 placement ---- */ + +static void pck09(void) +{ + int rc = run_parse( + "template struct S { B v; };\n"); + (void)rc; + CHECK(strstr(obuf, "error[E028]") != NULL); + PASS(); +} +TH_REG("pck", 9, "a class template pack must come last", pck09) + +/* ---- packs: no default argument on a pack ---- */ + +static void pck10(void) +{ + int rc = run_parse("template struct S { int v; };\n"); + (void)rc; + CHECK(strstr(obuf, "error[E029]") != NULL); + PASS(); +} +TH_REG("pck", 10, "a pack takes no default argument", pck10) + +/* ---- packs: pack indexing refuses by name ---- */ + +static void pck11(void) +{ + int rc = run_parse( + "template\n" + "__device__ int k(A... a) { return (int)a...[0]; }\n"); + (void)rc; + CHECK(strstr(obuf, "error[E030]") != NULL); + CHECK(strstr(obuf, "pack indexing") != NULL); + PASS(); +} +TH_REG("pck", 11, "pack indexing refuses by name", pck11) + +/* ---- packs: [dcl.fct]/27 keeps the C ellipsis ---- */ + +static void pck12(void) +{ + int rc = run_ir( + "__device__ int v(int a, ...) { return a; }\n" + "__global__ void k(int *o) { o[0] = v(1); }\n"); + CHEQ(rc, 0); + CHECK(strstr(obuf, "error") == NULL); + PASS(); +} +TH_REG("pck", 12, "a C ellipsis is still a C ellipsis", pck12) + +/* ---- packs: a forwarding-reference pack parses ---- */ + +static void pck13(void) +{ + int rc = run_parse( + "template\n" + "__host__ __device__ constexpr inline void u(A&&...) noexcept {}\n"); + (void)rc; + CHECK(strstr(obuf, "error[") == NULL); + PASS(); +} +TH_REG("pck", 13, "an unnamed forwarding pack parses", pck13) + +/* ---- packs: the pack need not be last in a function template ---- */ + +static void pck14(void) +{ + int rc = run_parse( + "template\n" + "__device__ T k(T x) { return x; }\n"); + (void)rc; + CHECK(strstr(obuf, "error[") == NULL); + PASS(); +} +TH_REG("pck", 14, "a function template may declare after a pack", pck14) + +/* ---- packs: an overload set is chosen by arity ---- */ + +static void pck15(void) +{ + /* Both overloads inline, so only the emitted constant tells them apart. */ + CHECK(wsrc("pack_tmp.cu", + "__device__ int g(int a, int b) { return a * 1000 + b; }\n" + "__device__ int g(int a) { return a + 7; }\n" + "__global__ void k(int *o) { o[0] = g(5); }\n" + "int main(void){ int *o; cudaMalloc(&o,16);" + " k<<<1,1>>>(o); return 0; }\n") == 0); + int rc = th_run(BC_BIN " --nvidia-ptx pack_tmp.cu -o pack_tmp.ptx", + obuf, TH_BUFSZ); + remove("pack_tmp.cu"); + CHEQ(rc, 0); + FILE *pf = fopen("pack_tmp.ptx", "rb"); + CHECK(pf != NULL); + size_t got = fread(obuf, 1, TH_BUFSZ - 1, pf); + obuf[got] = 0; + fclose(pf); + remove("pack_tmp.ptx"); + CHECK(strstr(obuf, "mov.u32 %r1, 12;") != NULL); + PASS(); +} +TH_REG("pck", 15, "a call binds to the overload with its arity", pck15) + +/* ---- packs: a pack that cannot be deduced refuses ---- */ + +static void pck16(void) +{ + int rc = run_ir( + "template\n" + "__global__ void k(float *o, A... a) { o[0] = (0.0f + ... + (float)a); }\n" + "int main(void){ float *o; float x = 1.0f; cudaMalloc(&o,16);" + " k<<<1,1>>>(o, x); return 0; }\n"); + (void)rc; + CHECK(strstr(obuf, "E030") != NULL); + PASS(); +} +TH_REG("pck", 16, "an undeducible pack refuses, never guesses", pck16) diff --git a/tests/tphase.c b/tests/tphase.c index 73a670f..82b05f8 100644 --- a/tests/tphase.c +++ b/tests/tphase.c @@ -5,6 +5,54 @@ static char obuf[TH_BUFSZ]; +/* ---- phase 2: line splicing ---- */ + +/* Fixtures are written here rather than committed. .gitattributes checks the + * tree out as eol=lf, so a CRLF file in tests/ would arrive as an LF one and + * every CRLF test below would pass without testing anything. */ + +#define SP_LF "tests/spl_lf.cu" +#define SP_CRLF "tests/spl_crlf.cu" + +static int sp_wr(const char *path, const char *body, int crlf) +{ + FILE *fp = fopen(path, "wb"); + if (!fp) return 0; + for (const char *p = body; *p != '\0'; p++) { + if (*p == '\n' && crlf) fputc('\r', fp); + fputc(*p, fp); + } + fclose(fp); + return 1; +} + +/* Token count out of the "N tokens, M error(s)" tail, -1 if it isn't there. */ +static int sp_ntok(const char *out) +{ + const char *p = strstr(out, " tokens,"); + if (p == NULL) return -1; + const char *q = p; + while (q > out && q[-1] >= '0' && q[-1] <= '9') q--; + if (q == p) return -1; + return atoi(q); +} + +/* Lex body under both line endings. n[i] is the token count where the lexer + * reported none of its own errors, -1 where it did. */ +static void sp_both(const char *body, int n[2]) +{ + static const char *const path[2] = { SP_LF, SP_CRLF }; + char cmd[512]; + for (int i = 0; i < 2; i++) { + n[i] = -1; + if (!sp_wr(path[i], body, i)) continue; + snprintf(cmd, sizeof cmd, "%s --lex %s", BC_BIN, path[i]); + th_run(cmd, obuf, TH_BUFSZ); + if (strstr(obuf, "0 error(s)") != NULL) + n[i] = sp_ntok(obuf); + } +} + /* ---- phase: preprocessor ---- */ static void pha01(void) @@ -94,11 +142,114 @@ static void pha08(void) } TH_REG("pha", 8, "semantic analysis runs", pha08) +/* ---- phase 2: line splicing ---- */ + +/* A splice is invisible: the same source with and without one has to give the + * same tokens, and CRLF has to agree with LF on both. */ +static void pha09(void) +{ + int spl[2], ref[2]; + sp_both("__device__ int f(int v)\n{\n v = v \\\n * 2;\n" + " return v;\n}\n", spl); + sp_both("__device__ int f(int v)\n{\n v = v * 2;\n" + " return v;\n}\n", ref); + CHECK(ref[0] > 0); + CHEQ(spl[0], ref[0]); + CHEQ(spl[1], ref[0]); + PASS(); +} +TH_REG("pha", 9, "a splice leaves the token stream alone", pha09) + +/* Phase 2 runs before tokenising, so it joins an identifier cut in half. */ +static void pha10(void) +{ + int spl[2], ref[2]; + sp_both("__device__ int foobar(void) { return 7; }\n" + "__device__ int g(void) { return foo\\\nbar(); }\n", spl); + sp_both("__device__ int foobar(void) { return 7; }\n" + "__device__ int g(void) { return foobar(); }\n", ref); + CHECK(ref[0] > 0); + CHEQ(spl[0], ref[0]); + CHEQ(spl[1], ref[0]); + PASS(); +} +TH_REG("pha", 10, "a splice joins a split identifier", pha10) + +/* And a string literal, which is why this cannot live in the lexer. */ +static void pha11(void) +{ + int spl[2]; + char cmd[512]; + sp_both("__device__ const char *s(void) { return \"ab\\\ncd\"; }\n", spl); + CHECK(spl[0] > 0); + CHEQ(spl[1], spl[0]); + snprintf(cmd, sizeof cmd, "%s --lex %s", BC_BIN, SP_CRLF); + th_run(cmd, obuf, TH_BUFSZ); + CHECK(strstr(obuf, "\"abcd\"") != NULL); + PASS(); +} +TH_REG("pha", 11, "a splice joins a string literal", pha11) + +/* The continuation a multi-line macro is actually written with. */ +static void pha12(void) +{ + char cmd[512]; + CHECK(sp_wr(SP_CRLF, "#define ADD(a,b) \\\n ((a) + \\\n (b))\n" + "__device__ int f(int x) { return ADD(x,2); }\n", 1)); + snprintf(cmd, sizeof cmd, "%s --pp %s", BC_BIN, SP_CRLF); + th_run(cmd, obuf, TH_BUFSZ); + CHECK(strstr(obuf, "((x) +") != NULL); + CHECK(strstr(obuf, "(2))") != NULL); + CHECK(strstr(obuf, "ADD") == NULL); + PASS(); +} +TH_REG("pha", 12, "a CRLF macro keeps its continuations", pha12) + +/* The one that stays invisible until it bites. A splice must not eat a line + * out of the buffer the renderer counts, or every diagnostic below one points + * at the wrong place. The '@' here sits on physical line 6. */ +static void pha13(void) +{ + static const char *const path[2] = { SP_LF, SP_CRLF }; + char cmd[512]; + for (int i = 0; i < 2; i++) { + CHECK(sp_wr(path[i], + "__device__ int f(int v)\n{\n v = v \\\n" + " * 2 \\\n + 1;\n return v @;\n}\n", i)); + snprintf(cmd, sizeof cmd, "%s --lex %s", BC_BIN, path[i]); + th_run(cmd, obuf, TH_BUFSZ); + CHECK(strstr(obuf, "spl_") != NULL); + CHECK(strstr(obuf, ".cu:6:15") != NULL); + CHECK(strstr(obuf, "return v @;") != NULL); + } + PASS(); +} +TH_REG("pha", 13, "a diagnostic below a splice keeps its line", pha13) + +/* Only a backslash immediately before the newline splices. Trailing space, or + * a CR with no LF behind it, is left alone and shows up as a stray backslash + * rather than quietly joining two lines. */ +static void pha14(void) +{ + int spc[2]; + char cmd[512]; + sp_both("__device__ int f(int v)\n{\n v = v \\ \n * 2;\n" + " return v;\n}\n", spc); + CHEQ(spc[0], -1); + CHEQ(spc[1], -1); + CHECK(sp_wr(SP_LF, "__device__ int f(int v)\n{\n v = v \\\r" + " * 2;\n return v;\n}\n", 0)); + snprintf(cmd, sizeof cmd, "%s --lex %s", BC_BIN, SP_LF); + th_run(cmd, obuf, TH_BUFSZ); + CHECK(strstr(obuf, "E005") != NULL); + PASS(); +} +TH_REG("pha", 14, "a backslash not against a newline is no splice", pha14) /* ---- phase: declaration specifiers Booth does not model ---- */ /* _Noreturn parsed as a type name, so the void after it was a syntax error, and 61 of llama.cpp's 67 CUDA files start with one via GGML_NORETURN. */ -static void pha09(void) +static void pha15(void) { int rc = th_run(BC_BIN " --parse tests/noretn.cu", obuf, TH_BUFSZ); CHEQ(rc, 0); @@ -107,13 +258,13 @@ static void pha09(void) CHECK(strstr(obuf, "(ident nrt_bump)") != NULL); PASS(); } -TH_REG("pha", 9, "_Noreturn and [[attributes]] are accepted", pha09) +TH_REG("pha", 15, "_Noreturn and [[attributes]] are accepted", pha15) /* ---- phase: preprocessor, again ---- */ /* ppcyc_a and ppcyc_b include each other. Without #pragma once the depth guard is all that stops them, and it stopped ggml-cuda too. */ -static void pha10(void) +static void pha16(void) { int rc = th_run(BC_BIN " --pp tests/ppcycle.cu", obuf, TH_BUFSZ); CHEQ(rc, 0); @@ -121,12 +272,12 @@ static void pha10(void) CHECK(strstr(obuf, "3 + 4") != NULL); PASS(); } -TH_REG("pha", 10, "#pragma once breaks an include cycle", pha10) +TH_REG("pha", 16, "#pragma once breaks an include cycle", pha16) /* Three that travel together: a body kept its trailing // comment, an argument list that closed on a later line was never joined, and ... in a parameter list ended up in the body rather than naming __VA_ARGS__. */ -static void pha11(void) +static void pha17(void) { int rc = th_run(BC_BIN " --pp tests/ppvarg.cu", obuf, TH_BUFSZ); CHEQ(rc, 0); @@ -138,12 +289,12 @@ static void pha11(void) CHECK(strstr(obuf, "\"a, b\"") != NULL); PASS(); } -TH_REG("pha", 11, "variadic macros and multi-line calls", pha11) +TH_REG("pha", 17, "variadic macros and multi-line calls", pha17) /* The output buffer used to fill and keep going, unterminated, so the lexer read on into the source buffer that follows it and reported whatever it found there. The symptom was an unterminated block comment. */ -static void pha12(void) +static void pha18(void) { int rc = th_run(BC_BIN " --nvidia-ptx tests/ppovfl.cu -o ppovfl_test.ptx", obuf, TH_BUFSZ); @@ -154,4 +305,4 @@ static void pha12(void) remove("ppovfl_test.ptx"); PASS(); } -TH_REG("pha", 12, "a truncated expansion is diagnosed", pha12) +TH_REG("pha", 18, "a truncated expansion is diagnosed", pha18) diff --git a/tests/trpi.c b/tests/trpi.c index 7770ec1..572bc6f 100644 --- a/tests/trpi.c +++ b/tests/trpi.c @@ -113,14 +113,351 @@ static void rpi03(void) } TH_REG("rpi", 3, "a constant above INT32_MAX is not clamped", rpi03) +/* Counts non-overlapping occurrences of needle in obuf. A dropped statement + * leaves the IR one instruction short rather than visibly wrong, so the count + * is what catches it. */ +static int occurs(const char *needle) +{ + int n = 0; + for (const char *p = strstr(obuf, needle); p; p = strstr(p + 1, needle)) + n++; + return n; +} + +/* Writes src to build/ and returns the path, so a regression carries its + * input with it instead of adding a fixture nobody can place later. */ +static const char *scratch(const char *name, const char *src) +{ + static char path[256]; + snprintf(path, sizeof path, "build/%s", name); + FILE *f = fopen(path, "w"); + if (!f) return NULL; + fputs(src, f); + fclose(f); + return path; +} + +/* #5 promoted every parameter to an alloca so it could be written to, and + * shipped without a test. Before it, only struct params were addressable and + * everything else refused with E108; the promotion is invisible in the IR once + * mem2reg folds it away, so what is checked here is the behaviour: a parameter + * is a local initialised from the argument, and reading it after a write sees + * the write. */ +static void rpi04(void) +{ + static const char *const src = + "__device__ int setc(int x) { x = 100; return x; }\n" + "__device__ int padd(const int *p, int n) {\n" + " p = p + 2; n = n * 3; return p[0] + n;\n" + "}\n" + "__device__ int ploop(int a, int b) {\n" + " for (int i = 0; i < 3; i++) { a = a + b; b = b * 2; }\n" + " return a;\n" + "}\n" + "__global__ void kmain(int *out, const int *in) {\n" + " int i = threadIdx.x;\n" + " i = i + 1;\n" + " out[0] = setc(7) + padd(in, 1) + ploop(in[0], 1) + i;\n" + "}\n"; + char cmd[512]; + const char *path = scratch("rpi04.cu", src); + CHNE(path, NULL); + + snprintf(cmd, sizeof cmd, "%s --ir %s", BC_BIN, path); + CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); + CHEQ(strstr(obuf, "E108"), NULL); + + /* x = 100 wins over the incoming argument, so the body is the constant. */ + CHNE(strstr(obuf, "ret i32 100"), NULL); + /* p = p + 2 has to move the pointer, not be dropped. */ + CHNE(strstr(obuf, "gep ptr, %0, 2"), NULL); + /* Both params are rewritten every trip, so the loop head carries a phi for + * each of them on top of the counter's. */ + CHEQ(occurs("phi i32"), 3); + /* A __global__ entry point's parameter is no different. */ + CHNE(strstr(obuf, "__global__"), NULL); + PASS(); +} +TH_REG("rpi", 4, "#5 a parameter can be assigned to", rpi04) + +/* The Triton frontend kept a name's value under the node that first bound it + * and never moved the binding on, so every assignment after the first was + * lowered and then dropped on the floor. Reading the name afterwards returned + * the original value and the compiler said nothing. On an RTX 4060 Ti a kernel + * doing v = v * 2.0 then v = v * 4.0 wrote x back unchanged. */ +static void rpi05(void) +{ + static const char *const src = + "import triton\n" + "import triton.language as tl\n" + "\n" + "@triton.jit\n" + "def kloc(x_ptr, out_ptr, n):\n" + " offs = tl.program_id(axis=0)\n" + " v = tl.load(x_ptr + offs)\n" + " v = v * 2.0\n" + " v = v * 4.0\n" + " tl.store(out_ptr + offs, v)\n"; + char cmd[512]; + const char *path = scratch("rpi05.py", src); + CHNE(path, NULL); + + snprintf(cmd, sizeof cmd, "%s --triton --ir %s", BC_BIN, path); + CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); + /* Both multiplies survive. Dropping the rebind leaves the second one dead + * and DCE takes it, so the count is one before the fix and two after. */ + CHEQ(occurs("fmul f32"), 2); + PASS(); +} +TH_REG("rpi", 5, "a Triton local keeps its latest value", rpi05) + +/* Same root cause on a parameter, where it could not be fixed by moving the + * binding: the name resolved straight back to the incoming argument, so both + * n = n + 5 and n += 5 vanished. Writing to a parameter makes the name local + * from that point, as it does in Python and in C. */ +static void rpi06(void) +{ + static const char *const tmpl = + "import triton\n" + "import triton.language as tl\n" + "\n" + "@triton.jit\n" + "def kpar(x_ptr, out_ptr, n):\n" + " offs = tl.program_id(axis=0)\n" + " %s\n" + " v = tl.load(x_ptr + n)\n" + " tl.store(out_ptr + offs, v)\n"; + static const char *const forms[] = { "n = n + 5", "n += 5", NULL }; + char src[1024], cmd[512]; + + for (int i = 0; forms[i]; i++) { + snprintf(src, sizeof src, tmpl, forms[i]); + const char *path = scratch("rpi06.py", src); + CHNE(path, NULL); + snprintf(cmd, sizeof cmd, "%s --triton --ir %s", BC_BIN, path); + CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); + /* The add is the whole statement. Dropped, it is dead and DCE takes + * it, and the load indexes off the raw argument instead. */ + CHNE(strstr(obuf, "add i32"), NULL); + } + PASS(); +} +TH_REG("rpi", 6, "a Triton parameter can be assigned to", rpi06) + +/* Rebinding across a loop back-edge needs a phi at the head, and the lowerer + * builds one only for the counter. The accumulator every reduction is written + * with therefore read its pre-loop value on every trip, and the sum came back + * as whatever the last iteration computed, or as the initialiser. Nothing in + * the pipeline objected. A refusal is the honest answer until the phi exists. */ +static void rpi07(void) +{ + static const char *const tmpl = + "import triton\n" + "import triton.language as tl\n" + "\n" + "@triton.jit\n" + "def kacc(x_ptr, out_ptr, n):\n" + " offs = tl.program_id(axis=0)\n" + " acc = 0.0\n" + " for i in range(0, 4):\n" + " %s\n" + " tl.store(out_ptr + offs, acc)\n"; + static const char *const forms[] = { + "acc = acc + tl.load(x_ptr + i)", "acc += tl.load(x_ptr + i)", NULL + }; + char src[1024], cmd[512]; + + for (int i = 0; forms[i]; i++) { + snprintf(src, sizeof src, tmpl, forms[i]); + const char *path = scratch("rpi07.py", src); + CHNE(path, NULL); + snprintf(cmd, sizeof cmd, "%s --triton --ir %s", BC_BIN, path); + CHNE(th_run(cmd, obuf, (int)sizeof obuf), 0); + CHNE(strstr(obuf, "E141"), NULL); + } + + /* A name whose whole life is inside the loop body crosses no back-edge and + * has to keep working, or the refusal has eaten the ordinary case with it. + * Both statements survive: the load, then the multiply that rebinds it. */ + static const char *const inner = + "import triton\n" + "import triton.language as tl\n" + "\n" + "@triton.jit\n" + "def kinner(x_ptr, out_ptr, n):\n" + " for i in range(0, 4):\n" + " t = tl.load(x_ptr + i)\n" + " t = t * 2.0\n" + " tl.store(out_ptr + i, t)\n"; + const char *ipath = scratch("rpi07b.py", inner); + CHNE(ipath, NULL); + snprintf(cmd, sizeof cmd, "%s --triton --ir %s", BC_BIN, ipath); + CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); + CHEQ(strstr(obuf, "E141"), NULL); + CHNE(strstr(obuf, "fmul f32"), NULL); + + /* The rank-2 accumulator in a tl.dot kernel is scratch-backed and unrolled + * rather than carried in a register, so it must still compile. */ + snprintf(cmd, sizeof cmd, "%s --triton --ir tests/tri_matmul_k.py", BC_BIN); + CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); + PASS(); +} +TH_REG("rpi", 7, "a loop-carried Triton rebind refuses", rpi07) + +/* The NVIDIA backend gave every i1 a %p, including values defined by adds, + * loads, phis and atomics, none of which can write one. `out[0] = (a==0)||(b==0)` + * came out as `st.global.u32 [%rd2], %p3`, which ptxas and the driver JIT both + * reject, so the kernel could not load at all. i1esc.cu walks the ways out: + * store, arithmetic, call, return, shared, atomic, float conversion and a phi. */ +static int pbad(const char *ptx, char *where, int wsz) +{ + const char *p = ptx; + while ((p = strstr(p, "%p")) != NULL) { + const char *ls = p; + while (ls > ptx && ls[-1] != '\n') ls--; + const char *le = strchr(p, '\n'); + if (!le) le = p + strlen(p); + + while (ls < le && (*ls == ' ' || *ls == '\t')) ls++; + + if (strncmp(ls, ".reg", 4) == 0 || strncmp(ls, "@%p", 3) == 0 + || strncmp(ls, "setp.", 5) == 0 || strncmp(ls, "selp.", 5) == 0 + || strncmp(ls, "vote.", 5) == 0 || strncmp(ls, "mov.pred", 8) == 0) { + p = le; + continue; + } + int n = (int)(le - ls); + if (n > wsz - 1) n = wsz - 1; + memcpy(where, ls, (size_t)n); + where[n] = '\0'; + return 1; + } + return 0; +} + +static char pbuf[1 << 16]; + +static void rpi08(void) +{ + static const char *const modes[] = { "", "--no-mem2reg", NULL }; + char cmd[512], bad[192]; + + for (int i = 0; modes[i] != NULL; i++) { + snprintf(cmd, sizeof cmd, + "%s --nvidia-ptx %s tests/i1esc.cu -o build/rpi04.ptx", + BC_BIN, modes[i]); + CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); + + FILE *f = fopen("build/rpi04.ptx", "r"); + CHNE(f, NULL); + size_t n = fread(pbuf, 1, sizeof pbuf - 1, f); + pbuf[n] = '\0'; + fclose(f); + + if (pbad(pbuf, bad, (int)sizeof bad)) + printf(" %s: %s\n", modes[i][0] ? modes[i] : "default", bad); + CHEQ(pbad(pbuf, bad, (int)sizeof bad), 0); + } + PASS(); +} +TH_REG("rpi", 8, "a predicate never reaches a wider slot", rpi08) + +/* A `bool` is one byte, so the stride between elements of a bool array is one + * byte, and five size functions disagreed about that. width/8 is zero for i1: + * the alloca sites clamped the zero to a minimum and the GEP sites did not, so + * AMD multiplied the index by zero and every element of a bool array resolved + * to the same address. x86-64 and RV64 substituted 4, Tensix refused. + * + * bir_bsz is the one answer now. These pin the answer rather than any one + * backend's arithmetic. */ +static void rpi09(void) +{ + int rc = th_run(BC_BIN " --ir tests/bstride.cu", obuf, (int)sizeof obuf); + CHEQ(rc, 0); + CHNE(strstr(obuf, "store i64 1,"), NULL); + CHEQ(strstr(obuf, "store i64 4,"), NULL); + PASS(); +} +TH_REG("rpi", 9, "sizeof(bool) is one byte", rpi09) + +/* The AMD case is the serious one because it was silent. A scaled GEP whose + * stride is zero collapses a whole array onto element zero, and the shape is + * a v_mul_lo_u32 against an immediate 0. */ +static void rpi10(void) +{ + const char *p; + int rc = th_run(BC_BIN " --amdgpu tests/bstride.cu", obuf, (int)sizeof obuf); + CHEQ(rc, 0); + + for (p = obuf; (p = strstr(p, "v_mul_lo_u32")) != NULL; p++) { + const char *nl = strchr(p, '\n'); + if (nl == NULL) break; + CHEQ(nl - p >= 3 && nl[-3] == ',' && nl[-1] == '0', 0); + } + + /* A bool[64] in LDS reserves 64 bytes, not the 4 the clamp used to give. */ + CHNE(strstr(obuf, "64 LDS bytes"), NULL); + CHNE(strstr(obuf, "8 scratch bytes"), NULL); + PASS(); +} +TH_REG("rpi", 10, "a bool array does not stride by zero on AMD", rpi10) + +/* Storage size, array stride and access width are three questions with one + * answer for a bool. PTX strides by 1 and must therefore touch one byte; + * a .u32 access at a one-byte stride writes over the next three elements. */ +static void rpi11(void) +{ + int rc = th_run(BC_BIN " --nvidia-ptx -o build/rpi11.ptx tests/bstride.cu", + obuf, (int)sizeof obuf); + CHEQ(rc, 0); + + FILE *f = fopen("build/rpi11.ptx", "r"); + CHNE(f, NULL); + if (f == NULL) return; + obuf[fread(obuf, 1, sizeof obuf - 1, f)] = '\0'; + fclose(f); + + CHNE(strstr(obuf, ", 1, %rd"), NULL); + CHEQ(strstr(obuf, ", 4, %rd"), NULL); + CHNE(strstr(obuf, "ld.global.u8"), NULL); + CHNE(strstr(obuf, "st.global.u8"), NULL); + CHNE(strstr(obuf, "ld.shared.u8"), NULL); + CHNE(strstr(obuf, "st.local.u8"), NULL); + CHEQ(strstr(obuf, "ld.global.u32"), NULL); + CHEQ(strstr(obuf, "st.shared.u32"), NULL); + PASS(); +} +TH_REG("rpi", 11, "a bool load reads the byte it strides by", rpi11) + +/* Sizing a struct by summing its fields is the same mistake in a different + * hat: an array of { char; int; } strides by 8 on the host and strode by 5 + * here, so element i past the first landed inside its predecessor. One size + * function that pads the way C does settles both. */ +static void rpi12(void) +{ + int rc = th_run(BC_BIN " --nvidia-ptx -o build/rpi12.ptx tests/bpad.cu", + obuf, (int)sizeof obuf); + CHEQ(rc, 0); + + FILE *f = fopen("build/rpi12.ptx", "r"); + CHNE(f, NULL); + if (f == NULL) return; + obuf[fread(obuf, 1, sizeof obuf - 1, f)] = '\0'; + fclose(f); + + CHNE(strstr(obuf, ", 8, %rd"), NULL); + CHEQ(strstr(obuf, ", 5, %rd"), NULL); + PASS(); +} +TH_REG("rpi", 12, "a struct array strides by its padded size", rpi12) /* looks_like_cast read `( ident )` before any prefix operator as a cast to a * type called ident, without ever asking whether ident named a type. Only `+` * and `-` are also infix, so `(a) + (b)` became a cast applied to `+(b)` and * the left operand vanished with no diagnostic. `-` left the tell, `sub 0, b`. * - * The fixture stays on disk after a failure, so build/rpi04.cu names the case + * The fixture stays on disk after a failure, so build/rpi13.cu names the case * that broke. */ -static void rpi04(void) +static void rpi13(void) { static const struct { const char *ex; const char *ir; } cs[] = { { "(a) + (b)", "add i32 %0, %1" }, @@ -141,23 +478,23 @@ static void rpi04(void) char cmd[512]; for (size_t i = 0; i < sizeof cs / sizeof cs[0]; i++) { - FILE *f = fopen("build/rpi04.cu", "w"); + FILE *f = fopen("build/rpi13.cu", "w"); CHNE(f, NULL); fprintf(f, "__device__ int f(int a, int b){ return %s; }\n", cs[i].ex); fclose(f); - snprintf(cmd, sizeof cmd, "%s --ir build/rpi04.cu", BC_BIN); + snprintf(cmd, sizeof cmd, "%s --ir build/rpi13.cu", BC_BIN); CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); CHNE(strstr(obuf, cs[i].ir), NULL); } PASS(); } -TH_REG("rpi", 4, "a parenthesised variable is not a cast", rpi04) +TH_REG("rpi", 13, "a parenthesised variable is not a cast", rpi13) /* The other half of the same test. Tightening it must not cost a real cast, * so every shape a typedef name reaches the parser in is checked here, and * `(pair){...}` is one the loose rule never recognised at all. */ -static void rpi05(void) +static void rpi14(void) { static const struct { const char *src; const char *ir; } cs[] = { { "typedef int myint;\n" @@ -190,37 +527,37 @@ static void rpi05(void) char cmd[512]; for (size_t i = 0; i < sizeof cs / sizeof cs[0]; i++) { - FILE *f = fopen("build/rpi05.cu", "w"); + FILE *f = fopen("build/rpi14.cu", "w"); CHNE(f, NULL); fputs(cs[i].src, f); fclose(f); - snprintf(cmd, sizeof cmd, "%s --ir build/rpi05.cu", BC_BIN); + snprintf(cmd, sizeof cmd, "%s --ir build/rpi14.cu", BC_BIN); CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); CHNE(strstr(obuf, cs[i].ir), NULL); } PASS(); } -TH_REG("rpi", 5, "a typedef name still casts", rpi05) +TH_REG("rpi", 14, "a typedef name still casts", rpi14) /* A template type parameter is a type name for the body below it. It is not a * typedef and never reached the registry, so it survived on the loose rule * alone and (T)a + b would have started returning b. */ -static void rpi06(void) +static void rpi15(void) { static const char *const src = "template __device__ T g(T a, T b){ return (T)a + b; }\n"; char cmd[512]; - FILE *f = fopen("build/rpi06.cu", "w"); + FILE *f = fopen("build/rpi15.cu", "w"); CHNE(f, NULL); fputs(src, f); fclose(f); - snprintf(cmd, sizeof cmd, "%s --parse build/rpi06.cu", BC_BIN); + snprintf(cmd, sizeof cmd, "%s --parse build/rpi15.cu", BC_BIN); CHEQ(th_run(cmd, obuf, (int)sizeof obuf), 0); CHNE(strstr(obuf, "(binary +"), NULL); CHNE(strstr(obuf, "(cast"), NULL); PASS(); } -TH_REG("rpi", 6, "a template type parameter names a type", rpi06) +TH_REG("rpi", 15, "a template type parameter names a type", rpi15) diff --git a/tests/tu_dup.cu b/tests/tu_dup.cu new file mode 100644 index 0000000..0e2e698 --- /dev/null +++ b/tests/tu_dup.cu @@ -0,0 +1 @@ +__global__ void tu_klib(int *o) { o[0] = 99; } diff --git a/tests/tu_g1.cu b/tests/tu_g1.cu new file mode 100644 index 0000000..3a06cb9 --- /dev/null +++ b/tests/tu_g1.cu @@ -0,0 +1,3 @@ +extern __constant__ float tu_gk; + +__global__ void tu_kg1(float *o) { o[0] = tu_gk * 2.0f; } diff --git a/tests/tu_g2.cu b/tests/tu_g2.cu new file mode 100644 index 0000000..5c7d3d8 --- /dev/null +++ b/tests/tu_g2.cu @@ -0,0 +1,3 @@ +__constant__ float tu_gk; + +__global__ void tu_kg2(float *o) { o[1] = tu_gk * 3.0f; } diff --git a/tests/tu_h1.cu b/tests/tu_h1.cu new file mode 100644 index 0000000..e856851 --- /dev/null +++ b/tests/tu_h1.cu @@ -0,0 +1,3 @@ +#include "tu_hdr.cuh" + +__global__ void tu_kh1(float *o) { o[0] = tu_sq(2.0f) + tu_bias; } diff --git a/tests/tu_h2.cu b/tests/tu_h2.cu new file mode 100644 index 0000000..445ab7c --- /dev/null +++ b/tests/tu_h2.cu @@ -0,0 +1,3 @@ +#include "tu_hdr.cuh" + +__global__ void tu_kh2(float *o) { o[1] = tu_sq(3.0f) + tu_bias; } diff --git a/tests/tu_hdr.cuh b/tests/tu_hdr.cuh new file mode 100644 index 0000000..c2c41ba --- /dev/null +++ b/tests/tu_hdr.cuh @@ -0,0 +1,5 @@ +#ifndef TU_HDR_CUH +#define TU_HDR_CUH +__device__ inline float tu_sq(float x) { return x * x; } +__constant__ float tu_bias; +#endif diff --git a/tests/tu_lib.cu b/tests/tu_lib.cu new file mode 100644 index 0000000..d76a64b --- /dev/null +++ b/tests/tu_lib.cu @@ -0,0 +1,3 @@ +__device__ int tu_mul7(int x) { return x * 7; } + +__global__ void tu_klib(int *o) { o[0] = tu_mul7(2); } diff --git a/tests/tu_sta.cu b/tests/tu_sta.cu new file mode 100644 index 0000000..d854a8e --- /dev/null +++ b/tests/tu_sta.cu @@ -0,0 +1,3 @@ +static __device__ int scale(int x) { return x + 1; } + +__global__ void tu_ka(int *o) { o[threadIdx.x] = scale(1); } diff --git a/tests/tu_stb.cu b/tests/tu_stb.cu new file mode 100644 index 0000000..a50ee6d --- /dev/null +++ b/tests/tu_stb.cu @@ -0,0 +1,3 @@ +static __device__ int scale(int x) { return x + 100; } + +__global__ void tu_kb(int *o) { o[threadIdx.x] = scale(1); } diff --git a/tests/tu_t1.cu b/tests/tu_t1.cu new file mode 100644 index 0000000..7d4ccc2 --- /dev/null +++ b/tests/tu_t1.cu @@ -0,0 +1,4 @@ +template +__global__ void tu_tsc(T *o, T k) { o[threadIdx.x] = o[threadIdx.x] * k; } + +void tu_h1(float *d) { tu_tsc<<<4, 256>>>(d, 2.0f); } diff --git a/tests/tu_t2.cu b/tests/tu_t2.cu new file mode 100644 index 0000000..a325fed --- /dev/null +++ b/tests/tu_t2.cu @@ -0,0 +1,4 @@ +template +__global__ void tu_tsc(T *o, T k) { o[threadIdx.x] = o[threadIdx.x] * k; } + +void tu_h2(float *d) { tu_tsc<<<8, 128>>>(d, 3.0f); } diff --git a/tests/tu_use.cu b/tests/tu_use.cu new file mode 100644 index 0000000..266e82a --- /dev/null +++ b/tests/tu_use.cu @@ -0,0 +1 @@ +__global__ void tu_kuse(int *o) { o[1] = tu_mul7(3); } diff --git a/tests/tu_use.opt b/tests/tu_use.opt new file mode 100644 index 0000000..e6a11dc --- /dev/null +++ b/tests/tu_use.opt @@ -0,0 +1,6 @@ +# Options for tu_use.cu. See test_diag.opt for the format. +# +# Half of a two-file fixture. It calls tu_mul7, which tu_lib.cu defines, so +# on its own it is meant to come back with E105. The pair is exercised by +# the mtu family instead. +xfail all one half of a multi-file fixture, tu_lib.cu holds the callee