| 5 | #include "ck_conv.h" |
| 6 | |
| 7 | OIDN_NAMESPACE_BEGIN |
| 8 | |
| 9 | Ref<Conv> newHIPConv(HIPEngine* engine, const ConvDesc& desc) |
| 10 | { |
| 11 | // Get the list of kernels supported by the engine |
| 12 | std::vector<CKConvFactory> kernels; |
| 13 | switch (engine->getArch()) |
| 14 | { |
| 15 | case HIPArch::DL: |
| 16 | kernels = getCKConvInstances<HIPArch::DL>(desc.activation); |
| 17 | break; |
| 18 | case HIPArch::WMMA: |
| 19 | kernels = getCKConvInstances<HIPArch::WMMA>(desc.activation); |
| 20 | break; |
| 21 | default: |
| 22 | throw std::runtime_error("unsupported architecture"); |
| 23 | } |
| 24 | |
| 25 | // Select the likely fastest compatible kernel based on the GEMM dimensions |
| 26 | const size_t M = desc.srcDesc.getH() * desc.srcDesc.getW(); // == destination dims |
| 27 | const size_t N = desc.weightDesc.getPaddedO(); |
| 28 | const size_t K = desc.weightDesc.getPaddedI() * desc.weightDesc.getH() * desc.weightDesc.getW(); |
| 29 | |
| 30 | const DataType accumType = DataType::Float32; |
| 31 | |
| 32 | const CKConvFactory* bestKernel = nullptr; |
| 33 | int bestBlockSize = 0; |
| 34 | size_t bestCost = std::numeric_limits<size_t>::max(); |
| 35 | |
| 36 | for (const auto& kernel : kernels) |
| 37 | { |
| 38 | if (kernel.dataType != desc.srcDesc.dataType || kernel.accumType < accumType) |
| 39 | continue; |
| 40 | |
| 41 | const int blockSize = kernel.blockM * kernel.blockN * kernel.blockK; |
| 42 | const size_t cost = round_up(M, kernel.blockM) * round_up(N, kernel.blockN) * round_up(K, kernel.blockK); |
| 43 | |
| 44 | if ((cost < bestCost) || |
| 45 | (cost == bestCost && blockSize > bestBlockSize) || |
| 46 | (cost == bestCost && blockSize == bestBlockSize && kernel.accumType == accumType)) |
| 47 | { |
| 48 | bestKernel = &kernel; |
| 49 | bestBlockSize = blockSize; |
| 50 | bestCost = cost; |
| 51 | } |
| 52 | } |
| 53 | |
| 54 | if (!bestKernel) |
| 55 | throw std::runtime_error("unsupported convolution"); |
| 56 | |
| 57 | return bestKernel->make(engine, desc); |
| 58 | } |
| 59 | |
| 60 | OIDN_NAMESPACE_END |
no test coverage detected