| 1433 | } |
| 1434 | |
| 1435 | Status SavedModelOptimizer::FindVariableParts( |
| 1436 | std::unordered_map<std::string, std::vector<Node*>>& var_parts) { |
| 1437 | // TODO: only support embedding variable currently |
| 1438 | for (Node* node : graph_.nodes()) { |
| 1439 | if (node->op_def().name() == "KvVarHandleOp") { |
| 1440 | for (auto sname : option_.shard_embedding_names) { |
| 1441 | if (node->name() == sname || |
| 1442 | node->name().find(sname+"/part_") != std::string::npos) { |
| 1443 | var_parts[sname].push_back(node); |
| 1444 | } |
| 1445 | } |
| 1446 | } |
| 1447 | } |
| 1448 | |
| 1449 | for (auto sname : option_.shard_embedding_names) { |
| 1450 | if (var_parts.find(sname) == var_parts.end() || |
| 1451 | var_parts[sname].size() == 0) { |
| 1452 | return tensorflow::errors::Internal( |
| 1453 | "Can not found variable info in graph: ", sname); |
| 1454 | } |
| 1455 | } |
| 1456 | |
| 1457 | return Status::OK(); |
| 1458 | } |
| 1459 | |
| 1460 | Status SavedModelOptimizer::GetFeature2IdAttr( |
| 1461 | const std::string& name, |