| 310 | } |
| 311 | |
| 312 | void Compute(OpKernelContext* context) override { |
| 313 | const Tensor& prefix = context->input(0); |
| 314 | const Tensor& tensor_names = context->input(1); |
| 315 | const Tensor& shape_and_slices = context->input(2); |
| 316 | const Tensor& ev_names = context->input(3); |
| 317 | const Tensor& ev_resources = context->input(4); |
| 318 | const int kFixedInputs = 5; |
| 319 | ValidateInputs(true /* is save op */, context, prefix, tensor_names, |
| 320 | shape_and_slices, kFixedInputs); |
| 321 | if (!context->status().ok()) return; |
| 322 | // Prefix, tensor names, shape_and_slices, ev names, ev resources. |
| 323 | const int num_tensors = static_cast<int>(tensor_names.NumElements()); |
| 324 | const int num_ev = static_cast<int>(ev_names.NumElements()); |
| 325 | const string& prefix_string = prefix.scalar<tstring>()(); |
| 326 | const auto& tensor_names_flat = tensor_names.flat<tstring>(); |
| 327 | const auto& ev_names_flat = ev_names.flat<tstring>(); |
| 328 | const auto& ev_resources_flat = ev_resources.flat<int64>(); |
| 329 | const auto& shape_and_slices_flat = shape_and_slices.flat<tstring>(); |
| 330 | |
| 331 | BundleWriter writer(Env::Default(), prefix_string); |
| 332 | OP_REQUIRES_OK(context, writer.status()); |
| 333 | VLOG(1) << "BundleWriter, prefix_string: " << prefix_string; |
| 334 | |
| 335 | int start_index = 0; |
| 336 | if (has_ev_) { |
| 337 | start_index = 1; |
| 338 | } |
| 339 | |
| 340 | for (int i = 0; i < num_ev; i++) { |
| 341 | const string& ev_name = ev_names_flat(i); |
| 342 | if (ev_key_types_[i] == DT_INT32) { |
| 343 | EmbeddingVar<int32, float>* ev = |
| 344 | reinterpret_cast< |
| 345 | EmbeddingVar<int32, float>*>(ev_resources_flat(i)); |
| 346 | DumpEvWithGlobalStep( |
| 347 | context, ev_name, ev, writer, tensor_types_[0]); |
| 348 | } else if (ev_key_types_[i] == DT_INT64) { |
| 349 | EmbeddingVar<int64, float>* ev = |
| 350 | reinterpret_cast< |
| 351 | EmbeddingVar<int64, float>*>(ev_resources_flat(i)); |
| 352 | DumpEvWithGlobalStep( |
| 353 | context, ev_name, ev, writer, tensor_types_[0]); |
| 354 | } |
| 355 | } |
| 356 | |
| 357 | for (int i = start_index; i < num_tensors; ++i) { |
| 358 | const string& tensor_name = tensor_names_flat(i); |
| 359 | if (tensor_types_[i] == DT_RESOURCE) { |
| 360 | auto& handle = HandleFromInput(context, i + kFixedInputs); |
| 361 | if (IsHandle<HashTableResource>(handle)) { |
| 362 | auto handles = |
| 363 | context->input(i + kFixedInputs).flat<ResourceHandle>(); |
| 364 | int tensible_size = handles.size() - 1; |
| 365 | std::vector<core::ScopedUnref> unrefs; |
| 366 | HashTable* hashtable; |
| 367 | std::vector<TensibleVariable*> tensibles; |
| 368 | |
| 369 | HashTableResource* htr; |
nothing calls this directly
no test coverage detected