如何在Vertex AI超参数调优作业中访问托管数据集
在Vertex AI Pipeline超参数调优作业中访问Vertex AI Dataset的方法
1. 获取Vertex AI Dataset的GCS存储路径
Vertex AI Dataset的实际数据存储在GCS中,首先需要在Pipeline里获取对应的GCS路径。可以用Vertex AI SDK封装成组件实现:
from google.cloud import aiplatform def get_dataset_gcs_path(dataset_id: str, project: str, region: str) -> str: aiplatform.init(project=project, region=region) # 根据数据集类型选择对应的类,比如TabularDataset/ImageDataset dataset = aiplatform.TabularDataset(dataset_name=dataset_id) # 提取GCS路径(多文件场景可根据需求调整逻辑) return dataset.gca_resource.data_source.gcs_source.uri[0]
2. 调整训练代码以读取GCS路径数据
修改超参数调优的训练逻辑,不再用tensorflow_datasets,而是接收GCS路径参数并读取数据。以表格数据为例:
def train_model(hparams, dataset_gcs_path): import pandas as pd from tensorflow import keras # 读取GCS上的数据集文件(格式根据实际情况调整,比如TFRecord) df = pd.read_csv(dataset_gcs_path) # 数据预处理、特征工程步骤 # ... # 模型定义与训练 model = keras.Sequential([...]) model.compile(optimizer=keras.optimizers.Adam(hparams["learning_rate"]), loss="binary_crossentropy", metrics=["accuracy"]) model.fit(train_data, epochs=10, batch_size=hparams["batch_size"]) # 保存模型到指定GCS路径 model.save("gs://your-bucket/tuned-models/") return "gs://your-bucket/tuned-models/"
3. 在Pipeline中传递数据集路径并配置调优作业
将获取路径的组件与HyperparameterTuningJobRunOp关联,把GCS路径作为参数传入训练容器:
from kfp.v2 import dsl from google_cloud_pipeline_components.v1.hyperparameter_tuning import HyperparameterTuningJobRunOp from google_cloud_pipeline_components.v1.model import ModelUploadOp @dsl.pipeline(name="hp-tuning-with-vertex-dataset") def pipeline(project: str, region: str, dataset_id: str): # 获取数据集GCS路径的组件 get_dataset_task = dsl.component(get_dataset_gcs_path)( dataset_id=dataset_id, project=project, region=region ) # 配置超参数调优作业 hp_tuning_task = HyperparameterTuningJobRunOp( project=project, region=region, display_name="hp-tuning-job", max_trial_count=10, parallel_trial_count=3, worker_pool_specs=[ { "machine_spec": { "machine_type": "n1-standard-4", "accelerator_type": "NVIDIA_TESLA_T4", "accelerator_count": 1 }, "replica_count": 1, "container_spec": { "image_uri": "us-docker.pkg.dev/vertex-ai/training/tf-cpu.2-8:latest", "command": ["python", "-m"], "args": [ "trainer.task", "--dataset_gcs_path", get_dataset_task.output, "--learning_rate", "{{$.hyperparameters.learning_rate}}", "--batch_size", "{{$.hyperparameters.batch_size}}" ] } } ], hyperparameter_spec={ "parameters": [ {"parameter_name": "learning_rate", "float_value_spec": {"min_value": 0.001, "max_value": 0.1}}, {"parameter_name": "batch_size", "integer_value_spec": {"min_value": 32, "max_value": 128}} ] }, metric_spec={"metric_name": "accuracy", "goal": "MAXIMIZE"} ) # 上传最优模型(可选,用于后续部署) model_upload_task = ModelUploadOp( project=project, display_name="tuned-model", artifact_uri=hp_tuning_task.outputs["best_model_artifact_uri"], serving_container_image_uri="us-docker.pkg.dev/vertex-ai/prediction/tf2-cpu.2-8:latest" )
4. 自动关联元数据
Vertex AI Pipeline会自动跟踪组件间的依赖关系:
- 数据集会被标记为超参数调优作业的输入
- 调优产出的模型会关联到原始数据集
- 后续部署端点时,端点也会自动关联到模型和数据集,可在Vertex AI控制台的元数据页面查看完整关联链路
注意事项
- 确保训练容器使用的服务账号有访问目标GCS路径的权限
- 根据数据集类型(图像、文本等)调整数据读取逻辑,比如图像数据集可读取GCS上的文件列表加载图片
- 若数据集包含标注信息,训练代码可直接读取GCS路径下的标注文件
内容的提问来源于stack exchange,提问作者Daruri
相关产品推荐
相关产品推荐

