本教程介绍如何使用 Google Kubernetes Engine (GKE) 上的张量处理单元 (TPU) 和 JAX 来对大语言模型 (LLM) 进行微调。借助微调,您可以调整基础模型(例如 Gemma 3),使其适合特定领域或任务。此过程通过使用您自己的专业数据集更新模型的参数,从而提高模型的精确度和准确度。
如果您在微调 AI/机器学习工作负载时需要利用托管式 Kubernetes 的精细控制、自定义、可伸缩性、弹性、可移植性和成本效益,那么本指南是一个很好的起点。
背景
通过在 GKE 上使用 TPU 和 Jax 来对 LLM 进行微调,您可以构建一个可用于生产用途的强大微调解决方案,具备托管式 Kubernetes 的所有优势。
Gemma
Gemma 是一组公开提供的轻量级生成式 AI/机器学习多模态模型(根据开放许可发布)。这些 AI 模型可以在应用、硬件、移动设备或托管服务中运行。Gemma 3 引入了多模态功能,支持视觉语言输入和文本输出。它可处理最多 128,000 个 token 的上下文窗口,并支持 140 多种语言。Gemma 3 还提供改进的数学、推理和聊天功能,包括结构化输出和函数调用。
您可以使用 Gemma 模型生成文本,也可以针对专门任务对这些模型进行调优。
如需了解详情,请参阅 Gemma 文档。
TPU
TPU 是 Google 定制开发的应用专用集成电路 (ASIC),用于加速使用 TensorFlow、PyTorch 和 JAX 等框架构建的机器学习和 AI 模型。
使用 GKE 中的 TPU 之前,我们建议您完成以下学习路线:
- 了解 Cloud TPU 系统架构中的当前 TPU 版本可用性。
- 了解 GKE 中的 TPU。
JAX
JAX 是一种高性能机器学习框架,旨在与 TPU 和 GPU 搭配使用。JAX 提供了一个用于构建和训练机器学习模型的 API。
如需了解详情,请参阅 JAX 代码库。
目标
本教程介绍以下步骤:
- 根据模型特征创建一个具有推荐 TPU 拓扑的 GKE Autopilot 或 Standard 集群。 在本教程中,您将在单主机节点池上执行微调。
- 将数据添加到 Cloud Storage 存储桶,并通过 Cloud Storage FUSE 将其装载到容器。
- 在 GKE 上部署 LLM 微调作业。
- 监控微调作业并查看日志。
准备工作
-
In the Google Cloud console, on the project selector page, select or create a Google Cloud project.
Roles required to select or create a project
- Select a project: Selecting a project doesn't require a specific IAM role—you can select any project that you've been granted a role on.
-
Create a project: To create a project, you need the Project Creator role
(
roles/resourcemanager.projectCreator), which contains theresourcemanager.projects.createpermission. Learn how to grant roles.
-
Verify that billing is enabled for your Google Cloud project.
Enable the required API.
Roles required to enable APIs
To enable APIs, you need the
serviceusage.services.enablepermission. If you created the project, then you likely already have this permission through the Owner role (roles/owner). Otherwise, you can get this permission through the Service Usage Admin role (roles/serviceusage.serviceUsageAdmin). Learn how to grant roles.-
确保您在项目中拥有以下一个或多个角色: roles/container.admin、roles/iam.serviceAccountAdmin、roles/storage.admin
检查角色
-
在 Google Cloud 控制台中,前往 IAM 页面。
转到 IAM - 选择项目。
-
在主账号列中,找到标识您或您所属群组的所有行。如需了解您属于哪些群组,请与您的管理员联系。
- 对于指定或包含您的所有行,请检查角色列以查看角色列表是否包含所需的角色。
授予角色
-
在 Google Cloud 控制台中,前往 IAM 页面。
转到 IAM - 选择项目。
- 点击 授予访问权限。
-
在新的主账号字段中,输入您的用户标识符。 这通常是员工身份池中的用户的标识符。如需了解详情,请参阅在 IAM 政策中表示员工池用户,或与您的管理员联系。
- 点击选择角色,然后搜索相应角色。
- 如需授予其他角色,请点击 添加其他角色,然后添加其他各个角色。
- 点击 Save(保存)。
-
- 确保您有足够的配额用于 16 个 TPU Trillium (v6e) 芯片。在本教程中,您将使用需要 16 个芯片和按需实例的节点池配置。
- 确保您拥有 Docker 代码库。如果您没有,请在 Artifact Registry 中创建一个标准代码库。
准备环境
在本教程中,您将使用 Cloud Shell 来管理 Google Cloud上托管的资源。Cloud Shell 中预安装了本教程所需的软件,包括 kubectl 和 Google Cloud CLI。
如需使用 Cloud Shell 设置您的环境,请按照以下步骤操作:
在 Google Cloud 控制台中,启动 Cloud Shell 会话,然后点击
激活 Cloud Shell。此操作会在 Google Cloud 控制台的底部窗格中启动会话。
设置默认环境变量:
gcloud config set project PROJECT_ID gcloud config set billing/quota_project PROJECT_ID export PROJECT_ID=$(gcloud config get project) export CLUSTER_NAME=CLUSTER_NAME export REGION=CONTROL_PLANE_LOCATION export ZONE=ZONE export GCS_BUCKET_NAME=BUCKET_NAME替换以下值:
PROJECT_ID:您的 Google Cloud 项目 ID。CLUSTER_NAME:GKE 集群的名称。CONTROL_PLANE_LOCATION:GKE 集群和 TPU 节点所在的 Compute Engine 区域。相应区域必须包含提供 TPU Trillium (v6e) 机器类型的可用区。ZONE:所选CONTROL_PLANE_LOCATION区域内可使用 TPU Trillium (v6e) 机器类型的可用区。如需列出提供 TPU Trillium (v6e) TPU 的地区,请运行以下命令:gcloud compute accelerator-types list --filter="name~ct6e" --format="value(zone)"BUCKET_NAME:包含训练数据的 Cloud Storage 存储桶的名称。
克隆示例代码库:
git clone https://github.com/GoogleCloudPlatform/kubernetes-engine-samples.git cd kubernetes-engine-samples导航到工作目录:
cd ai-ml/llm-training-jax-tpu-gemma3
创建和配置 Google Cloud 资源
在本部分中,您将创建和配置 Google Cloud 资源。
创建 GKE 集群
您可以在 GKE Autopilot 或 Standard 集群中的 TPU 上对 LLM 进行微调。我们建议您使用 Autopilot 集群获得全托管式 Kubernetes 体验。如需选择最适合您的工作负载的 GKE 操作模式,请参阅选择 GKE 操作模式。
Autopilot
创建一个使用 Workload Identity Federation for GKE并已启用 Cloud Storage FUSE 的 GKE Autopilot 集群。
gcloud container clusters create-auto ${CLUSTER_NAME} \
--location=${REGION}
集群创建可能需要几分钟的时间。
标准
创建使用Workload Identity Federation for GKE并已启用 Cloud Storage FUSE 的区域级 GKE Standard 集群。
gcloud container clusters create ${CLUSTER_NAME} \ --enable-ip-alias \ --addons GcsFuseCsiDriver \ --machine-type=n2-standard-4 \ --num-nodes=2 \ --workload-pool=${PROJECT_ID}. \ --location=${REGION}集群创建可能需要几分钟的时间。
创建单主机节点池:
gcloud container node-pools create jax-tpu-nodepool \ --cluster=${CLUSTER_NAME} \ --machine-type=ct6e-standard-1t \ --num-nodes=1 \ --location=${REGION} \ --node-locations=${ZONE} \ --workload-metadata=GKE_METADATA
GKE 会创建一个具有 1x1 拓扑和一个节点的 TPU Trillium 节点池。--workload-metadata=GKE_METADATA 标志将节点池配置为使用 GKE 元数据服务器。
安装 JobSet
配置
kubectl以与您的集群通信:gcloud container clusters get-credentials ${CLUSTER_NAME} --location=${REGION}安装最新发布的 JobSet 版本:
kubectl apply --server-side -f https://github.com/kubernetes-sigs/jobset/releases/download/JOBSET_VERSION/manifests.yaml将
JOBSET_VERSION替换为 JobSet 的最新发布版本。例如v0.11.0。验证 JobSet 安装:
kubectl get pods -n jobset-system输出类似于以下内容:
NAME READY STATUS RESTARTS AGE jobset-controller-manager-6c56668494-l4dhc 1/1 Running 0 4m45s如果 JobSet 正在等待资源,您可能需要添加更多节点。
配置 Cloud Storage FUSE
如需对 LLM 进行微调,您需要提供训练数据。在本教程中,您将使用 Hugging Face 中的 TinyStories 数据集。此数据集包含由 GPT-3.5 和 GPT-4 合成生成的短篇故事,这些故事使用有限的词汇。
本部分介绍如何配置 Cloud Storage FUSE 以从 Cloud Storage 存储桶中读取数据。
下载数据集:
wget https://huggingface.co/datasets/roneneldan/TinyStories/resolve/main/TinyStories-train.txt?download=true -O TinyStories-train.txt将数据上传到新的 Cloud Storage 存储桶:
gcloud storage buckets create gs://${GCS_BUCKET_NAME} \ --location=${REGION} \ --enable-hierarchical-namespace \ --uniform-bucket-level-access gcloud storage cp TinyStories-train.txt gs://${GCS_BUCKET_NAME}如需允许工作负载通过 Cloud Storage FUSE 读取数据,请创建 Kubernetes 服务账号 (KSA) 并添加所需权限。运行
permissionsetup.sh脚本:运行此脚本后,您的Google Cloud 项目和 GKE 集群中会配置以下资源:
- 系统会在您的项目中创建一个名为
gcs-fuse-sa的新 IAM 服务账号。 - 创建的 Google Cloud 服务账号 (GSA) (
gcs-fuse-sa) 会被授予${GCS_BUCKET_NAME}指定的 Cloud Storage 存储桶的roles/storage.objectViewer角色。此权限允许 GSA 从存储桶中读取对象。 - 系统会在 GKE 集群的
default命名空间中创建一个名为jaxserviceaccount的新 KSA。 - 更新 GSA 的 IAM 政策,以向 KSA 授予
roles/iam.workloadIdentityUser角色。此权限允许 KSA 模拟 GSA。 KSA 已添加注释,可将其与 GSA 相关联。此注解会告知 GKE,KSA 应使用 Workload Identity 模拟哪个 GSA。
现在,在 GKE 集群的
default命名空间中运行的任何使用jaxserviceaccount服务账号的 Pod 都将能够以gcs-fuse-saGSA 的身份进行身份验证。这些 Pod 将拥有对存储在gs://${GCS_BUCKET_NAME}存储桶中的对象的读取权限,这对于微调作业使用 Cloud Storage FUSE 访问数据集至关重要。
- 系统会在您的项目中创建一个名为
创建微调脚本
在本部分中,您将探索对 Gemma 3 模型执行微调操作的训练脚本。此脚本使用 Gemma3Tokenizer。
查看以下 Gemma3LLMTrain.py 微调脚本: