| 336 | class FusedMatMulOp : public OpKernel { |
| 337 | public: |
| 338 | explicit FusedMatMulOp(OpKernelConstruction* context) : OpKernel(context) { |
| 339 | OP_REQUIRES_OK(context, context->GetAttr("transpose_a", &transpose_a_)); |
| 340 | OP_REQUIRES_OK(context, context->GetAttr("transpose_b", &transpose_b_)); |
| 341 | |
| 342 | std::vector<FusedComputationPattern> patterns; |
| 343 | |
| 344 | using FCT = FusedComputationType; |
| 345 | if (std::is_same<Device, CPUDevice>::value) { |
| 346 | patterns = {{FCT::kBiasAdd, {"BiasAdd"}}, |
| 347 | {FCT::kBiasAddWithRelu, {"BiasAdd", "Relu"}}, |
| 348 | {FCT::kBiasAddWithRelu6, {"BiasAdd", "Relu6"}}, |
| 349 | {FCT::kBiasAddWithElu, {"BiasAdd", "Elu"}}}; |
| 350 | } else if (std::is_same<Device, GPUDevice>::value) { |
| 351 | patterns = {{FCT::kBiasAdd, {"BiasAdd"}}, |
| 352 | {FCT::kBiasAddWithRelu, {"BiasAdd", "Relu"}}}; |
| 353 | } |
| 354 | |
| 355 | OP_REQUIRES_OK(context, InitializeFusedComputation( |
| 356 | context, "MatMul", patterns, |
| 357 | &fused_computation_, &fused_computation_args_)); |
| 358 | use_autotune_ = MatmulAutotuneEnable(); |
| 359 | } |
| 360 | |
| 361 | void Compute(OpKernelContext* ctx) override { |
| 362 | const Tensor& a = ctx->input(0); |
nothing calls this directly
no test coverage detected