TACO:用于大语言模型微调的三元绝对最大逐列单稀疏优化器

TACO: Ternary Absolute-max Column-wise One-sparse Optimizer for LLM Fine-Tuning

arXiv: 2610.02199v1

论文信息

标题: TACO: Ternary Absolute-max Column-wise One-sparse Optimizer for LLM Fine-Tuning

作者: Jichao Jiang, Cristian McGee, El Houcine Bergou, et al.

发布日期: 2026-10-01

arXiv ID: 2610.02199v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:大语言模型全参数微调时,AdamW 的优化器状态内存开销过大,限制单卡可训练模型规模;Muon 虽降低状态但仍为稠密,且与 AdamW 预训练模型存在几何不匹配问题。
  • 核心方法:TACO 在维度归一化 1 ⁣→ ⁣11\!\to\!1 算子范数下求精确最速下降方向,每列只保留绝对值最大梯度的符号,形成三值、逐列 1-稀疏更新。
  • 关键结果:在 OPT-13B 上,TACO 的持久优化器状态从 AdamW8bit 的 27.7GB 降至 0.16GB,约 174 倍;峰值训练内存从 80.6GB 降至 27.5GB,约 2.9 倍,同时 SST-2 准确率为 94.22%(表 4)。
  • 主要局限:理论保证建立在光滑性、固定输入、平衡赢家等简化假设上;默认 2000 步预算并非性能上限;在 SQuAD、DROP 上 μ=0.95\mu=0.95 的梯度历史不如 μ=0\mu=0(表 7)。
  • 适合读者:关注大模型内存高效微调、优化器设计、算子范数优化或单卡大模型训练系统的研究者与工程师。

论文背景和研究动机

全参数微调仍是迁移大语言模型的有力基线,但 AdamW 需要维护两个 FP32 矩估计。论文指出,13B 参数模型在 BF16 下约需 26GB 权重内存,两个 FP32 矩缓冲约再增加 104GB,梯度约 26GB,尚未计算激活、临时工作区或 FP32 主副本(第 1 节)。因此 AdamW 的持久优化器状态约为参数内存的 4 倍。

已有方案包括状态量化、因子化二阶矩、低秩梯度投影和零阶方法。Muon 通过谱几何下的矩阵最速下降降低状态内存,但其状态仍是稠密的;且从 AdamW 预训练模型切换到 Muon 微调可能出现性能下降,论文称这一现象为优化器不匹配,并归因于两者更新几何与隐式偏差不同。作者由此提出核心问题:能否保留一阶梯度信息和计算效率,同时使持久优化器状态几乎可忽略。

核心方法和技术细节

TACO 从一阶 Taylor 展开下的局部最速下降视角出发,但选择与 Muon 不同的约束几何:限制更新方向 DD 在维度归一化 1 ⁣→ ⁣11\!\to\!1 算子范数下不超过 1。该范数等价于每列独立的 ℓ1\ell_1 预算:

∥D∥(1,n→1,m)=nmmax⁡j∥dj∥1.\|D\|_{(1,n\to 1,m)}=\frac{n}{m}\max_j\|d_j\|_1.

其对应的对偶范数为

∥G∥(1,n→1,m)∗=mn∑j∥gj∥∞.\|G\|_{(1,n\to 1,m)^*}=\frac{m}{n}\sum_j\|g_j\|_\infty.

因此局部问题的精确解是每个梯度列只保留绝对值最大项的符号:

D⋆=−mnS(G).D^\star=-\frac{m}{n}\mathcal{S}(G).

其中 S(G)\mathcal{S}(G) 在每列选择绝对值最大坐标的符号,其余位置为零。这样 m×nm\times n 更新最多只有 nn 个非零元。论文称这一稀疏性来自几何,而非事后梯度稀疏化(第 3 节)。

实际 minibatch 梯度噪声会使赢家坐标不稳定。TACO 不维护稠密梯度历史,而维护每列动态 heavy hitters:

Mt=Hk ⁣(μMt−1+(1−μ)Gt),M_t=\mathcal{H}_k\!\left(\mu M_{t-1}+(1-\mu)G_t\right),

其中 μ=0.95\mu=0.95,k=16k=16,Hk\mathcal{H}_k 每列只保留 kk 个绝对值最大项。保留值以 FP8 E4M3 存储,行索引用 int32 存储,状态复杂度从 O(mn)O(mn) 降至 O(kn)O(kn)。其思想是:TACO 只需要恢复每列的主导坐标,并不需要精确重建稠密历史。实现上,TACO 通过 PyTorch autograd hooks 在梯度计算后立即进行 in-place EMA、heavy hitter 提取和权重更新,然后释放稠密梯度;梯度累积设为 1,并配合 checkpointing 降低激活内存。矩阵参数使用 TACO,标量和向量参数使用辅助 AdamW;token embedding 的缩放因子取 1(第 5 节)。

