From 3f02c395d9a446038cedd9db048ab265e5fdaad0 Mon Sep 17 00:00:00 2001 From: graph-learning-team Date: Fri, 2 Oct 2026 12:29:05 -0700 Subject: [PATCH] Deploy GraphFlow model to Vertex AI. Adds `to_vertex_ai` function in `dgf.deploy` to deploy GraphFlow GNN models to Vertex AI Endpoints, including idempotent model registry versioning, GCS artifact verification, and zero-downtime hardware scaling. Also adds `extract_metadata_for_serving_function_signature` and `extract_prediction_schema_for_serving` to `NodePredictionModel` to generate Spanner TRAVERSE_GRAPH compatible signature JSON. PiperOrigin-RevId: 992473817 --- dgf/src/api/BUILD | 9 + dgf/src/api/__init__.py | 1 + dgf/src/api/deploy.py | 19 + dgf/src/deploy/BUILD | 49 ++ dgf/src/deploy/__init__.py | 15 + dgf/src/deploy/vertex.py | 447 +++++++++++ dgf/src/deploy/vertex_integration_test.py | 56 ++ dgf/src/deploy/vertex_test.py | 726 ++++++++++++++++++ dgf/src/learning/ten_lines/BUILD | 3 + dgf/src/learning/ten_lines/common.py | 234 +++++- dgf/src/learning/ten_lines/common_test.py | 67 ++ .../ten_lines/link_prediction_model.py | 20 + .../ten_lines/link_prediction_test.py | 171 +++++ .../ten_lines/node_prediction_model.py | 25 + .../ten_lines/node_prediction_test.py | 163 ++++ 15 files changed, 2004 insertions(+), 1 deletion(-) create mode 100644 dgf/src/api/deploy.py create mode 100644 dgf/src/deploy/BUILD create mode 100644 dgf/src/deploy/__init__.py create mode 100644 dgf/src/deploy/vertex.py create mode 100644 dgf/src/deploy/vertex_integration_test.py create mode 100644 dgf/src/deploy/vertex_test.py diff --git a/dgf/src/api/BUILD b/dgf/src/api/BUILD index 1b81cac0..36e9b567 100644 --- a/dgf/src/api/BUILD +++ b/dgf/src/api/BUILD @@ -15,6 +15,7 @@ py_library( ":analyse", ":convert", ":data", + ":deploy", ":exception", ":filesystem", ":generate", @@ -205,3 +206,11 @@ py_library( "//dgf/src/analyse:sampling", ], ) + +py_library( + name = "deploy", + srcs = ["deploy.py"], + deps = [ + "//dgf/src/deploy:vertex", + ], +) diff --git a/dgf/src/api/__init__.py b/dgf/src/api/__init__.py index 5d73c76e..e7976e68 100644 --- a/dgf/src/api/__init__.py +++ b/dgf/src/api/__init__.py @@ -30,6 +30,7 @@ from dgf.src.api import jax from dgf.src.api import exception from dgf.src.api import print +from dgf.src.api import deploy # TODO(gbm): Remove this alias. Instead, users have to do "from dgf import beam". from dgf.src.api import beam diff --git a/dgf/src/api/deploy.py b/dgf/src/api/deploy.py new file mode 100644 index 00000000..b04105e7 --- /dev/null +++ b/dgf/src/api/deploy.py @@ -0,0 +1,19 @@ +# Copyright 2022 Google LLC. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""API exports for deployment modules.""" + +# pylint: disable=unused-import,g-importing-member,g-import-not-at-top,g-bad-import-order,reimported,disable=attribute-error + +from dgf.src.deploy.vertex import to_vertex_ai diff --git a/dgf/src/deploy/BUILD b/dgf/src/deploy/BUILD new file mode 100644 index 00000000..3cf668b8 --- /dev/null +++ b/dgf/src/deploy/BUILD @@ -0,0 +1,49 @@ +load("@rules_python//python:py_library.bzl", "py_library") +load("@rules_python//python:py_test.bzl", "py_test") + +package( + + default_visibility = ["//visibility:public"], +) + +py_library( + name = "vertex", + srcs = ["vertex.py"], + deps = [ + "//dgf/src/learning/ten_lines:common", + "//dgf/src/util:filesystem", + "//dgf/src/util:log", + "//dgf/src/util/weak_dep:weak_dep_tensorflow", + # google/cloud/aiplatform dep, + # yaml dep, + ], +) + +py_test( + name = "vertex_test", + srcs = ["vertex_test.py"], + shard_count = 5, + deps = [ + ":vertex", + # absl/testing:absltest dep, + "//dgf/src/learning/ten_lines:common", + "//dgf/src/util/weak_dep:weak_dep_tensorflow", + # yaml dep, + ], +) + +py_test( + name = "vertex_integration_test", + srcs = ["vertex_integration_test.py"], + tags = [ + "manual", + "notap", + ], + deps = [ + ":vertex", + # absl/testing:absltest dep, + # absl/testing:parameterized dep, + "//dgf/src/learning/ten_lines:node_prediction", + "//dgf/src/util:gen_test_graph", + ], +) diff --git a/dgf/src/deploy/__init__.py b/dgf/src/deploy/__init__.py new file mode 100644 index 00000000..a644bd82 --- /dev/null +++ b/dgf/src/deploy/__init__.py @@ -0,0 +1,15 @@ +# Copyright 2022 Google LLC. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Deploy package.""" diff --git a/dgf/src/deploy/vertex.py b/dgf/src/deploy/vertex.py new file mode 100644 index 00000000..78004764 --- /dev/null +++ b/dgf/src/deploy/vertex.py @@ -0,0 +1,447 @@ +# Copyright 2022 Google LLC. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Vertex AI deployment utilities for GraphFlow models.""" + +import tempfile +from typing import Any, Optional + +from dgf.src.learning.ten_lines import common +from dgf.src.util import filesystem +from dgf.src.util import log +from dgf.src.util.weak_dep.weak_dep_tensorflow import tf # pylint: disable=g-importing-member +from google.cloud.aiplatform import aiplatform +import yaml + + +DEFAULT_SERVING_CONTAINER_URI = ( + "us-docker.pkg.dev/vertex-ai/prediction/tf2-cpu.2-12:latest" +) + + +def _create_batched_serving_model(tf_model: Any) -> Any: + """Wraps a GraphFlow TF model to handle Vertex AI's batched JSON instances. + + Vertex AI sends prediction requests as a JSON list of instances, which the + container parses into a batched tensor (e.g. adding a [None] outermost + dimension). GraphFlow models natively expect unbatched 1D arrays. This wrapper + strips the batch dimension on the input, passes it to the model, and adds a + batch dimension to the output to satisfy Vertex AI's response contract. + + Args: + tf_model: The raw TensorFlow callable exported from a GraphFlow model. + + Returns: + A wrapped `tf.Module` with a concrete `__call__` signature accepting batched + inputs. + """ + + class BatchedServingWrapper(tf.Module): + """Module wrapper stripping outer batch dimension.""" + + def __init__(self, raw_model): + super().__init__() + self.raw_model = raw_model + + @tf.function + def __call__(self, **kwargs): + # Vertex AI wraps instances in a batch dimension. We strip it here (v[0]) + unbatched_kwargs = {k: v[0] for k, v in kwargs.items()} + out = self.raw_model(**unbatched_kwargs) + # Re-add the batch dimension to the predictions + return {"predictions": tf.expand_dims(out, axis=0)} + + batched_model = BatchedServingWrapper(tf_model) + + # Dynamically trace tf.function using TensorSpecs to avoid needing dummy data + unbatched_spec = tf_model.__call__.structured_input_signature[1] + batched_spec = {} + for k, spec in unbatched_spec.items(): + batched_shape = [None] + spec.shape.as_list() + batched_spec[k] = tf.TensorSpec( + shape=batched_shape, dtype=spec.dtype, name=k + ) + + batched_model.__call__ = batched_model.__call__.get_concrete_function( + **batched_spec + ) + return batched_model + + +def _save_model_to_gcs( + model: common.Model, + target_gcs_dir: str, + instance_schema: dict[str, Any], + prediction_schema: dict[str, Any], + project: Optional[str] = None, +) -> None: + """Saves the model and serving schemas to GCS if not already present.""" + done_path = f"{target_gcs_dir}/DONE" + model_uuid = model.metadata.uuid + + gcs_save_needed = True + if filesystem.exists(done_path, project=project): + log.info( + "Step 1 (Save to GCS): SKIP - Model already exists in GCS (%s)", + model_uuid, + ) + gcs_save_needed = False + elif filesystem.exists(target_gcs_dir, project=project): + raise RuntimeError( + f"Corrupted GCS directory state: Directory {target_gcs_dir} already " + "contains files, but 'DONE' file is missing. This indicates a " + "previous upload crashed or was interrupted mid-way. " + "Please delete the directory in GCS and try again." + ) + + if gcs_save_needed: + log.info("Step 1 (Save to GCS): EXECUTE - Saving model to GCS...") + + with tempfile.TemporaryDirectory() as tmp_dir: + model.save(tmp_dir) + + # Export to TF SavedModel for Vertex AI + tf_model = model.to_tensorflow_function(consume_tf_graph_dict=True) + + batched_model = _create_batched_serving_model(tf_model) + tf.saved_model.save(batched_model, tmp_dir) + + # Upload all files from tmp_dir to GCS via filesystem utility + filesystem.copy_local_dir(tmp_dir, target_gcs_dir, project=project) + + # 1. Instance Schema (Request) + instance_schema_yaml = yaml.dump(instance_schema, sort_keys=False) + instance_schema_path = f"{target_gcs_dir}/instance_schema.yaml" + filesystem.write_text( + instance_schema_path, instance_schema_yaml, project=project + ) + log.info("Instance Schema YAML saved to %s", instance_schema_path) + + # 2. Prediction Schema (Response) + pred_yaml_str = yaml.dump(prediction_schema, sort_keys=False) + pred_schema_path = f"{target_gcs_dir}/prediction_schema.yaml" + filesystem.write_text(pred_schema_path, pred_yaml_str, project=project) + log.info("Prediction Schema YAML saved to %s", pred_schema_path) + + filesystem.write_text(done_path, "", project=project) + log.info("Model saved to %s", target_gcs_dir) + + +def _upload_to_vertex_ai( + model: common.Model, + target_gcs_dir: str, + display_name: str, + serving_container_image_uri: str, +) -> Any: + """Uploads the model to Vertex AI Model Registry or reuses an existing version.""" + model_uuid = model.metadata.uuid + log.info("Step 2 (Upload to Vertex): Checking existing models...") + existing_models = aiplatform.Model.list( + filter=f'display_name="{display_name}"', order_by="create_time desc" + ) + + vertex_model = None + parent_model_name = None + if existing_models: + parent_model_name = existing_models[0].resource_name + for existing_model in existing_models: + # Verify both the custom uuid label and GCS artifact URI match + existing_uri = (existing_model.uri or "").rstrip("/") + if ( + existing_model.labels + and existing_model.labels.get("uuid") == model_uuid + and existing_uri == target_gcs_dir.rstrip("/") + ): + log.info( + "Step 2 (Upload to Vertex): SKIP - Model UUID (%s) and GCS URI" + " match.", + model_uuid, + ) + vertex_model = existing_model + break + + if not vertex_model: + log.info( + "Step 2 (Upload to Vertex): EXECUTE - Uploading new model version..." + ) + upload_kwargs = { + "display_name": display_name, + "artifact_uri": target_gcs_dir, + "serving_container_image_uri": serving_container_image_uri, + "labels": {"uuid": model_uuid}, + "sync": True, + } + if parent_model_name: + upload_kwargs["parent_model"] = parent_model_name + upload_kwargs["is_default_version"] = True + + # Attach predict schemata + upload_kwargs["instance_schema_uri"] = ( + f"{target_gcs_dir}/instance_schema.yaml" + ) + upload_kwargs["prediction_schema_uri"] = ( + f"{target_gcs_dir}/prediction_schema.yaml" + ) + + vertex_model = aiplatform.Model.upload(**upload_kwargs) + + return vertex_model + + +def _get_or_create_endpoint( + endpoint_id: Optional[str], + display_name: str, + use_dedicated_endpoint: bool, +) -> Any: + """Retrieves an existing Vertex AI Endpoint or creates a new one.""" + if endpoint_id: + log.info( + "Step 3 (Create Endpoint): SKIP - Bypassing creation to use" + " provided endpoint_id %s", + endpoint_id, + ) + try: + endpoint = aiplatform.Endpoint(endpoint_name=endpoint_id) + except Exception as e: + raise RuntimeError( + f"Failed to load Vertex AI Endpoint with ID '{endpoint_id}'. Please" + " verify that the endpoint exists in your GCP project and location." + f" Original error: {e}" + ) from e + else: + log.info("Step 3 (Create Endpoint): Checking existing endpoints...") + existing_endpoints = aiplatform.Endpoint.list( + filter=f'display_name="{display_name}"', order_by="create_time desc" + ) + if existing_endpoints: + log.info( + "Step 3 (Create Endpoint): SKIP - Endpoint '%s' already exists.", + display_name, + ) + # Note: Reusing an existing endpoint does not automatically update the + # serving model version on Vertex AI's side. Instead, Step 4 + # (_deploy_model_to_endpoint) compares the target model's version_id + # against currently deployed models on this endpoint; if a new version + # was uploaded in Step 2, Step 4 will deploy it with 100% traffic and + # undeploy older model versions. + endpoint = existing_endpoints[0] + else: + log.info( + "Step 3 (Create Endpoint): EXECUTE - Creating new endpoint..." + ) + + # Vertex AI 'Dedicated Endpoints' increase max payload limits to 10MB + # and provide VPC isolation. + endpoint_kwargs = { + "display_name": display_name, + "sync": True, + } + if use_dedicated_endpoint: + endpoint_kwargs["dedicated_endpoint_enabled"] = True + + endpoint = aiplatform.Endpoint.create(**endpoint_kwargs) + + return endpoint + + +def _deploy_model_to_endpoint( + endpoint: Any, + vertex_model: Any, + machine_type: str, + blocking: bool, +) -> None: + """Deploys the target model version to the endpoint and cleans up old versions.""" + log.info("Step 4 (Deploy Model): Checking current deployments...") + deployed_models = endpoint.list_models() + already_deployed = False + matched_deployed_id = None + existing_deployed_ids = {d.id for d in deployed_models} + + target_version = vertex_model.version_id + + for deployed in deployed_models: + deployed_model_name = deployed.model or "" + if "@" in deployed_model_name: + deployed_base, deployed_version = deployed_model_name.split("@", 1) + else: + deployed_base = deployed_model_name + deployed_version = deployed.model_version_id + + same_version = (deployed_base == vertex_model.resource_name) and ( + target_version is None + or deployed_version is None + or str(target_version) == str(deployed_version) + ) + if same_version: + deployed_machine = None + if deployed.dedicated_resources: + deployed_machine = ( + deployed.dedicated_resources.machine_spec.machine_type + ) + + if deployed_machine == machine_type: + already_deployed = True + matched_deployed_id = deployed.id + log.info( + "Step 4 (Deploy Model): SKIP - Model version is already deployed" + " with matching hardware." + ) + break + + if not already_deployed: + log.info( + "Step 4 (Deploy Model): EXECUTE - Deploying model to endpoint (this" + " usually takes 10-15 minutes)..." + ) + endpoint.deploy( + model=vertex_model, + machine_type=machine_type, + traffic_percentage=100, + sync=blocking, + ) + log.info( + "Model deployed successfully to endpoint: %s", endpoint.resource_name + ) + + # Cleanup: Undeploy ALL old models to free up compute resources + if already_deployed: + for deployed in deployed_models: + if deployed.id != matched_deployed_id: + log.info( + "Undeploying old model (ID: %s) to prevent resource drain...", + deployed.id, + ) + endpoint.undeploy(deployed_model_id=deployed.id, sync=False) + elif blocking: + deployed_models = endpoint.list_models() + for deployed in deployed_models: + if deployed.id in existing_deployed_ids: + log.info( + "Undeploying old model (ID: %s) to prevent resource drain...", + deployed.id, + ) + endpoint.undeploy(deployed_model_id=deployed.id, sync=False) + else: + log.info( + "Skipping immediate undeployment of old models because blocking=False" + " (deployment is in progress)." + ) + + +def to_vertex_ai( + model: common.Model, + *, + model_dir_on_gcs: str, + display_name: str, + location: str, + project: Optional[str] = None, + machine_type: str = "n1-standard-4", + use_dedicated_endpoint: bool = True, + serving_container_image_uri: str = DEFAULT_SERVING_CONTAINER_URI, + endpoint_id: Optional[str] = None, + blocking: bool = True, +) -> Any: + """Deploy a GraphFlow GNN model as an inference endpoint on Vertex AI. + + Usage example: + ``` + graph, schema = dgf.io.read_graph("/tmp/my_hgraph") + model = dgf.learning.train_node_model(graph=graph, schema=schema, ...) + endpoint = dgf.deploy.to_vertex_ai( + model, + model_dir_on_gcs="gs://my-bucket/models", + display_name="my_gnn_endpoint", + location="us-central1", + project="my-gcp-project", + ) + # Run online predictions or inspect the deployed endpoint on Vertex AI: + response = endpoint.predict(instances=[{...}]) + ``` + + Args: + model: A trained DGF model (e.g. NodePredictionModel or + LinkPredictionModel). Note: calling this function populates + `model.serving_function_signature` so that it is persisted with the + SavedModel artifact. + model_dir_on_gcs: Base GCS directory where models should be saved. + display_name: Display name for the Vertex AI Model and Endpoint. Used + together with `model.metadata.uuid` to identify existing model versions in + Vertex Model Registry, and to look up or create the target Endpoint (if + `endpoint_id` is not provided). + location: GCP Region. + project: GCP Project ID. + machine_type: Machine type for the Vertex AI Endpoint. + use_dedicated_endpoint: Whether to use dedicated resources. + serving_container_image_uri: Docker image URI for Vertex AI prediction. + endpoint_id: If provided, deploy to this existing endpoint ID. + blocking: If True, waits for endpoint creation and model deployment to + finish. + + Returns: + The Vertex AI Endpoint object. + + Raises: + ValueError: If the model does not have a UUID in its metadata. + NotImplementedError: If the model does not support Vertex AI serving schema + extraction. + RuntimeError: If GCS state is corrupted or explicit endpoint lookup fails. + """ + + aiplatform.init(project=project, location=location) + + # Pre-flight Checks + if not model.metadata or not model.metadata.uuid: + raise ValueError( + "Model does not have a UUID in its metadata. A UUID must be assigned" + " during training to ensure proper lineage tracking." + ) + + instance_schema, prediction_schema = model._extract_serving_schemata() # pylint: disable=protected-access + # Populate serving_function_signature on the model instance so model.save() + # writes serving_function_signature.yaml into the SavedModel directory. + model.serving_function_signature = yaml.dump( + instance_schema, default_flow_style=False, sort_keys=False + ) + + # Step 1: Save Model to GCS + # Construct target GCS directory using UUID to prevent overwrites + model_dir_on_gcs = model_dir_on_gcs.rstrip("/") + target_gcs_dir = f"{model_dir_on_gcs}/{model.metadata.uuid}" + + _save_model_to_gcs( + model, + target_gcs_dir, + instance_schema=instance_schema, + prediction_schema=prediction_schema, + project=project, + ) + + # Step 2: Upload to Vertex AI + vertex_model = _upload_to_vertex_ai( + model, + target_gcs_dir, + display_name, + serving_container_image_uri, + ) + + # Step 3: Create Endpoint + endpoint = _get_or_create_endpoint( + endpoint_id, display_name, use_dedicated_endpoint + ) + + # Step 4: Deploy Model + _deploy_model_to_endpoint(endpoint, vertex_model, machine_type, blocking) + + return endpoint + diff --git a/dgf/src/deploy/vertex_integration_test.py b/dgf/src/deploy/vertex_integration_test.py new file mode 100644 index 00000000..f8a38429 --- /dev/null +++ b/dgf/src/deploy/vertex_integration_test.py @@ -0,0 +1,56 @@ +# Copyright 2022 Google LLC. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os + +from absl.testing import absltest +from absl.testing import parameterized +from dgf.src.deploy import vertex +from dgf.src.learning.ten_lines import node_prediction_train +from dgf.src.util import gen_test_graph + + +class VertexIntegrationTest(parameterized.TestCase): + + def test_deploy_model(self): + graph, schema = gen_test_graph.gen_toy_classification_dataset() + + # Extract training config + model = node_prediction_train.train_node_model( + graph=graph, + schema=schema, + target_nodeset="N1", + target_column="label", + num_train_steps=10, + batch_size=8, + ) + + test_dir = self.create_tempdir().full_path + endpoint = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs=os.path.join(test_dir, "models"), + display_name="test-integration-model", + location="us-central1", + machine_type="n1-standard-4", + blocking=True, + ) + self.assertIsNotNone(endpoint) + + # Clean up (undeploy) + endpoint.undeploy_all() + endpoint.delete() + + +if __name__ == "__main__": + absltest.main() diff --git a/dgf/src/deploy/vertex_test.py b/dgf/src/deploy/vertex_test.py new file mode 100644 index 00000000..8c42f9b8 --- /dev/null +++ b/dgf/src/deploy/vertex_test.py @@ -0,0 +1,726 @@ +# Copyright 2022 Google LLC. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Hermetic unit tests for dgf.src.deploy.vertex covering all 6 Design Doc cases.""" + +import types +from typing import Any, Optional +from unittest import mock +from absl.testing import absltest +from dgf.src.deploy import vertex +from dgf.src.learning.ten_lines import common +from dgf.src.util.weak_dep.weak_dep_tensorflow import tf # pylint: disable=g-importing-member +import yaml + + +class StatefulGCPMock: + """Stateful mock of GCS and Vertex AI (Model Registry + Endpoints) per region.""" + + def __init__(self): + self.current_location = "us-central1" + self.current_project = "test-project" + # (bucket_name, blob_path) -> str content + self.gcs_blobs = {} + self.gcs_upload_calls = 0 + + # (location, display_name) -> list of mock Model objects (newest first) + self.models_by_name = {} + self.model_upload_calls = [] + + # (location, endpoint_id) -> mock Endpoint object + self.endpoints_by_id = {} + self.endpoint_create_calls = [] + self.endpoint_list_calls = 0 + self.deploy_calls = [] + self.undeploy_calls = [] + self._id_counter = 1000 + + def _next_id(self) -> str: + self._id_counter += 1 + return str(self._id_counter) + + def init_aiplatform(self, project=None, location=None): + if project: + self.current_project = project + if location: + self.current_location = location + + def create_storage_client(self, project=None): + del project + mock_client = mock.MagicMock() + + def get_bucket(bucket_name): + mock_bucket = mock.MagicMock() + + def get_blob(blob_path): + mock_blob = mock.MagicMock() + key = (bucket_name, blob_path) + mock_blob.exists.side_effect = lambda: key in self.gcs_blobs + mock_blob.download_as_text.side_effect = lambda: self.gcs_blobs[key] + + def upload_str(content): + self.gcs_blobs[key] = content + self.gcs_upload_calls += 1 + + def upload_file(local_path): + self.gcs_blobs[key] = f"file:{local_path}" + self.gcs_upload_calls += 1 + + mock_blob.upload_from_string.side_effect = upload_str + mock_blob.upload_from_filename.side_effect = upload_file + return mock_blob + + def list_blobs(prefix="", max_results=None): + del max_results + matches = [] + for b_name, b_path in self.gcs_blobs: + if b_name == bucket_name and b_path.startswith(prefix): + matches.append(mock.MagicMock(name=b_path)) + return matches + + mock_bucket.blob.side_effect = get_blob + mock_bucket.list_blobs.side_effect = list_blobs + return mock_bucket + + mock_client.bucket.side_effect = get_bucket + return mock_client + + def model_list(self, filter=None, order_by=None): # pylint: disable=redefined-builtin + del order_by + # filter is of form: display_name="foo" + display_name = (filter or "").split('"')[1] + key = (self.current_location, display_name) + return list(self.models_by_name.get(key, [])) + + def model_upload( + self, + display_name, + artifact_uri, + serving_container_image_uri, + labels=None, + sync=True, + parent_model=None, + is_default_version=False, + instance_schema_uri=None, + prediction_schema_uri=None, + ): + self.model_upload_calls.append({ + "location": self.current_location, + "display_name": display_name, + "artifact_uri": artifact_uri, + "labels": labels, + "parent_model": parent_model, + "is_default_version": is_default_version, + "instance_schema_uri": instance_schema_uri, + "prediction_schema_uri": prediction_schema_uri, + "serving_container_image_uri": serving_container_image_uri, + "sync": sync, + }) + + key = (self.current_location, display_name) + existing = self.models_by_name.setdefault(key, []) + if parent_model: + base_resource_name = parent_model + version_id = str(len(existing) + 1) + else: + model_id = self._next_id() + base_resource_name = ( + f"projects/{self.current_project}/locations/{self.current_location}" + f"/models/{model_id}" + ) + version_id = "1" + + mock_model = types.SimpleNamespace( + resource_name=base_resource_name, + version_id=version_id, + uri=artifact_uri, + labels=dict(labels or {}), + display_name=display_name, + ) + # Insert at front so existing[0] is latest version + existing.insert(0, mock_model) + return mock_model + + def _make_endpoint_obj(self, endpoint_id, display_name, location): + resource_name = ( + f"projects/{self.current_project}/locations/{location}" + f"/endpoints/{endpoint_id}" + ) + ep = types.SimpleNamespace( + name=endpoint_id, + resource_name=resource_name, + display_name=display_name, + location=location, + _deployed_models=[], + ) + + def list_models(): + return list(ep._deployed_models) + + def deploy( + model, machine_type="n1-standard-4", traffic_percentage=100, sync=True + ): + dep_id = f"dep-{self._next_id()}" + self.deploy_calls.append({ + "endpoint": resource_name, + "model": model.resource_name, + "version_id": model.version_id, + "machine_type": machine_type, + "traffic_percentage": traffic_percentage, + "sync": sync, + }) + dep_obj = types.SimpleNamespace( + id=dep_id, + model=model.resource_name, + model_version_id=str(model.version_id), + dedicated_resources=types.SimpleNamespace( + machine_spec=types.SimpleNamespace(machine_type=machine_type) + ), + ) + ep._deployed_models.append(dep_obj) + + def undeploy(deployed_model_id, sync=False): + self.undeploy_calls.append({ + "endpoint": resource_name, + "deployed_model_id": deployed_model_id, + "sync": sync, + }) + ep._deployed_models = [ + d for d in ep._deployed_models if d.id != deployed_model_id + ] + + ep.list_models = list_models + ep.deploy = deploy + ep.undeploy = undeploy + return ep + + def endpoint_create( + self, display_name, sync=True, dedicated_endpoint_enabled=False + ): + endpoint_id = self._next_id() + self.endpoint_create_calls.append({ + "location": self.current_location, + "display_name": display_name, + "endpoint_id": endpoint_id, + "sync": sync, + "dedicated_endpoint_enabled": dedicated_endpoint_enabled, + }) + ep = self._make_endpoint_obj( + endpoint_id, display_name, self.current_location + ) + self.endpoints_by_id[(self.current_location, endpoint_id)] = ep + return ep + + def endpoint_list(self, filter=None, order_by=None): # pylint: disable=redefined-builtin + del order_by + self.endpoint_list_calls += 1 + display_name = (filter or "").split('"')[1] + matches = [ + ep + for (loc, _), ep in self.endpoints_by_id.items() + if loc == self.current_location and ep.display_name == display_name + ] + return matches + + def endpoint_ctor(self, endpoint_name): + key = (self.current_location, endpoint_name) + if key not in self.endpoints_by_id: + raise ValueError( + f"Endpoint {endpoint_name} not found in {self.current_location}" + ) + return self.endpoints_by_id[key] + + +class FakeGNNModel(common.Model): + """Lightweight fake DGF model for hermetic deployment testing.""" + + def __init__(self, model_uuid: str): + super().__init__(data=None) + self.metadata = common.Metadata(name="FakeGNNModel", uuid=model_uuid) + + @classmethod + def name(cls) -> str: + return "FakeGNNModel" + + def describe(self): + return None + + def data(self): + return None + + def _internal_save(self, path: str) -> None: + pass + + def _internal_load(self, path: str) -> None: + pass + + def save(self, tmp_dir: str): + # Write a dummy file into tmp_dir so os.walk finds an artifact to upload + with open(f"{tmp_dir}/saved_model.pb", "w") as f: + f.write("dummy_model") + + def to_tensorflow_function( + self, + *, + input_format: Any = None, + consume_tf_graph_dict: Optional[bool] = None, + ) -> Any: + del input_format, consume_tf_graph_dict + return mock.MagicMock() + + def _extract_serving_schemata( + self, + ) -> tuple[dict[str, Any], dict[str, Any]]: + sig = { + "title": f"NodePrediction_TestGraph_user_{self.metadata.uuid}", + "type": "object", + "required": ["gnn_user_seed_node_idxs"], + "x-google-graph": "TestGraph", + "x-google-gnn-input-graphs": [{ + "input_node": "user", + "sampling_plan": [{ + "edge": "follows", + "width": 5, + }], + }], + "properties": {"age": {"type": "integer"}}, + } + pred_schema = { + "type": "object", + "properties": {"predictions": {"type": "array"}}, + } + return sig, pred_schema + + +class VertexDeployTest(absltest.TestCase): + + def setUp(self): + super().setUp() + self.gcp = StatefulGCPMock() + + # Patch GCP & TF API calls with MagicMocks backed by StatefulGCPMock + self.mock_aiplatform_init = self.enter_context( + mock.patch.object( + vertex.aiplatform, "init", side_effect=self.gcp.init_aiplatform + ) + ) + self.mock_storage_client = self.enter_context( + mock.patch.object( + vertex.filesystem.storage, + "Client", + side_effect=self.gcp.create_storage_client, + ) + ) + self.mock_model_list = self.enter_context( + mock.patch.object( + vertex.aiplatform.Model, "list", side_effect=self.gcp.model_list + ) + ) + self.mock_model_upload = self.enter_context( + mock.patch.object( + vertex.aiplatform.Model, "upload", side_effect=self.gcp.model_upload + ) + ) + self.mock_endpoint_cls = self.enter_context( + mock.patch.object( + vertex.aiplatform, "Endpoint", side_effect=self.gcp.endpoint_ctor + ) + ) + self.mock_endpoint_create = mock.MagicMock( + side_effect=self.gcp.endpoint_create + ) + self.mock_endpoint_list = mock.MagicMock(side_effect=self.gcp.endpoint_list) + self.mock_endpoint_cls.create = self.mock_endpoint_create + self.mock_endpoint_cls.list = self.mock_endpoint_list + + self._real_create_batched_serving_model = ( + vertex._create_batched_serving_model + ) # pylint: disable=protected-access + self.mock_batched_wrapper = self.enter_context( + mock.patch.object( + vertex, + "_create_batched_serving_model", + return_value=mock.MagicMock(), + ) + ) + self.mock_tf_saved_model_save = self.enter_context( + mock.patch("tensorflow.saved_model.save") + ) + + def _reset_api_mocks(self): + """Resets call histories on all mocked GCP/TF API functions between steps.""" + self.mock_aiplatform_init.reset_mock() + self.mock_storage_client.reset_mock() + self.mock_model_list.reset_mock() + self.mock_model_upload.reset_mock() + self.mock_endpoint_cls.reset_mock() + self.mock_endpoint_create.reset_mock() + self.mock_endpoint_list.reset_mock() + self.mock_tf_saved_model_save.reset_mock() + + def _deploy_baseline(self): + """Helper to establish the Case 0 baseline state in the mock GCP backend.""" + model = FakeGNNModel("uuid-1") + ep0 = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-v1", + location="us-central1", + machine_type="n2-standard-4", + blocking=True, + ) + self._reset_api_mocks() + return model, ep0 + + def test_preflight_validation_checks(self): + model_no_uuid = FakeGNNModel("uuid-1") + model_no_uuid.metadata.uuid = None + with self.assertRaisesRegex(ValueError, "Model does not have a UUID"): + vertex.to_vertex_ai( + model=model_no_uuid, + model_dir_on_gcs="gs://bucket/dir", + display_name="test-model", + location="us-central1", + ) + + class UnsupportedModel(FakeGNNModel): + + def _extract_serving_schemata(self): + return super(FakeGNNModel, self)._extract_serving_schemata() + + model_no_sig = UnsupportedModel("uuid-1") + with self.assertRaisesRegex( + NotImplementedError, "does not support Vertex AI serving schema" + ): + vertex.to_vertex_ai( + model=model_no_sig, + model_dir_on_gcs="gs://bucket/dir", + display_name="test-model", + location="us-central1", + ) + + def test_corrupted_gcs_state_raises_error(self): + model = FakeGNNModel("uuid-1") + self.gcp.gcs_blobs[("my-bucket", "dir/uuid-1/partial_file.pb")] = ( + "partial-data" + ) + with self.assertRaisesRegex( + RuntimeError, + "Corrupted GCS directory state: Directory gs://my-bucket/dir/uuid-1" + " already contains files, but 'DONE' file is missing.", + ): + vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://my-bucket/dir", + display_name="test-model", + location="us-central1", + ) + + def test_case_0_initial_deploy_and_idempotency_check(self): + model = FakeGNNModel("uuid-1") + + # 1. Initial Deploy (All 4 steps EXECUTE via mocked APIs) + ep0 = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-v1", + location="us-central1", + machine_type="n2-standard-4", + blocking=True, + ) + self.mock_tf_saved_model_save.assert_called_once() + self.mock_model_upload.assert_called_once() + self.assertIsNone( + self.mock_model_upload.call_args.kwargs.get("parent_model") + ) + self.mock_endpoint_create.assert_called_once_with( + display_name="gnn-model-v1", + sync=True, + dedicated_endpoint_enabled=True, + ) + self.assertLen(self.gcp.deploy_calls, 1) + + # 2. Idempotency Check (Re-run with NO changes -> All 4 steps SKIP) + self._reset_api_mocks() + ep0b = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-v1", + location="us-central1", + machine_type="n2-standard-4", + blocking=True, + ) + self.assertEqual(ep0b.resource_name, ep0.resource_name) + # Step 1 SKIP: TF SavedModel save not called + self.mock_tf_saved_model_save.assert_not_called() + # Step 2 SKIP: Model.upload not called + self.mock_model_upload.assert_not_called() + # Step 3 SKIP: Endpoint.create not called + self.mock_endpoint_create.assert_not_called() + # Step 4 SKIP: No new deploy or undeploy calls + self.assertLen(self.gcp.deploy_calls, 1) + self.assertEmpty(self.gcp.undeploy_calls) + + def test_case_1_model_retrained_new_uuid(self): + """Case 1: `model` changes (new UUID) -> Step 1 EXEC, Step 2 EXEC (@2), Step 3 SKIP, Step 4 EXEC.""" + model, ep0 = self._deploy_baseline() + model.metadata.uuid = "uuid-2" + + ep1 = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-v1", + location="us-central1", + machine_type="n2-standard-4", + blocking=True, + ) + # Step 1 EXECUTED: Saved new model.UUID to GCS + self.mock_tf_saved_model_save.assert_called_once() + # Step 2 EXECUTED: Uploaded new model version under existing parent_model + self.mock_model_upload.assert_called_once() + self.assertIsNotNone( + self.mock_model_upload.call_args.kwargs.get("parent_model") + ) + self.assertTrue( + self.mock_model_upload.call_args.kwargs.get("is_default_version") + ) + # Step 3 SKIPPED: Reused existing endpoint + self.mock_endpoint_create.assert_not_called() + self.assertEqual(ep1.resource_name, ep0.resource_name) + # Step 4 EXECUTED: Deployed version 2 and undeployed version 1 + self.assertLen(self.gcp.deploy_calls, 2) + self.assertEqual(self.gcp.deploy_calls[-1]["version_id"], "2") + self.assertLen(self.gcp.undeploy_calls, 1) + + def test_case_2_model_dir_on_gcs_changes(self): + """Case 2: `model_dir_on_gcs` changes -> Step 1 EXEC, Step 2 EXEC (new URI), Step 3 SKIP, Step 4 EXEC.""" + model, ep0 = self._deploy_baseline() + + ep2 = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir2", # Changed GCS path + display_name="gnn-model-v1", + location="us-central1", + machine_type="n2-standard-4", + blocking=True, + ) + # Step 1 EXECUTED: Saved to new GCS directory + self.mock_tf_saved_model_save.assert_called_once() + # Step 2 EXECUTED: Detected URI mismatch and uploaded new version + self.mock_model_upload.assert_called_once() + self.assertEqual( + self.mock_model_upload.call_args.kwargs["artifact_uri"], + "gs://test-bucket/dir2/uuid-1", + ) + # Step 3 SKIPPED: Reused existing endpoint + self.mock_endpoint_create.assert_not_called() + self.assertEqual(ep2.resource_name, ep0.resource_name) + # Step 4 EXECUTED: Deployed new version and cleaned up old version + self.assertLen(self.gcp.deploy_calls, 2) + self.assertLen(self.gcp.undeploy_calls, 1) + + def test_case_3_display_name_changes(self): + """Case 3: `display_name` changes -> Step 1 SKIP, Step 2 EXEC, Step 3 EXEC, Step 4 EXEC.""" + model, ep0 = self._deploy_baseline() + + ep3 = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-v2", # Changed display_name + location="us-central1", + machine_type="n2-standard-4", + blocking=True, + ) + # Step 1 SKIPPED: GCS artifact already exists for uuid-1 + self.mock_tf_saved_model_save.assert_not_called() + # Step 2 EXECUTED: Uploaded brand new Model resource (parent_model=None) + self.mock_model_upload.assert_called_once() + self.assertIsNone( + self.mock_model_upload.call_args.kwargs.get("parent_model") + ) + self.assertEqual( + self.mock_model_upload.call_args.kwargs["display_name"], "gnn-model-v2" + ) + # Step 3 EXECUTED: Created brand new Endpoint resource + self.mock_endpoint_create.assert_called_once() + self.assertNotEqual(ep3.resource_name, ep0.resource_name) + # Step 4 EXECUTED: Deployed to new endpoint + self.assertLen(self.gcp.deploy_calls, 2) + self.assertEqual(self.gcp.deploy_calls[-1]["endpoint"], ep3.resource_name) + + def test_case_4_machine_type_changes(self): + """Case 4: `machine_type` changes -> Step 1 SKIP, Step 2 SKIP, Step 3 SKIP, Step 4 EXEC.""" + model, ep0 = self._deploy_baseline() + + ep4 = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-v1", + location="us-central1", + machine_type="n1-standard-8", # Scaled hardware + blocking=True, + ) + # Steps 1, 2, 3 all SKIPPED + self.mock_tf_saved_model_save.assert_not_called() + self.mock_model_upload.assert_not_called() + self.mock_endpoint_create.assert_not_called() + self.assertEqual(ep4.resource_name, ep0.resource_name) + # Step 4 EXECUTED: Redeployed on n1-standard-8 after hardware mismatch + self.assertLen(self.gcp.deploy_calls, 2) + self.assertEqual(self.gcp.deploy_calls[-1]["machine_type"], "n1-standard-8") + self.assertLen(self.gcp.undeploy_calls, 1) + + def test_case_5_explicit_endpoint_id_provided(self): + """Case 5: `endpoint_id` provided -> Step 1 SKIP, Step 2 SKIP, Step 3 SKIP (bypass search), Step 4 EXEC.""" + model, _ = self._deploy_baseline() + # Create a second standalone target endpoint to deploy onto via explicit ID + target_ep = self.gcp.endpoint_create(display_name="explicit-target-ep") + self._reset_api_mocks() + + ep5 = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-v1", + location="us-central1", + machine_type="n2-standard-4", + endpoint_id=target_ep.name, # Explicit endpoint_id + blocking=True, + ) + # Steps 1 & 2 SKIPPED + self.mock_tf_saved_model_save.assert_not_called() + self.mock_model_upload.assert_not_called() + # Step 3 SKIPPED creation AND bypassed Endpoint.list search + self.mock_endpoint_create.assert_not_called() + self.mock_endpoint_list.assert_not_called() + self.mock_endpoint_cls.assert_called_once_with(endpoint_name=target_ep.name) + self.assertEqual(ep5.resource_name, target_ep.resource_name) + # Step 4 EXECUTED: Deployed onto target_ep + self.assertLen(self.gcp.deploy_calls, 2) + self.assertEqual( + self.gcp.deploy_calls[-1]["endpoint"], target_ep.resource_name + ) + + def test_case_6_location_changes(self): + """Case 6: `location` changes -> Step 1 SKIP (global GCS), Step 2 EXEC, Step 3 EXEC, Step 4 EXEC.""" + model, _ = self._deploy_baseline() + + ep6 = vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-v1", + location="us-east1", # Changed GCP region + machine_type="n2-standard-4", + blocking=True, + ) + # Step 1 SKIPPED: GCS is global and uuid-1 already exists in bucket + self.mock_tf_saved_model_save.assert_not_called() + # Step 2 EXECUTED: Uploaded new model to us-east1 registry + self.mock_model_upload.assert_called_once() + self.assertEqual(self.gcp.model_upload_calls[-1]["location"], "us-east1") + # Step 3 EXECUTED: Created new endpoint in us-east1 + self.mock_endpoint_create.assert_called_once() + self.assertIn("/locations/us-east1/", ep6.resource_name) + # Step 4 EXECUTED: Deployed to us-east1 endpoint + self.assertLen(self.gcp.deploy_calls, 2) + self.assertEqual(self.gcp.deploy_calls[-1]["endpoint"], ep6.resource_name) + + def test_instance_schema_yaml_contains_graph_and_sampling_config(self): + """Verifies that instance_schema.yaml uploaded to GCS contains x-google-graph and x-google-gnn-input-graphs.""" + model = FakeGNNModel(model_uuid="uuid-graph-schema") + vertex.to_vertex_ai( + model=model, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="gnn-model-graph-schema", + location="us-central1", + blocking=True, + ) + yaml_str = self.gcp.gcs_blobs.get( + ("test-bucket", "dir1/uuid-graph-schema/instance_schema.yaml") + ) + self.assertIsNotNone(yaml_str) + parsed = yaml.safe_load(yaml_str) + expected_yaml = { + "title": f"NodePrediction_TestGraph_user_{model.metadata.uuid}", + "type": "object", + "required": ["gnn_user_seed_node_idxs"], + "x-google-graph": "TestGraph", + "x-google-gnn-input-graphs": [{ + "input_node": "user", + "sampling_plan": [{ + "edge": "follows", + "width": 5, + }], + }], + "properties": {"age": {"type": "integer"}}, + } + self.assertEqual(parsed, expected_yaml) + + def test_model_uuid_match_not_first_in_list(self): + """Verifies that _upload_to_vertex_ai searches all matching models by UUID.""" + model_v1 = FakeGNNModel("uuid-v1") + model_v2 = FakeGNNModel("uuid-v2") + + # Deploy v1 first, then v2 under the same display_name (so v2 is + # existing_models[0] and v1 is existing_models[1]). + vertex.to_vertex_ai( + model=model_v1, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="shared-display-name", + location="us-central1", + ) + vertex.to_vertex_ai( + model=model_v2, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="shared-display-name", + location="us-central1", + ) + + # Now re-deploy v1; it should find v1 at index 1 in existing_models and + # skip Model.upload. + self._reset_api_mocks() + vertex.to_vertex_ai( + model=model_v1, + model_dir_on_gcs="gs://test-bucket/dir1", + display_name="shared-display-name", + location="us-central1", + ) + self.mock_model_upload.assert_not_called() + + def test_create_batched_serving_model(self): + """Verifies _create_batched_serving_model strips input batch dim and adds output batch dim.""" + + class DummyRawModel(tf.Module): + + @tf.function + def __call__(self, **kwargs): + return kwargs["x"] * 2.0 + + raw_model = DummyRawModel() + raw_model.__call__ = raw_model.__call__.get_concrete_function( + x=tf.TensorSpec(shape=[3], dtype=tf.float32, name="x") + ) + + batched_model = self._real_create_batched_serving_model(raw_model) + out = batched_model.__call__( + x=tf.constant([[1.0, 2.0, 3.0]], dtype=tf.float32) + ) + self.assertIn("predictions", out) + self.assertEqual(out["predictions"].shape.as_list(), [1, 3]) + self.assertSequenceAlmostEqual( + out["predictions"].numpy().tolist()[0], [2.0, 4.0, 6.0] + ) + + +if __name__ == "__main__": + absltest.main() diff --git a/dgf/src/learning/ten_lines/BUILD b/dgf/src/learning/ten_lines/BUILD index b2aac055..d463eddc 100644 --- a/dgf/src/learning/ten_lines/BUILD +++ b/dgf/src/learning/ten_lines/BUILD @@ -21,6 +21,8 @@ py_library( # dataclasses_json dep, "//dgf/src/data:in_memory_graph", "//dgf/src/data:schema", + "//dgf/src/io:feature_format", + "//dgf/src/io:tf", "//dgf/src/learning:early_stopping_monitor", "//dgf/src/learning/jax:common", "//dgf/src/learning/jax/layers:hetero_gnn", @@ -300,6 +302,7 @@ py_test( ":common", # absl/testing:absltest dep, # absl/testing:parameterized dep, + "//dgf/src/data:schema", "//dgf/src/util:gen_test_graph", "//dgf/src/util:log", "//dgf/src/util:test_util", diff --git a/dgf/src/learning/ten_lines/common.py b/dgf/src/learning/ten_lines/common.py index 4806ee15..3eb6ea66 100644 --- a/dgf/src/learning/ten_lines/common.py +++ b/dgf/src/learning/ten_lines/common.py @@ -18,11 +18,13 @@ import dataclasses import enum import os -from typing import Any, Literal, TypeAlias +from typing import Any, Optional, Literal, TypeAlias import uuid import dataclasses_json from dgf.src.data import in_memory_graph from dgf.src.data import schema as schema_lib +from dgf.src.io import feature_format as feature_format_lib +from dgf.src.io import tf as io_tf_lib from dgf.src.learning import early_stopping_monitor from dgf.src.learning.jax import common as jax_common from dgf.src.learning.jax.layers import hetero_gnn @@ -39,6 +41,7 @@ import numpy as np import orbax.checkpoint as ocp + # The types of graphs supported. Graph = dataset.Graph @@ -277,6 +280,7 @@ def __init__(self, data: Any) -> None: method. """ self.metadata = Metadata(name=self.name()) + self.serving_function_signature: Optional[str] = None @abc.abstractmethod def describe(self) -> util.RichDisplay: @@ -341,6 +345,26 @@ def _internal_load(self, path: str) -> None: path: The directory path from which the model data should be loaded. """ + def to_tensorflow_function( + self, + *args: Any, + **kwargs: Any, + ) -> Any: + """Exports the model as a TensorFlow callable function.""" + raise NotImplementedError( + f"Model of type {type(self).__name__} does not support TensorFlow" + " export." + ) + + def _extract_serving_schemata( + self, + ) -> tuple[dict[str, Any], dict[str, Any]]: + """Extracts (instance_schema, prediction_schema) dicts for Vertex AI serving.""" + raise NotImplementedError( + f"Model of type {type(self).__name__} does not support Vertex AI" + " serving schema extraction." + ) + @dataclasses_json.dataclass_json @dataclasses.dataclass(kw_only=True) @@ -743,3 +767,211 @@ def extract_graph_metrics( metrics["num_edges"] = num_edges return metrics + + +def schema_to_serving_signature_dict( + schema_: schema_lib.GraphSchema, + target_nodeset: Optional[str], + model_name: str, + model_uuid: str, + target_edgeset: Optional[str] = None, + source_sampling_plan: Optional[Any] = None, + target_sampling_plan: Optional[Any] = None, +) -> dict[str, Any]: + """Converts a GraphSchema to a Spanner TRAVERSE_GRAPH serving signature dict. + + Args: + schema_: The graph schema. + target_nodeset: The target nodeset name for node prediction. + model_name: The registered name of the model class. + model_uuid: The unique identifier of the model instance. + target_edgeset: Optional target edgeset name for link prediction. + source_sampling_plan: Optional SamplingPlan for multi-hop GNN traversal + (used as primary plan for node prediction). + target_sampling_plan: Optional SamplingPlan for target node traversal in + edge prediction. + + Returns: + A dictionary representing the Spanner TRAVERSE_GRAPH instance schema. + """ + + def _convert_plan_edge(plan_edge: Any) -> dict[str, Any]: + step: dict[str, Any] = { + "edge": plan_edge.edgeset, + "width": plan_edge.hop_width, + } + if plan_edge.reversed: + step["reverse"] = True + if plan_edge.node and plan_edge.node.children: + step["children"] = [ + _convert_plan_edge(child) for child in plan_edge.node.children + ] + return step + + if target_nodeset is None and target_edgeset is None: + raise ValueError( + "Either target_nodeset or target_edgeset must be provided." + ) + if target_nodeset is not None and target_edgeset is not None: + raise ValueError( + "target_nodeset and target_edgeset are mutually exclusive." + ) + + target_name = target_nodeset if target_nodeset else target_edgeset + signature: dict[str, Any] = { + "title": f"{model_name}_{target_name}_{model_uuid}", + "type": "object", + "required": [], + } + + def _format_feature_spec( + feat_schema: schema_lib.FeatureSchema, + ) -> dict[str, str]: + tf_dtype = feature_format_lib.FEATURE_FORMAT_TO_TF_DTYPE[feat_schema.format] + dtype = f"tf.{tf_dtype.name}" + if feat_schema.semantic == schema_lib.FeatureSemantic.CATEGORICAL: + dtype = "tf.int64" + + if not feat_schema.shape: + shape_str = "(None,)" + else: + dims = ["-1" if d is None else str(d) for d in feat_schema.shape] + shape_str = f"({', '.join(dims)})" if len(dims) > 1 else f"({dims[0]},)" + return {"shape": shape_str, "dtype": dtype} + + def _add_graph_to_signature(prefix_str: str, input_node: str): + signature[f"{prefix_str}_seed_node_idxs"] = { + "shape": "(None,)", + "dtype": "tf.int32", + "input_node": input_node, + "field_kind": "seed_node_idxs", + } + + def _add_features( + entity_name: str, + features: dict[str, schema_lib.FeatureSchema], + is_edge: bool, + ): + for feat_name, feat_schema in features.items(): + if is_edge: + feat_key = io_tf_lib.tf_graph_dict_edge_key(entity_name, feat_name) + label_key = "edge_label" + else: + feat_key = io_tf_lib.tf_graph_dict_node_key(entity_name, feat_name) + label_key = "node_label" + signature[f"{prefix_str}_{feat_key}"] = { + **_format_feature_spec(feat_schema), + "input_node": input_node, + label_key: entity_name, + "property": feat_name, + "field_kind": "feature", + } + + for nodeset_name, nodeset_schema in schema_.node_sets.items(): + size_key = io_tf_lib.tf_graph_dict_node_key( + nodeset_name, io_tf_lib.TF_GRAPH_DICT_SIZE_KEY + ) + signature[f"{prefix_str}_{size_key}"] = { + "shape": "()", + "dtype": "tf.int32", + "input_node": input_node, + "node_label": nodeset_name, + "field_kind": "size", + } + _add_features(nodeset_name, nodeset_schema.features, is_edge=False) + + for edge_identifier, edgeset_schema in schema_.edge_sets.items(): + edgeset_name = ( + edge_identifier[1] + if isinstance(edge_identifier, tuple) + else edge_identifier + ) + size_key = io_tf_lib.tf_graph_dict_edge_key( + edgeset_name, io_tf_lib.TF_GRAPH_DICT_SIZE_KEY + ) + signature[f"{prefix_str}_{size_key}"] = { + "shape": "()", + "dtype": "tf.int32", + "input_node": input_node, + "edge_label": edgeset_name, + "field_kind": "size", + } + adj_key = io_tf_lib.tf_graph_dict_edge_key( + edgeset_name, io_tf_lib.TF_GRAPH_DICT_ADJACENCY_KEY + ) + signature[f"{prefix_str}_{adj_key}"] = { + "shape": "(2, None)", + "dtype": "tf.int64", + "input_node": input_node, + "edge_label": edgeset_name, + "field_kind": "adjacency", + } + _add_features(edgeset_name, edgeset_schema.features, is_edge=True) + + signature["x-google-gnn-input-graphs"] = [] + + if target_edgeset is not None: + # Link Prediction + source_input_node = schema_.edge_sets[target_edgeset].source + target_input_node = schema_.edge_sets[target_edgeset].target + if not source_input_node or not target_input_node: + raise ValueError("Source or target input node cannot be empty.") + source_input_node = str(source_input_node) + target_input_node = str(target_input_node) + + source_plan = [] + if source_sampling_plan: + source_plan = [ + _convert_plan_edge(edge) + for edge in source_sampling_plan.root.children + ] + target_plan = [] + if target_sampling_plan: + target_plan = [ + _convert_plan_edge(edge) + for edge in target_sampling_plan.root.children + ] + + signature["x-google-gnn-input-graphs"].append({ + "input_node": source_input_node, + "sampling_plan": source_plan, + }) + signature["x-google-gnn-input-graphs"].append({ + "input_node": target_input_node, + "sampling_plan": target_plan, + }) + + _add_graph_to_signature("source", source_input_node) + _add_graph_to_signature("target", target_input_node) + else: + # Node Prediction + input_node = target_nodeset + + if source_sampling_plan is not None: + signature["x-google-gnn-input-graphs"].append({ + "input_node": input_node, + "sampling_plan": [ + _convert_plan_edge(edge) + for edge in source_sampling_plan.root.children + ], + }) + else: + signature["x-google-gnn-input-graphs"].append({ + "input_node": input_node, + "sampling_plan": [], + }) + _add_graph_to_signature(f"gnn_{target_nodeset}", str(input_node)) + + signature["required"] = [ + k + for k in signature.keys() + if k + not in ( + "x-google-graph", + "x-google-gnn-input-graphs", + "title", + "type", + "required", + ) + ] + return signature diff --git a/dgf/src/learning/ten_lines/common_test.py b/dgf/src/learning/ten_lines/common_test.py index 6ecbd6ee..0aa95a96 100644 --- a/dgf/src/learning/ten_lines/common_test.py +++ b/dgf/src/learning/ten_lines/common_test.py @@ -19,6 +19,7 @@ from absl.testing import absltest from absl.testing import parameterized +from dgf.src.data import schema as schema_lib from dgf.src.learning.ten_lines import common from dgf.src.util import gen_test_graph from dgf.src.util import log @@ -351,6 +352,72 @@ def test_extract_graph_metrics(self): self.assertGreater(metrics["num_edges"], 0) self.assertGreater(metrics["num_features"], 0) + def test_schema_to_serving_signature_dict(self): + schema = schema_lib.GraphSchema( + node_sets={ + "n1": schema_lib.NodeSchema( + features={ + "f_none": schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.FLOAT_32, + shape=None, + ), + "f_none_23": schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.FLOAT_32, + shape=(None, 23), + ), + "f_none_1": schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.FLOAT_32, + shape=(None, 1), + ), + } + ) + }, + edge_sets={}, + ) + signature = common.schema_to_serving_signature_dict( + schema_=schema, + target_nodeset="n1", + model_name="test_model", + model_uuid="test_uuid", + ) + + f_none = signature["gnn_n1_nodes_n1_f_none"] + self.assertEqual(f_none["shape"], "(None,)") + + f_none_23 = signature["gnn_n1_nodes_n1_f_none_23"] + self.assertEqual(f_none_23["shape"], "(-1, 23)") + + f_none_1 = signature["gnn_n1_nodes_n1_f_none_1"] + self.assertEqual(f_none_1["shape"], "(-1, 1)") + + def test_schema_to_serving_signature_dict_link_prediction(self): + schema = schema_lib.GraphSchema( + node_sets={ + "n1": schema_lib.NodeSchema(), + }, + edge_sets={ + "e1": schema_lib.EdgeSchema( + source="n1", + target="n1", + features={ + "f_none": schema_lib.FeatureSchema( + format=schema_lib.FeatureFormat.FLOAT_32, + shape=None, + ), + } + ) + }, + ) + signature = common.schema_to_serving_signature_dict( + schema_=schema, + target_nodeset=None, + target_edgeset="e1", + model_name="test_model", + model_uuid="test_uuid", + ) + f_none = signature["source_edges_e1_f_none"] + self.assertEqual(f_none["shape"], "(None,)") + if __name__ == "__main__": absltest.main() diff --git a/dgf/src/learning/ten_lines/link_prediction_model.py b/dgf/src/learning/ten_lines/link_prediction_model.py index 3a967de8..904fa1a4 100644 --- a/dgf/src/learning/ten_lines/link_prediction_model.py +++ b/dgf/src/learning/ten_lines/link_prediction_model.py @@ -102,6 +102,7 @@ class ModelData: target_sampling_plan: sampling_config_lib.SamplingPlan training_stats: TrainingStats temporal_sampling: bool = False + nodeset_timestamp_features: dict[str, str] = dataclasses.field( default_factory=dict ) @@ -292,6 +293,25 @@ def _internal_save(self, path: str) -> None: def _internal_load(self, path: str) -> None: self._data.model_params = common.load_params(path) + def _extract_serving_schemata( + self, + ) -> tuple[dict[str, Any], dict[str, Any]]: + """Extracts (instance_schema, prediction_schema) dicts for Vertex AI serving.""" + if not self.metadata.uuid: + raise ValueError("Model UUID is not set.") + model_uuid = self.metadata.uuid + instance_schema = common.schema_to_serving_signature_dict( + schema_=self.data().schema, + target_nodeset=None, + target_edgeset=self.data().task.target_edgeset, + model_name=self.name(), + model_uuid=model_uuid, + source_sampling_plan=self.data().source_sampling_plan, + target_sampling_plan=self.data().target_sampling_plan, + ) + prediction_schema = {"type": "array", "items": {"type": "number"}} + return instance_schema, prediction_schema + def describe(self) -> util.RichDisplay: """Rich display for colab.""" tabs = [] diff --git a/dgf/src/learning/ten_lines/link_prediction_test.py b/dgf/src/learning/ten_lines/link_prediction_test.py index 0671d7c4..7223761c 100644 --- a/dgf/src/learning/ten_lines/link_prediction_test.py +++ b/dgf/src/learning/ten_lines/link_prediction_test.py @@ -901,6 +901,177 @@ def test_model(self): self.assertEqual(self.model.data().training_stats.num_train_seed_edges, 2) self.assertEqual(self.model.data().training_stats.num_valid_seed_edges, 2) + def test_extract_serving_schemata(self): + signature, pred_schema = self.model._extract_serving_schemata() + self.assertEqual( + pred_schema, {"type": "array", "items": {"type": "number"}} + ) + expected_dict = { + "x-google-gnn-input-graphs": [{ + "input_node": "A", + "sampling_plan": [], + }, { + "input_node": "B", + "sampling_plan": [], + }], + "title": f"LinkPrediction_A_to_B_{self.model.metadata.uuid}", + "type": "object", + "required": [ + "source_seed_node_idxs", + "source_nodes_A_reserved_size", + "source_nodes_A_#id", + "source_nodes_A_f", + "source_nodes_B_reserved_size", + "source_nodes_B_#id", + "source_nodes_B_f", + "source_edges_A_to_B_reserved_size", + "source_edges_A_to_B_reserved_adjacency", + "target_seed_node_idxs", + "target_nodes_A_reserved_size", + "target_nodes_A_#id", + "target_nodes_A_f", + "target_nodes_B_reserved_size", + "target_nodes_B_#id", + "target_nodes_B_f", + "target_edges_A_to_B_reserved_size", + "target_edges_A_to_B_reserved_adjacency", + ], + "source_seed_node_idxs": { + "shape": "(None,)", + "dtype": "tf.int32", + "input_node": "A", + "field_kind": "seed_node_idxs", + }, + "source_nodes_A_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "A", + "node_label": "A", + "field_kind": "size", + }, + "source_nodes_A_#id": { + "shape": "(None,)", + "dtype": "tf.int32", + "input_node": "A", + "node_label": "A", + "property": "#id", + "field_kind": "feature", + }, + "source_nodes_A_f": { + "shape": "(None,)", + "dtype": "tf.float32", + "input_node": "A", + "node_label": "A", + "property": "f", + "field_kind": "feature", + }, + "source_nodes_B_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "A", + "node_label": "B", + "field_kind": "size", + }, + "source_nodes_B_#id": { + "shape": "(None,)", + "dtype": "tf.int32", + "input_node": "A", + "node_label": "B", + "property": "#id", + "field_kind": "feature", + }, + "source_nodes_B_f": { + "shape": "(None,)", + "dtype": "tf.float32", + "input_node": "A", + "node_label": "B", + "property": "f", + "field_kind": "feature", + }, + "source_edges_A_to_B_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "A", + "edge_label": "A_to_B", + "field_kind": "size", + }, + "source_edges_A_to_B_reserved_adjacency": { + "shape": "(2, None)", + "dtype": "tf.int64", + "input_node": "A", + "edge_label": "A_to_B", + "field_kind": "adjacency", + }, + "target_seed_node_idxs": { + "shape": "(None,)", + "dtype": "tf.int32", + "input_node": "B", + "field_kind": "seed_node_idxs", + }, + "target_nodes_A_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "B", + "node_label": "A", + "field_kind": "size", + }, + "target_nodes_A_#id": { + "shape": "(None,)", + "dtype": "tf.int32", + "input_node": "B", + "node_label": "A", + "property": "#id", + "field_kind": "feature", + }, + "target_nodes_A_f": { + "shape": "(None,)", + "dtype": "tf.float32", + "input_node": "B", + "node_label": "A", + "property": "f", + "field_kind": "feature", + }, + "target_nodes_B_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "B", + "node_label": "B", + "field_kind": "size", + }, + "target_nodes_B_#id": { + "shape": "(None,)", + "dtype": "tf.int32", + "input_node": "B", + "node_label": "B", + "property": "#id", + "field_kind": "feature", + }, + "target_nodes_B_f": { + "shape": "(None,)", + "dtype": "tf.float32", + "input_node": "B", + "node_label": "B", + "property": "f", + "field_kind": "feature", + }, + "target_edges_A_to_B_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "B", + "edge_label": "A_to_B", + "field_kind": "size", + }, + "target_edges_A_to_B_reserved_adjacency": { + "shape": "(2, None)", + "dtype": "tf.int64", + "input_node": "B", + "edge_label": "A_to_B", + "field_kind": "adjacency", + }, + } + self.maxDiff = None + self.assertDictEqual(signature, expected_dict) + def test_predict(self): predictions = self.model.predict( self.graph, source_node_idxs=[0, 1], target_node_idxs=[0, 1, 1, 2] diff --git a/dgf/src/learning/ten_lines/node_prediction_model.py b/dgf/src/learning/ten_lines/node_prediction_model.py index f247cc37..807138ca 100644 --- a/dgf/src/learning/ten_lines/node_prediction_model.py +++ b/dgf/src/learning/ten_lines/node_prediction_model.py @@ -25,6 +25,7 @@ import enum import itertools import textwrap +from typing import Any import dataclasses_json from dgf.src.data import in_memory_graph @@ -157,6 +158,7 @@ class ModelData: feature_stats: statistics_lib.GraphFeatureStatistics training_stats: TrainingStats temporal_sampling: bool + nodeset_timestamp_features: dict[str, str] = dataclasses.field( default_factory=dict ) @@ -215,6 +217,29 @@ def _internal_save(self, path: str) -> None: def _internal_load(self, path: str) -> None: self._data.model_params = common.load_params(path) + def _extract_serving_schemata( + self, + ) -> tuple[dict[str, Any], dict[str, Any]]: + """Extracts (instance_schema, prediction_schema) dicts for Vertex AI serving.""" + if not self.metadata.uuid: + raise ValueError("Model UUID is not set.") + model_uuid = self.metadata.uuid + instance_schema = common.schema_to_serving_signature_dict( + schema_=self.data().schema, + target_nodeset=self.data().task.target_nodeset, + model_name=self.name(), + model_uuid=model_uuid, + source_sampling_plan=self.data().sampling_plan, + ) + task_type = self.data().task.task_type + if task_type == TaskType.NODE_REGRESSION: + prediction_schema = {"type": "number"} + elif task_type == TaskType.NODE_CLASSIFICATION: + prediction_schema = {"type": "array", "items": {"type": "number"}} + else: + raise ValueError(f"Unsupported task_type: {task_type}") + return instance_schema, prediction_schema + def describe(self) -> util.RichDisplay: # TODO(gbm): Make a good rich report. diff --git a/dgf/src/learning/ten_lines/node_prediction_test.py b/dgf/src/learning/ten_lines/node_prediction_test.py index 0916abd1..3604aafa 100644 --- a/dgf/src/learning/ten_lines/node_prediction_test.py +++ b/dgf/src/learning/ten_lines/node_prediction_test.py @@ -270,6 +270,169 @@ def setUpClass(cls): **RAPID_TRAINING_KWARGS, ) + def test_extract_serving_schemata(self): + signature, pred_schema = self.model._extract_serving_schemata() + self.assertEqual( + pred_schema, {"type": "array", "items": {"type": "number"}} + ) + + # Verify the entire signature dictionary matches expectations + expected_dict = { + "x-google-gnn-input-graphs": [{ + "input_node": "client", + "sampling_plan": [{ + "edge": "transation_to_client", + "width": 5, + "reverse": True, + }], + }], + "title": f"NodePrediction_client_{self.model.metadata.uuid}", + "type": "object", + "required": [ + "gnn_client_seed_node_idxs", + "gnn_client_nodes_client_reserved_size", + "gnn_client_nodes_client_#id", + "gnn_client_nodes_client_city", + "gnn_client_nodes_client_age", + "gnn_client_nodes_client_created_at", + "gnn_client_nodes_client_categorical_label", + "gnn_client_nodes_transaction_reserved_size", + "gnn_client_nodes_transaction_#id", + "gnn_client_nodes_transaction_date", + "gnn_client_nodes_transaction_amount", + "gnn_client_nodes_transaction_country", + "gnn_client_edges_transation_to_client_reserved_size", + "gnn_client_edges_transation_to_client_reserved_adjacency", + ], + "gnn_client_seed_node_idxs": { + "shape": "(None,)", + "dtype": "tf.int32", + "input_node": "client", + "field_kind": "seed_node_idxs", + }, + "gnn_client_nodes_client_#id": { + "shape": "(None,)", + "dtype": "tf.string", + "input_node": "client", + "node_label": "client", + "property": "#id", + "field_kind": "feature", + }, + "gnn_client_nodes_client_city": { + "shape": "(None,)", + "dtype": "tf.int64", + "input_node": "client", + "node_label": "client", + "property": "city", + "field_kind": "feature", + }, + "gnn_client_nodes_client_age": { + "shape": "(None,)", + "dtype": "tf.int64", + "input_node": "client", + "node_label": "client", + "property": "age", + "field_kind": "feature", + }, + "gnn_client_nodes_client_created_at": { + "shape": "(None,)", + "dtype": "tf.int64", + "input_node": "client", + "node_label": "client", + "property": "created_at", + "field_kind": "feature", + }, + "gnn_client_nodes_client_categorical_label": { + "shape": "(None,)", + "dtype": "tf.int64", + "input_node": "client", + "node_label": "client", + "property": "categorical_label", + "field_kind": "feature", + }, + "gnn_client_nodes_client_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "client", + "node_label": "client", + "field_kind": "size", + }, + "gnn_client_nodes_transaction_#id": { + "shape": "(None,)", + "dtype": "tf.string", + "input_node": "client", + "node_label": "transaction", + "property": "#id", + "field_kind": "feature", + }, + "gnn_client_nodes_transaction_date": { + "shape": "(None,)", + "dtype": "tf.int64", + "input_node": "client", + "node_label": "transaction", + "property": "date", + "field_kind": "feature", + }, + "gnn_client_nodes_transaction_amount": { + "shape": "(None,)", + "dtype": "tf.int64", + "input_node": "client", + "node_label": "transaction", + "property": "amount", + "field_kind": "feature", + }, + "gnn_client_nodes_transaction_country": { + "shape": "(None,)", + "dtype": "tf.int64", + "input_node": "client", + "node_label": "transaction", + "property": "country", + "field_kind": "feature", + }, + "gnn_client_nodes_transaction_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "client", + "node_label": "transaction", + "field_kind": "size", + }, + "gnn_client_edges_transation_to_client_reserved_size": { + "shape": "()", + "dtype": "tf.int32", + "input_node": "client", + "edge_label": "transation_to_client", + "field_kind": "size", + }, + "gnn_client_edges_transation_to_client_reserved_adjacency": { + "shape": "(2, None)", + "dtype": "tf.int64", + "input_node": "client", + "edge_label": "transation_to_client", + "field_kind": "adjacency", + }, + } + self.assertDictEqual(signature, expected_dict) + + # Manually hack the schema to test multi-dimensional shape + original_shape = ( + self.model.data().schema.node_sets["client"].features["age"].shape + ) + self.model.data().schema.node_sets["client"].features["age"].shape = ( + 1, + 128, + 64, + ) + try: + signature_with_shape, _ = self.model._extract_serving_schemata() + self.assertEqual( + signature_with_shape["gnn_client_nodes_client_age"]["shape"], + "(1, 128, 64)", + ) + finally: + self.model.data().schema.node_sets["client"].features[ + "age" + ].shape = original_shape + def test_predict(self): predictions = self.model.predict(graph=self.graph, seed_node_idxs=[0, 1, 2]) self.assertEqual(predictions.shape, (3, self.model.num_label_classes()))