这次我们来看一个优化器方向的新工作:MALT,全称 Lightweight Curvature-Aware Muon via Diagonal Preconditioning。严格说,它不是一个开箱即用的应用工具,而是一个用于大规模深度模型训练的优化器方案,和 AdamW、Muon、Shampoo、SOAP 属于同一类东西。它的核心卖点可以概括为三点:轻量、曲率感知、对角预条件。如果把 Muon 看成对更新方向做正交化约束的优化器,那 MALT 的思路就是在不引入完整矩阵预条件的前提下,用对角信息近似曲率,从而让二阶优化器的收益以更低的显存代价落地。本文会用一套通用流程,带你理解 MALT 的定位、如何在 PyTorch 训练循环中接入、如何做收敛性验证、显存占用观察,以及常见的踩坑排查。适合正在做大模型预训练、低资源微调、优化器选型对比和技术复现的读者。
MALT 这类优化器对显卡没有特殊门槛,只要 PyTorch 环境能正常训练模型,它就能以 optimizer 类的方式接入;它不依赖 WebUI、不依赖 API 服务,也没有一键启动包这回事。如果你关心的是“模型 API 调用”或“批量任务队列”,那在这个项目里对应的是训练实验批量跑数,而不是生成类接口。把预期放在训练复现上,这篇文章会更合适。
接下来按这样的顺序展开:核心能力速览、原理背景、适用边界、环境准备、训练接入、功能测试、批量实验、资源占用、排查清单、最佳实践。全程使用可复制的命令和代码模板,具体参数名需要按你拿到的官方实现调整,这一点后面会反复提醒。
1. 核心能力速览
先给一张速览表,方便你快速判断 MALT 适不适合自己当前的环境和项目。
| 能力项 | 说明 |
|---|---|
| 项目类型 | 深度学习优化器,属于 Muon 优化器家族的轻量改进方向 |
| 核心思路 | 用对角预条件(Diagonal Preconditioning)近似曲率信息,实现曲率感知(Curvature-Aware)更新 |
| 相对优势 | 相比 Shampoo、SOAP 等完整矩阵预条件优化器,显存和计算开销更低,规模更容易放大 |
| 适用模型 | Transformer 预训练、大模型微调、持续训练、大规模表征学习等场景 |
| 依赖框架 | 通常以 PyTorch 优化器形式接入,具体兼容范围以官方实现为准 |
| 硬件要求 | 有 NVIDIA GPU 更好,纯 CPU 也能做小规模验证,速度会明显下降 |
| 启动方式 | 无独立服务,直接在训练脚本中作为 optimizer 调用 |
| 接口能力 | 不提供 HTTP/API 服务,提供的是 Python 优化器接口 |
| 批量任务 | 无自带任务队列,可通过 shell 脚本或 Python 循环批量跑训练实验 |
| 显存占用 | 相对完整二阶优化器更低,具体数值取决于模型规模、batch size 和实现细节 |
| 适合读者 | 研究优化器、做大模型预训练、想把 AdamW 替换成曲率感知优化器的人 |
需要说明的是,表格里所有“通常”“取决于”的表述都表示这是一个方向性判断,不是已经验证过的固定数字。MALT 的具体实现、包名、参数名和显存数据,要以你实际拿到的源码、论文和 README 为准。
2. 从 Muon 到对角预条件:MALT 的原理定位
要判断一个优化器值不值得用,先要知道它在解决什么样的问题。
2.1 Muon 优化器解决了什么
AdamW 是目前最通用的优化器,它对每个参数维护一阶矩和二阶矩,更新方式本质上是逐元素的缩放。它的问题是:对参数之间的相关性考虑不够,在 Transformer 这类参数高度耦合的模型上,收敛步数通常较多。
Muon 的思路是引入矩阵级正交化。它对 2D 权重矩阵的更新方向做正交约束,让更新不再只是逐元素缩放,而是考虑整个矩阵的方向结构。这种约束通常通过 Newton-Schulz 迭代等对称正交化方法实现。实际操作中,Muon 往往对 embedding 和 bias 这类不适合矩阵正交化的参数继续走 AdamW 分支,主力更新交给正交化分支。
2.2 完整预条件的问题
Shampoo 这类优化器走得更远,它对每个参数维护真正的预条件矩阵,期望精确刻画曲率。这种做法的理论效果好,但工程代价很大:预条件矩阵本身要存储和更新,在高维参数下显存和算力都会快速增长。这也是二阶优化器长期停留在小规模场景的原因之一。
SOAP 等分块方案尝试用分块近似降低开销,但整体上仍然保留了较多的额外状态。在百亿参数模型上,这些额外状态会直接吃掉大量显存,甚至超过模型本身和激活值。
2.3 MALT 的轻量化路径
从命名看,MALT 是 Lightweight Curvature-Aware Muon,它要保留 Muon 的更新方向优势,同时又要有轻量级的曲率感知能力。一个合理的实现路径是:
- 保留 Muon 风格的动量更新方向构建;
- 不维护完整矩阵预条件,而是用对角近似统计量刻画曲率;
- 对 2D 权重矩阵和向量参数分别处理;
- 用少量额外状态换更好的收敛条件。
这是一个典型的“用一阶信息量近似二阶信息”的工程思路:不追求完整曲率矩阵,而是只保留曲率中对角线占比最高的部分。由于对角预条件的额外状态基本和梯度同尺度,显存增长可以控制在较低范围。
2.4 和主流优化器的粗略对比
| 优化器 | 预条件信息 | 额外状态量 | 显存趋势 | 适合场景 |
|---|---|---|---|---|
| AdamW | 一阶矩 + 二阶矩 | 约 2 倍梯度 | 低 | 通用训练、微调 |
| Muon | 动量 + 正交化更新 | 约 1 到 2 倍梯度 | 低到中 | 大模型预训练、矩阵权重场景 |
| Shampoo | 逐层矩阵预条件 | 多份小矩阵 | 高 | 小规模高精度优化 |
| SOAP | 分块矩阵预条件 | 多份分块矩阵 | 中到高 | 大模型、长上下文预训练 |
| MALT | 对角曲率近似 | 取决于实现,通常低于完整矩阵方案 | 较低 | 大规模训练、显存敏感场合 |
再次强调,这张表是方向性评估,不是精确 benchmark。如果你要做严谨选型,应该在相同模型、相同数据、相同步数下分别跑 AdamW、Muon 和 MALT,记录 loss、吞吐和显存。
3. 适用场景与使用边界
MALT 不是所有场景都能发挥优势。明确边界,比直接替换优化器更重要。
3.1 适合谁用
第一种是做大规模预训练和持续训练的工程师。他们通常已经对 AdamW 不满意,想在有限显存内获得更好的收敛性能,但又接受不了 Shampoo 的显存开销。MALT 这类“轻量曲率感知”方案是合理的中间选择。
第二种是做优化器研究和复现的研究者。MALT 提供了研究“对角预条件 + Muon 更新”如何影响收敛行为的实验接口。比起改模型结构,调优化器对现有代码的侵入更小。
第三种是在固定卡数下训练大模型的团队。显存是硬约束,只要能压缩优化器额外状态,就能把 batch size 调大,或者让模型规模再往上走一点。这也是曲率感知优化器的核心工程价值。
3.2 不适合什么场景
如果你只是做几十分钟的小实验、batch size 很小、模型只有几百万参数,那 AdamW 可能更省心。MALT 这类优化器在简单任务上未必有肉眼可见的优势,反而可能因为实现复杂引入额外 debug 成本。
如果项目不是 PyTorch 生态,比如纯 JAX/NumPy 上层,或者必须走 TensorFlow 的 SavedModel 流程,接入成本会变高。除非官方已经提供对应框架实现,否则不建议硬迁。
3.3 合规与安全边界
优化器本身不涉及生成内容,但使用过程中要关注三点:
- 代码许可证:MALT 如果以开源仓库发布,需要确认其许可证与你的项目是否兼容,尤其在公司内部或商用场景;
- 训练数据版权:用 MALT 训练模型时,数据集的获取、标注和授权要合规;
- 模型分发:基于 MALT 训练出的权重,在对外发布时要明确基座模型许可证、数据来源和二次使用限制。
4. 环境准备与前置条件
下面给出一套通用环境准备流程。由于 MALT 的官方依赖清单可能随版本变化,这里不写死具体版本号,以你手里的项目 README 为准。
4.1 基础依赖
你需要一个能正常训练模型的 PyTorch 环境,建议按以下顺序检查:
- 操作系统:Linux 最稳妥,Windows 和 macOS 也可以,但多卡训练建议 Linux;
- Python:3.9 以上比较常见,具体看项目要求;
- PyTorch:稳定版本即可,建议支持 AMP 的版本;
- CUDA:如果要用 GPU 训练,提前确认驱动和 CUDA 版本匹配;
- 磁盘:预训练模型、日志和 checkpoint 都会占空间,建议预留足够容量。
4.2 验证基础环境
先启动一个 Python 终端,确认 PyTorch 能看到显卡:
python -c "import torch; print(torch.__version__); print(torch.cuda.is_available())"如果输出True,说明 GPU 可用。如果输出False,需要先排查 CUDA、驱动或 PyTorch 安装问题。这一步不过,后面所有训练测试都跑不起来。
4.3 安装 MALT
安装方式大概率是以下两种之一,具体以官方说明为准:
# 方式一:如果项目发布到了 PyPI pip install malt-optimizer# 方式二:从源码安装 git clone <项目仓库地址> cd <项目目录> pip install -e .这里要特别注意:malt-optimizer是占位包名,实际包名可能不同,直接 pip install 前先查官方文档,避免装到同名无关包。源码安装能保证你拿到最新代码,也方便调试优化器内部实现。
安装完成后,用一个导入测试确认可用:
python -c "from malt_optimizer import MALT; print('MALT import ok')"如果导入失败,优先检查依赖版本冲突,再看是否缺少项目自定义的扩展模块。
5. 在 PyTorch 训练循环中接入 MALT
优化器类项目的核心操作不是“启动服务”,而是“替换优化器”。下面给出一套通用接入模板。
5.1 导入并按参数构造优化器
假设官方实现提供了MALT优化器类,通常可以这样替换:
import torch from torch import nn # 占位导入路径,以官方实现为准 from malt_optimizer import MALT model = nn.Linear(128, 128) optimizer = MALT( model.parameters(), lr=1e-3, weight_decay=0.1, )优化器构造参数中,lr、weight_decay一般都会有,其他参数要看具体实现。如果项目支持对 2D 权重走 Muon 风格更新、对向量参数走 AdamW 风格更新,那通常会有内部参数分组,不需要你手动分。
5.2 训练循环模板
接入方式和普通 PyTorch 优化器完全一致:
from torch.utils.data import DataLoader, TensorDataset inputs = torch.randn(1024, 128) targets = torch.randint(0, 10, (1024,)) dataset = TensorDataset(inputs, targets) dataloader = DataLoader(dataset, batch_size=64, shuffle=True) criterion = nn.CrossEntropyLoss() model = nn.Sequential( nn.Linear(128, 256), nn.ReLU(), nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, 128), nn.Linear(128, 10), ) optimizer = MALT(model.parameters(), lr=1e-3) for epoch in range(3): for batch_x, batch_y in dataloader: optimizer.zero_grad() logits = model(batch_x) loss = criterion(logits, batch_y) loss.backward() optimizer.step() print(f"epoch {epoch} loss {loss.item():.4f}")如果 loss 能稳定下降,说明基本接入流程已经跑通。
5.3 和混合精度配合
大模型训练基本离不开 AMP。在 PyTorch 中,混合精度训练用GradScaler包住反向和优化器更新:
scaler = torch.cuda.amp.GradScaler() optimizer.zero_grad() with torch.cuda.amp.autocast(): logits = model(batch_x) loss = criterion(logits, batch_y) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()这里要注意:如果你的优化器实现内部涉及矩阵正交化或复杂数值计算,AMP 下可能出现精度问题。第一次测试时建议分别在 FP32 和 AMP 下跑一遍,对比 loss 曲线是否异常。
6. 功能测试与效果验证
接入成功不等于效果正确。下面给出一套可复现的验证流程,从“能跑”到“有效”逐步确认。
6.1 冒烟测试:先确认更新正常
先用最小模型跑 50 到 100 步,确认 loss 不会发散、不会 NaN。这个阶段关注的是流程正确性,不是收敛效果。
python train_smoke.py --model mlp --batch-size 32 --steps 100判断标准:
- loss 在 100 步内没有出现 NaN;
- loss 相比初始值有下降趋势;
- 反向传播和 optimizer.step 没有报错。
如果出现 NaN,优先怀疑学习率过大、数值稳定性不足或 AMP 缩放问题。
6.2 和 AdamW 对比收敛曲线
功能测试的关键一步,是建立基线。用同一份数据、同一个模型结构,分别跑 AdamW 和 MALT,控制训练步数一致,记录每步 loss。
python train_compare.py --optimizer adamw --steps 3000 --run-name adamw_base python train_compare.py --optimizer malt --steps 3000 --run-name malt_test然后把两份 log 画成 loss 曲线,重点观察:
- MALT 达到 AdamW 相同 loss 所需步数;
- 两者最终 loss 的差距;
- loss 曲线是否稳定,有没有突然抖动。
这里不要只看一步的结果,建议至少跑 2000 到 5000 步再下结论。优化器的差异在小步数下经常被学习率噪声掩盖。
6.3 显存占用观察
显存是 MALT 的重要卖点,所以必须单独测。在训练脚本里加一段显存峰值记录:
import torch torch.cuda.reset_peak_memory_stats() # 训练循环 for step, (batch_x, batch_y) in enumerate(dataloader): optimizer.zero_grad() logits = model(batch_x) loss = criterion(logits, batch_y) loss.backward() optimizer.step() if step % 100 == 0: peak = torch.cuda.max_memory_allocated() / 1024**2 print(f"step {step} peak memory {peak:.1f} MB")同时可以用 nvidia-smi 观察整体显存:
nvidia-smi --query-gpu=name,memory.used,memory.total --format=csv得到的结果要分三部分看:
- 模型参数和激活值的显存;
- 优化器额外状态的显存;
- 是否因为 batch size 太大导致激活值占用过高。
如果你在对比 AdamW 和 MALT,建议两者使用完全相同的 batch size 和模型,只有优化器不同,这样才比较公平。
6.4 不同学习率的稳定性测试
曲率感知优化器对学习率的敏感度通常和 AdamW 不同。建议扫一组学习率:
lr=1e-4, 3e-4, 1e-3, 3e-3每组跑相同步数,记录最终 loss 和是否出现 NaN。通过这张学习率表,你才能判断 MALT 在你任务上是否需要调整默认 lr。
6.5 小规模多卡一致性测试
如果计划在真实大模型场景用 MALT,先做一次小规模多卡测试,确认 DDP 下梯度同步和优化器状态同步没有问题。
torchrun --nproc_per_node=2 train_ddp.py --model tiny-gpt --steps 500判断标准:单卡和双卡在相同 seed、相同参数下,最终 loss 应该一致,或者非常接近。如果两张卡的 loss 曲线明显分离,优先检查数据采样是否设置了相同 seed,以及梯度同步是否正确。
7. 批量实验与训练任务调度
MALT 没有 HTTP 接口,但训练侧可以方便地做批量实验。这里的“接口”指的是 Python 优化器 API,批量能力则体现在多配置、多卡、多机编排上。
7.1 用 shell 脚本批量跑配置
如果你把训练参数都做成命令行参数,可以用一个简单的 shell 脚本批量跑多组实验:
#!/bin/bash for lr in 1e-4 3e-4 1e-3 3e-3; do for opt in adamw malt; do python train.py \ --optimizer $opt \ --lr $lr \ --steps 3000 \ --save-dir runs/${opt}_lr${lr} done done这种做法的好处是:每个实验独立进程、独立日志、互不干扰,即使某个配置崩了,也不影响其他实验。
7.2 用配置文件驱动实验
另一种更工程化的方式是把实验参数写进 YAML,再统一读取:
model: name: tiny_gpt hidden_size: 256 num_layers: 4 training: optimizer: malt lr: 0.001 weight_decay: 0.1 batch_size: 64 steps: 3000 mixed_precision: true logging: save_every: 500 log_dir: runs/malt_tiny_gpt训练脚本里只需要读配置并构造优化器:
import yaml with open("config.yaml", "r") as f: config = yaml.safe_load(f) optimizer = MALT( model.parameters(), lr=config["training"]["lr"], weight_decay=config["training"]["weight_decay"], )配置化的好处是方便记录每次实验的完整参数,比手敲命令行更可追溯。
7.3 检查点保存与恢复
任何长训练任务都必须支持断点恢复。保存时要同时存模型参数、优化器状态、学习率调度器状态和当前步数。
torch.save({ "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict(), "step": step, "epoch": epoch, }, checkpoint_path)恢复时再重新构造优化器并加载:
checkpoint = torch.load(checkpoint_path, map_location="cpu") model.load_state_dict(checkpoint["model"]) optimizer.load_state_dict(checkpoint["optimizer"]) scheduler.load_state_dict(checkpoint["scheduler"]) step = checkpoint["step"]优化器状态字典经常被忽略,但恢复训练时缺失优化器状态会导致学习率调度和动量信息丢失,相当于半路重新热身。
7.4 批量任务的失败重试
如果批量跑训练的任务在某个配置上崩了,不要每次手动重启。可以在 shell 脚本里加简单的重试逻辑:
if [ -f "$save_dir/checkpoint.pt" ]; then RESUME_FLAG="--resume $save_dir/checkpoint.pt" else RESUME_FLAG="" fi CUDA_VISIBLE_DEVICES=$GPU_ID python train.py $RESUME_FLAG ...断点恢复配合批量循环,能很大程度上提升多组实验的稳定性。
8. 资源占用与性能观察
资源占用是优化器选型的核心观察项,但也是最容易被误报的部分。下面说清楚怎么看、怎么记、怎么避坑。
8.1 显存观察方法
推荐两种方式结合:
nvidia-smi:看整卡显存,能发现显存碎片和峰值问题;torch.cuda.max_memory_allocated():看当前进程 PyTorch 分配的峰值显存,更容易定位到优化器额外状态。
torch.cuda.reset_peak_memory_stats() # 跑完整训练循环 peak = torch.cuda.max_memory_allocated() / 1024**2 print(f"PyTorch peak memory: {peak:.1f} MB")对比不同优化器时,都以 PyTorch 峰值显存为准,不要只看 nvidia-smi 的整卡数值,因为框架缓存和其他进程会干扰判断。
8.2 什么因素会影响显存
- 模型参数量:参数越多,优化器额外状态越多;
- batch size:主要影响激活值显存,不是优化器状态;
- 混合精度:AMP 会降低激活和梯度显存,但优化器状态是否因此减少要看实现;
- 预条件结构:完整矩阵预条件会带来大量额外状态,对角预条件下理论额外状态接近梯度数量级;
- 梯度累积:不会降低单步峰值显存,但能降低单卡 batch size 压力。
8.3 与 AdamW 相比,重点看哪些指标
建议记录一张性能对比表:
| 指标 | AdamW | MALT | 差异说明 |
|---|---|---|---|
| 达到目标 loss 的步数 | 待实测 | 待实测 | 步数越少越好 |
| 每步训练耗时 | 待实测 | 待实测 | 受预条件计算影响 |
| 峰值显存 | 待实测 | 待实测 | 优化器状态差异 |
| 最终 loss | 待实测 | 待实测 | 判断收敛质量 |
| 是否出现 NaN | 待实测 | 待实测 | 数值稳定性 |
如果 MALT 的每步耗时明显更高,但达到目标 loss 的步数更少,你需要计算“总时间 = 每步耗时 × 步数”来判断整体收益。
8.4 如何降低显存占用
如果你的显存不够,可以先调整训练配置,而不是立刻否定优化器:
- 使用梯度累积,降低单卡 batch size;
- 开启梯度 checkpointing,减少激活值缓存;
- 开启 AMP,减少中间张量精度;
- 缩小模型输入序列长度或 hidden size;
- 检查优化器是否支持 fp16/bf16 状态存储。
需要注意,优化器状态的精度压缩通常会带来收敛精度损失。压缩前后建议各跑一小段,确认 loss 曲线没有明显恶化。
9. 常见问题与排查方法
下面是优化器接入实验中最常见的问题和排查思路。
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 导入 MALT 报 ModuleNotFoundError | 包名不对或未安装成功 | 检查安装命令和 Python 环境 | 确认官方包名,重新安装 |
| loss 直接变成 NaN | 学习率过大、AMP 缩放异常、实现数值不稳定 | 降低 lr,关闭 AMP 重跑 | 调整 lr,尝试 bf16 或 FP32 |
| loss 不下降 | 学习率过小、预热设置过长、优化器更新逻辑异常 | 看 loss 曲线,对比 AdamW | 调整 lr,检查优化器是否正确更新参数 |
| 显存不足 OOM | batch size 过大、模型过大、优化器状态过多 | 看峰值显存日志 | 降低 batch,开启梯度累积或 checkpointing |
| 多卡训练 loss 不一致 | DDP 未正确同步、数据采样未设置相同 seed | 单卡和双卡对比 | 固定 seed,检查 DDP 初始化 |
| 训练速度比 AdamW 慢很多 | 预条件计算开销大、临时张量多 | 用 profiler 统计各阶段耗时 | 观察是否只对 2D 权重启用,向量参数走轻量分支 |
| 断点恢复后 loss 突变 | 优化器状态未保存或未加载 | 检查 checkpoint 字段 | 完整保存 optimizer.state_dict 和 scheduler.state_dict |
| pip 安装到错误环境 | 当前 shell 激活了错误的虚拟环境 | 执行which python和pip show | 切换到正确环境重新安装 |
其中 NaN 问题是最常见的,也是最需要耐心排查的。建议先做一次“关闭 AMP + 极小 lr + 小模型”的测试,排除数值不稳定的情况,再逐步打开 AMP、提高 lr。
10. 最佳实践与使用建议
从过往优化器替换的经验看,最稳妥的路径不是直接在大模型上一步到位,而是分层验证。
第一,先在几百万参数的小模型上跑通流程。用固定数据集、固定步数,记录 AdamW 和 MALT 的 loss 曲线。这一步验证的是“优化器能够正常更新模型”,不是最终效果。
第二,再在目标模型的中等规模上测显存。确认 MALT 的额外状态没有超出显存预算。如果显存不是瓶颈,那 MALT 相对 AdamW 的收益就只体现在收敛步数上,你要评估的是时间成本换步数收益是否划算。
第三,最后才做大规模训练。大规模训练前一定要有断点恢复能力和日志采集能力。优化器实验的价值一半在最终结果,一半在你能不能准确解释结果。
日志记录建议每 50 到 100 步输出一条包含 loss、学习率、峰值显存、当前已用时间的记录,方便后期回放和对比。
整个流程中,要固定所有无关变量。对比实验时,模型结构、数据顺序、随机种子、batch size、训练步数、gradient clipping 设置必须完全一致,只允许优化器不同。否则你很难判断 loss 差异到底是优化器带来的,还是数据顺序带来的。
随机种子尤其重要。如果两个实验的 seed 不同,即使优化器完全一样,loss 也会因为数据采样差异产生波动。建议在训练脚本开头统一设置:
import random import numpy as np import torch def set_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)最后,关于代码和模型的合规使用:如果 MALT 以开源项目发布,请保留其许可证信息;如果你用 MALT 训练了模型权重并对外发布,需要确认基座模型权重、训练数据的授权范围,避免把不可再分发的数据或权重带入商用场景。
11. 总结与下一步
MALT 最值得关注的地方,是它在“Muon 风格更新”和“轻量曲率感知”之间找了一个折中点。对显存敏感的团队来说,它可能比 Shampoo、SOAP 更容易落地;对已经用惯 AdamW 的团队来说,它提供了一个不需要大规模改代码就能尝试的优化器替换方向。
如果只做一次验证,我建议先做这个实验:用一份固定数据集,把 AdamW 和 MALT 各跑 3000 步,记录 loss 下降曲线和峰值显存。关注两个数字:达到相同 loss 所需步数是变少了还是变多了,峰值显存比 AdamW 多了多少。如果步数明显变少、显存增加可控,MALT 就值得继续观察;如果两者 loss 几乎一样,显存也没有显著优势,那在你当前的模型规模下它可能不是最优选择。
最容易踩的坑有三个:一是学习率直接沿用 AdamW 的默认值,导致 loss 不稳定,需要重新扫 lr;二是在小模型上看到“没有提升”后就放弃,忽略了优化器差异在大规模训练中会被放大;三是忘了保存优化器状态,导致恢复训练时 loss 发生漂移,误判为优化器问题。
后续可以继续关注的方向包括:MALT 是否支持 bf16 优化器状态、是否兼容 FSDP 和 DeepSpeed ZeRO、是否提供对 1D 参数和 2D 参数的自动分流策略。如果这些能力都具备,那它离大规模预训练落地就更近了一步。