| 12 | using namespace BareMetal; |
| 13 | |
| 14 | bool BatchedMatrixMulKernel::IsAvailable(TContext* context) const { |
| 15 | bool ok_dtype = context->getAttrOprand("operand:0").dtype == "f32" && |
| 16 | context->getAttrOprand("operand:1").dtype == "f32" && |
| 17 | context->getAttrOprand("operand:2").dtype == "f32"; |
| 18 | bool ok_fp16 = context->getAttrOprand("operand:0").dtype == "f16" && |
| 19 | context->getAttrOprand("operand:1").dtype == "f16" && |
| 20 | context->getAttrOprand("operand:2").dtype == "f16"; |
| 21 | bool ok_mode = context->getAttrStr("format") == "DEFAULT" && |
| 22 | context->getAttrStr("compute_mode") == "DEFAULT"; |
| 23 | return (ok_dtype || ok_fp16) && ok_mode; |
| 24 | } |
| 25 | |
| 26 | //! kernel gen |
| 27 | std::string BatchedMatrixMulKernel::GetKernelSymbol(TContext* context) const { |
nothing calls this directly
no test coverage detected