Mesh Transformer JAX 实战指南:基于 JAX 与 Haiku 的 Transformer 模型并行训练、微调与推理 大模型基础模型深度学习分布式训练【免费下载链接】mesh-transformer-jaxModel parallel transformers in JAX and Haiku项目地址https://gitcode.com/gh_mirrors/me/mesh-transformer-jax点击查看免费下载导读本文围绕开源仓库 mesh-transformer-jax 的核心文档 README.md 展开系统讲解该库如何借助 JAX 的xmap/pjit算子实现 Megatron-LM 式及实验性 ZeRo 式的 Transformer 模型并行并完整覆盖其旗舰模型 GPT-J-6B 的架构细节、零样本评测结果以及基于 TPU VM Ray 的训练/微调/推理架构与 JAX 版本依赖约束。读者读完后将掌握从配置编写、数据准备、TPU 微调到权重瘦身与 HuggingFace 转换的完整实战方案。一、项目定位TPU 上的模型并行 Transformer 库根据 README.md 的定位说明mesh-transformer-jax 是一个基于 Haiku 的库使用 JAX 中的xmap/pjit算子来实现 Transformer 的模型并行。其并行方案与 Megatron-LM 原始论文 类似由于 TPU 具有高速 2D mesh 网络该方案在 TPU 上效率很高此外还有一个实现了 ZeRo 风格切分论文的实验性模型版本。README 明确指出该库的设计目标是在 TPUv3 上扩展到约 40B 参数规模超过这一规模后应改用其他并行策略README 建议参考 GPT-NeoX 或 DeepSpeed 等实现。这一规模边界是理解本库并行设计的前提它面向的是“单 Pod 内模型并行”的场景而不是超大规模跨 Pod 流水线并行。README 也提到未来研究方向之一是将本代码库与 swarm-jax 集成通过流水线并行实现进一步扩展——注意这是作者声明的“未来方向”并非当前已实现能力。二、GPT-J-6B预训练模型与权重README 的核心内容之一是对预训练模型 GPT-J-6B 的完整介绍这是一个在 The Pile 上训练的、拥有 60 亿参数的自回归文本生成模型。2.1 权重下载与演示资源README 提供了以下资源入口仓库文档原文列出本文仅转述其用途说明slim 权重仅 bf16 权重用于推理约 9GBstep_383500_slim.tar.zstd完整权重含优化器参数约 61GBstep_383500.tar.zstd部分训练的检查点step_383500/目录下的历史 stepColab 演示colab_demo.ipynb仓库内文件可直接打开实验Web 演示与作者博客这里需要特别注意一个仓库内可验证的事实完整权重含优化器状态是微调的前提。这一点在 howto_finetune.md 第 3 步中被明确强调“Youll need the full pretrained weights in order to fine-tune the model”微调模型需要完整的预训练权重。而 slim 权重只适合推理场景。2.2 模型超参数与实现细节README 的 Model Details 章节给出了 GPT-J-6B 的完整超参数表这是理解模型结构的第一手资料超参数值参数量 n_parameters6,053,381,344层数 n_layers28每层 1 个 FFN 块 1 个自注意力块模型维度 d_model4,096FFN 维度 d_ff16,384注意力头数 n_heads16每头维度 d_head256上下文长度 n_ctx2,048词表大小 n_vocab50,257与 GPT-2/3 相同的 tokenizer位置编码Rotary position encodingsRoPERoPE 维度64README 补充说明模型由 28 层组成模型维度 4096前馈维度 16384模型维度被切分为 16 个头每个头维度 256RoPE 应用于每个头的 64 个维度使用与 GPT-2/GPT-3 相同的 BPE 词表50257进行训练。仓库源码可以进一步印证上述实现细节在 configs/6B_roto_256.json 中layers: 28、d_model: 4096、n_heads: 16、pe: rotary、pe_rotary_dims: 64与 README 完全一致注意配置中n_vocab为 50400这是训练时的词表填充尺寸README 中的 50257 是 GPT-2/3 的实际 BPE 词表大小。RoPE 维度的实现位于 mesh_transformer/layers.py其中pe_rotary_dims默认取dim_per_head第 245 行RoPE 应用于 query/key 的前pe_rotary_dims维、其余维度原样透传第 262-266 行从而在注意力计算中实现旋转位置编码。2.3 Zero-Shot 评测表README 给出了一个按性能或 FLOPs大致排序的零样本评测对照表用于说明 GPT-J-6B 在其发布时的基准水平。表中带有*的数值由各自论文作者报告其余数值由作者用 lm-evaluation-harness 对发布权重或通过 API 运行得到。由于实现细节和任务框架差异这些数字可能不完全可比README 就此专门加了脚注说明。模型权重训练 FLOPsLAMBADA PPL ↓LAMBADA Acc ↑Winogrande ↑Hellaswag ↑PIQA ↑数据集大小 (GB)Chance✔0~a lot~0%50%25%25%0GPT-3-Ada‡✘-----9.9551.6%52.9%43.4%70.5%-----GPT-2-1.5B✔-----10.6351.21%59.4%50.9%70.8%40GPTNeo-1.3B‡✔3.0e217.5057.2%55.0%48.9%71.1%825Megatron-2.5B*✘2.4e21-----61.7%---------------174GPTNeo-2.7B‡✔6.8e215.6362.2%56.5%55.8%73.0%825GPT-3-1.3B*‡✘2.4e215.4463.6%58.7%54.7%75.1%~800GPT-3-Babbage‡✘-----5.5862.4%59.0%54.5%75.5%-----Megatron-8.3B*✘7.8e21-----66.5%---------------174GPT-3-2.7B*‡✘4.8e214.6067.1%62.3%62.8%75.6%~800Megatron-11B†✔1.0e22-------------------------161GPT-J-6B‡✔1.5e223.9969.7%65.3%66.1%76.5%825GPT-3-6.7B*‡✘1.2e224.0070.3%64.5%67.4%78.0%~800GPT-3-Curie‡✘-----4.0069.3%65.6%68.5%77.9%-----GPT-3-13B*‡✘2.3e223.5672.5%67.9%70.9%78.5%~800GPT-3-175B*‡✘3.1e233.0076.2%70.2%78.9%81.0%~800GPT-3-Davinci‡✘-----3.075%72%78%80%-----Gopher 230B*✘6.31E23-----74.50%70.10%79.20%81.80%1344MT-NLG 530B*‡✘----------76.6%73.0%80.2%82.0%-----README 的脚注含义需要原样保留避免误读*由论文作者报告的数字其他数字由运行 lm-evaluation-harness发布权重或 API 访问得出。由于实现细节与零样本任务框架的细微差异这些数字可能无法直接比较。†Megatron-11B 没有可比的指标且使用其发布权重的多个实现无法复现生成质量和评测结果因此未尝试评测。‡这些模型的训练数据可能包含测试集污染OpenAI GPT-3 模型未对某些测试集去重而 GPT-Neo 与本模型都训练于 The Pile未对任何测试集去重。关于评测工具的落地实现仓库内提供了 eval_harness.py 与 tasks/eval_harness.py并且配置项eval_harness_tasks如6B_roto_256.json中的lambada、piqa、hellaswag、winogrande等直接对应 README 表中出现的评测任务名。2.4 致谢与许可证README 说明项目计算资源由 TPU Research CloudTRC提供、EleutherAI 协助并感谢 Cloud TPU 团队提供 Cloud TPU VM alpha 早期访问。个人致谢名单Aran Komatsuzaki、James Bradbury、Janko Prester、Laurence Golding、Leo Gao 等对应各自贡献。GPT-J-6B 的权重采用 Apache License 2.0详见仓库 LICENSE.txt。三、架构与使用TPU VM Ray 的分工体系3.1 总体架构README 的 Architecture and Usage 章节描述了仓库的核心运行模型这是理解整个代码库的关键仓库中的大多数脚本设计为在TPU 上运行在 TPU-VM 架构下TPU VM 是能够运行任意代码的虚拟机。大多数脚本的工作方式是拉起一个 TPU → SSH 进入 TPU 安装依赖并拷贝本地代码 → 启动一个可接受 RPC 调用的 Ray worker。TPU VM 负责模型训练步、评测、checkpoint 的保存与加载driver Python 程序负责数据加载和整体编排例如何时保存 checkpoint。据此仓库脚本可分为两类在 GCE 虚拟机与 TPU 同区域上运行的脚本如 train.py、eval_harness.py 等。它们通过--tpu参数与 TPU 侧的 Ray worker 通信README 提醒它们要放在与 TPU 同区域以最小化 RPC 延迟和数据传输成本。直接在 TPU VM 上运行的脚本如 device_sample.py、device_serve.py、device_train.py这些脚本不接收--tpu参数。README 特别强调device_系列脚本只能在 v3-8 上工作不能在更大的 Pod 上运行*。从源码看这一约束确实存在以 device_train.py 为例其网格形状由mesh_shape (tpu_size // cores_per_replica, cores_per_replica)计算并断言cores_per_replica 8第 144 行即单个 replica 最多占 8 个设备正对应 v3-8 的 8 核结构。3.2 检查点转换从 8 分片到更少分片README 提到有一个 resharding_example.py 示例演示如何将官方提供的检查点GPT-J-6B 为 8 个分片转换到更少的分片数量例如在 GPU 上运行时。源码细节印证了这一点在 mesh_transformer/checkpoint.py 的read_ckpt中参数shards_in表示检查点原有分片数、shards_out表示目标分片数当二者不同时调用reshard函数第 157-158 行该函数按权重张量的形状特征区分 LayerNorm 参数、偏置、权重矩阵等不同情况分别做拼接/缩放处理。注意该函数对不同形状的张量分别处理一维张量取首片、二维张量区分“LN 类参数”各分片相同与“偏置/权重”需要按原形状拼接三维张量按两种 case 转置拼接无法匹配的形状会抛出unimplemented异常。这说明重分片只支持仓库已知的权重排布模式读者在自定义模型上使用时应先核对张量形状。resharding_example.py 还展示了单 GPU 推理的关键技巧将cores_per_replica设为 1、用optax.scale(0)构造“空优化器”以剔除优化器参数因为推理不需要、用read_ckpt(..., 8, shards_out1)把 8 分片权重合并到 1 个分片并把状态放到 CPU 上避免被 xmap 重复。文件头部注明该示例在 RTX 3090 上测试峰值显存约 22.4GB推理时、加载模型约 19GB并建议设置XLA_PYTHON_CLIENT_PREALLOCATEfalse与XLA_PYTHON_CLIENT_ALLOCATORplatform两个环境变量——这是 GPU 上跑 JAX 模型并行代码时非常实用的调优信息。3.3 微调Fine-tuningREADME 指出微调模型需要在 TPU VM 上运行device_train.py。使用 TPU v3-8 时微调速度约为~5000 tokens/second足以处理中小型数据集详细步骤见 howto_finetune.md。由于 README 将微调指南作为官方扩展文档README 的 Updates 一节明确记录了 12-07-21 新增该指南本文将把 howto_finetune.md 的核心步骤纳入构成完整的实操章节准备工作TRC 与 GCP 资源申请 TPU Research CloudTRC访问权限配合 Google Cloud 免费试用可以免费完成全部流程收到邮件后创建项目并填写表单。安装 Google Cloud SDK。创建 GCS bucket确保 bucket 与 TPU VM 所在区域一致TRC 邮件会告知可免费使用 TPU 的区域。下载完整预训练权重step_383500.tar.zstd含优化器参数约 61GB。上传权重与数据解压step_383500.tar.zstd得到包含分片检查点的未压缩目录不要上传压缩包TPU VM 本地存储不够解压指南作者因此不得不在 Colab 中解压重传。用gsutil -m cp -R LOCAL_PATH_TO/step_383500 gs://YOUR-BUCKET上传指南提到跨洋上传约需 12 小时建议选择地理位置更好的区域。训练数据的 tfrecords 也应上传到 bucket。准备索引文件与配置文件在仓库 data/ 目录下新建foo.train.indexfoo可自定义每个要训练的 tfrecord 在 index 中占一行GCS 路径如有验证集同样创建foo.val.index。可参照 data/example.train.index、data/pile.train.index 等现有文件。注意这些 index 文件在训练时由 device_train.py 以data/{params[train_set]}的形式读取第 236 行。复制 configs/6B_roto_256.json 并重命名按需修改以下字段指南原文要求本文结合仓库配置补充取值范围说明字段微调时的修改建议tpu_size从256改为8v3-8bucket改为你的 GCS bucketmodel_dir保存 checkpoint 的目录train_set/val_set指向上一步的 index 文件eval_harness_tasks不用评测时可删除/置空val_every/ckpt_every/keep_every按需设置不要设为 0否则除零报错没有val_set时把val_every设得比total_steps大val_batches等于验证集序列数可在create_finetune_tfrecords.py生成的 .tfrecords 文件末尾查到name模型名称对应 WandB run 名warmup_steps、lr、anneal_steps等见下文学习率说明仓库同时提供了微调配置模板 configs/example_config.jsontpu_size: 8、lr: 5e-5、end_lr: 1e-5、total_steps: 72、warmup_steps: 7、anneal_steps: 65、ckpt_every: 72、val_set: {}——这正是指南示例1147 序列 ÷ 16 梯度累积 72 步/epoch的成对实现可直接作为微调起点。将改动 push 到自己的 fork。部署 TPU VM 并启动微调按 Google Cloud TPU 官方 JAX 快速入门操作到“Connect to your Cloud TPU VM”步骤获得 VM 远程访问权限。在 VM 中git clone仓库或自己的 forkcd进入目录后pip install -r requirements.txt。注意 requirements.txt 未固定微调所需的确切 jax 版本需要额外执行pip install jax0.2.12详见下文“JAX 依赖”章节。运行微调python3 device_train.py --configYOUR_CONFIG.json --tune-model-pathgs://YOUR-BUCKET/step_383500/指南说明启动后模型先加载进内存控制台显示loading network后约 10-15 分钟进入下一步随后进入 WandB 日志设置选项 3 可在不使用 WandB 时跳过保存第 1 步的 checkpoint 后正式开始训练。小数据集会很快完成TPU VM 训练速率约 5000 tokens/second。结束后记得清理关闭 TPU VM、删除 bucket 中多余数据避免产生意外费用。关于--fresh-opt参数device_train.py提供--fresh-optactionstore_true用于忽略基础检查点中保存的优化器状态、使用新初始化的优化器而默认情况下read_ckpt(..., load_optnot args.fresh_opt)会连同优化器状态一起加载device_train.py 第 269 行。同时微调时会把加载进来的调度器 step 重置为 0第 271-274 行保证学习率调度从头开始——这也与指南中“global step 将重置为 0编写 lr 调度时要记住这一点”的说明一致。3.4 学习率调度建议Learning Rate Notes指南的学习率部分原作者感谢 nostalgebraist 的解释给出了可复用的调度经验值得完整保留确定每个 epoch 的步数gradient_accumulation_steps即 batch size默认16nostalgebraist 推荐32。.tfrecord 文件名中的数字是数据集序列数除以 batch size 即得到每 epoch 步数。lr推荐在1e-5~5e-5之间end_lr设为lr的1/5 或 1/10。weight_decay保持0.1即可。total_steps至少一个 epoch有验证集时可为更长。warmup_steps设为总步数的5-10%anneal_steps设为total_steps - warmup_steps。调度行为lr在warmup_steps anneal_steps之后降到end_lr并继续训练到total_steps但通常应在退火完成后停止。指南示例某小数据集 tokenize 后为 1147 序列除以gradient_accumulation_steps 16并向上取整得 72 步/epoch设lr 5e-5、end_lr 1e-5、total_steps 72无验证集、anneal_steps 65、warmup_steps 7。学习率调度的底层实现可追溯到 mesh_transformer/util.py 中的gpt3_schedule函数由 device_train.py 第 174 行创建并经optax.scale_by_schedule注入优化器链第 182 行。整个优化器链为optax.scale(1 / gradient_accumulation_steps)梯度累积归一化→clip_by_global_norm(1)梯度裁剪→scale_by_adamAdam→additive_weight_decay(weight_decay)L2 式权重衰减→scale(-1)→scale_by_schedule(scheduler)其语义对应仓库工具函数 mesh_transformer/util.py 中的clip_by_global_norm与additive_weight_decay。3.5 微调后的采样与模型导出指南的“Now what?”章节给出了微调完成后的三条路径README 与仓库源码均可佐证采样python3 device_sample.py --configconfigs/YOUR_CONFIG.json提供基础的采样交互界面对应仓库根目录 device_sample.py。权重瘦身使用 slim_model.py 将新权重转换为便于部署的 slim 版本去掉优化器状态。源码显示它支持--ckpt-step指定转换某个 step 的 checkpoint--f16切换为 float16默认转 bfloat16内部用optax.scale(0)之外的空优化器链构造推理状态与 resharding_example.py 的思路一致。HuggingFace 转换使用 to_hf_weights.py 将权重转换为 HuggingFacetransformers库可识别的 PyTorch 格式指南建议先运行slim_model.py再转换并可用python to_hf_weights.py --help查看用法。指南同时注明截至 2021-09-01GPT-J 已合并入transformers的main分支但尚未发布到生产版本需pip install githttps://github.com/huggingface/transformers#transformers安装 main 分支——该时间信息属于文档历史记载当前版本的 transformers 已稳定支持读者应按现网版本为准。四、JAX 版本依赖为什么必须固定jax0.2.12README 专门用一节强调了一个极易踩坑的约束本库对 JAX 版本有特定要求。具体而言要使用 v1 模型包括 GPT-J-6B需要jax0.2.12这又依赖于jaxlib0.1.68。如果不这样做你会得到令人费解的 xmap 错误。而 v2 模型代码没有公开发布的权重可以使用最新的 JAX 版本。仓库的 requirements.txt 中jax~0.2.12使用的是兼容范围符号而非精确锁定这正是 howto_finetune.md 第 11 步要求“在安装 requirements.txt 之后再显式pip install jax0.2.12”的原因。另外该版本组合还意味着配套依赖也被锁定在较老版本如optax0.0.9、dm-haiku0.0.5、ray[default]1.4.1等均为 requirements.txt 中的精确/近似固定。xmap 在后续 JAX 版本中 API 变化较大这是 README 警告的直接原因——读者在复现或微调 GPT-J-6B 时务必按此版本组合搭建环境。五、模型并行实现从配置到CausalTransformer为了更深入理解 README 所述“模型并行”在代码中的落点可结合源码查看核心模块 mesh_transformer/transformer_shard.pyCausalTransformerShard是每个分片上的模型本体由EmbeddingShard、TransformerLayerShard序列与ProjectionShard组成注意力头被切分heads_per_shard n_heads // shards第 31 行这正是“沿注意力头维度做模型并行”的 Megatron 式切分体现。层初始化缩放init_scale 2. / layer_count第 35 行对应 GPT-2 风格的小初始化用于深层网络的稳定性。位置编码分支pe t5时创建RelativePositionEmbs否则不创建第 42-45 行与 README 及配置中pe: rotary的选择相对应。生成推理路径拆分为generate_initial/generate_once用于增量解码并缓存 KV 状态第 74-115 行是device_sample.py与 mesh_transformer/sampling.py 采样器如nucleaus_sample的底层支撑。在 device_train.py 中网格以mesh_shape (tpu_size // cores_per_replica, cores_per_replica)构建轴命名为(dp, mp)即数据并行 × 模型并行二维 mesh第 195-196、257 行与 README 描述的“2d mesh 网络 Megatron 式并行”直接对应。配置文件中cores_per_replica决定每个模型副本占用多少设备核即模型并行度per_replica_batch决定每个副本的 batchgradient_accumulation_steps为梯度累积步数。训练吞吐可由这些量直接推算tokens_per_step seq * gradient_accumulation_steps * (per_replica_batch * tpu_size // cores_per_replica)device_train.py 第 253-254 行例如 6B 配置在 256 核 TPUv3 上单步处理2048 × 16 × (1 × 256 / 8) 1,048,576token。六、引用规范README 末尾提供了两段 BibTeX分别用于引用仓库本身和GPT-J-6B 权重使用本仓库或预训练权重做研究/论文时建议采用misc{mesh-transformer-jax, author {Wang, Ben}, title {{Mesh-Transformer-JAX: Model-Parallel Implementation of Transformer Language Model with JAX}}, howpublished {\url{https://github.com/kingoflolz/mesh-transformer-jax}}, year 2021, month May }misc{gpt-j, author {Wang, Ben and Komatsuzaki, Aran}, title {{GPT-J-6B: A 6 Billion Parameter Autoregressive Language Model}}, howpublished {\url{https://github.com/kingoflolz/mesh-transformer-jax}}, year 2021, month May }仓库根目录还提供 CITATION.bib 供直接引用。七、仓库配套资源速查围绕本文涉及的内容以下仓库文件可供进一步深入研究均为当前仓库相对路径核心文档README.md、howto_finetune.md模型并行核心实现mesh_transformer/transformer_shard.py、mesh_transformer/layers.py检查点读写与重分片mesh_transformer/checkpoint.py、resharding_example.py训练/采样/评测入口device_train.py、device_sample.py、device_serve.py、eval_harness.py、slim_model.py、to_hf_weights.py数据准备与加载create_finetune_tfrecords.py、tfrecord_loader.py配置与索引configs/6B_roto_256.json、configs/example_config.json、data/ 下的*.index示例环境依赖requirements.txtTPU 编排脚本位于 scripts/如 scripts/init_ray.sh、scripts/create_serve_tpu.sh结语mesh-transformer-jax 是 TPU 上训练与部署 GPT 级自回归模型的一套完整方案README 定义了其 Megatron 式模型并行的设计目标、GPT-J-6B 的规格与基准、TPU VM Ray 的运行架构以及严格的 JAX 版本约束配套的微调指南与仓库源码则把从数据准备、配置编写、学习率调度到权重导出与 HuggingFace 迁移的整条链路落到实处。复现或微调 GPT-J-6B 时请务必锁定jax0.2.12/jaxlib0.1.68并确认 device_* 脚本仅适用于 TPU v3-8即可规避 README 中列出的两大常见陷阱。赞分享大模型基础模型深度学习分布式训练【免费下载链接】mesh-transformer-jaxModel parallel transformers in JAX and Haiku项目地址https://gitcode.com/gh_mirrors/me/mesh-transformer-jax点击查看免费下载相关推荐基于 JAX/Flax 的双向布局 TransformerBLT训练与推理实战指南基于 JAX/Flax 的双向布局 TransformerBLT训练与推理实战指南 layout bltBidirectional Layout Tran人工智能深度学习NLP计算机视觉强化学习如何用mesh-transformer-jax进行模型微调从预训练到领域适配的完整指南如何用mesh transformer jax进行模型微调从预训练到领域适配的完整指南 想要让预训练的大语言模型更好地适应你的特定需求吗 mesh tran大模型基础模型深度学习分布式训练GLM-4.5微调实战LoRA策略高效训练GLM 4.5微调实战LoRA策略高效训练 引言为什么选择LoRA微调GLM 4.5 你是否曾面临这样的困境想要对大语言模型进行领域适配却发现全参数微大模型基础模型深度学习分布式训练上一篇Citra 3DS模拟器终极指南在电脑上畅玩任天堂经典游戏下一篇Flowframes实战指南免费AI视频插帧工具深度解析创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考