222 lines · plain
1; This is an excerpt from the tutorial of the Triton language converted into2; LLVM IR via the Triton XPU backend and cleaned of irrelevant details.3; The only pass criterion is that spirv-val considers output valid.4 5; Ths particular case is related to translation of <1 x Ty> vectors.6 7; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv64-unknown-unknown %s -o - -filetype=obj | spirv-val --target-env spv1.4 %}8 9define spir_kernel void @softmax_kernel(ptr addrspace(1) nocapture writeonly %0, ptr addrspace(1) nocapture readonly %1, i32 %2, i32 %3, i32 %4, i32 %5, ptr addrspace(3) nocapture %6) {10 %8 = tail call spir_func i64 @_Z12get_group_idj(i32 0)11 %9 = trunc i64 %8 to i3212 %10 = tail call spir_func i64 @_Z14get_num_groupsj(i32 0)13 %11 = trunc i64 %10 to i3214 %12 = tail call spir_func i64 @_Z12get_local_idj(i32 0)15 %13 = trunc i64 %12 to i3216 %14 = and i32 %13, 25517 %15 = or disjoint i32 %14, 25618 %16 = or disjoint i32 %14, 51219 %17 = or disjoint i32 %14, 76820 %18 = icmp slt i32 %14, %521 %19 = icmp slt i32 %15, %522 %20 = icmp slt i32 %16, %523 %21 = icmp slt i32 %17, %524 %22 = icmp sgt i32 %4, %925 br i1 %22, label %.lr.ph, label %._crit_edge26 27.lr.ph: ; preds = %728 %23 = lshr i64 %12, 529 %24 = and i32 %13, 3130 %25 = zext nneg i32 %15 to i6431 %26 = zext nneg i32 %16 to i6432 %27 = zext nneg i32 %17 to i6433 %28 = and i64 %12, 25534 %29 = and i64 %23, 735 %30 = icmp eq i32 %24, 036 %31 = getelementptr float, ptr addrspace(3) %6, i64 %2937 %32 = icmp slt i32 %13, 838 %sext = shl i64 %12, 3239 %33 = ashr exact i64 %sext, 3040 %34 = getelementptr i8, ptr addrspace(3) %6, i64 %3341 %35 = and i32 %13, 742 %36 = icmp eq i32 %35, 043 %37 = and i1 %32, %3644 br label %3845 4638: ; preds = %.lr.ph, %12347 %39 = phi i32 [ %9, %.lr.ph ], [ %124, %123 ]48 %40 = mul i32 %39, %249 %41 = sext i32 %40 to i6450 %42 = getelementptr float, ptr addrspace(1) %1, i64 %4151 %43 = getelementptr float, ptr addrspace(1) %42, i64 %2552 %44 = getelementptr float, ptr addrspace(1) %42, i64 %2653 %45 = getelementptr float, ptr addrspace(1) %42, i64 %2754 br i1 %18, label %46, label %4955 5646: ; preds = %3857 %47 = getelementptr float, ptr addrspace(1) %42, i64 %2858 %48 = load <1 x float>, ptr addrspace(1) %47, align 459 br label %4960 6149: ; preds = %46, %3862 %50 = phi <1 x float> [ %48, %46 ], [ splat (float 0xFFF0000000000000), %38 ]63 %51 = extractelement <1 x float> %50, i64 064 br i1 %19, label %52, label %5465 6652: ; preds = %4967 %53 = load <1 x float>, ptr addrspace(1) %43, align 468 br label %5469 7054: ; preds = %52, %4971 %55 = phi <1 x float> [ %53, %52 ], [ splat (float 0xFFF0000000000000), %49 ]72 %56 = extractelement <1 x float> %55, i64 073 br i1 %20, label %57, label %5974 7557: ; preds = %5476 %58 = load <1 x float>, ptr addrspace(1) %44, align 477 br label %5978 7959: ; preds = %57, %5480 %60 = phi <1 x float> [ %58, %57 ], [ splat (float 0xFFF0000000000000), %54 ]81 %61 = extractelement <1 x float> %60, i64 082 br i1 %21, label %62, label %6483 8462: ; preds = %5985 %63 = load <1 x float>, ptr addrspace(1) %45, align 486 br label %6487 8864: ; preds = %62, %5989 %65 = phi <1 x float> [ %63, %62 ], [ splat (float 0xFFF0000000000000), %59 ]90 %66 = extractelement <1 x float> %65, i64 091 tail call spir_func void @_Z7barrierj(i32 1)92 %67 = tail call float @llvm.maxnum.f32(float %51, float %56)93 %68 = tail call float @llvm.maxnum.f32(float %67, float %61)94 %69 = tail call float @llvm.maxnum.f32(float %68, float %66)95 %70 = tail call spir_func float @_Z27__spirv_GroupNonUniformFMaxiif(i32 3, i32 0, float %69)96 br i1 %30, label %71, label %7297 9871: ; preds = %6499 store float %70, ptr addrspace(3) %31, align 4100 br label %72101 10272: ; preds = %71, %64103 tail call spir_func void @_Z7barrierj(i32 1)104 br i1 %32, label %74, label %.thread1105 106.thread1: ; preds = %72107 %73 = tail call spir_func float @_Z27__spirv_GroupNonUniformFMaxiifj(i32 3, i32 3, float poison, i32 8)108 br label %78109 11074: ; preds = %72111 %75 = load float, ptr addrspace(3) %34, align 4112 %76 = tail call spir_func float @_Z27__spirv_GroupNonUniformFMaxiifj(i32 3, i32 3, float %75, i32 8)113 br i1 %37, label %77, label %78114 11577: ; preds = %74116 store float %76, ptr addrspace(3) %34, align 4117 br label %78118 11978: ; preds = %.thread1, %77, %74120 tail call spir_func void @_Z7barrierj(i32 1)121 %79 = load float, ptr addrspace(3) %6, align 4122 %80 = fsub float %51, %79123 %81 = fsub float %56, %79124 %82 = fsub float %61, %79125 %83 = fsub float %66, %79126 %84 = fmul float %80, 0x3FF7154760000000127 %85 = tail call float @llvm.exp2.f32(float %84)128 %86 = fmul float %81, 0x3FF7154760000000129 %87 = tail call float @llvm.exp2.f32(float %86)130 %88 = fmul float %82, 0x3FF7154760000000131 %89 = tail call float @llvm.exp2.f32(float %88)132 %90 = fmul float %83, 0x3FF7154760000000133 %91 = tail call float @llvm.exp2.f32(float %90)134 tail call spir_func void @_Z7barrierj(i32 1)135 %92 = fadd float %85, %87136 %93 = fadd float %89, %92137 %94 = fadd float %91, %93138 %95 = tail call spir_func float @_Z27__spirv_GroupNonUniformFAddiif(i32 3, i32 0, float %94)139 br i1 %30, label %96, label %97140 14196: ; preds = %78142 store float %95, ptr addrspace(3) %31, align 4143 br label %97144 14597: ; preds = %96, %78146 tail call spir_func void @_Z7barrierj(i32 1)147 br i1 %32, label %99, label %.thread148 149.thread: ; preds = %97150 %98 = tail call spir_func float @_Z27__spirv_GroupNonUniformFAddiifj(i32 3, i32 3, float poison, i32 8)151 br label %103152 15399: ; preds = %97154 %100 = load float, ptr addrspace(3) %34, align 4155 %101 = tail call spir_func float @_Z27__spirv_GroupNonUniformFAddiifj(i32 3, i32 3, float %100, i32 8)156 br i1 %37, label %102, label %103157 158102: ; preds = %99159 store float %101, ptr addrspace(3) %34, align 4160 br label %103161 162103: ; preds = %.thread, %102, %99163 tail call spir_func void @_Z7barrierj(i32 1)164 %104 = load float, ptr addrspace(3) %6, align 4165 %105 = fdiv float %87, %104166 %106 = fdiv float %89, %104167 %107 = fdiv float %91, %104168 %108 = mul i32 %39, %3169 %109 = sext i32 %108 to i64170 %110 = getelementptr float, ptr addrspace(1) %0, i64 %109171 %111 = getelementptr float, ptr addrspace(1) %110, i64 %25172 %112 = getelementptr float, ptr addrspace(1) %110, i64 %26173 %113 = getelementptr float, ptr addrspace(1) %110, i64 %27174 br i1 %18, label %114, label %117175 176114: ; preds = %103177 %115 = fdiv float %85, %104178 %116 = getelementptr float, ptr addrspace(1) %110, i64 %28179 store float %115, ptr addrspace(1) %116, align 4180 br label %117181 182117: ; preds = %114, %103183 br i1 %19, label %118, label %119184 185118: ; preds = %117186 store float %105, ptr addrspace(1) %111, align 4187 br label %119188 189119: ; preds = %118, %117190 br i1 %20, label %120, label %121191 192120: ; preds = %119193 store float %106, ptr addrspace(1) %112, align 4194 br label %121195 196121: ; preds = %120, %119197 br i1 %21, label %122, label %123198 199122: ; preds = %121200 store float %107, ptr addrspace(1) %113, align 4201 br label %123202 203123: ; preds = %122, %121204 %124 = add i32 %39, %11205 %125 = icmp slt i32 %124, %4206 br i1 %125, label %38, label %._crit_edge207 208._crit_edge: ; preds = %123, %7209 ret void210}211 212declare float @llvm.maxnum.f32(float, float)213declare spir_func float @_Z27__spirv_GroupNonUniformFAddiifj(i32, i32, float, i32)214declare spir_func float @_Z27__spirv_GroupNonUniformFAddiif(i32, i32, float)215declare spir_func float @_Z27__spirv_GroupNonUniformFMaxiifj(i32, i32, float, i32)216declare spir_func float @_Z27__spirv_GroupNonUniformFMaxiif(i32, i32, float)217declare spir_func void @_Z7barrierj(i32)218declare spir_func i64 @_Z12get_local_idj(i32)219declare spir_func i64 @_Z14get_num_groupsj(i32)220declare spir_func i64 @_Z12get_group_idj(i32)221declare float @llvm.exp2.f32(float)222