MESH:面向混合专家模型训练的内存高效Sinkhorn优化算法
MESH: Memory-Efficient Sinkhorn Optimization for Mixture-of-Experts Training
浏览论文内容
中文总结 AI 辅助
该研究针对MoE训练中直接应用内存高效Sinkhorn优化不可靠的问题,提出MESH算法,在降低优化器内存占用的同时保留接近AdamW的性能。
中文摘要 AI 辅助
内存高效的矩阵优化器(如Sinkhorn梯度下降)可为稠密Transformer矩阵移除大部分AdamW优化器状态,但直接应用于混合专家模型(Mixture-of-Experts,MoE)训练不可靠。我们在受控的1.1亿参数nanowhale DeepSeek风格MoE预训练设置中研究该失效问题:SAGE/Sinkhorn混合方案将优化器状态从0.883GB降至0.331GB,但评估损失降至3.8265,远高于相同设置下AdamW基准(研究的所有随机种子中为3.58至3.64)。我们发现路由式MoE专家矩阵是主要失效点:其梯度具有条件性、随时间变化的特性,无状态的Sinkhorn归一化难以适配。我们提出MESH,一种面向MoE专家的隐动量Sinkhorn更新,通过梯度缓冲生命周期恢复时间一阶矩信号,无需存储专家一阶矩作为优化器状态;MESH是可选的块预处理变体,添加了粗略神经元/块逆RMS乘数。消融实验显示,矩阵归一化前的时间平滑是核心因果因素,块/神经元预处理可优化内存-质量边界,但非普遍必需。在另外两个随机种子中,MESH和MESH-B相比AdamW分别减少62.5%的优化器状态内存、约12.6%的峰值PyTorch CUDA分配,且评估损失差距较小;全状态诊断变体在消融实验中恢复AdamW级性能,支持结论:MoE专家需要时间平滑,但未必需要全坐标AdamW状态。
英文摘要
Memory-efficient matrix optimizers such as Sinkhorn gradient descent remove most AdamW optimizer state for dense Transformer matrices, but direct application to Mixture-of-Experts (MoE) training is unreliable. We study this failure in a controlled 110M-parameter nanowhale DeepSeek-style MoE pretraining setting. A SAGE/Sinkhorn hybrid reduces optimizer state from 0.883GB to 0.331GB but degrades evaluation loss to 3.8265, far above the AdamW baselines observed in the same setup (3.58--3.64 across the seeds we study). We show that routed MoE expert matrices are the dominant failure point: their gradients are conditional, temporally varying, and poorly served by stateless Sinkhorn normalization. We propose MESH, a hidden-momentum Sinkhorn update for MoE experts. MESH restores a temporal first-moment signal through the gradient-buffer lifecycle, without storing the expert first moment as optimizer state. MESH is an optional block-preconditioned variant that adds a coarse neuron/block inverse-RMS multiplier. Across ablations, temporal smoothing before matrix normalization is the primary causal ingredient; block/neuron preconditioning can improve the memory-quality frontier, but is not established as universally necessary. In two additional seeds, MESH and MESH-B reduce optimizer-state memory by 62.5\% and peak PyTorch CUDA allocation by about 12.6\% relative to AdamW, with a modest evaluation-loss gap. Full-state diagnostic variants recover AdamW-like performance in ablations, supporting the conclusion that MoE experts need temporal smoothing, but not necessarily full coordinate-wise AdamW state.