创新点和贡献

论文的主要贡献可归纳为四点。第一,提出几何诱导的稀疏优化器:TACO 的逐列 Top-1 更新来自 1 ⁣→ ⁣11\!\to\!1 算子范数下局部问题的闭式解,而不是先计算稠密梯度再做事后压缩。第二,给出理论分析:定理 1 在标准光滑性假设下证明 vanilla TACO 达到 ε\varepsilon-稳定点需 O(ε−2)O(\varepsilon^{-2}) 次迭代;定理 2 和定理 3 在固定输入、共同目标输出设定下,证明 TACO 与 Adam proxy 共享延续几何和极限解,而 Muon 通常走向不同解;定理 4 在平衡赢家假设下证明 TACO 方向满足 μ\muP 谱缩放 Θ(m/n)\Theta(\sqrt{m/n})。第三,给出稀疏历史与 winner-margin 保持条件(定理 5),为低精度、稀疏状态设计提供理论依据。第四,系统实现将持久状态从 O(mn)O(mn) 降到 O(kn)O(kn),并实现单张 80GB H100 上全参数微调 30B–32B 级模型。

实验结果分析

在 OPT-13B 的 8 个下游任务比较中,TACO 在各任务上均使用最少峰值 GPU 内存,同时保持与基线可比的准确率或 F1(表 3)。例如 SST-2 效率对比中,TACO 的持久状态为 0.159GB,峰值内存 27.54GB,吞吐 275.8 tokens/s,准确率 94.22%;AdamW8bit 为 27.708GB 状态、80.60GB 峰值、95.25% 准确率;Adafactor 状态更小,为 0.013GB,但峰值内存达到 52.59GB;FlashAdamW 吞吐最高,为 546.7 tokens/s,但峰值内存 67.68GB(表 4)。这说明在该实验设置下,TACO 位于低内存 Pareto 前沿,但不一定是最快或最高精度方法。

缩放实验中,TACO 使用同一套超参数从 OPT-1.3B 扩展到 OPT-30B,以及从 Qwen3-8B 扩展到 Qwen3-32B,均保持在单张 80GB H100 内。OPT-30B 的峰值内存在 SST-2、RTE、BoolQ 上分别为 62.5GB、64.4GB、68.2GB;Qwen3-32B 相应约为 69.6GB、71.6GB、75.9GB(表 5)。在 OPT-30B 的多个下游任务中,优化器状态仅 0.267GB,不到各任务峰值内存的 0.5%(图 1 右)。架构迁移上,Llama-3.1-8B、Qwen3-32B 表现较好,Pythia-12B 和 Mistral-24B 在 RTE、BoolQ 上波动更大(图 6)。

消融实验给出若干有限实验结论:k=16k=16 在状态内存和 BoolQ/RTE 性能之间取得最佳平衡;μ=0.95\mu=0.95 在 6 个分类与多项选择任务上较 μ=0\mu=0 更好,但 SQuAD 和 DROP 例外;Top-1 更新在最终性能上优于 Top-2 及以上稠密更新;将步预算从 2K 增至 16K,OPT-1.3B 的 BoolQ、RTE 分别提高 6.5 和 4.3 个百分点,SQuAD 提高 2.1 F1(表 6)。所有评测的模型-数据集组合最佳学习率均为 3×10−53\times10^{-5}(图 7)。

实践建议

如果目标是在单张 80GB H100 上全参数微调 30B 以上模型,且希望保持较低峰值内存,TACO 是工程上值得尝试的候选。可直接沿用论文默认配置:学习率 3×10−53\times10^{-5},μ=0.95\mu=0.95,k=16k=16,batch size 4,线性学习率调度、50 步 warmup,并开启 gradient checkpointing。矩阵参数使用 TACO,标量和向量参数保留辅助 AdamW;对于 GPT-NeoX 类 fused QKV 参数,应按论文做法拆成 Q、K、V 三个逻辑矩阵分别处理。

如果硬件内存不紧张且吞吐优先,FlashAdamW 或 Adafactor 在 SST-2 上显示出更高的 tokens/s(表 4),可能更适合。但需要注意 FlashAdamW 的持久状态和峰值内存更高。

默认 2000 步预算不是 TACO 的性能天花板,若 BoolQ、RTE 类任务需要更强结果,可尝试提高到 4K–16K 步。对于 SQuAD、DROP 等任务,μ=0.95\mu=0.95 不一定总是有利,可小规模验证 μ=0\mu=0 或调整历史长度。学习率迁移在论文评测的 OPT 和 Qwen3 系列中表现一致,但换到新的模型架构或数据分布时,仍需以验证集复核。