本教程介绍如何使用 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 Service Usage Admin IAM role (
roles/serviceusage.serviceUsageAdmin), which contains theserviceusage.services.enablepermission. 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
创建使用 适用于 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 微调脚本:
在此脚本中,以下内容适用:
Gemma3Tokenizer将文本数据转换为模型可以处理的 token。load_and_preprocess_data函数从文件中读取训练数据,将其拆分为各个故事,并使用分词器将文本转换为填充后的词法单元序列。generate_text函数接受模型、其参数和提示,以生成文本。train_step函数定义了一次训练迭代,其中包括前向传递、损失计算(使用交叉熵)、梯度计算和参数更新。train_model函数会按指定的周期数遍历数据集,并针对每个批次调用train_step函数。run_training函数可协调整个流程,以加载数据、初始化 Gemma 3 模型 (Gemma3_270M) 和优化器、加载预训练的参数、设置用于并行处理的数据分片、运行测试生成、执行训练循环,并执行最终的文本生成来演示微调的效果。- 该脚本使用
argparse库来接受maxlen、batch_size和datacount参数的命令行实参。
现在,您已经探索了微调脚本,接下来将其容器化,以便在 GKE 上运行。
将微调脚本容器化
在 GKE 集群中运行微调脚本之前,您需要将其容器化。本教程使用 JAX AI 映像作为基础映像。
打开与
Gemma3LLMTrain.py文件位于同一目录中的Dockerfile:此 Dockerfile 会安装必要的依赖项,并将
Gemma3LLMTrain.py文件复制到容器中。构建 Docker 映像并将其推送到映像代码库:
export REPOSITORY=REPOSITORY_NAME export IMAGE_NAME="jax-gemma3-training" export IMAGE_TAG="latest" export DOCKERFILE_PATH="./Dockerfile" export IMAGE_URI="${REGION}-docker.pkg.dev/${PROJECT_ID}/${REPOSITORY}/${IMAGE_NAME}:${IMAGE_TAG}" docker build -t "${IMAGE_URI}" -f "${DOCKERFILE_PATH}" . gcloud auth configure-docker "${REGION}-docker.pkg.dev" -q docker push "${IMAGE_URI}"将
REPOSITORY_NAME替换为您的 Artifact Registry 制品库的名称。向服务账号添加角色绑定:
export PROJECT_NUMBER=$(gcloud projects describe $PROJECT_ID --format 'get(projectNumber)') gcloud artifacts repositories add-iam-policy-binding ${REPOSITORY} \ --project=${PROJECT_ID} \ --location=${REGION} \ --member="serviceAccount:${PROJECT_NUMBER}-compute@developer." \ --role="roles/artifactregistry.reader"
将映像放入仓库后,您现在可以将微调作业部署到 GKE 集群中。
部署 LLM 微调作业
本部分介绍了如何将 LLM 微调作业部署到 GKE 集群。
打开
training_singlehost.yaml清单:应用清单:
envsubst < training_singlehost.yaml | kubectl apply -f -
GKE 会创建一个作业,该作业在 TPU Trillium (v6e) 节点上启动一个 Pod。此 Pod 运行 Python 微调脚本,该脚本使用 Cloud Storage FUSE 从装载在 /data 路径的指定 Cloud Storage 存储桶中访问微调数据。然后,脚本会对 Gemma 模型进行微调。
监控训练作业
在本部分中,您将监控微调作业的进度及其性能。
查看微调进度
列出 Pod:
# Find the Pods kubectl get pods按照日志输出操作:
kubectl logs -f pods/POD_NAME将
POD_NAME替换为您的 Pod 名称。输出类似于以下内容:
Global device count: 1 Batch size: 128, Max length: 256, Data count: 96000 I1028 00:12:55.925999 1387 google_auth_provider.cc:181] Running on GCE, using service account ... Generating response for: Once upon a time, there was a girl named Amy. Response: Amy lived in a small house. The house was in a big field. Amy liked to play in the big field. She Start training model Loss after batch 0: 10.25 Loss after batch 10: 4.3125 . . . Loss after batch 740: 1.41406 Completed training model. Total time for training 294.6791355609894 seconds Generating response for: Once upon a time, there was a girl named Amy. Response: She loved to play with her toys. One day, Amy's mom told her that she had to go to the store to分析输出内容:
Global device count: 1线表示使用的 TPU 核心数。- 在运行此微调之前,模型会生成合理的文本,因为它会从预训练的检查点加载。
- 微调后生成的输出更像短篇故事的开头,这表明模型正在从新数据集中学习。
- 在完整数据集上进行微调应能生成更精细的输出。
观察指标
通过检查 TPU 和 CPU 指标,查看微调作业的性能。如需查看集群的可观测性指标,请按照查看集群和工作负载可观测性指标中的步骤操作。
其他微调配置
本部分概述了微调工作负载的替代配置。
模型选择
本教程使用了 Gemma3_270M 模型,这是一个小型模型,可放入单主机 TPU Trillium (v6e) 节点池中。对于需要更多内存和计算资源才能进行微调的较大模型,您可以使用多主机或多切片节点池配置。
如需查看可用模型的完整列表,请参阅 Gemma 文档。
节点池配置
本教程使用了单主机节点池。您还可以根据需要创建多主机 TPU 切片节点池或多切片节点池。
以下标签页展示了如何为多主机和多切片节点池创建节点池:
多主机
在 Cloud Shell 中,运行以下命令:
gcloud container node-pools create jax-tpu-multihost1 \ --cluster=${CLUSTER_NAME} \ --machine-type=ct6e-standard-4t \ --num-nodes=2 \ --tpu-topology=2x4 \ --location=${REGION} \ --node-locations=${ZONE}GKE 会创建一个具有
2x4拓扑和两个节点的 TPU Trillium 节点池。打开
training_multihost_jobset.yaml作业定义:部署微调作业:
envsubst < training_multihost_jobset.yaml | kubectl apply -f -
多切片
在 Cloud Shell 中,运行以下命令:
gcloud container node-pools create jax-tpu-multihost1 \ --cluster=${CLUSTER_NAME} \ --machine-type=ct6e-standard-4t \ --num-nodes=2 \ --tpu-topology=2x4 \ --location=${REGION} \ --node-locations=${ZONE} gcloud container node-pools create jax-tpu-multihost2 \ --cluster=${CLUSTER_NAME} \ --machine-type=ct6e-standard-4t \ --num-nodes=2 \ --tpu-topology=2x4 \ --location=${REGION} \ --node-locations=${ZONE}GKE 会创建两个 TPU Trillium 节点池。每个节点池都有一个
2x4拓扑和两个节点。打开
training_multislice_jobset.yaml作业定义:部署微调作业:
envsubst < training_multislice_jobset.yaml | kubectl apply -f -
性能分析和优化
如需分析和优化机器学习微调的性能,您可以使用 XProf。XProf 是一套工具,用于分析和检查使用 JAX、TensorFlow 或 PyTorch/XLA 构建的机器学习工作负载。通过显示执行轨迹、内存用量和其他数据,XProf 可让您微调模型和训练设置,以提高效率并加快训练速度。
如需使用 XProf 分析微调工作负载的性能,请完成本部分中的以下步骤:
- 安装
xprof软件包。 修改训练脚本以启动 XProf 服务器。 - 修改 Kubernetes 作业清单,以包含 XProf 日志的卷装载。
- 向服务账号授予将 XProf 日志写入 Cloud Storage 存储桶的权限。
- 在 Pod 内运行 XProf,并设置端口转发以访问 XProf 信息中心。
安装 XProf 软件包
导航到包含 XProf 样本的目录:
cd ai-ml/llm-training-jax-tpu-gemma3/xprof-enabled构建 Docker 映像并将其推送到映像代码库:
export REPOSITORY=REPOSITORY_NAME export IMAGE_NAME="jax-gemma3-training-xp" export IMAGE_TAG="latest" export DOCKERFILE_PATH="./Dockerfile" export IMAGE_URI="${REGION}-docker.pkg.dev/${PROJECT_ID}/${REPOSITORY}/${IMAGE_NAME}:${IMAGE_TAG}" docker build -t "${IMAGE_URI}" -f "${DOCKERFILE_PATH}" . gcloud auth configure-docker "${REGION}-docker.pkg.dev" -q docker push "${IMAGE_URI}"将
REPOSITORY_NAME替换为您的 Artifact Registry 制品库的名称。运行
Dockerfile脚本:此 Dockerfile 会安装 XProf 依赖项。
将微调脚本复制到容器中
在本部分中,您将创建并应用一个 Kubernetes 作业清单,其中包含 XProf 日志所需的卷装载。
打开
training_singlehost.yaml作业定义:应用清单:
envsubst < training_singlehost.yaml | kubectl apply -f -
向服务账号授予写入 XProf 日志的权限
如需使服务账号能够写入和读取,请添加
"roles/storage.objectUser"角色:export GSA_NAME="GSA_NAME" # Same as used in initial setup # Automatically get the current project ID export PROJECT_ID=$(gcloud config get-value project) # Cloud Storage Bucket details export XPROF_GCS_BUCKET_NAME="XPROF_GCS_BUCKET_NAME" # Derived Variables export GSA_EMAIL="${GSA_NAME}@${PROJECT_ID}.iam." gcloud storage buckets add-iam-policy-binding "gs://${XPROF_GCS_BUCKET_NAME}" \ --member="serviceAccount:${GSA_EMAIL}" \ --role="roles/storage.objectUser" \ --project="${PROJECT_ID}"替换以下内容:
GSA_NAME:要向其授予角色的 Google 服务账号的名称。XPROF_GCS_BUCKET_NAME:要向其授予角色的存储桶的名称。
在 Pod 中运行 XProf:
kubectl exec POD_NAME -c training-container -it -- bash # exec into the container xprof --port 9001 --logdir /xprof # start xprof将
POD_NAME替换为您的 Pod 名称。
访问 XProf 信息中心
设置到 Pod 中 XProf 服务器的端口转发:
kubectl port-forward POD_NAME 9001:9001在浏览器的地址栏中,输入以下内容:
http://localhost:9001/XProf Trace Viewer 随即打开。
在 TensorBoard 窗口中,点击捕获性能剖析文件。
在配置文件服务网址或 TPU 名称字段中,输入
localhost:9002。如需捕获更多详细信息,请在主机跟踪记录 (TraceMe) 级别中选择 verbose 并启用 Python 跟踪记录日志记录。
如需查看信息中心,请点击捕获。
TensorBoard 会捕获性能分析数据,并让您分析训练脚本的性能。该图显示了 TPU 和 CPU 性能配置的执行时间线:
如需了解更多用于分析训练工作负载性能的性能分析选项,请参阅有关计算性能分析的 JAX 文档。
在生产环境中进行微调
本教程介绍了如何在分布式环境中测试基于 JAX 的训练。如需在生产环境中优化 LLM 微调,请使用 Maxtext 库。如果您对扩散模型感兴趣,请使用 Maxdiffusion 实现。
对于生产环境中长时间运行的训练或微调工作负载,请设置工作负载检查点,以最大限度地减少故障期间的进度损失。如需详细了解如何设置多层级检查点,请参阅在 GKE 上使用多层级检查点机制训练大规模机器学习模型。
清理
为避免因本教程中使用的资源导致您的 Google Cloud 账号产生费用,请删除包含这些资源的项目,或者保留项目但删除各个资源。
逐个删除资源
为避免因本教程中使用的资源导致您的 Google Cloud 账号产生费用,请删除包含这些资源的项目,或者保留该项目但运行以下命令来删除各个资源:
删除您在本教程中创建的资源:
gcloud container clusters delete ${CLUSTER_NAME} --location=${REGION} gcloud storage rm --recursive gs://${GCS_BUCKET_NAME} gcloud artifacts docker images delete ${IMAGE_URI} --delete-tags如果您不需要 XProf 生成的数据,请移除 XProf 使用的 Cloud Storage 存储桶:
gcloud storage rm --recursive gs://${XPROF_GCS_BUCKET_NAME}
后续步骤
- 详细了解 GKE 中的 TPU。
- 探索 JAX 代码库。