就本文档而言,批处理工作负载是指执行到完成并部署在与 Pathways 集群相同的 GKE 集群中的 JAX 工作负载,具体来说,是与 Pathways 控制器组件(IFRT 代理服务器和 Pathways 资源管理器)一起部署。JAX 工作负载完成后,Pathways 集群组件会终止。 本指南使用 JAX 训练工作负载来演示这一点。
准备工作
请确保您已备妥:
使用 Maxtext 构建训练映像
MaxText 是 Google 开发的一款开源大语言模型 (LLM) 项目。它采用 JAX 编写,旨在实现高性能和可伸缩性,可在 Google Cloud TPU 和 GPU 上高效运行。
如需使用 OSS GitHub 代码库中的最新稳定版 JAX 构建 MaxText Docker 映像,请运行以下命令:
git clone https://github.com/AI-Hypercomputer/maxtext cd maxtext/dependencies/scripts gcloud config set project PROJECT_ID bash ./docker_build_dependency_image.sh MODE=stable gcloud auth configure-docker bash ./docker_upload_runner.sh CLOUD_IMAGE_NAME=USER_runner # This script needs bash version >= 4.2 to execute.
此命令会将 MaxText Kubernetes 映像推送到 gcr.io/$PROJECT_ID/${USER}_runner。您可以使用此 Docker 映像通过 Pathways 后端在 TPU 上运行训练。
使用 Cluster Toolkit 运行批处理工作负载
使用 gcluster job submit 命令提交预构建的 MaxText Docker 映像:
gcluster job submit \
--pathways \
--pathways-gcs-location="gs://BUCKET_NAME/pathways-artifacts" \
--name=WORKLOAD \
--cluster=CLUSTER \
--project=PROJECT_ID \
--location=ZONE \
--num-slices=WORKLOAD_NODEPOOL_COUNT \
--compute-type=COMPUTE_TYPE \
--topology=TOPOLOGY \
--image="gcr.io/PROJECT_ID/USER_runner" \
--command="python3 -m MaxText.train /deps/src/MaxText/configs/base.yml base_output_directory=gs://BUCKET_NAME per_device_batch_size=1 enable_checkpointing=false remat_policy=full global_parameter_scale=1 steps=20 max_target_length=2048 use_iota_embed=true reuse_example_batch=1 dataset_type=synthetic attention=flash gcs_metrics=True enable_single_controller=True run_name=RUN_NAME-pathways-job"
如需详细了解作业提交选项,请参阅 Cluster Toolkit 作业提交指南。
替换以下内容:
WORKLOAD:用于标识工作负载的唯一名称;由于 DNS 标签的限制,此名称必须不超过 22 个字符CLUSTER:GKE 集群的名称WORKLOAD_NODEPOOL_COUNT:TPU 切片节点池的数量COMPUTE_TYPE:TPU 机器类型(例如ct6e-standard-4t)。如需详细了解每个 TPU 版本支持的 TPU 类型,请参阅 TPU 版本。TOPOLOGY:TPU placement 拓扑(例如2x4)PROJECT_ID:您的 Google Cloud 项目 IDZONE:您计划运行工作负载的可用区USER:您的 Google Cloud 用户 IDBUCKET_NAME:用于输出的 Cloud Storage 存储桶RUN_NAME:用户分配的用于标识工作流运行的名称
使用 gcluster job logs 命令跟踪工作负载的进度:
gcluster job logs WORKLOAD \
--cluster=CLUSTER \
--project=PROJECT_ID \
--location=ZONE \
--main-only=false
completed step: 1, seconds: 0.484, TFLOP/s/device: 87.349, Tokens/s/device: 2117.382, total_weights: 2945, loss: 10.888 completed step: 2, seconds: 0.407, TFLOP/s/device: 103.699, Tokens/s/device: 2513.735, total_weights: 3253, loss: 9.697 completed step: 3, seconds: 0.248, TFLOP/s/device: 170.300, Tokens/s/device: 4128.167, total_weights: 3154, loss: 9.641 completed step: 4, seconds: 0.216, TFLOP/s/device: 195.122, Tokens/s/device: 4729.880, total_weights: 3119, loss: 9.547 completed step: 5, seconds: 0.272, TFLOP/s/device: 155.298, Tokens/s/device: 3764.512, total_weights: 2837, loss: 10.179 completed step: 6, seconds: 0.472, TFLOP/s/device: 89.489, Tokens/s/device: 2169.266, total_weights: 3069, loss: 9.776
如需在工作负载完成之前取消它,请使用 gcluster job cancel 命令:
gcluster job cancel WORKLOAD --cluster=CLUSTER --project=PROJECT_ID --location=ZONE
后续步骤
- 创建 GKE 集群 - Pathways
- 使用 Pathways 进行多主机推理
- Pathway 互动模式
- 将 JAX 工作负载移植到 Pathways
- 通过 Pathways 打造弹性训练
- 排查 Pathways on Cloud 问题