mirror of
https://github.com/GoogleCloudPlatform/vertex-ai-samples.git
synced 2026-09-26 14:42:04 +00:00
Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
32a5015179 | ||
|
|
4be8b0a59a | ||
|
|
aa09d46265 | ||
|
|
5667967131 | ||
|
|
100c47a197 | ||
|
|
dc4c04346c | ||
|
|
d2797cb77c | ||
|
|
07c918da07 | ||
|
|
074afa32e7 | ||
|
|
76638f8bb4 | ||
|
|
e960c6efda |
@@ -2,4 +2,5 @@ cpr_model_server.py
|
||||
entrypoint.py
|
||||
state_dict.pth
|
||||
config.json
|
||||
**/__pycache__
|
||||
**/__pycache__
|
||||
!testdata/**
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
## About CPR
|
||||
|
||||
CPR ([custom prediction routines](https://github.com/googleapis/python-aiplatform/blob/custom-prediction-routine/google/cloud/aiplatform/prediction/README.md)) is a framework designed by Google Cloud developers to make it easier to combine machine learning models with custom preprocessing and postprocessing logic in a real-time serving application.
|
||||
CPR ([custom prediction routines](https://github.com/googleapis/python-aiplatform/blob/main/google/cloud/aiplatform/prediction/README.md)) is a framework designed by Google Cloud developers to make it easier to combine machine learning models with custom preprocessing and postprocessing logic in a real-time serving application.
|
||||
|
||||
## Using this example
|
||||
|
||||
@@ -34,6 +34,23 @@ Finally, install the Python modules required to build and run the model server:
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
### Auth
|
||||
|
||||
This example uses Google Cloud Storage for hosting model artifacts and Artifact Registry to store the container image.
|
||||
You'll need to authorize yourself before you can interact with these.
|
||||
|
||||
First, log in to GCP with application default credentials:
|
||||
```sh
|
||||
gcloud auth application-default login
|
||||
```
|
||||
|
||||
Next, if you haven't done so already, set up the [gcloud credential helper](https://cloud.google.com/artifact-registry/docs/docker/authentication)
|
||||
for the Artifact Registry region where you intend to host the image.
|
||||
```
|
||||
gcloud auth configure-docker <region>-docker.pkg.dev
|
||||
```
|
||||
|
||||
|
||||
### Predictor
|
||||
|
||||
The `TimmPredictor` class in `timm_serving/predictor.py` implements most of the important logic for the server.
|
||||
|
||||
@@ -60,9 +60,9 @@ class CPRConfig(object):
|
||||
image: str = "timm_predictor:latest"
|
||||
artifact_local_dir: str = ""
|
||||
region: str = "us-central1"
|
||||
project_id: str = "samthrasher-experimental"
|
||||
project_id: str = "<your project ID here>"
|
||||
repository: str = "cpr-images"
|
||||
artifact_gcs_dir: str = "gs://samthrasher-cpr-example/timm-vit224/"
|
||||
artifact_gcs_dir: str = "gs://<your bucket ID here>/timm-vit224/"
|
||||
model_name: str = ""
|
||||
endpoint_name: str = ""
|
||||
machine_type: str = "n1-standard-2"
|
||||
|
||||
@@ -5,4 +5,4 @@ timm==0.5.4
|
||||
smart_open==6.0.0
|
||||
|
||||
google-cloud-storage>=1.26.0,<2.0.0dev
|
||||
google-cloud-aiplatform[prediction] @ git+https://github.com/googleapis/python-aiplatform.git@custom-prediction-routine
|
||||
google-cloud-aiplatform[prediction]>=1.16.0
|
||||
@@ -70,7 +70,10 @@ class PredictorUnitTests(absltest.TestCase):
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.config = CPRConfig()
|
||||
self.config.load()
|
||||
try:
|
||||
self.config.load()
|
||||
except FileNotFoundError:
|
||||
logging.info("No saved config file found, using default values.")
|
||||
self.predictor = predictor.TimmPredictor()
|
||||
|
||||
def test_load_from_saved_state_dict_ok(self):
|
||||
@@ -170,7 +173,10 @@ class ServerEndToEndTests(absltest.TestCase):
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.config = CPRConfig()
|
||||
self.config.load()
|
||||
try:
|
||||
self.config.load()
|
||||
except FileNotFoundError:
|
||||
logging.info("No saved config file found, using default values.")
|
||||
self.local_model = cpr.LocalModel(
|
||||
serving_container_spec=aiplatform.gapic.ModelContainerSpec(
|
||||
image_uri=self.config.image
|
||||
|
||||
+1
@@ -0,0 +1 @@
|
||||
blah
|
||||
BIN
Binary file not shown.
@@ -8,7 +8,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Copyright 2021 Google LLC\n",
|
||||
"# Copyright 2022 Google LLC\n",
|
||||
"#\n",
|
||||
"# Licensed under the Apache License, Version 2.0 (the \"License\");\n",
|
||||
"# you may not use this file except in compliance with the License.\n",
|
||||
@@ -44,8 +44,9 @@
|
||||
" </a>\n",
|
||||
" </td>\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/ml_ops/stage6/get_started_with_automl_tabular_model_batch.ipynb\">\n",
|
||||
" Open in Google Cloud Notebooks\n",
|
||||
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/ml_ops/stage6/get_started_with_automl_image_model_batch.ipynb\">\n",
|
||||
" <img src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" alt=\"Vertex AI logo\">\n",
|
||||
" Open in Vertex AI Workbench\n",
|
||||
" </a>\n",
|
||||
" </td>\n",
|
||||
"</table>\n",
|
||||
@@ -61,7 +62,7 @@
|
||||
"## Overview\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"This tutorial demonstrates how to use the Vertex SDK to create image classification models and do batch prediction using a Google Cloud [AutoML](https://cloud.google.com/vertex-ai/docs/start/automl-users) model."
|
||||
"This tutorial demonstrates how to use the Vertex AI SDK to create image classification models and do batch prediction using a Google Cloud [AutoML](https://cloud.google.com/vertex-ai/docs/start/automl-users) model."
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -752,7 +753,7 @@
|
||||
"\n",
|
||||
"- JSONL\n",
|
||||
"\n",
|
||||
"The batch server accepts the following input formats for AutoML image models:\n",
|
||||
"The batch server accepts the following output formats for AutoML image models:\n",
|
||||
"\n",
|
||||
"- JSONL\n",
|
||||
"\n",
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Copyright 2021 Google LLC\n",
|
||||
"# Copyright 2022 Google LLC\n",
|
||||
"#\n",
|
||||
"# Licensed under the Apache License, Version 2.0 (the \"License\");\n",
|
||||
"# you may not use this file except in compliance with the License.\n",
|
||||
@@ -786,7 +786,9 @@
|
||||
"- CSV\n",
|
||||
"- Big Query table\n",
|
||||
"\n",
|
||||
"The batch server accepts the following input formats for AutoML tabular models:\n",
|
||||
"### Output format for batch prediction jobs\n",
|
||||
"\n",
|
||||
"The batch server accepts the following output formats for AutoML tabular models:\n",
|
||||
"\n",
|
||||
"- JSONL\n",
|
||||
"- CSV\n",
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -741,13 +741,16 @@
|
||||
"\n",
|
||||
"### Input format for batch prediction jobs\n",
|
||||
"\n",
|
||||
"The batch server accepts the following input formats:\n",
|
||||
"The batch server accepts the following input formats for custom image models:\n",
|
||||
"\n",
|
||||
"- JSONL\n",
|
||||
"- CSV\n",
|
||||
"- TFRecords\n",
|
||||
"- File-List\n",
|
||||
"- BigQuery table\n",
|
||||
"\n",
|
||||
"### Output format for batch prediction jobs\n",
|
||||
"\n",
|
||||
"The batch server accepts the following output formats for custom image models:\n",
|
||||
"\n",
|
||||
"- JSONL\n",
|
||||
"\n",
|
||||
"### Pivot format\n",
|
||||
"\n",
|
||||
@@ -1306,8 +1309,6 @@
|
||||
"source": [
|
||||
"### Send the prediction request\n",
|
||||
"\n",
|
||||
"BLAH\n",
|
||||
"\n",
|
||||
"To make a batch prediction request, call the model object's `batch_predict` method with the following parameters: \n",
|
||||
"- `instances_format`: The format of the batch prediction request file: \"jsonl\", \"csv\", \"bigquery\", \"tf-record\", \"tf-record-gzip\" or \"file-list\"\n",
|
||||
"- `prediction_format`: The format of the batch prediction response file: \"jsonl\", \"csv\", \"bigquery\", \"tf-record\", \"tf-record-gzip\" or \"file-list\"\n",
|
||||
|
||||
+269
-16
@@ -716,14 +716,19 @@
|
||||
"\n",
|
||||
"### Input format for batch prediction jobs\n",
|
||||
"\n",
|
||||
"The batch server accepts the following input formats:\n",
|
||||
"The batch server accepts the following input formats for custom tabular models:\n",
|
||||
"\n",
|
||||
"- JSONL\n",
|
||||
"- CSV\n",
|
||||
"- TFRecords\n",
|
||||
"- File-List\n",
|
||||
"- BigQuery table\n",
|
||||
"\n",
|
||||
"### Output format for batch prediction jobs\n",
|
||||
"\n",
|
||||
"The batch server accepts the following output formats for custom tabular models:\n",
|
||||
"\n",
|
||||
"- JSONL\n",
|
||||
"- BigQuery table (when input is BigQuery table)\n",
|
||||
"\n",
|
||||
"### Pivot format\n",
|
||||
"\n",
|
||||
"The batch server converts the input format to the `pivot` (JSONL) format as follows:\n",
|
||||
@@ -1098,12 +1103,12 @@
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "0d078e5e8953"
|
||||
"id": "776346c90074"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"DATA_DIR = \"gs://cloud-samples-data/ai-platform/iris/iris_data.csv\"\n",
|
||||
"! gsutil cat $data_dir | head -n 10 > test.csv\n",
|
||||
"data_file = \"gs://cloud-samples-data/ai-platform/iris/iris_data.csv\"\n",
|
||||
"! gsutil cat $data_file | head -n 10 > test.csv\n",
|
||||
"\n",
|
||||
"! cat test.csv\n",
|
||||
"\n",
|
||||
@@ -1209,6 +1214,264 @@
|
||||
" break"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "9ce65fff362b"
|
||||
},
|
||||
"source": [
|
||||
"#### Delete the batch prediction job\n",
|
||||
"\n",
|
||||
"You can delete your batch prediction job using the `delete()` method."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "5943dfde5123"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"batch_prediction_job.delete()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "d61289b87070"
|
||||
},
|
||||
"source": [
|
||||
"## Batch prediction with BigQuery input format\n",
|
||||
"\n",
|
||||
"Next, you do the same batch job, except the input format is a BigQuery table. When the input format is a BigQuery table, the output format has to be a BigQuery table as well. To use BigQuery as the input format, your model must take its input as a list (array) of values. "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "csv_to_bq"
|
||||
},
|
||||
"source": [
|
||||
"### Create a BigQuery dataset from CSV files\n",
|
||||
"\n",
|
||||
"You can create a BigQuery dataset from CSV files using the BigQuery `create_dataset()` and `load_table_from_uri()` methods, as follows:\n",
|
||||
"\n",
|
||||
"- `create_dataset()`: Creates an empty BigQuery dataset, with the following parameters:\n",
|
||||
" - `dataset_ref`: The `DatasetReference` created from the dataset_id -- e.g., samples.\n",
|
||||
"- `load_table_from_uri()`: Loads one or more CSV files into a table within the corresponding dataset, with the following parameters:\n",
|
||||
" - `url`: A set of one or more CVS files in Cloud Storage storage.\n",
|
||||
" - `table`: The `TableReference` for the table.\n",
|
||||
" - `job_config`: Specifications on how to load the CSV data.\n",
|
||||
"\n",
|
||||
"Learn more about [Importing CSV data into BigQuery](https://www.tensorflow.org/io/tutorials/bigquery#import_census_data_into_bigquery)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "csv_to_bq"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"LOCATION = \"us\"\n",
|
||||
"\n",
|
||||
"CSV_SCHEMA = [\n",
|
||||
" bigquery.SchemaField(\"sepal_length\", \"FLOAT\"),\n",
|
||||
" bigquery.SchemaField(\"sepal_width\", \"FLOAT\"),\n",
|
||||
" bigquery.SchemaField(\"petal_length\", \"FLOAT\"),\n",
|
||||
" bigquery.SchemaField(\"petal_width\", \"FLOAT\"),\n",
|
||||
"]\n",
|
||||
"\n",
|
||||
"DATASET_ID = \"batch\"\n",
|
||||
"TABLE_ID = \"slice1\"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def create_bigquery_dataset(dataset_id):\n",
|
||||
" dataset = bigquery.Dataset(\n",
|
||||
" bigquery.dataset.DatasetReference(PROJECT_ID, dataset_id)\n",
|
||||
" )\n",
|
||||
" dataset.location = \"us\"\n",
|
||||
"\n",
|
||||
" try:\n",
|
||||
" dataset = bqclient.create_dataset(dataset) # API request\n",
|
||||
" return True\n",
|
||||
" except Exception as err:\n",
|
||||
" print(err)\n",
|
||||
" if err.code != 409: # http_client.CONFLICT\n",
|
||||
" raise\n",
|
||||
" return False\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def load_data_into_bigquery(url, dataset_id, table_id):\n",
|
||||
" create_bigquery_dataset(dataset_id)\n",
|
||||
" dataset = bqclient.dataset(dataset_id)\n",
|
||||
" table = dataset.table(table_id)\n",
|
||||
"\n",
|
||||
" job_config = bigquery.LoadJobConfig()\n",
|
||||
" job_config.write_disposition = bigquery.WriteDisposition.WRITE_TRUNCATE\n",
|
||||
" job_config.source_format = bigquery.SourceFormat.CSV\n",
|
||||
" job_config.schema = CSV_SCHEMA\n",
|
||||
" job_config.skip_leading_rows = 1 # heading\n",
|
||||
"\n",
|
||||
" load_job = bqclient.load_table_from_uri(url, table, job_config=job_config)\n",
|
||||
" print(\"Starting job {}\".format(load_job.job_id))\n",
|
||||
"\n",
|
||||
" load_job.result() # Waits for table load to complete.\n",
|
||||
" print(\"Job finished.\")\n",
|
||||
"\n",
|
||||
" destination_table = bqclient.get_table(table)\n",
|
||||
" print(\"Loaded {} rows.\".format(destination_table.num_rows))\n",
|
||||
"\n",
|
||||
" return destination_table\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"bq_table = load_data_into_bigquery(data_file, DATASET_ID, TABLE_ID)\n",
|
||||
"print(bq_table)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "send_prediction_request:image"
|
||||
},
|
||||
"source": [
|
||||
"### Send the batch prediction request\n",
|
||||
"\n",
|
||||
"Again, you make a batch prediction request with the `batch_predict()` method, but with the following changes in parameters:\n",
|
||||
"\n",
|
||||
"- `instances_format`: Set to 'bigquery'\n",
|
||||
"- `predictions_format`: Set to 'bigquery'\n",
|
||||
"- `bigquery_source`: Used instead of `gcs_source`.\n",
|
||||
"- `bigquery_destination_prefix`: Used instead of `gcs_destination_prefix`"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "1cf1076178fc"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"MIN_NODES = 1\n",
|
||||
"MAX_NODES = 1\n",
|
||||
"\n",
|
||||
"# The name of the job\n",
|
||||
"BATCH_PREDICTION_JOB_NAME = \"churn_batch-\" + UUID\n",
|
||||
"\n",
|
||||
"# Folder in the bucket to write results to\n",
|
||||
"DESTINATION_FOLDER = \"batch_prediction_results_bq\"\n",
|
||||
"\n",
|
||||
"# The Cloud Storage bucket to upload results to\n",
|
||||
"BATCH_PREDICTION_GCS_DEST_PREFIX = BUCKET_URI + \"/\" + DESTINATION_FOLDER\n",
|
||||
"\n",
|
||||
"BQ_TABLE_SOURCE = f\"bq://{PROJECT_ID}.{bq_table.dataset_id}.{bq_table.table_id}\"\n",
|
||||
"\n",
|
||||
"# Make SDK batch_predict method call\n",
|
||||
"batch_prediction_job = model.batch_predict(\n",
|
||||
" instances_format=\"bigquery\",\n",
|
||||
" predictions_format=\"bigquery\",\n",
|
||||
" job_display_name=BATCH_PREDICTION_JOB_NAME,\n",
|
||||
" bigquery_source=BQ_TABLE_SOURCE,\n",
|
||||
" bigquery_destination_prefix=f\"bq://{PROJECT_ID}.batch\",\n",
|
||||
" model_parameters=None,\n",
|
||||
" machine_type=DEPLOY_COMPUTE,\n",
|
||||
" accelerator_type=DEPLOY_GPU,\n",
|
||||
" accelerator_count=DEPLOY_NGPU,\n",
|
||||
" starting_replica_count=MIN_NODES,\n",
|
||||
" max_replica_count=MAX_NODES,\n",
|
||||
" sync=True,\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "get_batch_prediction:mbsdk,custom,icn"
|
||||
},
|
||||
"source": [
|
||||
"### Get the predictions\n",
|
||||
"\n",
|
||||
"Next, get the results from the completed batch prediction job.\n",
|
||||
"\n",
|
||||
"The results are written to a BigQuery table at the BigQuery dataset path you specified as the destination. The batch server creates the table, where the table location is specified by:\n",
|
||||
"\n",
|
||||
"`batch_prediction_job.output_info.bigquery_output_dataset`: The project and dataset components.\n",
|
||||
"`batch_prediction_job.output_info.bigquery_output_table`: The table component.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "cd515ed616c8"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"BQ_RESULTS_TABLE = (\n",
|
||||
" batch_prediction_job.output_info.bigquery_output_dataset\n",
|
||||
" + \".\"\n",
|
||||
" + batch_prediction_job.output_info.bigquery_output_table\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(BQ_RESULTS_TABLE)\n",
|
||||
"\n",
|
||||
"table = bigquery.TableReference.from_string(BQ_RESULTS_TABLE[5:])\n",
|
||||
"\n",
|
||||
"rows = bqclient.list_rows(table, max_results=10)\n",
|
||||
"\n",
|
||||
"for row in rows:\n",
|
||||
" print(row)\n",
|
||||
" for key, value in row.items():\n",
|
||||
" pass"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "c7cda79d0e42"
|
||||
},
|
||||
"source": [
|
||||
"#### Delete the batch prediction job\n",
|
||||
"\n",
|
||||
"You can delete your batch prediction job using the `delete()` method."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "de4b5c638c2a"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"batch_prediction_job.delete()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "e5d4b9c51e86"
|
||||
},
|
||||
"source": [
|
||||
"#### Delete the model\n",
|
||||
"\n",
|
||||
"You can delete your model using the `delete()` method."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "00fa6a7b4f24"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"model.delete()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
@@ -1232,16 +1495,6 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"delete_bucket = False\n",
|
||||
"delete_model = True\n",
|
||||
"delete_batch_job = True\n",
|
||||
"\n",
|
||||
"if delete_model:\n",
|
||||
" try:\n",
|
||||
" model.delete()\n",
|
||||
" except Exception as e:\n",
|
||||
" print(e)\n",
|
||||
"if delete_batch_job:\n",
|
||||
" batch_prediction_job.delete()\n",
|
||||
"\n",
|
||||
"if delete_bucket or os.getenv(\"IS_TESTING\"):\n",
|
||||
" ! gsutil rm -rf {BUCKET_URI}"
|
||||
|
||||
Reference in New Issue
Block a user