| 48 | } |
| 49 | |
| 50 | void CopyHostToDevice(const Tensor* input, Allocator* cpu_allocator, |
| 51 | Allocator* out_allocator, StringPiece edge_name, |
| 52 | Device* dst, Tensor* output, |
| 53 | DeviceContext* recv_dev_context, StatusCallback done, |
| 54 | bool sync_dst_compute) { |
| 55 | if (input->dtype() == DT_VARIANT) { |
| 56 | Tensor copy(cpu_allocator, DT_VARIANT, input->shape()); |
| 57 | auto* status_cb = new ReffedStatusCallback(std::move(done)); |
| 58 | core::ScopedUnref status_cb_unref(status_cb); |
| 59 | |
| 60 | auto wrapped_done = [status_cb](const Status& s) { |
| 61 | status_cb->UpdateStatus(s); |
| 62 | status_cb->Unref(); |
| 63 | }; |
| 64 | auto copier = std::bind( |
| 65 | [dst, recv_dev_context, out_allocator, status_cb, cpu_allocator, |
| 66 | edge_name, sync_dst_compute](StatusCallback wrapped_done_, |
| 67 | // Begin unbound arguments |
| 68 | const Tensor& from, Tensor* to) { |
| 69 | if (from.dtype() == DT_VARIANT) { |
| 70 | status_cb->Ref(); |
| 71 | CopyHostToDevice(&from, cpu_allocator, out_allocator, edge_name, |
| 72 | dst, to, recv_dev_context, wrapped_done_, |
| 73 | sync_dst_compute); |
| 74 | return Status::OK(); |
| 75 | } else { |
| 76 | if (!DMAHelper::CanUseDMA(&from)) { |
| 77 | Status err = errors::InvalidArgument( |
| 78 | "During Variant Host->Device Copy: " |
| 79 | "non-DMA-copy attempted of tensor type: ", |
| 80 | DataTypeString(from.dtype())); |
| 81 | status_cb->UpdateStatus(err); |
| 82 | return err; |
| 83 | } |
| 84 | if (status_cb->ok()) { |
| 85 | status_cb->Ref(); |
| 86 | *to = Tensor(out_allocator, from.dtype(), from.shape()); |
| 87 | recv_dev_context->CopyCPUTensorToDevice( |
| 88 | &from, dst, to, wrapped_done_, sync_dst_compute); |
| 89 | return Status::OK(); |
| 90 | } else { |
| 91 | return status_cb->status(); |
| 92 | } |
| 93 | } |
| 94 | }, |
| 95 | std::move(wrapped_done), std::placeholders::_1, std::placeholders::_2); |
| 96 | |
| 97 | const Variant* v = input->flat<Variant>().data(); |
| 98 | Variant* v_out = copy.flat<Variant>().data(); |
| 99 | Status s_copy_init; |
| 100 | for (int64 i = 0; i < input->NumElements(); ++i) { |
| 101 | s_copy_init = VariantDeviceCopy( |
| 102 | VariantDeviceCopyDirection::HOST_TO_DEVICE, v[i], &v_out[i], copier); |
| 103 | if (!s_copy_init.ok()) { |
| 104 | status_cb->UpdateStatus(s_copy_init); |
| 105 | break; |
| 106 | } |
| 107 | } |
no test coverage detected