| 96 | REGISTER_UNARY_VARIANT_DECODE_FUNCTION(TensorList, TensorList::kTypeName); |
| 97 | |
| 98 | bool TensorList::Decode(const VariantTensorData& data) { |
| 99 | // TODO(srbs): Change the signature to Decode(VariantTensorData data) so |
| 100 | // that we do not have to copy each tensor individually below. This would |
| 101 | // require changing VariantTensorData::tensors() as well. |
| 102 | string metadata; |
| 103 | data.get_metadata(&metadata); |
| 104 | uint64 scratch; |
| 105 | StringPiece iter(metadata); |
| 106 | std::vector<size_t> invalid_indices; |
| 107 | core::GetVarint64(&iter, &scratch); |
| 108 | size_t num_invalid_tensors = static_cast<size_t>(scratch); |
| 109 | invalid_indices.resize(num_invalid_tensors); |
| 110 | for (size_t i = 0; i < num_invalid_tensors; i++) { |
| 111 | core::GetVarint64(&iter, &scratch); |
| 112 | invalid_indices[i] = static_cast<size_t>(scratch); |
| 113 | } |
| 114 | |
| 115 | size_t total_num_tensors = data.tensors().size() + num_invalid_tensors; |
| 116 | tensors().reserve(total_num_tensors); |
| 117 | std::vector<size_t>::iterator invalid_indices_it = invalid_indices.begin(); |
| 118 | std::vector<Tensor>::const_iterator tensors_it = data.tensors().begin(); |
| 119 | for (size_t i = 0; i < total_num_tensors; i++) { |
| 120 | if (invalid_indices_it != invalid_indices.end() && |
| 121 | *invalid_indices_it == i) { |
| 122 | tensors().emplace_back(Tensor(DT_INVALID)); |
| 123 | invalid_indices_it++; |
| 124 | } else if (tensors_it != data.tensors().end()) { |
| 125 | tensors().emplace_back(*tensors_it); |
| 126 | tensors_it++; |
| 127 | } else { |
| 128 | // VariantTensorData is corrupted. |
| 129 | return false; |
| 130 | } |
| 131 | } |
| 132 | |
| 133 | core::GetVarint64(&iter, &scratch); |
| 134 | element_dtype = static_cast<DataType>(scratch); |
| 135 | core::GetVarint64(&iter, &scratch); |
| 136 | max_num_elements = static_cast<int>(scratch); |
| 137 | TensorShapeProto element_shape_proto; |
| 138 | element_shape_proto.ParseFromString(string(iter.data(), iter.size())); |
| 139 | element_shape = PartialTensorShape(element_shape_proto); |
| 140 | return true; |
| 141 | } |
| 142 | |
| 143 | Status TensorShapeFromTensor(const Tensor& t, PartialTensorShape* out) { |
| 144 | if (t.shape() == TensorShape({})) { |
nothing calls this directly
no test coverage detected