| 457 | self._server.init(self._edge_sources, self._node_sources) |
| 458 | |
| 459 | def deploy_in_server_mode(self, task_index, cluster, job_name): |
| 460 | if isinstance(cluster, dict): |
| 461 | cluster_spec = cluster |
| 462 | elif isinstance(cluster, str): |
| 463 | cluster_spec = json.loads(cluster) |
| 464 | else: |
| 465 | raise ValueError("cluster must be dict or json string.") |
| 466 | |
| 467 | tracker = cluster_spec.get("tracker", "root://graphlearn") |
| 468 | |
| 469 | # parse servers |
| 470 | server_count = cluster_spec.get("server_count") |
| 471 | servers = cluster_spec.get("server") |
| 472 | if servers: |
| 473 | pywrap.set_server_hosts(servers) |
| 474 | servers = servers.split(',') |
| 475 | server_count = len(servers) |
| 476 | |
| 477 | # parse clients |
| 478 | client_count = cluster_spec.get("client_count") |
| 479 | clients = cluster_spec.get("client") |
| 480 | if clients: |
| 481 | client_count = len(clients.split(',')) |
| 482 | |
| 483 | if not server_count or not client_count: |
| 484 | raise ValueError("Invalid cluster schema") |
| 485 | |
| 486 | pywrap.set_server_count(server_count) |
| 487 | pywrap.set_client_count(client_count) |
| 488 | |
| 489 | if job_name == "client": |
| 490 | pywrap.set_tracker(tracker) |
| 491 | pywrap.set_client_id(task_index) |
| 492 | self._client = Client(client_id=task_index, in_memory=False) |
| 493 | self._server = None |
| 494 | elif job_name == "server": |
| 495 | self._client = None |
| 496 | server_host = "0.0.0.0:0" if not servers else servers[task_index] |
| 497 | self._server = Server(task_index, server_count, server_host, tracker) |
| 498 | self._server.start() |
| 499 | self._server.init(self._edge_sources, self._node_sources) |
| 500 | else: |
| 501 | raise ValueError("Only support client and server job name in SERVER mode.") |
| 502 | |
| 503 | def add_dataset(self, ds): |
| 504 | self._datasets.append(ds) |