Flax 云端训练实战:用 launch_gce.py 在 Google Cloud 上启动、监控与自动回收训练任务
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
本篇指南以 Flax 仓库 examples/cloud 目录为核心,讲解如何用其中的launch_gce.py一键在 Google Cloud Compute Engine(GCE)上创建虚拟机(VM)、拉取 Flax 仓库并运行训练示例、通过gcloud storage rsync将训练产物同步到 GCS 存储桶,并在任务结束后自动关机。读完本文,你将掌握完整的云端训练流程:环境准备、命令行参数、MNIST 与 ImageNet(8×V100)两类实战启动方式,以及如何进入 tmux 会话进行调试、用 TensorBoard 观察训练进度。
设计思路:单 VM 单任务,天然隔离互不干扰
examples/cloud/README.md(原文)说明这套方案的核心思想是:每个训练任务运行在一台独立的 VM 上,这台 VM 内含该任务所需的全部代码与配置。这种"单任务单机"的模式带来两个直接好处:
- 并行实验互不干扰:可以同时创建多台 VM 运行多个实验(例如不同的超参数配置),彼此完全隔离,无需担心资源抢占或依赖冲突。
- 便于人工介入调试:任何一台机器上的训练都可以通过 SSH 登录并附加(attach)到 tmux 会话,实时查看日志、htop 与 GPU 状态。
整个编排由两个文件完成:
- launch_gce.py:本地执行的启动器。它负责校验参数、渲染启动脚本模板、调用
gcloud compute instances create创建 VM,并在创建后打印登录/监控所需的命令。 - startup_script.sh:VM 启动时由 GCE 元数据(
startup-script)触发的 Shell 模板。其中所有__XXX__占位符会被launch_gce.py在创建实例前逐一代换为实际值(见 launch_gce.py#L137-L159 的generate_startup_file函数)。
launch_gce.py会针对每个 VM 生成一份形如flax-<example>-<timestamp>-startup_script.sh的渲染后脚本(存于 examples/cloud 目录),再通过--metadata-from-file=startup-script=...传给 GCE。注意:VM 无论训练成功还是失败,都会在等待 5 分钟后自动关机,避免闲置计费;该等待时长可通过--shutdown_secs调整。
前置准备:账号、计费、存储桶与配额
在开始之前,需要完成以下准备工作(对应 README 的 Preparation 章节):
- 创建 Google Cloud 账号:拥有一个可用的 Google Cloud 项目(Project)。
- 开通计费:在 Google Cloud 控制台的 Billing 页面为项目绑定结算账户,这是创建 VM 和使用 GCS 的前提。
- 创建存储桶(GCS Bucket):用于存放训练输出(产物、最终 checkpoint)。该存储桶由
$GCS_BUCKET环境变量指定。 - (可选)申请加速器配额:若计划使用 GPU(如 V100),需要先在 IAM & Admin 的 Quotas 页面申请加速器配额。配额通常会在较短延迟内自动审批。
此外,launch_gce.py在本地运行,依赖gcloud命令行工具,并且需要在本地完成gcloud auth login或服务账号认证,确保有创建实例、写入 GCS 的权限。
环境变量约定
文档中的命令统一依赖以下环境变量(见 README 的 Setting up your environment 章节):
必填变量:
| 变量 | 含义 |
|---|---|
$PROJECT | 你的 Google Cloud 项目名(Project ID)。 |
$GCS_BUCKET | Google Cloud Storage 存储桶名,模型输出(产物、最终 checkpoint)存放于此。 |
$ZONE | 计算区域(Compute Zone),例如us-west1-a、central1-a。 |
可选变量:
| 变量 | 含义 |
|---|---|
$REPO | 替代默认仓库https://github.com/google/flax的 Git 仓库地址,便于开发时指向自己的 fork。 |
$BRANCH | 替代默认分支main的分支名,例如你自己的开发分支。 |
$REPO与$BRANCH在文档中的用法是${REPO:-https://github.com/google/flax}形式——即 Shell 参数展开,未设置时自动回退到默认值。
实战一:在云端训练 MNIST
MNIST 是最轻量的入门示例,一条命令即可完成"建机 → 装环境 → 训练 → 同步产物 → 关机"全流程(命令见 README 的 Training the MNIST example 章节)。运行前请确保$PROJECT与$GCS_BUCKET已正确设置:
python examples/cloud/launch_gce.py \ --project=$PROJECT \ --zone=us-west1-a \ --machine_type=n2-standard-2 \ --gcs_workdir_base=gs://$GCS_BUCKET/workdir_base \ --repo=${REPO:-https://github.com/google/flax} \ --branch=${BRANCH:-main} \ --example=mnist \ --args='--config=configs/default.py' \ --name=default上述参数对应 launch_gce.py 中定义的 flags,含义如下:
--project、--zone:项目名与区域,两者与--machine_type、--gcs_workdir_base、--example、--name一起被flags.mark_flags_as_required标记为必填(见 launch_gce.py#L130-L132)。--machine_type=n2-standard-2:VM 机型,可用gcloud compute machine-types list查看可选列表。--gcs_workdir_base=gs://$GCS_BUCKET/workdir_base:GCS 上的工作目录基址。实际的--workdir会由脚本自动拼接为{gcs_workdir_base}/{example}/{name}/{timestamp},形如gs://my-bucket/workdir_base/mnist/default/20240916_101530,无需手动指定(见 launch_gce.py#L81-L88)。--example=mnist:要运行的示例名,对应仓库 examples/mnist 目录。脚本会校验该目录确实存在(见 launch_gce.py#L237-L242)。--args='--config=configs/default.py':透传给示例main.py的额外命令行参数。脚本只负责补全--workdir,其余参数原样透传。MNIST 的 configs/default.py 定义了learning_rate=0.1、momentum=0.9、batch_size=128、num_epochs=10等超参数。--name=default:实验名,会被扩展为{example}/{name}/{timestamp}路径段。--repo、--branch:Git 仓库与分支,默认分别为https://github.com/google/flax与main。
此外还有几个常用 flags 未在上例中出现,在后面的 ImageNet 实战中会用到:--accelerator_type(加速器类型)、--accelerator_count(加速器数量,默认 8)、--tfds_data_dir(预置 TFDS 数据集目录)、--shutdown_secs(自动关机等待秒数,默认 300,设为 0 可禁用)、--dry_run(只打印将执行的 gcloud 命令而不真正建机)、--wait(等待 VM 就绪,可选执行VM_READY_CMD)、--connect(就绪后直接 SSH 进入训练会话)。
VM 内发生了什么:startup script 的执行流程
VM 启动后,渲染后的 startup_script.sh 依次完成以下步骤:
- 创建
/train工作目录并进入。 - 生成
sudo_tmux_a.sh/tmux_a.sh两个辅助脚本,让用户可以通过gcloud compute ssh <vm> -- /sudo_tmux_a.sh一键附加到 tmux 会话(见 startup_script.sh#L10-L15)。 - 写入主训练脚本
/install_train_stop.sh,其逻辑为:激活conda环境flax→ 浅克隆(--depth 1)指定分支的 Flax 仓库 → 用 Python 3.9 创建flaxconda 环境 →pip install -e .安装 Flax → 进入examples/<example>目录安装requirements.txt→ 执行python main.py --workdir=$WORKDIR <args>,全程日志通过tee写入$WORKDIR/setup_train_log_<timestamp>.txt(见 startup_script.sh#L17-L50)。 - 若
__SHUTDOWN_SECS__ > 0,打印倒计时提示后sleep对应秒数再执行shutdown now自动关机(见 startup_script.sh#L46-L50)。
TMUX 四窗格布局:训练、监控、同步三线并行
启动脚本随即创建名为flax的 tmux 会话,并编排为四个窗格(见 startup_script.sh#L56-L77):
- 左上:
htop,实时查看 CPU/内存。 - 右上:
watch nvidia-smi,轮询 GPU 利用率与显存。 - 左下:执行
/install_train_stop.sh主训练脚本。 - 右下:死循环执行
gcloud storage rsync --recursive workdir_base <gcs_workdir_base>,每 60 秒将本地工作目录增量同步到 GCS 存储桶,日志写入$WORKDIR/gcs_rsync_<timestamp>.txt。
这套布局的好处是:训练与同步并行进行,训练日志和 checkpoint 会"实时"出现在 GCS 中;即使 VM 意外终止,已同步的产物也不会丢失。用快捷键CTRL-B后按A即可从 tmux 会话中脱出而不中断训练(参考 launch_gce.py#L216-L217 的提示)。
建机后打印的监控信息
创建 VM 成功后,print_howto会输出一段操作指引(见 launch_gce.py#L204-L229),包括:
- 在 GCE 控制台的实例页面启停实例;
- SSH 登录并附加训练会话的完整命令:
gcloud compute ssh --project <project> --zone <zone> <vm> -- /sudo_tmux_a.sh; - 在本地启动 TensorBoard 观察训练:
tensorboard --logdir=<gcs_workdir_base>(TensorBoard 可直接读取 GCS 路径); - 通过 GCS 控制台的存储桶浏览器查看已同步的文件。
VM 的命名规则为flax-<example>-<timestamp>,并将非法字符替换为-(见 launch_gce.py#L246-L251)。
实战二:在 8×V100 上训练 ImageNet
ImageNet 属于大规模训练场景,需要 GPU 与预置数据集。完整步骤见 README 的 Training the imagenet example 章节。
第一步:准备 ImageNet 数据集
ImageNet 数据无法自动下载,必须先手动准备:
- 从 image-net.org 官网下载
imagenet2012原始数据(具体下载方式以 TensorFlow Datasets 的 imagenet2012 目录页说明为准)。 - 设置环境变量
$IMAGENET_DOWNLOAD_PATH指向下载文件所在目录,然后执行以下命令让tensorflow_datasets完成数据集的构建:
python -c " import tensorflow_datasets as tfds tfds.builder('imagenet2012').download_and_prepare( download_config=tfds.download.DownloadConfig( manual_dir='$IMAGENET_DOWNLOAD_PATH')) "- 将生成的
~/tensorflow_datasets目录内容复制到gs://$GCS_TFDS_BUCKET/datasets。$GCS_TFDS_BUCKET与$GCS_BUCKET可以是同一个存储桶。
第二步:启动训练
python examples/cloud/launch_gce.py \ --project=$PROJECT \ --zone=us-west1-a \ --machine_type=n1-standard-96 \ --accelerator_type=nvidia-tesla-v100 --accelerator_count=8 \ --gcs_workdir_base=gs://$GCS_BUCKET/workdir_base \ --tfds_data_dir=gs://$GCS_TFDS_BUCKET/datasets \ --repo=${REPO:-https://github.com/google/flax} \ --branch=${BRANCH:-main} \ --example=imagenet \ --args='--config=configs/v100_x8_mixed_precision.py' \ --name=v100_x8_mixed_precision与 MNIST 相比的关键差异:
--machine_type=n1-standard-96:96 vCPU 的高配机型,匹配 8 卡 GPU 的数据吞吐需求。--accelerator_type=nvidia-tesla-v100 --accelerator_count=8:挂载 8 张 V100 GPU。当二者非空时,脚本会额外追加--maintenance-policy=TERMINATE与--accelerator=type=...,count=8参数(见 launch_gce.py#L182-L186)。--maintenance-policy=TERMINATE保证发生维护事件时实例直接终止而非迁移,避免 GPU 实例迁移失败的问题。--tfds_data_dir=gs://$GCS_TFDS_BUCKET/datasets:指向 GCS 上预置的数据集目录,训练时通过环境变量TFDS_DATA_DIR注入,避免每台 VM 重复从外网下载(见 startup_script.sh#L42)。若留空,数据集会从网络下载。--args='--config=configs/v100_x8_mixed_precision.py':使用 ImageNet 的 8 卡混合精度配置 configs/v100_x8_mixed_precision.py。该配置继承 default.py(ResNet50、learning_rate=0.1、warmup_epochs=5.0、num_epochs=100等),并覆盖为batch_size=2048、shuffle_buffer_size=16*2048、cache=True、half_precision=True。仓库还提供了非混合精度的 v100_x8.py(batch_size=512、cache=True),可按需选用。
训练入口 examples/imagenet/main.py 与 examples/mnist/main.py 结构一致:都通过--workdir指定输出目录、--config指定 ml_collections 配置文件,并在入口处将 GPU 对 TensorFlow 隐藏(tf.config.experimental.set_visible_devices([], 'GPU')),确保 TF 不抢占显存、把 GPU 完整留给 JAX。
调试与优化技巧(Tips)
README 的 Tips 章节 给出两条高频实用技巧:
--connect直达训练现场:在启动命令后追加--connect,脚本会轮询 VM 就绪状态(对 "connection refused"、HTTP 502 等瞬态错误自动重试,每次等待 20 秒,见 launch_gce.py#L273-L301),就绪后直接 SSH 进入训练 tmux 会话。修改配置或脚本后调试时非常高效。同样地,--wait只等待不登录,此时若设置了VM_READY_CMD(例如 macOS 下VM_READY_CMD="osascript -e 'display notification \"VM ready\"'"),VM 就绪时会执行该命令弹出通知,避免干等。若同时使用--connect与--dry_run,脚本会直接报错拒绝执行(见 launch_gce.py#L243-L244)。- 手工微调启动脚本:当需要反复调试 startup script 或单个参数时,可以 SSH 登录 VM → 停止正在运行的脚本并结束 tmux 会话 → 把
launch_gce.py生成的flax-<example>-<timestamp>-startup_script.sh内容复制出来,修改后再手动执行。由于生成脚本中的__XXX__占位符已被替换为真实值,直接编辑它比改动模板再重跑建机流程更快。
安全与参数校验说明
launch_gce.py在建机前会做两类校验(见 launch_gce.py#L232-L244):
- 正则校验:
repo、branch、example、name、gcs_workdir_base五个参数若包含\w、:、/、_、-之外的字符(正则[^\w:/_-]命中),直接抛ValueError,防止注入异常内容到生成的 Shell 脚本。 - 目录存在性校验:
--example必须在 examples 目录下有对应子目录,否则报Could not find --example=...。
另外建机时会固定使用 Deep Learning VM 镜像c1-deeplearning-tf-2-10-cu113-v20221107-debian-10(来自ml-images项目,带预装 CUDA 与 TF 驱动,见 launch_gce.py#L173-L174),并附带--scopes=cloud-platform,storage-full(授予云平台与存储桶全量访问权限)、--boot-disk-size=256GB、--boot-disk-type=pd-ssd、--metadata=install-nvidia-driver=True(自动安装 NVIDIA 驱动)。可用gcloud compute images list --project ml-images查看该镜像项目下可用的镜像列表(见 launch_gce.py#L163-L164)。
小结
Flax 的 examples/cloud 提供了一套"零常驻资源"的云端训练方案:launch_gce.py负责建机与参数注入,startup_script.sh负责环境搭建、训练执行、产物同步与自动关机,tmux 四窗格让训练/监控/同步并行可见。无论你是想快速验证 MNIST 小实验,还是在 8×V100 上跑 ImageNet 大规模训练,都可以在此骨架之上扩展新的示例与配置,实现多实验并行、低成本自动回收的云端训练工作流。相关核心文件:launch_gce.py、startup_script.sh、README.md。
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考