| 14 | #if MGB_CUDA |
| 15 | namespace { |
| 16 | std::unique_ptr<LayoutTransformContext> make_ctx() { |
| 17 | using OprFormatConfigID = LayoutTransformContext::OprFormatConfigID; |
| 18 | using OprList = LayoutTransformContext::OprList; |
| 19 | using Attribute = LayoutTransformContext::Attribute; |
| 20 | using Target = LayoutTransformContext::Target; |
| 21 | OprList opr_list = { |
| 22 | opr::ConvBiasForward::typeinfo(), |
| 23 | opr::ConvolutionForward::typeinfo(), |
| 24 | opr::ConvolutionBackwardData::typeinfo(), |
| 25 | opr::ElemwiseMultiType::typeinfo(), |
| 26 | opr::Elemwise::typeinfo(), |
| 27 | opr::TypeCvt::typeinfo(), |
| 28 | opr::PoolingForward::typeinfo(), |
| 29 | opr::WarpPerspectiveForward::typeinfo(), |
| 30 | }; |
| 31 | |
| 32 | SmallVector<TensorFormats> available_tensor_formats = { |
| 33 | TensorFormats::NCHW, TensorFormats::NHWC, TensorFormats::NCHWc4, |
| 34 | TensorFormats::NCHWc32, TensorFormats::NCHWc64, TensorFormats::CHWNc4}; |
| 35 | Attribute attribute = {OprFormatConfigID::NCHW, TensorFormats::NCHW, Target::CUDA}; |
| 36 | auto ctx = std::make_unique<LayoutTransformContext>( |
| 37 | std::move(opr_list), std::move(available_tensor_formats), attribute); |
| 38 | ctx->add_opr_config( |
| 39 | opr::ConvBiasForward::typeinfo(), |
| 40 | {OprFormatConfigID::NCHW, OprFormatConfigID::NHWC, |
| 41 | OprFormatConfigID::NCHW4, OprFormatConfigID::NCHW32, |
| 42 | OprFormatConfigID::NCHW64, OprFormatConfigID::CHWN4}) |
| 43 | .add_opr_config( |
| 44 | opr::ConvolutionForward::typeinfo(), |
| 45 | {OprFormatConfigID::NCHW, OprFormatConfigID::NCHW4}) |
| 46 | .add_opr_config( |
| 47 | opr::ConvolutionBackwardData::typeinfo(), |
| 48 | {OprFormatConfigID::NCHW, OprFormatConfigID::NCHW4}) |
| 49 | .add_opr_config( |
| 50 | opr::PoolingForward::typeinfo(), |
| 51 | {OprFormatConfigID::NCHW4, OprFormatConfigID::NCHW32, |
| 52 | OprFormatConfigID::NHWC, OprFormatConfigID::NCHW64, |
| 53 | OprFormatConfigID::CHWN4}) |
| 54 | .add_opr_config( |
| 55 | opr::WarpPerspectiveForward::typeinfo(), |
| 56 | {OprFormatConfigID::NHWC, OprFormatConfigID::NCHW4, |
| 57 | OprFormatConfigID::NCHW64}); |
| 58 | return ctx; |
| 59 | } |
| 60 | } // namespace |
| 61 | |
| 62 | #if CUDA_VERSION >= 10020 |