mirror of
https://github.com/GoogleCloudPlatform/vertex-ai-samples.git
synced 2026-09-26 22:51:56 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
45071a5fac | ||
|
|
c6a2c8e673 | ||
|
|
3e090ce61d | ||
|
|
8061ce96ed | ||
|
|
5eeb4b1fd9 | ||
|
|
bc14aa72b9 | ||
|
|
c9faf848b2 | ||
|
|
ea89e8e1f6 | ||
|
|
aa8e2c6af3 | ||
|
|
2529586683 | ||
|
|
66601678e0 | ||
|
|
75da4cfb99 | ||
|
|
7d406848ea | ||
|
|
2aea69a022 | ||
|
|
4a56ab4138 | ||
|
|
b7c7691c52 | ||
|
|
57da2a745f | ||
|
|
d0ea048385 | ||
|
|
15537b2ed8 | ||
|
|
a3239d8b71 | ||
|
|
6ecebd2973 | ||
|
|
c4fe6fad23 | ||
|
|
a162051187 | ||
|
|
0d7d96e09c | ||
|
|
5a1549d5cc | ||
|
|
aed9434c7f | ||
|
|
5ce45d1222 | ||
|
|
58b8cba1e9 | ||
|
|
da1e2b17f3 | ||
|
|
f98a7aa4d6 | ||
|
|
84393171c2 | ||
|
|
d7f07167aa | ||
|
|
2511bb1246 | ||
|
|
4e3fbed9f7 | ||
|
|
6726b90892 | ||
|
|
f257b2f924 | ||
|
|
2ace347737 | ||
|
|
572cbaf8b7 | ||
|
|
eb93058f8e | ||
|
|
45e9a2db9e | ||
|
|
530b4a6978 | ||
|
|
3a9d88d642 | ||
|
|
fe5dc2bfbe | ||
|
|
d4c96f334d | ||
|
|
487c04a216 | ||
|
|
45536c644c | ||
|
|
297d7b7e71 | ||
|
|
4a9e61a007 | ||
|
|
ef6c7457c3 | ||
|
|
578287d3ef | ||
|
|
0fb27b69b6 | ||
|
|
f327bd4bec | ||
|
|
efc37b883e | ||
|
|
5e8011040d | ||
|
|
7189513c27 | ||
|
|
ed560893d9 | ||
|
|
0df92f8127 | ||
|
|
b43d97e2c6 | ||
|
|
5b4f20a1ef | ||
|
|
cac8816db7 | ||
|
|
6d00ad9e95 | ||
|
|
65a5f952c4 | ||
|
|
6662639e30 | ||
|
|
461e1e6f8b | ||
|
|
e27db74595 | ||
|
|
e4895da43b | ||
|
|
a331687d86 | ||
|
|
54cb49c269 | ||
|
|
a7a106db51 | ||
|
|
0d080632d8 | ||
|
|
3b1249eb41 | ||
|
|
2e9d41266a | ||
|
|
b8b6a17836 | ||
|
|
d1b68b2d19 | ||
|
|
9c6ea2571c | ||
|
|
3a4c6f6c3a | ||
|
|
e77aadf64e | ||
|
|
dc6949043f | ||
|
|
d9ac568f07 | ||
|
|
c32505330c | ||
|
|
f5354ad8e8 | ||
|
|
0647c1c790 | ||
|
|
bda6fb06d8 | ||
|
|
66472e2642 | ||
|
|
887ed4c9e2 | ||
|
|
c9bea5fa06 | ||
|
|
c8941953a7 | ||
|
|
ba91df54ac | ||
|
|
8f5f5a6b69 | ||
|
|
ff1c126df5 | ||
|
|
d35b3d08c8 | ||
|
|
5dd9acd84b | ||
|
|
7703378a58 | ||
|
|
f74425e740 | ||
|
|
ee651d1f22 | ||
|
|
5128b8c6f2 | ||
|
|
80752fb7b8 | ||
|
|
0a4421504e | ||
|
|
fe18c65c5b | ||
|
|
27ad9ef273 | ||
|
|
5c6d4b89a0 | ||
|
|
20b69adc5c | ||
|
|
d7e5cf0f85 | ||
|
|
d94e1b0edf | ||
|
|
fab75315ae | ||
|
|
f3be7fac74 | ||
|
|
df20a2fe50 | ||
|
|
1a9c7011f0 | ||
|
|
3eb27ebf71 | ||
|
|
edb90d4255 | ||
|
|
c7332c647b | ||
|
|
60e0aafbbc | ||
|
|
a89991e159 | ||
|
|
437c23bbdf | ||
|
|
8730fd6fec | ||
|
|
d7caba028c | ||
|
|
9ce9cec0a8 | ||
|
|
42bc870ee3 | ||
|
|
85f4e2b294 | ||
|
|
734836b928 | ||
|
|
0b13475152 | ||
|
|
4320bf500c | ||
|
|
07ec84687e | ||
|
|
06926f8318 | ||
|
|
067fab6aba | ||
|
|
b239467901 | ||
|
|
9cb60dc7f8 | ||
|
|
8ad0e435e8 | ||
|
|
e229ba997b | ||
|
|
8f1c79684f | ||
|
|
1ecb182603 | ||
|
|
68452d30ce | ||
|
|
2c12bcb257 | ||
|
|
688f748c1e | ||
|
|
57061d7a7b | ||
|
|
c26c240570 | ||
|
|
ce1f9080ee | ||
|
|
5909a3dbb1 | ||
|
|
5ec6512c3e | ||
|
|
c9ff35db22 | ||
|
|
2dd8729326 | ||
|
|
b9e07d9400 | ||
|
|
9d3c84dbd5 | ||
|
|
d2b07abdea | ||
|
|
fff45ff60a | ||
|
|
f115e52637 | ||
|
|
3c7c3f8b3a | ||
|
|
188525acc9 | ||
|
|
90da7214c7 | ||
|
|
ac4bf93914 | ||
|
|
80eefe2043 | ||
|
|
d573c9e7f5 | ||
|
|
62f49b91ec | ||
|
|
0901306cf5 | ||
|
|
2dd47e8c70 | ||
|
|
3e89a23166 | ||
|
|
d48692bd4b | ||
|
|
9fa9fb078e | ||
|
|
e170a5cb5a | ||
|
|
a6794907e4 | ||
|
|
fed657b8fb | ||
|
|
c2ca773c27 | ||
|
|
228cad82c2 | ||
|
|
d83ef25cc6 | ||
|
|
a6439ecb5e | ||
|
|
2cbebe604c | ||
|
|
432ce2aeb1 | ||
|
|
654907ad4d | ||
|
|
8d0ad548b2 | ||
|
|
0bb5343dca | ||
|
|
97a18feba0 | ||
|
|
7200238f4f | ||
|
|
d9058c2e4e | ||
|
|
1b6e663af0 | ||
|
|
5887f400c8 | ||
|
|
5e509423a6 | ||
|
|
bb61d92f80 | ||
|
|
34431b6511 | ||
|
|
ec3ec5a2c1 | ||
|
|
d9f5a40088 | ||
|
|
713a54815b | ||
|
|
06c87bc24d | ||
|
|
75c37416d8 | ||
|
|
ad99d0d0c0 | ||
|
|
c7b3e67989 |
@@ -4,7 +4,7 @@
|
||||
# 2. To lint specific notebooks:
|
||||
# docker run -v ${PWD}:/setup/app gcr.io/python-docs-samples-tests/notebook_linter:latest notebooks/1.ipynb notebooks/2.ipynb
|
||||
|
||||
FROM python:3.10
|
||||
FROM python:3.11
|
||||
|
||||
WORKDIR setup
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ ipython
|
||||
jupyter
|
||||
nbconvert
|
||||
black==23.3.0
|
||||
pyupgrade==3.7.0
|
||||
pyupgrade==3.13.0
|
||||
isort==5.12.0
|
||||
flake8==6.0.0
|
||||
nbqa==1.7.0
|
||||
|
||||
@@ -12,4 +12,14 @@
|
||||
/prediction_featurestore_integration @googleapis/vertex-prediction-team
|
||||
/vertex_vision_model_garden/model_oss/util @weigary
|
||||
/vertex_vision_model_garden/model_oss/diffusers @weigary
|
||||
/vertex_vision_model_garden/model_oss/keras @dstnluong-google
|
||||
/vertex_vision_model_garden/model_oss/transformers @dstnluong-google
|
||||
/vertex_vision_model_garden/model_oss/pic2word @jismailyan-google
|
||||
/vertex_vision_model_garden/model_oss/open_clip @lydhr
|
||||
/vertex_vision_model_garden/model_oss/movinet @KCFindstr
|
||||
/vertex_vision_model_garden/model_oss/data_converter @KCFindstr
|
||||
/vertex_vision_model_garden/model_oss/peft @weigary
|
||||
/vertex_vision_model_garden/model_oss/lm-evaluation-harness @kathyyu-google
|
||||
/vertex_vision_model_garden/model_oss/tfvision @dstnluong-google
|
||||
/vertex_vision_model_garden/model_oss/fvlm @minwoo33park
|
||||
|
||||
|
||||
+1
-1
@@ -9,7 +9,7 @@ binarize_column_using_Pandas_on_CSV_data_op = components.load_component_from_url
|
||||
split_rows_into_subsets_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/dataset_manipulation/Split_rows_into_subsets/in_CSV/component.yaml")
|
||||
train_XGBoost_model_on_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/XGBoost/Train/component.yaml")
|
||||
xgboost_predict_on_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/XGBoost/Predict/component.yaml")
|
||||
upload_XGBoost_model_to_Google_Cloud_Vertex_AI_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Upload_XGBoost_model/component.yaml")
|
||||
upload_XGBoost_model_to_Google_Cloud_Vertex_AI_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/d5c9918850a6cc70004c4269dae066cfe2e664eb/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Upload_XGBoost_model/component.yaml")
|
||||
deploy_model_to_endpoint_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Deploy_to_endpoint/component.yaml")
|
||||
|
||||
# %% Pipeline definition
|
||||
|
||||
+1
-1
@@ -23,7 +23,7 @@ upload_PyTorch_model_archive_to_Google_Cloud_Vertex_AI_op = components.load_comp
|
||||
# XGBoost
|
||||
train_XGBoost_model_on_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/XGBoost/Train/component.yaml")
|
||||
xgboost_predict_on_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/XGBoost/Predict/component.yaml")
|
||||
upload_XGBoost_model_to_Google_Cloud_Vertex_AI_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Upload_XGBoost_model/component.yaml")
|
||||
upload_XGBoost_model_to_Google_Cloud_Vertex_AI_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/d5c9918850a6cc70004c4269dae066cfe2e664eb/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Upload_XGBoost_model/component.yaml")
|
||||
|
||||
# Scikit-learn
|
||||
#train_linear_regression_model_using_scikit_learn_from_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/1f5cf6e06409b704064b2086c0a705e4e6b4fcde/community-content/pipeline_components/ML_frameworks/Scikit_learn/Train_linear_regression_model/from_CSV/component.yaml")
|
||||
|
||||
+1
-1
@@ -8,7 +8,7 @@ fill_all_missing_values_using_Pandas_on_CSV_data_op = components.load_component_
|
||||
split_rows_into_subsets_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/dataset_manipulation/Split_rows_into_subsets/in_CSV/component.yaml")
|
||||
train_XGBoost_model_on_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/XGBoost/Train/component.yaml")
|
||||
xgboost_predict_on_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/XGBoost/Predict/component.yaml")
|
||||
upload_XGBoost_model_to_Google_Cloud_Vertex_AI_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Upload_XGBoost_model/component.yaml")
|
||||
upload_XGBoost_model_to_Google_Cloud_Vertex_AI_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/d5c9918850a6cc70004c4269dae066cfe2e664eb/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Upload_XGBoost_model/component.yaml")
|
||||
deploy_model_to_endpoint_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Deploy_to_endpoint/component.yaml")
|
||||
|
||||
# %% Pipeline definition
|
||||
|
||||
+1
-1
@@ -22,7 +22,7 @@ upload_PyTorch_model_archive_to_Google_Cloud_Vertex_AI_op = components.load_comp
|
||||
# XGBoost
|
||||
train_XGBoost_model_on_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/XGBoost/Train/component.yaml")
|
||||
xgboost_predict_on_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/XGBoost/Predict/component.yaml")
|
||||
upload_XGBoost_model_to_Google_Cloud_Vertex_AI_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/399405402d95f4a011e2d2e967c96f8508ba5688/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Upload_XGBoost_model/component.yaml")
|
||||
upload_XGBoost_model_to_Google_Cloud_Vertex_AI_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/d5c9918850a6cc70004c4269dae066cfe2e664eb/community-content/pipeline_components/google-cloud/Vertex_AI/Models/Upload_XGBoost_model/component.yaml")
|
||||
|
||||
# Scikit-learn
|
||||
train_linear_regression_model_using_scikit_learn_from_CSV_op = components.load_component_from_url("https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/1f5cf6e06409b704064b2086c0a705e4e6b4fcde/community-content/pipeline_components/ML_frameworks/Scikit_learn/Train_linear_regression_model/from_CSV/component.yaml")
|
||||
|
||||
+2
-2
@@ -64,8 +64,8 @@ implementation:
|
||||
labels["component-source"] = "github-com-ark-kun-pipeline-components"
|
||||
|
||||
# The serving container decides the model type based on the model file extension.
|
||||
# So we need to rename the mode file (e.g. /tmp/inputs/model/data) to *.pkl
|
||||
_, renamed_model_path = tempfile.mkstemp(suffix=".pkl")
|
||||
# So we need to rename the mode file (e.g. /tmp/inputs/model/data) to *.bst
|
||||
_, renamed_model_path = tempfile.mkstemp(suffix=".bst")
|
||||
shutil.copyfile(src=model_path, dst=renamed_model_path)
|
||||
|
||||
model = aiplatform.Model.upload_xgboost_model_file(
|
||||
|
||||
+2
-2
@@ -87,7 +87,7 @@ outputs:
|
||||
- {name: image_size_path, type: HeightWidth}
|
||||
implementation:
|
||||
container:
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.1
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.2
|
||||
# command is a list of strings (command-line arguments).
|
||||
# The YAML language has two syntaxes for lists and you can use either of them.
|
||||
# Here we use the "flow syntax" - comma-separated strings inside square brackets.
|
||||
@@ -109,4 +109,4 @@ implementation:
|
||||
{inputValue: l2_regularization_penalty},
|
||||
--image-size-path,
|
||||
{outputPath: image_size_path},
|
||||
]
|
||||
]
|
||||
|
||||
+1
-1
@@ -34,7 +34,7 @@ outputs:
|
||||
path for the validation data,'}
|
||||
implementation:
|
||||
container:
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.1
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.2
|
||||
# command is a list of strings (command-line arguments).
|
||||
# The YAML language has two syntaxes for lists and you can use either of them.
|
||||
# Here we use the "flow syntax" - comma-separated strings inside square brackets.
|
||||
|
||||
+1
-1
@@ -55,7 +55,7 @@ outputs:
|
||||
for the saved model,'}
|
||||
implementation:
|
||||
container:
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.1
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.2
|
||||
# command is a list of strings (command-line arguments).
|
||||
# The YAML language has two syntaxes for lists and you can use either of them.
|
||||
# Here we use the "flow syntax" - comma-separated strings inside square brackets.
|
||||
|
||||
+1
-1
@@ -20,7 +20,7 @@ outputs:
|
||||
path for the TFRecord image data}
|
||||
implementation:
|
||||
container:
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.1
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.2
|
||||
# command is a list of strings (command-line arguments).
|
||||
# The YAML language has two syntaxes for lists and you can use either of them.
|
||||
# Here we use the "flow syntax" - comma-separated strings inside square brackets.
|
||||
|
||||
+1
-1
@@ -22,7 +22,7 @@ outputs:
|
||||
path for the TFRecord image data}
|
||||
implementation:
|
||||
container:
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.1
|
||||
image: us-docker.pkg.dev/vertex-ai/ready-to-go-image-classification/image-components:v0.2
|
||||
# command is a list of strings (command-line arguments).
|
||||
# The YAML language has two syntaxes for lists and you can use either of them.
|
||||
# Here we use the "flow syntax" - comma-separated strings inside square brackets.
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
google-cloud-bigquery==2.20.0
|
||||
tensorflow==2.7.2
|
||||
pillow==9.0.1
|
||||
pillow==10.0.1
|
||||
tf-agents==0.8.0
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
google-cloud-pubsub==2.5.0
|
||||
pillow==9.0.1
|
||||
pillow==10.0.1
|
||||
tf-agents==0.8.0
|
||||
tensorflow==2.7.2
|
||||
|
||||
+1
-1
@@ -1,5 +1,5 @@
|
||||
dataclasses==0.6
|
||||
google-cloud-aiplatform==1.8.1
|
||||
tensorflow==2.7.2
|
||||
pillow==9.0.1
|
||||
pillow==10.0.1
|
||||
tf-agents==0.8.0
|
||||
@@ -0,0 +1,623 @@
|
||||
"""Library with functions to use for data conversion."""
|
||||
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
from typing import Any, Callable, Dict, Iterable, List, Optional, Sequence, Tuple, Union
|
||||
import uuid
|
||||
|
||||
from absl import logging
|
||||
import apache_beam as beam
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import PIL
|
||||
from PIL import Image
|
||||
import tensorflow as tf
|
||||
import yaml
|
||||
|
||||
from util import constants
|
||||
from util import fileutils
|
||||
from apache_beam.options import pipeline_options
|
||||
|
||||
REFORMATTED_CSV_SUFFIX = '-reformatted.csv'
|
||||
|
||||
LABEL_MAP_NAME = 'label_map.yaml'
|
||||
|
||||
_SPLIT_RATIO_ERROR_THRESHOLD = 1e-5
|
||||
# Internal constant. Only for distinguishing rows without ML use.
|
||||
ML_USE_UNASSIGNED = 'unassigned'
|
||||
ALL_ML_USES = (
|
||||
constants.ML_USE_TRAINING,
|
||||
constants.ML_USE_VALIDATION,
|
||||
constants.ML_USE_TEST,
|
||||
ML_USE_UNASSIGNED,
|
||||
)
|
||||
COLUMN_NAME_ML_USE = 'ml_use'
|
||||
COLUMN_NAME_GCS_FILE_PATH = 'gcs_file_path'
|
||||
COLUMN_NAME_LABEL = 'label'
|
||||
COLUMN_NAME_START_SEC = 'start_sec'
|
||||
COLUMN_NAME_END_SEC = 'end_sec'
|
||||
# Output filenames
|
||||
TRAIN_TFRECORD_NAME = 'train.tfrecord'
|
||||
VALIDATION_TFRECORD_NAME = 'val.tfrecord'
|
||||
TEST_TFRECORD_NAME = 'test.tfrecord'
|
||||
# Jsonl keys
|
||||
JSON_GCS_URI_KEY = 'imageGcsUri'
|
||||
JSON_RESOURCE_LABEL_KEY = 'dataItemResourceLabels'
|
||||
JSON_ML_USE_KEY = 'aiplatform.googleapis.com/ml_use'
|
||||
# I/O parameters
|
||||
READ_CHUNK_SIZE = 1024 * 1024 * 1024 # 1GB
|
||||
|
||||
|
||||
class WriteToTFRecord(beam.DoFn):
|
||||
"""DoFn to write TF examples to sharded TF record files."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
output_prefix: str,
|
||||
num_shards: int,
|
||||
convert_fn: Callable[[Dict[str, Any]], tf.train.Example],
|
||||
):
|
||||
self.output_prefix = output_prefix
|
||||
self.num_shards = num_shards
|
||||
self.writer: list[tf.io.TFRecordWriter] = []
|
||||
self.sharded_files: list[str] = []
|
||||
self.convert_fn = convert_fn
|
||||
self.success_counter = beam.metrics.Metrics.counter(
|
||||
self.__class__.__name__, 'Success'
|
||||
)
|
||||
self.failure_counter = beam.metrics.Metrics.counter(
|
||||
self.__class__.__name__, 'Failure'
|
||||
)
|
||||
|
||||
def start_bundle(self):
|
||||
logging.info('Start writing TF Record to %s.', self.output_prefix)
|
||||
unique_str = uuid.uuid4().hex
|
||||
for i in range(self.num_shards):
|
||||
uri = f'{self.output_prefix}-{i}-{unique_str}'
|
||||
self.sharded_files.append(uri)
|
||||
self.writer.append(tf.io.TFRecordWriter(uri))
|
||||
|
||||
def process(self, data: Dict[str, Any]) -> Iterable[Tuple[int, str]]:
|
||||
try:
|
||||
example = self.convert_fn(data)
|
||||
data = example.SerializeToString()
|
||||
idx = hash(data) % self.num_shards
|
||||
self.writer[idx].write(data)
|
||||
self.success_counter.inc()
|
||||
yield (idx, self.sharded_files[idx])
|
||||
# pylint: disable-next=broad-exception-caught
|
||||
except Exception as err:
|
||||
logging.error('Failed to process %s', data)
|
||||
logging.exception(err)
|
||||
self.failure_counter.inc()
|
||||
|
||||
def finish_bundle(self):
|
||||
logging.info('Finish writing TF Record to %s.', self.output_prefix)
|
||||
for writer in self.writer:
|
||||
writer.close()
|
||||
self.writer = []
|
||||
|
||||
|
||||
def convert_to_feature(
|
||||
value: Union[List[Union[int, float, bytes]], int, float, bytes],
|
||||
value_type: Optional[str] = None,
|
||||
) -> tf.train.Feature:
|
||||
"""Converts the given python object to a tf.train.Feature.
|
||||
|
||||
This is copied from tensorflow_models/official/vision/data/tfrecord_lib.py.
|
||||
|
||||
Args:
|
||||
value: int, float, bytes or a list of them.
|
||||
value_type: optional, if specified, forces the feature to be of the given
|
||||
type. Otherwise, type is inferred automatically. Can be one of ['bytes',
|
||||
'int64', 'float', 'bytes_list', 'int64_list', 'float_list']
|
||||
|
||||
Returns:
|
||||
feature: A tf.train.Feature object.
|
||||
"""
|
||||
|
||||
if value_type is None:
|
||||
element = value[0] if isinstance(value, list) else value
|
||||
|
||||
if isinstance(element, bytes):
|
||||
value_type = 'bytes'
|
||||
|
||||
elif isinstance(element, (int, np.integer)):
|
||||
value_type = 'int64'
|
||||
|
||||
elif isinstance(element, (float, np.floating)):
|
||||
value_type = 'float'
|
||||
|
||||
else:
|
||||
raise ValueError(
|
||||
'Cannot convert type {} to feature'.format(type(element))
|
||||
)
|
||||
|
||||
if isinstance(value, list):
|
||||
value_type = value_type + '_list'
|
||||
|
||||
if value_type == 'int64':
|
||||
return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))
|
||||
|
||||
elif value_type == 'int64_list':
|
||||
value = np.asarray(value).astype(np.int64).reshape(-1)
|
||||
return tf.train.Feature(int64_list=tf.train.Int64List(value=value))
|
||||
|
||||
elif value_type == 'float':
|
||||
return tf.train.Feature(float_list=tf.train.FloatList(value=[value]))
|
||||
|
||||
elif value_type == 'float_list':
|
||||
value = np.asarray(value).astype(np.float32).reshape(-1)
|
||||
return tf.train.Feature(float_list=tf.train.FloatList(value=value))
|
||||
|
||||
elif value_type == 'bytes':
|
||||
return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))
|
||||
|
||||
elif value_type == 'bytes_list':
|
||||
return tf.train.Feature(bytes_list=tf.train.BytesList(value=value))
|
||||
|
||||
else:
|
||||
raise ValueError('Unknown value_type parameter - {}'.format(value_type))
|
||||
|
||||
|
||||
def convert_to_string_feature(
|
||||
value: str, encoding: str = 'utf-8'
|
||||
) -> tf.train.Feature:
|
||||
"""Returns a bytes_list from an encoded string."""
|
||||
return convert_to_feature(value.encode(encoding))
|
||||
|
||||
|
||||
def convert_to_list_string_feature(
|
||||
lst: list[str], encoding: str = 'utf-8'
|
||||
) -> tf.train.Feature:
|
||||
"""Returns a bytes_list from a list of encoded strings."""
|
||||
return convert_to_feature([value.encode(encoding) for value in lst])
|
||||
|
||||
|
||||
def create_ml_use_array_with_split(
|
||||
total_size: int,
|
||||
split_ratio: Sequence[float],
|
||||
) -> list[str]:
|
||||
"""Create randomized list of 'training', 'validation', 'test'.
|
||||
|
||||
The list of will be of length total_size with ratios according to train_size,
|
||||
validation_size, and test_size.
|
||||
|
||||
Args:
|
||||
total_size: Length of sequence to return
|
||||
split_ratio: Proportions to split into 'training', 'validation', and 'test'
|
||||
|
||||
Returns:
|
||||
List containing 'training', 'validation', and 'test'
|
||||
"""
|
||||
train_size, validation_size, _ = split_ratio
|
||||
num_train = round(train_size * total_size)
|
||||
num_validation = round(validation_size * total_size)
|
||||
num_test = total_size - num_train - num_validation
|
||||
ml_use_row = (
|
||||
[constants.ML_USE_TRAINING] * num_train
|
||||
+ [constants.ML_USE_VALIDATION] * num_validation
|
||||
+ [constants.ML_USE_TEST] * num_test
|
||||
)
|
||||
random.shuffle(ml_use_row)
|
||||
return ml_use_row
|
||||
|
||||
|
||||
def format_ml_use_column(df: pd.DataFrame):
|
||||
df[COLUMN_NAME_ML_USE].replace(
|
||||
# We need to support non-standard ML uses other than documented ones,
|
||||
# since they are used by some existing datasets.
|
||||
[r'(?i)^train(ing)?$', r'(?i)^test$', r'(?i)^validat(ion|e)$'],
|
||||
[
|
||||
constants.ML_USE_TRAINING,
|
||||
constants.ML_USE_TEST,
|
||||
constants.ML_USE_VALIDATION,
|
||||
],
|
||||
inplace=True,
|
||||
regex=True,
|
||||
)
|
||||
|
||||
|
||||
def insert_missing_ml_use(df: pd.DataFrame) -> None:
|
||||
"""For every row that does not have ml_use as the first column, insert a column containing 'unassigned' to the front.
|
||||
|
||||
Args:
|
||||
df: The DataFrame to process. The first column should be 'ml_use'.
|
||||
"""
|
||||
df[COLUMN_NAME_ML_USE].fillna(ML_USE_UNASSIGNED, inplace=True)
|
||||
rows_to_fill = ~df[COLUMN_NAME_ML_USE].isin(ALL_ML_USES)
|
||||
df.loc[rows_to_fill] = df[rows_to_fill].shift(
|
||||
axis=1, fill_value=ML_USE_UNASSIGNED
|
||||
)
|
||||
|
||||
|
||||
def replace_unassigned_ml_use(
|
||||
ml_uses: List[str],
|
||||
split_ratio: Sequence[float],
|
||||
):
|
||||
"""Replace `unassigned` in ml_uses with `training`, `validation`, and `test` with ratios according to split_ratio.
|
||||
|
||||
Args:
|
||||
ml_uses: List of ml_use string values.
|
||||
split_ratio: Proportions to split into `training`, `validation`, and `test`.
|
||||
"""
|
||||
unassigned_indices = [
|
||||
i for i, ml_use in enumerate(ml_uses) if ml_use == ML_USE_UNASSIGNED
|
||||
]
|
||||
ml_use_arr = create_ml_use_array_with_split(
|
||||
len(unassigned_indices), split_ratio
|
||||
)
|
||||
for unassigned_index, ml_use in zip(unassigned_indices, ml_use_arr):
|
||||
ml_uses[unassigned_index] = ml_use
|
||||
|
||||
|
||||
def merge_seq_into_dicts(
|
||||
key: str, values: Sequence[Any], dicts: Sequence[Dict[Any, Any]]
|
||||
):
|
||||
"""Merges a list of values into a list of dicts, inserted with the given key.
|
||||
|
||||
Args:
|
||||
key: Key to insert or overwrite in the dictionary.
|
||||
values: A list of values to insert.
|
||||
dicts: A list of dictionaries. Each value will be inserted into the
|
||||
corresponding dictionary. The original value will be overwritten if the
|
||||
key already existed.
|
||||
|
||||
Raises:
|
||||
ValueError: The values and dicts have different lengths.
|
||||
"""
|
||||
if len(values) != len(dicts):
|
||||
raise ValueError(
|
||||
f'Length of values and dicts must match, got {len(values)} and'
|
||||
f' {len(dicts)}'
|
||||
)
|
||||
for val, d in zip(values, dicts):
|
||||
d[key] = val
|
||||
|
||||
|
||||
def drop_invalid_rows(df: pd.DataFrame) -> int:
|
||||
"""Drops DataFrame rows missing the gcs_file_path column or the label column.
|
||||
|
||||
Args:
|
||||
df: The DataFrame to process in place.
|
||||
|
||||
Returns:
|
||||
The number of rows dropped.
|
||||
"""
|
||||
original_rows = df.shape[0]
|
||||
df.dropna(subset=[COLUMN_NAME_GCS_FILE_PATH, COLUMN_NAME_LABEL], inplace=True)
|
||||
dropped_num = original_rows - df.shape[0]
|
||||
if dropped_num > 0:
|
||||
df.reset_index(drop=True, inplace=True)
|
||||
return dropped_num
|
||||
|
||||
|
||||
def check_split_ratio(split_ratio: Sequence[float]):
|
||||
"""Checks if the give split ratio is valid.
|
||||
|
||||
Args:
|
||||
split_ratio: Proportions to split into 'training', 'validation', and 'test'
|
||||
|
||||
Raises:
|
||||
ValueError: Must have valid entries, correct length, and sum to 1.
|
||||
"""
|
||||
if len(split_ratio) != 3:
|
||||
raise ValueError('split_ratio must contain exactly 3 values.')
|
||||
if abs(sum(split_ratio) - 1) > _SPLIT_RATIO_ERROR_THRESHOLD:
|
||||
raise ValueError('split_ratio must sum to 1.')
|
||||
if not all([0 <= val <= 1 for val in split_ratio]):
|
||||
raise ValueError('Entries of split_ratio must be in the range [0, 1].')
|
||||
|
||||
|
||||
def check_num_shard(num_shard: Sequence[int]):
|
||||
"""Checks if the number of shards is valid.
|
||||
|
||||
Args:
|
||||
num_shard: The number of shards for each tfrecord.
|
||||
|
||||
Raises:
|
||||
ValueError: Must have valid entries and correct length.
|
||||
"""
|
||||
if len(num_shard) != 3:
|
||||
raise ValueError('num_shard must contain exactly 3 values.')
|
||||
if not all([val >= 1 for val in num_shard]):
|
||||
raise ValueError('Shards must be at least 1.')
|
||||
|
||||
|
||||
def create_label_map_yaml(meta_data_path: str, output_dir: str) -> None:
|
||||
"""Generate label_map.yaml from meta_data.yaml.
|
||||
|
||||
Args:
|
||||
meta_data_path: Path to a meta_data.yaml file.
|
||||
output_dir: Directory to output label_map.yaml.
|
||||
"""
|
||||
tf.io.gfile.copy(
|
||||
meta_data_path, os.path.join(output_dir, LABEL_MAP_NAME), overwrite=True
|
||||
)
|
||||
|
||||
|
||||
def reformat_bbox(
|
||||
bbox: Sequence[int], img_width: int, img_height: int
|
||||
) -> Tuple[float, float, float, float]:
|
||||
"""Converts XYWH unnormalized bounding box with to a normalized XYXY bounding box.
|
||||
|
||||
Args:
|
||||
bbox: Relative bounding box with unnormalized coordinates as [x, y, width,
|
||||
height].
|
||||
img_width: Image's pixel width.
|
||||
img_height: Image's pixel height.
|
||||
|
||||
Returns:
|
||||
Absolute bounding box with normalized coordinates as
|
||||
[xmin, ymin, xmax, ymax].
|
||||
"""
|
||||
x, y, width, height = bbox
|
||||
xmin = x / img_width
|
||||
ymin = y / img_height
|
||||
xmax = (x + width) / img_width
|
||||
ymax = (y + height) / img_height
|
||||
return xmin, ymin, xmax, ymax
|
||||
|
||||
|
||||
def encode_image(
|
||||
filepath: str,
|
||||
output_shape: Optional[Sequence[int]] = None,
|
||||
image_format: str = 'png',
|
||||
) -> Tuple[bytes, Sequence[int]]:
|
||||
"""Encodes an image at the given path.
|
||||
|
||||
Args:
|
||||
filepath: Path to the image.
|
||||
output_shape: The output shape of the image, (height, width).
|
||||
image_format: The format of the output image.
|
||||
|
||||
Returns:
|
||||
The encoded image data in bytes and the shape of the image, (height, width).
|
||||
|
||||
Raises:
|
||||
IOError: The image file is corrupt.
|
||||
"""
|
||||
filepath = fileutils.force_gcs_fuse_path(filepath)
|
||||
with open(filepath, 'rb') as f:
|
||||
# If an output_shape is specified, resize the image and set data to the new
|
||||
# bytes.
|
||||
try:
|
||||
img = Image.open(f)
|
||||
except PIL.UnidentifiedImageError as e:
|
||||
raise IOError(f'Failed to open {filepath}') from e
|
||||
|
||||
try:
|
||||
if output_shape is not None:
|
||||
rgb_img = img.resize((output_shape[1], output_shape[0])).convert('RGB')
|
||||
else:
|
||||
rgb_img = img.convert('RGB')
|
||||
rgb_img = np.array(rgb_img)
|
||||
|
||||
_, data = cv2.imencode(f'.{image_format}', rgb_img)
|
||||
data = data.tobytes()
|
||||
return data, rgb_img.shape
|
||||
except cv2.error as e:
|
||||
raise IOError(f'Failed to encode {filepath}') from e
|
||||
finally:
|
||||
img.close()
|
||||
|
||||
|
||||
def encode_video(
|
||||
filepath: str,
|
||||
start_sec: float,
|
||||
end_sec: float,
|
||||
output_fps: int = 5,
|
||||
output_shape: Optional[Sequence[int]] = None,
|
||||
image_format: str = 'jpg',
|
||||
) -> Sequence[bytes]:
|
||||
"""Encodes a video clip at the given path with start and end timestamps.
|
||||
|
||||
Args:
|
||||
filepath: Path to the video.
|
||||
start_sec: Start timestamp of the video clip in seconds.
|
||||
end_sec: End timestamp of the video clip in seconds.
|
||||
output_fps: The output frame rate per second.
|
||||
output_shape: The output shape of each frame, (height, width).
|
||||
image_format: The format of the encoded frames.
|
||||
|
||||
Returns:
|
||||
A list of the encoded frames data in bytes.
|
||||
|
||||
Raises:
|
||||
IOError if the video file is corrupt.
|
||||
"""
|
||||
filepath = fileutils.force_gcs_fuse_path(filepath)
|
||||
video = None
|
||||
|
||||
try:
|
||||
video = cv2.VideoCapture(filepath)
|
||||
frames = []
|
||||
frame_interval = 1 / output_fps
|
||||
total_frames = video.get(cv2.CAP_PROP_FRAME_COUNT)
|
||||
original_fps = video.get(cv2.CAP_PROP_FPS)
|
||||
if not original_fps:
|
||||
# 0 or None indicates the video is invalid
|
||||
raise IOError(f'Failed to load {filepath}')
|
||||
video_length = total_frames / original_fps
|
||||
start_sec = max(start_sec, 0)
|
||||
end_sec = min(end_sec, video_length)
|
||||
for t in np.arange(start_sec, end_sec, frame_interval):
|
||||
frame_idx = min(total_frames - 1, round(t * original_fps))
|
||||
video.set(cv2.CAP_PROP_POS_FRAMES, frame_idx)
|
||||
ret, frame = video.read()
|
||||
if not ret:
|
||||
raise IOError(f'Failed to load {filepath} at frame {frame_idx}')
|
||||
if output_shape is not None:
|
||||
frame = cv2.resize(frame, (output_shape[1], output_shape[0]))
|
||||
_, data = cv2.imencode(f'.{image_format}', frame)
|
||||
frames.append(data.tobytes())
|
||||
except cv2.error as e:
|
||||
raise IOError(f'Failed to load {filepath}') from e
|
||||
finally:
|
||||
if video:
|
||||
video.release()
|
||||
return frames
|
||||
|
||||
|
||||
def create_label_map(
|
||||
labels: Sequence[str],
|
||||
) -> Tuple[Sequence[int], Dict[int, str]]:
|
||||
"""Creates a label map from a sequence of label strings.
|
||||
|
||||
Args:
|
||||
labels: The sequence of labels to create label map from. Must not contain
|
||||
invalid values, which means data without labels should be filtered first.
|
||||
|
||||
Returns:
|
||||
The integer labels and the mapping from integers to the original strings.
|
||||
"""
|
||||
inverse_label_map: Dict[str, int] = dict()
|
||||
num_labels = 0
|
||||
for label in labels:
|
||||
if label not in inverse_label_map:
|
||||
num_labels += 1
|
||||
inverse_label_map[label] = num_labels
|
||||
int_labels = [inverse_label_map[label] for label in labels]
|
||||
label_map = {value: key for key, value in inverse_label_map.items()}
|
||||
return int_labels, label_map
|
||||
|
||||
|
||||
def write_label_map(output_file: str, label_map: Dict[int, str]) -> None:
|
||||
"""Writes a label map to the output file, which can be a GCS uri."""
|
||||
with tf.io.gfile.GFile(output_file, 'w') as f:
|
||||
yaml.dump({'label_map': label_map}, f)
|
||||
|
||||
|
||||
def detectron_json_to_image_rows(input_json: str) -> list[Dict[str, Any]]:
|
||||
"""Converts a Detectron JSON file to a list of image rows.
|
||||
|
||||
Args:
|
||||
input_json: A path to a Detectron JSON or JSONL file.
|
||||
|
||||
Returns:
|
||||
A list of dictionaries, where each dictionary contains Detectron format
|
||||
entry.
|
||||
|
||||
Raises:
|
||||
ValueError: If the input JSON is invalid.
|
||||
"""
|
||||
|
||||
image_rows = []
|
||||
with tf.io.gfile.GFile(input_json, 'r') as f:
|
||||
for line in f:
|
||||
json_data = json.loads(line)
|
||||
if isinstance(json_data, dict):
|
||||
image_rows.append(json_data)
|
||||
elif isinstance(json_data, list):
|
||||
image_rows.extend(json_data)
|
||||
else:
|
||||
raise ValueError(
|
||||
'The input JSON is invalid. Dict or list is expected, but got '
|
||||
f'{type(json_data)}.'
|
||||
)
|
||||
return image_rows
|
||||
|
||||
|
||||
def coco_json_to_image_rows(
|
||||
input_json: str,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Converts a COCO JSON file to a list of image rows.
|
||||
|
||||
Args:
|
||||
input_json: A path to a COCO JSON or JSONL file.
|
||||
|
||||
Returns:
|
||||
A list of dictionaries, where each dictionary contains COCO format entry.
|
||||
|
||||
Raises:
|
||||
ValueError: If the input JSON is invalid.
|
||||
"""
|
||||
|
||||
with tf.io.gfile.GFile(input_json, 'r') as f:
|
||||
coco_json = json.load(f)
|
||||
if 'annotations' not in coco_json:
|
||||
raise ValueError('"annotations" is not in the dataset.')
|
||||
if 'images' not in coco_json:
|
||||
raise ValueError('"images" is not in the dataset.')
|
||||
|
||||
images = coco_json['images']
|
||||
return images
|
||||
|
||||
|
||||
def partition_by_ml_use(element: Dict[str, Any], num_partitions: int) -> int:
|
||||
"""Beam partition function to split data by ml_use."""
|
||||
del num_partitions
|
||||
try:
|
||||
partition = ALL_ML_USES.index(element[COLUMN_NAME_ML_USE])
|
||||
except Exception as e:
|
||||
raise ValueError(f'Invalid ML use: {element[COLUMN_NAME_ML_USE]}') from e
|
||||
return partition
|
||||
|
||||
|
||||
def run_beam_pipeline(pipeline: Any) -> None:
|
||||
"""Runs a beam pipeline. Works in both internal and docker environment."""
|
||||
options = pipeline_options.PipelineOptions([
|
||||
'--runner=FlinkRunner',
|
||||
'--faster_copy',
|
||||
'--max_parallelism', '8',
|
||||
])
|
||||
p = beam.Pipeline(options=options)
|
||||
pipeline(p)
|
||||
result = p.run()
|
||||
result.wait_until_finish()
|
||||
for counter in result.metrics().query()['counters']:
|
||||
logging.info('%s counter: %s.', counter.key.metric.name, counter)
|
||||
logging.info('Completing beam pipeline.')
|
||||
|
||||
|
||||
def beam_convert_tfexamples(
|
||||
root: beam.Pipeline,
|
||||
data_list: Sequence[Dict[str, Any]],
|
||||
convert_fn: Callable[[Dict[str, Any]], tf.train.Example],
|
||||
output_dir: str,
|
||||
num_shards: Sequence[int],
|
||||
) -> None:
|
||||
"""Constructs beam pipelines to convert train, val, test TF Examples."""
|
||||
names = [TRAIN_TFRECORD_NAME, VALIDATION_TFRECORD_NAME, TEST_TFRECORD_NAME]
|
||||
split_data = (
|
||||
root
|
||||
| 'Create PCollection' >> beam.Create(data_list)
|
||||
| 'Data split' >> beam.Partition(partition_by_ml_use, 3)
|
||||
)
|
||||
for i in range(3):
|
||||
ml_use: str = ALL_ML_USES[i]
|
||||
num_shard = num_shards[i]
|
||||
output_prefix = os.path.join(output_dir, names[i])
|
||||
_ = (
|
||||
split_data[i]
|
||||
| f'Convert {ml_use} TF Examples'
|
||||
>> beam.ParDo(WriteToTFRecord(output_prefix, num_shard, convert_fn))
|
||||
| f'Group {ml_use} TF Record files' >> beam.GroupBy(lambda x: x[0])
|
||||
| f'Merge {ml_use} TF Record files'
|
||||
>> beam.Map(merge_tfrecords_func(output_prefix, num_shard))
|
||||
)
|
||||
|
||||
|
||||
def merge_tfrecords_func(output_prefix: str, num_shard: int) -> ...:
|
||||
"""Returns a function to merge sharded worker output into expected shards."""
|
||||
output_prefix = fileutils.force_gcs_fuse_path(output_prefix)
|
||||
|
||||
def merge_tfrecords(worker_output: Tuple[int, Sequence[Tuple[int, str]]]):
|
||||
idx = worker_output[0]
|
||||
files: Sequence[str] = np.unique([x[1] for x in worker_output[1]])
|
||||
output_file = f'{output_prefix}-{idx:05d}-of-{num_shard:05d}'
|
||||
with open(output_file, 'wb') as f:
|
||||
for file in files:
|
||||
logging.info('Merging %s.', file)
|
||||
file = fileutils.force_gcs_fuse_path(file)
|
||||
with open(file, 'rb') as fin:
|
||||
while True:
|
||||
data = fin.read(READ_CHUNK_SIZE)
|
||||
if not data:
|
||||
break
|
||||
f.write(data)
|
||||
os.remove(file)
|
||||
|
||||
return merge_tfrecords
|
||||
+111
@@ -0,0 +1,111 @@
|
||||
r"""Converts COCO labels as yamls for model garden playground (IOD).
|
||||
"""
|
||||
|
||||
import os
|
||||
import urllib.request
|
||||
|
||||
from absl import app
|
||||
from absl import flags
|
||||
import tensorflow as tf
|
||||
import yaml
|
||||
|
||||
from object_detection.utils import label_map_util
|
||||
|
||||
_CONVERT_LABEL_TYPE_COCO_80 = 'coco_80'
|
||||
_CONVERT_LABEL_TYPE_COCO_91 = 'coco_91'
|
||||
|
||||
_CONVERT_LABEL_TYPE = flags.DEFINE_enum(
|
||||
'convert_label_type',
|
||||
None,
|
||||
[
|
||||
_CONVERT_LABEL_TYPE_COCO_80,
|
||||
_CONVERT_LABEL_TYPE_COCO_91,
|
||||
],
|
||||
'Different types of label type conversion.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_TEMPORARY_PATH = flags.DEFINE_string(
|
||||
'temporary_path',
|
||||
None,
|
||||
'The tempory path.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_OUTPUT_YAML_FILEPATH = flags.DEFINE_string(
|
||||
'output_yaml_filepath',
|
||||
None,
|
||||
'The output yaml filepath.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
|
||||
def convert_coco_label_map_91(
|
||||
output_yaml_filepath: str,
|
||||
) -> None:
|
||||
"""Converts coco label map 91."""
|
||||
input_proto_filepath = 'https://raw.githubusercontent.com/tensorflow/models/master/research/object_detection/data/mscoco_label_map.pbtxt'
|
||||
local_input_proto_filepath = os.path.join(
|
||||
_TEMPORARY_PATH.value, 'mscoco_label_map.pbtxt'
|
||||
)
|
||||
with open(local_input_proto_filepath, 'w') as writer:
|
||||
contents = (
|
||||
urllib.request.urlopen(input_proto_filepath).read().decode('utf-8')
|
||||
)
|
||||
writer.write(contents)
|
||||
|
||||
label_map = label_map_util.load_labelmap(local_input_proto_filepath)
|
||||
label_map_dict = label_map_util.get_label_map_dict(
|
||||
label_map, use_display_name=True
|
||||
)
|
||||
swapped_label_map_dict = {v: k for k, v in label_map_dict.items()}
|
||||
print(swapped_label_map_dict)
|
||||
|
||||
# Saves new label maps as yamls.
|
||||
with tf.io.gfile.GFile(output_yaml_filepath, 'w') as writer:
|
||||
writer.write(yaml.dump(swapped_label_map_dict))
|
||||
|
||||
|
||||
def convert_coco_label_map_80(
|
||||
output_yaml_filepath: str,
|
||||
) -> None:
|
||||
"""Converts coco label map 80."""
|
||||
# Loads label maps from texts.
|
||||
input_text_filepath = 'https://gist.githubusercontent.com/AruniRC/7b3dadd004da04c80198557db5da4bda/raw/2f10965ace1e36c4a9dca76ead19b744f5eb7e88/ms_coco_classnames.txt'
|
||||
local_input_text_filepath = os.path.join(
|
||||
_TEMPORARY_PATH.value, 'ms_coco_classnames.txt'
|
||||
)
|
||||
with open(local_input_text_filepath, 'w') as writer:
|
||||
contents = (
|
||||
urllib.request.urlopen(input_text_filepath).read().decode('utf-8')
|
||||
)
|
||||
writer.write(contents)
|
||||
with open(local_input_text_filepath, 'r') as file:
|
||||
content = file.read()
|
||||
label_map = yaml.safe_load(content)
|
||||
|
||||
# Removes background in label maps.
|
||||
new_label_map = {}
|
||||
for k, v in label_map.items():
|
||||
if k == 0:
|
||||
continue
|
||||
new_label_map[k - 1] = v
|
||||
print(new_label_map)
|
||||
# Saves new label maps as yamls.
|
||||
with tf.io.gfile.GFile(output_yaml_filepath, 'w') as writer:
|
||||
writer.write(yaml.dump(new_label_map))
|
||||
|
||||
|
||||
def main(_) -> None:
|
||||
if _CONVERT_LABEL_TYPE.value == _CONVERT_LABEL_TYPE_COCO_80:
|
||||
convert_coco_label_map_80(_OUTPUT_YAML_FILEPATH.value)
|
||||
elif _CONVERT_LABEL_TYPE.value == _CONVERT_LABEL_TYPE_COCO_91:
|
||||
convert_coco_label_map_91(
|
||||
_OUTPUT_YAML_FILEPATH.value,
|
||||
)
|
||||
else:
|
||||
print('Not supported convert label type: ', _CONVERT_LABEL_TYPE.value)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(main)
|
||||
+86
@@ -0,0 +1,86 @@
|
||||
r"""Converts ImageNet label texts as yamls for model garden playground.
|
||||
|
||||
# ImageNet1K will have label maps with background.
|
||||
"""
|
||||
|
||||
import urllib.request
|
||||
from absl import app
|
||||
from absl import flags
|
||||
import tensorflow as tf
|
||||
import yaml
|
||||
|
||||
|
||||
_INPUT_TEXT_FILEPATH = flags.DEFINE_string(
|
||||
'input_text_filepath',
|
||||
None,
|
||||
'The input text filepath.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_ADD_BACKGROUND_LABEL = flags.DEFINE_boolean(
|
||||
'add_background_label',
|
||||
None,
|
||||
'Whether or not add background labels.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_ADD_IDS = flags.DEFINE_boolean(
|
||||
'add_ids',
|
||||
None,
|
||||
'Whether or not add ids.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_OUTPUT_YAML_FILEPATH = flags.DEFINE_string(
|
||||
'output_yaml_filepath',
|
||||
None,
|
||||
'The output yaml filepath.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
|
||||
def convert_imagenet_label_map_from_text_to_yaml(
|
||||
input_text_filepath: str,
|
||||
add_background_label: bool,
|
||||
add_ids: bool,
|
||||
output_yaml_filepath: str,
|
||||
) -> None:
|
||||
"""Converts imagenet label map from text to yamls."""
|
||||
label_map = {}
|
||||
|
||||
# Shifts all keys by 1, and add 0 as 'background'.
|
||||
if add_background_label:
|
||||
label_map = yaml.safe_load(
|
||||
urllib.request.urlopen(input_text_filepath).read()
|
||||
)
|
||||
new_label_map = {}
|
||||
for key, value in label_map.items():
|
||||
new_label_map[key + 1] = value
|
||||
new_label_map[0] = 'background'
|
||||
label_map = new_label_map
|
||||
|
||||
# Adds maps from id to each line.
|
||||
if add_ids:
|
||||
lines = urllib.request.urlopen(input_text_filepath).readlines()
|
||||
current_id = 0
|
||||
for line in lines:
|
||||
label_map[current_id] = line.decode('ascii').strip()
|
||||
print(label_map[current_id])
|
||||
current_id += 1
|
||||
|
||||
# Saves new label maps as yamls.
|
||||
with tf.io.gfile.GFile(output_yaml_filepath, 'w') as writer:
|
||||
writer.write(yaml.dump(label_map))
|
||||
|
||||
|
||||
def main(_) -> None:
|
||||
convert_imagenet_label_map_from_text_to_yaml(
|
||||
_INPUT_TEXT_FILEPATH.value,
|
||||
_ADD_BACKGROUND_LABEL.value,
|
||||
_ADD_IDS.value,
|
||||
_OUTPUT_YAML_FILEPATH.value,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(main)
|
||||
+199
@@ -0,0 +1,199 @@
|
||||
"""Converts ICN CSV/JSONL files to TFRecord with apache beam."""
|
||||
|
||||
import json
|
||||
from os import path
|
||||
from typing import Any, Dict, Sequence, Union, cast
|
||||
|
||||
from absl import logging
|
||||
import apache_beam as beam
|
||||
import pandas as pd
|
||||
import tensorflow as tf
|
||||
|
||||
from data_converter import common_lib
|
||||
|
||||
|
||||
_COLUMN_NAMES = [
|
||||
common_lib.COLUMN_NAME_ML_USE,
|
||||
common_lib.COLUMN_NAME_GCS_FILE_PATH,
|
||||
common_lib.COLUMN_NAME_LABEL,
|
||||
]
|
||||
_JSON_GCS_URI_KEY = 'imageGcsUri'
|
||||
_JSON_CLASS_ANNOTATION_KEY = 'classificationAnnotation'
|
||||
_JSON_RESOURCE_LABEL_KEY = 'dataItemResourceLabels'
|
||||
_JSON_CLASS_NAME_KEY = 'displayName'
|
||||
_JSON_ML_USE_KEY = 'aiplatform.googleapis.com/ml_use'
|
||||
|
||||
|
||||
def build_tf_example(element: Dict[str, Union[str, int]]) -> tf.train.Example:
|
||||
"""Builds a TF Example from an image uri and label.
|
||||
|
||||
Args:
|
||||
element: A dict with the keys gcs_file_path and label.
|
||||
|
||||
Returns:
|
||||
The created TF Example.
|
||||
"""
|
||||
image_uri = cast(str, element[common_lib.COLUMN_NAME_GCS_FILE_PATH])
|
||||
label = cast(int, element[common_lib.COLUMN_NAME_LABEL])
|
||||
image_bytes, shape = common_lib.encode_image(image_uri, image_format='jpeg')
|
||||
features = tf.train.Features(
|
||||
feature={
|
||||
'image/encoded': common_lib.convert_to_feature(image_bytes),
|
||||
'image/format': common_lib.convert_to_string_feature('jpeg'),
|
||||
'image/height': common_lib.convert_to_feature(shape[0]),
|
||||
'image/width': common_lib.convert_to_feature(shape[1]),
|
||||
'image/class/label': common_lib.convert_to_feature(label),
|
||||
},
|
||||
)
|
||||
return tf.train.Example(features=features)
|
||||
|
||||
|
||||
def _run_convert_pipeline(
|
||||
output_dir: str, df: pd.DataFrame, num_shards: Sequence[int]
|
||||
) -> None:
|
||||
"""Starts a Beam pipeline to write DataFrame as TF Records.
|
||||
|
||||
Args:
|
||||
output_dir: TF Records output directory.
|
||||
df: DataFrame to convert from.
|
||||
num_shards: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
images_list = df.to_dict('records')
|
||||
|
||||
def pipeline(root: beam.Pipeline):
|
||||
common_lib.beam_convert_tfexamples(
|
||||
root,
|
||||
images_list,
|
||||
build_tf_example,
|
||||
output_dir,
|
||||
num_shards,
|
||||
)
|
||||
|
||||
common_lib.run_beam_pipeline(pipeline)
|
||||
|
||||
|
||||
def _convert_df_to_tfrecord(
|
||||
df: pd.DataFrame,
|
||||
output_dir: str,
|
||||
split_ratio: Sequence[float],
|
||||
num_shard: Sequence[int],
|
||||
) -> None:
|
||||
"""Converts a DataFrame into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
Args:
|
||||
df: DataFrame to convert.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
split_ratio: List specifying the training, validation, and testing splits
|
||||
for unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
# Replaces ml_use with common_lib string constants for consistency.
|
||||
common_lib.format_ml_use_column(df)
|
||||
common_lib.insert_missing_ml_use(df)
|
||||
|
||||
# Ignores invalid rows.
|
||||
dropped_row_num = common_lib.drop_invalid_rows(df)
|
||||
if dropped_row_num > 0:
|
||||
logging.warning('Ignored %d invalid rows.', dropped_row_num)
|
||||
|
||||
common_lib.replace_unassigned_ml_use(
|
||||
df[common_lib.COLUMN_NAME_ML_USE], split_ratio
|
||||
)
|
||||
|
||||
# Converts labels to integers as required by training.
|
||||
new_labels, label_map = common_lib.create_label_map(
|
||||
df[common_lib.COLUMN_NAME_LABEL]
|
||||
)
|
||||
df[common_lib.COLUMN_NAME_LABEL] = new_labels
|
||||
label_map_path = path.join(output_dir, common_lib.LABEL_MAP_NAME)
|
||||
logging.info('Writing label map to %s.', label_map_path)
|
||||
common_lib.write_label_map(label_map_path, label_map)
|
||||
|
||||
_run_convert_pipeline(output_dir, df, num_shard)
|
||||
|
||||
|
||||
def convert_csv_to_tfrecord(
|
||||
input_csv: str,
|
||||
output_dir: str,
|
||||
split_ratio: Sequence[float] = (0.8, 0.1, 0.1),
|
||||
num_shard: Sequence[int] = (10, 10, 10),
|
||||
) -> None:
|
||||
"""Parses input_csv file into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
The csv format is shown in
|
||||
https://cloud.google.com/vertex-ai/docs/image-data/classification/prepare-data#csv.
|
||||
|
||||
If an ml_use column is not provided, one will be created.
|
||||
|
||||
label_map.yaml containing the label map will be placed in output_dir.
|
||||
|
||||
Args:
|
||||
input_csv: Name of the csv file.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
split_ratio: List specifying the training, validation, and testing splits
|
||||
for unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
with tf.io.gfile.GFile(input_csv, 'r') as f:
|
||||
df: pd.DataFrame = pd.read_csv(
|
||||
f, header=None, names=_COLUMN_NAMES, on_bad_lines='warn'
|
||||
)
|
||||
|
||||
_convert_df_to_tfrecord(df, output_dir, split_ratio, num_shard)
|
||||
|
||||
|
||||
def convert_jsonl_to_tfrecord(
|
||||
input_jsonl: str,
|
||||
output_dir: str,
|
||||
split_ratio: Sequence[float] = (0.8, 0.1, 0.1),
|
||||
num_shard: Sequence[int] = (10, 10, 10),
|
||||
) -> None:
|
||||
"""Parses input_jsonl file into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
The JSONL format is shown in
|
||||
https://cloud.google.com/vertex-ai/docs/image-data/classification/prepare-data#json-lines.
|
||||
|
||||
If an ml_use column is not provided, one will be created.
|
||||
|
||||
label_map.yaml containing the label map will be placed in output_dir.
|
||||
|
||||
Args:
|
||||
input_jsonl: Name of the JSONL file.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
split_ratio: List specifying the training, validation, and testing splits
|
||||
for unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
df_rows = []
|
||||
with tf.io.gfile.GFile(input_jsonl, 'r') as f:
|
||||
lines = f.read().rstrip().splitlines()
|
||||
|
||||
for i, line in enumerate(lines, 1):
|
||||
try:
|
||||
item: Dict[str, Any] = json.loads(line)
|
||||
|
||||
gcs_uri = item.get(_JSON_GCS_URI_KEY)
|
||||
label = item.get(_JSON_CLASS_ANNOTATION_KEY, {}).get(_JSON_CLASS_NAME_KEY)
|
||||
if not gcs_uri or not label:
|
||||
logging.warning('Invalid JSON at line %d, skipped.', i)
|
||||
continue
|
||||
|
||||
ml_use = item.get(_JSON_RESOURCE_LABEL_KEY, {}).get(
|
||||
_JSON_ML_USE_KEY, common_lib.ML_USE_UNASSIGNED
|
||||
)
|
||||
except (json.JSONDecodeError, AttributeError):
|
||||
logging.warning('Invalid JSON at line %d, skipped.', i)
|
||||
continue
|
||||
|
||||
df_rows.append([ml_use, gcs_uri, label])
|
||||
|
||||
df = pd.DataFrame(
|
||||
data=df_rows,
|
||||
columns=[
|
||||
common_lib.COLUMN_NAME_ML_USE,
|
||||
common_lib.COLUMN_NAME_GCS_FILE_PATH,
|
||||
common_lib.COLUMN_NAME_LABEL,
|
||||
],
|
||||
)
|
||||
|
||||
_convert_df_to_tfrecord(df, output_dir, split_ratio, num_shard)
|
||||
+430
@@ -0,0 +1,430 @@
|
||||
"""Converts IOD dataset files to TFRecord with apache beam."""
|
||||
|
||||
import collections
|
||||
import json
|
||||
from os import path
|
||||
from typing import Any, Dict, Sequence
|
||||
|
||||
from absl import logging
|
||||
import apache_beam as beam
|
||||
import pandas as pd
|
||||
import tensorflow as tf
|
||||
|
||||
from data_converter import common_lib
|
||||
from util import constants
|
||||
|
||||
COLUMN_NAME_LABEL_INT = 'label_int'
|
||||
_COLUMN_NAME_XMIN = 'X_MIN'
|
||||
_COLUMN_NAME_YMIN = 'Y_MIN'
|
||||
_COLUMN_NAME_XMAX = 'X_MAX'
|
||||
_COLUMN_NAME_YMAX = 'Y_MAX'
|
||||
COLUMN_NAMES = [
|
||||
common_lib.COLUMN_NAME_ML_USE,
|
||||
common_lib.COLUMN_NAME_GCS_FILE_PATH,
|
||||
common_lib.COLUMN_NAME_LABEL,
|
||||
_COLUMN_NAME_XMIN,
|
||||
_COLUMN_NAME_YMIN,
|
||||
'XMAX_NOT_USED',
|
||||
'YMIN_NOT_USED',
|
||||
_COLUMN_NAME_XMAX,
|
||||
_COLUMN_NAME_YMAX,
|
||||
'XMIN_NOT_USED',
|
||||
'YMAX_NOT_USED',
|
||||
]
|
||||
_BOUNDING_BOX_COLUMNS = [
|
||||
_COLUMN_NAME_XMIN,
|
||||
_COLUMN_NAME_YMIN,
|
||||
_COLUMN_NAME_XMAX,
|
||||
_COLUMN_NAME_YMAX,
|
||||
]
|
||||
_JSON_BBOX_ANNOTATIONS_KEY = 'boundingBoxAnnotations'
|
||||
_JSON_DISPLAY_NAME_KEY = 'displayName'
|
||||
_JSON_X_MIN_KEY = 'xMin'
|
||||
_JSON_X_MAX_KEY = 'xMax'
|
||||
_JSON_Y_MIN_KEY = 'yMin'
|
||||
_JSON_Y_MAX_KEY = 'yMax'
|
||||
|
||||
|
||||
def build_tf_example(image_row: Dict[str, Any]) -> tf.train.Example:
|
||||
"""Builds a TF Example from an image row.
|
||||
|
||||
Args:
|
||||
image_row: A dictionary containing information about the image, such as its
|
||||
GCS uri, labels, and bounding box coordinates.
|
||||
|
||||
Returns:
|
||||
A tf.train.Example containing the encoded image and optionally a
|
||||
bounding box and label.
|
||||
"""
|
||||
image_uri = image_row[common_lib.COLUMN_NAME_GCS_FILE_PATH]
|
||||
image_bytes, shape = common_lib.encode_image(image_uri, image_format='jpeg')
|
||||
feature = {
|
||||
'image/encoded': common_lib.convert_to_feature(image_bytes),
|
||||
'image/format': common_lib.convert_to_string_feature('jpeg'),
|
||||
'image/height': common_lib.convert_to_feature(shape[0]),
|
||||
'image/width': common_lib.convert_to_feature(shape[1]),
|
||||
'image/source_id': common_lib.convert_to_string_feature(image_uri),
|
||||
'image/object/bbox/xmin': common_lib.convert_to_feature(
|
||||
image_row[_COLUMN_NAME_XMIN]
|
||||
),
|
||||
'image/object/bbox/ymin': common_lib.convert_to_feature(
|
||||
image_row[_COLUMN_NAME_YMIN]
|
||||
),
|
||||
'image/object/bbox/xmax': common_lib.convert_to_feature(
|
||||
image_row[_COLUMN_NAME_XMAX]
|
||||
),
|
||||
'image/object/bbox/ymax': common_lib.convert_to_feature(
|
||||
image_row[_COLUMN_NAME_YMAX]
|
||||
),
|
||||
'image/object/class/text': common_lib.convert_to_list_string_feature(
|
||||
image_row[common_lib.COLUMN_NAME_LABEL]
|
||||
),
|
||||
'image/object/class/label': common_lib.convert_to_feature(
|
||||
image_row[COLUMN_NAME_LABEL_INT]
|
||||
),
|
||||
}
|
||||
return tf.train.Example(features=tf.train.Features(feature=feature))
|
||||
|
||||
|
||||
def _run_convert_pipeline(
|
||||
output_dir: str,
|
||||
image_rows: Sequence[Dict[str, Any]],
|
||||
num_shards: Sequence[int],
|
||||
) -> None:
|
||||
"""Starts a Beam pipeline to write DataFrame as TF Records.
|
||||
|
||||
Args:
|
||||
output_dir: TF Records output directory.
|
||||
image_rows: Contains all necessary information to create a TF Example.
|
||||
num_shards: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
|
||||
def pipeline(root: beam.Pipeline):
|
||||
common_lib.beam_convert_tfexamples(
|
||||
root,
|
||||
image_rows,
|
||||
build_tf_example,
|
||||
output_dir,
|
||||
num_shards,
|
||||
)
|
||||
|
||||
common_lib.run_beam_pipeline(pipeline)
|
||||
|
||||
|
||||
def _convert_df_to_tfrecord(
|
||||
df: pd.DataFrame,
|
||||
output_dir: str,
|
||||
split_ratio: Sequence[float],
|
||||
num_shard: Sequence[int],
|
||||
) -> None:
|
||||
"""Converts a DataFrame into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
Args:
|
||||
df: DataFrame to convert.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
split_ratio: List specifying the training, validation, and testing splits
|
||||
for unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
# Replaces ml_use with common_lib string constants for consistency.
|
||||
common_lib.format_ml_use_column(df)
|
||||
common_lib.insert_missing_ml_use(df)
|
||||
|
||||
# Specify bounding box columns to be numeric.
|
||||
df[_BOUNDING_BOX_COLUMNS] = df[_BOUNDING_BOX_COLUMNS].apply(pd.to_numeric)
|
||||
|
||||
# Ignores invalid rows.
|
||||
dropped_row_num = common_lib.drop_invalid_rows(df)
|
||||
dropped_row_num += drop_rows_without_bbox(df)
|
||||
if dropped_row_num > 0:
|
||||
logging.warning('Ignored %d invalid rows.', dropped_row_num)
|
||||
|
||||
# Converts labels to integers as required by training.
|
||||
int_labels, label_map = common_lib.create_label_map(
|
||||
df[common_lib.COLUMN_NAME_LABEL]
|
||||
)
|
||||
df[COLUMN_NAME_LABEL_INT] = int_labels
|
||||
label_map_path = path.join(output_dir, common_lib.LABEL_MAP_NAME)
|
||||
logging.info('Writing label map to %s.', label_map_path)
|
||||
common_lib.write_label_map(label_map_path, label_map)
|
||||
|
||||
image_rows = _condense_bounding_boxes(df.to_dict(orient='records'))
|
||||
ml_uses = [row[common_lib.COLUMN_NAME_ML_USE] for row in image_rows]
|
||||
common_lib.replace_unassigned_ml_use(ml_uses, split_ratio)
|
||||
common_lib.merge_seq_into_dicts(
|
||||
common_lib.COLUMN_NAME_ML_USE, ml_uses, image_rows
|
||||
)
|
||||
|
||||
_run_convert_pipeline(output_dir, image_rows, num_shard)
|
||||
|
||||
|
||||
def _condense_bounding_boxes(
|
||||
image_rows: Sequence[Dict[str, Any]]
|
||||
) -> Sequence[Dict[str, Any]]:
|
||||
"""Gather all the bounding boxes in an image and put them in the same dictionary.
|
||||
|
||||
Args:
|
||||
image_rows: List of dictionaries, each containing information about the
|
||||
image, such as its GCS uri, labels, and bounding box coordinates.
|
||||
|
||||
Returns:
|
||||
List of dictionaries such that each contains all the bounding boxes for a
|
||||
given gcs_file_path.
|
||||
|
||||
Raises:
|
||||
RuntimeError: This is raised when the input data contains images that have
|
||||
annotations in different ml_use classes.
|
||||
"""
|
||||
output = {}
|
||||
for image_row in image_rows:
|
||||
ml_use = image_row[common_lib.COLUMN_NAME_ML_USE]
|
||||
gcs_file_path = image_row[common_lib.COLUMN_NAME_GCS_FILE_PATH]
|
||||
label = image_row[common_lib.COLUMN_NAME_LABEL]
|
||||
xmin = image_row[_COLUMN_NAME_XMIN]
|
||||
ymin = image_row[_COLUMN_NAME_YMIN]
|
||||
xmax = image_row[_COLUMN_NAME_XMAX]
|
||||
ymax = image_row[_COLUMN_NAME_YMAX]
|
||||
label_int = image_row[COLUMN_NAME_LABEL_INT]
|
||||
if gcs_file_path in output:
|
||||
d = output[gcs_file_path]
|
||||
if ml_use != common_lib.ML_USE_UNASSIGNED:
|
||||
if d[common_lib.COLUMN_NAME_ML_USE] == common_lib.ML_USE_UNASSIGNED:
|
||||
d[common_lib.COLUMN_NAME_ML_USE] = ml_use
|
||||
elif ml_use != d[common_lib.COLUMN_NAME_ML_USE]:
|
||||
raise RuntimeError(
|
||||
f'Image {gcs_file_path} can only be placed in one of'
|
||||
f' training/validation/test. It is currently in {ml_use} and'
|
||||
f' {d[common_lib.COLUMN_NAME_ML_USE]}.'
|
||||
)
|
||||
d[common_lib.COLUMN_NAME_LABEL].append(label)
|
||||
d[_COLUMN_NAME_XMIN].append(xmin)
|
||||
d[_COLUMN_NAME_YMIN].append(ymin)
|
||||
d[_COLUMN_NAME_XMAX].append(xmax)
|
||||
d[_COLUMN_NAME_YMAX].append(ymax)
|
||||
d[COLUMN_NAME_LABEL_INT].append(label_int)
|
||||
else:
|
||||
output[gcs_file_path] = {
|
||||
common_lib.COLUMN_NAME_ML_USE: ml_use,
|
||||
common_lib.COLUMN_NAME_GCS_FILE_PATH: gcs_file_path,
|
||||
common_lib.COLUMN_NAME_LABEL: [label],
|
||||
_COLUMN_NAME_XMIN: [xmin],
|
||||
_COLUMN_NAME_YMIN: [ymin],
|
||||
_COLUMN_NAME_XMAX: [xmax],
|
||||
_COLUMN_NAME_YMAX: [ymax],
|
||||
COLUMN_NAME_LABEL_INT: [label_int],
|
||||
}
|
||||
return list(output.values())
|
||||
|
||||
|
||||
def convert_csv_to_tfrecord(
|
||||
input_csv: str,
|
||||
output_dir: str,
|
||||
split_ratio: Sequence[float] = (0.8, 0.1, 0.1),
|
||||
num_shard: Sequence[int] = (10, 10, 10),
|
||||
) -> None:
|
||||
"""Parses input_csv file into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
The csv format is shown in
|
||||
https://cloud.google.com/vertex-ai/docs/image-data/object-detection/prepare-data#csv.
|
||||
|
||||
If an ml_use column is not provided, one will be created.
|
||||
|
||||
label_map.yaml containing the label map will be placed in output_dir.
|
||||
|
||||
Args:
|
||||
input_csv: Name of the csv file.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
split_ratio: List specifying the train, validation, and test splits for
|
||||
unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
with tf.io.gfile.GFile(input_csv, 'r') as f:
|
||||
df: pd.DataFrame = pd.read_csv(
|
||||
f, header=None, names=COLUMN_NAMES, on_bad_lines='warn'
|
||||
)
|
||||
|
||||
_convert_df_to_tfrecord(df, output_dir, split_ratio, num_shard)
|
||||
|
||||
|
||||
def drop_rows_without_bbox(df: pd.DataFrame) -> int:
|
||||
"""Drops DataFrame rows without bounding_boxes.
|
||||
|
||||
Args:
|
||||
df: The DataFrame to process in place.
|
||||
|
||||
Returns:
|
||||
The number of rows dropped.
|
||||
"""
|
||||
invalid_rows = df.index[~(df[_BOUNDING_BOX_COLUMNS].notnull().all(axis=1))]
|
||||
dropped_num = len(invalid_rows)
|
||||
if dropped_num > 0:
|
||||
invalid_df = df.loc[invalid_rows].to_dict(orient='records')
|
||||
for entry in invalid_df:
|
||||
logging.warning('Skipping entry due to missing bounding box: %s.', entry)
|
||||
df.drop(invalid_rows, inplace=True)
|
||||
df.reset_index(drop=True, inplace=True)
|
||||
return dropped_num
|
||||
|
||||
|
||||
def convert_coco_json_categories_to_label_map(
|
||||
categories: Sequence[Dict[str, Any]]
|
||||
) -> Dict[int, str]:
|
||||
return {category['id']: category['name'] for category in categories}
|
||||
|
||||
|
||||
def convert_coco_json_to_tfrecord(
|
||||
input_coco_json: str,
|
||||
output_dir: str,
|
||||
split_ratio: Sequence[float] = (0.8, 0.1, 0.1),
|
||||
num_shard: Sequence[int] = (10, 10, 10),
|
||||
) -> None:
|
||||
"""Parses input_csv file into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
The COCO json format is shown here: https://cocodataset.org/#format-data.
|
||||
|
||||
label_map.yaml containing the label map will be placed in output_dir.
|
||||
|
||||
Args:
|
||||
input_coco_json: Name of coco json file.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
split_ratio: List specifying the train, validation, and test splits for
|
||||
dataset.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
with tf.io.gfile.GFile(input_coco_json, 'r') as f:
|
||||
coco_json = json.load(f)
|
||||
# Writes label map from coco json categories.
|
||||
label_map = convert_coco_json_categories_to_label_map(
|
||||
coco_json[constants.COCO_JSON_CATEGORIES]
|
||||
)
|
||||
label_map_path = path.join(output_dir, common_lib.LABEL_MAP_NAME)
|
||||
logging.info('Writes label map to %s.', label_map_path)
|
||||
common_lib.write_label_map(label_map_path, label_map)
|
||||
|
||||
img_to_anns = collections.defaultdict(list)
|
||||
imgs = {}
|
||||
if constants.COCO_JSON_ANNOTATIONS in coco_json:
|
||||
for ann in coco_json[constants.COCO_JSON_ANNOTATIONS]:
|
||||
img_to_anns[ann[constants.COCO_JSON_ANNOTATION_IMAGE_ID]].append(ann)
|
||||
|
||||
if constants.COCO_JSON_IMAGES in coco_json:
|
||||
for img in coco_json[constants.COCO_JSON_IMAGES]:
|
||||
imgs[img[constants.COCO_JSON_IMAGE_ID]] = img
|
||||
|
||||
df_rows = []
|
||||
|
||||
for image_id, annotations in img_to_anns.items():
|
||||
img = imgs[image_id]
|
||||
for ann in annotations:
|
||||
xmin, ymin, xmax, ymax = common_lib.reformat_bbox(
|
||||
ann[constants.COCO_ANNOTATION_BBOX],
|
||||
img[constants.COCO_JSON_IMAGE_WIDTH],
|
||||
img[constants.COCO_JSON_IMAGE_HEIGHT],
|
||||
)
|
||||
df_rows.append([
|
||||
common_lib.ML_USE_UNASSIGNED,
|
||||
img[constants.COCO_JSON_IMAGE_COCO_URL],
|
||||
label_map[ann[constants.COCO_JSON_ANNOTATION_CATEGORY_ID]],
|
||||
xmin,
|
||||
ymin,
|
||||
xmax,
|
||||
ymin,
|
||||
xmax,
|
||||
ymax,
|
||||
xmin,
|
||||
ymax,
|
||||
ann[constants.COCO_JSON_ANNOTATION_CATEGORY_ID],
|
||||
])
|
||||
df = pd.DataFrame(
|
||||
data=df_rows,
|
||||
columns=COLUMN_NAMES + [COLUMN_NAME_LABEL_INT],
|
||||
)
|
||||
|
||||
# Replaces ml_use with common_lib string constants for consistency.
|
||||
common_lib.format_ml_use_column(df)
|
||||
common_lib.insert_missing_ml_use(df)
|
||||
|
||||
# Species bounding box columns to be numeric.
|
||||
df[_BOUNDING_BOX_COLUMNS] = df[_BOUNDING_BOX_COLUMNS].apply(pd.to_numeric)
|
||||
|
||||
# Ignores invalid rows.
|
||||
dropped_row_num = common_lib.drop_invalid_rows(df)
|
||||
dropped_row_num += drop_rows_without_bbox(df)
|
||||
if dropped_row_num > 0:
|
||||
logging.warning('Ignored %d invalid rows.', dropped_row_num)
|
||||
|
||||
image_rows = _condense_bounding_boxes(df.to_dict(orient='records'))
|
||||
ml_uses = [row[common_lib.COLUMN_NAME_ML_USE] for row in image_rows]
|
||||
common_lib.replace_unassigned_ml_use(ml_uses, split_ratio)
|
||||
common_lib.merge_seq_into_dicts(
|
||||
common_lib.COLUMN_NAME_ML_USE, ml_uses, image_rows
|
||||
)
|
||||
|
||||
_run_convert_pipeline(output_dir, image_rows, num_shard)
|
||||
|
||||
|
||||
def convert_jsonl_to_tfrecord(
|
||||
input_jsonl: str,
|
||||
output_dir: str,
|
||||
split_ratio: Sequence[float] = (0.8, 0.1, 0.1),
|
||||
num_shard: Sequence[int] = (10, 10, 10),
|
||||
) -> None:
|
||||
"""Parses input_jsonl file into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
The JSONL format is shown in
|
||||
https://cloud.google.com/vertex-ai/docs/image-data/object-detection/prepare-data#json-lines.
|
||||
|
||||
If an ml_use column is not provided, one will be created.
|
||||
|
||||
label_map.yaml containing the label map will be placed in output_dir.
|
||||
|
||||
Args:
|
||||
input_jsonl: Name of the JSONL file.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
split_ratio: List specifying the training, validation, and testing splits
|
||||
for unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
df_rows = []
|
||||
with tf.io.gfile.GFile(input_jsonl, 'r') as f:
|
||||
lines = f.read().rstrip().splitlines()
|
||||
|
||||
for i, line in enumerate(lines, start=1):
|
||||
try:
|
||||
item: Dict[str, Any] = json.loads(line)
|
||||
except (json.JSONDecodeError, AttributeError):
|
||||
logging.warning('Invalid JSON at line %d skipped.', i)
|
||||
continue
|
||||
|
||||
gcs_uri = item.get(common_lib.JSON_GCS_URI_KEY)
|
||||
if not gcs_uri:
|
||||
logging.warning(
|
||||
'Invalid JSON at line %d skipped. Missing gcs_uri_key.', i
|
||||
)
|
||||
continue
|
||||
ml_use = item.get(common_lib.JSON_RESOURCE_LABEL_KEY, {}).get(
|
||||
common_lib.JSON_ML_USE_KEY, common_lib.ML_USE_UNASSIGNED
|
||||
)
|
||||
|
||||
for bbox in item.get(_JSON_BBOX_ANNOTATIONS_KEY, []):
|
||||
label = bbox.get(_JSON_DISPLAY_NAME_KEY)
|
||||
xmin = bbox.get(_JSON_X_MIN_KEY)
|
||||
ymin = bbox.get(_JSON_Y_MIN_KEY)
|
||||
xmax = bbox.get(_JSON_X_MAX_KEY)
|
||||
ymax = bbox.get(_JSON_Y_MAX_KEY)
|
||||
|
||||
df_rows.append([ml_use, gcs_uri, label, xmin, ymin, xmax, ymax])
|
||||
|
||||
df = pd.DataFrame(
|
||||
data=df_rows,
|
||||
columns=[
|
||||
common_lib.COLUMN_NAME_ML_USE,
|
||||
common_lib.COLUMN_NAME_GCS_FILE_PATH,
|
||||
common_lib.COLUMN_NAME_LABEL,
|
||||
_COLUMN_NAME_XMIN,
|
||||
_COLUMN_NAME_YMIN,
|
||||
_COLUMN_NAME_XMAX,
|
||||
_COLUMN_NAME_YMAX,
|
||||
],
|
||||
)
|
||||
_convert_df_to_tfrecord(df, output_dir, split_ratio, num_shard)
|
||||
+328
@@ -0,0 +1,328 @@
|
||||
"""Python script to convert different file formats for ISG to tfrecords."""
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
from typing import Any, Dict, Iterator, List, Optional, Tuple, Union
|
||||
|
||||
from absl import logging
|
||||
import apache_beam as beam
|
||||
from apache_beam.io import tfrecordio
|
||||
import cv2
|
||||
import numpy as np
|
||||
from pycocotools import coco
|
||||
import tensorflow as tf
|
||||
import yaml
|
||||
|
||||
from data_converter import common_lib
|
||||
from util import constants
|
||||
from util import fileutils
|
||||
|
||||
_IMAGE_FORMAT = 'PNG'
|
||||
|
||||
|
||||
def build_tf_example(
|
||||
image_info: dict[str, Union[str, int]],
|
||||
segmentation_image: List[List[int]],
|
||||
output_shape: Optional[Tuple[int, int]] = None,
|
||||
) -> tf.train.Example:
|
||||
"""Encodes an image and its segmentation mask into a tf.train.Example.
|
||||
|
||||
Args:
|
||||
image_info: A dictionary containing information about the image, such as its
|
||||
file name, height, and width.
|
||||
segmentation_image: 2D image in list of lists having category ids.
|
||||
output_shape: The desired output shape of the image. If None, the original
|
||||
image shape will be used.
|
||||
|
||||
Returns:
|
||||
A tf.train.Example containing the encoded image and segmentation mask.
|
||||
|
||||
Raises:
|
||||
IOError: If image cannot be found in the path.
|
||||
"""
|
||||
file_name = image_info[constants.COCO_JSON_FILE_NAME]
|
||||
height = int(image_info[constants.COCO_JSON_IMAGE_HEIGHT])
|
||||
width = int(image_info[constants.COCO_JSON_IMAGE_WIDTH])
|
||||
|
||||
segmentation_image = np.expand_dims(
|
||||
np.asarray(segmentation_image, dtype=np.int32), axis=-1
|
||||
)
|
||||
_, encoded_seg = cv2.imencode(f'.{_IMAGE_FORMAT.lower()}', segmentation_image)
|
||||
encoded_seg = encoded_seg.tobytes()
|
||||
|
||||
encoded_img, _ = common_lib.encode_image(
|
||||
image_info[constants.COCO_JSON_IMAGE_COCO_URL],
|
||||
output_shape=output_shape,
|
||||
image_format=_IMAGE_FORMAT.lower(),
|
||||
)
|
||||
|
||||
key = hashlib.sha256(encoded_img).hexdigest()
|
||||
|
||||
return tf.train.Example(
|
||||
features=tf.train.Features(
|
||||
feature={
|
||||
'image/height': common_lib.convert_to_feature(height),
|
||||
'image/width': common_lib.convert_to_feature(width),
|
||||
'image/filename': common_lib.convert_to_string_feature(file_name),
|
||||
'image/sha256': common_lib.convert_to_string_feature(key),
|
||||
'image/encoded': common_lib.convert_to_feature(encoded_img),
|
||||
'image/format': common_lib.convert_to_string_feature(
|
||||
_IMAGE_FORMAT
|
||||
),
|
||||
'image/segmentation/class/encoded': common_lib.convert_to_feature(
|
||||
encoded_seg
|
||||
),
|
||||
'image/segmentation/class/format': (
|
||||
common_lib.convert_to_string_feature(_IMAGE_FORMAT)
|
||||
),
|
||||
'image/segmentation/class/height': common_lib.convert_to_feature(
|
||||
height
|
||||
),
|
||||
'image/segmentation/class/width': common_lib.convert_to_feature(
|
||||
width
|
||||
),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class AcquireTFExampleDoFn(beam.DoFn):
|
||||
"""Beam DoFn to build TF Examples from a single row of image_info data."""
|
||||
|
||||
# These tags will be used to tag the outputs of this DoFn.
|
||||
output_tag_train = constants.ML_USE_TRAINING
|
||||
output_tag_validation = constants.ML_USE_VALIDATION
|
||||
output_tag_test = constants.ML_USE_TEST
|
||||
|
||||
valid_ml_use_set = set(
|
||||
[output_tag_train, output_tag_validation, output_tag_test]
|
||||
)
|
||||
|
||||
def __init__(self, output_shape: Optional[Tuple[int, int]] = None):
|
||||
self.acquired_examples_counter = beam.metrics.Metrics.counter(
|
||||
self.__class__.__name__, 'Success'
|
||||
)
|
||||
self.failure_counter = beam.metrics.Metrics.counter(
|
||||
self.__class__.__name__, 'Failure'
|
||||
)
|
||||
self.output_shape = output_shape
|
||||
|
||||
def process(
|
||||
self,
|
||||
row: Tuple[str, Dict[str, Union[str, int]], List[List[int]]],
|
||||
) -> Iterator[tf.train.Example]:
|
||||
ml_use, image_info, annotation_info = row
|
||||
if ml_use not in self.valid_ml_use_set:
|
||||
logging.warning('ml_use invalid: %s', ml_use)
|
||||
self.failure_counter.inc()
|
||||
return
|
||||
|
||||
try:
|
||||
tf_example = build_tf_example(
|
||||
image_info, annotation_info, self.output_shape
|
||||
)
|
||||
except IOError as e:
|
||||
logging.warning('Failed to build TF Example: %s', e)
|
||||
self.failure_counter.inc()
|
||||
else:
|
||||
self.acquired_examples_counter.inc()
|
||||
yield beam.pvalue.TaggedOutput(ml_use, tf_example)
|
||||
|
||||
|
||||
def _define_data_conversion_pipeline(
|
||||
root: beam.Pipeline,
|
||||
ml_use_rows: List[str],
|
||||
image_rows: List[Dict[str, Union[str, int]]],
|
||||
segmentation_rows: List[List[List[int]]],
|
||||
output_dir: str,
|
||||
output_shape: Optional[Tuple[int, int]],
|
||||
num_shard_list: List[int],
|
||||
):
|
||||
"""Define a data conversion pipeline.
|
||||
|
||||
Args:
|
||||
root: A Beam pipeline.
|
||||
ml_use_rows: List containing the ml_use.
|
||||
image_rows: List of dictionaries containing information about the image,
|
||||
such as its file name, height, and width.
|
||||
segmentation_rows: List of 2D images of integers representing segmentation
|
||||
masks.
|
||||
output_dir: Directory where the output TFRecords will be written.
|
||||
output_shape: Desired output shape of the image. If None, the original image
|
||||
shape will be used.
|
||||
num_shard_list: Number of shards to write to each output TFRecord.
|
||||
|
||||
Returns:
|
||||
A Beam pipeline.
|
||||
"""
|
||||
train, validation, test = (
|
||||
root
|
||||
| 'Load ml use and image rows to beam'
|
||||
>> beam.Create(zip(ml_use_rows, image_rows, segmentation_rows))
|
||||
| 'Build TF Examples'
|
||||
>> beam.ParDo(AcquireTFExampleDoFn(output_shape)).with_outputs(
|
||||
AcquireTFExampleDoFn.output_tag_train,
|
||||
AcquireTFExampleDoFn.output_tag_validation,
|
||||
AcquireTFExampleDoFn.output_tag_test,
|
||||
)
|
||||
)
|
||||
|
||||
# Save each split to TFRecord.
|
||||
_ = train | 'Save train split to TFRecord' >> tfrecordio.WriteToTFRecord(
|
||||
os.path.join(output_dir, common_lib.TRAIN_TFRECORD_NAME),
|
||||
coder=beam.coders.ProtoCoder(tf.train.Example),
|
||||
num_shards=num_shard_list[0],
|
||||
)
|
||||
_ = (
|
||||
validation
|
||||
| 'Save validation split to TFRecord'
|
||||
>> tfrecordio.WriteToTFRecord(
|
||||
os.path.join(output_dir, common_lib.VALIDATION_TFRECORD_NAME),
|
||||
coder=beam.coders.ProtoCoder(tf.train.Example),
|
||||
num_shards=num_shard_list[1],
|
||||
)
|
||||
)
|
||||
_ = test | 'Save test split to TFRecord' >> tfrecordio.WriteToTFRecord(
|
||||
os.path.join(output_dir, common_lib.TEST_TFRECORD_NAME),
|
||||
coder=beam.coders.ProtoCoder(tf.train.Example),
|
||||
num_shards=num_shard_list[2],
|
||||
)
|
||||
|
||||
|
||||
def _image_info_to_segmentation_image(
|
||||
img: Dict[str, Any],
|
||||
coco_dataset: coco.COCO,
|
||||
label_id_by_category_id: Dict[int, int],
|
||||
) -> List[List[int]]:
|
||||
"""Convert image information to a segmentation image.
|
||||
|
||||
Args:
|
||||
img: The image information.
|
||||
coco_dataset: The COCO dataset.
|
||||
label_id_by_category_id: The mapping from label id used for training to
|
||||
category_id defined in dataset.
|
||||
|
||||
Returns:
|
||||
The segmentation image.
|
||||
|
||||
Raises:
|
||||
ValueError: If the mask size does not match the image or if a pixel has
|
||||
multiple labels.
|
||||
"""
|
||||
seg_img = np.zeros(
|
||||
shape=(
|
||||
img[constants.COCO_JSON_IMAGE_HEIGHT],
|
||||
img[constants.COCO_JSON_IMAGE_WIDTH],
|
||||
),
|
||||
dtype=np.int32,
|
||||
)
|
||||
for ann in coco_dataset.imgToAnns[img[constants.COCO_JSON_IMAGE_ID]]:
|
||||
new_category_id = ann[constants.COCO_JSON_ANNOTATION_CATEGORY_ID]
|
||||
binary_mask = coco_dataset.annToMask(ann)
|
||||
if seg_img.shape != binary_mask.shape:
|
||||
raise ValueError(
|
||||
'Binary mask does not have the same shape as image. image_id:'
|
||||
f' {img["id"]}'
|
||||
)
|
||||
boolean_mask = binary_mask == 1
|
||||
if (seg_img[boolean_mask] != 0).any():
|
||||
raise ValueError(
|
||||
'Error: Some pixels have more than one label in image_id:'
|
||||
f' {img["id"]}.'
|
||||
)
|
||||
seg_img[boolean_mask] = label_id_by_category_id[new_category_id]
|
||||
|
||||
return seg_img.tolist()
|
||||
|
||||
|
||||
def get_input_rows(
|
||||
coco_dataset: coco.COCO,
|
||||
split_ratio: List[float],
|
||||
label_id_by_category_id: Dict[int, int],
|
||||
) -> Tuple[List[str], List[Dict[str, Union[str, int]]], List[List[List[int]]]]:
|
||||
"""Get input rows for training and validation.
|
||||
|
||||
Args:
|
||||
coco_dataset: The COCO dataset.
|
||||
split_ratio: The split ratio for training and validation.
|
||||
label_id_by_category_id: The mapping from label id used for training to
|
||||
category_id defined in dataset.
|
||||
|
||||
Returns:
|
||||
- A list of ml_use strings.
|
||||
- A list of image informations.
|
||||
- A list of segmentation images for the corresponding images.
|
||||
"""
|
||||
image_rows = coco_dataset.dataset[constants.COCO_JSON_IMAGES]
|
||||
|
||||
segmentation_rows = [
|
||||
_image_info_to_segmentation_image(
|
||||
img, coco_dataset, label_id_by_category_id
|
||||
)
|
||||
for img in image_rows
|
||||
]
|
||||
|
||||
ml_use_rows = common_lib.create_ml_use_array_with_split(
|
||||
len(image_rows), split_ratio
|
||||
)
|
||||
return ml_use_rows, image_rows, segmentation_rows
|
||||
|
||||
|
||||
def beam_build_tfrecord_from_coco_json(
|
||||
input_json: str,
|
||||
output_dir: str,
|
||||
split_ratio: List[float],
|
||||
num_shard_list: List[int],
|
||||
output_shape: Optional[Tuple[int, int]] = None,
|
||||
) -> None:
|
||||
"""Builds TFRecord files from COCO dataset.
|
||||
|
||||
The output file names are `_TRAIN_TFRECORD_NAME`, `_VALIDATION_TFRECORD_NAME`,
|
||||
and `_TEST_TFRECORD_NAME`.
|
||||
|
||||
Args:
|
||||
input_json: Path to a COCO JSON or JSONL file.
|
||||
output_dir: Directory to output the TFRecord files.
|
||||
split_ratio: List of how to split entries to train, validation, and test
|
||||
TFRecords.
|
||||
num_shard_list: List of the number of shards for each TFRecord file.
|
||||
output_shape: The desired output shape of the image. If None, the original
|
||||
image shape will be used.
|
||||
"""
|
||||
# `coco` cannot access gcs uri. Use gcsfuse, it is faster.
|
||||
input_json = fileutils.force_gcs_fuse_path(input_json)
|
||||
coco_dataset = coco.COCO(input_json)
|
||||
|
||||
label_map = {}
|
||||
label_id_by_category_id = {}
|
||||
for idx, category in enumerate(
|
||||
coco_dataset.dataset[constants.COCO_JSON_CATEGORIES], start=1
|
||||
):
|
||||
label_map[idx] = category[constants.COCO_JSON_CATEGORY_NAME]
|
||||
label_id_by_category_id[category[constants.COCO_JSON_CATEGORY_ID]] = idx
|
||||
label_map_path = os.path.join(output_dir, common_lib.LABEL_MAP_NAME)
|
||||
logging.info('Writing label map to %s.', label_map_path)
|
||||
common_lib.write_label_map(label_map_path, label_map)
|
||||
|
||||
with tf.io.gfile.GFile(
|
||||
os.path.join(output_dir, 'label_id_by_category_id.yaml'), 'w'
|
||||
) as f:
|
||||
yaml.dump(label_id_by_category_id, f)
|
||||
|
||||
ml_use_rows, image_rows, segmentation_rows = get_input_rows(
|
||||
coco_dataset, split_ratio, label_id_by_category_id
|
||||
)
|
||||
|
||||
def pipeline(root):
|
||||
_define_data_conversion_pipeline(
|
||||
root,
|
||||
ml_use_rows,
|
||||
image_rows,
|
||||
segmentation_rows,
|
||||
output_dir,
|
||||
output_shape,
|
||||
num_shard_list,
|
||||
)
|
||||
|
||||
logging.info('Beginning beam pipeline to acquire tfrecords.')
|
||||
common_lib.run_beam_pipeline(pipeline)
|
||||
+166
@@ -0,0 +1,166 @@
|
||||
r"""Python script to convert user input data to training docker format.
|
||||
|
||||
|
||||
Note: the training format is designed to be tfrecord as in the design doc.
|
||||
If there are training efficiency issues for pytorch algorithms, we will also
|
||||
support pytorch formats as well.
|
||||
"""
|
||||
|
||||
from absl import app
|
||||
from absl import flags
|
||||
from absl import logging
|
||||
from data_converter import common_lib
|
||||
from data_converter import data_converter_icn_lib
|
||||
from data_converter import data_converter_iod_lib
|
||||
from data_converter import data_converter_isg_lib
|
||||
from data_converter import data_converter_vcn_lib
|
||||
from util import constants
|
||||
|
||||
|
||||
_INPUT_FILE_PATH = flags.DEFINE_string(
|
||||
'input_file_path',
|
||||
None,
|
||||
'Input file path.',
|
||||
required=True,
|
||||
)
|
||||
_INPUT_FILE_TYPE = flags.DEFINE_enum(
|
||||
'input_file_type',
|
||||
None,
|
||||
[
|
||||
constants.INPUT_FILE_TYPE_CSV,
|
||||
constants.INPUT_FILE_TYPE_JSONL,
|
||||
constants.INPUT_FILE_TYPE_COCO_JSON,
|
||||
],
|
||||
'Input file type.',
|
||||
required=True,
|
||||
)
|
||||
_OBJECTIVE = flags.DEFINE_enum(
|
||||
'objective',
|
||||
None,
|
||||
[
|
||||
constants.OBJECTIVE_IMAGE_CLASSIFICATION,
|
||||
constants.OBJECTIVE_IMAGE_OBJECT_DETECTION,
|
||||
constants.OBJECTIVE_IMAGE_SEGMENTATION,
|
||||
constants.OBJECTIVE_VIDEO_CLASSIFICATION,
|
||||
],
|
||||
'The objective of this training job.',
|
||||
required=True,
|
||||
)
|
||||
_OUTPUT_DIR = flags.DEFINE_string(
|
||||
'output_dir',
|
||||
None,
|
||||
'The output directory for converted data and label map files.',
|
||||
required=True,
|
||||
)
|
||||
_SPLIT_RATIO = flags.DEFINE_list(
|
||||
'split_ratio',
|
||||
'0.8,0.1,0.1',
|
||||
'Proportion of data to split into train/validation/test.',
|
||||
)
|
||||
_NUM_SHARD = flags.DEFINE_list(
|
||||
'num_shard', '10,10,10', 'The number of shards for train/validation/test.'
|
||||
)
|
||||
_OUTPUT_FPS = flags.DEFINE_integer(
|
||||
'output_fps', 5, 'For videos only. The output frames rate per second.'
|
||||
)
|
||||
|
||||
|
||||
def main(_) -> None:
|
||||
logging.info(
|
||||
(
|
||||
'Start data converter on: %s (type: %s) with split: %s for %s'
|
||||
' (shard=%s), and output to %s.'
|
||||
),
|
||||
_INPUT_FILE_PATH.value,
|
||||
_INPUT_FILE_TYPE.value,
|
||||
_SPLIT_RATIO.value,
|
||||
_OBJECTIVE.value,
|
||||
_NUM_SHARD.value,
|
||||
_OUTPUT_DIR.value,
|
||||
)
|
||||
split_ratio = list(map(float, _SPLIT_RATIO.value))
|
||||
num_shard = list(map(int, _NUM_SHARD.value))
|
||||
common_lib.check_split_ratio(split_ratio)
|
||||
common_lib.check_num_shard(num_shard)
|
||||
if (
|
||||
_OBJECTIVE.value == constants.OBJECTIVE_IMAGE_OBJECT_DETECTION
|
||||
and _INPUT_FILE_TYPE.value == constants.INPUT_FILE_TYPE_CSV
|
||||
):
|
||||
data_converter_iod_lib.convert_csv_to_tfrecord(
|
||||
_INPUT_FILE_PATH.value,
|
||||
_OUTPUT_DIR.value,
|
||||
split_ratio,
|
||||
num_shard,
|
||||
)
|
||||
elif (
|
||||
_OBJECTIVE.value == constants.OBJECTIVE_IMAGE_OBJECT_DETECTION
|
||||
and _INPUT_FILE_TYPE.value == constants.INPUT_FILE_TYPE_JSONL
|
||||
):
|
||||
data_converter_iod_lib.convert_jsonl_to_tfrecord(
|
||||
_INPUT_FILE_PATH.value, _OUTPUT_DIR.value, split_ratio, num_shard
|
||||
)
|
||||
elif (
|
||||
_OBJECTIVE.value == constants.OBJECTIVE_IMAGE_OBJECT_DETECTION
|
||||
and _INPUT_FILE_TYPE.value == constants.INPUT_FILE_TYPE_COCO_JSON
|
||||
):
|
||||
data_converter_iod_lib.convert_coco_json_to_tfrecord(
|
||||
_INPUT_FILE_PATH.value,
|
||||
_OUTPUT_DIR.value,
|
||||
split_ratio,
|
||||
num_shard,
|
||||
)
|
||||
elif _OBJECTIVE.value == constants.OBJECTIVE_IMAGE_SEGMENTATION:
|
||||
data_converter_isg_lib.beam_build_tfrecord_from_coco_json(
|
||||
_INPUT_FILE_PATH.value,
|
||||
_OUTPUT_DIR.value,
|
||||
split_ratio,
|
||||
num_shard,
|
||||
)
|
||||
elif (
|
||||
_OBJECTIVE.value == constants.OBJECTIVE_IMAGE_CLASSIFICATION
|
||||
and _INPUT_FILE_TYPE.value == constants.INPUT_FILE_TYPE_CSV
|
||||
):
|
||||
data_converter_icn_lib.convert_csv_to_tfrecord(
|
||||
_INPUT_FILE_PATH.value,
|
||||
_OUTPUT_DIR.value,
|
||||
split_ratio,
|
||||
num_shard,
|
||||
)
|
||||
elif (
|
||||
_OBJECTIVE.value == constants.OBJECTIVE_IMAGE_CLASSIFICATION
|
||||
and _INPUT_FILE_TYPE.value == constants.INPUT_FILE_TYPE_JSONL
|
||||
):
|
||||
data_converter_icn_lib.convert_jsonl_to_tfrecord(
|
||||
_INPUT_FILE_PATH.value, _OUTPUT_DIR.value, split_ratio, num_shard
|
||||
)
|
||||
elif (
|
||||
_OBJECTIVE.value == constants.OBJECTIVE_VIDEO_CLASSIFICATION
|
||||
and _INPUT_FILE_TYPE.value == constants.INPUT_FILE_TYPE_CSV
|
||||
):
|
||||
data_converter_vcn_lib.convert_csv_to_tfrecord(
|
||||
_INPUT_FILE_PATH.value,
|
||||
_OUTPUT_DIR.value,
|
||||
_OUTPUT_FPS.value,
|
||||
split_ratio,
|
||||
num_shard,
|
||||
)
|
||||
elif (
|
||||
_OBJECTIVE.value == constants.OBJECTIVE_VIDEO_CLASSIFICATION
|
||||
and _INPUT_FILE_TYPE.value == constants.INPUT_FILE_TYPE_JSONL
|
||||
):
|
||||
data_converter_vcn_lib.convert_jsonl_to_tfrecord(
|
||||
_INPUT_FILE_PATH.value,
|
||||
_OUTPUT_DIR.value,
|
||||
_OUTPUT_FPS.value,
|
||||
split_ratio,
|
||||
num_shard,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'File format {_INPUT_FILE_TYPE.value} is not supported for'
|
||||
f' {_OBJECTIVE.value}.'
|
||||
)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(main)
|
||||
+289
@@ -0,0 +1,289 @@
|
||||
"""Converts VCN CSV/JSONL files to TFRecord with apache beam."""
|
||||
|
||||
import json
|
||||
from os import path
|
||||
from typing import Any, Dict, Iterator, Sequence, Union, cast
|
||||
|
||||
from absl import logging
|
||||
import apache_beam as beam
|
||||
from apache_beam.io import tfrecordio
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import tensorflow as tf
|
||||
|
||||
from data_converter import common_lib
|
||||
from util import constants
|
||||
|
||||
|
||||
_COLUMN_NAMES = [
|
||||
common_lib.COLUMN_NAME_ML_USE,
|
||||
common_lib.COLUMN_NAME_GCS_FILE_PATH,
|
||||
common_lib.COLUMN_NAME_LABEL,
|
||||
common_lib.COLUMN_NAME_START_SEC,
|
||||
common_lib.COLUMN_NAME_END_SEC,
|
||||
]
|
||||
_JSON_GCS_URI_KEY = 'videoGcsUri'
|
||||
_JSON_CLASS_ANNOTATION_KEY = 'timeSegmentAnnotations'
|
||||
_JSON_CLASS_NAME_KEY = 'displayName'
|
||||
_JSON_START_TIME_KEY = 'startTime'
|
||||
_JSON_END_TIME_KEY = 'endTime'
|
||||
_JSON_RESOURCE_LABEL_KEY = 'dataItemResourceLabels'
|
||||
_JSON_ML_USE_KEY = 'aiplatform.googleapis.com/ml_use'
|
||||
|
||||
|
||||
def build_tf_example(
|
||||
video_uri: str,
|
||||
label: int,
|
||||
start_sec: float,
|
||||
end_sec: float,
|
||||
output_fps: int,
|
||||
) -> tf.train.SequenceExample:
|
||||
"""Builds a TF Example from a video clip.
|
||||
|
||||
Args:
|
||||
video_uri: GCS URI to the video file.
|
||||
label: Class label as an integer.
|
||||
start_sec: Start timestamp of the video clip in seconds.
|
||||
end_sec: End timestamp of the video clip in seconds.
|
||||
output_fps: The output frame rate per second.
|
||||
|
||||
Returns:
|
||||
The created TF Example.
|
||||
"""
|
||||
frame_bytes = common_lib.encode_video(
|
||||
video_uri, start_sec, end_sec, output_fps, image_format='jpg'
|
||||
)
|
||||
seq_example = tf.train.SequenceExample()
|
||||
seq_example.context.feature['clip/label/index'].int64_list.value[:] = [label]
|
||||
for frame in frame_bytes:
|
||||
seq_example.feature_lists.feature_list.get_or_create(
|
||||
'image/encoded'
|
||||
).feature.add().bytes_list.value[:] = [frame]
|
||||
|
||||
return seq_example
|
||||
|
||||
|
||||
class AcquireTFExampleDoFn(beam.DoFn):
|
||||
"""Beam DoFn to build TF Examples from a DataFrame row dict for VCN."""
|
||||
|
||||
def __init__(self, output_fps: int):
|
||||
self._success_counter = beam.metrics.Metrics.counter(
|
||||
self.__class__.__name__, 'Success'
|
||||
)
|
||||
self._failure_counter = beam.metrics.Metrics.counter(
|
||||
self.__class__.__name__, 'Failure'
|
||||
)
|
||||
self._output_fps = output_fps
|
||||
|
||||
def process(
|
||||
self, element: Dict[str, Union[float, int, str]]
|
||||
) -> Iterator[tf.train.SequenceExample]:
|
||||
ml_use: str = cast(str, element[common_lib.COLUMN_NAME_ML_USE])
|
||||
video_uri: str = cast(str, element[common_lib.COLUMN_NAME_GCS_FILE_PATH])
|
||||
|
||||
try:
|
||||
label: int = int(element[common_lib.COLUMN_NAME_LABEL])
|
||||
start_sec: float = float(element[common_lib.COLUMN_NAME_START_SEC])
|
||||
end_sec: float = float(element[common_lib.COLUMN_NAME_END_SEC])
|
||||
|
||||
tf_example = build_tf_example(
|
||||
video_uri,
|
||||
label,
|
||||
start_sec,
|
||||
end_sec,
|
||||
self._output_fps,
|
||||
)
|
||||
self._success_counter.inc()
|
||||
yield beam.pvalue.TaggedOutput(ml_use, tf_example)
|
||||
except (ValueError, IOError) as err:
|
||||
logging.error('Failed to process %s', video_uri)
|
||||
logging.exception(err)
|
||||
self._failure_counter.inc()
|
||||
|
||||
|
||||
def _run_convert_pipeline(
|
||||
output_dir: str,
|
||||
df: pd.DataFrame,
|
||||
num_shards: Sequence[int],
|
||||
output_fps: int,
|
||||
) -> None:
|
||||
"""Starts a Beam pipeline to write DataFrame as TF Records.
|
||||
|
||||
Args:
|
||||
output_dir: TF Records output directory.
|
||||
df: DataFrame to convert from.
|
||||
num_shards: Number of shards for train/validation/test TFRecord files.
|
||||
output_fps: The output frame rate per second.
|
||||
"""
|
||||
clip_list = df.to_dict('records')
|
||||
|
||||
def pipeline(root):
|
||||
train, val, test = (
|
||||
root
|
||||
| 'Create PCollection' >> beam.Create(clip_list)
|
||||
| 'Convert to TF Example'
|
||||
>> beam.ParDo(AcquireTFExampleDoFn(output_fps)).with_outputs(
|
||||
constants.ML_USE_TRAINING,
|
||||
constants.ML_USE_VALIDATION,
|
||||
constants.ML_USE_TEST,
|
||||
)
|
||||
)
|
||||
_ = train | 'Save train TF Record' >> tfrecordio.WriteToTFRecord(
|
||||
path.join(output_dir, common_lib.TRAIN_TFRECORD_NAME),
|
||||
coder=beam.coders.ProtoCoder(tf.train.Example),
|
||||
num_shards=num_shards[0],
|
||||
)
|
||||
_ = val | 'Save val TF Record' >> tfrecordio.WriteToTFRecord(
|
||||
path.join(output_dir, common_lib.VALIDATION_TFRECORD_NAME),
|
||||
coder=beam.coders.ProtoCoder(tf.train.Example),
|
||||
num_shards=num_shards[1],
|
||||
)
|
||||
_ = test | 'Save test TF Record' >> tfrecordio.WriteToTFRecord(
|
||||
path.join(output_dir, common_lib.TEST_TFRECORD_NAME),
|
||||
coder=beam.coders.ProtoCoder(tf.train.Example),
|
||||
num_shards=num_shards[2],
|
||||
)
|
||||
|
||||
common_lib.run_beam_pipeline(pipeline)
|
||||
|
||||
|
||||
def _convert_df_to_tfrecord(
|
||||
df: pd.DataFrame,
|
||||
output_dir: str,
|
||||
split_ratio: Sequence[float],
|
||||
num_shard: Sequence[int],
|
||||
output_fps: int,
|
||||
) -> None:
|
||||
"""Converts a DataFrame into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
Args:
|
||||
df: DataFrame to convert.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
split_ratio: List specifying the training, validation, and testing splits
|
||||
for unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
output_fps: The output frame rate per second.
|
||||
"""
|
||||
# Replaces ml_use with common_lib string constants for consistency.
|
||||
common_lib.format_ml_use_column(df)
|
||||
common_lib.insert_missing_ml_use(df)
|
||||
|
||||
# Ignores invalid rows.
|
||||
dropped_row_num = common_lib.drop_invalid_rows(df)
|
||||
if dropped_row_num > 0:
|
||||
logging.warning('Ignored %d invalid rows.', dropped_row_num)
|
||||
|
||||
common_lib.replace_unassigned_ml_use(
|
||||
df[common_lib.COLUMN_NAME_ML_USE], split_ratio
|
||||
)
|
||||
|
||||
# Converts labels to integers as required by training.
|
||||
new_labels, label_map = common_lib.create_label_map(
|
||||
df[common_lib.COLUMN_NAME_LABEL]
|
||||
)
|
||||
df[common_lib.COLUMN_NAME_LABEL] = new_labels
|
||||
label_map_path = path.join(output_dir, common_lib.LABEL_MAP_NAME)
|
||||
logging.info('Writing label map to %s.', label_map_path)
|
||||
common_lib.write_label_map(label_map_path, label_map)
|
||||
|
||||
# Missing start / end times are treated as 0, inf, respectively.
|
||||
df[common_lib.COLUMN_NAME_START_SEC].fillna(0, inplace=True)
|
||||
df[common_lib.COLUMN_NAME_END_SEC].fillna(np.inf, inplace=True)
|
||||
|
||||
_run_convert_pipeline(output_dir, df, num_shard, output_fps)
|
||||
|
||||
|
||||
def convert_csv_to_tfrecord(
|
||||
input_csv: str,
|
||||
output_dir: str,
|
||||
output_fps: int,
|
||||
split_ratio: Sequence[float] = (0.8, 0.1, 0.1),
|
||||
num_shard: Sequence[int] = (10, 10, 10),
|
||||
) -> None:
|
||||
"""Parses input_csv file into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
The csv format is shown in
|
||||
https://cloud.google.com/vertex-ai/docs/video-data/classification/prepare-data#csv
|
||||
|
||||
If an ml_use column is not provided, one will be created.
|
||||
|
||||
label_map.yaml containing the label map will be placed in output_dir.
|
||||
|
||||
Args:
|
||||
input_csv: Name of the csv file.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
output_fps: The output frame rate per second.
|
||||
split_ratio: List specifying the training, validation, and testing splits
|
||||
for unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
with tf.io.gfile.GFile(input_csv, 'r') as f:
|
||||
df: pd.DataFrame = pd.read_csv(
|
||||
f, header=None, names=_COLUMN_NAMES, on_bad_lines='warn'
|
||||
)
|
||||
|
||||
_convert_df_to_tfrecord(df, output_dir, split_ratio, num_shard, output_fps)
|
||||
|
||||
|
||||
def convert_jsonl_to_tfrecord(
|
||||
input_jsonl: str,
|
||||
output_dir: str,
|
||||
output_fps: int,
|
||||
split_ratio: Sequence[float] = (0.8, 0.1, 0.1),
|
||||
num_shard: Sequence[int] = (10, 10, 10),
|
||||
) -> None:
|
||||
"""Parses input_jsonl file into three separate tfrecords for training, validation, and testing into output_dir.
|
||||
|
||||
The JSONL format is shown in
|
||||
https://cloud.google.com/vertex-ai/docs/video-data/classification/prepare-data#jsonl.
|
||||
|
||||
If an ml_use column is not provided, one will be created.
|
||||
|
||||
label_map.yaml containing the label map will be placed in output_dir.
|
||||
|
||||
Args:
|
||||
input_jsonl: Name of the JSONL file.
|
||||
output_dir: The directory to save TFRecords and label_map.yaml.
|
||||
output_fps: The output frame rate per second.
|
||||
split_ratio: List specifying the training, validation, and testing splits
|
||||
for unassigned TFRecords.
|
||||
num_shard: Number of shards for train/validation/test TFRecord files.
|
||||
"""
|
||||
df_rows = []
|
||||
with tf.io.gfile.GFile(input_jsonl, 'r') as f:
|
||||
lines = f.read().rstrip().splitlines()
|
||||
|
||||
for i, line in enumerate(lines, 1):
|
||||
try:
|
||||
item: Dict[str, Any] = json.loads(line)
|
||||
|
||||
gcs_uri = item.get(_JSON_GCS_URI_KEY)
|
||||
if not gcs_uri:
|
||||
logging.warning('Invalid JSON at line %d, skipped.', i)
|
||||
continue
|
||||
|
||||
annotations = item.get(_JSON_CLASS_ANNOTATION_KEY, [])
|
||||
ml_use = item.get(_JSON_RESOURCE_LABEL_KEY, {}).get(
|
||||
_JSON_ML_USE_KEY, common_lib.ML_USE_UNASSIGNED
|
||||
)
|
||||
|
||||
for j, annotation in enumerate(annotations):
|
||||
label = annotation.get(_JSON_CLASS_NAME_KEY)
|
||||
if not label:
|
||||
logging.warning('Invalid annotation #%d at line %d, skipped.', j, i)
|
||||
continue
|
||||
# The example in external documentation uses strings like "1.0s", so we
|
||||
# need to remove the "s" suffix.
|
||||
start_time = annotation.get(_JSON_START_TIME_KEY, '0').removesuffix('s')
|
||||
end_time = annotation.get(_JSON_END_TIME_KEY, 'inf').removesuffix('s')
|
||||
df_rows.append([ml_use, gcs_uri, label, start_time, end_time])
|
||||
except (json.JSONDecodeError, AttributeError):
|
||||
logging.warning('Invalid JSON at line %d, skipped.', i)
|
||||
continue
|
||||
|
||||
df = pd.DataFrame(
|
||||
data=df_rows,
|
||||
columns=_COLUMN_NAMES,
|
||||
)
|
||||
|
||||
_convert_df_to_tfrecord(df, output_dir, split_ratio, num_shard, output_fps)
|
||||
+50
@@ -0,0 +1,50 @@
|
||||
FROM python:3.9
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Install basic libs.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake \
|
||||
curl \
|
||||
wget \
|
||||
sudo \
|
||||
gnupg \
|
||||
python3-opencv \
|
||||
lsb-release \
|
||||
ca-certificates \
|
||||
build-essential \
|
||||
git \
|
||||
vim \
|
||||
screen \
|
||||
libportaudio2 \
|
||||
libusb-1.0-0-dev \
|
||||
openjdk-17-jre
|
||||
|
||||
# Add gcsfuse distribution URL as a package source and import its public key.
|
||||
RUN echo "deb https://packages.cloud.google.com/apt gcsfuse-`lsb_release -c -s` main" | sudo tee /etc/apt/sources.list.d/gcsfuse.list
|
||||
RUN curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | sudo apt-key add -
|
||||
|
||||
# Install gcsfuse.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends gcsfuse
|
||||
|
||||
# Install google cloud SDK.
|
||||
RUN wget -q https://dl.google.com/dl/cloudsdk/channels/rapid/downloads/google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN tar xzf google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN ./google-cloud-sdk/install.sh -q
|
||||
# Make sure gsutil will use the default service account.
|
||||
RUN echo '[GoogleCompute]\nservice_account = default' > /etc/boto.cfg
|
||||
|
||||
|
||||
# Install required libs.
|
||||
RUN pip install --upgrade pip
|
||||
RUN pip install pyyaml==5.4.1
|
||||
RUN pip install pycocotools==2.0.6
|
||||
RUN pip install opencv-python-headless==4.7.0.72
|
||||
RUN pip install numpy==1.24.2
|
||||
RUN pip install pandas==1.5.3
|
||||
RUN pip install Pillow==9.4.0
|
||||
RUN pip install apache-beam[gcp]==2.45.0
|
||||
RUN pip install object-detection==0.0.3
|
||||
RUN pip install google-cloud-storage==1.42.3
|
||||
RUN pip install gcsfs==2021.10.1
|
||||
RUN pip install pylint==2.17.2
|
||||
+23
@@ -0,0 +1,23 @@
|
||||
FROM gcr.io/automl-migration-test/automl-vision-data-converter-base:latest
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
COPY model_oss/data_converter /automl_vision/data_converter
|
||||
COPY model_oss/util /automl_vision/util
|
||||
|
||||
WORKDIR /automl_vision
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision"
|
||||
|
||||
# Run pylint to validate code.
|
||||
COPY .pylintrc /automl_vision/.pylintrc
|
||||
RUN find . -type f -name "*.py" | xargs pylint --rcfile=./.pylintrc --errors-only
|
||||
|
||||
ENTRYPOINT ["python3","data_converter/data_converter_main.py"]
|
||||
|
||||
CMD ["--input_file_path=YOUR_INPUT_FILE",\
|
||||
"--input_file_type=csv",\
|
||||
"--objective=iod",\
|
||||
"--output_dir=YOUR_OUTPUT_DIR",\
|
||||
"--num_shard=10,10,10",\
|
||||
"--split_ratio=0.8,0.1,0.1"]
|
||||
+90
@@ -0,0 +1,90 @@
|
||||
# Dockerfile for Detectron2 serving.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/detectron2/dockerfile/serving.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
FROM pytorch/torchserve:0.7.0-cpu
|
||||
|
||||
USER root
|
||||
|
||||
# Install tools.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
wget \
|
||||
vim
|
||||
|
||||
# run and update some basic packages software packages, including security libs
|
||||
RUN apt-get update && apt-get install -y \
|
||||
software-properties-common && \
|
||||
add-apt-repository -y ppa:ubuntu-toolchain-r/test && \
|
||||
apt-get update && apt-get install -y \
|
||||
gcc-9 g++-9 apt-transport-https ca-certificates gnupg curl
|
||||
|
||||
# Install gcloud tools for gsutil as well as debugging
|
||||
RUN echo "deb [signed-by=/usr/share/keyrings/cloud.google.gpg] http://packages.cloud.google.com/apt cloud-sdk main" | \
|
||||
tee -a /etc/apt/sources.list.d/google-cloud-sdk.list && \
|
||||
curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | \
|
||||
apt-key --keyring /usr/share/keyrings/cloud.google.gpg add - && \
|
||||
apt-get update -y && apt-get install google-cloud-sdk -y
|
||||
|
||||
USER model-server
|
||||
|
||||
# install detectron2 dependencies
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
RUN python3 -m pip install --user numpy==1.24.2
|
||||
RUN python3 -m pip install --user opencv-python==4.7.0.72
|
||||
RUN python3 -m pip install --user 'git+https://github.com/facebookresearch/detectron2.git@v0.6'
|
||||
|
||||
# Install GCS storage library.
|
||||
RUN pip install google-cloud-storage==2.6.0
|
||||
|
||||
# For mask encoding.
|
||||
RUN pip install --upgrade pycocotools==2.0.6
|
||||
|
||||
ARG MODEL_NAME=detectron2_serving
|
||||
ENV MODEL_NAME="${MODEL_NAME}"
|
||||
|
||||
# health and prediction listener ports
|
||||
ARG AIP_HTTP_PORT=7080
|
||||
ENV AIP_HTTP_PORT="${AIP_HTTP_PORT}"
|
||||
|
||||
ARG MODEL_MGMT_PORT=7081
|
||||
|
||||
# expose health and prediction listener ports from the image
|
||||
EXPOSE "${AIP_HTTP_PORT}"
|
||||
EXPOSE "${MODEL_MGMT_PORT}"
|
||||
EXPOSE 8080 8081 8082 7070 7071
|
||||
|
||||
# create torchserve configuration file
|
||||
USER root
|
||||
RUN echo "service_envelope=json\n" \
|
||||
"inference_address=http://0.0.0.0:${AIP_HTTP_PORT}\n" \
|
||||
"management_address=http://0.0.0.0:${MODEL_MGMT_PORT}" >> /home/model-server/config.properties
|
||||
USER model-server
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Copy model artifacts.
|
||||
COPY ./model_oss/detectron2/handler.py /home/model-server/handler.py
|
||||
WORKDIR /home/model-server/
|
||||
|
||||
# Create model archive file packaging model artifacts and dependencies.
|
||||
# Note(lavrai): The model `.pth` file and `cfg.yaml` file will be set by the
|
||||
# customer as an environment variable and will be later loaded by the
|
||||
# `handler.py` file.
|
||||
RUN torch-model-archiver \
|
||||
--model-name="${MODEL_NAME}" \
|
||||
--version=1.0 \
|
||||
--handler=/home/model-server/handler.py \
|
||||
--export-path=/home/model-server/model-store \
|
||||
-f
|
||||
|
||||
# run Torchserve HTTP serve to respond to prediction requests
|
||||
CMD ["ls", "-ltr", "/home/model-server/model-store/", ";", \
|
||||
"torchserve", "--start", "--ts-config=/home/model-server/config.properties", \
|
||||
"--models", "${MODEL_NAME}=${MODEL_NAME}.mar", \
|
||||
"--model-store", "/home/model-server/model-store"]
|
||||
+100
@@ -0,0 +1,100 @@
|
||||
# Dockerfile for Detectron2 training.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/detectron2/dockerfile/train.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM nvidia/cuda:11.1.1-cudnn8-devel-ubuntu18.04
|
||||
# Using an older system (18.04) to avoid opencv incompatibility (issue#3524).
|
||||
|
||||
ENV DEBIAN_FRONTEND noninteractive
|
||||
RUN apt-get update && apt-get install -y \
|
||||
python3.7 python3.7-dev python3.7-distutils \
|
||||
python3-opencv ca-certificates git wget sudo ninja-build \
|
||||
curl wget vim
|
||||
|
||||
# Make python3 available for python3.7.
|
||||
RUN update-alternatives --install /usr/bin/python3 python3 /usr/bin/python3.6 1
|
||||
RUN update-alternatives --install /usr/bin/python3 python3 /usr/bin/python3.7 2
|
||||
RUN update-alternatives --config python3
|
||||
# Make python available for python3.7.
|
||||
RUN ln -sv /usr/bin/python3.7 /usr/bin/python
|
||||
|
||||
# Create a non-root user.
|
||||
ARG USER_ID=1000
|
||||
RUN useradd -m --no-log-init --system --uid ${USER_ID} appuser -g sudo
|
||||
RUN echo '%sudo ALL=(ALL) NOPASSWD:ALL' >> /etc/sudoers
|
||||
USER appuser
|
||||
WORKDIR /home/appuser
|
||||
|
||||
ENV PATH="/home/appuser/.local/bin:${PATH}"
|
||||
RUN wget https://bootstrap.pypa.io/pip/get-pip.py && \
|
||||
python3.7 get-pip.py --user && \
|
||||
rm get-pip.py
|
||||
|
||||
# Important! Otherwise, it uses existing numpy from host-modules
|
||||
# which throws error.
|
||||
RUN pip install --user numpy==1.20.3
|
||||
|
||||
# Install dependencies:
|
||||
# See https://pytorch.org/ for other options if you use
|
||||
# a different version of CUDA.
|
||||
RUN pip install --user tensorboard==2.11.0
|
||||
# cmake from apt-get is too old.
|
||||
RUN pip install --user cmake==3.25.2
|
||||
RUN pip install --user torch==1.10.0+cu111 torchvision==0.11.0+cu111 -f https://download.pytorch.org/whl/torch_stable.html
|
||||
RUN pip install --user setuptools==59.5.0
|
||||
RUN pip install --user opencv-python==4.7.0.72
|
||||
RUN pip install --user cloudml-hypertune==0.1.0.dev6
|
||||
RUN pip install --user fvcore==0.1.5.post20221221
|
||||
# Install detectron2.
|
||||
RUN git clone -b v0.6 https://github.com/facebookresearch/detectron2 detectron2_repo
|
||||
# Set FORCE_CUDA because during `docker build` cuda is not accessible.
|
||||
ENV FORCE_CUDA="1"
|
||||
# This will by default build detectron2 for all common cuda
|
||||
# architectures and take a lot more time,
|
||||
# because inside `docker build`, there is no way to tell
|
||||
# which architecture will be used.
|
||||
ARG TORCH_CUDA_ARCH_LIST="Kepler;Kepler+Tesla;Maxwell;Maxwell+Tegra;Pascal;Volta;Turing"
|
||||
ENV TORCH_CUDA_ARCH_LIST="${TORCH_CUDA_ARCH_LIST}"
|
||||
RUN pip install --user -e detectron2_repo
|
||||
|
||||
# Set a fixed model cache directory.
|
||||
ENV FVCORE_CACHE="/tmp"
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Copy model-garden detectron2 files to '/home/appuser/trainer' folder.
|
||||
ADD ./model_oss/detectron2 /home/appuser/trainer
|
||||
|
||||
################ Copy plain_train_net.py to task.py and
|
||||
# then modify it using sed commands. ###################
|
||||
# Src: https://github.com/facebookresearch/detectron2/blob/v0.6/tools/plain_train_net.py
|
||||
RUN sudo cp /home/appuser/detectron2_repo/tools/plain_train_net.py /home/appuser/trainer/task.py
|
||||
# Make additional changes to task.py.
|
||||
# Note(lavrai): Start adding SED commands from end of file towards the top
|
||||
# so that the line numbers do not keep changing for the source file.
|
||||
# For entry-point:
|
||||
RUN sudo sed -i "214 d" /home/appuser/trainer/task.py
|
||||
RUN sudo sed -i "213 a\ default_arg_parser = default_argument_parser()" /home/appuser/trainer/task.py
|
||||
RUN sudo sed -i "214 a\ extended_parser = trainer_utils.extend_parser_arguments(default_arg_parser)" /home/appuser/trainer/task.py
|
||||
RUN sudo sed -i "215 a\ args = extended_parser.parse_args()" /home/appuser/trainer/task.py
|
||||
# For main() function:
|
||||
RUN sudo sed -i "192 a\ trainer_utils.register_dataset(args)" /home/appuser/trainer/task.py
|
||||
# For setup() function:
|
||||
RUN sudo sed -i "184 a\ cfg.SOLVER.BASE_LR = args.lr" /home/appuser/trainer/task.py
|
||||
RUN sudo sed -i "185 a\ cfg.OUTPUT_DIR = args.output_dir" /home/appuser/trainer/task.py
|
||||
RUN sudo sed -i "186 a\ cfg.MODEL.WEIGHTS = model_zoo.get_checkpoint_url(config_file_copy)" /home/appuser/trainer/task.py
|
||||
RUN sudo sed -i "182 a\ config_file_copy = args.config_file" /home/appuser/trainer/task.py
|
||||
RUN sudo sed -i "183 a\ args.config_file = model_zoo.get_config_file(args.config_file)" /home/appuser/trainer/task.py
|
||||
# For new import:
|
||||
RUN sudo sed -i "27 a\from detectron2 import model_zoo" /home/appuser/trainer/task.py
|
||||
RUN sudo sed -i "21 a\import trainer_utils" /home/appuser/trainer/task.py
|
||||
|
||||
ENV PYTHONPATH /home/appuser/trainer
|
||||
|
||||
ENTRYPOINT ["python", "-m", "trainer.task"]
|
||||
@@ -0,0 +1,154 @@
|
||||
"""Custom handler for Detectron2 serving."""
|
||||
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
from typing import Any, List, Tuple
|
||||
|
||||
import cv2
|
||||
from detectron2.config import get_cfg
|
||||
from detectron2.engine import DefaultPredictor
|
||||
from google.cloud import storage
|
||||
import numpy as np
|
||||
import pycocotools.mask as mask_util
|
||||
import torch
|
||||
|
||||
|
||||
def get_bucket_and_blob_name(gcs_filepath: str) -> Tuple[str, str]:
|
||||
"""Gets bucket and blob name from gcs path."""
|
||||
# The gcs path is of the form gs://<bucket-name>/<blob-name>
|
||||
gs_suffix = gcs_filepath.split("gs://", 1)[1]
|
||||
return tuple(gs_suffix.split("/", 1))
|
||||
|
||||
|
||||
def download_gcs_file(src_file_path: str, dst_file_path: str):
|
||||
"""Downloads gcs-file to local folder."""
|
||||
src_bucket_name, src_blob_name = get_bucket_and_blob_name(src_file_path)
|
||||
client = storage.Client()
|
||||
src_bucket = client.get_bucket(src_bucket_name)
|
||||
src_blob = src_bucket.blob(src_blob_name)
|
||||
src_blob.download_to_filename(dst_file_path)
|
||||
|
||||
|
||||
class ModelHandler:
|
||||
"""Custom model handler for Detectron2."""
|
||||
|
||||
def __init__(self):
|
||||
self.error = None
|
||||
self._batch_size = 0
|
||||
self.initialized = False
|
||||
self.predictor = None
|
||||
self.test_threshold = 0.5
|
||||
|
||||
def initialize(self, context: Any):
|
||||
"""Initialize."""
|
||||
print("context.system_properties: ", context.system_properties)
|
||||
print("context.manifest: ", context.manifest)
|
||||
self.manifest = context.manifest
|
||||
properties = context.system_properties
|
||||
# Get threshold from environment variable.
|
||||
# This will be set by customer.
|
||||
self.test_threshold = float(os.environ.get("TEST_THRESHOLD"))
|
||||
print("test_threshold: ", self.test_threshold)
|
||||
# Get model and config file location from environment variables.
|
||||
# These will be set by customer when doing model upload.
|
||||
gcs_model_file = os.environ["MODEL_PTH_FILE"]
|
||||
gcs_config_file = os.environ["CONFIG_YAML_FILE"]
|
||||
print("Copying gcs_model_file: ", gcs_model_file)
|
||||
print("Copying gcs_config_file: ", gcs_config_file)
|
||||
# Copy these files from GCS location to local file.
|
||||
# Note(lavrai): GCSFuse path does not seem to work here for now.
|
||||
model_file = "./model.pth"
|
||||
config_file = "./cfg.yaml"
|
||||
download_gcs_file(src_file_path=gcs_model_file, dst_file_path=model_file)
|
||||
if not os.path.exists(model_file):
|
||||
raise RuntimeError("Missing model_file: %s" % model_file)
|
||||
download_gcs_file(src_file_path=gcs_config_file, dst_file_path=config_file)
|
||||
if not os.path.exists(config_file):
|
||||
raise RuntimeError("Missing config_file: %s" % config_file)
|
||||
|
||||
# Set up config file.
|
||||
cfg = get_cfg()
|
||||
cfg.merge_from_file(config_file)
|
||||
cfg.MODEL.WEIGHTS = model_file
|
||||
cfg.MODEL.DEVICE = (
|
||||
cfg.MODEL.DEVICE + str(properties.get("gpu_id"))
|
||||
if torch.cuda.is_available()
|
||||
else "cpu"
|
||||
)
|
||||
cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST = self.test_threshold
|
||||
|
||||
# Build predictor from config.
|
||||
self.predictor = DefaultPredictor(cfg)
|
||||
self._batch_size = context.system_properties["batch_size"]
|
||||
self.initialized = True
|
||||
|
||||
def preprocess(self, batch: List[Any]) -> List[Any]:
|
||||
"""Preprocess raw input and return as list of images."""
|
||||
print("Running pre-processing.")
|
||||
images = []
|
||||
for request in batch:
|
||||
request_data = request.get("data")
|
||||
input_bytes = io.BytesIO(request_data)
|
||||
img = cv2.imdecode(np.fromstring(input_bytes.read(), np.uint8), 1)
|
||||
images.append(img)
|
||||
return images
|
||||
|
||||
def inference(self, model_input: List[Any]) -> List[Any]:
|
||||
"""Runs inference."""
|
||||
print("Running model-inference.")
|
||||
return [self.predictor(image) for image in model_input]
|
||||
|
||||
def postprocess(self, inference_result: List[Any]) -> List[Any]:
|
||||
"""Post process inference result."""
|
||||
response_list = []
|
||||
print("Num inference_items are:", len(inference_result))
|
||||
for inference_item in inference_result:
|
||||
predictions = inference_item["instances"].to("cpu")
|
||||
print("Predictions are:", predictions)
|
||||
boxes = None
|
||||
if predictions.has("pred_boxes"):
|
||||
boxes = predictions.pred_boxes.tensor.numpy().tolist()
|
||||
scores = None
|
||||
if predictions.has("scores"):
|
||||
scores = predictions.scores.numpy().tolist()
|
||||
classes = None
|
||||
if predictions.has("pred_classes"):
|
||||
classes = predictions.pred_classes.numpy().tolist()
|
||||
masks_rle = None
|
||||
if predictions.has("pred_masks"):
|
||||
# Do run length encoding, else the mask output becomes huge.
|
||||
masks_rle = [
|
||||
mask_util.encode(np.asfortranarray(mask))
|
||||
for mask in predictions.pred_masks
|
||||
]
|
||||
for rle in masks_rle:
|
||||
rle["counts"] = rle["counts"].decode("utf-8")
|
||||
response = {
|
||||
"classes": classes,
|
||||
"scores": scores,
|
||||
"boxes": boxes,
|
||||
"masks_rle": masks_rle,
|
||||
}
|
||||
response_list.append(json.dumps(response))
|
||||
print("response_list: ", response_list)
|
||||
return response_list
|
||||
|
||||
def handle(self, data: Any, context: Any) -> List[Any]: # pylint: disable=unused-argument
|
||||
"""Runs preprocess, inference, and post-processing."""
|
||||
model_input = self.preprocess(data)
|
||||
model_out = self.inference(model_input)
|
||||
output = self.postprocess(model_out)
|
||||
print("Done handling input.")
|
||||
return output
|
||||
|
||||
|
||||
_service = ModelHandler()
|
||||
|
||||
|
||||
def handle(data: Any, context: Any) -> List[Any]:
|
||||
if not _service.initialized:
|
||||
_service.initialize(context)
|
||||
if data is None:
|
||||
return None
|
||||
return _service.handle(data, context)
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Detectron2 trainer helper functions."""
|
||||
|
||||
import argparse
|
||||
from detectron2.data.datasets import register_coco_instances
|
||||
|
||||
|
||||
def extend_parser_arguments(
|
||||
parser: argparse.ArgumentParser,
|
||||
) -> argparse.ArgumentParser:
|
||||
"""Adds additional model-garden related arguments."""
|
||||
parser.add_argument(
|
||||
"--train_dataset_name",
|
||||
required=False,
|
||||
default="",
|
||||
type=str,
|
||||
help=(
|
||||
"The training dataset name for registration. "
|
||||
"For example: 'balloon_train'."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_coco_json_file",
|
||||
required=False,
|
||||
default="",
|
||||
type=str,
|
||||
help="The path to the training coco-json format file.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_image_root",
|
||||
required=False,
|
||||
default="",
|
||||
type=str,
|
||||
help="The path to the root folder containing the training images.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--val_dataset_name",
|
||||
required=False,
|
||||
default="",
|
||||
type=str,
|
||||
help=(
|
||||
"The validation dataset name for registration. "
|
||||
"For example: 'balloon_val'."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--val_coco_json_file",
|
||||
required=False,
|
||||
default="",
|
||||
type=str,
|
||||
help="The path to the validation coco-json format file.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--val_image_root",
|
||||
required=False,
|
||||
default="",
|
||||
type=str,
|
||||
help="The path to the root folder containing the validation images.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
required=True,
|
||||
type=str,
|
||||
help="The path to the output directory.",
|
||||
)
|
||||
# Add hyper-parameter tuning related variables.
|
||||
parser.add_argument(
|
||||
"--lr",
|
||||
type=float,
|
||||
default=0.00025,
|
||||
help="The learning rate to be tuned.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--hp_eval_task",
|
||||
type=str,
|
||||
choices=["bbox", "segm"],
|
||||
default="bbox",
|
||||
help="The task choice for HP tuning.",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def register_dataset(args: argparse.Namespace):
|
||||
"""Register the input dataset in Detectron2 Coco format."""
|
||||
if args.train_dataset_name:
|
||||
register_coco_instances(
|
||||
name=args.train_dataset_name,
|
||||
metadata={},
|
||||
json_file=args.train_coco_json_file,
|
||||
image_root=args.train_image_root,
|
||||
)
|
||||
if args.val_dataset_name:
|
||||
register_coco_instances(
|
||||
name=args.val_dataset_name,
|
||||
metadata={},
|
||||
json_file=args.val_coco_json_file,
|
||||
image_root=args.val_image_root,
|
||||
)
|
||||
@@ -4,9 +4,10 @@
|
||||
# pylint: disable=logging-fstring-interpolation
|
||||
|
||||
import base64
|
||||
import io
|
||||
import logging
|
||||
import os
|
||||
from typing import Any, List, Tuple
|
||||
from typing import Any, List, Sequence, Tuple
|
||||
|
||||
from diffusers import ControlNetModel
|
||||
from diffusers import DiffusionPipeline
|
||||
@@ -20,6 +21,7 @@ from diffusers import StableDiffusionPipeline
|
||||
from diffusers import StableDiffusionUpscalePipeline
|
||||
from diffusers import TextToVideoZeroPipeline
|
||||
from diffusers import UniPCMultistepScheduler
|
||||
import imageio
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import torch
|
||||
@@ -43,6 +45,13 @@ TEXT_TO_VIDEO_ZERO_SHOT = "text-to-video-zero-shot"
|
||||
TEXT_TO_VIDEO = "text-to-video"
|
||||
|
||||
|
||||
def frames_to_video_bytes(frames: Sequence[np.ndarray], fps: int) -> bytes:
|
||||
images = [Image.fromarray(array) for array in frames]
|
||||
io_obj = io.BytesIO()
|
||||
imageio.mimsave(io_obj, images, format=".mp4", fps=fps)
|
||||
return io_obj.getvalue()
|
||||
|
||||
|
||||
class DiffusersHandler(BaseHandler):
|
||||
"""Custom handler for TIMM models."""
|
||||
|
||||
@@ -214,7 +223,7 @@ class DiffusersHandler(BaseHandler):
|
||||
numpy_arrays = self.pipeline(prompt=prompt).images
|
||||
numpy_arrays = [(i * 255).astype("uint8") for i in numpy_arrays]
|
||||
videos.append(
|
||||
video_format_converter.frames_to_video_bytes(numpy_arrays, fps=4)
|
||||
frames_to_video_bytes(numpy_arrays, fps=4)
|
||||
)
|
||||
return videos
|
||||
elif self.task == TEXT_TO_VIDEO:
|
||||
@@ -224,7 +233,7 @@ class DiffusersHandler(BaseHandler):
|
||||
# Therefore we need to split the output into different videos.
|
||||
predicted_images = np.array_split(predicted_images, len(prompts), axis=2)
|
||||
videos = [
|
||||
video_format_converter.frames_to_video_bytes(images, fps=8)
|
||||
frames_to_video_bytes(images, fps=8)
|
||||
for images in predicted_images
|
||||
]
|
||||
return videos
|
||||
|
||||
+1
File diff suppressed because one or more lines are too long
+1
File diff suppressed because one or more lines are too long
+151
@@ -0,0 +1,151 @@
|
||||
# This Dockerfile converts JAX vision transformer model to
|
||||
# tensorflow saved model format.
|
||||
# Here is an example to build this dockerfile:
|
||||
# PROJECT="your gcp project"
|
||||
# IMAGE_TAG="jax-vit-model-conversion:${USER}-test"
|
||||
# docker build -f model_oss/jax_vision_transformer/dockerfile/jax_vit_model_conversion.Dockerfile . -t "${IMAGE_TAG}"
|
||||
# docker tag "${IMAGE_TAG}" "gcr.io/${PROJECT}/${IMAGE_TAG}"
|
||||
# docker push "gcr.io/${PROJECT}/${IMAGE_TAG}"
|
||||
|
||||
|
||||
FROM tensorflow/tensorflow:2.12.0-gpu
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Install basic libs
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake \
|
||||
curl \
|
||||
wget \
|
||||
sudo \
|
||||
gnupg \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender-dev \
|
||||
lsb-release \
|
||||
ca-certificates \
|
||||
build-essential \
|
||||
git
|
||||
|
||||
# Copy Apache license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Get 'vision_transformer' repository from github.
|
||||
RUN git clone https://github.com/google-research/vision_transformer
|
||||
# Set current directory to the downloaded 'vision_transformer' repository.
|
||||
WORKDIR ./vision_transformer
|
||||
# Using git reset command to pin it down to a specific version.
|
||||
RUN git reset --hard e66b4732d44504251197a3da3f5949f3f3ce9ca6
|
||||
|
||||
# Install required libs
|
||||
RUN pip install --upgrade pip
|
||||
# The following pip installs are pinned down versions of those inside
|
||||
# vit_jax/requirements.txt file.
|
||||
# NOTE: Using `no-deps` flag to avoid overwriting of
|
||||
# dependent library versions. For example,
|
||||
# both `chex` and `jax` can overwrite each others
|
||||
# `jax-lib` version.
|
||||
RUN pip install --no-deps absl-py==1.4.0
|
||||
RUN pip install --no-deps aqtp==0.0.10
|
||||
RUN pip install --no-deps array-record==0.2.0
|
||||
RUN pip install --no-deps astunparse==1.6.3
|
||||
RUN pip install --no-deps cached-property==1.5.2
|
||||
RUN pip install --no-deps cachetools==5.3.0
|
||||
RUN pip install --no-deps certifi==2019.11.28
|
||||
RUN pip install --no-deps chardet==3.0.4
|
||||
RUN pip install --no-deps chex==0.1.7
|
||||
RUN pip install --no-deps click==8.1.3
|
||||
RUN pip install --no-deps cloudpickle==2.2.1
|
||||
RUN pip install --no-deps clu==0.0.9
|
||||
RUN pip install --no-deps contextlib2==21.6.0
|
||||
RUN pip install --no-deps dacite==1.8.1
|
||||
RUN pip install --no-deps dbus-python==1.2.16
|
||||
RUN pip install --no-deps decorator==5.1.1
|
||||
RUN pip install --no-deps dm-tree==0.1.8
|
||||
RUN pip install --no-deps einops==0.6.1
|
||||
RUN pip install --no-deps etils==1.3.0
|
||||
RUN pip install --no-deps flatbuffers==23.3.3
|
||||
RUN pip install --no-deps flax==0.6.10
|
||||
RUN pip install --no-deps git+https://github.com/google/flaxformer@9adaa4467cf17703949b9f537c3566b99de1b416
|
||||
RUN pip install --no-deps gast==0.4.0
|
||||
RUN pip install --no-deps google-auth==2.16.2
|
||||
RUN pip install --no-deps google-auth-oauthlib==0.4.6
|
||||
RUN pip install --no-deps google-pasta==0.2.0
|
||||
RUN pip install --no-deps googleapis-common-protos==1.59.0
|
||||
RUN pip install --no-deps grpcio==1.51.3
|
||||
RUN pip install --no-deps h5py==3.8.0
|
||||
RUN pip install --no-deps idna==2.8
|
||||
RUN pip install --no-deps importlib-metadata==6.1.0
|
||||
RUN pip install --no-deps importlib-resources==5.12.0
|
||||
RUN pip install --no-deps keras==2.12.0
|
||||
RUN pip install --no-deps libclang==16.0.0
|
||||
RUN pip install --no-deps Markdown==3.4.3
|
||||
RUN pip install --no-deps markdown-it-py==2.2.0
|
||||
RUN pip install --no-deps MarkupSafe==2.1.2
|
||||
RUN pip install --no-deps mdurl==0.1.2
|
||||
RUN pip install --no-deps ml-collections==0.1.1
|
||||
RUN pip install --no-deps msgpack==1.0.5
|
||||
RUN pip install --no-deps nest-asyncio==1.5.6
|
||||
RUN pip install --no-deps numpy==1.23.5
|
||||
RUN pip install --no-deps oauthlib==3.2.2
|
||||
RUN pip install --no-deps opt-einsum==3.3.0
|
||||
RUN pip install --no-deps optax==0.1.5
|
||||
RUN pip install --no-deps orbax-checkpoint==0.1.6
|
||||
RUN pip install --no-deps packaging==23.0
|
||||
RUN pip install --no-deps pandas==2.0.1
|
||||
RUN pip install --no-deps pip==23.1.2
|
||||
RUN pip install --no-deps promise==2.3
|
||||
RUN pip install --no-deps protobuf==4.22.1
|
||||
RUN pip install --no-deps psutil==5.9.5
|
||||
RUN pip install --no-deps pyasn1==0.4.8
|
||||
RUN pip install --no-deps pyasn1-modules==0.2.8
|
||||
RUN pip install --no-deps Pygments==2.15.1
|
||||
RUN pip install --no-deps PyGObject==3.36.0
|
||||
RUN pip install --no-deps python-apt==2.0.1+ubuntu0.20.4.1
|
||||
RUN pip install --no-deps python-dateutil==2.8.2
|
||||
RUN pip install --no-deps pytz==2023.3
|
||||
RUN pip install --no-deps PyYAML==6.0
|
||||
RUN pip install --no-deps requests==2.22.0
|
||||
RUN pip install --no-deps requests-oauthlib==1.3.1
|
||||
RUN pip install --no-deps requests-unixsocket==0.2.0
|
||||
RUN pip install --no-deps rich==13.3.5
|
||||
RUN pip install --no-deps rsa==4.9
|
||||
RUN pip install --no-deps scipy==1.10.1
|
||||
RUN pip install --no-deps setuptools==67.6.0
|
||||
RUN pip install --no-deps six==1.14.0
|
||||
RUN pip install --no-deps tensorboard==2.12.0
|
||||
RUN pip install --no-deps tensorboard-data-server==0.7.0
|
||||
RUN pip install --no-deps tensorboard-plugin-wit==1.8.1
|
||||
RUN pip install --no-deps tensorflow==2.12.0
|
||||
RUN pip install --no-deps tensorflow-cpu==2.12.0
|
||||
RUN pip install --no-deps tensorflow-datasets==4.9.2
|
||||
RUN pip install --no-deps tensorflow-estimator==2.12.0
|
||||
RUN pip install --no-deps tensorflow-hub==0.13.0
|
||||
RUN pip install --no-deps tensorflow-io-gcs-filesystem==0.31.0
|
||||
RUN pip install --no-deps tensorflow-metadata==1.13.1
|
||||
RUN pip install --no-deps tensorflow-probability==0.20.0
|
||||
RUN pip install --no-deps tensorflow-text==2.12.1
|
||||
RUN pip install --no-deps tensorstore==0.1.36
|
||||
RUN pip install --no-deps termcolor==2.2.0
|
||||
RUN pip install --no-deps toml==0.10.2
|
||||
RUN pip install --no-deps toolz==0.12.0
|
||||
RUN pip install --no-deps tqdm==4.65.0
|
||||
RUN pip install --no-deps typing_extensions==4.5.0
|
||||
RUN pip install --no-deps tzdata==2023.3
|
||||
RUN pip install --no-deps urllib3==1.25.8
|
||||
RUN pip install --no-deps Werkzeug==2.2.3
|
||||
RUN pip install --no-deps wheel==0.40.0
|
||||
RUN pip install --no-deps wrapt==1.14.1
|
||||
RUN pip install --no-deps zipp==3.15.0
|
||||
# Installing jax at the very end with GPU support.
|
||||
# NOTE: Not using `no-deps` flag here because
|
||||
# we need CUDA support.
|
||||
RUN pip install jax[cuda11_cudnn82]==0.4.6 \
|
||||
--find-links https://storage.googleapis.com/jax-releases/jax_cuda_releases.html
|
||||
|
||||
ENV PYTHONPATH ./vit_jax
|
||||
|
||||
COPY ./model_oss/jax_vision_transformer/vit_jax2tf.py ./
|
||||
COPY ./model_oss/jax_vision_transformer/vit_config_without_data.py vit_jax/configs/vit.py
|
||||
|
||||
ENTRYPOINT ["python", "vit_jax2tf.py"]
|
||||
+149
@@ -0,0 +1,149 @@
|
||||
# This Dockerfile runs the JAX based Vision transformer training on GPU.
|
||||
# See https://github.com/google-research/vision_transformer#running-on-cloud
|
||||
# for more details.
|
||||
# Here is an example to build this dockerfile:
|
||||
# PROJECT="your gcp project"
|
||||
# IMAGE_TAG="trainn_vit_gpu:${USER}-test"
|
||||
# docker build -f model_oss/jax_vision_transformer/dockerfile/train_vit_gpu.Dockerfile . -t "${IMAGE_TAG}"
|
||||
# docker tag "${IMAGE_TAG}" "gcr.io/${PROJECT}/${IMAGE_TAG}"
|
||||
# docker push "gcr.io/${PROJECT}/${IMAGE_TAG}"
|
||||
|
||||
FROM tensorflow/tensorflow:2.12.0-gpu
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Install basic libs
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake \
|
||||
curl \
|
||||
wget \
|
||||
sudo \
|
||||
gnupg \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender-dev \
|
||||
lsb-release \
|
||||
ca-certificates \
|
||||
build-essential \
|
||||
git
|
||||
|
||||
# Copy Apache license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Get 'vision_transformer' repository from github.
|
||||
RUN git clone https://github.com/google-research/vision_transformer
|
||||
# Ser current directory to the downloaded 'vision_transformer' repository.
|
||||
WORKDIR ./vision_transformer
|
||||
# Using git reset command to pin it down to a specific version.
|
||||
RUN git reset --hard e66b4732d44504251197a3da3f5949f3f3ce9ca6
|
||||
|
||||
# Install required libs
|
||||
RUN pip install --upgrade pip
|
||||
# The following pip installs are pinned down versions of those inside
|
||||
# vit_jax/requirements.txt file.
|
||||
# NOTE: Using `no-deps` flag to avoid overwriting of
|
||||
# dependent library versions. For example,
|
||||
# both `chex` and `jax` can overwrite each others
|
||||
# `jax-lib` version.
|
||||
RUN pip install --no-deps absl-py==1.4.0
|
||||
RUN pip install --no-deps aqtp==0.0.10
|
||||
RUN pip install --no-deps array-record==0.2.0
|
||||
RUN pip install --no-deps astunparse==1.6.3
|
||||
RUN pip install --no-deps cached-property==1.5.2
|
||||
RUN pip install --no-deps cachetools==5.3.0
|
||||
RUN pip install --no-deps certifi==2019.11.28
|
||||
RUN pip install --no-deps chardet==3.0.4
|
||||
RUN pip install --no-deps chex==0.1.7
|
||||
RUN pip install --no-deps click==8.1.3
|
||||
RUN pip install --no-deps cloudpickle==2.2.1
|
||||
RUN pip install --no-deps clu==0.0.9
|
||||
RUN pip install --no-deps contextlib2==21.6.0
|
||||
RUN pip install --no-deps dacite==1.8.1
|
||||
RUN pip install --no-deps dbus-python==1.2.16
|
||||
RUN pip install --no-deps decorator==5.1.1
|
||||
RUN pip install --no-deps dm-tree==0.1.8
|
||||
RUN pip install --no-deps einops==0.6.1
|
||||
RUN pip install --no-deps etils==1.3.0
|
||||
RUN pip install --no-deps flatbuffers==23.3.3
|
||||
RUN pip install --no-deps flax==0.6.10
|
||||
RUN pip install --no-deps git+https://github.com/google/flaxformer@9adaa4467cf17703949b9f537c3566b99de1b416
|
||||
RUN pip install --no-deps gast==0.4.0
|
||||
RUN pip install --no-deps google-auth==2.16.2
|
||||
RUN pip install --no-deps google-auth-oauthlib==0.4.6
|
||||
RUN pip install --no-deps google-pasta==0.2.0
|
||||
RUN pip install --no-deps googleapis-common-protos==1.59.0
|
||||
RUN pip install --no-deps grpcio==1.51.3
|
||||
RUN pip install --no-deps h5py==3.8.0
|
||||
RUN pip install --no-deps idna==2.8
|
||||
RUN pip install --no-deps importlib-metadata==6.1.0
|
||||
RUN pip install --no-deps importlib-resources==5.12.0
|
||||
RUN pip install --no-deps keras==2.12.0
|
||||
RUN pip install --no-deps libclang==16.0.0
|
||||
RUN pip install --no-deps Markdown==3.4.3
|
||||
RUN pip install --no-deps markdown-it-py==2.2.0
|
||||
RUN pip install --no-deps MarkupSafe==2.1.2
|
||||
RUN pip install --no-deps mdurl==0.1.2
|
||||
RUN pip install --no-deps ml-collections==0.1.1
|
||||
RUN pip install --no-deps msgpack==1.0.5
|
||||
RUN pip install --no-deps nest-asyncio==1.5.6
|
||||
RUN pip install --no-deps numpy==1.23.5
|
||||
RUN pip install --no-deps oauthlib==3.2.2
|
||||
RUN pip install --no-deps opt-einsum==3.3.0
|
||||
RUN pip install --no-deps optax==0.1.5
|
||||
RUN pip install --no-deps orbax-checkpoint==0.1.6
|
||||
RUN pip install --no-deps packaging==23.0
|
||||
RUN pip install --no-deps pandas==2.0.1
|
||||
RUN pip install --no-deps pip==23.1.2
|
||||
RUN pip install --no-deps promise==2.3
|
||||
RUN pip install --no-deps protobuf==4.22.1
|
||||
RUN pip install --no-deps psutil==5.9.5
|
||||
RUN pip install --no-deps pyasn1==0.4.8
|
||||
RUN pip install --no-deps pyasn1-modules==0.2.8
|
||||
RUN pip install --no-deps Pygments==2.15.1
|
||||
RUN pip install --no-deps PyGObject==3.36.0
|
||||
RUN pip install --no-deps python-apt==2.0.1+ubuntu0.20.4.1
|
||||
RUN pip install --no-deps python-dateutil==2.8.2
|
||||
RUN pip install --no-deps pytz==2023.3
|
||||
RUN pip install --no-deps PyYAML==6.0
|
||||
RUN pip install --no-deps requests==2.22.0
|
||||
RUN pip install --no-deps requests-oauthlib==1.3.1
|
||||
RUN pip install --no-deps requests-unixsocket==0.2.0
|
||||
RUN pip install --no-deps rich==13.3.5
|
||||
RUN pip install --no-deps rsa==4.9
|
||||
RUN pip install --no-deps scipy==1.10.1
|
||||
RUN pip install --no-deps setuptools==67.6.0
|
||||
RUN pip install --no-deps six==1.14.0
|
||||
RUN pip install --no-deps tensorboard==2.12.0
|
||||
RUN pip install --no-deps tensorboard-data-server==0.7.0
|
||||
RUN pip install --no-deps tensorboard-plugin-wit==1.8.1
|
||||
RUN pip install --no-deps tensorflow==2.12.0
|
||||
RUN pip install --no-deps tensorflow-cpu==2.12.0
|
||||
RUN pip install --no-deps tensorflow-datasets==4.9.2
|
||||
RUN pip install --no-deps tensorflow-estimator==2.12.0
|
||||
RUN pip install --no-deps tensorflow-hub==0.13.0
|
||||
RUN pip install --no-deps tensorflow-io-gcs-filesystem==0.31.0
|
||||
RUN pip install --no-deps tensorflow-metadata==1.13.1
|
||||
RUN pip install --no-deps tensorflow-probability==0.20.0
|
||||
RUN pip install --no-deps tensorflow-text==2.12.1
|
||||
RUN pip install --no-deps tensorstore==0.1.36
|
||||
RUN pip install --no-deps termcolor==2.2.0
|
||||
RUN pip install --no-deps toml==0.10.2
|
||||
RUN pip install --no-deps toolz==0.12.0
|
||||
RUN pip install --no-deps tqdm==4.65.0
|
||||
RUN pip install --no-deps typing_extensions==4.5.0
|
||||
RUN pip install --no-deps tzdata==2023.3
|
||||
RUN pip install --no-deps urllib3==1.25.8
|
||||
RUN pip install --no-deps Werkzeug==2.2.3
|
||||
RUN pip install --no-deps wheel==0.40.0
|
||||
RUN pip install --no-deps wrapt==1.14.1
|
||||
RUN pip install --no-deps zipp==3.15.0
|
||||
# Installing jax at the very end with GPU support.
|
||||
# NOTE: Not using `no-deps` flag here because
|
||||
# we need CUDA support.
|
||||
RUN pip install jax[cuda11_cudnn82]==0.4.6 \
|
||||
--find-links https://storage.googleapis.com/jax-releases/jax_cuda_releases.html
|
||||
|
||||
COPY ./model_oss/jax_vision_transformer/vit_config_without_data.py vit_jax/configs/vit.py
|
||||
|
||||
ENV PYTHONPATH ./vit_jax
|
||||
ENTRYPOINT ["python", "-m", "vit_jax.main"]
|
||||
+27
@@ -0,0 +1,27 @@
|
||||
"""Returns a config for a Vision Transformer model without asking for data."""
|
||||
import ml_collections
|
||||
from vit_jax.configs import common
|
||||
from vit_jax.configs import models
|
||||
|
||||
|
||||
def get_config(model: str) -> ml_collections.ConfigDict:
|
||||
"""Returns default parameters for finetuning ViT `model`."""
|
||||
config = common.get_config()
|
||||
|
||||
get_model_config = getattr(models, f'get_{model}_config')
|
||||
config.model = get_model_config()
|
||||
|
||||
# These values are often overridden on the command line.
|
||||
config.base_lr = 0.03
|
||||
config.total_steps = 500
|
||||
config.warmup_steps = 100
|
||||
config.pp = ml_collections.ConfigDict()
|
||||
config.pp.train = 'train'
|
||||
config.pp.test = 'test'
|
||||
config.pp.resize = 448
|
||||
config.pp.crop = 384
|
||||
|
||||
# This value MUST be overridden on the command line.
|
||||
config.dataset = ''
|
||||
|
||||
return config
|
||||
+118
@@ -0,0 +1,118 @@
|
||||
# Dockerfile for basic serving dockers with Keras.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/keras/dockerfile/serve.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM tensorflow/tensorflow:2.12.0-gpu
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# This is added to fix docker build error related to Nvidia key update.
|
||||
RUN rm -f /etc/apt/sources.list.d/cuda.list
|
||||
RUN curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add -
|
||||
|
||||
# Install basic libs.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake \
|
||||
curl \
|
||||
wget \
|
||||
sudo \
|
||||
gnupg \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender-dev \
|
||||
lsb-release \
|
||||
ca-certificates \
|
||||
build-essential \
|
||||
git \
|
||||
vim \
|
||||
screen \
|
||||
libtcmalloc-minimal4
|
||||
|
||||
|
||||
# Install google cloud SDK.
|
||||
RUN wget -q https://dl.google.com/dl/cloudsdk/channels/rapid/downloads/google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN tar xzf google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN ./google-cloud-sdk/install.sh -q
|
||||
# Make sure gsutil will use the default service account.
|
||||
RUN echo '[GoogleCompute]\nservice_account = default' > /etc/boto.cfg
|
||||
|
||||
|
||||
# Install required libs.
|
||||
RUN pip install --upgrade pip
|
||||
RUN pip install cloud-tpu-client==0.10
|
||||
RUN pip install pyyaml==5.4.1
|
||||
RUN pip install fsspec==2021.10.1
|
||||
RUN pip install gcsfs==2021.10.1
|
||||
RUN pip install tensorflow-text==2.11.0
|
||||
RUN pip install pyglove==0.1.0
|
||||
RUN pip install cloudml-hypertune==0.1.0.dev6
|
||||
RUN pip install pylint==2.17.2
|
||||
RUN pip install keras-cv==0.4.0
|
||||
RUN pip install tensorflow-datasets==4.8.3
|
||||
RUN pip install protobuf==3.20.3
|
||||
RUN pip install Pillow==9.5.0
|
||||
RUN pip install flask==2.3.2
|
||||
RUN pip install waitress==2.1.2
|
||||
|
||||
# Installs Reduction Server NCCL plugin.
|
||||
RUN echo "deb https://packages.cloud.google.com/apt google-fast-socket main" | tee /etc/apt/sources.list.d/google-fast-socket.list \
|
||||
&& curl -s -L https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add - \
|
||||
&& apt update && apt install -y google-reduction-server
|
||||
|
||||
# Downloading gcloud package
|
||||
RUN curl https://dl.google.com/dl/cloudsdk/release/google-cloud-sdk.tar.gz > /tmp/google-cloud-sdk.tar.gz
|
||||
|
||||
# Installing the package
|
||||
RUN mkdir -p /usr/local/gcloud \
|
||||
&& tar -C /usr/local/gcloud -xvf /tmp/google-cloud-sdk.tar.gz \
|
||||
&& /usr/local/gcloud/google-cloud-sdk/install.sh
|
||||
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Adding the package path to local
|
||||
ENV PATH $PATH:/usr/local/gcloud/google-cloud-sdk/bin
|
||||
|
||||
|
||||
ENV PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=cpp
|
||||
|
||||
# Lower the memory fragmentation, and speed up the training.
|
||||
# https://github.com/tensorflow/tensorflow/issues/44176#issuecomment-783768033
|
||||
ENV LD_PRELOAD=/usr/lib/x86_64-linux-gnu/libtcmalloc_minimal.so.4
|
||||
|
||||
# Enable userspace DNS cache
|
||||
ENV GCS_RESOLVE_REFRESH_SECS=60
|
||||
ENV GCS_REQUEST_CONNECTION_TIMEOUT_SECS=300
|
||||
ENV GCS_METADATA_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_READ_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_WRITE_REQUEST_TIMEOUT_SECS=600
|
||||
# Each opened GCS file takes GCS_READ_CACHE_BLOCK_SIZE_MB of RAM, reduce the
|
||||
# value from the default 64MB to 8MB to decrease memory footprint.
|
||||
ENV GCS_READ_CACHE_BLOCK_SIZE_MB=8
|
||||
|
||||
EXPOSE 8501
|
||||
|
||||
WORKDIR /usr/local/lib/python3.8/dist-packages/official/vision
|
||||
|
||||
COPY model_oss/keras /automl_vision/keras
|
||||
COPY model_oss/util /automl_vision/util
|
||||
|
||||
WORKDIR /automl_vision
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/util"
|
||||
ENV MODEL_PATH ""
|
||||
ENV IMAGE_WIDTH "512"
|
||||
ENV IMAGE_HEIGHT "512"
|
||||
|
||||
COPY model_oss/keras/serve.py ./app.py
|
||||
|
||||
# Run pylint to validate code.
|
||||
COPY .pylintrc /automl_vision/.pylintrc
|
||||
RUN find . -type f -name "*.py" | xargs pylint --rcfile=./.pylintrc --errors-only
|
||||
|
||||
ENTRYPOINT ["flask","run"]
|
||||
CMD ["--host=0.0.0.0", "--port=8501"]
|
||||
+111
@@ -0,0 +1,111 @@
|
||||
# Dockerfile for basic training dockers with Keras.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/keras/dockerfile/train.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM tensorflow/tensorflow:2.12.0-gpu
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# This is added to fix docker build error related to Nvidia key update.
|
||||
RUN rm -f /etc/apt/sources.list.d/cuda.list
|
||||
RUN curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add -
|
||||
|
||||
# Install basic libs.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake \
|
||||
curl \
|
||||
wget \
|
||||
sudo \
|
||||
gnupg \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender-dev \
|
||||
lsb-release \
|
||||
ca-certificates \
|
||||
build-essential \
|
||||
git \
|
||||
vim \
|
||||
screen \
|
||||
libtcmalloc-minimal4
|
||||
|
||||
|
||||
# Install google cloud SDK.
|
||||
RUN wget -q https://dl.google.com/dl/cloudsdk/channels/rapid/downloads/google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN tar xzf google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN ./google-cloud-sdk/install.sh -q
|
||||
# Make sure gsutil will use the default service account.
|
||||
RUN echo '[GoogleCompute]\nservice_account = default' > /etc/boto.cfg
|
||||
|
||||
|
||||
# Install required libs.
|
||||
RUN pip install --upgrade pip
|
||||
RUN pip install cloud-tpu-client==0.10
|
||||
RUN pip install pyyaml==5.4.1
|
||||
RUN pip install fsspec==2021.10.1
|
||||
RUN pip install gcsfs==2021.10.1
|
||||
RUN pip install tensorflow-text==2.11.0
|
||||
RUN pip install pyglove==0.1.0
|
||||
RUN pip install cloudml-hypertune==0.1.0.dev6
|
||||
RUN pip install pylint==2.17.2
|
||||
RUN pip install keras-cv==0.4.0
|
||||
RUN pip install tensorflow-datasets==4.8.3
|
||||
RUN pip install tensorflow-estimator==2.12.0
|
||||
RUN pip install tensorflow-gcs-config==2.12.0
|
||||
RUN pip install tensorflow-hub==0.13.0
|
||||
RUN pip install tensorflow-io-gcs-filesystem==0.32.0
|
||||
RUN pip install tensorflow-metadata==1.13.1
|
||||
RUN pip install tensorflow-probability==0.19.0
|
||||
RUN pip install tensorboard==2.12.2
|
||||
RUN pip install tensorboard-data-server==0.7.0
|
||||
RUN pip install tensorboard-plugin-wit==1.8.1
|
||||
RUN pip install protobuf==3.20.3
|
||||
RUN pip install pandas==1.5.3
|
||||
RUN pip install pandas-datareader==0.10.0
|
||||
RUN pip install pandas-gbq==0.17.9
|
||||
RUN pip install pycocotools==2.0.6
|
||||
|
||||
# Installs Reduction Server NCCL plugin.
|
||||
RUN echo "deb https://packages.cloud.google.com/apt google-fast-socket main" | tee /etc/apt/sources.list.d/google-fast-socket.list \
|
||||
&& curl -s -L https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add - \
|
||||
&& apt update && apt install -y google-reduction-server
|
||||
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
ENV PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=cpp
|
||||
|
||||
# Lower the memory fragmentation, and speed up the training.
|
||||
# https://github.com/tensorflow/tensorflow/issues/44176#issuecomment-783768033
|
||||
ENV LD_PRELOAD=/usr/lib/x86_64-linux-gnu/libtcmalloc_minimal.so.4
|
||||
|
||||
# Enable userspace DNS cache
|
||||
ENV GCS_RESOLVE_REFRESH_SECS=60
|
||||
ENV GCS_REQUEST_CONNECTION_TIMEOUT_SECS=300
|
||||
ENV GCS_METADATA_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_READ_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_WRITE_REQUEST_TIMEOUT_SECS=600
|
||||
# Each opened GCS file takes GCS_READ_CACHE_BLOCK_SIZE_MB of RAM, reduce the
|
||||
# value from the default 64MB to 8MB to decrease memory footprint.
|
||||
ENV GCS_READ_CACHE_BLOCK_SIZE_MB=8
|
||||
|
||||
WORKDIR /usr/local/lib/python3.8/dist-packages/official/vision
|
||||
|
||||
COPY model_oss/keras /automl_vision/keras
|
||||
COPY model_oss/util /automl_vision/util
|
||||
|
||||
WORKDIR /automl_vision
|
||||
|
||||
# Keras stable diffusion training codes set width and height as RESOLUTION.
|
||||
ENV RESOLUTION "512"
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/util"
|
||||
|
||||
# Run pylint to validate code.
|
||||
COPY .pylintrc /automl_vision/.pylintrc
|
||||
RUN find . -type f -name "*.py" | xargs pylint --rcfile=./.pylintrc --errors-only
|
||||
|
||||
ENTRYPOINT ["python3","keras/train.py"]
|
||||
@@ -0,0 +1,184 @@
|
||||
r"""Servers Keras Stable Diffusion models.
|
||||
|
||||
python serve.py --model_path=<model path in gcs>
|
||||
|
||||
curl -d \
|
||||
'{"prompt":"Hello Kitty"}' \
|
||||
-H "Content-Type: application/json" \
|
||||
-X POST http://localhost:8501/predict
|
||||
"""
|
||||
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
from typing import List, Tuple
|
||||
|
||||
from absl import app
|
||||
# The docker builds could not find flask and waitress.
|
||||
# pylint: disable=import-error
|
||||
from flask import Flask
|
||||
from flask import request
|
||||
from flask import Response
|
||||
import keras_cv
|
||||
from PIL import Image
|
||||
from waitress import serve
|
||||
|
||||
from util import constants
|
||||
from util import fileutils
|
||||
|
||||
|
||||
flask_app = Flask(__name__)
|
||||
|
||||
stable_diffusion_model = None
|
||||
|
||||
|
||||
model_path = os.environ.get('MODEL_PATH', '')
|
||||
if model_path.startswith(constants.GCS_URI_PREFIX):
|
||||
print('Downloading models from gcs to local.')
|
||||
os.makedirs(constants.LOCAL_MODEL_DIR, exist_ok=True)
|
||||
fileutils.download_gcs_dir_to_local(
|
||||
os.path.dirname(model_path), constants.LOCAL_MODEL_DIR
|
||||
)
|
||||
model_path = os.path.join(
|
||||
constants.LOCAL_MODEL_DIR, os.path.basename(model_path)
|
||||
)
|
||||
|
||||
image_width = int(os.environ.get('IMAGE_WIDTH', 512))
|
||||
image_height = int(os.environ.get('IMAGE_HEIGHT', 512))
|
||||
|
||||
print('image_width=', image_width, 'image_height=', image_height)
|
||||
print('Create Keras stable diffusion models.')
|
||||
stable_diffusion_model = keras_cv.models.StableDiffusion(
|
||||
img_width=image_width,
|
||||
img_height=image_height,
|
||||
jit_compile=True,
|
||||
)
|
||||
|
||||
if model_path:
|
||||
# We just reload the weights of the fine-tuned diffusion model.
|
||||
print('Initialize finetuned models from: ', model_path)
|
||||
stable_diffusion_model.diffusion_model.load_weights(model_path)
|
||||
|
||||
|
||||
def error(message: str) -> str:
|
||||
"""Returns a JSON representing an error response."""
|
||||
return json.dumps({
|
||||
'success': False,
|
||||
'error': message,
|
||||
})
|
||||
|
||||
|
||||
def check_key_in_json(content: str, keys: List[str]) -> str:
|
||||
for key in keys:
|
||||
if key not in content:
|
||||
return error('No {} in request {}.'.format(key, content))
|
||||
return None
|
||||
|
||||
|
||||
def validate_json_key(json_key_string: str) -> Tuple[str, bool]:
|
||||
try:
|
||||
json_key = json.loads(json_key_string)
|
||||
except (ValueError, TypeError):
|
||||
return (error('Invalid key found in request'), False)
|
||||
return (json_key, True)
|
||||
|
||||
|
||||
# The health check route is required for docker deployment in google cloud.
|
||||
@flask_app.route('/ping')
|
||||
def ping() -> Response:
|
||||
"""Health checks."""
|
||||
return Response(status=200)
|
||||
|
||||
|
||||
# The return should be `Response` for docker deployment in google cloud.
|
||||
@flask_app.route('/predict', methods=['GET', 'POST'])
|
||||
def predict_model() -> Response:
|
||||
"""Predictions."""
|
||||
if request.method == 'POST':
|
||||
contents = request.get_json(force=True)
|
||||
|
||||
print('The input contents are:', contents)
|
||||
batch_size = 1
|
||||
num_steps = 25
|
||||
seed = 1234
|
||||
if 'parameters' in contents:
|
||||
parameters = contents['parameters']
|
||||
if 'batch_size' in parameters:
|
||||
batch_size = int(parameters['batch_size'])
|
||||
if 'num_steps' in parameters:
|
||||
num_steps = int(parameters['num_steps'])
|
||||
if 'seed' in parameters:
|
||||
seed = int(parameters['seed'])
|
||||
print('batch_size=', batch_size, 'num_steps=', num_steps, 'seed=', seed)
|
||||
if batch_size < 1:
|
||||
return Response(
|
||||
response=error('The batch size must be a positive integar.'),
|
||||
status=200,
|
||||
mimetype='text/plain',
|
||||
)
|
||||
if num_steps < 1:
|
||||
return Response(
|
||||
response=error('The num steps must be a positive integar.'),
|
||||
status=200,
|
||||
mimetype='text/plain',
|
||||
)
|
||||
predictions = []
|
||||
for content in contents['instances']:
|
||||
print('Processing:', content)
|
||||
prompt = content['prompt']
|
||||
generated_image_array = stable_diffusion_model.text_to_image(
|
||||
prompt=prompt,
|
||||
batch_size=batch_size,
|
||||
num_steps=num_steps,
|
||||
seed=seed,
|
||||
)
|
||||
|
||||
generated_image_bytes_array = []
|
||||
for i in range(batch_size):
|
||||
generated_image = Image.fromarray(generated_image_array[i])
|
||||
# Converts the image to a base64-encoded string.
|
||||
buffered_image = io.BytesIO()
|
||||
generated_image.save(buffered_image, format='JPEG')
|
||||
generated_image_bytes = base64.b64encode(
|
||||
buffered_image.getvalue()
|
||||
).decode('utf-8')
|
||||
generated_image_bytes_array.append(generated_image_bytes)
|
||||
prediction = {
|
||||
'prompt': prompt,
|
||||
'predicted_image': generated_image_bytes_array,
|
||||
}
|
||||
predictions.append(prediction)
|
||||
|
||||
return Response(
|
||||
response=json.dumps({
|
||||
'success': True,
|
||||
'predictions': predictions,
|
||||
}),
|
||||
status=200,
|
||||
mimetype='text/plain',
|
||||
)
|
||||
else:
|
||||
return Response(
|
||||
response=json.dumps({
|
||||
'success': True,
|
||||
'isalive': stable_diffusion_model is not None,
|
||||
}),
|
||||
status=200,
|
||||
mimetype='text/plain',
|
||||
)
|
||||
|
||||
|
||||
def serve_main(unused_argv):
|
||||
"""The main function to serve Keras models."""
|
||||
del unused_argv
|
||||
# This is used when running locally only. When deploying to Google App
|
||||
# Engine, a webserver process such as Gunicorn will serve the app.
|
||||
# # Debug deployment.
|
||||
# flask_app.run(host='0.0.0.0', port=8501, debug=True)
|
||||
# Prod deployment.
|
||||
serve(flask_app, host='0.0.0.0', port=8501)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(serve_main)
|
||||
@@ -0,0 +1,363 @@
|
||||
"""Train Keras Stable Diffusion.
|
||||
|
||||
Most the codes below are from
|
||||
https://keras.io/examples/generative/finetune_stable_diffusion/.
|
||||
"""
|
||||
import os
|
||||
|
||||
from absl import app
|
||||
from absl import flags
|
||||
from absl import logging
|
||||
import keras_cv
|
||||
# pylint: disable=g-importing-member
|
||||
from keras_cv.models.stable_diffusion.clip_tokenizer import SimpleTokenizer
|
||||
from keras_cv.models.stable_diffusion.diffusion_model import DiffusionModel
|
||||
from keras_cv.models.stable_diffusion.image_encoder import ImageEncoder
|
||||
from keras_cv.models.stable_diffusion.noise_scheduler import NoiseScheduler
|
||||
from keras_cv.models.stable_diffusion.text_encoder import TextEncoder
|
||||
import numpy as np
|
||||
# The docker builds could not find pandas.
|
||||
# pylint: disable=import-error
|
||||
import pandas as pd
|
||||
import tensorflow as tf
|
||||
from tensorflow import keras
|
||||
import tensorflow.experimental.numpy as tnp
|
||||
|
||||
from util import constants
|
||||
from util import fileutils
|
||||
|
||||
_INPUT_CSV_PATH = flags.DEFINE_string(
|
||||
'input_csv_path',
|
||||
None,
|
||||
'The input csv path.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_USE_MP = flags.DEFINE_bool(
|
||||
'use_mp',
|
||||
True,
|
||||
'Enable mixed-precision training if the underlying GPU has tensor cores.',
|
||||
)
|
||||
|
||||
_EPOCHS = flags.DEFINE_integer('epochs', 1, 'The number of epochs.')
|
||||
|
||||
_OUTPUT_MODEL_DIR = flags.DEFINE_string(
|
||||
'output_model_dir',
|
||||
None,
|
||||
'The output model dir.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
# These hyperparameters defaults come from this tutorial by Hugging Face:
|
||||
# https://huggingface.co/docs/diffusers/training/text2image
|
||||
_LEARNING_RATE = flags.DEFINE_float(
|
||||
'learning_rate', 1e-5, 'The learning rate parameter for AdamW optimizer.'
|
||||
)
|
||||
|
||||
_BETA_1 = flags.DEFINE_float(
|
||||
'beta_1', 0.9, 'The beta_1 parameter for AdamW optimizer.'
|
||||
)
|
||||
|
||||
_BETA_2 = flags.DEFINE_float(
|
||||
'beta_2', 0.999, 'The beta_2 parameter for AdamW optimizer.'
|
||||
)
|
||||
|
||||
_WEIGHT_DECAY = flags.DEFINE_float(
|
||||
'weight_decay', 1e-2, 'The weight decay parameter for AdamW optimizer.'
|
||||
)
|
||||
|
||||
_EPSILON = flags.DEFINE_float(
|
||||
'epsilon', 1e-08, 'The epsilon parameter for AdamW optimizer.'
|
||||
)
|
||||
|
||||
RESOLUTION = int(os.environ.get('RESOLUTION', 512))
|
||||
|
||||
# The padding token and maximum prompt length are specific to the text encoder.
|
||||
# If you're using a different text encoder be sure to change them accordingly.
|
||||
PADDING_TOKEN = 49407
|
||||
MAX_PROMPT_LENGTH = 77
|
||||
|
||||
AUTO = tf.data.AUTOTUNE
|
||||
POS_IDS = tf.convert_to_tensor([list(range(MAX_PROMPT_LENGTH))], dtype=tf.int32)
|
||||
|
||||
|
||||
augmenter = keras.Sequential(
|
||||
layers=[
|
||||
keras_cv.layers.CenterCrop(RESOLUTION, RESOLUTION),
|
||||
keras_cv.layers.RandomFlip(),
|
||||
tf.keras.layers.Rescaling(scale=1.0 / 127.5, offset=-1),
|
||||
]
|
||||
)
|
||||
text_encoder = TextEncoder(MAX_PROMPT_LENGTH)
|
||||
|
||||
|
||||
def process_image(image_path, tokenized_text):
|
||||
image = tf.io.read_file(image_path)
|
||||
image = tf.io.decode_png(image, 3)
|
||||
image = tf.image.resize(image, (RESOLUTION, RESOLUTION))
|
||||
return image, tokenized_text
|
||||
|
||||
|
||||
def apply_augmentation(image_batch, token_batch):
|
||||
return augmenter(image_batch), token_batch
|
||||
|
||||
|
||||
def run_text_encoder(image_batch, token_batch):
|
||||
return (
|
||||
image_batch,
|
||||
token_batch,
|
||||
text_encoder([token_batch, POS_IDS], training=False),
|
||||
)
|
||||
|
||||
|
||||
def prepare_dict(image_batch, token_batch, encoded_text_batch):
|
||||
return {
|
||||
'images': image_batch,
|
||||
'tokens': token_batch,
|
||||
'encoded_text': encoded_text_batch,
|
||||
}
|
||||
|
||||
|
||||
def prepare_dataset(image_paths, tokenized_texts, batch_size=1):
|
||||
dataset = tf.data.Dataset.from_tensor_slices((image_paths, tokenized_texts))
|
||||
dataset = dataset.shuffle(batch_size * 10)
|
||||
dataset = dataset.map(process_image, num_parallel_calls=AUTO).batch(
|
||||
batch_size
|
||||
)
|
||||
dataset = dataset.map(apply_augmentation, num_parallel_calls=AUTO)
|
||||
dataset = dataset.map(run_text_encoder, num_parallel_calls=AUTO)
|
||||
dataset = dataset.map(prepare_dict, num_parallel_calls=AUTO)
|
||||
return dataset.prefetch(AUTO)
|
||||
|
||||
|
||||
def prepare_training_dataset(dataset_csv):
|
||||
"""Prepares training datasets."""
|
||||
if dataset_csv.startswith(constants.GCS_URI_PREFIX):
|
||||
if not os.path.exists(constants.LOCAL_DATA_DIR):
|
||||
os.makedirs(constants.LOCAL_DATA_DIR)
|
||||
logging.info(
|
||||
'Start to download data from %s to %s.',
|
||||
os.path.dirname(dataset_csv),
|
||||
constants.LOCAL_DATA_DIR,
|
||||
)
|
||||
fileutils.download_gcs_dir_to_local(
|
||||
os.path.dirname(dataset_csv), constants.LOCAL_DATA_DIR
|
||||
)
|
||||
data_frame = pd.read_csv(
|
||||
os.path.join(constants.LOCAL_DATA_DIR, os.path.basename(dataset_csv))
|
||||
)
|
||||
data_frame['image_path'] = data_frame['image_path'].apply(
|
||||
lambda x: os.path.join(constants.LOCAL_DATA_DIR, x)
|
||||
)
|
||||
else:
|
||||
# Keeps the following codes for experiments with
|
||||
# https://keras.io/examples/generative/finetune_stable_diffusion/.
|
||||
data_path = tf.keras.utils.get_file(origin=dataset_csv, untar=True)
|
||||
data_frame = pd.read_csv(os.path.join(data_path, 'data.csv'))
|
||||
data_frame['image_path'] = data_frame['image_path'].apply(
|
||||
lambda x: os.path.join(data_path, x)
|
||||
)
|
||||
data_frame.head()
|
||||
|
||||
# Load the tokenizer.
|
||||
tokenizer = SimpleTokenizer()
|
||||
|
||||
# Method to tokenize and pad the tokens.
|
||||
def process_text(caption):
|
||||
tokens = tokenizer.encode(caption)
|
||||
tokens = tokens + [PADDING_TOKEN] * (MAX_PROMPT_LENGTH - len(tokens))
|
||||
return np.array(tokens)
|
||||
|
||||
# Collate the tokenized captions into an array.
|
||||
tokenized_texts = np.empty((len(data_frame), MAX_PROMPT_LENGTH))
|
||||
|
||||
all_captions = list(data_frame['caption'].values)
|
||||
for i, caption in enumerate(all_captions):
|
||||
tokenized_texts[i] = process_text(caption)
|
||||
|
||||
# Prepare the dataset.
|
||||
training_dataset = prepare_dataset(
|
||||
np.array(data_frame['image_path']), tokenized_texts, batch_size=4
|
||||
)
|
||||
|
||||
return training_dataset
|
||||
|
||||
|
||||
class Trainer(tf.keras.Model):
|
||||
"""The trainer for Keras Stable Diffusion."""
|
||||
|
||||
# Reference:
|
||||
# https://github.com/huggingface/diffusers/blob/main/examples/text_to_image/train_text_to_image.py
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
diffusion_model,
|
||||
vae,
|
||||
noise_scheduler,
|
||||
use_mixed_precision=False,
|
||||
max_grad_norm=1.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
self.diffusion_model = diffusion_model
|
||||
self.vae = vae
|
||||
self.noise_scheduler = noise_scheduler
|
||||
self.max_grad_norm = max_grad_norm
|
||||
|
||||
self.use_mixed_precision = use_mixed_precision
|
||||
self.vae.trainable = False
|
||||
|
||||
def train_step(self, inputs):
|
||||
images = inputs['images']
|
||||
encoded_text = inputs['encoded_text']
|
||||
batch_size = tf.shape(images)[0]
|
||||
|
||||
with tf.GradientTape() as tape:
|
||||
# Project image into the latent space and sample from it.
|
||||
latents = self.sample_from_encoder_outputs(
|
||||
self.vae(images, training=False)
|
||||
)
|
||||
# Know more about the magic number here:
|
||||
# https://keras.io/examples/generative/fine_tune_via_textual_inversion/
|
||||
latents = latents * 0.18215
|
||||
|
||||
# Sample noise that we'll add to the latents.
|
||||
noise = tf.random.normal(tf.shape(latents))
|
||||
|
||||
# Sample a random timestep for each image.
|
||||
timesteps = tnp.random.randint(
|
||||
0, self.noise_scheduler.train_timesteps, (batch_size,)
|
||||
)
|
||||
|
||||
# Add noise to the latents according to the noise magnitude at each
|
||||
# timestep (this is the forward diffusion process).
|
||||
noisy_latents = self.noise_scheduler.add_noise(
|
||||
tf.cast(latents, noise.dtype), noise, timesteps
|
||||
)
|
||||
|
||||
# Get the target for loss depending on the prediction type
|
||||
# just the sampled noise for now.
|
||||
target = noise # noise_schedule.predict_epsilon == True
|
||||
|
||||
# Predict the noise residual and compute loss.
|
||||
# pylint: disable=unnecessary-lambda
|
||||
timestep_embedding = tf.map_fn(
|
||||
lambda t: self.get_timestep_embedding(t), timesteps, dtype=tf.float32
|
||||
)
|
||||
timestep_embedding = tf.squeeze(timestep_embedding, 1)
|
||||
model_pred = self.diffusion_model(
|
||||
[noisy_latents, timestep_embedding, encoded_text], training=True
|
||||
)
|
||||
loss = self.compiled_loss(target, model_pred)
|
||||
if self.use_mixed_precision:
|
||||
loss = self.optimizer.get_scaled_loss(loss)
|
||||
|
||||
# Update parameters of the diffusion model.
|
||||
trainable_vars = self.diffusion_model.trainable_variables
|
||||
gradients = tape.gradient(loss, trainable_vars)
|
||||
if self.use_mixed_precision:
|
||||
gradients = self.optimizer.get_unscaled_gradients(gradients)
|
||||
gradients = [tf.clip_by_norm(g, self.max_grad_norm) for g in gradients]
|
||||
self.optimizer.apply_gradients(zip(gradients, trainable_vars))
|
||||
|
||||
return {m.name: m.result() for m in self.metrics}
|
||||
|
||||
def get_timestep_embedding(self, timestep, dim=320, max_period=10000):
|
||||
half = dim // 2
|
||||
log_max_preiod = tf.math.log(tf.cast(max_period, tf.float32))
|
||||
# The docker builds could not support unary `-`.
|
||||
# pylint: disable=invalid-unary-operand-type
|
||||
freqs = tf.math.exp(
|
||||
-log_max_preiod * tf.range(0, half, dtype=tf.float32) / half
|
||||
)
|
||||
args = tf.convert_to_tensor([timestep], dtype=tf.float32) * freqs
|
||||
embedding = tf.concat([tf.math.cos(args), tf.math.sin(args)], 0)
|
||||
embedding = tf.reshape(embedding, [1, -1])
|
||||
return embedding
|
||||
|
||||
def sample_from_encoder_outputs(self, outputs):
|
||||
mean, logvar = tf.split(outputs, 2, axis=-1)
|
||||
logvar = tf.clip_by_value(logvar, -30.0, 20.0)
|
||||
std = tf.exp(0.5 * logvar)
|
||||
sample = tf.random.normal(tf.shape(mean), dtype=mean.dtype)
|
||||
return mean + std * sample
|
||||
|
||||
def save_weights(
|
||||
self, filepath, overwrite=True, save_format=None, options=None
|
||||
):
|
||||
# Overriding this method will allow us to use the `ModelCheckpoint`
|
||||
# callback directly with this trainer class. In this case, it will
|
||||
# only checkpoint the `diffusion_model` since that's what we're training
|
||||
# during fine-tuning.
|
||||
self.diffusion_model.save_weights(
|
||||
filepath=filepath,
|
||||
overwrite=overwrite,
|
||||
save_format=save_format,
|
||||
options=options,
|
||||
)
|
||||
|
||||
|
||||
def main(_) -> None:
|
||||
# _INPUT_CSV_PATH and _OUTPUT_MODEL_DIR should have the format as
|
||||
# gs://<bucket_name>/<object_name>.
|
||||
if _INPUT_CSV_PATH.value:
|
||||
if not _INPUT_CSV_PATH.value.startswith(constants.GCS_URI_PREFIX):
|
||||
raise ValueError('The input csv path should be a gcs path like gs://<>')
|
||||
if _OUTPUT_MODEL_DIR.value:
|
||||
if not _OUTPUT_MODEL_DIR.value.startswith(constants.GCS_URI_PREFIX):
|
||||
raise ValueError('The output model dir should be a gcs path like gs://<>')
|
||||
|
||||
if _USE_MP.value:
|
||||
keras.mixed_precision.set_global_policy('mixed_float16')
|
||||
|
||||
image_encoder = ImageEncoder(RESOLUTION, RESOLUTION)
|
||||
diffusion_ft_trainer = Trainer(
|
||||
diffusion_model=DiffusionModel(RESOLUTION, RESOLUTION, MAX_PROMPT_LENGTH),
|
||||
# Remove the top layer from the encoder, which cuts off the variance and
|
||||
# only returns the mean.
|
||||
vae=tf.keras.Model(
|
||||
image_encoder.input,
|
||||
image_encoder.layers[-2].output,
|
||||
),
|
||||
noise_scheduler=NoiseScheduler(),
|
||||
use_mixed_precision=_USE_MP.value,
|
||||
)
|
||||
|
||||
optimizer = tf.keras.optimizers.experimental.AdamW(
|
||||
learning_rate=_LEARNING_RATE.value,
|
||||
weight_decay=_WEIGHT_DECAY.value,
|
||||
beta_1=_BETA_1.value,
|
||||
beta_2=_BETA_2.value,
|
||||
epsilon=_EPSILON.value,
|
||||
)
|
||||
diffusion_ft_trainer.compile(optimizer=optimizer, loss='mse')
|
||||
|
||||
training_dataset = prepare_training_dataset(_INPUT_CSV_PATH.value)
|
||||
|
||||
# Note: gcsfuse does not work for Keras. We saves the trained models locally
|
||||
# first, and then copy to gcs storages.
|
||||
if not os.path.exists(constants.LOCAL_MODEL_DIR):
|
||||
os.makedirs(constants.LOCAL_MODEL_DIR)
|
||||
# The default saved model is in HDF5.
|
||||
ckpt_path = os.path.join(constants.LOCAL_MODEL_DIR, 'saved_model.h5')
|
||||
ckpt_callback = tf.keras.callbacks.ModelCheckpoint(
|
||||
ckpt_path,
|
||||
save_weights_only=True,
|
||||
monitor='loss',
|
||||
mode='min',
|
||||
)
|
||||
diffusion_ft_trainer.fit(
|
||||
training_dataset, epochs=_EPOCHS.value, callbacks=[ckpt_callback]
|
||||
)
|
||||
|
||||
# Copies the files in constants.LOCAL_MODEL_DIR to output_model_dir.
|
||||
fileutils.upload_local_dir_to_gcs(
|
||||
constants.LOCAL_MODEL_DIR, _OUTPUT_MODEL_DIR.value
|
||||
)
|
||||
|
||||
return
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(main)
|
||||
+40
@@ -0,0 +1,40 @@
|
||||
# Dockerfile for lm-evaluation-harness evaluation.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/lm-evaluation-harness/dockerfile/eval.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/{YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/{YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM pytorch/pytorch:2.0.0-cuda11.7-cudnn8-devel
|
||||
|
||||
USER root
|
||||
|
||||
# Install tools.
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
RUN apt-get update
|
||||
RUN apt-get install -y --no-install-recommends apt-utils
|
||||
RUN apt-get install -y --no-install-recommends curl
|
||||
RUN apt-get install -y --no-install-recommends wget
|
||||
RUN apt-get install -y --no-install-recommends git
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Install libraries.
|
||||
ENV PIP_ROOT_USER_ACTION=ignore
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
RUN pip install google-cloud-storage==2.7.0
|
||||
RUN pip install absl-py==1.4.0
|
||||
|
||||
# Install lm-evaluation-harness
|
||||
RUN git clone https://github.com/EleutherAI/lm-evaluation-harness
|
||||
WORKDIR lm-evaluation-harness
|
||||
# Pin version up to date 08/08/2023
|
||||
RUN git reset --hard b952a206de210b72b1bf750fbab38c26121e0dc0
|
||||
# Edit tokenizer loading function to avoid using fast tokenizer for OpenLLaMA
|
||||
RUN sed -i '355 i\ use_fast = not pretrained.startswith("openlm-research/open_llama")' lm_eval/models/huggingface.py
|
||||
RUN sed -i '360 i\ use_fast=use_fast,' lm_eval/models/huggingface.py
|
||||
# Install from source while including the sentencepiece dependency
|
||||
RUN pip install -e ".[sentencepiece]"
|
||||
+64
@@ -0,0 +1,64 @@
|
||||
FROM tensorflow/build:2.12-python3.9
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# This is added to fix docker build error related to Nvidia key update.
|
||||
RUN rm -f /etc/apt/sources.list.d/cuda.list
|
||||
RUN curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add -
|
||||
|
||||
# Install basic libs.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake \
|
||||
curl \
|
||||
wget \
|
||||
sudo \
|
||||
gnupg \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender-dev \
|
||||
lsb-release \
|
||||
ca-certificates \
|
||||
build-essential \
|
||||
git \
|
||||
vim \
|
||||
libtcmalloc-minimal4
|
||||
|
||||
|
||||
# Install google cloud CLI.
|
||||
RUN wget -q https://dl.google.com/dl/cloudsdk/channels/rapid/downloads/google-cloud-cli-430.0.0-linux-x86.tar.gz
|
||||
RUN tar xzf google-cloud-cli-430.0.0-linux-x86.tar.gz
|
||||
RUN ./google-cloud-sdk/install.sh -q
|
||||
# Make sure gsutil will use the default service account.
|
||||
RUN echo '[GoogleCompute]\nservice_account = default' > /etc/boto.cfg
|
||||
|
||||
|
||||
# Install required libs.
|
||||
RUN pip install --upgrade pip
|
||||
RUN pip install cloud-tpu-client==0.10
|
||||
RUN pip install pyyaml==6.0
|
||||
RUN pip install fsspec==2023.4.0
|
||||
RUN pip install gcsfs==2023.4.0
|
||||
RUN pip install tf-models-official==2.12.0
|
||||
RUN pip install cloudml-hypertune==0.1.0.dev6
|
||||
RUN pip install pylint==2.17.3
|
||||
|
||||
# Installs Reduction Server NCCL plugin.
|
||||
RUN echo "deb https://packages.cloud.google.com/apt google-fast-socket main" | tee /etc/apt/sources.list.d/google-fast-socket.list \
|
||||
&& curl -s -L https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add - \
|
||||
&& apt update && apt install -y google-reduction-server
|
||||
|
||||
ENV PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=cpp
|
||||
|
||||
# Lower the memory fragmentation, and speed up the training.
|
||||
# https://github.com/tensorflow/tensorflow/issues/44176#issuecomment-783768033
|
||||
ENV LD_PRELOAD=/usr/lib/x86_64-linux-gnu/libtcmalloc_minimal.so.4
|
||||
|
||||
# Enable userspace DNS cache
|
||||
ENV GCS_RESOLVE_REFRESH_SECS=60
|
||||
ENV GCS_REQUEST_CONNECTION_TIMEOUT_SECS=300
|
||||
ENV GCS_METADATA_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_READ_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_WRITE_REQUEST_TIMEOUT_SECS=600
|
||||
# Each opened GCS file takes GCS_READ_CACHE_BLOCK_SIZE_MB of RAM, reduce the
|
||||
# value from the default 64MB to 8MB to decrease memory footprint.
|
||||
ENV GCS_READ_CACHE_BLOCK_SIZE_MB=8
|
||||
+13
@@ -0,0 +1,13 @@
|
||||
FROM gcr.io/automl-migration-test/movinet-base:latest
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
RUN wget https://raw.githubusercontent.com/tensorflow/models/954dd73bffd43174bd3ca26a4a34abebe4147570/official/projects/movinet/tools/export_saved_model.py \
|
||||
-O /usr/local/lib/python3.9/dist-packages/official/projects/movinet/tools/export_saved_model.py
|
||||
|
||||
WORKDIR /automl_vision
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/util"
|
||||
|
||||
ENTRYPOINT ["python3", "-m", "official.projects.movinet.tools.export_saved_model"]
|
||||
+18
@@ -0,0 +1,18 @@
|
||||
FROM gcr.io/automl-migration-test/movinet-base:latest
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
RUN pip install flask==2.3.2
|
||||
RUN pip install waitress==2.1.2
|
||||
|
||||
RUN mkdir -p /automl_vision/movinet/serving
|
||||
COPY model_oss/movinet/serving /automl_vision/movinet/serving
|
||||
COPY model_oss/util /automl_vision/util
|
||||
|
||||
WORKDIR /automl_vision
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/util"
|
||||
|
||||
ENTRYPOINT ["flask", "--app", "movinet.serving.serving_main", "run"]
|
||||
CMD ["--host=0.0.0.0", "--port=8501"]
|
||||
+18
@@ -0,0 +1,18 @@
|
||||
FROM gcr.io/automl-migration-test/movinet-base:latest
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
RUN mkdir -p /automl_vision/movinet
|
||||
COPY model_oss/movinet/*.py /automl_vision/movinet/
|
||||
COPY model_oss/util /automl_vision/util
|
||||
|
||||
WORKDIR /automl_vision
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/util"
|
||||
|
||||
# Run pylint to validate code.
|
||||
COPY .pylintrc /automl_vision/.pylintrc
|
||||
RUN find . -type f -name "*.py" | xargs pylint --rcfile=./.pylintrc --errors-only
|
||||
|
||||
ENTRYPOINT ["python3","movinet/train.py"]
|
||||
+142
@@ -0,0 +1,142 @@
|
||||
"""Main executable for MoViNet online / batch predictions."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
import json
|
||||
import os
|
||||
|
||||
from absl import app
|
||||
from absl import logging
|
||||
import flask
|
||||
import tensorflow as tf
|
||||
import waitress
|
||||
|
||||
from movinet.serving import video_serving_lib
|
||||
from util import constants
|
||||
|
||||
|
||||
flask_app = flask.Flask(__name__)
|
||||
logging.set_verbosity(logging.INFO)
|
||||
|
||||
movinet_model = None
|
||||
|
||||
_BATCH_SIZE = int(os.environ.get('BATCH_SIZE', '1'))
|
||||
_NUM_FRAMES = int(os.environ.get('NUM_FRAMES', '32'))
|
||||
_FPS = float(os.environ.get('FPS', '5'))
|
||||
_OVERLAP_FRAMES = int(os.environ.get('OVERLAP_FRAMES', '24'))
|
||||
_OBJECTIVE = os.environ.get(
|
||||
'OBJECTIVE', constants.OBJECTIVE_VIDEO_CLASSIFICATION
|
||||
).lower()
|
||||
|
||||
# VAR parameters.
|
||||
_CONFIDENCE_THRESHOLD = float(os.environ.get('CONFIDENCE_THRESHOLD', '0.5'))
|
||||
_MIN_GAP_TIME = float(os.environ.get('MIN_GAP_TIME', '1.5'))
|
||||
|
||||
|
||||
def load_movinet_model() -> None:
|
||||
model_path = os.environ.get('MODEL_PATH')
|
||||
|
||||
if not model_path:
|
||||
raise app.UsageError('Missing MODEL_PATH environment variable.')
|
||||
|
||||
# We just reload the weights of the fine-tuned diffusion model.
|
||||
logging.info('Initialize finetuned models from: %s', model_path)
|
||||
global movinet_model
|
||||
movinet_model = tf.saved_model.load(model_path)
|
||||
|
||||
|
||||
load_movinet_model()
|
||||
|
||||
|
||||
def error(message: str) -> str:
|
||||
"""Returns a JSON representing an error response."""
|
||||
return json.dumps({
|
||||
'success': False,
|
||||
'error': message,
|
||||
})
|
||||
|
||||
|
||||
# The health check route is required for docker deployment in google cloud.
|
||||
@flask_app.route('/ping')
|
||||
def ping() -> flask.Response:
|
||||
"""Health checks."""
|
||||
return flask.Response(status=200)
|
||||
|
||||
|
||||
# The return should be `Response` for docker deployment in google cloud.
|
||||
@flask_app.route('/predict', methods=['GET', 'POST'])
|
||||
def predict_model() -> flask.Response:
|
||||
"""Predictions."""
|
||||
if flask.request.method == 'POST':
|
||||
contents = flask.request.get_json(force=True)
|
||||
|
||||
logging.info('The input contents are: %s', contents)
|
||||
instances = contents.get('instances', [])
|
||||
|
||||
try:
|
||||
predictions = []
|
||||
for instance in instances:
|
||||
executor = video_serving_lib.parse_request(instance)
|
||||
prediction = executor.get_prediction(
|
||||
movinet_model,
|
||||
_BATCH_SIZE,
|
||||
_FPS,
|
||||
_NUM_FRAMES,
|
||||
_OVERLAP_FRAMES,
|
||||
_OBJECTIVE,
|
||||
)
|
||||
if _OBJECTIVE == constants.OBJECTIVE_VIDEO_CLASSIFICATION:
|
||||
prediction = video_serving_lib.postprocess_vcn(prediction)
|
||||
elif _OBJECTIVE == constants.OBJECTIVE_VIDEO_ACTION_RECOGNITION:
|
||||
prediction = video_serving_lib.postprocess_var(
|
||||
executor.windows, prediction, _CONFIDENCE_THRESHOLD, _MIN_GAP_TIME
|
||||
)
|
||||
predictions.append(prediction)
|
||||
except ValueError as e:
|
||||
return flask.Response(
|
||||
error(str(e)), status=500, mimetype='application/json'
|
||||
)
|
||||
|
||||
return flask.Response(
|
||||
response=json.dumps({
|
||||
'success': True,
|
||||
'predictions': predictions,
|
||||
}),
|
||||
status=200,
|
||||
mimetype='application/json',
|
||||
)
|
||||
else:
|
||||
return flask.Response(
|
||||
response=json.dumps({
|
||||
'success': True,
|
||||
'isalive': movinet_model is not None,
|
||||
}),
|
||||
status=200,
|
||||
mimetype='application/json',
|
||||
)
|
||||
|
||||
|
||||
def main(argv: Sequence[str]) -> None:
|
||||
if len(argv) > 1:
|
||||
raise app.UsageError('Too many command-line arguments.')
|
||||
# This is used when running locally only. When deploying to Google App
|
||||
# Engine, a webserver process such as Gunicorn will serve the app.
|
||||
# # Debug deployment.
|
||||
# flask_app.run(host='0.0.0.0', port=8501, debug=True)
|
||||
# Prod deployment.
|
||||
if _OBJECTIVE not in [
|
||||
constants.OBJECTIVE_VIDEO_CLASSIFICATION,
|
||||
constants.OBJECTIVE_VIDEO_ACTION_RECOGNITION,
|
||||
]:
|
||||
raise app.UsageError('Objective must be vcn or var.')
|
||||
logging.info(
|
||||
'Env: batch_size: %s, num_frames: %s, fps: %s, overlap_frames: %s',
|
||||
_BATCH_SIZE,
|
||||
_NUM_FRAMES,
|
||||
_FPS,
|
||||
_OVERLAP_FRAMES,
|
||||
)
|
||||
waitress.serve(flask_app, host='0.0.0.0', port=8501)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(main)
|
||||
+462
@@ -0,0 +1,462 @@
|
||||
"""Lib for handling video prediction requests.
|
||||
|
||||
The VCN inference algorithm is as follows:
|
||||
1. Find all video frames within the given clip according to the sampling FPS.
|
||||
2. Create possibly overlapping sliding windows according to the num_frames and
|
||||
overlap_frames parameters. The last window might have a larger overlap if it
|
||||
doesn't exactly fit.
|
||||
3. Run model inference on each sliding window and compute softmax to obtain
|
||||
probabilities.
|
||||
4. Average the probabilities over all sliding windows.
|
||||
|
||||
The VAR inference algorithm is very similar to VCN, with a few differences:
|
||||
1. The last sliding window is discarded if it does not exactly fit.
|
||||
2. Instead of averaging, the postprocessing consists of temporal nonmaximal
|
||||
suppression and removing background and low-confidence labels.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import os
|
||||
from typing import Any, Dict, Optional, Sequence, Union, cast
|
||||
|
||||
from absl import logging
|
||||
import cv2
|
||||
import numpy as np
|
||||
import tensorflow as tf
|
||||
|
||||
from util import constants
|
||||
from util import fileutils
|
||||
|
||||
|
||||
_JSON_LABEL_KEY = 'label'
|
||||
_JSON_GCS_URI_KEY = 'content'
|
||||
_JSON_CONFIDENCE_KEY = 'confidence'
|
||||
_JSON_START_TIME_KEY = 'timeSegmentStart'
|
||||
_JSON_END_TIME_KEY = 'timeSegmentEnd'
|
||||
_BACKGROUND_LABEL = 0
|
||||
_JSON_REQUIRED_KEYS = [
|
||||
_JSON_GCS_URI_KEY,
|
||||
_JSON_START_TIME_KEY,
|
||||
_JSON_END_TIME_KEY,
|
||||
]
|
||||
_IMAGE_WIDTH = int(os.environ.get('IMAGE_WIDTH', '172'))
|
||||
_IMAGE_HEIGHT = int(os.environ.get('IMAGE_HEIGHT', '172'))
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class DetectionOutput:
|
||||
timestamp: float
|
||||
label: int
|
||||
confidence: float
|
||||
|
||||
def to_json_obj(self) -> Dict[str, Union[int, float]]:
|
||||
"""Encodes self as a dict for JSON serialization."""
|
||||
return {
|
||||
_JSON_LABEL_KEY: self.label,
|
||||
_JSON_START_TIME_KEY: self.timestamp,
|
||||
_JSON_END_TIME_KEY: self.timestamp,
|
||||
_JSON_CONFIDENCE_KEY: self.confidence,
|
||||
}
|
||||
|
||||
|
||||
def create_detection_output(
|
||||
timestamp: float, predictions: np.ndarray
|
||||
) -> DetectionOutput:
|
||||
label = np.argmax(predictions).item()
|
||||
confidence: float = predictions[label].item()
|
||||
return DetectionOutput(timestamp, label, confidence)
|
||||
|
||||
|
||||
class SlidingWindow:
|
||||
"""Represents a sliding window with start / end timestamps."""
|
||||
|
||||
def __init__(self, fps: float, frames: Sequence[int]):
|
||||
if not frames:
|
||||
raise ValueError('Sliding window cannot be empty.')
|
||||
self.frames = frames
|
||||
self.start_time = frames[0] / fps
|
||||
self.end_time = frames[-1] / fps
|
||||
self.frame_data: list[Optional[np.ndarray]] = []
|
||||
self.clear_frame_data()
|
||||
|
||||
def load_cache_from(self, other: SlidingWindow) -> int:
|
||||
"""Loads cache from another sliding window if possible."""
|
||||
cache_count = 0
|
||||
for i, frame in enumerate(self.frames):
|
||||
try:
|
||||
other_idx = other.frames.index(frame)
|
||||
self.frame_data[i] = other.frame_data[other_idx]
|
||||
cache_count += 1
|
||||
except ValueError:
|
||||
# Cache miss.
|
||||
pass
|
||||
return cache_count
|
||||
|
||||
def load_frames(self, video: Any) -> Sequence[np.ndarray]:
|
||||
"""Loads frames of this sliding window from a video."""
|
||||
for i, frame in enumerate(self.frames):
|
||||
if self.frame_data[i] is None:
|
||||
video.set(cv2.CAP_PROP_POS_FRAMES, frame)
|
||||
ret, frame = video.read()
|
||||
if not ret:
|
||||
raise IOError(f'Failed to read video at frame {frame}.')
|
||||
self.frame_data[i] = cv2.resize(frame, (_IMAGE_WIDTH, _IMAGE_HEIGHT))
|
||||
return cast(Sequence[np.ndarray], self.frame_data)
|
||||
|
||||
def clear_frame_data(self) -> None:
|
||||
"""Clears frame data of this sliding window to reduce memory usage."""
|
||||
self.frame_data: list[Optional[np.ndarray]] = [None] * len(self)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.frames)
|
||||
|
||||
@property
|
||||
def middle_timestamp(self) -> float:
|
||||
return (self.start_time + self.end_time) / 2
|
||||
|
||||
|
||||
def _get_sliding_windows(
|
||||
frames: Sequence[int],
|
||||
original_fps: float,
|
||||
window_size: int,
|
||||
overlap: int,
|
||||
flush_last_window: bool,
|
||||
) -> Sequence[SlidingWindow]:
|
||||
"""Computes a list of sliding windows from frames.
|
||||
|
||||
Args:
|
||||
frames: A list of frame indices.
|
||||
original_fps: Frames per second of the original video.
|
||||
window_size: Number of frames in a single window.
|
||||
overlap: Number of overlapping frames in adjacent windows.
|
||||
flush_last_window: Where to flush the last window if there are not enough
|
||||
frames left.
|
||||
|
||||
Returns:
|
||||
A list of sliding windows, each has a list of frame indices. The last two
|
||||
windows might have a larger overlap if the last window does not exactly fit
|
||||
and flush_last_window is set to True.
|
||||
|
||||
Raises:
|
||||
ValueError: Arguments are invalid.
|
||||
"""
|
||||
if window_size <= overlap:
|
||||
raise ValueError(f'Window size {window_size} <= overlap {overlap}')
|
||||
total_frames = len(frames)
|
||||
windows: list[SlidingWindow] = []
|
||||
for i in range(0, total_frames, window_size - overlap):
|
||||
if i == 0 or i + window_size <= total_frames:
|
||||
windows.append(SlidingWindow(original_fps, frames[i : i + window_size]))
|
||||
elif i + overlap < total_frames and flush_last_window:
|
||||
# Some frames in this window are not covered by the previous window.
|
||||
windows.append(
|
||||
SlidingWindow(
|
||||
original_fps, frames[total_frames - window_size : total_frames]
|
||||
)
|
||||
)
|
||||
return windows
|
||||
|
||||
|
||||
def _sample_frame_indices(
|
||||
start_time: float,
|
||||
end_time: float,
|
||||
original_fps: float,
|
||||
sample_fps: float,
|
||||
max_frames: int,
|
||||
padding_left: int = 0,
|
||||
padding_right: int = 0,
|
||||
) -> Sequence[int]:
|
||||
"""Samples frames from start_time to end_time by sample_fps.
|
||||
|
||||
Args:
|
||||
start_time: Start timestamp in seconds.
|
||||
end_time: End timestamp in seconds.
|
||||
original_fps: Frames per second of the original video.
|
||||
sample_fps: Number of frames to sample per second.
|
||||
max_frames: Total number of frames in the video.
|
||||
padding_left: Padding to add to the start in frames. Padded frames will be
|
||||
duplicates of the first frame.
|
||||
padding_right: Padding to add to the end in frames. Padded frames will be
|
||||
duplicates of the last frame.
|
||||
|
||||
Returns:
|
||||
A list of sampled frame indices.
|
||||
"""
|
||||
ret = [
|
||||
min(max_frames - 1, round(t * original_fps))
|
||||
for t in np.arange(start_time, end_time, 1 / sample_fps)
|
||||
]
|
||||
if ret:
|
||||
ret = [ret[0]] * padding_left + ret + [ret[-1]] * padding_right
|
||||
return ret
|
||||
|
||||
|
||||
class VideoPredictionExecutor:
|
||||
"""Represents a Video prediction request with a video clip."""
|
||||
|
||||
def __init__(self, gcs_uri: str, start_time: float, end_time: float):
|
||||
self._gcs_uri = gcs_uri
|
||||
self._start_time = start_time
|
||||
self._end_time = end_time
|
||||
self.windows: Sequence[SlidingWindow] = []
|
||||
self._last_window: SlidingWindow = None
|
||||
|
||||
def _read_frames_from_window(
|
||||
self, video: Any, new_window: SlidingWindow
|
||||
) -> Sequence[np.ndarray]:
|
||||
"""Reads video frames from the new window.
|
||||
|
||||
Args:
|
||||
video: Video loaded with cv2.
|
||||
new_window: A list of sorted frame indices in the new window.
|
||||
|
||||
Returns:
|
||||
Frame data from the video as a list of numpy arrays.
|
||||
|
||||
Raises:
|
||||
IOError: Failed to read video.
|
||||
"""
|
||||
# Caches frames as much as possible.
|
||||
if self._last_window is not None:
|
||||
cache_count = new_window.load_cache_from(self._last_window)
|
||||
logging.info('Cached %d frames.', cache_count)
|
||||
self._last_window.clear_frame_data()
|
||||
self._last_window = new_window
|
||||
return new_window.load_frames(video)
|
||||
|
||||
def _predict(
|
||||
self, model: Any, video: Any, batched_windows: Sequence[SlidingWindow]
|
||||
) -> np.ndarray:
|
||||
"""Run model inference on specific frames of a video.
|
||||
|
||||
Args:
|
||||
model: MoViNet model.
|
||||
video: Video loaded with cv2.
|
||||
batched_windows: A batch of sliding windows to predict. Each element is an
|
||||
integer frame index. Must have equal number of frames in each window.
|
||||
|
||||
Returns:
|
||||
Prediction results.
|
||||
|
||||
Raises:
|
||||
ValueError: Batched windows are not sorted, or do not have equal number of
|
||||
frames in each window.
|
||||
IOError: Failed to read video.
|
||||
"""
|
||||
if any(
|
||||
(
|
||||
len(window) != len(batched_windows[0])
|
||||
for window in batched_windows[1:]
|
||||
)
|
||||
):
|
||||
raise ValueError(
|
||||
'Batched windows do not have equal number of frames in each window.'
|
||||
)
|
||||
batch = []
|
||||
logging.info('Loading video frames...')
|
||||
for window in batched_windows:
|
||||
logging.info('Predict frames: %s', window.frames)
|
||||
frames = self._read_frames_from_window(video, window)
|
||||
batch.append(frames)
|
||||
input_tensor = tf.convert_to_tensor(batch, dtype=tf.float32) / 255.0
|
||||
logging.info('Predict: Input tensor shape %s', input_tensor.shape)
|
||||
predictions = model({'image': input_tensor})
|
||||
logging.info('Running softmax on predictions...')
|
||||
predictions = tf.nn.softmax(predictions, axis=1)
|
||||
return predictions.numpy()
|
||||
|
||||
def get_prediction(
|
||||
self,
|
||||
model: Any,
|
||||
batch_size: int,
|
||||
fps: float,
|
||||
num_frames: int,
|
||||
overlap_frames: int,
|
||||
objective: str,
|
||||
) -> Sequence[np.ndarray]:
|
||||
"""Predicts the video clip with the model.
|
||||
|
||||
Args:
|
||||
model: The loaded MoViNet model.
|
||||
batch_size: Batch size for prediction.
|
||||
fps: Video sampling FPS.
|
||||
num_frames: Number of frames in a single predictions. If the model is
|
||||
exported with a fixed input shape, this must match its num_frames
|
||||
dimension.
|
||||
overlap_frames: Number of overlapping frames of consecutive sliding
|
||||
windows.
|
||||
objective: A string `vcn` or `var`.
|
||||
|
||||
Returns:
|
||||
A list of floats as the prediction response.
|
||||
|
||||
Raises:
|
||||
IOError: The video fails to load.
|
||||
ValueError: Some arguments are invalid.
|
||||
"""
|
||||
if objective not in [
|
||||
constants.OBJECTIVE_VIDEO_CLASSIFICATION,
|
||||
constants.OBJECTIVE_VIDEO_ACTION_RECOGNITION,
|
||||
]:
|
||||
raise ValueError(f'{objective} objective is not supported.')
|
||||
|
||||
# cv2 expects a local path so we need to download the video from GCS.
|
||||
local_file_path = fileutils.generate_tmp_path(
|
||||
os.path.splitext(self._gcs_uri)[1]
|
||||
)
|
||||
logging.info('Downloading %s to %s...', self._gcs_uri, local_file_path)
|
||||
fileutils.download_gcs_file_to_local(self._gcs_uri, local_file_path)
|
||||
logging.info('Download %s complete.', self._gcs_uri)
|
||||
|
||||
# Loads video.
|
||||
video = cv2.VideoCapture(local_file_path)
|
||||
total_frames = video.get(cv2.CAP_PROP_FRAME_COUNT)
|
||||
original_fps = video.get(cv2.CAP_PROP_FPS)
|
||||
if not original_fps:
|
||||
# 0 or None indicates the video is invalid.
|
||||
raise IOError(f'Failed to load {self._gcs_uri}.')
|
||||
video_length = total_frames / original_fps
|
||||
self._start_time = max(0, self._start_time)
|
||||
self._end_time = min(video_length, self._end_time)
|
||||
padding = (
|
||||
(num_frames // 2)
|
||||
if objective == constants.OBJECTIVE_VIDEO_ACTION_RECOGNITION
|
||||
else 0
|
||||
)
|
||||
|
||||
# Computes sliding windows.
|
||||
frame_indices = _sample_frame_indices(
|
||||
self._start_time,
|
||||
self._end_time,
|
||||
original_fps,
|
||||
fps,
|
||||
total_frames,
|
||||
padding,
|
||||
padding,
|
||||
)
|
||||
logging.info('Frame indices: %s', frame_indices)
|
||||
self.windows = _get_sliding_windows(
|
||||
frame_indices,
|
||||
original_fps,
|
||||
num_frames,
|
||||
overlap_frames,
|
||||
objective != 'var',
|
||||
)
|
||||
if not self.windows:
|
||||
raise ValueError(
|
||||
f'No sliding windows found from {self._start_time} to'
|
||||
f' {self._end_time}.'
|
||||
)
|
||||
self._last_window = None
|
||||
|
||||
# Runs inference.
|
||||
predictions = []
|
||||
for i in range(0, len(self.windows), batch_size):
|
||||
predictions.extend(
|
||||
self._predict(model, video, self.windows[i : i + batch_size])
|
||||
)
|
||||
return predictions
|
||||
|
||||
|
||||
def parse_request(req_json: Any) -> VideoPredictionExecutor:
|
||||
"""Parses VideoPredictionExecutor from request JSON object.
|
||||
|
||||
Args:
|
||||
req_json: Request JSON object.
|
||||
|
||||
Returns:
|
||||
Parsed VideoPredictionExecutor.
|
||||
|
||||
Raises:
|
||||
ValueError: Request JSON object is invalid.
|
||||
"""
|
||||
for key in _JSON_REQUIRED_KEYS:
|
||||
if key not in req_json:
|
||||
raise ValueError(f'{key} not found in {req_json}.')
|
||||
gcs_uri = req_json[_JSON_GCS_URI_KEY]
|
||||
start_time = float(req_json[_JSON_START_TIME_KEY].removesuffix('s'))
|
||||
end_time = float(req_json[_JSON_END_TIME_KEY].removesuffix('s'))
|
||||
return VideoPredictionExecutor(gcs_uri, start_time, end_time)
|
||||
|
||||
|
||||
def postprocess_vcn(predictions: Sequence[np.ndarray]) -> Sequence[float]:
|
||||
"""Aggregates VCN predictions of sliding windows."""
|
||||
return np.mean(predictions, axis=0).tolist()
|
||||
|
||||
|
||||
def temporal_nonmaximal_suppression(
|
||||
detections: Sequence[DetectionOutput], min_gap_time: float
|
||||
) -> Sequence[DetectionOutput]:
|
||||
"""Nonmaximal suppression for key frame detection.
|
||||
|
||||
For consecutive packets of the same label within a pre-defined duration, we
|
||||
only keep the one with the highest confidence score. Such duration can be
|
||||
determined by performing data analysis on users' dataset.
|
||||
|
||||
Args:
|
||||
detections: A list of DetectionOutputs.
|
||||
min_gap_time: Minimum time between consecutive key frames of the same label
|
||||
in seconds.
|
||||
|
||||
Returns:
|
||||
DetectionOutput after nonmaximal suppression sorted in ascending timestamps.
|
||||
"""
|
||||
max_label = max([detection.label for detection in detections])
|
||||
prev_detections: list[Optional[DetectionOutput]] = [None] * (max_label + 1)
|
||||
ret: list[DetectionOutput] = []
|
||||
by_time = lambda x: x.timestamp
|
||||
for detection in sorted(detections, key=by_time):
|
||||
prev_detection = prev_detections[detection.label]
|
||||
prev_detections[detection.label] = detection
|
||||
if not prev_detection:
|
||||
continue
|
||||
if detection.timestamp - prev_detection.timestamp > min_gap_time:
|
||||
ret.append(prev_detection)
|
||||
continue
|
||||
detection.confidence = max(detection.confidence, prev_detection.confidence)
|
||||
ret.extend((d for d in prev_detections if d is not None))
|
||||
return sorted(ret, key=by_time)
|
||||
|
||||
|
||||
def postprocess_var(
|
||||
windows: Sequence[SlidingWindow],
|
||||
predictions: Sequence[np.ndarray],
|
||||
confidence_threshold: float,
|
||||
min_gap_time: float,
|
||||
) -> Sequence[Dict[str, Any]]:
|
||||
"""Generates a list of detected keyframes from sliding window predictions.
|
||||
|
||||
Args:
|
||||
windows: Sliding windows.
|
||||
predictions: A list of predictions of sliding windows.
|
||||
confidence_threshold: Only probabilities greater than this threshold will
|
||||
contribute to the final result.
|
||||
min_gap_time: Minimum time between consecutive key frames of the same label
|
||||
in seconds. Used in temporal nonmaximal suppression.
|
||||
|
||||
Returns:
|
||||
A sequence of dictionaries, each item has the following keys:
|
||||
- label: Integer label of the detection result.
|
||||
- timeSegmentStart: Start timestamp in seconds.
|
||||
- timeSegmentEnd: End timestamp in seconds. Always equals timeSegmentStart.
|
||||
"""
|
||||
if len(windows) != len(predictions):
|
||||
raise ValueError('Mismatched # of windows with # of predictions.')
|
||||
|
||||
# Creates detection results from windows, filtering out the background label.
|
||||
detections = [
|
||||
create_detection_output(window.middle_timestamp, predictions[i])
|
||||
for i, window in enumerate(windows)
|
||||
]
|
||||
|
||||
# Temporal nonmaximal suppression.
|
||||
detections = temporal_nonmaximal_suppression(detections, min_gap_time)
|
||||
|
||||
# Filters out ones with low confidence and the background label.
|
||||
return [
|
||||
x.to_json_obj()
|
||||
for x in detections
|
||||
if x.label != _BACKGROUND_LABEL and x.confidence > confidence_threshold
|
||||
]
|
||||
@@ -0,0 +1,210 @@
|
||||
"""Main executable for MoViNet docker."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Sequence, Any
|
||||
|
||||
from absl import app
|
||||
from absl import flags
|
||||
from absl import logging
|
||||
import gin
|
||||
import hypertune
|
||||
import tensorflow as tf
|
||||
|
||||
from util import constants
|
||||
from util import hypertune_utils
|
||||
from official.common import distribute_utils
|
||||
from official.common import flags as tfm_flags
|
||||
from official.core import task_factory
|
||||
from official.core import train_lib
|
||||
from official.core import train_utils
|
||||
from official.modeling import performance
|
||||
# Import movinet libraries to register the backbone and model into tf.vision
|
||||
# model garden factory.
|
||||
# pylint: disable=unused-import
|
||||
from official.projects.movinet.modeling import movinet
|
||||
from official.projects.movinet.modeling import movinet_model
|
||||
from official.vision import registry_imports
|
||||
# pylint: enable=unused-import
|
||||
|
||||
|
||||
FLAGS = flags.FLAGS
|
||||
|
||||
_FILE_TYPE_TFRECORD = 'tfrecord'
|
||||
|
||||
_LEARNING_RATE = flags.DEFINE_float(
|
||||
'learning_rate', None, 'The learning rate of this training job.'
|
||||
)
|
||||
|
||||
_NUM_CLASSES = flags.DEFINE_integer(
|
||||
'num_classes', None, 'The number of classes.'
|
||||
)
|
||||
|
||||
_INIT_CHECKPOINT = flags.DEFINE_string(
|
||||
'init_checkpoint', None, 'The initial checkpoint of this training job.'
|
||||
)
|
||||
|
||||
_INPUT_TRAIN_DATA_PATH = flags.DEFINE_string(
|
||||
'input_train_data_path', None, 'Input train data path.'
|
||||
)
|
||||
|
||||
_INPUT_VALIDATION_DATA_PATH = flags.DEFINE_string(
|
||||
'input_validation_data_path', None, 'Input validation data path.'
|
||||
)
|
||||
|
||||
_GLOBAL_BATCH_SIZE = flags.DEFINE_integer(
|
||||
'global_batch_size', None, 'Global batch size.'
|
||||
)
|
||||
|
||||
_PREFETCH_BUFFER_SIZE = flags.DEFINE_integer(
|
||||
'prefetch_buffer_size', None, 'Prefetch buffer size.'
|
||||
)
|
||||
|
||||
_SHUFFLE_BUFFER_SIZE = flags.DEFINE_integer(
|
||||
'shuffle_buffer_size', None, 'Shuffle buffer size.'
|
||||
)
|
||||
|
||||
_TRAIN_STEPS = flags.DEFINE_integer('train_steps', None, 'Train steps.')
|
||||
_LOG_LEVEL = flags.DEFINE_enum(
|
||||
'log_level',
|
||||
'INFO',
|
||||
['FATAL', 'ERROR', 'WARNING', 'INFO', 'DEBUG'],
|
||||
'Log level.',
|
||||
)
|
||||
|
||||
|
||||
def parse_params() -> Any:
|
||||
"""Parses parameters."""
|
||||
gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
|
||||
params = train_utils.parse_configuration(FLAGS, lock_return=False)
|
||||
if _INIT_CHECKPOINT.value:
|
||||
params.task.init_checkpoint = _INIT_CHECKPOINT.value
|
||||
params.task.init_checkpoint_modules = 'backbone'
|
||||
if _NUM_CLASSES.value:
|
||||
params.task.model.num_classes = _NUM_CLASSES.value
|
||||
params.task.train_data.num_classes = _NUM_CLASSES.value
|
||||
params.task.validation_data.num_classes = _NUM_CLASSES.value
|
||||
# If users set input train/validation data path, we assume the data are
|
||||
# converted from data converter as tfrecord. Users can use tfds by writing
|
||||
# their own config directly, and no need to override this parameter.
|
||||
if _INPUT_TRAIN_DATA_PATH.value:
|
||||
params.task.train_data.input_path = _INPUT_TRAIN_DATA_PATH.value
|
||||
params.task.train_data.file_type = _FILE_TYPE_TFRECORD
|
||||
params.task.train_data.tfds_name = ''
|
||||
if _INPUT_VALIDATION_DATA_PATH.value:
|
||||
params.task.validation_data.input_path = _INPUT_VALIDATION_DATA_PATH.value
|
||||
params.task.validation_data.file_type = _FILE_TYPE_TFRECORD
|
||||
params.task.validation_data.tfds_name = ''
|
||||
if _GLOBAL_BATCH_SIZE.value:
|
||||
params.task.train_data.global_batch_size = _GLOBAL_BATCH_SIZE.value
|
||||
params.task.validation_data.global_batch_size = _GLOBAL_BATCH_SIZE.value
|
||||
if _PREFETCH_BUFFER_SIZE.value:
|
||||
params.task.train_data.prefetch_buffer_size = _PREFETCH_BUFFER_SIZE.value
|
||||
params.task.validation_data.prefetch_buffer_size = (
|
||||
_PREFETCH_BUFFER_SIZE.value
|
||||
)
|
||||
if _SHUFFLE_BUFFER_SIZE.value:
|
||||
params.task.train_data.shuffle_buffer_size = _SHUFFLE_BUFFER_SIZE.value
|
||||
if _TRAIN_STEPS.value:
|
||||
params.trainer.train_steps = _TRAIN_STEPS.value
|
||||
if _LEARNING_RATE.value:
|
||||
logging.info('Updating learning_rate: %s', _LEARNING_RATE.value)
|
||||
# Use `get` method of train_utils.hyperparams.OneOfConfig to get learning
|
||||
# rate config.
|
||||
learning_rate = params.trainer.optimizer_config.learning_rate.get()
|
||||
if hasattr(learning_rate, 'initial_learning_rate'):
|
||||
learning_rate.initial_learning_rate = _LEARNING_RATE.value
|
||||
else:
|
||||
logging.warning('Cannot set learning rate for %s', learning_rate)
|
||||
# Set default params for best checkpoints.
|
||||
params.trainer.best_checkpoint_export_subdir = constants.BEST_CKPT_DIRNAME
|
||||
params.trainer.best_checkpoint_metric_comp = constants.BEST_CKPT_METRIC_COMP
|
||||
params.trainer.best_checkpoint_eval_metric = (
|
||||
constants.VIDEO_CLASSIFICATION_BEST_EVAL_METRIC
|
||||
)
|
||||
return params
|
||||
|
||||
|
||||
def main(argv: Sequence[str]) -> None:
|
||||
logging.set_verbosity(_LOG_LEVEL.value)
|
||||
if len(argv) > 1:
|
||||
raise app.UsageError('Too many command-line arguments.')
|
||||
params = parse_params()
|
||||
logging.info('The actual training parameters are:\n%s', params.as_dict())
|
||||
model_dir: str = os.path.join(
|
||||
FLAGS.model_dir,
|
||||
constants.TRIAL_PREFIX + hypertune_utils.get_trial_id_from_environment(),
|
||||
)
|
||||
logging.info('model_dir: %s', model_dir)
|
||||
|
||||
if 'train' in FLAGS.mode:
|
||||
# Pure eval modes do not output yaml files. Otherwise continuous eval job
|
||||
# may race against the train job for writing the same file.
|
||||
train_utils.serialize_config(params, model_dir)
|
||||
|
||||
# Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
|
||||
# can have significant impact on model speeds by utilizing float16 in case of
|
||||
# GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
|
||||
# dtype is float16
|
||||
if params.runtime.mixed_precision_dtype:
|
||||
performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
|
||||
distribution_strategy = distribute_utils.get_distribution_strategy(
|
||||
distribution_strategy=params.runtime.distribution_strategy,
|
||||
all_reduce_alg=params.runtime.all_reduce_alg,
|
||||
num_gpus=params.runtime.num_gpus,
|
||||
tpu_address=params.runtime.tpu,
|
||||
)
|
||||
|
||||
# Create task and run experiment.
|
||||
with distribution_strategy.scope():
|
||||
task = task_factory.get_task(params.task, logging_dir=model_dir)
|
||||
|
||||
train_lib.run_experiment(
|
||||
distribution_strategy=distribution_strategy,
|
||||
task=task,
|
||||
mode=FLAGS.mode,
|
||||
params=params,
|
||||
model_dir=model_dir,
|
||||
)
|
||||
|
||||
train_utils.save_gin_config(FLAGS.mode, model_dir)
|
||||
|
||||
eval_metric_name = constants.VIDEO_CLASSIFICATION_BEST_EVAL_METRIC
|
||||
|
||||
eval_filepath = os.path.join(
|
||||
model_dir, constants.BEST_CKPT_DIRNAME, constants.BEST_CKPT_EVAL_FILENAME
|
||||
)
|
||||
logging.info('Load eval metrics from: %s.', eval_filepath)
|
||||
|
||||
with tf.io.gfile.GFile(eval_filepath, 'rb') as f:
|
||||
eval_metric_results = json.load(f)
|
||||
logging.info('eval metrics are: %s.', eval_metric_results)
|
||||
if (
|
||||
eval_metric_name in eval_metric_results
|
||||
and constants.BEST_CKPT_STEP_NAME in eval_metric_results
|
||||
):
|
||||
hp_metric = eval_metric_results[eval_metric_name]
|
||||
hp_step = int(eval_metric_results[constants.BEST_CKPT_STEP_NAME])
|
||||
hpt = hypertune.HyperTune()
|
||||
hpt.report_hyperparameter_tuning_metric(
|
||||
hyperparameter_metric_tag=constants.HP_METRIC_TAG,
|
||||
metric_value=hp_metric,
|
||||
global_step=hp_step,
|
||||
)
|
||||
logging.info(
|
||||
'Send HP metric: %f and steps %d to hyperparameter tuning.',
|
||||
hp_metric,
|
||||
hp_step,
|
||||
)
|
||||
else:
|
||||
logging.info(
|
||||
'Either %s or %s is not included in the evaluation results: %s.',
|
||||
eval_metric_name,
|
||||
constants.BEST_CKPT_STEP_NAME,
|
||||
eval_metric_results,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
tfm_flags.define_flags()
|
||||
app.run(main)
|
||||
+67
@@ -0,0 +1,67 @@
|
||||
# Dockerfile for basic serving dockers for OpenCLIP.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/open_clip/dockerfile/serve.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
# Switch to this base image for gpu serve.
|
||||
FROM pytorch/torchserve:0.7.1-gpu
|
||||
|
||||
USER root
|
||||
|
||||
# Install tools.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
wget \
|
||||
vim
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
ENV infer_port=7080
|
||||
ENV mng_port=7081
|
||||
ENV model_name="transformers_serving"
|
||||
ENV PATH="/home/model-server/:${PATH}"
|
||||
|
||||
# Install libraries.
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
RUN pip install torch==1.13.1
|
||||
RUN pip install open_clip_torch==2.20.0
|
||||
RUN pip install pillow==9.5.0
|
||||
RUN pip install google-cloud-storage==2.7.0
|
||||
|
||||
# Copy model artifacts.
|
||||
COPY model_oss/open_clip/handler.py /home/model-server/handler.py
|
||||
COPY model_oss/util/ /home/model-server/util/
|
||||
ENV PYTHONPATH /home/model-server/
|
||||
|
||||
# Create torchserve configuration file.
|
||||
RUN echo \
|
||||
"default_response_timeout=1800\n" \
|
||||
"service_envelope=json\n" \
|
||||
"inference_address=http://0.0.0.0:${infer_port}\n" \
|
||||
"management_address=http://0.0.0.0:${mng_port}" >> /home/model-server/config.properties
|
||||
|
||||
# Expose ports.
|
||||
EXPOSE ${infer_port}
|
||||
EXPOSE ${mng_port}
|
||||
|
||||
# Archive model artifacts and dependencies.
|
||||
# Do not set --model-file and --serialized-file because model and checkpoint will be dynamically loaded in handler.py.
|
||||
RUN torch-model-archiver \
|
||||
--model-name=${model_name} \
|
||||
--version=1.0 \
|
||||
--handler=/home/model-server/handler.py \
|
||||
--runtime=python3 \
|
||||
--export-path=/home/model-server/model-store \
|
||||
--archive-format=default \
|
||||
--force
|
||||
|
||||
# Run Torchserve HTTP serve to respond to prediction requests.
|
||||
CMD ["torchserve", "--start", \
|
||||
"--ts-config", "/home/model-server/config.properties", \
|
||||
"--models", "${model_name}=${model_name}.mar", \
|
||||
"--model-store", "/home/model-server/model-store"]
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
# Dockerfile for training dockers with OpenCLIP.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/open_clilp/dockerfile/train.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM pytorch/pytorch:2.0.0-cuda11.7-cudnn8-devel
|
||||
|
||||
# Install tools.
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
RUN apt-get update
|
||||
RUN apt-get install -y --no-install-recommends apt-utils
|
||||
RUN apt-get install -y --no-install-recommends curl
|
||||
RUN apt-get install -y --no-install-recommends wget
|
||||
RUN apt-get install -y --no-install-recommends git
|
||||
RUN apt-get install -y --no-install-recommends jq
|
||||
RUN apt-get install -y --no-install-recommends gnupg
|
||||
RUN apt-get install -y --no-install-recommends build-essential
|
||||
|
||||
ENV PIP_ROOT_USER_ACTION=ignore
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
wget \
|
||||
vim
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Prepare artifacts.
|
||||
WORKDIR /workspace
|
||||
RUN git clone --branch main https://github.com/mlfoundations/open_clip.git
|
||||
WORKDIR ./open_clip
|
||||
RUN git reset --hard 67e5e5ec8741281eb9b30f640c26f91c666308b7
|
||||
|
||||
# Install libraries.
|
||||
RUN pip install webdataset==0.2.5
|
||||
RUN pip install regex==2023.6.3
|
||||
RUN pip install ftfy==6.1.1
|
||||
RUN pip install pandas==2.0.3
|
||||
RUN pip install braceexpand==0.1.7
|
||||
RUN pip install huggingface_hub==0.16.4
|
||||
RUN pip install transformers==4.31.0
|
||||
RUN pip install timm==0.9.2
|
||||
RUN pip install fsspec==2023.6.0
|
||||
RUN pip install sentencepiece==0.1.99
|
||||
RUN pip install protobuf==3.20.3
|
||||
RUN pip install tensorboard==2.12.2
|
||||
|
||||
# Switch work folder for training.
|
||||
WORKDIR ./src
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Custom handler for OpenCLIP model."""
|
||||
|
||||
# pylint:disable=g-importing-member
|
||||
import enum
|
||||
import logging
|
||||
import os
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import open_clip
|
||||
import torch
|
||||
from ts.torch_handler.base_handler import BaseHandler
|
||||
|
||||
from google3.cloud.ml.applications.vision.model_garden.model_oss.util import constants
|
||||
from google3.cloud.ml.applications.vision.model_garden.model_oss.util import fileutils
|
||||
from google3.cloud.ml.applications.vision.model_garden.model_oss.util import image_format_converter
|
||||
|
||||
|
||||
@enum.unique
|
||||
class Precision(enum.Enum):
|
||||
AMP = "amp"
|
||||
AMP_BF16 = "amp_bf16"
|
||||
AMP_BFLOAT16 = "amp_bfloat16"
|
||||
BF16 = "bf16"
|
||||
FP16 = "fp16"
|
||||
PURE_BF16 = "pure_bf16"
|
||||
PURE_FP16 = "pure_fp16"
|
||||
FP32 = "fp32"
|
||||
|
||||
|
||||
# Supported checkpoint&model pairs:
|
||||
# https://github.com/mlfoundations/open_clip#pretrained-model-interface
|
||||
_DEFAULT_CHECKPOINT = "openai"
|
||||
_DEFAULT_MODEL = "RN50"
|
||||
_DEFAULT_PRECISION = Precision.AMP
|
||||
_ZERO_CLASSIFICATION = "zero-shot-image-classification"
|
||||
_FEATURE_EMBEDDING = "feature-embedding"
|
||||
_VALID_TASKS = frozenset([_ZERO_CLASSIFICATION, _FEATURE_EMBEDDING])
|
||||
|
||||
_IMAGE_KEY = "image"
|
||||
_TEXT_KEY = "text"
|
||||
_IMAGE_FEATURES_KEY = "image_features"
|
||||
_TEXT_FEATURES_KEY = "text_features"
|
||||
|
||||
|
||||
class OpenclipHandler(BaseHandler):
|
||||
"""Custom handler for OpenCLIP."""
|
||||
|
||||
def initialize(self, context: Any):
|
||||
"""Custom initialize."""
|
||||
|
||||
properties = context.system_properties
|
||||
self.map_location = (
|
||||
"cuda"
|
||||
if torch.cuda.is_available() and properties.get("gpu_id") is not None
|
||||
else "cpu"
|
||||
)
|
||||
self.device = torch.device(
|
||||
self.map_location + ":" + str(properties.get("gpu_id"))
|
||||
if torch.cuda.is_available() and properties.get("gpu_id") is not None
|
||||
else self.map_location
|
||||
)
|
||||
self.manifest = context.manifest
|
||||
|
||||
model_name = os.environ.get("MODEL", _DEFAULT_MODEL)
|
||||
precision = os.environ.get("PRECISION", _DEFAULT_PRECISION)
|
||||
checkpoint = os.environ.get("CHECKPOINT", _DEFAULT_CHECKPOINT)
|
||||
self.task = os.environ.get("TASK", _FEATURE_EMBEDDING)
|
||||
if self.task not in _VALID_TASKS:
|
||||
raise ValueError(f"Invalid task: {self.task}.")
|
||||
logging.info(
|
||||
"Handler initializing task:%s, model:%s, precision:%s, checkpoint:%s",
|
||||
self.task,
|
||||
model_name,
|
||||
precision,
|
||||
checkpoint,
|
||||
)
|
||||
|
||||
if checkpoint != _DEFAULT_CHECKPOINT:
|
||||
local_fname = os.path.join(constants.LOCAL_MODEL_DIR, "model.pt")
|
||||
fileutils.download_gcs_file_to_local(checkpoint, local_fname)
|
||||
checkpoint = local_fname
|
||||
|
||||
self.model, _, self.preprocessor = open_clip.create_model_and_transforms(
|
||||
model_name, pretrained=checkpoint, precision=precision
|
||||
)
|
||||
self.tokenizer = open_clip.get_tokenizer(model_name)
|
||||
|
||||
self.initialized = True
|
||||
|
||||
def preprocess(self, data: Any) -> List[Dict[str, Any]]:
|
||||
"""Preprocess input data."""
|
||||
logging.info("preprocessing: %d instances received.", len(data))
|
||||
processed_list = []
|
||||
for item in data:
|
||||
sample = {}
|
||||
if _IMAGE_KEY in item:
|
||||
sample[_IMAGE_KEY] = self.preprocessor(
|
||||
image_format_converter.base64_to_image(item[_IMAGE_KEY])
|
||||
).unsqueeze(0)
|
||||
if _TEXT_KEY in item:
|
||||
sample[_TEXT_KEY] = self.tokenizer(item[_TEXT_KEY])
|
||||
processed_list.append(sample)
|
||||
return processed_list
|
||||
|
||||
def inference(
|
||||
self, data: List[Dict[str, Any]], *args, **kwargs
|
||||
) -> List[Dict[str, Any]]:
|
||||
feature_list = []
|
||||
with torch.no_grad(), torch.cuda.amp.autocast():
|
||||
for item in data:
|
||||
sample = {}
|
||||
if _IMAGE_KEY in item:
|
||||
sample[_IMAGE_FEATURES_KEY] = self.model.encode_image(
|
||||
item[_IMAGE_KEY]
|
||||
)
|
||||
if _TEXT_KEY in item:
|
||||
sample[_TEXT_FEATURES_KEY] = self.model.encode_text(item[_TEXT_KEY])
|
||||
feature_list.append(sample)
|
||||
return feature_list
|
||||
|
||||
def postprocess(self, features: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
"""Postprocess the image/text featreus for downstream task."""
|
||||
preds = []
|
||||
if self.task == _FEATURE_EMBEDDING:
|
||||
for item in features:
|
||||
preds.append({k: v.tolist() for k, v in item.items()})
|
||||
elif self.task == _ZERO_CLASSIFICATION:
|
||||
for item in features:
|
||||
image_features = item.get(_IMAGE_FEATURES_KEY, None)
|
||||
text_features = item.get(_TEXT_FEATURES_KEY, None)
|
||||
if image_features is None or text_features is None:
|
||||
raise ValueError(
|
||||
"Missing input for {} task. {} received.".format(
|
||||
_ZERO_CLASSIFICATION, item.keys()
|
||||
)
|
||||
)
|
||||
image_features /= image_features.norm(dim=-1, keepdim=True)
|
||||
text_features /= text_features.norm(dim=-1, keepdim=True)
|
||||
text_probs = (100.0 * image_features @ text_features.T).softmax(dim=-1)
|
||||
preds.append(text_probs.tolist())
|
||||
|
||||
return preds
|
||||
+142
@@ -0,0 +1,142 @@
|
||||
"""Causal language modeling with LoRA models."""
|
||||
|
||||
# pylint: disable=g-importing-member
|
||||
|
||||
from datasets import load_dataset
|
||||
from peft import get_peft_model
|
||||
from peft import LoraConfig
|
||||
import torch
|
||||
from torch import nn
|
||||
import transformers
|
||||
from transformers import AutoModelForCausalLM
|
||||
from transformers import AutoTokenizer
|
||||
from transformers import BitsAndBytesConfig
|
||||
from transformers import TrainingArguments
|
||||
from util import constants
|
||||
|
||||
|
||||
def finetune_causal_language_modeling(
|
||||
pretrained_model_id: str,
|
||||
dataset_name: str,
|
||||
output_dir: str,
|
||||
precision_mode: str = None,
|
||||
lora_rank: int = 16,
|
||||
lora_alpha: int = 32,
|
||||
lora_dropout: float = 0.05,
|
||||
warmup_steps: int = 10,
|
||||
max_steps: int = 10,
|
||||
learning_rate: float = 2e-4,
|
||||
local_pretrained_model_id: str = None,
|
||||
) -> None:
|
||||
"""Finetunes causal language modelings."""
|
||||
if precision_mode == constants.PRECISION_MODE_32:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
local_pretrained_model_id
|
||||
if local_pretrained_model_id
|
||||
else pretrained_model_id,
|
||||
torch_dtype=torch.float32,
|
||||
device_map="auto",
|
||||
)
|
||||
elif precision_mode == constants.PRECISION_MODE_16:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
local_pretrained_model_id
|
||||
if local_pretrained_model_id
|
||||
else pretrained_model_id,
|
||||
torch_dtype=torch.bfloat16,
|
||||
device_map="auto",
|
||||
)
|
||||
elif precision_mode == constants.PRECISION_MODE_8:
|
||||
quantization_config = BitsAndBytesConfig(
|
||||
load_in_8bit=True, int8_threshold=0
|
||||
)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
local_pretrained_model_id
|
||||
if local_pretrained_model_id
|
||||
else pretrained_model_id,
|
||||
torch_dtype=torch.float16,
|
||||
device_map="auto",
|
||||
quantization_config=quantization_config,
|
||||
)
|
||||
else:
|
||||
quantization_config = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_quant_type="nf4",
|
||||
bnb_4bit_compute_dtype=torch.bfloat16,
|
||||
)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
local_pretrained_model_id
|
||||
if local_pretrained_model_id
|
||||
else pretrained_model_id,
|
||||
device_map="auto",
|
||||
torch_dtype=torch.bfloat16,
|
||||
quantization_config=quantization_config,
|
||||
)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
local_pretrained_model_id
|
||||
if local_pretrained_model_id
|
||||
else pretrained_model_id
|
||||
)
|
||||
if "llama" in pretrained_model_id:
|
||||
tokenizer.pad_token = "[PAD]"
|
||||
|
||||
for param in model.parameters():
|
||||
# Freezes the model - train adapters later.
|
||||
param.requires_grad = False
|
||||
if param.ndim == 1:
|
||||
# Casts the small parameters (e.g. layernorm) to fp32 for stability.
|
||||
param.data = param.data.to(torch.float32)
|
||||
|
||||
# Reduces the number of stored activations.
|
||||
model.gradient_checkpointing_enable()
|
||||
model.enable_input_require_grads()
|
||||
|
||||
class CastOutputToFloat(nn.Sequential):
|
||||
|
||||
def forward(self, x):
|
||||
return super().forward(x).to(torch.float32)
|
||||
|
||||
model.lm_head = CastOutputToFloat(model.lm_head)
|
||||
|
||||
config = LoraConfig(
|
||||
r=lora_rank,
|
||||
lora_alpha=lora_alpha,
|
||||
target_modules=["q_proj", "v_proj"],
|
||||
lora_dropout=lora_dropout,
|
||||
bias="none",
|
||||
task_type="CAUSAL_LM",
|
||||
)
|
||||
|
||||
model = get_peft_model(model, config)
|
||||
model.print_trainable_parameters()
|
||||
|
||||
data = load_dataset(dataset_name)
|
||||
data = data.map(
|
||||
lambda samples: tokenizer(samples["quote"]),
|
||||
batched=True,
|
||||
)
|
||||
|
||||
trainer = transformers.Trainer(
|
||||
model=model,
|
||||
train_dataset=data["train"],
|
||||
args=TrainingArguments(
|
||||
per_device_train_batch_size=4,
|
||||
gradient_accumulation_steps=4,
|
||||
warmup_steps=warmup_steps,
|
||||
max_steps=max_steps,
|
||||
learning_rate=learning_rate,
|
||||
fp16=True,
|
||||
logging_steps=1,
|
||||
output_dir=output_dir,
|
||||
ddp_find_unused_parameters=False,
|
||||
),
|
||||
data_collator=transformers.DataCollatorForLanguageModeling(
|
||||
tokenizer,
|
||||
mlm=False,
|
||||
),
|
||||
)
|
||||
# Silence the warnings. Please re-enable for inference!
|
||||
model.config.use_cache = False
|
||||
trainer.train()
|
||||
|
||||
model.save_pretrained(output_dir)
|
||||
@@ -0,0 +1,21 @@
|
||||
number_of_netty_threads=32
|
||||
job_queue_size=1000
|
||||
model_store=/home/model-server/model-store
|
||||
workflow_store=/home/model-server/wf-store
|
||||
default_response_timeout=1800
|
||||
service_envelope=json
|
||||
inference_address=http://0.0.0.0:7080
|
||||
management_address=http://0.0.0.0:7081
|
||||
metrics_address=http://0.0.0.0:7082
|
||||
|
||||
models={\
|
||||
"peft_serving": {\
|
||||
"1.0": {\
|
||||
"defaultVersion": true,\
|
||||
"marName": "peft_serving.mar",\
|
||||
"minWorkers": 1,\
|
||||
"maxWorkers": 1,\
|
||||
"batchSize": 1\
|
||||
}\
|
||||
}\
|
||||
}
|
||||
+107
@@ -0,0 +1,107 @@
|
||||
# Dockerfile for PEFT Serving.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/peft/dockerfile/serve.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM pytorch/torchserve:0.7.0-gpu
|
||||
|
||||
USER root
|
||||
|
||||
ENV infer_port=7080
|
||||
ENV mng_port=7081
|
||||
ENV model_name="peft_serving"
|
||||
ENV PATH="/home/model-server/:${PATH}"
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
wget \
|
||||
vim \
|
||||
git \
|
||||
git-lfs
|
||||
RUN git lfs install
|
||||
|
||||
# Install libraries.
|
||||
ENV PIP_ROOT_USER_ACTION=ignore
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
RUN pip install --upgrade torch==2.0.1
|
||||
RUN pip install torchvision==0.15.2
|
||||
RUN pip install tokenizers==0.13.3
|
||||
RUN pip install accelerate==0.21.0
|
||||
RUN pip install sentencepiece==0.1.99
|
||||
RUN pip install grpcio-status==1.33.2
|
||||
RUN pip install protobuf==3.19.6
|
||||
RUN python3 -m pip install --no-cache-dir git+https://github.com/huggingface/peft.git
|
||||
RUN pip install datasets==2.14.4
|
||||
RUN pip install triton==2.0.0.dev20221120
|
||||
RUN pip install xformers==0.0.20
|
||||
RUN pip install google-cloud-storage==2.7.0
|
||||
RUN pip install absl-py==1.4.0
|
||||
RUN pip install scipy==1.10.1
|
||||
RUN pip install evaluate==0.4.0
|
||||
RUN pip install scikit-learn==1.2.2
|
||||
RUN pip install loralib==0.1.1
|
||||
RUN pip install bitsandbytes==0.39.0
|
||||
RUN pip install trl==0.4.4
|
||||
RUN pip install einops==0.6.1
|
||||
|
||||
# Install diffusers from source.
|
||||
RUN git clone --depth 1 --branch v0.16.1 https://github.com/huggingface/diffusers.git
|
||||
WORKDIR diffusers
|
||||
RUN pip install -e .
|
||||
WORKDIR /home/model-server
|
||||
|
||||
# Install transformers from source.
|
||||
RUN git clone --depth 1 --branch v4.31.0 https://github.com/huggingface/transformers.git
|
||||
# The patch is used to change the transformers loading model behavior:
|
||||
# 1) For models on Huggingface hub: if the model has multiple shards, each shard
|
||||
# will be downloaded separately and get deleted after loading to GPU.
|
||||
# 2) For models on local disk: if a model bin file is actually a text file
|
||||
# recording a GCS path, the model file will be downloaded and get deleted
|
||||
# after loading to GPU.
|
||||
COPY model_oss/peft/hf_transformers_lazy_download.patch /home/model-server/hf_transformers_lazy_download.patch
|
||||
WORKDIR transformers
|
||||
RUN git apply /home/model-server/hf_transformers_lazy_download.patch
|
||||
RUN pip install -e .
|
||||
WORKDIR /home/model-server
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Copy model artifacts.
|
||||
COPY model_oss/peft/handler.py /home/model-server/handler.py
|
||||
COPY model_oss/peft/config.properties /home/model-server/config.properties
|
||||
COPY model_oss/util/ /home/model-server/util/
|
||||
ENV PYTHONPATH /home/model-server/
|
||||
|
||||
# Expose ports.
|
||||
EXPOSE ${infer_port}
|
||||
EXPOSE ${mng_port}
|
||||
|
||||
# Set environments.
|
||||
ENV TASK "causal-language-modeling-lora"
|
||||
ENV BASE_MODEL_ID "openlm-research/open_llama_7b"
|
||||
ENV PRECISION_LOADING_MODE "float16"
|
||||
ENV FINETUNED_LORA_MODEL_PATH ""
|
||||
|
||||
|
||||
# Archive model artifacts and dependencies.
|
||||
# Do not set --model-file and --serialized-file because model and checkpoint
|
||||
# will be dynamically loaded in handler.py.
|
||||
RUN torch-model-archiver \
|
||||
--model-name=${model_name} \
|
||||
--version=1.0 \
|
||||
--handler=/home/model-server/handler.py \
|
||||
--runtime=python3 \
|
||||
--export-path=/home/model-server/model-store \
|
||||
--archive-format=default \
|
||||
--force
|
||||
|
||||
# Run Torchserve HTTP serve to respond to prediction requests.
|
||||
CMD ["torchserve", "--start", \
|
||||
"--ts-config", "/home/model-server/config.properties", \
|
||||
"--models", "${model_name}=${model_name}.mar", \
|
||||
"--model-store", "/home/model-server/model-store"]
|
||||
+111
@@ -0,0 +1,111 @@
|
||||
# Dockerfile for PEFT Training.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/peft/dockerfile/train.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
# Builds GPU docker image of PyTorch
|
||||
# Uses multi-staged approach to reduce size
|
||||
# Stage 1
|
||||
# Use base conda image to reduce time
|
||||
FROM continuumio/miniconda3:latest AS compile-image
|
||||
# Specify py version
|
||||
ENV PYTHON_VERSION=3.8
|
||||
# Install apt libs - copied from https://github.com/huggingface/accelerate/blob/main/docker/accelerate-gpu/Dockerfile
|
||||
RUN apt-get update && \
|
||||
apt-get install -y curl git wget software-properties-common git-lfs && \
|
||||
apt-get clean && \
|
||||
rm -rf /var/lib/apt/lists*
|
||||
|
||||
# Install audio-related libraries
|
||||
RUN apt-get update && \
|
||||
apt install -y ffmpeg
|
||||
|
||||
RUN apt install -y libsndfile1-dev
|
||||
RUN git lfs install
|
||||
|
||||
# Create our conda env - copied from https://github.com/huggingface/accelerate/blob/main/docker/accelerate-gpu/Dockerfile
|
||||
RUN conda create --name peft python=${PYTHON_VERSION} ipython jupyter pip
|
||||
RUN python3 -m pip install --no-cache-dir --upgrade pip
|
||||
|
||||
# Below is copied from https://github.com/huggingface/accelerate/blob/main/docker/accelerate-gpu/Dockerfile
|
||||
# We don't install pytorch here yet since CUDA isn't available
|
||||
# instead we use the direct torch wheel
|
||||
ENV PATH /opt/conda/envs/peft/bin:$PATH
|
||||
# Activate our bash shell
|
||||
RUN chsh -s /bin/bash
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
# Activate the conda env and install transformers + accelerate from source
|
||||
RUN source activate peft
|
||||
RUN python3 -m pip install --no-cache-dir git+https://github.com/huggingface/transformers
|
||||
RUN python3 -m pip install --no-cache-dir git+https://github.com/huggingface/accelerate
|
||||
RUN python3 -m pip install --no-cache-dir git+https://github.com/huggingface/peft#egg=peft[test]
|
||||
RUN python3 -m pip install --no-cache-dir bitsandbytes
|
||||
|
||||
# Stage 2
|
||||
FROM nvidia/cuda:11.2.2-cudnn8-devel-ubuntu20.04 AS build-image
|
||||
COPY --from=compile-image /opt/conda /opt/conda
|
||||
ENV PATH /opt/conda/bin:$PATH
|
||||
|
||||
# Install apt libs
|
||||
RUN apt-get update && \
|
||||
apt-get install -y curl git wget vim && \
|
||||
apt-get clean && \
|
||||
rm -rf /var/lib/apt/lists*
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
RUN echo "source activate peft" >> ~/.profile
|
||||
|
||||
# Install libraries.
|
||||
RUN pip install --upgrade torch==2.0.1
|
||||
RUN pip install torchvision==0.15.2
|
||||
RUN pip install git+https://github.com/huggingface/transformers@de9255de27abfcae4a1f816b904915f0b1e23cd9
|
||||
RUN pip install transformers -U
|
||||
RUN pip install accelerate==0.21.0
|
||||
RUN pip install sentencepiece==0.1.99
|
||||
RUN pip install grpcio-status==1.33.2
|
||||
RUN pip install protobuf==3.19.6
|
||||
RUN python3 -m pip install --no-cache-dir git+https://github.com/huggingface/peft.git
|
||||
RUN pip install datasets==2.9.0
|
||||
RUN pip install triton==2.0.0.dev20221120
|
||||
RUN pip install xformers==0.0.20
|
||||
RUN pip install Jinja2==3.1.2
|
||||
RUN pip install ftfy==6.1.1
|
||||
RUN pip install cloudml-hypertune==0.1.0.dev6
|
||||
RUN pip install tensorboard==2.12.0
|
||||
RUN pip install scipy==1.10.1
|
||||
RUN pip install evaluate==0.4.0
|
||||
RUN pip install scikit-learn==1.2.2
|
||||
RUN pip install loralib==0.1.1
|
||||
RUN pip install bitsandbytes==0.39.0
|
||||
RUN pip install trl==0.4.4
|
||||
RUN pip install einops==0.6.1
|
||||
RUN pip install google-cloud-storage==2.7.0
|
||||
|
||||
RUN git clone --depth 1 --branch v0.16.1 https://github.com/huggingface/diffusers.git
|
||||
WORKDIR diffusers
|
||||
RUN pip install -e .
|
||||
|
||||
# Switch to diffusers examples folder.
|
||||
WORKDIR examples
|
||||
|
||||
# NOTE: use 'sed' to modify train_text_to_image_lora.py to
|
||||
# fix the bug for accelerator.
|
||||
RUN sed -i \
|
||||
"s#logging_dir=logging_dir#project_dir=logging_dir#g" \
|
||||
text_to_image/train_text_to_image_lora.py
|
||||
|
||||
# Config accelerate.
|
||||
RUN mkdir -p ./vertex_vision_model_garden_peft/
|
||||
COPY model_oss/peft/train.sh ./vertex_vision_model_garden_peft/train.sh
|
||||
COPY model_oss/peft/*.py ./vertex_vision_model_garden_peft/
|
||||
COPY model_oss/util /diffusers/examples/util
|
||||
ENV PYTHONPATH /diffusers/examples/
|
||||
|
||||
# Generate accelerate config at the beginning of docker run.
|
||||
ENTRYPOINT ["python3", "vertex_vision_model_garden_peft/main.py"]
|
||||
@@ -0,0 +1,250 @@
|
||||
"""Custom handler for huggingface/peft models."""
|
||||
|
||||
# pylint: disable=g-importing-member
|
||||
# pylint: disable=logging-fstring-interpolation
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Any, List
|
||||
|
||||
from absl import logging
|
||||
from diffusers import DPMSolverMultistepScheduler
|
||||
from diffusers import StableDiffusionPipeline
|
||||
from peft import PeftModel
|
||||
from PIL import Image
|
||||
import torch
|
||||
import transformers
|
||||
from transformers import AutoModelForCausalLM
|
||||
from transformers import AutoModelForSequenceClassification
|
||||
from transformers import AutoTokenizer
|
||||
from transformers import BitsAndBytesConfig
|
||||
from ts.torch_handler.base_handler import BaseHandler
|
||||
|
||||
from util import constants
|
||||
from util import fileutils
|
||||
from util import image_format_converter
|
||||
|
||||
# Tasks
|
||||
TEXT_TO_IMAGE_LORA = "text-to-image-lora"
|
||||
SEQUENCE_CLASSIFICATION_LORA = "sequence-classification-lora"
|
||||
CAUSAL_LANGUAGE_MODELING_LORA = "causal-language-modeling-lora"
|
||||
INSTRUCT_LORA = "instruct-lora"
|
||||
|
||||
# Inference parameters.
|
||||
_NUM_INFERENCE_STEPS = 25
|
||||
_MAX_LENGTH_DEFAULT = 200
|
||||
_TOP_K_DEFAULT = 10
|
||||
|
||||
|
||||
class PeftHandler(BaseHandler):
|
||||
"""Custom handler for Peft models."""
|
||||
|
||||
def initialize(self, context: Any):
|
||||
"""Initializes the handler."""
|
||||
logging.info("Start to initialize the PEFT handler.")
|
||||
properties = context.system_properties
|
||||
self.map_location = (
|
||||
"cuda"
|
||||
if torch.cuda.is_available() and properties.get("gpu_id") is not None
|
||||
else "cpu"
|
||||
)
|
||||
|
||||
self.device = torch.device(
|
||||
self.map_location + ":" + str(properties.get("gpu_id"))
|
||||
if torch.cuda.is_available() and properties.get("gpu_id") is not None
|
||||
else self.map_location
|
||||
)
|
||||
self.manifest = context.manifest
|
||||
self.precision_mode = os.environ.get(
|
||||
"PRECISION_LOADING_MODE", constants.PRECISION_MODE_16
|
||||
)
|
||||
self.task = os.environ.get("TASK", CAUSAL_LANGUAGE_MODELING_LORA)
|
||||
self.base_model_id = os.environ.get(
|
||||
"BASE_MODEL_ID", "openlm-research/open_llama_7b"
|
||||
)
|
||||
if fileutils.is_gcs_path(self.base_model_id):
|
||||
fileutils.download_gcs_dir_to_local(
|
||||
self.base_model_id,
|
||||
constants.LOCAL_BASE_MODEL_DIR,
|
||||
skip_hf_model_bin=True,
|
||||
)
|
||||
self.base_model_id = constants.LOCAL_BASE_MODEL_DIR
|
||||
self.finetuned_lora_model_path = os.environ.get(
|
||||
"FINETUNED_LORA_MODEL_PATH", ""
|
||||
)
|
||||
if fileutils.is_gcs_path(self.finetuned_lora_model_path):
|
||||
fileutils.download_gcs_dir_to_local(
|
||||
self.finetuned_lora_model_path, constants.LOCAL_MODEL_DIR
|
||||
)
|
||||
self.finetuned_lora_model_path = constants.LOCAL_MODEL_DIR
|
||||
|
||||
logging.info(
|
||||
f"Using task:{self.task}, base model:{self.base_model_id}, lora model:"
|
||||
f" {self.finetuned_lora_model_path}, and precision"
|
||||
f" {self.precision_mode}."
|
||||
)
|
||||
|
||||
self.pipeline = None
|
||||
self.model = None
|
||||
self.tokenizer = None
|
||||
if self.task == TEXT_TO_IMAGE_LORA:
|
||||
pipeline = StableDiffusionPipeline.from_pretrained(
|
||||
self.base_model_id, torch_dtype=torch.float16
|
||||
)
|
||||
logging.debug("Initialized the base model for text to image.")
|
||||
pipeline.scheduler = DPMSolverMultistepScheduler.from_config(
|
||||
pipeline.scheduler.config
|
||||
)
|
||||
logging.debug("Initialized the scheduler for text to image.")
|
||||
if self.finetuned_lora_model_path:
|
||||
pipeline.unet.load_attn_procs(self.finetuned_lora_model_path)
|
||||
logging.debug("Initialized the LoRA model for text to image.")
|
||||
# This is to reduce GPU memory requirements.
|
||||
pipeline.enable_xformers_memory_efficient_attention()
|
||||
pipeline = pipeline.to(self.map_location)
|
||||
# Reduces memory footprint.
|
||||
pipeline.enable_attention_slicing()
|
||||
self.pipeline = pipeline
|
||||
logging.info("Initialized the text to image pipelines.")
|
||||
elif self.task == SEQUENCE_CLASSIFICATION_LORA:
|
||||
tokenizer = AutoTokenizer.from_pretrained(self.base_model_id)
|
||||
logging.debug("Initialized the tokenizer for sequence classification.")
|
||||
model = AutoModelForSequenceClassification.from_pretrained(
|
||||
self.base_model_id, torch_dtype=torch.float16
|
||||
)
|
||||
logging.debug("Initialized the base model for sequence classification.")
|
||||
if self.finetuned_lora_model_path:
|
||||
model = PeftModel.from_pretrained(model, self.finetuned_lora_model_path)
|
||||
logging.debug("Initialized the LoRA model for sequence classification.")
|
||||
model.to(self.map_location)
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
elif (
|
||||
self.task == CAUSAL_LANGUAGE_MODELING_LORA or self.task == INSTRUCT_LORA
|
||||
):
|
||||
tokenizer = AutoTokenizer.from_pretrained(self.base_model_id)
|
||||
logging.debug("Initialized the tokenizer.")
|
||||
if self.task == CAUSAL_LANGUAGE_MODELING_LORA:
|
||||
if self.precision_mode == constants.PRECISION_MODE_32:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
self.base_model_id,
|
||||
return_dict=True,
|
||||
torch_dtype=torch.float32,
|
||||
device_map="auto",
|
||||
)
|
||||
elif self.precision_mode == constants.PRECISION_MODE_16:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
self.base_model_id,
|
||||
return_dict=True,
|
||||
torch_dtype=torch.bfloat16,
|
||||
device_map="auto",
|
||||
)
|
||||
elif self.precision_mode == constants.PRECISION_MODE_8:
|
||||
quantization_config = BitsAndBytesConfig(
|
||||
load_in_8bit=True, int8_threshold=0
|
||||
)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
self.base_model_id,
|
||||
return_dict=True,
|
||||
torch_dtype=torch.float16,
|
||||
device_map="auto",
|
||||
quantization_config=quantization_config,
|
||||
)
|
||||
else:
|
||||
quantization_config = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_quant_type="nf4",
|
||||
bnb_4bit_compute_dtype=torch.bfloat16,
|
||||
)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
self.base_model_id,
|
||||
return_dict=True,
|
||||
device_map="auto",
|
||||
torch_dtype=torch.bfloat16,
|
||||
quantization_config=quantization_config,
|
||||
)
|
||||
else:
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
self.base_model_id,
|
||||
return_dict=True,
|
||||
torch_dtype=torch.bfloat16,
|
||||
trust_remote_code=True,
|
||||
device_map="auto",
|
||||
)
|
||||
logging.debug("Initialized the base model.")
|
||||
if self.finetuned_lora_model_path:
|
||||
model = PeftModel.from_pretrained(model, self.finetuned_lora_model_path)
|
||||
logging.debug("Initialized the LoRA model.")
|
||||
pipeline = transformers.pipeline(
|
||||
"text-generation",
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
self.tokenizer = tokenizer
|
||||
self.pipeline = pipeline
|
||||
else:
|
||||
raise ValueError(f"Invalid TASK: {self.task}")
|
||||
|
||||
self.initialized = True
|
||||
logging.info("The PEFT handler was initialized.")
|
||||
|
||||
def preprocess(self, data: Any) -> Any:
|
||||
"""Preprocesses input data."""
|
||||
# Assumes that the parameters are same in one request. We parse the
|
||||
# parameters from the first instance for all instances in one request.
|
||||
max_length = _MAX_LENGTH_DEFAULT
|
||||
top_k = _TOP_K_DEFAULT
|
||||
|
||||
prompts = [item["prompt"] for item in data]
|
||||
if "max_length" in data[0]:
|
||||
max_length = data[0]["max_length"]
|
||||
if "top_k" in data[0]:
|
||||
top_k = data[0]["top_k"]
|
||||
|
||||
return prompts, max_length, top_k
|
||||
|
||||
def inference(self, data: Any, *args, **kwargs) -> List[Image.Image]:
|
||||
"""Runs the inference."""
|
||||
prompts, max_length, top_k = data
|
||||
logging.debug(
|
||||
f"Inference prompts={prompts}, max_length={max_length}, top_k={top_k}."
|
||||
)
|
||||
if self.task == TEXT_TO_IMAGE_LORA:
|
||||
predicted_results = self.pipeline(
|
||||
prompt=prompts, num_inference_steps=_NUM_INFERENCE_STEPS
|
||||
).images
|
||||
elif self.task == SEQUENCE_CLASSIFICATION_LORA:
|
||||
encoded_input = self.tokenizer(prompts, return_tensors="pt")
|
||||
encoded_input.to(self.map_location)
|
||||
with torch.no_grad():
|
||||
outputs = self.model(**encoded_input)
|
||||
predictions = outputs.logits.argmax(dim=-1)
|
||||
predicted_results = predictions.tolist()
|
||||
elif (
|
||||
self.task == CAUSAL_LANGUAGE_MODELING_LORA or self.task == INSTRUCT_LORA
|
||||
):
|
||||
predicted_results = self.pipeline(
|
||||
prompts,
|
||||
max_length=max_length,
|
||||
do_sample=True,
|
||||
top_k=top_k,
|
||||
num_return_sequences=1,
|
||||
eos_token_id=self.tokenizer.eos_token_id,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Invalid TASK: {self.task}")
|
||||
return predicted_results
|
||||
|
||||
def postprocess(self, data: Any) -> List[str]:
|
||||
"""Postprocesses output data."""
|
||||
if self.task == TEXT_TO_IMAGE_LORA:
|
||||
# Converts the images to base64 string.
|
||||
outputs = [
|
||||
image_format_converter.image_to_base64(image) for image in data
|
||||
]
|
||||
else:
|
||||
outputs = data
|
||||
return outputs
|
||||
|
||||
|
||||
# pylint: enable=logging-fstring-interpolation
|
||||
+131
@@ -0,0 +1,131 @@
|
||||
diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py
|
||||
index 45459ed..32527f4 100644
|
||||
--- a/src/transformers/modeling_utils.py
|
||||
+++ b/src/transformers/modeling_utils.py
|
||||
@@ -32,6 +32,8 @@ import torch
|
||||
from packaging import version
|
||||
from torch import Tensor, nn
|
||||
from torch.nn import CrossEntropyLoss
|
||||
+from huggingface_hub import hf_hub_download
|
||||
+from google.cloud import storage
|
||||
|
||||
from .activations import get_activation
|
||||
from .configuration_utils import PretrainedConfig
|
||||
@@ -442,6 +444,29 @@ def load_state_dict(checkpoint_file: Union[str, os.PathLike]):
|
||||
"""
|
||||
Reads a PyTorch checkpoint file, returning properly formatted errors if they arise.
|
||||
"""
|
||||
+ delete_download = False
|
||||
+ tmp_dir = "/tmp/model"
|
||||
+ os.makedirs(tmp_dir, exist_ok=True)
|
||||
+ if isinstance(checkpoint_file, dict):
|
||||
+ # Download model file from huggingface
|
||||
+ print(f"==> Download model from HF: {checkpoint_file}")
|
||||
+ checkpoint_file = hf_hub_download(
|
||||
+ local_dir=tmp_dir, local_dir_use_symlinks=False, force_download=True, resume_download=True, **checkpoint_file)
|
||||
+ delete_download = True
|
||||
+ else:
|
||||
+ with open(checkpoint_file, "rb") as f:
|
||||
+ is_gcs_file = (f.read(2) == b"gs")
|
||||
+ if is_gcs_file:
|
||||
+ # Download model file from GCS
|
||||
+ with open(checkpoint_file, "r") as f:
|
||||
+ gcs_file = f.read()
|
||||
+ checkpoint_file = os.path.join(tmp_dir, gcs_file.split("/")[-1])
|
||||
+ print(f"==> Download model from GCS: {gcs_file} to: {checkpoint_file}")
|
||||
+ client = storage.Client()
|
||||
+ with open(checkpoint_file, 'wb') as f:
|
||||
+ client.download_blob_to_file(gcs_file, f)
|
||||
+ delete_download = True
|
||||
+
|
||||
if checkpoint_file.endswith(".safetensors") and is_safetensors_available():
|
||||
# Check format of the archive
|
||||
with safe_open(checkpoint_file, framework="pt") as f:
|
||||
@@ -455,9 +480,9 @@ def load_state_dict(checkpoint_file: Union[str, os.PathLike]):
|
||||
raise NotImplementedError(
|
||||
f"Conversion from a {metadata['format']} safetensors archive to PyTorch is not implemented yet."
|
||||
)
|
||||
- return safe_load_file(checkpoint_file)
|
||||
+ state_dict = safe_load_file(checkpoint_file)
|
||||
try:
|
||||
- return torch.load(checkpoint_file, map_location="cpu")
|
||||
+ state_dict = torch.load(checkpoint_file, map_location="cpu")
|
||||
except Exception as e:
|
||||
try:
|
||||
with open(checkpoint_file) as f:
|
||||
@@ -478,6 +503,10 @@ def load_state_dict(checkpoint_file: Union[str, os.PathLike]):
|
||||
f"at '{checkpoint_file}'. "
|
||||
"If you tried to load a PyTorch model from a TF 2.0 checkpoint, please set from_tf=True."
|
||||
)
|
||||
+ if delete_download:
|
||||
+ print(f"==> Delete downloaded model: {checkpoint_file}")
|
||||
+ os.remove(checkpoint_file)
|
||||
+ return state_dict
|
||||
|
||||
|
||||
def set_initialized_submodules(model, state_dict_keys):
|
||||
@@ -3179,7 +3208,10 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin, GenerationMixin, PushToHubMix
|
||||
return mismatched_keys
|
||||
|
||||
if resolved_archive_file is not None:
|
||||
- folder = os.path.sep.join(resolved_archive_file[0].split(os.path.sep)[:-1])
|
||||
+ if isinstance(resolved_archive_file, str):
|
||||
+ folder = os.path.sep.join(resolved_archive_file[0].split(os.path.sep)[:-1])
|
||||
+ else:
|
||||
+ folder = None
|
||||
else:
|
||||
folder = None
|
||||
if device_map is not None and is_safetensors:
|
||||
diff --git a/src/transformers/utils/hub.py b/src/transformers/utils/hub.py
|
||||
index ffed743..4b15770 100644
|
||||
--- a/src/transformers/utils/hub.py
|
||||
+++ b/src/transformers/utils/hub.py
|
||||
@@ -414,20 +414,34 @@ def cached_file(
|
||||
user_agent = http_user_agent(user_agent)
|
||||
try:
|
||||
# Load from URL or cache if already cached
|
||||
- resolved_file = hf_hub_download(
|
||||
- path_or_repo_id,
|
||||
- filename,
|
||||
- subfolder=None if len(subfolder) == 0 else subfolder,
|
||||
- repo_type=repo_type,
|
||||
- revision=revision,
|
||||
- cache_dir=cache_dir,
|
||||
- user_agent=user_agent,
|
||||
- force_download=force_download,
|
||||
- proxies=proxies,
|
||||
- resume_download=resume_download,
|
||||
- use_auth_token=use_auth_token,
|
||||
- local_files_only=local_files_only,
|
||||
- )
|
||||
+ if filename.endswith(".bin"):
|
||||
+ # NOTE: To save disk we do not download bin file eagerly. Do not support safetensors.
|
||||
+ resolved_file = dict(
|
||||
+ repo_id=path_or_repo_id,
|
||||
+ filename=filename,
|
||||
+ subfolder=None if len(subfolder) == 0 else subfolder,
|
||||
+ repo_type=repo_type,
|
||||
+ revision=revision,
|
||||
+ user_agent=user_agent,
|
||||
+ proxies=proxies,
|
||||
+ use_auth_token=use_auth_token,
|
||||
+ )
|
||||
+ print(f"--> Apply lazy download to bin file: {resolved_file}")
|
||||
+ else:
|
||||
+ resolved_file = hf_hub_download(
|
||||
+ path_or_repo_id,
|
||||
+ filename,
|
||||
+ subfolder=None if len(subfolder) == 0 else subfolder,
|
||||
+ repo_type=repo_type,
|
||||
+ revision=revision,
|
||||
+ cache_dir=cache_dir,
|
||||
+ user_agent=user_agent,
|
||||
+ force_download=force_download,
|
||||
+ proxies=proxies,
|
||||
+ resume_download=resume_download,
|
||||
+ use_auth_token=use_auth_token,
|
||||
+ local_files_only=local_files_only,
|
||||
+ )
|
||||
|
||||
except RepositoryNotFoundError:
|
||||
raise EnvironmentError(
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Instruct/Chat with LoRA models."""
|
||||
|
||||
# pylint: disable=g-importing-member
|
||||
from datasets import load_dataset
|
||||
from peft import LoraConfig
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM
|
||||
from transformers import AutoTokenizer
|
||||
from transformers import BitsAndBytesConfig
|
||||
from transformers import TrainingArguments
|
||||
from trl import SFTTrainer
|
||||
|
||||
|
||||
def finetune_instruct(
|
||||
pretrained_model_id: str,
|
||||
dataset_name: str,
|
||||
output_dir: str,
|
||||
lora_rank: int = 64,
|
||||
lora_alpha: int = 16,
|
||||
lora_dropout: float = 0.1,
|
||||
warmup_ratio: int = 0.03,
|
||||
max_steps: int = 10,
|
||||
max_seq_length: int = 512,
|
||||
learning_rate: float = 2e-4,
|
||||
) -> None:
|
||||
"""Finetunes instruct."""
|
||||
dataset = load_dataset(dataset_name, split="train")
|
||||
|
||||
bnb_config = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_quant_type="nf4",
|
||||
bnb_4bit_compute_dtype=torch.float16,
|
||||
)
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
pretrained_model_id,
|
||||
quantization_config=bnb_config,
|
||||
trust_remote_code=True,
|
||||
)
|
||||
model.config.use_cache = False
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
pretrained_model_id, trust_remote_code=True
|
||||
)
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
peft_config = LoraConfig(
|
||||
lora_alpha=lora_alpha,
|
||||
lora_dropout=lora_dropout,
|
||||
r=lora_rank,
|
||||
bias="none",
|
||||
task_type="CAUSAL_LM",
|
||||
target_modules=[
|
||||
"query_key_value",
|
||||
"dense",
|
||||
"dense_h_to_4h",
|
||||
"dense_4h_to_h",
|
||||
],
|
||||
)
|
||||
|
||||
per_device_train_batch_size = 4
|
||||
gradient_accumulation_steps = 4
|
||||
optim = "paged_adamw_32bit"
|
||||
save_steps = 10
|
||||
logging_steps = 10
|
||||
max_grad_norm = 0.3
|
||||
lr_scheduler_type = "constant"
|
||||
|
||||
training_arguments = TrainingArguments(
|
||||
output_dir=output_dir,
|
||||
per_device_train_batch_size=per_device_train_batch_size,
|
||||
gradient_accumulation_steps=gradient_accumulation_steps,
|
||||
optim=optim,
|
||||
save_steps=save_steps,
|
||||
logging_steps=logging_steps,
|
||||
learning_rate=learning_rate,
|
||||
fp16=True,
|
||||
max_grad_norm=max_grad_norm,
|
||||
max_steps=max_steps,
|
||||
warmup_ratio=warmup_ratio,
|
||||
group_by_length=True,
|
||||
lr_scheduler_type=lr_scheduler_type,
|
||||
)
|
||||
|
||||
trainer = SFTTrainer(
|
||||
model=model,
|
||||
train_dataset=dataset,
|
||||
peft_config=peft_config,
|
||||
dataset_text_field="text",
|
||||
max_seq_length=max_seq_length,
|
||||
tokenizer=tokenizer,
|
||||
args=training_arguments,
|
||||
)
|
||||
for name, module in trainer.model.named_modules():
|
||||
if "norm" in name:
|
||||
module = module.to(torch.float32)
|
||||
trainer.train()
|
||||
@@ -0,0 +1,177 @@
|
||||
"""Main function to start PEFT finetuning."""
|
||||
import subprocess
|
||||
|
||||
from absl import app
|
||||
from absl import flags
|
||||
from absl import logging
|
||||
|
||||
from peft import causal_language_modeling_lora
|
||||
from peft import instruct_lora
|
||||
from peft import sequence_classification_lora
|
||||
from util import constants
|
||||
from util import fileutils
|
||||
|
||||
_TASK = flags.DEFINE_string(
|
||||
'task',
|
||||
constants.CAUSAL_LANGUAGE_MODELING_LORA,
|
||||
'The supported PEFT tasks.',
|
||||
)
|
||||
|
||||
_PRETRAINED_MODEL_ID = flags.DEFINE_string(
|
||||
'pretrained_model_id',
|
||||
None,
|
||||
'The pretrained model id. Supported models can be causal language modeling'
|
||||
' models from https://github.com/huggingface/peft/tree/main.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_DATASET_NAME = flags.DEFINE_string(
|
||||
'dataset_name',
|
||||
None,
|
||||
'The dataset name in huggingface.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_OUTPUT_DIR = flags.DEFINE_string(
|
||||
'output_dir',
|
||||
None,
|
||||
'The output directory.',
|
||||
required=True,
|
||||
)
|
||||
|
||||
_PRECISION_MODE = flags.DEFINE_string(
|
||||
'precision_mode',
|
||||
constants.PRECISION_MODE_16,
|
||||
'Supported finetuning precision_modes are `{}` and `{}`.'.format(
|
||||
constants.PRECISION_MODE_8, constants.PRECISION_MODE_16
|
||||
),
|
||||
)
|
||||
|
||||
_LORA_RANK = flags.DEFINE_integer(
|
||||
'lora_rank',
|
||||
16,
|
||||
'The rank of the update matrices, expressed in int. Lower rank results in'
|
||||
' smaller update matrices with fewer trainable parameters, referring to'
|
||||
' https://huggingface.co/docs/peft/conceptual_guides/lora.',
|
||||
)
|
||||
|
||||
_LORA_ALPHA = flags.DEFINE_integer(
|
||||
'lora_alpha',
|
||||
32,
|
||||
'LoRA scaling factor, referring to'
|
||||
' https://huggingface.co/docs/peft/conceptual_guides/lora.',
|
||||
)
|
||||
|
||||
_LORA_DROPOUT = flags.DEFINE_float(
|
||||
'lora_dropout',
|
||||
0.05,
|
||||
'dropout probability of the LoRA layers, referring to'
|
||||
' https://huggingface.co/docs/peft/task_guides/token-classification-lora.',
|
||||
)
|
||||
|
||||
_WARMUP_STEPS = flags.DEFINE_integer(
|
||||
'warmup_steps',
|
||||
10,
|
||||
'Number of steps for the warmup in the learning rate scheduler.',
|
||||
)
|
||||
|
||||
_WARMUP_RATIO = flags.DEFINE_float(
|
||||
'warmup_ratio',
|
||||
0.03,
|
||||
'The warmup ratio in the learning rate scheduler.',
|
||||
)
|
||||
|
||||
_MAX_STEPS = flags.DEFINE_integer(
|
||||
'max_steps',
|
||||
10,
|
||||
'Total number of training steps.',
|
||||
)
|
||||
|
||||
_MAX_SEQ_LENGTH = flags.DEFINE_integer(
|
||||
'max_seq_length',
|
||||
512,
|
||||
'The maximum sequence length.',
|
||||
)
|
||||
|
||||
_NUM_EPOCHS = flags.DEFINE_integer(
|
||||
'num_epochs',
|
||||
20,
|
||||
'The number of training epochs.',
|
||||
)
|
||||
|
||||
_BATCH_SIZE = flags.DEFINE_integer(
|
||||
'batch_size',
|
||||
32,
|
||||
'The batch size.',
|
||||
)
|
||||
|
||||
_LEARNING_RATE = flags.DEFINE_float(
|
||||
'learning_rate',
|
||||
2e-4,
|
||||
'The learning rate after the potential warmup period.',
|
||||
)
|
||||
|
||||
|
||||
def main(_) -> None:
|
||||
task = _TASK.value
|
||||
pretrained_model_id = _PRETRAINED_MODEL_ID.value
|
||||
local_pretrained_model_id = None
|
||||
if pretrained_model_id.startswith(constants.GCS_URI_PREFIX):
|
||||
logging.info(
|
||||
'Start to copy pretrained models locally: %s.', pretrained_model_id
|
||||
)
|
||||
fileutils.download_gcs_dir_to_local(
|
||||
pretrained_model_id, constants.LOCAL_BASE_MODEL_DIR
|
||||
)
|
||||
local_pretrained_model_id = constants.LOCAL_BASE_MODEL_DIR
|
||||
logging.info(
|
||||
'Finished copying pretrained models locally to: %s.',
|
||||
local_pretrained_model_id,
|
||||
)
|
||||
if task == constants.TEXT_TO_IMAGE_LORA:
|
||||
subprocess.run(['/bin/bash', 'train.sh'], check=True)
|
||||
elif task == constants.SEQUENCE_CLASSIFICATION_LORA:
|
||||
sequence_classification_lora.finetune_sequence_classification(
|
||||
pretrained_model_id=pretrained_model_id,
|
||||
dataset_name=_DATASET_NAME.value,
|
||||
output_dir=_OUTPUT_DIR.value,
|
||||
lora_rank=_LORA_RANK.value,
|
||||
lora_alpha=_LORA_ALPHA.value,
|
||||
lora_dropout=_LORA_DROPOUT.value,
|
||||
num_epochs=_NUM_EPOCHS.value,
|
||||
batch_size=_BATCH_SIZE.value,
|
||||
learning_rate=_LEARNING_RATE.value,
|
||||
)
|
||||
elif task == constants.CAUSAL_LANGUAGE_MODELING_LORA:
|
||||
causal_language_modeling_lora.finetune_causal_language_modeling(
|
||||
pretrained_model_id=pretrained_model_id,
|
||||
dataset_name=_DATASET_NAME.value,
|
||||
output_dir=_OUTPUT_DIR.value,
|
||||
precision_mode=_PRECISION_MODE.value,
|
||||
lora_rank=_LORA_RANK.value,
|
||||
lora_alpha=_LORA_ALPHA.value,
|
||||
lora_dropout=_LORA_DROPOUT.value,
|
||||
warmup_steps=_WARMUP_STEPS.value,
|
||||
max_steps=_MAX_STEPS.value,
|
||||
learning_rate=_LEARNING_RATE.value,
|
||||
local_pretrained_model_id=local_pretrained_model_id,
|
||||
)
|
||||
elif task == constants.INSTRUCT_LORA:
|
||||
instruct_lora.finetune_instruct(
|
||||
pretrained_model_id=pretrained_model_id,
|
||||
dataset_name=_DATASET_NAME.value,
|
||||
output_dir=_OUTPUT_DIR.value,
|
||||
lora_rank=_LORA_RANK.value,
|
||||
lora_alpha=_LORA_ALPHA.value,
|
||||
lora_dropout=_LORA_DROPOUT.value,
|
||||
warmup_ratio=_WARMUP_RATIO.value,
|
||||
max_steps=_MAX_STEPS.value,
|
||||
max_seq_length=_MAX_SEQ_LENGTH.value,
|
||||
learning_rate=_LEARNING_RATE.value,
|
||||
)
|
||||
else:
|
||||
raise ValueError('The task {} is not supported.'.format(task))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(main)
|
||||
+133
@@ -0,0 +1,133 @@
|
||||
"""Sequence classification with LoRA models."""
|
||||
|
||||
# pylint: disable=g-importing-member
|
||||
|
||||
from datasets import load_dataset
|
||||
import evaluate
|
||||
from peft import get_peft_model
|
||||
from peft import LoraConfig
|
||||
import torch
|
||||
from torch.optim import AdamW
|
||||
from torch.utils.data import DataLoader
|
||||
from tqdm import tqdm
|
||||
from transformers import AutoModelForSequenceClassification
|
||||
from transformers import AutoTokenizer
|
||||
from transformers import get_linear_schedule_with_warmup
|
||||
|
||||
|
||||
def finetune_sequence_classification(
|
||||
pretrained_model_id: str,
|
||||
dataset_name: str,
|
||||
output_dir: str,
|
||||
lora_rank: int = 8,
|
||||
lora_alpha: int = 16,
|
||||
lora_dropout: float = 0.1,
|
||||
num_epochs: int = 20,
|
||||
batch_size: int = 32,
|
||||
learning_rate: float = 3e-4,
|
||||
) -> None:
|
||||
"""Finetunes sequence classification."""
|
||||
task = "mrpc"
|
||||
device = "cuda"
|
||||
|
||||
peft_config = LoraConfig(
|
||||
task_type="SEQ_CLS",
|
||||
inference_mode=False,
|
||||
r=lora_rank,
|
||||
lora_alpha=lora_alpha,
|
||||
lora_dropout=lora_dropout,
|
||||
)
|
||||
if any(k in pretrained_model_id for k in ("gpt", "opt", "bloom")):
|
||||
padding_side = "left"
|
||||
else:
|
||||
padding_side = "right"
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
pretrained_model_id, padding_side=padding_side
|
||||
)
|
||||
if getattr(tokenizer, "pad_token_id") is None:
|
||||
tokenizer.pad_token_id = tokenizer.eos_token_id
|
||||
|
||||
datasets = load_dataset(dataset_name, task)
|
||||
metric = evaluate.load(dataset_name, task)
|
||||
|
||||
def tokenize_function(examples):
|
||||
# max_length=None => use the model max length (it's actually the default)
|
||||
outputs = tokenizer(
|
||||
examples["sentence1"],
|
||||
examples["sentence2"],
|
||||
truncation=True,
|
||||
max_length=None,
|
||||
)
|
||||
return outputs
|
||||
|
||||
tokenized_datasets = datasets.map(
|
||||
tokenize_function,
|
||||
batched=True,
|
||||
remove_columns=["idx", "sentence1", "sentence2"],
|
||||
)
|
||||
|
||||
# We also rename the 'label' column to 'labels' which is the expected name for
|
||||
# labels by the models of the transformers library.
|
||||
tokenized_datasets = tokenized_datasets.rename_column("label", "labels")
|
||||
|
||||
def collate_fn(examples):
|
||||
return tokenizer.pad(examples, padding="longest", return_tensors="pt")
|
||||
|
||||
# Instantiate dataloaders.
|
||||
train_dataloader = DataLoader(
|
||||
tokenized_datasets["train"],
|
||||
shuffle=True,
|
||||
collate_fn=collate_fn,
|
||||
batch_size=batch_size,
|
||||
)
|
||||
eval_dataloader = DataLoader(
|
||||
tokenized_datasets["validation"],
|
||||
shuffle=False,
|
||||
collate_fn=collate_fn,
|
||||
batch_size=batch_size,
|
||||
)
|
||||
|
||||
model = AutoModelForSequenceClassification.from_pretrained(
|
||||
pretrained_model_id, return_dict=True
|
||||
)
|
||||
model = get_peft_model(model, peft_config)
|
||||
model.print_trainable_parameters()
|
||||
|
||||
optimizer = AdamW(params=model.parameters(), lr=learning_rate)
|
||||
|
||||
# Instantiate scheduler
|
||||
lr_scheduler = get_linear_schedule_with_warmup(
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=0.06 * (len(train_dataloader) * num_epochs),
|
||||
num_training_steps=(len(train_dataloader) * num_epochs),
|
||||
)
|
||||
|
||||
model.to(device)
|
||||
for epoch in range(num_epochs):
|
||||
model.train()
|
||||
for _, batch in enumerate(tqdm(train_dataloader)):
|
||||
batch.to(device)
|
||||
outputs = model(**batch)
|
||||
loss = outputs.loss
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
model.eval()
|
||||
for _, batch in enumerate(tqdm(eval_dataloader)):
|
||||
batch.to(device)
|
||||
with torch.no_grad():
|
||||
outputs = model(**batch)
|
||||
predictions = outputs.logits.argmax(dim=-1)
|
||||
references = batch["labels"]
|
||||
metric.add_batch(
|
||||
predictions=predictions,
|
||||
references=references,
|
||||
)
|
||||
|
||||
eval_metric = metric.compute()
|
||||
print(f"epoch {epoch}:", eval_metric)
|
||||
|
||||
model.save_pretrained(output_dir)
|
||||
@@ -0,0 +1,6 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Setup accelerate config before running trainer.
|
||||
python -c "from accelerate.utils import write_basic_config; write_basic_config(mixed_precision='fp16')"
|
||||
|
||||
accelerate launch "$@"
|
||||
+115
@@ -0,0 +1,115 @@
|
||||
FROM pytorch/torchserve:0.7.1-gpu
|
||||
|
||||
USER root
|
||||
|
||||
ENV infer_port=7080
|
||||
ENV mng_port=7081
|
||||
ENV model_name="pic2word"
|
||||
ENV PATH="/home/model-server/:${PATH}"
|
||||
|
||||
# Copy license.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget
|
||||
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Install dependencies.
|
||||
ENV PIP_ROOT_USER_ACTION=ignore
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
RUN pip install google-cloud-storage==2.7.0
|
||||
RUN pip install open_clip_torch==2.20.0
|
||||
RUN pip install numpy==1.22.0
|
||||
RUN pip install scikit-image==0.21.0
|
||||
RUN pip install scikit-learn==1.0.2
|
||||
RUN pip install torch==2.0.0
|
||||
RUN pip install torchvision==0.15.2
|
||||
RUN pip install tensorboard==2.13.0
|
||||
RUN pip install ase==3.21.1
|
||||
RUN pip install braceexpand==0.1.7
|
||||
RUN pip install cached-property==1.5.2
|
||||
RUN pip install configparser==5.0.2
|
||||
RUN pip install cycler==0.10.0
|
||||
RUN pip install decorator==4.4.2
|
||||
RUN pip install docker-pycreds==0.4.0
|
||||
RUN pip install gitdb==4.0.7
|
||||
RUN pip install gitpython==3.1.30
|
||||
RUN pip install googledrivedownloader==0.4
|
||||
RUN pip install h5py==3.1.0
|
||||
RUN pip install isodate==0.6.0
|
||||
RUN pip install jinja2==3.0.1
|
||||
RUN pip install kiwisolver==1.3.1
|
||||
RUN pip install littleutils==0.2.2
|
||||
RUN pip install llvmlite==0.36.0
|
||||
RUN pip install markupsafe==2.0.1
|
||||
RUN pip install matplotlib==3.3.4
|
||||
RUN pip install networkx==2.5.1
|
||||
RUN pip install numba==0.53.1
|
||||
RUN pip install ogb==1.3.1
|
||||
RUN pip install outdated==0.2.1
|
||||
RUN pip install pathtools==0.1.2
|
||||
RUN pip install promise==2.3
|
||||
RUN pip install psutil==5.8.0
|
||||
RUN pip install pyarrow==4.0.0
|
||||
RUN pip install pyparsing==2.4.7
|
||||
RUN pip install python-louvain==0.15
|
||||
RUN pip install pyyaml==5.4.1
|
||||
RUN pip install rdflib==5.0.0
|
||||
RUN pip install sentry-sdk==1.14.0
|
||||
RUN pip install shortuuid==1.0.1
|
||||
RUN pip install sklearn==0.0
|
||||
RUN pip install smmap==4.0.0
|
||||
RUN pip install subprocess32==3.5.4
|
||||
RUN pip install torch-geometric==1.7.0
|
||||
RUN pip install wandb==0.10.30
|
||||
RUN pip install wilds==1.1.0
|
||||
RUN pip install ftfy==6.1.1
|
||||
RUN pip install regex==2023.6.3
|
||||
RUN pip install webdataset==0.2.48
|
||||
RUN pip install requests==2.31.0
|
||||
RUN pip install hydra-core==1.3.2
|
||||
RUN pip install omegaconf==2.3.0
|
||||
RUN pip install fairseq==0.10.0
|
||||
RUN pip install bitarray==2.7.6
|
||||
|
||||
# Get 'composed_image_retrieval' repository from github.
|
||||
RUN git clone https://github.com/google-research/composed_image_retrieval
|
||||
# Set workdir to composed_image_retrieval.
|
||||
WORKDIR ./composed_image_retrieval
|
||||
# Using git reset command to pin it down to a specific version.
|
||||
RUN git reset --hard 8c053297c2fae9cd17ddcded48445a4f47208dbd
|
||||
|
||||
# Fix issue introduced by installing composed_image_retrieval
|
||||
# https://github.com/huggingface/transformers/issues/8638#issuecomment-790772391
|
||||
RUN pip uninstall dataclasses -y
|
||||
|
||||
# Copy model artifacts.
|
||||
COPY model_oss/pic2word/handler.py /home/model-server/handler.py
|
||||
|
||||
# Create torchserve configuration file.
|
||||
RUN echo \
|
||||
"default_response_timeout=1800\n" \
|
||||
"service_envelope=json\n" \
|
||||
"inference_address=http://0.0.0.0:${infer_port}\n" \
|
||||
"management_address=http://0.0.0.0:${mng_port}" >> /home/model-server/config.properties
|
||||
|
||||
# Expose ports.
|
||||
EXPOSE ${infer_port}
|
||||
EXPOSE ${mng_port}
|
||||
|
||||
# Archive model artifacts and dependencies.
|
||||
# Do not set --model-file and --serialized-file because model and checkpoint
|
||||
# will be dynamically loaded in handler.py.
|
||||
RUN torch-model-archiver \
|
||||
--model-name=${model_name} \
|
||||
--version=1.0 \
|
||||
--handler=/home/model-server/handler.py \
|
||||
--runtime=python3 \
|
||||
--export-path=/home/model-server/model-store \
|
||||
--archive-format=default \
|
||||
--force
|
||||
|
||||
# Run Torchserve HTTP serve to respond to prediction requests.
|
||||
CMD ["torchserve", "--start", \
|
||||
"--ts-config", "/home/model-server/config.properties", \
|
||||
"--models", "${model_name}=${model_name}.mar", \
|
||||
"--model-store", "/home/model-server/model-store"]
|
||||
@@ -0,0 +1,167 @@
|
||||
"""Custom handler for Pic2Word."""
|
||||
|
||||
from argparse import Namespace # pylint: disable=g-importing-member
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from absl import logging
|
||||
from data import CustomFolder
|
||||
from eval_utils import visualize_results
|
||||
from model.clip import load
|
||||
from model.model import convert_weights
|
||||
from model.model import IM2TEXT
|
||||
from params import get_project_root
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from ts.torch_handler.base_handler import BaseHandler
|
||||
|
||||
from util import fileutils
|
||||
|
||||
# The COCO dataset is stored in a publicly accessible bucket.
|
||||
_COCO_STORAGE_DIR = "gs://pic2word-bucket/data/coco/"
|
||||
_COCO_LOCAL_DIR = "/home/model-server/composed_image_retrieval/data/coco/"
|
||||
_COCO_VAL2017_PATH = "coco/val2017"
|
||||
_COCO_DATASET_NAME = "coco"
|
||||
_MODEL_NAME = "ViT-L/14"
|
||||
_LOCAL_QUERY_PATH = "./query/"
|
||||
_IMAGE_OUTPUT_LOCAL_DIR = "demo_out/images"
|
||||
_OUTPUT_LOCAL_DIR = "/demo_out/"
|
||||
_DATA_DIR = "data"
|
||||
_CHECKPOINT_DIR = "checkpoint/pic2word_model.pt"
|
||||
_REQUEST_PROMPTS = "prompts"
|
||||
_REQUEST_OUTPUT_STORAGE_DIR = "output_storage_dir"
|
||||
_REQUEST_IMAGE_PATH = "image_path"
|
||||
_REQUEST_IMAGE_FILE_NAME = "image_file_name"
|
||||
_RESPONSE_MSG = "Successfully retrieved images."
|
||||
|
||||
|
||||
class ModelHandler(BaseHandler):
|
||||
"""A custom model handler implementation."""
|
||||
|
||||
def __init__(self):
|
||||
self.initialized = False
|
||||
self.gpu = 0
|
||||
self.model = None
|
||||
self.dataloader = None
|
||||
self.prompt = None
|
||||
self.output_storage_dir = None
|
||||
|
||||
def initialize(self, context: Any):
|
||||
"""Initialize."""
|
||||
logging.info("Initializing pic2word.")
|
||||
|
||||
# Download COCO dataset. The model looks for this folder specifically
|
||||
# during image retrieval to generate a response for each request.
|
||||
# This is a publicly accessible bucket.
|
||||
fileutils.download_gcs_dir_to_local(
|
||||
_COCO_STORAGE_DIR,
|
||||
_COCO_LOCAL_DIR,
|
||||
)
|
||||
|
||||
# Load the model.
|
||||
|
||||
self.initialized = True
|
||||
|
||||
torch.cuda.set_device(self.gpu)
|
||||
model, _, preprocess_val = load(_MODEL_NAME, jit=False)
|
||||
|
||||
img2text = IM2TEXT(
|
||||
embed_dim=model.embed_dim,
|
||||
output_dim=model.token_embedding.weight.shape[1],
|
||||
)
|
||||
|
||||
model.cuda(self.gpu)
|
||||
img2text.cuda(self.gpu)
|
||||
|
||||
convert_weights(model)
|
||||
convert_weights(img2text)
|
||||
|
||||
self.model = model
|
||||
self.img2text = img2text
|
||||
|
||||
# Load the dataset
|
||||
logging.info("Loading dataset.")
|
||||
|
||||
root_project = os.path.join(get_project_root(), _DATA_DIR)
|
||||
dataset = CustomFolder(
|
||||
os.path.join(root_project, _COCO_VAL2017_PATH), transform=preprocess_val
|
||||
)
|
||||
|
||||
# Initialize the dataloader. This is used to create the pickle file from
|
||||
# the dataset.
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=64,
|
||||
shuffle=False,
|
||||
num_workers=1,
|
||||
pin_memory=True,
|
||||
drop_last=False,
|
||||
)
|
||||
|
||||
self.dataloader = dataloader
|
||||
|
||||
logging.info("Finished initializing Pic2Word server.")
|
||||
|
||||
def preprocess(self, data: Any) -> str:
|
||||
"""Preprocess input data."""
|
||||
logging.info("Preprocessing Pic2Word inference request.")
|
||||
query = data[0]
|
||||
|
||||
self.output_storage_dir = query[_REQUEST_OUTPUT_STORAGE_DIR]
|
||||
prompts = query[_REQUEST_PROMPTS]
|
||||
prompts = prompts.split(",")
|
||||
self.prompt = prompts
|
||||
|
||||
image_path = query[_REQUEST_IMAGE_PATH]
|
||||
# The query image is only supported via GCS bucket upload.
|
||||
fileutils.download_gcs_dir_to_local(image_path, _LOCAL_QUERY_PATH)
|
||||
image_file_name = query[_REQUEST_IMAGE_FILE_NAME]
|
||||
|
||||
query_file = f"./query/{image_file_name}"
|
||||
|
||||
logging.info("Setting model args.")
|
||||
|
||||
args = {
|
||||
"openai-pretrained": True,
|
||||
"resume": _CHECKPOINT_DIR,
|
||||
"retrieval_data": _COCO_DATASET_NAME,
|
||||
"query_file": query_file,
|
||||
"demo_out": _OUTPUT_LOCAL_DIR,
|
||||
"prompts": prompts,
|
||||
"distributed": False,
|
||||
"dp": False,
|
||||
"gpu": 0,
|
||||
"model": _MODEL_NAME,
|
||||
"world_size": 1,
|
||||
}
|
||||
model_input = Namespace(**args)
|
||||
|
||||
logging.info("Finished preprocessing Pic2Word inference request.")
|
||||
return model_input
|
||||
|
||||
def inference(self, model_input: Any):
|
||||
"""Runs inference."""
|
||||
logging.info("Running model-inference.")
|
||||
visualize_results(
|
||||
model=self.model,
|
||||
img2text=self.img2text,
|
||||
args=model_input,
|
||||
prompt=self.prompt,
|
||||
dataloader=self.dataloader,
|
||||
)
|
||||
|
||||
def postprocess(self):
|
||||
"""Upload the output images to the bucket."""
|
||||
logging.info("Running request postprocess.")
|
||||
fileutils.upload_local_dir_to_gcs(
|
||||
_IMAGE_OUTPUT_LOCAL_DIR, self.output_storage_dir
|
||||
)
|
||||
|
||||
def handle(self, data: Any, context: Any) -> str: # pylint: disable=unused-argument
|
||||
"""Runs preprocess, inference, and post-processing."""
|
||||
logging.info("Received Pic2Word inference request")
|
||||
model_input = self.preprocess(data)
|
||||
self.inference(model_input)
|
||||
self.postprocess()
|
||||
logging.info("Done handling input.")
|
||||
return _RESPONSE_MSG
|
||||
@@ -0,0 +1,4 @@
|
||||
"""AutoML Vision Tfvision configs package definition."""
|
||||
|
||||
from tfvision.configs import backbones
|
||||
from tfvision.configs import hub_model
|
||||
@@ -0,0 +1,28 @@
|
||||
"""Backbones configurations."""
|
||||
import dataclasses
|
||||
from typing import Optional
|
||||
|
||||
from official.modeling import hyperparams
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class HubModel(hyperparams.Config):
|
||||
"""Tf-hub model config."""
|
||||
handle: Optional[str] = None
|
||||
trainable: bool = True
|
||||
mean_rgb: Optional[float] = None
|
||||
stddev_rgb: Optional[float] = None
|
||||
signature: Optional[str] = None
|
||||
output_key: Optional[str] = None
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class Backbone(hyperparams.OneOfConfig):
|
||||
"""Configuration for backbones.
|
||||
|
||||
Attributes:
|
||||
type: The type of a backbone, such as 'hub_model'.
|
||||
hub_model: hub model backbone config.
|
||||
"""
|
||||
type: Optional[str] = 'hub_model'
|
||||
hub_model: HubModel = dataclasses.field(default_factory=HubModel)
|
||||
@@ -0,0 +1,166 @@
|
||||
"""Tf-hub model configuration definition for AutoML Vision ICN.."""
|
||||
|
||||
import os
|
||||
|
||||
from tfvision.configs import backbones
|
||||
from official.core import config_definitions as cfg
|
||||
from official.core import exp_factory
|
||||
from official.modeling import optimization
|
||||
from official.vision.configs import image_classification
|
||||
|
||||
_HANDLE = 'https://tfhub.dev/google/imagenet/efficientnet_v2_imagenet21k_m/feature_vector/2' # pylint: disable=line-too-long
|
||||
_COCA_HANDLE = None
|
||||
_INPUT_SIZE = [480, 480, 3]
|
||||
_MEAN_RGB = 0.0
|
||||
_STDDEV_RGB = 255.0
|
||||
|
||||
|
||||
# pylint is unable to handle dataclasses constructor arguments correctly.
|
||||
# pylint: disable=unexpected-keyword-arg
|
||||
@exp_factory.register_config_factory('hub_model')
|
||||
def hub_model() -> cfg.ExperimentConfig:
|
||||
"""Gets experimental configs for tf-hub models."""
|
||||
|
||||
batch_size = 8
|
||||
train_steps = 625000
|
||||
steps_per_loop = 1250
|
||||
return cfg.ExperimentConfig(
|
||||
task=image_classification.ImageClassificationTask(
|
||||
model=image_classification.ImageClassificationModel(
|
||||
num_classes=1000,
|
||||
input_size=_INPUT_SIZE,
|
||||
backbone=backbones.Backbone(
|
||||
type='hub_model',
|
||||
hub_model=backbones.HubModel(
|
||||
handle=_HANDLE, mean_rgb=_MEAN_RGB, stddev_rgb=_STDDEV_RGB
|
||||
),
|
||||
),
|
||||
dropout_rate=0.0,
|
||||
),
|
||||
losses=image_classification.Losses(
|
||||
l2_weight_decay=0.0, label_smoothing=0.1, one_hot=True
|
||||
),
|
||||
train_data=image_classification.DataConfig(
|
||||
input_path=os.path.join(
|
||||
image_classification.IMAGENET_INPUT_PATH_BASE, 'train*'
|
||||
),
|
||||
aug_type=None,
|
||||
dtype='float32',
|
||||
global_batch_size=batch_size,
|
||||
is_training=True,
|
||||
decode_jpeg_only=False,
|
||||
),
|
||||
validation_data=image_classification.DataConfig(
|
||||
input_path=os.path.join(
|
||||
image_classification.IMAGENET_INPUT_PATH_BASE, 'valid*'
|
||||
),
|
||||
dtype='float32',
|
||||
global_batch_size=batch_size,
|
||||
is_training=False,
|
||||
decode_jpeg_only=False,
|
||||
drop_remainder=False,
|
||||
),
|
||||
),
|
||||
trainer=cfg.TrainerConfig(
|
||||
best_checkpoint_eval_metric='accuracy',
|
||||
best_checkpoint_export_subdir='best_ckpt',
|
||||
best_checkpoint_metric_comp='higher',
|
||||
optimizer_config=optimization.OptimizationConfig(
|
||||
learning_rate=optimization.LrConfig(
|
||||
type='cosine',
|
||||
cosine=optimization.lr_cfg.CosineLrConfig(
|
||||
decay_steps=train_steps, initial_learning_rate=0.001
|
||||
),
|
||||
),
|
||||
optimizer=optimization.OptimizerConfig(
|
||||
type='sgd', sgd=optimization.SGDConfig(momentum=0.9)
|
||||
),
|
||||
),
|
||||
checkpoint_interval=steps_per_loop,
|
||||
steps_per_loop=steps_per_loop,
|
||||
summary_interval=steps_per_loop,
|
||||
validation_interval=steps_per_loop,
|
||||
train_steps=train_steps,
|
||||
validation_steps=-1,
|
||||
),
|
||||
restrictions=[
|
||||
'task.train_data.is_training != None',
|
||||
'task.validation_data.is_training != None',
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@exp_factory.register_config_factory('coca')
|
||||
def coca() -> cfg.ExperimentConfig:
|
||||
"""Gets experimental configs for tf-hub models."""
|
||||
|
||||
batch_size = 8
|
||||
train_steps = 625000
|
||||
steps_per_loop = 1250
|
||||
return cfg.ExperimentConfig(
|
||||
task=image_classification.ImageClassificationTask(
|
||||
model=image_classification.ImageClassificationModel(
|
||||
num_classes=1000,
|
||||
input_size=[288, 288, 3],
|
||||
backbone=backbones.Backbone(
|
||||
type='hub_model',
|
||||
hub_model=backbones.HubModel(
|
||||
handle=_COCA_HANDLE,
|
||||
trainable=False,
|
||||
mean_rgb=0.0,
|
||||
stddev_rgb=255.0,
|
||||
),
|
||||
),
|
||||
dropout_rate=0.0,
|
||||
),
|
||||
losses=image_classification.Losses(
|
||||
l2_weight_decay=0.0, label_smoothing=0.1, one_hot=True
|
||||
),
|
||||
train_data=image_classification.DataConfig(
|
||||
input_path=os.path.join(
|
||||
image_classification.IMAGENET_INPUT_PATH_BASE, 'train*'
|
||||
),
|
||||
aug_type=None,
|
||||
dtype='float32',
|
||||
global_batch_size=batch_size,
|
||||
is_training=True,
|
||||
decode_jpeg_only=False,
|
||||
),
|
||||
validation_data=image_classification.DataConfig(
|
||||
input_path=os.path.join(
|
||||
image_classification.IMAGENET_INPUT_PATH_BASE, 'valid*'
|
||||
),
|
||||
dtype='float32',
|
||||
global_batch_size=batch_size,
|
||||
is_training=False,
|
||||
decode_jpeg_only=False,
|
||||
drop_remainder=False,
|
||||
),
|
||||
),
|
||||
trainer=cfg.TrainerConfig(
|
||||
best_checkpoint_eval_metric='accuracy',
|
||||
best_checkpoint_export_subdir='best_ckpt',
|
||||
best_checkpoint_metric_comp='higher',
|
||||
optimizer_config=optimization.OptimizationConfig(
|
||||
learning_rate=optimization.LrConfig(
|
||||
type='cosine',
|
||||
cosine=optimization.lr_cfg.CosineLrConfig(
|
||||
decay_steps=train_steps, initial_learning_rate=0.001
|
||||
),
|
||||
),
|
||||
optimizer=optimization.OptimizerConfig(
|
||||
type='sgd', sgd=optimization.SGDConfig(momentum=0.9)
|
||||
),
|
||||
),
|
||||
checkpoint_interval=steps_per_loop,
|
||||
steps_per_loop=steps_per_loop,
|
||||
summary_interval=steps_per_loop,
|
||||
validation_interval=steps_per_loop,
|
||||
train_steps=train_steps,
|
||||
validation_steps=-1,
|
||||
),
|
||||
restrictions=[
|
||||
'task.train_data.is_training != None',
|
||||
'task.validation_data.is_training != None',
|
||||
],
|
||||
)
|
||||
+86
@@ -0,0 +1,86 @@
|
||||
# Dockerfile for basic training dockers with tfvision.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/tfvision/dockerfile/base.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM tensorflow/tensorflow:2.11.0-gpu
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# This is added to fix docker build error related to Nvidia key update.
|
||||
RUN rm -f /etc/apt/sources.list.d/cuda.list
|
||||
RUN curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add -
|
||||
|
||||
# Install basic libs.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake \
|
||||
curl \
|
||||
wget \
|
||||
sudo \
|
||||
gnupg \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender-dev \
|
||||
lsb-release \
|
||||
ca-certificates \
|
||||
build-essential \
|
||||
git \
|
||||
vim \
|
||||
screen \
|
||||
libtcmalloc-minimal4
|
||||
|
||||
|
||||
# Install google cloud SDK.
|
||||
RUN wget -q https://dl.google.com/dl/cloudsdk/channels/rapid/downloads/google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN tar xzf google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN ./google-cloud-sdk/install.sh -q
|
||||
# Make sure gsutil will use the default service account.
|
||||
RUN echo '[GoogleCompute]\nservice_account = default' > /etc/boto.cfg
|
||||
|
||||
|
||||
# Install required libs.
|
||||
RUN pip install --upgrade pip
|
||||
RUN pip install cloud-tpu-client==0.10
|
||||
RUN pip install pyyaml==5.4.1
|
||||
RUN pip install fsspec==2021.10.1
|
||||
RUN pip install gcsfs==2021.10.1
|
||||
RUN pip install tensorflow-text==2.11.0
|
||||
RUN pip install tf-models-official==2.11.3
|
||||
RUN pip install pyglove==0.1.0
|
||||
RUN pip install cloudml-hypertune==0.1.0.dev6
|
||||
RUN pip install object-detection==0.0.3
|
||||
RUN pip install pylint==2.17.2
|
||||
|
||||
# Installs Reduction Server NCCL plugin.
|
||||
RUN echo "deb https://packages.cloud.google.com/apt google-fast-socket main" | tee /etc/apt/sources.list.d/google-fast-socket.list \
|
||||
&& curl -s -L https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add - \
|
||||
&& apt update && apt install -y google-reduction-server
|
||||
|
||||
ENV PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=cpp
|
||||
|
||||
# Lower the memory fragmentation, and speed up the training.
|
||||
# https://github.com/tensorflow/tensorflow/issues/44176#issuecomment-783768033
|
||||
ENV LD_PRELOAD=/usr/lib/x86_64-linux-gnu/libtcmalloc_minimal.so.4
|
||||
|
||||
# Enable userspace DNS cache
|
||||
ENV GCS_RESOLVE_REFRESH_SECS=60
|
||||
ENV GCS_REQUEST_CONNECTION_TIMEOUT_SECS=300
|
||||
ENV GCS_METADATA_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_READ_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_WRITE_REQUEST_TIMEOUT_SECS=600
|
||||
# Each opened GCS file takes GCS_READ_CACHE_BLOCK_SIZE_MB of RAM, reduce the
|
||||
# value from the default 64MB to 8MB to decrease memory footprint.
|
||||
ENV GCS_READ_CACHE_BLOCK_SIZE_MB=8
|
||||
|
||||
WORKDIR /usr/local/lib/python3.8/dist-packages/official/vision
|
||||
|
||||
ENTRYPOINT ["python3","train.py"]
|
||||
|
||||
CMD ["--experiment=YOUR_EXPERIMENT",\
|
||||
"--config_file=YOUR_CONFIG_FILE",\
|
||||
"--mode=YOUR_MODE",\
|
||||
"--model_dir=YOUR_MODEL_DIR"]
|
||||
+86
@@ -0,0 +1,86 @@
|
||||
# Dockerfile for basic training dockers with tfvision.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/tfvision/dockerfile/base_v2.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM tensorflow/build:2.12-python3.9
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# This is added to fix docker build error related to Nvidia key update.
|
||||
RUN rm -f /etc/apt/sources.list.d/cuda.list
|
||||
RUN curl https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add -
|
||||
|
||||
# Install basic libs.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake \
|
||||
curl \
|
||||
wget \
|
||||
sudo \
|
||||
gnupg \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender-dev \
|
||||
lsb-release \
|
||||
ca-certificates \
|
||||
build-essential \
|
||||
git \
|
||||
vim \
|
||||
screen \
|
||||
libtcmalloc-minimal4
|
||||
|
||||
|
||||
# Install google cloud SDK.
|
||||
RUN wget -q https://dl.google.com/dl/cloudsdk/channels/rapid/downloads/google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN tar xzf google-cloud-sdk-359.0.0-linux-x86_64.tar.gz
|
||||
RUN ./google-cloud-sdk/install.sh -q
|
||||
# Make sure gsutil will use the default service account.
|
||||
RUN echo '[GoogleCompute]\nservice_account = default' > /etc/boto.cfg
|
||||
|
||||
|
||||
# Install required libs.
|
||||
RUN pip install --upgrade pip
|
||||
RUN pip install cloud-tpu-client==0.10
|
||||
RUN pip install pyyaml==5.4.1
|
||||
RUN pip install fsspec==2021.10.1
|
||||
RUN pip install gcsfs==2021.10.1
|
||||
RUN pip install tensorflow-text==2.12.1
|
||||
RUN pip install tf-models-official==2.12.0
|
||||
RUN pip install pyglove==0.1.0
|
||||
RUN pip install cloudml-hypertune==0.1.0.dev6
|
||||
RUN pip install object-detection==0.0.3
|
||||
RUN pip install pylint==2.17.2
|
||||
|
||||
# Installs Reduction Server NCCL plugin.
|
||||
RUN echo "deb https://packages.cloud.google.com/apt google-fast-socket main" | tee /etc/apt/sources.list.d/google-fast-socket.list \
|
||||
&& curl -s -L https://packages.cloud.google.com/apt/doc/apt-key.gpg | apt-key add - \
|
||||
&& apt update && apt install -y google-reduction-server
|
||||
|
||||
ENV PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=cpp
|
||||
|
||||
# Lower the memory fragmentation, and speed up the training.
|
||||
# https://github.com/tensorflow/tensorflow/issues/44176#issuecomment-783768033
|
||||
ENV LD_PRELOAD=/usr/lib/x86_64-linux-gnu/libtcmalloc_minimal.so.4
|
||||
|
||||
# Enable userspace DNS cache
|
||||
ENV GCS_RESOLVE_REFRESH_SECS=60
|
||||
ENV GCS_REQUEST_CONNECTION_TIMEOUT_SECS=300
|
||||
ENV GCS_METADATA_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_READ_REQUEST_TIMEOUT_SECS=300
|
||||
ENV GCS_WRITE_REQUEST_TIMEOUT_SECS=600
|
||||
# Each opened GCS file takes GCS_READ_CACHE_BLOCK_SIZE_MB of RAM, reduce the
|
||||
# value from the default 64MB to 8MB to decrease memory footprint.
|
||||
ENV GCS_READ_CACHE_BLOCK_SIZE_MB=8
|
||||
|
||||
WORKDIR /usr/local/lib/python3.9/dist-packages/official/vision
|
||||
|
||||
ENTRYPOINT ["python3","train.py"]
|
||||
|
||||
CMD ["--experiment=YOUR_EXPERIMENT",\
|
||||
"--config_file=YOUR_CONFIG_FILE",\
|
||||
"--mode=YOUR_MODE",\
|
||||
"--model_dir=YOUR_MODEL_DIR"]
|
||||
+68
@@ -0,0 +1,68 @@
|
||||
# Dockerfile for AutoML vision model export dockers with tfvision.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/tfvision/dockerfile/model_export.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/tfvision-base-v2:latest
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
RUN PROTOC_ZIP=protoc-3.9.2-linux-x86_64.zip && \
|
||||
curl -OL https://github.com/google/protobuf/releases/download/v3.9.2/$PROTOC_ZIP && \
|
||||
unzip -o $PROTOC_ZIP -d /usr/local bin/protoc && \
|
||||
unzip -o $PROTOC_ZIP -d /usr/local include/* && \
|
||||
rm -f $PROTOC_ZIP
|
||||
|
||||
COPY model_oss/tfvision /automl_vision/tfvision
|
||||
COPY model_oss/util /automl_vision/util
|
||||
|
||||
# Install tensorflow models following:
|
||||
# https://github.com/tensorflow/models/blob/master/research/object_detection/g3doc/tf2.md.
|
||||
# https://github.com/tensorflow/models/blob/master/research/object_detection/colab_tutorials/object_detection_tutorial.ipynb.
|
||||
RUN cd /automl_vision && \
|
||||
git clone --depth 1 https://github.com/tensorflow/models && \
|
||||
cd models/research && \
|
||||
protoc object_detection/protos/*.proto --python_out=. && \
|
||||
cp object_detection/packages/tf2/setup.py . && \
|
||||
pip install . && \
|
||||
cd /automl_vision && \
|
||||
rm -rf ./models
|
||||
|
||||
RUN pip install tensorflow-io==0.25.0
|
||||
|
||||
RUN pip install "opencv-python-headless<4.3"
|
||||
RUN pip install google-cloud-aiplatform==1.23.0
|
||||
|
||||
# Install yolov4, yolov7, and maxvit
|
||||
RUN mkdir /tmp/buffer && \
|
||||
cd /tmp/buffer && \
|
||||
git clone https://github.com/tensorflow/models.git && \
|
||||
cd models && \
|
||||
git reset --hard 6138633a41097a3c0f320bd895ac5da65c33016f && \
|
||||
cd /usr/local/lib/python3.9/dist-packages/official/projects/ && \
|
||||
cp -R /tmp/buffer/models/official/projects/yolo/ ./ && \
|
||||
cp -R /tmp/buffer/models/official/projects/maxvit/ ./ && \
|
||||
rm -rf /tmp/buffer
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/tfvision"
|
||||
|
||||
WORKDIR /automl_vision
|
||||
|
||||
# Run pylint to validate code.
|
||||
COPY .pylintrc /automl_vision/.pylintrc
|
||||
RUN find . -type f -name "*.py" | xargs pylint --rcfile=./.pylintrc --errors-only
|
||||
|
||||
ENTRYPOINT ["python3","tfvision/serving/export_oss_saved_model.py"]
|
||||
|
||||
CMD ["--experiment=YOUR_EXPERIMENT",\
|
||||
"--objective=YOUR_OBJECTIVE",\
|
||||
"--config_file=YOUR_CONFIG_FILE",\
|
||||
"--checkpoint_path=YOUR_CHECKPOINT_DIR",\
|
||||
"--label_map_path=YOUR_LABEL_MAP_PATH",\
|
||||
"--input_image_size=YOUR_INPUT_IMAGE_SIZE",\
|
||||
"--export_dir=YOUR_EXPORT_DIR"]
|
||||
+52
@@ -0,0 +1,52 @@
|
||||
# Dockerfile for AutoML vision training dockers with tfvision.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/tfvision/dockerfile/train_oss.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/tfvision-base:latest
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Fix yolo and retinanet issues.
|
||||
RUN mkdir /tmp/buffer && \
|
||||
cd /tmp/buffer && \
|
||||
git clone https://github.com/tensorflow/models.git && \
|
||||
cd models && \
|
||||
git checkout fbd4c57fd7e9f7d73da30ed3fc755b8c4c682df7 && \
|
||||
cd /usr/local/lib/python3.8/dist-packages/official/projects/yolo && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/optimization/optimizer_factory.py ./optimization/ && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/configs/yolo.py ./configs && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/factory.py ./modeling && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/layers/detection_generator.py ./modeling/layers && \
|
||||
cd /usr/local/lib/python3.8/dist-packages/official/vision && \
|
||||
cp /tmp/buffer/models/official/vision/configs/retinanet.py ./configs && \
|
||||
cp /tmp/buffer/models/official/vision/modeling/layers/detection_generator.py ./modeling/layers && \
|
||||
cp /tmp/buffer/models/official/vision/modeling/layers/edgetpu.py ./modeling/layers && \
|
||||
rm -rf /tmp/buffer
|
||||
|
||||
COPY model_oss/tfvision /automl_vision/tfvision
|
||||
COPY model_oss/util /automl_vision/util
|
||||
RUN rm -rf /automl_vision/tfvision/serving
|
||||
|
||||
WORKDIR /automl_vision
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/util"
|
||||
|
||||
# Run pylint to validate code.
|
||||
COPY .pylintrc /automl_vision/.pylintrc
|
||||
RUN find . -type f -name "*.py" | xargs pylint --rcfile=./.pylintrc --errors-only
|
||||
|
||||
ENTRYPOINT ["python3","tfvision/train_hpt_oss.py"]
|
||||
|
||||
CMD ["--experiment=YOUR_EXPERIMENT",\
|
||||
"--config_file=",\
|
||||
"--mode=YOUR_MODE",\
|
||||
"--model_dir=YOUR_MODEL_DIR",\
|
||||
"--objective=YOUR_OBJECTIVE",\
|
||||
"--learning_rate=",\
|
||||
"--anchor_size="]
|
||||
+80
@@ -0,0 +1,80 @@
|
||||
# Dockerfile for AutoML vision training dockers with tfvision.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/tfvision/dockerfile/train_oss_v2.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/tfvision-base-v2:latest
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Fix yolo and retinanet issues.
|
||||
RUN mkdir /tmp/buffer && \
|
||||
cd /tmp/buffer && \
|
||||
git clone https://github.com/tensorflow/models.git && \
|
||||
cd models && \
|
||||
# Add support for newly added config options.
|
||||
git reset --hard ed6d4d220b86237980d3f7563d261d19e040ef1a && \
|
||||
cd /usr/local/lib/python3.9/dist-packages/official/projects/yolo && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/dataloaders/yolo_input.py ./dataloaders/ && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/optimization/optimizer_factory.py ./optimization/ && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/configs/yolo.py ./configs && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/factory.py ./modeling && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/layers/detection_generator.py ./modeling/layers && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/common/registry_imports.py ./common && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/configs/yolov7.py ./configs && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/configs/decoders.py ./configs && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/configs/backbones.py ./configs && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/yolov7_model.py ./modeling && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/backbones/yolov7.py ./modeling/backbones && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/decoders/yolov7.py ./modeling/decoders && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/heads/yolov7_head.py ./modeling/heads && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/modeling/layers/nn_blocks.py ./modeling/layers && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/losses/yolov7_loss.py ./losses && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/tasks/yolov7.py ./tasks && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/ops/initializer_ops.py ./ops && \
|
||||
cp /tmp/buffer/models/official/projects/yolo/ops/mosaic.py ./ops && \
|
||||
cd /usr/local/lib/python3.9/dist-packages/official/vision && \
|
||||
cp /tmp/buffer/models/official/vision/configs/retinanet.py ./configs && \
|
||||
cp /tmp/buffer/models/official/vision/modeling/layers/detection_generator.py ./modeling/layers && \
|
||||
cp /tmp/buffer/models/official/vision/modeling/layers/edgetpu.py ./modeling/layers && \
|
||||
cp /tmp/buffer/models/official/vision/ops/augment.py ./ops && \
|
||||
rm -rf /tmp/buffer
|
||||
|
||||
# Add MaxViT
|
||||
RUN mkdir /tmp/buffer && \
|
||||
cd /tmp/buffer && \
|
||||
git clone https://github.com/tensorflow/models.git && \
|
||||
cd models && \
|
||||
git reset --hard 6138633a41097a3c0f320bd895ac5da65c33016f && \
|
||||
cd /usr/local/lib/python3.9/dist-packages/official/projects/ && \
|
||||
cp -R /tmp/buffer/models/official/projects/maxvit/ ./ && \
|
||||
rm -rf /tmp/buffer
|
||||
ENV ENABLE_MAX_VIT "True"
|
||||
|
||||
|
||||
COPY model_oss/tfvision /automl_vision/tfvision
|
||||
COPY model_oss/util /automl_vision/util
|
||||
RUN rm -rf /automl_vision/tfvision/serving
|
||||
|
||||
WORKDIR /automl_vision
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/util"
|
||||
|
||||
# Run pylint to validate code.
|
||||
COPY .pylintrc /automl_vision/.pylintrc
|
||||
RUN find . -type f -name "*.py" | xargs pylint --rcfile=./.pylintrc --errors-only
|
||||
|
||||
ENTRYPOINT ["python3","tfvision/train_hpt_oss.py"]
|
||||
|
||||
CMD ["--experiment=YOUR_EXPERIMENT",\
|
||||
"--config_file=",\
|
||||
"--mode=YOUR_MODE",\
|
||||
"--model_dir=YOUR_MODEL_DIR",\
|
||||
"--objective=YOUR_OBJECTIVE",\
|
||||
"--learning_rate=",\
|
||||
"--anchor_size="]
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
"""Backbones package definition."""
|
||||
|
||||
from tfvision.modeling.backbones import hub_model
|
||||
+165
@@ -0,0 +1,165 @@
|
||||
"""Loads a tf-hub model."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Mapping, Optional
|
||||
|
||||
from absl import logging
|
||||
import tensorflow as tf
|
||||
import tensorflow_hub as hub
|
||||
|
||||
from official.modeling import hyperparams
|
||||
from official.vision.modeling.backbones import factory
|
||||
from official.vision.ops import preprocess_ops
|
||||
|
||||
layers = tf.keras.layers
|
||||
|
||||
|
||||
@tf.keras.utils.register_keras_serializable(package='Vision')
|
||||
class HubModel(tf.keras.Model):
|
||||
"""A tf-hub model wrapper."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
handle: str,
|
||||
input_specs: tf.keras.layers.InputSpec = layers.InputSpec(
|
||||
shape=[None, None, None, 3]
|
||||
),
|
||||
trainable: bool = True,
|
||||
mean_rgb: Optional[float] = None,
|
||||
stddev_rgb: Optional[float] = None,
|
||||
kernel_regularizer: Optional[tf.keras.regularizers.Regularizer] = None,
|
||||
signature: Optional[str] = None,
|
||||
output_key: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initializes a tf-hub model.
|
||||
|
||||
Args:
|
||||
handle: A handle to load a saved model via hub.load().
|
||||
input_specs: A input_spec of the input tensor.
|
||||
trainable: Controls whether this layer is trainable. Must not be set to
|
||||
True when using a signature (raises ValueError), including the use of
|
||||
legacy TF1 Hub format.
|
||||
mean_rgb: The mean rgb value used for normalization.
|
||||
stddev_rgb: The standard deviation of rgb values used for normalization.
|
||||
kernel_regularizer: A regularizer object for kernel weights.
|
||||
signature: Optional. If set, KerasLayer will use the requested signature.
|
||||
For legacy models in TF1 Hub format leaving unset means to use the
|
||||
`default` signature. When using a signature, output_key have to set.
|
||||
output_key: Name of the output item to return if the layer returns a dict.
|
||||
For legacy models in TF1 Hub format leaving unset means to return the
|
||||
`default` output.
|
||||
**kwargs: Additional keyword arguments to be passed.
|
||||
"""
|
||||
self._handle = handle
|
||||
self._mean_rgb = mean_rgb
|
||||
self._stddev_rgb = stddev_rgb
|
||||
self._kernel_regularizer = kernel_regularizer
|
||||
self._signature = signature
|
||||
self._output_key = output_key
|
||||
|
||||
inputs = tf.keras.Input(shape=input_specs.shape[1:])
|
||||
x = inputs
|
||||
if mean_rgb or stddev_rgb:
|
||||
x = layers.Lambda(self.re_normalize)(x)
|
||||
|
||||
model = hub.KerasLayer(
|
||||
handle=handle,
|
||||
trainable=trainable,
|
||||
signature=signature,
|
||||
output_key=output_key,
|
||||
)
|
||||
if trainable and kernel_regularizer:
|
||||
if hasattr(model, 'regularization_losses'):
|
||||
logging.warning('regularization_losses already defined in the model.')
|
||||
|
||||
def reg_loss(x):
|
||||
return lambda: kernel_regularizer(x)
|
||||
|
||||
for v in model.trainable_variables:
|
||||
if 'kernel' in v.name:
|
||||
model.add_loss(reg_loss(v))
|
||||
x = model(x)
|
||||
if not trainable:
|
||||
# Solves backpropagation errors when loading CoCa.
|
||||
x = tf.stop_gradient(x)
|
||||
endpoints = {'0': x[:, tf.newaxis, tf.newaxis, :]}
|
||||
|
||||
self._output_specs = {l: endpoints[l].get_shape() for l in endpoints}
|
||||
|
||||
super().__init__(
|
||||
inputs=inputs, outputs=endpoints, trainable=trainable, **kwargs
|
||||
)
|
||||
|
||||
def re_normalize(self, x: tf.Tensor) -> tf.Tensor:
|
||||
"""Re-normalizes the input image.
|
||||
|
||||
Tf-vision normalizes the images from [0, 255] to normal distribution. The
|
||||
tf-hub models are usually normalized to [0.0, 1.0]. This function converts
|
||||
the input image to proper scale.
|
||||
|
||||
Args:
|
||||
x: The input image.
|
||||
|
||||
Returns:
|
||||
The re-normalized image.
|
||||
"""
|
||||
offset = tf.constant(preprocess_ops.MEAN_RGB)
|
||||
scale = tf.constant(preprocess_ops.STDDEV_RGB)
|
||||
x = x * scale + offset
|
||||
|
||||
if self._mean_rgb:
|
||||
x -= self._mean_rgb
|
||||
if self._stddev_rgb:
|
||||
x /= self._stddev_rgb
|
||||
return x
|
||||
|
||||
def get_config(self) -> Mapping[str, Any]:
|
||||
config_dict = {
|
||||
'handle': self._handle,
|
||||
'trainable': self.trainable,
|
||||
'mean_rgb': self._mean_rgb,
|
||||
'stddev_rgb': self._stddev_rgb,
|
||||
'kernel_regularizer': self._kernel_regularizer,
|
||||
'signature': self._signature,
|
||||
'output_key': self._output_key,
|
||||
}
|
||||
return config_dict
|
||||
|
||||
@classmethod
|
||||
def from_config(cls,
|
||||
config: Mapping[str, Any],
|
||||
custom_objects: Optional[Any] = None) -> HubModel:
|
||||
return cls(**config)
|
||||
|
||||
@property
|
||||
def output_specs(self) -> Mapping[str, tf.TensorShape]:
|
||||
"""A dict of {level: TensorShape} pairs for the model output."""
|
||||
return self._output_specs
|
||||
|
||||
|
||||
@factory.register_backbone_builder('hub_model')
|
||||
def build_hub_model(
|
||||
input_specs: tf.keras.layers.InputSpec,
|
||||
backbone_config: hyperparams.Config,
|
||||
l2_regularizer: tf.keras.regularizers.Regularizer = None,
|
||||
**kwargs: Any,
|
||||
) -> tf.keras.Model: # pytype: disable=annotation-type-mismatch # typed-keras
|
||||
"""Builds ResNet backbone from a config."""
|
||||
del kwargs
|
||||
backbone_type = backbone_config.type
|
||||
backbone_cfg = backbone_config.get()
|
||||
assert backbone_type == 'hub_model', (f'Inconsistent backbone type '
|
||||
f'{backbone_type}')
|
||||
|
||||
return HubModel(
|
||||
input_specs=input_specs,
|
||||
handle=backbone_cfg.handle,
|
||||
trainable=backbone_cfg.trainable,
|
||||
mean_rgb=backbone_cfg.mean_rgb,
|
||||
stddev_rgb=backbone_cfg.stddev_rgb,
|
||||
kernel_regularizer=l2_regularizer,
|
||||
signature=backbone_cfg.signature,
|
||||
output_key=backbone_cfg.output_key,
|
||||
)
|
||||
@@ -0,0 +1,4 @@
|
||||
"""AutoML Vision tf-vision custom code import."""
|
||||
# pylint: disable=unused-import
|
||||
from tfvision import configs
|
||||
from tfvision.modeling import backbones
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
"""AutoML tfvision saved_model constants."""
|
||||
|
||||
# Tfvision training artifact marcos.
|
||||
# Exported parameter.yaml in the model directory.
|
||||
CFG_FILENAME = 'params.yaml'
|
||||
|
||||
# Common automl saved_model marcos.
|
||||
## Type of input to automl saved_model, fixed as image bytes string.
|
||||
INPUT_TYPE = 'image_bytes'
|
||||
IMAGE_TENSOR = 'image_tensor'
|
||||
## Automl IOD saved_model signature input image argument name.
|
||||
IOD_INPUT_NAME = 'encoded_image'
|
||||
## ICN saved_model input name.
|
||||
ICN_INPUT_NAME = 'image_bytes'
|
||||
## Automl saved_model signature input key argument name.
|
||||
INPUT_KEY_NAME = 'key'
|
||||
OUTPUT_KEY_NAME = 'key'
|
||||
|
||||
# IOD saved_model marcos.
|
||||
## IOD class as text output
|
||||
DETECTION_CLASSES_AS_TEXT = 'detection_classes_as_text'
|
||||
## Default value for labelmap text lookup table.
|
||||
LOOKUP_DEFAULT_VALUE = 'unknown'
|
||||
## Suffix for signature def without input key tensor.
|
||||
NO_KEY_SIG_DEF_SUFFIX = '_without_key'
|
||||
@@ -0,0 +1,475 @@
|
||||
"""Detection input and model functions for serving/inference."""
|
||||
|
||||
import functools
|
||||
import heapq
|
||||
from typing import Any, Callable, Dict, List, Optional, Text
|
||||
|
||||
import tensorflow as tf
|
||||
|
||||
from tfvision.serving import automl_constants
|
||||
from object_detection.utils import label_map_util
|
||||
from official.core import config_definitions as cfg
|
||||
from official.projects.yolo.modeling import factory as yolo_factory
|
||||
from official.projects.yolo.modeling.decoders import yolo_decoder # pylint: disable=unused-import
|
||||
from official.projects.yolo.serving import model_fn as yolo_model_fn
|
||||
from official.vision import configs
|
||||
from official.vision.ops import box_ops
|
||||
from official.vision.serving import detection as detection_module
|
||||
|
||||
|
||||
def load_label_map_to_string_list(label_map_path: str,
|
||||
fill_in_gaps_and_background: bool = True
|
||||
) -> List[str]:
|
||||
"""Loads class labels as string list ordered by class id.
|
||||
|
||||
Args:
|
||||
label_map_path: the path to label_map.pbtxt with string_int_label_map_pb2
|
||||
proto format.
|
||||
fill_in_gaps_and_background: whether to fill in gaps and background with
|
||||
respect to the id field in the proto. The id: 0 is reserved for the
|
||||
'background' class and will be added if it is missing. All other missing
|
||||
ids in range(1, max(id)) will be added with a dummy class name
|
||||
("class_<id>") if they are missing.
|
||||
|
||||
Returns:
|
||||
The class labels as text string lists in the order of the class numeric id.
|
||||
"""
|
||||
|
||||
labelmap = label_map_util.get_label_map_dict(
|
||||
label_map_path, fill_in_gaps_and_background=fill_in_gaps_and_background)
|
||||
heap = []
|
||||
for label_name, label_id in labelmap.items():
|
||||
heapq.heappush(heap, (label_id, label_name))
|
||||
label_list = [heapq.heappop(heap)[1] for _ in range(len(heap))]
|
||||
|
||||
return label_list
|
||||
|
||||
|
||||
class DetectionModule(detection_module.DetectionModule):
|
||||
"""Detection Module."""
|
||||
|
||||
def __init__(self,
|
||||
params: cfg.ExperimentConfig,
|
||||
*,
|
||||
batch_size: int,
|
||||
input_image_size: List[int],
|
||||
input_type: str = automl_constants.INPUT_TYPE,
|
||||
num_channels: int = 3,
|
||||
model: Optional[tf.keras.Model] = None,
|
||||
label_map_path: Optional[str] = None,
|
||||
input_name: str = automl_constants.IOD_INPUT_NAME,
|
||||
key_name: str = automl_constants.INPUT_KEY_NAME):
|
||||
"""Initializes a module for export.
|
||||
|
||||
Args:
|
||||
params: Experiment params.
|
||||
batch_size: The batch size of the model input. Can be `int` or None.
|
||||
input_image_size: List or Tuple of size of the input image. For 2D image,
|
||||
it is [height, width].
|
||||
input_type: The input signature type.
|
||||
num_channels: The number of the image channels.
|
||||
model: A tf.keras.Model instance to be exported.
|
||||
label_map_path: A labelmap proto file path.
|
||||
input_name: A customized input tensor name. This will be used as the
|
||||
signature's input image argument name.
|
||||
key_name: A name to the automl model input key.
|
||||
"""
|
||||
self._key_name = key_name
|
||||
if label_map_path is not None:
|
||||
self._label_map_table = self._generate_label_map_list(label_map_path)
|
||||
else:
|
||||
self._label_map_table = None
|
||||
super().__init__(
|
||||
params=params,
|
||||
model=model,
|
||||
batch_size=batch_size,
|
||||
input_image_size=input_image_size,
|
||||
input_name=input_name,
|
||||
input_type=input_type)
|
||||
|
||||
def _generate_label_map_list(self, label_map_path: str) -> tf.Tensor:
|
||||
"""Generates a list of label texts from a labelmap path."""
|
||||
mapping_string = tf.convert_to_tensor(
|
||||
load_label_map_to_string_list(label_map_path))
|
||||
return tf.lookup.index_to_string_table_from_tensor(
|
||||
mapping_string, default_value=automl_constants.LOOKUP_DEFAULT_VALUE)
|
||||
|
||||
def _generate_class_text_output(self, detection_classes) -> tf.Tensor:
|
||||
"""Converts class index to class text."""
|
||||
if self._label_map_table is None:
|
||||
raise ValueError('_label_map_table is None.')
|
||||
indices = tf.cast(detection_classes, tf.int64)
|
||||
indices = tf.reshape(indices, [-1])
|
||||
values = self._label_map_table.lookup(indices)
|
||||
return tf.reshape(
|
||||
values, [-1, tf.array_ops.shape(detection_classes)[1]],
|
||||
name=automl_constants.DETECTION_CLASSES_AS_TEXT)
|
||||
|
||||
def serve(self,
|
||||
images: tf.Tensor,
|
||||
key: Optional[tf.Tensor] = None) -> Dict[Text, tf.Tensor]:
|
||||
"""Cast image to float and run inference.
|
||||
|
||||
Args:
|
||||
images: uint8 Tensor of input images. For input type image tensor, the
|
||||
shape is [batch_size, None, None, 3], for image_bytes, the shape is
|
||||
[batch_size].
|
||||
key: Optional string Tensor of shape [batch_size]. If not provided
|
||||
output tensors will not contain it as well.
|
||||
|
||||
Returns:
|
||||
Tensor holding detection output logits.
|
||||
"""
|
||||
|
||||
images, anchor_boxes, image_info = self.preprocess(images)
|
||||
input_image_shape = image_info[:, 1, :]
|
||||
|
||||
# To overcome keras.Model extra limitation to save a model with layers that
|
||||
# have multiple inputs, we use `model.call` here to trigger the forward
|
||||
# path. Note that, this disables some keras magics happens in `__call__`.
|
||||
detections = self.model.call(
|
||||
images=images,
|
||||
image_shape=input_image_shape,
|
||||
anchor_boxes=anchor_boxes,
|
||||
training=False)
|
||||
|
||||
if self.params.task.model.detection_generator.apply_nms:
|
||||
# For RetinaNet model, apply export_config.
|
||||
if isinstance(self.params.task.model, configs.retinanet.RetinaNet):
|
||||
export_config = self.params.task.export_config
|
||||
# Normalize detection box coordinates to [0, 1].
|
||||
if export_config.output_normalized_coordinates:
|
||||
detection_boxes = (
|
||||
detections['detection_boxes'] /
|
||||
tf.tile(image_info[:, 2:3, :], [1, 1, 2]))
|
||||
detections['detection_boxes'] = box_ops.normalize_boxes(
|
||||
detection_boxes, image_info[:, 0:1, :])
|
||||
|
||||
# Cast num_detections and detection_classes to float. This allows the
|
||||
# model inference to work on chain (go/chain) as chain requires floating
|
||||
# point outputs.
|
||||
if export_config.cast_num_detections_to_float:
|
||||
detections['num_detections'] = tf.cast(
|
||||
detections['num_detections'], dtype=tf.float32)
|
||||
if export_config.cast_detection_classes_to_float:
|
||||
detections['detection_classes'] = tf.cast(
|
||||
detections['detection_classes'], dtype=tf.float32)
|
||||
|
||||
final_outputs = {
|
||||
'detection_boxes': detections['detection_boxes'],
|
||||
'detection_scores': detections['detection_scores'],
|
||||
'detection_classes': detections['detection_classes'],
|
||||
'num_detections': detections['num_detections']
|
||||
}
|
||||
else:
|
||||
final_outputs = {
|
||||
'decoded_boxes': detections['decoded_boxes'],
|
||||
'decoded_box_scores': detections['decoded_box_scores']
|
||||
}
|
||||
|
||||
if 'detection_masks' in detections.keys():
|
||||
final_outputs['detection_masks'] = detections['detection_masks']
|
||||
|
||||
# Adding AutoML specific outputs.
|
||||
if self._label_map_table is not None:
|
||||
final_outputs.update({
|
||||
automl_constants.DETECTION_CLASSES_AS_TEXT:
|
||||
self._generate_class_text_output(detections['detection_classes'])
|
||||
})
|
||||
|
||||
final_outputs.update({'image_info': image_info})
|
||||
if key is not None:
|
||||
final_outputs.update({automl_constants.OUTPUT_KEY_NAME: key})
|
||||
|
||||
return final_outputs
|
||||
|
||||
@tf.function
|
||||
def inference_from_image_bytes(
|
||||
self,
|
||||
inputs: tf.Tensor,
|
||||
key: tf.Tensor,
|
||||
) -> Dict[Text, tf.Tensor]:
|
||||
"""Entry point for model input.
|
||||
|
||||
Raw image tensor will be decoded to the desired image format.
|
||||
|
||||
Args:
|
||||
inputs: Image tensor to be feed to the model.
|
||||
key: AutoML specific input key to track image names or image ids.
|
||||
|
||||
Returns:
|
||||
A dictionary of Tensor that contains model outputs.
|
||||
"""
|
||||
with tf.device('cpu:0'):
|
||||
images = tf.nest.map_structure(
|
||||
tf.identity,
|
||||
tf.map_fn(
|
||||
self._decode_image,
|
||||
elems=inputs,
|
||||
fn_output_signature=tf.TensorSpec(
|
||||
shape=[None] * len(self._input_image_size) +
|
||||
[self._num_channels],
|
||||
dtype=tf.uint8),
|
||||
parallel_iterations=32))
|
||||
images = tf.stack(images)
|
||||
|
||||
return self.serve(images, key)
|
||||
|
||||
@tf.function
|
||||
def inference_from_image_bytes_wo_key(
|
||||
self, inputs: tf.Tensor) -> Dict[Text, tf.Tensor]:
|
||||
"""Entry point for model inference without input key tensor.
|
||||
|
||||
Raw image tensor will be decoded to the desired image format.
|
||||
|
||||
Args:
|
||||
inputs: Image tensor to be feed to the model.
|
||||
|
||||
Returns:
|
||||
A dictionary of Tensor that contains model outputs.
|
||||
"""
|
||||
with tf.device('cpu:0'):
|
||||
images = tf.nest.map_structure(
|
||||
tf.identity,
|
||||
tf.map_fn(
|
||||
self._decode_image,
|
||||
elems=inputs,
|
||||
fn_output_signature=tf.TensorSpec(
|
||||
shape=[None] * len(self._input_image_size) +
|
||||
[self._num_channels],
|
||||
dtype=tf.uint8),
|
||||
parallel_iterations=32))
|
||||
images = tf.stack(images)
|
||||
|
||||
return self.serve(images)
|
||||
|
||||
def get_inference_signatures(
|
||||
self, function_keys: Dict[Text, Text]
|
||||
) -> Dict[Text, Callable[[tf.Tensor, tf.Tensor], Dict[Text, tf.Tensor]]]:
|
||||
"""Gets defined function signatures.
|
||||
|
||||
Args:
|
||||
function_keys: A dictionary with keys as the function to create signature
|
||||
for and values as the signature keys when returns.
|
||||
|
||||
Returns:
|
||||
A dictionary with key as signature key and value as concrete functions
|
||||
that can be used for tf.saved_model.save.
|
||||
"""
|
||||
signatures = {}
|
||||
for key, def_name in function_keys.items():
|
||||
# Adds input string 'key' to image_bytes input type.
|
||||
if key == automl_constants.INPUT_TYPE:
|
||||
input_images = tf.TensorSpec(
|
||||
shape=[self._batch_size], dtype=tf.string, name=self._input_name)
|
||||
input_key = tf.TensorSpec(
|
||||
shape=[self._batch_size], dtype=tf.string, name=self._key_name)
|
||||
signatures[
|
||||
def_name] = self.inference_from_image_bytes.get_concrete_function(
|
||||
input_images, input_key)
|
||||
# For each input type, create a signature without input key tensor.
|
||||
def_name_wo_key = def_name + automl_constants.NO_KEY_SIG_DEF_SUFFIX
|
||||
signatures[def_name_wo_key] = (
|
||||
self.inference_from_image_bytes_wo_key.get_concrete_function(
|
||||
input_images))
|
||||
else:
|
||||
raise ValueError('Unrecognized `input_type`')
|
||||
return signatures
|
||||
|
||||
|
||||
class YoloDetectionModule(DetectionModule):
|
||||
"""Yolo detection module for Model Garden."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params: cfg.ExperimentConfig,
|
||||
*,
|
||||
batch_size: int,
|
||||
input_image_size: List[int],
|
||||
preprocessor: Callable[..., Any],
|
||||
inference_step: Callable[..., Any],
|
||||
input_type: str = automl_constants.INPUT_TYPE,
|
||||
num_channels: int = 3,
|
||||
model: Optional[tf.keras.Model] = None,
|
||||
label_map_path: Optional[str] = None,
|
||||
input_name: str = automl_constants.IOD_INPUT_NAME,
|
||||
key_name: str = automl_constants.INPUT_KEY_NAME,
|
||||
):
|
||||
"""Initializes a module for export.
|
||||
|
||||
Args:
|
||||
params: Experiment params.
|
||||
batch_size: The batch size of the model input. Can be `int` or None.
|
||||
input_image_size: List or Tuple of size of the input image. For 2D image,
|
||||
it is [height, width].
|
||||
preprocessor: An optional callable to preprocess the inputs.
|
||||
inference_step: An optional callable to forward-pass the model.
|
||||
input_type: The input signature type.
|
||||
num_channels: The number of the image channels.
|
||||
model: A tf.keras.Model instance to be exported.
|
||||
label_map_path: A labelmap proto file path.
|
||||
input_name: A customized input tensor name. This will be used as the
|
||||
signature's input image argument name.
|
||||
key_name: A name to the automl model input key.
|
||||
"""
|
||||
super().__init__(
|
||||
params=params,
|
||||
batch_size=batch_size,
|
||||
input_image_size=input_image_size,
|
||||
input_type=input_type,
|
||||
num_channels=num_channels,
|
||||
model=model,
|
||||
label_map_path=label_map_path,
|
||||
input_name=input_name,
|
||||
key_name=key_name,
|
||||
)
|
||||
|
||||
self.preprocessor = preprocessor
|
||||
self.inference_step = functools.partial(inference_step, model=self.model)
|
||||
|
||||
def preprocess(self, images: tf.Tensor) -> None:
|
||||
raise NotImplementedError('Use self.preprocessor instead.')
|
||||
|
||||
def serve(
|
||||
self, images: tf.Tensor, key: Optional[tf.Tensor] = None
|
||||
) -> Dict[Text, tf.Tensor]:
|
||||
"""Cast image to float and run inference.
|
||||
|
||||
Args:
|
||||
images: uint8 Tensor of input images. For input type image tensor, the
|
||||
shape is [batch_size, None, None, 3], for image_bytes, the shape is
|
||||
[batch_size].
|
||||
key: Optional string Tensor of shape [batch_size]. If not provided output
|
||||
tensors will not contain it as well.
|
||||
|
||||
Returns:
|
||||
Tensor holding detection output logits.
|
||||
"""
|
||||
images, image_info = self.preprocessor(images)
|
||||
final_outputs = self.inference_step((images, image_info))
|
||||
|
||||
# Normalize detection box coordinates to [0, 1].
|
||||
detection_boxes = final_outputs['detection_boxes'] / tf.tile(
|
||||
image_info[:, 2:3, :], [1, 1, 2]
|
||||
)
|
||||
final_outputs['detection_boxes'] = box_ops.normalize_boxes(
|
||||
detection_boxes, image_info[:, 0:1, :]
|
||||
)
|
||||
|
||||
# Cast num_detections and detection_classes to float. This allows the
|
||||
# model inference to work on chain (go/chain) as chain requires floating
|
||||
# point outputs.
|
||||
final_outputs['num_detections'] = tf.cast(
|
||||
final_outputs['num_detections'], dtype=tf.float32
|
||||
)
|
||||
final_outputs['detection_classes'] = tf.cast(
|
||||
final_outputs['detection_classes'], dtype=tf.float32
|
||||
)
|
||||
|
||||
# Adding AutoML specific outputs.
|
||||
if self._label_map_table is not None:
|
||||
final_outputs.update(
|
||||
{
|
||||
automl_constants.DETECTION_CLASSES_AS_TEXT: (
|
||||
self._generate_class_text_output(
|
||||
final_outputs['detection_classes']
|
||||
)
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
final_outputs.update({'image_info': image_info})
|
||||
if key is not None:
|
||||
final_outputs.update({automl_constants.OUTPUT_KEY_NAME: key})
|
||||
|
||||
return final_outputs
|
||||
|
||||
|
||||
def create_yolov7_export_module(
|
||||
params: cfg.ExperimentConfig,
|
||||
input_type: str,
|
||||
batch_size: int,
|
||||
input_image_size: List[int],
|
||||
num_channels: int = 3,
|
||||
input_name: Optional[str] = None,
|
||||
label_map_path: Optional[str] = None,
|
||||
) -> YoloDetectionModule:
|
||||
"""Creates YOLO export module for Model Garden."""
|
||||
input_specs = tf.keras.layers.InputSpec(
|
||||
shape=[batch_size] + input_image_size + [num_channels]
|
||||
)
|
||||
model = yolo_factory.build_yolov7(
|
||||
input_specs=input_specs,
|
||||
model_config=params.task.model,
|
||||
l2_regularization=None,
|
||||
)
|
||||
|
||||
def preprocess_fn(image_tensor):
|
||||
def normalize_image_fn(inputs):
|
||||
image = tf.cast(inputs, dtype=tf.float32)
|
||||
return image / 255.0
|
||||
|
||||
# If input_type is `tflite`, do not apply image preprocessing. Only apply
|
||||
# normalization.
|
||||
if input_type == 'tflite':
|
||||
return normalize_image_fn(image_tensor), None
|
||||
|
||||
def preprocess_image_fn(inputs):
|
||||
image = normalize_image_fn(inputs)
|
||||
(image, image_info) = yolo_model_fn.letterbox(
|
||||
image,
|
||||
input_image_size,
|
||||
letter_box=params.task.validation_data.parser.letter_box,
|
||||
)
|
||||
return image, image_info
|
||||
|
||||
images_spec = tf.TensorSpec(shape=input_image_size + [3], dtype=tf.float32)
|
||||
|
||||
image_info_spec = tf.TensorSpec(shape=[4, 2], dtype=tf.float32)
|
||||
|
||||
images, image_info = tf.nest.map_structure(
|
||||
tf.identity,
|
||||
tf.map_fn(
|
||||
preprocess_image_fn,
|
||||
elems=image_tensor,
|
||||
fn_output_signature=(images_spec, image_info_spec),
|
||||
parallel_iterations=32,
|
||||
),
|
||||
)
|
||||
|
||||
return images, image_info
|
||||
|
||||
def inference_steps(inputs, model):
|
||||
images, image_info = inputs
|
||||
detection = model.call(images, training=False)
|
||||
if input_type != 'tflite':
|
||||
detection['bbox'] = yolo_model_fn.undo_info(
|
||||
detection['bbox'],
|
||||
detection['num_detections'],
|
||||
image_info,
|
||||
expand=False,
|
||||
)
|
||||
|
||||
final_outputs = {
|
||||
'detection_boxes': detection['bbox'],
|
||||
'detection_scores': detection['confidence'],
|
||||
'detection_classes': detection['classes'],
|
||||
'num_detections': detection['num_detections'],
|
||||
}
|
||||
|
||||
return final_outputs
|
||||
|
||||
export_module = YoloDetectionModule(
|
||||
params=params,
|
||||
model=model,
|
||||
batch_size=batch_size,
|
||||
input_image_size=input_image_size,
|
||||
input_type=input_type,
|
||||
num_channels=num_channels,
|
||||
input_name=input_name,
|
||||
label_map_path=label_map_path,
|
||||
preprocessor=preprocess_fn,
|
||||
inference_step=inference_steps,
|
||||
)
|
||||
|
||||
return export_module
|
||||
Executable
+293
@@ -0,0 +1,293 @@
|
||||
"""Export OSS TfVision models."""
|
||||
import os
|
||||
|
||||
from absl import app
|
||||
from absl import flags
|
||||
from absl import logging
|
||||
|
||||
from google.cloud import aiplatform as aip
|
||||
|
||||
# pylint: disable=line-too-long,unused-import
|
||||
from tfvision import registry_imports as vision_registry_imports
|
||||
from tfvision.serving import automl_constants
|
||||
from tfvision.serving import export_oss_saved_model_lib as export_automl_oss_saved_model_lib
|
||||
from util import constants
|
||||
from official.core import exp_factory
|
||||
from official.modeling import hyperparams
|
||||
from official.projects.maxvit import registry_imports as maxvit_imports
|
||||
from official.projects.yolo.common import registry_imports as yolo_imports
|
||||
from official.vision import registry_imports
|
||||
from official.vision.serving import export_saved_model_lib as export_oss_saved_model_lib
|
||||
# pylint: enable=line-too-long, unused-import
|
||||
|
||||
_PARAMS_OVERRIDE_IOD = """
|
||||
task:
|
||||
export_config:
|
||||
output_normalized_coordinates: true
|
||||
cast_num_detections_to_float: true
|
||||
cast_detection_classes_to_float: true
|
||||
model:
|
||||
detection_generator:
|
||||
nms_version: batched"""
|
||||
|
||||
_PARAMS_OVERRIDE_YOLO = """
|
||||
task:
|
||||
export_config:
|
||||
output_normalized_coordinates: true
|
||||
cast_num_detections_to_float: true
|
||||
cast_detection_classes_to_float: true
|
||||
model:
|
||||
detection_generator:
|
||||
nms_version: v2"""
|
||||
|
||||
|
||||
_PARAMS_OVERRIDE_ISG = """
|
||||
task:
|
||||
export_config:
|
||||
rescale_output: true"""
|
||||
|
||||
_YOLO_KEY = 'yolo'
|
||||
|
||||
_OBJECTIVE = flags.DEFINE_enum(
|
||||
'objective',
|
||||
None,
|
||||
[
|
||||
constants.OBJECTIVE_IMAGE_CLASSIFICATION,
|
||||
constants.OBJECTIVE_IMAGE_OBJECT_DETECTION,
|
||||
constants.OBJECTIVE_IMAGE_SEGMENTATION,
|
||||
],
|
||||
'The objective of this training job.',
|
||||
)
|
||||
|
||||
# Cloud AI platform HPT related parameter
|
||||
_PROJECT_NAME = flags.DEFINE_string(
|
||||
'project_name', None, 'Training vizier study name.'
|
||||
)
|
||||
_LOCATION = flags.DEFINE_string('location', None, 'Vizier study owner.')
|
||||
_HPT_JOB_ID = flags.DEFINE_string('hpt_job_id', None, 'HPT job id.')
|
||||
_HPT_RESULT_DIR = flags.DEFINE_string(
|
||||
'hpt_result_dir', None, 'HPT job result directory.'
|
||||
)
|
||||
_USE_BIGSTORE = flags.DEFINE_bool(
|
||||
'use_bigstore', None, 'Whether to use bigstore in hub model path.'
|
||||
)
|
||||
|
||||
# TfVision related inputs.
|
||||
_EXPERIMENT = flags.DEFINE_string(
|
||||
'experiment', None, 'experiment type, e.g. retinanet_resnetfpn_coco')
|
||||
_EXPORT_DIR = flags.DEFINE_string('export_dir', None, 'The export directory.')
|
||||
_CHECKPOINT_PATH = flags.DEFINE_string('checkpoint_path', None,
|
||||
'Checkpoint path.')
|
||||
_LABEL_MAP_PATH = flags.DEFINE_string('label_map_path', None,
|
||||
'Path to the labelmap proto file.')
|
||||
_LABEL_PATH = flags.DEFINE_string(
|
||||
'label_path', None, 'Path to the image classification label file.')
|
||||
_CONFIG_FILE = flags.DEFINE_multi_string(
|
||||
'config_file',
|
||||
default=None,
|
||||
help=(
|
||||
'YAML/JSON files which specifies overrides. The override order follows'
|
||||
' the order of args. Note that each file can be used as an override'
|
||||
' template to override the default parameters specified in Python. If'
|
||||
' the same parameter is specified in both `--config_file`.'
|
||||
),
|
||||
)
|
||||
_INPUT_IMAGE_SIZE = flags.DEFINE_string(
|
||||
'input_image_size', '224,224',
|
||||
'The comma-separated string of two integers representing the height,width '
|
||||
'of the input to the model.')
|
||||
|
||||
# Fixed inputs.
|
||||
_IMAGE_TYPE = flags.DEFINE_string(
|
||||
'input_type',
|
||||
'image_bytes',
|
||||
'One of `image_tensor`, `image_bytes`, `tf_example` and `tflite`.',
|
||||
)
|
||||
_EXPORT_SAVED_MODEL_SUBDIR = flags.DEFINE_string(
|
||||
'export_saved_model_subdir', 'saved_model',
|
||||
'The subdirectory for saved model.')
|
||||
_BATCH_SIZE = flags.DEFINE_integer('batch_size', 1, 'The batch size.')
|
||||
_INPUT_NAME = flags.DEFINE_string(
|
||||
'input_name',
|
||||
'encoded_image',
|
||||
(
|
||||
'Input tensor name in signature def. Default at None which'
|
||||
'produces input tensor name `inputs`.'
|
||||
),
|
||||
)
|
||||
_MAX_TRIAL_COUNT = flags.DEFINE_integer(
|
||||
'max_trial_count', None, 'The desired total number of trials.'
|
||||
)
|
||||
_EVALUATION_METRIC = flags.DEFINE_string(
|
||||
'evaluation_metric',
|
||||
None,
|
||||
'The evaluation metric to use (e.g. accuracy).',
|
||||
)
|
||||
|
||||
|
||||
def change_handle(params: hyperparams.ParamsDict) -> hyperparams.ParamsDict:
|
||||
"""Changes the prefix of the `handle` path in the `model.backbone.hub_model` sub-dictionary from gs:// to /bigstore/.
|
||||
|
||||
Args:
|
||||
params: hyperparams.ParamsDict object containing experiment config
|
||||
information.
|
||||
|
||||
Returns:
|
||||
params: hyperparams.ParamsDict.
|
||||
"""
|
||||
|
||||
params.task.model.backbone.hub_model.handle = (
|
||||
params.task.model.backbone.hub_model.handle.replace(
|
||||
'gs://', '/bigstore/', 1
|
||||
)
|
||||
)
|
||||
|
||||
return params
|
||||
|
||||
|
||||
def get_best_hpt_trials(
|
||||
project: str, location: str, hpt_job_id: str, hpt_result_dir: str
|
||||
) -> str:
|
||||
"""Select best trials by cloud ai platorm hyperparameter tuning.
|
||||
|
||||
Args:
|
||||
project: GCP project name.
|
||||
location: Hyperparameter job location.
|
||||
hpt_job_id: Hyperparameter job id.
|
||||
hpt_result_dir: HPT job result GCS directory.
|
||||
|
||||
Returns:
|
||||
Trial Id of the best performing trial.
|
||||
"""
|
||||
|
||||
aip.init(project=project, location=location)
|
||||
job_response = aip.HyperparameterTuningJob.get(resource_name=hpt_job_id)
|
||||
max_value = -1
|
||||
best_trial_id = -1
|
||||
trials = list(job_response._gca_resource.trials) # pylint: disable=protected-access
|
||||
for trial in trials:
|
||||
if trial.final_measurement.metrics[0].metric_id != constants.HP_METRIC_TAG:
|
||||
continue
|
||||
if trial.final_measurement.metrics[0].value > max_value:
|
||||
best_trial_id = trial.id
|
||||
max_value = trial.final_measurement.metrics[0].value
|
||||
if best_trial_id == -1:
|
||||
raise ValueError('No valid completed trials.')
|
||||
best_model_dir = os.path.join(
|
||||
hpt_result_dir, constants.TRIAL_PREFIX + str(best_trial_id)
|
||||
)
|
||||
logging.info(
|
||||
'Best model directory: %s with performance: %s', best_model_dir, max_value
|
||||
)
|
||||
return best_model_dir
|
||||
|
||||
|
||||
def main(_) -> None:
|
||||
if (
|
||||
_MAX_TRIAL_COUNT.present
|
||||
and _EVALUATION_METRIC.present
|
||||
and _CONFIG_FILE.present
|
||||
):
|
||||
best_ckpt_dir, _ = export_automl_oss_saved_model_lib.get_best_oss_trial(
|
||||
_CHECKPOINT_PATH.value, _MAX_TRIAL_COUNT.value, _EVALUATION_METRIC.value
|
||||
)
|
||||
config_filepath = _CONFIG_FILE.value
|
||||
elif _CHECKPOINT_PATH.present and _CONFIG_FILE.present:
|
||||
best_ckpt_dir = _CHECKPOINT_PATH.value
|
||||
config_filepath = _CONFIG_FILE.value
|
||||
elif (
|
||||
_PROJECT_NAME.present
|
||||
and _LOCATION.present
|
||||
and _HPT_JOB_ID.present
|
||||
and _HPT_RESULT_DIR.present
|
||||
):
|
||||
# Reads HPT results by project and location and hpt_job_id.
|
||||
best_ckpt_dir = get_best_hpt_trials(
|
||||
_PROJECT_NAME.value,
|
||||
_LOCATION.value,
|
||||
_HPT_JOB_ID.value,
|
||||
_HPT_RESULT_DIR.value,
|
||||
)
|
||||
config_filepath = [
|
||||
os.path.join(best_ckpt_dir, automl_constants.CFG_FILENAME)
|
||||
]
|
||||
else:
|
||||
raise ValueError('No checkpoint path or HTP Job parameters given.')
|
||||
|
||||
params = exp_factory.get_exp_config(_EXPERIMENT.value)
|
||||
for config_file in config_filepath or []:
|
||||
params = hyperparams.override_params_dict(
|
||||
params, config_file, is_strict=False
|
||||
)
|
||||
if _OBJECTIVE.value == constants.OBJECTIVE_IMAGE_OBJECT_DETECTION:
|
||||
if _YOLO_KEY in _EXPERIMENT.value:
|
||||
params = hyperparams.override_params_dict(
|
||||
params, _PARAMS_OVERRIDE_YOLO, is_strict=False
|
||||
)
|
||||
else:
|
||||
params = hyperparams.override_params_dict(
|
||||
params, _PARAMS_OVERRIDE_IOD, is_strict=False
|
||||
)
|
||||
elif _OBJECTIVE.value == constants.OBJECTIVE_IMAGE_SEGMENTATION:
|
||||
params = hyperparams.override_params_dict(
|
||||
params, _PARAMS_OVERRIDE_ISG, is_strict=True
|
||||
)
|
||||
|
||||
if _USE_BIGSTORE.value:
|
||||
params = change_handle(params)
|
||||
|
||||
params.validate()
|
||||
params.lock()
|
||||
|
||||
if best_ckpt_dir and not best_ckpt_dir.endswith(
|
||||
params.trainer.best_checkpoint_export_subdir
|
||||
):
|
||||
best_ckpt_dir = os.path.join(
|
||||
best_ckpt_dir, params.trainer.best_checkpoint_export_subdir
|
||||
)
|
||||
|
||||
if (
|
||||
_LABEL_MAP_PATH.value
|
||||
or _LABEL_PATH.value
|
||||
or _OBJECTIVE.value == constants.OBJECTIVE_IMAGE_SEGMENTATION
|
||||
):
|
||||
export_automl_oss_saved_model_lib.export_inference_graph(
|
||||
input_type=_IMAGE_TYPE.value,
|
||||
batch_size=_BATCH_SIZE.value,
|
||||
input_image_size=[int(x) for x in _INPUT_IMAGE_SIZE.value.split(',')],
|
||||
params=params,
|
||||
checkpoint_path=best_ckpt_dir,
|
||||
label_map_path=_LABEL_MAP_PATH.value,
|
||||
label_path=_LABEL_PATH.value,
|
||||
export_dir=_EXPORT_DIR.value,
|
||||
export_saved_model_subdir=_EXPORT_SAVED_MODEL_SUBDIR.value,
|
||||
input_name=_INPUT_NAME.value,
|
||||
objective=_OBJECTIVE.value,
|
||||
)
|
||||
elif _YOLO_KEY in _EXPERIMENT.value:
|
||||
export_automl_oss_saved_model_lib.export_inference_graph(
|
||||
input_type=_IMAGE_TYPE.value,
|
||||
batch_size=_BATCH_SIZE.value,
|
||||
input_image_size=[int(x) for x in _INPUT_IMAGE_SIZE.value.split(',')],
|
||||
params=params,
|
||||
checkpoint_path=best_ckpt_dir,
|
||||
export_dir=_EXPORT_DIR.value,
|
||||
export_saved_model_subdir=_EXPORT_SAVED_MODEL_SUBDIR.value,
|
||||
input_name=_INPUT_NAME.value,
|
||||
objective=_OBJECTIVE.value,
|
||||
)
|
||||
else:
|
||||
export_oss_saved_model_lib.export_inference_graph(
|
||||
input_type=_IMAGE_TYPE.value,
|
||||
batch_size=_BATCH_SIZE.value,
|
||||
input_image_size=[int(x) for x in _INPUT_IMAGE_SIZE.value.split(',')],
|
||||
params=params,
|
||||
checkpoint_path=best_ckpt_dir,
|
||||
export_dir=_EXPORT_DIR.value,
|
||||
export_saved_model_subdir=_EXPORT_SAVED_MODEL_SUBDIR.value,
|
||||
input_name=_INPUT_NAME.value,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
app.run(main)
|
||||
+169
@@ -0,0 +1,169 @@
|
||||
r"""Vision models export utility function for serving/inference."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import tensorflow as tf
|
||||
|
||||
from tfvision.serving import detection
|
||||
from tfvision.serving import image_classification
|
||||
from tfvision.serving import semantic_segmentation_export_module_lib as isg_export_lib
|
||||
from util import constants
|
||||
from official.core import config_definitions as cfg
|
||||
from official.core import export_base
|
||||
from official.projects.yolo.configs import yolo as yolo_config
|
||||
from official.projects.yolo.configs import yolov7 as yolov7_config
|
||||
from official.projects.yolo.serving import export_module_factory as yolo_export_module_factory
|
||||
|
||||
|
||||
def export_inference_graph(
|
||||
input_type: str,
|
||||
batch_size: Optional[int],
|
||||
input_image_size: List[int],
|
||||
params: cfg.ExperimentConfig,
|
||||
checkpoint_path: str,
|
||||
export_dir: str,
|
||||
label_map_path: Optional[str] = None,
|
||||
label_path: Optional[str] = None,
|
||||
num_channels: Optional[int] = 3,
|
||||
export_module: Optional[export_base.ExportModule] = None,
|
||||
export_saved_model_subdir: Optional[str] = None,
|
||||
save_options: Optional[tf.saved_model.SaveOptions] = None,
|
||||
checkpoint: Optional[tf.train.Checkpoint] = None,
|
||||
input_name: Optional[str] = None,
|
||||
function_keys: Optional[Union[List[str], Dict[str, str]]] = None,
|
||||
objective: Optional[str] = None,
|
||||
):
|
||||
"""Exports inference graph for the model specified in the exp config.
|
||||
|
||||
Saved model is stored at export_dir/saved_model, checkpoint is saved
|
||||
at export_dir/checkpoint, and params is saved at export_dir/params.yaml.
|
||||
|
||||
Args:
|
||||
input_type: Input type must be `image_bytes`.
|
||||
batch_size: 'int', or None.
|
||||
input_image_size: List or Tuple of height and width.
|
||||
params: Experiment params.
|
||||
checkpoint_path: Trained checkpoint path or directory.
|
||||
export_dir: CNS export directory path.
|
||||
label_map_path: Labelmap proto file path.
|
||||
label_path: Label file path.
|
||||
num_channels: The number of input image channels.
|
||||
export_module: Optional export module to be used instead of using params to
|
||||
create one. If None, the params will be used to create an export module.
|
||||
export_saved_model_subdir: Optional subdirectory under export_dir to store
|
||||
saved model.
|
||||
save_options: `SaveOptions` for `tf.saved_model.save`.
|
||||
checkpoint: An optional tf.train.Checkpoint. If provided, the export module
|
||||
will use it to read the weights.
|
||||
input_name: The input tensor name, default at `None` which produces input
|
||||
tensor name `inputs`.
|
||||
function_keys: a list of string keys to retrieve pre-defined serving
|
||||
signatures. The signaute keys will be set with defaults. If a dictionary
|
||||
is provided, the values will be used as signature keys.
|
||||
objective: The objective of the training job.
|
||||
"""
|
||||
if export_saved_model_subdir:
|
||||
output_saved_model_directory = os.path.join(export_dir,
|
||||
export_saved_model_subdir)
|
||||
else:
|
||||
output_saved_model_directory = export_dir
|
||||
|
||||
if not export_module:
|
||||
if objective == constants.OBJECTIVE_IMAGE_CLASSIFICATION:
|
||||
export_module = image_classification.ClassificationModule(
|
||||
params=params,
|
||||
batch_size=batch_size,
|
||||
input_image_size=input_image_size,
|
||||
input_type=input_type,
|
||||
num_channels=num_channels,
|
||||
input_name=input_name,
|
||||
label_path=label_path,
|
||||
)
|
||||
elif objective == constants.OBJECTIVE_IMAGE_OBJECT_DETECTION:
|
||||
# If experiment is YOLO object detection, loads Yolo detection module.
|
||||
if isinstance(
|
||||
params.task, (yolo_config.YoloTask, yolov7_config.YoloV7Task)
|
||||
):
|
||||
export_module = yolo_export_module_factory.get_export_module(
|
||||
params=params,
|
||||
batch_size=batch_size,
|
||||
input_image_size=input_image_size,
|
||||
input_type=input_type,
|
||||
num_channels=num_channels,
|
||||
input_name=input_name,
|
||||
)
|
||||
else:
|
||||
export_module = detection.DetectionModule(
|
||||
params=params,
|
||||
batch_size=batch_size,
|
||||
input_image_size=input_image_size,
|
||||
input_type=input_type,
|
||||
num_channels=num_channels,
|
||||
input_name=input_name,
|
||||
label_map_path=label_map_path,
|
||||
)
|
||||
elif objective == constants.OBJECTIVE_IMAGE_SEGMENTATION:
|
||||
export_module = isg_export_lib.OssSegmentationModule(
|
||||
params=params,
|
||||
batch_size=batch_size,
|
||||
input_image_size=input_image_size,
|
||||
input_type=input_type,
|
||||
num_channels=num_channels,
|
||||
input_name=input_name,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
'Export module not implemented for objective {}.'.format(objective)
|
||||
)
|
||||
|
||||
export_base.export(
|
||||
export_module,
|
||||
function_keys=function_keys if function_keys else [input_type],
|
||||
export_savedmodel_dir=output_saved_model_directory,
|
||||
checkpoint=checkpoint,
|
||||
checkpoint_path=checkpoint_path,
|
||||
timestamped=False,
|
||||
save_options=save_options)
|
||||
|
||||
|
||||
def get_best_oss_trial(
|
||||
model_dir: str, max_trial_count: int, evaluation_metric: str
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""Export models from TF checkpoints to TF saved model format.
|
||||
|
||||
Args:
|
||||
model_dir: Path of directory to store checkpoints and metric summaries.
|
||||
max_trial_count: The desired total number of trials.
|
||||
evaluation_metric: The evaluation metric to use (ie. accuracy).
|
||||
|
||||
Returns:
|
||||
"""
|
||||
best_trial_dir = ''
|
||||
best_trial_evaluation_results = {}
|
||||
best_performance = -1
|
||||
trial_file_count = 0
|
||||
for i in range(max_trial_count):
|
||||
current_trial = i + 1
|
||||
current_trial_dir = os.path.join(model_dir, 'trial_' + str(current_trial))
|
||||
current_trial_best_ckpt_dir = os.path.join(current_trial_dir, 'best_ckpt')
|
||||
current_trial_best_ckpt_evaluation_filepath = os.path.join(
|
||||
current_trial_best_ckpt_dir, 'info.json'
|
||||
)
|
||||
if tf.io.gfile.exists(current_trial_best_ckpt_evaluation_filepath):
|
||||
trial_file_count += 1
|
||||
with tf.io.gfile.GFile(
|
||||
current_trial_best_ckpt_evaluation_filepath, 'rb'
|
||||
) as f:
|
||||
eval_metric_results = json.load(f)
|
||||
current_performance = eval_metric_results[evaluation_metric]
|
||||
if current_performance > best_performance:
|
||||
best_performance = current_performance
|
||||
best_trial_dir = current_trial_dir
|
||||
best_trial_evaluation_results = eval_metric_results
|
||||
|
||||
if not trial_file_count:
|
||||
raise ValueError('None of the best checkpoint paths exist.')
|
||||
|
||||
return best_trial_dir, best_trial_evaluation_results
|
||||
+159
@@ -0,0 +1,159 @@
|
||||
"""Image classification input and model functions for serving/inference."""
|
||||
|
||||
from typing import Callable, List, Mapping, Optional
|
||||
|
||||
import tensorflow as tf
|
||||
from tensorflow.io import gfile
|
||||
|
||||
from tfvision.serving import automl_constants
|
||||
from official.core import config_definitions as cfg
|
||||
from official.vision.serving import image_classification
|
||||
|
||||
|
||||
class ClassificationModule(image_classification.ClassificationModule):
|
||||
"""classification Module."""
|
||||
|
||||
def __init__(self,
|
||||
params: cfg.ExperimentConfig,
|
||||
*,
|
||||
batch_size: Optional[int] = None,
|
||||
input_image_size: List[int],
|
||||
input_type: str = automl_constants.INPUT_TYPE,
|
||||
num_channels: int = 3,
|
||||
model: Optional[tf.keras.Model] = None,
|
||||
input_name: str = automl_constants.ICN_INPUT_NAME,
|
||||
label_path: Optional[str] = None,
|
||||
key_name: str = automl_constants.INPUT_KEY_NAME):
|
||||
"""Initializes a module for export.
|
||||
|
||||
Args:
|
||||
params: Experiment params.
|
||||
batch_size: The batch size of the model input. Can be `int` or None.
|
||||
input_image_size: List or Tuple of size of the input image. For 2D image,
|
||||
it is [height, width].
|
||||
input_type: The input signature type.
|
||||
num_channels: The number of the image channels.
|
||||
model: A tf.keras.Model instance to be exported.
|
||||
input_name: A customized input tensor name.
|
||||
label_path: A label file path.
|
||||
key_name: A name to the automl model input key.
|
||||
"""
|
||||
super().__init__(
|
||||
params=params,
|
||||
model=model,
|
||||
batch_size=batch_size,
|
||||
input_image_size=input_image_size,
|
||||
num_channels=num_channels,
|
||||
input_name=input_name,
|
||||
input_type=input_type,
|
||||
)
|
||||
|
||||
self._key_name = key_name
|
||||
if label_path is not None:
|
||||
self._label = self._read_label(label_path)
|
||||
else:
|
||||
self._label = None
|
||||
|
||||
def _read_label(self, label_path: str) -> tf.Tensor:
|
||||
"""Reads the labels from a label file."""
|
||||
with gfile.GFile(label_path, 'r') as f:
|
||||
labels = [i.strip() for i in f.readlines()]
|
||||
labels = tf.convert_to_tensor([labels])
|
||||
return labels
|
||||
|
||||
def serve(self, images: tf.Tensor, key: tf.Tensor) -> Mapping[str, tf.Tensor]:
|
||||
"""Cast image to float and run inference.
|
||||
|
||||
Args:
|
||||
images: uint8 Tensor of shape [batch_size, None, None, 3]
|
||||
key: string Tensor of shape [batch_size].
|
||||
|
||||
Returns:
|
||||
Dictionary holding classification outputs.
|
||||
"""
|
||||
with tf.device('cpu:0'):
|
||||
images = tf.cast(images, dtype=tf.float32)
|
||||
|
||||
images = tf.nest.map_structure(
|
||||
tf.identity,
|
||||
tf.map_fn(
|
||||
self._build_inputs,
|
||||
elems=images,
|
||||
fn_output_signature=tf.TensorSpec(
|
||||
shape=self._input_image_size + [3], dtype=tf.float32),
|
||||
parallel_iterations=32))
|
||||
|
||||
logits = self.inference_step(images)
|
||||
if self.params.task.train_data.is_multilabel:
|
||||
probs = tf.math.sigmoid(logits)
|
||||
else:
|
||||
probs = tf.nn.softmax(logits)
|
||||
|
||||
outputs = {'scores': probs, automl_constants.OUTPUT_KEY_NAME: key}
|
||||
if self._label is not None:
|
||||
outputs['labels'] = tf.tile(self._label, [tf.shape(images)[0], 1])
|
||||
return outputs
|
||||
|
||||
@tf.function
|
||||
def inference_from_image_bytes(self, inputs: tf.Tensor,
|
||||
key: tf.Tensor) -> Mapping[str, tf.Tensor]:
|
||||
with tf.device('cpu:0'):
|
||||
images = tf.nest.map_structure(
|
||||
tf.identity,
|
||||
tf.map_fn(
|
||||
self._decode_image,
|
||||
elems=inputs,
|
||||
fn_output_signature=tf.TensorSpec(
|
||||
shape=[None] * len(self._input_image_size) +
|
||||
[self._num_channels],
|
||||
dtype=tf.uint8),
|
||||
parallel_iterations=32))
|
||||
images = tf.stack(images)
|
||||
return self.serve(images, key)
|
||||
|
||||
@tf.function
|
||||
def inference_from_image_tensors(
|
||||
self, inputs: tf.Tensor
|
||||
) -> Mapping[str, tf.Tensor]:
|
||||
return self.serve(inputs, tf.zeros(tf.shape(inputs)[0], dtype=tf.string))
|
||||
|
||||
def get_inference_signatures(
|
||||
self, function_keys: Mapping[str, str]
|
||||
) -> Mapping[str, Callable[[tf.Tensor, tf.Tensor], Mapping[str, tf.Tensor]]]:
|
||||
"""Gets defined function signatures.
|
||||
|
||||
Args:
|
||||
function_keys: A dictionary with keys as the function to create signature
|
||||
for and values as the signature keys when returns.
|
||||
|
||||
Returns:
|
||||
A dictionary with key as signature key and value as concrete functions
|
||||
that can be used for tf.saved_model.save.
|
||||
"""
|
||||
signatures = {}
|
||||
for key, def_name in function_keys.items():
|
||||
# Adds input string 'key' to image_bytes input type.
|
||||
if key == automl_constants.INPUT_TYPE:
|
||||
input_images = tf.TensorSpec(
|
||||
shape=[self._batch_size], dtype=tf.string, name=self._input_name)
|
||||
input_key = tf.TensorSpec(
|
||||
shape=[self._batch_size], dtype=tf.string, name=self._key_name)
|
||||
signatures[
|
||||
def_name] = self.inference_from_image_bytes.get_concrete_function(
|
||||
input_images, input_key)
|
||||
elif key == automl_constants.IMAGE_TENSOR:
|
||||
input_signature = tf.TensorSpec(
|
||||
shape=[self._batch_size]
|
||||
+ [None] * len(self._input_image_size)
|
||||
+ [self._num_channels],
|
||||
dtype=tf.uint8,
|
||||
name=self._input_name,
|
||||
)
|
||||
signatures[def_name] = (
|
||||
self.inference_from_image_tensors.get_concrete_function(
|
||||
input_signature
|
||||
)
|
||||
)
|
||||
else:
|
||||
raise ValueError('Unrecognized `input_type`')
|
||||
return signatures
|
||||
+52
@@ -0,0 +1,52 @@
|
||||
"""Semantic segmentation input and model functions for serving/inference."""
|
||||
|
||||
|
||||
import tensorflow as tf
|
||||
|
||||
from official.vision.serving import semantic_segmentation
|
||||
|
||||
|
||||
class OssSegmentationModule(semantic_segmentation.SegmentationModule):
|
||||
"""OSS Segmentation Module."""
|
||||
|
||||
def serve(self, images):
|
||||
"""Cast image to float and run inference.
|
||||
|
||||
Overrides the method in the super class, and changes the output format.
|
||||
|
||||
Args:
|
||||
images: uint8 Tensor of shape [batch_size, None, None, 3]
|
||||
|
||||
Returns:
|
||||
Dict containing the following key value pairs:
|
||||
category_bytes: Encoded PNG image of grayscale output categories.
|
||||
score_bytes: Encoded PNG image of grayscale probability scores mapped to
|
||||
[0, 255].
|
||||
"""
|
||||
result = super().serve(images)
|
||||
logits = result['logits']
|
||||
|
||||
probabilities = tf.nn.softmax(logits)
|
||||
scores = tf.reduce_max(probabilities, 3, keepdims=True)
|
||||
scores = tf.cast(tf.minimum(scores * 255.0, 255), dtype=tf.uint8)
|
||||
|
||||
categories = tf.cast(
|
||||
tf.expand_dims(tf.argmax(logits, 3), -1), dtype=tf.int32
|
||||
)
|
||||
|
||||
score_bytes = tf.map_fn(
|
||||
tf.image.encode_png, scores, back_prop=False, dtype=tf.string
|
||||
)
|
||||
category_bytes = tf.map_fn(
|
||||
tf.image.encode_png,
|
||||
tf.cast(categories, dtype=tf.uint8),
|
||||
back_prop=False,
|
||||
dtype=tf.string,
|
||||
)
|
||||
|
||||
outputs = {
|
||||
'category_bytes': tf.identity(category_bytes, name='category_bytes'),
|
||||
'score_bytes': tf.identity(score_bytes, name='score_bytes'),
|
||||
}
|
||||
|
||||
return outputs
|
||||
@@ -0,0 +1,385 @@
|
||||
"""TensorFlow Model Garden Vision training driver.
|
||||
|
||||
This is the main function to start OSS vision training dockers, and will run in
|
||||
external environment.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from absl import app
|
||||
from absl import flags
|
||||
from absl import logging
|
||||
import gin
|
||||
import hypertune
|
||||
import tensorflow as tf
|
||||
|
||||
from util import constants
|
||||
from util import hypertune_utils
|
||||
from official.common import distribute_utils
|
||||
from official.common import flags as tfm_flags
|
||||
from official.core import task_factory
|
||||
from official.core import train_lib
|
||||
from official.core import train_utils
|
||||
from official.modeling import performance
|
||||
# pylint: disable=unused-import
|
||||
from tfvision import registry_imports as vision_registry_imports
|
||||
from official.projects.yolo.common import registry_imports as yolo_imports
|
||||
from official.vision import registry_imports
|
||||
|
||||
if os.environ.get('ENABLE_MAX_VIT', ''):
|
||||
# pylint: disable=g-import-not-at-top
|
||||
# pylint: disable=import-error
|
||||
# pylint: disable=no-name-in-module
|
||||
from official.projects.maxvit import registry_imports as maxvit_imports
|
||||
# pylint: enable=unused-import
|
||||
|
||||
# File type tfrecord.
|
||||
_FILE_TYPE_TFRECORD = 'tfrecord'
|
||||
|
||||
|
||||
_OBJECTIVE = flags.DEFINE_enum(
|
||||
'objective',
|
||||
constants.OBJECTIVE_IMAGE_CLASSIFICATION,
|
||||
[
|
||||
constants.OBJECTIVE_IMAGE_CLASSIFICATION,
|
||||
constants.OBJECTIVE_IMAGE_OBJECT_DETECTION,
|
||||
constants.OBJECTIVE_IMAGE_SEGMENTATION,
|
||||
],
|
||||
'The objective of this training job.',
|
||||
)
|
||||
|
||||
_MODEL_NAME = flags.DEFINE_string(
|
||||
'model_name',
|
||||
None,
|
||||
(
|
||||
'The model name for backbones. e.g.: the model names can be `vit-ti16`,'
|
||||
'`vit-b16`, `vit-s16`, `vit-l16`, for `deit_imagenet_pretrain`.'
|
||||
),
|
||||
)
|
||||
|
||||
_INIT_CHECKPOINT = flags.DEFINE_string(
|
||||
'init_checkpoint', None, 'The initial checkpoint of this training job.'
|
||||
)
|
||||
|
||||
_BACKBONE_TRAINABLE = flags.DEFINE_bool(
|
||||
'backbone_trainable', None, 'Whether to train the backbone.'
|
||||
)
|
||||
|
||||
_LEARNING_RATE = flags.DEFINE_float(
|
||||
'learning_rate', None, 'The learning rate of this training job.'
|
||||
)
|
||||
|
||||
_WEIGHT_DECAY = flags.DEFINE_float(
|
||||
'weight_decay', None, 'The weight decay of this training job.'
|
||||
)
|
||||
|
||||
_NUM_CLASSES = flags.DEFINE_integer(
|
||||
'num_classes', None, 'The number of classes.'
|
||||
)
|
||||
|
||||
_INPUT_SIZE = flags.DEFINE_list(
|
||||
'input_size', None, 'Expected width and height of the input image.'
|
||||
)
|
||||
|
||||
_INPUT_TRAIN_DATA_PATH = flags.DEFINE_string(
|
||||
'input_train_data_path', None, 'Input train data path.'
|
||||
)
|
||||
|
||||
_INPUT_VALIDATION_DATA_PATH = flags.DEFINE_string(
|
||||
'input_validation_data_path', None, 'Input validation data path.'
|
||||
)
|
||||
|
||||
_GLOBAL_BATCH_SIZE = flags.DEFINE_integer(
|
||||
'global_batch_size', None, 'Global batch size.'
|
||||
)
|
||||
|
||||
_PREFETCH_BUFFER_SIZE = flags.DEFINE_integer(
|
||||
'prefetch_buffer_size', None, 'Prefetch buffer size.'
|
||||
)
|
||||
|
||||
_TRAIN_STEPS = flags.DEFINE_integer('train_steps', None, 'Train steps.')
|
||||
|
||||
|
||||
_ANCHOR_SIZE = flags.DEFINE_integer(
|
||||
'anchor_size', None, 'IOD model anchor size.'
|
||||
)
|
||||
|
||||
_OUTPUT_SIZE = flags.DEFINE_list(
|
||||
'output_size',
|
||||
None,
|
||||
'Expected width and height of the output image for ISG models.',
|
||||
)
|
||||
|
||||
_MAX_EVAL_WAIT_TIME = flags.DEFINE_integer(
|
||||
'max_eval_wait_time',
|
||||
0,
|
||||
(
|
||||
'Maximum duration to wait for evaluation result file after finishing'
|
||||
' the training job in seconds. Defaults to 0, immediately looking for'
|
||||
' the evaluation file.'
|
||||
),
|
||||
)
|
||||
|
||||
_LOG_LEVEL = flags.DEFINE_string('log_level', 'INFO', 'Log level.')
|
||||
|
||||
FLAGS = flags.FLAGS
|
||||
|
||||
|
||||
def get_best_eval_metric(objective: str, params: Any) -> str:
|
||||
"""Gets best eval metric.
|
||||
|
||||
Args:
|
||||
objective: The objective of this training job.
|
||||
params: Experiment config.
|
||||
|
||||
Returns:
|
||||
Eval metric to use.
|
||||
|
||||
Raises:
|
||||
ValueError: If params does not have best_checkpoint_eval_metric set and the
|
||||
objective is not valid.
|
||||
"""
|
||||
try:
|
||||
eval_metric_name = params.trainer.best_checkpoint_eval_metric
|
||||
except AttributeError:
|
||||
eval_metric_name = None
|
||||
|
||||
if not eval_metric_name:
|
||||
# If eval metric is not given in params, use the default value.
|
||||
if objective == constants.OBJECTIVE_IMAGE_CLASSIFICATION:
|
||||
try:
|
||||
is_multilabel = params.task.train_data.is_multilabel
|
||||
except AttributeError:
|
||||
# Set default.
|
||||
is_multilabel = False
|
||||
if is_multilabel:
|
||||
eval_metric_name = (
|
||||
constants.IMAGE_CLASSIFICATION_MULTI_LABEL_BEST_EVAL_METRIC
|
||||
)
|
||||
else:
|
||||
eval_metric_name = (
|
||||
constants.IMAGE_CLASSIFICATION_SINGLE_LABEL_BEST_EVAL_METRIC
|
||||
)
|
||||
elif objective == constants.OBJECTIVE_IMAGE_OBJECT_DETECTION:
|
||||
eval_metric_name = constants.IMAGE_OBJECT_DETECTION_BEST_EVAL_METRIC
|
||||
elif objective == constants.OBJECTIVE_IMAGE_SEGMENTATION:
|
||||
eval_metric_name = constants.IMAGE_SEGMENTATION_BEST_EVAL_METRIC
|
||||
else:
|
||||
raise ValueError(
|
||||
'The objective must be {}, {}, or {}.'.format(
|
||||
constants.OBJECTIVE_IMAGE_CLASSIFICATION,
|
||||
constants.OBJECTIVE_IMAGE_OBJECT_DETECTION,
|
||||
constants.OBJECTIVE_IMAGE_SEGMENTATION,
|
||||
)
|
||||
)
|
||||
return eval_metric_name
|
||||
|
||||
|
||||
def parse_params() -> Any:
|
||||
"""Parses parameters."""
|
||||
gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)
|
||||
params = train_utils.parse_configuration(FLAGS, lock_return=False)
|
||||
if _INIT_CHECKPOINT.value:
|
||||
params.task.init_checkpoint = _INIT_CHECKPOINT.value
|
||||
if 'yolov7' in FLAGS.experiment:
|
||||
params.task.init_checkpoint_modules = ['backbone', 'decoder']
|
||||
else:
|
||||
params.task.init_checkpoint_modules = 'backbone'
|
||||
if _MODEL_NAME.value:
|
||||
if FLAGS.experiment in [
|
||||
'deit_imagenet_pretrain',
|
||||
'vit_imagenet_pretrain',
|
||||
'vit_imagenet_finetune',
|
||||
]:
|
||||
params.task.model.backbone.vit.model_name = _MODEL_NAME.value
|
||||
if _NUM_CLASSES.value:
|
||||
params.task.model.num_classes = _NUM_CLASSES.value
|
||||
if _INPUT_SIZE.value:
|
||||
input_size = [int(elem) for elem in _INPUT_SIZE.value]
|
||||
if len(input_size) != 2:
|
||||
raise ValueError('The input size must contain 2 integers.')
|
||||
if input_size[0] < 0 or input_size[1] < 0:
|
||||
raise ValueError('The input size must be positive.')
|
||||
params.task.model.input_size = [input_size[0], input_size[1], 3]
|
||||
# If users set input train/validation data path, we assume the data are
|
||||
# converted from data converter as tfrecord. Users can use tfds by writing
|
||||
# their own config directly, and no need to override this parameter.
|
||||
if _INPUT_TRAIN_DATA_PATH.value:
|
||||
params.task.train_data.input_path = _INPUT_TRAIN_DATA_PATH.value
|
||||
params.task.train_data.file_type = _FILE_TYPE_TFRECORD
|
||||
params.task.train_data.tfds_name = ''
|
||||
if _INPUT_VALIDATION_DATA_PATH.value:
|
||||
params.task.validation_data.input_path = _INPUT_VALIDATION_DATA_PATH.value
|
||||
params.task.validation_data.file_type = _FILE_TYPE_TFRECORD
|
||||
params.task.validation_data.tfds_name = ''
|
||||
if _GLOBAL_BATCH_SIZE.value:
|
||||
params.task.train_data.global_batch_size = _GLOBAL_BATCH_SIZE.value
|
||||
params.task.validation_data.global_batch_size = _GLOBAL_BATCH_SIZE.value
|
||||
if _PREFETCH_BUFFER_SIZE.value:
|
||||
params.task.train_data.prefetch_buffer_size = _PREFETCH_BUFFER_SIZE.value
|
||||
params.task.validation_data.prefetch_buffer_size = (
|
||||
_PREFETCH_BUFFER_SIZE.value
|
||||
)
|
||||
|
||||
# Use `get` method of train_utils.hyperparams.OneOfConfig to get learning
|
||||
# rate config.
|
||||
learning_rate = params.trainer.optimizer_config.learning_rate.get()
|
||||
|
||||
if _TRAIN_STEPS.value:
|
||||
params.trainer.train_steps = _TRAIN_STEPS.value
|
||||
if hasattr(learning_rate, 'decay_steps'):
|
||||
learning_rate.decay_steps = _TRAIN_STEPS.value
|
||||
if (
|
||||
_BACKBONE_TRAINABLE.value is not None
|
||||
and params.task.model.backbone.type == 'hub_model'
|
||||
):
|
||||
params.task.model.backbone.hub_model.trainable = _BACKBONE_TRAINABLE.value
|
||||
if _LEARNING_RATE.value:
|
||||
logging.info('Updating learning_rate: %s', _LEARNING_RATE.value)
|
||||
if hasattr(learning_rate, 'initial_learning_rate'):
|
||||
learning_rate.initial_learning_rate = _LEARNING_RATE.value
|
||||
|
||||
if _WEIGHT_DECAY.value and 'yolo' in FLAGS.experiment:
|
||||
if 'sgd_torch' == params.trainer.optimizer_config.optimizer.type:
|
||||
params.trainer.optimizer_config.optimizer.sgd_torch.weight_decay = (
|
||||
_WEIGHT_DECAY.value
|
||||
)
|
||||
elif 'adamw' == params.trainer.optimizer_config.optimizer.type:
|
||||
params.trainer.optimizer_config.optimizer.adamw.weight_decay_rate = (
|
||||
_WEIGHT_DECAY.value
|
||||
)
|
||||
|
||||
# Yolo models does not support anchor size.
|
||||
if _ANCHOR_SIZE.value and 'yolo' not in FLAGS.experiment:
|
||||
params.task.model.anchor.anchor_size = _ANCHOR_SIZE.value
|
||||
|
||||
# Segmentation models will also set output size.
|
||||
if _OUTPUT_SIZE.value:
|
||||
output_size = [int(elem) for elem in _OUTPUT_SIZE.value]
|
||||
if len(output_size) != 2:
|
||||
raise ValueError('The output size must contain 2 integers.')
|
||||
if output_size[0] < 0 or output_size[1] < 0:
|
||||
raise ValueError('The output size must be positive.')
|
||||
params.task.train_data.output_size = output_size
|
||||
params.task.validation_data.output_size = output_size
|
||||
|
||||
# Set default params for best checkpoints.
|
||||
params.trainer.best_checkpoint_export_subdir = constants.BEST_CKPT_DIRNAME
|
||||
params.trainer.best_checkpoint_metric_comp = constants.BEST_CKPT_METRIC_COMP
|
||||
params.trainer.best_checkpoint_eval_metric = get_best_eval_metric(
|
||||
_OBJECTIVE.value, params
|
||||
)
|
||||
return params
|
||||
|
||||
|
||||
def wait_for_evaluation_file(
|
||||
eval_filepath: str,
|
||||
max_eval_wait_time: int,
|
||||
eval_wait_interval: int = 30,
|
||||
) -> None:
|
||||
"""Waits for the evaluation file to be created.
|
||||
|
||||
Args:
|
||||
eval_filepath: The path to the evaluation file.
|
||||
max_eval_wait_time: The maximum amount of time to wait for the evaluation
|
||||
file to be created, in seconds.
|
||||
eval_wait_interval: The interval at which to check for the existence of the
|
||||
evaluation file, in seconds. Defaults to 30 seconds.
|
||||
|
||||
Raises:
|
||||
ValueError: If the evaluation file does not exist after the maximum amount
|
||||
of time has passed.
|
||||
"""
|
||||
eval_wait_start_time = time.time()
|
||||
while not tf.io.gfile.exists(eval_filepath):
|
||||
if time.time() - eval_wait_start_time >= max_eval_wait_time:
|
||||
raise ValueError('The eval file {} does not exist.'.format(eval_filepath))
|
||||
time.sleep(eval_wait_interval)
|
||||
return
|
||||
|
||||
|
||||
def main(_):
|
||||
log_level = _LOG_LEVEL.value
|
||||
if log_level and log_level in ['FATAL', 'ERROR', 'WARNING', 'INFO', 'DEBUG']:
|
||||
logging.set_verbosity(log_level)
|
||||
params = parse_params()
|
||||
logging.info('The actual training parameters are:\n%s', params.as_dict())
|
||||
model_dir = os.path.join(
|
||||
FLAGS.model_dir,
|
||||
'trial_' + hypertune_utils.get_trial_id_from_environment(),
|
||||
)
|
||||
logging.info('model_dir in this trial is: %s', model_dir)
|
||||
if 'train' in FLAGS.mode:
|
||||
# Pure eval modes do not output yaml files. Otherwise continuous eval job
|
||||
# may race against the train job for writing the same file.
|
||||
train_utils.serialize_config(params, model_dir)
|
||||
|
||||
# Sets mixed_precision policy. Using 'mixed_float16' or 'mixed_bfloat16'
|
||||
# can have significant impact on model speeds by utilizing float16 in case of
|
||||
# GPUs, and bfloat16 in the case of TPUs. loss_scale takes effect only when
|
||||
# dtype is float16
|
||||
if params.runtime.mixed_precision_dtype:
|
||||
performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype)
|
||||
distribution_strategy = distribute_utils.get_distribution_strategy(
|
||||
distribution_strategy=params.runtime.distribution_strategy,
|
||||
all_reduce_alg=params.runtime.all_reduce_alg,
|
||||
num_gpus=params.runtime.num_gpus,
|
||||
tpu_address=params.runtime.tpu,
|
||||
)
|
||||
with distribution_strategy.scope():
|
||||
task = task_factory.get_task(params.task, logging_dir=model_dir)
|
||||
|
||||
train_lib.run_experiment(
|
||||
distribution_strategy=distribution_strategy,
|
||||
task=task,
|
||||
mode=FLAGS.mode,
|
||||
params=params,
|
||||
model_dir=model_dir,
|
||||
)
|
||||
|
||||
train_utils.save_gin_config(FLAGS.mode, model_dir)
|
||||
|
||||
eval_metric_name = get_best_eval_metric(_OBJECTIVE.value, params)
|
||||
|
||||
eval_filepath = os.path.join(
|
||||
model_dir, constants.BEST_CKPT_DIRNAME, constants.BEST_CKPT_EVAL_FILENAME
|
||||
)
|
||||
logging.info('Load eval metrics from: %s.', eval_filepath)
|
||||
wait_for_evaluation_file(eval_filepath, _MAX_EVAL_WAIT_TIME.value)
|
||||
|
||||
with tf.io.gfile.GFile(eval_filepath, 'rb') as f:
|
||||
eval_metric_results = json.load(f)
|
||||
logging.info('eval metrics are: %s.', eval_metric_results)
|
||||
if (
|
||||
eval_metric_name in eval_metric_results
|
||||
and constants.BEST_CKPT_STEP_NAME in eval_metric_results
|
||||
):
|
||||
hp_metric = eval_metric_results[eval_metric_name]
|
||||
hp_step = int(eval_metric_results[constants.BEST_CKPT_STEP_NAME])
|
||||
hpt = hypertune.HyperTune()
|
||||
hpt.report_hyperparameter_tuning_metric(
|
||||
hyperparameter_metric_tag=constants.HP_METRIC_TAG,
|
||||
metric_value=hp_metric,
|
||||
global_step=hp_step,
|
||||
)
|
||||
logging.info(
|
||||
'Send HP metric: %f and steps %d to hyperparameter tuning.',
|
||||
hp_metric,
|
||||
hp_step,
|
||||
)
|
||||
else:
|
||||
logging.info(
|
||||
'Either %s or %s is not included in the evaluation results: %s.',
|
||||
eval_metric_name,
|
||||
constants.BEST_CKPT_STEP_NAME,
|
||||
eval_metric_results,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
tfm_flags.define_flags()
|
||||
flags.mark_flags_as_required(['experiment', 'mode', 'model_dir'])
|
||||
app.run(main)
|
||||
+59
@@ -0,0 +1,59 @@
|
||||
# Dockerfile for serving dockers with timm.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/timm/dockerfile/serve.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
FROM pytorch/torchserve:0.7.0-gpu
|
||||
|
||||
USER root
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
wget \
|
||||
vim
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
ENV infer_port=7080
|
||||
ENV mng_port=7081
|
||||
ENV model_name="timm_serving"
|
||||
|
||||
# Install timm.
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
RUN python3 -m pip install timm==0.6.12
|
||||
RUN python3 -m pip install google-cloud-storage==2.9.0
|
||||
|
||||
# Copy model artifacts.
|
||||
COPY model_oss/timm/handler.py /home/model-server/handler.py
|
||||
|
||||
# Create torchserve configuration file.
|
||||
RUN echo \
|
||||
"default_response_timeout=1200\n" \
|
||||
"service_envelope=json\n" \
|
||||
"inference_address=http://0.0.0.0:${infer_port}\n" \
|
||||
"management_address=http://0.0.0.0:${mng_port}" >> /home/model-server/config.properties
|
||||
|
||||
# Expose ports.
|
||||
EXPOSE ${infer_port}
|
||||
EXPOSE ${mng_port}
|
||||
|
||||
# Archive eager mode model artifacts and dependencies.
|
||||
# Do not set --model-file and --serialized-file because model and checkpoint will be dynamically loaded in handler.py.
|
||||
RUN torch-model-archiver \
|
||||
--model-name=${model_name} \
|
||||
--version=1.0 \
|
||||
--handler=/home/model-server/handler.py \
|
||||
--runtime=python3 \
|
||||
--export-path=/home/model-server/model-store \
|
||||
--archive-format=default \
|
||||
--force
|
||||
|
||||
# Run Torchserve HTTP serve to respond to prediction requests.
|
||||
CMD ["torchserve", "--start", \
|
||||
"--ts-config", "/home/model-server/config.properties", \
|
||||
"--models", "${model_name}=${model_name}.mar", \
|
||||
"--model-store", "/home/model-server/model-store"]
|
||||
+47
@@ -0,0 +1,47 @@
|
||||
# Dockerfile for basic training dockers with timm.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/timm/dockerfile/train.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
# Base on pytorch-cuda image.
|
||||
FROM pytorch/pytorch:1.13.0-cuda11.6-cudnn8-runtime
|
||||
|
||||
# Install tools.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
wget \
|
||||
vim
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Download timm source code with pinned version.
|
||||
RUN wget -q https://github.com/rwightman/pytorch-image-models/archive/refs/tags/v0.6.12.tar.gz
|
||||
RUN tar xzf v0.6.12.tar.gz
|
||||
|
||||
# Install libraries.
|
||||
RUN pip install cloudml-hypertune==0.1.0.dev6
|
||||
|
||||
# Switch to timm repo.
|
||||
WORKDIR /workspace/pytorch-image-models-0.6.12
|
||||
|
||||
# NOTE: use 'sed' to modify the timm source code to
|
||||
# make timm CheckpointSaver can work with gcsfuse.
|
||||
RUN sed -i "1 i\import shutil" timm/utils/checkpoint_saver.py
|
||||
RUN sed -i "s#os.link#shutil.copyfile#g" timm/utils/checkpoint_saver.py
|
||||
RUN sed -i "s#os.unlink#os.remove#g" timm/utils/checkpoint_saver.py
|
||||
|
||||
# NOTE: use 'sed' to modify the timm source code to
|
||||
# add hp training support to timm trainer.
|
||||
RUN sed -i "693 a\ if saver is not None: hpt = hypertune.HyperTune(); hpt.report_hyperparameter_tuning_metric(hyperparameter_metric_tag='top1_accuracy', metric_value=best_metric, global_step=best_epoch)" train.py
|
||||
RUN sed -i "1 i\import hypertune" train.py
|
||||
|
||||
# Install timm from source code.
|
||||
RUN pip install -e .
|
||||
|
||||
# https://pytorch.org/docs/stable/elastic/run.html
|
||||
ENTRYPOINT ["torchrun"]
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Custom handler for TIMM models."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from google.cloud import storage
|
||||
import timm
|
||||
import torch
|
||||
from ts.torch_handler.base_handler import load_label_mapping
|
||||
from ts.torch_handler.image_classifier import ImageClassifier
|
||||
|
||||
|
||||
GCS_PREFIX = "gs://"
|
||||
DOWNLOAD_DIR = "/tmp/download"
|
||||
|
||||
|
||||
def download_gcs_file(gcs_uri: str, local_dir: str) -> str:
|
||||
"""Download a GCS file to a local directory.
|
||||
|
||||
Arguments:
|
||||
gcs_uri: A string of file path on GCS.
|
||||
local_dir: A string of local directory path.
|
||||
|
||||
Returns:
|
||||
Local path to downloaded file.
|
||||
"""
|
||||
if not gcs_uri.startswith(GCS_PREFIX):
|
||||
raise ValueError(f"{gcs_uri} is not a GCS path starting with gs://.")
|
||||
|
||||
file_name = os.path.basename(gcs_uri)
|
||||
local_file_path = os.path.join(local_dir, file_name)
|
||||
os.makedirs(local_dir, exist_ok=True)
|
||||
client = storage.Client()
|
||||
with open(local_file_path, "wb") as f:
|
||||
client.download_blob_to_file(gcs_uri, f)
|
||||
return local_file_path
|
||||
|
||||
|
||||
class TimmHandler(ImageClassifier):
|
||||
"""Custom handler for TIMM models."""
|
||||
|
||||
def initialize(self, context: Any):
|
||||
"""Custom initialize."""
|
||||
|
||||
properties = context.system_properties
|
||||
self.map_location = (
|
||||
"cuda"
|
||||
if torch.cuda.is_available() and properties.get("gpu_id") is not None
|
||||
else "cpu"
|
||||
)
|
||||
self.device = torch.device(
|
||||
self.map_location + ":" + str(properties.get("gpu_id"))
|
||||
if torch.cuda.is_available() and properties.get("gpu_id") is not None
|
||||
else self.map_location
|
||||
)
|
||||
self.manifest = context.manifest
|
||||
|
||||
# Load timm model by model name.
|
||||
self.model_name = os.environ["MODEL_NAME"]
|
||||
# Whether to use timm pretrained weights, MODEL_PT_PATH overrides this.
|
||||
timm_pretrained = True if os.environ.get("TIMM_PRETRAINED") else False
|
||||
# Load custom checkpoint, it overrides TIMM_PRETRAINED model.
|
||||
self.model_pt_path = os.environ.get("MODEL_PT_PATH")
|
||||
if self.model_pt_path and self.model_pt_path.startswith(GCS_PREFIX):
|
||||
self.model_pt_path = download_gcs_file(self.model_pt_path, DOWNLOAD_DIR)
|
||||
|
||||
if self.model_pt_path and self.model_pt_path.endswith(".pt"):
|
||||
logging.info(
|
||||
"Load model with .pt in jit mode, not working for all timm models"
|
||||
" yet."
|
||||
)
|
||||
self.model = self._load_torchscript_model(self.model_pt_path)
|
||||
else:
|
||||
logging.info("Load model with .pth in eager mode.")
|
||||
self.model = timm.create_model(
|
||||
self.model_name, pretrained=timm_pretrained
|
||||
)
|
||||
if self.model_pt_path and (
|
||||
self.model_pt_path.endswith(".pth")
|
||||
or self.model_pt_path.endswith(".pth.tar")
|
||||
):
|
||||
checkpoint = torch.load(self.model_pt_path, map_location=self.device)
|
||||
state_dict = checkpoint["state_dict"]
|
||||
self.model.load_state_dict(state_dict)
|
||||
self.model.to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
mapping_file_path = os.environ.get("INDEX_TO_NAME_FILE")
|
||||
if mapping_file_path:
|
||||
if mapping_file_path.startswith(GCS_PREFIX):
|
||||
mapping_file_path = download_gcs_file(mapping_file_path, DOWNLOAD_DIR)
|
||||
self.mapping = load_label_mapping(mapping_file_path)
|
||||
|
||||
self.initialized = True
|
||||
|
||||
# NOTE: Preprocess and postprocess are implemented by ImageClassifier.
|
||||
@@ -0,0 +1,79 @@
|
||||
"""Common utility lib for prediction on images."""
|
||||
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import tensorflow as tf
|
||||
import yaml
|
||||
|
||||
from util import image_format_converter
|
||||
|
||||
|
||||
def get_prediction_instances(image: Image.Image) -> List[Dict[str, Any]]:
|
||||
"""Gets prediction instances.
|
||||
|
||||
Args:
|
||||
image: Image instance.
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: List of prediction instances.
|
||||
"""
|
||||
instances = [{
|
||||
"encoded_image": {"b64": image_format_converter.image_to_base64(image)},
|
||||
}]
|
||||
return instances
|
||||
|
||||
|
||||
def get_label_map(label_map_yaml_filepath: str) -> Dict[str, Any]:
|
||||
"""Gets the label map from a YAML file.
|
||||
|
||||
Args:
|
||||
label_map_yaml_filepath: Filepath to the label map YAML file.
|
||||
|
||||
Returns:
|
||||
dict: Label map.
|
||||
"""
|
||||
with tf.io.gfile.GFile(label_map_yaml_filepath, "rb") as input_file:
|
||||
label_map = yaml.safe_load(input_file.read())
|
||||
return label_map
|
||||
|
||||
|
||||
def get_object_detection_endpoint_predictions(
|
||||
detection_endpoint: ...,
|
||||
input_image: np.ndarray,
|
||||
detection_thresh: float = 0.2,
|
||||
) -> np.ndarray:
|
||||
"""Gets endpoint predictions.
|
||||
|
||||
Args:
|
||||
detection_endpoint: image object detection endpoint.
|
||||
input_image: Input image.
|
||||
detection_thresh: Detection threshold.
|
||||
|
||||
Returns:
|
||||
Object detection predictions from endpoints.
|
||||
"""
|
||||
height, width, _ = input_image.shape
|
||||
predictions = detection_endpoint.predict(
|
||||
get_prediction_instances(Image.fromarray(input_image))
|
||||
).predictions
|
||||
detection_scores = np.array(predictions[0]["detection_scores"])
|
||||
detection_classes = np.array(predictions[0]["detection_classes"])
|
||||
detection_boxes = np.array(
|
||||
[
|
||||
[b[1] * width, b[0] * height, b[3] * width, b[2] * height]
|
||||
for b in predictions[0]["detection_boxes"]
|
||||
]
|
||||
)
|
||||
thresh_indices = [
|
||||
x for x, val in enumerate(detection_scores) if val > detection_thresh
|
||||
]
|
||||
preds_merge_conf = np.column_stack((
|
||||
detection_boxes[thresh_indices],
|
||||
detection_scores[thresh_indices],
|
||||
))
|
||||
preds_merge_cls = np.column_stack(
|
||||
(preds_merge_conf, detection_classes[thresh_indices])
|
||||
)
|
||||
return preds_merge_cls
|
||||
@@ -1,9 +1,14 @@
|
||||
"""Vertex vision model garden util constants."""
|
||||
|
||||
# Objectives.
|
||||
# TfVision Objectives.
|
||||
OBJECTIVE_IMAGE_CLASSIFICATION = 'icn'
|
||||
OBJECTIVE_IMAGE_OBJECT_DETECTION = 'iod'
|
||||
OBJECTIVE_IMAGE_SEGMENTATION = 'isg'
|
||||
OBJECTIVE_VIDEO_CLASSIFICATION = 'vcn'
|
||||
OBJECTIVE_VIDEO_ACTION_RECOGNITION = 'var'
|
||||
|
||||
# PyTorch Models.
|
||||
OBJECTIVE_TIMM = 'timm'
|
||||
|
||||
# Input file types.
|
||||
INPUT_FILE_TYPE_CSV = 'csv'
|
||||
@@ -61,4 +66,20 @@ GCSFUSE_URI_PREFIX = '/gcs/'
|
||||
|
||||
LOCAL_EVALUATION_RESULT_DIR = '/tmp/evaluation_result_dir'
|
||||
LOCAL_MODEL_DIR = '/tmp/model_dir'
|
||||
LOCAL_BASE_MODEL_DIR = '/tmp/base_model_dir'
|
||||
LOCAL_DATA_DIR = '/tmp/data'
|
||||
|
||||
# Huggingface files.
|
||||
HF_MODEL_WEIGHTS_SUFFIX = '.bin'
|
||||
|
||||
# PEFT finetuning constants.
|
||||
TEXT_TO_IMAGE_LORA = 'text-to-image-lora'
|
||||
SEQUENCE_CLASSIFICATION_LORA = 'sequence-classification-lora'
|
||||
CAUSAL_LANGUAGE_MODELING_LORA = 'causal-language-modeling-lora'
|
||||
INSTRUCT_LORA = 'instruct-lora'
|
||||
|
||||
# Precision modes for loading model weights.
|
||||
PRECISION_MODE_4 = '4bit'
|
||||
PRECISION_MODE_8 = '8bit'
|
||||
PRECISION_MODE_16 = 'float16'
|
||||
PRECISION_MODE_32 = 'float32'
|
||||
|
||||
@@ -2,6 +2,10 @@
|
||||
|
||||
import glob
|
||||
import os
|
||||
import pathlib
|
||||
import shutil
|
||||
from typing import Tuple
|
||||
import uuid
|
||||
|
||||
from absl import logging
|
||||
from google.cloud import storage
|
||||
@@ -9,6 +13,44 @@ from google.cloud import storage
|
||||
from util import constants
|
||||
|
||||
|
||||
def generate_tmp_path(extension: str = '') -> str:
|
||||
"""Generates a temporary file path with UUID.
|
||||
|
||||
Args:
|
||||
extension: File extension, e.g. '.jpg', '.avi'. If not given, no extension
|
||||
will be appended to the filename.
|
||||
|
||||
Returns:
|
||||
Generated file path.
|
||||
"""
|
||||
return os.path.join(constants.LOCAL_DATA_DIR, uuid.uuid1().hex) + extension
|
||||
|
||||
|
||||
def force_gcs_fuse_path(gcs_uri: str) -> str:
|
||||
"""Converts gs:// uris to their /gcs/ equivalents. No-op for other uris."""
|
||||
if is_gcs_path(gcs_uri):
|
||||
return (
|
||||
constants.GCSFUSE_URI_PREFIX + gcs_uri[len(constants.GCS_URI_PREFIX) :]
|
||||
)
|
||||
else:
|
||||
return gcs_uri
|
||||
|
||||
|
||||
def download_gcs_file_to_local_dir(gcs_uri: str, local_dir: str):
|
||||
"""Download a gcs file to a local dir.
|
||||
|
||||
Args:
|
||||
gcs_uri: A string of file path on GCS.
|
||||
local_dir: A string of local directory.
|
||||
"""
|
||||
if not is_gcs_path(gcs_uri):
|
||||
raise ValueError(
|
||||
f'{gcs_uri} is not a GCS path starting with {constants.GCS_URI_PREFIX}.'
|
||||
)
|
||||
filename = os.path.basename(gcs_uri)
|
||||
download_gcs_file_to_local(gcs_uri, os.path.join(local_dir, filename))
|
||||
|
||||
|
||||
def download_gcs_file_to_local(gcs_uri: str, local_path: str):
|
||||
"""Download a gcs file to a local path.
|
||||
|
||||
@@ -16,7 +58,7 @@ def download_gcs_file_to_local(gcs_uri: str, local_path: str):
|
||||
gcs_uri: A string of file path on GCS.
|
||||
local_path: A string of local file path.
|
||||
"""
|
||||
if not gcs_uri.startswith(constants.GCS_URI_PREFIX):
|
||||
if not is_gcs_path(gcs_uri):
|
||||
raise ValueError(
|
||||
f'{gcs_uri} is not a GCS path starting with {constants.GCS_URI_PREFIX}.'
|
||||
)
|
||||
@@ -26,7 +68,9 @@ def download_gcs_file_to_local(gcs_uri: str, local_path: str):
|
||||
client.download_blob_to_file(gcs_uri, f)
|
||||
|
||||
|
||||
def download_gcs_dir_to_local(gcs_dir: str, local_dir: str):
|
||||
def download_gcs_dir_to_local(
|
||||
gcs_dir: str, local_dir: str, skip_hf_model_bin: bool = False
|
||||
):
|
||||
"""Downloads files in a GCS directory to a local directory.
|
||||
|
||||
For example:
|
||||
@@ -37,7 +81,10 @@ def download_gcs_dir_to_local(gcs_dir: str, local_dir: str):
|
||||
Arguments:
|
||||
gcs_dir: A string of directory path on GCS.
|
||||
local_dir: A string of local directory path.
|
||||
skip_hf_model_bin: True to skip downloading HF model bin files.
|
||||
"""
|
||||
if not is_gcs_path(gcs_dir):
|
||||
raise ValueError(f'{gcs_dir} is not a GCS path starting with gs://.')
|
||||
bucket_name = gcs_dir.split('/')[2]
|
||||
prefix = gcs_dir[len(constants.GCS_URI_PREFIX + bucket_name) :].strip('/')
|
||||
client = storage.Client()
|
||||
@@ -48,8 +95,16 @@ def download_gcs_dir_to_local(gcs_dir: str, local_dir: str):
|
||||
file_path = blob.name[len(prefix) :].strip('/')
|
||||
local_file_path = os.path.join(local_dir, file_path)
|
||||
os.makedirs(os.path.dirname(local_file_path), exist_ok=True)
|
||||
logging.info('Downloading %s to %s', file_path, local_file_path)
|
||||
blob.download_to_filename(local_file_path)
|
||||
if (
|
||||
file_path.endswith(constants.HF_MODEL_WEIGHTS_SUFFIX)
|
||||
and skip_hf_model_bin
|
||||
):
|
||||
logging.info('Skip downloading model bin %s', file_path)
|
||||
with open(local_file_path, 'w') as f:
|
||||
f.write(f'{constants.GCS_URI_PREFIX}{bucket_name}/{prefix}/{file_path}')
|
||||
else:
|
||||
logging.info('Downloading %s to %s', file_path, local_file_path)
|
||||
blob.download_to_filename(local_file_path)
|
||||
|
||||
|
||||
def upload_local_dir_to_gcs(local_dir: str, gcs_dir: str):
|
||||
@@ -77,3 +132,126 @@ def upload_local_dir_to_gcs(local_dir: str, gcs_dir: str):
|
||||
)
|
||||
blob = bucket.blob(os.path.join(blob_dir, os.path.basename(local_file)))
|
||||
blob.upload_from_filename(local_file)
|
||||
|
||||
|
||||
def upload_file_to_gcs_path(
|
||||
source_path: str,
|
||||
destination_uri: str,
|
||||
):
|
||||
"""Uploads local files to GCS uri.
|
||||
|
||||
After upload the destination_uri will contain the same data as the
|
||||
source_path.
|
||||
|
||||
Args:
|
||||
source_path: Required. Path of the local data to copy to GCS.
|
||||
destination_uri: Required. GCS URI where the data should be uploaded.
|
||||
|
||||
Raises:
|
||||
RuntimeError: When source_path does not exist.
|
||||
GoogleCloudError: When the upload process fails.
|
||||
"""
|
||||
source_path_obj = pathlib.Path(source_path)
|
||||
if not source_path_obj.exists():
|
||||
raise RuntimeError(f'Source path does not exist: {source_path}')
|
||||
|
||||
storage_client = storage.Client()
|
||||
source_file_path = source_path
|
||||
destination_file_uri = destination_uri
|
||||
logging.info('Uploading "%s" to "%s"', source_file_path, destination_file_uri)
|
||||
destination_blob = storage.Blob.from_string(
|
||||
destination_file_uri, client=storage_client
|
||||
)
|
||||
destination_blob.upload_from_filename(filename=source_file_path)
|
||||
|
||||
|
||||
def is_gcs_path(input_path: str) -> bool:
|
||||
"""Checks if the input path is a Google Cloud Storage (GCS) path.
|
||||
|
||||
Args:
|
||||
input_path: The input path to be checked.
|
||||
|
||||
Returns:
|
||||
True if the input path is a GCS path, False otherwise.
|
||||
"""
|
||||
return input_path.startswith(constants.GCS_URI_PREFIX)
|
||||
|
||||
|
||||
def release_text_assets(
|
||||
output_bucket: str, local_text_file_name: str, remote_text_file_name: str
|
||||
) -> None:
|
||||
"""Releases text assets.
|
||||
|
||||
Args:
|
||||
output_bucket: gcs output bucket.
|
||||
local_text_file_name: Local text file name.
|
||||
remote_text_file_name: Remote text file name.
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
remote_file_path = '{}/{}'.format(output_bucket, remote_text_file_name)
|
||||
logging.info('Uploading "%s" to "%s"', local_text_file_name, remote_file_path)
|
||||
upload_file_to_gcs_path(local_text_file_name, remote_file_path)
|
||||
os.remove(local_text_file_name)
|
||||
|
||||
|
||||
def upload_video_from_local_to_gcs(
|
||||
output_bucket: str,
|
||||
local_video_file_name: str,
|
||||
remote_video_file_name: str,
|
||||
temp_local_video_file_name: str,
|
||||
) -> None:
|
||||
"""Uploads video from local to gcs buckent and releases video assets.
|
||||
|
||||
Args:
|
||||
output_bucket: GCS bucket address.
|
||||
local_video_file_name: Local video file name.
|
||||
remote_video_file_name: Remote video file name.
|
||||
temp_local_video_file_name: Temporary local video file name.
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
upload_file_to_gcs_path(
|
||||
temp_local_video_file_name,
|
||||
'{}/{}'.format(output_bucket, remote_video_file_name),
|
||||
)
|
||||
shutil.rmtree(local_video_file_name, ignore_errors=True)
|
||||
shutil.rmtree(temp_local_video_file_name, ignore_errors=True)
|
||||
|
||||
|
||||
def download_video_from_gcs_to_local(video_file_path: str) -> Tuple[str, str]:
|
||||
"""Downloads video from gcs to local folders.
|
||||
|
||||
Args:
|
||||
video_file_path: Path to the video file.
|
||||
|
||||
Returns:
|
||||
Local and remote video file paths.
|
||||
"""
|
||||
_, local_video_file_name = os.path.split(video_file_path)
|
||||
file_extension = os.path.splitext(video_file_path)[1]
|
||||
remote_video_file_name = local_video_file_name.replace(
|
||||
file_extension, '_overlay.mp4'
|
||||
)
|
||||
local_file_path = generate_tmp_path(os.path.splitext(video_file_path)[1])
|
||||
logging.info('Downloading %s to %s...', video_file_path, local_file_path)
|
||||
download_gcs_file_to_local(video_file_path, local_file_path)
|
||||
return local_file_path, remote_video_file_name
|
||||
|
||||
|
||||
def get_output_video_file(video_output_file_path: str) -> str:
|
||||
"""Gets the output video file name for writing video.
|
||||
|
||||
Args:
|
||||
video_output_file_path: Path to the video output file.
|
||||
|
||||
Returns:
|
||||
str: Local video output file path.
|
||||
"""
|
||||
file_extension = os.path.splitext(video_output_file_path)[1]
|
||||
out_local_video_file_name = video_output_file_path.replace(
|
||||
file_extension, '_overlay' + file_extension
|
||||
)
|
||||
return out_local_video_file_name
|
||||
|
||||
+1
@@ -4,6 +4,7 @@ import os
|
||||
|
||||
from absl import logging
|
||||
|
||||
|
||||
_ENVIRONMENT_VARIABLE_FOR_TRIAL_ID = 'CLOUD_ML_TRIAL_ID'
|
||||
|
||||
|
||||
-14
@@ -1,14 +0,0 @@
|
||||
"""Video format converter util lib."""
|
||||
|
||||
import io
|
||||
from typing import Sequence
|
||||
import imageio
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def frames_to_video_bytes(frames: Sequence[np.ndarray], fps: int) -> bytes:
|
||||
images = [Image.fromarray(array) for array in frames]
|
||||
io_obj = io.BytesIO()
|
||||
imageio.mimsave(io_obj, images, format=".mp4", fps=fps)
|
||||
return io_obj.getvalue()
|
||||
+73
@@ -0,0 +1,73 @@
|
||||
# Dockerfile for vLLM serving.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/vllm/dockerfile/serve.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/{YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/{YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
# The base image is required by vllm
|
||||
# https://vllm.readthedocs.io/en/latest/getting_started/installation.html
|
||||
FROM nvcr.io/nvidia/pytorch:22.12-py3
|
||||
|
||||
USER root
|
||||
|
||||
# Install tools.
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
RUN apt-get update
|
||||
RUN apt-get install -y --no-install-recommends apt-utils
|
||||
RUN apt-get install -y --no-install-recommends curl
|
||||
RUN apt-get install -y --no-install-recommends wget
|
||||
RUN apt-get install -y --no-install-recommends git
|
||||
RUN apt-get install -y --no-install-recommends jq
|
||||
RUN apt-get install -y --no-install-recommends gnupg
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Install libraries.
|
||||
ENV PIP_ROOT_USER_ACTION=ignore
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
RUN pip install google-cloud-storage==2.7.0
|
||||
RUN pip install absl-py==1.4.0
|
||||
|
||||
# Install pytorch
|
||||
RUN pip install --upgrade torch==2.0.1
|
||||
|
||||
# Install vllm deps.
|
||||
RUN pip install xformers==0.0.20
|
||||
RUN pip install ninja==1.11.1
|
||||
RUN pip install psutil==5.9.5
|
||||
RUN pip install ray==2.6.2
|
||||
RUN pip install sentencepiece==0.1.99
|
||||
RUN pip install fastapi==0.100.1
|
||||
RUN pip install uvicorn==0.23.2
|
||||
RUN pip install pydantic==1.10.12
|
||||
|
||||
# Install transformers from source.
|
||||
WORKDIR /workspace
|
||||
RUN git clone https://github.com/huggingface/transformers.git
|
||||
WORKDIR transformers
|
||||
# Pin the commit to add-code-llama at 08/25/2023
|
||||
RUN git reset --hard 015f8e110d270a0ad42de4ae5b98198d69eb1964
|
||||
RUN pip install -e .
|
||||
WORKDIR /workspace
|
||||
|
||||
# Install vllm from source.
|
||||
RUN git clone https://github.com/vllm-project/vllm.git
|
||||
WORKDIR vllm
|
||||
# Pin the version to a fixed git commit on 08/16/2023.
|
||||
RUN git reset --hard d1744376ae9fdbfa6a2dc763e1c67309e138fa3d
|
||||
# Apply a patch to vllm source:
|
||||
# 1) For models on Huggingface hub: if the model has multiple bin files, each
|
||||
# bin file is downloaded separately and gets deleted after loading to GPU
|
||||
# 2) For models on GCS bucket: each model bin files is download separately
|
||||
# and gets deleted after loading to GPU.
|
||||
# 3) Support code-llama model loading.
|
||||
COPY model_oss/vllm/vllm.patch /tmp/vllm.patch
|
||||
RUN git apply /tmp/vllm.patch
|
||||
RUN pip install -e .
|
||||
|
||||
# Expose port 7080 for host serving.
|
||||
EXPOSE 7080
|
||||
@@ -0,0 +1,311 @@
|
||||
diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py
|
||||
index 99fe593..e11246b 100644
|
||||
--- a/vllm/engine/arg_utils.py
|
||||
+++ b/vllm/engine/arg_utils.py
|
||||
@@ -1,12 +1,43 @@
|
||||
import argparse
|
||||
import dataclasses
|
||||
from dataclasses import dataclass
|
||||
+import os
|
||||
from typing import Optional, Tuple
|
||||
|
||||
+from google.cloud import storage
|
||||
from vllm.config import (CacheConfig, ModelConfig, ParallelConfig,
|
||||
SchedulerConfig)
|
||||
|
||||
|
||||
+GCS_PREFIX = "gs://"
|
||||
+
|
||||
+
|
||||
+def is_gcs_path(input_path: str) -> bool:
|
||||
+ return input_path.startswith(GCS_PREFIX)
|
||||
+
|
||||
+
|
||||
+def download_gcs_dir_to_local(gcs_dir: str, local_dir: str):
|
||||
+ if os.path.isdir(local_dir):
|
||||
+ return
|
||||
+ # gs://bucket_name/dir
|
||||
+ bucket_name = gcs_dir.split('/')[2]
|
||||
+ prefix = gcs_dir[len(GCS_PREFIX + bucket_name) :].strip('/')
|
||||
+ client = storage.Client()
|
||||
+ blobs = client.list_blobs(bucket_name, prefix=prefix)
|
||||
+ for blob in blobs:
|
||||
+ if blob.name[-1] == '/':
|
||||
+ continue
|
||||
+ file_path = blob.name[len(prefix) :].strip('/')
|
||||
+ local_file_path = os.path.join(local_dir, file_path)
|
||||
+ os.makedirs(os.path.dirname(local_file_path), exist_ok=True)
|
||||
+ if file_path.endswith(".bin"):
|
||||
+ with open(local_file_path, 'w') as f:
|
||||
+ f.write(f'{GCS_PREFIX}{bucket_name}/{prefix}/{file_path}')
|
||||
+ else:
|
||||
+ print(f"==> Download {gcs_dir}/{file_path} to {local_file_path}")
|
||||
+ blob.download_to_filename(local_file_path)
|
||||
+
|
||||
+
|
||||
@dataclass
|
||||
class EngineArgs:
|
||||
"""Arguments for vLLM engine."""
|
||||
@@ -143,6 +174,19 @@ class EngineArgs:
|
||||
def create_engine_configs(
|
||||
self,
|
||||
) -> Tuple[ModelConfig, CacheConfig, ParallelConfig, SchedulerConfig]:
|
||||
+ # Preprocess GCS paths.
|
||||
+ if is_gcs_path(self.tokenizer) and self.tokenizer != self.model:
|
||||
+ local_dir = "/tmp/gcs_tokenizer"
|
||||
+ download_gcs_dir_to_local(self.tokenizer, local_dir)
|
||||
+ self.tokenizer = local_dir
|
||||
+ if is_gcs_path(self.model):
|
||||
+ # Download GCS model without bin files.
|
||||
+ local_dir = "/tmp/gcs_model"
|
||||
+ download_gcs_dir_to_local(self.model, local_dir)
|
||||
+ if self.tokenizer == self.model:
|
||||
+ self.tokenizer = local_dir
|
||||
+ self.model = local_dir
|
||||
+
|
||||
# Initialize the configs.
|
||||
model_config = ModelConfig(self.model, self.tokenizer,
|
||||
self.tokenizer_mode, self.trust_remote_code,
|
||||
diff --git a/vllm/entrypoints/api_server.py b/vllm/entrypoints/api_server.py
|
||||
index 58ea2e2..350e209 100644
|
||||
--- a/vllm/entrypoints/api_server.py
|
||||
+++ b/vllm/entrypoints/api_server.py
|
||||
@@ -15,6 +15,10 @@ TIMEOUT_KEEP_ALIVE = 5 # seconds.
|
||||
TIMEOUT_TO_PREVENT_DEADLOCK = 1 # seconds.
|
||||
app = FastAPI()
|
||||
|
||||
+# Required by Vertex deployment.
|
||||
+@app.get("/ping")
|
||||
+async def ping() -> Response:
|
||||
+ return Response(status_code=200)
|
||||
|
||||
@app.post("/generate")
|
||||
async def generate(request: Request) -> Response:
|
||||
@@ -26,6 +30,9 @@ async def generate(request: Request) -> Response:
|
||||
- other fields: the sampling parameters (See `SamplingParams` for details).
|
||||
"""
|
||||
request_dict = await request.json()
|
||||
+ is_on_vertex = "instances" in request_dict
|
||||
+ if is_on_vertex:
|
||||
+ request_dict = request_dict["instances"][0]
|
||||
prompt = request_dict.pop("prompt")
|
||||
stream = request_dict.pop("stream", False)
|
||||
sampling_params = SamplingParams(**request_dict)
|
||||
@@ -63,7 +70,10 @@ async def generate(request: Request) -> Response:
|
||||
assert final_output is not None
|
||||
prompt = final_output.prompt
|
||||
text_outputs = [prompt + output.text for output in final_output.outputs]
|
||||
- ret = {"text": text_outputs}
|
||||
+ if is_on_vertex:
|
||||
+ ret = {"predictions": text_outputs}
|
||||
+ else:
|
||||
+ ret = {"text": text_outputs}
|
||||
return JSONResponse(ret)
|
||||
|
||||
|
||||
diff --git a/vllm/model_executor/models/llama.py b/vllm/model_executor/models/llama.py
|
||||
index 93ab499..eca1b89 100644
|
||||
--- a/vllm/model_executor/models/llama.py
|
||||
+++ b/vllm/model_executor/models/llama.py
|
||||
@@ -85,6 +85,7 @@ class LlamaAttention(nn.Module):
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
num_kv_heads: int,
|
||||
+ rope_theta: float = 10000,
|
||||
):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
@@ -99,6 +100,7 @@ class LlamaAttention(nn.Module):
|
||||
self.q_size = self.num_heads * self.head_dim
|
||||
self.kv_size = self.num_kv_heads * self.head_dim
|
||||
self.scaling = self.head_dim**-0.5
|
||||
+ self.rope_theta = rope_theta
|
||||
|
||||
self.qkv_proj = ColumnParallelLinear(
|
||||
hidden_size,
|
||||
@@ -118,6 +120,7 @@ class LlamaAttention(nn.Module):
|
||||
self.attn = PagedAttentionWithRoPE(self.num_heads,
|
||||
self.head_dim,
|
||||
self.scaling,
|
||||
+ base=self.rope_theta,
|
||||
rotary_dim=self.head_dim,
|
||||
num_kv_heads=self.num_kv_heads)
|
||||
|
||||
@@ -143,10 +146,15 @@ class LlamaDecoderLayer(nn.Module):
|
||||
def __init__(self, config: LlamaConfig):
|
||||
super().__init__()
|
||||
self.hidden_size = config.hidden_size
|
||||
+ try:
|
||||
+ rope_theta = config.rope_theta
|
||||
+ except AttributeError:
|
||||
+ rope_theta = 10000
|
||||
self.self_attn = LlamaAttention(
|
||||
hidden_size=self.hidden_size,
|
||||
num_heads=config.num_attention_heads,
|
||||
num_kv_heads=config.num_key_value_heads,
|
||||
+ rope_theta=rope_theta,
|
||||
)
|
||||
self.mlp = LlamaMLP(
|
||||
hidden_size=self.hidden_size,
|
||||
diff --git a/vllm/model_executor/weight_utils.py b/vllm/model_executor/weight_utils.py
|
||||
index a9d899a..57f39b5 100644
|
||||
--- a/vllm/model_executor/weight_utils.py
|
||||
+++ b/vllm/model_executor/weight_utils.py
|
||||
@@ -3,13 +3,17 @@ import filelock
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
+import time
|
||||
from typing import Iterator, List, Optional, Tuple
|
||||
|
||||
-from huggingface_hub import snapshot_download
|
||||
+from google.cloud import storage
|
||||
+from huggingface_hub import hf_hub_download, snapshot_download
|
||||
import numpy as np
|
||||
import torch
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
+HF_PREFIX = "hf://"
|
||||
+
|
||||
|
||||
class Disabledtqdm(tqdm):
|
||||
|
||||
@@ -22,60 +26,90 @@ def hf_model_weights_iterator(
|
||||
cache_dir: Optional[str] = None,
|
||||
use_np_cache: bool = False,
|
||||
) -> Iterator[Tuple[str, torch.Tensor]]:
|
||||
+ if use_np_cache:
|
||||
+ raise ValueError("Do not support use_np_cache for lazy download.")
|
||||
+
|
||||
# Prepare file lock directory to prevent multiple processes from
|
||||
# downloading the same model weights at the same time.
|
||||
lock_dir = cache_dir if cache_dir is not None else "/tmp"
|
||||
lock_file_name = model_name_or_path.replace("/", "-") + ".lock"
|
||||
lock = filelock.FileLock(os.path.join(lock_dir, lock_file_name))
|
||||
|
||||
- # Download model weights from huggingface.
|
||||
- is_local = os.path.isdir(model_name_or_path)
|
||||
- if not is_local:
|
||||
- with lock:
|
||||
- hf_folder = snapshot_download(model_name_or_path,
|
||||
- allow_patterns="*.bin",
|
||||
- cache_dir=cache_dir,
|
||||
- tqdm_class=Disabledtqdm)
|
||||
- else:
|
||||
- hf_folder = model_name_or_path
|
||||
-
|
||||
- hf_bin_files = [
|
||||
- x for x in glob.glob(os.path.join(hf_folder, "*.bin"))
|
||||
- if not x.endswith("training_args.bin")
|
||||
- ]
|
||||
-
|
||||
- if use_np_cache:
|
||||
- # Convert the model weights from torch tensors to numpy arrays for
|
||||
- # faster loading.
|
||||
- np_folder = os.path.join(hf_folder, "np")
|
||||
- os.makedirs(np_folder, exist_ok=True)
|
||||
- weight_names_file = os.path.join(np_folder, "weight_names.json")
|
||||
- with lock:
|
||||
- if not os.path.exists(weight_names_file):
|
||||
- weight_names = []
|
||||
- for bin_file in hf_bin_files:
|
||||
- state = torch.load(bin_file, map_location="cpu")
|
||||
- for name, param in state.items():
|
||||
- param_path = os.path.join(np_folder, name)
|
||||
- with open(param_path, "wb") as f:
|
||||
- np.save(f, param.cpu().detach().numpy())
|
||||
- weight_names.append(name)
|
||||
- with open(weight_names_file, "w") as f:
|
||||
- json.dump(weight_names, f)
|
||||
-
|
||||
- with open(weight_names_file, "r") as f:
|
||||
- weight_names = json.load(f)
|
||||
-
|
||||
- for name in weight_names:
|
||||
- param_path = os.path.join(np_folder, name)
|
||||
- with open(param_path, "rb") as f:
|
||||
- param = np.load(f)
|
||||
- yield name, torch.from_numpy(param)
|
||||
+ bin_files = []
|
||||
+ if not os.path.isdir(model_name_or_path):
|
||||
+ try:
|
||||
+ with lock:
|
||||
+ index_file = hf_hub_download(repo_id=model_name_or_path,
|
||||
+ filename="pytorch_model.bin.index.json",
|
||||
+ cache_dir=cache_dir)
|
||||
+ except:
|
||||
+ print("==> The model is in HF hub with 1 bin file, download it directly.", flush=True)
|
||||
+ with lock:
|
||||
+ hf_folder = snapshot_download(repo_id=model_name_or_path,
|
||||
+ allow_patterns="*.bin",
|
||||
+ cache_dir=cache_dir,
|
||||
+ tqdm_class=Disabledtqdm)
|
||||
+ bin_files = [x for x in glob.glob(os.path.join(hf_folder, "*.bin"))]
|
||||
+ else:
|
||||
+ print("==> The model is in HF hub with multiple bin file, do not download it now.", flush=True)
|
||||
+ with open(index_file, "r") as f:
|
||||
+ index = json.loads(f.read())
|
||||
+ bin_filenames = set(index["weight_map"].values())
|
||||
+ bin_files = [f"{HF_PREFIX}{model_name_or_path}/{bin_filename}" for bin_filename in bin_filenames]
|
||||
else:
|
||||
- for bin_file in hf_bin_files:
|
||||
- state = torch.load(bin_file, map_location="cpu")
|
||||
- for name, param in state.items():
|
||||
- yield name, param
|
||||
+ print("==> The model is in local disk.", flush=True)
|
||||
+ bin_files = [x for x in glob.glob(os.path.join(model_name_or_path, "*.bin"))]
|
||||
+
|
||||
+ if "training_args.bin" in bin_files:
|
||||
+ bin_files.remove("training_args.bin")
|
||||
+ bin_files.sort()
|
||||
+ print(f"==> Fetched bin files: {bin_files}", flush=True)
|
||||
+
|
||||
+ model_dir = "/tmp/model"
|
||||
+ os.makedirs(model_dir, exist_ok=True)
|
||||
+ for bin_file in bin_files:
|
||||
+ delete_download = False
|
||||
+
|
||||
+ if os.path.exists(bin_file):
|
||||
+ if open(bin_file, "rb").read(2) == b"gs":
|
||||
+ gcs_path = open(bin_file).read()
|
||||
+ bin_filename = gcs_path.split("/")[-1]
|
||||
+ local_file = os.path.join(model_dir, bin_filename)
|
||||
+ with lock:
|
||||
+ if not os.path.exists(local_file):
|
||||
+ client = storage.Client()
|
||||
+ with open(local_file, 'wb') as f:
|
||||
+ print(f"==> Download {gcs_path} to {bin_file}", flush=True)
|
||||
+ client.download_blob_to_file(gcs_path, f)
|
||||
+ bin_file = local_file
|
||||
+ delete_download = True
|
||||
+ else:
|
||||
+ assert bin_file.startswith(HF_PREFIX)
|
||||
+ bin_filename = os.path.basename(bin_file)
|
||||
+ local_file = os.path.join(model_dir, bin_filename)
|
||||
+ with lock:
|
||||
+ if not os.path.exists(local_file):
|
||||
+ print(f"==> Download {model_name_or_path}/{bin_filename} to {local_file}", flush=True)
|
||||
+ hf_hub_download(repo_id=model_name_or_path,
|
||||
+ filename=bin_filename,
|
||||
+ local_dir=model_dir,
|
||||
+ local_dir_use_symlinks=False,
|
||||
+ force_download=True)
|
||||
+ bin_file = local_file
|
||||
+ delete_download = True
|
||||
+
|
||||
+ torch.distributed.barrier()
|
||||
+ print(f"==> Load {bin_file} to memory.", flush=True)
|
||||
+ state = torch.load(bin_file, map_location="cpu")
|
||||
+ for name, param in state.items():
|
||||
+ yield name, param
|
||||
+ torch.distributed.barrier()
|
||||
+
|
||||
+ if delete_download:
|
||||
+ with lock:
|
||||
+ if os.path.exists(bin_file):
|
||||
+ print(f"==> Delete {bin_file}", flush=True)
|
||||
+ os.remove(bin_file)
|
||||
|
||||
|
||||
def load_tensor_parallel_weights(
|
||||
|
||||
+87
@@ -0,0 +1,87 @@
|
||||
# Dockerfile for basic serving dockers for vot.
|
||||
#
|
||||
# To build:
|
||||
# docker build -f model_oss/vot/dockerfile/serving.Dockerfile . -t ${YOUR_IMAGE_TAG}
|
||||
#
|
||||
# To push to gcr:
|
||||
# docker tag ${YOUR_IMAGE_TAG} gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
# docker push gcr.io/${YOUR_PROJECT}/${YOUR_IMAGE_TAG}
|
||||
|
||||
# Switch to this base image for gpu serve.
|
||||
FROM pytorch/torchserve:0.7.0-gpu
|
||||
|
||||
USER root
|
||||
|
||||
ENV infer_port=7080
|
||||
ENV mng_port=7081
|
||||
ENV model_name="vot_serving"
|
||||
ENV PATH="/home/model-server/:${PATH}"
|
||||
|
||||
ENV PYTHONPATH "${PYTHONPATH}:/automl_vision/vot"
|
||||
|
||||
# Install libraries.
|
||||
ENV PIP_ROOT_USER_ACTION=ignore
|
||||
RUN python3 -m pip install --upgrade pip
|
||||
RUN pip install accelerate==0.17.0
|
||||
RUN pip install datasets==2.9.0
|
||||
RUN pip install bytetracker==0.3.2
|
||||
RUN pip install imageio[ffmpeg]==2.31.1
|
||||
RUN pip install google-cloud-aiplatform==1.25.0
|
||||
RUN pip install google-cloud-storage==2.9.0
|
||||
RUN pip install fastapi==0.96.0
|
||||
RUN pip install lap==0.4.0
|
||||
RUN pip install numpy==1.24.3
|
||||
RUN pip install opencv-python==4.7.0.72
|
||||
RUN pip install Pillow==9.5.0
|
||||
RUN pip install protobuf==3.19.6
|
||||
RUN pip install pandas==2.0.2
|
||||
RUN pip install pycocotools==2.0.6
|
||||
RUN pip install scipy==1.10.1
|
||||
RUN pip install tensorflow==2.11.1
|
||||
RUN pip install torch==2.0.1
|
||||
RUN pip install torchvision==0.15.2
|
||||
RUN pip install triton==2.0.0.dev20221120
|
||||
RUN pip install uvicorn==0.22.0
|
||||
|
||||
# Install tools.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl \
|
||||
wget \
|
||||
vim
|
||||
|
||||
# Copy license.
|
||||
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
|
||||
|
||||
# Copy model artifacts.
|
||||
COPY model_oss/vot/handler.py /home/model-server/handler.py
|
||||
COPY model_oss/vot/visualization_utils.py /home/model-server/vot/
|
||||
COPY model_oss/util/ /home/model-server/util/
|
||||
ENV PYTHONPATH /home/model-server/
|
||||
|
||||
# Create torchserve configuration file.
|
||||
RUN echo \
|
||||
"default_response_timeout=3600\n" \
|
||||
"service_envelope=json\n" \
|
||||
"inference_address=http://0.0.0.0:${infer_port}\n" \
|
||||
"management_address=http://0.0.0.0:${mng_port}" >> /home/model-server/config.properties
|
||||
|
||||
# Expose ports.
|
||||
EXPOSE ${infer_port}
|
||||
EXPOSE ${mng_port}
|
||||
|
||||
# Archive model artifacts and dependencies.
|
||||
# Do not set --model-file and --serialized-file because model and checkpoint will be dynamically loaded in handler.py.
|
||||
RUN torch-model-archiver \
|
||||
--model-name=${model_name} \
|
||||
--version=1.0 \
|
||||
--handler=/home/model-server/handler.py \
|
||||
--runtime=python3 \
|
||||
--export-path=/home/model-server/model-store \
|
||||
--archive-format=default \
|
||||
--force
|
||||
|
||||
# Run Torchserve HTTP serve to respond to prediction requests.
|
||||
CMD ["torchserve", "--start", \
|
||||
"--ts-config", "/home/model-server/config.properties", \
|
||||
"--models", "${model_name}=${model_name}.mar", \
|
||||
"--model-store", "/home/model-server/model-store"]
|
||||
@@ -0,0 +1,196 @@
|
||||
"""Custom handler for video object tracking models."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
from typing import Any, List, Optional, Tuple
|
||||
|
||||
from bytetracker import BYTETracker
|
||||
import cv2
|
||||
from google.cloud import aiplatform
|
||||
import imageio.v2 as iio
|
||||
from PIL import Image
|
||||
import tensorflow as tf
|
||||
import torch
|
||||
from ts.torch_handler.base_handler import BaseHandler
|
||||
|
||||
from util import commons
|
||||
from util import fileutils
|
||||
import visualization_utils
|
||||
|
||||
_VIDEO_URI = "video_uri"
|
||||
_DATA = "data"
|
||||
_TRACK_THRESHOLD = 0.45
|
||||
_TRACK_BUFFER = 25
|
||||
_MATCH_THRESHOLD = 0.8
|
||||
|
||||
|
||||
class VideoObjectTrackingHandler(BaseHandler):
|
||||
"""Custom handler for video object tracking models."""
|
||||
|
||||
def initialize(self, context: Any) -> None:
|
||||
properties = context.system_properties
|
||||
self.map_location = (
|
||||
"cuda"
|
||||
if torch.cuda.is_available() and properties.get("gpu_id") is not None
|
||||
else "cpu"
|
||||
)
|
||||
self.device = torch.device(
|
||||
self.map_location + ":" + str(properties.get("gpu_id"))
|
||||
if torch.cuda.is_available() and properties.get("gpu_id") is not None
|
||||
else self.map_location
|
||||
)
|
||||
self.manifest = context.manifest
|
||||
|
||||
detection_endpoint_id = os.environ.get("DETECTION_ENDPOINT", None)
|
||||
if detection_endpoint_id:
|
||||
self.detection_endpoint = aiplatform.Endpoint(detection_endpoint_id)
|
||||
endpoint_label_map = os.environ.get("LABEL_MAP", None)
|
||||
if endpoint_label_map:
|
||||
endpoint_label_map_file = endpoint_label_map
|
||||
self.label_map = commons.get_label_map(endpoint_label_map_file)
|
||||
else:
|
||||
raise ValueError(
|
||||
"LABEL MAP must be provided with DETECTION ENDPOINT:"
|
||||
f" {self.detection_endpoint}"
|
||||
)
|
||||
|
||||
self.track_thresh = os.environ.get("TRACK_THRESHOLD", _TRACK_THRESHOLD)
|
||||
self.track_buffer = os.environ.get("TRACK_BUFFER", _TRACK_BUFFER)
|
||||
self.match_thresh = os.environ.get("MATCH_THRESHOLD", _MATCH_THRESHOLD)
|
||||
self.save_video_results = bool(int(os.environ.get("SAVE_VIDEO_RESULTS", 0)))
|
||||
self.output_bucket = os.environ.get("OUTPUT_BUCKET", None)
|
||||
if not self.output_bucket:
|
||||
raise ValueError("Empty Output Bucket.")
|
||||
self.initialized = True
|
||||
logging.info("Handler initialization done.")
|
||||
|
||||
def preprocess(
|
||||
self, data: Any
|
||||
) -> Tuple[Optional[List[str]], Optional[List[Image.Image]]]:
|
||||
"""Preprocesses the input data.
|
||||
|
||||
Args:
|
||||
data (Any): Input data.
|
||||
|
||||
Returns:
|
||||
List of videos uris.
|
||||
"""
|
||||
video_uris = None
|
||||
if _VIDEO_URI in data[0]:
|
||||
video_uris = [item[_VIDEO_URI] for item in data]
|
||||
# TorchServe's default handlers expect each instance
|
||||
# to be wrapped in a data field for batch prediction.
|
||||
if _DATA in data[0]:
|
||||
video_uris = [item[_DATA][_VIDEO_URI] for item in data]
|
||||
return video_uris
|
||||
|
||||
def inference(self, data: Any, *args, **kwargs) -> List[Any]:
|
||||
"""Runs object detection and tracking inference on a video frame by frame.
|
||||
|
||||
If using yolo detection, the function uses the ultralytics yolo models for
|
||||
IOD, otherwise it uses the provided IOD endpoint and associated the selected
|
||||
tracking method to the detections.
|
||||
|
||||
Args:
|
||||
data: List of video files.
|
||||
*args: Additional arguments.
|
||||
**kwargs: Additional keyword arguments.
|
||||
|
||||
Returns:
|
||||
List of video frame annotations and/or output decorated video uris.
|
||||
"""
|
||||
gcs_video_files = data
|
||||
video_preds = []
|
||||
for gcs_video_file in gcs_video_files:
|
||||
results_info = {}
|
||||
temp_text_file = tempfile.NamedTemporaryFile(delete=False, mode="w+t")
|
||||
local_video_file_name, remote_video_file_name = (
|
||||
fileutils.download_video_from_gcs_to_local(gcs_video_file)
|
||||
)
|
||||
remote_text_file_name = remote_video_file_name.replace(
|
||||
"overlay.mp4", "annotations.txt"
|
||||
)
|
||||
|
||||
cap = cv2.VideoCapture(local_video_file_name)
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
temp_local_video_file_name = fileutils.get_output_video_file(
|
||||
local_video_file_name
|
||||
)
|
||||
if self.save_video_results:
|
||||
self.video_writer = iio.get_writer(
|
||||
temp_local_video_file_name,
|
||||
format="FFMPEG",
|
||||
mode="I",
|
||||
fps=float(fps),
|
||||
codec="h264",
|
||||
)
|
||||
self.tracker = BYTETracker(
|
||||
track_thresh=self.track_thresh,
|
||||
track_buffer=self.track_buffer,
|
||||
match_thresh=self.match_thresh,
|
||||
frame_rate=fps,
|
||||
)
|
||||
frame_idx = 1
|
||||
while cap.isOpened():
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
dets_np = commons.get_object_detection_endpoint_predictions(
|
||||
self.detection_endpoint, frame
|
||||
)
|
||||
dets_tf = tf.convert_to_tensor(dets_np)
|
||||
online_targets = self.tracker.update(dets_tf, None)
|
||||
if online_targets.size > 0:
|
||||
frame = visualization_utils.overlay_tracking_results(
|
||||
frame_idx,
|
||||
frame,
|
||||
online_targets,
|
||||
label_map=self.label_map,
|
||||
temp_text_file_path=temp_text_file.name,
|
||||
)
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
if self.save_video_results:
|
||||
self.video_writer.append_data(frame)
|
||||
logging.info(
|
||||
"Finished processing frame %s for video %s.",
|
||||
frame_idx,
|
||||
gcs_video_file,
|
||||
)
|
||||
frame_idx += 1
|
||||
|
||||
self.video_writer.close()
|
||||
cap.release()
|
||||
if self.save_video_results:
|
||||
fileutils.upload_video_from_local_to_gcs(
|
||||
self.output_bucket,
|
||||
local_video_file_name,
|
||||
remote_video_file_name,
|
||||
temp_local_video_file_name,
|
||||
)
|
||||
results_info["output_video"] = "{}/{}".format(
|
||||
self.output_bucket, remote_video_file_name
|
||||
)
|
||||
|
||||
fileutils.release_text_assets(
|
||||
self.output_bucket,
|
||||
temp_text_file.name,
|
||||
remote_text_file_name,
|
||||
)
|
||||
results_info["annotations"] = "{}/{}".format(
|
||||
self.output_bucket, remote_text_file_name
|
||||
)
|
||||
video_preds.append(results_info)
|
||||
|
||||
return video_preds
|
||||
|
||||
def handle(self, data: Any, context: Any) -> List[Any]:
|
||||
model_input = self.preprocess(data)
|
||||
model_out = self.inference(model_input)
|
||||
output = self.postprocess(model_out)
|
||||
return output
|
||||
|
||||
def postprocess(self, inference_result: List[Any]) -> List[Any]:
|
||||
return inference_result
|
||||
@@ -0,0 +1,168 @@
|
||||
"""Image and bounding box visualization util lib."""
|
||||
|
||||
from typing import Dict, List, Optional
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import ImageColor
|
||||
|
||||
|
||||
def draw_bounding_box_on_image(
|
||||
image: np.ndarray,
|
||||
ymin: float,
|
||||
xmin: float,
|
||||
ymax: float,
|
||||
xmax: float,
|
||||
color: str,
|
||||
thickness: int = 4,
|
||||
display_str_list: Optional[List[str]] = None,
|
||||
) -> np.ndarray:
|
||||
"""Draws a bounding box on an image.
|
||||
|
||||
Args:
|
||||
image: The image to draw the bounding box on.
|
||||
ymin: The minimum y-coordinate of the bounding box.
|
||||
xmin: The minimum x-coordinate of the bounding box.
|
||||
ymax: The maximum y-coordinate of the bounding box.
|
||||
xmax: The maximum x-coordinate of the bounding box.
|
||||
color: The color of the bounding box.
|
||||
thickness: The thickness of the bounding box lines. Defaults to 4.
|
||||
display_str_list: List of strings to display in new line inside the bounding
|
||||
box.
|
||||
|
||||
Returns:
|
||||
An image with a bounding box.
|
||||
"""
|
||||
color = ImageColor.getrgb(color)
|
||||
cv2.rectangle(
|
||||
image, (int(xmin), int(ymin)), (int(xmax), int(ymax)), color, thickness
|
||||
)
|
||||
# Display the strings below the bounding box
|
||||
for i, display_str in enumerate(display_str_list):
|
||||
font = cv2.FONT_HERSHEY_SIMPLEX
|
||||
scale = 0.4
|
||||
thickness = 1
|
||||
text_width, text_height = cv2.getTextSize(
|
||||
display_str, font, scale, thickness
|
||||
)[0]
|
||||
text_bottom = int(ymin - i * text_height)
|
||||
text_left = int(xmin)
|
||||
cv2.rectangle(
|
||||
image,
|
||||
(text_left, text_bottom - text_height),
|
||||
(text_left + text_width, text_bottom),
|
||||
color,
|
||||
-1,
|
||||
)
|
||||
cv2.putText(
|
||||
image,
|
||||
display_str,
|
||||
(text_left, text_bottom),
|
||||
font,
|
||||
scale,
|
||||
(0, 0, 0),
|
||||
thickness,
|
||||
)
|
||||
|
||||
return image
|
||||
|
||||
|
||||
def draw_boxes(
|
||||
image: np.ndarray,
|
||||
boxes: List[List[float]],
|
||||
track_ids: List[int],
|
||||
class_names: List[str],
|
||||
scores: List[float],
|
||||
max_boxes: int = 40,
|
||||
min_score: float = 0.05,
|
||||
) -> np.ndarray:
|
||||
"""Overlays labeled boxes on an image with formatted scores and label names.
|
||||
|
||||
Args:
|
||||
image: The image to overlay the boxes on.
|
||||
boxes: List of bounding box coordinates [xmin, ymin, xmax, ymax].
|
||||
track_ids: List of track IDs corresponding to each box.
|
||||
class_names: List of class names corresponding to each box.
|
||||
scores: List of scores corresponding to each box.
|
||||
max_boxes: Maximum number of boxes to draw. Defaults to 40.
|
||||
min_score: Minimum score threshold for displaying a box. Defaults to 0.05.
|
||||
|
||||
Returns:
|
||||
PIL.Image.Image: The image with the labeled boxes overlay.
|
||||
"""
|
||||
colors = list(ImageColor.colormap.values())
|
||||
for i in range(min(len(boxes), max_boxes)):
|
||||
if scores[i] >= min_score:
|
||||
xmin, ymin, xmax, ymax = boxes[i]
|
||||
display_str = "{}-{}: {}%".format(
|
||||
track_ids[i], class_names[i], int(100 * scores[i])
|
||||
)
|
||||
color = colors[hash(class_names[i]) % len(colors)]
|
||||
image = draw_bounding_box_on_image(
|
||||
image,
|
||||
ymin,
|
||||
xmin,
|
||||
ymax,
|
||||
xmax,
|
||||
color,
|
||||
display_str_list=[display_str],
|
||||
)
|
||||
return image
|
||||
|
||||
|
||||
def overlay_tracking_results(
|
||||
frame_idx: int,
|
||||
image_np: np.ndarray,
|
||||
tracker_outputs: np.ndarray,
|
||||
model_names: Optional[Dict[int, str]] = None,
|
||||
label_map: Optional[Dict[str, Dict[int, str]]] = None,
|
||||
temp_text_file_path: Optional[str] = None,
|
||||
) -> np.ndarray:
|
||||
"""Overlays the results on the image.
|
||||
|
||||
Args:
|
||||
frame_idx: frame index.
|
||||
image_np: Input image.
|
||||
tracker_outputs: Tracker outputs.
|
||||
model_names: label map for yolo models.
|
||||
label_map: label map for IOD detector model.
|
||||
temp_text_file_path: tempfile to save annotations.
|
||||
|
||||
Returns:
|
||||
Decorated output frame.
|
||||
"""
|
||||
dboxes = tracker_outputs[:, :4]
|
||||
dtracks = tracker_outputs[:, 4]
|
||||
dclasses = tracker_outputs[:, 5]
|
||||
dscores = tracker_outputs[:, 6]
|
||||
dclasses_as_text = []
|
||||
for detection_class in dclasses:
|
||||
if model_names:
|
||||
dclasses_as_text.append(model_names[int(detection_class)])
|
||||
elif label_map:
|
||||
dclasses_as_text.append(label_map["label_map"][int(detection_class)])
|
||||
else:
|
||||
dclasses_as_text.append("")
|
||||
|
||||
plotted_img = np.array(
|
||||
draw_boxes(
|
||||
image=image_np,
|
||||
boxes=dboxes,
|
||||
track_ids=dtracks,
|
||||
class_names=dclasses_as_text,
|
||||
scores=dscores,
|
||||
)
|
||||
)
|
||||
track_anno_list = []
|
||||
for i, box in enumerate(dboxes):
|
||||
result_list = [dtracks[i], dscores[i], dclasses[i]]
|
||||
xyxy_anno = [np.round(item.item(), 2) for item in box]
|
||||
tracks_anno = (
|
||||
[frame_idx]
|
||||
+ [np.round(item.item(), 2) for item in result_list]
|
||||
+ xyxy_anno
|
||||
)
|
||||
track_anno_list.append(tracks_anno)
|
||||
with open(temp_text_file_path, "a") as file:
|
||||
file.write(", ".join([str(item) for item in tracks_anno]) + "\n")
|
||||
|
||||
return plotted_img
|
||||
@@ -43,8 +43,12 @@
|
||||
/notebooks/community/pipelines/google_cloud_pipeline_components_ready_to_go_text_classification_pipeline.ipynb @Narwhalprime
|
||||
/notebooks/community/feature_store/get_started_vertex_feature_store.ipynb @junkourata
|
||||
/notebooks/community/model_garden/model_garden_huggingface_local_inference.ipynb @dstnluong-google
|
||||
/notebooks/community/model_garden/model_garden_mediapipe_face_stylizer.pynb @schmidt-sebastian
|
||||
/notebooks/community/model_garden/model_garden_mediapipe_gesture_recognition.ipynb @schmidt-sebastian
|
||||
/notebooks/community/model_garden/model_garden_mediapipe_image_classification.ipynb @schmidt-sebastian
|
||||
/notebooks/community/model_garden/model_garden_mediapipe_image_generation.ipynb @schmidt-sebastian
|
||||
/notebooks/community/model_garden/model_garden_mediapipe_object_detection.ipynb @schmidt-sebastian
|
||||
/notebooks/community/model_garden/model_garden_mediapipe_text_classification.ipynb @schmidt-sebastian
|
||||
/notebooks/community/model_garden/model_garden_proprietary_image_classification.ipynb @weigary
|
||||
/notebooks/community/model_garden/model_garden_proprietary_image_object_detection.ipynb @weigary
|
||||
/notebooks/community/model_garden/model_garden_tfvision_image_classification.ipynb @genquan9
|
||||
@@ -53,6 +57,7 @@
|
||||
/notebooks/community/model_garden/model_garden_pytorch_stable_diffusion.ipynb @xiangxu-google
|
||||
/notebooks/community/model_garden/model_garden_pytorch_stable_diffusion_2_1.ipynb @bingatgoogle
|
||||
/notebooks/community/model_garden/model_garden_pytorch_stable_diffusion_inpainting.ipynb @xiangxu-google
|
||||
/notebooks/community/model_garden/model_garden_pytorch_stable_diffusion_xl_1_0.ipynb @bingatgoogle
|
||||
/notebooks/community/model_garden/model_garden_pytorch_instructpix2pix.ipynb @xiangxu-google
|
||||
/notebooks/community/model_garden/model_garden_pytorch_controlnet.ipynb @xiangxu-google
|
||||
/notebooks/community/model_garden/model_garden_pytorch_blip_image_captioning.ipynb @xiangxu-google
|
||||
@@ -66,11 +71,24 @@
|
||||
/notebooks/community/model_garden/model_garden_pytorch_detectron2.ipynb @lavraicse
|
||||
/notebooks/community/model_garden/model_garden_pytorch_dolly_v2.ipynb @lavraicse
|
||||
/notebooks/community/model_garden/model_garden_pytorch_bart_large_cnn.ipynb @lavraicse
|
||||
/notebooks/community/model_garden/model_garden_pytorch_starcoder.ipynb @xcchen1
|
||||
/notebooks/community/model_garden/model_garden_jax_vision_transformer.ipynb @lavraicse
|
||||
/notebooks/community/model_garden/model_garden_jax_fvlm.ipynb @lavraicse
|
||||
/notebooks/community/model_garden/model_garden_pytorch_text_to_video_zero_shot.ipynb @bingatgoogle
|
||||
/notebooks/community/model_garden/model_garden_pytorch_text_to_video.ipynb @KCFindstr
|
||||
/notebooks/community/generative_ai/text_embedding_api_cloud_next_new_models.ipynb @xqr-g
|
||||
/notebooks/community/generative_ai/text_embedding_api_semantic_search_with_scann.ipynb @henrytansetiawan
|
||||
/notebooks/community/bigquery_ml_inference/bq_ml_with_vision_translation_nlp.ipynb @deaconsmith
|
||||
/notebooks/community/model_garden/model_garden_keras_stable_diffusion.ipynb @genquan9
|
||||
/notebooks/community/model_garden/model_garden_keras_yolov8.ipynb @@dstnluong-google
|
||||
/notebooks/community/model_garden/model_garden_pytorch_sam.ipynb @huguensjean
|
||||
/notebooks/community/model_garden/model_garden_pytorch_pic2word.ipynb @jismailyan
|
||||
/notebooks/community/model_garden/model_garden_pytorch_peft.ipynb @genquan9
|
||||
/notebooks/community/model_garden/model_garden_pytorch_openllama_peft.ipynb @genquan9
|
||||
/notebooks/community/model_garden/model_garden_pytorch_falcon_instruct_peft.ipynb @genquan9
|
||||
/notebooks/community/model_garden/model_garden_movinet_clip_classification.ipynb @KCFindstr
|
||||
/notebooks/community/model_garden/model_garden_movinet_action_recognition.ipynb @KCFindstr
|
||||
/notebooks/community/model_garden/model_garden_pytorch_open_clip.ipynb @lydhr
|
||||
/notebooks/community/model_garden/model_garden_pytorch_llama2_peft.ipynb @genquan9
|
||||
/notebooks/community/model_garden/model_garden_pytorch_codellama.ipynb @xiangxu-google
|
||||
/notebooks/community/model_garden/model_garden_pytorch_nllb.ipynb @weigary
|
||||
|
||||
@@ -0,0 +1,894 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "ur8xi4C7S06n"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Copyright 2023 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",
|
||||
"# You may obtain a copy of the License at\n",
|
||||
"#\n",
|
||||
"# https://www.apache.org/licenses/LICENSE-2.0\n",
|
||||
"#\n",
|
||||
"# Unless required by applicable law or agreed to in writing, software\n",
|
||||
"# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
|
||||
"# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
|
||||
"# See the License for the specific language governing permissions and\n",
|
||||
"# limitations under the License."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "JAPoU8Sm5E6e"
|
||||
},
|
||||
"source": [
|
||||
"## Use BigQuery DataFrames with Generative AI for code generation\n",
|
||||
"\n",
|
||||
"<table align=\"left\">\n",
|
||||
"\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://colab.research.google.com/github/googleapis/python-bigquery-dataframes/tree/main/notebooks/getting_started/bq_dataframes_llm_code_generation.ipynb\">\n",
|
||||
" <img src=\"https://cloud.google.com/ml-engine/images/colab-logo-32px.png\" alt=\"Colab logo\"> Run in Colab\n",
|
||||
" </a>\n",
|
||||
" </td>\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://github.com/googleapis/python-bigquery-dataframes/tree/main/notebooks/getting_started/bq_dataframes_llm_code_generation.ipynb\">\n",
|
||||
" <img src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" alt=\"GitHub logo\">\n",
|
||||
" View on GitHub\n",
|
||||
" </a>\n",
|
||||
" </td>\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/deploy-notebook?download_url=https://raw.githubusercontent.com/googleapis/python-bigquery-dataframes/tree/main/notebooks/getting_started/bq_dataframes_llm_code_generation.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>"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "24743cf4a1e1"
|
||||
},
|
||||
"source": [
|
||||
"**_NOTE_**: This notebook has been tested in the following environment:\n",
|
||||
"\n",
|
||||
"* Python version = 3.10"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "tvgnzT1CKxrO"
|
||||
},
|
||||
"source": [
|
||||
"## Overview\n",
|
||||
"\n",
|
||||
"Use this notebook to walk through an example use case of generating sample code by using BigQuery DataFrames and its integration with Generative AI support on Vertex AI.\n",
|
||||
"\n",
|
||||
"Learn more about [BigQuery DataFrames](https://cloud.google.com/python/docs/reference/bigframes/latest)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "d975e698c9a4"
|
||||
},
|
||||
"source": [
|
||||
"### Objective\n",
|
||||
"\n",
|
||||
"In this tutorial, you create a CSV file containing sample code for calling a given set of APIs.\n",
|
||||
"\n",
|
||||
"The steps include:\n",
|
||||
"\n",
|
||||
"- Defining an LLM model in BigQuery DataFrames, specifically the [`text-bison` model of the PaLM API](https://cloud.google.com/vertex-ai/docs/generative-ai/model-reference/text), using `bigframes.ml.llm`.\n",
|
||||
"- Creating a DataFrame by reading in data from Cloud Storage.\n",
|
||||
"- Manipulating data in the DataFrame to build LLM prompts.\n",
|
||||
"- Sending DataFrame prompts to the LLM model using the `predict` method.\n",
|
||||
"- Creating and using a custom function to transform the output provided by the LLM model response.\n",
|
||||
"- Exporting the resulting transformed DataFrame as a CSV file."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "08d289fa873f"
|
||||
},
|
||||
"source": [
|
||||
"### Dataset\n",
|
||||
"\n",
|
||||
"This tutorial uses a dataset listing the names of various pandas DataFrame and Series APIs."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "aed92deeb4a0"
|
||||
},
|
||||
"source": [
|
||||
"### Costs\n",
|
||||
"\n",
|
||||
"This tutorial uses billable components of Google Cloud:\n",
|
||||
"\n",
|
||||
"* BigQuery\n",
|
||||
"* Generative AI support on Vertex AI\n",
|
||||
"* Cloud Functions\n",
|
||||
"\n",
|
||||
"Learn about [BigQuery compute pricing](https://cloud.google.com/bigquery/pricing#analysis_pricing_models),\n",
|
||||
"[Generative AI support on Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing#generative_ai_models), and [Cloud Functions pricing](https://cloud.google.com/functions/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/)\n",
|
||||
"to generate a cost estimate based on your projected usage."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "i7EUnXsZhAGF"
|
||||
},
|
||||
"source": [
|
||||
"## Installation\n",
|
||||
"\n",
|
||||
"Install the following packages, which are required to run this notebook:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "2b4ef9b72d43"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install bigframes --upgrade --quiet"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "BF1j6f9HApxa"
|
||||
},
|
||||
"source": [
|
||||
"## Before you begin\n",
|
||||
"\n",
|
||||
"Complete the tasks in this section to set up your environment."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Wbr2aVtFQBcg"
|
||||
},
|
||||
"source": [
|
||||
"### Set up your Google Cloud project\n",
|
||||
"\n",
|
||||
"**The following steps are required, regardless of your notebook environment.**\n",
|
||||
"\n",
|
||||
"1. [Select or create a Google Cloud project](https://console.cloud.google.com/cloud-resource-manager). When you first create an account, you get a $300 credit towards your compute/storage costs.\n",
|
||||
"\n",
|
||||
"2. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
|
||||
"\n",
|
||||
"3. [Click here](https://console.cloud.google.com/flows/enableapi?apiid=bigquery.googleapis.com,bigqueryconnection.googleapis.com,cloudfunctions.googleapis.com,run.googleapis.com,artifactregistry.googleapis.com,cloudbuild.googleapis.com,cloudresourcemanager.googleapis.com) to enable the following APIs:\n",
|
||||
"\n",
|
||||
" * BigQuery API\n",
|
||||
" * BigQuery Connection API\n",
|
||||
" * Cloud Functions API\n",
|
||||
" * Cloud Run API\n",
|
||||
" * Artifact Registry API\n",
|
||||
" * Cloud Build API\n",
|
||||
" * Cloud Resource Manager API\n",
|
||||
" * Vertex AI API\n",
|
||||
"\n",
|
||||
"4. If you are running this notebook locally, install the [Cloud SDK](https://cloud.google.com/sdk)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "WReHDGG5g0XY"
|
||||
},
|
||||
"source": [
|
||||
"#### Set your project ID\n",
|
||||
"\n",
|
||||
"If you don't know your project ID, try the following:\n",
|
||||
"* Run `gcloud config list`.\n",
|
||||
"* Run `gcloud projects list`.\n",
|
||||
"* See the support page: [Locate the project ID](https://support.google.com/googleapi/answer/7014113)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "oM1iC_MfAts1"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"PROJECT_ID = \"\" # @param {type:\"string\"}\n",
|
||||
"\n",
|
||||
"# Set the project id\n",
|
||||
"! gcloud config set project {PROJECT_ID}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "region"
|
||||
},
|
||||
"source": [
|
||||
"#### Set the region\n",
|
||||
"\n",
|
||||
"You can also change the `REGION` variable used by BigQuery. Learn more about [BigQuery regions](https://cloud.google.com/bigquery/docs/locations#supported_locations)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "eF-Twtc4XGem"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"REGION = \"US\" # @param {type: \"string\"}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "sBCra4QMA2wR"
|
||||
},
|
||||
"source": [
|
||||
"### Authenticate your Google Cloud account\n",
|
||||
"\n",
|
||||
"Depending on your Jupyter environment, you might have to manually authenticate. Follow the relevant instructions below."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "74ccc9e52986"
|
||||
},
|
||||
"source": [
|
||||
"**Vertex AI Workbench**\n",
|
||||
"\n",
|
||||
"Do nothing, you are already authenticated."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "de775a3773ba"
|
||||
},
|
||||
"source": [
|
||||
"**Local JupyterLab instance**\n",
|
||||
"\n",
|
||||
"Uncomment and run the following cell:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "254614fa0c46"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# ! gcloud auth login"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "ef21552ccea8"
|
||||
},
|
||||
"source": [
|
||||
"**Colab**\n",
|
||||
"\n",
|
||||
"Uncomment and run the following cell:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "603adbbf0532"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# from google.colab import auth\n",
|
||||
"# auth.authenticate_user()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "960505627ddf"
|
||||
},
|
||||
"source": [
|
||||
"### Import libraries"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "PyQmSRbKA8r-"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import bigframes.pandas as bf\n",
|
||||
"from google.cloud import bigquery_connection_v1 as bq_connection"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "init_aip:mbsdk,all"
|
||||
},
|
||||
"source": [
|
||||
"### Set BigQuery DataFrames options"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "NPPMuw2PXGeo"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"bf.options.bigquery.project = PROJECT_ID\n",
|
||||
"bf.options.bigquery.location = REGION"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "DTVtFlqeFbrU"
|
||||
},
|
||||
"source": [
|
||||
"If you want to reset the location of the created DataFrame or Series objects, reset the session by executing `bf.reset_session()`. After that, you can reuse `bf.options.bigquery.location` to specify another location."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "6eytf4xQHzcF"
|
||||
},
|
||||
"source": [
|
||||
"# Define the LLM model\n",
|
||||
"\n",
|
||||
"BigQuery DataFrames provides integration with [`text-bison` model of the PaLM API](https://cloud.google.com/vertex-ai/docs/generative-ai/model-reference/text) via Vertex AI.\n",
|
||||
"\n",
|
||||
"This section walks through a few steps required in order to use the model in your notebook."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "rS4VO1TGiO4G"
|
||||
},
|
||||
"source": [
|
||||
"## Create a BigQuery Cloud resource connection\n",
|
||||
"\n",
|
||||
"You need to create a [Cloud resource connection](https://cloud.google.com/bigquery/docs/create-cloud-resource-connection) to enable BigQuery DataFrames to interact with Vertex AI services."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "KFPjDM4LVh96"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"CONN_NAME = \"bqdf-llm\"\n",
|
||||
"\n",
|
||||
"client = bq_connection.ConnectionServiceClient()\n",
|
||||
"new_conn_parent = f\"projects/{PROJECT_ID}/locations/{REGION}\"\n",
|
||||
"exists_conn_parent = f\"projects/{PROJECT_ID}/locations/{REGION}/connections/{CONN_NAME}\"\n",
|
||||
"cloud_resource_properties = bq_connection.CloudResourceProperties({})\n",
|
||||
"\n",
|
||||
"try:\n",
|
||||
" request = client.get_connection(\n",
|
||||
" request=bq_connection.GetConnectionRequest(name=exists_conn_parent)\n",
|
||||
" )\n",
|
||||
" CONN_SERVICE_ACCOUNT = f\"serviceAccount:{request.cloud_resource.service_account_id}\"\n",
|
||||
"except Exception:\n",
|
||||
" connection = bq_connection.types.Connection(\n",
|
||||
" {\"friendly_name\": CONN_NAME, \"cloud_resource\": cloud_resource_properties}\n",
|
||||
" )\n",
|
||||
" request = bq_connection.CreateConnectionRequest(\n",
|
||||
" {\n",
|
||||
" \"parent\": new_conn_parent,\n",
|
||||
" \"connection_id\": CONN_NAME,\n",
|
||||
" \"connection\": connection,\n",
|
||||
" }\n",
|
||||
" )\n",
|
||||
" response = client.create_connection(request)\n",
|
||||
" CONN_SERVICE_ACCOUNT = (\n",
|
||||
" f\"serviceAccount:{response.cloud_resource.service_account_id}\"\n",
|
||||
" )\n",
|
||||
"print(CONN_SERVICE_ACCOUNT)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "W6l6Ol2biU9h"
|
||||
},
|
||||
"source": [
|
||||
"## Set permissions for the service account\n",
|
||||
"\n",
|
||||
"The resource connection service account requires certain project-level permissions:\n",
|
||||
" - `roles/aiplatform.user` and `roles/bigquery.connectionUser`: These roles are required for the connection to create a model definition using the LLM model in Vertex AI ([documentation](https://cloud.google.com/bigquery/docs/generate-text#give_the_service_account_access)).\n",
|
||||
" - `roles/run.invoker`: This role is required for the connection to have read-only access to Cloud Run services that back custom/remote functions ([documentation](https://cloud.google.com/bigquery/docs/remote-functions#grant_permission_on_function)).\n",
|
||||
"\n",
|
||||
"Set these permissions by running the following `gcloud` commands:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "d8wja24SVq6s"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!gcloud projects add-iam-policy-binding {PROJECT_ID} --condition=None --no-user-output-enabled --member={CONN_SERVICE_ACCOUNT} --role='roles/bigquery.connectionUser'\n",
|
||||
"!gcloud projects add-iam-policy-binding {PROJECT_ID} --condition=None --no-user-output-enabled --member={CONN_SERVICE_ACCOUNT} --role='roles/aiplatform.user'\n",
|
||||
"!gcloud projects add-iam-policy-binding {PROJECT_ID} --condition=None --no-user-output-enabled --member={CONN_SERVICE_ACCOUNT} --role='roles/run.invoker'"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "qUjT8nw-jIXp"
|
||||
},
|
||||
"source": [
|
||||
"## Define the model\n",
|
||||
"\n",
|
||||
"Use `bigframes.ml.llm` to define the model:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "sdjeXFwcHfl7"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from bigframes.ml.llm import PaLM2TextGenerator\n",
|
||||
"\n",
|
||||
"session = bf.get_global_session()\n",
|
||||
"connection = f\"{PROJECT_ID}.{REGION}.{CONN_NAME}\"\n",
|
||||
"model = PaLM2TextGenerator(session=session, connection_name=connection)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "GbW0oCnU1s1N"
|
||||
},
|
||||
"source": [
|
||||
"# Read data from Cloud Storage into BigQuery DataFrames\n",
|
||||
"\n",
|
||||
"You can create a BigQuery DataFrames DataFrame by reading data from any of the following locations:\n",
|
||||
"\n",
|
||||
"* A local data file\n",
|
||||
"* Data stored in a BigQuery table\n",
|
||||
"* A data file stored in Cloud Storage\n",
|
||||
"* An in-memory pandas DataFrame\n",
|
||||
"\n",
|
||||
"In this tutorial, you create BigQuery DataFrames DataFrames by reading two CSV files stored in Cloud Storage, one containing a list of DataFrame API names and one containing a list of Series API names."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "SchiTkQGIJog"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_api = bf.read_csv(\"gs://cloud-samples-data/vertex-ai/bigframe/df.csv\")\n",
|
||||
"series_api = bf.read_csv(\"gs://cloud-samples-data/vertex-ai/bigframe/series.csv\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "7OBjw2nmQY3-"
|
||||
},
|
||||
"source": [
|
||||
"Take a peek at a few rows of data for each file:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "QCqgVCIsGGuv"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_api.head(2)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "BGJnZbgEGS5-"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"series_api.head(2)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "m3ZJEsi7SUKV"
|
||||
},
|
||||
"source": [
|
||||
"# Generate code using the LLM model\n",
|
||||
"\n",
|
||||
"Prepare the prompts and send them to the LLM model for prediction."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "9EMAqR37AfLS"
|
||||
},
|
||||
"source": [
|
||||
"## Prompt design in BigQuery DataFrames\n",
|
||||
"\n",
|
||||
"Designing prompts for LLMs is a fast growing area and you can read more in [this documentation](https://cloud.google.com/vertex-ai/docs/generative-ai/learn/introduction-prompt-design).\n",
|
||||
"\n",
|
||||
"For this tutorial, you use a simple prompt to ask the LLM model for sample code for each of the API methods (or rows) from the last step's DataFrames. The output is the new DataFrames `df_prompt` and `series_prompt`, which contain the full prompt text."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "EDAaIwHpQCDZ"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_prompt_prefix = \"Generate Pandas sample code for DataFrame.\"\n",
|
||||
"series_prompt_prefix = \"Generate Pandas sample code for Series.\"\n",
|
||||
"\n",
|
||||
"df_prompt = df_prompt_prefix + df_api[\"API\"]\n",
|
||||
"series_prompt = series_prompt_prefix + series_api[\"API\"]\n",
|
||||
"\n",
|
||||
"df_prompt.head(2)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "rwPLjqW2Ajzh"
|
||||
},
|
||||
"source": [
|
||||
"## Make predictions using the LLM model\n",
|
||||
"\n",
|
||||
"Use the BigQuery DataFrames DataFrame containing the full prompt text as the input to the `predict` method. The `predict` method calls the LLM model and returns its generated text output back to two new BigQuery DataFrames DataFrames, `df_pred` and `series_pred`.\n",
|
||||
"\n",
|
||||
"Note: The predictions might take a few minutes to run."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "6i6HkFJZa8na"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_pred = model.predict(df_prompt.to_frame(), max_output_tokens=1024)\n",
|
||||
"series_pred = model.predict(series_prompt.to_frame(), max_output_tokens=1024)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "89cB8MW4UIdV"
|
||||
},
|
||||
"source": [
|
||||
"Once the predictions are processed, take a look at the sample output from the LLM, which provides code samples for the API names listed in the DataFrames dataset."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "9A2gw6hP_2nX"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(df_pred[\"ml_generate_text_llm_result\"].iloc[0])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Fx4lsNqMorJ-"
|
||||
},
|
||||
"source": [
|
||||
"# Manipulate LLM output using a remote function\n",
|
||||
"\n",
|
||||
"The output that the LLM provides often contains additional text beyond the code sample itself. Using BigQuery DataFrames, you can deploy custom Python functions that process and transform this output.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "d8L7SN03VByG"
|
||||
},
|
||||
"source": [
|
||||
"Running the cell below creates a custom function that you can use to process the LLM output data in two ways:\n",
|
||||
"1. Strip the LLM text output to include only the code block.\n",
|
||||
"2. Substitute `import pandas as pd` with `import bigframes.pandas as bf` so that the resulting code block works with BigQuery DataFrames."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "GskyyUQPowBT"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"@bf.remote_function([str], str, bigquery_connection=CONN_NAME)\n",
|
||||
"def extract_code(text: str):\n",
|
||||
" try:\n",
|
||||
" res = text[text.find(\"\\n\") + 1 : text.find(\"```\", 3)]\n",
|
||||
" res = res.replace(\"import pandas as pd\", \"import bigframes.pandas as bf\")\n",
|
||||
" if \"import bigframes.pandas as bf\" not in res:\n",
|
||||
" res = \"import bigframes.pandas as bf\\n\" + res\n",
|
||||
" return res\n",
|
||||
" except:\n",
|
||||
" return \"\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "hVQAoqBUOJQf"
|
||||
},
|
||||
"source": [
|
||||
"The custom function is deployed as a Cloud Function, and then integrated with BigQuery as a [remote function](https://cloud.google.com/bigquery/docs/remote-functions). Save both of the function names so that you can clean them up at the end of this notebook."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "PBlp-C-DOHRO"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"CLOUD_FUNCTION_NAME = format(extract_code.bigframes_cloud_function)\n",
|
||||
"print(\"Cloud Function Name \" + CLOUD_FUNCTION_NAME)\n",
|
||||
"REMOTE_FUNCTION_NAME = format(extract_code.bigframes_remote_function)\n",
|
||||
"print(\"Remote Function Name \" + REMOTE_FUNCTION_NAME)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "4FEucaiqVs3H"
|
||||
},
|
||||
"source": [
|
||||
"Apply the custom function to each LLM output DataFrame to get the processed results:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "bsQ9cmoWo0Ps"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_code = df_pred.assign(\n",
|
||||
" code=df_pred[\"ml_generate_text_llm_result\"].apply(extract_code)\n",
|
||||
")\n",
|
||||
"series_code = series_pred.assign(\n",
|
||||
" code=series_pred[\"ml_generate_text_llm_result\"].apply(extract_code)\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "ujQVVuhfWA3y"
|
||||
},
|
||||
"source": [
|
||||
"You can see the differences by inspecting the first row of data:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "7yWzjhGy_zcy"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(df_code[\"code\"].iloc[0])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "GTRdUw-Ro5R1"
|
||||
},
|
||||
"source": [
|
||||
"# Save the results to Cloud Storage\n",
|
||||
"\n",
|
||||
"BigQuery DataFrames lets you save a BigQuery DataFrames DataFrame as a CSV file in Cloud Storage for further use. Try that now with your processed LLM output data."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "9DQ7eiQxPTi3"
|
||||
},
|
||||
"source": [
|
||||
"Create a new Cloud Storage bucket with a unique name:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "-J5LHgS6LLZ0"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import uuid\n",
|
||||
"\n",
|
||||
"BUCKET_ID = \"code-samples-\" + str(uuid.uuid1())\n",
|
||||
"\n",
|
||||
"!gsutil mb gs://{BUCKET_ID}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "tyxZXj0UPYUv"
|
||||
},
|
||||
"source": [
|
||||
"Use `to_csv` to write each BigQuery DataFrames DataFrame as a CSV file in the Cloud Storage bucket:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "Zs_b5L-4IvER"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_code[[\"code\"]].to_csv(f\"gs://{BUCKET_ID}/df_code*.csv\")\n",
|
||||
"series_code[[\"code\"]].to_csv(f\"gs://{BUCKET_ID}/series_code*.csv\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "UDBtDlrTuuh8"
|
||||
},
|
||||
"source": [
|
||||
"You can navigate to the Cloud Storage bucket browser to download the two files and view them.\n",
|
||||
"\n",
|
||||
"Run the following cell, and then follow the link to your Cloud Storage bucket browser:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "PspCXu-qu_ND"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(f\"https://console.developers.google.com/storage/browser/{BUCKET_ID}/\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "RGSvUk48RK20"
|
||||
},
|
||||
"source": [
|
||||
"# Summary and next steps\n",
|
||||
"\n",
|
||||
"You've used BigQuery DataFrames' integration with LLM models (`bigframes.ml.llm`) to generate code samples, and have tranformed LLM output by creating and using a custom function in BigQuery DataFrames.\n",
|
||||
"\n",
|
||||
"Learn more about BigQuery DataFrames in the [documentation](https://cloud.google.com/python/docs/reference/bigframes/latest) and find more sample notebooks in the [GitHub repo](https://github.com/googleapis/python-bigquery-dataframes/tree/main/notebooks)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "TpV-iwP9qw9c"
|
||||
},
|
||||
"source": [
|
||||
"## Cleaning up\n",
|
||||
"\n",
|
||||
"To clean up all Google Cloud resources used in this project, you can [delete the Google Cloud\n",
|
||||
"project](https://cloud.google.com/resource-manager/docs/creating-managing-projects#shutting_down_projects) you used for the tutorial.\n",
|
||||
"\n",
|
||||
"Otherwise, you can uncomment the remaining cells and run them to delete the individual resources you created in this tutorial:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "yw7A461XLjvW"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# # Delete the BigQuery Connection\n",
|
||||
"# from google.cloud import bigquery_connection_v1 as bq_connection\n",
|
||||
"# client = bq_connection.ConnectionServiceClient()\n",
|
||||
"# CONNECTION_ID = f\"projects/{PROJECT_ID}/locations/{REGION}/connections/{CONN_NAME}\"\n",
|
||||
"# client.delete_connection(name=CONNECTION_ID)\n",
|
||||
"# print(f\"Deleted connection '{CONNECTION_ID}'.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "sx_vKniMq9ZX"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# # Delete the Cloud Function\n",
|
||||
"# ! gcloud functions delete {CLOUD_FUNCTION_NAME} --quiet\n",
|
||||
"# # Delete the Remote Function\n",
|
||||
"# REMOTE_FUNCTION_NAME = REMOTE_FUNCTION_NAME.replace(PROJECT_ID + \".\", \"\")\n",
|
||||
"# ! bq rm --routine --force=true {REMOTE_FUNCTION_NAME}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "iQFo6OUBLmi3"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# # Delete the Google Cloud Storage bucket and files\n",
|
||||
"# ! gsutil rm -r gs://{BUCKET_ID}\n",
|
||||
"# print(f\"Deleted bucket '{BUCKET_ID}'.\")"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"name": "bq_dataframes_llm_code_generation.ipynb",
|
||||
"toc_visible": true
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"name": "python3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
@@ -0,0 +1,989 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "ur8xi4C7S06n"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Copyright 2023 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",
|
||||
"# You may obtain a copy of the License at\n",
|
||||
"#\n",
|
||||
"# https://www.apache.org/licenses/LICENSE-2.0\n",
|
||||
"#\n",
|
||||
"# Unless required by applicable law or agreed to in writing, software\n",
|
||||
"# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
|
||||
"# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
|
||||
"# See the License for the specific language governing permissions and\n",
|
||||
"# limitations under the License."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "JAPoU8Sm5E6e"
|
||||
},
|
||||
"source": [
|
||||
"# BigQuery DataFrames ML: Drug Name Generation\n",
|
||||
"\n",
|
||||
"<table align=\"left\">\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://colab.research.google.com/github/googleapis/python-bigquery-dataframes/blob/main/notebooks/generative_ai/bq_dataframes_ml_drug_name_generation.ipynb\">\n",
|
||||
" <img src=\"https://cloud.google.com/ml-engine/images/colab-logo-32px.png\" alt=\"Colab logo\"> Run in Colab\n",
|
||||
" </a>\n",
|
||||
" </td>\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://github.com/googleapis/python-bigquery-dataframes/blob/main/notebooks/generative_ai/bq_dataframes_ml_drug_name_generation.ipynb\">\n",
|
||||
" <img src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" alt=\"GitHub logo\">\n",
|
||||
" View on GitHub\n",
|
||||
" </a>\n",
|
||||
" </td>\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/deploy-notebook?download_url=https://raw.githubusercontent.com/googleapis/python-bigquery-dataframes/blob/main/notebooks/generative_ai/bq_dataframes_ml_drug_name_generation.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>"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "24743cf4a1e1"
|
||||
},
|
||||
"source": [
|
||||
"**_NOTE_**: This notebook has been tested in the following environment:\n",
|
||||
"\n",
|
||||
"* Python version = 3.9"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "tvgnzT1CKxrO"
|
||||
},
|
||||
"source": [
|
||||
"## Overview\n",
|
||||
"\n",
|
||||
"The goal of this notebook is to demonstrate an enterprise generative AI use case. A marketing user can provide information about a new pharmaceutical drug and its generic name, and receive ideas on marketing-oriented brand names for that drug.\n",
|
||||
"\n",
|
||||
"Learn more about [BigQuery DataFrames](https://cloud.google.com/bigquery/docs/dataframes-quickstart)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "d975e698c9a4"
|
||||
},
|
||||
"source": [
|
||||
"### Objective\n",
|
||||
"\n",
|
||||
"In this tutorial, you learn about Generative AI concepts such as prompting and few-shot learning, as well as how to use BigFrames ML for performing these tasks simply using an intuitive dataframe API.\n",
|
||||
"\n",
|
||||
"The steps performed include:\n",
|
||||
"\n",
|
||||
"1. Ask the user for the generic name and usage for the drug.\n",
|
||||
"1. Use `bigframes` to query the FDA dataset of over 100,000 drugs, filtered on the brand name, generic name, and indications & usage columns.\n",
|
||||
"1. Filter this dataset to find prototypical brand names that can be used as examples in prompt tuning.\n",
|
||||
"1. Create a prompt with the user input, general instructions, examples and counter-examples for the desired brand name.\n",
|
||||
"1. Use the `bigframes.ml.llm.PaLM2TextGenerator` to generate choices of brand names."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "08d289fa873f"
|
||||
},
|
||||
"source": [
|
||||
"### Dataset\n",
|
||||
"\n",
|
||||
"This notebook uses the [FDA dataset](https://cloud.google.com/blog/topics/healthcare-life-sciences/fda-mystudies-comes-to-google-cloud) available at [`bigquery-public-data.fda_drug`](https://console.cloud.google.com/bigquery?ws=!1m4!1m3!3m2!1sbigquery-public-data!2sfda_drug)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "aed92deeb4a0"
|
||||
},
|
||||
"source": [
|
||||
"### Costs\n",
|
||||
"\n",
|
||||
"This tutorial uses billable components of Google Cloud:\n",
|
||||
"\n",
|
||||
"* BigQuery (compute)\n",
|
||||
"* BigQuery ML\n",
|
||||
"\n",
|
||||
"Learn about [BigQuery compute pricing](https://cloud.google.com/bigquery/pricing#analysis_pricing_models),\n",
|
||||
"and [BigQuery ML pricing](https://cloud.google.com/bigquery/pricing#bqml),\n",
|
||||
"and use the [Pricing Calculator](https://cloud.google.com/products/calculator/)\n",
|
||||
"to generate a cost estimate based on your projected usage."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "i7EUnXsZhAGF"
|
||||
},
|
||||
"source": [
|
||||
"## Installation\n",
|
||||
"\n",
|
||||
"Install the following packages required to execute this notebook."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "2b4ef9b72d43"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install -U --quiet bigframes"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "58707a750154"
|
||||
},
|
||||
"source": [
|
||||
"### Colab only: Uncomment the following cell to restart the kernel."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "f200f10a1da3"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# # Automatically restart kernel after installs so that your environment can access the new packages\n",
|
||||
"# import IPython\n",
|
||||
"\n",
|
||||
"# app = IPython.Application.instance()\n",
|
||||
"# app.kernel.do_shutdown(True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "960505627ddf"
|
||||
},
|
||||
"source": [
|
||||
"### Import libraries"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "PyQmSRbKA8r-"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import bigframes.pandas as bpd\n",
|
||||
"from bigframes.ml.llm import PaLM2TextGenerator\n",
|
||||
"from google.cloud import bigquery_connection_v1 as bq_connection\n",
|
||||
"from IPython.display import Markdown"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "sBCra4QMA2wR"
|
||||
},
|
||||
"source": [
|
||||
"### Authenticate your Google Cloud account\n",
|
||||
"\n",
|
||||
"Depending on your Jupyter environment, you may have to manually authenticate. Follow the relevant instructions below."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "74ccc9e52986"
|
||||
},
|
||||
"source": [
|
||||
"**1. Vertex AI Workbench**\n",
|
||||
"* Do nothing as you are already authenticated."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "de775a3773ba"
|
||||
},
|
||||
"source": [
|
||||
"**2. Local JupyterLab instance, uncomment and run:**"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "254614fa0c46"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# ! gcloud auth login"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "ef21552ccea8"
|
||||
},
|
||||
"source": [
|
||||
"**3. Colab, uncomment and run:**"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "603adbbf0532"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# from google.colab import auth\n",
|
||||
"\n",
|
||||
"# auth.authenticate_user()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "BF1j6f9HApxa"
|
||||
},
|
||||
"source": [
|
||||
"## Before you begin\n",
|
||||
"\n",
|
||||
"### Set up your Google Cloud project\n",
|
||||
"\n",
|
||||
"**The following steps are required, regardless of your notebook environment.**\n",
|
||||
"\n",
|
||||
"1. [Select or create a Google Cloud project](https://console.cloud.google.com/cloud-resource-manager). When you first create an account, you get a $300 free credit towards your compute/storage costs.\n",
|
||||
"\n",
|
||||
"2. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
|
||||
"\n",
|
||||
"3. [Enable the BigQuery API](https://console.cloud.google.com/flows/enableapi?apiid=bigquery.googleapis.com).\n",
|
||||
"\n",
|
||||
"4. If you are running this notebook locally, you need to install the [Cloud SDK](https://cloud.google.com/sdk)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "WReHDGG5g0XY"
|
||||
},
|
||||
"source": [
|
||||
"#### Set your project ID\n",
|
||||
"\n",
|
||||
"**If you don't know your project ID**, try the following:\n",
|
||||
"* Run `gcloud config list`.\n",
|
||||
"* Run `gcloud projects list`.\n",
|
||||
"* See the support page: [Locate the project ID](https://support.google.com/googleapi/answer/7014113)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "oM1iC_MfAts1"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"PROJECT_ID = \"<your-project-id>\" # @param {type:\"string\"}\n",
|
||||
"\n",
|
||||
"# Set the project id\n",
|
||||
"! gcloud config set project {PROJECT_ID}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "evsJaAj5te0X"
|
||||
},
|
||||
"source": [
|
||||
"#### BigFrames configuration\n",
|
||||
"\n",
|
||||
"Next, we will specify a [BigQuery connection](https://cloud.google.com/bigquery/docs/working-with-connections). If you already have a connection, you can simplify provide the name and skip the following creation steps.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "G1vVsPiMsL2X"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Please fill in these values.\n",
|
||||
"LOCATION = \"us\" # @param {type:\"string\"}\n",
|
||||
"CONNECTION = \"<your-connection>\" # @param {type:\"string\"}\n",
|
||||
"\n",
|
||||
"connection_name = f\"{PROJECT_ID}.{LOCATION}.{CONNECTION}\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "WGS_TzhWlPBN"
|
||||
},
|
||||
"source": [
|
||||
"We will now try to use the provided connection, and if it doesn't exist, create a new one. We will also print the service account used."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "56Hw42m6kFrj"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Initialize client and set request parameters\n",
|
||||
"client = bq_connection.ConnectionServiceClient()\n",
|
||||
"new_conn_parent = f\"projects/{PROJECT_ID}/locations/{LOCATION}\"\n",
|
||||
"exists_conn_parent = (\n",
|
||||
" f\"projects/{PROJECT_ID}/locations/{LOCATION}/connections/{CONNECTION}\"\n",
|
||||
")\n",
|
||||
"cloud_resource_properties = bq_connection.CloudResourceProperties({})\n",
|
||||
"\n",
|
||||
"# Try to connect using provided connection\n",
|
||||
"try:\n",
|
||||
" request = client.get_connection(\n",
|
||||
" request=bq_connection.GetConnectionRequest(name=exists_conn_parent)\n",
|
||||
" )\n",
|
||||
" CONN_SERVICE_ACCOUNT = f\"serviceAccount:{request.cloud_resource.service_account_id}\"\n",
|
||||
"# Create a new connection on error\n",
|
||||
"except Exception:\n",
|
||||
" connection = bq_connection.types.Connection(\n",
|
||||
" {\"friendly_name\": CONNECTION, \"cloud_resource\": cloud_resource_properties}\n",
|
||||
" )\n",
|
||||
" request = bq_connection.CreateConnectionRequest(\n",
|
||||
" {\n",
|
||||
" \"parent\": new_conn_parent,\n",
|
||||
" \"connection_id\": CONNECTION,\n",
|
||||
" \"connection\": connection,\n",
|
||||
" }\n",
|
||||
" )\n",
|
||||
" response = client.create_connection(request)\n",
|
||||
" CONN_SERVICE_ACCOUNT = (\n",
|
||||
" f\"serviceAccount:{response.cloud_resource.service_account_id}\"\n",
|
||||
" )\n",
|
||||
"# Set service account permissions\n",
|
||||
"!gcloud projects add-iam-policy-binding {PROJECT_ID} --condition=None --no-user-output-enabled --member={CONN_SERVICE_ACCOUNT} --role='roles/bigquery.connectionUser'\n",
|
||||
"!gcloud projects add-iam-policy-binding {PROJECT_ID} --condition=None --no-user-output-enabled --member={CONN_SERVICE_ACCOUNT} --role='roles/aiplatform.user'\n",
|
||||
"!gcloud projects add-iam-policy-binding {PROJECT_ID} --condition=None --no-user-output-enabled --member={CONN_SERVICE_ACCOUNT} --role='roles/run.invoker'\n",
|
||||
"\n",
|
||||
"print(CONN_SERVICE_ACCOUNT)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "init_aip:mbsdk,all"
|
||||
},
|
||||
"source": [
|
||||
"### Initialize BigFrames client\n",
|
||||
"\n",
|
||||
"Here, we set the project configuration based on the provided parameters."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "OCccLirpkSRz"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"bpd.options.bigquery.project = PROJECT_ID\n",
|
||||
"bpd.options.bigquery.location = LOCATION"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "m8UCEtX9uLn6"
|
||||
},
|
||||
"source": [
|
||||
"## Generate a name\n",
|
||||
"\n",
|
||||
"Let's start with entering a generic name and description of the drug."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "oxphj2gnuKou"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"GENERIC_NAME = \"Entropofloxacin\" # @param {type:\"string\"}\n",
|
||||
"USAGE = \"Entropofloxacin is a fluoroquinolone antibiotic that is used to treat a variety of bacterial infections, including: pneumonia, streptococcus infections, salmonella infections, escherichia coli infections, and pseudomonas aeruginosa infections It is taken by mouth or by injection. The dosage and frequency of administration will vary depending on the type of infection being treated. It should be taken for the full course of treatment, even if symptoms improve after a few days. Stopping the medication early may increase the risk of the infection coming back.\" # @param {type:\"string\"}\n",
|
||||
"NUM_NAMES = 10 # @param {type:\"integer\"}\n",
|
||||
"TEMPERATURE = 0.5 # @param {type: \"number\"}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "1q-vlbalzu1Q"
|
||||
},
|
||||
"source": [
|
||||
"We can now create a prompt string, and populate it with the name and description."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "0knz5ZWMzed-"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"zero_shot_prompt = f\"\"\"Provide {NUM_NAMES} unique and modern brand names in Markdown bullet point format. Do not provide any additional explanation.\n",
|
||||
"\n",
|
||||
"Be creative with the brand names. Don't use English words directly; use variants or invented words.\n",
|
||||
"\n",
|
||||
"The generic name is: {GENERIC_NAME}\n",
|
||||
"\n",
|
||||
"The indications and usage are: {USAGE}.\"\"\"\n",
|
||||
"\n",
|
||||
"print(zero_shot_prompt)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "LCRE2L720f5y"
|
||||
},
|
||||
"source": [
|
||||
"Next, let's create a helper function to predict with our model. It will take a string input, and add it to a temporary BigFrames `DataFrame`. It will also return the string extracted from the response `DataFrame`."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "LB3xgDroIxlx"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def predict(prompt: str, temperature: float = TEMPERATURE) -> str:\n",
|
||||
" # Create dataframe\n",
|
||||
" input = bpd.DataFrame(\n",
|
||||
" {\n",
|
||||
" \"prompt\": [prompt],\n",
|
||||
" }\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" # Return response\n",
|
||||
" return model.predict(input, temperature).ml_generate_text_llm_result.iloc[0]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "b1ZapNZsJW2p"
|
||||
},
|
||||
"source": [
|
||||
"We can now initialize the model, and get a response to our prompt!"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "UW2fQ2k5Hsic"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Get BigFrames session\n",
|
||||
"session = bpd.get_global_session()\n",
|
||||
"\n",
|
||||
"# Define the model\n",
|
||||
"model = PaLM2TextGenerator(session=session, connection_name=connection_name)\n",
|
||||
"\n",
|
||||
"# Invoke LLM with prompt\n",
|
||||
"response = predict(zero_shot_prompt)\n",
|
||||
"\n",
|
||||
"# Print results as Markdown\n",
|
||||
"Markdown(response)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "o3yIhHV2jsUT"
|
||||
},
|
||||
"source": [
|
||||
"We're off to a great start! Let's see if we can refine our response."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "mBroUzWS8xOL"
|
||||
},
|
||||
"source": [
|
||||
"## Few-shot learning\n",
|
||||
"\n",
|
||||
"Let's try using [few-shot learning](https://paperswithcode.com/task/few-shot-learning). We will provide a few examples of what we're looking for along with our prompt.\n",
|
||||
"\n",
|
||||
"Our prompt will consist of 3 parts:\n",
|
||||
"* General instructions (e.g. generate $n$ brand names)\n",
|
||||
"* Multiple examples\n",
|
||||
"* Information about the drug we'd like to generate a name for\n",
|
||||
"\n",
|
||||
"Let's walk through how to construct this prompt.\n",
|
||||
"\n",
|
||||
"Our first step will be to define how many examples we want to provide in the prompt."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "MXdI78SOElyt"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Specify number of examples to include\n",
|
||||
"\n",
|
||||
"NUM_EXAMPLES = 3 # @param {type:\"integer\"}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "U8w4puVM_892"
|
||||
},
|
||||
"source": [
|
||||
"Next, let's define a prefix that will set the overall context."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "aQ2iscnhF2cx"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"prefix_prompt = f\"\"\"Provide {NUM_NAMES} unique and modern brand names in Markdown bullet point format, related to the drug at the bottom of this prompt.\n",
|
||||
"\n",
|
||||
"Be creative with the brand names. Don't use English words directly; use variants or invented words.\n",
|
||||
"\n",
|
||||
"First, we will provide {NUM_EXAMPLES} examples to help with your thought process.\n",
|
||||
"\n",
|
||||
"Then, we will provide the generic name and usage for the drug we'd like you to generate brand names for.\n",
|
||||
"\"\"\"\n",
|
||||
"\n",
|
||||
"print(prefix_prompt)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "VI0Spv-axN7d"
|
||||
},
|
||||
"source": [
|
||||
"Our next step will be to include examples into the prompt.\n",
|
||||
"\n",
|
||||
"We will start out by retrieving the raw data for the examples, by querying the BigQuery public dataset."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "IoO_Bp8wA07N"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Query 3 columns of interest from drug label dataset\n",
|
||||
"df = bpd.read_gbq(\n",
|
||||
" \"bigquery-public-data.fda_drug.drug_label\",\n",
|
||||
" col_order=[\"openfda_generic_name\", \"openfda_brand_name\", \"indications_and_usage\"],\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Exclude any rows with missing data\n",
|
||||
"df = df.dropna()\n",
|
||||
"\n",
|
||||
"# Drop duplicate rows\n",
|
||||
"df = df.drop_duplicates()\n",
|
||||
"\n",
|
||||
"# Print values\n",
|
||||
"df.head()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "W5kOtbNGBTI2"
|
||||
},
|
||||
"source": [
|
||||
"Let's now filter the results to remove atypical names."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "95WDe2eCCeLx"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Remove names with spaces\n",
|
||||
"df = df[df[\"openfda_brand_name\"].str.find(\" \") == -1]\n",
|
||||
"\n",
|
||||
"# Remove names with 5 or fewer characters\n",
|
||||
"df = df[df[\"openfda_brand_name\"].str.len() > 5]\n",
|
||||
"\n",
|
||||
"# Remove names where the generic and brand name match (case-insensitive)\n",
|
||||
"df = df[df[\"openfda_generic_name\"].str.lower() != df[\"openfda_brand_name\"].str.lower()]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "FZD89ep4EyYc"
|
||||
},
|
||||
"source": [
|
||||
"Let's take `NUM_EXAMPLES` samples to include in the prompt."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "2ohZYg7QEyJV"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Take a sample and convert to a Pandas dataframe for local usage.\n",
|
||||
"df_examples = df.sample(NUM_EXAMPLES, random_state=3).to_pandas()\n",
|
||||
"\n",
|
||||
"df_examples"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "J-Qa1_SCImXy"
|
||||
},
|
||||
"source": [
|
||||
"Let's now convert the data to a JSON structure, to enable embedding into a prompt. For consistency, we'll capitalize each example brand name."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "PcJdSaw0EGcW"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"examples = [\n",
|
||||
" {\n",
|
||||
" \"brand_name\": brand_name.capitalize(),\n",
|
||||
" \"generic_name\": generic_name,\n",
|
||||
" \"usage\": usage,\n",
|
||||
" }\n",
|
||||
" for brand_name, generic_name, usage in zip(\n",
|
||||
" df_examples[\"openfda_brand_name\"],\n",
|
||||
" df_examples[\"openfda_generic_name\"],\n",
|
||||
" df_examples[\"indications_and_usage\"],\n",
|
||||
" )\n",
|
||||
"]\n",
|
||||
"\n",
|
||||
"print(examples)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "oU4mb1Dwgq64"
|
||||
},
|
||||
"source": [
|
||||
"We'll create a prompt template for each example, and view the first one."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "kzAVsF6wJ93S"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"example_prompt = \"\"\n",
|
||||
"for example in examples:\n",
|
||||
" example_prompt += f\"Generic name: {example['generic_name']}\\nUsage: {example['usage']}\\nBrand name: {example['brand_name']}\\n\\n\"\n",
|
||||
"\n",
|
||||
"example_prompt"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "kbV2X1CXAyLV"
|
||||
},
|
||||
"source": [
|
||||
"Finally, we can create a suffix to our prompt. This will contain the generic name of the drug, its usage, ending with a request for brand names."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "OYp6W_XfHTlo"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"suffix_prompt = f\"\"\"Generic name: {GENERIC_NAME}\n",
|
||||
"Usage: {USAGE}\n",
|
||||
"Brand names:\"\"\"\n",
|
||||
"\n",
|
||||
"print(suffix_prompt)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "RiaisW1nihJP"
|
||||
},
|
||||
"source": [
|
||||
"Let's pull it altogether into a few shot prompt."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "99xdU7l8C1h8"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Define the prompt\n",
|
||||
"few_shot_prompt = prefix_prompt + example_prompt + suffix_prompt\n",
|
||||
"\n",
|
||||
"# Print the prompt\n",
|
||||
"print(few_shot_prompt)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "nbUWdHtfitWn"
|
||||
},
|
||||
"source": [
|
||||
"Now, let's pass our prompt to the LLM, and get a response!"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "d4ODRJdvLhlQ"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"response = predict(few_shot_prompt)\n",
|
||||
"\n",
|
||||
"Markdown(response)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "pFakjrTElOBs"
|
||||
},
|
||||
"source": [
|
||||
"# Bulk generation\n",
|
||||
"\n",
|
||||
"Let's take these experiments to the next level by generating many names in bulk. We'll see how to leverage BigFrames at scale!\n",
|
||||
"\n",
|
||||
"We can start by finding drugs that are missing brand names. There are approximately 4,000 drugs that meet this criteria. We'll put a limit of 100 in this notebook."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "8eAutS41mx6U"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Query 3 columns of interest from drug label dataset\n",
|
||||
"df_missing = bpd.read_gbq(\n",
|
||||
" \"bigquery-public-data.fda_drug.drug_label\",\n",
|
||||
" col_order=[\"openfda_generic_name\", \"openfda_brand_name\", \"indications_and_usage\"],\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Exclude any rows with missing data\n",
|
||||
"df_missing = df_missing.dropna()\n",
|
||||
"\n",
|
||||
"# Include rows in which openfda_brand_name equals openfda_generic_name\n",
|
||||
"df_missing = df_missing[\n",
|
||||
" df_missing[\"openfda_generic_name\"] == df_missing[\"openfda_brand_name\"]\n",
|
||||
"]\n",
|
||||
"\n",
|
||||
"# Limit the number of rows for demonstration purposes\n",
|
||||
"df_missing = df_missing.head(100)\n",
|
||||
"\n",
|
||||
"# Print values\n",
|
||||
"df_missing.head()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Fm6L8S7eVnCI"
|
||||
},
|
||||
"source": [
|
||||
"We will create a column `prompt` with a customized prompt for each row."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "19TvGN1PVmVX"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_missing[\"prompt\"] = (\n",
|
||||
" \"Provide a unique and modern brand name related to this pharmaceutical drug.\"\n",
|
||||
" + \"Don't use English words directly; use variants or invented words. The generic name is: \"\n",
|
||||
" + df_missing[\"openfda_generic_name\"]\n",
|
||||
" + \". The indications and usage are: \"\n",
|
||||
" + df_missing[\"indications_and_usage\"]\n",
|
||||
" + \".\"\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "njxwBvCKgMPE"
|
||||
},
|
||||
"source": [
|
||||
"We'll create a new helper method, `batch_predict()` and query the LLM. The job may take a couple minutes to execute."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "tiSHa5B4aFhw"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def batch_predict(\n",
|
||||
" input: bpd.DataFrame, temperature: float = TEMPERATURE\n",
|
||||
") -> bpd.DataFrame:\n",
|
||||
" return model.predict(input, temperature).ml_generate_text_llm_result\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"response = batch_predict(df_missing[\"prompt\"])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "K5a2nHdLgZEj"
|
||||
},
|
||||
"source": [
|
||||
"Let's check the results for one of our responses!"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "TnizdeqBdbZj"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Pick a sample\n",
|
||||
"k = 0\n",
|
||||
"\n",
|
||||
"# Gather the prompt and response details\n",
|
||||
"prompt_generic = df_missing[\"openfda_generic_name\"][k].iloc[0]\n",
|
||||
"prompt_usage = df_missing[\"indications_and_usage\"][k].iloc[0]\n",
|
||||
"response_str = response[k].iloc[0]\n",
|
||||
"\n",
|
||||
"# Print details\n",
|
||||
"print(f\"Generic name: {prompt_generic}\")\n",
|
||||
"print(f\"Brand name: {prompt_usage}\")\n",
|
||||
"print(f\"Response: {response_str}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "W4MviwyMI-Qh"
|
||||
},
|
||||
"source": [
|
||||
"Congratulations! You have learned how to use generative AI to jumpstart the creative process.\n",
|
||||
"\n",
|
||||
"You've also seen how BigFrames can manage each step of the process, including gathering data, data manipulation, and querying the LLM."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Bys6--dVmq7R"
|
||||
},
|
||||
"source": [
|
||||
"## Cleaning up\n",
|
||||
"\n",
|
||||
"To clean up all Google Cloud resources used in this project, you can [delete the Google Cloud\n",
|
||||
"project](https://cloud.google.com/resource-manager/docs/creating-managing-projects#shutting_down_projects) you used for the tutorial.\n",
|
||||
"\n",
|
||||
"Otherwise, you can uncomment the remaining cells and run them to delete the individual resources you created in this tutorial:"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "cIODjOLump_-"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Delete the BigQuery Connection\n",
|
||||
"from google.cloud import bigquery_connection_v1 as bq_connection\n",
|
||||
"\n",
|
||||
"client = bq_connection.ConnectionServiceClient()\n",
|
||||
"CONNECTION_ID = f\"projects/{PROJECT_ID}/locations/{LOCATION}/connections/{CONNECTION}\"\n",
|
||||
"client.delete_connection(name=CONNECTION_ID)\n",
|
||||
"print(f\"Deleted connection {CONNECTION_ID}.\")"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"name": "bq_dataframes_ml_drug_name_generation.ipynb",
|
||||
"toc_visible": true
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"name": "python3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
@@ -0,0 +1,363 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "view-in-github"
|
||||
},
|
||||
"source": [
|
||||
"<a href=\"https://colab.research.google.com/github/xqr-g/vertex-ai-samples/blob/main/notebooks/community/generative_ai/text_embedding_api_cloud_next_new_models.ipynb\" target=\"_parent\"><img src=\"https://colab.research.google.com/assets/colab-badge.svg\" alt=\"Open In Colab\"/></a>"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "KSP1duKDeaDR"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Copyright 2023 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",
|
||||
"# You may obtain a copy of the License at\n",
|
||||
"#\n",
|
||||
"# https://www.apache.org/licenses/LICENSE-2.0\n",
|
||||
"#\n",
|
||||
"# Unless required by applicable law or agreed to in writing, software\n",
|
||||
"# distributed under the License is distributed on an \"AS IS\" BASIS,\n",
|
||||
"# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n",
|
||||
"# See the License for the specific language governing permissions and\n",
|
||||
"# limitations under the License."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "JAPoU8Sm5E6e"
|
||||
},
|
||||
"source": [
|
||||
"# Cloud Next Embedding models\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"<table align=\"left\">\n",
|
||||
"\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://colab.research.google.com/github/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/generative_ai/text_embedding_api_cloud_next_new_models.ipynb\">\n",
|
||||
" <img src=\"https://cloud.google.com/ml-engine/images/colab-logo-32px.png\" alt=\"Colab logo\"> Run in Colab\n",
|
||||
" </a>\n",
|
||||
" </td>\n",
|
||||
" <td>\n",
|
||||
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/generative_ai/text_embedding_api_cloud_next_new_models.ipynb\">\n",
|
||||
" <img src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" alt=\"GitHub logo\">\n",
|
||||
" View on GitHub\n",
|
||||
" </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/generative_ai/text_embedding_api_cloud_next_new_models.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>"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "24743cf4a1e1"
|
||||
},
|
||||
"source": [
|
||||
"**_NOTE_**: This notebook has been tested in the following environment:\n",
|
||||
"\n",
|
||||
"* Python version = 3.10"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "tvgnzT1CKxrO"
|
||||
},
|
||||
"source": [
|
||||
"## Overview\n",
|
||||
"\n",
|
||||
"This colab is used as a code example for how to call our newly released text embedding models (textembedding-gecko@latest and textembedding-gecko-multilingual@latest).\n",
|
||||
"\n",
|
||||
"Learn more about [text embedding api](https://cloud.google.com/vertex-ai/docs/generative-ai/embeddings/get-text-embeddings).\n",
|
||||
"\n",
|
||||
"This tutorial uses the following Google Cloud ML services and resources:\n",
|
||||
"- Vertex LLM SDK\n",
|
||||
"\n",
|
||||
"The steps performed include:\n",
|
||||
"- Installation and imports\n",
|
||||
"- Generate embeddings\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "aed92deeb4a0"
|
||||
},
|
||||
"source": [
|
||||
"### Costs\n",
|
||||
"\n",
|
||||
"This tutorial uses billable components of Google Cloud:\n",
|
||||
"\n",
|
||||
"* Vertex AI\n",
|
||||
"\n",
|
||||
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing),\n",
|
||||
"and use the [Pricing Calculator](https://cloud.google.com/products/calculator/)\n",
|
||||
"to generate a cost estimate based on your projected usage."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "BF1j6f9HApxa"
|
||||
},
|
||||
"source": [
|
||||
"## Before you begin\n",
|
||||
"\n",
|
||||
"### Set up your Google Cloud project\n",
|
||||
"\n",
|
||||
"**The following steps are required, regardless of your notebook environment.**\n",
|
||||
"\n",
|
||||
"1. [Select or create a Google Cloud project](https://console.cloud.google.com/cloud-resource-manager). When you first create an account, you get a $300 free credit towards your compute/storage costs.\n",
|
||||
"\n",
|
||||
"2. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
|
||||
"\n",
|
||||
"3. [Enable the Vertex AI API](https://console.cloud.google.com/flows/enableapi?apiid=aiplatform.googleapis.com).\n",
|
||||
"\n",
|
||||
"4. If you are running this notebook locally, you need to install the [Cloud SDK](https://cloud.google.com/sdk)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "sBCra4QMA2wR"
|
||||
},
|
||||
"source": [
|
||||
"### Authenticate your Google Cloud account\n",
|
||||
"\n",
|
||||
"Depending on your Jupyter environment, you may have to manually authenticate. Follow the relevant instructions below."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "74ccc9e52986"
|
||||
},
|
||||
"source": [
|
||||
"**1. Vertex AI Workbench**\n",
|
||||
"* Do nothing as you are already authenticated."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "de775a3773ba"
|
||||
},
|
||||
"source": [
|
||||
"**2. Local JupyterLab instance, uncomment and run:**"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "254614fa0c46"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# ! gcloud auth login"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "ef21552ccea8"
|
||||
},
|
||||
"source": [
|
||||
"**3. Colab, uncomment and run:**"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "603adbbf0532"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# from google.colab import auth\n",
|
||||
"# auth.authenticate_user()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "FyyMdUeAJIVv"
|
||||
},
|
||||
"source": [
|
||||
"## Installation\n",
|
||||
"\n",
|
||||
"Install the following packages required to execute this notebook.\n",
|
||||
"\n",
|
||||
"**Remember to restart the runtime after installation.**"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "snBUuUamoJPz"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install git+https://github.com/googleapis/python-aiplatform.git"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "WX3CHZitmSJM"
|
||||
},
|
||||
"source": [
|
||||
"### Please restart the runtime."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "dae340cb-0583-4e7e-a562-6817ee4d7f6d"
|
||||
},
|
||||
"source": [
|
||||
"### Imports libraries"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "412d00f1-08db-4880-8ced-52a9583757b8"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import vertexai\n",
|
||||
"from vertexai.language_models import TextEmbeddingInput, TextEmbeddingModel"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "MyMXIZoRlUcR"
|
||||
},
|
||||
"source": [
|
||||
"#### Set your project ID and initiate Vertex AI\n",
|
||||
"\n",
|
||||
"**If you don't know your project ID**, try the following:\n",
|
||||
"* Run `gcloud config list`.\n",
|
||||
"* Run `gcloud projects list`.\n",
|
||||
"* See the support page: [Locate the project ID](https://support.google.com/googleapi/answer/7014113)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "3EdtdqnoldX4"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"PROJECT_ID = \"\" # @param {type:\"string\"}\n",
|
||||
"REGION = \"us-central1\"\n",
|
||||
"\n",
|
||||
"# Set the project id\n",
|
||||
"! gcloud config set project {PROJECT_ID}\n",
|
||||
"\n",
|
||||
"# Initiate Vertex AI\n",
|
||||
"vertexai.init(project=PROJECT_ID, location=REGION)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "f50f22f3-ec85-463e-b6fe-5c8e6b80b07b"
|
||||
},
|
||||
"source": [
|
||||
"## Generate embeddings"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "hQZoBXNGjizH"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Set the model name.\n",
|
||||
"MODEL_NAME = \"textembedding-gecko@latest\" # @param [\"textembedding-gecko@latest\", \"textembedding-gecko-multilingual@latest\"]\n",
|
||||
"\n",
|
||||
"# Set the task_type, text and optional title as the model inputs.\n",
|
||||
"TASK_TYPE = \"RETRIEVAL_DOCUMENT\" # @param [\"RETRIEVAL_QUERY\", \"RETRIEVAL_DOCUMENT\", \"SEMANTIC_SIMILARITY\", \"CLASSIFICATION\", \"CLUSTERING\"]\n",
|
||||
"TITLE = \"Google\" # @param {type:\"string\"}\n",
|
||||
"TEXT = \"Embed text.\" # @param {type:\"string\"}\n",
|
||||
"\n",
|
||||
"# Verify the input is valid.\n",
|
||||
"if not MODEL_NAME:\n",
|
||||
" raise ValueError(\"Please set MODEL_NAME.\")\n",
|
||||
"if not TASK_TYPE:\n",
|
||||
" raise ValueError(\"Please set TASK_TYPE.\")\n",
|
||||
"if not TEXT:\n",
|
||||
" raise ValueError(\"Please set TEXT.\")\n",
|
||||
"if TITLE and TASK_TYPE != \"RETRIEVAL_DOCUMENT\":\n",
|
||||
" raise ValueError(\"Title can only be provided if the task_type is RETRIEVAL_DOCUMENT\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "BNPapKXviHlE"
|
||||
},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"768\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"def text_embedding(\n",
|
||||
" model_name: str, task_type: str, text: str, title: str = \"\") -> list:\n",
|
||||
" \"\"\"Generate text embedding with a Large Language Model.\"\"\"\n",
|
||||
" model = TextEmbeddingModel.from_pretrained(model_name)\n",
|
||||
"\n",
|
||||
" text_embedding_input = TextEmbeddingInput(\n",
|
||||
" task_type=task_type, title=title, text=text)\n",
|
||||
" embeddings = model.get_embeddings([text_embedding_input])\n",
|
||||
" return embeddings[0].values\n",
|
||||
"\n",
|
||||
"embedding = text_embedding(\n",
|
||||
" model_name=MODEL_NAME, task_type=TASK_TYPE, text=TEXT, title=TITLE)\n",
|
||||
"print(len(embedding))"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"name": "text_embedding_api_cloud_next_new_models.ipynb",
|
||||
"toc_visible": true
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"name": "python3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user