| 407 | } |
| 408 | |
| 409 | void CgmodeRunOp::FallbackToTF(OpKernelContext* ctx, const NameAttrList func) { |
| 410 | VLOG(1) << "Fallback to run TF subgraph when CUDA Graph execution fails"; |
| 411 | FunctionLibraryRuntime::Handle handle; |
| 412 | flib_->Instantiate(func.name(), AttrSlice(&func.attr()), &handle); |
| 413 | const FunctionBody* fbody; |
| 414 | fbody = flib_->GetFunctionBody(handle); |
| 415 | FunctionLibraryRuntime::Options opts; |
| 416 | std::vector<Tensor> in; |
| 417 | for (int i = 0; i < ctx->num_inputs() - 1; i++) { |
| 418 | in.emplace_back(ctx->input(i)); |
| 419 | } |
| 420 | Status s_tf_run; |
| 421 | Notification done_tf_run; |
| 422 | std::vector<Tensor> out(fbody->ret_types.size()); |
| 423 | flib_->Run(opts, handle, in, &out, |
| 424 | [&s_tf_run, &done_tf_run](const Status& s) { |
| 425 | s_tf_run = s; |
| 426 | done_tf_run.Notify(); |
| 427 | }); |
| 428 | done_tf_run.WaitForNotification(); |
| 429 | OP_REQUIRES(ctx, s_tf_run.ok(), errors::Internal(s_tf_run.ToString())); |
| 430 | for (int i = 0; i < out.size(); i++) { |
| 431 | ctx->set_output(i, out[i]); |
| 432 | } |
| 433 | } |
| 434 | REGISTER_KERNEL_BUILDER(Name("_CgmodeCompile") |
| 435 | .Device(DEVICE_GPU) |
| 436 | .HostMemory("key") |
nothing calls this directly
no test coverage detected