Compare commits

...
Author SHA1 Message Date
samthrasherandGitHub 32a5015179 Merge branch 'main' into samthrasher-cpr-examples 2022-09-07 10:59:58 -07:00
Andrew FerlitschandGitHub 4be8b0a59a fix: finetuning of batch notebooks (#930)
* fix: fine-tuning

* fix: fine-tuning
2022-09-07 10:44:44 -07:00
Andrew FerlitschandGitHub aa09d46265 feat: batch prediction for AutoML text models (#929)
* feat: Automl text model batch predict

* feat: Automl text model batch predict

* feat: Automl text model batch predict
2022-09-07 13:14:16 -04:00
Andrew FerlitschandGitHub 5667967131 feat: add BQ input example (#926)
* feat: add notebook for custom tabular batch predict

* feat: add notebook for custom tabular batch predict

* feat: add example for BQ input

* feat: add example for BQ input

* feat: add example for BQ input

* feat: add example for BQ input
2022-09-07 08:38:36 -07:00
Andrew FerlitschandGitHub 100c47a197 Merge branch 'main' into samthrasher-cpr-examples 2022-08-09 14:58:40 -07:00
Sam Thrasher dc4c04346c Point CPR links to main branch of SDK repo. 2022-07-28 10:08:29 -07:00
Sam Thrasher d2797cb77c fix typo 2022-07-26 09:12:05 -07:00
Sam Thrasher 07c918da07 Fix merge conflicts 2022-07-25 12:46:25 -07:00
Sam Thrasher 074afa32e7 Fix merge conflicts 2022-07-25 12:44:20 -07:00
Sam Thrasher 76638f8bb4 Minor fixes for CPR Pytorch sample: Add missing test data, add auth info to readme, scrub private project and bucket names from config, tolerate missing config.json in unit tests. 2022-07-25 12:42:16 -07:00
Sam Thrasher e960c6efda Minor fixes for CPR Pytorch sample: Add missing test data, add auth info to readme, scrub private project and bucket names from config, tolerate missing config.json in unit tests. 2022-07-25 12:39:17 -07:00
12 changed files with 1398 additions and 36 deletions
@@ -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
@@ -0,0 +1 @@
blah
@@ -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",
@@ -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}"