Flax 云端训练实战:用 launch_gce.py 在 Google Cloud 上启动、监控与自动回收训练任务
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 EngineGCE上创建虚拟机VM、拉取 Flax 仓库并运行训练示例、通过gcloud storage rsync将训练产物同步到 GCS 存储桶并在任务结束后自动关机。读完本文你将掌握完整的云端训练流程环境准备、命令行参数、MNIST 与 ImageNet8×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.shVM 启动时由 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-filestartup-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_BUCKETGoogle 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 参数展开未设置时自动回退到默认值。实战一在云端训练 MNISTMNIST 是最轻量的入门示例一条命令即可完成建机 → 装环境 → 训练 → 同步产物 → 关机全流程命令见 README 的 Training the MNIST example 章节。运行前请确保$PROJECT与$GCS_BUCKET已正确设置python examples/cloud/launch_gce.py \ --project$PROJECT \ --zoneus-west1-a \ --machine_typen2-standard-2 \ --gcs_workdir_basegs://$GCS_BUCKET/workdir_base \ --repo${REPO:-https://github.com/google/flax} \ --branch${BRANCH:-main} \ --examplemnist \ --args--configconfigs/default.py \ --namedefault上述参数对应 launch_gce.py 中定义的 flags含义如下--project、--zone项目名与区域两者与--machine_type、--gcs_workdir_base、--example、--name一起被flags.mark_flags_as_required标记为必填见 launch_gce.py#L130-L132。--machine_typen2-standard-2VM 机型可用gcloud compute machine-types list查看可选列表。--gcs_workdir_basegs://$GCS_BUCKET/workdir_baseGCS 上的工作目录基址。实际的--workdir会由脚本自动拼接为{gcs_workdir_base}/{example}/{name}/{timestamp}形如gs://my-bucket/workdir_base/mnist/default/20240916_101530无需手动指定见 launch_gce.py#L81-L88。--examplemnist要运行的示例名对应仓库 examples/mnist 目录。脚本会校验该目录确实存在见 launch_gce.py#L237-L242。--args--configconfigs/default.py透传给示例main.py的额外命令行参数。脚本只负责补全--workdir其余参数原样透传。MNIST 的 configs/default.py 定义了learning_rate0.1、momentum0.9、batch_size128、num_epochs10等超参数。--namedefault实验名会被扩展为{example}/{name}/{timestamp}路径段。--repo、--branchGit 仓库与分支默认分别为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 --logdirgcs_workdir_baseTensorBoard 可直接读取 GCS 路径通过 GCS 控制台的存储桶浏览器查看已同步的文件。VM 的命名规则为flax-example-timestamp并将非法字符替换为-见 launch_gce.py#L246-L251。实战二在 8×V100 上训练 ImageNetImageNet 属于大规模训练场景需要 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_configtfds.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 \ --zoneus-west1-a \ --machine_typen1-standard-96 \ --accelerator_typenvidia-tesla-v100 --accelerator_count8 \ --gcs_workdir_basegs://$GCS_BUCKET/workdir_base \ --tfds_data_dirgs://$GCS_TFDS_BUCKET/datasets \ --repo${REPO:-https://github.com/google/flax} \ --branch${BRANCH:-main} \ --exampleimagenet \ --args--configconfigs/v100_x8_mixed_precision.py \ --namev100_x8_mixed_precision与 MNIST 相比的关键差异--machine_typen1-standard-9696 vCPU 的高配机型匹配 8 卡 GPU 的数据吞吐需求。--accelerator_typenvidia-tesla-v100 --accelerator_count8挂载 8 张 V100 GPU。当二者非空时脚本会额外追加--maintenance-policyTERMINATE与--acceleratortype...,count8参数见 launch_gce.py#L182-L186。--maintenance-policyTERMINATE保证发生维护事件时实例直接终止而非迁移避免 GPU 实例迁移失败的问题。--tfds_data_dirgs://$GCS_TFDS_BUCKET/datasets指向 GCS 上预置的数据集目录训练时通过环境变量TFDS_DATA_DIR注入避免每台 VM 重复从外网下载见 startup_script.sh#L42。若留空数据集会从网络下载。--args--configconfigs/v100_x8_mixed_precision.py使用 ImageNet 的 8 卡混合精度配置 configs/v100_x8_mixed_precision.py。该配置继承 default.pyResNet50、learning_rate0.1、warmup_epochs5.0、num_epochs100等并覆盖为batch_size2048、shuffle_buffer_size16*2048、cacheTrue、half_precisionTrue。仓库还提供了非混合精度的 v100_x8.pybatch_size512、cacheTrue可按需选用。训练入口 examples/imagenet/main.py 与 examples/mnist/main.py 结构一致都通过--workdir指定输出目录、--config指定 ml_collections 配置文件并在入口处将 GPU 对 TensorFlow 隐藏tf.config.experimental.set_visible_devices([], GPU)确保 TF 不抢占显存、把 GPU 完整留给 JAX。调试与优化技巧TipsREADME 的 Tips 章节 给出两条高频实用技巧--connect直达训练现场在启动命令后追加--connect脚本会轮询 VM 就绪状态对 connection refused、HTTP 502 等瞬态错误自动重试每次等待 20 秒见 launch_gce.py#L273-L301就绪后直接 SSH 进入训练 tmux 会话。修改配置或脚本后调试时非常高效。同样地--wait只等待不登录此时若设置了VM_READY_CMD例如 macOS 下VM_READY_CMDosascript -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并附带--scopescloud-platform,storage-full授予云平台与存储桶全量访问权限、--boot-disk-size256GB、--boot-disk-typepd-ssd、--metadatainstall-nvidia-driverTrue自动安装 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),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →