| 189 | } |
| 190 | |
| 191 | int64 GetTotalKeysNum(const std::string& part_var_name, |
| 192 | std::string& tensor_key, |
| 193 | std::string& tensor_value, |
| 194 | BundleReader* reader, |
| 195 | OpKernelContext* ctx, |
| 196 | int dim_len, |
| 197 | int& curr_part_index, |
| 198 | int partition_num, |
| 199 | size_t key_size, size_t value_size) { |
| 200 | while(curr_part_index < partition_num) { |
| 201 | TensorShape key_shape, value_shape; |
| 202 | reader->LookupTensorShape(tensor_key, &key_shape); |
| 203 | reader->LookupTensorShape(tensor_value, &value_shape); |
| 204 | if (value_shape.dim_size(1) != dim_len) { |
| 205 | ctx->CtxFailure(errors::InvalidArgument( |
| 206 | strings::StrCat("value_shape.dim_size(1) not equal " |
| 207 | "the dim_len attr value. ", |
| 208 | std::to_string(value_shape.dim_size(1)), |
| 209 | " vs ", std::to_string(dim_len)))); |
| 210 | return -1; |
| 211 | } |
| 212 | |
| 213 | Status s = reader->LookupHeader(tensor_key, |
| 214 | key_size * key_shape.dim_size(0)); |
| 215 | if (!s.ok()) { |
| 216 | ctx->CtxFailure(s); |
| 217 | return -1; |
| 218 | } |
| 219 | |
| 220 | s = reader->LookupHeader(tensor_value, |
| 221 | value_size * value_shape.dim_size(0) * dim_len); |
| 222 | if (!s.ok()) { |
| 223 | ctx->CtxFailure(s); |
| 224 | return -1; |
| 225 | } |
| 226 | |
| 227 | int64 total_keys_num = key_shape.dim_size(0); |
| 228 | if (total_keys_num > 0) return total_keys_num; |
| 229 | |
| 230 | LOG(WARNING) << "Current variable partitions' key num is 0. " |
| 231 | << tensor_key << ", " << tensor_value; |
| 232 | |
| 233 | ++curr_part_index; |
| 234 | // try next variable partition |
| 235 | tensor_key = strings::StrCat( |
| 236 | part_var_name, std::to_string(curr_part_index), "-keys"); |
| 237 | tensor_value = strings::StrCat( |
| 238 | part_var_name, std::to_string(curr_part_index), "-values"); |
| 239 | } |
| 240 | |
| 241 | return 0; |
| 242 | } |
| 243 | |
| 244 | BatchSetCallback make_import_callback( |
| 245 | OpKernelContext* ctx, |
no test coverage detected