Compare commits

..
Author SHA1 Message Date
denisj3030andGitHub 8f9a263df6 Update CODEOWNERS 2025-04-07 14:50:50 -04:00
261 changed files with 10316 additions and 31450 deletions
@@ -238,7 +238,7 @@ def _get_notebook_python_version(notebook_path: str) -> str:
# Look for the python version specification pattern
re_match = re.search(
r"python version = (\d+\.\d+)", markdown, flags=re.IGNORECASE
"python version = (\d+\.\d+)", markdown, flags=re.IGNORECASE
)
if re_match:
# get the version number
+2
View File
@@ -1,3 +1,5 @@
notebooks/official/vizier/gapic-vizier-multi-objective-optimization.ipynb
notebooks/official/pipelines/lightweight_functions_component_io_kfp.ipynb
notebooks/official/ml_metadata/sdk-metric-parameter-tracking-for-locally-trained-models.ipynb
notebooks/official/custom/custom-tabular-bq-managed-dataset.ipynb
.cloud-build/tests/python_version_test.ipynb
+2 -2
View File
@@ -3,8 +3,8 @@ ipython
jupyter
nbconvert
black==25.1.0
pyupgrade==3.20.0
pyupgrade==3.19.1
isort==6.0.1
flake8==7.3.0
flake8==7.2.0
nbqa==1.9.1
-1
View File
@@ -29,5 +29,4 @@
/vertex_model_garden/model_oss/vllm @kathyyu-google
/vertex_model_garden/benchmarking_reports @lavraicse
/vertex_model_garden/model_oss/autogluon @lavraicse
/vertex_distributed_training/a3mega/llama-3-8b-nemo-pretraining @mstyer-google @erwinh85 @mchrestkha
@@ -1,3 +1,3 @@
torch==2.7.0
torch==2.2.0
torchvision==0.9.1
tensorboard==2.5.0
@@ -1,126 +0,0 @@
# Vertex AI Training: Llama 3.1 8B pre-training using Nvidia A3 Mega VMs (H100)
This document provides a step-by-step guide for pre-training a Llama 3.1 8B model on the `en-wiki` dataset using multiple [Vertex AI Custom Training](https://cloud.google.com/vertex-ai/docs/training/overview) `a3-megagpu-8g` nodes.
We will use a custom container based on NVIDIA's [NeMo Framework](https://docs.nvidia.com/nemo-framework/user-guide/24.07/overview.html) to demonstrate a scalable, multi-node training workflow. All required artifacts and commands are included.
## 1. Prerequisites
### 1.1. Google Cloud Project setup
- **Enable APIs:** Ensure the Vertex AI API is [enabled for your project](http://console.cloud.google.com/flows/enableapi?apiid=aiplatform.googleapis.com).
- **H100 Mega Quota:** A3 Mega VMs are powered by H100 GPUs. Request quota for `custom_model_training_nvidia_h100_mega_gpus` in one of the [supported regions](https://cloud.google.com/vertex-ai/docs/general/locations#accelerator_support). If using Spot VMs, request `custom_model_training_preemptible_nvidia_h100_mega_gpus` quota instead.
- **Reservations (Optional but recommended):** For guaranteed capacity, [create a reservation](https://cloud.google.com/compute/docs/instances/reservations-shared) and ensure the reservation is shared with the Vertex AI service account. This guide requires a minimum of **16 H100 GPUs** (2 full A3 Mega nodes).
### 1.2. GCS bucket
Create a [Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) in the same region where you have quota. If you're using Hierarchical Namespace for your bucket, you may need to update permissions of the Vertex AI Custom Code Service Agent .
This bucket is used for:
- Staging the training application.
- Storing model checkpoints and logs.
- Storing data if you use your own data.
## 2. Setup & configuration
### 2.1. Clone the repo
First clone the repo into your development environment.
```bash
git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git
```
Navigate to the root folder for this sample.
### 2.2. Environment Setup
First, configure your local environment. These variables are used in subsequent commands.
```bash
# Required: Update with your values
export PROJECT_ID="<your-project-id>"
export REPOSITORY="<your-artifact-registry-repo-name>" # e.g., "my-containers"
export BUCKET="<your-gcs-bucket-name>"
# Optional: Change if needed
export REGION="us-central1"
# --- Do not change the lines below ---
export ARTIFACT_REGISTRY="${REGION}-docker.pkg.dev/${PROJECT_ID}/${REPOSITORY}"
export REPO_ROOT=$(git rev-parse --show-toplevel)
```
## 3. Build and push a docker container image to Artifact Registry
Normally, you can use any custom training container on Vertex AI Training. In this example you build a NeMo Docker image that is based on the [Nvidia’s NeMo 24.09](https://catalog.ngc.nvidia.com/orgs/nvidia/containers/nemo/tags) image. Use Cloud Build to build and push the container image.
This document picked NeMo as the demonstrating container since it’s a widely adopted GPU LLM training framework providing high performance and versatile training functionalities.
In addition to the base image, some customizations are included to form the final prebuilt image:
- Some dependencies are installed to integrate with Vertex AI Training.
- An entrypoint script that sets up required environments and calls the training job.
- Some patches are applied to the NeMo code to let it load the dataset from a GCS bucket.
Run this command to build the container and push the container into the Google Artifact Registry.
```bash
cd "${REPO_ROOT}/community-content/vertex-distributed-training/a3mega/llama-3-8b-nemo-pretraining"
export IMAGE_NAME="vertex-nemo-llama"
gcloud builds submit . \
--project="${PROJECT_ID}" \
--region="${REGION}" \
--config=docker/cloudbuild.yml \
--substitutions="_ARTIFACT_REGISTRY=${ARTIFACT_REGISTRY},_IMAGE_NAME=${IMAGE_NAME}" \
--timeout="2h" \
--machine-type="e2-highcpu-32"
```
## 4. Launch the Training Job
### 4.1. Job Configuration File
Once the container is built, update the job_config.json to set up the training job.
File: job_config.json
```json
{
"project_id": "<project-id>",
"region": "<region>",
"zone": "<zone if using reservation>",
"bucket": "<bucket>",
"dataset_bucket": "github-repo/data/third-party/enwiki-latest-pages-articles",
"image_uri": "<docker image uri from artifact registry>",
"strategy": "spot",
"nodes": "2",
"machine_type": "a3-megagpu-8g",
"gpu_type": "NVIDIA_H100_MEGA_80GB",
"gpus_per_node": "8",
"recipe_name": "llama3_1_8b_pretrain_a3mega",
"job_prefix": "vertex-spot-",
"reservation_name": ""
}
```
### 4.2 Launch the Training Job
First, create a Python virtual environment using your tool of choice, then install
the requirements specified in `requirements.txt`. Using `pip`, the command would be:
```bash
pip install -r requirements.txt
```
Now launch the Vertex AI training job using the provided Python script.
```bash
python3 scripts/launch.py --config_file=job_config.json
```
This script reads job_config.json, defines the cluster specification (2 nodes, 8 GPUs each), and submits the custom training job to Vertex AI.
## 5. Monitor and Clean Up
### 5.1. Monitoring
Vertex AI Console: Track the job's status in the Google Cloud Console under Vertex AI > Training > Custom Jobs.
Logs: View detailed logs in Cloud Logging by filtering for your job name.
Checkpoints: Model checkpoints are saved to your GCS bucket at the path specified in your training script's configuration.
### 5.2. Cleaning Up
To avoid ongoing charges, delete the resources you created:
- The Artifact Registry image.
- The contents of the GCS bucket (checkpoints, logs).
- The Vertex AI Custom Job will eventually complete or fail, incurring no further cost.
@@ -1,265 +0,0 @@
# Reference:
# https://github.com/NVIDIA/NeMo-Framework-Launcher/blob/24.07/launcher_scripts/conf/training/llama/llama3_1_8b.yaml
name: llama3_1_8b_pretrain_a3mega
restore_from_path: null # used when starting from a .nemo file
trainer:
devices: 8
num_nodes: 1
accelerator: gpu
precision: bf16
logger: false # logger provided by exp_manager
enable_checkpointing: false
use_distributed_sampler: false
max_epochs: -1 # PTL default. In practice, max_steps will be reached first.
max_steps: 30 # consumed_samples = global_step * micro_batch_size * data_parallel_size * accumulate_grad_batches
log_every_n_steps: 1
val_check_interval: null
limit_val_batches: 1
limit_test_batches: 1
accumulate_grad_batches: 1 # do not modify, grad acc is automatic for training megatron models
gradient_clip_val: 1.0
benchmark: false
enable_model_summary: false # default PTL callback for this does not support model parallelism, instead we log manually
exp_manager:
explicit_log_dir: null
exp_dir: /data
name: ${name}
create_dllogger_logger: true
dllogger_logger_kwargs:
verbose: true
stdout: true
json_file: "/data/dllogger.json"
create_wandb_logger: false
wandb_logger_kwargs:
project: null
name: null
resume_if_exists: true
resume_ignore_no_checkpoint: true
create_checkpoint_callback: false
checkpoint_callback_params:
monitor: val_loss
save_top_k: 3
mode: min
always_save_nemo: false # saves nemo file during validation, not implemented for model parallel
save_nemo_on_train_end: false # not recommended when training large models on clusters with short time limits
filename: 'megatron_gpt--{val_loss:.2f}-{step}-{consumed_samples}'
model_parallel_size: ${multiply:${model.tensor_model_parallel_size}, ${model.pipeline_model_parallel_size}}
seconds_to_sleep: 5 # Allows node_rank!=0 to sleep and let node0 to init, like preparing data
model:
mcore_gpt: true
# specify micro_batch_size, global_batch_size, and model parallelism
# gradient accumulation will be done automatically based on data_parallel_size
micro_batch_size: 1 # limited by GPU memory
global_batch_size: 1024 # will use more micro batches to reach global batch size
tensor_model_parallel_size: 1 # intra-layer model parallelism
pipeline_model_parallel_size: 2 # inter-layer model parallelism
context_parallel_size: 1
virtual_pipeline_model_parallel_size: null # interleaved pipeline
## Sequence Parallelism
# Makes tensor parallelism more memory efficient for LLMs (20B+) by parallelizing layer norms and dropout sequentially
# See Reducing Activation Recomputation in Large Transformer Models: https://arxiv.org/abs/2205.05198 for more details.
sequence_parallel: false
fsdp: false
fsdp_cpu_offload: true
fsdp_sharding_strategy: "full" # Method to shard model states. Available options are 'full', 'hybrid', and 'grad'.
fsdp_grad_reduce_dtype: "16" # Gradient reduction data type.
fsdp_sharded_checkpoint: false # Store and load FSDP shared checkpoint.
fsdp_use_orig_params: false # Set to True to use FSDP for specific peft scheme.
# Distributed checkpoint setup
dist_ckpt_format: "torch_dist" # Set to 'torch_dist' to use PyTorch distributed checkpoint format.
dist_ckpt_load_on_device: true # whether to load checkpoint weights directly on GPU or to CPU
dist_ckpt_parallel_save: true # if true, each worker will write its own part of the dist checkpoint
dist_ckpt_parallel_save_within_dp: false # if true, save will be parallelized only within a DP group (whole world otherwise), which might slightly reduce the save overhead
dist_ckpt_parallel_load: false # if true, each worker will load part of the dist checkpoint and exchange with NCCL. Might use some extra GPU memory
dist_ckpt_torch_dist_multiproc: 2 # number of extra processes per rank used during ckpt save with PyTorch distributed format
dist_ckpt_assume_constant_structure: false # set to True only if the state dict structure doesn't change within a single job. Allows caching some computation across checkpoint saves.
dist_ckpt_parallel_dist_opt: true # parallel save/load of a DistributedOptimizer. 'True' allows performant save and reshardable checkpoints. Set to 'False' only in order to minimize the number of checkpoint files.
dist_ckpt_load_strictness: null # defines checkpoint keys mismatch behavior (only during dist-ckpt load). Choices: assume_ok_unexpected (default - try loading without any check), log_all (log mismatches), raise_all (raise mismatches)
# model architecture
encoder_seq_length: 8192
max_position_embeddings: ${.encoder_seq_length}
num_layers: 32 # 8b: 32 | 70b: 80 | 405b: 126
hidden_size: 4096 # 8b: 4096 | 70b: 8192 | 405b: 16384
ffn_hidden_size: 14336 # 8b: 14336 | 70b: 28672 | 405b: 53248
num_attention_heads: 32 # 8b: 32 | 70b: 64 | 405b: 128
num_query_groups: 8 # Number of query groups for group query attention. If None, normal attention is used. 8b: 8 | 70b: 8 | 405b: 16
init_method_std: 0.01 # Standard deviation of the zero mean normal distribution used for weight initialization. 8b: 0.01 | 70b: 0.008944 | 405b: 0.02
use_scaled_init_method: true # use scaled residuals initialization
hidden_dropout: 0.0 # Dropout probability for hidden state transformer.
attention_dropout: 0.0 # Dropout probability for attention
ffn_dropout: 0.0 # Dropout probability in the feed-forward layer.
kv_channels: null # Projection weights dimension in multi-head attention. Set to hidden_size // num_attention_heads if null
apply_query_key_layer_scaling: true # scale Q * K^T by 1 / layer-number.
normalization: 'rmsnorm' # Normalization layer to use. Options are 'layernorm', 'rmsnorm'
layernorm_epsilon: 1e-5
do_layer_norm_weight_decay: false # True means weight decay on all params
make_vocab_size_divisible_by: 128 # Pad the vocab size to be divisible by this value for computation efficiency.
pre_process: true # add embedding
post_process: true # add pooler
persist_layer_norm: true # Use of persistent fused layer norm kernel.
bias: false # Whether to use bias terms in all weight matrices.
activation: 'fast-swiglu' # Options ['gelu', 'geglu', 'swiglu', 'reglu', 'squared-relu', 'fast-geglu', 'fast-swiglu', 'fast-reglu']
headscale: false # Whether to learn extra parameters that scale the output of the each self-attention head.
transformer_block_type: 'pre_ln' # Options ['pre_ln', 'post_ln', 'normformer']
openai_gelu: false # Use OpenAI's GELU instead of the default GeLU
normalize_attention_scores: true # Whether to scale the output Q * K^T by 1 / sqrt(hidden_size_per_head). This arg is provided as a configuration option mostly for compatibility with models that have been weight-converted from HF. You almost always want to se this to True.
position_embedding_type: 'rope' # Position embedding type. Options ['learned_absolute', 'rope']
rotary_percentage: 1.0 # If using position_embedding_type=rope, then the per head dim is multiplied by this.
attention_type: 'multihead' # Attention type. Options ['multihead']
share_embeddings_and_output_weights: false # Share embedding and output layer weights.
scale_positional_embedding: true # This is false for llama3 models. Only used for >= llama3.1.
# Use GPT2BPETokenizer for test, because the testing dataset is tokenized by this tokenizer.
# https://docs.nvidia.com/nemo-framework/user-guide/24.07/playbooks/singlenodepretrain.html#data-download-and-pre-processing
tokenizer:
library: megatron
type: GPT2BPETokenizer
model: null # /path/to/tokenizer.model
vocab_file: null
merge_file: null
delimiter: null # only used for tabular tokenizer
sentencepiece_legacy: false # Legacy=True allows you to add special tokens to sentencepiece tokenizers.
# Mixed precision
native_amp_init_scale: 4294967296 # 2 ** 32
native_amp_growth_interval: 1000
hysteresis: 2 # Gradient scale hysteresis
fp32_residual_connection: false # Move residual connections to fp32
fp16_lm_cross_entropy: false # Move the cross entropy unreduced loss calculation for lm head to fp16
# Megatron O2-style half-precision
megatron_amp_O2: true # Enable O2-level automatic mixed precision using main parameters
grad_allreduce_chunk_size_mb: 125
# Fusion
grad_div_ar_fusion: true # Fuse grad division into torch.distributed.all_reduce. Only used with O2 and no pipeline parallelism..
gradient_accumulation_fusion: true # Fuse weight gradient accumulation to GEMMs. Only used with pipeline parallelism and O2.
bias_activation_fusion: true # Use a kernel that fuses the bias addition from weight matrices with the subsequent activation function.
bias_dropout_add_fusion: true # Use a kernel that fuses the bias addition, dropout and residual connection addition.
masked_softmax_fusion: true # Use a kernel that fuses the attention softmax with it's mask.
apply_rope_fusion: true # Use a kernel to add rotary positional embeddings. Only used if position_embedding_type=rope
cross_entropy_loss_fusion: true
# Miscellaneous
seed: 1234
resume_from_checkpoint: null # manually set the checkpoint file to load from
use_cpu_initialization: false # Init weights on the CPU (slow for large models)
onnx_safe: false # Use work-arounds for known problems with Torch ONNX exporter.
apex_transformer_log_level: 30 # Python logging level displays logs with severity greater than or equal to this
gradient_as_bucket_view: true # PyTorch DDP argument. Allocate gradients in a contiguous bucket to save memory (less fragmentation and buffer memory)
sync_batch_comm: false # Enable stream synchronization after each p2p communication between pipeline stages
## Activation Checkpointing
# NeMo Megatron supports 'selective' activation checkpointing where only the memory intensive part of attention is checkpointed.
# These memory intensive activations are also less compute intensive which makes activation checkpointing more efficient for LLMs (20B+).
# See Reducing Activation Recomputation in Large Transformer Models: https://arxiv.org/abs/2205.05198 for more details.
# 'full' will checkpoint the entire transformer layer.
activations_checkpoint_granularity: null # 'selective' or 'full'
activations_checkpoint_method: null # 'uniform', 'block'
# 'uniform' divides the total number of transformer layers and checkpoints the input activation
# of each chunk at the specified granularity. When used with 'selective', 'uniform' checkpoints all attention blocks in the model.
# 'block' checkpoints the specified number of layers per pipeline stage at the specified granularity
activations_checkpoint_num_layers: null
# when using 'uniform' this creates groups of transformer layers to checkpoint. Usually set to 1. Increase to save more memory.
# when using 'block' this this will checkpoint the first activations_checkpoint_num_layers per pipeline stage.
num_micro_batches_with_partial_activation_checkpoints: null
# This feature is valid only when used with pipeline-model-parallelism.
# When an integer value is provided, it sets the number of micro-batches where only a partial number of Transformer layers get checkpointed
# and recomputed within a window of micro-batches. The rest of micro-batches in the window checkpoint all Transformer layers. The size of window is
# set by the maximum outstanding micro-batch backpropagations, which varies at different pipeline stages. The number of partial layers to checkpoint
# per micro-batch is set by 'activations_checkpoint_num_layers' with 'activations_checkpoint_method' of 'block'.
# This feature enables using activation checkpoint at a fraction of micro-batches up to the point of full GPU memory usage.
activations_checkpoint_layers_per_pipeline: null
# This feature is valid only when used with pipeline-model-parallelism.
# When an integer value (rounded down when float is given) is provided, it sets the number of Transformer layers to skip checkpointing at later
# pipeline stages. For example, 'activations_checkpoint_layers_per_pipeline' of 3 makes pipeline stage 1 to checkpoint 3 layers less than
# stage 0 and stage 2 to checkpoint 6 layers less stage 0, and so on. This is possible because later pipeline stage
# uses less GPU memory with fewer outstanding micro-batch backpropagations. Used with 'num_micro_batches_with_partial_activation_checkpoints',
# this feature removes most of activation checkpoints at the last pipeline stage, which is the critical execution path.
## Transformer Engine
transformer_engine: true
fp8: false # enables fp8 in TransformerLayer forward
fp8_e4m3: false # sets fp8_format = recipe.Format.E4M3
fp8_hybrid: false # sets fp8_format = recipe.Format.HYBRID
fp8_margin: 0 # scaling margin
fp8_interval: 1 # scaling update interval
fp8_amax_history_len: 1024 # Number of steps for which amax history is recorded per tensor
fp8_amax_compute_algo: 'max' # 'most_recent' or 'max'. Algorithm for computing amax from history
ub_tp_comm_overlap: false # do not turn on because of b/397797926
use_flash_attention: true
gc_interval: 100
## Offloading Activations/Weights to CPU
cpu_offloading: false
cpu_offloading_num_layers: ${sum:${.num_layers},-1} # This value should be between [1,num_layers-1] as we don't want to offload the final layer's activations and expose any offloading duration for the final layer
cpu_offloading_activations: true
cpu_offloading_weights: true
data:
# Path to data must be specified by the user.
# Supports List, String and Dictionary
# List : can override from the CLI: "model.data.data_prefix=[.5,/raid/data/pile/my-gpt3_00_text_document,.5,/raid/data/pile/my-gpt3_01_text_document]",
# Or see example below:
# data_prefix:
# - .5
# - /raid/data/pile/my-gpt3_00_text_document
# - .5
# - /raid/data/pile/my-gpt3_01_text_document
# Dictionary: can override from CLI "model.data.data_prefix"={"train":[1.0, /path/to/data], "validation":/path/to/data, "test":/path/to/test}
# Or see example below:
# "model.data.data_prefix: {train:[1.0,/path/to/data], validation:[/path/to/data], test:[/path/to/test]}"
data_prefix: [1.0, /data/hfbpe_gpt_training_data_text_document]
index_mapping_dir: null # path to save index mapping .npy files, by default will save in the same location as data_prefix
data_impl: mmap
splits_string: 900,50,50
seq_length: ${model.encoder_seq_length}
skip_warmup: true
num_workers: 2
dataloader_type: single # cyclic
reset_position_ids: false # Reset position ids after end-of-document token
reset_attention_mask: false # Reset attention mask after end-of-document token
eod_mask_loss: false # Mask loss for the end of document tokens
validation_drop_last: true # Set to false if the last partial validation samples is to be consumed
no_seqlen_plus_one_input_tokens: false # Set to True to disable fetching (sequence length + 1) input tokens, instead get (sequence length) input tokens and mask the last token
pad_samples_to_global_batch_size: false # Set to True if you want to pad the last partial batch with -1's to equal global batch size
shuffle_documents: true # Set to False to disable documents shuffling. Sample index will still be shuffled
# Nsys profiling options
nsys_profile:
enabled: false
start_step: 0 # Global batch to start profiling
end_step: 1 # Global batch to end profiling
ranks: [0] # Global rank IDs to profile
gen_shape: false # Generate model and kernel details including input shapes
memory_profile:
enabled: false
start_step: 0
end_step: 1
ranks: [0]
output_path: /data # Must be a dir
optim:
name: distributed_fused_adam # E.g., fused_adam or set _target_: torch.optim.AdamW field
lr: 2e-5
weight_decay: 0.01
betas:
- 0.9
- 0.98
bucket_cap_mb: 125
overlap_grad_sync: true
overlap_param_sync: true
contiguous_grad_buffer: true
contiguous_param_buffer: true
sched:
name: CosineAnnealing
warmup_steps: 400
constant_steps: 0
min_lr: 2e-6
@@ -1,26 +0,0 @@
# Copyright 2024 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
steps:
- name: 'gcr.io/cloud-builders/docker'
args:
- 'build'
- '--tag=${_ARTIFACT_REGISTRY}/${_IMAGE_NAME}'
- '--file=docker/vertex-dist-recipes.Dockerfile'
- '.'
automapSubstitutions: true
env:
- 'DOCKER_BUILDKIT=1'
images:
- '${_ARTIFACT_REGISTRY}/${_IMAGE_NAME}'
@@ -1,41 +0,0 @@
diff --git a/nemo/collections/nlp/parts/megatron_trainer_builder.py b/nemo/collections/nlp/parts/megatron_trainer_builder.py
index b2c85cde4..a3a9670c3 100644
--- a/nemo/collections/nlp/parts/megatron_trainer_builder.py
+++ b/nemo/collections/nlp/parts/megatron_trainer_builder.py
@@ -19,6 +19,7 @@ from lightning_fabric.utilities.exceptions import MisconfigurationException
from omegaconf import DictConfig
from pytorch_lightning import Trainer
from pytorch_lightning.callbacks import ModelSummary
+from pytorch_lightning.callbacks import Callback
from pytorch_lightning.plugins.environments import TorchElasticEnvironment
from nemo.collections.common.metrics.perf_metrics import FLOPsMeasurementCallback
@@ -38,6 +39,23 @@ from nemo.utils.callbacks.dist_ckpt_io import (
AsyncFinalizerCallback,
DistributedCheckpointIO,
)
+from vmg.util.device_stats import gpu_stats_str
+
+class GpuStatsMon(Callback):
+ def on_train_start(self, trainer, pl_module) -> None:
+ rank=pl_module.global_rank
+ print(f'train_start: {rank=} {gpu_stats_str()}', flush=True)
+
+ def on_train_batch_start(self, trainer, pl_module, batch, batch_idx) -> None:
+ rank=pl_module.global_rank
+ print(f'batch_start: {rank=} {gpu_stats_str()}', flush=True)
+
+ def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx) -> None:
+ rank=pl_module.global_rank
+ print(f'batch_end: {rank=} {gpu_stats_str()}', flush=True)
class MegatronTrainerBuilder:
@@ -178,6 +196,7 @@ class MegatronTrainerBuilder:
if self.cfg.get('exp_manager', {}).get('log_tflops_per_sec_per_gpu', True):
callbacks.append(FLOPsMeasurementCallback(self.cfg))
+ callbacks.append(GpuStatsMon())
return callbacks
def create_trainer(self, callbacks=None) -> Trainer:
@@ -1,41 +0,0 @@
diff -ruN old-datasets/blended_megatron_dataset_builder.py datasets/blended_megatron_dataset_builder.py
--- old-datasets/blended_megatron_dataset_builder.py 2025-05-02 04:08:45.369199665 +0000
+++ datasets/blended_megatron_dataset_builder.py 2025-05-02 04:10:47.369119891 +0000
@@ -2,6 +2,7 @@
import logging
import math
+import os
from concurrent.futures import ThreadPoolExecutor
from typing import Any, Callable, Iterable, List, Optional, Type, Union
@@ -353,7 +354,7 @@
num_dataset_builder_threads = self.config.num_dataset_builder_threads
if torch.distributed.is_initialized():
- rank = torch.distributed.get_rank()
+ rank = int(os.getenv("LOCAL_RANK", "0"))
# First, build on rank 0
if rank == 0:
num_workers = num_dataset_builder_threads
@@ -475,7 +476,7 @@
Optional[Union[DistributedDataset, Iterable]]: The DistributedDataset instantion, the Iterable instantiation, or None
"""
if torch.distributed.is_initialized():
- rank = torch.distributed.get_rank()
+ rank = int(os.getenv("LOCAL_RANK", "0"))
dataset = None
diff -ruN old-datasets/gpt_dataset.py datasets/gpt_dataset.py
--- old-datasets/gpt_dataset.py 2025-05-02 04:08:45.369199665 +0000
+++ datasets/gpt_dataset.py 2025-05-02 04:09:30.309170278 +0000
@@ -351,7 +351,7 @@
if not path_to_cache or (
not cache_hit
- and (not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0)
+ and (not torch.distributed.is_initialized() or int(os.getenv("LOCAL_RANK", "0")) == 0)
):
log_single_rank(
@@ -1,13 +0,0 @@
diff --git a/scripts/checkpoint_converters/convert_llama_nemo_to_hf.py b/scripts/checkpoint_converters/convert_llama_nemo_to_hf.py
index 8da15148d..005cae6c9 100644
--- a/scripts/checkpoint_converters/convert_llama_nemo_to_hf.py
+++ b/scripts/checkpoint_converters/convert_llama_nemo_to_hf.py
@@ -104,6 +104,8 @@ def convert(input_nemo_file, output_hf_file, precision=None, cpu_only=False) ->
dummy_trainer = Trainer(devices=1, accelerator='cpu', strategy=NLPDDPStrategy())
model_config = MegatronGPTModel.restore_from(input_nemo_file, trainer=dummy_trainer, return_config=True)
model_config.tensor_model_parallel_size = 1
+ model_config.virtual_pipeline_model_parallel_size = None
+ model_config.sequence_parallel = False
model_config.pipeline_model_parallel_size = 1
if cpu_only:
map_location = torch.device('cpu')
@@ -1,24 +0,0 @@
diff --git a/examples/nlp/language_modeling/tuning/megatron_gpt_finetuning.py b/examples/nlp/language_modeling/tuning/megatron_gpt_finetuning.py
index bfe8ea359..dfeaf93b5 100644
--- a/examples/nlp/language_modeling/tuning/megatron_gpt_finetuning.py
+++ b/examples/nlp/language_modeling/tuning/megatron_gpt_finetuning.py
@@ -13,6 +13,8 @@
# limitations under the License.
import torch.multiprocessing as mp
+import torch.distributed as dist
+
from omegaconf.omegaconf import OmegaConf
from nemo.collections.nlp.models.language_modeling.megatron_gpt_sft_model import MegatronGPTSFTModel
@@ -76,6 +78,10 @@ def main(cfg) -> None:
trainer.fit(model)
+ if dist.is_available() and dist.is_initialized():
+ dist.barrier()
+ dist.destroy_process_group()
+
if __name__ == '__main__':
main()
@@ -1,13 +0,0 @@
diff --git a/src/utils/training_metrics/process_training_results.py b/src/utils/training_metrics/process_training_results.py
index 3e82a66..e61e1d8 100644
--- a/src/utils/training_metrics/process_training_results.py
+++ b/src/utils/training_metrics/process_training_results.py
@@ -134,7 +134,7 @@ def get_average_step_time(file: str, start_step: int, end_step: int) -> float:
for line in datajson:
if line.get("step") != "PARAMETER":
step = line.get("step")
- if step >= start_step and step <= end_step:
+ if step >= start_step and step <= end_step and "train_step_timing in s" in line["data"]:
time_step_accumulator += line["data"].get("train_step_timing in s")
num_steps += 1
if num_steps == 0:
@@ -1,10 +0,0 @@
dllogger@git+https://github.com/NVIDIA/dllogger@v1.0.0
# Fixing these libraries versions to avoid conflicting or broken packages.
immutabledict==4.2.1
protobuf==3.20.3
opencv-python-headless==4.11.0.86
docutils==0.16
urllib3==2.0.7
google-cloud-storage==3.0.0
retrying
@@ -1,18 +0,0 @@
# cuml-cu12==24.8.0 was installed in nemo:24.09
# Removing cuml=24.4.0 to avoid conflicting packages.
cudf==24.4.0
cugraph==24.4.0
cugraph-service-server==24.4.0
cuml==24.4.0
dask-cudf==24.4.0
raft-dask==24.4.0
cugraph-dgl==24.4.0
cugraph-pyg==24.4.0
# The following packages are removed temporarily to avoid conflicting packages
# and can be brought back if needed.
tensorrt-llm==0.12.0
img2dataset==1.45.0
Sphinx==8.1.3
sphinxcontrib-bibtex==2.6.3
torchx==0.7.0
nemo-run
@@ -1,66 +0,0 @@
# Dockerfile wrapping NeMo.
#
# To workaround base nemo docker image using too many layers, we use Multi-stage
# build to first collect the additional files we'll need.
FROM alpine:latest AS prep_files
WORKDIR /workspace
RUN mkdir -p configs vdt vdt/util
COPY scripts/*.py vdt/
COPY scripts/util/*.py vdt/util/
COPY configs/* configs/
COPY docker/patches/24.09/* vdt/patches/
RUN chmod a+rwX -R vdt
# Copy license.
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
# Available tags
# https://catalog.ngc.nvidia.com/orgs/nvidia/containers/nemo/tags
# It installs NeMo source code in /opt/NeMo folder, with tag=r2.0.0
FROM nvcr.io/nvidia/nemo:24.09
RUN apt-get update && apt-get install -y sudo zsh tmux && \
rm -rf /var/lib/apt/lists*
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 && \
rm -rf /var/lib/apt/lists*
# Install libraries with pip
ENV PIP_ROOT_USER_ACTION=ignore
# We expect this will be run in the root directory of the vertex-dist-recipes repo
ARG HOST_SRC_DIR="."
# The pre-installed NeMo introduces a lot of deps conflicts.
# We uninstall the confilicting libs and reinstall some of them as needed.
COPY ${HOST_SRC_DIR}/docker/uninstall.txt /tmp/uninstall.txt
RUN cat /tmp/uninstall.txt | grep -v '#' | xargs pip uninstall -y
COPY ${HOST_SRC_DIR}/docker/requirements.txt /tmp/requirements.txt
RUN pip install -r /tmp/requirements.txt
# Make sure there's no inconsistent pip libraries.
RUN pip check
WORKDIR /workspace
# Copy configs
COPY ${HOST_SRC_DIR}/configs/* /opt/NeMo/examples/nlp/language_modeling/conf/
# Copy all additional files we need from `prep_files` image.
COPY --from=prep_files /workspace/ .
# Install for `src/utils/training_metrics/process_training_results.py` to report
# throughput and MFU numbers.
RUN git clone https://github.com/AI-Hypercomputer/gpu-recipes.git
# This hack is needed for multi-node training while not using a sharing file system.
RUN patch --verbose -l -d /opt/megatron-lm/megatron/core/datasets -p1 -i /workspace/vdt/patches/local_rank.patch; \
git -C /workspace/gpu-recipes apply /workspace/vdt/patches/throughput_calc.patch; \
git -C /opt/NeMo apply /workspace/vdt/patches/nemo2hf.patch; \
git -C /opt/NeMo apply /workspace/vdt/patches/sigabort.patch;
# git -C /opt/NeMo apply /workspace/vdt/patches/gpu_stats.patch;
# Do not put an entrypoint here. Specify the entrypoint in the docker run script.
@@ -1,16 +0,0 @@
{
"project_id": "<your_project_id>",
"region": "us-central1",
"zone": "us-central1-c",
"bucket": "<your_bucket",
"dataset_bucket": "github-repo/data/third-party/enwiki-latest-pages-articles",
"image_uri": "<your_image_uri>",
"strategy": "spot",
"nodes": "2",
"machine_type": "a3-megagpu-8g",
"gpu_type": "NVIDIA_H100_MEGA_80GB",
"gpus_per_node": "8",
"recipe_name": "llama3_1_8b_pretrain_a3mega",
"job_prefix": "vertex-ai",
"reservation_name": ""
}
@@ -1,49 +0,0 @@
absl-py==2.2.2
annotated-types==0.7.0
anyio==4.9.0
black==25.1.0
cachetools==5.5.2
certifi==2025.4.26
charset-normalizer==3.4.2
click==8.1.8
docstring_parser==0.16
google-api-core==2.24.2
google-auth==2.40.1
google-cloud-aiplatform==1.92.0
google-cloud-bigquery==3.31.0
google-cloud-core==2.4.3
google-cloud-resource-manager==1.14.2
google-cloud-storage==2.19.0
google-crc32c==1.7.1
google-genai==1.14.0
google-resumable-media==2.7.2
googleapis-common-protos==1.70.0
grpc-google-iam-v1==0.14.2
grpcio==1.71.0
grpcio-status==1.71.0
h11==0.16.0
httpcore==1.0.9
httpx==0.28.1
idna==3.10
mypy_extensions==1.1.0
numpy==2.2.5
packaging==25.0
pathspec==0.12.1
platformdirs==4.3.8
proto-plus==1.26.1
protobuf==5.29.4
pyasn1==0.6.1
pyasn1_modules==0.4.2
pydantic==2.11.4
pydantic_core==2.33.2
python-dateutil==2.9.0.post0
pytz==2025.2
requests==2.32.3
rsa==4.9.1
shapely==2.1.0
six==1.17.0
sniffio==1.3.1
typing-inspection==0.4.0
typing_extensions==4.13.2
urllib3==2.4.0
websockets==15.0.1
@@ -1,173 +0,0 @@
"""Launch script for Vertex distributed training"""
# Copy the sample_job_config.json file to job_config.json
# to define the job parameters.
#
# Run like this:
#
# python3 vertex_dist_train/launch.py --config_file=job_config.json
#
import datetime
import json
import os
import pprint
from collections.abc import Sequence
from typing import Any, List
from absl import app, flags
from google.cloud import aiplatform
from google.cloud.aiplatform_v1.types.custom_job import Scheduling
from pytz import timezone
FLAGS = flags.FLAGS
flags.DEFINE_string("config_file", None, "Path to JSON config file")
flags.DEFINE_boolean(
"debug", False, "Debug mode: just print the command, don't run it."
)
def launch_job(
job_name: str,
project: str,
region: str,
gcs_bucket: str,
image_uri: str,
entrypoint_cmd: List[str],
trainer_args: List[Any],
num_nodes: int,
machine_type: str,
num_gpus_per_node: int,
gpu_type: str,
strategy: str,
reservation_name: str = "",
):
assert strategy in ("dws", "spot", "reservation")
aiplatform.init(
project=project, location=region, staging_bucket=gcs_bucket
)
train_job = aiplatform.CustomContainerTrainingJob(
display_name=job_name,
container_uri=image_uri,
command=entrypoint_cmd,
)
job_args = dict(
args=trainer_args,
enable_web_access=True,
replica_count=num_nodes,
machine_type=machine_type,
accelerator_type=gpu_type,
accelerator_count=num_gpus_per_node,
boot_disk_size_gb=1000,
restart_job_on_worker_restart=True,
#restart_job_on_worker_restart=False,
)
if strategy == "spot":
job_args.update({"scheduling_strategy": Scheduling.Strategy.SPOT.name})
elif strategy == "dws":
job_args.update(
{"scheduling_strategy": Scheduling.Strategy.FLEX_START.name}
)
elif strategy == "reservation":
assert reservation_name != "", (
"If using a reservation, provide the reservation_name in the "
"format `projects/{project_id_or_number}/zones/{zone}/"
"reservations/{reservation_name}`"
)
job_args.update(
{
"reservation_affinity_type": "SPECIFIC_RESERVATION",
"reservation_affinity_key": "compute.googleapis.com/reservation-name",
"reservation_affinity_values": [reservation_name],
}
)
pprint.pprint(job_args)
if not FLAGS.debug:
train_job.submit(**job_args)
def main(argv: Sequence[str]) -> None:
config_file_path = FLAGS.config_file
print(f"Reading job config from {config_file_path}")
with open(config_file_path, encoding="utf-8") as config_file:
config = json.load(config_file)
project_id = config["project_id"]
region = config["region"]
zone = config["zone"]
bucket = config["bucket"]
dataset_bucket = config["dataset_bucket"]
n_nodes = int(config["nodes"])
machine_type = config["machine_type"]
num_gpus_per_node = int(config["gpus_per_node"])
gpu_type = config["gpu_type"]
reservation_name = config.get("reservation_name")
reservation_full_name = (
f"projects/{project_id}/zones/{zone}/reservations/{reservation_name}"
if "reservation_name" in config
else ""
)
strategy = config["strategy"]
recipe_name = config["recipe_name"]
job_prefix = config["job_prefix"]
image_uri = config["image_uri"]
# Job name
timestamp = (
datetime.datetime.now()
.astimezone(timezone("US/Pacific"))
.strftime("%Y%m%d_%H%M%S")
)
job_name = f"{recipe_name}-{timestamp}"
if job_prefix:
job_name = f"{job_prefix}-{job_name}"
base_output_dir = os.path.join("/gcs", bucket, job_name)
# Training command and args
entrypoint_cmd = ["python3", "vdt/run.py"]
dataset_bucket = f"gs://{config['dataset_bucket']}"
trainer_args = [
f"--train_data_gcs={dataset_bucket}",
"/opt/NeMo/examples/nlp/language_modeling/megatron_gpt_pretraining.py",
"--config-path=conf/",
f"--config-name={recipe_name}.yaml",
f"exp_manager.explicit_log_dir={base_output_dir}",
f"exp_manager.dllogger_logger_kwargs.json_file={base_output_dir}/dllogger.json",
"+exp_manager.create_tensorboard_logger=true",
"exp_manager.create_checkpoint_callback=false",
f"trainer.num_nodes={n_nodes}",
f"trainer.devices={num_gpus_per_node}",
"trainer.max_steps=10",
"trainer.log_every_n_steps=1",
"model.tokenizer.vocab_file=/data/gpt2-vocab.json",
"model.tokenizer.merge_file=/data/gpt2-merges.txt",
"model.data.data_prefix=[1.0,/data/hfbpe_gpt_training_data_text_document]",
]
launch_job(
job_name=job_name,
project=project_id,
region=region,
gcs_bucket=bucket,
image_uri=image_uri,
entrypoint_cmd=entrypoint_cmd,
trainer_args=trainer_args,
num_nodes=n_nodes,
machine_type=machine_type,
num_gpus_per_node=num_gpus_per_node,
gpu_type=gpu_type,
strategy=strategy,
reservation_name=reservation_full_name,
)
if __name__ == "__main__":
app.run(main)
@@ -1,85 +0,0 @@
"""Entrypoint for Vertex Distributed Training container."""
import argparse
import os
import sys
from collections.abc import Sequence
from subprocess import STDOUT, check_output, run
from absl import app, flags, logging
from util import cluster_spec
from retrying import retry
# PyTorch barrier call which synchronizes all of the nodes before launching the training process.
# This makes sure that processes will block until all processes are ready.
# Improves the reliability of spot VM usage for multi-node training jobs
@retry(stop_max_attempt_number=100, wait_exponential_multiplier=1000)
def barrier_with_retry() -> None:
import torch
logging.info("Starting barrier on RANK {}".format(os.environ["RANK"]))
torch.distributed.init_process_group()
torch.distributed.barrier()
torch.distributed.destroy_process_group()
logging.info("Finished barrier on RANK {}".format(os.environ["RANK"]))
def main(unused_argv: Sequence[str]) -> None:
parser = argparse.ArgumentParser()
parser.add_argument(
"--train_data_gcs",
type=str,
help="Download training data from gcs path",
)
args, unknown = parser.parse_known_args()
for key, val in os.environ.items():
logging.info("ENV %s=%s", key, val)
if args.train_data_gcs:
local_dir = "/data"
if not os.path.exists(local_dir):
os.mkdir(local_dir)
logging.info("downloading %s to %s...", args.train_data_gcs, local_dir)
check_output(
[
"gcloud",
"storage",
"cp",
"-r",
f"{args.train_data_gcs}/*",
local_dir,
],
stderr=STDOUT,
)
logging.info("%s downloaded.", args.train_data_gcs)
primary_node_addr, primary_node_port, node_rank, num_nodes = (
cluster_spec.get_cluster_spec()
)
cmd = [
"torchrun",
"--nproc-per-node=8",
f"--nnodes={num_nodes}",
f"--node_rank={node_rank}",
]
if num_nodes > 1:
cmd += [
"--max-restarts=3",
"--rdzv-backend=static",
f'--rdzv_id={os.getenv("CLOUD_ML_JOB_ID", primary_node_port)}',
f"--rdzv-endpoint={primary_node_addr}:{primary_node_port}",
]
cmd += unknown
logging.info("launching with cmd: \n%s", " \\\n".join(cmd))
barrier_with_retry()
run(cmd, stdout=sys.stdout, stderr=sys.stdout, check=True)
if __name__ == "__main__":
logging.get_absl_handler().python_handler.stream = sys.stdout
app.run(
main, flags_parser=lambda _args: flags.FLAGS(_args, known_only=True)
)
@@ -1,81 +0,0 @@
"""Get cluster info from environment variables."""
import dataclasses
import json
import os
from absl import logging
@dataclasses.dataclass
class ClusterInfo:
"""Contains information about the cluster.
Attributes:
primary_node_addr: The address of the primary node.
primary_node_port: The port of the primary node.
node_rank: The rank of the node.
num_nodes: The number of nodes in the cluster.
"""
primary_node_addr: str | None = None
primary_node_port: str | None = None
node_rank: int = 0
num_nodes: int = 1
# Allows unpacking operation like
# primary_node_addr, primary_node_port, _, _ = ClusterInfo()
# See https://stackoverflow.com/a/70753113
def __iter__(self):
return iter(dataclasses.astuple(self))
def get_cluster_spec() -> ClusterInfo:
"""Parses CLUSTER_SPEC environment variable and returns the cluster info.
Returns:
A ClusterInfo object.
"""
cluster_spec = os.getenv("CLUSTER_SPEC", None)
# If CLUSTER_SPEC is not set, use individual vars to construct cluster info.
if not cluster_spec:
cluster_info = ClusterInfo(
primary_node_addr=os.getenv("MASTER_ADDR", None),
primary_node_port=os.getenv("MASTER_PORT", None),
node_rank=int(os.getenv("RANK", "0")),
num_nodes=int(os.getenv("NNODES", "1")),
)
return cluster_info
cluster_data = json.loads(cluster_spec)
# Get primary node info
primary_node = cluster_data["cluster"]["workerpool0"][0]
logging.info("primary node: %s", primary_node)
primary_node_addr, primary_node_port = primary_node.split(":")
logging.info("primary node address: %s", primary_node_addr)
logging.info("primary node port: %s", primary_node_port)
# Determine node rank of this machine
workerpool = cluster_data["task"]["type"]
if workerpool == "workerpool0":
node_rank = 0
elif workerpool == "workerpool1":
# Add 1 for the primary node, since `index` is the index of workerpool1.
node_rank = cluster_data["task"]["index"] + 1
else:
raise ValueError(
"Only workerpool0 and workerpool1 are supported. Unknown workerpool:"
f" {workerpool}"
)
logging.info("node rank: %s", node_rank)
# Calculate total nodes.
num_nodes = 1 # For the primary node.
if "workerpool1" in cluster_data["cluster"]:
num_nodes += len(cluster_data["cluster"]["workerpool1"])
logging.info("num nodes: %s", num_nodes)
return ClusterInfo(
primary_node_addr, primary_node_port, node_rank, num_nodes
)
@@ -1,59 +0,0 @@
"""Add tests for cluster_spec.py."""
import os
from . import cluster_spec
# TODO(styer): Use pytest instead
class ClusterSpecTest(googletest.TestCase):
def setUp(self):
super().setUp()
self.curr_env_var = os.environ.copy()
def tearDown(self):
super().tearDown()
os.environ = self.curr_env_var
def test_get_cluster_spec_from_env_vars(self):
os.environ["CLUSTER_SPEC"] = ""
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = "8080"
os.environ["RANK"] = "0"
os.environ["NNODES"] = "2"
cluster_info = cluster_spec.get_cluster_spec()
self.assertEqual(cluster_info.primary_node_addr, "127.0.0.1")
self.assertEqual(cluster_info.primary_node_port, "8080")
self.assertEqual(cluster_info.node_rank, 0)
self.assertEqual(cluster_info.num_nodes, 2)
def test_get_cluster_spec_from_cluster_spec(self):
os.environ[
"CLUSTER_SPEC"
] = """
{
"cluster": {
"workerpool0": [
"127.0.0.1:8080"
],
"workerpool1": [
"127.0.0.2:8080",
"127.0.0.3:8080"
]
},
"task": {
"type": "workerpool1",
"index": 0
}
}
"""
cluster_info = cluster_spec.get_cluster_spec()
self.assertEqual(cluster_info.primary_node_addr, "127.0.0.1")
self.assertEqual(cluster_info.primary_node_port, "8080")
self.assertEqual(cluster_info.node_rank, 1)
self.assertEqual(cluster_info.num_nodes, 3)
if __name__ == "__main__":
googletest.main()
@@ -1,14 +1,13 @@
"""Common util functions for notebook."""
import base64
from collections.abc import Sequence
import datetime
import io
import json
import os
import subprocess
import time
from typing import Any
from typing import Any, Dict, Sequence
from google import auth
from google.cloud import storage
@@ -284,7 +283,7 @@ def decode_image(
return image
def get_label_map(label_map_yaml_filepath: str) -> dict[int, str]:
def get_label_map(label_map_yaml_filepath: str) -> Dict[int, str]:
"""Returns class id to label mapping given a filepath to the label map.
Args:
@@ -334,7 +333,6 @@ def vqa_predict(
image: Any,
language_code: str = "en",
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> Sequence[str]:
"""Predicts the answer to a question about an image using an Endpoint."""
# Resize and convert image to base64 string.
@@ -358,9 +356,7 @@ def vqa_predict(
"image": resized_image_base64,
})
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
response = endpoint.predict(instances=instances)
return [pred.get("response") for pred in response.predictions]
@@ -370,7 +366,6 @@ def caption_predict(
image: Any,
caption_prompt: bool = False,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Predicts a caption for a given image using an Endpoint."""
# Resize and convert image to base64 string.
@@ -385,9 +380,7 @@ def caption_predict(
instance["prompt"] = caption_prompt_format.format(language_code)
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
response = endpoint.predict(instances=instances)
return response.predictions[0].get("response")
@@ -396,7 +389,6 @@ def ocr_predict(
ocr_prompt: str,
image: Any,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Extracts text from a given image using an Endpoint."""
# Resize and convert image to base64 string.
@@ -408,9 +400,7 @@ def ocr_predict(
instance["prompt"] = ocr_prompt
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
response = endpoint.predict(instances=instances)
return response.predictions[0].get("response")
@@ -419,7 +409,6 @@ def detect_predict(
detect_prompt: str,
image: Any,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Predicts the answer to a question about an image using an Endpoint."""
# Resize and convert image to base64 string.
@@ -431,9 +420,7 @@ def detect_predict(
instance["prompt"] = detect_prompt
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
response = endpoint.predict(instances=instances)
return response.predictions[0].get("response")
@@ -510,17 +497,6 @@ def get_quota(project_id: str, region: str, resource_id: str) -> int:
):
return -1
all_regions_data = quota_data[0]["consumerQuotaLimits"][0]["quotaBuckets"]
# If the quota data does not have dimensions, it is global quota. However,
# global quota may be overridden by regional quota. So we need to check the
# global quota first.
global_quota = -1
if (
all_regions_data
and "dimensions" not in all_regions_data[0]
and "effectiveLimit" in all_regions_data[0]
):
global_quota = int(all_regions_data[0]["effectiveLimit"])
for region_data in all_regions_data:
if (
region_data.get("dimensions")
@@ -530,13 +506,12 @@ def get_quota(project_id: str, region: str, resource_id: str) -> int:
return int(region_data["effectiveLimit"])
else:
return 0
return global_quota
return -1
def get_resource_id(
accelerator_type: str,
is_for_training: bool,
is_spot: bool = False,
is_restricted_image: bool = False,
is_dynamic_workload_scheduler: bool = False,
) -> str:
@@ -546,7 +521,6 @@ def get_resource_id(
accelerator_type: The accelerator type.
is_for_training: Whether the resource is used for training. Set false for
serving use case.
is_spot: Whether the resource is used with Spot.
is_restricted_image: Whether the image is hosted in `vertex-ai-restricted`.
is_dynamic_workload_scheduler: Whether the resource is used with Dynamic
Workload Scheduler.
@@ -562,9 +536,7 @@ def get_resource_id(
"NVIDIA_A100_80GB": "nvidia_a100_80gb_gpus",
"NVIDIA_H100_80GB": "nvidia_h100_gpus",
"NVIDIA_H100_MEGA_80GB": "nvidia_h100_mega_gpus",
"NVIDIA_H200_141GB": "nvidia_h200_gpus",
"NVIDIA_TESLA_T4": "nvidia_t4_gpus",
"TPU_V6e": "tpu_v6e",
"TPU_V5e": "tpu_v5e",
"TPU_V3": "tpu_v3",
}
@@ -579,10 +551,6 @@ def get_resource_id(
restricted_image_training_accelerator_map = {
"NVIDIA_A100_80GB": "restricted_image_training_nvidia_a100_80gb_gpus",
}
spot_serving_accelerator_map = {
key: f"custom_model_serving_preemptible_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
serving_accelerator_map = {
key: f"custom_model_serving_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
@@ -611,11 +579,8 @@ def get_resource_id(
else:
if is_dynamic_workload_scheduler:
raise ValueError("Dynamic Workload Scheduler does not work for serving.")
accelerator_map = (
spot_serving_accelerator_map if is_spot else serving_accelerator_map
)
if accelerator_type in accelerator_map:
return accelerator_map[accelerator_type]
if accelerator_type in serving_accelerator_map:
return serving_accelerator_map[accelerator_type]
else:
raise ValueError(
f"Could not find accelerator type: {accelerator_type} for serving."
@@ -628,28 +593,13 @@ def check_quota(
accelerator_type: str,
accelerator_count: int,
is_for_training: bool,
is_spot: bool = False,
is_restricted_image: bool = False,
is_dynamic_workload_scheduler: bool = False,
) -> None:
"""Checks if the project and the region has the required quota.
Args:
project_id: The project id.
region: The region.
accelerator_type: The accelerator type.
accelerator_count: The number of accelerators to check quota for.
is_for_training: Whether the resource is used for training. Set false for
serving use case.
is_spot: Whether the resource is used with Spot.
is_restricted_image: Whether the image is hosted in `vertex-ai-restricted`.
is_dynamic_workload_scheduler: Whether the resource is used with Dynamic
Workload Scheduler.
"""
):
"""Checks if the project and the region has the required quota."""
resource_id = get_resource_id(
accelerator_type,
is_for_training=is_for_training,
is_spot=is_spot,
is_restricted_image=is_restricted_image,
is_dynamic_workload_scheduler=is_dynamic_workload_scheduler,
)
@@ -7,7 +7,7 @@ import json
import multiprocessing
import os
import subprocess
from typing import Any, Callable, Dict, Tuple, Union
from typing import Any, Callable, Dict, Union
from absl import logging
import accelerate
import datasets
@@ -70,9 +70,7 @@ def force_gcs_fuse_path(gcs_uri: str) -> str:
def download_gcs_uri_to_local(
gcs_uri: str,
destination_dir: str = LOCAL_BASE_MODEL_DIR,
check_path_exists: bool = True,
gcs_uri: str, destination_dir: str = LOCAL_BASE_MODEL_DIR
) -> str:
"""Downloads GCS URI to local.
@@ -83,7 +81,6 @@ def download_gcs_uri_to_local(
Args:
gcs_uri: GCS URI to download.
destination_dir: Local directory directory.
check_path_exists: Whether to check if the path exists.
Returns:
Local path to target folder/file.
@@ -92,7 +89,7 @@ def download_gcs_uri_to_local(
destination_dir,
os.path.basename(os.path.normpath(gcs_uri)),
)
if check_path_exists and os.path.exists(target):
if os.path.exists(target):
logging.info("File %s already exists.", target)
return target
if accelerate.PartialState().is_local_main_process:
@@ -418,42 +415,13 @@ def get_filtered_dataset(
return filtered_dataset
def format_dataset(
dataset: datasets.Dataset,
input_column: str,
template: str = None,
tokenizer: transformers.PreTrainedTokenizer | None = None,
) -> datasets.Dataset:
"""Takes a raw dataset and formats it using a template and tokenizer.
Args:
dataset: The raw (unprocessed) dataset to format.
input_column: The input column in the dataset to be used or updaded by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
tokenizer: The tokenizer to use for chat_template templates.
Returns:
A dataset compatible with the template.
"""
return dataset.map(
_format_template_fn(
template,
input_column=input_column,
tokenizer=tokenizer,
)
)
def load_dataset_with_template(
dataset_name: str,
split: str,
input_column: str,
template: str = None,
tokenizer: transformers.PreTrainedTokenizer | None = None,
) -> Tuple[Any, Any]:
) -> Any:
"""Loads dataset with templates.
Args:
@@ -467,15 +435,19 @@ def load_dataset_with_template(
tokenizer: The tokenizer to use for chat_template templates.
Returns:
The raw dataset and the dataset compatible with the template.
A dataset compatible with the template.
"""
raw = _get_dataset(dataset_name, split=split)
dataset = _get_dataset(dataset_name, split=split)
if template:
templated = format_dataset(raw, input_column, template, tokenizer)
else:
templated = None
dataset = dataset.map(
_format_template_fn(
template,
input_column=input_column,
tokenizer=tokenizer,
)
)
return raw, templated
return dataset
def validate_dataset_with_template(
@@ -549,11 +521,12 @@ def validate_dataset_with_template(
f" https://github.com/GoogleCloudPlatform/{_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME}/tree/main/{_VERTEX_AI_SAMPLES_GITHUB_TEMPLATE_DIR}."
)
dataset = format_dataset(
_get_dataset(dataset_name, split, num_proc),
input_column,
template_path,
tokenizer,
dataset = _get_dataset(dataset_name, split, num_proc).map(
_format_template_fn(
template_path,
input_column=input_column,
tokenizer=tokenizer,
)
)
if tokenizer is not None:
@@ -3,7 +3,6 @@
import copy
import dataclasses
import datetime
import inspect
import os
import signal
import subprocess
@@ -12,7 +11,7 @@ from absl import flags
from absl import logging
from absl.testing import parameterized
import command_builder
import immutabledict
import frozendict
import torch
_DOCKER_URI = flags.DEFINE_string('docker_uri', None, 'docker image uri')
@@ -34,19 +33,19 @@ _LOCAL_OUTPUT_DIR = flags.DEFINE_string(
_GCS_INPUT_DIR = flags.DEFINE_string(
'gcs_input_dir',
'gs://vmg-tuning-docker-test',
'gs://peft-docker-test',
'GCS directory that stores model checkpoint, dataset and etc.',
)
_GCS_OUTPUT_DIR = flags.DEFINE_string(
'gcs_output_dir',
'gs://vmg-tuning-docker-test/output',
'gs://peft-docker-test/output',
'GCS directory that stores test output.',
)
_GCS_TESTDATA_DIR = 'peft-train-image-test'
_THROUGHPUT_TEST_EXCEPTIONS = immutabledict.immutabledict({
_THROUGHPUT_TEST_EXCEPTIONS = frozendict.frozendict({
('bm_deepspeed_zero3_8gpu_gemma-2-9b-it_4bit.txt', '12.0'): float('inf'),
('bm_fsdp_8gpu_llama3.1-70b-hf_4bit.txt', '20.0'): float('inf'),
('bm_deepspeed_zero2_8gpu_gemma-2-2b-it_bfloat16.txt', '12.0'): 20.0,
@@ -102,25 +101,26 @@ class TestBase(parameterized.TestCase):
return self.command_builder.build_cmd() + self.task_cmd_builder.build_cmd()
def run_cmd(self) -> int:
return run_cmd(self.cmd(), output_file=None)
logging.info('running command: \n%s', ' \\\n'.join(self.cmd()))
if _DRY_RUN.value:
return 0
p = subprocess.Popen(self.cmd(), stdout=sys.stdout, stderr=sys.stderr)
try:
unused_output, unused_error = p.communicate()
return p.returncode
except KeyboardInterrupt:
p.send_signal(signal.SIGINT)
return 0
def gcs_output_dir(self):
return _GCS_OUTPUT_DIR.value
def local_input_dir(self):
"""Returns local input dir in host/docker."""
return _LOCAL_INPUT_DIR.value
def local_output_dir(self):
"""Returns local output dir in host/docker."""
return _LOCAL_OUTPUT_DIR.value
def get_testcase_name(self):
"""Returns the function name at the calling site."""
# https://docs.python.org/3/library/inspect.html#inspect.FrameInfo
cur_frame = inspect.currentframe()
# https://stackoverflow.com/a/17366561
return cur_frame.f_back.f_code.co_name
def local_input_dir(self):
return _LOCAL_INPUT_DIR.value
def get_timestamp():
@@ -157,44 +157,13 @@ def get_test_data_path(name: str, download: bool = True) -> str:
local_data = os.path.join(_LOCAL_INPUT_DIR.value, name)
if not os.path.exists(local_data):
# If `name` is a file in sub-folders, then create the sub-folders under
# `_LOCAL_INPUT_DIR`.
local_data_dir = os.path.dirname(local_data)
if not os.path.exists(local_data_dir):
os.makedirs(local_data_dir)
download_from_gcs(os.path.join(_GCS_INPUT_DIR.value, name), local_data_dir)
download_from_gcs(
os.path.join(_GCS_INPUT_DIR.value, name), _LOCAL_INPUT_DIR.value
)
return local_data
def run_cmd(cmd: list[str], output_file: str = None) -> int:
"""Runs the command and returns the return code.
Args:
cmd: The command to run.
output_file: The file to write the output to.
Returns:
The return code of the command.
"""
logging.info('running command: \n%s', ' \\\n'.join(cmd))
if _DRY_RUN.value:
return 0
stdout = sys.stdout if output_file is None else open(output_file, 'w')
p = subprocess.Popen(cmd, stdout=stdout, stderr=sys.stderr)
try:
unused_output, unused_error = p.communicate()
return_code = p.returncode
except KeyboardInterrupt:
p.send_signal(signal.SIGINT)
return_code = 0
finally:
if output_file is not None:
stdout.close()
return return_code
def get_pretrained_model_name_or_path(model_id: str) -> str:
# If `model_id` contains `/`, it is assumed to be HF model or model from GCS.
if '/' in model_id:
@@ -1,79 +0,0 @@
"""Get cluster info from environment variables."""
import dataclasses
import json
import os
from absl import logging
@dataclasses.dataclass
class ClusterInfo:
"""Contains information about the cluster.
Attributes:
primary_node_addr: The address of the primary node.
primary_node_port: The port of the primary node.
node_rank: The rank of the node.
num_nodes: The number of nodes in the cluster.
"""
primary_node_addr: str | None = None
primary_node_port: str | None = None
node_rank: int = 0
num_nodes: int = 1
# Allows unpacking operation like
# primary_node_addr, primary_node_port, _, _ = ClusterInfo()
# See https://stackoverflow.com/a/70753113
def __iter__(self):
return iter(dataclasses.astuple(self))
def get_cluster_spec() -> ClusterInfo:
"""Parses CLUSTER_SPEC environment variable and returns the cluster info.
Returns:
A ClusterInfo object.
"""
cluster_spec = os.getenv('CLUSTER_SPEC', None)
# If CLUSTER_SPEC is not set, use individual vars to construct cluster info.
if not cluster_spec:
cluster_info = ClusterInfo(
primary_node_addr=os.getenv('MASTER_ADDR', None),
primary_node_port=os.getenv('MASTER_PORT', None),
node_rank=int(os.getenv('RANK', '0')),
num_nodes=int(os.getenv('NNODES', '1')),
)
return cluster_info
cluster_data = json.loads(cluster_spec)
# Get primary node info
primary_node = cluster_data['cluster']['workerpool0'][0]
logging.info('primary node: %s', primary_node)
primary_node_addr, primary_node_port = primary_node.split(':')
logging.info('primary node address: %s', primary_node_addr)
logging.info('primary node port: %s', primary_node_port)
# Determine node rank of this machine
workerpool = cluster_data['task']['type']
if workerpool == 'workerpool0':
node_rank = 0
elif workerpool == 'workerpool1':
# Add 1 for the primary node, since `index` is the index of workerpool1.
node_rank = cluster_data['task']['index'] + 1
else:
raise ValueError(
'Only workerpool0 and workerpool1 are supported. Unknown workerpool:'
f' {workerpool}'
)
logging.info('node rank: %s', node_rank)
# Calculate total nodes.
num_nodes = 1 # For the primary node.
if 'workerpool1' in cluster_data['cluster']:
num_nodes += len(cluster_data['cluster']['workerpool1'])
logging.info('num nodes: %s', num_nodes)
return ClusterInfo(primary_node_addr, primary_node_port, node_rank, num_nodes)
@@ -1,24 +0,0 @@
"""Utility functions."""
import logging
import subprocess
import sys
import time
def run_cmd(cmd: list[str]) -> float:
"""Runs the command and logs the output.
Args:
cmd: The command to run.
Returns:
The time it took to run the command.
"""
cmd_str = ' \\\n'.join(cmd)
logging.info('launching cmd: \n%s', cmd_str)
start_time = time.time()
subprocess.run(cmd, stdout=sys.stdout, stderr=sys.stdout, check=True)
elapsed_time = round(time.time() - start_time, 2)
logging.info('Command %s finished in %0.2f seconds.', cmd_str, elapsed_time)
return elapsed_time
@@ -1,197 +0,0 @@
"""Calculate dataset statistics like token, example and character counts."""
from collections.abc import Mapping, Sequence
import dataclasses
import json
from typing import Any
import datasets
import numpy as np
import transformers
from util import dataset_validation_util
_MAX_NUM_DATASET_SAMPLES = 6
@dataclasses.dataclass
class SupervisedTuningDatasetBucket:
"""Represents a histogram bucket for tuning dataset distribution stats."""
count: float = 0
left: float = 0
right: float = 0
@dataclasses.dataclass
class SupervisedTuningDatasetDistribution:
"""Represents a histogram with summary statistics for tuning dataset distribution stats."""
sum: int = 0
billable_sum: int = 0
min: float = 0
max: float = 0
mean: float = 0
median: float = 0
p5: float = 0
p95: float = 0
buckets: list[SupervisedTuningDatasetBucket] = dataclasses.field(
default_factory=list
)
# Represents detailed tuning dataset statistics.
@dataclasses.dataclass
class SupervisedTuningDataStats:
"""Represents detailed tuning dataset stats."""
tuning_dataset_example_count: int = 0
total_tuning_character_count: int = 0
total_billable_token_count: int = 0
tuning_step_count: int = 0
# Represents a histogram and some summary statistics of the number of input
# tokens across examples.
user_input_token_distribution: SupervisedTuningDatasetDistribution | None = (
None
)
# Represents a histogram and some summary statistics for the number of output
# tokens across examples.
user_output_token_distribution: SupervisedTuningDatasetDistribution | None = (
None
)
# Represents the number of "messages" (a single-turn conversation will have a
# single message) across examples.
user_message_per_example_distribution: (
SupervisedTuningDatasetDistribution | None
) = None
user_dataset_examples: list[str] = dataclasses.field(default_factory=list)
def get_dataset_stats(
*,
raw: Any,
templated: Any,
template: str,
tokenizer: transformers.PreTrainedTokenizer,
column: str,
effective_batch_size: int,
) -> Mapping[str, Any]:
"""Calculates dataset statistics for managed fine-tuning, e.g., total number of tokens."""
tokenized_dataset = templated.map(lambda x: tokenizer(x[column]))
inputs = tokenized_dataset["input_ids"]
tuning_dataset_example_count = int(len(inputs))
total_billable_token_count = int(np.sum([len(ex) for ex in inputs]))
total_tuning_character_count = int(
np.sum([len(ex[column]) for ex in templated])
)
tuning_step_count = (
tuning_dataset_example_count + effective_batch_size - 1
) // effective_batch_size
# Assume that data is represented as ChatCompletions or Vertex Text-Bison
# formats to extract per-example input/output tokens.
user_inputs = []
user_outputs = []
user_input_messages_counts = []
for ex in raw:
if "messages" in ex:
messages = ex["messages"]
if messages:
# For ChatCompletions assume the last turn (i.e. the instruction
# response) is the expected output.
user_inputs.append({**ex, "messages": messages[:-1]})
user_outputs.append({**ex, "messages": messages[-1:]})
# Exclude everything but the last message for the number of input
# messages.
user_input_messages_counts.append(len(messages[:-1]))
elif "input_text" in ex:
# For Vertex Text-Bison, the `output_text` field is the expected output.
user_inputs.append({**ex, "output_text": ""})
user_outputs.append(
{**ex, "input_text": ex["output_text"], "output_text": ""}
)
# Vertex Text-Bison goes from input -> output; i.e. there is only a single
# input "message".
user_input_messages_counts.append(1)
def calc_histogram(
counts: Sequence[int],
) -> SupervisedTuningDatasetDistribution:
mean = np.mean(counts)
median = np.median(counts).item()
max_count = np.max(counts).item()
min_count = np.min(counts).item()
count_sum = np.sum(counts).item()
p5 = np.percentile(counts, 0.05).item()
p95 = np.percentile(counts, 0.95).item()
hist, bin_edges = np.histogram(counts, bins=10)
return SupervisedTuningDatasetDistribution(
sum=count_sum,
billable_sum=count_sum,
min=min_count,
max=max_count,
mean=mean,
median=median,
p5=p5,
p95=p95,
buckets=[
SupervisedTuningDatasetBucket(
count=hist[i].item(),
left=bin_edges[i].item(),
right=bin_edges[i + 1].item(),
)
for i in range(len(hist))
],
)
# Tokenize input and output messages separately to generate separate summary
# statistics about them.
user_input_token_distribution = None
if user_inputs:
user_input_dataset = dataset_validation_util.format_dataset(
datasets.Dataset.from_list(user_inputs), column, template, tokenizer
)
user_input_tokenized_dataset = user_input_dataset.map(
lambda x: tokenizer(x[column])
)
user_input_tokens = user_input_tokenized_dataset["input_ids"]
user_input_token_counts = np.array([len(ex) for ex in user_input_tokens])
user_input_token_distribution = calc_histogram(user_input_token_counts)
user_output_token_distribution = None
if user_outputs:
user_output_dataset = dataset_validation_util.format_dataset(
datasets.Dataset.from_list(user_outputs), column, template, tokenizer
)
user_output_tokenized_dataset = user_output_dataset.map(
lambda x: tokenizer(x[column])
)
user_output_tokens = user_output_tokenized_dataset["input_ids"]
user_output_token_counts = np.array([len(ex) for ex in user_output_tokens])
user_output_token_distribution = calc_histogram(user_output_token_counts)
user_messages_per_example_distribution = None
if user_input_messages_counts:
user_input_messages_counts = np.array(user_input_messages_counts)
user_messages_per_example_distribution = calc_histogram(
user_input_messages_counts
)
user_dataset_examples = [
json.dumps(ex)
for ex in raw.shuffle().select(
range(min(len(raw), _MAX_NUM_DATASET_SAMPLES))
)
]
dataset_stats = SupervisedTuningDataStats(
tuning_dataset_example_count=tuning_dataset_example_count,
total_tuning_character_count=total_tuning_character_count,
total_billable_token_count=total_billable_token_count,
tuning_step_count=tuning_step_count,
user_input_token_distribution=user_input_token_distribution,
user_output_token_distribution=user_output_token_distribution,
user_message_per_example_distribution=user_messages_per_example_distribution,
user_dataset_examples=user_dataset_examples,
)
return dataclasses.asdict(dataset_stats)
@@ -1,140 +0,0 @@
"""Util functions for reporting device (GPU, CPU) stats."""
import dataclasses
import psutil
import pynvml
import torch
@dataclasses.dataclass
class GpuStats:
"""Holds information about GPU usage stats.
For memory related, see
https://pytorch.org/docs/stable/notes/cuda.html#cuda-memory-management
"""
# device id
device_id: int
# memory reserved.
reserved: float
# memory occupied.
occupied: float
# memory reserved, but not used.
unused: float
# nvidia-smi usually reports more memory usages than pytorch (for driver,
# kernel and etc). `smi_diff` tracks this difference.
smi_diff: float
# Gpu utilization.
util: float
# Allows unpacking operation like
# device_id, reserved, occupied, unused, smi_diff, util = GpuStats(...)
# See https://stackoverflow.com/a/70753113
def __iter__(self):
return iter(dataclasses.astuple(self))
def gpu_stats() -> GpuStats:
"""Reports GPU memory usage and utilization."""
# See https://pytorch.org/docs/stable/notes/cuda.html#memory-management
bytes_per_gb = 1024.0**3
device = torch.cuda.current_device()
occupied = torch.cuda.memory_allocated(device) / bytes_per_gb
reserved = torch.cuda.memory_reserved(device) / bytes_per_gb
unused = reserved - occupied
def smi_mem(device):
try:
pynvml.nvmlInit()
handle = pynvml.nvmlDeviceGetHandleByIndex(device)
info = pynvml.nvmlDeviceGetMemoryInfo(handle)
return info.used / bytes_per_gb
except pynvml.NVMLError:
return 0.0
mem_used_smi = smi_mem(device)
smi_diff = mem_used_smi - reserved
util = torch.cuda.utilization(device)
return GpuStats(device, reserved, occupied, unused, smi_diff, util)
def gpu_stats_str(stats: GpuStats | None = None) -> str:
if stats is None:
stats = gpu_stats()
device, reserved, occupied, unused, smi_diff, util = stats
return (
f"GPU ({device=}) memory: {reserved:.2f}({occupied=:.2f}, {unused=:.2f}),"
f" {smi_diff=:.2f} GB. Utilization: {util:.2f}%"
)
@dataclasses.dataclass
class CpuStats:
"""Holds information about CPU usage stats."""
# Total CPU virtual memory i.e. virtual memory allocated + unallocated.
total_virtual_mem: float
# CPU virtual memory available for use.
unallocated_virtual_mem: float
# CPU virtual memory already used.
allocated_virtual_mem: float
# Total CPU swap memory i.e. swap memory allocated + unallocated.
total_swap_mem: float
# CPU swap memory available for use.
unallocated_swap_mem: float
# CPU swap memory already used.
allocated_swap_mem: float
# CPU utilization percentage.
utilization: float
def cpu_stats() -> CpuStats:
"""Reports CPU memory usage and utilization."""
# https://psutil.readthedocs.io/en/latest/#memory
gb = 1024.0**3
vmem = psutil.virtual_memory()
vmem_total = vmem.total / gb
vmem_available = vmem.available / gb
vmem_used = vmem_total - vmem_available
smem = psutil.swap_memory()
swap_total = smem.total / gb
swap_free = smem.free / gb
swap_used = smem.used / gb
# https://psutil.readthedocs.io/en/latest/#psutil.cpu_percent
cpu_util = psutil.cpu_percent(interval=1e-6)
return CpuStats(
total_virtual_mem=vmem_total,
unallocated_virtual_mem=vmem_available,
allocated_virtual_mem=vmem_used,
total_swap_mem=swap_total,
unallocated_swap_mem=swap_free,
allocated_swap_mem=swap_used,
utilization=cpu_util,
)
def cpu_stats_str(stats: CpuStats | None = None) -> str:
"""Returns a string representation of the CPU stats."""
if stats is None:
stats = cpu_stats()
total, occupied, unused = (
stats.total_virtual_mem,
stats.allocated_virtual_mem,
stats.unallocated_virtual_mem,
)
virtual_mem = (
f"CPU virtual memory: {total:.2f}({occupied=:.2f}, {unused=:.2f}) GB"
)
total, occupied, unused = (
stats.total_swap_mem,
stats.allocated_swap_mem,
stats.unallocated_swap_mem,
)
swap_mem = f"CPU swap memory: {total:.2f}({occupied=:.2f}, {unused=:.2f}) GB"
percent = stats.utilization
return f"{virtual_mem} {swap_mem} CPU Utilization: {percent:.2f}%"
@@ -11,7 +11,7 @@ from transformers.trainer_callback import TrainerCallback
from transformers.trainer_callback import TrainerControl
from transformers.trainer_callback import TrainerState
from util import device_stats
from vertex_vision_model_garden_peft.train.vmg import utils
class TrainerStatsCallback(TrainerCallback):
@@ -75,15 +75,13 @@ class TrainerStatsCallback(TrainerCallback):
state.global_step - 1
)
gpu_stats = device_stats.gpu_stats()
self._peak_mem = max(
gpu_stats.reserved + gpu_stats.smi_diff, self._peak_mem
)
gpu_stats = utils.gpu_stats()
self._peak_mem = max(gpu_stats.total_mem, self._peak_mem)
logging.info(
'on_step_end: Throughput: %.2f token/s. %s, %s',
throughput,
device_stats.gpu_stats_str(gpu_stats),
device_stats.cpu_stats_str(),
utils.gpu_stats_str(gpu_stats),
utils.cpu_stats_str(),
)
def on_train_begin(
@@ -97,8 +95,8 @@ class TrainerStatsCallback(TrainerCallback):
self._start_time = time.time()
logging.info(
'on_train_begin: %s, %s',
device_stats.gpu_stats_str(),
device_stats.cpu_stats_str(),
utils.gpu_stats_str(),
utils.cpu_stats_str(),
)
def on_train_end(
@@ -1,28 +0,0 @@
compute_environment: LOCAL_MACHINE
debug: false
distributed_type: FSDP
downcast_bf16: 'no'
enable_cpu_affinity: false
fsdp_config:
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
fsdp_transformer_layer_cls_to_wrap: Qwen2DecoderLayer
fsdp_backward_prefetch: NO_PREFETCH
fsdp_cpu_ram_efficient_loading: true
fsdp_forward_prefetch: false
fsdp_offload_params: true
fsdp_sharding_strategy: FULL_SHARD
fsdp_state_dict_type: SHARDED_STATE_DICT
fsdp_sync_module_states: true
fsdp_use_orig_params: false
fsdp_activation_checkpointing: false
main_training_function: main
mixed_precision: bf16
machine_rank: 0
num_machines: 1
num_processes: 8
rdzv_backend: static
same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
@@ -16,7 +16,6 @@ diffusers==0.25.1
evaluate==0.4.3
fsspec==2024.3.1
gcsfs==2024.3.1
immutabledict==4.2.1
ninja==1.11.1 # Needed to avoid `ninja 1.11.1.1 is not supported on this platform` error
nltk==3.9.1
optimum==1.17.1
@@ -68,12 +68,10 @@ RUN mkdir -p ./vertex_vision_model_garden_peft/
COPY model_oss/peft/train/vmg/configs/* ./vertex_vision_model_garden_peft/
COPY model_oss/peft/train/vmg/*.py ./vertex_vision_model_garden_peft/train/vmg/
COPY model_oss/peft/train/vmg/templates /diffusers/examples/util/templates
COPY model_oss/peft/train/util/*.py /diffusers/examples/util/
COPY model_oss/util/* /diffusers/examples/util/
COPY model_oss/util /diffusers/examples/util
COPY model_oss/notebook_util/dataset_validation_util.py /diffusers/examples/util
COPY model_oss/peft/train/vmg/tests/*.py ./vertex_vision_model_garden_peft/tests/
COPY model_oss/peft/train/test_utils/test_util.py ./vertex_vision_model_garden_peft/tests/
COPY model_oss/peft/train/test_utils/command_builder.py ./vertex_vision_model_garden_peft/tests/
RUN chmod a+rwX -R /diffusers/examples/
ENV PYTHONPATH /diffusers/examples/
@@ -37,6 +37,7 @@ class EvalConfig:
steps: The number of steps to run evaluation.
tasks: The list of tasks to run evaluation on.
per_device_batch_size: The per device batch size for evaluation.
num_fewshot: The number of few-shot examples to use for evaluation.
limit: The maximum number of examples to evaluate.
metric_name: The name of the metric to compute.
tokenize_dataset: Whether to tokenize the dataset.
@@ -49,6 +50,7 @@ class EvalConfig:
steps: int
per_device_batch_size: int
num_fewshot: int | None
limit: float | None
metric_name: Sequence[str]
tokenize_dataset: bool
@@ -97,7 +99,7 @@ def create_trainer(
kwargs["tokenizer"] = tokenizer
try:
_, eval_dataset = dataset_validation_util.load_dataset_with_template(
eval_dataset = dataset_validation_util.load_dataset_with_template(
dataset_name=eval_config.dataset_path,
split=eval_config.split,
input_column=eval_config.column,
@@ -1,107 +1,17 @@
"""Sync local directory to GCS directory using rsync."""
from collections.abc import Sequence
import multiprocessing
import os
import subprocess
import time
from typing import Optional, Sequence, Tuple
from absl import logging
from util import constants
from util import fileutils
_GCS_COMMAND_RETRIES = 3
_RSYNC_RETRY_INTERVAL_SECS = 30
def is_gcs_or_gcsfuse_path(path: str) -> bool:
"""Returns if the path is a GCS or gcsfuse path.
Args:
path: The path to check.
Returns:
True if the path is a GCS or gcsfuse path.
"""
return path.startswith(
(constants.GCS_URI_PREFIX, constants.GCSFUSE_URI_PREFIX)
)
def manage_sync_path(
path: str, node_rank: Optional[int] = None
) -> Tuple[str, str]:
"""Returns local dir and GCS location for the given path if the given path is a GCS or gcsfuse path.
It will also create a local directory if it does not exist. Otherwise, it
returns the same path.
Args:
path: The local or GCS path to manage.
node_rank: The node rank to be appended to the GCS path.
Returns:
The local and GCS paths.
"""
local_dir = path
gcs_dir = path
if is_gcs_or_gcsfuse_path(path):
local_dir = os.path.join(
constants.LOCAL_OUTPUT_DIR,
fileutils.force_gcs_fuse_path(path)[1:],
)
gcs_dir = fileutils.force_gcs_path(path)
if not os.path.exists(local_dir):
os.makedirs(local_dir, exist_ok=True)
if node_rank is None:
return local_dir, gcs_dir
return local_dir, os.path.join(gcs_dir, f"node-{node_rank}")
def setup_gcs_rsync(
dirs_to_sync: Sequence[Tuple[str, str]],
mp_queue: multiprocessing.Queue,
gcs_rsync_interval_secs: int,
) -> multiprocessing.Process:
"""Sets up the GCS rsync process.
Args:
dirs_to_sync: The absolute directory paths which will be synced to GCS.
mp_queue: The multiprocessing queue to check if the training is finished.
gcs_rsync_interval_secs: Integer, interval in seconds to run gcs rsync.
Returns:
The GCS rsync process.
"""
rsync_process = multiprocessing.Process(
target=start_gcs_rsync,
args=(dirs_to_sync, mp_queue, gcs_rsync_interval_secs),
)
rsync_process.start()
return rsync_process
def cleanup_gcs_rsync(
rsync_process: multiprocessing.Process, mp_queue: multiprocessing.Queue
) -> None:
"""Cleans up the GCS rsync process.
Args:
rsync_process: The GCS rsync process.
mp_queue: The multiprocessing queue.
"""
mp_queue.put("finish rsync process")
rsync_process.join()
if rsync_process.exitcode == 0:
logging.info("Artifacts have been uploaded to GCS.")
else:
logging.error(
"GCS rsync process failed with exit code %d.", rsync_process.exitcode
)
def _rsync_local_to_gcs(local_dir: str, gcs_dir: str) -> None:
"""Syncs the local directory to GCS.
@@ -147,7 +57,7 @@ def _rsync_local_to_gcs(local_dir: str, gcs_dir: str) -> None:
def start_gcs_rsync(
dirs_to_sync: Sequence[Tuple[str, str]],
dirs_to_sync: Sequence[tuple[str, str]],
mp_queue: multiprocessing.Queue,
gcs_rsync_interval_secs: int,
) -> None:
@@ -1,6 +1,7 @@
"""Instruct/Chat with LoRA models."""
from collections.abc import Callable, Mapping, Sequence
import dataclasses
import datetime
import json
import os
@@ -22,8 +23,6 @@ import trl
import wandb
from util import dataset_validation_util
from util import dataset_stats
from util import device_stats
from vertex_vision_model_garden_peft.train.vmg import callbacks
from vertex_vision_model_garden_peft.train.vmg import eval_lib
from vertex_vision_model_garden_peft.train.vmg import utils
@@ -232,6 +231,11 @@ _PER_DEVICE_EVAL_BATCH_SIZE = flags.DEFINE_integer(
'The per device batch size for model evaluation.',
)
_EVAL_NUM_FEWSHOT = flags.DEFINE_integer(
'eval_num_fewshot',
None,
'Run N-shot language model evaluation. Not implemented in `builtin_eval`.',
)
_EVAL_LIMIT = flags.DEFINE_float(
'eval_limit',
@@ -251,7 +255,8 @@ _EVAL_METRIC_NAME = flags.DEFINE_list(
_EVAL_DATASET = flags.DEFINE_string(
'eval_dataset',
None,
'The Hugging Face dataset name or path to use for evaluation.',
'Overrides the default evaluation dataset path. In `builtin_eval` mode,'
' this can be any Hugging Face dataset name or path.',
)
# We set the default eval split as `test`, based on observation from
@@ -259,13 +264,13 @@ _EVAL_DATASET = flags.DEFINE_string(
_EVAL_SPLIT = flags.DEFINE_string(
'eval_split',
'test',
'Eval split name in the eval dataset.',
'Eval split name in the eval dataset for `builtin_eval`.',
)
_EVAL_TEMPLATE = flags.DEFINE_string(
'eval_template',
None,
'Template for formatting language model evaluation data.'
'Template for formatting language model evaluation data for `builtin_eval`.'
' Must be a filename under `templates` folder, without `.json` extension,'
' e.g. `alpaca`, or a Cloud Storage URI to a JSON file.',
)
@@ -273,7 +278,7 @@ _EVAL_TEMPLATE = flags.DEFINE_string(
_EVAL_COLUMN = flags.DEFINE_string(
'eval_column',
None,
'Eval column name in the eval dataset.',
'Eval column name in the eval dataset for `builtin_eval`.',
)
_METRIC_FOR_BEST_MODEL = flags.DEFINE_string(
@@ -571,8 +576,8 @@ def finetune_instruct(
"""Finetunes instruct."""
logging.info(
'on entering instruct_lora, %s,\n%s',
device_stats.gpu_stats_str(),
device_stats.cpu_stats_str(),
utils.gpu_stats_str(),
utils.cpu_stats_str(),
)
gradient_checkpointing_kwargs = {}
# DDP provides limited support with the reentrant variant of gradient
@@ -589,7 +594,7 @@ def finetune_instruct(
access_token=access_token,
)
train_dataset, train_dataset_with_template = (
train_dataset_with_template = (
dataset_validation_util.load_dataset_with_template(
train_dataset,
split=train_split,
@@ -616,20 +621,18 @@ def finetune_instruct(
'getting tuning data stats with effective batch size %s',
effective_batch_size,
)
train_dataset_stats = dataset_stats.get_dataset_stats(
raw=train_dataset,
templated=train_dataset_with_template,
template=train_template,
tokenizer=tokenizer,
column=train_column,
effective_batch_size=effective_batch_size,
train_dataset_stats = utils.get_dataset_stats(
train_dataset_with_template,
tokenizer,
train_column,
effective_batch_size,
)
logging.info('stats: %s', train_dataset_stats)
tuning_data_stats_file = dataset_validation_util.force_gcs_fuse_path(
tuning_data_stats_file
)
with open(tuning_data_stats_file, 'w') as out_f:
json.dump(train_dataset_stats, out_f)
json.dump(dataclasses.asdict(train_dataset_stats), out_f)
model = utils.load_model(
pretrained_model_name_or_path=pretrained_model_name_or_path,
@@ -660,9 +663,7 @@ def finetune_instruct(
# `get_peft_model`, which may revert other changes we did before. That's why
# we are calling `get_peft_model` explicitly here.
model = get_peft_model(model, peft_config)
adapter_for_eval_dir = os.path.join(output_dir, 'adapter_for_eval')
logging.info('saving adapter for evaluation to %s...', adapter_for_eval_dir)
peft_config.save_pretrained(adapter_for_eval_dir)
# This is to work-around mix-precision training. This issue is not fixed as
# of transformers==4.41.2.
# See b/332760883#comment30 for more details.
@@ -839,6 +840,7 @@ def main(unused_argv: Sequence[str]) -> None:
if _EVAL_DATASET.value:
eval_config = eval_lib.EvalConfig(
per_device_batch_size=_PER_DEVICE_EVAL_BATCH_SIZE.value,
num_fewshot=_EVAL_NUM_FEWSHOT.value,
limit=_EVAL_LIMIT.value,
metric_name=_EVAL_METRIC_NAME.value,
steps=_EVAL_STEPS.value,
@@ -31,7 +31,7 @@ _MERGE_BASE_AND_LORA_OUTPUT_DIR = flags.DEFINE_string(
_MERGE_MODEL_PRECISION_MODE = flags.DEFINE_enum(
'merge_model_precision_mode',
constants.PRECISION_MODE_16B,
constants.PRECISION_MODE_16,
[
constants.PRECISION_MODE_4,
constants.PRECISION_MODE_8,
@@ -86,19 +86,10 @@ def main(unused_argv: Sequence[str]) -> None:
)
)
finetuned_lora_model_dir = fileutils.force_gcs_path(
_FINETUNED_LORA_MODEL_DIR.value
)
if dataset_validation_util.is_gcs_path(finetuned_lora_model_dir):
finetuned_lora_model_dir = (
dataset_validation_util.download_gcs_uri_to_local(
finetuned_lora_model_dir
)
)
utils.merge_causal_language_model_with_lora(
pretrained_model_name_or_path=pretrained_model_name_or_path,
precision_mode=_MERGE_MODEL_PRECISION_MODE.value,
finetuned_lora_model_dir=finetuned_lora_model_dir,
finetuned_lora_model_dir=_FINETUNED_LORA_MODEL_DIR.value,
merged_model_output_dir=_MERGE_BASE_AND_LORA_OUTPUT_DIR.value,
access_token=_HUGGINGFACE_ACCESS_TOKEN.value,
)
@@ -0,0 +1,232 @@
"""Sequence classification with LoRA models."""
from typing import Sequence
from absl import app
from absl import flags
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
from util import dataset_validation_util
_PRETRAINED_MODEL_NAME_OR_PATH = flags.DEFINE_string(
"pretrained_model_name_or_path",
None,
"The pretrained model name or path. Supported models can be causal language"
" modeling models from https://github.com/huggingface/peft/tree/main. Note,"
" there might be different paddings for different models. This tool assumes"
" the pretrained_model_name_or_path contains model name, and then choose"
" proper padding methods. e.g. it must contain `llama` for `Llama2"
" models`.",
)
_OUTPUT_DIR = flags.DEFINE_string(
"output_dir",
None,
"The output directory.",
)
_DATASET_NAME = flags.DEFINE_string(
"dataset_name",
None,
"The dataset name in huggingface.",
)
_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.",
)
_NUM_TRAIN_EPOCHS = flags.DEFINE_integer(
"num_train_epochs",
None,
"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 finetune_sequence_classification(
pretrained_model_name_or_path: str,
dataset_name: str,
output_dir: str,
lora_rank: int = 8,
lora_alpha: int = 16,
lora_dropout: float = 0.1,
num_train_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_name_or_path for k in ("gpt", "opt", "bloom")):
padding_side = "left"
else:
padding_side = "right"
tokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path, 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_name_or_path, 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_train_epochs),
num_training_steps=(len(train_dataloader) * num_train_epochs),
)
model.to(device)
for epoch in range(num_train_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)
def main(unused_argv: Sequence[str]) -> None:
if dataset_validation_util.is_gcs_path(_PRETRAINED_MODEL_NAME_OR_PATH.value):
pretrained_model_name_or_path = (
dataset_validation_util.download_gcs_uri_to_local(
_PRETRAINED_MODEL_NAME_OR_PATH.value
)
)
else:
pretrained_model_name_or_path = _PRETRAINED_MODEL_NAME_OR_PATH.value
pretrained_model_path = dataset_validation_util.force_gcs_fuse_path(
pretrained_model_name_or_path
)
output_dir = dataset_validation_util.force_gcs_fuse_path(_OUTPUT_DIR.value)
finetune_sequence_classification(
pretrained_model_name_or_path=pretrained_model_path,
dataset_name=_DATASET_NAME.value,
output_dir=output_dir,
lora_rank=_LORA_RANK.value,
lora_alpha=_LORA_ALPHA.value,
lora_dropout=_LORA_DROPOUT.value,
num_train_epochs=int(_NUM_TRAIN_EPOCHS.value),
batch_size=_BATCH_SIZE.value,
learning_rate=_LEARNING_RATE.value,
)
if __name__ == "__main__":
app.run(main)
@@ -1,7 +0,0 @@
{
"description": "Chat template used by Qwen 2.5.",
"source": "https://huggingface.co/Qwen/Qwen2.5-72B-Instruct/blob/main/tokenizer_config.json#L198",
"chat_template": "{%- if tools %}\n {{- '<|im_start|>system\\n' }}\n {%- if messages[0]['role'] == 'system' %}\n {{- messages[0]['content'] }}\n {%- else %}\n {{- 'You are Qwen, created by Alibaba Cloud. You are a helpful assistant.' }}\n {%- endif %}\n {{- \"\\n\\n# Tools\\n\\nYou may call one or more functions to assist with the user query.\\n\\nYou are provided with function signatures within <tools></tools> XML tags:\\n<tools>\" }}\n {%- for tool in tools %}\n {{- \"\\n\" }}\n {{- tool | tojson }}\n {%- endfor %}\n {{- \"\\n</tools>\\n\\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\\n<tool_call>\\n{\\\"name\\\": <function-name>, \\\"arguments\\\": <args-json-object>}\\n</tool_call><|im_end|>\\n\" }}\n{%- else %}\n {%- if messages[0]['role'] == 'system' %}\n {{- '<|im_start|>system\\n' + messages[0]['content'] + '<|im_end|>\\n' }}\n {%- else %}\n {{- '<|im_start|>system\\nYou are Qwen, created by Alibaba Cloud. You are a helpful assistant.<|im_end|>\\n' }}\n {%- endif %}\n{%- endif %}\n{%- for message in messages %}\n {%- if (message.role == \"user\") or (message.role == \"system\" and not loop.first) or (message.role == \"assistant\" and not message.tool_calls) %}\n {{- '<|im_start|>' + message.role + '\\n' + message.content + '<|im_end|>' + '\\n' }}\n {%- elif message.role == \"assistant\" %}\n {{- '<|im_start|>' + message.role }}\n {%- if message.content %}\n {{- '\\n' + message.content }}\n {%- endif %}\n {%- for tool_call in message.tool_calls %}\n {%- if tool_call.function is defined %}\n {%- set tool_call = tool_call.function %}\n {%- endif %}\n {{- '\\n<tool_call>\\n{\"name\": \"' }}\n {{- tool_call.name }}\n {{- '\", \"arguments\": ' }}\n {{- tool_call.arguments | tojson }}\n {{- '}\\n</tool_call>' }}\n {%- endfor %}\n {{- '<|im_end|>\\n' }}\n {%- elif message.role == \"tool\" %}\n {%- if (loop.index0 == 0) or (messages[loop.index0 - 1].role != \"tool\") %}\n {{- '<|im_start|>user' }}\n {%- endif %}\n {{- '\\n<tool_response>\\n' }}\n {{- message.content }}\n {{- '\\n</tool_response>' }}\n {%- if loop.last or (messages[loop.index0 + 1].role != \"tool\") %}\n {{- '<|im_end|>\\n' }}\n {%- endif %}\n {%- endif %}\n{%- endfor %}\n{%- if add_generation_prompt %}\n {{- '<|im_start|>assistant\\n' }}\n{%- endif %}\n",
"instruction_separator": "<|im_start|>user\n",
"response_separator": "<|im_start|>assistant\n"
}
@@ -32,7 +32,6 @@ class DockerCommandBuilder(CommandBuilder):
super().__init__()
self._docker_uri = [docker_uri]
self.privilege_mode = []
self.entrypoint = []
self._defaults = [
'docker',
@@ -63,9 +62,6 @@ class DockerCommandBuilder(CommandBuilder):
def add_privilege_mode(self):
self.privilege_mode = ['--privileged']
def add_entrypoint(self, entrypoint: list[str]):
self.entrypoint = entrypoint
def build_cmd(self) -> str:
return (
self._defaults
@@ -73,7 +69,6 @@ class DockerCommandBuilder(CommandBuilder):
+ self._mount_maps
+ self.privilege_mode
+ self._docker_uri
+ self.entrypoint
)
@@ -90,6 +85,3 @@ class PythonCommandBuilder(CommandBuilder):
def build_cmd(self) -> str:
os.environ.update(self._env_vars)
return self._defaults
def add_entrypoint(self, entrypoint: list[str]):
self._defaults = entrypoint
@@ -64,7 +64,6 @@ class TrainerThroughputTest(test_util.TestBase):
'Mistral-7B-v0.1',
'Mixtral-8x7B-v0.1',
'gemma-2-9b-it',
'Qwen2.5-32B-Instruct',
],
precision=['4bit', '8bit', 'bfloat16'],
max_seq_length=list(range(4 * 1024, 24 * 1024 + 1, 4 * 1024)),
@@ -90,7 +89,6 @@ class TrainerThroughputTest(test_util.TestBase):
'Mistral-7B-v0.1',
'Mixtral-8x7B-v0.1',
'gemma-2-9b-it',
'Qwen2.5-32B-Instruct',
],
precision=['4bit', '8bit', 'bfloat16'],
max_seq_length=list(range(4 * 1024, 24 * 1024 + 1, 4 * 1024)),
@@ -121,11 +119,7 @@ class TrainerThroughputTest(test_util.TestBase):
self.assertEqual(self.run_cmd_and_handle_failure(), 0)
@parameterized.product(
model_name=[
'llama3.1-8b-hf',
'llama3.1-70b-hf',
'Qwen2.5-32B-Instruct',
],
model_name=['llama3.1-8b-hf', 'llama3.1-70b-hf'],
precision=['4bit', '8bit', 'bfloat16'],
max_seq_length=list(range(4 * 1024, 24 * 1024 + 1, 4 * 1024)),
num_gpus=[8],
@@ -142,16 +136,9 @@ class TrainerThroughputTest(test_util.TestBase):
self.test_suite_output_dir,
f'bm_fsdp_{num_gpus}gpu_{model_name}_{precision}.txt',
)
if 'llama' in model_name.lower():
self.task_cmd_builder.config_file = (
'vertex_vision_model_garden_peft/llama2_fsdp_8gpu.yaml'
)
elif 'qwen' in model_name.lower():
self.task_cmd_builder.config_file = (
'vertex_vision_model_garden_peft/qwen2_fsdp_8gpu.yaml'
)
else:
self.fail(f'Unsupported model: {model_name}')
self.task_cmd_builder.config_file = (
'vertex_vision_model_garden_peft/llama_fsdp_8gpu.yaml'
)
self.command_builder.add_env_var(
'CUDA_VISIBLE_DEVICES', ','.join([str(x) for x in range(0, num_gpus)])
@@ -140,70 +140,6 @@ class TrainedModelQualityTest(test_util.TestBase):
self.assertEqual(self.run_cmd(), 0)
@parameterized.named_parameters(
('Qwen2.5-32B-Instruct', 'Qwen2.5-32B-Instruct'),
)
def test_qwen_model_deepspeed(self, model_name):
self.setup_output_dir(f'test_deepspeed_{model_name}')
self.task_cmd_builder.pretrained_model_name_or_path = model_name
self.task_cmd_builder.config_file = (
'vertex_vision_model_garden_peft/deepspeed_zero2_8gpu.yaml'
)
self.task_cmd_builder.train_dataset = test_util.get_test_data_path(
'llama-tuning-test/opposite-examples-train.jsonl'
)
self.task_cmd_builder.train_split = 'train'
self.task_cmd_builder.train_column = 'messages'
self.task_cmd_builder.train_template = 'qwen2_5'
self.task_cmd_builder.eval_dataset = test_util.get_test_data_path(
'llama-tuning-test/opposite-examples-eval.jsonl'
)
self.task_cmd_builder.eval_split = 'train'
self.task_cmd_builder.eval_column = self.task_cmd_builder.train_column
self.task_cmd_builder.eval_template = self.task_cmd_builder.train_template
self.command_builder.add_env_var('CUDA_VISIBLE_DEVICES', '0,1,2,3,4,5,6,7')
# Note(lavrai): The following parameters are needed for the opposite-word
# dataset to converge properly.
self.task_cmd_builder.gradient_accumulation_steps = 1
self.task_cmd_builder.num_train_epochs = 10.0
self.task_cmd_builder.logging_steps = 1
self.assertEqual(self.run_cmd(), 0)
@parameterized.named_parameters(
('Qwen2.5-32B-Instruct', 'Qwen2.5-32B-Instruct'),
)
def test_qwen_model_fsdp(self, model_name):
self.setup_output_dir(f'test_fsdp_{model_name}')
self.task_cmd_builder.pretrained_model_name_or_path = (
test_util.get_pretrained_model_name_or_path(model_name)
)
self.task_cmd_builder.config_file = (
'vertex_vision_model_garden_peft/qwen2_fsdp_8gpu.yaml'
)
self.task_cmd_builder.train_dataset = test_util.get_test_data_path(
'llama-tuning-test/opposite-examples-train.jsonl'
)
self.task_cmd_builder.train_split = 'train'
self.task_cmd_builder.train_column = 'messages'
self.task_cmd_builder.train_template = 'qwen2_5'
self.task_cmd_builder.eval_dataset = test_util.get_test_data_path(
'llama-tuning-test/opposite-examples-eval.jsonl'
)
self.task_cmd_builder.eval_split = 'train'
self.task_cmd_builder.eval_column = self.task_cmd_builder.train_column
self.task_cmd_builder.eval_template = self.task_cmd_builder.train_template
self.command_builder.add_env_var('CUDA_VISIBLE_DEVICES', '0,1,2,3,4,5,6,7')
# Note(lavrai): The following parameters are needed for the opposite-word
# dataset to converge properly.
self.task_cmd_builder.gradient_accumulation_steps = 1
self.task_cmd_builder.num_train_epochs = 10.0
self.task_cmd_builder.logging_steps = 1
self.assertEqual(self.run_cmd(), 0)
if __name__ == '__main__':
absltest.main()
@@ -9,17 +9,19 @@ environment. Otherwise, `python3` is used.
import argparse
from collections.abc import MutableSequence, Sequence
import json
import multiprocessing
import os
import subprocess
import sys
from absl import app
from absl import flags
from absl import logging
from util import dataset_validation_util
from util import cluster_spec
from vertex_vision_model_garden_peft.train.vmg import gcs_syncer
from vertex_vision_model_garden_peft.train.vmg import utils
from util import constants
from util import gcs_syncer
from util import fileutils
from util import hypertune_utils
@@ -39,12 +41,9 @@ _TASK_TO_SCRIPT = {
constants.INSTRUCT_LORA: (
'vertex_vision_model_garden_peft/train/vmg/instruct_lora.py'
),
constants.MERGE_CAUSAL_LANGUAGE_MODEL_LORA: (
'vertex_vision_model_garden_peft/train/vmg/merge_causal_language_model_lora.py'
),
constants.VALIDATE_DATASET_WITH_TEMPLATE: (
'vertex_vision_model_garden_peft/train/vmg/validate_dataset_with_template.py'
),
constants.MERGE_CAUSAL_LANGUAGE_MODEL_LORA: 'vertex_vision_model_garden_peft/train/vmg/merge_causal_language_model_lora.py',
constants.SEQUENCE_CLASSIFICATION_LORA: 'vertex_vision_model_garden_peft/train/vmg/sequence_classification_lora.py',
constants.VALIDATE_DATASET_WITH_TEMPLATE: 'vertex_vision_model_garden_peft/train/vmg/validate_dataset_with_template.py',
constants.RUN_TESTS: 'vertex_vision_model_garden_peft/tests/run_tests.py',
}
@@ -72,17 +71,53 @@ def launch_script_cmd(
def _get_accelerate_args() -> argparse.Namespace:
"""Returns the accelerate args."""
primary_node_addr, primary_node_port, node_rank, num_nodes = (
cluster_spec.get_cluster_spec()
)
# For the format of the cluster spec, see
# https://cloud.google.com/vertex-ai/docs/training/distributed-training#cluster-spec-format # pylint: disable=line-too-long
cluster_spec = os.getenv('CLUSTER_SPEC', default=None)
if not cluster_spec:
return argparse.Namespace()
logging.info('CLUSTER_SPEC: %s', cluster_spec)
cluster_data = json.loads(cluster_spec)
if (
'workerpool1' not in cluster_data['cluster']
or not cluster_data['cluster']['workerpool1']
):
return argparse.Namespace()
# Get primary node info
primary_node = cluster_data['cluster']['workerpool0'][0]
logging.info('primary node: %s', primary_node)
primary_node_addr, primary_node_port = primary_node.split(':')
logging.info('primary node address: %s', primary_node_addr)
logging.info('primary node port: %s', primary_node_port)
# Determine node rank of this machine
workerpool = cluster_data['task']['type']
if workerpool == 'workerpool0':
node_rank = 0
elif workerpool == 'workerpool1':
# Add 1 for the primary node, since `index` is the index of workerpool1.
node_rank = cluster_data['task']['index'] + 1
else:
raise ValueError(
'Only workerpool0 and workerpool1 are supported. Unknown workerpool:'
f' {workerpool}'
)
logging.info('node rank: %s', node_rank)
# Calculate total nodes
num_worker_nodes = len(cluster_data['cluster']['workerpool1'])
num_nodes = num_worker_nodes + 1 # Add 1 for the primary node
logging.info('num nodes: %s', num_nodes)
accelerate_args = argparse.Namespace()
if num_nodes > 1:
accelerate_args.machine_rank = node_rank
accelerate_args.num_machines = num_nodes
accelerate_args.main_process_ip = primary_node_addr
accelerate_args.main_process_port = primary_node_port
accelerate_args.max_restarts = 0
accelerate_args.monitor_interval = 120
accelerate_args.machine_rank = node_rank
accelerate_args.num_machines = num_nodes
accelerate_args.main_process_ip = primary_node_addr
accelerate_args.main_process_port = primary_node_port
accelerate_args.max_restarts = 0
accelerate_args.monitor_interval = 120
return accelerate_args
@@ -96,6 +131,45 @@ def _append_args_to_command_in_place(
command.append(f'--{key}={value}')
def _is_gcs_or_gcsfuse_path(path: str) -> bool:
"""Returns if the path is a GCS or gcsfuse path.
Args:
path: The path to check.
Returns:
True if the path is a GCS or gcsfuse path.
"""
return path.startswith(
(constants.GCS_URI_PREFIX, constants.GCSFUSE_URI_PREFIX)
)
def _manage_training_path(path: str, node_rank: int) -> tuple[str, str]:
"""Returns local dir and GCS location for the given path if the given path is a GCS or gcsfuse path.
It will also create a local directory if it does not exist. Othereise, it
returns the same path.
Args:
path: The local or GCS path to manage.
node_rank: The node rank to be appended to the GCS path.
Returns:
The local and GCS paths.
"""
local_dir = path
gcs_dir = path
if _is_gcs_or_gcsfuse_path(path):
local_dir = os.path.join(
constants.LOCAL_OUTPUT_DIR,
dataset_validation_util.force_gcs_fuse_path(path)[1:],
)
gcs_dir = fileutils.force_gcs_path(path)
os.makedirs(local_dir, exist_ok=True)
return local_dir, os.path.join(gcs_dir, f'node-{node_rank}')
def _get_train_and_maybe_merge_cmd_and_dirs_to_sync(
task_type: str, config_file: str, unknown: Sequence[str]
) -> Sequence[Sequence[str]]:
@@ -129,11 +203,11 @@ def _get_train_and_maybe_merge_cmd_and_dirs_to_sync(
dataset_validation_util.force_gcs_fuse_path(training_args.output_dir)
)
local_output_dir, gcs_output_dir = gcs_syncer.manage_sync_path(
local_output_dir, gcs_output_dir = _manage_training_path(
training_args.output_dir, node_rank
)
training_args.output_dir = local_output_dir
if gcs_syncer.is_gcs_or_gcsfuse_path(gcs_output_dir):
if _is_gcs_or_gcsfuse_path(gcs_output_dir):
dirs_to_sync.append((local_output_dir, gcs_output_dir))
# Merge only flags.
@@ -143,11 +217,11 @@ def _get_train_and_maybe_merge_cmd_and_dirs_to_sync(
merge_args, unknown = merge_parser.parse_known_args(unknown)
if merge_args.merge_base_and_lora_output_dir:
merge_local_dir, merge_gcs_dir = gcs_syncer.manage_sync_path(
merge_args.merge_base_and_lora_output_dir, None
merge_local_dir, merge_gcs_dir = _manage_training_path(
merge_args.merge_base_and_lora_output_dir, node_rank
)
merge_args.merge_base_and_lora_output_dir = merge_local_dir
if gcs_syncer.is_gcs_or_gcsfuse_path(merge_gcs_dir):
if _is_gcs_or_gcsfuse_path(merge_gcs_dir):
dirs_to_sync.append((merge_local_dir, merge_gcs_dir))
# Common flags shared by merging and training.
@@ -165,10 +239,8 @@ def _get_train_and_maybe_merge_cmd_and_dirs_to_sync(
# Only the main node runs merging.
if merge_args.merge_base_and_lora_output_dir and node_rank == 0:
lora_dir = utils.get_final_checkpoint_path(training_args.output_dir)
lora_local_dir, lora_gcs_dir = gcs_syncer.manage_sync_path(
lora_dir, node_rank
)
if gcs_syncer.is_gcs_or_gcsfuse_path(lora_gcs_dir):
lora_local_dir, lora_gcs_dir = _manage_training_path(lora_dir, node_rank)
if _is_gcs_or_gcsfuse_path(lora_gcs_dir):
dirs_to_sync.append((lora_local_dir, lora_gcs_dir))
merge_cmd = [
@@ -191,37 +263,46 @@ def _get_train_and_maybe_merge_cmd_and_dirs_to_sync(
return commands, dirs_to_sync
def _get_merge_cmd_and_dirs_to_sync(
task_type: str, config_file: str, unknown: Sequence[str]
) -> Sequence[Sequence[str]]:
"""Returns the merge command and dirs to sync.
def _setup_gcs_rsync(
dirs_to_sync: Sequence[tuple[str, str]],
mp_queue: multiprocessing.Queue,
gcs_rsync_interval_secs: int,
) -> multiprocessing.Process:
"""Sets up the GCS rsync process.
Args:
task_type: The task type.
config_file: The accelerate config file path.
unknown: The unknown args which are not recognised by the parser.
dirs_to_sync: The absolute directory paths which will be synced to GCS.
mp_queue: The multiprocessing queue to check if the training is finished.
gcs_rsync_interval_secs: Integer, interval in seconds to run gcs rsync.
Returns:
The bash commands to execute and the directories to sync.
The GCS rsync process.
"""
# Merge only flags.
merge_parser = argparse.ArgumentParser()
merge_parser.add_argument('--merge_base_and_lora_output_dir')
merge_args, unknown = merge_parser.parse_known_args(unknown)
rsync_process = multiprocessing.Process(
target=gcs_syncer.start_gcs_rsync,
args=(dirs_to_sync, mp_queue, gcs_rsync_interval_secs),
)
rsync_process.start()
return rsync_process
dirs_to_sync = []
if merge_args.merge_base_and_lora_output_dir:
merge_local_dir, merge_gcs_dir = gcs_syncer.manage_sync_path(
merge_args.merge_base_and_lora_output_dir, None
def _cleanup_gcs_rsync(
rsync_process: multiprocessing.Process, mp_queue: multiprocessing.Queue
) -> None:
"""Cleans up the GCS rsync process.
Args:
rsync_process: The GCS rsync process.
mp_queue: The multiprocessing queue.
"""
mp_queue.put('training finished')
rsync_process.join()
if rsync_process.exitcode == 0:
logging.info('Artifacts have been uploaded to GCS.')
else:
logging.error(
'GCS rsync process failed with exit code %d.', rsync_process.exitcode
)
merge_args.merge_base_and_lora_output_dir = merge_local_dir
if gcs_syncer.is_gcs_or_gcsfuse_path(merge_gcs_dir):
dirs_to_sync.append((merge_local_dir, merge_gcs_dir))
cmd = launch_script_cmd(_TASK_TO_SCRIPT[task_type], config_file)
_append_args_to_command_in_place(merge_args, cmd)
cmd.extend(unknown)
return [cmd], dirs_to_sync
def main(unused_argv: Sequence[str]) -> None:
@@ -254,10 +335,6 @@ def main(unused_argv: Sequence[str]) -> None:
commands, dirs_to_sync = _get_train_and_maybe_merge_cmd_and_dirs_to_sync(
task_type=task, config_file=args.config_file, unknown=unknown
)
elif task in [constants.MERGE_CAUSAL_LANGUAGE_MODEL_LORA]:
commands, dirs_to_sync = _get_merge_cmd_and_dirs_to_sync(
task_type=task, config_file=args.config_file, unknown=unknown
)
else:
assert task in _TASK_TO_SCRIPT
cmd = launch_script_cmd(_TASK_TO_SCRIPT[task], args.config_file)
@@ -267,7 +344,7 @@ def main(unused_argv: Sequence[str]) -> None:
rsync_process = None
mp_queue = multiprocessing.Queue(maxsize=1)
if dirs_to_sync:
rsync_process = gcs_syncer.setup_gcs_rsync(
rsync_process = _setup_gcs_rsync(
dirs_to_sync, mp_queue, args.gcs_rsync_interval_secs
)
@@ -284,7 +361,7 @@ def main(unused_argv: Sequence[str]) -> None:
rsync_process.terminate()
raise e
if rsync_process is not None:
gcs_syncer.cleanup_gcs_rsync(rsync_process, mp_queue)
_cleanup_gcs_rsync(rsync_process, mp_queue)
if __name__ == '__main__':
@@ -1,6 +1,7 @@
"""Common libraries for PEFT."""
from collections.abc import Mapping, Sequence
import dataclasses
import datetime
import gc
import os
@@ -10,9 +11,12 @@ from absl import logging
import accelerate
from accelerate import DistributedType
from accelerate import PartialState
import numpy as np
import peft
from peft import PeftModel
from peft import prepare_model_for_kbit_training
import psutil
import pynvml
import torch
import transformers
from transformers import AutoModelForCausalLM
@@ -24,6 +28,7 @@ import trl
from util import dataset_validation_util
from util import constants
_LLAMA_3_1_405B_MODEL_ID = "Meta-Llama-3.1-405B"
_LOCAL_MERGED_MODEL_DIR = "/tmp/merged_model"
_GEMMA2_MODEL = "gemma-2"
@@ -121,7 +126,7 @@ def load_model(
"device_map": device_map,
"torch_dtype": torch_dtype,
"quantization_config": quantization_config,
"trust_remote_code": False,
"trust_remote_code": True,
"token": access_token,
"attn_implementation": attn_implementation,
}
@@ -305,12 +310,171 @@ def convert_model_to_fp8(
PartialState().wait_for_everyone()
@dataclasses.dataclass
class TuningDataStats:
tuning_dataset_example_count: int
total_billable_token_count: int
tuning_step_count: int
def get_dataset_stats(
dataset: Any,
tokenizer: transformers.PreTrainedTokenizer,
column: str,
effective_batch_size: int,
) -> TuningDataStats:
"""Calculates dataset statistics, e.g., total number of tokens."""
tokenized_dataset = dataset.map(lambda x: tokenizer(x[column]))
inputs = tokenized_dataset["input_ids"]
tuning_dataset_example_count = int(len(inputs))
total_billable_token_count = int(np.sum([len(ex) for ex in inputs]))
tuning_step_count = (
tuning_dataset_example_count + effective_batch_size - 1
) // effective_batch_size
return TuningDataStats(
tuning_dataset_example_count,
total_billable_token_count,
tuning_step_count,
)
def force_gc():
"""Collects garbage immediately to release unused CPU/GPU resources."""
gc.collect()
torch.cuda.empty_cache()
@dataclasses.dataclass
class GpuStats:
"""Holds information about GPU usage stats.
For memory related, see
https://pytorch.org/docs/stable/notes/cuda.html#cuda-memory-management
"""
# total memory
total_mem: float
# memory occupied.
occupied: float
# memory reserved, but not used.
unused: float
# nvidia-smi usually reports more memory usages than pytorch (for driver,
# kernel and etc). `smi_diff` tracks this difference.
smi_diff: float
# Gpu utilization.
util: float
# Allows unpacking operation like
# total_mem, occupied, unused, smi_diff, util = GpuStats(...)
# See https://stackoverflow.com/a/70753113
def __iter__(self):
return iter(dataclasses.astuple(self))
def gpu_stats() -> GpuStats:
"""Reports GPU memory usage and utilization."""
# See https://pytorch.org/docs/stable/notes/cuda.html#memory-management
bytes_per_gb = 1024.0**3
device = torch.cuda.current_device()
occupied = torch.cuda.memory_allocated(device) / bytes_per_gb
reserved = torch.cuda.memory_reserved(device) / bytes_per_gb
unused = reserved - occupied
def smi_mem(device):
try:
pynvml.nvmlInit()
handle = pynvml.nvmlDeviceGetHandleByIndex(device)
info = pynvml.nvmlDeviceGetMemoryInfo(handle)
return info.used / bytes_per_gb
except pynvml.NVMLError:
return 0.0
mem_used_smi = smi_mem(device)
smi_diff = mem_used_smi - reserved
util = torch.cuda.utilization(device)
return GpuStats(mem_used_smi, occupied, unused, smi_diff, util)
def gpu_stats_str(stats: GpuStats | None = None) -> str:
if stats is None:
stats = gpu_stats()
total, occupied, unused, smi_diff, util = stats
return (
f"GPU memory: {total:.2f}({occupied=:.2f}, {unused=:.2f},"
f" {smi_diff=:.2f}) GB. Utilization: {util:.2f}%"
)
@dataclasses.dataclass
class CpuStats:
"""Holds information about CPU usage stats."""
# Total CPU virtual memory i.e. virtual memory allocated + unallocated.
total_virtual_mem: float
# CPU virtual memory available for use.
unallocated_virtual_mem: float
# CPU virtual memory already used.
allocated_virtual_mem: float
# Total CPU swap memory i.e. swap memory allocated + unallocated.
total_swap_mem: float
# CPU swap memory available for use.
unallocated_swap_mem: float
# CPU swap memory already used.
allocated_swap_mem: float
# CPU utilization percentage.
utilization: float
def cpu_stats() -> CpuStats:
"""Reports CPU memory usage and utilization."""
# https://psutil.readthedocs.io/en/latest/#memory
gb = 1024.0**3
vmem = psutil.virtual_memory()
vmem_total = vmem.total / gb
vmem_available = vmem.available / gb
vmem_used = vmem_total - vmem_available
smem = psutil.swap_memory()
swap_total = smem.total / gb
swap_free = smem.free / gb
swap_used = smem.used / gb
# https://psutil.readthedocs.io/en/latest/#psutil.cpu_percent
cpu_util = psutil.cpu_percent(interval=1e-6)
return CpuStats(
total_virtual_mem=vmem_total,
unallocated_virtual_mem=vmem_available,
allocated_virtual_mem=vmem_used,
total_swap_mem=swap_total,
unallocated_swap_mem=swap_free,
allocated_swap_mem=swap_used,
utilization=cpu_util,
)
def cpu_stats_str(stats: CpuStats | None = None) -> str:
"""Returns a string representation of the CPU stats."""
if stats is None:
stats = cpu_stats()
total, occupied, unused = (
stats.total_virtual_mem,
stats.allocated_virtual_mem,
stats.unallocated_virtual_mem,
)
virtual_mem = (
f"CPU virtual memory: {total:.2f}({occupied=:.2f}, {unused=:.2f}) GB"
)
total, occupied, unused = (
stats.total_swap_mem,
stats.allocated_swap_mem,
stats.unallocated_swap_mem,
)
swap_mem = f"CPU swap memory: {total:.2f}({occupied=:.2f}, {unused=:.2f}) GB"
percent = stats.utilization
return f"{virtual_mem} {swap_mem} CPU Utilization: {percent:.2f}%"
def init_partial_state(
timeout: datetime.timedelta = datetime.timedelta(seconds=600),
) -> None:
@@ -1,12 +1,9 @@
"""Fileutil lib to copy files between gcs and local."""
import filecmp
import fnmatch
import os
import pathlib
import shutil
import subprocess
import time
from typing import List, Optional, Tuple
import uuid
@@ -60,96 +57,6 @@ def force_gcs_path(uri: str) -> str:
return uri
def is_file_available(
file_path: str, retry_interval_secs: int = 60, timeout_secs: int = 3600
) -> bool:
"""Checks and waits for a file to be available in GCS.
Args:
file_path: The file path to check.
retry_interval_secs: The interval in seconds to check the file.
timeout_secs: The timeout in seconds to wait for the file.
Returns:
True if the file is available, False otherwise.
"""
start_time = time.time()
while True:
try:
file_check_cmd = ['gcloud', 'storage', 'ls', file_path]
result = subprocess.run(
file_check_cmd, capture_output=True, text=True, check=True
)
if file_path in result.stdout:
logging.info('File %s exists.', file_path)
return True
except subprocess.CalledProcessError as e:
elapsed_time = time.time() - start_time
if elapsed_time > timeout_secs:
logging.info(
"Timeout: File '%s' not found after %d seconds. Error: %s",
file_path,
elapsed_time,
e,
)
return False
logging.info(
"File '%s' not found yet. Checking again in %d seconds. Error: %s",
file_path,
retry_interval_secs,
e,
)
time.sleep(retry_interval_secs)
def compare_dirs(
local_dir: str,
gcsfuse_dir: str,
retry_interval_secs: int = 30,
timeout_secs: int = 3600,
) -> bool:
"""Compares two directories and returns True if they are the same.
Args:
local_dir: The local directory.
gcsfuse_dir: The gcsfuse directory.
retry_interval_secs: The interval in seconds to check the directories.
timeout_secs: The timeout in seconds to wait for the directories.
Returns:
True if the directories are the same, False otherwise.
"""
start_time = time.time()
while True:
if os.path.exists(local_dir) and os.path.exists(gcsfuse_dir):
comparison = filecmp.dircmp(local_dir, gcsfuse_dir)
if (
not comparison.left_only
and not comparison.right_only
and not comparison.diff_files
):
return True
elapsed_time = time.time() - start_time
if elapsed_time > timeout_secs:
logging.info(
"Timeout: Directories '%s' and '%s' do not match after %d seconds.",
local_dir,
gcsfuse_dir,
elapsed_time,
)
return False
logging.info(
"Directories '%s' and '%s' do not match yet. Checking again in %d"
' seconds.',
local_dir,
gcsfuse_dir,
retry_interval_secs,
)
time.sleep(retry_interval_secs)
def download_gcs_file_to_memory(gcs_uri: str) -> bytes:
"""Downloads a gcs file to in memory.
@@ -445,15 +352,3 @@ def get_output_video_file(video_output_file_path: str) -> str:
file_extension, '_overlay' + file_extension
)
return out_local_video_file_name
def delete_local_file(local_file_path: str) -> None:
"""Deletes a local file."""
if os.path.exists(local_file_path):
os.remove(local_file_path)
def delete_local_dir(local_dir: str) -> None:
"""Deletes a local directory recursively."""
if os.path.exists(local_dir):
shutil.rmtree(local_dir)
@@ -1,119 +0,0 @@
#!/bin/bash
#
# This launcher downloads model files from GCS to local model directory before
# launching the actual command.
#
# If GCS URI is passed as an environment variable, set GCS_URI_ENV_KEY to the
# environment variable name.
# If GCS URI is passed as an argument, set GCS_URI_ARG_KEY to the argument name.
# The argument must be in the format of '--$GCS_URI_ARG_KEY=gs://*'. Do not
# separate argument name and value with spaces.
# This script will also try reading from AIP_STORAGE_URI or AIP_STORAGE_DIR.
# Note that AIP_STORAGE_DIR is expected to be a local path, so it bypasses the
# download process.
#
# Input priority: AIP_STORAGE_DIR > AIP_STORAGE_URI > GCS_URI_ENV_KEY > GCS_URI_ARG_KEY.
# Will output the local model directory to GCS_URI_ENV_KEY and GCS_URI_ARG_KEY
# if they are set. Both will be updated if both set.
#
# Requires google-cloud-sdk as a dependency (for gcloud storage CLI).
set -e
readonly LOCAL_MODEL_DIR=${LOCAL_MODEL_DIR:-"/tmp/model_dir"}
readonly LOCAL_ARGS_FILE=${LOCAL_ARGS_FILE:-"/tmp/args.txt"}
update_model_id() {
if [[ ! -z "$GCS_URI_ENV_KEY" ]]; then
echo "Updating env var $GCS_URI_ENV_KEY to $AIP_STORAGE_DIR."
export "$GCS_URI_ENV_KEY"="$AIP_STORAGE_DIR"
fi
if [[ ! -z "$GCS_URI_ARG_KEY" ]]; then
echo "Updating args $GCS_URI_ARG_KEY to $AIP_STORAGE_DIR."
updated=0
for (( i=1; i <= $#; i++)); do
arg="${!i}"
if [[ "$arg" == "--$GCS_URI_ARG_KEY="* ]]; then
echo "Found $arg, updating to $AIP_STORAGE_DIR."
set -- "${@:1:(($i-1))}" "--$GCS_URI_ARG_KEY=$AIP_STORAGE_DIR" "${@:$(($i+1))}";
updated=1
break
fi
done
if [[ $updated -eq 0 ]]; then
echo "Appending args $GCS_URI_ARG_KEY to $AIP_STORAGE_DIR."
set -- "$@" "--$GCS_URI_ARG_KEY=$AIP_STORAGE_DIR";
fi
fi
echo "$*" > "$LOCAL_ARGS_FILE"
}
maybe_download_model() {
if [[ -z "$GCS_URI_ENV_KEY" ]] && [[ -z "$GCS_URI_ARG_KEY" ]]; then
echo "Internal error: Required GCS_URI_ENV_KEY or GCS_URI_ARG_KEY."
exit 1
fi
echo "$*" > "$LOCAL_ARGS_FILE"
gcs_uri=""
if [[ ! -z "$AIP_STORAGE_DIR" ]]; then
# AIP_STORAGE_DIR is expected to be a local path.
echo "AIP_STORAGE_DIR set, proceeding to run the launcher."
update_model_id "$@"
return
elif [[ $AIP_STORAGE_URI == gs://* ]]; then
# Check AIP_STORAGE_URI environment variable.
echo "AIP_STORAGE_URI set and starts with 'gs://', proceeding to download from GCS."
gcs_uri="$AIP_STORAGE_URI"
elif [[ ! -z "$GCS_URI_ENV_KEY" ]] && [[ ${!GCS_URI_ENV_KEY} == gs://* ]]; then
# Check custom environment variable.
echo "Custom environment variable ${GCS_URI_ENV_KEY} set and starts with 'gs://', proceeding to download from GCS."
gcs_uri="${!GCS_URI_ENV_KEY}"
elif [[ ! -z "$GCS_URI_ARG_KEY" ]]; then
# Check custom args.
for arg in "$@"; do
if [[ "$arg" == "--$GCS_URI_ARG_KEY=gs://"* ]]; then
gcs_uri="${arg#*=}"
echo "Custom args ${GCS_URI_ARG_KEY} set and starts with 'gs://', proceeding to download from GCS."
break
elif [[ "$arg" == "--$GCS_URI_ARG_KEY" ]]; then
echo "Found $GCS_URI_ARG_KEY, but it's not in the format of '--$GCS_URI_ARG_KEY=gs://*'."
echo "Ensure the value of $GCS_URI_ARG_KEY is within the same arg, separated by '='."
exit 1
fi
done
fi
if [[ -z "$gcs_uri" ]]; then
echo "No GCS URI found, proceeding to run the launcher."
return
fi
# Remove trailing '/' if any.
gcs_uri="${gcs_uri%%/}"
export AIP_STORAGE_DIR="$LOCAL_MODEL_DIR/${gcs_uri##gs://}"
# Create the target directory.
mkdir -p "$AIP_STORAGE_DIR"
echo "Downloading model from ${gcs_uri} to ${AIP_STORAGE_DIR}."
# Use gcloud storage CLI to copy the content from GCS to the target directory.
if gcloud storage cp -r "$gcs_uri/*" "$AIP_STORAGE_DIR"; then
echo "Model downloaded successfully to ${AIP_STORAGE_DIR}."
update_model_id "$@"
else
echo "Failed to download model from GCS."
exit 1
fi
}
run_local_command() {
command=$(cat "$LOCAL_ARGS_FILE")
rm -f "$LOCAL_ARGS_FILE"
echo "Launch command: $command"
eval "$command"
}
maybe_download_model "$@"
run_local_command
@@ -1,35 +0,0 @@
#!/bin/bash
# !/bin/bash
# The Startup prober built to check whether models listed in local disk are
# loaded in memory and are ready to serve traffic. The script returns 0 if
# succeed. Any other returned value are consider as an error. More detail could be
# found from [shell script Exit codes](http://shellscript.sh/exitcodes.html).
#
# TorchServe: The Management API listens on port 8081 and is only accessible
# from localhost by default.
if [[ -z "${MNG_PORT}" ]]; then
MNG_PORT=7081 # We default the management_port to 7081.
else
MNG_PORT="${MNG_PORT}"
fi
check_model_availability(){
local MODEL_NAME=$1
# Returns whether "READY" is found in the model status.
# Reference: https://pytorch.org/serve/management_api.html#describe-model.
curl -s "http://localhost:${MNG_PORT}/models/${MODEL_NAME}" | grep "READY" -q
}
main(){
check_model_availability "$MODEL" # Assume Dockerfile sets MODEL environment parameter.
local available=$?
if [[ $available -gt 0 ]]
then
echo "Warning: Model(${MODEL}) is not yet available."
return 1
fi
return 0
}
main
@@ -1,746 +0,0 @@
"""Common util functions for notebook."""
import base64
from collections.abc import Sequence
import datetime
import io
import json
import os
import subprocess
import time
from typing import Any
from google import auth
from google.cloud import storage
import matplotlib.pyplot as plt
import numpy as np
from PIL import Image
import requests
import tensorflow as tf
import yaml
GCS_URI_PREFIX = "gs://"
CHECKPOINT_BUCKET = "gs://model_garden_checkpoints"
def convert_numpy_array_to_byte_string_via_tf_tensor(
np_array: np.ndarray,
) -> str:
"""Serializes a numpy array to tensor bytes.
Args:
np_array: A numpy array.
Returns:
A tensor bytes.
"""
tensor_array = tf.convert_to_tensor(np_array)
tensor_byte_string = tf.io.serialize_tensor(tensor_array)
return tensor_byte_string.numpy()
def get_jpeg_bytes(local_image_path: str, new_width: int = -1) -> bytes:
"""Returns jpeg bytes given an image path and resizes if required.
Args:
local_image_path: A string of local image path.
new_width: An integer of new image width.
Returns:
A jpeg bytes.
"""
image = Image.open(local_image_path)
if new_width <= 0:
new_image = image
else:
width, height = image.size
print("original input image size: ", width, " , ", height)
new_height = int(height * new_width / width)
print("new input image size: ", new_width, " , ", new_height)
new_image = image.resize((new_width, new_height))
buffered = io.BytesIO()
new_image.save(buffered, format="JPEG")
return buffered.getvalue()
def gcs_fuse_path(path: str) -> str:
"""Try to convert path to gcsfuse path if it starts with gs:// else do not modify it.
Args:
path: A string of path.
Returns:
A gcsfuse path.
"""
path = path.strip()
if path.startswith("gs://"):
return "/gcs/" + path[5:]
return path
def get_job_name_with_datetime(prefix: str) -> str:
"""Gets a job name by adding current time to prefix.
Args:
prefix: A string of job name prefix.
Returns:
A job name.
"""
now = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
job_name = f"{prefix}-{now}".replace("_", "-")
return job_name
def create_job_name(prefix: str) -> str:
"""Creates a job name.
Args:
prefix: A string of job name prefix.
Returns:
A job name.
"""
user = os.environ.get("USER")
now = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
job_name = f"{prefix}-{user}-{now}".replace("_", "-")
return job_name
def save_subset_annotation(
input_annotation_path: str, output_annotation_path: str
):
"""Saves a subset of COCO annotation json file with CCA 4.0 license.
Args:
input_annotation_path: A string of input annotation path.
output_annotation_path: A string of output annotation path.
"""
with open(input_annotation_path) as f:
coco_json = json.load(f)
img_ids = set()
images = []
annotations = []
for img in coco_json["images"]:
if img["license"] in [4, 5]: # CCA 4.0 license.
img_ids.add(img["id"])
images.append(img)
for ann in coco_json["annotations"]:
if ann["image_id"] in img_ids:
annotations.append(ann)
new_json = {
"info": coco_json["info"],
"licenses": coco_json["licenses"],
"images": images,
"annotations": annotations,
"categories": coco_json["categories"],
}
with open(output_annotation_path, "w") as f:
json.dump(new_json, f)
def image_to_base64(image: Any, image_format: str = "JPEG") -> str:
"""Converts an image to base64.
Args:
image: A PIL.Image instance.
image_format: A string of image format.
Returns:
A base64 string.
"""
buffer = io.BytesIO()
image.save(buffer, format=image_format)
image_str = base64.b64encode(buffer.getvalue()).decode("utf-8")
return image_str
def base64_to_image(image_str: str) -> Any:
"""Convert base64 encoded string to an image.
Args:
image_str: A string of base64 encoded image.
Returns:
A PIL.Image instance.
"""
image = Image.open(io.BytesIO(base64.b64decode(image_str)))
return image
def image_grid(imgs: Sequence[Any], rows: int = 2, cols: int = 2) -> Any:
"""Creates an image grid.
Args:
imgs: A list of PIL.Image instances.
rows: An integer of number of rows.
cols: An integer of number of columns.
Returns:
A PIL.Image instance.
"""
w, h = imgs[0].size
grid = Image.new(
mode="RGB", size=(cols * w + 10 * cols, rows * h), color=(255, 255, 255)
)
for i, img in enumerate(imgs):
grid.paste(img, box=(i % cols * w + 10 * i, i // cols * h))
return grid
def display_image(image: Any):
"""Displays an image.
Args:
image: A PIL.Image instance.
"""
_ = plt.figure(figsize=(20, 15))
plt.grid(False)
plt.imshow(image)
def download_gcs_file_to_local(gcs_uri: str, local_path: str):
"""Download a gcs file to a local path.
Args:
gcs_uri: A string of file path on GCS.
local_path: A string of local file path.
"""
if not gcs_uri.startswith(GCS_URI_PREFIX):
raise ValueError(
f"{gcs_uri} is not a GCS path starting with {GCS_URI_PREFIX}."
)
client = storage.Client()
os.makedirs(os.path.dirname(local_path), exist_ok=True)
with open(local_path, "wb") as f:
client.download_blob_to_file(gcs_uri, f)
def download_image(url: str) -> str:
"""Downloads an image from the given URL.
Args:
url: The URL of the image to download.
Returns:
base64 encoded image.
"""
response = requests.get(url)
return Image.open(io.BytesIO(response.content)) # pytype: disable=bad-return-type # pillow-102-upgrade
def resize_image(image: Any, new_width: int = 1000) -> Any:
"""Resizes an image to a certain width.
Args:
image: The image which has to be resized.
new_width: New width of the image.
Returns:
New resized image.
"""
width, height = image.size
new_height = int(height * new_width / width)
new_img = image.resize((new_width, new_height))
return new_img
def load_img(path: str) -> Any:
"""Reads image from path and return PIL.Image instance.
Args:
path: A string of image path.
Returns:
A PIL.Image instance.
"""
img = tf.io.read_file(path)
img = tf.image.decode_jpeg(img, channels=3)
return Image.fromarray(np.uint8(img)).convert("RGB")
def decode_image(
image_str_tensor: tf.string, new_height: int, new_width: int
) -> tf.float32:
"""Converts and resizes image bytes to image tensor.
Args:
image_str_tensor: A string of image bytes.
new_height: An integer of new image height.
new_width: An integer of new image width.
Returns:
An image tensor.
"""
image = tf.io.decode_image(image_str_tensor, 3, expand_animations=False)
image = tf.image.resize(image, (new_height, new_width))
return image
def get_label_map(label_map_yaml_filepath: str) -> dict[int, str]:
"""Returns class id to label mapping given a filepath to the label map.
Args:
label_map_yaml_filepath: A string of label map yaml file path.
Returns:
A dictionary of class id to label mapping.
"""
with tf.io.gfile.GFile(label_map_yaml_filepath, "rb") as input_file:
label_map = yaml.safe_load(input_file.read())["label_map"]
return label_map
def get_prediction_instances(test_filepath: str, new_width: int = -1) -> Any:
"""Generate instance from image path to pass to Vertex AI Endpoint for prediction.
Args:
test_filepath: A string of test image path.
new_width: An integer of new image width.
Returns:
A list of instances.
"""
if new_width <= 0:
with tf.io.gfile.GFile(test_filepath, "rb") as input_file:
encoded_string = base64.b64encode(input_file.read()).decode("utf-8")
else:
img = load_img(test_filepath)
width, height = img.size
print("original input image size: ", width, " , ", height)
new_height = int(height * new_width / width)
new_img = img.resize((new_width, new_height))
print("resized input image size: ", new_width, " , ", new_height)
buffered = io.BytesIO()
new_img.save(buffered, format="JPEG")
encoded_string = base64.b64encode(buffered.getvalue()).decode("utf-8")
instances = [{
"encoded_image": {"b64": encoded_string},
}]
return instances
def vqa_predict(
endpoint: Any,
question_prompts: Sequence[str],
image: Any,
language_code: str = "en",
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> Sequence[str]:
"""Predicts the answer to a question about an image using an Endpoint."""
# Resize and convert image to base64 string.
resized_image = resize_image(image, new_width)
resized_image_base64 = image_to_base64(resized_image)
instances = []
if question_prompts:
# Format question prompt
question_prompt_format = "answer {} {}\n"
for question_prompt in question_prompts:
if question_prompt:
instances.append({
"prompt": question_prompt_format.format(
language_code, question_prompt
),
"image": resized_image_base64,
})
else:
instances.append({
"image": resized_image_base64,
})
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
return [pred.get("response") for pred in response.predictions]
def caption_predict(
endpoint: Any,
language_code: str,
image: Any,
caption_prompt: bool = False,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Predicts a caption for a given image using an Endpoint."""
# Resize and convert image to base64 string.
resized_image = resize_image(image, new_width)
resized_image_base64 = image_to_base64(resized_image)
instance = {"image": resized_image_base64}
if caption_prompt:
# Format caption prompt
caption_prompt_format = "caption {}\n"
instance["prompt"] = caption_prompt_format.format(language_code)
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
return response.predictions[0].get("response")
def ocr_predict(
endpoint: Any,
ocr_prompt: str,
image: Any,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Extracts text from a given image using an Endpoint."""
# Resize and convert image to base64 string.
resized_image = resize_image(image, new_width)
resized_image_base64 = image_to_base64(resized_image)
instance = {"image": resized_image_base64}
if ocr_prompt:
instance["prompt"] = ocr_prompt
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
return response.predictions[0].get("response")
def detect_predict(
endpoint: Any,
detect_prompt: str,
image: Any,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Predicts the answer to a question about an image using an Endpoint."""
# Resize and convert image to base64 string.
resized_image = resize_image(image, new_width)
resized_image_base64 = image_to_base64(resized_image)
instance = {"image": resized_image_base64}
if detect_prompt:
instance["prompt"] = detect_prompt
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
return response.predictions[0].get("response")
def copy_model_artifacts(
model_id: str,
model_source: str,
model_destination: str,
) -> None:
"""Copies model artifacts from model_source to model_destination.
model_source and model_destination should be GCS path.
Args:
model_id: The model id.
model_source: The source of the model artifact.
model_destination: The destination of the model artifact.
"""
if not model_source.startswith(GCS_URI_PREFIX):
raise ValueError(
f"{model_source} is not a GCS path starting with {GCS_URI_PREFIX}."
)
if not model_destination.startswith(GCS_URI_PREFIX):
raise ValueError(
f"{model_destination} is not a GCS path starting with {GCS_URI_PREFIX}."
)
model_source = f"{model_source}/{model_id}"
model_destination = f"{model_destination}/{model_id}"
print("Copying model artifact from ", model_source, " to ", model_destination)
subprocess.check_output([
"gcloud",
"storage",
"cp",
"-r",
model_source,
model_destination,
])
def get_quota(project_id: str, region: str, resource_id: str) -> int:
"""Returns the quota for a resource in a region.
Args:
project_id: The project id.
region: The region.
resource_id: The resource id.
Returns:
The quota for the resource in the region. Returns -1 if can not figure out
the quota.
Raises:
RuntimeError: If the command to get quota fails.
"""
service_endpoint = "aiplatform.googleapis.com"
command = (
"gcloud alpha services quota list"
f" --service={service_endpoint} --consumer=projects/{project_id}"
f" --filter='{service_endpoint}/{resource_id}' --format=json"
)
process = subprocess.run(
command, shell=True, capture_output=True, text=True, check=True
)
if process.returncode == 0:
quota_data = json.loads(process.stdout)
else:
raise RuntimeError(f"Error fetching quota data: {process.stderr}")
if not quota_data or "consumerQuotaLimits" not in quota_data[0]:
return -1
if (
not quota_data[0]["consumerQuotaLimits"]
or "quotaBuckets" not in quota_data[0]["consumerQuotaLimits"][0]
):
return -1
all_regions_data = quota_data[0]["consumerQuotaLimits"][0]["quotaBuckets"]
# If the quota data does not have dimensions, it is global quota. However,
# global quota may be overridden by regional quota. So we need to check the
# global quota first.
global_quota = -1
if (
all_regions_data
and "dimensions" not in all_regions_data[0]
and "effectiveLimit" in all_regions_data[0]
):
global_quota = int(all_regions_data[0]["effectiveLimit"])
for region_data in all_regions_data:
if (
region_data.get("dimensions")
and region_data["dimensions"]["region"] == region
):
if "effectiveLimit" in region_data:
return int(region_data["effectiveLimit"])
else:
return 0
return global_quota
def get_resource_id(
accelerator_type: str,
is_for_training: bool,
is_spot: bool = False,
is_restricted_image: bool = False,
is_dynamic_workload_scheduler: bool = False,
) -> str:
"""Returns the resource id for a given accelerator type and the use case.
Args:
accelerator_type: The accelerator type.
is_for_training: Whether the resource is used for training. Set false for
serving use case.
is_spot: Whether the resource is used with Spot.
is_restricted_image: Whether the image is hosted in `vertex-ai-restricted`.
is_dynamic_workload_scheduler: Whether the resource is used with Dynamic
Workload Scheduler.
Returns:
The resource id.
"""
accelerator_suffix_map = {
"NVIDIA_TESLA_V100": "nvidia_v100_gpus",
"NVIDIA_TESLA_P100": "nvidia_p100_gpus",
"NVIDIA_L4": "nvidia_l4_gpus",
"NVIDIA_TESLA_A100": "nvidia_a100_gpus",
"NVIDIA_A100_80GB": "nvidia_a100_80gb_gpus",
"NVIDIA_H100_80GB": "nvidia_h100_gpus",
"NVIDIA_H100_MEGA_80GB": "nvidia_h100_mega_gpus",
"NVIDIA_H200_141GB": "nvidia_h200_gpus",
"NVIDIA_TESLA_T4": "nvidia_t4_gpus",
"TPU_V6e": "tpu_v6e",
"TPU_V5e": "tpu_v5e",
"TPU_V3": "tpu_v3",
}
default_training_accelerator_map = {
key: f"custom_model_training_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
dws_training_accelerator_map = {
key: f"custom_model_training_preemptible_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
restricted_image_training_accelerator_map = {
"NVIDIA_A100_80GB": "restricted_image_training_nvidia_a100_80gb_gpus",
}
spot_serving_accelerator_map = {
key: f"custom_model_serving_preemptible_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
serving_accelerator_map = {
key: f"custom_model_serving_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
if is_for_training:
if is_restricted_image and is_dynamic_workload_scheduler:
raise ValueError(
"Dynamic Workload Scheduler does not work for restricted image"
" training."
)
training_accelerator_map = (
restricted_image_training_accelerator_map
if is_restricted_image
else default_training_accelerator_map
)
if accelerator_type in training_accelerator_map:
if is_dynamic_workload_scheduler:
return dws_training_accelerator_map[accelerator_type]
else:
return training_accelerator_map[accelerator_type]
else:
raise ValueError(
f"Could not find accelerator type: {accelerator_type} for training."
)
else:
if is_dynamic_workload_scheduler:
raise ValueError("Dynamic Workload Scheduler does not work for serving.")
accelerator_map = (
spot_serving_accelerator_map if is_spot else serving_accelerator_map
)
if accelerator_type in accelerator_map:
return accelerator_map[accelerator_type]
else:
raise ValueError(
f"Could not find accelerator type: {accelerator_type} for serving."
)
def check_quota(
project_id: str,
region: str,
accelerator_type: str,
accelerator_count: int,
is_for_training: bool,
is_spot: bool = False,
is_restricted_image: bool = False,
is_dynamic_workload_scheduler: bool = False,
) -> None:
"""Checks if the project and the region has the required quota.
Args:
project_id: The project id.
region: The region.
accelerator_type: The accelerator type.
accelerator_count: The number of accelerators to check quota for.
is_for_training: Whether the resource is used for training. Set false for
serving use case.
is_spot: Whether the resource is used with Spot.
is_restricted_image: Whether the image is hosted in `vertex-ai-restricted`.
is_dynamic_workload_scheduler: Whether the resource is used with Dynamic
Workload Scheduler.
"""
resource_id = get_resource_id(
accelerator_type,
is_for_training=is_for_training,
is_spot=is_spot,
is_restricted_image=is_restricted_image,
is_dynamic_workload_scheduler=is_dynamic_workload_scheduler,
)
quota = get_quota(project_id, region, resource_id)
quota_request_instruction = (
"Either use "
"a different region or request additional quota. Follow "
"instructions here "
"https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota"
" to check quota in a region or request additional quota for "
"your project."
)
if quota == -1:
raise ValueError(
f"Quota not found for: {resource_id} in {region}."
f" {quota_request_instruction}"
)
if quota < accelerator_count:
raise ValueError(
f"Quota not enough for {resource_id} in {region}: {quota} <"
f" {accelerator_count}. {quota_request_instruction}"
)
def get_deploy_source() -> str:
"""Gets deploy_source string based on running environment."""
vertex_product = os.environ.get("VERTEX_PRODUCT", "")
match vertex_product:
case "COLAB_ENTERPRISE":
return "notebook_colab_enterprise"
case "WORKBENCH_INSTANCE":
return "notebook_workbench"
case _:
# Legacy workbench, legacy colab, or other custom environments.
return "notebook_environment_unspecified"
def _is_operation_done(op_name: str, region: str) -> bool:
"""Checks if the operation is done.
Args:
op_name: The name of the operation to poll.
region: The region of the operation.
Returns:
True if the operation is done, False otherwise.
Raises:
ValueError: If the operation failed.
"""
creds, _ = auth.default()
auth_req = auth.transport.requests.Request()
creds.refresh(auth_req)
headers = {
"Authorization": f"Bearer {creds.token}",
}
url = f"https://{region}-aiplatform.googleapis.com/ui/{op_name}"
response = requests.get(url, headers=headers)
operation_data = response.json()
if "error" in operation_data:
raise ValueError(f"Operation failed: {operation_data['error']}")
return operation_data.get("done", False)
def poll_and_wait(
op_name: str, region: str, total_wait: int, interval: int = 60
) -> None:
"""Polls the operation and waits for it to complete.
Args:
op_name: The name of the operation to poll.
region: The region of the operation.
total_wait: The total wait time in seconds.
interval: The interval between each poll in seconds.
Raises:
TimeoutError: If the operation times out.
"""
start_time = time.time()
while True:
if _is_operation_done(op_name, region):
break
time_elapsed = time.time() - start_time
if time_elapsed > total_wait:
raise TimeoutError(
f"Operation timed out after {int(time_elapsed)} seconds."
)
print(
"\rStill waiting for operation... Elapsed time in seconds:"
f" {int(time_elapsed):<6}",
end="",
flush=True,
)
time.sleep(interval)
@@ -1,598 +0,0 @@
"""Functions for dataset validation.
This tool is used to validate the dataset against the given template.
"""
from collections.abc import Callable
import json
import multiprocessing
import os
import subprocess
from typing import Any, Union
from absl import logging
import accelerate
import datasets
import transformers
GCS_URI_PREFIX = "gs://"
GCSFUSE_URI_PREFIX = "/gcs/"
LOCAL_BASE_MODEL_DIR = "/tmp/base_model_dir"
LOCAL_TEMPLATE_DIR = "/tmp/template_dir"
_TEMPLATE_DIRNAME = "templates"
_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME = "vertex-ai-samples"
_VERTEX_AI_SAMPLES_GITHUB_TEMPLATE_DIR = (
"community-content/vertex_model_garden/model_oss/peft/train/vmg/templates"
)
_MODELS_REQUIRING_PAD_TOKEN = ("llama", "falcon", "mistral", "mixtral")
_MODELS_REQUIRING_EOS_TOEKN = ("gemma-2b", "gemma-7b")
_DESCRIPTION_KEY = "description"
_SOURCE_KEY = "source"
_PROMPT_INPUT_KEY = "prompt_input"
_PROMPT_NO_INPUT_KEY = "prompt_no_input"
_RESPONSE_SEPARATOR = "response_separator"
_INSTRUCTION_SEPARATOR = "instruction_separator"
_CHAT_TEMPLATE_KEY = "chat_template"
_KNOWN_KEYS = (
_DESCRIPTION_KEY,
_SOURCE_KEY,
_PROMPT_INPUT_KEY,
_PROMPT_NO_INPUT_KEY,
_RESPONSE_SEPARATOR,
_INSTRUCTION_SEPARATOR,
_CHAT_TEMPLATE_KEY,
)
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 is not None and input_path.startswith(GCS_URI_PREFIX)
def force_gcs_fuse_path(gcs_uri: str) -> str:
"""Converts gs:// uris to their /gcs/ equivalents. No-op for other uris.
Args:
gcs_uri: The GCS URI to convert.
Returns:
The converted GCS URI.
"""
if is_gcs_path(gcs_uri):
return GCSFUSE_URI_PREFIX + gcs_uri[len(GCS_URI_PREFIX) :]
else:
return gcs_uri
def download_gcs_uri_to_local(
gcs_uri: str,
destination_dir: str = LOCAL_BASE_MODEL_DIR,
check_path_exists: bool = True,
) -> str:
"""Downloads GCS URI to local.
If GCS URI is a directory, gs://some/folder is downloaded to
/destination_dir/folder. If GCS URI is a file, gs://some/file is downloaded to
/destination_dir/file.
Args:
gcs_uri: GCS URI to download.
destination_dir: Local directory directory.
check_path_exists: Whether to check if the path exists.
Returns:
Local path to target folder/file.
"""
target = os.path.join(
destination_dir,
os.path.basename(os.path.normpath(gcs_uri)),
)
if check_path_exists and os.path.exists(target):
logging.info("File %s already exists.", target)
return target
if accelerate.PartialState().is_local_main_process:
logging.info(
"Downloading file(s) from %s to %s...", gcs_uri, destination_dir
)
if not os.path.exists(destination_dir):
os.mkdir(destination_dir)
subprocess.check_output([
"gsutil",
"-m",
"cp",
"-r",
gcs_uri,
destination_dir,
])
logging.info("Downloaded file(s) from %s to %s.", gcs_uri, destination_dir)
# Make sure ALL processes process to next step after data downloading is done.
# It matters for the main process to wait for other processes as well.
accelerate.PartialState().wait_for_everyone()
return target
def get_template(template_path: str) -> dict[str, str]:
"""Gets the template dictionary given the file path.
Args:
template_path: Path to the template file.
Returns:
A dictionary of the template.
Raises:
ValueError: If the template file does not exist or contains unknown keys.
"""
if is_gcs_path(template_path):
template_path = force_gcs_fuse_path(template_path)
elif not os.path.isfile(template_path):
template_path = os.path.join(
os.path.dirname(__file__),
_TEMPLATE_DIRNAME,
template_path + ".json",
)
if not os.path.isfile(template_path):
raise ValueError(f"Template file {template_path} does not exist.")
with open(template_path, "r") as f:
template_json: dict[str, str] = json.load(f)
for key in template_json:
if key not in _KNOWN_KEYS:
raise ValueError(f"Unknown key {key} in template {template_path}.")
return template_json
def get_response_separator(template_json: dict[str, str]) -> Union[str, None]:
return template_json.get(_RESPONSE_SEPARATOR, None)
def get_instruction_separator(
template_json: dict[str, str],
) -> Union[str, None]:
return template_json.get(_INSTRUCTION_SEPARATOR, None)
def _format_template_fn(
template: str,
input_column: str,
tokenizer: transformers.PreTrainedTokenizer | None = None,
) -> Callable[[dict[str, str]], dict[str, str]]:
"""Formats a dataset example according to a template.
Args:
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
input_column: The input column in the dataset to be used or updated by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
tokenizer: The tokenizer to use for chat_template templates.
Returns:
A function that formats data according to the template.
"""
template_json = get_template(template)
if _CHAT_TEMPLATE_KEY not in template_json:
def format_fn(example: dict[str, str]) -> dict[str, str]:
format_dict = {key: value for key, value in example.items()}
if format_dict.get(input_column):
format_str = template_json[_PROMPT_INPUT_KEY]
elif _PROMPT_NO_INPUT_KEY in template_json:
format_str = template_json[_PROMPT_NO_INPUT_KEY]
else:
raise KeyError(
f"The template {os.path.basename(template)} does not contain"
f" {_PROMPT_INPUT_KEY} or {_PROMPT_NO_INPUT_KEY} key."
)
try:
return {input_column: format_str.format(**format_dict)}
except KeyError as e:
raise KeyError(
f"The template {os.path.basename(template)} contains a key {e} in"
f" {_PROMPT_INPUT_KEY} or {_PROMPT_NO_INPUT_KEY} that does not"
" exist in the dataset example. The dataset example looks like"
f" {format_dict}."
) from e
return format_fn
elif (
_PROMPT_INPUT_KEY in template_json
or _PROMPT_NO_INPUT_KEY in template_json
):
raise ValueError(
f"chat_template templates do not support {_PROMPT_INPUT_KEY} or"
f" {_PROMPT_NO_INPUT_KEY} templates."
)
else:
if tokenizer is None:
raise ValueError("A tokenizer is required for chat_template templates.")
# Assign HuggingFace jinja template.
tokenizer.chat_template = template_json[_CHAT_TEMPLATE_KEY]
def format_fn(example: dict[str, str]) -> dict[str, str]:
try:
return {
input_column: tokenizer.apply_chat_template(
example[input_column],
tokenize=False,
add_generation_prompt=False,
)
}
except KeyError as e:
raise KeyError(
f"The template {os.path.basename(template)} contains a key {e} in"
f" {_CHAT_TEMPLATE_KEY} that does not exist in the dataset example."
) from e
return format_fn
def _get_split_string(
split: str,
dataset_percent: int | None = None,
dataset_k_rows: int | None = None,
) -> str:
"""Gets the formatted split string for the dataset.
This is used to format the split string as per
https://huggingface.co/docs/datasets/v2.21.0/loading#slice-splits. Also, this
function will only be used to load the partial dataset for validating the
dataset against the template.
Args:
split: Split of the dataset.
dataset_percent: The percentage of the dataset to load.
dataset_k_rows: The top k sequences to load from the dataset.
Returns:
A formatted split string.
"""
# Validate the dataset_percent and dataset_k_rows values.
if dataset_percent and dataset_k_rows:
raise ValueError(
"You can set either validate_percentage_of_dataset or"
" validate_k_rows_of_dataset, but not both."
)
if dataset_percent:
logging.info("Loading %d percent of the dataset...", dataset_percent)
return f"{split}[:{dataset_percent}%]"
if dataset_k_rows:
logging.info("Loading top %d rows of the dataset...", dataset_k_rows)
return f"{split}[:{dataset_k_rows}]"
return split
def _github_template_path(template: str) -> str:
"""Generates the path to the template in the Vertex AI Samples GitHub repo.
Args:
template: Name of the template.
Returns:
The path to the template in the Vertex AI Samples GitHub repo.
"""
# vertex-ai-samples directory may lie under separate directory depending on
# the scratch_dir parameter in the notebook execution environment.
vertex_ai_samples_abs_path = os.getcwd().split(
_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME
)[0]
return os.path.join(
vertex_ai_samples_abs_path,
_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME,
_VERTEX_AI_SAMPLES_GITHUB_TEMPLATE_DIR,
template + ".json",
)
def _get_dataset(
dataset_name: str,
split: str,
num_proc: int | None = None,
) -> datasets.DatasetDict:
"""Gets a dataset.
Args:
dataset_name: Name of the dataset or path to a custom dataset.
split: Split of the dataset.
num_proc: Number of processors to use.
Returns:
A dataset.
"""
dataset_name = force_gcs_fuse_path(dataset_name)
if os.path.isfile(dataset_name):
# Custom dataset.
return datasets.load_dataset(
"json",
data_files=[dataset_name],
split=split,
num_proc=num_proc,
)
# HF dataset.
return datasets.load_dataset(dataset_name, split=split, num_proc=num_proc)
def should_add_pad_token(model_id: str) -> bool:
"""Returns whether the model requires adding a special pad token.
Args:
model_id: The name of the model.
Returns:
True if the model requires adding a special pad token, False otherwise.
"""
model_config = transformers.AutoConfig.from_pretrained(model_id)
if model_config.model_type is None:
return False
return any(
s.lower() in model_config.model_type.lower()
for s in _MODELS_REQUIRING_PAD_TOKEN
)
def should_add_eos_token(model_id: str) -> bool:
"""Returns whether the model requires adding a special eos token.
Args:
model_id: The name of the model.
Returns:
True if the model requires adding a special eos token, False otherwise.
"""
return any(m in model_id for m in _MODELS_REQUIRING_EOS_TOEKN)
def load_tokenizer(
pretrained_model_id: str,
padding_side: str | None = None,
access_token: str | None = None,
) -> transformers.AutoTokenizer:
"""Loads tokenizer based on `pretrained_model_id`.
Args:
pretrained_model_id: The name of the pretrained model.
padding_side: The side to pad the input on.
access_token: The access token to use for the tokenizer.
Returns:
The tokenizer.
"""
tokenizer_kwargs = {}
if should_add_eos_token(pretrained_model_id):
tokenizer_kwargs["add_eos_token"] = True
if padding_side:
tokenizer_kwargs["padding_side"] = padding_side
with accelerate.PartialState().local_main_process_first():
tokenizer = transformers.AutoTokenizer.from_pretrained(
pretrained_model_id,
trust_remote_code=False,
use_fast=True,
token=access_token,
**tokenizer_kwargs,
)
if should_add_pad_token(pretrained_model_id):
tokenizer.add_special_tokens({"pad_token": "[PAD]"})
return tokenizer
def get_filtered_dataset(
dataset: Any,
input_column: str,
max_seq_length: int,
tokenizer: transformers.PreTrainedTokenizer,
example_removed_threshold: float = 50.0,
) -> Any:
"""Returns the dataset by removing examples that are longer than max_seq_length.
Args:
dataset: The dataset to filter.
input_column: The input column in the dataset to be used.
max_seq_length: The maximum sequence length.
tokenizer: The tokenizer.
example_removed_threshold: The percent threshold for the number of examples
removed from the dataset. It should be in the range of [0, 100].
Returns:
The filtered dataset.
Raises:
ValueError: If more than `example_removed_threshold` of the dataset is
filtered out.
"""
actual_dataset_length = len(dataset)
filtered_dataset = dataset.filter(
lambda x: len(tokenizer(x[input_column])["input_ids"]) <= max_seq_length
)
filtered_dataset_length = len(filtered_dataset)
if actual_dataset_length != filtered_dataset_length:
examples_removed_percent = (
(actual_dataset_length - filtered_dataset_length)
* 100
/ actual_dataset_length
)
logging.info(
"(%.2f%%) of examples token length is <= max-seq-length(%d); (%.2f%%) >"
" max-seq-length. Filtering out %d example(s) which are longer than"
" max-seq-length.",
100 - examples_removed_percent,
max_seq_length,
examples_removed_percent,
actual_dataset_length - filtered_dataset_length,
)
if examples_removed_percent > example_removed_threshold:
raise ValueError(
"More than %.2f%% of the dataset is filtered out. This may be due to"
" small value of max-seq-length(%d) or incorrect template. Please"
" increase the max-seq-length or check the template."
% (examples_removed_percent, max_seq_length)
)
print(f"Some formatted examples from the dataset are: {filtered_dataset[:5]}")
return filtered_dataset
def format_dataset(
dataset: datasets.Dataset,
input_column: str,
template: str = None,
tokenizer: transformers.PreTrainedTokenizer | None = None,
) -> datasets.Dataset:
"""Takes a raw dataset and formats it using a template and tokenizer.
Args:
dataset: The raw (unprocessed) dataset to format.
input_column: The input column in the dataset to be used or updaded by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
tokenizer: The tokenizer to use for chat_template templates.
Returns:
A dataset compatible with the template.
"""
return dataset.map(
_format_template_fn(
template,
input_column=input_column,
tokenizer=tokenizer,
)
)
def load_dataset_with_template(
dataset_name: str,
split: str,
input_column: str,
template: str = None,
tokenizer: transformers.PreTrainedTokenizer | None = None,
) -> tuple[Any, Any]:
"""Loads dataset with templates.
Args:
dataset_name: Name of the dataset or path to a custom dataset.
split: Split of the dataset.
input_column: The input column in the dataset to be used or updaded by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
tokenizer: The tokenizer to use for chat_template templates.
Returns:
The raw dataset and the dataset compatible with the template.
"""
raw = _get_dataset(dataset_name, split=split)
if template:
templated = format_dataset(raw, input_column, template, tokenizer)
else:
templated = None
return raw, templated
def validate_dataset_with_template(
dataset_name: str,
split: str,
input_column: str,
template: str,
tokenizer: transformers.PreTrainedTokenizer | None = None,
max_seq_length: int | None = None,
use_multiprocessing: bool = False,
validate_percentage_of_dataset: int | None = None,
validate_k_rows_of_dataset: int | None = None,
example_removed_threshold: float = 50.0,
) -> Any:
"""Validates dataset with templates.
This function will be used to load the dataset and validate it against the
template. In case of validation, we also allow the users to load the dataset
partially by allowing them to read x% or top k rows of the dataset. To
validate the dataset, the template file must be available in the GCS bucket
and the dataset must be available either in the GCS bucket or Hugging Face.
Args:
dataset_name: Name of the dataset or path to a custom dataset.
split: Split of the dataset.
input_column: The input column in the dataset to be used or updaded by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
tokenizer: The tokenizer to use for chat_template templates.
max_seq_length: The maximum sequence length.
use_multiprocessing: If True, it will use multiprocessing to load the
dataset.
validate_percentage_of_dataset: The percentage of the dataset to load.
validate_k_rows_of_dataset: The top k sequences to load from the dataset.
example_removed_threshold: The threshold for the number of examples removed
from the dataset.
Returns:
None if the validation is successful, otherwise returns the error message.
"""
if not template:
raise ValueError("template is required for validate_dataset.")
if not dataset_name:
raise ValueError("dataset_name is empty.")
if not split:
raise ValueError("split is empty.")
split = _get_split_string(
split,
validate_percentage_of_dataset,
validate_k_rows_of_dataset,
)
num_proc = multiprocessing.cpu_count() if use_multiprocessing else 1
# gcsfuse cannot be used from the notebook runtime env. Hence, we have
# to download dataset and template from gcs to local.
if is_gcs_path(dataset_name):
dataset_name = download_gcs_uri_to_local(dataset_name, LOCAL_BASE_MODEL_DIR)
if is_gcs_path(template):
template_path = download_gcs_uri_to_local(template, LOCAL_TEMPLATE_DIR)
elif os.path.isfile(_github_template_path(template)):
template_path = _github_template_path(template)
else:
raise ValueError(
f"Template file {template} does not exist. To validate the"
" dataset, please provide a valid GCS path for the template or a valid"
" template name from"
f" https://github.com/GoogleCloudPlatform/{_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME}/tree/main/{_VERTEX_AI_SAMPLES_GITHUB_TEMPLATE_DIR}."
)
dataset = format_dataset(
_get_dataset(dataset_name, split, num_proc),
input_column,
template_path,
tokenizer,
)
if tokenizer is not None:
get_filtered_dataset(
dataset=dataset,
input_column=input_column,
max_seq_length=max_seq_length,
tokenizer=tokenizer,
example_removed_threshold=example_removed_threshold,
)
print(
"Dataset {} is compatible with the {} template.".format(
os.path.basename(dataset_name), os.path.basename(template)
)
)
@@ -1,275 +0,0 @@
"""Utility functions for interacting with Google Cloud Platform."""
import datetime
import logging
import os
import subprocess
import uuid
from google.cloud import aiplatform
import requests
# Configure logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
def get_project_id() -> str:
"""Read cloud project id from metadata service."""
project_request = requests.get(
"http://metadata.google.internal/computeMetadata/v1/project/project-id",
headers={"Metadata-Flavor": "Google"},
)
return project_request.text
def get_region() -> str:
"""Read region from metadata service."""
region_request = requests.get(
"http://metadata.google.internal/computeMetadata/v1/instance/region",
headers={"Metadata-Flavor": "Google"},
)
return region_request.text.split("/")[-1]
# Get the default cloud project id and region
PROJECT_ID = get_project_id()
REGION = get_region()
def init_aiplatform(project: str = None, location: str = None) -> None:
"""Initialize the Vertex AI SDK.
Args:
project: The Google Cloud project ID.
location: The Google Cloud location.
"""
project = PROJECT_ID if project is None else project
location = REGION if location is None else location
aiplatform.init(project=project, location=location)
subprocess.call([
"gcloud",
"services",
"enable",
"aiplatform.googleapis.com",
"compute.googleapis.com",
])
def run_command(command: list[str]) -> str:
"""Runs a shell command and returns the output.
Args:
command: The shell command to run as a list.
Returns:
The output of the command.
"""
try:
result = subprocess.run(
command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
)
return result.stdout
except subprocess.CalledProcessError as e:
logger.error("Error: %s", e.stderr)
raise e
def enable_apis() -> None:
"""Enable the Vertex AI API and Compute Engine API."""
logger.info("Enabling Vertex AI API and Compute Engine API.")
run_command([
"gcloud",
"services",
"enable",
"aiplatform.googleapis.com",
"compute.googleapis.com",
])
def setup_buckets(bucket_uri: str, model_bucket_name: str) -> tuple[str, str]:
"""Set up Cloud Storage buckets for storing experiment artifacts.
Args:
bucket_uri: The bucket URI provided by the user.
model_bucket_name: The name of the model bucket.
Returns:
A tuple containing the bucket name and model bucket path.
"""
if not bucket_uri.strip():
# Generate a default bucket URI if none provided
now = datetime.datetime.now().strftime("%Y%m%d%H%M%S")
bucket_uri = f"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}"
logger.info("No bucket URI provided. Using default bucket: %s", bucket_uri)
else:
if not bucket_uri.startswith("gs://"):
raise ValueError("Bucket URI must start with 'gs://'.")
# Remove any trailing slashes
bucket_uri = bucket_uri.rstrip("/")
bucket_name = "/".join(bucket_uri.split("/")[:3])
# Check if bucket exists
try:
run_command(["gsutil", "ls", "-b", bucket_uri])
logger.info("Bucket %s already exists.", bucket_uri)
except subprocess.CalledProcessError:
logger.info("Creating bucket %s.", bucket_uri)
# Create the bucket in the same region as the project
run_command(["gsutil", "mb", "-l", REGION, bucket_uri])
# Construct the model bucket path
model_bucket = os.path.join(bucket_uri, model_bucket_name)
# Check if the model bucket exists (as a folder within the main bucket)
try:
run_command(["gsutil", "ls", model_bucket])
logger.info("Model bucket %s already exists.", model_bucket)
except subprocess.CalledProcessError:
logger.info("Creating model bucket %s.", model_bucket)
# Create the model bucket folder
run_command(["gsutil", "cp", "/dev/null", model_bucket + "/"])
return bucket_name, model_bucket
def get_service_account() -> str:
"""Get the default service account."""
shell_output = run_command(["gcloud", "projects", "describe", PROJECT_ID])
project_number_line = next(
(line for line in shell_output.splitlines() if "projectNumber" in line),
None,
)
if project_number_line:
project_number = project_number_line.split(":")[1].strip().replace("'", "")
service_account = f"{project_number}-compute@developer.gserviceaccount.com"
logger.info("Using default Service Account: %s", service_account)
return service_account
else:
raise ValueError("Could not find project number in gcloud output.")
def get_project_number() -> str:
"""Get the default project number."""
shell_output = run_command(["gcloud", "projects", "describe", PROJECT_ID])
project_number_line = next(
(line for line in shell_output.splitlines() if "projectNumber" in line),
None,
)
if project_number_line:
project_number = project_number_line.split(":")[1].strip().replace("'", "")
logger.info("Using default Project Number: %s", project_number)
return project_number
else:
raise ValueError("Could not find project number in gcloud output.")
def provision_permissions(service_account: str, bucket_name: str) -> None:
"""Provision permissions to the service account with the GCS bucket."""
if bucket_name:
run_command([
"gsutil",
"iam",
"ch",
f"serviceAccount:{service_account}:roles/storage.admin",
bucket_name,
])
def set_gcloud_project() -> None:
"""Set gcloud config project."""
run_command(["gcloud", "config", "set", "project", PROJECT_ID])
def initialize(
bucket_uri: str, model_bucket_name: str, create_bucket: bool
) -> tuple[str, str]:
"""Initialize the environment.
Args:
bucket_uri: The bucket URI provided by the user.
model_bucket_name: The name of the model bucket.
create_bucket: Whether to create the bucket or not.
Returns:
A tuple containing the model bucket path and service account.
"""
enable_apis()
bucket_name = None
if create_bucket:
bucket_name, model_bucket = setup_buckets(bucket_uri, model_bucket_name)
else:
model_bucket = None
service_account = get_service_account()
provision_permissions(service_account, bucket_name)
set_gcloud_project()
return model_bucket, service_account
def clean_resources_ui(
project_id: str,
region: str,
endpoint_name: str,
delete_bucket: bool,
bucket_name: str = None,
) -> str:
"""UI function for cleaning a specific Vertex AI endpoint and its model."""
if delete_bucket and not bucket_name:
raise ValueError("Bucket name is required when 'Delete Bucket' is checked.")
try:
delete_endpoint_and_model(project_id, region, endpoint_name)
bucket_status_message = ""
if delete_bucket:
bucket_status_message = delete_gcs_bucket(bucket_name)
if endpoint_name:
return (
f"Endpoint {endpoint_name} and associated model deleted successfully!"
f" {bucket_status_message}"
)
else:
return (
"There are currently no endpoints available for deletion."
f" {bucket_status_message}"
)
except Exception as e: # pylint: disable=broad-exception-caught
return f"Error cleaning up resources: {e}"
def delete_endpoint_and_model(
project_id: str, region: str, endpoint_name: str
) -> None:
"""Deletes a specific Vertex AI endpoint and its associated model."""
if endpoint_name:
endpoint_id = endpoint_name.split(" - ")[0]
endpoint_resource_name = (
f"projects/{project_id}/locations/{region}/endpoints/{endpoint_id}"
)
endpoint = aiplatform.Endpoint(
endpoint_resource_name, project=project_id, location=region
)
deployed_models = endpoint.list_models()
for deployed_model in deployed_models:
endpoint.undeploy(deployed_model_id=deployed_model.id)
model = aiplatform.Model(deployed_model.model)
model.delete()
endpoint.delete()
def delete_gcs_bucket(bucket_name: str) -> str:
"""Deletes a GCS bucket using gsutil."""
try:
run_command(["gsutil", "-m", "rm", "-r", bucket_name])
logger.info("Bucket %s deleted using gsutil.", bucket_name)
return f"Bucket {bucket_name} deleted successfully!"
except subprocess.CalledProcessError as e:
logger.error(
"Error deleting bucket %s using gsutil: %s", bucket_name, str(e)
)
return f"Bucket {bucket_name} could not be found or deleted. "
@@ -1,58 +0,0 @@
FROM nvidia/cuda:12.3.2-devel-ubuntu22.04
# Install basic libs
RUN apt-get update && apt-get upgrade -y && apt-get install -y --no-install-recommends \
cmake \
curl \
wget \
sudo \
gnupg \
libsm6 \
libxext6 \
libxrender-dev \
lsb-release \
ca-certificates \
build-essential \
git \
software-properties-common \
cuda-toolkit \
libcudnn8 \
apt-transport-https
RUN apt install -y --no-install-recommends python3.10 \
python3.10-venv \
python3.10-dev \
python3-pip
Run apt-get autoremove -y
RUN pip install --upgrade pip
RUN pip install --upgrade --ignore-installed \
"jax[cuda12]==0.4.26" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html \
numpy==1.26.4 \
paxml==1.4.0 \
praxis==1.4.0 \
jaxlib==0.4.26 \
pandas==2.1.4 \
einshape==1.0.0 \
utilsforecast==0.1.10 \
huggingface_hub[cli]==0.23.0 \
google-cloud-aiplatform[prediction]==1.51.0 \
fastapi==0.109.1 \
flask==3.0.3 \
smart_open[gcs]==7.0.4 \
protobuf==3.19.6 \
scikit-learn==1.0.2 \
timesfm==1.0.1
# Download license.
RUN wget https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/LICENSE
# Move scaffold.
COPY model_oss/timesfm/main.py /app/main.py
COPY model_oss/timesfm/predictor.py /app/predictor.py
WORKDIR ..
# Spin off inference server.
CMD ["python3", "/app/main.py"]
@@ -1,71 +0,0 @@
"""Predict server for TimesFM."""
import json
import os
import flask
import predictor
from predictor import PredictionError
# Create the flask app.
app = flask.Flask(__name__)
_OK_STATUS = 200
_INTERNAL_ERROR_STATUS = 500
_BAD_REQUEST_STATUS = 400
_HOST = '0.0.0.0'
# Define the predictor and load the checkpoints.
predictor = predictor.TimesFMPredictor()
predictor.load(os.environ['AIP_STORAGE_URI'])
@app.route(os.environ['AIP_HEALTH_ROUTE'], methods=['GET'])
def health() -> flask.Response:
return flask.Response(status=_OK_STATUS)
@app.route(os.environ['AIP_PREDICT_ROUTE'], methods=['GET', 'POST'])
def predict() -> flask.Response:
"""Calls TimesFM for prediction.
Returns:
A `flask.Response` containing the prediction result in JSON.
"""
try:
body = flask.request.get_json(silent=True, force=True)
preprocessed_inputs = predictor.preprocess(body)
outputs = predictor.predict(preprocessed_inputs)
conf_level = preprocessed_inputs.get('conf_level')
if conf_level is not None:
postprocessed_outputs = predictor.postprocess_with_conf_level(
outputs, preprocessed_inputs['conf_level']
)
else:
postprocessed_outputs = predictor.postprocess(outputs)
return flask.Response(
json.dumps(postprocessed_outputs),
status=_OK_STATUS,
mimetype='application/json',
)
except PredictionError as e:
return flask.Response(
json.dumps({'error': str(e)}),
status=e.status_code,
mimetype='application/json',
)
except ValueError as e:
return flask.Response(
json.dumps({'error': str(e)}),
status=_BAD_REQUEST_STATUS,
mimetype='application/json',
)
except Exception as e: # pylint: disable=broad-exception-caught
return flask.Response(
json.dumps({'error': str(e)}),
status=_INTERNAL_ERROR_STATUS,
mimetype='application/json',
)
if __name__ == '__main__':
app.run(host=_HOST, port=os.environ['AIP_HTTP_PORT'])
@@ -1,609 +0,0 @@
"""Adapts a pretrained TimesFM to the CPR framework.
Documentation for the model is here:
https://github.com/google-research/timesfm
Model checkpoints can be found here:
https://www.huggingface.co/google/timesfm-1.0-200m
"""
from collections.abc import Sequence
import datetime
import os
from typing import Any
import fastapi
from google.cloud.aiplatform.utils import prediction_utils
from jax._src import config
import numpy as np
import scipy.stats as st
import timesfm
HTTPException = fastapi.HTTPException
_BACKEND = os.getenv("TIMESFM_BACKEND", default="cpu")
config.update(
"jax_platforms", {"cpu": "cpu", "gpu": "cuda", "tpu": ""}[_BACKEND]
)
TsArray = None | float | int | str | list["TsArray"]
_BAD_REQUEST_STATUS = 400
_EXPECTED_FORMAT = """
[NOTICE] TimesFM inference server expects input format:
{
"instances": [
{
"input": [0.0, 0.1, 0.2, ...],
"freq": 0, # optional, 0/1/2
"horizon": 12, # optional
"timestamp": ["2024-01-01", "2024-01-02", ...], # optional
"timestamp_format": "%Y-%m-%d", # optional
"dynamic_numerical_covariates": {
"dncov1": [1.0, 2.0, 1.5, ...],
"dncov2": [3.0, 1.1, 2.4, ...],
}, # optional
"dynamic_categorical_covariates": {
"dccov1": ["a", "b", "a", ...],
"dccov2": [0, 1, 0, ...],
}, # optional
"static_numerical_covariates": {
"sncov1": 1.0,
"sncov2": 2.0,
}, # optional
"static_categorical_covariates": {
"sccov1": "a",
"sccov2": "b",
}, # optional
"xreg_kwargs": {...}, # optional
},
{"input": [113.2, 15.0, 65.4, ...], ...},
{"input": [ 0.0, 10.0, 20.0, ...], ...},
...
]
}
"""
class PredictionError(Exception):
"""Custom exception for prediction errors."""
def __init__(self, message: str, status_code: int = _BAD_REQUEST_STATUS):
super().__init__(message)
self.status_code = status_code
self.message = message
def _raise_bad_request(message: str):
message = message + "\n" + _EXPECTED_FORMAT
raise PredictionError(
message=message,
status_code=_BAD_REQUEST_STATUS,
)
def _datetime_to_freq(dt1: datetime.datetime, dt2: datetime.datetime) -> int:
delta = dt2 - dt1
if delta.days <= 1:
return 0
elif delta.days <= 31:
return 1
else:
return 2
def _add_cov_to_dict(
index: int,
cov_input: dict[str, TsArray],
cov_dict: dict[str, list[TsArray]],
):
"""Adds covariates to the dictionary of covariates.
Args:
index: Index of the instance.
cov_input: Dictionary of covariates for the current instance.
cov_dict: Dictionary of covariates for all instances.
"""
if index == 0:
cov_dict.update({k: [v] for k, v in cov_input.items()})
else:
if set(cov_input.keys()) != set(cov_dict.keys()):
_raise_bad_request(
f"Instance {index}:"
" All instances must have the same set of covariates if any."
)
for k, v in cov_input.items():
cov_dict[k].append(v)
def _linear_interpolate_missing_timepoints(
timestamp: list[datetime.datetime],
value: list[float],
) -> tuple[list[datetime.datetime], list[TsArray]]:
"""Linearly interpolates missing timepoints in a timeseries."""
def _gcd_timelapse(t1, t2):
if (w := t2 % t1) == datetime.timedelta(0):
return t1
if t1 > t2:
return _gcd_timelapse(t2, t1)
return _gcd_timelapse(w, t1)
if len(timestamp) < 3:
return timestamp, value, False
no_missing = True
delta = timestamp[1] - timestamp[0]
if delta <= datetime.timedelta(0):
_raise_bad_request(
f"Timestamps must be in ascending order. Got {timestamp}"
)
for i in range(2, len(timestamp)):
delta_next = timestamp[i] - timestamp[i - 1]
if delta_next <= datetime.timedelta(0):
_raise_bad_request(
f"Timestamps must be in ascending order. Got {timestamp}"
)
delta_new = _gcd_timelapse(delta, delta_next)
if delta_new != delta:
no_missing = False
delta = delta_new
if no_missing:
return timestamp, value, False
new_timestamp = []
new_value = []
for i in range(len(timestamp) - 1):
new_timestamp.append(timestamp[i])
new_value.append(value[i])
if (num_deltas := int((timestamp[i + 1] - timestamp[i]) / delta + 0.5)) > 1:
value_delta = (value[i + 1] - value[i]) / num_deltas
for j in range(1, num_deltas):
new_timestamp.append(timestamp[i] + j * delta)
new_value.append(value[i] + j * value_delta)
new_timestamp.append(timestamp[-1])
new_value.append(value[-1])
return new_timestamp, new_value, True
class TimesFMPredictor:
"""Predictor class for time-series foundation model TimesFM."""
TIMESFM_MODEL_NAME = os.getenv(
"TIMESFM_MODEL_NAME", default="timesfm-1.0-200m"
)
CONTEXT_LEN = 512
INPUT_PATCH_LEN = 32
OUTPUT_PATCH_LEN = 128
NUM_LAYERS = 20
MODEL_DIMS = 1280
BACKEND = os.getenv("TIMESFM_BACKEND", default="cpu")
MAX_HORIZON = int(os.getenv("TIMESFM_HORIZON", default="128"))
def load(self, artifacts_uri: str = ""):
"""Initializes the model and preprocessing transforms.
Args:
artifacts_uri: Directory where state dict is stored. Can be a GCS URI or
local path.
"""
if not (os.path.isdir(artifacts_uri) or artifacts_uri.startswith("gs://")):
raise ValueError(
f"Provided artifact_uri is not a directory: {artifacts_uri}"
)
print(f"Downloading checkpoints from {artifacts_uri}")
prediction_utils.download_model_artifacts(artifacts_uri)
artifact_path = os.getcwd()
print(f"Loading checkpoints from {artifact_path}")
self._model = timesfm.TimesFm(
context_len=self.CONTEXT_LEN,
horizon_len=(
((self.MAX_HORIZON - 1) // self.OUTPUT_PATCH_LEN + 1)
* self.OUTPUT_PATCH_LEN
),
input_patch_len=self.INPUT_PATCH_LEN,
output_patch_len=self.OUTPUT_PATCH_LEN,
num_layers=self.NUM_LAYERS,
model_dims=self.MODEL_DIMS,
backend=self.BACKEND,
)
self._model.load_from_checkpoint(artifact_path)
print(f"Loaded TimesFM model from {artifact_path}")
def preprocess(
self, request_dict: dict[str, Sequence[dict[str, TsArray]]]
) -> dict[str, TsArray]:
"""Performs preprocessing.
By default, the server expects a request body consisting of a valid JSON
object. This will be parsed by the handler before it's evaluated by the
preprocess method.
Args:
request_dict: Parsed request body. We expect that the input consists of a
list of time-series forecast contexts. Each context should be in a
format convertible to JTensor by `jnp.array`.
Returns:
Time-series forecast contexts are passed as is from the input as a list.
"""
if "instances" not in request_dict:
_raise_bad_request('Request must contain "instances" as a top-level key.')
input_instances = request_dict["instances"]
if not input_instances or not isinstance(input_instances, list):
_raise_bad_request(
f"Received `instances` not a list. Got {type(input_instances)}"
)
inputs, freqs, timestamps, timestamp_formats = [], [], [], []
horizon_lens = []
conf_level = None
static_numerical_covariates, static_categorical_covariates = {}, {}
dynamic_numerical_covariates, dynamic_categorical_covariates = {}, {}
xreg_kwargs = {}
exists_missing = False
for index, each_input in enumerate(input_instances):
# 1. Add input time-series context.
if (
(not isinstance(each_input, dict))
or ("input" not in each_input)
or (len(each_input["input"]) < 2)
):
_raise_bad_request(
f"Instance {index}:"
" Invalid datatype. Each input example must have `input` key"
" mapped to a list of time-series forecast context with length > 1."
)
new_input = each_input["input"]
# 2. Process timestamps.
if "timestamp" not in each_input:
timestamps.append(None)
else:
if len(each_input["timestamp"]) != len(each_input["input"]):
_raise_bad_request(
f"Instance {index}:"
" Invalid datatype. `timestamp` if given must have same length as"
"`input`."
)
new_timestamp = [
datetime.datetime.fromisoformat(s) for s in each_input["timestamp"]
]
# Linearly interpolate missing timepoints and values.
new_timestamp, new_input, new_exists_missing = (
_linear_interpolate_missing_timepoints(new_timestamp, new_input)
)
exists_missing = exists_missing or new_exists_missing
timestamps.append(new_timestamp)
if "timestamp_format" in each_input:
timestamp_formats.append(each_input["timestamp_format"])
else:
timestamp_formats.append(None)
inputs.append(new_input)
# 3. Process frequency.
if "freq" in each_input:
freqs.append(each_input["freq"])
elif timestamps[index]:
freqs.append(
_datetime_to_freq(timestamps[index][0], timestamps[index][1])
)
else:
freqs.append(0)
# 4. Process covariate data.
for cov_category, cov_dict in [
("static_numerical_covariates", static_numerical_covariates),
("static_categorical_covariates", static_categorical_covariates),
("dynamic_numerical_covariates", dynamic_numerical_covariates),
("dynamic_categorical_covariates", dynamic_categorical_covariates),
]:
if cov_category in each_input:
_add_cov_to_dict(index, each_input[cov_category], cov_dict)
# 5. Process xreg config. Power user option. If nothing set we apply
# TimesFM default.
if "xreg_kwargs" in each_input:
if not xreg_kwargs:
xreg_kwargs = each_input["xreg_kwargs"]
elif xreg_kwargs != each_input["xreg_kwargs"]:
_raise_bad_request(
f"Instance {index}:"
" All instances must have the same xreg_kwargs if any."
)
# 6. Process horizon length.
if "horizon" in each_input:
if (w := each_input["horizon"]) > self.MAX_HORIZON:
_raise_bad_request(
f"Instance {index}: `horizon` must be <= maximum horizon"
f" {self.MAX_HORIZON}. Got {w}. To increase the maximum horizon,"
" recreate the endpoint with a higher `TIMESFM_HORIZON` env"
" value."
)
horizon_lens.append(w)
else:
horizon_lens.append(self.MAX_HORIZON)
# 7. Process conf level.
all_conf_levels = [
each_input.get("conf_level", None) for each_input in input_instances
]
defined_conf_levels = [cl for cl in all_conf_levels if cl is not None]
undefined_conf_levels = [cl for cl in all_conf_levels if cl is None]
if defined_conf_levels and undefined_conf_levels:
_raise_bad_request(
"Either all or none of the instances must define `conf_level`."
)
if defined_conf_levels:
unique_conf_levels = set(defined_conf_levels)
if len(unique_conf_levels) > 1:
_raise_bad_request("All instances must have the same `conf_level`.")
conf_level = unique_conf_levels.pop()
if not 0 <= conf_level <= 1:
_raise_bad_request(
f"`conf_level` must be between 0 and 1. Got {conf_level}."
)
else:
conf_level = None
return {
"inputs": inputs,
"freqs": freqs,
"timestamps": timestamps,
"timestamp_formats": timestamp_formats,
"exists_missing": exists_missing,
"static_numerical_covariates": static_numerical_covariates,
"static_categorical_covariates": static_categorical_covariates,
"dynamic_numerical_covariates": dynamic_numerical_covariates,
"dynamic_categorical_covariates": dynamic_categorical_covariates,
"xreg_kwargs": xreg_kwargs,
"horizon_lens": horizon_lens,
"conf_level": conf_level,
}
def predict(self, instances: dict[str, Any]) -> Any:
"""Performs prediction.
Args:
instances: A dictionary with two keys - `inputs` and `freq` where `inputs`
is list of time series forecast contexts. Each context time series
should be in a format convertible to JTensor by `jnp.array`. `freq` is
frequencies of each forecast context with values as 0 (high), 1 (medium)
and 2 (low). If not provided, all contexts are assumed to be high
frequency.
Returns:
A tuple of List:
- the mean forecast of size (# inputs, # forecast horizon),
- the full forecast (mean + quantiles) of size
(# inputs, # forecast horizon, 1 + # quantiles).
"""
(
inputs,
freqs,
timestamps,
timestamp_formats,
exists_missing,
static_numerical_covariates,
static_categorical_covariates,
dynamic_numerical_covariates,
dynamic_categorical_covariates,
xreg_kwargs,
horizon_lens,
) = (
instances["inputs"],
instances["freqs"],
instances["timestamps"],
instances["timestamp_formats"],
instances["exists_missing"],
instances["static_numerical_covariates"],
instances["static_categorical_covariates"],
instances["dynamic_numerical_covariates"],
instances["dynamic_categorical_covariates"],
instances["xreg_kwargs"],
instances["horizon_lens"],
)
if (
static_numerical_covariates
or static_categorical_covariates
or dynamic_numerical_covariates
or dynamic_categorical_covariates
):
if (
dynamic_categorical_covariates or dynamic_numerical_covariates
) and exists_missing:
_raise_bad_request(
"Dynamic covariates are not supported when input has missing"
" timestamps."
)
print("Detected covariates. Callng model.forecast_with_covariates.")
try:
point_forecast, _ = self._model.forecast_with_covariates(
inputs=inputs,
dynamic_numerical_covariates=dynamic_numerical_covariates,
dynamic_categorical_covariates=dynamic_categorical_covariates,
static_numerical_covariates=static_numerical_covariates,
static_categorical_covariates=static_categorical_covariates,
freq=freqs,
**xreg_kwargs,
)
# point_forecast is a list of np.ndarrays.
point_forecast = [p.tolist() for p in point_forecast]
quantile_forecast = None
except ValueError as e:
_raise_bad_request(f"model.forecast_with_covariates failed from {e}.")
return
else:
print("Calling model.forecast.")
point_forecast, quantile_forecast = self._model.forecast(
inputs=inputs, freq=freqs
)
# point_forecast and quantile_forecast are JTensors (np.ndarrays).
point_forecast = point_forecast.tolist()
quantile_forecast = quantile_forecast.tolist()
return (
point_forecast,
quantile_forecast,
timestamps,
timestamp_formats,
horizon_lens,
)
def postprocess(
self, forecasts: tuple[TsArray, TsArray, TsArray, TsArray]
) -> dict[str, list[dict[str, TsArray]]]:
"""Translates the model output.
Args:
forecasts: A tuple of List - the mean forecast of size (# inputs, #
forecast horizon), - the full forecast (mean + quantiles) of size (#
inputs, # forecast horizon, 1 + # quantiles).
Returns:
Dictionary containing the list of point forecasts and quantile forecasts
for each of the input time-series context.
"""
(
point_forecasts,
quantile_forecasts,
timestamps,
timestamp_formats,
horizon_lens,
) = forecasts
predictions = []
quantile_names = ["mean"] + [
f"p{int(quantile * 100)}" for quantile in self._model.model_p.quantiles
]
for i, point_forecast in enumerate(point_forecasts):
response = {"point_forecast": point_forecast[: horizon_lens[i]]}
if quantile_forecasts:
for j, quantile_name in enumerate(quantile_names):
response[quantile_name] = [x[j] for x in quantile_forecasts[i]][
: horizon_lens[i]
]
if timestamps[i]:
last_timestamp = timestamps[i][-1]
timestamp_delta = timestamps[i][-1] - timestamps[i][-2]
response["timestamp"] = []
for _ in range(len(point_forecast)):
last_timestamp = last_timestamp + timestamp_delta
response["timestamp"].append(
datetime.datetime.strftime(last_timestamp, timestamp_formats[i])
if timestamp_formats[i]
else last_timestamp.isoformat()
)
response["timestamp"] = response["timestamp"][: horizon_lens[i]]
predictions.append(response)
return {"predictions": predictions}
def postprocess_with_conf_level(
self,
forecasts: tuple[TsArray, TsArray, TsArray, TsArray, TsArray],
conf_level: float | None,
) -> dict[str, list[dict[str, TsArray]]]:
"""Translates the model output."""
lower_quantile = (1 - conf_level) / 2
higher_quantile = (1 + conf_level) / 2
_, quantile_forecast, _, _, horizon_lens = forecasts
response = self.postprocess(forecasts)
if quantile_forecast is None:
return response
# Note: The raw quantile forecast from TimesFM has the mean as the 0-th
# element. We strip it before passing to extend_quantiles.
quantile_forecast_np_array = np.array(quantile_forecast)
extended_forecasts = extend_quantiles(
quantile_forecast_np_array[..., 1:],
lower_quantile,
higher_quantile,
model_quantiles=self._model.model_p.quantiles,
)
lower_bounds = extended_forecasts["lower_bound"]
upper_bounds = extended_forecasts["upper_bound"]
for i, prediction in enumerate(response["predictions"]):
horizon = horizon_lens[i]
prediction["lower_bound"] = lower_bounds[i][:horizon].tolist()
prediction["upper_bound"] = upper_bounds[i][:horizon].tolist()
return response
def extend_quantiles(
quantile_forecast: np.ndarray,
lower_quantile: float,
higher_quantile: float,
model_quantiles: list[float],
) -> dict[str, np.ndarray]:
"""Extends the quantile forecast to the lower and upper bounds.
Args:
quantile_forecast: The quantile forecast from TimesFM.
lower_quantile: The lower quantile to extend to.
higher_quantile: The higher quantile to extend to.
model_quantiles: The quantiles used by the model.
Returns:
A dictionary containing the lower and upper bounds.
"""
if quantile_forecast.shape[2] != len(model_quantiles):
raise ValueError(
"Number of model quantiles should match the last dimension of the"
" quantile forecast. If you are using the raw TimesFM quantile forecast"
"output, you likely need to strip the 0-index which is the mean."
)
idx_median = model_quantiles.index(0.5)
idx_low_q = np.argmin(model_quantiles)
low_q = model_quantiles[idx_low_q]
if not (low_q < 0.5):
raise ValueError(
f"The lowest quantile {low_q=} provided in the forecast must be less"
" than 0.5."
)
idx_high_q = np.argmax(model_quantiles)
high_q = model_quantiles[idx_high_q]
if not (high_q > 0.5):
raise ValueError(
f"The highest quantile {high_q=} provided in the forecast must be"
" greater than 0.5."
)
positive_sigma = np.maximum(
0, quantile_forecast[..., idx_high_q] - quantile_forecast[..., idx_median]
) / st.norm.ppf(high_q)
negative_sigma = np.minimum(
0, quantile_forecast[..., idx_low_q] - quantile_forecast[..., idx_median]
) / st.norm.ppf(low_q)
lower_bound = quantile_forecast[
..., idx_median
] + negative_sigma * st.norm.ppf(lower_quantile)
upper_bound = quantile_forecast[
..., idx_median
] + positive_sigma * st.norm.ppf(higher_quantile)
return {"lower_bound": lower_bound, "upper_bound": upper_bound}
@@ -1,746 +0,0 @@
"""Common util functions for notebook."""
import base64
from collections.abc import Sequence
import datetime
import io
import json
import os
import subprocess
import time
from typing import Any
from google import auth
from google.cloud import storage
import matplotlib.pyplot as plt
import numpy as np
from PIL import Image
import requests
import tensorflow as tf
import yaml
GCS_URI_PREFIX = "gs://"
CHECKPOINT_BUCKET = "gs://model_garden_checkpoints"
def convert_numpy_array_to_byte_string_via_tf_tensor(
np_array: np.ndarray,
) -> str:
"""Serializes a numpy array to tensor bytes.
Args:
np_array: A numpy array.
Returns:
A tensor bytes.
"""
tensor_array = tf.convert_to_tensor(np_array)
tensor_byte_string = tf.io.serialize_tensor(tensor_array)
return tensor_byte_string.numpy()
def get_jpeg_bytes(local_image_path: str, new_width: int = -1) -> bytes:
"""Returns jpeg bytes given an image path and resizes if required.
Args:
local_image_path: A string of local image path.
new_width: An integer of new image width.
Returns:
A jpeg bytes.
"""
image = Image.open(local_image_path)
if new_width <= 0:
new_image = image
else:
width, height = image.size
print("original input image size: ", width, " , ", height)
new_height = int(height * new_width / width)
print("new input image size: ", new_width, " , ", new_height)
new_image = image.resize((new_width, new_height))
buffered = io.BytesIO()
new_image.save(buffered, format="JPEG")
return buffered.getvalue()
def gcs_fuse_path(path: str) -> str:
"""Try to convert path to gcsfuse path if it starts with gs:// else do not modify it.
Args:
path: A string of path.
Returns:
A gcsfuse path.
"""
path = path.strip()
if path.startswith("gs://"):
return "/gcs/" + path[5:]
return path
def get_job_name_with_datetime(prefix: str) -> str:
"""Gets a job name by adding current time to prefix.
Args:
prefix: A string of job name prefix.
Returns:
A job name.
"""
now = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
job_name = f"{prefix}-{now}".replace("_", "-")
return job_name
def create_job_name(prefix: str) -> str:
"""Creates a job name.
Args:
prefix: A string of job name prefix.
Returns:
A job name.
"""
user = os.environ.get("USER")
now = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
job_name = f"{prefix}-{user}-{now}".replace("_", "-")
return job_name
def save_subset_annotation(
input_annotation_path: str, output_annotation_path: str
):
"""Saves a subset of COCO annotation json file with CCA 4.0 license.
Args:
input_annotation_path: A string of input annotation path.
output_annotation_path: A string of output annotation path.
"""
with open(input_annotation_path) as f:
coco_json = json.load(f)
img_ids = set()
images = []
annotations = []
for img in coco_json["images"]:
if img["license"] in [4, 5]: # CCA 4.0 license.
img_ids.add(img["id"])
images.append(img)
for ann in coco_json["annotations"]:
if ann["image_id"] in img_ids:
annotations.append(ann)
new_json = {
"info": coco_json["info"],
"licenses": coco_json["licenses"],
"images": images,
"annotations": annotations,
"categories": coco_json["categories"],
}
with open(output_annotation_path, "w") as f:
json.dump(new_json, f)
def image_to_base64(image: Any, image_format: str = "JPEG") -> str:
"""Converts an image to base64.
Args:
image: A PIL.Image instance.
image_format: A string of image format.
Returns:
A base64 string.
"""
buffer = io.BytesIO()
image.save(buffer, format=image_format)
image_str = base64.b64encode(buffer.getvalue()).decode("utf-8")
return image_str
def base64_to_image(image_str: str) -> Any:
"""Convert base64 encoded string to an image.
Args:
image_str: A string of base64 encoded image.
Returns:
A PIL.Image instance.
"""
image = Image.open(io.BytesIO(base64.b64decode(image_str)))
return image
def image_grid(imgs: Sequence[Any], rows: int = 2, cols: int = 2) -> Any:
"""Creates an image grid.
Args:
imgs: A list of PIL.Image instances.
rows: An integer of number of rows.
cols: An integer of number of columns.
Returns:
A PIL.Image instance.
"""
w, h = imgs[0].size
grid = Image.new(
mode="RGB", size=(cols * w + 10 * cols, rows * h), color=(255, 255, 255)
)
for i, img in enumerate(imgs):
grid.paste(img, box=(i % cols * w + 10 * i, i // cols * h))
return grid
def display_image(image: Any):
"""Displays an image.
Args:
image: A PIL.Image instance.
"""
_ = plt.figure(figsize=(20, 15))
plt.grid(False)
plt.imshow(image)
def download_gcs_file_to_local(gcs_uri: str, local_path: str):
"""Download a gcs file to a local path.
Args:
gcs_uri: A string of file path on GCS.
local_path: A string of local file path.
"""
if not gcs_uri.startswith(GCS_URI_PREFIX):
raise ValueError(
f"{gcs_uri} is not a GCS path starting with {GCS_URI_PREFIX}."
)
client = storage.Client()
os.makedirs(os.path.dirname(local_path), exist_ok=True)
with open(local_path, "wb") as f:
client.download_blob_to_file(gcs_uri, f)
def download_image(url: str) -> str:
"""Downloads an image from the given URL.
Args:
url: The URL of the image to download.
Returns:
base64 encoded image.
"""
response = requests.get(url)
return Image.open(io.BytesIO(response.content)) # pytype: disable=bad-return-type # pillow-102-upgrade
def resize_image(image: Any, new_width: int = 1000) -> Any:
"""Resizes an image to a certain width.
Args:
image: The image which has to be resized.
new_width: New width of the image.
Returns:
New resized image.
"""
width, height = image.size
new_height = int(height * new_width / width)
new_img = image.resize((new_width, new_height))
return new_img
def load_img(path: str) -> Any:
"""Reads image from path and return PIL.Image instance.
Args:
path: A string of image path.
Returns:
A PIL.Image instance.
"""
img = tf.io.read_file(path)
img = tf.image.decode_jpeg(img, channels=3)
return Image.fromarray(np.uint8(img)).convert("RGB")
def decode_image(
image_str_tensor: tf.string, new_height: int, new_width: int
) -> tf.float32:
"""Converts and resizes image bytes to image tensor.
Args:
image_str_tensor: A string of image bytes.
new_height: An integer of new image height.
new_width: An integer of new image width.
Returns:
An image tensor.
"""
image = tf.io.decode_image(image_str_tensor, 3, expand_animations=False)
image = tf.image.resize(image, (new_height, new_width))
return image
def get_label_map(label_map_yaml_filepath: str) -> dict[int, str]:
"""Returns class id to label mapping given a filepath to the label map.
Args:
label_map_yaml_filepath: A string of label map yaml file path.
Returns:
A dictionary of class id to label mapping.
"""
with tf.io.gfile.GFile(label_map_yaml_filepath, "rb") as input_file:
label_map = yaml.safe_load(input_file.read())["label_map"]
return label_map
def get_prediction_instances(test_filepath: str, new_width: int = -1) -> Any:
"""Generate instance from image path to pass to Vertex AI Endpoint for prediction.
Args:
test_filepath: A string of test image path.
new_width: An integer of new image width.
Returns:
A list of instances.
"""
if new_width <= 0:
with tf.io.gfile.GFile(test_filepath, "rb") as input_file:
encoded_string = base64.b64encode(input_file.read()).decode("utf-8")
else:
img = load_img(test_filepath)
width, height = img.size
print("original input image size: ", width, " , ", height)
new_height = int(height * new_width / width)
new_img = img.resize((new_width, new_height))
print("resized input image size: ", new_width, " , ", new_height)
buffered = io.BytesIO()
new_img.save(buffered, format="JPEG")
encoded_string = base64.b64encode(buffered.getvalue()).decode("utf-8")
instances = [{
"encoded_image": {"b64": encoded_string},
}]
return instances
def vqa_predict(
endpoint: Any,
question_prompts: Sequence[str],
image: Any,
language_code: str = "en",
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> Sequence[str]:
"""Predicts the answer to a question about an image using an Endpoint."""
# Resize and convert image to base64 string.
resized_image = resize_image(image, new_width)
resized_image_base64 = image_to_base64(resized_image)
instances = []
if question_prompts:
# Format question prompt
question_prompt_format = "answer {} {}\n"
for question_prompt in question_prompts:
if question_prompt:
instances.append({
"prompt": question_prompt_format.format(
language_code, question_prompt
),
"image": resized_image_base64,
})
else:
instances.append({
"image": resized_image_base64,
})
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
return [pred.get("response") for pred in response.predictions]
def caption_predict(
endpoint: Any,
language_code: str,
image: Any,
caption_prompt: bool = False,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Predicts a caption for a given image using an Endpoint."""
# Resize and convert image to base64 string.
resized_image = resize_image(image, new_width)
resized_image_base64 = image_to_base64(resized_image)
instance = {"image": resized_image_base64}
if caption_prompt:
# Format caption prompt
caption_prompt_format = "caption {}\n"
instance["prompt"] = caption_prompt_format.format(language_code)
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
return response.predictions[0].get("response")
def ocr_predict(
endpoint: Any,
ocr_prompt: str,
image: Any,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Extracts text from a given image using an Endpoint."""
# Resize and convert image to base64 string.
resized_image = resize_image(image, new_width)
resized_image_base64 = image_to_base64(resized_image)
instance = {"image": resized_image_base64}
if ocr_prompt:
instance["prompt"] = ocr_prompt
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
return response.predictions[0].get("response")
def detect_predict(
endpoint: Any,
detect_prompt: str,
image: Any,
new_width: int = 1000,
use_dedicated_endpoint: bool = False,
) -> str:
"""Predicts the answer to a question about an image using an Endpoint."""
# Resize and convert image to base64 string.
resized_image = resize_image(image, new_width)
resized_image_base64 = image_to_base64(resized_image)
instance = {"image": resized_image_base64}
if detect_prompt:
instance["prompt"] = detect_prompt
instances = [instance]
response = endpoint.predict(
instances=instances, use_dedicated_endpoint=use_dedicated_endpoint
)
return response.predictions[0].get("response")
def copy_model_artifacts(
model_id: str,
model_source: str,
model_destination: str,
) -> None:
"""Copies model artifacts from model_source to model_destination.
model_source and model_destination should be GCS path.
Args:
model_id: The model id.
model_source: The source of the model artifact.
model_destination: The destination of the model artifact.
"""
if not model_source.startswith(GCS_URI_PREFIX):
raise ValueError(
f"{model_source} is not a GCS path starting with {GCS_URI_PREFIX}."
)
if not model_destination.startswith(GCS_URI_PREFIX):
raise ValueError(
f"{model_destination} is not a GCS path starting with {GCS_URI_PREFIX}."
)
model_source = f"{model_source}/{model_id}"
model_destination = f"{model_destination}/{model_id}"
print("Copying model artifact from ", model_source, " to ", model_destination)
subprocess.check_output([
"gcloud",
"storage",
"cp",
"-r",
model_source,
model_destination,
])
def get_quota(project_id: str, region: str, resource_id: str) -> int:
"""Returns the quota for a resource in a region.
Args:
project_id: The project id.
region: The region.
resource_id: The resource id.
Returns:
The quota for the resource in the region. Returns -1 if can not figure out
the quota.
Raises:
RuntimeError: If the command to get quota fails.
"""
service_endpoint = "aiplatform.googleapis.com"
command = (
"gcloud alpha services quota list"
f" --service={service_endpoint} --consumer=projects/{project_id}"
f" --filter='{service_endpoint}/{resource_id}' --format=json"
)
process = subprocess.run(
command, shell=True, capture_output=True, text=True, check=True
)
if process.returncode == 0:
quota_data = json.loads(process.stdout)
else:
raise RuntimeError(f"Error fetching quota data: {process.stderr}")
if not quota_data or "consumerQuotaLimits" not in quota_data[0]:
return -1
if (
not quota_data[0]["consumerQuotaLimits"]
or "quotaBuckets" not in quota_data[0]["consumerQuotaLimits"][0]
):
return -1
all_regions_data = quota_data[0]["consumerQuotaLimits"][0]["quotaBuckets"]
# If the quota data does not have dimensions, it is global quota. However,
# global quota may be overridden by regional quota. So we need to check the
# global quota first.
global_quota = -1
if (
all_regions_data
and "dimensions" not in all_regions_data[0]
and "effectiveLimit" in all_regions_data[0]
):
global_quota = int(all_regions_data[0]["effectiveLimit"])
for region_data in all_regions_data:
if (
region_data.get("dimensions")
and region_data["dimensions"]["region"] == region
):
if "effectiveLimit" in region_data:
return int(region_data["effectiveLimit"])
else:
return 0
return global_quota
def get_resource_id(
accelerator_type: str,
is_for_training: bool,
is_spot: bool = False,
is_restricted_image: bool = False,
is_dynamic_workload_scheduler: bool = False,
) -> str:
"""Returns the resource id for a given accelerator type and the use case.
Args:
accelerator_type: The accelerator type.
is_for_training: Whether the resource is used for training. Set false for
serving use case.
is_spot: Whether the resource is used with Spot.
is_restricted_image: Whether the image is hosted in `vertex-ai-restricted`.
is_dynamic_workload_scheduler: Whether the resource is used with Dynamic
Workload Scheduler.
Returns:
The resource id.
"""
accelerator_suffix_map = {
"NVIDIA_TESLA_V100": "nvidia_v100_gpus",
"NVIDIA_TESLA_P100": "nvidia_p100_gpus",
"NVIDIA_L4": "nvidia_l4_gpus",
"NVIDIA_TESLA_A100": "nvidia_a100_gpus",
"NVIDIA_A100_80GB": "nvidia_a100_80gb_gpus",
"NVIDIA_H100_80GB": "nvidia_h100_gpus",
"NVIDIA_H100_MEGA_80GB": "nvidia_h100_mega_gpus",
"NVIDIA_H200_141GB": "nvidia_h200_gpus",
"NVIDIA_TESLA_T4": "nvidia_t4_gpus",
"TPU_V6e": "tpu_v6e",
"TPU_V5e": "tpu_v5e",
"TPU_V3": "tpu_v3",
}
default_training_accelerator_map = {
key: f"custom_model_training_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
dws_training_accelerator_map = {
key: f"custom_model_training_preemptible_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
restricted_image_training_accelerator_map = {
"NVIDIA_A100_80GB": "restricted_image_training_nvidia_a100_80gb_gpus",
}
spot_serving_accelerator_map = {
key: f"custom_model_serving_preemptible_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
serving_accelerator_map = {
key: f"custom_model_serving_{accelerator_suffix_map[key]}"
for key in accelerator_suffix_map
}
if is_for_training:
if is_restricted_image and is_dynamic_workload_scheduler:
raise ValueError(
"Dynamic Workload Scheduler does not work for restricted image"
" training."
)
training_accelerator_map = (
restricted_image_training_accelerator_map
if is_restricted_image
else default_training_accelerator_map
)
if accelerator_type in training_accelerator_map:
if is_dynamic_workload_scheduler:
return dws_training_accelerator_map[accelerator_type]
else:
return training_accelerator_map[accelerator_type]
else:
raise ValueError(
f"Could not find accelerator type: {accelerator_type} for training."
)
else:
if is_dynamic_workload_scheduler:
raise ValueError("Dynamic Workload Scheduler does not work for serving.")
accelerator_map = (
spot_serving_accelerator_map if is_spot else serving_accelerator_map
)
if accelerator_type in accelerator_map:
return accelerator_map[accelerator_type]
else:
raise ValueError(
f"Could not find accelerator type: {accelerator_type} for serving."
)
def check_quota(
project_id: str,
region: str,
accelerator_type: str,
accelerator_count: int,
is_for_training: bool,
is_spot: bool = False,
is_restricted_image: bool = False,
is_dynamic_workload_scheduler: bool = False,
) -> None:
"""Checks if the project and the region has the required quota.
Args:
project_id: The project id.
region: The region.
accelerator_type: The accelerator type.
accelerator_count: The number of accelerators to check quota for.
is_for_training: Whether the resource is used for training. Set false for
serving use case.
is_spot: Whether the resource is used with Spot.
is_restricted_image: Whether the image is hosted in `vertex-ai-restricted`.
is_dynamic_workload_scheduler: Whether the resource is used with Dynamic
Workload Scheduler.
"""
resource_id = get_resource_id(
accelerator_type,
is_for_training=is_for_training,
is_spot=is_spot,
is_restricted_image=is_restricted_image,
is_dynamic_workload_scheduler=is_dynamic_workload_scheduler,
)
quota = get_quota(project_id, region, resource_id)
quota_request_instruction = (
"Either use "
"a different region or request additional quota. Follow "
"instructions here "
"https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota"
" to check quota in a region or request additional quota for "
"your project."
)
if quota == -1:
raise ValueError(
f"Quota not found for: {resource_id} in {region}."
f" {quota_request_instruction}"
)
if quota < accelerator_count:
raise ValueError(
f"Quota not enough for {resource_id} in {region}: {quota} <"
f" {accelerator_count}. {quota_request_instruction}"
)
def get_deploy_source() -> str:
"""Gets deploy_source string based on running environment."""
vertex_product = os.environ.get("VERTEX_PRODUCT", "")
match vertex_product:
case "COLAB_ENTERPRISE":
return "notebook_colab_enterprise"
case "WORKBENCH_INSTANCE":
return "notebook_workbench"
case _:
# Legacy workbench, legacy colab, or other custom environments.
return "notebook_environment_unspecified"
def _is_operation_done(op_name: str, region: str) -> bool:
"""Checks if the operation is done.
Args:
op_name: The name of the operation to poll.
region: The region of the operation.
Returns:
True if the operation is done, False otherwise.
Raises:
ValueError: If the operation failed.
"""
creds, _ = auth.default()
auth_req = auth.transport.requests.Request()
creds.refresh(auth_req)
headers = {
"Authorization": f"Bearer {creds.token}",
}
url = f"https://{region}-aiplatform.googleapis.com/ui/{op_name}"
response = requests.get(url, headers=headers)
operation_data = response.json()
if "error" in operation_data:
raise ValueError(f"Operation failed: {operation_data['error']}")
return operation_data.get("done", False)
def poll_and_wait(
op_name: str, region: str, total_wait: int, interval: int = 60
) -> None:
"""Polls the operation and waits for it to complete.
Args:
op_name: The name of the operation to poll.
region: The region of the operation.
total_wait: The total wait time in seconds.
interval: The interval between each poll in seconds.
Raises:
TimeoutError: If the operation times out.
"""
start_time = time.time()
while True:
if _is_operation_done(op_name, region):
break
time_elapsed = time.time() - start_time
if time_elapsed > total_wait:
raise TimeoutError(
f"Operation timed out after {int(time_elapsed)} seconds."
)
print(
"\rStill waiting for operation... Elapsed time in seconds:"
f" {int(time_elapsed):<6}",
end="",
flush=True,
)
time.sleep(interval)
@@ -1,598 +0,0 @@
"""Functions for dataset validation.
This tool is used to validate the dataset against the given template.
"""
from collections.abc import Callable
import json
import multiprocessing
import os
import subprocess
from typing import Any, Union
from absl import logging
import accelerate
import datasets
import transformers
GCS_URI_PREFIX = "gs://"
GCSFUSE_URI_PREFIX = "/gcs/"
LOCAL_BASE_MODEL_DIR = "/tmp/base_model_dir"
LOCAL_TEMPLATE_DIR = "/tmp/template_dir"
_TEMPLATE_DIRNAME = "templates"
_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME = "vertex-ai-samples"
_VERTEX_AI_SAMPLES_GITHUB_TEMPLATE_DIR = (
"community-content/vertex_model_garden/model_oss/peft/train/vmg/templates"
)
_MODELS_REQUIRING_PAD_TOKEN = ("llama", "falcon", "mistral", "mixtral")
_MODELS_REQUIRING_EOS_TOEKN = ("gemma-2b", "gemma-7b")
_DESCRIPTION_KEY = "description"
_SOURCE_KEY = "source"
_PROMPT_INPUT_KEY = "prompt_input"
_PROMPT_NO_INPUT_KEY = "prompt_no_input"
_RESPONSE_SEPARATOR = "response_separator"
_INSTRUCTION_SEPARATOR = "instruction_separator"
_CHAT_TEMPLATE_KEY = "chat_template"
_KNOWN_KEYS = (
_DESCRIPTION_KEY,
_SOURCE_KEY,
_PROMPT_INPUT_KEY,
_PROMPT_NO_INPUT_KEY,
_RESPONSE_SEPARATOR,
_INSTRUCTION_SEPARATOR,
_CHAT_TEMPLATE_KEY,
)
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 is not None and input_path.startswith(GCS_URI_PREFIX)
def force_gcs_fuse_path(gcs_uri: str) -> str:
"""Converts gs:// uris to their /gcs/ equivalents. No-op for other uris.
Args:
gcs_uri: The GCS URI to convert.
Returns:
The converted GCS URI.
"""
if is_gcs_path(gcs_uri):
return GCSFUSE_URI_PREFIX + gcs_uri[len(GCS_URI_PREFIX) :]
else:
return gcs_uri
def download_gcs_uri_to_local(
gcs_uri: str,
destination_dir: str = LOCAL_BASE_MODEL_DIR,
check_path_exists: bool = True,
) -> str:
"""Downloads GCS URI to local.
If GCS URI is a directory, gs://some/folder is downloaded to
/destination_dir/folder. If GCS URI is a file, gs://some/file is downloaded to
/destination_dir/file.
Args:
gcs_uri: GCS URI to download.
destination_dir: Local directory directory.
check_path_exists: Whether to check if the path exists.
Returns:
Local path to target folder/file.
"""
target = os.path.join(
destination_dir,
os.path.basename(os.path.normpath(gcs_uri)),
)
if check_path_exists and os.path.exists(target):
logging.info("File %s already exists.", target)
return target
if accelerate.PartialState().is_local_main_process:
logging.info(
"Downloading file(s) from %s to %s...", gcs_uri, destination_dir
)
if not os.path.exists(destination_dir):
os.mkdir(destination_dir)
subprocess.check_output([
"gsutil",
"-m",
"cp",
"-r",
gcs_uri,
destination_dir,
])
logging.info("Downloaded file(s) from %s to %s.", gcs_uri, destination_dir)
# Make sure ALL processes process to next step after data downloading is done.
# It matters for the main process to wait for other processes as well.
accelerate.PartialState().wait_for_everyone()
return target
def get_template(template_path: str) -> dict[str, str]:
"""Gets the template dictionary given the file path.
Args:
template_path: Path to the template file.
Returns:
A dictionary of the template.
Raises:
ValueError: If the template file does not exist or contains unknown keys.
"""
if is_gcs_path(template_path):
template_path = force_gcs_fuse_path(template_path)
elif not os.path.isfile(template_path):
template_path = os.path.join(
os.path.dirname(__file__),
_TEMPLATE_DIRNAME,
template_path + ".json",
)
if not os.path.isfile(template_path):
raise ValueError(f"Template file {template_path} does not exist.")
with open(template_path, "r") as f:
template_json: dict[str, str] = json.load(f)
for key in template_json:
if key not in _KNOWN_KEYS:
raise ValueError(f"Unknown key {key} in template {template_path}.")
return template_json
def get_response_separator(template_json: dict[str, str]) -> Union[str, None]:
return template_json.get(_RESPONSE_SEPARATOR, None)
def get_instruction_separator(
template_json: dict[str, str],
) -> Union[str, None]:
return template_json.get(_INSTRUCTION_SEPARATOR, None)
def _format_template_fn(
template: str,
input_column: str,
tokenizer: transformers.PreTrainedTokenizer | None = None,
) -> Callable[[dict[str, str]], dict[str, str]]:
"""Formats a dataset example according to a template.
Args:
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
input_column: The input column in the dataset to be used or updated by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
tokenizer: The tokenizer to use for chat_template templates.
Returns:
A function that formats data according to the template.
"""
template_json = get_template(template)
if _CHAT_TEMPLATE_KEY not in template_json:
def format_fn(example: dict[str, str]) -> dict[str, str]:
format_dict = {key: value for key, value in example.items()}
if format_dict.get(input_column):
format_str = template_json[_PROMPT_INPUT_KEY]
elif _PROMPT_NO_INPUT_KEY in template_json:
format_str = template_json[_PROMPT_NO_INPUT_KEY]
else:
raise KeyError(
f"The template {os.path.basename(template)} does not contain"
f" {_PROMPT_INPUT_KEY} or {_PROMPT_NO_INPUT_KEY} key."
)
try:
return {input_column: format_str.format(**format_dict)}
except KeyError as e:
raise KeyError(
f"The template {os.path.basename(template)} contains a key {e} in"
f" {_PROMPT_INPUT_KEY} or {_PROMPT_NO_INPUT_KEY} that does not"
" exist in the dataset example. The dataset example looks like"
f" {format_dict}."
) from e
return format_fn
elif (
_PROMPT_INPUT_KEY in template_json
or _PROMPT_NO_INPUT_KEY in template_json
):
raise ValueError(
f"chat_template templates do not support {_PROMPT_INPUT_KEY} or"
f" {_PROMPT_NO_INPUT_KEY} templates."
)
else:
if tokenizer is None:
raise ValueError("A tokenizer is required for chat_template templates.")
# Assign HuggingFace jinja template.
tokenizer.chat_template = template_json[_CHAT_TEMPLATE_KEY]
def format_fn(example: dict[str, str]) -> dict[str, str]:
try:
return {
input_column: tokenizer.apply_chat_template(
example[input_column],
tokenize=False,
add_generation_prompt=False,
)
}
except KeyError as e:
raise KeyError(
f"The template {os.path.basename(template)} contains a key {e} in"
f" {_CHAT_TEMPLATE_KEY} that does not exist in the dataset example."
) from e
return format_fn
def _get_split_string(
split: str,
dataset_percent: int | None = None,
dataset_k_rows: int | None = None,
) -> str:
"""Gets the formatted split string for the dataset.
This is used to format the split string as per
https://huggingface.co/docs/datasets/v2.21.0/loading#slice-splits. Also, this
function will only be used to load the partial dataset for validating the
dataset against the template.
Args:
split: Split of the dataset.
dataset_percent: The percentage of the dataset to load.
dataset_k_rows: The top k sequences to load from the dataset.
Returns:
A formatted split string.
"""
# Validate the dataset_percent and dataset_k_rows values.
if dataset_percent and dataset_k_rows:
raise ValueError(
"You can set either validate_percentage_of_dataset or"
" validate_k_rows_of_dataset, but not both."
)
if dataset_percent:
logging.info("Loading %d percent of the dataset...", dataset_percent)
return f"{split}[:{dataset_percent}%]"
if dataset_k_rows:
logging.info("Loading top %d rows of the dataset...", dataset_k_rows)
return f"{split}[:{dataset_k_rows}]"
return split
def _github_template_path(template: str) -> str:
"""Generates the path to the template in the Vertex AI Samples GitHub repo.
Args:
template: Name of the template.
Returns:
The path to the template in the Vertex AI Samples GitHub repo.
"""
# vertex-ai-samples directory may lie under separate directory depending on
# the scratch_dir parameter in the notebook execution environment.
vertex_ai_samples_abs_path = os.getcwd().split(
_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME
)[0]
return os.path.join(
vertex_ai_samples_abs_path,
_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME,
_VERTEX_AI_SAMPLES_GITHUB_TEMPLATE_DIR,
template + ".json",
)
def _get_dataset(
dataset_name: str,
split: str,
num_proc: int | None = None,
) -> datasets.DatasetDict:
"""Gets a dataset.
Args:
dataset_name: Name of the dataset or path to a custom dataset.
split: Split of the dataset.
num_proc: Number of processors to use.
Returns:
A dataset.
"""
dataset_name = force_gcs_fuse_path(dataset_name)
if os.path.isfile(dataset_name):
# Custom dataset.
return datasets.load_dataset(
"json",
data_files=[dataset_name],
split=split,
num_proc=num_proc,
)
# HF dataset.
return datasets.load_dataset(dataset_name, split=split, num_proc=num_proc)
def should_add_pad_token(model_id: str) -> bool:
"""Returns whether the model requires adding a special pad token.
Args:
model_id: The name of the model.
Returns:
True if the model requires adding a special pad token, False otherwise.
"""
model_config = transformers.AutoConfig.from_pretrained(model_id)
if model_config.model_type is None:
return False
return any(
s.lower() in model_config.model_type.lower()
for s in _MODELS_REQUIRING_PAD_TOKEN
)
def should_add_eos_token(model_id: str) -> bool:
"""Returns whether the model requires adding a special eos token.
Args:
model_id: The name of the model.
Returns:
True if the model requires adding a special eos token, False otherwise.
"""
return any(m in model_id for m in _MODELS_REQUIRING_EOS_TOEKN)
def load_tokenizer(
pretrained_model_id: str,
padding_side: str | None = None,
access_token: str | None = None,
) -> transformers.AutoTokenizer:
"""Loads tokenizer based on `pretrained_model_id`.
Args:
pretrained_model_id: The name of the pretrained model.
padding_side: The side to pad the input on.
access_token: The access token to use for the tokenizer.
Returns:
The tokenizer.
"""
tokenizer_kwargs = {}
if should_add_eos_token(pretrained_model_id):
tokenizer_kwargs["add_eos_token"] = True
if padding_side:
tokenizer_kwargs["padding_side"] = padding_side
with accelerate.PartialState().local_main_process_first():
tokenizer = transformers.AutoTokenizer.from_pretrained(
pretrained_model_id,
trust_remote_code=False,
use_fast=True,
token=access_token,
**tokenizer_kwargs,
)
if should_add_pad_token(pretrained_model_id):
tokenizer.add_special_tokens({"pad_token": "[PAD]"})
return tokenizer
def get_filtered_dataset(
dataset: Any,
input_column: str,
max_seq_length: int,
tokenizer: transformers.PreTrainedTokenizer,
example_removed_threshold: float = 50.0,
) -> Any:
"""Returns the dataset by removing examples that are longer than max_seq_length.
Args:
dataset: The dataset to filter.
input_column: The input column in the dataset to be used.
max_seq_length: The maximum sequence length.
tokenizer: The tokenizer.
example_removed_threshold: The percent threshold for the number of examples
removed from the dataset. It should be in the range of [0, 100].
Returns:
The filtered dataset.
Raises:
ValueError: If more than `example_removed_threshold` of the dataset is
filtered out.
"""
actual_dataset_length = len(dataset)
filtered_dataset = dataset.filter(
lambda x: len(tokenizer(x[input_column])["input_ids"]) <= max_seq_length
)
filtered_dataset_length = len(filtered_dataset)
if actual_dataset_length != filtered_dataset_length:
examples_removed_percent = (
(actual_dataset_length - filtered_dataset_length)
* 100
/ actual_dataset_length
)
logging.info(
"(%.2f%%) of examples token length is <= max-seq-length(%d); (%.2f%%) >"
" max-seq-length. Filtering out %d example(s) which are longer than"
" max-seq-length.",
100 - examples_removed_percent,
max_seq_length,
examples_removed_percent,
actual_dataset_length - filtered_dataset_length,
)
if examples_removed_percent > example_removed_threshold:
raise ValueError(
"More than %.2f%% of the dataset is filtered out. This may be due to"
" small value of max-seq-length(%d) or incorrect template. Please"
" increase the max-seq-length or check the template."
% (examples_removed_percent, max_seq_length)
)
print(f"Some formatted examples from the dataset are: {filtered_dataset[:5]}")
return filtered_dataset
def format_dataset(
dataset: datasets.Dataset,
input_column: str,
template: str = None,
tokenizer: transformers.PreTrainedTokenizer | None = None,
) -> datasets.Dataset:
"""Takes a raw dataset and formats it using a template and tokenizer.
Args:
dataset: The raw (unprocessed) dataset to format.
input_column: The input column in the dataset to be used or updaded by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
tokenizer: The tokenizer to use for chat_template templates.
Returns:
A dataset compatible with the template.
"""
return dataset.map(
_format_template_fn(
template,
input_column=input_column,
tokenizer=tokenizer,
)
)
def load_dataset_with_template(
dataset_name: str,
split: str,
input_column: str,
template: str = None,
tokenizer: transformers.PreTrainedTokenizer | None = None,
) -> tuple[Any, Any]:
"""Loads dataset with templates.
Args:
dataset_name: Name of the dataset or path to a custom dataset.
split: Split of the dataset.
input_column: The input column in the dataset to be used or updaded by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
tokenizer: The tokenizer to use for chat_template templates.
Returns:
The raw dataset and the dataset compatible with the template.
"""
raw = _get_dataset(dataset_name, split=split)
if template:
templated = format_dataset(raw, input_column, template, tokenizer)
else:
templated = None
return raw, templated
def validate_dataset_with_template(
dataset_name: str,
split: str,
input_column: str,
template: str,
tokenizer: transformers.PreTrainedTokenizer | None = None,
max_seq_length: int | None = None,
use_multiprocessing: bool = False,
validate_percentage_of_dataset: int | None = None,
validate_k_rows_of_dataset: int | None = None,
example_removed_threshold: float = 50.0,
) -> Any:
"""Validates dataset with templates.
This function will be used to load the dataset and validate it against the
template. In case of validation, we also allow the users to load the dataset
partially by allowing them to read x% or top k rows of the dataset. To
validate the dataset, the template file must be available in the GCS bucket
and the dataset must be available either in the GCS bucket or Hugging Face.
Args:
dataset_name: Name of the dataset or path to a custom dataset.
split: Split of the dataset.
input_column: The input column in the dataset to be used or updaded by the
template. If it does not exist, the template's `prompt_no_input` will be
used, and the input_column will be created.
template: Name of the JSON template file under `templates/` or GCS path to
the template file.
tokenizer: The tokenizer to use for chat_template templates.
max_seq_length: The maximum sequence length.
use_multiprocessing: If True, it will use multiprocessing to load the
dataset.
validate_percentage_of_dataset: The percentage of the dataset to load.
validate_k_rows_of_dataset: The top k sequences to load from the dataset.
example_removed_threshold: The threshold for the number of examples removed
from the dataset.
Returns:
None if the validation is successful, otherwise returns the error message.
"""
if not template:
raise ValueError("template is required for validate_dataset.")
if not dataset_name:
raise ValueError("dataset_name is empty.")
if not split:
raise ValueError("split is empty.")
split = _get_split_string(
split,
validate_percentage_of_dataset,
validate_k_rows_of_dataset,
)
num_proc = multiprocessing.cpu_count() if use_multiprocessing else 1
# gcsfuse cannot be used from the notebook runtime env. Hence, we have
# to download dataset and template from gcs to local.
if is_gcs_path(dataset_name):
dataset_name = download_gcs_uri_to_local(dataset_name, LOCAL_BASE_MODEL_DIR)
if is_gcs_path(template):
template_path = download_gcs_uri_to_local(template, LOCAL_TEMPLATE_DIR)
elif os.path.isfile(_github_template_path(template)):
template_path = _github_template_path(template)
else:
raise ValueError(
f"Template file {template} does not exist. To validate the"
" dataset, please provide a valid GCS path for the template or a valid"
" template name from"
f" https://github.com/GoogleCloudPlatform/{_VERTEX_AI_SAMPLES_GITHUB_REPO_NAME}/tree/main/{_VERTEX_AI_SAMPLES_GITHUB_TEMPLATE_DIR}."
)
dataset = format_dataset(
_get_dataset(dataset_name, split, num_proc),
input_column,
template_path,
tokenizer,
)
if tokenizer is not None:
get_filtered_dataset(
dataset=dataset,
input_column=input_column,
max_seq_length=max_seq_length,
tokenizer=tokenizer,
example_removed_threshold=example_removed_threshold,
)
print(
"Dataset {} is compatible with the {} template.".format(
os.path.basename(dataset_name), os.path.basename(template)
)
)
@@ -1,275 +0,0 @@
"""Utility functions for interacting with Google Cloud Platform."""
import datetime
import logging
import os
import subprocess
import uuid
from google.cloud import aiplatform
import requests
# Configure logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
def get_project_id() -> str:
"""Read cloud project id from metadata service."""
project_request = requests.get(
"http://metadata.google.internal/computeMetadata/v1/project/project-id",
headers={"Metadata-Flavor": "Google"},
)
return project_request.text
def get_region() -> str:
"""Read region from metadata service."""
region_request = requests.get(
"http://metadata.google.internal/computeMetadata/v1/instance/region",
headers={"Metadata-Flavor": "Google"},
)
return region_request.text.split("/")[-1]
# Get the default cloud project id and region
PROJECT_ID = get_project_id()
REGION = get_region()
def init_aiplatform(project: str = None, location: str = None) -> None:
"""Initialize the Vertex AI SDK.
Args:
project: The Google Cloud project ID.
location: The Google Cloud location.
"""
project = PROJECT_ID if project is None else project
location = REGION if location is None else location
aiplatform.init(project=project, location=location)
subprocess.call([
"gcloud",
"services",
"enable",
"aiplatform.googleapis.com",
"compute.googleapis.com",
])
def run_command(command: list[str]) -> str:
"""Runs a shell command and returns the output.
Args:
command: The shell command to run as a list.
Returns:
The output of the command.
"""
try:
result = subprocess.run(
command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
)
return result.stdout
except subprocess.CalledProcessError as e:
logger.error("Error: %s", e.stderr)
raise e
def enable_apis() -> None:
"""Enable the Vertex AI API and Compute Engine API."""
logger.info("Enabling Vertex AI API and Compute Engine API.")
run_command([
"gcloud",
"services",
"enable",
"aiplatform.googleapis.com",
"compute.googleapis.com",
])
def setup_buckets(bucket_uri: str, model_bucket_name: str) -> tuple[str, str]:
"""Set up Cloud Storage buckets for storing experiment artifacts.
Args:
bucket_uri: The bucket URI provided by the user.
model_bucket_name: The name of the model bucket.
Returns:
A tuple containing the bucket name and model bucket path.
"""
if not bucket_uri.strip():
# Generate a default bucket URI if none provided
now = datetime.datetime.now().strftime("%Y%m%d%H%M%S")
bucket_uri = f"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}"
logger.info("No bucket URI provided. Using default bucket: %s", bucket_uri)
else:
if not bucket_uri.startswith("gs://"):
raise ValueError("Bucket URI must start with 'gs://'.")
# Remove any trailing slashes
bucket_uri = bucket_uri.rstrip("/")
bucket_name = "/".join(bucket_uri.split("/")[:3])
# Check if bucket exists
try:
run_command(["gsutil", "ls", "-b", bucket_uri])
logger.info("Bucket %s already exists.", bucket_uri)
except subprocess.CalledProcessError:
logger.info("Creating bucket %s.", bucket_uri)
# Create the bucket in the same region as the project
run_command(["gsutil", "mb", "-l", REGION, bucket_uri])
# Construct the model bucket path
model_bucket = os.path.join(bucket_uri, model_bucket_name)
# Check if the model bucket exists (as a folder within the main bucket)
try:
run_command(["gsutil", "ls", model_bucket])
logger.info("Model bucket %s already exists.", model_bucket)
except subprocess.CalledProcessError:
logger.info("Creating model bucket %s.", model_bucket)
# Create the model bucket folder
run_command(["gsutil", "cp", "/dev/null", model_bucket + "/"])
return bucket_name, model_bucket
def get_service_account() -> str:
"""Get the default service account."""
shell_output = run_command(["gcloud", "projects", "describe", PROJECT_ID])
project_number_line = next(
(line for line in shell_output.splitlines() if "projectNumber" in line),
None,
)
if project_number_line:
project_number = project_number_line.split(":")[1].strip().replace("'", "")
service_account = f"{project_number}-compute@developer.gserviceaccount.com"
logger.info("Using default Service Account: %s", service_account)
return service_account
else:
raise ValueError("Could not find project number in gcloud output.")
def get_project_number() -> str:
"""Get the default project number."""
shell_output = run_command(["gcloud", "projects", "describe", PROJECT_ID])
project_number_line = next(
(line for line in shell_output.splitlines() if "projectNumber" in line),
None,
)
if project_number_line:
project_number = project_number_line.split(":")[1].strip().replace("'", "")
logger.info("Using default Project Number: %s", project_number)
return project_number
else:
raise ValueError("Could not find project number in gcloud output.")
def provision_permissions(service_account: str, bucket_name: str) -> None:
"""Provision permissions to the service account with the GCS bucket."""
if bucket_name:
run_command([
"gsutil",
"iam",
"ch",
f"serviceAccount:{service_account}:roles/storage.admin",
bucket_name,
])
def set_gcloud_project() -> None:
"""Set gcloud config project."""
run_command(["gcloud", "config", "set", "project", PROJECT_ID])
def initialize(
bucket_uri: str, model_bucket_name: str, create_bucket: bool
) -> tuple[str, str]:
"""Initialize the environment.
Args:
bucket_uri: The bucket URI provided by the user.
model_bucket_name: The name of the model bucket.
create_bucket: Whether to create the bucket or not.
Returns:
A tuple containing the model bucket path and service account.
"""
enable_apis()
bucket_name = None
if create_bucket:
bucket_name, model_bucket = setup_buckets(bucket_uri, model_bucket_name)
else:
model_bucket = None
service_account = get_service_account()
provision_permissions(service_account, bucket_name)
set_gcloud_project()
return model_bucket, service_account
def clean_resources_ui(
project_id: str,
region: str,
endpoint_name: str,
delete_bucket: bool,
bucket_name: str = None,
) -> str:
"""UI function for cleaning a specific Vertex AI endpoint and its model."""
if delete_bucket and not bucket_name:
raise ValueError("Bucket name is required when 'Delete Bucket' is checked.")
try:
delete_endpoint_and_model(project_id, region, endpoint_name)
bucket_status_message = ""
if delete_bucket:
bucket_status_message = delete_gcs_bucket(bucket_name)
if endpoint_name:
return (
f"Endpoint {endpoint_name} and associated model deleted successfully!"
f" {bucket_status_message}"
)
else:
return (
"There are currently no endpoints available for deletion."
f" {bucket_status_message}"
)
except Exception as e: # pylint: disable=broad-exception-caught
return f"Error cleaning up resources: {e}"
def delete_endpoint_and_model(
project_id: str, region: str, endpoint_name: str
) -> None:
"""Deletes a specific Vertex AI endpoint and its associated model."""
if endpoint_name:
endpoint_id = endpoint_name.split(" - ")[0]
endpoint_resource_name = (
f"projects/{project_id}/locations/{region}/endpoints/{endpoint_id}"
)
endpoint = aiplatform.Endpoint(
endpoint_resource_name, project=project_id, location=region
)
deployed_models = endpoint.list_models()
for deployed_model in deployed_models:
endpoint.undeploy(deployed_model_id=deployed_model.id)
model = aiplatform.Model(deployed_model.model)
model.delete()
endpoint.delete()
def delete_gcs_bucket(bucket_name: str) -> str:
"""Deletes a GCS bucket using gsutil."""
try:
run_command(["gsutil", "-m", "rm", "-r", bucket_name])
logger.info("Bucket %s deleted using gsutil.", bucket_name)
return f"Bucket {bucket_name} deleted successfully!"
except subprocess.CalledProcessError as e:
logger.error(
"Error deleting bucket %s using gsutil: %s", bucket_name, str(e)
)
return f"Bucket {bucket_name} could not be found or deleted. "
@@ -42,7 +42,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/gke_model_ui_deployment_notebook.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -57,21 +57,12 @@
"source": [
"# Overview\n",
"\n",
"This notebook will guide you through the initial step of testing your recently\n",
"deployed model with text prompts. Depending on your deployed model's inference\n",
"setup, the notebook utilizes either Text Generation Inference\n",
"[TGI](https://huggingface.co/docs/text-generation-inference/en/index) or\n",
"[vLLM](https://developers.googleblog.com/en/inference-with-gemma-using-dataflow-and-vllm/#:~:text=model%20frameworks%20simple.-,What%20is%20vLLM%3F,-vLLM%20is%20an),\n",
"two efficient serving frameworks that enhance the performance of your GPU model.\n",
"Ready to see your deployed model respond? Run the cells below and start\n",
"experimenting with different prompts!\n",
"This notebook will guide you through the initial step of testing your recently deployed model with text prompts. Depending on your deployed model's inference setup, the notebook utilizes either Text Generation Inference [TGI](https://huggingface.co/docs/text-generation-inference/en/index) or [vLLM](https://developers.googleblog.com/en/inference-with-gemma-using-dataflow-and-vllm/#:~:text=model%20frameworks%20simple.-,What%20is%20vLLM%3F,-vLLM%20is%20an), two efficient serving frameworks that enhance the performance of your GPU model. Ready to see your deployed model respond? Run the cells below and start experimenting with different prompts!\n",
"\n",
"### Prerequisites\n",
"\n",
"Before proceeding with this notebook, ensure you have already deployed a model\n",
"using the Google Cloud Console. You can find an overview of AI and Machine\n",
"Learning services on\n",
"[GKE AI/ML](https://console.cloud.google.com/kubernetes/aiml/overview).\n",
"Before proceeding with this notebook, ensure you have already deployed a model using the Google Cloud Console. You can find an overview of AI and Machine Learning services on [GKE AI/ML](https://console.cloud.google.com/kubernetes/aiml/overview).\n",
"\n",
"\n",
"### Objective\n",
"\n",
@@ -79,46 +70,33 @@
"\n",
"### GPUs\n",
"\n",
"GPUs let you accelerate specific workloads running on your nodes, such as\n",
"machine learning and data processing. GKE provides a range of machine type\n",
"options for node configuration, including machine types with NVIDIA H100, L4,\n",
"and A100 GPUs.\n",
"GPUs let you accelerate specific workloads running on your nodes, such as machine learning and data processing. GKE provides a range of machine type options for node configuration, including machine types with NVIDIA H100, L4, and A100 GPUs.\n",
"\n",
"### Understanding the Inference Frameworks\n",
"\n",
"Your model is running on one of two popular and efficient serving frameworks:\n",
"vLLM or Text Generation Inference (TGI). The following sections provide a brief\n",
"overview of each to give you context on the underlying technology powering your\n",
"model.\n",
"Your model is running on one of two popular and efficient serving frameworks: vLLM or Text Generation Inference (TGI). The following sections provide a brief overview of each to give you context on the underlying technology powering your model.\n",
"\n",
"\n",
"#### TGI\n",
"\n",
"TGI is a highly optimized open-source LLM serving framework that can increase\n",
"serving throughput on GPUs. TGI includes features such as:\n",
"TGI is a highly optimized open-source LLM serving framework that can increase serving throughput on GPUs. TGI includes features such as:\n",
"\n",
"* Optimized transformer implementation with PagedAttention\n",
"* Continuous batching to improve the overall serving throughput\n",
"* Tensor parallelism and distributed serving on multiple GPUs\n",
"* Optimized transformer implementation with PagedAttention\n",
"* Continuous batching to improve the overall serving throughput\n",
"* Tensor parallelism and distributed serving on multiple GPUs\n",
"\n",
"To learn more, refer to the\n",
"[TGI documentation](https://github.com/huggingface/text-generation-inference/blob/main/README.md)\n",
"To learn more, refer to the [TGI documentation](https://github.com/huggingface/text-generation-inference/blob/main/README.md)\n",
"\n",
"#### vLLM\n",
"\n",
"vLLM is another fast and easy-to-use library for LLM inference and serving. It's\n",
"known for its high throughput and efficiency, and it leverages PagedAttention.\n",
"Key features include:\n",
"vLLM is another fast and easy-to-use library for LLM inference and serving. It's known for its high throughput and efficiency, and it leverages PagedAttention. Key features include:\n",
"\n",
"* PagedAttention: Efficient memory management for handling long sequences and\n",
" dynamic workloads.\n",
"* Continuous batching: Maximizes GPU utilization by batching incoming\n",
" requests.\n",
"* High-throughput serving: Designed for production-level serving with low\n",
" latency.\n",
"* Optimized CUDA kernels.\n",
"* PagedAttention: Efficient memory management for handling long sequences and dynamic workloads.\n",
"* Continuous batching: Maximizes GPU utilization by batching incoming requests.\n",
"* High-throughput serving: Designed for production-level serving with low latency.\n",
"* Optimized CUDA kernels.\n",
"\n",
"To learn more, refer to the\n",
"[vLLM documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/vllm/use-vllm)"
"To learn more, refer to the [vLLM documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/vllm/use-vllm)"
]
},
{
@@ -133,9 +111,9 @@
"source": [
"# @title # Connect to Google Cloud Project\n",
"# @markdown #### Run this cell to configure your Google Cloud environment for Kubernetes (GKE) operations.\n",
"# @markdown\n",
"\n",
"# @markdown #### Actions:\n",
"# @markdown 1. **Connects to Project:** Retrieves and sets your Google Cloud project ID.\n",
"# @markdown 1. **Connects to Project & Region:** Retrieves and sets your Google Cloud project ID and region.\n",
"# @markdown 3. **Installs `kubectl`:** Installs the Kubernetes command-line tool.\n",
"\n",
"import os\n",
@@ -143,6 +121,9 @@
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Set up gcloud.\n",
"! gcloud config set project \"$PROJECT_ID\"\n",
"! gcloud services enable container.googleapis.com\n",
@@ -163,406 +144,124 @@
"outputs": [],
"source": [
"# @title # Select Cluster and Deployment { vertical-output: true }\n",
"# @markdown **Instructions:**\n",
"# @markdown\n",
"# @markdown Run this cell using the ▶ button. Then, use the interactive widgets that appear below:\n",
"# @markdown 1. **Select Cluster:** From the first dropdown, choose the GKE cluster where your model deployment is running. Note: the list only contains autopilot clusters.\n",
"# @markdown 2. **Select Namespace:** After selecting a cluster, choose the Kubernetes *Namespace* where your deployment resides within that cluster.\n",
"# @markdown 3. **Select Deployment:** After selecting a cluster, this dropdown will populate with the names of deployments found.\n",
"\n",
"# @markdown ## Instruction:\n",
"\n",
"# @markdown This cell provides interactive dropdown menus to select a Google Kubernetes Engine (GKE) cluster and a deployment within that cluster.\n",
"\n",
"# @markdown ***Please select a cluster and deployment before proceeding.***\n",
"\n",
"import json\n",
"import subprocess\n",
"\n",
"import ipywidgets as widgets\n",
"from IPython.display import Markdown, clear_output, display\n",
"from IPython.display import display\n",
"\n",
"# --- Globals and Configuration ---\n",
"DEFAULT_NAMESPACE = \"default\"\n",
"SELECTED_DEPLOYMENT = None\n",
"SELECTED_NAMESPACE = DEFAULT_NAMESPACE\n",
"deployment_dropdown = None\n",
"namespace_dropdown = None\n",
"cluster_dropdown = None\n",
"output_area = widgets.Output()\n",
"\n",
"\n",
"# --- Data Fetching Functions ---\n",
"def get_clusters(project_id):\n",
" \"\"\"Fetches autopilot GKE clusters for a given project.\"\"\"\n",
" # Note: Uses broad exception handling as per original code.\n",
"def get_clusters(p, r):\n",
" try:\n",
" cmd = f\"gcloud container clusters list --filter=autopilot.enabled=true --format=json --project={project_id}\"\n",
" result = subprocess.run(\n",
" cmd, shell=True, capture_output=True, text=True, check=True, timeout=60\n",
" return (\n",
" subprocess.run(\n",
" [\n",
" \"gcloud\",\n",
" \"container\",\n",
" \"clusters\",\n",
" \"list\",\n",
" \"--project\",\n",
" p,\n",
" \"--region\",\n",
" r,\n",
" \"--format=value(name)\",\n",
" ],\n",
" capture_output=True,\n",
" text=True,\n",
" check=True,\n",
" )\n",
" .stdout.strip()\n",
" .split(\"\\n\")\n",
" )\n",
" clusters_data = json.loads(result.stdout)\n",
" # Create a map of cluster name to its region/location\n",
" return {c[\"name\"]: c[\"location\"] for c in clusters_data}\n",
" except Exception as e:\n",
" # Original code prints error and returns empty dict\n",
" print(f\"Error getting clusters: {e}\")\n",
" return {}\n",
"\n",
"\n",
"# Fetch clusters immediately using PROJECT_ID assumed to be globally defined\n",
"# Note: This relies on PROJECT_ID being set *before* this cell runs.\n",
"try:\n",
" CLUSTER_REGION_MAP = get_clusters(PROJECT_ID)\n",
"except NameError:\n",
" print(\n",
" \"Error: PROJECT_ID variable is not defined. Please define it in a previous cell.\"\n",
" )\n",
" CLUSTER_REGION_MAP = {} # Define as empty to prevent errors later\n",
"\n",
"\n",
"def get_deployments(cluster, region, namespace):\n",
" \"\"\"Fetches deployments from a specific namespace in a cluster.\"\"\"\n",
" # Note: Uses PROJECT_ID as a global variable as per original code.\n",
" # Note: Uses broad exception handling as per original code.\n",
" target_namespace = namespace if namespace else DEFAULT_NAMESPACE\n",
" try:\n",
" # Ensure credentials for the target cluster\n",
" cred_cmd = [\n",
" \"gcloud\",\n",
" \"container\",\n",
" \"clusters\",\n",
" \"get-credentials\",\n",
" cluster,\n",
" f\"--location={region}\",\n",
" f\"--project={PROJECT_ID}\",\n",
" ]\n",
" subprocess.run(cred_cmd, capture_output=True, text=True, check=True, timeout=60)\n",
"\n",
" # Fetch deployments using kubectl\n",
" kubectl_cmd = [\n",
" \"kubectl\",\n",
" \"get\",\n",
" \"deployments\",\n",
" f\"--namespace={target_namespace}\",\n",
" \"-o\",\n",
" \"json\",\n",
" ]\n",
" result = subprocess.run(\n",
" kubectl_cmd, capture_output=True, text=True, check=True, timeout=60\n",
" )\n",
" deployments_data = json.loads(result.stdout)\n",
" # Extract deployment names\n",
" return [item[\"metadata\"][\"name\"] for item in deployments_data.get(\"items\", [])]\n",
" except Exception as e:\n",
" # Original code prints error and returns empty list\n",
" print(f\"Error fetching deployments from namespace '{target_namespace}': {e}\")\n",
" except subprocess.CalledProcessError as e:\n",
" print(f\"Error: {e}\")\n",
" return []\n",
"\n",
"\n",
"def get_namespaces(cluster, region, project_id):\n",
" \"\"\"Fetches namespaces for a given cluster.\"\"\"\n",
" # Note: Uses broad exception handling as per original code.\n",
"def get_deployments(c, r):\n",
" try:\n",
" # Ensure credentials for the target cluster\n",
" cred_cmd = [\n",
" \"gcloud\",\n",
" \"container\",\n",
" \"clusters\",\n",
" \"get-credentials\",\n",
" cluster,\n",
" f\"--location={region}\",\n",
" f\"--project={project_id}\",\n",
" ]\n",
" subprocess.run(cred_cmd, capture_output=True, text=True, check=True, timeout=60)\n",
"\n",
" # Fetch namespaces using kubectl\n",
" kubectl_cmd = [\"kubectl\", \"get\", \"namespaces\", \"-o\", \"json\"]\n",
" result = subprocess.run(\n",
" kubectl_cmd, capture_output=True, text=True, check=True, timeout=60\n",
" subprocess.run(\n",
" [\n",
" \"gcloud\",\n",
" \"container\",\n",
" \"clusters\",\n",
" \"get-credentials\",\n",
" c,\n",
" \"--location\",\n",
" r,\n",
" ],\n",
" capture_output=True,\n",
" text=True,\n",
" check=True,\n",
" )\n",
" namespaces_data = json.loads(result.stdout)\n",
" # Extract namespace names\n",
" all_ns = [item[\"metadata\"][\"name\"] for item in namespaces_data.get(\"items\", [])]\n",
" return all_ns\n",
" except Exception as e:\n",
" # Original code displays error in output_area and returns None\n",
" with output_area:\n",
" # Clear previous output before showing error\n",
" clear_output(wait=True)\n",
" display(\n",
" Markdown(\n",
" f\"<font color='red'>Error processing namespaces for **{cluster}**: {e}</font>\"\n",
" )\n",
" )\n",
" return None\n",
" deployments = json.loads(\n",
" subprocess.run(\n",
" [\"kubectl\", \"get\", \"deployments\", \"-o\", \"json\"],\n",
" capture_output=True,\n",
" text=True,\n",
" check=True,\n",
" ).stdout\n",
" )\n",
" return [i[\"metadata\"][\"name\"] for i in deployments[\"items\"]]\n",
" except subprocess.CalledProcessError as e:\n",
" print(f\"Error: {e}\")\n",
" return []\n",
"\n",
"\n",
"# --- Event Handlers ---\n",
"def on_deployment_select(change):\n",
" \"\"\"Handles changes in the deployment selection.\"\"\"\n",
"def create_deployment_dropdown(cluster_name, region, on_select_deployment):\n",
" deployments = get_deployments(cluster_name, region)\n",
" deployments_with_prompt = [\"Select Deployment\"] + deployments\n",
" deployment_dropdown = widgets.Dropdown(\n",
" options=deployments_with_prompt,\n",
" description=\"Deployments\",\n",
" disabled=False,\n",
" width=\"4000px\",\n",
" )\n",
" deployment_dropdown.observe(\n",
" lambda c: on_select_deployment(c[\"new\"])\n",
" if c[\"type\"] == \"change\" and c[\"name\"] == \"value\"\n",
" else None,\n",
" names=\"value\",\n",
" )\n",
" return deployment_dropdown\n",
"\n",
"\n",
"def on_deployment_select(deployment_name):\n",
" global SELECTED_DEPLOYMENT\n",
" if change[\"type\"] == \"change\" and change[\"name\"] == \"value\":\n",
" SELECTED_DEPLOYMENT = change[\"new\"]\n",
" with output_area:\n",
" clear_output(wait=True)\n",
" current_cluster = cluster_dropdown.value\n",
"\n",
" # Display context message\n",
" if current_cluster != \"Select Cluster\":\n",
" # Use SELECTED_NAMESPACE global which should be set by on_namespace_change\n",
" # or default if namespace hasn't been selected yet.\n",
" ns_context = SELECTED_NAMESPACE or DEFAULT_NAMESPACE\n",
" ns_info = f\"Cluster: **{current_cluster}**, Namespace: **{ns_context}**\"\n",
" display(Markdown(ns_info))\n",
"\n",
" # Display selection message if a valid deployment is chosen\n",
" if (\n",
" SELECTED_DEPLOYMENT\n",
" and SELECTED_DEPLOYMENT != \"Select Deployment\"\n",
" and SELECTED_DEPLOYMENT != \"Loading...\"\n",
" ):\n",
" mes = f\"\"\"Selected deployment: **{SELECTED_DEPLOYMENT}**\"\"\"\n",
" display(Markdown(mes))\n",
"\n",
"\n",
"def update_deployment_dropdown(cluster_name, namespace_to_use):\n",
" \"\"\"Updates the deployment list based on cluster/namespace change.\"\"\"\n",
" global deployment_dropdown, SELECTED_DEPLOYMENT\n",
" target_namespace = namespace_to_use if namespace_to_use else DEFAULT_NAMESPACE\n",
"\n",
" # Reset selection before fetching/updating\n",
" SELECTED_DEPLOYMENT = None\n",
" deployment_dropdown.disabled = True # Disable while loading/updating\n",
" deployment_dropdown.options = [\"Loading...\"]\n",
" deployment_dropdown.value = \"Loading...\"\n",
"\n",
" # Clear output area and show loading context\n",
" with output_area:\n",
" clear_output(wait=True)\n",
" display(Markdown(f\"Cluster: **{cluster_name}**\"))\n",
" if namespace_to_use:\n",
" display(Markdown(f\"Namespace: **{namespace_to_use}**\"))\n",
" display(Markdown(\"Fetching deployments...\"))\n",
"\n",
" # Fetch deployments (assuming CLUSTER_REGION_MAP and PROJECT_ID are available)\n",
" region = CLUSTER_REGION_MAP.get(cluster_name)\n",
" if not region:\n",
" with output_area:\n",
" clear_output(wait=True)\n",
" display(\n",
" Markdown(\n",
" f\"<font color='red'>Error: Region not found for cluster {cluster_name}.</font>\"\n",
" )\n",
" )\n",
" deployment_dropdown.options = [\"Error loading\"]\n",
" deployment_dropdown.value = \"Error loading\"\n",
" return # Stop if region is missing\n",
"\n",
" deployments = get_deployments(cluster_name, region, target_namespace)\n",
"\n",
" # Update dropdown options\n",
" new_options = [\"Select Deployment\"] + deployments\n",
" deployment_dropdown.options = new_options\n",
"\n",
" # Set final state based on results\n",
" if deployments:\n",
" deployment_dropdown.value = \"Select Deployment\"\n",
" deployment_dropdown.disabled = False\n",
" status_message = f\"Found {len(deployments)} deployment(s) in namespace **{target_namespace}**.\"\n",
" else:\n",
" deployment_dropdown.value = \"Select Deployment\" # Keep prompt\n",
" deployment_dropdown.disabled = True # No valid options to select\n",
" # Check if get_deployments printed an error or if it just returned empty\n",
" if not output_area.outputs: # If no error printed by get_deployments\n",
" status_message = (\n",
" f\"No deployments found in namespace **{target_namespace}**.\"\n",
" )\n",
" else:\n",
" status_message = None # Error likely already shown\n",
"\n",
" # Update output area with final status\n",
" with output_area:\n",
" clear_output(wait=True)\n",
" display(Markdown(f\"Cluster: **{cluster_name}**\"))\n",
" if namespace_to_use:\n",
" display(Markdown(f\"Namespace: **{namespace_to_use}**\"))\n",
" if status_message:\n",
" display(Markdown(status_message))\n",
"\n",
"\n",
"def update_namespace_dropdown(cluster_name):\n",
" \"\"\"Updates the namespace list based on cluster change.\"\"\"\n",
" global namespace_dropdown, SELECTED_NAMESPACE\n",
" global deployment_dropdown, SELECTED_DEPLOYMENT # Need to reset deployment too\n",
"\n",
" # Reset namespace state and dependent deployment dropdown\n",
" SELECTED_NAMESPACE = None # Reset selection\n",
" SELECTED_DEPLOYMENT = None\n",
" namespace_dropdown.disabled = True\n",
" namespace_dropdown.options = [\"Loading...\"]\n",
" namespace_dropdown.value = \"Loading...\"\n",
" deployment_dropdown.options = [\"Select Deployment\"]\n",
" deployment_dropdown.value = \"Select Deployment\"\n",
" deployment_dropdown.disabled = True\n",
"\n",
" # Clear output area and show loading context\n",
" with output_area:\n",
" clear_output(wait=True)\n",
" display(Markdown(f\"Cluster: **{cluster_name}**\"))\n",
" display(Markdown(\"Fetching namespaces...\"))\n",
"\n",
" # Fetch namespaces (assuming CLUSTER_REGION_MAP and PROJECT_ID are available)\n",
" region = CLUSTER_REGION_MAP.get(cluster_name)\n",
" if not region:\n",
" with output_area:\n",
" clear_output(wait=True)\n",
" display(\n",
" Markdown(\n",
" f\"<font color='red'>Error: Region not found for cluster {cluster_name}.</font>\"\n",
" )\n",
" )\n",
" namespace_dropdown.options = [\"Error loading\"]\n",
" namespace_dropdown.value = \"Error loading\"\n",
" return # Stop if region is missing\n",
"\n",
" # Assuming PROJECT_ID is globally available\n",
" namespaces = get_namespaces(cluster_name, region, PROJECT_ID)\n",
"\n",
" # Update dropdown options based on fetch result\n",
" if namespaces is not None: # Success (get_namespaces returns None on error)\n",
" new_options = [\"Select Namespace\"] + namespaces # Use \"Select Namespace\" prompt\n",
" namespace_dropdown.options = new_options\n",
" namespace_dropdown.value = \"Select Namespace\"\n",
" namespace_dropdown.disabled = False\n",
" status_message = (\n",
" f\"Found {len(namespaces)} namespace(s). Select one to list deployments.\"\n",
" )\n",
" else: # Error occurred during fetch\n",
" namespace_dropdown.options = [\"Error loading\"] # Keep error state\n",
" namespace_dropdown.value = \"Error loading\"\n",
" namespace_dropdown.disabled = True\n",
" status_message = None # Error already displayed by get_namespaces\n",
"\n",
" # Update output area with final status\n",
" with output_area:\n",
" clear_output(wait=True)\n",
" display(Markdown(f\"Cluster: **{cluster_name}**\"))\n",
" if status_message:\n",
" display(Markdown(status_message))\n",
" SELECTED_DEPLOYMENT = deployment_name\n",
" print(f\"Selected deployment: {SELECTED_DEPLOYMENT}\")\n",
"\n",
"\n",
"def on_cluster_change(change):\n",
" \"\"\"Handles cluster selection changes.\"\"\"\n",
" # Globals not strictly needed here as it calls update_namespace_dropdown which uses them\n",
" if change[\"type\"] == \"change\" and change[\"name\"] == \"value\":\n",
" cluster = change[\"new\"]\n",
"\n",
" # Clear output area for new selection process\n",
" with output_area:\n",
" clear_output(wait=True)\n",
"\n",
" if cluster == \"Select Cluster\":\n",
" # Reset namespace dropdown\n",
" namespace_dropdown.options = [\"Select Namespace\"] # Correct prompt\n",
" namespace_dropdown.value = \"Select Namespace\"\n",
" namespace_dropdown.disabled = True\n",
" # Reset deployment dropdown\n",
" deployment_dropdown.options = [\"Select Deployment\"]\n",
" deployment_dropdown.value = \"Select Deployment\"\n",
" deployment_dropdown.disabled = True\n",
" # Clear globals\n",
" global SELECTED_NAMESPACE, SELECTED_DEPLOYMENT\n",
" SELECTED_NAMESPACE = None\n",
" SELECTED_DEPLOYMENT = None\n",
" else:\n",
" # Trigger update for the namespace dropdown\n",
" update_namespace_dropdown(cluster)\n",
" if change[\"new\"] == \"Select Cluster\":\n",
" return\n",
" deployment_dropdown = create_deployment_dropdown(\n",
" change[\"new\"], REGION, on_deployment_select\n",
" )\n",
" display(deployment_dropdown)\n",
"\n",
"\n",
"def on_namespace_change(change):\n",
" \"\"\"Handles namespace selection: fetches deployments.\"\"\"\n",
" global SELECTED_NAMESPACE, cluster_dropdown, deployment_dropdown # Added deployment_dropdown\n",
" if change[\"type\"] == \"change\" and change[\"name\"] == \"value\":\n",
" new_namespace = change[\"new\"]\n",
"\n",
" # Get current cluster value\n",
" current_cluster = cluster_dropdown.value\n",
"\n",
" # Handle placeholder/loading/error values or if cluster isn't selected\n",
" if (\n",
" new_namespace in [\"Select Namespace\", \"Loading...\", \"Error loading\"]\n",
" or current_cluster == \"Select Cluster\"\n",
" ):\n",
" SELECTED_NAMESPACE = None\n",
" # Reset deployment dropdown state\n",
" deployment_dropdown.options = [\"Select Deployment\"]\n",
" deployment_dropdown.value = \"Select Deployment\"\n",
" deployment_dropdown.disabled = True\n",
" global SELECTED_DEPLOYMENT\n",
" SELECTED_DEPLOYMENT = None\n",
" # Clear output area for clean state\n",
" with output_area:\n",
" clear_output(wait=True)\n",
" if current_cluster != \"Select Cluster\": # Keep cluster context\n",
" display(Markdown(f\"Cluster: **{current_cluster}**\"))\n",
" if new_namespace == \"Select Namespace\":\n",
" display(Markdown(\"Select a namespace to list deployments.\"))\n",
" return # Don't proceed to fetch deployments\n",
"\n",
" # Valid namespace selected\n",
" SELECTED_NAMESPACE = new_namespace\n",
"\n",
" # Trigger update for the deployment dropdown\n",
" if current_cluster != \"Select Cluster\":\n",
" update_deployment_dropdown(current_cluster, SELECTED_NAMESPACE)\n",
"\n",
"\n",
"# --- Main Widget Setup ---\n",
"if CLUSTER_REGION_MAP:\n",
" clusters_with_prompt = [\"Select Cluster\"] + sorted(list(CLUSTER_REGION_MAP.keys()))\n",
"clusters = get_clusters(PROJECT_ID, REGION)\n",
"if clusters:\n",
" # @markdown Run this cell to display the Cluster dropdown menu:\n",
" clusters_with_prompt = [\"Select Cluster\"] + clusters\n",
" cluster_dropdown = widgets.Dropdown(\n",
" options=clusters_with_prompt,\n",
" value=\"Select Cluster\", # Set initial value\n",
" description=\"Cluster:\",\n",
" style={\"description_width\": \"initial\"},\n",
" layout=widgets.Layout(width=\"auto\"), # Auto width\n",
" options=clusters_with_prompt, description=\"Clusters\", disabled=False\n",
" )\n",
"\n",
" namespace_dropdown = widgets.Dropdown(\n",
" options=[\"Select Namespace\"], # Correct initial prompt\n",
" value=\"Select Namespace\",\n",
" description=\"Namespace:\",\n",
" disabled=True, # Initially disabled\n",
" style={\"description_width\": \"initial\"},\n",
" layout=widgets.Layout(width=\"auto\"),\n",
" )\n",
"\n",
" deployment_dropdown = widgets.Dropdown(\n",
" options=[\"Select Deployment\"],\n",
" value=\"Select Deployment\",\n",
" description=\"Deployment:\",\n",
" disabled=True, # Initially disabled\n",
" style={\"description_width\": \"initial\"},\n",
" layout=widgets.Layout(width=\"auto\"),\n",
" )\n",
"\n",
" # Observe changes\n",
" cluster_dropdown.observe(on_cluster_change, names=\"value\")\n",
" namespace_dropdown.observe(on_namespace_change, names=\"value\")\n",
" deployment_dropdown.observe(on_deployment_select, names=\"value\")\n",
"\n",
" # Display initial status and widgets\n",
" print(\n",
" f\"Found {len(CLUSTER_REGION_MAP)} Autopilot Cluster(s) in Project '{PROJECT_ID}'.\\n\"\n",
" )\n",
" display(cluster_dropdown, namespace_dropdown, deployment_dropdown, output_area)\n",
"\n",
" display(cluster_dropdown)\n",
"else:\n",
" # Handle case where PROJECT_ID might be missing or no clusters found\n",
" if \"PROJECT_ID\" not in globals() or not PROJECT_ID:\n",
" error_message = \"Error: PROJECT_ID variable is not defined or empty. Please define it in a previous cell.\"\n",
" else:\n",
" error_message = f\"Error: No Autopilot clusters found or accessible in project '{PROJECT_ID}'. Check Project ID, permissions, and ensure Autopilot clusters exist.\"\n",
" print(error_message)\n",
" # Display error message using a widget for better integration in notebook\n",
" display(widgets.HTML(f\"<font color='red'>{error_message}</font>\"))\n",
" # Keep output_area widget displayed even on error for potential messages from retries etc.\n",
" display(output_area)"
" print(f\"No clusters found in {PROJECT_ID}/{REGION}.\")"
]
},
{
@@ -575,202 +274,91 @@
},
"outputs": [],
"source": [
"# @title # Chat completion for text-only models { vertical-output: true}\n",
"# @title # Chat completion for text-only models {run:\"auto\", vertical-output: true}\n",
"\n",
"# @markdown You may send prompts to the model server for prediction.\n",
"# @markdown\n",
"# @markdown * **user_prompt (string):** This is the text prompt you provide to the language model. It's the question or instruction e (e.g., \"Explain neural networks\").\n",
"\n",
"# @markdown * **temperature (number):** This parameter controls the randomness of the model's output. It influences how the model selects the next token in the sequence it generates. Typical values range from 0.2 to 1.0.\n",
"\n",
"# @markdown * **max_tokens (number):** This parameter refers to the maximum number of tokens (words or sub-word units) that the model is allowed to generate in its response.\n",
"\n",
"import ipywidgets as widgets\n",
"from IPython.display import HTML\n",
"\n",
"\n",
"def _run_kubectl(cmd):\n",
" \"\"\"Executes a kubectl command and returns its stdout.\"\"\"\n",
" result = subprocess.run(cmd, capture_output=True, text=True, check=True, timeout=60)\n",
" return result.stdout.strip()\n",
"\n",
"\n",
"def get_deployment_pod_name(deployment, namespace):\n",
" \"\"\"Finds the running pod name for a given deployment and namespace.\"\"\"\n",
" cmd = [\n",
" \"kubectl\",\n",
" \"get\",\n",
" \"pods\",\n",
" \"-n\",\n",
" namespace,\n",
" \"-o\",\n",
" \"json\",\n",
" \"-l\",\n",
" f\"app={deployment}-app\",\n",
" \"--field-selector=status.phase=Running\",\n",
" ]\n",
"def get_deployment_pod_name(deployment):\n",
" try:\n",
" pods_json = _run_kubectl(cmd)\n",
" pods = json.loads(pods_json)\n",
" if pods.get(\"items\"):\n",
" return pods[\"items\"][0][\"metadata\"][\"name\"]\n",
" print(f\"No running pods found for {deployment} in {namespace}.\")\n",
" return None\n",
" label = deployment + \"-app\"\n",
" pods = json.loads(\n",
" subprocess.run(\n",
" [\"kubectl\", \"get\", \"pods\", \"-o\", \"json\", \"-l\", f\"app={label}\"],\n",
" capture_output=True,\n",
" check=True,\n",
" ).stdout\n",
" )\n",
" return pods[\"items\"][0][\"metadata\"][\"name\"] if pods[\"items\"] else None\n",
" except (\n",
" subprocess.CalledProcessError,\n",
" json.JSONDecodeError,\n",
" IndexError,\n",
" KeyError,\n",
" ) as e:\n",
" print(f\"Error getting pod name for {deployment} in {namespace}: {e}\")\n",
" IndexError,\n",
" ):\n",
" return None\n",
"\n",
"\n",
"def check_inference_label(pod_name, namespace):\n",
" \"\"\"Checks if the specified pod has the vLLM inference server label.\"\"\"\n",
" cmd = [\"kubectl\", \"get\", \"pod\", pod_name, \"-n\", namespace, \"-o\", \"json\"]\n",
"def check_vllm_label(pod_name):\n",
" \"\"\"Checks if the pod has the 'ai.gke.io/inference-server=vllm' label.\"\"\"\n",
" try:\n",
" pod_json = _run_kubectl(cmd)\n",
" labels = json.loads(pod_json).get(\"metadata\", {}).get(\"labels\", {})\n",
" result = subprocess.run(\n",
" [\"kubectl\", \"get\", \"pod\", pod_name, \"-o\", \"json\"],\n",
" capture_output=True,\n",
" check=True,\n",
" )\n",
" labels = json.loads(result.stdout)[\"metadata\"][\"labels\"]\n",
" return labels.get(\"ai.gke.io/inference-server\") == \"vllm\"\n",
" except (subprocess.CalledProcessError, json.JSONDecodeError, KeyError) as e:\n",
" print(f\"Error checking labels for pod {pod_name} in {namespace}: {e}\")\n",
" except (subprocess.CalledProcessError, KeyError, json.JSONDecodeError):\n",
" return False\n",
"\n",
"\n",
"def process_response(request, pod_name, pod_endpoint, is_vllm_inference, namespace):\n",
" \"\"\"Sends a request to the pod and processes the response.\"\"\"\n",
" json_data_escaped = json.dumps(request).replace(\"'\", \"'\\\\''\")\n",
" curl_cmd = f\"kubectl exec -n {namespace} -t {pod_name} -- curl -s -X POST http://{pod_endpoint}/generate -H \\\"Content-Type: application/json\\\" -d '{json_data_escaped}' 2> /dev/null\"\n",
"def process_response(request, pod_name, pod_endpoint, is_vllm):\n",
" response = !kubectl exec -t {pod_name} -- curl -X POST http://{pod_endpoint}/generate -H \"Content-Type: application/json\" -d '{json.dumps(request)}' 2> /dev/null\n",
" try:\n",
" response_raw = _run_kubectl([\"bash\", \"-c\", curl_cmd])\n",
" if not response_raw:\n",
" return f\"Error: Empty response from pod {pod_name}.\"\n",
" first_line = response_raw.splitlines()[0]\n",
" data = json.loads(first_line)\n",
"\n",
" if is_vllm_inference:\n",
" predictions = data.get(\"predictions\")\n",
" if isinstance(predictions, (list, tuple)) and predictions:\n",
" return predictions[0]\n",
" return f\"Error: Unexpected vLLM format. Raw: {first_line}\"\n",
" else: # TGI format\n",
" generated_text = data.get(\"generated_text\")\n",
" if generated_text is not None:\n",
" return generated_text\n",
" return f\"Error: Unexpected TGI format. Raw: {first_line}\"\n",
"\n",
" except json.JSONDecodeError as e:\n",
" raw_response = (\n",
" response_raw.splitlines()[0]\n",
" if \"response_raw\" in locals() and response_raw\n",
" else \"N/A\"\n",
" )\n",
" return f\"Error decoding JSON: {e}. Raw: {raw_response}\"\n",
" except (subprocess.CalledProcessError, IndexError, KeyError, TypeError) as e:\n",
" raw_response = (\n",
" response_raw.splitlines()[0]\n",
" if \"response_raw\" in locals() and response_raw\n",
" else \"N/A\"\n",
" )\n",
" return f\"Error processing response: {e}. Raw: {raw_response}\"\n",
" except Exception as e:\n",
" return f\"Unexpected error during response processing: {e}\"\n",
" data = json.loads(response[0])\n",
" if is_vllm:\n",
" return data[\"predictions\"][0]\n",
" else:\n",
" return data[\"generated_text\"]\n",
" except (json.JSONDecodeError, KeyError, IndexError) as e:\n",
" return f\"Error: {e}, Raw: {response}\"\n",
"\n",
"\n",
"# --- Widgets Setup ---\n",
"user_prompt_widget = widgets.Textarea(\n",
" value=\"What is AI?\",\n",
" description=\"User Prompt:\",\n",
" layout=widgets.Layout(width=\"95%\", height=\"100px\"),\n",
")\n",
"temperature_widget = widgets.FloatSlider(\n",
" value=0.50, min=0.0, max=1.0, step=0.01, description=\"Temperature:\"\n",
")\n",
"max_tokens_widget = widgets.IntSlider(\n",
" value=250, min=1, max=2048, step=1, description=\"Max Tokens:\"\n",
")\n",
"submit_button = widgets.Button(description=\"Submit\")\n",
"output_area_response = widgets.Output()\n",
"deployment_pod = get_deployment_pod_name(SELECTED_DEPLOYMENT)\n",
"is_vllm_inference = check_vllm_label(deployment_pod)\n",
"\n",
"user_prompt = \"What is AI?\" # @param {type: \"string\"}\n",
"temperature = 0.50 # @param {type: \"number\"}\n",
"max_tokens = 250 # @param {type: \"number\"}\n",
"\n",
"# --- Submit Button Logic ---\n",
"def on_submit_clicked(b):\n",
" \"\"\"Handles the submit button click event.\"\"\"\n",
" with output_area_response:\n",
" clear_output()\n",
" if (\n",
" \"SELECTED_DEPLOYMENT\" not in globals()\n",
" or \"SELECTED_NAMESPACE\" not in globals()\n",
" ):\n",
" display(\n",
" Markdown(\n",
" \"**Error:** `SELECTED_DEPLOYMENT` or `SELECTED_NAMESPACE` not defined.\"\n",
" )\n",
" )\n",
" return\n",
"request = {\n",
" \"max_tokens\": 250 if max_tokens is None else max_tokens,\n",
" \"temperature\": 0.5 if temperature is None else temperature,\n",
"}\n",
"\n",
" print(\n",
" f\"Target: {SELECTED_DEPLOYMENT} in {SELECTED_NAMESPACE}. \\n\\nRequesting response...\"\n",
" )\n",
"if is_vllm_inference:\n",
" request[\"prompt\"] = user_prompt\n",
"else:\n",
" request[\"inputs\"] = user_prompt\n",
"\n",
" pod_name = get_deployment_pod_name(SELECTED_DEPLOYMENT, SELECTED_NAMESPACE)\n",
" if not pod_name:\n",
" display(\n",
" Markdown(\n",
" f\"**Error:** Could not find running pod for `{SELECTED_DEPLOYMENT}`.\"\n",
" )\n",
" )\n",
" return\n",
"model_service = SELECTED_DEPLOYMENT + \"-service\"\n",
"output = !kubectl get endpoints {model_service}\n",
"pod_endpoint = output[1].split()[1]\n",
"\n",
" is_vllm = check_inference_label(pod_name, SELECTED_NAMESPACE)\n",
" request = {\n",
" \"max_tokens\": max_tokens_widget.value,\n",
" \"temperature\": temperature_widget.value,\n",
" \"prompt\" if is_vllm else \"inputs\": user_prompt_widget.value,\n",
" }\n",
" service = f\"{SELECTED_DEPLOYMENT}-service\"\n",
" endpoint_cmd = [\n",
" \"kubectl\",\n",
" \"get\",\n",
" \"endpoints\",\n",
" service,\n",
" \"-n\",\n",
" SELECTED_NAMESPACE,\n",
" ]\n",
"\n",
" try:\n",
" endpoint_output = _run_kubectl(endpoint_cmd).splitlines()\n",
" if len(endpoint_output) < 2 or len(endpoint_output[1].split()) < 2:\n",
" display(\n",
" Markdown(\n",
" f\"**Error:** Endpoint data incomplete for service `{service}`.\"\n",
" )\n",
" )\n",
" print(\"kubectl output:\\n\", \"\\n\".join(endpoint_output))\n",
" return\n",
" endpoint = endpoint_output[1].split()[\n",
" 1\n",
" ] # Assumes format: NAME ENDPOINTS AGE -> service ip:port,... age\n",
" response = process_response(\n",
" request, pod_name, endpoint, is_vllm, SELECTED_NAMESPACE\n",
" )\n",
" display(Markdown(f\"**Response:**\\n\\n{response}\"))\n",
"\n",
" except subprocess.CalledProcessError as e:\n",
" display(\n",
" Markdown(\n",
" f\"**Error getting endpoints for `{service}`:**\\n```\\n{e.stderr}\\n```\"\n",
" )\n",
" )\n",
" except Exception as e:\n",
" display(Markdown(f\"**Unexpected Error:**\\n```\\n{e}\\n```\"))\n",
"\n",
"\n",
"# --- Display Widgets ---\n",
"submit_button.on_click(on_submit_clicked)\n",
"display(\n",
" user_prompt_widget,\n",
" temperature_widget,\n",
" max_tokens_widget,\n",
" submit_button,\n",
" output_area_response,\n",
"# @markdown ### Response:\n",
"response = process_response(request, deployment_pod, pod_endpoint, is_vllm_inference)\n",
"HTML(\n",
" '<div style=\"overflow-x: auto; font-size: 16px; line-height:'\n",
" f' 1.8;\">{response}</div>'\n",
")"
]
},
@@ -783,67 +371,39 @@
"source": [
"# Next Steps: Integrating the GKE Service Endpoint\n",
"\n",
"After successfully deploying a model on Google Kubernetes Engine (GKE) and\n",
"verifying it via a notebook, the next step is to integrate it into various\n",
"applications. This involves making HTTP requests to the service's endpoint from\n",
"your application code.\n",
"After successfully deploying a model on Google Kubernetes Engine (GKE) and verifying it via a notebook, the next step is to integrate it into various applications. This involves making HTTP requests to the service's endpoint from your application code.\n",
"\n",
"### Exposing the Service\n",
"\n",
"To make your deployed model accessible to applications, you'll need to expose\n",
"its service endpoint. Google Kubernetes Engine offers several ways to do this:\n",
"To make your deployed model accessible to applications, you'll need to expose its service endpoint. Google Kubernetes Engine offers several ways to do this:\n",
"\n",
"1. **Ingress:** Configure an Ingress resource to route external HTTP(S) traffic\n",
" to your service. Set up Ingress for either an internal Load Balancer\n",
" (accessible only within your VPC) or an external Load Balancer (accessible\n",
" from the internet).\n",
" [Learn more about GKE Ingress](https://cloud.google.com/kubernetes-engine/docs/concepts/ingress).\n",
"2. **Gateway API:** A more modern and feature-rich API for managing traffic\n",
" routing in Kubernetes. Similar to Ingress, Gateway API allows you to define\n",
" how external and internal traffic should be directed to your services.\n",
" [Explore GKE Gateway API](https://cloud.google.com/kubernetes-engine/docs/concepts/gateway-api).\n",
"1. **Ingress:** Configure an Ingress resource to route external HTTP(S) traffic to your service. Set up Ingress for either an internal Load Balancer (accessible only within your VPC) or an external Load Balancer (accessible from the internet). [Learn more about GKE Ingress](https://cloud.google.com/kubernetes-engine/docs/concepts/ingress).\n",
"2. **Gateway API:** A more modern and feature-rich API for managing traffic routing in Kubernetes. Similar to Ingress, Gateway API allows you to define how external and internal traffic should be directed to your services. [Explore GKE Gateway API](https://cloud.google.com/kubernetes-engine/docs/concepts/gateway-api).\n",
"\n",
"### Setting Up Autoscaling\n",
"\n",
"Ensure your model serving can handle varying traffic by configuring the\n",
"Horizontal Pod Autoscaler (HPA). HPA automatically scales the number of Pods\n",
"based on resource utilization or custom metrics, optimizing performance and\n",
"cost.\n",
"[See how to configure HPA](https://cloud.google.com/kubernetes-engine/docs/how-to/horizontal-pod-autoscaling).\n",
"Ensure your model serving can handle varying traffic by configuring the Horizontal Pod Autoscaler (HPA). HPA automatically scales the number of Pods based on resource utilization or custom metrics, optimizing performance and cost. [See how to configure HPA](https://cloud.google.com/kubernetes-engine/docs/how-to/horizontal-pod-autoscaling).\n",
"\n",
"### Setting Up Monitoring\n",
"\n",
"Monitor the health and performance of your deployed model using Google Cloud\n",
"Managed Service for Prometheus. Configure your model serving to expose\n",
"Prometheus metrics for comprehensive insights.\n",
"[Get started with Google Cloud Managed Prometheus](https://cloud.google.com/kubernetes-engine/docs/how-to/configure-automatic-application-monitoring).\n",
"Monitor the health and performance of your deployed model using Google Cloud Managed Service for Prometheus. Configure your model serving to expose Prometheus metrics for comprehensive insights. [Get started with Google Cloud Managed Prometheus](https://cloud.google.com/kubernetes-engine/docs/how-to/configure-automatic-application-monitoring).\n",
"\n",
"### Additional Resources:\n",
"\n",
"* #### Kubernetes Documentation:\n",
"* #### Kubernetes Documentation:\n",
" * Services: https://kubernetes.io/docs/concepts/services-networking/service/\n",
"\n",
" * Services:\n",
" https://kubernetes.io/docs/concepts/services-networking/service/\n",
"* #### Google Cloud Documentation:\n",
" * Google Kubernetes Engine (GKE): https://cloud.google.com/kubernetes-engine\n",
" * Cloud Load Balancing: https://cloud.google.com/load-balancing/docs/ingress\n",
" * Gateway API on GKE: https://cloud.google.com/kubernetes-engine/docs/concepts/gateway-api\n",
" * Learn about GPUs in GKE: https://cloud.google.com/kubernetes-engine/docs/concepts/gpus\n",
"\n",
"* #### Google Cloud Documentation:\n",
"* #### Python requests Library:\n",
" * https://requests.readthedocs.io/en/latest/\n",
"\n",
" * Google Kubernetes Engine (GKE):\n",
" https://cloud.google.com/kubernetes-engine\n",
" * Cloud Load Balancing:\n",
" https://cloud.google.com/load-balancing/docs/ingress\n",
" * Gateway API on GKE:\n",
" https://cloud.google.com/kubernetes-engine/docs/concepts/gateway-api\n",
" * Learn about GPUs in GKE:\n",
" https://cloud.google.com/kubernetes-engine/docs/concepts/gpus\n",
"\n",
"* #### Python requests Library:\n",
"\n",
" * https://requests.readthedocs.io/en/latest/\n",
"\n",
"* #### LangChain with Google Integrations:\n",
"\n",
" * The Langchain documentation is very useful:\n",
" https://python.langchain.com/docs/integrations/providers/google/"
"* #### LangChain with Google Integrations:\n",
" * The Langchain documentation is very useful: https://python.langchain.com/docs/integrations/providers/google/"
]
}
],
@@ -1,487 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"id": "Pr9TgOcV9vAXeqGiyTaTI5kS",
"metadata": {
"cellView": "form",
"id": "Pr9TgOcV9vAXeqGiyTaTI5kS"
},
"outputs": [],
"source": [
"# Copyright 2025 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",
"id": "M1CpgYundFwz",
"metadata": {
"id": "M1CpgYundFwz"
},
"source": [
"# Get started with your deployed model on GKE\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fgke_model_ui_deployment_notebook_auto.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/gke_model_ui_deployment_notebook_auto.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"id": "t2jj2XOgkS4F",
"metadata": {
"id": "t2jj2XOgkS4F"
},
"source": [
"# Overview\n",
"\n",
"This notebook will guide you through the initial step of testing your recently\n",
"deployed model with text prompts. Depending on your deployed model's inference\n",
"setup, the notebook utilizes either Text Generation Inference\n",
"[TGI](https://huggingface.co/docs/text-generation-inference/en/index) or\n",
"[vLLM](https://developers.googleblog.com/en/inference-with-gemma-using-dataflow-and-vllm/#:~:text=model%20frameworks%20simple.-,What%20is%20vLLM%3F,-vLLM%20is%20an),\n",
"two efficient serving frameworks that enhance the performance of your GPU model.\n",
"Ready to see your deployed model respond? Run the cells below and start\n",
"experimenting with different prompts!\n",
"\n",
"### Prerequisites\n",
"\n",
"Before proceeding with this notebook, ensure you have already deployed a model\n",
"using the Google Cloud Console. You can find an overview of AI and Machine\n",
"Learning services on\n",
"[GKE AI/ML](https://console.cloud.google.com/kubernetes/aiml/overview).\n",
"\n",
"### Objective\n",
"\n",
"Enable prompt-based testing of the AI model deployed on GKE\n",
"\n",
"### GPUs\n",
"\n",
"GPUs let you accelerate specific workloads running on your nodes, such as\n",
"machine learning and data processing. GKE provides a range of machine type\n",
"options for node configuration, including machine types with NVIDIA H100, L4,\n",
"and A100 GPUs.\n",
"\n",
"### Understanding the Inference Frameworks\n",
"\n",
"Your model is running on one of two popular and efficient serving frameworks:\n",
"vLLM or Text Generation Inference (TGI). The following sections provide a brief\n",
"overview of each to give you context on the underlying technology powering your\n",
"model.\n",
"\n",
"#### TGI\n",
"\n",
"TGI is a highly optimized open-source LLM serving framework that can increase\n",
"serving throughput on GPUs. TGI includes features such as:\n",
"\n",
"* Optimized transformer implementation with PagedAttention\n",
"* Continuous batching to improve the overall serving throughput\n",
"* Tensor parallelism and distributed serving on multiple GPUs\n",
"\n",
"To learn more, refer to the\n",
"[TGI documentation](https://github.com/huggingface/text-generation-inference/blob/main/README.md)\n",
"\n",
"#### vLLM\n",
"\n",
"vLLM is another fast and easy-to-use library for LLM inference and serving. It's\n",
"known for its high throughput and efficiency, and it leverages PagedAttention.\n",
"Key features include:\n",
"\n",
"* PagedAttention: Efficient memory management for handling long sequences and\n",
" dynamic workloads.\n",
"* Continuous batching: Maximizes GPU utilization by batching incoming\n",
" requests.\n",
"* High-throughput serving: Designed for production-level serving with low\n",
" latency.\n",
"* Optimized CUDA kernels.\n",
"\n",
"To learn more, refer to the\n",
"[vLLM documentation](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/vllm/use-vllm)"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "XMf-T58TkDy1",
"metadata": {
"cellView": "form",
"id": "XMf-T58TkDy1"
},
"outputs": [],
"source": [
"# @title # Connect to Google Cloud Project\n",
"# @markdown #### Run this cell to configure your Google Cloud environment for Kubernetes (GKE) operations.\n",
"# @markdown\n",
"# @markdown #### Actions:\n",
"# @markdown 1. **Connects to Project:** Retrieves and sets your Google Cloud project ID.\n",
"# @markdown 3. **Installs `kubectl`:** Installs the Kubernetes command-line tool.\n",
"\n",
"import os\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Set up gcloud.\n",
"! gcloud config set project \"$PROJECT_ID\"\n",
"! gcloud services enable container.googleapis.com\n",
"\n",
"# Add kubectl to the set of available tools.\n",
"! mkdir -p /tools/google-cloud-sdk/.install\n",
"! gcloud components install kubectl --quiet"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "IKGTaN84p8rX",
"metadata": {
"cellView": "form",
"id": "IKGTaN84p8rX"
},
"outputs": [],
"source": [
"# @title # Chat completion for text-only models {vertical-output: true}\n",
"# @markdown Run cell to prompt the model server for prediction.\n",
"# @markdown\n",
"# @markdown * **user_prompt (string):** This is the text prompt you provide to the language model. It's the question or instruction e (e.g., \"Explain neural networks\").\n",
"# @markdown * **temperature (number):** This parameter controls the randomness of the model's output. It influences how the model selects the next token in the sequence it generates. Typical values range from 0.2 to 1.0.\n",
"# @markdown * **max_tokens (number):** This parameter refers to the maximum number of tokens (words or sub-word units) that the model is allowed to generate in its response.\n",
"# @markdown\n",
"\n",
"import json\n",
"import subprocess\n",
"\n",
"import ipywidgets as widgets\n",
"from IPython.display import Markdown, clear_output, display\n",
"\n",
"CLUSTER = \"\" # @param {type:\"string\", isTemplate:true}\n",
"REGION = \"\" # @param {type:\"string\", isTemplate:true}\n",
"NAMESPACE = \"\" # @param {type:\"string\", isTemplate:true}\n",
"DEPLOYMENT = \"\" # @param {type:\"string\", isTemplate:true}\n",
"POD_PORT = \"\" # @param {type:\"string\", isTemplate:true}\n",
"\n",
"\n",
"def _run_kubectl(cmd, timeout=60):\n",
" \"\"\"Executes a kubectl command.\"\"\"\n",
" try:\n",
" result = subprocess.run(\n",
" cmd, capture_output=True, text=True, check=True, timeout=timeout\n",
" )\n",
" return result.stdout.strip()\n",
" except subprocess.CalledProcessError as e:\n",
" raise RuntimeError(\n",
" f\"Kubectl command failed: {' '.join(e.cmd)}\\nStderr: {e.stderr}\"\n",
" ) from e\n",
" except subprocess.TimeoutExpired as e:\n",
" raise RuntimeError(f\"Kubectl command timed out: {' '.join(e.cmd)}\") from e\n",
"\n",
"\n",
"def fetch_cluster_credentials(cluster, region, project_id):\n",
" \"\"\"Ensures credentials for the target GKE cluster.\"\"\"\n",
" cred_cmd = [\n",
" \"gcloud\",\n",
" \"container\",\n",
" \"clusters\",\n",
" \"get-credentials\",\n",
" cluster,\n",
" f\"--location={region}\",\n",
" f\"--project={project_id}\",\n",
" ]\n",
" _run_kubectl(cred_cmd)\n",
"\n",
"\n",
"def get_deployment_selector_labels(deployment_name, namespace):\n",
" \"\"\"Retrieves the selector labels for a given Kubernetes deployment.\"\"\"\n",
" cmd = [\n",
" \"kubectl\",\n",
" \"get\",\n",
" \"deployment\",\n",
" deployment_name,\n",
" \"-n\",\n",
" namespace,\n",
" \"-o\",\n",
" \"json\",\n",
" ]\n",
" deployment_json = _run_kubectl(cmd)\n",
" deployment_data = json.loads(deployment_json)\n",
"\n",
" selector_labels = (\n",
" deployment_data.get(\"spec\", {}).get(\"selector\", {}).get(\"matchLabels\")\n",
" )\n",
" if not selector_labels:\n",
" raise RuntimeError(\n",
" f\"No selector labels found for deployment '{deployment_name}' in\"\n",
" f\" namespace '{namespace}'.\"\n",
" )\n",
" return selector_labels\n",
"\n",
"\n",
"def get_running_pod_name(deployment_name, namespace):\n",
" \"\"\"Retrieves the name of a running pod associated with a deployment.\"\"\"\n",
" selector_labels = get_deployment_selector_labels(deployment_name, namespace)\n",
" label_selector_str = \",\".join(f\"{k}={v}\" for k, v in selector_labels.items())\n",
"\n",
" cmd = [\n",
" \"kubectl\",\n",
" \"get\",\n",
" \"pods\",\n",
" \"-n\",\n",
" namespace,\n",
" \"-o\",\n",
" \"json\",\n",
" \"-l\",\n",
" label_selector_str,\n",
" \"--field-selector=status.phase=Running\",\n",
" ]\n",
" pods_json = _run_kubectl(cmd)\n",
" pods_data = json.loads(pods_json)\n",
"\n",
" if not pods_data.get(\"items\"):\n",
" raise RuntimeError(\n",
" f\"No running pods found for deployment '{deployment_name}' in namespace\"\n",
" f\" '{namespace}' with selector '{label_selector_str}'.\"\n",
" )\n",
" return pods_data[\"items\"][0][\"metadata\"][\"name\"]\n",
"\n",
"\n",
"def check_vllm_inference_label(pod_name, namespace):\n",
" \"\"\"Checks if the specified pod has the vLLM inference server label.\"\"\"\n",
" cmd = [\"kubectl\", \"get\", \"pod\", pod_name, \"-n\", namespace, \"-o\", \"json\"]\n",
" pod_json = _run_kubectl(cmd)\n",
" labels = json.loads(pod_json).get(\"metadata\", {}).get(\"labels\", {})\n",
" return labels.get(\"ai.gke.io/inference-server\") == \"vllm\"\n",
"\n",
"\n",
"def send_inference_request(\n",
" request_payload, pod_name, pod_port, is_vllm_inference, namespace\n",
"):\n",
" \"\"\"Sends an inference request to the specified pod and returns the model's response.\"\"\"\n",
" json_data_escaped = json.dumps(request_payload).replace(\"'\", \"'\\\\''\")\n",
" curl_cmd = (\n",
" f\"kubectl exec -n {namespace} -t {pod_name} -- curl -s -X POST\"\n",
" f' http://localhost:{pod_port}/generate -H \"Content-Type:'\n",
" ' application/json\"'\n",
" f\" -d '{json_data_escaped}' 2> /dev/null\"\n",
" )\n",
"\n",
" response_raw = _run_kubectl([\"bash\", \"-c\", curl_cmd])\n",
"\n",
" if not response_raw:\n",
" raise RuntimeError(f\"Empty response received from pod '{pod_name}'.\")\n",
"\n",
" try:\n",
" first_line = response_raw.splitlines()[0]\n",
" data = json.loads(first_line)\n",
" except json.JSONDecodeError as e:\n",
" raise RuntimeError(\n",
" f\"Failed to decode JSON response from pod: {e}. Raw: {response_raw}\"\n",
" ) from e\n",
" except IndexError:\n",
" raise RuntimeError(\n",
" f\"Unexpected empty response line from pod. Raw: {response_raw}\"\n",
" )\n",
"\n",
" if is_vllm_inference:\n",
" predictions = data.get(\"predictions\")\n",
" if isinstance(predictions, list) and predictions:\n",
" return predictions[0]\n",
" raise RuntimeError(f\"Unexpected vLLM response format. Raw data: {data}\")\n",
" else: # TGI format\n",
" generated_text = data.get(\"generated_text\")\n",
" if generated_text is not None:\n",
" return generated_text\n",
" raise RuntimeError(f\"Unexpected TGI response format. Raw data: {data}\")\n",
"\n",
"\n",
"# --- Main Execution Logic ---\n",
"\n",
"\n",
"def execute_chat_completion(\n",
" deployment_name, namespace, pod_port, user_prompt, temperature, max_tokens\n",
"):\n",
" \"\"\"Executes the full chat completion process: fetches credentials, finds a pod,\n",
"\n",
" determines inference type, sends a request, and returns the response.\n",
" \"\"\"\n",
" display(Markdown(\"Establishing cluster credentials...\"))\n",
" fetch_cluster_credentials(CLUSTER, REGION, PROJECT_ID)\n",
"\n",
" display(Markdown(\"Retrieving pod information...\"))\n",
" pod_name = get_running_pod_name(deployment_name, namespace)\n",
" display(Markdown(f\"Successfully identified pod: `{pod_name}`\"))\n",
"\n",
" is_vllm = check_vllm_inference_label(pod_name, namespace)\n",
"\n",
" request_payload = {\n",
" \"max_tokens\": max_tokens,\n",
" \"temperature\": temperature,\n",
" \"prompt\" if is_vllm else \"inputs\": user_prompt,\n",
" }\n",
" display(Markdown(\"Sending inference request...\"))\n",
" response = send_inference_request(\n",
" request_payload, pod_name, pod_port, is_vllm, namespace\n",
" )\n",
"\n",
" return response\n",
"\n",
"\n",
"# --- Widgets Setup ---\n",
"user_prompt_widget = widgets.Textarea(\n",
" value=\"What is AI?\",\n",
" description=\"User Prompt:\",\n",
" layout=widgets.Layout(width=\"95%\", height=\"100px\"),\n",
")\n",
"\n",
"temperature_widget = widgets.FloatSlider(\n",
" value=0.50, min=0.0, max=1.0, step=0.01, description=\"Temperature:\"\n",
")\n",
"\n",
"max_tokens_widget = widgets.IntSlider(\n",
" value=250, min=1, max=2048, step=1, description=\"Max Tokens:\"\n",
")\n",
"\n",
"submit_button = widgets.Button(description=\"Submit\")\n",
"output_area_response = widgets.Output()\n",
"\n",
"\n",
"# --- Submit Button Logic ---\n",
"def on_submit_clicked(b):\n",
" with output_area_response:\n",
" clear_output()\n",
" display(Markdown(\"Loading...\"))\n",
"\n",
" try:\n",
" model_response = execute_chat_completion(\n",
" DEPLOYMENT,\n",
" NAMESPACE,\n",
" POD_PORT,\n",
" user_prompt_widget.value,\n",
" temperature_widget.value,\n",
" max_tokens_widget.value,\n",
" )\n",
" clear_output()\n",
" display(Markdown(f\"**Response:**\\n\\n{model_response}\"))\n",
" except Exception as e:\n",
" clear_output()\n",
" display(Markdown(f\"**An error occurred:**\\n```\\n{e}\\n```\"))\n",
"\n",
"\n",
"# --- Display Widgets ---\n",
"submit_button.on_click(on_submit_clicked)\n",
"display(\n",
" user_prompt_widget,\n",
" temperature_widget,\n",
" max_tokens_widget,\n",
" submit_button,\n",
" output_area_response,\n",
")"
]
},
{
"cell_type": "markdown",
"id": "5b6ZM2K3fux0",
"metadata": {
"id": "5b6ZM2K3fux0"
},
"source": [
"# Next Steps: Integrating the GKE Service Endpoint\n",
"\n",
"After successfully deploying a model on Google Kubernetes Engine (GKE) and\n",
"verifying it via a notebook, the next step is to integrate it into various\n",
"applications. This involves making HTTP requests to the service's endpoint from\n",
"your application code.\n",
"\n",
"### Exposing the Service\n",
"\n",
"To make your deployed model accessible to applications, you'll need to expose\n",
"its service endpoint. Google Kubernetes Engine offers several ways to do this:\n",
"\n",
"1. **Ingress:** Configure an Ingress resource to route external HTTP(S) traffic\n",
" to your service. Set up Ingress for either an internal Load Balancer\n",
" (accessible only within your VPC) or an external Load Balancer (accessible\n",
" from the internet).\n",
" [Learn more about GKE Ingress](https://cloud.google.com/kubernetes-engine/docs/concepts/ingress).\n",
"2. **Gateway API:** A more modern and feature-rich API for managing traffic\n",
" routing in Kubernetes. Similar to Ingress, Gateway API allows you to define\n",
" how external and internal traffic should be directed to your services.\n",
" [Explore GKE Gateway API](https://cloud.google.com/kubernetes-engine/docs/concepts/gateway-api).\n",
"\n",
"### Setting Up Autoscaling\n",
"\n",
"Ensure your model serving can handle varying traffic by configuring the\n",
"Horizontal Pod Autoscaler (HPA). HPA automatically scales the number of Pods\n",
"based on resource utilization or custom metrics, optimizing performance and\n",
"cost.\n",
"[See how to configure HPA](https://cloud.google.com/kubernetes-engine/docs/how-to/horizontal-pod-autoscaling).\n",
"\n",
"### Setting Up Monitoring\n",
"\n",
"Monitor the health and performance of your deployed model using Google Cloud\n",
"Managed Service for Prometheus. Configure your model serving to expose\n",
"Prometheus metrics for comprehensive insights.\n",
"[Get started with Google Cloud Managed Prometheus](https://cloud.google.com/kubernetes-engine/docs/how-to/configure-automatic-application-monitoring).\n",
"\n",
"### Additional Resources:\n",
"\n",
"* #### Kubernetes Documentation:\n",
"\n",
" * Services:\n",
" https://kubernetes.io/docs/concepts/services-networking/service/\n",
"\n",
"* #### Google Cloud Documentation:\n",
"\n",
" * Google Kubernetes Engine (GKE):\n",
" https://cloud.google.com/kubernetes-engine\n",
" * Cloud Load Balancing:\n",
" https://cloud.google.com/load-balancing/docs/ingress\n",
" * Gateway API on GKE:\n",
" https://cloud.google.com/kubernetes-engine/docs/concepts/gateway-api\n",
" * Learn about GPUs in GKE:\n",
" https://cloud.google.com/kubernetes-engine/docs/concepts/gpus\n",
"\n",
"* #### Python requests Library:\n",
"\n",
" * https://requests.readthedocs.io/en/latest/\n",
"\n",
"* #### LangChain with Google Integrations:\n",
"\n",
" * The Langchain documentation is very useful:\n",
" https://python.langchain.com/docs/integrations/providers/google/"
]
}
],
"metadata": {
"colab": {
"name": "gke_model_ui_deployment_notebook_auto.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -6,7 +6,7 @@
"id": "DZ1j6RRg-Td6",
"metadata": {
"cellView": "form",
"id": "f705f4be70e9"
"id": "DZ1j6RRg-Td6"
},
"outputs": [],
"source": [
@@ -29,7 +29,7 @@
"cell_type": "markdown",
"id": "99c1c3fc2ca5",
"metadata": {
"id": "71a642b5575a"
"id": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - Advanced Features\n",
@@ -42,7 +42,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_advanced_features.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -52,7 +52,7 @@
"cell_type": "markdown",
"id": "f9-tJ6RfDLIs",
"metadata": {
"id": "0779b48f654e"
"id": "f9-tJ6RfDLIs"
},
"source": [
"## Overview\n",
@@ -90,7 +90,7 @@
"cell_type": "markdown",
"id": "47GcOrZjosOx",
"metadata": {
"id": "69453bf7230e"
"id": "47GcOrZjosOx"
},
"source": [
"## Before you begin"
@@ -100,7 +100,7 @@
"cell_type": "markdown",
"id": "1D_pWejJPHP3",
"metadata": {
"id": "bf3706e69f61"
"id": "1D_pWejJPHP3"
},
"source": [
"### Request for quota\n",
@@ -118,9 +118,10 @@
"cell_type": "code",
"execution_count": null,
"id": "L3dqbxovo5t6",
"language": "python",
"metadata": {
"cellView": "form",
"id": "86a3d4d4d3f5"
"id": "L3dqbxovo5t6"
},
"outputs": [],
"source": [
@@ -148,8 +149,6 @@
"# Install and import the necessary packages\n",
"! pip install -q openai google-auth requests\n",
"\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.97.0'\n",
"\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"import datetime\n",
@@ -225,7 +224,7 @@
"cell_type": "markdown",
"id": "SeGqxuMfRBS5",
"metadata": {
"id": "4782dd003acb"
"id": "SeGqxuMfRBS5"
},
"source": [
"### Access Llama 3.1, 3.2, and 3.3 models on Vertex AI for serving"
@@ -237,7 +236,7 @@
"id": "BxlzWU2KQqmw",
"metadata": {
"cellView": "form",
"id": "798068fc0355"
"id": "BxlzWU2KQqmw"
},
"outputs": [],
"source": [
@@ -276,7 +275,7 @@
"cell_type": "markdown",
"id": "JpNBJJgjWL7j",
"metadata": {
"id": "10ed490e28e5"
"id": "JpNBJJgjWL7j"
},
"source": [
"## Prefix Caching <a name=\"prefix-caching\"></a>\n",
@@ -304,7 +303,7 @@
"cell_type": "markdown",
"id": "9gZJ8cB27e1m",
"metadata": {
"id": "30ddb93fdd7b"
"id": "9gZJ8cB27e1m"
},
"source": [
"### Try out Prefix Caching with Hex-LLM\n",
@@ -320,7 +319,7 @@
"id": "RpmoA2nXjdCd",
"metadata": {
"cellView": "form",
"id": "b56d82c1aa6f"
"id": "RpmoA2nXjdCd"
},
"outputs": [],
"source": [
@@ -509,7 +508,7 @@
"id": "5QoK8c0R9U3B",
"metadata": {
"cellView": "form",
"id": "96c5afed49b4"
"id": "5QoK8c0R9U3B"
},
"outputs": [],
"source": [
@@ -520,7 +519,9 @@
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"hexllm_tpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"hexllm_tpu\"].resource_name\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"hexllm_tpu\"].name\n",
")\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint using the OpenAI SDK.\n",
"\n",
@@ -575,7 +576,7 @@
"id": "29rn5ATmB2YC",
"metadata": {
"cellView": "form",
"id": "9a95c9f90358"
"id": "29rn5ATmB2YC"
},
"outputs": [],
"source": [
@@ -648,7 +649,7 @@
"cell_type": "markdown",
"id": "KjbM8E9DGuuR",
"metadata": {
"id": "12ad6d1ff725"
"id": "KjbM8E9DGuuR"
},
"source": [
"#### Delete the models and endpoints"
@@ -660,7 +661,7 @@
"id": "JpLU7GRQGuuR",
"metadata": {
"cellView": "form",
"id": "1ab4e3bb74b4"
"id": "JpLU7GRQGuuR"
},
"outputs": [],
"source": [
@@ -686,7 +687,7 @@
"cell_type": "markdown",
"id": "XZ33HhYmOxCS",
"metadata": {
"id": "7a8a9a1b2ddf"
"id": "XZ33HhYmOxCS"
},
"source": [
"### Try out Prefix Caching with vLLM\n",
@@ -709,7 +710,7 @@
"id": "E8OiHHNNE_wj",
"metadata": {
"cellView": "form",
"id": "4425cc0bdedc"
"id": "E8OiHHNNE_wj"
},
"outputs": [],
"source": [
@@ -913,7 +914,7 @@
"id": "zex1oXl36A70",
"metadata": {
"cellView": "form",
"id": "bcbafec839cd"
"id": "zex1oXl36A70"
},
"outputs": [],
"source": [
@@ -921,7 +922,9 @@
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"vllm_gpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"vllm_gpu\"].resource_name\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"vllm_gpu\"].name\n",
")\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint using the OpenAI SDK.\n",
"\n",
@@ -976,7 +979,7 @@
"id": "gDOC_nfsJeUR",
"metadata": {
"cellView": "form",
"id": "e984f43422d5"
"id": "gDOC_nfsJeUR"
},
"outputs": [],
"source": [
@@ -1049,7 +1052,7 @@
"cell_type": "markdown",
"id": "GdGxaTirJeUR",
"metadata": {
"id": "dff0d10dcc20"
"id": "GdGxaTirJeUR"
},
"source": [
"#### Delete the models and endpoints"
@@ -1061,7 +1064,7 @@
"id": "OgoqXE-VJeUR",
"metadata": {
"cellView": "form",
"id": "5b8751773e7f"
"id": "OgoqXE-VJeUR"
},
"outputs": [],
"source": [
@@ -1087,7 +1090,7 @@
"cell_type": "markdown",
"id": "w4Guijaw_NEs",
"metadata": {
"id": "863775857a46"
"id": "w4Guijaw_NEs"
},
"source": [
"### Best practices\n",
@@ -1102,7 +1105,7 @@
"cell_type": "markdown",
"id": "ml8fgoIQWSbY",
"metadata": {
"id": "565cbdc3a06b"
"id": "ml8fgoIQWSbY"
},
"source": [
"## Speculative Decoding <a name=\"spec-decoding\"></a>\n",
@@ -1144,7 +1147,7 @@
"cell_type": "markdown",
"id": "NmWRro8Q-Td6",
"metadata": {
"id": "94eaa9050abb"
"id": "NmWRro8Q-Td6"
},
"source": [
"### Try out Speculative Decoding with vLLM"
@@ -1156,7 +1159,7 @@
"id": "72d1GlrYifKU",
"metadata": {
"cellView": "form",
"id": "5f358cc230a6"
"id": "72d1GlrYifKU"
},
"outputs": [],
"source": [
@@ -1473,7 +1476,7 @@
"id": "CNiItf5hdVFU",
"metadata": {
"cellView": "form",
"id": "be3170e0e05a"
"id": "CNiItf5hdVFU"
},
"outputs": [],
"source": [
@@ -1499,7 +1502,9 @@
" DEDICATED_ENDPOINT_DNS = endpoints[\n",
" \"vllm_gpu_spec\"\n",
" ].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"vllm_gpu_spec\"].resource_name\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"vllm_gpu_spec\"].name\n",
")\n",
"\n",
"BASE_URL = (\n",
" f\"https://{REGION}-aiplatform.googleapis.com/v1beta1/{ENDPOINT_RESOURCE_NAME}\"\n",
@@ -1548,7 +1553,7 @@
"cell_type": "markdown",
"id": "WahYGAZyq6Gl",
"metadata": {
"id": "30c5d2535df3"
"id": "WahYGAZyq6Gl"
},
"source": [
"## Clean up resources"
@@ -1558,7 +1563,7 @@
"cell_type": "markdown",
"id": "bV5Yjkgav9BZ",
"metadata": {
"id": "63c10917ff95"
"id": "bV5Yjkgav9BZ"
},
"source": [
"### Delete the models and endpoints"
@@ -1570,7 +1575,7 @@
"id": "qsks36cOH9rb",
"metadata": {
"cellView": "form",
"id": "92892e1b1730"
"id": "qsks36cOH9rb"
},
"outputs": [],
"source": [
@@ -50,7 +50,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_agent_engine_llama3_1.ipynb\">\n",
" <img src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" alt=\"GitHub logo\"><br> View on GitHub\n",
" <img src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" alt=\"GitHub logo\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</table>"
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_axolotl_finetuning.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -55,47 +55,34 @@
"## Overview\n",
"This notebook demonstrates fine-tuning using [Axolotl](https://github.com/axolotl-ai-cloud/axolotl). Axolotl streamlines AI model fine-tuning by providing a wide range of training recipes and supporting multiple configurations and architectures.\n",
"\n",
"The notebook shows two modes of running the fine-tuning:\n",
"- Local fine-tuning using the Enterprise Colab runtime.\n",
"- Fine-tuning on cloud using the Vertex AI training.\n",
"\n",
"Both of these modes can be run independent of each other and the local fine-tuning is optional.\n",
"\n",
"Local fine-tuning with the Enterprise Colab runtime has the following advantages:\n",
"- **Debugging**: Use Enterprise Colab runtime to debug axolotl fine-tuning. This can be more efficient because debugging on the Vertex AI training involves waiting for resources to be provisioned, which can add delays. Also it is easier to debug on Enterprise Colab runtime compared to Vertex AI training.\n",
"We can use either Enterprise Colab runtime or Vertex AI training for fine-tuning using axolotl.\n",
"Colab runtime has below advantages:\n",
"- **Sanity check for flags**: Use Enterprise Colab runtime to do sanity check for Axolotl flags before running it on Vertex AI training directly.\n",
"- **Quick experimentations**: Use Enterprise Colab runtime to do quick experimentations with Axolotl flags.\n",
"- **Quick experimentations**: Use Enterprise Colab runtime to do quick experimentations with Axolotl flags. The [max-steps](https://github.com/axolotl-ai-cloud/axolotl/blob/8fb72cbc0b94129141bae5fa4d84edd23b648af6/docs/config.qmd#L360) flag is useful to limit the training time.\n",
"- **Debugging**: Use Enterprise Colab runtime to debug axolotl fine-tuning. This can be more efficient because debugging on the Vertex AI training involves waiting for resources to be provisioned, which can add delays. Also it is easier to add debug statements on Enterprise Colab runtime compared to Vertex AI training.\n",
"\n",
"Once the local fine-tuning is verified, the Vertex AI training is the recommended way to run the fine-tuning. Vertex AI training has several advantages, including:\n",
"- **Running multiple training jobs in parallel**: This can be useful for hyperparameter tuning or running experiments with different datasets etc.\n",
"- **Availability of Higher-end GPUs**: Vertex AI training provides access to higher-end GPUs like the H100, which can be crucial if you encounter out-of-memory (OOM) errors.\n",
"- **For High End GPU**: Vertex AI training provides access to higher-end GPUs like the H100, which can be crucial if you encounter out-of-memory (OOM) errors.\n",
"- **[DWS support](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws)**: DWS makes Vertex AI training more cost-effective, and easier to manage, especially in scenarios where GPU availability is a concern.\n",
"Refer to [this documentation](https://cloud.google.com/vertex-ai/docs/training/overview#vertexi-ai-operationalizes-training-at-scale) for more details on Vertex AI training advantages.\n",
"\n",
"### Objective\n",
"- Train model using Axolotl in local Enterprise Colab runtime.\n",
"- Run local prediction with Enterprise Colab runtime for the trained model.\n",
"- Train model using Axolotl with Vertex AI Training.\n",
"- Deploy the trained model on Vertex AI and run predictions on cloud.\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook."
"- Train model using Axolotl with Vertex AI Training."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "Xq8JgAE4BQTj"
"id": "K-YsE6oUoxjY"
},
"source": [
"## [Optional] Setup Colab Runtime\n",
"**You need to setup the Colab Runtime with GPU if you want to run local finetuning. The following sections perform the setup for A100 GPU.**\n",
"## Setup Colab Runtime\n",
"**You need to setup the Colab Runtime with L4 GPU or A100 GPU if you want to run local finetuning. The following sections perform the setup for L4 GPU.**\n",
"To learn more about creating runtime, you can optionally read [this](https://cloud.google.com/colab/docs/create-runtime).\n",
"\n",
"**Note: make sure to create a runtime with appropriate machine type and gpu type to avoid out of memory issues. [Refer this](https://huggingface.co/spaces/hf-accelerate/model-memory-usage) to decide which machine type and gpu type to select.**\n",
"\n",
"**Note: We recommend using a runtime environment configured with NVIDIA_TESLA_A100 with 4 GPUs or any other multi-GPU machine with higher GPU memory.**"
"**Note: make sure to create a runtime with appropriate machine type and gpu type to avoid out of memory issues. [Refer this](https://huggingface.co/spaces/hf-accelerate/model-memory-usage) to decide which machine type and gpu type to select.**"
]
},
{
@@ -103,85 +90,71 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "ggq03Jf5nNFp"
"id": "LGO6kScKoxjY"
},
"outputs": [],
"source": [
"# @title Create runtime\n",
"# @markdown This cell creates a local GPU runtime.\n",
"# @markdown **If you have already created a runtime previously, then you can skip this cell.** Optionally, read [this](https://cloud.google.com/colab/docs/create-runtime) to learn how to manually create a runtime.\n",
"# @markdown This cell creates a runtime template and then creates a runtime using that template.\n",
"# @markdown **If you have already created a runtime, you can skip this cell.**\n",
"# @markdown This cell can take up to 5 minutes to run.\n",
"# @markdown After the cell execution finishes, you have to connect manually to the runtime by following [the instructions here](https://cloud.google.com/colab/docs/connect-to-runtime).\n",
"\n",
"import os\n",
"import uuid\n",
"import re\n",
"import subprocess\n",
"\n",
"RUNTIME_PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"RUNTIME_REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"RUNTIME_ACCELERATOR_TYPE = \"NVIDIA_TESLA_A100\" # @param {type:\"string\"}\n",
"RUNTIME_ACCELERATOR_COUNT = \"4\" # @param [1, 2, 4, 8, 16]\n",
"RUNTIME_ACCELERATOR_TYPE = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"NVIDIA_TESLA_A100\", \"NVIDIA_A100_80GB\"]\n",
"RUNTIME_ACCELERATOR_COUNT = \"1\" # @param [1, 2, 4, 8, 16]\n",
"RUNTIME_ACCELERATOR_COUNT = int(RUNTIME_ACCELERATOR_COUNT)\n",
"\n",
"if not RUNTIME_ACCELERATOR_TYPE:\n",
" print(\"Warning: No accelerator type specified. Skipping runtime creation.\")\n",
" try:\n",
" subprocess.check_output(\"nvidia-smi\")\n",
" print(\"Nvidia GPU detected!\")\n",
" except Exception:\n",
" print(\"Warning: Nvidia GPU not detected. Use GPU runtime for local fine-tuning.\")\n",
"else:\n",
" if RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_TESLA_A100\" and RUNTIME_ACCELERATOR_COUNT != 16:\n",
" RUNTIME_MACHINE_TYPE = f\"a2-highgpu-{RUNTIME_ACCELERATOR_COUNT}g\"\n",
" elif RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_TESLA_A100\" and RUNTIME_ACCELERATOR_COUNT == 16:\n",
" RUNTIME_MACHINE_TYPE = \"a2-megagpu-16g\"\n",
" else:\n",
" raise ValueError(f\"Invalid GPU type {RUNTIME_ACCELERATOR_TYPE}, and count {RUNTIME_ACCELERATOR_COUNT} combination.\")\n",
"if RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_L4\" and RUNTIME_ACCELERATOR_COUNT == 1:\n",
" RUNTIME_MACHINE_TYPE = \"g2-standard-8\"\n",
"elif RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_L4\" and RUNTIME_ACCELERATOR_COUNT == 2:\n",
" RUNTIME_MACHINE_TYPE = \"g2-standard-24\"\n",
"elif RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_L4\" and RUNTIME_ACCELERATOR_COUNT == 4:\n",
" RUNTIME_MACHINE_TYPE = \"g2-standard-48\"\n",
"elif RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_L4\" and RUNTIME_ACCELERATOR_COUNT == 8:\n",
" RUNTIME_MACHINE_TYPE = \"g2-standard-96\"\n",
"elif RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_TESLA_A100\" and RUNTIME_ACCELERATOR_COUNT != 16:\n",
" RUNTIME_MACHINE_TYPE = f\"a2-highgpu-{RUNTIME_ACCELERATOR_COUNT}g\"\n",
"elif RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_TESLA_A100\" and RUNTIME_ACCELERATOR_COUNT == 16:\n",
" RUNTIME_MACHINE_TYPE = \"a2-megagpu-16g\"\n",
"elif RUNTIME_ACCELERATOR_TYPE == \"NVIDIA_A100_80GB\":\n",
" assert RUNTIME_ACCELERATOR_COUNT in [1, 2, 4, 8], \"Only 1, 2, 4, 8 A100-80GB are supported.\"\n",
" RUNTIME_MACHINE_TYPE = f\"a2-ultragpu-{RUNTIME_ACCELERATOR_COUNT}g\"\n",
"\n",
" print(f\"Machine type: {RUNTIME_MACHINE_TYPE}\")\n",
"uuid = uuid.uuid4()\n",
"RUNTIME_DISPLAY_NAME = f\"axolotl-{RUNTIME_ACCELERATOR_TYPE}-{RUNTIME_ACCELERATOR_COUNT}-{uuid}\"\n",
"\n",
" uuid = uuid.uuid4()\n",
" RUNTIME_DISPLAY_NAME = f\"axolotl-{RUNTIME_ACCELERATOR_TYPE}-{RUNTIME_ACCELERATOR_COUNT}-{uuid}\"\n",
" print(f\"Creating runtime with display name: {RUNTIME_DISPLAY_NAME}\")\n",
"# create runtime template\n",
"shell_output = ! gcloud colab runtime-templates create --display-name=$RUNTIME_DISPLAY_NAME \\\n",
" --project=$RUNTIME_PROJECT_ID --region=$RUNTIME_REGION \\\n",
" --machine-type=$RUNTIME_MACHINE_TYPE --accelerator-type=$RUNTIME_ACCELERATOR_TYPE \\\n",
" --accelerator-count=$RUNTIME_ACCELERATOR_COUNT --disk-type=PD_BALANCED\n",
"shell_output = \"\\n\".join(shell_output)\n",
"print(shell_output)\n",
"RUNTIME_TEMPLATE_ID = re.search(r\"projects/.*/locations/.*/notebookRuntimeTemplates/(\\d+)\", shell_output).group(1)\n",
"\n",
" # create runtime template\n",
" shell_output = ! gcloud colab runtime-templates create --display-name=$RUNTIME_DISPLAY_NAME \\\n",
" --project=$RUNTIME_PROJECT_ID --region=$RUNTIME_REGION \\\n",
" --machine-type=$RUNTIME_MACHINE_TYPE --accelerator-type=$RUNTIME_ACCELERATOR_TYPE \\\n",
" --accelerator-count=$RUNTIME_ACCELERATOR_COUNT --disk-type=PD_BALANCED\n",
" shell_output = \"\\n\".join(shell_output)\n",
" print(shell_output)\n",
" RUNTIME_TEMPLATE_ID = re.search(r\"projects/.*/locations/.*/notebookRuntimeTemplates/(\\d+)\", shell_output).group(1)\n",
"# create runtime\n",
"shell_output = ! gcloud colab runtimes create --display-name=$RUNTIME_DISPLAY_NAME \\\n",
" --runtime-template=$RUNTIME_TEMPLATE_ID --project=$RUNTIME_PROJECT_ID \\\n",
" --region=$RUNTIME_REGION\n",
"shell_output = \"\\n\".join(shell_output)\n",
"print(shell_output)\n",
"RUNTIME_ID = re.search(r\"projects/.*/locations/.*/notebookRuntimes/(\\d+)\", shell_output).group(1)\n",
"\n",
" # create runtime\n",
" shell_output = ! gcloud colab runtimes create --display-name=$RUNTIME_DISPLAY_NAME \\\n",
" --runtime-template=$RUNTIME_TEMPLATE_ID --project=$RUNTIME_PROJECT_ID \\\n",
" --region=$RUNTIME_REGION\n",
" shell_output = \"\\n\".join(shell_output)\n",
" print(shell_output)\n",
" RUNTIME_ID = re.search(r\"projects/.*/locations/.*/notebookRuntimes/(\\d+)\", shell_output).group(1)\n",
"\n",
" # start runtime\n",
" ! gcloud colab runtimes start $RUNTIME_ID --project=$RUNTIME_PROJECT_ID --region=$RUNTIME_REGION\n",
"\n",
" print(f\"Runtime: {RUNTIME_ID} created successfully.\")"
"# start runtime\n",
"! gcloud colab runtimes start $RUNTIME_ID --project=$RUNTIME_PROJECT_ID --region=$RUNTIME_REGION"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "uGDpB0NnnNFp"
},
"source": [
"### Connect to runtime manually\n",
"Although the previous step created a runtime, you still have to **connect manually to this runtime by following [the instructions here](https://cloud.google.com/colab/docs/connect-to-runtime).** This also applies if you want to use a previously created GPU runtime."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "kik-beBAnNFq"
"id": "_Gq2bkOMsYFJ"
},
"source": [
"## Before you begin"
@@ -192,33 +165,12 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "XeJdEWe9nNFq"
},
"outputs": [],
"source": [
"# @title [Optional for Vertex AI fine-tuning] Setup Pytorch for local fine-tuning\n",
"# @markdown **Note: This section must be run before performing local fine-tuning.** This section installs correct pytorch dependency needed for the Axolotl local run.\n",
"import os\n",
"\n",
"os.environ[\"TORCH_CUDA_ARCH_LIST\"] = \"7.0 7.5 8.0 8.6 9.0+PTX\"\n",
"! pip install torch==2.4.1 torchvision\n",
"os.environ[\"PYTORCH_INSTALLATION\"] = \"done\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "_cAAzli5nNFq"
"id": "emRjTZVEsYFJ"
},
"outputs": [],
"source": [
"# @title Import utility packages for fine-tuning\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"\n",
"# Import the necessary packages.\n",
"! rm -rf vertex-ai-samples && git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples\n",
@@ -229,7 +181,6 @@
"import importlib\n",
"import os\n",
"import pathlib\n",
"import time\n",
"import uuid\n",
"from typing import Tuple\n",
"\n",
@@ -241,38 +192,8 @@
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"\n",
"def run_cmd_and_check_output(\n",
" cmd: list[str], env: dict[str, str] = None, input: str = \"\", cwd: str = None\n",
"):\n",
" \"\"\"Runs the given command and raises exception if the command fails.\"\"\"\n",
" with subprocess.Popen(\n",
" cmd,\n",
" stdin=subprocess.PIPE,\n",
" stdout=subprocess.PIPE,\n",
" stderr=subprocess.STDOUT,\n",
" text=True,\n",
" bufsize=1,\n",
" env=env,\n",
" cwd=cwd,\n",
" ) as p:\n",
" if input:\n",
" p.stdin.write(input)\n",
" p.stdin.flush()\n",
" p.stdin.close()\n",
" for line in p.stdout:\n",
" print(line, end=\"\", flush=True)\n",
" if p.returncode:\n",
" raise ValueError(\n",
" f\"Command '{' '.join(cmd)}' execution failed with return code {p.returncode}\"\n",
" )\n",
"\n",
"\n",
"train_job = None\n",
"models, endpoints = {}, {}\n",
"HF_TOKEN = \"\"\n",
"WORKING_DIR = os.getcwd()\n",
"print(f\"Current working directory for notebook: {WORKING_DIR}\")"
"models, endpoints = {}, {}"
]
},
{
@@ -280,7 +201,7 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "k32BrMnWnNFq"
"id": "E0LS8jpwyUFu"
},
"outputs": [],
"source": [
@@ -288,9 +209,9 @@
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. For finetuning using Vertex AI, we will use Dynamic Workload Scheduler. Learn more about Dynamic workload scheduler [here](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws). For Dynamic Workload Scheduler, check the [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) or [europe-west4](https://console.cloud.google.com/iam-admin/quotas?location=europe-west4&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) quota for Nvidia H100 GPUs, [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_a100_80gb_gpus) quota for Nvidia A100 80GB GPUs. If you do not have enough GPUs, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request quota.\n",
"# @markdown 2. For finetuning, follow [these instructions](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws) to use Dynamic Workload Scheduler. For Dynamic Workload Scheduler, check the [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) or [europe-west4](https://console.cloud.google.com/iam-admin/quotas?location=europe-west4&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) quota for Nvidia H100 GPUs, and [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_a100_gpus) quota for Nvidia Tesla A100 GPUs. To train using L4 gpus with default quota, check [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_l4_gpus) quota for Nvidia L4 GPUs. If you do not have enough GPUs, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request quota.\n",
"\n",
"# @markdown 3. For serving, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_l4_gpus)** to check if your project already has the required L4 GPUs in the us-central1 region. If you need more L4 GPUs for your project, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request more. Alternatively, if you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"# @markdown 3. For serving, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_l4_gpus)** to check if your project already has the required 1 L4 GPU in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you need more L4 GPUs for your project, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request more. Alternatively, if you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
@@ -372,7 +293,7 @@
{
"cell_type": "markdown",
"metadata": {
"id": "g7AfA9UsnNFq"
"id": "FcVDnCvbFZUM"
},
"source": [
"## Finetune with Axolotl"
@@ -383,13 +304,13 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "-oKISrDMnNFr"
"id": "K8uSeXw9f1xs"
},
"outputs": [],
"source": [
"# @title Set Axolotl config\n",
"\n",
"# @markdown You can use below axolotl configs taken from [examples directory](https://github.com/axolotl-ai-cloud/axolotl/tree/c7d07de6b47b1b11d2098589e4bb15c6ed1066c3/examples), which have been verified by model garden team through internal testing. Note that we have used A100 80GB and H100 80GB GPU for testing.\n",
"# @markdown You can use below axolotl configs taken from [examples directory](https://github.com/axolotl-ai-cloud/axolotl/tree/8fb72cbc0b94129141bae5fa4d84edd23b648af6/examples), which have been verified by model garden team through internal testing. Note that we have used A100 80GB and H100 80GB GPU for testing.\n",
"# @markdown > | Model Name | Base Model | Axolotl Config |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | code-llama | codellama/CodeLlama-7b-hf | examples/code-llama/7b/lora.yml |\n",
@@ -401,60 +322,49 @@
"# @markdown | falcon | tiiuae/falcon-7b | examples/falcon/config-7b-lora.yml |\n",
"# @markdown | falcon | tiiuae/falcon-7b | examples/falcon/config-7b.yml |\n",
"# @markdown | gemma | google/gemma-7b | examples/gemma/qlora.yml |\n",
"# @markdown | gemma2 | google/gemma-7b | examples/gemma2/qlora.yml |\n",
"# @markdown | gemma3 | google/gemma-3-1b-it | examples/gemma3/gemma-3-1b-qlora.yml |\n",
"# @markdown | gemma3 | google/gemma-3-4b-it | examples/gemma3/gemma-3-4b-qlora.yml |\n",
"# @markdown | llama-2 | NousResearch/Llama-2-7b-hf | examples/llama-2/fft_optimized.yml |\n",
"# @markdown | llama-2 | NousResearch/Llama-2-7b-hf | examples/llama-2/loftq.yml |\n",
"# @markdown | llama-2 | NousResearch/Llama-2-7b-hf | examples/llama-2/lora.yml |\n",
"# @markdown | llama-2 | NousResearch/Llama-2-7b-hf | examples/llama-2/qlora-fsdp.yml |\n",
"# @markdown | llama-2 | NousResearch/Llama-2-7b-hf | examples/llama-2/qlora.yml |\n",
"# @markdown | llama-3 | hugging-quants/Meta-Llama-3.1-405B-BNB-NF4-BF16 | examples/llama-3/qlora-fsdp-405b.yaml |\n",
"# @markdown | llama-3 | NousResearch/Llama-3.2-1B | examples/llama-3/lora-1b-kernels.yml |\n",
"# @markdown | llama-3 | NousResearch/Meta-Llama-3.1-8B | examples/llama-3/fft-8b.yaml |\n",
"# @markdown | llama-3 | NousResearch/Meta-Llama-3-8B-Instruct | examples/llama-3/instruct-lora-8b.yml |\n",
"# @markdown | llama-3 | NousResearch/Meta-Llama-3-8B | examples/llama-3/lora-8b.yml |\n",
"# @markdown | llama-3 | NousResearch/Llama-3.2-1B | examples/llama-3/qlora-1b.yml |\n",
"# @markdown | llama-3 | NousResearch/Llama-3.2-1B | examples/llama-3/lora-1b.yml |\n",
"# @markdown | llama-3 | casperhansen/llama-3-70b-fp16 | examples/llama-3/qlora-fsdp-70b.yaml |\n",
"# @markdown | llama-3 | meta-llama/Llama-3.2-1B | examples/llama-3/lora-1b-deduplicate-dpo.yml |\n",
"# @markdown | llama-3 | meta-llama/Llama-3.2-1B | examples/llama-3/lora-1b-deduplicate-sft.yml |\n",
"# @markdown | llama-3 | meta-llama/Llama-3.2-1B | examples/llama-3/lora-1b-sample-packing-sequentially.yml |\n",
"# @markdown | mistral | mistralai/Mistral-7B-v0.1 | examples/mistral/config.yml |\n",
"# @markdown | mistral | mistralai/Mistral-7B-v0.1 | examples/mistral/lora-mps.yml |\n",
"# @markdown | mistral | mistralai/Mistral-7B-v0.1 | examples/mistral/lora.yml |\n",
"# @markdown | mistral | mistralai/Mistral-7B-v0.1 | examples/mistral/mistral-qlora-orpo.yml |\n",
"# @markdown | mistral | mistralai/Mistral-7B-Instruct-v0.2 | examples/mistral/mistral-dpo-qlora.yml |\n",
"# @markdown | mistral | mistral-community/Mixtral-8x22B-v0.1 | examples/mistral/mixtral-8x22b-qlora-fsdp.yml |\n",
"# @markdown | mistral | mistralai/Mixtral-8x7B-v0.1 | examples/mistral/mixtral.yml |\n",
"# @markdown | mistral | mistralai/Mixtral-8x7B-v0.1 | examples/mistral/mixtral-qlora-fsdp.yml |\n",
"# @markdown | mistral | mistralai/Mistral-7B-v0.1 | examples/mistral/qlora.yml |\n",
"# @markdown | openllama-3b | openlm-research/open_llama_3b_v2 | examples/openllama-3b/config.yml |\n",
"# @markdown | openllama-3b | openlm-research/open_llama_3b_v2 | examples/openllama-3b/lora.yml |\n",
"# @markdown | openllama-3b | openlm-research/open_llama_3b_v2 | examples/openllama-3b/qlora.yml |\n",
"# @markdown | phi | microsoft/Phi-3.5-mini-instruct | examples/phi/lora-3.5.yaml |\n",
"# @markdown | phi | microsoft/phi-1_5 | examples/phi/phi-ft.yml |\n",
"# @markdown | phi | microsoft/phi-1_5 | examples/phi/phi-qlora.yml |\n",
"# @markdown | phi | microsoft/phi-2 | examples/phi/phi2-ft.yml |\n",
"# @markdown | phi | microsoft/Phi-3-mini-4k-instruct | examples/phi/phi3-ft.yml |\n",
"# @markdown | qwen | Qwen/Qwen1.5-MoE-A2.7B | examples/qwen/qwen2-moe-lora.yaml |\n",
"# @markdown | qwen | Qwen/Qwen1.5-MoE-A2.7B | examples/qwen/qwen2-moe-qlora.yaml |\n",
"# @markdown | qwen2 | Qwen/Qwen2.5-0.5B | examples/qwen2/dpo.yaml |\n",
"# @markdown | qwen2 | Qwen/Qwen2.5-3B | examples/qwen2/prm.yaml |\n",
"# @markdown | qwen2 | Qwen/Qwen2-7B | examples/qwen2/qlora-fsdp.yaml |\n",
"# @markdown | tiny-llama | TinyLlama/TinyLlama_v1.1 | examples/tiny-llama/lora-mps.yml |\n",
"# @markdown | tiny-llama | TinyLlama/TinyLlama_v1.1 | examples/tiny-llama/lora.yml |\n",
"# @markdown | tiny-llama | TinyLlama/TinyLlama-1.1B-Chat-v1.0 | examples/tiny-llama/pretrain.yml |\n",
"# @markdown | tiny-llama | TinyLlama/TinyLlama_v1.1 | examples/tiny-llama/qlora.yml |\n",
"\n",
"# @markdown You can also customize the Axolotl config as per your requirements. To use a custom Axolotl config you can use `LOCAL` or `GCS` source option below.\n",
"# @markdown Alternatively, you can specify github axolotl config and override flags using `Setup Axolotl Flags` section below.\n",
"\n",
"# @markdown 1. Set Axolotl config source.<br>\n",
"# @markdown For **GITHUB** as source, you can explore different Axolotl configurations in the [examples directory](https://github.com/axolotl-ai-cloud/axolotl/tree/6ba5c0ed2c42a0e069b28c83646ee5a2a6904430/examples). For `GITHUB` source, `AXOLOTL_CONFIG_PATH` should start with `examples/`. e.g. \"examples/tiny-llama/qlora.yml\".<br>\n",
"# @markdown For **LOCAL** as source, create Axolotl config yaml file and specify correct path below. Note that, the local file will be copied to GCS bucket before running Vertex AI training job. For `LOCAL` source, `AXOLOTL_CONFIG_PATH` should be a absolute path of the config file, e.g. /content/lora.yml.<br>\n",
"# @markdown For **GCS** as source, specify the GCS URI to the Axolotl config file. Make sure the file is accessible to service account used in the notebook. For `GCS` source, `AXOLOTL_CONFIG_PATH` should be a complete GCS URI of the config file, e.g. gs://bucket/path/to/config/file.yml.\n",
"# @markdown For `GITHUB` as source, you can explore different Axolotl configurations in the [examples directory](https://github.com/axolotl-ai-cloud/axolotl/tree/8fb72cbc0b94129141bae5fa4d84edd23b648af6/examples). For `GITHUB` source, `AXOLOTL_CONFIG_PATH` should start with `examples/`. e.g. examples/tiny-llama/lora.yml.<br>\n",
"# @markdown For `LOCAL` as source, create Axolotl config yaml file and specify correct path below. Note that, the local file will be copied to GCS bucket before running Vertex AI training job. For `LOCAL` source, `AXOLOTL_CONFIG_PATH` should be a complete path of the config file. e.g. /content/lora.yml.<br>\n",
"# @markdown For `GCS` as source, specify the GCS URI to the Axolotl config file. Make sure the file is accessible to service account used in the notebook. For `GCS` source, `AXOLOTL_CONFIG_PATH` should be a complete GCS URI of the config file. e.g. gs://bucket/path/to/config/file.yml.\n",
"\n",
"AXOLOTL_SOURCE = \"GITHUB\" # @param [\"GITHUB\", \"LOCAL\", \"GCS\"]\n",
"\n",
"# @markdown 2. Set the Axolotl config file path.\n",
"AXOLOTL_CONFIG_PATH = \"examples/tiny-llama/qlora.yml\" # @param [\"examples/tiny-llama/qlora.yml\"] {allow-input: true}\n",
"AXOLOTL_CONFIG_PATH = \"examples/tiny-llama/lora.yml\" # @param {type:\"string\"}\n",
"\n",
"assert AXOLOTL_CONFIG_PATH, \"AXOLOTL_CONFIG_PATH must be set.\"\n",
"\n",
@@ -462,7 +372,7 @@
" assert AXOLOTL_CONFIG_PATH.startswith(\n",
" \"examples/\"\n",
" ), \"AXOLOTL_CONFIG_PATH must start with examples/ for GITHUB source.\"\n",
" github_url = f\"https://github.com/axolotl-ai-cloud/axolotl/raw/6ba5c0ed2c42a0e069b28c83646ee5a2a6904430/{AXOLOTL_CONFIG_PATH}\"\n",
" github_url = f\"https://github.com/axolotl-ai-cloud/axolotl/raw/8fb72cbc0b94129141bae5fa4d84edd23b648af6/{AXOLOTL_CONFIG_PATH}\"\n",
" r = requests.get(github_url)\n",
" axolotl_config = r.content.decode(\"utf-8\")\n",
" axolotl_config = yaml.safe_load(axolotl_config)\n",
@@ -472,7 +382,7 @@
" file_content = config_path.read_text()\n",
" axolotl_config = yaml.safe_load(file_content)\n",
"elif AXOLOTL_SOURCE == \"GCS\":\n",
" local_path = pathlib.Path(f\"{WORKING_DIR}/tmp/axolotl_config.yml\")\n",
" local_path = pathlib.Path(\"/content/tmp/axolotl_config.yml\")\n",
" common_util.download_gcs_file_to_local(AXOLOTL_CONFIG_PATH, local_path.absolute())\n",
" file_content = local_path.read_text()\n",
" axolotl_config = yaml.safe_load(file_content)\n",
@@ -483,17 +393,7 @@
"OUTPUT_GCS_URI = MODEL_BUCKET\n",
"\n",
"if not OUTPUT_GCS_URI.startswith(\"gs://\"):\n",
" OUTPUT_GCS_URI = f\"gs://{OUTPUT_GCS_URI}\"\n",
"\n",
"output_sub_dir = (\n",
" AXOLOTL_CONFIG_PATH.replace(\"/\", \"_\").replace(\".yaml\", \"\").replace(\".yml\", \"\")\n",
")\n",
"BASE_AXOLOTL_OUTPUT_GCS_URI = f\"{OUTPUT_GCS_URI}/{output_sub_dir}/axolotl_output\"\n",
"BASE_AXOLOTL_OUTPUT_DIR = common_util.gcs_fuse_path(BASE_AXOLOTL_OUTPUT_GCS_URI)\n",
"\n",
"# Placeholders for dataset settings.\n",
"datasets = []\n",
"test_datasets = []"
" OUTPUT_GCS_URI = f\"gs://{OUTPUT_GCS_URI}\""
]
},
{
@@ -501,11 +401,12 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "_iTQ4nSlODZJ"
"id": "dySqVhK8cFoO"
},
"outputs": [],
"source": [
"# @title Setup HF token\n",
"# @title **[Optional]** Setup HF token\n",
"# @markdown Some models like Gemma2, Mistral, Llama3 etc require a token to access with [gated access from huggingface](https://huggingface.co/docs/hub/en/models-gated).\n",
"HF_TOKEN = \"\" # @param {type:\"string\"}"
]
},
@@ -514,15 +415,13 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "4bp13RSoODZJ"
"id": "Z3xVT_VtFZUM"
},
"outputs": [],
"source": [
"# @title **[Optional]** Setup dataset\n",
"\n",
"# @markdown This section configures the dataset used for fine-tuning.\n",
"\n",
"# @markdown **Note: If you don't fill any of the dataset options given below, then the dataset used will be the one defined in the Axolotl config file.** You have two options to configure the dataset:\n",
"# @markdown This section configures the dataset used for fine-tuning. **Note: If you don't fill any of the dataset options given below, then the dataset used will be the one defined in the Axolotl config file.** You have two options to configure the dataset:\n",
"\n",
"# @markdown **1. Use a Hugging Face Dataset**\n",
"# @markdown - Requires specifying the dataset name and type.\n",
@@ -537,7 +436,7 @@
"\n",
"# @markdown **Hugging Face Dataset Name:**\n",
"HF_DATASET = \"\" # @param {type:\"string\", placeholder: \"e.g. timdettmers/openassistant-guanaco\"}\n",
"# @markdown **Set the dataset type:** Refer to [Axolotl config file](https://github.com/axolotl-ai-cloud/axolotl/blob/6ba5c0ed2c42a0e069b28c83646ee5a2a6904430/docs/config.qmd#L102) for more details.\n",
"# @markdown **Set the dataset type:** Refer to [Axolotl config file](https://github.com/axolotl-ai-cloud/axolotl/blob/8fb72cbc0b94129141bae5fa4d84edd23b648af6/docs/config.qmd#L87) for more details.\n",
"HF_DATASET_TYPE = \"\" # @param {type:\"string\", placeholder: \"e.g. completion\"}\n",
"if HF_DATASET:\n",
" assert HF_DATASET_TYPE, \"HF_DATASET_TYPE must be set if HF_DATASET is set.\"\n",
@@ -547,9 +446,9 @@
"\n",
"# @markdown **Bucket Name:**\n",
"DATASET_BUCKET_NAME = \"\" # @param {type:\"string\"}\n",
"# @markdown **Dataset Type:** Refer to the [Axolotl config file](https://github.com/axolotl-ai-cloud/axolotl/blob/6ba5c0ed2c42a0e069b28c83646ee5a2a6904430/docs/config.qmd#L102) for more details.\n",
"# @markdown **Dataset Type:** Refer to the [Axolotl config file](https://github.com/axolotl-ai-cloud/axolotl/blob/8fb72cbc0b94129141bae5fa4d84edd23b648af6/docs/config.qmd#L181) for more details.\n",
"DATASET_TYPE = \"\" # @param {type:\"string\"}\n",
"# @markdown **File Type**. Refer to the [Axolotl config file](https://github.com/axolotl-ai-cloud/axolotl/blob/6ba5c0ed2c42a0e069b28c83646ee5a2a6904430/docs/config.qmd#L103).\n",
"# @markdown **File Type**. Refer to the [Axolotl config file](https://github.com/axolotl-ai-cloud/axolotl/blob/8fb72cbc0b94129141bae5fa4d84edd23b648af6/docs/config.qmd#L178).\n",
"FILE_TYPE = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown **Path to Training Data (relative to bucket):**\n",
@@ -569,6 +468,7 @@
" HF_DATASET and DATASET_BUCKET_NAME\n",
"), \"Only one of HF_DATASET or DATASET_BUCKET_NAME can be set.\"\n",
"\n",
"datasets = []\n",
"if DATASET_BUCKET_NAME:\n",
" paths = TRAIN_DATAFILES_PATH.split(\",\")\n",
" dataset = {\n",
@@ -584,6 +484,7 @@
" dataset[\"split\"] = \"train\"\n",
" datasets.append(dataset)\n",
"\n",
"test_datasets = []\n",
"if TEST_DATAFILES_PATH:\n",
" paths = TEST_DATAFILES_PATH.split(\",\")\n",
" dataset = {\n",
@@ -608,49 +509,26 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "3SOjc9q8ODZJ"
"id": "7dj8WuXRWGn8"
},
"outputs": [],
"source": [
"# @title Setup Axolotl Flags\n",
"# @markdown This section configures additional Axolotl flags. You can explore different Axolotl flags in the [Axolotl config file](https://github.com/axolotl-ai-cloud/axolotl/blob/6ba5c0ed2c42a0e069b28c83646ee5a2a6904430/docs/config.qmd).\n",
"# @markdown This section configures additional Axolotl flags. You can explore different Axolotl flags in the [Axolotl config file](https://github.com/axolotl-ai-cloud/axolotl/blob/8fb72cbc0b94129141bae5fa4d84edd23b648af6/docs/config.qmd).\n",
"\n",
"# @markdown **To avoid OOM, you can reduce sequence length.** This can be done by setting `sequence_len` flag to some smaller value. But reducing sequence length might also reduce the fine-tuned model's quality.\n",
"# @markdown **Another alternative to avoid OOM is to use higher memory GPU.** It is recommended to use Vertex AI training for higher memory GPUs like A100 and H100. Vertex AI training offers greater availability of high-end GPUs.\n",
"# @markdown **To avoid OOM, you can reduce sequence length.** This can be done by setting `sequence_len` flag to some smaller value. But reducing sequence length will also reduce the model performance.\n",
"# @markdown **Another alternative to avoid OOM is to use higher memory gpu.** It is recommended to use vertex ai training for Higher memory gpu like A100 and H100. Vertex AI training offers greater availability of high-end GPUs.\n",
"\n",
"# @markdown **Training can take a long time (20+ hours) to complete depending on the model, dataset and axololt config.** You can reduce the training time by reducing the max training steps. This can be done by setting `max_steps` flag to some smaller value. Note that, this might also reduce the fine-tuned model's quality.\n",
"\n",
"# @markdown If you want to override base model then you can use `base_model` flag.\n",
"\n",
"# @markdown For example, let's say you want to log results to tensorboard and also want to use Qwen/Qwen3-32B model, then you can set [\"--use-tensorboard=True\", \"--base_model=Qwen/Qwen3-32B\"] in below `axolotl_flag_overrides` to achieve that.\n",
"# @markdown **Training can take a long time (20+ hours) to complete depending on the model, dataset and axololt config.** You can reduce the training time by reducing the max training steps. This can be done by setting `max_steps` flag to some smaller value. Note that this will also reduce the model performance.\n",
"\n",
"axolotl_flag_overrides = [\"--use-tensorboard=True\"] # @param {type:\"raw\"}\n",
"assert type(axolotl_flag_overrides) is list, \"axolotl_flag_overrides must be a list.\"\n",
"\n",
"# Set model_id and publisher. This is required for Vertex AI fine-tuning job and Vertex AI model deployment.\n",
"\n",
"\n",
"# Check if duplicate flags are passed.\n",
"flags_seen = set()\n",
"for flag in axolotl_flag_overrides:\n",
" if flag in flags_seen:\n",
" raise ValueError(f\"Duplicate flag: {flag}\")\n",
" flags_seen.add(flag)\n",
"\n",
"base_model = axolotl_config[\"base_model\"]\n",
"for overrides in axolotl_flag_overrides:\n",
" if overrides.startswith(\"--base_model=\"):\n",
" base_model = overrides.split(\"=\")[1]\n",
" break\n",
"publisher = base_model.split(\"/\")[0]\n",
"model_id = base_model.split(\"/\")[1]\n",
"model_id = model_id.replace(\".\", \"-\")"
"assert type(axolotl_flag_overrides) is list, \"axolotl_flag_overrides must be a list.\""
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "NrZ8eQJKODZJ"
"id": "uLMDYHuPJhOg"
},
"source": [
"### Finetune with Local Run"
@@ -661,36 +539,21 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "aQ0v6wsyODZJ"
"id": "D8kZmun8Ov3t"
},
"outputs": [],
"source": [
"# @title Install Axolotl And gscfuse\n",
"# @markdown 1. Check machine type\n",
"import subprocess\n",
"\n",
"try:\n",
" subprocess.check_output(\"nvidia-smi\")\n",
" print(\"Nvidia GPU detected!\")\n",
"except Exception:\n",
" raise ValueError(\"Nvidia GPU not detected. Use GPU runtime for local fine-tuning.\")\n",
"\n",
"# @markdown 2. Check if correct pytorch is installed.\n",
"if \"PYTORCH_INSTALLATION\" not in os.environ:\n",
" raise ValueError(\n",
" \"pytorch is not installed. Install it from `Setup Pytorch for local fine-tuning` section of the notebook.\"\n",
" )\n",
"\n",
"# @markdown 3. Install Axolotl\n",
"# @title Install Axolotl\n",
"! rm -rf axolotl\n",
"! git clone https://github.com/axolotl-ai-cloud/axolotl.git\n",
"! cd axolotl && git reset --hard 6ba5c0ed2c42a0e069b28c83646ee5a2a6904430\n",
"! cd axolotl && git reset --hard 8fb72cbc0b94129141bae5fa4d84edd23b648af6\n",
"! pip3 install packaging ninja\n",
"! cd axolotl && pip3 install --no-build-isolation -e '.[flash-attn,deepspeed,llmcompressor,ring-flash-attn,optimizers]'\n",
"! cd axolotl && python scripts/unsloth_install.py | sh\n",
"! cd axolotl && python scripts/cutcrossentropy_install.py | sh\n",
"! cd axolotl && pip3 install --no-build-isolation -e '.[flash-attn,deepspeed]'\n",
"\n",
"# @markdown 4. Install gscfuse\n",
"# This is needed because of this issue: https://github.com/bitsandbytes-foundation/bitsandbytes/issues/1492\n",
"! pip3 install bitsandbytes==0.45.1\n",
"\n",
"# @title Install GCSFUSE\n",
"! apt-get install gcsfuse -y"
]
},
@@ -699,33 +562,25 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "GREqaOn1ODZJ"
"id": "3tYJV5ScyscK"
},
"outputs": [],
"source": [
"# @title Run Local fine-tuning\n",
"# @markdown This section runs the Axolotl training locally (i.e. colab runtime).\n",
"# @markdown **Note: This section can take a long time to run. You can reduce the training time by reducing the max training steps as mentioned in `Setup Axolotl Flags` section.**\n",
"# @markdown Model trained using Axolotl will be saved in the GCS bucket with the help of gscfuse.\n",
"# @markdown Model trained using Axolotl will be saved in the GCS bucket with the help of GCSFUSE.\n",
"\n",
"# @markdown 1. Run gscfuse so that Axolotl can store the training output in the GCS bucket.\n",
"assert OUTPUT_GCS_URI, \"OUTPUT_GCS_URI must be set for local fine-tuning.\"\n",
"\n",
"# @markdown 1. Run GCSFUSE so that axolotl can store the training output in the GCS bucket.\n",
"! mkdir -p /gcs/\n",
"! gcsfuse /gcs\n",
"\n",
"# @markdown 2. Set up huggingface cache dir and access token.\n",
"os.environ[\"HF_HOME\"] = f\"{WORKING_DIR}/hf\"\n",
"os.environ[\"HF_TOKEN\"] = HF_TOKEN\n",
"\n",
"# @markdown 3. Run Axolotl training.\n",
"\n",
"local_config_path = AXOLOTL_CONFIG_PATH\n",
"if AXOLOTL_SOURCE == \"GITHUB\":\n",
" local_config_path = f\"{WORKING_DIR}/axolotl/{AXOLOTL_CONFIG_PATH}\"\n",
"finetuning_time = time.time_ns()\n",
"AXOLOTL_OUTPUT_GCS_URI = (\n",
" f\"{BASE_AXOLOTL_OUTPUT_GCS_URI}/local/time_ns_{finetuning_time}\"\n",
")\n",
"# @markdown 2. Run Axolotl training.\n",
"AXOLOTL_OUTPUT_GCS_URI = f\"{OUTPUT_GCS_URI}/axolotl_output\"\n",
"AXOLOTL_OUTPUT_DIR = common_util.gcs_fuse_path(AXOLOTL_OUTPUT_GCS_URI)\n",
"\n",
"axolotl_args = f\" --output-dir={AXOLOTL_OUTPUT_DIR}\"\n",
"if len(datasets) > 0:\n",
" axolotl_args += f' --datasets=\"{datasets}\"'\n",
@@ -734,9 +589,9 @@
" axolotl_args += \" --val-set-size=0\"\n",
"additional_flags = \" \".join(axolotl_flag_overrides)\n",
"axolotl_args += f\" {additional_flags}\"\n",
"! accelerate launch -m axolotl.cli.train $axolotl_args $local_config_path\n",
"! accelerate launch -m axolotl.cli.train $axolotl_args /content/axolotl/$AXOLOTL_CONFIG_PATH\n",
"\n",
"# @markdown 4. Check the output in the bucket.\n",
"# @markdown 3. Check the output in the bucket.\n",
"! gsutil ls $AXOLOTL_OUTPUT_GCS_URI"
]
},
@@ -745,40 +600,21 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "eD_66aBjODZJ"
"id": "HhcXl_JpyUFu"
},
"outputs": [],
"source": [
"# @title Run Local inference\n",
"# @markdown This section performs inference using the finetuned model.\n",
"# @markdown There are two options for inference:\n",
"# @markdown 1. Gradio: This option provides a URL for the playground to test the model.\n",
"# @markdown 2. CLI: This is option outputs the inference results in the console.\n",
"\n",
"INFERENCE_METHOD = \"gradio\" # @param [\"gradio\", \"cli\"]\n",
"# @markdown 1. Copy the finetuned model from GCS to local.\n",
"! mkdir -p /tmp/axolotl_output\n",
"! gsutil -m cp -r $AXOLOTL_OUTPUT_GCS_URI/* /tmp/axolotl_output/\n",
"\n",
"# @markdown **Note: `CLI_PROMPT` will be only used if `INFERENCE_METHOD` is `cli`.**\n",
"CLI_PROMPT = \"What is car?\" # @param {type:\"string\"}\n",
"# @markdown 2. Run Axolotl inference using gradio on local finetuned model.\n",
"! cd axolotl && axolotl inference examples/tiny-llama/lora.yml --output-dir=/tmp/axolotl_output/ --gradio\n",
"\n",
"if INFERENCE_METHOD == \"gradio\":\n",
" ! cd axolotl && export CUDA_VISIBLE_DEVICES=0 && axolotl inference --base-model=$base_model $local_config_path --lora-model-dir=$AXOLOTL_OUTPUT_DIR --gradio\n",
"elif INFERENCE_METHOD == \"cli\":\n",
" assert CLI_PROMPT, \"CLI_PROMPT must be set if INFERENCE_METHOD is 'cli'.\"\n",
" env = os.environ.copy()\n",
" env[\"CUDA_VISIBLE_DEVICES\"] = \"0\"\n",
" cmd = [\n",
" \"axolotl\",\n",
" \"inference\",\n",
" local_config_path,\n",
" f\"--base-model={base_model}\",\n",
" f\"--lora-model-dir={AXOLOTL_OUTPUT_DIR}\",\n",
" ]\n",
" run_cmd_and_check_output(cmd, env, f\"{CLI_PROMPT}\\x04\", f\"{WORKING_DIR}/axolotl/\")\n",
"else:\n",
" raise ValueError(f\"Unsupported inference method: {INFERENCE_METHOD}\")\n",
"\n",
"\n",
"# @markdown For Gradio, after running the cell, a public URL ([\"https://*.gradio.live\"](#)) will appear in the cell output. The playground is available in a separate browser tab when you click the URL."
"# @markdown 3. After running the cell, a public URL ([\"https://*.gradio.live\"](#)) will appear in the cell output. The playground is available in a separate browser tab when you click the URL."
]
},
{
@@ -786,12 +622,12 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "VEGLvAlnODZK"
"id": "z3EvrUpxNbtN"
},
"outputs": [],
"source": [
"# @title Create merged model\n",
"# @markdown This section merges the finetuned adapter with the base model.\n",
"# @markdown **Note: This is only needed for lora and qlora. In case of full finetuning you can skip this cell.**\n",
"\n",
"if (\n",
" \"adapter\" in axolotl_config\n",
@@ -800,22 +636,21 @@
"):\n",
" raise ValueError(\"This cell is only needed for lora and qlora.\")\n",
"\n",
"# @markdown 1. Run Axolotl merge. **Note: Based on model size, this step can take 5-20 minutes to complete.**\n",
"cmd = [\n",
" \"python3\",\n",
" \"-m\",\n",
" \"axolotl.cli.merge_lora\",\n",
" f\"--base-model={base_model}\",\n",
" f\"--output-dir={AXOLOTL_OUTPUT_DIR}\",\n",
" local_config_path,\n",
"]\n",
"run_cmd_and_check_output(cmd, None, None, f\"{WORKING_DIR}/axolotl/\")"
"# @markdown 1. Copy the finetuned model from GCS to local.\n",
"! mkdir -p /tmp/axolotl_output\n",
"! gsutil -m cp -r $AXOLOTL_OUTPUT_GCS_URI/* /tmp/axolotl_output/\n",
"\n",
"# @markdown 2. Run Axolotl merge.\n",
"! cd axolotl && python3 -m axolotl.cli.merge_lora $AXOLOTL_CONFIG_PATH --output-dir=/tmp/axolotl_output/\n",
"\n",
"# @markdown 3. Copy the merged model to GCS.\n",
"! gsutil -m cp -r /tmp/axolotl_output/merged /* $AXOLOTL_OUTPUT_GCS_URI/merged/"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "1H7V8fEjODZK"
"id": "LcMgeq_0CYDJ"
},
"source": [
"### Finetune with Vertex AI Training"
@@ -826,7 +661,7 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "0XqH_Mvelhon"
"id": "ZHKMHpOGFZUM"
},
"outputs": [],
"source": [
@@ -839,28 +674,33 @@
" custom_job as gca_custom_job_compat\n",
"\n",
"# @markdown Acceletor type to use for training.\n",
"training_accelerator_type = \"NVIDIA_H100_80GB\" # @param [\"NVIDIA_H100_80GB\", \"NVIDIA_A100_80GB\"]\n",
"training_accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"NVIDIA_TESLA_A100\", \"NVIDIA_H100_80GB\"]\n",
"\n",
"\n",
"replica_count = 1\n",
"repo = \"us-docker.pkg.dev/vertex-ai\"\n",
"per_node_accelerator_count = 8\n",
"per_node_accelerator_count = 1\n",
"boot_disk_size_gb = 500\n",
"dws_kwargs = {\n",
" \"max_wait_duration\": 1800, # 30 minutes\n",
" \"scheduling_strategy\": gca_custom_job_compat.Scheduling.Strategy.FLEX_START,\n",
"}\n",
"is_dynamic_workload_scheduler = True\n",
"if training_accelerator_type == \"NVIDIA_A100_80GB\":\n",
" training_machine_type = \"a2-ultragpu-8g\"\n",
"if training_accelerator_type == \"NVIDIA_L4\":\n",
" training_machine_type = \"g2-standard-8\"\n",
" is_dynamic_workload_scheduler = False\n",
" dws_kwargs = {}\n",
"elif training_accelerator_type == \"NVIDIA_TESLA_A100\":\n",
" training_machine_type = \"a2-highgpu-1g\"\n",
"elif training_accelerator_type == \"NVIDIA_H100_80GB\":\n",
" training_machine_type = \"a3-highgpu-8g\"\n",
" per_node_accelerator_count = 8\n",
" boot_disk_size_gb = 2000\n",
"else:\n",
" raise ValueError(f\"Unsupported accelerator type: {training_accelerator_type}\")\n",
"\n",
"TRAIN_DOCKER_URI = (\n",
" f\"{repo}/vertex-vision-model-garden-dockers/axolotl-train-dws:20250515-1800-rc0\"\n",
" f\"{repo}/vertex-vision-model-garden-dockers/axolotl-train:20250225-1800-rc0\"\n",
")\n",
"\n",
"common_util.check_quota(\n",
@@ -873,33 +713,73 @@
" is_dynamic_workload_scheduler=is_dynamic_workload_scheduler,\n",
")\n",
"\n",
"vertex_ai_config_path = AXOLOTL_CONFIG_PATH\n",
"# @markdown Run Vertex AI job.\n",
"\n",
"# Copy the config file to the bucket.\n",
"if AXOLOTL_SOURCE == \"LOCAL\":\n",
" ! gsutil -m cp $AXOLOTL_CONFIG_PATH $MODEL_BUCKET/config/\n",
" vertex_ai_config_path = f\"{common_util.gcs_fuse_path(MODEL_BUCKET)}/config/{pathlib.Path(AXOLOTL_CONFIG_PATH).name}\"\n",
" AXOLOTL_CONFIG_PATH = f\"{common_util.gcs_fuse_path(MODEL_BUCKET)}/config/{pathlib.Path(AXOLOTL_CONFIG_PATH).name}\"\n",
"\n",
"job_name = common_util.get_job_name_with_datetime(\"axolotl-train\")\n",
"AXOLOTL_OUTPUT_GCS_URI = f\"{BASE_AXOLOTL_OUTPUT_GCS_URI}/{job_name}\"\n",
"AXOLOTL_OUTPUT_DIR = f\"{BASE_AXOLOTL_OUTPUT_DIR}/{job_name}\"\n",
"# Set axolotl flags.\n",
"datasets = []\n",
"if DATASET_BUCKET_NAME:\n",
" paths = TRAIN_DATAFILES_PATH.split(\",\")\n",
" dataset = {\n",
" \"path\": f\"/gcs/{DATASET_BUCKET_NAME}/\",\n",
" \"type\": DATASET_TYPE,\n",
" \"data_files\": [],\n",
" \"ds_type\": FILE_TYPE,\n",
" }\n",
" for path in paths:\n",
" if path.startswith(\"/\"):\n",
" path = path[1:]\n",
" dataset[\"data_files\"].append(f\"/gcs/{DATASET_BUCKET_NAME}/{path}\")\n",
" dataset[\"split\"] = \"train\"\n",
" datasets.append(dataset)\n",
"\n",
"test_datasets = []\n",
"if TEST_DATAFILES_PATH:\n",
" paths = TEST_DATAFILES_PATH.split(\",\")\n",
" dataset = {\n",
" \"path\": f\"/gcs/{DATASET_BUCKET_NAME}/\",\n",
" \"type\": DATASET_TYPE,\n",
" \"data_files\": [],\n",
" \"ds_type\": FILE_TYPE,\n",
" }\n",
" for path in paths:\n",
" if path.startswith(\"/\"):\n",
" path = path[1:]\n",
" dataset[\"data_files\"].append(f\"/gcs/{DATASET_BUCKET_NAME}/{path}\")\n",
" dataset[\"split\"] = \"train\"\n",
" test_datasets.append(dataset)\n",
"\n",
"if HF_DATASET:\n",
" datasets.append({\"path\": HF_DATASET, \"type\": HF_DATASET_TYPE})\n",
"\n",
"if not OUTPUT_GCS_URI:\n",
" OUTPUT_GCS_URI = MODEL_BUCKET\n",
"AXOLOTL_OUTPUT_GCS_URI = f\"{OUTPUT_GCS_URI}/axolotl_output\"\n",
"AXOLOTL_OUTPUT_DIR = common_util.gcs_fuse_path(AXOLOTL_OUTPUT_GCS_URI)\n",
"TRAINING_JOB_OUTPUT_DIR = f\"{AXOLOTL_OUTPUT_GCS_URI}/training_job_output\"\n",
"\n",
"# Set Axolotl flags.\n",
"\n",
"axolotl_config_overwrites = []\n",
"axolotl_config_overwrites.append(f\"--output_dir={AXOLOTL_OUTPUT_DIR}\")\n",
"if len(datasets) > 0:\n",
" axolotl_config_overwrites.append(f\"--datasets={datasets}\")\n",
" axolotl_config_overwrites.append(f'--datasets=\"{datasets}\"')\n",
"if len(test_datasets) > 0:\n",
" axolotl_config_overwrites.append(f\"--test_datasets={test_datasets}\")\n",
" axolotl_config_overwrites.append(f'--test_datasets=\"{test_datasets}\"')\n",
" axolotl_config_overwrites.append(\"--val_set_size=0\")\n",
"axolotl_config_overwrites += axolotl_flag_overrides\n",
"\n",
"train_job_args = []\n",
"train_job_args.append(f\"--axolotl_config_path={vertex_ai_config_path}\")\n",
"train_job_args.append(f\"--axolotl_config_path={AXOLOTL_CONFIG_PATH}\")\n",
"train_job_args += axolotl_config_overwrites\n",
"\n",
"\n",
"train_job_envs = {}\n",
"if HF_TOKEN:\n",
" train_job_args.append(f\"--huggingface_access_token={HF_TOKEN}\")\n",
" train_job_envs[\"HF_TOKEN\"] = HF_TOKEN\n",
"\n",
"job_name = common_util.get_job_name_with_datetime(\"axolotl-train\")\n",
"\n",
@@ -910,11 +790,11 @@
"}\n",
"\n",
"model_name = AXOLOTL_CONFIG_PATH.split(\"/\")[1]\n",
"publisher = axolotl_config[\"base_model\"].split(\"/\")[0]\n",
"model_id = axolotl_config[\"base_model\"].split(\"/\")[1]\n",
"model_id = model_id.replace(\".\", \"-\")\n",
"labels[\"mg-tune\"] = f\"publishers-{publisher}-models-{model_name}\".lower()\n",
"labels[\"versioned-mg-tune\"] = f\"{labels['mg-tune']}-{model_id}\".lower()\n",
"labels[\"versioned-mg-tune\"] = labels[\"versioned-mg-tune\"][\n",
" : min(len(labels[\"versioned-mg-tune\"]), 63)\n",
"]\n",
"\n",
"\n",
"# Pass training arguments and launch job.\n",
@@ -924,7 +804,6 @@
" labels=labels,\n",
")\n",
"\n",
"# Run Vertex AI job.\n",
"print(\"Running training job with args:\")\n",
"print(\" \\\\\\n\".join(train_job_args))\n",
"train_job.run(\n",
@@ -962,6 +841,8 @@
},
"outputs": [],
"source": [
"base_output_dir = AXOLOTL_OUTPUT_DIR\n",
"\n",
"# @markdown This section shows how to launch TensorBoard in a [Cloud Shell](https://cloud.google.com/shell/docs).\n",
"# @markdown 1. Click the Cloud Shell icon(![terminal](https://github.com/google/material-design-icons/blob/master/png/action/terminal/materialicons/24dp/1x/baseline_terminal_black_24dp.png?raw=true)) on the top right to open the Cloud Shell.\n",
"# @markdown 2. Copy the `tensorboard` command shown below by running this cell.\n",
@@ -969,7 +850,7 @@
"# @markdown 4. Once the command runs (You may have to click `Authorize` if prompted), click the link starting with `http://localhost`.\n",
"\n",
"# @markdown Note: You may need to wait around 10 minutes after the job starts in order for the TensorBoard logs to be written to the GCS bucket.\n",
"print(f\"Command to copy: tensorboard --logdir {AXOLOTL_OUTPUT_GCS_URI}\")"
"print(f\"Command to copy: tensorboard --logdir {base_output_dir}/logs\")"
]
},
{
@@ -978,7 +859,7 @@
"id": "BO1BYNnfXy9_"
},
"source": [
"## Deploy using VLLM"
"## Deploy using vllm"
]
},
{
@@ -998,7 +879,7 @@
"\n",
"# @markdown 2. Set up VLLM docker URI and model gcs uri.\n",
"\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20250405_1205_RC01\"\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20241001_0916_RC00\"\n",
"VLLM_MODEL_GCS_URI = AXOLOTL_OUTPUT_GCS_URI\n",
"\n",
"if \"adapter\" in axolotl_config and (\n",
@@ -1040,25 +921,11 @@
" is_for_training=False,\n",
")\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"use_dedicated_endpoint = False\n",
"gpu_memory_utilization = 0.95\n",
"max_model_len = 2048\n",
"\n",
"\n",
"def get_deploy_source() -> str:\n",
" \"\"\"Gets deploy_source string based on running environment.\"\"\"\n",
" vertex_product = os.environ.get(\"VERTEX_PRODUCT\", \"\")\n",
" if vertex_product == \"COLAB_ENTERPRISE\":\n",
" return \"notebook_colab_enterprise\"\n",
" elif vertex_product == \"WORKBENCH_INSTANCE\":\n",
" return \"notebook_workbench\"\n",
" else:\n",
" # Legacy workbench, legacy colab, or other custom environments.\n",
" return \"notebook_environment_unspecified\"\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
" model_id: str,\n",
@@ -1178,7 +1045,7 @@
" service_account=service_account,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_axolotl_finetuning.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": get_deploy_source(),\n",
" \"NOTEBOOK_ENVIRONMENT\": common_util.get_deploy_source(),\n",
" },\n",
" )\n",
" print(\"endpoint_name:\", endpoint.name)\n",
@@ -1188,8 +1055,8 @@
"\n",
"models[\"vllm_gpu\"], endpoints[\"vllm_gpu\"] = deploy_model_vllm(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"axolotl-vllm-serve\"),\n",
" publisher=publisher.lower(),\n",
" publisher_model_id=model_id.lower(),\n",
" publisher=publisher,\n",
" publisher_model_id=model_id,\n",
" model_id=VLLM_MODEL_GCS_URI,\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=machine_type,\n",
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -38,7 +38,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_camp_zipnerf.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -38,7 +38,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_camp_zipnerf_gradio.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_codegemma_deployment_on_vertex.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -123,10 +123,8 @@
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# Dedicated endpoint not supported yet\n",
"use_dedicated_endpoint = False\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
@@ -221,7 +219,7 @@
"# @markdown *--- Or ---*\n",
"\n",
"# @markdown #### Access CodeGemma models on HuggingFace\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the CodeGemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"# @markdown You must provide a Hugging Face User Access Token (read) to access the CodeGemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"if LOAD_MODEL_FROM == \"Hugging Face\":\n",
" assert (\n",
@@ -339,7 +337,6 @@
" disagg_topology: str = None,\n",
" hbm_utilization_factor: float = 0.6,\n",
" max_running_seqs: int = 256,\n",
" decode_seqs_padding: int = None,\n",
" max_model_len: int = 4096,\n",
" enable_prefix_cache_hbm: bool = False,\n",
" endpoint_id: str = \"\",\n",
@@ -380,10 +377,6 @@
" f\"--max_running_seqs={max_running_seqs}\",\n",
" f\"--max_model_len={max_model_len}\",\n",
" ]\n",
"\n",
" if decode_seqs_padding is not None:\n",
" hexllm_args.append(f\"--decode_seqs_padding={decode_seqs_padding}\")\n",
"\n",
" if disagg_topology:\n",
" hexllm_args.append(f\"--disagg_topo={disagg_topology}\")\n",
" if enable_prefix_cache_hbm and not disagg_topology:\n",
@@ -34,7 +34,7 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_deployment_tutorial.ipynb\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/instances\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
@@ -45,7 +45,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_deployment_tutorial.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -107,7 +107,7 @@
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.84.0'\n",
"\n",
"# Import the necessary packages\n",
"import os\n",
@@ -166,7 +166,7 @@
"outputs": [],
"source": [
"# @title Choose the model to deploy\n",
"from vertexai import model_garden\n",
"from vertexai.preview import model_garden\n",
"\n",
"# @markdown List all deployable models and then get the ID of the model to deploy.\n",
"\n",
@@ -222,16 +222,14 @@
"# @title Deploy with Model Garden SDK\n",
"\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"from vertexai import model_garden\n",
"from vertexai.preview import model_garden\n",
"\n",
"model = model_garden.OpenModel(MODEL_ID)\n",
"endpoints[LABEL] = model.deploy(\n",
" hugging_face_access_token=HF_TOKEN,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]"
")"
]
},
{
@@ -314,7 +312,9 @@
" DEDICATED_ENDPOINT_DNS = endpoints[\n",
" \"my-endpoint\"\n",
" ].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"my-endpoint\"].resource_name\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"my-endpoint\"].name\n",
")\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
@@ -4,12 +4,11 @@
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "7d9bbf86da5e"
},
"outputs": [],
"source": [
"# Copyright 2025 Google LLC\n",
"# Copyright 2024 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",
@@ -34,18 +33,13 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_e5.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_e5.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_e5.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -73,10 +67,6 @@
"- Run inference on the deployed Vertex AI Endpoint\n",
"\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
@@ -84,7 +74,7 @@
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), [Cloud NL API pricing](https://cloud.google.com/natural-language/pricing) and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
@@ -109,61 +99,133 @@
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"# @markdown 2. [Optional] [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus). You can request for quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# Import the necessary packages\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"! pip3 install --quiet torchvision\n",
"\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from datetime import datetime\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"from torch import Tensor\n",
"\n",
"if os.environ.get(\"VERTEX_PRODUCT\") != \"COLAB_ENTERPRISE\":\n",
" ! pip install --upgrade tensorflow\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_URI = \"gs://\" # @param {type: \"string\"}\n",
"assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
"\n",
"# @markdown Click \"Show code\" to see more details.\n",
"\n",
"# Create a unique GCS bucket for this notebook, if not specified by the user.\n",
"assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"else:\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION)\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"import vertexai\n",
"! gcloud services enable language.googleapis.com\n",
"\n",
"vertexai.init(\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
")"
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"\n",
"# Gets the default BUCKET_URI and SERVICE_ACCOUNT if they were not specified by the user.\n",
"\n",
"SERVICE_ACCOUNT = None\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"\n",
"def create_name_with_datetime(prefix: str) -> str:\n",
" \"\"\"Creates a name with date time when triggering training or deployment\n",
" jobs in Vertex AI.\n",
" \"\"\"\n",
" return prefix + datetime.now().strftime(\"_%Y%m%d_%H%M%S\")\n",
"\n",
"\n",
"def deploy_model_tei(\n",
" model_name: str,\n",
" model_id: str,\n",
" service_account: str,\n",
" docker_uri: str,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" max_model_len: int = 512,\n",
" gpu_memory_utilization: float = 0.9,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys E5 models with TEI on Vertex AI.\n",
"\n",
" Args:\n",
" model_name: Display name of the model.\n",
" model_id: Model ID or path to model weights.\n",
" service_account: Service account for model uploading and deployment.\n",
" machine_type: Deployment machine type.\n",
" accelerator_type: Deployment accelerator type.\n",
" accelerator_count: Number of accelerators to use.\n",
" max_model_len: Maximum model length.\n",
" gpu_memory_utilization: Fraction of GPU memory to be used for the model\n",
" executor.\n",
"\n",
" Returns:\n",
" Model instance and endpoint instance.\n",
" \"\"\"\n",
" endpoint = aiplatform.Endpoint.create(display_name=f\"{model_name}-endpoint\")\n",
"\n",
" tei_args = [\n",
" f\"--model-id={model_id}\",\n",
" ]\n",
" serving_env = {\n",
" \"MODEL_ID\": model_id,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=docker_uri,\n",
" serving_container_args=tei_args,\n",
" serving_container_ports=[7080],\n",
" serving_container_environment_variables=serving_env,\n",
" model_garden_source_model_name=\"publishers/intfloat/models/e5\"\n",
" )\n",
"\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_e5.ipynb\"\n",
" },\n",
" )\n",
" return model, endpoint"
]
},
{
@@ -180,25 +242,19 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "I1u2FLa9XgVD"
"id": "kg5MwMIfB9Uj"
},
"outputs": [],
"source": [
"# @title Select the model variants\n",
"# @title Deploy\n",
"# @markdown This section uploads a prebuilt model to Model Registry and deploys it on the Endpoint. The model deployment step will take ~15 minutes to complete.\n",
"\n",
"prebuilt_model_id = \"intfloat/e5-small-v2\" # @param [\"intfloat/multilingual-e5-large-instruct\", \"intfloat/multilingual-e5-large\", \"intfloat/e5-large-v2\", \"intfloat/multilingual-e5-small\", \"intfloat/e5-base-v2\", \"intfloat/e5-small-v2\"]\n",
"\n",
"# @markdown Specify a processor for the TEI docker image. E5 models can be run on either GPU or CPU.\n",
"# Find Vertex AI prediction supported accelerators and regions in\n",
"# https://cloud.google.com/vertex-ai/docs/predictions/configure-compute.\n",
"processor = \"NVIDIA_L4\" # @param[\"NVIDIA_TESLA_T4\", \"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\", \"CPU\"]\n",
"processor = \"NVIDIA_L4\" # @param[\"NVIDIA_TESLA_V100\", \"NVIDIA_L4\", \"NVIDIA_TESLA_A100\", \"CPU\"]\n",
"\n",
"if processor == \"NVIDIA_TESLA_T4\":\n",
" accelerator_type = \"NVIDIA_TESLA_T4\"\n",
" machine_type = \"n1-highmem-16\"\n",
" accelerator_count = 1\n",
"elif processor == \"NVIDIA_TESLA_V100\":\n",
"if processor == \"NVIDIA_TESLA_V100\":\n",
" accelerator_type = \"NVIDIA_TESLA_V100\"\n",
" machine_type = \"n1-highmem-16\"\n",
" accelerator_count = 2\n",
@@ -224,108 +280,24 @@
"else:\n",
" TEI_DOCKER_URI = \"us-docker.pkg.dev/deeplearning-platform-release/gcr.io/huggingface-text-embeddings-inference-cu122.1-2.ubuntu2204\"\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"# @markdown Click \"Show code\" to see more details.\n",
"\n",
"# @markdown Click \"Show code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "6dY6_ppyObQy"
},
"outputs": [],
"source": [
"# @title Deploy model using custom configuration\n",
"# @markdown This section uploads prebuilt E5 models to Model Registry and deploys it to a Vertex AI Endpoint. It might take ~15 minutes to 1 hour to finish depending on the size of the model.\n",
"# Finds Vertex AI prediction supported accelerators and regions in\n",
"# https://cloud.google.com/vertex-ai/docs/predictions/configure-compute.\n",
"\n",
"\n",
"def deploy_model_tei(\n",
" model_name: str,\n",
" model_id: str,\n",
" docker_uri: str,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" max_model_len: int = 512,\n",
" gpu_memory_utilization: float = 0.9,\n",
" use_dedicated_endpoint: bool = True,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys E5 models with TEI on Vertex AI.\n",
"\n",
" Args:\n",
" model_name: Display name of the model.\n",
" model_id: Model ID or path to model weights.\n",
" machine_type: Deployment machine type.\n",
" accelerator_type: Deployment accelerator type.\n",
" accelerator_count: Number of accelerators to use.\n",
" max_model_len: Maximum model length.\n",
" gpu_memory_utilization: Fraction of GPU memory to be used for the model\n",
" executor.\n",
" use_dedicated_endpoint: A dedicated endpoint is an endpoint for online\n",
" prediction,provide a secure connection for private communication\n",
" between on-premises and Google Cloud.\n",
"\n",
" Returns:\n",
" Model instance and endpoint instance.\n",
" \"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
"\n",
" tei_args = [\n",
" f\"--model-id={model_id}\",\n",
" ]\n",
" serving_env = {\n",
" \"MODEL_ID\": model_id,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=docker_uri,\n",
" serving_container_args=tei_args,\n",
" serving_container_ports=[7080],\n",
" serving_container_environment_variables=serving_env,\n",
" model_garden_source_model_name=\"publishers/intfloat/models/e5\",\n",
" )\n",
"\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" system_labels={\"NOTEBOOK_NAME\": \"model_garden_e5.ipynb\"},\n",
" )\n",
" return model, endpoint\n",
"\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"\n",
"LABEL = \"tei\"\n",
"models[LABEL], endpoints[LABEL] = deploy_model_tei(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"e5-serve-tei\"),\n",
"model, endpoint = deploy_model_tei(\n",
" model_name=create_name_with_datetime(prefix=\"e5-serve-tei\"),\n",
" model_id=prebuilt_model_id,\n",
" service_account=SERVICE_ACCOUNT,\n",
" docker_uri=TEI_DOCKER_URI,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"model = models[LABEL]\n",
"endpoint = endpoints[LABEL]"
"print(\"endpoint_name:\", endpoint.name)\n",
"print(\"model_name:\", model.display_name)\n",
"print(\"model_id:\", model.resource_name)"
]
},
{
@@ -383,6 +355,7 @@
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"from torch import Tensor\n",
"\n",
"# Each input text should start with \"query: \" or \"passage: \".\n",
"# For tasks other than retrieval, you can simply use the \"query: \" prefix.\n",
@@ -396,9 +369,7 @@
" ],\n",
" },\n",
"]\n",
"response = endpoint.predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"response = endpoint.predict(instances=instances)\n",
"\n",
"for prediction in response.predictions:\n",
" embeddings = Tensor(prediction)\n",
@@ -426,6 +397,9 @@
"# @markdown Instruct: Given a web search query, retrieve relevant passages that answer the query\n",
"# @markdown Query: how much protein should a female eat\n",
"# @markdown Instruct: Given a web search query, retrieve relevant passages that answer the query\n",
"# @markdown Query: 南瓜的家常做法\n",
"# @markdown As a general guideline, the CDC's average requirement of protein for women ages 19 to 70 is 46 grams per day. But, as you can see from this chart, you'll need to increase that if you're expecting or training for a marathon. Check out the chart below to see how much protein you should be eating each day.\n",
"# @markdown 1.清炒南瓜丝 原料:嫩南瓜半个 调料:葱、盐、白糖、鸡精 做法: 1、南瓜用刀薄薄的削去表面一层皮,用勺子刮去瓤 2、擦成细丝(没有擦菜板就用刀慢慢切成细丝) 3、锅烧热放油,入葱花煸出香味 4、入南瓜丝快速翻炒一分钟左右,放盐、一点白糖和鸡精调味出锅 2.香葱炒南瓜 原料:南瓜1只 调料:香葱、蒜末、橄榄油、盐 做法: 1、将南瓜去皮,切成片 2、油锅8成热后,将蒜末放入爆香 3、爆香后,将南瓜片放入,翻炒 4、在翻炒的同时,可以不时地往锅里加水,但不要太多 5、放入盐,炒匀 6、南瓜差不多软和绵了之后,就可以关火 7、撒入香葱,即可出锅\n",
"# @markdown ```\n",
"\n",
"# @markdown API reference link to HuggingFace : [Text Embeddings Inference API](https://huggingface.github.io/text-embeddings-inference/#/).\n",
@@ -454,6 +428,8 @@
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"from torch import Tensor\n",
"\n",
"\n",
"def get_detailed_instruct(task_description: str, query: str) -> str:\n",
" return f\"Instruct: {task_description}\\nQuery: {query}\"\n",
@@ -472,9 +448,7 @@
"\n",
"instances = [{\"inputs\": queries + documents}]\n",
"\n",
"response = endpoint.predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"response = endpoint.predict(instances=instances)\n",
"\n",
"for prediction in response.predictions:\n",
" embeddings = Tensor(prediction)\n",
@@ -505,13 +479,12 @@
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"for endpoint in endpoints.values():\n",
" endpoint.delete(force=True)\n",
"endpoint.delete(force=True)\n",
"model.delete()\n",
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()"
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_URI"
]
}
],
@@ -32,18 +32,18 @@
"source": [
"# Vertex AI Model Garden - Finetuning Tutorial\n",
"\n",
"\u003ctable\u003e\u003ctbody\u003e\u003ctr\u003e\n",
" \u003ctd style=\"text-align: center\"\u003e\n",
" \u003ca href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_finetuning_tutorial.ipynb\"\u003e\n",
" \u003cimg alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"\u003e\u003cbr\u003e Run in Colab Enterprise\n",
" \u003c/a\u003e\n",
" \u003c/td\u003e\n",
" \u003ctd style=\"text-align: center\"\u003e\n",
" \u003ca href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_finetuning_tutorial.ipynb\"\u003e\n",
" \u003cimg alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"\u003e\u003cbr\u003e View on GitHub\n",
" \u003c/a\u003e\n",
" \u003c/td\u003e\n",
"\u003c/tr\u003e\u003c/tbody\u003e\u003c/table\u003e"
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_finetuning_tutorial.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_finetuning_tutorial.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
@@ -168,11 +168,11 @@
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. For finetuning, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1\u0026metric=aiplatform.googleapis.com%2Frestricted_image_training_nvidia_a100_80gb_gpus)** to check if your project already has the required 8 Nvidia A100 80 GB GPUs in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you do not have 8 Nvidia A100 80 GPUs or have more GPU requirements than this, then schedule your job with Nvidia H100 GPUs via Dynamic Workload Scheduler using [these instructions](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws). For Dynamic Workload Scheduler, check the [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1\u0026metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) or [europe-west4](https://console.cloud.google.com/iam-admin/quotas?location=europe-west4\u0026metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) quota for Nvidia H100 GPUs. If you do not have enough GPUs, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request quota.\n",
"# @markdown 2. For finetuning, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Frestricted_image_training_nvidia_a100_80gb_gpus)** to check if your project already has the required 8 Nvidia A100 80 GB GPUs in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you do not have 8 Nvidia A100 80 GPUs or have more GPU requirements than this, then schedule your job with Nvidia H100 GPUs via Dynamic Workload Scheduler using [these instructions](https://cloud.google.com/vertex-ai/docs/training/schedule-jobs-dws). For Dynamic Workload Scheduler, check the [us-central1](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) or [europe-west4](https://console.cloud.google.com/iam-admin/quotas?location=europe-west4&metric=aiplatform.googleapis.com%2Fcustom_model_training_preemptible_nvidia_h100_gpus) quota for Nvidia H100 GPUs. If you do not have enough GPUs, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request quota.\n",
"\n",
"# @markdown 3. For evaluation and deployment, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1\u0026metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_l4_gpus)** to check if your project already has the required 1 L4 GPU in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you need more L4 GPUs for your project, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request more. Alternatively, if you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"# @markdown 3. For evaluation and deployment, **[click here](https://console.cloud.google.com/iam-admin/quotas?location=us-central1&metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_l4_gpus)** to check if your project already has the required 1 L4 GPU in the us-central1 region. If yes, then run this notebook in the us-central1 region. If you need more L4 GPUs for your project, then you can follow [these instructions](https://cloud.google.com/docs/quotas/view-manage#viewing_your_quota_console) to request more. Alternatively, if you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"\n",
"# @markdown \u003e | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
@@ -188,8 +188,8 @@
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# Import the necessary packages.\n",
"! rm -rf vertex-ai-samples \u0026\u0026 git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples \u0026\u0026 git reset --hard 0727e19520cf7957bceb701c248221bd3dbe4f1f\n",
"! rm -rf vertex-ai-samples && git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples && git reset --hard 0727e19520cf7957bceb701c248221bd3dbe4f1f\n",
"\n",
"import datetime\n",
"import importlib\n",
@@ -269,10 +269,13 @@
]
},
{
"metadata": {
"id": "VUSi9jUcvdBC"
},
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "36c21f10355f"
},
"outputs": [],
"source": [
"# @title Access Llama 3.1 models\n",
"\n",
@@ -285,7 +288,7 @@
"base_model_id = \"meta-llama/Llama-3.1-8B-Instruct\" # @param {type: \"string\"}\n",
"pretrained_model_id = base_model_id\n",
"\n",
"# @markdown Additionally, you must provide a Hugging Face User Access Token (with read access) to access the Llama 3.1 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"# @markdown Additionally, you must provide a Hugging Face User Access Token (read) to access the Llama 3.1 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"\n",
@@ -293,9 +296,7 @@
" assert (\n",
" HF_TOKEN\n",
" ), \"Provide a read access HF_TOKEN to load models from Hugging Face, or select a different model source. You can comment out this assert statement to skip this check.\""
],
"outputs": [],
"execution_count": null
]
},
{
"cell_type": "markdown",
@@ -322,7 +323,7 @@
"| coqa | 0.1158 | 0.0137 | 0.1872 | 0.0150 |\n"
],
"text/plain": [
"\u003cIPython.core.display.Markdown object\u003e"
"<IPython.core.display.Markdown object>"
]
},
"metadata": {},
@@ -358,7 +359,7 @@
" lora_path: str = None,\n",
" max_num_seqs: int = 64,\n",
" eval_task: str = \"coqa\",\n",
") -\u003e str:\n",
") -> str:\n",
" \"\"\"Run lm-evaluation-harness to evaluate the model, and returns .\"\"\"\n",
"\n",
" if accelerator_type == \"NVIDIA_L4\":\n",
@@ -606,7 +607,7 @@
"version_minor": 0
},
"text/plain": [
"Downloading readme: 0%| | 0.00/8.20k [00:00\u003c?, ?B/s]"
"Downloading readme: 0%| | 0.00/8.20k [00:00<?, ?B/s]"
]
},
"metadata": {},
@@ -620,7 +621,7 @@
"version_minor": 0
},
"text/plain": [
"Downloading data: 0%| | 0.00/13.1M [00:00\u003c?, ?B/s]"
"Downloading data: 0%| | 0.00/13.1M [00:00<?, ?B/s]"
]
},
"metadata": {},
@@ -634,7 +635,7 @@
"version_minor": 0
},
"text/plain": [
"Generating train split: 0%| | 0/15011 [00:00\u003c?, ? examples/s]"
"Generating train split: 0%| | 0/15011 [00:00<?, ? examples/s]"
]
},
"metadata": {},
@@ -652,7 +653,7 @@
"| When was Tomoaki Komorida born? | Komorida was born in Kumamoto Prefecture on July 10, 1981. After graduating from high school, he joined the J1 League club Avispa Fukuoka in 2000. Although he debuted as a midfielder in 2001, he did not play much and the club was relegated to the J2 League at the end of the 2001 season. In 2002, he moved to the J2 club Oita Trinita. He became a regular player as a defensive midfielder and the club won the championship in 2002 and was promoted in 2003. He played many matches until 2005. In September 2005, he moved to the J2 club Montedio Yamagata. In 2006, he moved to the J2 club Vissel Kobe. Although he became a regular player as a defensive midfielder, his gradually was played less during the summer. In 2007, he moved to the Japan Football League club Rosso Kumamoto (later Roasso Kumamoto) based in his local region. He played as a regular player and the club was promoted to J2 in 2008. Although he did not play as much, he still played in many matches. In 2010, he moved to Indonesia and joined Persela Lamongan. In July 2010, he returned to Japan and joined the J2 club Giravanz Kitakyushu. He played often as a defensive midfielder and center back until 2012 when he retired. | Tomoaki Komorida was born on July 10,1981. | closed_qa |"
],
"text/plain": [
"\u003cIPython.core.display.Markdown object\u003e"
"<IPython.core.display.Markdown object>"
]
},
"metadata": {},
@@ -767,14 +768,14 @@
"name": "stdout",
"output_type": "stream",
"text": [
"('\u003c|begin_of_text|\u003e\u003c|start_header_id|\u003esystem\u003c|end_header_id|\u003e\\n'\n",
"('<|begin_of_text|><|start_header_id|>system<|end_header_id|>\\n'\n",
" '\\n'\n",
" 'You are a helpful '\n",
" 'assistant.\u003c|eot_id|\u003e\u003c|start_header_id|\u003euser\u003c|end_header_id|\u003e\\n'\n",
" 'assistant.<|eot_id|><|start_header_id|>user<|end_header_id|>\\n'\n",
" '\\n'\n",
" 'Hello, how are you?\u003c|eot_id|\u003e\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n'\n",
" 'Hello, how are you?<|eot_id|><|start_header_id|>assistant<|end_header_id|>\\n'\n",
" '\\n'\n",
" 'I am doing well, thank you.\u003c|eot_id|\u003e')\n"
" 'I am doing well, thank you.<|eot_id|>')\n"
]
}
],
@@ -793,7 +794,7 @@
"# @markdown to translate the Jinja template for you. For example, you can ask:\n",
"\n",
"# @markdown ````\n",
"# @markdown Translate the following Jinja template to Python, where bos_token is \"\u003c|begin_of_text|\u003e\":\n",
"# @markdown Translate the following Jinja template to Python, where bos_token is \"<|begin_of_text|>\":\n",
"# @markdown ```\n",
"# @markdown {{- bos_token }}\n",
"# @markdown {#- This block extracts the system message, so we can slot it into the right place. #}\n",
@@ -804,14 +805,14 @@
"# @markdown {%- set system_message = \"\" %}\n",
"# @markdown {%- endif %}\n",
"# @markdown {#- System message #}\n",
"# @markdown {{- \"\u003c|start_header_id|\u003esystem\u003c|end_header_id|\u003e\\n\\n\" }}\n",
"# @markdown {{- \"<|start_header_id|>system<|end_header_id|>\\n\\n\" }}\n",
"# @markdown {{- system_message }}\n",
"# @markdown {{- \"\u003c|eot_id|\u003e\" }}\n",
"# @markdown {{- \"<|eot_id|>\" }}\n",
"# @markdown {%- for message in messages %}\n",
"# @markdown {{- '\u003c|start_header_id|\u003e' + message['role'] + '\u003c|end_header_id|\u003e\\n\\n'+ message['content'] | trim + '\u003c|eot_id|\u003e' }}\n",
"# @markdown {{- '<|start_header_id|>' + message['role'] + '<|end_header_id|>\\n\\n'+ message['content'] | trim + '<|eot_id|>' }}\n",
"# @markdown {%- endfor %}\n",
"# @markdown {%- if add_generation_prompt %}\n",
"# @markdown {{- '\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n' }}\n",
"# @markdown {{- '<|start_header_id|>assistant<|end_header_id|>\\n\\n' }}\n",
"# @markdown {%- endif %}\n",
"# @markdown ```\n",
"# @markdown ````\n",
@@ -820,7 +821,7 @@
"# @markdown `bos_token` is specified in the [tokenizer_config.json](https://huggingface.co/meta-llama/Llama-3.1-8B-Instruct/blob/main/tokenizer_config.json#L2052).\n",
"# @markdown ```\n",
"# @markdown def render_template(messages, add_generation_prompt=False):\n",
"# @markdown bos_token = \"\u003c|begin_of_text|\u003e\"\n",
"# @markdown bos_token = \"<|begin_of_text|>\"\n",
"# @markdown output = bos_token\n",
"# @markdown\n",
"# @markdown system_message = \"\"\n",
@@ -828,22 +829,22 @@
"# @markdown system_message = messages[0]['content'].strip()\n",
"# @markdown messages = messages[1:]\n",
"# @markdown\n",
"# @markdown output += \"\u003c|start_header_id|\u003esystem\u003c|end_header_id|\u003e\\n\\n\"\n",
"# @markdown output += \"<|start_header_id|>system<|end_header_id|>\\n\\n\"\n",
"# @markdown output += system_message\n",
"# @markdown output += \"\u003c|eot_id|\u003e\"\n",
"# @markdown output += \"<|eot_id|>\"\n",
"# @markdown\n",
"# @markdown for message in messages:\n",
"# @markdown output += f\"\u003c|start_header_id|\u003e{message['role']}\u003c|end_header_id|\u003e\\n\\n{message['content'].strip()}\u003c|eot_id|\u003e\"\n",
"# @markdown output += f\"<|start_header_id|>{message['role']}<|end_header_id|>\\n\\n{message['content'].strip()}<|eot_id|>\"\n",
"# @markdown\n",
"# @markdown if add_generation_prompt:\n",
"# @markdown output += \"\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n\"\n",
"# @markdown output += \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"# @markdown\n",
"# @markdown return output\n",
"# @markdown ```\n",
"\n",
"\n",
"def render_template(messages, add_generation_prompt=False):\n",
" bos_token = \"\u003c|begin_of_text|\u003e\"\n",
" bos_token = \"<|begin_of_text|>\"\n",
" output = bos_token\n",
"\n",
" system_message = \"\"\n",
@@ -851,26 +852,26 @@
" system_message = messages[0][\"content\"].strip()\n",
" messages = messages[1:]\n",
"\n",
" output += \"\u003c|start_header_id|\u003esystem\u003c|end_header_id|\u003e\\n\\n\"\n",
" output += \"<|start_header_id|>system<|end_header_id|>\\n\\n\"\n",
" output += system_message\n",
" output += \"\u003c|eot_id|\u003e\"\n",
" output += \"<|eot_id|>\"\n",
"\n",
" for message in messages:\n",
" output += (\n",
" \"\u003c|start_header_id|\u003e\"\n",
" \"<|start_header_id|>\"\n",
" + message[\"role\"]\n",
" + \"\u003c|end_header_id|\u003e\\n\\n\"\n",
" + \"<|end_header_id|>\\n\\n\"\n",
" + message[\"content\"].strip()\n",
" + \"\u003c|eot_id|\u003e\"\n",
" + \"<|eot_id|>\"\n",
" )\n",
"\n",
" if add_generation_prompt:\n",
" output += \"\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n\"\n",
" output += \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"\n",
" return output\n",
"\n",
"\n",
"# @markdown The `\u003c|begin_of_text|\u003e`, `\u003c|start_header_id|\u003e`, `\u003c|end_header_id|\u003e`, and `\u003c|eot_id|\u003e` tokens serve specific purposes in the context of text generation models.\n",
"# @markdown The `<|begin_of_text|>`, `<|start_header_id|>`, `<|end_header_id|>`, and `<|eot_id|>` tokens serve specific purposes in the context of text generation models.\n",
"# @markdown These tokens help the model delineate the boundaries of a text generation task. They provide clear markers for the start and end points, enabling the model to function effectively and produce coherent text.\n",
"\n",
"# @markdown Run this cell to show an example output of the template given the\n",
@@ -1029,15 +1030,15 @@
"name": "stdout",
"output_type": "stream",
"text": [
"('\u003c|begin_of_text|\u003e\u003c|start_header_id|\u003euser\u003c|end_header_id|\u003e\\n'\n",
"('<|begin_of_text|><|start_header_id|>user<|end_header_id|>\\n'\n",
" '\\n'\n",
" 'Hello, how are you? Context: This is a test '\n",
" 'context.\u003c|eot_id|\u003e\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n'\n",
" 'context.<|eot_id|><|start_header_id|>assistant<|end_header_id|>\\n'\n",
" '\\n'\n",
" \"I'm doing well, thank \"\n",
" 'you!\u003c|eot_id|\u003e\u003c|start_header_id|\u003euser\u003c|end_header_id|\u003e\\n'\n",
" 'you!<|eot_id|><|start_header_id|>user<|end_header_id|>\\n'\n",
" '\\n'\n",
" 'Another question without context.\u003c|eot_id|\u003e')\n"
" 'Another question without context.<|eot_id|>')\n"
]
}
],
@@ -1051,15 +1052,15 @@
"# @markdown chat_template_string = r\"\"\"{{- bos_token }}\n",
"# @markdown\n",
"# @markdown {% for message in messages %}\n",
"# @markdown {{- '\u003c|start_header_id|\u003e' + message.role + '\u003c|end_header_id|\u003e\\n\\n' + message.content | trim }}\n",
"# @markdown {% if message.context and message.context | length \u003e 0 %}\n",
"# @markdown {{- '<|start_header_id|>' + message.role + '<|end_header_id|>\\n\\n' + message.content | trim }}\n",
"# @markdown {% if message.context and message.context | length > 0 %}\n",
"# @markdown {{- ' Context: ' + message.context }}\n",
"# @markdown {% endif %}\n",
"# @markdown {{- '\u003c|eot_id|\u003e' }}\n",
"# @markdown {{- '<|eot_id|>' }}\n",
"# @markdown {% endfor %}\n",
"# @markdown\n",
"# @markdown {% if add_generation_prompt %}\n",
"# @markdown {{- '\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n' }}\n",
"# @markdown {{- '<|start_header_id|>assistant<|end_header_id|>\\n\\n' }}\n",
"# @markdown {% endif %}\n",
"# @markdown \"\"\"\n",
"# @markdown ```\n",
@@ -1070,15 +1071,15 @@
"chat_template_string = r\"\"\"{{- bos_token }}\n",
"\n",
"{% for message in messages %}\n",
" {{- '\u003c|start_header_id|\u003e' + message.role + '\u003c|end_header_id|\u003e\\n\\n' + message.content | trim }}\n",
" {% if message.context and message.context | length \u003e 0 %}\n",
" {{- '<|start_header_id|>' + message.role + '<|end_header_id|>\\n\\n' + message.content | trim }}\n",
" {% if message.context and message.context | length > 0 %}\n",
" {{- ' Context: ' + message.context }}\n",
" {% endif %}\n",
" {{- '\u003c|eot_id|\u003e' }}\n",
" {{- '<|eot_id|>' }}\n",
"{% endfor %}\n",
"\n",
"{% if add_generation_prompt %}\n",
" {{- '\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n' }}\n",
" {{- '<|start_header_id|>assistant<|end_header_id|>\\n\\n' }}\n",
"{% endif %}\n",
"\"\"\"\n",
"\n",
@@ -1088,19 +1089,19 @@
"# @markdown output = bos_token\n",
"# @markdown\n",
"# @markdown for message in messages:\n",
"# @markdown output += f\"\u003c|start_header_id|\u003e{message['role']}\u003c|end_header_id|\u003e\\n\\n{message['content'].strip()}\"\n",
"# @markdown output += f\"<|start_header_id|>{message['role']}<|end_header_id|>\\n\\n{message['content'].strip()}\"\n",
"# @markdown if 'context' in message and message['context']:\n",
"# @markdown output += f\" Context: {message['context']}\"\n",
"# @markdown output += \"\u003c|eot_id|\u003e\"\n",
"# @markdown output += \"<|eot_id|>\"\n",
"# @markdown\n",
"# @markdown if add_generation_prompt:\n",
"# @markdown output += \"\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n\"\n",
"# @markdown output += \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"# @markdown\n",
"# @markdown return output\n",
"# @markdown ```\n",
"\n",
"# @markdown Run this cell to show an example output of the template given the\n",
"# @markdown below `messages` and `bos_token=\"\u003c|begin_of_text|\u003e\"`.\n",
"# @markdown below `messages` and `bos_token=\"<|begin_of_text|>\"`.\n",
"# @markdown ```\n",
"# @markdown messages = [\n",
"# @markdown {\"role\": \"user\", \"content\": \"Hello, how are you?\", \"context\": \"This is a test context.\"},\n",
@@ -1116,13 +1117,13 @@
" output = bos_token\n",
"\n",
" for message in messages:\n",
" output += f\"\u003c|start_header_id|\u003e{message['role']}\u003c|end_header_id|\u003e\\n\\n{message['content'].strip()}\"\n",
" output += f\"<|start_header_id|>{message['role']}<|end_header_id|>\\n\\n{message['content'].strip()}\"\n",
" if \"context\" in message and message[\"context\"]:\n",
" output += f\" Context: {message['context']}\"\n",
" output += \"\u003c|eot_id|\u003e\"\n",
" output += \"<|eot_id|>\"\n",
"\n",
" if add_generation_prompt:\n",
" output += \"\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n\"\n",
" output += \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"\n",
" return output\n",
"\n",
@@ -1137,7 +1138,7 @@
" {\"role\": \"user\", \"content\": \"Another question without context.\", \"context\": \"\"},\n",
"]\n",
"\n",
"rendered_text = render_template(messages, bos_token=\"\u003c|begin_of_text|\u003e\")\n",
"rendered_text = render_template(messages, bos_token=\"<|begin_of_text|>\")\n",
"pprint.pprint(rendered_text, width=80)"
]
},
@@ -1172,16 +1173,16 @@
"# @markdown template = {\n",
"# @markdown \"description\": \"Template used by Llama 3.1, accepting databricks dolly dataset.\",\n",
"# @markdown \"chat_template\": chat_template_string,\n",
"# @markdown \"instruction_separator\": \"\u003c|start_header_id|\u003euser\u003c|end_header_id|\u003e\\n\\n\",\n",
"# @markdown \"response_separator\": \"\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n\"\n",
"# @markdown \"instruction_separator\": \"<|start_header_id|>user<|end_header_id|>\\n\\n\",\n",
"# @markdown \"response_separator\": \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
"# @markdown }\n",
"# @markdown ```\n",
"\n",
"template_data = {\n",
" \"description\": \"Template used by Llama 3.1, accepting databricks dolly dataset.\",\n",
" \"chat_template\": chat_template_string,\n",
" \"instruction_separator\": \"\u003c|start_header_id|\u003euser\u003c|end_header_id|\u003e\\n\\n\",\n",
" \"response_separator\": \"\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n\",\n",
" \"instruction_separator\": \"<|start_header_id|>user<|end_header_id|>\\n\\n\",\n",
" \"response_separator\": \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\",\n",
"}\n",
"\n",
"template_filename = \"template.json\"\n",
@@ -1345,7 +1346,7 @@
"\n",
"# @markdown **Note**:\n",
"# @markdown 1. We recommend setting `finetuning_precision_mode` to `float16`.\n",
"# @markdown 1. If `max_steps\u003e0`, it takes precedence over `epochs`. One can set a small `max_steps`\n",
"# @markdown 1. If `max_steps>0`, it takes precedence over `epochs`. One can set a small `max_steps`\n",
"# @markdown value to quickly check the pipeline.\n",
"\n",
"# @markdown Acceletor type to use for training.\n",
@@ -1392,7 +1393,7 @@
"# Set config file.\n",
"if replica_count == 1:\n",
" config_file = \"vertex_vision_model_garden_peft/llama_fsdp_8gpu.yaml\"\n",
"elif replica_count \u003c= 4:\n",
"elif replica_count <= 4:\n",
" config_file = (\n",
" \"vertex_vision_model_garden_peft/\"\n",
" f\"llama_hsdp_{replica_count * per_node_accelerator_count}gpu.yaml\"\n",
@@ -1653,7 +1654,7 @@
"\n",
"\n",
"# @markdown Expected evaluation results:\n",
"# @markdown \u003e | alias | exact_match | exact_match_stderr | f1 | f1_stderr |\n",
"# @markdown > | alias | exact_match | exact_match_stderr | f1 | f1_stderr |\n",
"# @markdown | --- | --- | --- | --- | --- |\n",
"# @markdown | coqa | 0.3213 | 0.0197 | 0.4660 | 0.0187 |\n",
"\n",
@@ -1709,14 +1710,14 @@
" is_for_training=False,\n",
")\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"# Dedicated endpoint not supported yet.\n",
"use_dedicated_endpoint = False\n",
"\n",
"gpu_memory_utilization = 0.85\n",
"max_model_len = 8192 # Maximum context length.\n",
"\n",
"# Ensure max_model_len does not exceed the limit.\n",
"if max_model_len \u003e 8192:\n",
"if max_model_len > 8192:\n",
" raise ValueError(\"max_model_len cannot exceed 8192\")\n",
"\n",
"\n",
@@ -1742,7 +1743,7 @@
" use_dedicated_endpoint: bool = False,\n",
" max_num_seqs: int = 256,\n",
" model_type: str = None,\n",
") -\u003e Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with vLLM into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
@@ -1786,7 +1787,7 @@
" if enable_prefix_cache:\n",
" vllm_args.append(\"--enable-prefix-caching\")\n",
"\n",
" if 0 \u003c host_prefix_kv_cache_utilization_target \u003c 1:\n",
" if 0 < host_prefix_kv_cache_utilization_target < 1:\n",
" vllm_args.append(\n",
" f\"--host-prefix-kv-cache-utilization-target={host_prefix_kv_cache_utilization_target}\"\n",
" )\n",
@@ -1845,7 +1846,6 @@
" top_k: int,\n",
" raw_response: bool,\n",
" lora_weight: str = \"\",\n",
" use_dedicated_endpoint: bool = False,\n",
"):\n",
" # Parameters for inference.\n",
" instance = {\n",
@@ -1901,7 +1901,7 @@
"output_type": "stream",
"text": [
"Prompt:\n",
"\u003c|start_header_id|\u003euser\u003c|end_header_id|\u003e\\n\\nWhat was Anya looking for? Context: Anya clutched the worn teddy bear, its button eye dangling precariously. She'd lost it in the park yesterday, and the thought of never seeing Mr. Snuggles again made her tummy ache. She retraced her steps, her eyes scanning the colorful playground equipment and the sprawling green lawn.\u003c|eot_id|\u003e\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\n",
"<|start_header_id|>user<|end_header_id|>\\n\\nWhat was Anya looking for? Context: Anya clutched the worn teddy bear, its button eye dangling precariously. She'd lost it in the park yesterday, and the thought of never seeing Mr. Snuggles again made her tummy ache. She retraced her steps, her eyes scanning the colorful playground equipment and the sprawling green lawn.<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n",
"Output:\n",
"Anya was looking for her teddy bear, Mr. Snuggles.\n"
]
@@ -1925,8 +1925,8 @@
"\n",
"prompt = \"What was Anya looking for? Context: Anya clutched the worn teddy bear, its button eye dangling precariously. She'd lost it in the park yesterday, and the thought of never seeing Mr. Snuggles again made her tummy ache. She retraced her steps, her eyes scanning the colorful playground equipment and the sprawling green lawn.\" # @param {type: \"string\"}\n",
"prompt_with_headers = (\n",
" f\"\u003c|start_header_id|\u003euser\u003c|end_header_id|\u003e\\\\n\\\\n{prompt}\u003c|eot_id|\u003e\"\n",
" \"\u003c|start_header_id|\u003eassistant\u003c|end_header_id|\u003e\\n\\n\"\n",
" f\"<|start_header_id|>user<|end_header_id|>\\\\n\\\\n{prompt}<|eot_id|>\"\n",
" \"<|start_header_id|>assistant<|end_header_id|>\\n\\n\"\n",
")\n",
"\n",
"# @markdown If you encounter the issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, such as set `max_tokens` as 20.\n",
@@ -1945,7 +1945,6 @@
" top_k=top_k,\n",
" raw_response=raw_response,\n",
" lora_weight=lora_output_dir,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
@@ -3,6 +3,7 @@
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "7d9bbf86da5e"
@@ -34,18 +35,13 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_gemma2_deployment_on_vertex.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_gemma2_deployment_on_vertex.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma2_deployment_on_vertex.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -69,10 +65,6 @@
"- Deploy Gemma 2 with Hex-LLM on TPU\n",
"- Deploy Gemma with [TGI](https://github.com/huggingface/text-generation-inference) on GPU\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
@@ -119,38 +111,32 @@
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"# @markdown 2. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus). You can request for quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"# Import the necessary packages\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.64.0'\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"if os.environ.get(\"VERTEX_PRODUCT\") != \"COLAB_ENTERPRISE\":\n",
" ! pip install --upgrade tensorflow\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"LABEL = \"tgi\"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
@@ -158,26 +144,64 @@
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"gemma2\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION)\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\"\n",
"\n",
"import vertexai\n",
"\n",
"vertexai.init(\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
")\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"# @markdown ## Access Gemma 2 Models\n",
"\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Gemma 2 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"# @markdown You must provide a Hugging Face User Access Token (read) to access the Gemma 2 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"assert (\n",
@@ -205,22 +229,21 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "B7bg9nM0S0Mp"
"id": "E8OiHHNNE_wj"
},
"outputs": [],
"source": [
"# @title Select the model variants\n",
"# @title Deploy\n",
"# @markdown Set the model ID. Model weights can be loaded from HuggingFace or from a GCS bucket.\n",
"\n",
"# The pre-built serving docker images.\n",
"HEXLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai-restricted/vertex-vision-model-garden-dockers/hex-llm-serve:20241210_2323_RC00\"\n",
"\n",
"# @markdown Select one of the four model variations.\n",
"MODEL_ID = \"gemma-2-2b-it\" # @param [\"gemma-2-2b\", \"gemma-2-2b-it\", \"gemma-2-9b\", \"gemma-2-9b-it\", \"gemma-2-27b\", \"gemma-2-27b-it\"] {allow-input: true, isTemplate: true}\n",
"version_id = f\"publishers/google/models/gemma2/@{MODEL_ID}\"\n",
"\n",
"TPU_DEPLOYMENT_REGION = \"us-west1\" # @param [\"us-west1\"] {isTemplate:true}\n",
"model_id = os.path.join(model_path_prefix, MODEL_ID)\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# @markdown Find Vertex AI prediction TPUv5e machine types in\n",
"# @markdown https://cloud.google.com/vertex-ai/docs/predictions/use-tpu#deploy_a_model.\n",
"if \"2b\" in model_id:\n",
@@ -248,29 +271,16 @@
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "E8OiHHNNE_wj"
},
"outputs": [],
"source": [
"# @title Deploy Gemma2 models with Hex-LLM on TPU\n",
"# @markdown Set the model ID. Model weights can be loaded from HuggingFace or from a GCS bucket.\n",
"\n",
"# The pre-built serving docker images.\n",
"HEXLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai-restricted/vertex-vision-model-garden-dockers/hex-llm-serve:20241210_2323_RC00\"\n",
")\n",
"\n",
"# Server parameters.\n",
"tensor_parallel_size = accelerator_count\n",
"hbm_utilization_factor = 0.6 # Fraction of HBM memory allocated for KV cache after model loading. A larger value improves throughput but gives higher risk of TPU out-of-memory errors with long prompts.\n",
"max_running_seqs = 256 # Maximum number of running sequences in a continuous batch.\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# Endpoint configurations.\n",
"min_replica_count = 1\n",
"max_replica_count = 1\n",
@@ -281,6 +291,7 @@
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str = None,\n",
" base_model_id: str = None,\n",
" data_parallel_size: int = 1,\n",
" tensor_parallel_size: int = 1,\n",
@@ -369,6 +380,7 @@
" machine_type=machine_type,\n",
" tpu_topology=tpu_topology if num_hosts > 1 else None,\n",
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" min_replica_count=min_replica_count,\n",
" max_replica_count=max_replica_count,\n",
" system_labels={\n",
@@ -384,6 +396,7 @@
" model_id=model_id,\n",
" publisher=\"google\",\n",
" publisher_model_id=\"gemma2\",\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=machine_type,\n",
" tensor_parallel_size=tensor_parallel_size,\n",
" hbm_utilization_factor=hbm_utilization_factor,\n",
@@ -476,28 +489,23 @@
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "eYst8GHqcGco"
"id": "TBNJYZMlBNwZ"
},
"outputs": [],
"source": [
"# @title Select the model variants\n",
"# @title Deploy\n",
"\n",
"# The pre-built serving docker image.\n",
"TGI_DOCKER_URI = \"us-docker.pkg.dev/deeplearning-platform-release/gcr.io/huggingface-text-generation-inference-cu121.2-1.ubuntu2204.py310\"\n",
"\n",
"MODEL_ID = \"gemma-2-2b\" # @param [\"gemma-2-2b\", \"gemma-2-2b-it\", \"gemma-2-9b\", \"gemma-2-9b-it\", \"gemma-2-27b\", \"gemma-2-27b-it\"] {allow-input: true, isTemplate: true}\n",
"model_id = os.path.join(model_path_prefix, MODEL_ID)\n",
"PUBLISHER_MODEL_NAME = f\"publishers/google/models/gemma2@{MODEL_ID}\"\n",
"\n",
"# @markdown Finds Vertex AI prediction supported accelerators and regions in\n",
"# @markdown https://cloud.google.com/vertex-ai/docs/predictions/configure-compute.\n",
"\n",
"accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\"] {isTemplate: true}\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"\n",
"if \"2b\" in MODEL_ID:\n",
" if accelerator_type == \"NVIDIA_L4\":\n",
" # Sets 1 L4 (24G) to deploy Gemma 2 2B models.\n",
@@ -539,46 +547,6 @@
" is_for_training=False,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "nlRmOQmZhjvp"
},
"outputs": [],
"source": [
"# @title [Option 1] Deploy with Model Garden SDK\n",
"\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"from vertexai import model_garden\n",
"\n",
"model = model_garden.OpenModel(PUBLISHER_MODEL_NAME)\n",
"endpoints[LABEL] = model.deploy(\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "TBNJYZMlBNwZ"
},
"outputs": [],
"source": [
"# @title [Option 2] Deploy Gemma models with TGI on GPU\n",
"\n",
"# Note that larger token counts will require more GPU memory. For example, if you'd\n",
"# like to increase the `max_total_tokens` and `max_batch_prefill_tokens` to 8192,\n",
"# you may need 1 L4 for 2b model, 4 L4s for the 9b model, and 8 L4s for the 27b model.\n",
@@ -592,7 +560,7 @@
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str = None,\n",
" service_account: str,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
@@ -623,9 +591,6 @@
" except NameError:\n",
" pass\n",
"\n",
" if service_account:\n",
" env_vars[\"SERVICE_ACCOUNT\"] = service_account\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=TGI_DOCKER_URI,\n",
@@ -652,11 +617,12 @@
" return model, endpoint\n",
"\n",
"\n",
"models[LABEL], endpoints[LABEL] = deploy_model_tgi(\n",
"models[\"tgi\"], endpoints[\"tgi\"] = deploy_model_tgi(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=MODEL_ID),\n",
" model_id=model_id,\n",
" publisher=\"google\",\n",
" publisher_model_id=\"gemma2\",\n",
" service_account=SERVICE_ACCOUNT,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
@@ -664,9 +630,7 @@
" max_total_tokens=max_total_tokens,\n",
" max_batch_prefill_tokens=max_batch_prefill_tokens,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
")"
]
},
{
@@ -749,8 +713,6 @@
},
"outputs": [],
"source": [
"# @title Delete the models and endpoints\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
@@ -760,7 +722,11 @@
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()"
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
]
}
],
@@ -3,6 +3,7 @@
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "7d9bbf86da5e"
@@ -26,6 +27,7 @@
},
{
"cell_type": "markdown",
"language": "markdown",
"metadata": {
"id": "99c1c3fc2ca5"
},
@@ -34,18 +36,13 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_gemma2_finetuning_on_vertex.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_gemma2_finetuning_on_vertex.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma2_finetuning_on_vertex.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -68,7 +65,6 @@
"### Objective\n",
"\n",
"- Finetune and deploy Gemma 2 models with Vertex AI Custom Training Jobs.\n",
"- Evaluate the finetuned model using [lm-evaluation-harness](https://github.com/EleutherAI/lm-evaluation-harness).\n",
"- Send prediction requests to your finetuned Gemma 2 model.\n",
"\n",
"### File a bug\n",
@@ -115,6 +111,7 @@
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "855d6b96f291"
@@ -146,7 +143,7 @@
"\n",
"# Import the necessary packages\n",
"! rm -rf vertex-ai-samples && git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples && git reset --hard c45f6a4f4d32e31a050f0e4ba52824b0caf4eda3\n",
"! cd vertex-ai-samples && git reset --hard 80320a9a1b818534ca785444e704f6953f2a9dd9\n",
"\n",
"import datetime\n",
"import importlib\n",
@@ -158,9 +155,6 @@
"from google.cloud.aiplatform.compat.types import \\\n",
" custom_job as gca_custom_job_compat\n",
"\n",
"if os.environ.get(\"VERTEX_PRODUCT\") != \"COLAB_ENTERPRISE\":\n",
" ! pip install --upgrade tensorflow\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
@@ -229,7 +223,7 @@
"\n",
"# @markdown ## Access Gemma 2 Models\n",
"\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Gemma 2 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"# @markdown You must provide a Hugging Face User Access Token (read) to access the Gemma 2 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"assert HF_TOKEN, \"Provide a read HF_TOKEN to load models from Hugging Face.\"\n",
@@ -395,6 +389,7 @@
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "ivVGS9dHXPOz"
@@ -412,7 +407,9 @@
"# @markdown 1. If `max_steps > 0`, it takes precedence over `epochs`. One can set a small `max_steps` value to quickly check the pipeline.\n",
"\n",
"# @markdown Accelerator type to use for training.\n",
"# fmt: off\n",
"training_accelerator_type = \"NVIDIA_A100_80GB\" # @param [\"NVIDIA_A100_80GB\", \"NVIDIA_H100_80GB\"]\n",
"# fmt: on\n",
"\n",
"# The pre-built training docker image.\n",
"if training_accelerator_type == \"NVIDIA_A100_80GB\":\n",
@@ -430,7 +427,7 @@
" }\n",
"\n",
"TRAIN_DOCKER_URI = (\n",
" f\"{repo}/vertex-vision-model-garden-dockers/pytorch-peft-train:stable_20250705\"\n",
" f\"{repo}/vertex-vision-model-garden-dockers/pytorch-peft-train:stable_20250213\"\n",
")\n",
"\n",
"# Worker pool spec.\n",
@@ -506,9 +503,8 @@
"merged_model_output_dir = os.path.join(base_output_dir, \"merged-model\")\n",
"\n",
"# Add labels for the finetuning job.\n",
"\n",
"labels = {\n",
" \"mg-source\": common_util.get_deploy_source(),\n",
" \"mg-source\": \"notebook\",\n",
" \"mg-notebook-name\": \"model_garden_gemma2_finetuning_on_vertex.ipynb\".split(\".\")[0],\n",
"}\n",
"\n",
@@ -586,6 +582,8 @@
"# Wait until resource has been created.\n",
"train_job.wait_for_resource_creation()\n",
"\n",
"merged_model_output_dir = os.path.join(merged_model_output_dir, \"node-0\")\n",
"\n",
"print(\"LoRA adapter will be saved in:\", lora_output_dir)\n",
"print(\"Trained and merged models will be saved in:\", merged_model_output_dir)\n",
"\n",
@@ -615,126 +613,7 @@
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "KdtcMGHgtrVC"
},
"outputs": [],
"source": [
"# @title Select Evaluation Checkpoint\n",
"\n",
"if train_job.end_time is None:\n",
" print(\"Waiting for the training job to finish...\")\n",
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"# @markdown The following checkpoints are available for evaluation:\n",
"! gcloud storage ls \"{lora_output_dir}/node-0\" | grep \"checkpoint-\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "1LBADPr6tTqy"
},
"outputs": [],
"source": [
"# @title Run Evaluation Job\n",
"# @markdown This section runs the evaluation using [lm-evaluation-harness](https://github.com/EleutherAI/lm-evaluation-harness) on the finetuned model. The evaluation takes approximately 20 mins to finish.\n",
"\n",
"# The pre-built evaluation docker image for LM Evaluation Harness.\n",
"LM_EVAL_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-lm-evaluation-harness:20250410_1035_RC00\"\n",
"\n",
"# @markdown Set `RUN_EVALUATION` to False to skip the evaluation job.\n",
"RUN_EVALUATION = True # @param {type:\"boolean\"}\n",
"\n",
"eval_accelerator_type = \"NVIDIA_L4\"\n",
"gpu_memory_utilization = 0.85\n",
"\n",
"if \"2b\" in base_model_id:\n",
" eval_machine_type = \"g2-standard-12\"\n",
" eval_accelerator_count = 1\n",
"elif \"9b\" in base_model_id:\n",
" eval_machine_type = \"g2-standard-48\"\n",
" eval_accelerator_count = 4\n",
"elif \"27b\" in base_model_id:\n",
" eval_machine_type = \"g2-standard-96\"\n",
" eval_accelerator_count = 8\n",
" gpu_memory_utilization = 0.8\n",
"else:\n",
" raise ValueError(\n",
" \"Recommended machine settings not found for model: %s\" % base_model_id\n",
" )\n",
"\n",
"# @markdown Set `evaluation_checkpoint_dir` to an intermediate checkpoint from the above training job. If not set, the evaluation job will use the merged model.\n",
"evaluation_checkpoint_dir = \"\" # @param {type:\"string\"}\n",
"if evaluation_checkpoint_dir:\n",
" pretrained = pretrained_model_id\n",
"else:\n",
" pretrained = merged_model_output_dir\n",
"\n",
"# @markdown Evaluation tasks to run.\n",
"eval_tasks = \"coqa\" # @param {type:\"string\"}\n",
"# @markdown Model to use for evaluation.\n",
"model = \"vllm\" # @param {type:\"string\"}\n",
"# @markdown Batch size for evaluation.\n",
"batch_size = \"auto\" # @param {type:\"string\"}\n",
"apply_chat_template = True if \"-it\" in pretrained_model_id else False\n",
"max_model_len = 4096 # Maximum context length.\n",
"\n",
"model_args = f\"tensor_parallel_size={eval_accelerator_count},max_model_len={max_model_len},gpu_memory_utilization={gpu_memory_utilization},enforce_eager=True\"\n",
"eval_output_dir = os.path.join(base_output_dir, \"lm_eval\")\n",
"\n",
"lm_eval_job_args = [\n",
" \"--task=lm_eval\",\n",
" f\"--model={model}\",\n",
" f\"--eval_tasks={eval_tasks}\",\n",
" f\"--pretrained_model_name_or_path={pretrained}\",\n",
" f\"--model_args={model_args}\",\n",
" f\"--output_dir={eval_output_dir}\",\n",
" f\"--apply_chat_template={apply_chat_template}\",\n",
" f\"--batch_size={batch_size}\",\n",
" f\"--huggingface_access_token={HF_TOKEN}\",\n",
"]\n",
"\n",
"if evaluation_checkpoint_dir:\n",
" lm_eval_job_args.append(f\"--lora_path={evaluation_checkpoint_dir}\")\n",
"\n",
"if RUN_EVALUATION:\n",
" common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=eval_accelerator_type,\n",
" accelerator_count=eval_accelerator_count,\n",
" is_for_training=True,\n",
" )\n",
" lm_eval_job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=common_util.get_job_name_with_datetime(\"gemma2-lm-eval\"),\n",
" container_uri=LM_EVAL_DOCKER_URI,\n",
" labels=labels,\n",
" )\n",
"\n",
" print(\"Running evaluation job with args:\")\n",
" print(\" \\\\\\n\".join(lm_eval_job_args))\n",
" lm_eval_job.run(\n",
" args=lm_eval_job_args,\n",
" replica_count=1,\n",
" machine_type=eval_machine_type,\n",
" accelerator_type=eval_accelerator_type,\n",
" accelerator_count=eval_accelerator_count,\n",
" boot_disk_size_gb=500,\n",
" service_account=SERVICE_ACCOUNT,\n",
" base_output_dir=base_output_dir,\n",
" )\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "qmHW6m8xG_4U"
@@ -743,12 +622,18 @@
"source": [
"# @title Deploy\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint. It takes 15 minutes to 1 hour to finish.\n",
"\n",
"if train_job.end_time is None:\n",
" print(\"Waiting for the training job to finish...\")\n",
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"print(\"Deploying models in:\", merged_model_output_dir)\n",
"\n",
"# The pre-built serving docker image for vLLM.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20250116_0916_RC00\"\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# Find Vertex AI prediction supported accelerators and regions [here](https://cloud.google.com/vertex-ai/docs/predictions/configure-compute).\n",
@@ -1042,12 +927,7 @@
"outputs": [],
"source": [
"# Delete the train job.\n",
"\n",
"if train_job:\n",
" train_job.delete()\n",
"if RUN_EVALUATION and lm_eval_job:\n",
" lm_eval_job.delete()\n",
"\n",
"train_job.delete()\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
@@ -34,7 +34,7 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_gemma3_deployment_on_vertex.ipynb\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/instances\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
@@ -45,7 +45,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma3_deployment_on_vertex.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -66,10 +66,6 @@
"\n",
"- Deploy Gemma 3 with vLLM on GPU\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
@@ -116,7 +112,7 @@
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.84.0'\n",
"\n",
"# Import the necessary packages\n",
"import importlib\n",
@@ -134,9 +130,8 @@
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"# Initialize models and endpoints as a dict\n",
"models, endpoints = {}, {}\n",
"\n",
"LABEL = \"vllm_gpu\"\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
@@ -224,9 +219,8 @@
"source": [
"# @title [Option 1] Deploy with Model Garden SDK\n",
"\n",
"LABEL = \"sdk-deploy-1b\"\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"from vertexai import model_garden\n",
"from vertexai.preview import model_garden\n",
"\n",
"model = model_garden.OpenModel(PUBLISHER_MODEL_NAME)\n",
"endpoints[LABEL] = model.deploy(\n",
@@ -235,9 +229,7 @@
" accelerator_count=accelerator_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]"
")"
]
},
{
@@ -382,9 +374,7 @@
" return model, endpoint\n",
"\n",
"\n",
"LABEL = \"custom-deploy-1b\"\n",
"\n",
"models[LABEL], endpoints[LABEL] = deploy_model_vllm(\n",
"models[\"vllm_gpu\"], endpoints[\"vllm_gpu\"] = deploy_model_vllm(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"gemma3-serve\"),\n",
" model_id=model_id,\n",
" publisher=\"google\",\n",
@@ -396,10 +386,7 @@
" gpu_memory_utilization=gpu_memory_utilization,\n",
" max_model_len=max_model_len,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"model = models[LABEL]\n",
"endpoint = endpoints[LABEL]"
")"
]
},
{
@@ -457,7 +444,7 @@
" \"raw_response\": raw_response,\n",
" },\n",
"]\n",
"response = endpoint.predict(\n",
"response = endpoints[\"vllm_gpu\"].predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"\n",
@@ -479,8 +466,10 @@
"# @title Chat completion\n",
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoint.gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoint.resource_name\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"vllm_gpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"vllm_gpu\"].name\n",
")\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
@@ -607,9 +596,8 @@
"source": [
"# @title [Option 1] Deploy with Model Garden SDK\n",
"\n",
"LABEL = \"sdk-deploy\"\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"from vertexai import model_garden\n",
"from vertexai.preview import model_garden\n",
"\n",
"model = model_garden.OpenModel(PUBLISHER_MODEL_NAME)\n",
"endpoints[LABEL] = model.deploy(\n",
@@ -618,9 +606,7 @@
" accelerator_count=accelerator_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]"
")"
]
},
{
@@ -639,8 +625,6 @@
"gpu_memory_utilization = 0.95\n",
"max_model_len = 131072\n",
"\n",
"LABEL = \"multimodal-deploy\"\n",
"\n",
"\n",
"def deploy_model_vllm(\n",
" model_name: str,\n",
@@ -767,7 +751,7 @@
" return model, endpoint\n",
"\n",
"\n",
"models[LABEL], endpoints[LABEL] = deploy_model_vllm(\n",
"models[\"vllm_gpu\"], endpoints[\"vllm_gpu\"] = deploy_model_vllm(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"gemma3-serve\"),\n",
" model_id=model_id,\n",
" publisher=\"google\",\n",
@@ -779,10 +763,7 @@
" gpu_memory_utilization=gpu_memory_utilization,\n",
" max_model_len=max_model_len,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"model = models[LABEL]\n",
"endpoint = endpoints[LABEL]"
")"
]
},
{
@@ -797,8 +778,10 @@
"# @title Chat completion with text-only requests\n",
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoint.gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoint.resource_name\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"vllm_gpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"vllm_gpu\"].name\n",
")\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
@@ -872,8 +855,10 @@
"# @title Chat completion with multimodal requests\n",
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoint.gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoint.resource_name\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"vllm_gpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"vllm_gpu\"].name\n",
")\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
@@ -885,7 +870,7 @@
"\n",
"# @markdown Next fill out some request parameters:\n",
"\n",
"user_image = \"https://upload.wikimedia.org/wikipedia/commons/thumb/d/dd/Gfp-wisconsin-madison-the-nature-boardwalk.jpg/2560px-Gfp-wisconsin-madison-the-nature-boardwalk.jpg\"\n",
"user_image = \"https://upload.wikimedia.org/wikipedia/commons/thumb/c/cb/The_Blue_Marble_%28remastered%29.jpg/580px-The_Blue_Marble_%28remastered%29.jpg\" # @param {type: \"string\"}\n",
"user_message = \"What is in the image?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter the issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, such as set `max_tokens` as 20.\n",
"max_tokens = 50 # @param {type: \"integer\"}\n",
@@ -3,6 +3,7 @@
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "7d9bbf86da5e"
@@ -26,6 +27,7 @@
},
{
"cell_type": "markdown",
"language": "markdown",
"metadata": {
"id": "99c1c3fc2ca5"
},
@@ -34,18 +36,13 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_gemma3_finetuning_on_vertex.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_gemma3_finetuning_on_vertex.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma3_finetuning_on_vertex.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -68,7 +65,6 @@
"### Objective\n",
"\n",
"- Finetune and deploy Gemma 3 models with Vertex AI Custom Training Jobs.\n",
"- Evaluate the finetuned model using [lm-evaluation-harness](https://github.com/EleutherAI/lm-evaluation-harness).\n",
"- Send prediction requests to your finetuned Gemma 3 model.\n",
"\n",
"### File a bug\n",
@@ -115,8 +111,9 @@
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "code",
"cellView": "form",
"id": "855d6b96f291"
},
"outputs": [],
@@ -146,7 +143,7 @@
"\n",
"# Import the necessary packages\n",
"! rm -rf vertex-ai-samples && git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"! cd vertex-ai-samples && git reset --hard c45f6a4f4d32e31a050f0e4ba52824b0caf4eda3\n",
"! cd vertex-ai-samples && git reset --hard 80320a9a1b818534ca785444e704f6953f2a9dd9\n",
"\n",
"import datetime\n",
"import importlib\n",
@@ -158,19 +155,12 @@
"from google.cloud.aiplatform.compat.types import \\\n",
" custom_job as gca_custom_job_compat\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"if os.environ.get(\"VERTEX_PRODUCT\") != \"COLAB_ENTERPRISE\":\n",
" ! pip install --upgrade tensorflow\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"# Initialize models and endpoints as a dict\n",
"models, endpoints = {}, {}\n",
"\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
@@ -233,7 +223,7 @@
"\n",
"# @markdown ## Access Gemma 3 Models\n",
"\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Gemma 3 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"# @markdown You must provide a Hugging Face User Access Token (read) to access the Gemma 3 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"assert HF_TOKEN, \"Provide a read HF_TOKEN to load models from Hugging Face.\"\n",
@@ -399,6 +389,7 @@
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "ivVGS9dHXPOz"
@@ -416,7 +407,9 @@
"# @markdown 1. If `max_steps > 0`, it takes precedence over `epochs`. One can set a small `max_steps` value to quickly check the pipeline.\n",
"\n",
"# @markdown Accelerator type to use for training.\n",
"# fmt: off\n",
"training_accelerator_type = \"NVIDIA_A100_80GB\" # @param [\"NVIDIA_A100_80GB\", \"NVIDIA_H100_80GB\"]\n",
"# fmt: on\n",
"\n",
"# The pre-built training docker image.\n",
"if training_accelerator_type == \"NVIDIA_A100_80GB\":\n",
@@ -550,7 +543,6 @@
" f\"--lr_scheduler_type={lr_scheduler_type}\",\n",
" f\"--precision_mode={finetuning_precision_mode}\",\n",
" f\"--train_precision={train_precision}\",\n",
" f\"--merge_model_precision_mode={train_precision}\",\n",
" f\"--gradient_checkpointing={gradient_checkpointing}\",\n",
" f\"--num_train_epochs={num_train_epochs}\",\n",
" f\"--attn_implementation={attn_implementation}\",\n",
@@ -618,113 +610,7 @@
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "4g0woSqhvF9O"
},
"outputs": [],
"source": [
"# @title Select Evaluation Checkpoint\n",
"\n",
"if train_job.end_time is None:\n",
" print(\"Waiting for the training job to finish...\")\n",
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"# @markdown The following checkpoints are available for evaluation:\n",
"! gcloud storage ls \"{lora_output_dir}/node-0\" | grep \"checkpoint-\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "635Bdo0Pt6iq"
},
"outputs": [],
"source": [
"# @title Run Evaluation Job\n",
"# @markdown This section runs the evaluation using [lm-evaluation-harness](https://github.com/EleutherAI/lm-evaluation-harness) on the finetuned model. The evaluation takes approximately 20 mins to finish.\n",
"\n",
"# The pre-built evaluation docker image for LM Evaluation Harness.\n",
"LM_EVAL_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-lm-evaluation-harness:20250410_1035_RC00\"\n",
"\n",
"# @markdown Set `RUN_EVALUATION` to False to skip the evaluation job.\n",
"RUN_EVALUATION = True # @param {type:\"boolean\"}\n",
"\n",
"eval_accelerator_type = \"NVIDIA_L4\"\n",
"eval_machine_type = \"g2-standard-12\"\n",
"eval_accelerator_count = 1\n",
"\n",
"# @markdown Set `evaluation_checkpoint_dir` to an intermediate checkpoint from the above training job. If not set, the evaluation job will use the merged model.\n",
"evaluation_checkpoint_dir = \"\" # @param {type:\"string\"}\n",
"if evaluation_checkpoint_dir:\n",
" pretrained = pretrained_model_id\n",
"else:\n",
" pretrained = merged_model_output_dir\n",
"\n",
"# @markdown Evaluation tasks to run.\n",
"eval_tasks = \"coqa\" # @param {type:\"string\"}\n",
"# @markdown Model to use for evaluation.\n",
"model = \"vllm\" # @param {type:\"string\"}\n",
"# @markdown Batch size for evaluation.\n",
"batch_size = \"auto\" # @param {type:\"string\"}\n",
"apply_chat_template = True if \"-it\" in pretrained_model_id else False\n",
"gpu_memory_utilization = 0.85\n",
"max_model_len = 4096 # Maximum context length.\n",
"\n",
"model_args = f\"tensor_parallel_size={eval_accelerator_count},max_model_len={max_model_len},gpu_memory_utilization={gpu_memory_utilization},enforce_eager=True\"\n",
"eval_output_dir = os.path.join(base_output_dir, \"lm_eval\")\n",
"\n",
"lm_eval_job_args = [\n",
" \"--task=lm_eval\",\n",
" f\"--model={model}\",\n",
" f\"--eval_tasks={eval_tasks}\",\n",
" f\"--pretrained_model_name_or_path={pretrained}\",\n",
" f\"--model_args={model_args}\",\n",
" f\"--output_dir={eval_output_dir}\",\n",
" f\"--apply_chat_template={apply_chat_template}\",\n",
" f\"--batch_size={batch_size}\",\n",
" f\"--huggingface_access_token={HF_TOKEN}\",\n",
"]\n",
"\n",
"if evaluation_checkpoint_dir:\n",
" lm_eval_job_args.append(f\"--lora_path={evaluation_checkpoint_dir}\")\n",
"\n",
"if RUN_EVALUATION:\n",
" common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=eval_accelerator_type,\n",
" accelerator_count=eval_accelerator_count,\n",
" is_for_training=True,\n",
" )\n",
" lm_eval_job = aiplatform.CustomContainerTrainingJob(\n",
" display_name=common_util.get_job_name_with_datetime(\"gemma3-lm-eval\"),\n",
" container_uri=LM_EVAL_DOCKER_URI,\n",
" labels=labels,\n",
" )\n",
"\n",
" print(\"Running evaluation job with args:\")\n",
" print(\" \\\\\\n\".join(lm_eval_job_args))\n",
" lm_eval_job.run(\n",
" args=lm_eval_job_args,\n",
" replica_count=1,\n",
" machine_type=eval_machine_type,\n",
" accelerator_type=eval_accelerator_type,\n",
" accelerator_count=eval_accelerator_count,\n",
" boot_disk_size_gb=500,\n",
" service_account=SERVICE_ACCOUNT,\n",
" base_output_dir=base_output_dir,\n",
" )\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "qmHW6m8xG_4U"
@@ -733,12 +619,18 @@
"source": [
"# @title Deploy\n",
"# @markdown This section uploads the model to Model Registry and deploys it on the Endpoint. It takes 15 minutes to 1 hour to finish.\n",
"\n",
"if train_job.end_time is None:\n",
" print(\"Waiting for the training job to finish...\")\n",
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"print(\"Deploying models in:\", merged_model_output_dir)\n",
"\n",
"# The pre-built serving docker image for vLLM.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20250312_0916_RC01\"\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# Find Vertex AI prediction supported accelerators and regions [here](https://cloud.google.com/vertex-ai/docs/predictions/configure-compute).\n",
@@ -912,9 +804,6 @@
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"model = models[\"vllm_gpu\"]\n",
"endpoint = endpoints[\"vllm_gpu\"]\n",
"\n",
"# @markdown Click \"Show code\" to see more details."
]
},
@@ -973,7 +862,7 @@
" \"raw_response\": raw_response,\n",
" },\n",
"]\n",
"response = endpoint.predict(\n",
"response = endpoints[\"vllm_gpu\"].predict(\n",
" instances=instances, use_dedicated_endpoint=use_dedicated_endpoint\n",
")\n",
"\n",
@@ -1002,12 +891,7 @@
"outputs": [],
"source": [
"# Delete the train job.\n",
"\n",
"if train_job:\n",
" train_job.delete()\n",
"if RUN_EVALUATION and lm_eval_job:\n",
" lm_eval_job.delete()\n",
"\n",
"train_job.delete()\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
@@ -1,868 +0,0 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "SgQ6t5bqZVlH"
},
"outputs": [],
"source": [
"# Copyright 2025 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": "99c1c3fc2ca5"
},
"source": [
"# Vertex AI Model Garden - Gemma 3n (Deployment)\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_gemma3n_deployment_on_vertex.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_gemma3n_deployment_on_vertex.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma3n_deployment_on_vertex.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3de7470326a2"
},
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates serving Gemma3n models with [SGLang](https://github.com/sgl-project/sglang). Gemma 3n models use selective parameter activation technology to reduce resource requirements. This technique allows the models to operate at an effective size of 2B and 4B parameters, which is lower than the total number of parameters they contain. For more information on Gemma 3n's efficient parameter management technology, see the [Gemma 3n](https://ai.google.dev/gemma/docs/gemma-3n#parameters) page.\n",
"\n",
"\n",
"### Objective\n",
"\n",
"- Deploy Gemma 3n with SGLang on GPU.\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
"\n",
"* Vertex AI\n",
"* Cloud Storage\n",
"\n",
"Learn about [Vertex AI pricing](https://cloud.google.com/vertex-ai/pricing), [Cloud Storage pricing](https://cloud.google.com/storage/pricing), and use the [Pricing Calculator](https://cloud.google.com/products/calculator/) to generate a cost estimate based on your projected usage."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "264c07757582"
},
"source": [
"## Before you begin"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "ax7zWynUDcjk"
},
"outputs": [],
"source": [
"# @title Request for quota\n",
"\n",
"# @markdown To deploy Gemma 3n models, you need 1 host of 1 x H100 machine. Check that you have sufficient quota:\n",
"# @markdown - For Spot VM quota, check [`CustomModelServingPreemptibleH100GPUsPerProjectPerRegion`](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_preemptible_nvidia_h100_gpus).\n",
"# @markdown - For regular VM quota, check [`CustomModelServingH100GPUsPerProjectPerRegion`](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus).\n",
"#\n",
"# @markdown If you don't have sufficient quota, request for quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota)."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "YXFGIp1l-qtT"
},
"outputs": [],
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus). You can request for quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"\n",
"# Import the necessary packages\n",
"import importlib\n",
"import os\n",
"import time\n",
"from typing import Tuple\n",
"\n",
"import requests\n",
"from google import auth\n",
"from google.cloud import aiplatform\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"if os.environ.get(\"VERTEX_PRODUCT\") != \"COLAB_ENTERPRISE\":\n",
" ! pip install --upgrade tensorflow\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"\n",
"def check_quota(\n",
" project_id: str,\n",
" region: str,\n",
" resource_id: str,\n",
" accelerator_count: int,\n",
"):\n",
" \"\"\"Checks if the project and the region has the required quota.\"\"\"\n",
" quota = common_util.get_quota(project_id, region, resource_id)\n",
" quota_request_instruction = (\n",
" \"Either use \"\n",
" \"a different region or request additional quota. Follow \"\n",
" \"instructions here \"\n",
" \"https://cloud.google.com/docs/quotas/view-manage#requesting_higher_quota\"\n",
" \" to check quota in a region or request additional quota for \"\n",
" \"your project.\"\n",
" )\n",
" if quota == -1:\n",
" raise ValueError(\n",
" f\"Quota not found for: {resource_id} in {region}.\"\n",
" f\" {quota_request_instruction}\"\n",
" )\n",
" if quota < accelerator_count:\n",
" raise ValueError(\n",
" f\"Quota not enough for {resource_id} in {region}: {quota} <\"\n",
" f\" {accelerator_count}. {quota_request_instruction}\"\n",
" )\n",
"\n",
"\n",
"LABEL = \"sglang_gpu\"\n",
"models, endpoints = {}, {}\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION)\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"\n",
"import vertexai\n",
"\n",
"vertexai.init(\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
")"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "3zJJDmldn7rw"
},
"source": [
"## Deploy Gemma 3n with SGLang"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "_3Swj3pxn7rw"
},
"outputs": [],
"source": [
"# @title Select the model variants\n",
"\n",
"# @markdown Set the model to deploy.\n",
"\n",
"base_model_name = \"gemma-3n-E4B-it\" # @param [\"gemma-3n-E4B-it\", \"gemma-3n-E4B\", \"gemma-3n-E2B-it\", \"gemma-3n-E2B\"] {isTemplate:true}\n",
"model_id = \"gs://vertex-model-garden-restricted-us/gemma3n/\" + base_model_name\n",
"hf_model_id = \"google/\" + base_model_name\n",
"\n",
"# The pre-built serving docker images.\n",
"SGLANG_DOCKER_URI = \"us-docker.pkg.dev/deeplearning-platform-release/vertex-model-garden/sglang-serve.cu124.0-4.ubuntu2204.py310:20250626-1121-rc0\"\n",
"\n",
"# @markdown Choose whether to use a [Spot VM](https://cloud.google.com/compute/docs/instances/spot) for the deployment.\n",
"is_spot = False # @param {type:\"boolean\"}\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# @markdown Find Vertex AI prediction supported accelerators and regions at https://cloud.google.com/vertex-ai/docs/predictions/configure-compute.\n",
"accelerator_type = \"NVIDIA_H100_80GB\" # @param [\"NVIDIA_H100_80GB\", \"NVIDIA_A100_80GB\"] {isTemplate:true}\n",
"\n",
"PUBLISHER_MODEL_NAME = f\"publishers/google/models/gemma3n@{base_model_name.lower()}\"\n",
"\n",
"if accelerator_type == \"NVIDIA_H100_80GB\":\n",
" if is_spot:\n",
" resource_id = \"custom_model_serving_preemptible_nvidia_h100_gpus\"\n",
" else:\n",
" resource_id = \"custom_model_serving_nvidia_h100_gpus\"\n",
" machine_type = \"a3-highgpu-1g\"\n",
" accelerator_count = 1\n",
"elif accelerator_type == \"NVIDIA_A100_80GB\":\n",
" if is_spot:\n",
" resource_id = \"custom_model_serving_preemptible_nvidia_a100_80gb_gpus\"\n",
" else:\n",
" resource_id = \"custom_model_serving_nvidia_a100_80gb_gpus\"\n",
" machine_type = \"a2-ultragpu-1g\"\n",
" accelerator_count = 1\n",
"else:\n",
" raise ValueError(f\"Recommended GPU setting not found for: {base_model_name}.\")\n",
"\n",
"check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" resource_id=resource_id,\n",
" accelerator_count=accelerator_count,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "omW0LaC8wWz5"
},
"outputs": [],
"source": [
"# @title [Option 1] Deploy with Model Garden SDK\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"deploy_request_timeout = 1800 # 30 minutes\n",
"from vertexai import model_garden\n",
"\n",
"model = model_garden.OpenModel(PUBLISHER_MODEL_NAME)\n",
"endpoints[LABEL] = model.deploy(\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" spot=is_spot,\n",
" deploy_request_timeout=deploy_request_timeout,\n",
" accept_eula=False,\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "3m-tDxgawYhU"
},
"outputs": [],
"source": [
"# @title [Option 2] Deploy with customized configs\n",
"\n",
"# @markdown This section uploads Gemma 3n models to Model Registry and deploys them to a Vertex Prediction Endpoint. It takes ~30 minutes to finish.\n",
"\n",
"# @markdown It's recommended to use the region selected by the deployment button on the model card. If the deployment button is not available, it's recommended to stay with the default region of the notebook.\n",
"\n",
"\n",
"def poll_operation(op_name: str) -> bool: # noqa: F811\n",
" creds, _ = auth.default()\n",
" auth_req = auth.transport.requests.Request()\n",
" creds.refresh(auth_req)\n",
" headers = {\n",
" \"Authorization\": f\"Bearer {creds.token}\",\n",
" }\n",
" get_resp = requests.get(\n",
" f\"https://{REGION}-aiplatform.googleapis.com/ui/{op_name}\",\n",
" headers=headers,\n",
" )\n",
" opjs = get_resp.json()\n",
" if \"error\" in opjs:\n",
" raise ValueError(f\"Operation failed: {opjs['error']}\")\n",
" return opjs.get(\"done\", False)\n",
"\n",
"\n",
"def poll_and_wait(op_name: str, total_wait: int, interval: int = 60): # noqa: F811\n",
" waited = 0\n",
" while not poll_operation(op_name):\n",
" if waited > total_wait:\n",
" raise TimeoutError(\"Operation timed out\")\n",
" print(\n",
" f\"\\rStill waiting for operation... Waited time in second: {waited:<6}\",\n",
" end=\"\",\n",
" flush=True,\n",
" )\n",
" waited += interval\n",
" time.sleep(interval)\n",
"\n",
"\n",
"def deploy_model_sglang_multihost(\n",
" model_name: str,\n",
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str = \"\",\n",
" base_model_id: str = \"\",\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" multihost_gpu_node_count: int = 1,\n",
" gpu_memory_utilization: float | None = None,\n",
" context_length: int | None = None,\n",
" dtype: str | None = None,\n",
" enable_trust_remote_code: bool = False,\n",
" enable_torch_compile: bool = False,\n",
" torch_compile_max_bs: int | None = None,\n",
" attention_backend: str = \"\",\n",
" enable_flashinfer_mla: bool = False,\n",
" disable_cuda_graph: bool = False,\n",
" speculative_algorithm: str | None = None,\n",
" speculative_draft_model_path: str = \"\",\n",
" speculative_num_steps: int = 3,\n",
" speculative_eagle_topk: int = 1,\n",
" speculative_num_draft_tokens: int = 4,\n",
" enable_jit_deepgemm: bool = False,\n",
" enable_dp_attention: bool = False,\n",
" dp_size: int = 1,\n",
" enable_multimodal: bool = False,\n",
" use_dedicated_endpoint: bool = False,\n",
" max_num_seqs: int | None = None,\n",
" is_spot: bool = True,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys trained models with SGLang into Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
"\n",
" if not base_model_id:\n",
" base_model_id = model_id\n",
"\n",
" # See https://docs.sglang.ai/backend/server_arguments.html for a list of possible arguments with descriptions.\n",
" sglang_args = [\n",
" f\"--model={model_id}\",\n",
" f\"--tp={accelerator_count * multihost_gpu_node_count}\",\n",
" f\"--dp={dp_size}\",\n",
" ]\n",
"\n",
" if context_length:\n",
" sglang_args.append(f\"--context-length={context_length}\")\n",
"\n",
" if gpu_memory_utilization:\n",
" sglang_args.append(f\"--mem-fraction-static={gpu_memory_utilization}\")\n",
"\n",
" if max_num_seqs:\n",
" sglang_args.append(f\"--max-running-requests={max_num_seqs}\")\n",
"\n",
" if dtype:\n",
" sglang_args.append(f\"--dtype={dtype}\")\n",
"\n",
" if enable_trust_remote_code:\n",
" sglang_args.append(\"--trust-remote-code\")\n",
"\n",
" if enable_torch_compile:\n",
" sglang_args.append(\"--enable-torch-compile\")\n",
" if torch_compile_max_bs:\n",
" sglang_args.append(f\"--torch-compile-max-bs={torch_compile_max_bs}\")\n",
"\n",
" if attention_backend:\n",
" sglang_args.append(f\"--attention-backend={attention_backend}\")\n",
"\n",
" if enable_flashinfer_mla:\n",
" sglang_args.append(\"--enable-flashinfer-mla\")\n",
"\n",
" if disable_cuda_graph:\n",
" sglang_args.append(\"--disable-cuda-graph\")\n",
"\n",
" if speculative_algorithm:\n",
" sglang_args.append(f\"--speculative-algorithm={speculative_algorithm}\")\n",
" sglang_args.append(\n",
" f\"--speculative-draft-model-path={speculative_draft_model_path}\"\n",
" )\n",
" sglang_args.append(f\"--speculative-num-steps={speculative_num_steps}\")\n",
" sglang_args.append(f\"--speculative-eagle-topk={speculative_eagle_topk}\")\n",
" sglang_args.append(\n",
" f\"--speculative-num-draft-tokens={speculative_num_draft_tokens}\"\n",
" )\n",
"\n",
" if enable_dp_attention:\n",
" sglang_args.append(\"--enable-dp-attention\")\n",
"\n",
" if enable_multimodal:\n",
" sglang_args.append(\"--enable-multimodal\")\n",
"\n",
" env_vars = {\n",
" \"MODEL_ID\": base_model_id,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" }\n",
"\n",
" if enable_jit_deepgemm:\n",
" env_vars[\"SGL_ENABLE_JIT_DEEPGEMM\"] = \"1\"\n",
"\n",
" # HF_TOKEN is not a compulsory field and may not be defined.\n",
" try:\n",
" if HF_TOKEN:\n",
" env_vars[\"HF_TOKEN\"] = HF_TOKEN\n",
" except NameError:\n",
" pass\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=SGLANG_DOCKER_URI,\n",
" serving_container_args=sglang_args,\n",
" serving_container_ports=[30000],\n",
" serving_container_predict_route=\"/vertex_generate\",\n",
" serving_container_health_route=\"/health\",\n",
" serving_container_environment_variables=env_vars,\n",
" serving_container_shared_memory_size_mb=(16 * 1024), # 16 GB\n",
" serving_container_deployment_timeout=7200,\n",
" model_garden_source_model_name=(\n",
" f\"publishers/{publisher}/models/{publisher_model_id}\"\n",
" ),\n",
" )\n",
" print(\n",
" f\"Deploying {model_name} on {machine_type} with {int(accelerator_count * multihost_gpu_node_count)} {accelerator_type} GPU(s).\"\n",
" )\n",
"\n",
" creds, _ = auth.default()\n",
" auth_req = auth.transport.requests.Request()\n",
" creds.refresh(auth_req)\n",
"\n",
" url = f\"https://{REGION}-aiplatform.googleapis.com/ui/projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint.name}:deployModel\"\n",
" headers = {\n",
" \"Content-Type\": \"application/json\",\n",
" \"Authorization\": f\"Bearer {creds.token}\",\n",
" }\n",
" data = {\n",
" \"deployedModel\": {\n",
" \"model\": model.resource_name,\n",
" \"displayName\": model_name,\n",
" \"dedicatedResources\": {\n",
" \"machineSpec\": {\n",
" \"machineType\": machine_type,\n",
" \"multihostGpuNodeCount\": multihost_gpu_node_count,\n",
" \"acceleratorType\": accelerator_type,\n",
" \"acceleratorCount\": accelerator_count,\n",
" },\n",
" \"minReplicaCount\": 1,\n",
" \"maxReplicaCount\": 1,\n",
" },\n",
" \"system_labels\": {\n",
" \"NOTEBOOK_NAME\": \"model_garden_gemma3n_deployment_on_vertex.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": common_util.get_deploy_source(),\n",
" },\n",
" },\n",
" }\n",
" if service_account:\n",
" data[\"deployedModel\"][\"serviceAccount\"] = service_account\n",
" if is_spot:\n",
" data[\"deployedModel\"][\"dedicatedResources\"][\"spot\"] = True\n",
" response = requests.post(url, headers=headers, json=data)\n",
" print(f\"Deploy Model response: {response.json()}\")\n",
" if response.status_code != 200 or \"name\" not in response.json():\n",
" raise ValueError(f\"Failed to deploy model: {response.text}\")\n",
" poll_and_wait(response.json()[\"name\"], 7200)\n",
" print(\"endpoint_name:\", endpoint.name)\n",
"\n",
" return model, endpoint\n",
"\n",
"\n",
"models[LABEL], endpoints[LABEL] = deploy_model_sglang_multihost(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=\"gemma3n-serve\"),\n",
" model_id=model_id,\n",
" publisher=\"google\",\n",
" publisher_model_id=\"gemma3n\",\n",
" base_model_id=hf_model_id,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" attention_backend=\"fa3\",\n",
" enable_multimodal=True,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" is_spot=is_spot,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "AGVPzwHkn7rw"
},
"outputs": [],
"source": [
"# @title Raw predict\n",
"\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with text prompts. Sampling parameters supported by SGLang can be found [here](https://docs.sglang.ai/backend/sampling_params.html).\n",
"\n",
"# @markdown Example:\n",
"\n",
"# @markdown ```\n",
"# @markdown User: What is the best way to diagnose and fix a flickering light in my house?\n",
"# @markdown Assistant: Okay, so I need to figure out how to diagnose and fix a flickering light in my house. Hmm, where do I start? Let's think. First, I remember that flickering lights can be caused by various issues. Maybe the bulb is loose? That's a common problem. Let me start with the simplest things first.\n",
"# @markdown ```\n",
"# @markdown Additionally, you can moderate the generated text with Vertex AI. See [Moderate text documentation](https://cloud.google.com/natural-language/docs/moderating-text) for more details.\n",
"\n",
"# Loads an existing endpoint instance using the endpoint name:\n",
"# - Using `endpoint_name = endpoint.name` allows us to get the\n",
"# endpoint name of the endpoint `endpoint` created in the cell\n",
"# above.\n",
"# - Alternatively, you can set `endpoint_name = \"1234567890123456789\"` to load\n",
"# an existing endpoint with the ID 1234567890123456789.\n",
"# You may uncomment the code below to load an existing endpoint.\n",
"\n",
"# endpoint_name = \"\" # @param {type:\"string\"}\n",
"# aip_endpoint_name = (\n",
"# f\"projects/{PROJECT_ID}/locations/{REGION}/endpoints/{endpoint_name}\"\n",
"# )\n",
"# endpoint = aiplatform.Endpoint(aip_endpoint_name)\n",
"\n",
"# @markdown A chat template formatted prompt for Gemma 3n is shown below as an example.\n",
"prompt = \"<start_of_turn>user\\nWhat is a car?<end_of_turn>\\n<start_of_turn>model\\n\" # @param {type: \"string\"}\n",
"\n",
"max_new_tokens = 128 # @param {type:\"integer\"}\n",
"temperature = 0.95 # @param {type:\"number\"}\n",
"top_k = 64 # @param {type:\"number\"}\n",
"\n",
"# Overrides parameters for inferences.\n",
"instances = [{\"text\": prompt}]\n",
"parameters = {\n",
" \"sampling_params\": {\n",
" \"max_new_tokens\": max_new_tokens,\n",
" \"temperature\": temperature,\n",
" \"top_k\": top_k,\n",
" }\n",
"}\n",
"response = endpoints[\"sglang_gpu\"].predict(\n",
" instances=instances,\n",
" parameters=parameters,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"for prediction in response.predictions:\n",
" print(prediction)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "ZauMzfXJzAKZ"
},
"outputs": [],
"source": [
"# @title Chat completion with text-only requests\n",
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"sglang_gpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"sglang_gpu\"].resource_name\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint using the OpenAI SDK.\n",
"\n",
"# @markdown First you will need to install the SDK and some auth-related dependencies.\n",
"\n",
"! pip install -qU openai google-auth requests\n",
"\n",
"# @markdown Next fill out some request parameters:\n",
"\n",
"user_message = \"How is your day going?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter the issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, such as set `max_tokens` as 20.\n",
"max_tokens = 50 # @param {type: \"integer\"}\n",
"temperature = 1.0 # @param {type: \"number\"}\n",
"stream = False # @param {type: \"boolean\"}\n",
"\n",
"# @markdown Now we can send a request.\n",
"\n",
"import google.auth\n",
"import openai\n",
"\n",
"creds, project = google.auth.default()\n",
"auth_req = google.auth.transport.requests.Request()\n",
"creds.refresh(auth_req)\n",
"\n",
"BASE_URL = (\n",
" f\"https://{REGION}-aiplatform.googleapis.com/v1beta1/{ENDPOINT_RESOURCE_NAME}\"\n",
")\n",
"try:\n",
" if use_dedicated_endpoint:\n",
" BASE_URL = f\"https://{DEDICATED_ENDPOINT_DNS}/v1beta1/{ENDPOINT_RESOURCE_NAME}\"\n",
"except NameError:\n",
" pass\n",
"\n",
"client = openai.OpenAI(base_url=BASE_URL, api_key=creds.token)\n",
"\n",
"model_response = client.chat.completions.create(\n",
" model=\"\",\n",
" messages=[{\"role\": \"user\", \"content\": user_message}],\n",
" temperature=temperature,\n",
" max_tokens=max_tokens,\n",
" stream=stream,\n",
")\n",
"\n",
"if stream:\n",
" usage = None\n",
" contents = []\n",
" for chunk in model_response:\n",
" if chunk.usage is not None:\n",
" usage = chunk.usage\n",
" continue\n",
" print(chunk.choices[0].delta.content, end=\"\")\n",
" contents.append(chunk.choices[0].delta.content)\n",
" print(f\"\\n\\n{usage}\")\n",
"else:\n",
" print(model_response)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "baa783359502"
},
"outputs": [],
"source": [
"# @title Chat completion with text+image requests\n",
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"sglang_gpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"sglang_gpu\"].resource_name\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint using the OpenAI SDK.\n",
"\n",
"# @markdown First you will need to install the SDK and some auth-related dependencies.\n",
"\n",
"! pip install -qU openai google-auth requests\n",
"\n",
"# @markdown Next fill out some request parameters:\n",
"\n",
"user_image = \"https://upload.wikimedia.org/wikipedia/commons/thumb/d/dd/Gfp-wisconsin-madison-the-nature-boardwalk.jpg/2560px-Gfp-wisconsin-madison-the-nature-boardwalk.jpg\"\n",
"user_message = \"What is in the image?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter the issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, such as set `max_tokens` as 20.\n",
"max_tokens = 50 # @param {type: \"integer\"}\n",
"temperature = 1.0 # @param {type: \"number\"}\n",
"\n",
"# @markdown Now we can send a request.\n",
"\n",
"import google.auth\n",
"import openai\n",
"\n",
"creds, project = google.auth.default()\n",
"auth_req = google.auth.transport.requests.Request()\n",
"creds.refresh(auth_req)\n",
"\n",
"BASE_URL = (\n",
" f\"https://{REGION}-aiplatform.googleapis.com/v1beta1/{ENDPOINT_RESOURCE_NAME}\"\n",
")\n",
"try:\n",
" if use_dedicated_endpoint:\n",
" BASE_URL = f\"https://{DEDICATED_ENDPOINT_DNS}/v1beta1/{ENDPOINT_RESOURCE_NAME}\"\n",
"except NameError:\n",
" pass\n",
"\n",
"client = openai.OpenAI(base_url=BASE_URL, api_key=creds.token)\n",
"\n",
"model_response = client.chat.completions.create(\n",
" model=\"\",\n",
" messages=[\n",
" {\n",
" \"role\": \"user\",\n",
" \"content\": [\n",
" {\"type\": \"image_url\", \"image_url\": {\"url\": user_image}},\n",
" {\"type\": \"text\", \"text\": user_message},\n",
" ],\n",
" }\n",
" ],\n",
" temperature=temperature,\n",
" max_tokens=max_tokens,\n",
")\n",
"print(model_response)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "9b74bb8e4e5f"
},
"outputs": [],
"source": [
"# @title Chat completion with text+audio requests\n",
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"sglang_gpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"sglang_gpu\"].resource_name\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint using the OpenAI SDK.\n",
"\n",
"# @markdown First you will need to install the SDK and some auth-related dependencies.\n",
"\n",
"! pip install -qU openai google-auth requests\n",
"\n",
"# @markdown Next fill out some request parameters:\n",
"\n",
"user_audio = \"https://freewavesamples.com/files/Cat-Meow.wav\" # @param {type: \"string\"}\n",
"user_message = \"What animal is making the sound?\" # @param {type: \"string\"}\n",
"# @markdown If you encounter the issue like `ServiceUnavailable: 503 Took too long to respond when processing`, you can reduce the maximum number of output tokens, such as set `max_tokens` as 20.\n",
"max_tokens = 50 # @param {type: \"integer\"}\n",
"temperature = 0.95 # @param {type: \"number\"}\n",
"\n",
"# @markdown Now we can send a request.\n",
"\n",
"import google.auth\n",
"import openai\n",
"\n",
"creds, project = google.auth.default()\n",
"auth_req = google.auth.transport.requests.Request()\n",
"creds.refresh(auth_req)\n",
"\n",
"BASE_URL = (\n",
" f\"https://{REGION}-aiplatform.googleapis.com/v1beta1/{ENDPOINT_RESOURCE_NAME}\"\n",
")\n",
"try:\n",
" if use_dedicated_endpoint:\n",
" BASE_URL = f\"https://{DEDICATED_ENDPOINT_DNS}/v1beta1/{ENDPOINT_RESOURCE_NAME}\"\n",
"except NameError:\n",
" pass\n",
"\n",
"client = openai.OpenAI(base_url=BASE_URL, api_key=creds.token)\n",
"\n",
"model_response = client.chat.completions.create(\n",
" model=\"\",\n",
" messages=[\n",
" {\n",
" \"role\": \"user\",\n",
" \"content\": [\n",
" {\"type\": \"audio_url\", \"audio_url\": {\"url\": user_audio}},\n",
" {\"type\": \"text\", \"text\": user_message},\n",
" ],\n",
" }\n",
" ],\n",
" temperature=temperature,\n",
" max_tokens=max_tokens,\n",
")\n",
"print(model_response)\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "JETd33jIDcjm"
},
"source": [
"## Clean up resources"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "911406c1561e"
},
"outputs": [],
"source": [
"# @title Delete the models and endpoints\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
"# Undeploy model and delete endpoint.\n",
"for endpoint in endpoints.values():\n",
" endpoint.delete(force=True)\n",
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()"
]
}
],
"metadata": {
"colab": {
"name": "model_garden_gemma3n_deployment_on_vertex.ipynb",
"toc_visible": true
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
}
},
"nbformat": 4,
"nbformat_minor": 0
}
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma_deployment_on_gke.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma_deployment_on_vertex.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -199,7 +199,7 @@
"# @markdown ---\n",
"\n",
"# @markdown ### Access Gemma models on Hugging Face\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Gemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"# @markdown You must provide a Hugging Face User Access Token (read) to access the Gemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"if LOAD_MODEL_FROM == \"Hugging Face\":\n",
@@ -303,7 +303,7 @@
"hbm_utilization_factor = 0.6 # A larger value improves throughput but gives higher risk of TPU out-of-memory errors with long prompts.\n",
"max_running_seqs = 256\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# Endpoint configurations.\n",
@@ -325,7 +325,6 @@
" disagg_topology: str = None,\n",
" hbm_utilization_factor: float = 0.6,\n",
" max_running_seqs: int = 256,\n",
" decode_seqs_padding: int = None,\n",
" max_model_len: int = 4096,\n",
" enable_prefix_cache_hbm: bool = False,\n",
" endpoint_id: str = \"\",\n",
@@ -366,10 +365,6 @@
" f\"--max_running_seqs={max_running_seqs}\",\n",
" f\"--max_model_len={max_model_len}\",\n",
" ]\n",
"\n",
" if decode_seqs_padding is not None:\n",
" hexllm_args.append(f\"--decode_seqs_padding={decode_seqs_padding}\")\n",
"\n",
" if disagg_topology:\n",
" hexllm_args.append(f\"--disagg_topo={disagg_topology}\")\n",
" if enable_prefix_cache_hbm and not disagg_topology:\n",
@@ -518,7 +513,9 @@
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"hexllm_tpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"hexllm_tpu\"].resource_name\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"hexllm_tpu\"].name\n",
")\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
@@ -678,7 +675,7 @@
"# Note that a larger max_model_len will require more GPU memory.\n",
"max_model_len = 2048\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"\n",
@@ -912,7 +909,9 @@
"\n",
"if use_dedicated_endpoint:\n",
" DEDICATED_ENDPOINT_DNS = endpoints[\"vllm_gpu\"].gca_resource.dedicated_endpoint_dns\n",
"ENDPOINT_RESOURCE_NAME = endpoints[\"vllm_gpu\"].resource_name\n",
"ENDPOINT_RESOURCE_NAME = \"projects/{}/locations/{}/endpoints/{}\".format(\n",
" PROJECT_ID, REGION, endpoints[\"vllm_gpu\"].name\n",
")\n",
"\n",
"# @title Chat Completions Inference\n",
"\n",
@@ -41,7 +41,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gemma_evaluation.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -203,7 +203,7 @@
"\n",
"# @markdown This section demonstrates how to evaluate the Gemma models with and without finetuned LoRA adapters using EleutherAI's [Language Model Evaluation Harness (lm-evaluation-harness)](https://github.com/EleutherAI/lm-evaluation-harness) with Vertex CustomJob. Refer the peak GPU memory usage for serving and adjust the machine type, accelerator type and accelerator count accordingly.\n",
"\n",
"# @markdown You must provide a Hugging Face User Access Token (with read access) to access the Gemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"# @markdown You must provide a Hugging Face User Access Token (read) to access the Gemma models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate:true}\n",
"\n",
"# @markdown This example uses the dataset [HellaSwag](https://arxiv.org/abs/1905.07830). All supported tasks are listed in [this task table](https://github.com/EleutherAI/lm-evaluation-harness/blob/master/docs/task_table.md).\n",
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_gradio_streaming_chat_completions.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -316,7 +316,7 @@
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str = None,\n",
" service_account: str,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
@@ -347,9 +347,6 @@
" except NameError:\n",
" pass\n",
"\n",
" if service_account:\n",
" env_vars[\"SERVICE_ACCOUNT\"] = service_account\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=TGI_DOCKER_URI,\n",
File diff suppressed because one or more lines are too long
@@ -34,18 +34,13 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_hf_paligemma2_deployment.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_hf_paligemma2_deployment.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_hf_paligemma2_deployment.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -68,10 +63,6 @@
"- Make predictions to the endpoint including:\n",
" - Answering questions about a given image.\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
@@ -102,6 +93,16 @@
"source": [
"# @title Setup Google Cloud project\n",
"\n",
"# Used for common utilities.\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"import importlib\n",
"# Import the necessary packages\n",
"import os\n",
"from typing import Any, Dict, Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
@@ -117,27 +118,6 @@
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"\n",
"import importlib\n",
"# Import the necessary packages\n",
"import os\n",
"from typing import Any, Dict, Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
"\n",
"if os.environ.get(\"VERTEX_PRODUCT\") != \"COLAB_ENTERPRISE\":\n",
" ! pip install --upgrade tensorflow\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"models, endpoints = {}, {}\n",
"LABEL = \"paligemma2\"\n",
"\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
@@ -146,118 +126,27 @@
"if not REGION:\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"import vertexai\n",
"\n",
"vertexai.init(\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
")"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "kyMJXkfviWgl"
},
"source": [
"## Deploy Model to a Vertex AI Endpoint"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "toY-WPKDFesF"
},
"outputs": [],
"source": [
"# @title Deploy\n",
"\n",
"MODEL_NAME = \"paligemma2-3b-pt-224\" # @param [\"paligemma2-3b-pt-224\", \"paligemma2-3b-mix-224\", \"paligemma2-3b-ft-docci-448\", \"paligemma2-3b-mix-448\", \"paligemma2-3b-pt-448\", \"paligemma2-3b-pt-896\", \"paligemma2-10b-mix-224\", \"paligemma2-10b-pt-224\", \"paligemma2-10b-ft-docci-448\", \"paligemma2-10b-mix-448\", \"paligemma2-10b-pt-448\", \"paligemma2-10b-pt-896\", \"paligemma2-28b-mix-224\", \"paligemma2-28b-pt-224\", \"paligemma2-28b-mix-448\", \"paligemma2-28b-pt-448\", \"paligemma2-28b-pt-896\"]\n",
"GCS_PREFIX = \"gs://vertex-model-garden-restricted-us/paligemma2\"\n",
"\n",
"MODEL_ID = os.path.join(GCS_PREFIX, MODEL_NAME)\n",
"\n",
"PUBLISHER_MODEL_NAME = f\"publishers/google/models/paligemma@{MODEL_NAME}\"\n",
"\n",
"\n",
"# @markdown If you want to use other accelerator types not listed above, then check other Vertex AI prediction supported accelerators and regions at https://cloud.google.com/vertex-ai/docs/predictions/configure-compute. You may need to manually set the `machine_type`, `accelerator_type`, and `accelerator_count` in the code by clicking `Show code` first.\n",
"\n",
"if \"3b\" in MODEL_NAME:\n",
" accelerator_type = \"NVIDIA_L4\"\n",
" machine_type = \"g2-standard-16\"\n",
" accelerator_count = 1\n",
"elif \"10b\" in MODEL_NAME:\n",
" accelerator_type = \"NVIDIA_TESLA_A100\"\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
"elif \"28b\" in MODEL_NAME:\n",
" accelerator_type = \"NVIDIA_H100_80GB\"\n",
" machine_type = \"a3-highgpu-8g\"\n",
" accelerator_count = 8\n",
"else:\n",
" raise ValueError(f\"Recommended GPU setting not found for: {MODEL_NAME}.\")\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "pe_qbTCA6nKf"
},
"outputs": [],
"source": [
"# @title [Option 1] Deploy with Model Garden SDK\n",
"\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"from vertexai import model_garden\n",
"\n",
"model = model_garden.OpenModel(PUBLISHER_MODEL_NAME)\n",
"endpoints[LABEL] = model.deploy(\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "jbeLl-9C6nKf"
},
"outputs": [],
"source": [
"# @title [Option 2] Deploy with customized configs\n",
"\n",
"# @markdown This section uploads the prebuilt PaliGemma 2 models to Model Registry and deploys it to a Vertex AI Endpoint. It takes approximately 15 minutes to finish.\n",
"\n",
"# @markdown Select the desired resolution and precision of prebuilt model to deploy, leaving the optional `custom_paligemma_model_uri` as is. Higher resolution and precision_type can result in better inference results, but may require additional GPU.\n",
"\n",
"TASK = \"paligemma_VQA\"\n",
"models, endpoints = {}, {}\n",
"\n",
"# The pre-built serving docker images.\n",
"SERVE_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-one-serve:20250205_0822_RC00\"\n",
@@ -270,6 +159,7 @@
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" service_account: str = None,\n",
" serving_port: int = 7080,\n",
" serving_route: str = \"/predict\",\n",
" serving_docker_uri: str = SERVE_DOCKER_URI,\n",
@@ -283,6 +173,7 @@
" machine_type: The machine type.\n",
" accelerator_type: The accelerator type.\n",
" accelerator_count: The accelerator count.\n",
" service_account: The service account.\n",
" serving_port: The serving port.\n",
" serving_route: The serving route.\n",
" hf_token: HuggingFace token for model access.\n",
@@ -313,15 +204,111 @@
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" service_account=service_account,\n",
" sync=False,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_hf_paligemma2_deployment.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": common_util.get_deploy_source(),\n",
" },\n",
" system_labels={\"NOTEBOOK_NAME\": \"model_garden_hf_paligemma2_deployment.ipynb\"},\n",
" )\n",
" return endpoint, model\n",
"\n",
"\n",
"def vqa_predict(\n",
" endpoint: aiplatform.Endpoint,\n",
" image_url: str,\n",
" text_prompt: str,\n",
" parameters: Dict[str, Any] = None,\n",
") -> str:\n",
" \"\"\"Predicts the answer to a question about an image using an Endpoint,\n",
"\n",
" and passes parameters in the payload.\n",
"\n",
" Args:\n",
" endpoint: The deployed Vertex AI endpoint.\n",
" image_url: URL of the image to ask about.\n",
" text_prompt: The text prompt question.\n",
" parameters: Additional parameters for the prediction request.\n",
"\n",
" Returns:\n",
" The predicted answer string or None if no prediction.\n",
" \"\"\"\n",
"\n",
" instances = []\n",
" if text_prompt:\n",
" instances.append(\n",
" {\n",
" \"text_prompt\": text_prompt,\n",
" \"image_url\": image_url,\n",
" }\n",
" )\n",
"\n",
" # Construct the prediction payload\n",
" payload = {\"instances\": instances}\n",
" if parameters:\n",
" payload[\"parameters\"] = parameters\n",
"\n",
" response = endpoint.predict(instances=instances, parameters=parameters)\n",
" answer = None\n",
" if response.predictions:\n",
" answer = response.predictions[0][\"text\"].split(\"\\n\")[1]\n",
" return answer"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "kyMJXkfviWgl"
},
"source": [
"## Deploy Model to a Vertex AI Endpoint"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "toY-WPKDFesF"
},
"outputs": [],
"source": [
"# @title Deploy\n",
"\n",
"# @markdown This section uploads the prebuilt PaliGemma 2 models to Model Registry and deploys it to a Vertex AI Endpoint. It takes approximately 15 minutes to finish.\n",
"\n",
"# @markdown Select the desired resolution and precision of prebuilt model to deploy, leaving the optional `custom_paligemma_model_uri` as is. Higher resolution and precision_type can result in better inference results, but may require additional GPU.\n",
"\n",
"MODEL_NAME = \"paligemma2-3b-pt-224\" # @param [\"paligemma2-3b-pt-224\", \"paligemma2-3b-mix-224\", \"paligemma2-3b-ft-docci-448\", \"paligemma2-3b-mix-448\", \"paligemma2-3b-pt-448\", \"paligemma2-3b-pt-896\", \"paligemma2-10b-mix-224\", \"paligemma2-10b-pt-224\", \"paligemma2-10b-ft-docci-448\", \"paligemma2-10b-mix-448\", \"paligemma2-10b-pt-448\", \"paligemma2-10b-pt-896\", \"paligemma2-28b-mix-224\", \"paligemma2-28b-pt-224\", \"paligemma2-28b-mix-448\", \"paligemma2-28b-pt-448\", \"paligemma2-28b-pt-896\"]\n",
"GCS_PREFIX = \"gs://vertex-model-garden-restricted-us/paligemma2\"\n",
"\n",
"MODEL_ID = os.path.join(GCS_PREFIX, MODEL_NAME)\n",
"\n",
"\n",
"# @markdown If you want to use other accelerator types not listed above, then check other Vertex AI prediction supported accelerators and regions at https://cloud.google.com/vertex-ai/docs/predictions/configure-compute. You may need to manually set the `machine_type`, `accelerator_type`, and `accelerator_count` in the code by clicking `Show code` first.\n",
"\n",
"if \"3b\" in MODEL_NAME:\n",
" accelerator_type = \"NVIDIA_L4\"\n",
" machine_type = \"g2-standard-8\"\n",
" accelerator_count = 1\n",
"elif \"10b\" in MODEL_NAME:\n",
" accelerator_type = \"NVIDIA_TESLA_A100\"\n",
" machine_type = \"a2-highgpu-1g\"\n",
" accelerator_count = 1\n",
"elif \"28b\" in MODEL_NAME:\n",
" accelerator_type = \"NVIDIA_H100_80GB\"\n",
" machine_type = \"a3-highgpu-8g\"\n",
" accelerator_count = 8\n",
"else:\n",
" raise ValueError(f\"Recommended GPU setting not found for: {MODEL_NAME}.\")\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"TASK = \"paligemma_VQA\"\n",
"\n",
"endpoints[\"paligemma2\"], models[\"paligemma2\"] = deploy_model(\n",
" model_name=MODEL_NAME,\n",
" model_id=MODEL_ID,\n",
@@ -329,6 +316,7 @@
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" service_account=SERVICE_ACCOUNT,\n",
" serving_port=7080,\n",
" serving_route=\"/predict\",\n",
" serving_docker_uri=SERVE_DOCKER_URI,\n",
@@ -391,48 +379,6 @@
"\n",
"# @markdown The question prompt can be non-English languages.\n",
"\n",
"\n",
"def vqa_predict(\n",
" endpoint: aiplatform.Endpoint,\n",
" image_url: str,\n",
" text_prompt: str,\n",
" parameters: Dict[str, Any] = None,\n",
") -> str:\n",
" \"\"\"Predicts the answer to a question about an image using an Endpoint,\n",
"\n",
" and passes parameters in the payload.\n",
"\n",
" Args:\n",
" endpoint: The deployed Vertex AI endpoint.\n",
" image_url: URL of the image to ask about.\n",
" text_prompt: The text prompt question.\n",
" parameters: Additional parameters for the prediction request.\n",
"\n",
" Returns:\n",
" The predicted answer string or None if no prediction.\n",
" \"\"\"\n",
"\n",
" instances = []\n",
" if text_prompt:\n",
" instances.append(\n",
" {\n",
" \"text_prompt\": text_prompt,\n",
" \"image_url\": image_url,\n",
" }\n",
" )\n",
"\n",
" # Construct the prediction payload\n",
" payload = {\"instances\": instances}\n",
" if parameters:\n",
" payload[\"parameters\"] = parameters\n",
"\n",
" response = endpoint.predict(instances=instances, parameters=parameters)\n",
" answer = None\n",
" if response.predictions:\n",
" answer = response.predictions[0][\"text\"].split(\"\\n\")[1]\n",
" return answer\n",
"\n",
"\n",
"# Using max_new_tokens along with other parameters\n",
"parameters_with_tokens = {\"max_new_tokens\": 50}\n",
"predictions_with_tokens = vqa_predict(\n",
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_huggingface_local_inference.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_huggingface_pytorch_inference_deployment.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -138,10 +138,8 @@
"! gcloud config set project $PROJECT_ID\n",
"\n",
"HF_TOKEN = \"\"\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"# Dedicated endpoint not supported yet\n",
"use_dedicated_endpoint = False\n",
"SERVICE_ACCOUNT = \"\""
]
},
@@ -163,7 +161,7 @@
"TASK = \"text-classification\" # @param {type: \"string\", isTemplate: true}\n",
"\n",
"# The pre-built serving docker images for Hugging Face Pytorch Inference.\n",
"SERVE_DOCKER_URI = \"us-docker.pkg.dev/deeplearning-platform-release/vertex-model-garden/hf-inference-toolkit.cu125.0-1.ubuntu2204.py311\"\n",
"SERVE_DOCKER_URI = \"us-docker.pkg.dev/deeplearning-platform-release/gcr.io/huggingface-pytorch-inference-cu121.2-3.transformers.4-46.ubuntu2204.py311\"\n",
"\n",
"machine_type = \"g2-standard-8\" # @param {type: \"string\", isTemplate: true}\n",
"accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"None\"] {isTemplate: true}\n",
@@ -34,18 +34,13 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_huggingface_tei_deployment.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_huggingface_tei_deployment.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_huggingface_tei_deployment.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -67,10 +62,6 @@
"- Download and deploy the `nomic-ai/nomic-embed-text-v1` model with TEI\n",
"- Send prediction request to the deployed endpoint\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
@@ -102,30 +93,16 @@
"\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"# @markdown 2. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 4. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus). You can request for quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.64.0'\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"import datetime\n",
"import importlib\n",
"import os\n",
"import uuid\n",
"from typing import Tuple\n",
"\n",
"from google.cloud import aiplatform\n",
@@ -151,67 +128,31 @@
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"tei\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"aiplatform.init(project=PROJECT_ID, location=REGION)\n",
"! gcloud config set project $PROJECT_ID\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\"\n",
"\n",
"import vertexai\n",
"models, endpoints = {}, {}\n",
"\n",
"vertexai.init(\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
")\n",
"HF_TOKEN = \"\""
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "USB7dvYqvNdu"
},
"outputs": [],
"source": [
"# @title Deploy with TEI from Hugging Face\n",
"\n",
"# @markdown Set Hugging Face access token. It is strongly recommended for\n",
"# @markdown mitigating model artifact download errors. Hugging Face has recently\n",
"# @markdown enforced rate limits for anonymous callers, which can disrupt model\n",
"# @markdown artifact downloads. See [this discussion](https://discuss.huggingface.co/t/hugging-face-api-rate-limits/16746)\n",
"# @markdown for additional details.\n",
"# @markdown This section downloads the `nomic-ai/nomic-embed-text-v1` model from Hugging Face and deploys it to a Vertex AI Endpoint.\n",
"# @markdown It takes ~20 minutes to complete the deployment.\n",
"\n",
"HF_TOKEN = \"\" # @param {type: \"string\", isTemplate: true}\n",
"\n",
"# @markdown Set Hugging Face model id and deployment configs.\n",
"\n",
"HUGGING_FACE_MODEL_ID = \"nomic-ai/nomic-embed-text-v1\" # @param {type: \"string\", isTemplate: true}\n",
"MODEL_ID = \"nomic-ai/nomic-embed-text-v1\" # @param {type: \"string\", isTemplate: true}\n",
"\n",
"# The pre-built serving docker images for TEI.\n",
"TEI_CPU_DOCKER_URI = \"us-docker.pkg.dev/deeplearning-platform-release/gcr.io/huggingface-text-embeddings-inference-cpu.1-4\"\n",
@@ -232,136 +173,6 @@
" is_for_training=False,\n",
" )\n",
"\n",
"LABEL = \"tei\"\n",
"models, endpoints = {}, {}\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "obbeTtMJ5C8j"
},
"outputs": [],
"source": [
"# @title [Option 1] Deploy with Model Garden SDK\n",
"\n",
"accelerator_count = 1\n",
"\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"from vertexai import model_garden\n",
"\n",
"model = model_garden.OpenModel(HUGGING_FACE_MODEL_ID)\n",
"endpoints[LABEL] = model.deploy(\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" hugging_face_access_token=HF_TOKEN,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "USB7dvYqvNdu"
},
"outputs": [],
"source": [
"# @title [Option 2] Deploy with customized configs\n",
"\n",
"# @markdown This section downloads the model from Hugging Face and deploys it to a Vertex AI Endpoint.\n",
"# @markdown It takes ~20 minutes to complete the deployment.\n",
"\n",
"\n",
"# @markdown This notebook downloads model artifacts from Hugging Face\n",
"# @markdown repository, and uploads them to a gcs bucket to avoid model download\n",
"# @markdown errors during deployment.\n",
"\n",
"# @markdown **[Optional]** If you want to skip uploading model, set `UPLOAD_MODEL_ARTIFACT_TO_GCS` to False.\n",
"\n",
"UPLOAD_MODEL_ARTIFACT_TO_GCS = True # @param {type:\"boolean\"}\n",
"\n",
"\n",
"import os\n",
"\n",
"from google.cloud import storage\n",
"from huggingface_hub import snapshot_download\n",
"\n",
"\n",
"def download_hugging_face_model_artifacts(model_id: str) -> str:\n",
" \"\"\"Downloads model artifacts from Hugging Face repository.\n",
"\n",
" Args:\n",
" model_id: The model ID to download.\n",
"\n",
" Returns:\n",
" The absolute path to the downloaded model .\n",
" \"\"\"\n",
" os.environ[\"HF_TOKEN\"] = HF_TOKEN\n",
"\n",
" folder_name = f\"{model_id.replace('/', '_')}_artifacts\"\n",
" local_dir = f\"./{folder_name}\"\n",
" print(f\"Downloading model '{model_id}' to '{local_dir}'\")\n",
"\n",
" downloaded_path = snapshot_download(\n",
" repo_id=model_id,\n",
" local_dir=local_dir,\n",
" local_dir_use_symlinks=False,\n",
" )\n",
" print(\"Download complete.\")\n",
" print(f\"Model artifacts saved to: {downloaded_path}\\n\")\n",
" return downloaded_path\n",
"\n",
"\n",
"def upload_model_artifacts_to_gcs(local_dir: str, bucket_uri: str) -> str:\n",
" \"\"\"Uploads model artifacts to a GCS bucket.\n",
"\n",
" Args:\n",
" local_dir: The absolute path to the local model dir.\n",
" bucket_name: The GCS bucket uri with \"gs://\" prefix.\n",
"\n",
" Returns:\n",
" The GCS uri to the model directory.\n",
" \"\"\"\n",
" assert bucket_uri.startswith(\"gs://\"), \"bucket_uri must start with `gs://`.\"\n",
"\n",
" storage_client = storage.Client()\n",
" bucket_name = bucket_uri[len(\"gs://\") :]\n",
" bucket = storage_client.bucket(bucket_name)\n",
" folder_name = os.path.basename(local_dir)\n",
"\n",
" print(\n",
" f\"Uploading model artifacts '{local_dir}' to GCS bucket 'gs://{bucket_name}/{folder_name}'\"\n",
" )\n",
"\n",
" for root, _, files in os.walk(local_dir):\n",
" for file_name in files:\n",
" absolute_path = os.path.join(root, file_name)\n",
" relative_path = os.path.relpath(absolute_path, local_dir)\n",
" blob_name = os.path.join(folder_name, relative_path).replace(os.sep, \"/\")\n",
" blob = bucket.blob(blob_name)\n",
" blob.upload_from_filename(absolute_path)\n",
"\n",
" print(\"Upload complete.\\n\")\n",
" return f\"gs://{bucket_name}/{folder_name}\"\n",
"\n",
"\n",
"if UPLOAD_MODEL_ARTIFACT_TO_GCS:\n",
" local_dir = download_hugging_face_model_artifacts(model_id=HUGGING_FACE_MODEL_ID)\n",
" aip_storage_uri = upload_model_artifacts_to_gcs(local_dir, BUCKET_URI)\n",
"else:\n",
" aip_storage_uri = \"\"\n",
"\n",
"\n",
"def deploy_model_tei(\n",
" model_name: str,\n",
@@ -372,7 +183,6 @@
" machine_type: str = \"g2-standard-4\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" use_dedicated_endpoint: bool = False,\n",
" aip_storage_uri: str = \"\",\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Deploys models with TEI on Vertex AI.\"\"\"\n",
" endpoint = aiplatform.Endpoint.create(\n",
@@ -385,7 +195,6 @@
" \"MODEL_ID\": model_id,\n",
" \"JSON_OUTPUT\": \"true\",\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" \"AIP_STORAGE_URI\": aip_storage_uri,\n",
" }\n",
"\n",
" # HF_TOKEN is not a compulsory field and may not be defined.\n",
@@ -420,16 +229,19 @@
" return model, endpoint\n",
"\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"\n",
"models[\"tei\"], endpoints[\"tei\"] = deploy_model_tei(\n",
" model_name=common_util.get_job_name_with_datetime(prefix=HUGGING_FACE_MODEL_ID),\n",
" model_id=HUGGING_FACE_MODEL_ID,\n",
" model_name=common_util.get_job_name_with_datetime(prefix=MODEL_ID),\n",
" model_id=MODEL_ID,\n",
" publisher=\"hf-nomic-ai\",\n",
" publisher_model_id=\"nomic-embed-text-v1\",\n",
" service_account=SERVICE_ACCOUNT,\n",
" service_account=\"\",\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" aip_storage_uri=aip_storage_uri,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
@@ -505,11 +317,7 @@
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
" model.delete()"
]
}
],
@@ -34,7 +34,7 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_huggingface_tgi_deployment.ipynb\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/instances\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
@@ -45,7 +45,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_huggingface_tgi_deployment.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -110,7 +110,7 @@
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.84.0'\n",
"\n",
"# @markdown 5. You must agree to the license on the [model card](https://huggingface.co/google/gemma-2-2b-it) before accessing the Gemma 2 models.\n",
"\n",
@@ -186,7 +186,7 @@
"SERVING_CONTAINER_IMAGE_URI = TGI_DOCKER_URI\n",
"LABEL = \"tgi\"\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}"
]
},
@@ -202,7 +202,7 @@
"# @title [Option 1] Deploy with Model Garden SDK\n",
"\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"from vertexai import model_garden\n",
"from vertexai.preview import model_garden\n",
"\n",
"model = model_garden.OpenModel(HUGGING_FACE_MODEL_ID)\n",
"endpoints[LABEL] = model.deploy(\n",
@@ -213,8 +213,6 @@
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
@@ -238,7 +236,7 @@
" model_id: str,\n",
" publisher: str,\n",
" publisher_model_id: str,\n",
" service_account: str = None,\n",
" service_account: str,\n",
" machine_type: str = \"g2-standard-8\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
@@ -269,9 +267,6 @@
" except NameError:\n",
" pass\n",
"\n",
" if service_account:\n",
" env_vars[\"SERVICE_ACCOUNT\"] = service_account\n",
"\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name,\n",
" serving_container_image_uri=TGI_DOCKER_URI,\n",
@@ -30,22 +30,22 @@
"id": "08f4AuF5eXzO"
},
"source": [
"# Vertex AI Model Garden - Hugging Face Deployment with vLLM Container\n",
"# Vertex AI Model Garden - Hugging Face Text Generation with vLLM Container Deployment\n",
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_huggingface_vllm_deployment.ipynb\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/instances\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_huggingface_vllm_deployment.ipynb\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_huggingface_tgi_vllm_deployment.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_huggingface_vllm_deployment.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_huggingface_tgi_vllm_deployment.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -59,13 +59,13 @@
"source": [
"## Overview\n",
"\n",
"This notebook demonstrates deploying [qwen/qwq-32b](https://huggingface.co/Qwen/QwQ-32B) model with vLLM container from Hugging Face. In additional to `qwen/qwq-32b`, You can view and change the code to deploy a different Hugging Face model with appropriate machine specs.\n",
"This notebook demonstrates deploying [qwen/qwq-32b](https://huggingface.co/Qwen/QwQ-32B) model with vLLM container from Hugging Face. In additional to `qwen/qwq-32b`, You can view and change the code to deploy a different Hugging Face `text-generation` model with appropriate machine specs. **Note that some models might fail to deploy, even if they have `text-generation` tags on the Hugging Face model card page.**\n",
"\n",
"\n",
"### Objective\n",
"\n",
"- Download and deploy the `qwen/qwq-32b` model with vLLM container.\n",
"- Send prediction request to the deployed endpoint.\n",
"- Download and deploy the `qwen/qwq-32b` model with TGI\n",
"- Send prediction request to the deployed endpoint\n",
"\n",
"### Costs\n",
"\n",
@@ -113,7 +113,7 @@
"\n",
"# @markdown 4. Follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below.\n",
"\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform>=1.84.0'\n",
"\n",
"import importlib\n",
"import os\n",
@@ -163,8 +163,8 @@
"\n",
"HF_TOKEN = \"\" # @param {type:\"string\", isTemplate: true}\n",
"\n",
"# The pre-built vLLM serving docker image.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/pytorch-vllm-serve:20250506_0916_RC01\"\n",
"# The pre-built serving docker image for TGI with vLLM.\n",
"VLLM_DOCKER_URI = \"us-docker.pkg.dev/deeplearning-platform-release/vertex-model-garden/vllm-inference.cu121.0-6.ubuntu2204.py310\"\n",
"SERVING_CONTAINER_IMAGE_URI = VLLM_DOCKER_URI\n",
"LABEL = \"vllm\"\n",
"\n",
@@ -222,7 +222,7 @@
" is_for_training=False,\n",
")\n",
"\n",
"from vertexai import model_garden\n",
"from vertexai.preview import model_garden\n",
"\n",
"model = model_garden.OpenModel(HUGGING_FACE_MODEL_ID)\n",
"endpoints[LABEL] = model.deploy(\n",
@@ -233,8 +233,6 @@
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
@@ -386,7 +384,7 @@
" deploy_request_timeout=1800,\n",
" service_account=service_account,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_huggingface_vllm_deployment.ipynb\",\n",
" \"NOTEBOOK_NAME\": \"model_garden_huggingface_tgi_vllm_deployment.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": common_util.get_deploy_source(),\n",
" },\n",
" )\n",
@@ -405,8 +403,6 @@
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" enforce_eager=True,\n",
" max_num_seqs=5,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
@@ -595,7 +591,7 @@
],
"metadata": {
"colab": {
"name": "model_garden_huggingface_vllm_deployment.ipynb",
"name": "model_garden_huggingface_tgi_vllm_deployment.ipynb",
"toc_visible": true
},
"kernelspec": {
File diff suppressed because it is too large Load Diff
@@ -40,7 +40,7 @@
"\n",
" <td>\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_jax_dito.ipynb\">\n",
" <img src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" alt=\"GitHub logo\">\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",
@@ -40,7 +40,7 @@
"\n",
" <td>\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_jax_fvlm.ipynb\">\n",
" <img src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" alt=\"GitHub logo\">\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",
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_jax_owl_vit_v2.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -26,6 +26,7 @@
},
{
"cell_type": "markdown",
"language": "markdown",
"metadata": {
"id": "VJWDivOv3OWy"
},
@@ -34,18 +35,13 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_jax_paligemma_deployment.ipynb\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/colab/import/https:%2F%2Fraw.githubusercontent.com%2FGoogleCloudPlatform%2Fvertex-ai-samples%2Fmain%2Fnotebooks%2Fcommunity%2Fmodel_garden%2Fmodel_garden_jax_paligemma_deployment.ipynb\">\n",
" <img alt=\"Google Cloud Colab Enterprise logo\" src=\"https://lh3.googleusercontent.com/JmcxdQi-qOpctIvWKgPtrzZdJJK-J3sWE1RsfjZNwshCFgE_9fULcNpuXYTilIR2hjwN\" width=\"32px\"><br> Run in Colab Enterprise\n",
" </a>\n",
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_jax_paligemma_deployment.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -72,10 +68,6 @@
" - Detecting objects.\n",
"- Create a playground website to use with the PaliGemma Vertex AI Endpoint.\n",
"\n",
"### File a bug\n",
"\n",
"File a bug on [GitHub](https://github.com/GoogleCloudPlatform/vertex-ai-samples/issues/new) if you encounter any issue with the notebook.\n",
"\n",
"### Costs\n",
"\n",
"This tutorial uses billable components of Google Cloud:\n",
@@ -98,6 +90,7 @@
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "QvQjsmIJ6Y3f"
@@ -107,29 +100,25 @@
"# @title Setup Google Cloud project\n",
"# @markdown 1. [Make sure that billing is enabled for your project](https://cloud.google.com/billing/docs/how-to/modify-project).\n",
"\n",
"# @markdown 2. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"# @markdown 2. **[Optional]** [Create a Cloud Storage bucket](https://cloud.google.com/storage/docs/creating-buckets) for storing experiment outputs. Set the BUCKET_URI for the experiment environment. The specified Cloud Storage bucket (`BUCKET_URI`) should be located in the same region as where the notebook was launched. Note that a multi-region bucket (eg. \"us\") is not considered a match for a single region covered by the multi-region range (eg. \"us-central1\"). If not set, a unique GCS bucket will be created instead.\n",
"\n",
"BUCKET_URI = \"gs://\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. **[Optional]** Set region. If not set, the region will be set automatically according to Colab Enterprise environment.\n",
"\n",
"REGION = \"\" # @param {type:\"string\"}\n",
"\n",
"# @markdown 3. If you want to run predictions with A100 80GB or H100 GPUs, we recommend using the regions listed below. **NOTE:** Make sure you have associated quota in selected regions. Click the links to see your current quota for each GPU type: [Nvidia A100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_a100_80gb_gpus), [Nvidia H100 80GB](https://console.cloud.google.com/iam-admin/quotas?metric=aiplatform.googleapis.com%2Fcustom_model_serving_nvidia_h100_gpus). You can request for quota following the instructions at [\"Request a higher quota\"](https://cloud.google.com/docs/quota/view-manage#requesting_higher_quota).\n",
"\n",
"# @markdown > | Machine Type | Accelerator Type | Recommended Regions |\n",
"# @markdown | ----------- | ----------- | ----------- |\n",
"# @markdown | a2-ultragpu-1g | 1 NVIDIA_A100_80GB | us-central1, us-east4, europe-west4, asia-southeast1, us-east4 |\n",
"# @markdown | a3-highgpu-2g | 2 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-4g | 4 NVIDIA_H100_80GB | us-west1, asia-southeast1, europe-west4 |\n",
"# @markdown | a3-highgpu-8g | 8 NVIDIA_H100_80GB | us-central1, europe-west4, us-west1, asia-southeast1 |\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"! pip3 install --upgrade --quiet 'google-cloud-aiplatform==1.103.0'\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"# Import the necessary packages\n",
"! pip install -q gradio==4.21.0\n",
"import datetime\n",
"import enum\n",
"import importlib\n",
"import io\n",
"import os\n",
"import re\n",
"import uuid\n",
"from typing import Sequence, Tuple\n",
"\n",
"import gradio as gr\n",
@@ -139,38 +128,71 @@
"from google.cloud import aiplatform\n",
"from PIL import Image\n",
"\n",
"# Upgrade Vertex AI SDK.\n",
"if os.environ.get(\"VERTEX_PRODUCT\") != \"COLAB_ENTERPRISE\":\n",
" ! pip install --upgrade tensorflow\n",
"! git clone https://github.com/GoogleCloudPlatform/vertex-ai-samples.git\n",
"\n",
"common_util = importlib.import_module(\n",
" \"vertex-ai-samples.community-content.vertex_model_garden.model_oss.notebook_util.common_util\"\n",
")\n",
"\n",
"LABEL = \"endpoint\"\n",
"models, endpoints = {}, {}\n",
"\n",
"\n",
"# Get the default cloud project id.\n",
"PROJECT_ID = os.environ[\"GOOGLE_CLOUD_PROJECT\"]\n",
"\n",
"# Get the default region for launching jobs.\n",
"if not REGION:\n",
" if not os.environ.get(\"GOOGLE_CLOUD_REGION\"):\n",
" raise ValueError(\n",
" \"REGION must be set. See\"\n",
" \" https://cloud.google.com/vertex-ai/docs/general/locations for\"\n",
" \" available cloud locations.\"\n",
" )\n",
" REGION = os.environ[\"GOOGLE_CLOUD_REGION\"]\n",
"\n",
"# Enable the Vertex AI API and Compute Engine API, if not already.\n",
"print(\"Enabling Vertex AI API and Compute Engine API.\")\n",
"! gcloud services enable aiplatform.googleapis.com compute.googleapis.com\n",
"\n",
"# Cloud Storage bucket for storing the experiment artifacts.\n",
"# A unique GCS bucket will be created for the purpose of this notebook. If you\n",
"# prefer using your own GCS bucket, change the value yourself below.\n",
"now = datetime.datetime.now().strftime(\"%Y%m%d%H%M%S\")\n",
"BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
"\n",
"if BUCKET_URI is None or BUCKET_URI.strip() == \"\" or BUCKET_URI == \"gs://\":\n",
" BUCKET_URI = f\"gs://{PROJECT_ID}-tmp-{now}-{str(uuid.uuid4())[:4]}\"\n",
" BUCKET_NAME = \"/\".join(BUCKET_URI.split(\"/\")[:3])\n",
" ! gsutil mb -l {REGION} {BUCKET_URI}\n",
"else:\n",
" assert BUCKET_URI.startswith(\"gs://\"), \"BUCKET_URI must start with `gs://`.\"\n",
" shell_output = ! gsutil ls -Lb {BUCKET_NAME} | grep \"Location constraint:\" | sed \"s/Location constraint://\"\n",
" bucket_region = shell_output[0].strip().lower()\n",
" if bucket_region != REGION:\n",
" raise ValueError(\n",
" \"Bucket region %s is different from notebook region %s\"\n",
" % (bucket_region, REGION)\n",
" )\n",
"print(f\"Using this GCS Bucket: {BUCKET_URI}\")\n",
"\n",
"STAGING_BUCKET = os.path.join(BUCKET_URI, \"temporal\")\n",
"MODEL_BUCKET = os.path.join(BUCKET_URI, \"paligemma\")\n",
"\n",
"\n",
"# Initialize Vertex AI API.\n",
"print(\"Initializing Vertex AI API.\")\n",
"aiplatform.init(project=PROJECT_ID, location=REGION)\n",
"aiplatform.init(project=PROJECT_ID, location=REGION, staging_bucket=STAGING_BUCKET)\n",
"\n",
"# Gets the default SERVICE_ACCOUNT.\n",
"shell_output = ! gcloud projects describe $PROJECT_ID\n",
"project_number = shell_output[-1].split(\":\")[1].strip().replace(\"'\", \"\")\n",
"SERVICE_ACCOUNT = f\"{project_number}-compute@developer.gserviceaccount.com\"\n",
"print(\"Using this default Service Account:\", SERVICE_ACCOUNT)\n",
"\n",
"\n",
"# Provision permissions to the SERVICE_ACCOUNT with the GCS bucket\n",
"! gsutil iam ch serviceAccount:{SERVICE_ACCOUNT}:roles/storage.admin $BUCKET_NAME\n",
"\n",
"! gcloud config set project $PROJECT_ID\n",
"\n",
"import vertexai\n",
"\n",
"vertexai.init(\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
")\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/storage.admin\"\n",
"! gcloud projects add-iam-policy-binding --no-user-output-enabled {PROJECT_ID} --member=serviceAccount:{SERVICE_ACCOUNT} --role=\"roles/aiplatform.user\"\n",
"\n",
"# @markdown ### Access PaliGemma models on Vertex AI for GPU based serving\n",
"# @markdown Accept the model agreement to access the models:\n",
@@ -183,28 +205,17 @@
"VERTEX_AI_MODEL_GARDEN_PALIGEMMA = \"gs://\" # @param {type:\"string\", isTemplate:true}\n",
"assert (\n",
" VERTEX_AI_MODEL_GARDEN_PALIGEMMA\n",
"), \"Click the agreement of PaliGemma in Vertex AI Model Garden, and get the GCS path of PaliGemma model artifacts.\""
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "kyMJXkfviWgl"
},
"source": [
"## Deploy PaliGemma to a Vertex AI Endpoint"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "JThvioAxy8-a"
},
"outputs": [],
"source": [
"# @title Select the model variants\n",
"), \"Click the agreement of PaliGemma in Vertex AI Model Garden, and get the GCS path of PaliGemma model artifacts.\"\n",
"print(\n",
" \"Copying PaliGemma model artifacts from\",\n",
" VERTEX_AI_MODEL_GARDEN_PALIGEMMA,\n",
" \"to \",\n",
" MODEL_BUCKET,\n",
")\n",
"\n",
"! gsutil -m cp -R $VERTEX_AI_MODEL_GARDEN_PALIGEMMA/* $MODEL_BUCKET\n",
"\n",
"model_path_prefix = MODEL_BUCKET\n",
"\n",
"pretrained_filename_lookup = {\n",
" \"paligemma-224-float32\": \"pt_224.npz\",\n",
@@ -224,341 +235,6 @@
" \"paligemma-mix-448-bfloat16\": \"mix_448.bf16.npz\",\n",
"}\n",
"\n",
"# @markdown Select the desired resolution and precision of prebuilt model to deploy, leaving the optional `custom_paligemma_model_uri` as is. Higher resolution and precision_type can result in better inference results, but may require additional GPU.\n",
"\n",
"# @markdown You can also serve a finetuned PaliGemma model by setting `resolution` and `precision_type` to the resolution and precision type of the original base model and then setting `custom_paligemma_model_uri` to the GCS URI containing the model.\n",
"\n",
"# @markdown **Note**: You cannot use accelerator type `NVIDIA_TESLA_V100` to serve prebuilt or finetuned PaliGemma models with resolution `896` and precision_type `float32`.\n",
"\n",
"model_variant = \"mix\" # @param [\"mix\", \"pt\"]\n",
"resolution = 224 # @param [224, 448, 896]\n",
"precision_type = \"float32\" # @param [\"float32\", \"float16\", \"bfloat16\"]\n",
"custom_paligemma_model_uri = \"gs://\" # @param {type: \"string\"}\n",
"\n",
"if model_variant == \"mix\":\n",
" model_name_prefix = \"paligemma-mix\"\n",
"else:\n",
" model_name_prefix = \"paligemma\"\n",
"\n",
"\n",
"if custom_paligemma_model_uri == \"gs://\" or not custom_paligemma_model_uri:\n",
" model_name = f\"{model_name_prefix}-{resolution}-{precision_type}\"\n",
" checkpoint_filename = pretrained_filename_lookup[model_name]\n",
" checkpoint_path = os.path.join(\n",
" VERTEX_AI_MODEL_GARDEN_PALIGEMMA, checkpoint_filename\n",
" )\n",
" PUBLISHER_MODEL_NAME = f\"publishers/google/models/paligemma@{model_name}\"\n",
"else:\n",
" model_name = f\"{model_name_prefix}-{resolution}-{precision_type}-custom\"\n",
" checkpoint_path = custom_paligemma_model_uri\n",
"\n",
"# @markdown If you want to use other accelerator types not listed below, then check other Vertex AI prediction supported accelerators and regions at https://cloud.google.com/vertex-ai/docs/predictions/configure-compute. You may need to manually set the `machine_type`, `accelerator_type`, and `accelerator_count` in the code by clicking `Show code` first.\n",
"# @markdown Select the accelerator type to use to deploy the model:\n",
"accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"NVIDIA_TESLA_V100\"]\n",
"if accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-16\"\n",
" accelerator_count = 1\n",
"elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" if resolution == 896 and precision_type == \"float32\":\n",
" raise ValueError(\n",
" \"NVIDIA_TESLA_V100 is not sufficient. Multi-gpu is not supported for PaLIGemma.\"\n",
" )\n",
" else:\n",
" machine_type = \"n1-highmem-8\"\n",
" accelerator_count = 1\n",
"else:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another another accelerator, edit this code block to pass in an appropriate `machine_type`, `accelerator_type`, and `accelerator_count` to the deploy_model function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "144IKkHrzrMs"
},
"outputs": [],
"source": [
"# @title [Option 1] Deploy with Model Garden SDK\n",
"\n",
"# @markdown Kindly note that the deployment using custom_paligemma_model_uri is not supported.\n",
"\n",
"# @markdown Deploy with Gen AI model-centric SDK. This section uploads the prebuilt model to Model Registry and deploys it to a Vertex AI Endpoint. It takes 15 minutes to 1 hour to finish depending on the size of the model. See [use open models with Vertex AI](https://cloud.google.com/vertex-ai/generative-ai/docs/open-models/use-open-models) for documentation on other use cases.\n",
"from vertexai import model_garden\n",
"\n",
"model = model_garden.OpenModel(PUBLISHER_MODEL_NAME)\n",
"endpoints[LABEL] = model.deploy(\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" accept_eula=True, # Accept the End User License Agreement (EULA) on the model card before deploy. Otherwise, the deployment will be forbidden.\n",
")\n",
"\n",
"endpoint = endpoints[LABEL]"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "toY-WPKDFesF"
},
"outputs": [],
"source": [
"# @title [Option 2] Deploy with custom configs\n",
"\n",
"# @markdown This section uploads the prebuilt PaliGemma model to Model Registry and deploys it to a Vertex AI Endpoint. It takes approximately 15 minutes to finish.\n",
"\n",
"# The pre-built serving docker image.\n",
"SERVE_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/jax-paligemma-serve-gpu:20240807_0916_RC00\"\n",
"\n",
"\n",
"def deploy_model(\n",
" model_name: str,\n",
" checkpoint_path: str,\n",
" machine_type: str = \"g2-standard-32\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" resolution: int = 224,\n",
" use_dedicated_endpoint: bool = False,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Create a Vertex AI Endpoint and deploy the specified model to the endpoint.\"\"\"\n",
" model_name_with_time = common_util.get_job_name_with_datetime(model_name)\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name_with_time}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" )\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name_with_time,\n",
" serving_container_image_uri=SERVE_DOCKER_URI,\n",
" serving_container_ports=[8080],\n",
" serving_container_predict_route=\"/predict\",\n",
" serving_container_health_route=\"/health\",\n",
" serving_container_environment_variables={\n",
" \"CKPT_PATH\": checkpoint_path,\n",
" \"RESOLUTION\": resolution,\n",
" \"MODEL_ID\": \"google/\" + model_name,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" },\n",
" model_garden_source_model_name=\"publishers/google/models/paligemma\",\n",
" )\n",
" print(\n",
" f\"Deploying {model_name_with_time} on {machine_type} with {accelerator_count} {accelerator_type} GPU(s).\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" enable_access_logging=True,\n",
" min_replica_count=1,\n",
" sync=True,\n",
" system_labels={\n",
" \"NOTEBOOK_NAME\": \"model_garden_jax_paligemma_deployment.ipynb\",\n",
" \"NOTEBOOK_ENVIRONMENT\": common_util.get_deploy_source(),\n",
" },\n",
" )\n",
" return model, endpoint\n",
"\n",
"\n",
"models[LABEL], endpoints[LABEL] = deploy_model(\n",
" model_name=model_name,\n",
" checkpoint_path=checkpoint_path,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" resolution=resolution,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "tOtYOhZa3lsx"
},
"outputs": [],
"source": [
"# @title [Optional] Loading an existing Endpoint\n",
"# @markdown If you've already deployed an Endpoint, you can load it by filling in the Endpoint's ID below.\n",
"# @markdown You can view deployed Endpoints at [Vertex Online Prediction](https://console.cloud.google.com/vertex-ai/online-prediction/endpoints).\n",
"endpoint_id = \"\" # @param {type: \"string\"}\n",
"\n",
"if endpoint_id:\n",
" endpoints[LABEL] = aiplatform.Endpoint(\n",
" endpoint_name=endpoint_id,\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
" )"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "MlP2Y7XE4SS5"
},
"source": [
"### Predict\n",
"\n",
"The following sections will use images from [pexels.com](https://www.pexels.com/) for demoing purposes. All the images have the following license: https://www.pexels.com/license/.\n",
"\n",
"Images will be resized to a width of 1000 pixels by default since requests made to a Vertex Endpoint are limited to 1.500MB."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "xnZw8wNyQhmN"
},
"outputs": [],
"source": [
"# @title Visual Question Answering\n",
"\n",
"# @markdown This section uses the deployed PaliGemma model to answer questions about a given image.\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with images and questions.\n",
"# @markdown ![](https://images.pexels.com/photos/4012966/pexels-photo-4012966.jpeg?w=1260&h=750)\n",
"image_url = \"https://images.pexels.com/photos/4012966/pexels-photo-4012966.jpeg\" # @param {type:\"string\"}\n",
"\n",
"image = common_util.download_image(image_url)\n",
"display(image)\n",
"\n",
"# @markdown You may leave question prompts empty and they will be ignored.\n",
"question_prompt_1 = \"Which of laptop, book, pencil, clock, flower are in the image?\" # @param {type: \"string\"}\n",
"question_prompt_2 = \"Do the book and the cup have the same color?\" # @param {type: \"string\"}\n",
"question_prompt_3 = \"Is there a person in the image?\" # @param {type: \"string\"}\n",
"question_prompt_4 = \"How many laptop are in the image?\" # @param {type: \"string\"}\n",
"question_prompt_5 = \"桌子是什么颜色的?\" # @param {type: \"string\"}\n",
"\n",
"# @markdown The question prompt can be non-English languages.\n",
"questions_list = [\n",
" question_prompt_1,\n",
" question_prompt_2,\n",
" question_prompt_3,\n",
" question_prompt_4,\n",
" question_prompt_5,\n",
"]\n",
"questions = [question for question in questions_list if question]\n",
"\n",
"answers = common_util.vqa_predict(\n",
" endpoints[\"endpoint\"],\n",
" questions,\n",
" image,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"for question, answer in zip(questions, answers):\n",
" print(f\"Question: {question}\")\n",
" print(f\"Answer: {answer}\")\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "mF1MxC1ouzqj"
},
"outputs": [],
"source": [
"# @title Image Captioning\n",
"# @markdown This section uses the deployed PaliGemma model to caption and describe an image in a chosen language.\n",
"\n",
"caption_prompt = True\n",
"\n",
"# @markdown <img src=\"https://storage.googleapis.com/longcap100/91.jpeg\" width=\"400\" >\n",
"\n",
"image_url = \"https://storage.googleapis.com/longcap100/91.jpeg\" # @param {type:\"string\"}\n",
"\n",
"language_code = \"en\" # @param {type: \"string\"}\n",
"\n",
"image = common_util.download_image(image_url)\n",
"display(image)\n",
"\n",
"# Make a prediction.\n",
"image_base64 = common_util.image_to_base64(image)\n",
"\n",
"caption = common_util.caption_predict(\n",
" endpoints[\"endpoint\"],\n",
" language_code,\n",
" image,\n",
" caption_prompt,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"print(\"Caption: \", caption)\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "TtkXMZTIegLq"
},
"outputs": [],
"source": [
"# @title OCR\n",
"# @markdown This section uses the deployed PaliGemma model to extract text from an image, starting from the top left.\n",
"ocr_prompt = \"ocr\"\n",
"\n",
"# @markdown ![](https://images.pexels.com/photos/8919535/pexels-photo-8919535.jpeg?auto=compress&cs=tinysrgb&w=630&h=375&dpr=2)\n",
"image_url = \"https://images.pexels.com/photos/8919535/pexels-photo-8919535.jpeg?auto=compress&cs=tinysrgb&w=1260&h=750&dpr=2\" # @param {type:\"string\"}\n",
"\n",
"image = common_util.download_image(image_url)\n",
"display(image)\n",
"text_found = common_util.ocr_predict(\n",
" endpoints[\"endpoint\"],\n",
" ocr_prompt,\n",
" image,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")\n",
"\n",
"print(f\"Text found: {text_found}\")\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "JlLr3nu-YEon"
},
"outputs": [],
"source": [
"# @title Object Detection\n",
"# @markdown This section uses the deployed PaliGemma model to output bounding boxes for specified object image in a given image.\n",
"# @markdown The text output will be parsed into bounding boxes and overlaid on the original image.\n",
"\n",
"# @markdown Specify what object to detect. To specify multiple objects, enter them as a semicolon separated list as shown below.\n",
"objects = \"plant ; pineapple ; glasses\" # @param {type:\"string\"}\n",
"detect_promt = f\"detect {objects}\"\n",
"\n",
"\n",
"def parse_detections(txt):\n",
" \"\"\"Parses bounding boxes from a detection string.\"\"\"\n",
@@ -601,8 +277,296 @@
" buf = io.BytesIO()\n",
" fig.savefig(buf)\n",
" buf.seek(0)\n",
" return Image.open(buf)\n",
" return Image.open(buf)"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "kyMJXkfviWgl"
},
"source": [
"## Deploy PaliGemma to a Vertex AI Endpoint"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "toY-WPKDFesF"
},
"outputs": [],
"source": [
"# @title Deploy\n",
"\n",
"# @markdown This section uploads the prebuilt PaliGemma model to Model Registry and deploys it to a Vertex AI Endpoint. It takes approximately 15 minutes to finish.\n",
"\n",
"# @markdown Select the desired resolution and precision of prebuilt model to deploy, leaving the optional `custom_paligemma_model_uri` as is. Higher resolution and precision_type can result in better inference results, but may require additional GPU.\n",
"\n",
"# @markdown You can also serve a finetuned PaliGemma model by setting `resolution` and `precision_type` to the resolution and precision type of the original base model and then setting `custom_paligemma_model_uri` to the GCS URI containing the model.\n",
"\n",
"# @markdown **Note**: You cannot use accelerator type `NVIDIA_TESLA_V100` to serve prebuilt or finetuned PaliGemma models with resolution `896` and precision_type `float32`.\n",
"\n",
"model_variant = \"mix\" # @param [\"mix\", \"pt\"]\n",
"resolution = 224 # @param [224, 448, 896]\n",
"precision_type = \"float32\" # @param [\"float32\", \"float16\", \"bfloat16\"]\n",
"custom_paligemma_model_uri = \"gs://\" # @param {type: \"string\"}\n",
"\n",
"if model_variant == \"mix\":\n",
" model_name_prefix = \"paligemma-mix\"\n",
"else:\n",
" model_name_prefix = \"paligemma\"\n",
"\n",
"if custom_paligemma_model_uri == \"gs://\" or not custom_paligemma_model_uri:\n",
" print(\"Deploying prebuilt PaliGemma model.\")\n",
" model_name = f\"{model_name_prefix}-{resolution}-{precision_type}\"\n",
" checkpoint_filename = pretrained_filename_lookup[model_name]\n",
" checkpoint_path = os.path.join(model_path_prefix, checkpoint_filename)\n",
"else:\n",
" print(\"Deploying custom PaliGemma model.\")\n",
" model_name = f\"{model_name_prefix}-{resolution}-{precision_type}-custom\"\n",
" checkpoint_path = custom_paligemma_model_uri\n",
"\n",
"# The pre-built serving docker image.\n",
"SERVE_DOCKER_URI = \"us-docker.pkg.dev/vertex-ai/vertex-vision-model-garden-dockers/jax-paligemma-serve-gpu:20240807_0916_RC00\"\n",
"\n",
"# @markdown If you want to use other accelerator types not listed below, then check other Vertex AI prediction supported accelerators and regions at https://cloud.google.com/vertex-ai/docs/predictions/configure-compute. You may need to manually set the `machine_type`, `accelerator_type`, and `accelerator_count` in the code by clicking `Show code` first.\n",
"# @markdown Select the accelerator type to use to deploy the model:\n",
"accelerator_type = \"NVIDIA_L4\" # @param [\"NVIDIA_L4\", \"NVIDIA_TESLA_V100\"]\n",
"if accelerator_type == \"NVIDIA_L4\":\n",
" machine_type = \"g2-standard-16\"\n",
" accelerator_count = 1\n",
"elif accelerator_type == \"NVIDIA_TESLA_V100\":\n",
" if resolution == 896 and precision_type == \"float32\":\n",
" raise ValueError(\n",
" \"NVIDIA_TESLA_V100 is not sufficient. Multi-gpu is not supported for PaLIGemma.\"\n",
" )\n",
" else:\n",
" machine_type = \"n1-highmem-8\"\n",
" accelerator_count = 1\n",
"else:\n",
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another another accelerator, edit this code block to pass in an appropriate `machine_type`, `accelerator_type`, and `accelerator_count` to the deploy_model function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" is_for_training=False,\n",
")\n",
"\n",
"\n",
"def deploy_model(\n",
" model_name: str,\n",
" checkpoint_path: str,\n",
" machine_type: str = \"g2-standard-32\",\n",
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" resolution: int = 224,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Create a Vertex AI Endpoint and deploy the specified model to the endpoint.\"\"\"\n",
" model_name_with_time = common_util.get_job_name_with_datetime(model_name)\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name_with_time}-endpoint\"\n",
" )\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name_with_time,\n",
" serving_container_image_uri=SERVE_DOCKER_URI,\n",
" serving_container_ports=[8080],\n",
" serving_container_predict_route=\"/predict\",\n",
" serving_container_health_route=\"/health\",\n",
" serving_container_environment_variables={\n",
" \"CKPT_PATH\": checkpoint_path,\n",
" \"RESOLUTION\": resolution,\n",
" \"MODEL_ID\": \"google/\" + model_name,\n",
" },\n",
" model_garden_source_model_name=\"publishers/google/models/paligemma\",\n",
" )\n",
" print(\n",
" f\"Deploying {model_name_with_time} on {machine_type} with {accelerator_count} {accelerator_type} GPU(s).\"\n",
" )\n",
" model.deploy(\n",
" endpoint=endpoint,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" deploy_request_timeout=1800,\n",
" service_account=SERVICE_ACCOUNT,\n",
" enable_access_logging=True,\n",
" min_replica_count=1,\n",
" sync=True,\n",
" system_labels={\"NOTEBOOK_NAME\": \"model_garden_jax_paligemma_deployment.ipynb\"},\n",
" )\n",
" return model, endpoint\n",
"\n",
"\n",
"models[\"model\"], endpoints[\"endpoint\"] = deploy_model(\n",
" model_name=model_name,\n",
" checkpoint_path=checkpoint_path,\n",
" machine_type=machine_type,\n",
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" resolution=resolution,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "tOtYOhZa3lsx"
},
"outputs": [],
"source": [
"# @title [Optional] Loading an existing Endpoint\n",
"# @markdown If you've already deployed an Endpoint, you can load it by filling in the Endpoint's ID below.\n",
"# @markdown You can view deployed Endpoints at [Vertex Online Prediction](https://console.cloud.google.com/vertex-ai/online-prediction/endpoints).\n",
"endpoint_id = \"\" # @param {type: \"string\"}\n",
"\n",
"if endpoint_id:\n",
" endpoint = aiplatform.Endpoint(\n",
" endpoint_name=endpoint_id,\n",
" project=PROJECT_ID,\n",
" location=REGION,\n",
" )"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "MlP2Y7XE4SS5"
},
"source": [
"### Predict\n",
"\n",
"The following sections will use images from [pexels.com](https://www.pexels.com/) for demoing purposes. All the images have the following license: https://www.pexels.com/license/.\n",
"\n",
"Images will be resized to a width of 1000 pixels by default since requests made to a Vertex Endpoint are limited to 1.500MB."
]
},
{
"cell_type": "code",
"execution_count": null,
"language": "python",
"metadata": {
"cellView": "form",
"id": "xnZw8wNyQhmN"
},
"outputs": [],
"source": [
"# @title Visual Question Answering\n",
"\n",
"# @markdown This section uses the deployed PaliGemma model to answer questions about a given image.\n",
"\n",
"# @markdown Once deployment succeeds, you can send requests to the endpoint with images and questions.\n",
"# @markdown ![](https://images.pexels.com/photos/4012966/pexels-photo-4012966.jpeg?w=1260&h=750)\n",
"image_url = \"https://images.pexels.com/photos/4012966/pexels-photo-4012966.jpeg\" # @param {type:\"string\"}\n",
"\n",
"image = common_util.download_image(image_url)\n",
"display(image)\n",
"\n",
"# @markdown You may leave question prompts empty and they will be ignored.\n",
"question_prompt_1 = \"Which of laptop, book, pencil, clock, flower are in the image?\" # @param {type: \"string\"}\n",
"question_prompt_2 = \"Do the book and the cup have the same color?\" # @param {type: \"string\"}\n",
"question_prompt_3 = \"Is there a person in the image?\" # @param {type: \"string\"}\n",
"question_prompt_4 = \"How many laptop are in the image?\" # @param {type: \"string\"}\n",
"question_prompt_5 = \"桌子是什么颜色的?\" # @param {type: \"string\"}\n",
"\n",
"# @markdown The question prompt can be non-English languages.\n",
"questions_list = [\n",
" question_prompt_1,\n",
" question_prompt_2,\n",
" question_prompt_3,\n",
" question_prompt_4,\n",
" question_prompt_5,\n",
"]\n",
"questions = [question for question in questions_list if question]\n",
"\n",
"answers = common_util.vqa_predict(endpoints[\"endpoint\"], questions, image)\n",
"\n",
"for question, answer in zip(questions, answers):\n",
" print(f\"Question: {question}\")\n",
" print(f\"Answer: {answer}\")\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "mF1MxC1ouzqj"
},
"outputs": [],
"source": [
"# @title Image Captioning\n",
"# @markdown This section uses the deployed PaliGemma model to caption and describe an image in a chosen language.\n",
"\n",
"caption_prompt = True\n",
"\n",
"# @markdown <img src=\"https://storage.googleapis.com/longcap100/91.jpeg\" width=\"400\" >\n",
"\n",
"image_url = \"https://storage.googleapis.com/longcap100/91.jpeg\" # @param {type:\"string\"}\n",
"language_code = \"en\" # @param {type: \"string\"}\n",
"\n",
"image = common_util.download_image(image_url)\n",
"display(image)\n",
"\n",
"# Make a prediction.\n",
"image_base64 = common_util.image_to_base64(image)\n",
"\n",
"caption = common_util.caption_predict(\n",
" endpoints[\"endpoint\"], language_code, image, caption_prompt\n",
")\n",
"\n",
"print(\"Caption: \", caption)\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "TtkXMZTIegLq"
},
"outputs": [],
"source": [
"# @title OCR\n",
"# @markdown This section uses the deployed PaliGemma model to extract text from an image, starting from the top left.\n",
"ocr_prompt = \"ocr\"\n",
"\n",
"# @markdown ![](https://images.pexels.com/photos/8919535/pexels-photo-8919535.jpeg?auto=compress&cs=tinysrgb&w=630&h=375&dpr=2)\n",
"image_url = \"https://images.pexels.com/photos/8919535/pexels-photo-8919535.jpeg?auto=compress&cs=tinysrgb&w=1260&h=750&dpr=2\" # @param {type:\"string\"}\n",
"\n",
"image = common_util.download_image(image_url)\n",
"display(image)\n",
"text_found = common_util.ocr_predict(endpoints[\"endpoint\"], ocr_prompt, image)\n",
"\n",
"print(f\"Text found: {text_found}\")\n",
"# @markdown Click \"Show Code\" to see more details."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"cellView": "form",
"id": "JlLr3nu-YEon"
},
"outputs": [],
"source": [
"# @title Object Detection\n",
"# @markdown This section uses the deployed PaliGemma model to output bounding boxes for specified object image in a given image.\n",
"# @markdown The text output will be parsed into bounding boxes and overlaid on the original image.\n",
"\n",
"# @markdown Specify what object to detect. To specify multiple objects, enter them as a semicolon separated list as shown below.\n",
"objects = \"plant ; pineapple ; glasses\" # @param {type:\"string\"}\n",
"detect_promt = f\"detect {objects}\"\n",
"\n",
"# @markdown ![](https://images.pexels.com/photos/1006293/pexels-photo-1006293.jpeg?auto=compress&cs=tinysrgb&w=630&h=375&dpr=2)\n",
"image_url = \"https://images.pexels.com/photos/1006293/pexels-photo-1006293.jpeg?auto=compress&cs=tinysrgb&w=1260&h=750&dpr=2\" # @param {type:\"string\"}\n",
@@ -612,15 +576,10 @@
"\n",
"# Make a prediction.\n",
"detection_response = common_util.detect_predict(\n",
" endpoints[\"endpoint\"],\n",
" detect_promt,\n",
" image,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" endpoints[\"endpoint\"], detect_promt, image\n",
")\n",
"\n",
"print(\"Output: \", detection_response)\n",
"\n",
"\n",
"bboxes = parse_detections(detection_response)\n",
"plot_bounding_boxes(image, bboxes)\n",
"# @markdown Click \"Show Code\" to see more details."
@@ -747,9 +706,7 @@
" resolution = int(resolution)\n",
" model, endpoint = deploy_model(\n",
" model_name=model_choice,\n",
" checkpoint_path=os.path.join(\n",
" VERTEX_AI_MODEL_GARDEN_PALIGEMMA, checkpoint_filename\n",
" ),\n",
" checkpoint_path=os.path.join(model_path_prefix, checkpoint_filename),\n",
" machine_type=\"g2-standard-16\",\n",
" accelerator_type=\"NVIDIA_L4\",\n",
" accelerator_count=1,\n",
@@ -934,8 +891,6 @@
},
"outputs": [],
"source": [
"# @title Delete the models and endpoints\n",
"\n",
"# @markdown Delete the experiment models and endpoints to recycle the resources\n",
"# @markdown and avoid unnecessary continuous charges that may incur.\n",
"\n",
@@ -945,7 +900,11 @@
"\n",
"# Delete models.\n",
"for model in models.values():\n",
" model.delete()"
" model.delete()\n",
"\n",
"delete_bucket = False # @param {type:\"boolean\"}\n",
"if delete_bucket:\n",
" ! gsutil -m rm -r $BUCKET_NAME"
]
}
],
@@ -43,7 +43,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_jax_paligemma_finetuning.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -277,17 +277,9 @@
"\n",
"dataset_gcs_uri = \"gs://longcap100/data_train90.jsonl\" # @param {type: \"string\"}\n",
"\n",
"# @markdown [Optional] You can specify the `image` fields in the JSONL file to\n",
"# @markdown contain only filenames. In this case, you must also provide the\n",
"# @markdown image storage location in `dataset_image_dir`. If the JSONL file\n",
"# @markdown already contains full paths to the images, leave\n",
"# @markdown `dataset_image_dir` blank. Note that the `SERVICE_ACCOUNT` defined\n",
"# @markdown above must have read access to the images.\n",
"dataset_image_dir = \"\" # @param {type:\"string\"}\n",
"\n",
"# Set defaults for the example dataset.\n",
"if dataset_gcs_uri == \"gs://longcap100/data_train90.jsonl\" and not dataset_image_dir:\n",
" dataset_image_dir = \"gs://longcap100\""
"# @markdown [Optional] You can optionally specify the image fields in the JSONL file to use the\n",
"# @markdown filename and fill in the `dataset_image_dir` with the location where the images are stored.\n",
"dataset_image_dir = \"\" # @param {type:\"string\"}"
]
},
{
@@ -423,7 +415,8 @@
"if learning_rate:\n",
" train_args.append(f\"--config.lr={learning_rate}\")\n",
"\n",
"train_args.append(f\"--config.input.data.fopen_keys.image={dataset_image_dir}\")\n",
"if dataset_image_dir:\n",
" train_args.append(f\"--config.input.data.fopen_keys.image={dataset_image_dir}\")\n",
"train_job.run(\n",
" args=train_args,\n",
" replica_count=replica_count,\n",
@@ -536,10 +529,6 @@
" raise ValueError(\n",
" f\"Recommended machine settings not found for: {accelerator_type}. To use another another accelerator, edit this code block to pass in an appropriate `machine_type`, `accelerator_type`, and `accelerator_count` to the deploy_model function by clicking `Show Code` and then modifying the code.\"\n",
" )\n",
"\n",
"# @markdown Set use_dedicated_endpoint to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint). Note that [dedicated endpoint does not support VPC Service Controls](https://cloud.google.com/vertex-ai/docs/predictions/choose-endpoint-type), uncheck the box if you are using VPC-SC.\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"\n",
"common_util.check_quota(\n",
" project_id=PROJECT_ID,\n",
" region=REGION,\n",
@@ -556,13 +545,11 @@
" accelerator_type: str = \"NVIDIA_L4\",\n",
" accelerator_count: int = 1,\n",
" resolution: int = 224,\n",
" use_dedicated_endpoint: bool = False,\n",
") -> Tuple[aiplatform.Model, aiplatform.Endpoint]:\n",
" \"\"\"Create a Vertex AI Endpoint and deploy the specified model to the endpoint.\"\"\"\n",
" model_name_with_time = common_util.get_job_name_with_datetime(model_name)\n",
" endpoint = aiplatform.Endpoint.create(\n",
" display_name=f\"{model_name_with_time}-endpoint\",\n",
" dedicated_endpoint_enabled=use_dedicated_endpoint,\n",
" display_name=f\"{model_name_with_time}-endpoint\"\n",
" )\n",
" model = aiplatform.Model.upload(\n",
" display_name=model_name_with_time,\n",
@@ -574,7 +561,6 @@
" \"CKPT_PATH\": checkpoint_path,\n",
" \"RESOLUTION\": resolution,\n",
" \"MODEL_ID\": model_name,\n",
" \"DEPLOY_SOURCE\": \"notebook\",\n",
" },\n",
" model_garden_source_model_name=\"publishers/google/models/paligemma\",\n",
" )\n",
@@ -604,7 +590,6 @@
" accelerator_type=accelerator_type,\n",
" accelerator_count=accelerator_count,\n",
" resolution=model_resolution,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
")"
]
},
@@ -634,7 +619,6 @@
"# @markdown <img src=\"https://storage.googleapis.com/longcap100/91.jpeg\" width=\"400\" >\n",
"\n",
"image_url = \"https://storage.googleapis.com/longcap100/91.jpeg\" # @param {type:\"string\"}\n",
"\n",
"language_code = \"en\" # @param {type: \"string\"}\n",
"\n",
"image = common_util.download_image(image_url)\n",
@@ -644,11 +628,7 @@
"image_base64 = common_util.image_to_base64(image)\n",
"\n",
"caption = common_util.caption_predict(\n",
" endpoints[\"endpoint\"],\n",
" language_code,\n",
" image,\n",
" caption_prompt,\n",
" use_dedicated_endpoint=use_dedicated_endpoint,\n",
" endpoints[\"endpoint\"], language_code, image, caption_prompt\n",
")\n",
"\n",
"print(\"Caption: \", caption)\n",
@@ -39,7 +39,7 @@
" </td>\n",
" <td>\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_jax_stable_diffusion_xl.ipynb\">\n",
" <img src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" alt=\"GitHub logo\">\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",
@@ -40,7 +40,7 @@
"\n",
" <td>\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_jax_vision_transformer.ipynb\">\n",
" <img src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" alt=\"GitHub logo\">\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",
@@ -40,7 +40,7 @@
"\n",
" <td>\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_keras_stable_diffusion.ipynb\">\n",
" <img src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" alt=\"GitHub logo\">\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",
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_keras_yolov8.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -34,7 +34,7 @@
"\n",
"<table><tbody><tr>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/notebooks/deploy-notebook?download_url=https://raw.githubusercontent.com/GoogleCloudPlatform/vertex-ai-samples/main/notebooks/community/model_garden/model_garden_llama3_1_finetuning_with_workbench.ipynb\">\n",
" <a href=\"https://console.cloud.google.com/vertex-ai/workbench/instances\">\n",
" <img alt=\"Workbench logo\" src=\"https://lh3.googleusercontent.com/UiNooY4LUgW_oTvpsNhPpQzsstV5W8F7rYgxgGBD85cWJoLmrOzhVs_ksK_vgx40SHs7jCqkTkCk=e14-rj-sc0xffffff-h130-w32\" width=\"32px\"><br> Run in Workbench\n",
" </a>\n",
" </td>\n",
@@ -45,7 +45,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_llama3_1_finetuning_with_workbench.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"
@@ -427,7 +427,7 @@
"id": "L_q9h-SArI0c"
},
"source": [
"You must provide a Hugging Face User Access Token (with read access) to access the Llama 3.1 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below."
"You must provide a Hugging Face User Access Token (read) to access the Llama 3.1 models. You can follow the [Hugging Face documentation](https://huggingface.co/docs/hub/en/security-tokens) to create a **read** access token and put it in the `HF_TOKEN` field below."
]
},
{
@@ -1309,8 +1309,7 @@
" train_job.wait()\n",
" print(\"The training job has finished.\")\n",
"\n",
"# @markdown Set `use_dedicated_endpoint` to False if you don't want to use [dedicated endpoint](https://cloud.google.com/vertex-ai/docs/general/deployment#create-dedicated-endpoint).\n",
"use_dedicated_endpoint = True # @param {type:\"boolean\"}\n",
"use_dedicated_endpoint = False\n",
"\n",
"# Find Vertex AI prediction supported accelerators and regions [here](https://cloud.google.com/vertex-ai/docs/predictions/configure-compute).\n",
"if \"8b\" in base_model_id.lower():\n",
@@ -40,7 +40,7 @@
" </td>\n",
" <td style=\"text-align: center\">\n",
" <a href=\"https://github.com/GoogleCloudPlatform/vertex-ai-samples/blob/main/notebooks/community/model_garden/model_garden_llama3_2_deployment_on_gke.ipynb\">\n",
" <img alt=\"GitHub logo\" src=\"https://github.githubassets.com/assets/GitHub-Mark-ea2971cee799.png\" width=\"32px\"><br> View on GitHub\n",
" <img alt=\"GitHub logo\" src=\"https://cloud.google.com/ml-engine/images/github-logo-32px.png\" width=\"32px\"><br> View on GitHub\n",
" </a>\n",
" </td>\n",
"</tr></tbody></table>"

Some files were not shown because too many files have changed in this diff Show More