MCPcopy Create free account
hub / github.com/RenderKit/oidn / newHIPConv

Function newHIPConv

devices/hip/hip_conv.cpp:7–58  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5#include "ck_conv.h"
6
7OIDN_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
60OIDN_NAMESPACE_END

Callers 1

newConvMethod · 0.85

Calls 8

getPaddedOMethod · 0.80
getPaddedIMethod · 0.80
makeMethod · 0.80
maxFunction · 0.50
round_upFunction · 0.50
getArchMethod · 0.45
getHMethod · 0.45
getWMethod · 0.45

Tested by

no test coverage detected