MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / Compute

Method Compute

tensorflow/core/kernels/save_restore_v2_ops.cc:312–451  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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;

Callers

nothing calls this directly

Calls 15

DefaultFunction · 0.85
HandleFromInputFunction · 0.85
LookupResourceFunction · 0.85
handlesClass · 0.85
ParseShapeAndSliceFunction · 0.85
InvalidArgumentFunction · 0.85
SaveHashTableFunction · 0.85
SaveBloomFilterFunction · 0.85
AddSliceMethod · 0.80
ValidateInputsFunction · 0.70
SplitClass · 0.70
inputMethod · 0.45

Tested by

no test coverage detected