Thanks for the response @kuhar, I’m having another look.
Yes, type converter references options.use64bitIndex which I can’t find an interface to set for the cf-to-spirv pass. Currently it only converts to i32 because the option is false by default.
I’m trying to lower the gpu.func which has index type argument used by spirv.branch. The argument is converted into integer type using unrealized conversion cast and doesn’t seem to be cancelled before gpu-to-spirv which can’t handle it either.
These two issues are connected and probably there’s another (correct) way to address this, I’d appreciate if I can find it.
Here’s an IR snippet,
module attributes {gpu.container_module, spirv.target_env = #spirv.target_env<#spirv.vce<v1.0, [Addresses, Int64, Kernel], []>, #spirv.resource_limits<>>} {
gpu.module @forward_kernel_1 [#spirv.target_env<#spirv.vce<v1.0, [Addresses, Int64, Kernel], []>, #spirv.resource_limits<>>] attributes {spirv.target_env = #spirv.target_env<#spirv.vce<v1.0, [Addresses, Int64, Kernel], []>, api=OpenCL, #spirv.resource_limits<>>} {
gpu.func @forward_kernel_1_forward_kernel(%arg0: index, %arg1: memref<f32>, %arg2: index, %arg3: memref<f32>, %arg4: index, %arg5: memref<f32>, %arg6: index, %arg7: memref<f32>, %arg8: index, %arg9: index, %arg10: index) kernel attributes {gpu.known_block_size = array<i32: 1, 1, 1>, spirv.entry_point_abi = #spirv.entry_point_abi<>} {
%0 = gpu.block_id x
%1 = gpu.block_id y
cf.br ^bb1(%arg8 : index)
^bb1(%2: index): // 2 preds: ^bb0, ^bb2
%3 = arith.cmpi slt, %2, %arg9 : index
cf.cond_br %3, ^bb2, ^bb3
^bb2: // pred: ^bb1
%4 = arith.muli %0, %arg0 : index
%5 = arith.addi %4, %2 : index
%reinterpret_cast = memref.reinterpret_cast %arg1 to offset: [%5], sizes: [], strides: [] : memref<f32> to memref<f32>
%6 = memref.load %reinterpret_cast[] : memref<f32>
%7 = arith.muli %2, %arg2 : index
%8 = arith.addi %7, %1 : index
%reinterpret_cast_0 = memref.reinterpret_cast %arg3 to offset: [%8], sizes: [], strides: [] : memref<f32> to memref<f32>
%9 = memref.load %reinterpret_cast_0[] : memref<f32>
%10 = arith.muli %0, %arg4 : index
%11 = arith.addi %10, %1 : index
%reinterpret_cast_1 = memref.reinterpret_cast %arg5 to offset: [%11], sizes: [], strides: [] : memref<f32> to memref<f32>
%12 = memref.load %reinterpret_cast_1[] : memref<f32>
%13 = arith.mulf %6, %9 : f32
%14 = arith.addf %12, %13 : f32
%15 = arith.muli %0, %arg6 : index
%16 = arith.addi %15, %1 : index
%reinterpret_cast_2 = memref.reinterpret_cast %arg7 to offset: [%16], sizes: [], strides: [] : memref<f32> to memref<f32>
memref.store %14, %reinterpret_cast_2[] : memref<f32>
%17 = arith.addi %2, %arg10 : index
cf.br ^bb1(%17 : index)
^bb3: // pred: ^bb1
gpu.return
}
}
}