| 13 | using namespace ArmCommon; |
| 14 | |
| 15 | bool RelayoutKernel::IsAvailable(TContext* context) const { |
| 16 | auto src_dtype_str = context->getAttrOprand("operand:0").dtype; |
| 17 | int type_size = Utils::get_dtype_size(src_dtype_str); |
| 18 | bool ok_dtype = (src_dtype_str == context->getAttrOprand("operand:1").dtype) && |
| 19 | (type_size == 1 || type_size == 4); |
| 20 | std::vector<size_t> shape_in = context->getAttrOprand("operand:0").shape; |
| 21 | std::vector<size_t> shape_out = context->getAttrOprand("operand:0").shape; |
| 22 | bool ok_shape = shape_in.size() == shape_out.size(); |
| 23 | for (size_t i = 0; ok_shape && i < shape_in.size(); ++i) { |
| 24 | ok_shape = ok_shape && shape_in.at(i) == shape_out.at(i); |
| 25 | } |
| 26 | return ok_dtype && ok_shape; |
| 27 | } |
| 28 | |
| 29 | //! kernel gen |
| 30 | std::string RelayoutKernel::GetKernelSymbol(TContext* context) const { |
nothing calls this directly
no test coverage detected