Flax 云端训练实战:用 launch_gce.py 在 Google Cloud 上启动、监控与自动回收训练任务
2026/9/17 16:26:55 网站建设 项目流程

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 内含该任务所需的全部代码与配置。这种"单任务单机"的模式带来两个直接好处:

  1. 并行实验互不干扰:可以同时创建多台 VM 运行多个实验(例如不同的超参数配置),彼此完全隔离,无需担心资源抢占或依赖冲突。
  2. 便于人工介入调试:任何一台机器上的训练都可以通过 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 章节):

  1. 创建 Google Cloud 账号:拥有一个可用的 Google Cloud 项目(Project)。
  2. 开通计费:在 Google Cloud 控制台的 Billing 页面为项目绑定结算账户,这是创建 VM 和使用 GCS 的前提。
  3. 创建存储桶(GCS Bucket):用于存放训练输出(产物、最终 checkpoint)。该存储桶由$GCS_BUCKET环境变量指定。
  4. (可选)申请加速器配额:若计划使用 GPU(如 V100),需要先在 IAM & Admin 的 Quotas 页面申请加速器配额。配额通常会在较短延迟内自动审批。

此外,launch_gce.py在本地运行,依赖gcloud命令行工具,并且需要在本地完成gcloud auth login或服务账号认证,确保有创建实例、写入 GCS 的权限。

环境变量约定

文档中的命令统一依赖以下环境变量(见 README 的 Setting up your environment 章节):

必填变量:

变量含义
$PROJECT你的 Google Cloud 项目名(Project ID)。
$GCS_BUCKETGoogle Cloud Storage 存储桶名,模型输出(产物、最终 checkpoint)存放于此。
$ZONE计算区域(Compute Zone),例如us-west1-acentral1-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.1momentum=0.9batch_size=128num_epochs=10等超参数。
  • --name=default:实验名,会被扩展为{example}/{name}/{timestamp}路径段。
  • --repo--branch:Git 仓库与分支,默认分别为https://github.com/google/flaxmain

此外还有几个常用 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 依次完成以下步骤:

  1. 创建/train工作目录并进入。
  2. 生成sudo_tmux_a.sh/tmux_a.sh两个辅助脚本,让用户可以通过gcloud compute ssh <vm> -- /sudo_tmux_a.sh一键附加到 tmux 会话(见 startup_script.sh#L10-L15)。
  3. 写入主训练脚本/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)。
  4. __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 数据无法自动下载,必须先手动准备:

  1. 从 image-net.org 官网下载imagenet2012原始数据(具体下载方式以 TensorFlow Datasets 的 imagenet2012 目录页说明为准)。
  2. 设置环境变量$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')) "
  1. 将生成的~/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.1warmup_epochs=5.0num_epochs=100等),并覆盖为batch_size=2048shuffle_buffer_size=16*2048cache=Truehalf_precision=True。仓库还提供了非混合精度的 v100_x8.py(batch_size=512cache=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 章节 给出两条高频实用技巧:

  1. --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)。
  2. 手工微调启动脚本:当需要反复调试 startup script 或单个参数时,可以 SSH 登录 VM → 停止正在运行的脚本并结束 tmux 会话 → 把launch_gce.py生成的flax-<example>-<timestamp>-startup_script.sh内容复制出来,修改后再手动执行。由于生成脚本中的__XXX__占位符已被替换为真实值,直接编辑它比改动模板再重跑建机流程更快。

安全与参数校验说明

launch_gce.py在建机前会做两类校验(见 launch_gce.py#L232-L244):

  • 正则校验repobranchexamplenamegcs_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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询