| 119 | |
| 120 | |
| 121 | void CopyDeviceToDevice(CopyTensor::CopyFunction copy_function, |
| 122 | Allocator* cpu_allocator, Allocator* out_allocator, |
| 123 | DeviceContext* send_dev_context, |
| 124 | DeviceContext* recv_dev_context, Device* src, |
| 125 | Device* dst, const AllocatorAttributes src_alloc_attr, |
| 126 | const AllocatorAttributes dst_alloc_attr, |
| 127 | const Tensor* input, Tensor* output, |
| 128 | int dev_to_dev_stream_index, StatusCallback done) { |
| 129 | if (input->dtype() == DT_VARIANT) { |
| 130 | Tensor copy(cpu_allocator, DT_VARIANT, input->shape()); |
| 131 | auto* status_cb = new ReffedStatusCallback(std::move(done)); |
| 132 | core::ScopedUnref status_cb_unref(status_cb); |
| 133 | |
| 134 | auto wrapped_done = [status_cb](const Status& s) { |
| 135 | status_cb->UpdateStatus(s); |
| 136 | status_cb->Unref(); |
| 137 | }; |
| 138 | auto copier = std::bind( |
| 139 | [copy_function, cpu_allocator, src, dst, src_alloc_attr, dst_alloc_attr, |
| 140 | recv_dev_context, send_dev_context, out_allocator, status_cb, |
| 141 | dev_to_dev_stream_index](StatusCallback wrapped_done_, |
| 142 | // Begin unbound arguments |
| 143 | const Tensor& from, Tensor* to) { |
| 144 | if (from.dtype() == DT_VARIANT) { |
| 145 | status_cb->Ref(); |
| 146 | CopyDeviceToDevice(copy_function, cpu_allocator, out_allocator, |
| 147 | send_dev_context, recv_dev_context, src, dst, |
| 148 | src_alloc_attr, dst_alloc_attr, &from, to, |
| 149 | dev_to_dev_stream_index, wrapped_done_); |
| 150 | return Status::OK(); |
| 151 | } else { |
| 152 | if (!DMAHelper::CanUseDMA(&from)) { |
| 153 | Status err = errors::InvalidArgument( |
| 154 | "During Variant Device->Device Copy: ", src->name(), " to ", |
| 155 | dst->name(), " non-DMA-copy attempted of tensor type: ", |
| 156 | DataTypeString(from.dtype())); |
| 157 | status_cb->UpdateStatus(err); |
| 158 | return err; |
| 159 | } |
| 160 | if (status_cb->ok()) { |
| 161 | status_cb->Ref(); |
| 162 | *to = Tensor(out_allocator, from.dtype(), from.shape()); |
| 163 | copy_function(send_dev_context, recv_dev_context, src, dst, |
| 164 | src_alloc_attr, dst_alloc_attr, &from, to, |
| 165 | dev_to_dev_stream_index, std::move(wrapped_done_)); |
| 166 | return Status::OK(); |
| 167 | } else { |
| 168 | return status_cb->status(); |
| 169 | } |
| 170 | } |
| 171 | }, |
| 172 | std::move(wrapped_done), std::placeholders::_1, std::placeholders::_2); |
| 173 | |
| 174 | const Variant* v = input->flat<Variant>().data(); |
| 175 | Variant* v_out = copy.flat<Variant>().data(); |
| 176 | Status s_copy_init; |
| 177 | for (int64 i = 0; i < input->NumElements(); ++i) { |
| 178 | s_copy_init = |
no test coverage detected