Shared validations of the inputs to the SaveV2 and RestoreV2 ops.
| 48 | |
| 49 | // Shared validations of the inputs to the SaveV2 and RestoreV2 ops. |
| 50 | void ValidateInputs(bool is_save_op, OpKernelContext* context, |
| 51 | const Tensor& prefix, const Tensor& tensor_names, |
| 52 | const Tensor& shape_and_slices, |
| 53 | const int kFixedInputs) { |
| 54 | const int num_tensors = static_cast<int>(tensor_names.NumElements()); |
| 55 | OP_REQUIRES( |
| 56 | context, prefix.NumElements() == 1, |
| 57 | errors::InvalidArgument("Input prefix should have a single element, got ", |
| 58 | prefix.NumElements(), " instead.")); |
| 59 | OP_REQUIRES(context, |
| 60 | TensorShapeUtils::IsVector(tensor_names.shape()) && |
| 61 | TensorShapeUtils::IsVector(shape_and_slices.shape()), |
| 62 | errors::InvalidArgument( |
| 63 | "Input tensor_names and shape_and_slices " |
| 64 | "should be an 1-D tensors, got ", |
| 65 | tensor_names.shape().DebugString(), " and ", |
| 66 | shape_and_slices.shape().DebugString(), " instead.")); |
| 67 | OP_REQUIRES(context, |
| 68 | tensor_names.NumElements() == shape_and_slices.NumElements(), |
| 69 | errors::InvalidArgument("tensor_names and shape_and_slices " |
| 70 | "have different number of elements: ", |
| 71 | tensor_names.NumElements(), " vs. ", |
| 72 | shape_and_slices.NumElements())); |
| 73 | OP_REQUIRES(context, |
| 74 | FastBoundsCheck(tensor_names.NumElements() + kFixedInputs, |
| 75 | std::numeric_limits<int>::max()), |
| 76 | errors::InvalidArgument("Too many inputs to the op")); |
| 77 | OP_REQUIRES( |
| 78 | context, shape_and_slices.NumElements() == num_tensors, |
| 79 | errors::InvalidArgument("Expected ", num_tensors, |
| 80 | " elements in shapes_and_slices, but got ", |
| 81 | context->input(2).NumElements())); |
| 82 | if (is_save_op) { |
| 83 | OP_REQUIRES(context, context->num_inputs() == num_tensors + kFixedInputs, |
| 84 | errors::InvalidArgument( |
| 85 | "Got ", num_tensors, " tensor names but ", |
| 86 | context->num_inputs() - kFixedInputs, " tensors.")); |
| 87 | OP_REQUIRES(context, context->num_inputs() == num_tensors + kFixedInputs, |
| 88 | errors::InvalidArgument( |
| 89 | "Expected a total of ", num_tensors + kFixedInputs, |
| 90 | " inputs as input #1 (which is a string " |
| 91 | "tensor of saved names) contains ", |
| 92 | num_tensors, " names, but received ", context->num_inputs(), |
| 93 | " inputs")); |
| 94 | } |
| 95 | } |
| 96 | |
| 97 | } // namespace |
| 98 |
no test coverage detected