| 110 | } |
| 111 | |
| 112 | CPUConcatConvAMX::CPUConcatConvAMX(CPUEngine* engine, const ConcatConvDesc& desc) |
| 113 | : ConcatConv(desc), |
| 114 | engine(engine) |
| 115 | { |
| 116 | if (src0Desc.layout != TensorLayout::Chw32c || src0Desc.dataType != DataType::Float16) |
| 117 | throw std::invalid_argument("unsupported convolution source layout/data type"); |
| 118 | if (src1Desc.layout != TensorLayout::Chw32c || src1Desc.dataType != DataType::Float16) |
| 119 | throw std::invalid_argument("unsupported convolution source layout/data type"); |
| 120 | if (weightDesc.getW() != 3 || weightDesc.getH() != 3) |
| 121 | throw std::invalid_argument("unsupported convolution kernel size"); |
| 122 | if (weightDesc.layout != TensorLayout::OIhw2o16i16o2i || weightDesc.dataType != DataType::Float16) |
| 123 | throw std::invalid_argument("unsupported convolution weight layout/data type"); |
| 124 | if (biasDesc.layout != TensorLayout::x || biasDesc.dataType != DataType::Float16) |
| 125 | throw std::invalid_argument("unsupported convolution bias layout/data type"); |
| 126 | } |
| 127 | |
| 128 | void CPUConcatConvAMX::submitKernels(const Ref<CancellationToken>& ct) |
| 129 | { |