Making Scalable Meta Learning Practical

📄 arXiv: 2310.05674v2 📥 PDF

作者: Sang Keun Choe, Sanket Vaibhav Mehta, Hwijeen Ahn, Willie Neiswanger, Pengtao Xie, Emma Strubell, Eric Xing

分类: cs.LG, cs.AI

发布日期: 2023-10-09 (更新: 2023-10-23)


💡 一句话要点

提出SAMA以解决可扩展元学习的实际应用问题

🎯 匹配领域: 支柱九:具身大模型 (Embodied Foundation Models)

关键词: 元学习 可扩展性 隐式微分 自适应优化器 分布式训练 文本分类 图像分类

📋 核心要点

  1. 现有元学习方法面临计算和内存成本高、训练不稳定及缺乏高效分布式训练支持等挑战。
  2. 本文提出SAMA,通过避免显式计算二阶梯度信息,结合隐式微分算法和高效分布式训练技术,提升元学习的可扩展性。
  3. 实验结果表明,SAMA在单/多GPU设置下吞吐量提升高达1.7/4.8倍,内存消耗降低2.0/3.8倍,并在文本分类和图像分类任务中实现了最先进的效果。

📝 摘要(中文)

尽管元学习在机器学习程序中具有灵活学习多样化归纳偏置的能力,但由于计算/内存成本高、训练不稳定以及缺乏高效的分布式训练支持,元学习的可扩展性一直受到限制。本文提出SAMA,结合隐式微分算法和系统的进展,旨在使可扩展元学习更具实用性。SAMA灵活支持多种自适应优化器,减少计算负担,避免显式计算二阶梯度信息,并利用针对一阶梯度的高效分布式训练技术。在多个大规模元学习基准测试中,SAMA在单/多GPU设置下分别展示了高达1.7/4.8倍的吞吐量提升和2.0/3.8倍的内存消耗降低。此外,基于SAMA的数据优化在BERT和RoBERTa大语言模型的文本分类准确性上也取得了一致性提升,并在图像分类任务的小规模和大规模数据剪枝中实现了最先进的结果,展示了可扩展元学习在语言和视觉领域的实际应用潜力。

🔬 方法详解

问题定义:本文旨在解决元学习在可扩展性方面的不足,特别是计算和内存成本高、训练不稳定及缺乏高效分布式训练支持的问题。

核心思路:SAMA的核心思路是结合隐式微分算法和高效的分布式训练技术,灵活支持多种自适应优化器,同时减少计算负担,避免显式计算二阶梯度信息。

技术框架:SAMA的整体架构包括多个模块,首先是自适应优化器的选择模块,其次是隐式微分计算模块,最后是分布式训练模块,确保在不同硬件环境下的高效性。

关键创新:SAMA的主要创新在于其避免了显式计算二阶梯度信息的设计,使得计算效率显著提升,并且能够灵活适应多种优化器,这与现有方法的设计思路有本质区别。

关键设计:在关键设计上,SAMA采用了针对一阶梯度的高效分布式训练技术,并在参数设置和损失函数上进行了优化,以适应不同的元学习任务和数据集。

🖼️ 关键图片

fig_0
fig_1
fig_2

📊 实验亮点

实验结果显示,SAMA在单GPU和多GPU设置下分别实现了1.7倍和4.8倍的吞吐量提升,同时内存消耗分别降低了2.0倍和3.8倍。此外,SAMA在文本分类任务中提升了BERT和RoBERTa模型的准确性,并在图像分类任务中实现了最先进的剪枝效果。

🎯 应用场景

该研究的潜在应用领域包括自然语言处理和计算机视觉等多个领域。SAMA的设计使得元学习在大规模数据集上的应用变得更加可行,能够有效提升模型的训练效率和准确性,具有重要的实际价值和未来影响。

📄 摘要(原文)

Despite its flexibility to learn diverse inductive biases in machine learning programs, meta learning (i.e., learning to learn) has long been recognized to suffer from poor scalability due to its tremendous compute/memory costs, training instability, and a lack of efficient distributed training support. In this work, we focus on making scalable meta learning practical by introducing SAMA, which combines advances in both implicit differentiation algorithms and systems. Specifically, SAMA is designed to flexibly support a broad range of adaptive optimizers in the base level of meta learning programs, while reducing computational burden by avoiding explicit computation of second-order gradient information, and exploiting efficient distributed training techniques implemented for first-order gradients. Evaluated on multiple large-scale meta learning benchmarks, SAMA showcases up to 1.7/4.8x increase in throughput and 2.0/3.8x decrease in memory consumption respectively on single-/multi-GPU setups compared to other baseline meta learning algorithms. Furthermore, we show that SAMA-based data optimization leads to consistent improvements in text classification accuracy with BERT and RoBERTa large language models, and achieves state-of-the-art results in both small- and large-scale data pruning on image classification tasks, demonstrating the practical applicability of scalable meta learning across language and vision domains.